Compare commits

..
Author SHA1 Message Date
elky 28f61ec45b test: align concurrency fixtures with runtime boundaries
Centralize isolated Redis test connections in the shared test module, use the canonical Claude Messages format, and correct the borrowed entry comparison.
2026-09-10 08:36:53 +08:00
elky 6aeadcd1d7 fix: resolve concurrency hardening lint failures
Use typed connection admission errors, group HTTP limits, and make test lock lifetimes explicit. Handle fixture reads and remove unnecessary cloning and manual divisibility checks.
2026-09-10 08:31:47 +08:00
elky 3a8dadcd6b Merge remote-tracking branch 'origin/main' into codex/concurrency-hardening 2026-09-10 08:16:50 +08:00
elky ecc16673eb fix: harden concurrency limits and high-RPM runtime paths
Bound request, stream, queue, and shutdown resource lifetimes. Reduce scheduler and Redis hot-path work and isolate database maintenance. Include regression coverage, load probes, and concurrency audit results.
2026-09-10 08:14:58 +08:00
elky d28dd89039 fix(providers): restore legacy endpoint health defaults 2026-09-09 15:51:51 +08:00
elky 8260a87215 ci: build Linux-only gateway releases 2026-09-09 13:01:05 +08:00
elky 361952ada9 fix: resolve workspace lint and regression test failures 2026-09-09 11:34:45 +08:00
elky 6630856061 fix: harden routing failover, model testing, and wallet queries 2026-09-09 10:38:25 +08:00
elky a893bd0557 refactor(data): reuse payment order query 2026-09-09 09:21:09 +08:00
elky f2839ae6a7 feat(routing): add strategy failover controls 2026-09-09 09:12:09 +08:00
elky e58570d79d feat(routing): make client disconnect behavior strategy-scoped 2026-09-08 23:11:37 +08:00
elky 99f6499b2b fix(conversion): improve stream failures and diagnostic exports 2026-09-08 21:04:22 +08:00
elky 17d01d7fe0 fix(dns): unify provider resolution and bound SMTP and tunnel egress
Share provider DNS policy across WebSocket and connection probes, handle bracketed IPv6 literals, and preserve bounded address sets for outbound clients.

Bound SMTP DNS and TCP setup with multi-address fallback. Add opt-in trusted proxy DNS for tunnel upstreams while retaining default IP ACLs and origin isolation.

Document DNS policy boundaries and verify 809 gateway, tunnel, and HTTP regression tests.
2026-09-08 17:44:59 +08:00
elky 8b766930b0 fix(ci): resolve formatting, clippy and migration alias checks 2026-09-08 12:41:19 +08:00
elky c7e403b410 fix: restore container logging compatibility and normalize legacy policies 2026-09-08 11:43:51 +08:00
elky cf8ea19856 fix: harden OAuth identity and cookies and correct quota and JSON display 2026-09-08 10:51:25 +08:00
elky 7113d04f8a fix(usage): preserve original captured HTTP headers 2026-09-08 08:49:35 +08:00
elky 099b810a2f feat: optimize usage body viewing and provider card layout 2026-09-08 02:49:06 +08:00
elky 7aa0c89244 fix(gateway): restore HTTP and WS upstream support 2026-09-07 22:15:05 +08:00
elky 7847ae98c6 fix(gateway): reset stream first-byte timeout per candidate 2026-09-07 21:56:16 +08:00
elky a90d564931 fix: restore security hardening compatibility and validation
Restore authorized rule reveal, explicit full HTTP capture and retention, video task business fields, and valid payment URLs. Add opt-in credential preservation for trusted recovery, fix frontend type contracts and async races, and eliminate PostgreSQL test fixture resource leaks. Document audit coverage and successful fmt and CI-scoped Clippy checks.
2026-09-07 21:14:27 +08:00
github-actions[bot] a5c3699ae9 chore(tunnel): update download links for tunnel-v0.3.17 2026-09-07 08:06:59 +00:00
495 changed files with 48941 additions and 6230 deletions
+58 -12
View File
@@ -15,11 +15,6 @@ APP_PORT=8084
# APP_IMAGE=ghcr.io/fawney19/aether:beta
# APP_IMAGE=ghcr.io/fawney19/aether:0.7.0-rc.1
# Compose 应用容器的非 root 数字身份。
# install.sh 会自动写入安装用户的 UID/GID。
AETHER_CONTAINER_UID=65532
AETHER_CONTAINER_GID=65532
# API Key 前缀(默认 sk)
API_KEY_PREFIX=sk
@@ -31,7 +26,11 @@ RUST_LOG=aether_gateway=info
# 示例: http://localhost:5173,https://app.example.com
# CORS_ORIGINS=http://localhost:5173
# CORS_ALLOW_CREDENTIALS=true
# 如果前后端跨站并依赖登录刷新 Cookie,还要配合:
# 登录刷新 Cookie 对同源浏览器请求和可信反代自动适配 HTTP/HTTPS。
# HTTP 自动使用兼容的 SameSite=Lax(显式 Strict 保留);HTTPS 保留原有 SameSite 配置。
# 无法确认访问协议时保留安全默认值;HTTPS 反代请正确传递 X-Forwarded-Proto。
# AUTH_REFRESH_COOKIE_SECURE 可显式覆盖自动判断,公网部署仍建议使用 HTTPS。
# 如果前后端跨站并依赖登录刷新 Cookie,必须使用 HTTPS,并配合:
# AUTH_REFRESH_COOKIE_SAMESITE=None
# AUTH_REFRESH_COOKIE_SECURE=true
@@ -84,11 +83,58 @@ ADMIN_USERNAME=admin123456
# PostgreSQL 连接池配置(默认每核 4 条、总池至少 32 条且最多 100 条;多实例部署应显式分配每实例预算)
# AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=12
# AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=80
# 普通 PostgreSQL 连接默认语句超时 30 秒、锁等待超时 3 秒;0 显式关闭。
# 这是单条 SQL 的期限,不是整个事务总期限;迁移和历史 backfill 使用独立连接放宽。
# AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_TIMEOUT_MS=30000
# AETHER_GATEWAY_DATA_POSTGRES_LOCK_TIMEOUT_MS=3000
# AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS=2048
# 所有监听分片共用的入站 TCP 连接上限,包含握手、空闲 keep-alive 和升级后的 WebSocket。
# 未设置或 0 时按请求上限 + WebSocket 上限推导,最大 65536;显式配置也受 FD 余量限制。
# 已知 FD soft limit 时最多 max(1, (FD - 256) / 2),不是整个进程的 FD/内存保证。
# 满额的新连接在 HTTP 解析前关闭,不排队创建任务,也不会返回 HTTP 429/503。
# AETHER_GATEWAY_MAX_HTTP_CONNECTIONS=4096
# 停机先等待 HTTP 请求,再排空本地用量写入;以下期限单位为毫秒。
# 进程管理器的强杀期限应覆盖两阶段之和,再预留至少 10 秒收尾。
# AETHER_GATEWAY_HTTP_SHUTDOWN_TIMEOUT_MS=30000
# AETHER_GATEWAY_USAGE_SHUTDOWN_TIMEOUT_MS=30000
# 每客户端、每 origin 的上游空闲连接缓存;不限制活动请求或流持续时间。
# AETHER_GATEWAY_UPSTREAM_POOL_MAX_IDLE_PER_HOST=32
# AETHER_GATEWAY_UPSTREAM_POOL_IDLE_TIMEOUT_MS=15000
# AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB=256
# 请求体完整读取总超时默认关闭;确需限制时配置 1000-600000 毫秒的非零值。
# AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS=0
# 单请求解压后 Payload 上限(MiB),默认 256;显式设为 0 才表示不限制。
# 请求体按实际缓冲增长申请额度,解压同时计入输入和输出;额度不足返回 503。
# 请求体完整读取总超时默认 120000 毫秒;非零值限制在 1000-600000,显式 0 关闭。
# AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS=120000
# 上游流首包后空闲超时默认 300000 毫秒;执行配置 read_ms 优先,显式 0 关闭。
# AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS=300000
# 流式响应诊断捕获共享预算默认 128 MiB,包含 provider/client 分配;不足时截断审计副本。
# 显式 0 关闭此类捕获;不限制协议解析、终态编码和 usage 队列的总内存。
# AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES=134217728
# usage 诊断正文共享预算,按 JSON 堆内存估算,默认 128 MiB。
# 覆盖终态队列 seed、Redis 解码事件、数据库写入 DTO 及正文副本;额度随正文释放。
# 不足或显式 0 时先保留计费事实再舍弃正文;已有清空/禁用状态保留,其余标为截断。
# 不包含原始 Redis 批次、解码临时分配、序列化和压缩结果、协议观察缓冲或进程总内存。
# AETHER_USAGE_EVENT_CAPTURE_MEMORY_BUDGET_BYTES=134217728
# 新 usage 队列消息的完整 JSON payload 上限,默认 1 MiB;显式 0 非法。
# 超限先保留计费事实并舍弃诊断字段;仍超限或无法保留计费语义则拒绝入队,终态尝试受限落库,失败即明确失败。
# 不限制存量 Redis 消息、整个读取批次、DLQ 或进程总内存。
# usage_runtime_queue_payload_* 降级/拒绝计数包含入队和重试预校验的编码尝试,不代表唯一事件数。
# AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES=1048576
# usage worker 读取/重领共用的逻辑 payload 预留:全进程默认 128 MiB,单批目标 8 MiB。
# 按当前消息上限推导 COUNT,默认最多 8 条;预留覆盖整批处理和确认,额度不足等待。
# 当前消息上限不能超过总预留额度;单批目标不足一条时仍读一条。0 或非法值回退默认。
# 历史/其他生产者的大消息继续处理并计数,不是 RESP、实际堆内存或 DLQ 的硬上限。
# AETHER_USAGE_QUEUE_READ_PAYLOAD_BUDGET_BYTES=134217728
# AETHER_USAGE_QUEUE_READ_BATCH_PAYLOAD_BYTES=8388608
# DLQ 原文和最坏 JSON 编码独立预留默认 64 MiB,后台编码及写入最多同时 4 个任务。
# 入场前一次预留,额度占满或单条超预算立即失败并保留 pending 原消息,后续重领。
# 编码失败不阻塞同批其余正常消息;批次仍报告失败,只有成功项会被确认。
# JSON 按字符串最多 6 倍转义保守估算;不截原账务字段,不包含 Redis 命令/连接副本或 RSS。
# 0/非法值回退默认;bytes 最大约 4 GiB,jobs 最大 128。超大存量可能需调高额度后恢复。
# AETHER_USAGE_DLQ_ENCODING_BUDGET_BYTES=67108864
# AETHER_USAGE_DLQ_ENCODING_MAX_JOBS=4
# 内置 Redis 死信转移要求 Redis 7+ 及 EVAL/TYPE/XPENDING/XADD/XACK/XDEL 权限。
# stream 与 DLQ 不能同名;Cluster 还要求两键同 slot,现有默认键未自动迁移。
# 单请求解压后 Payload 上限(MiB),默认 256;显式 0 仍受 256 MiB 硬上限保护。
# AETHER_MAX_REQUEST_BODY_MB=256
# AETHER_GATEWAY_SECURITY_CACHE_TTL_MS=1000
# AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB=64
@@ -115,9 +161,9 @@ ADMIN_USERNAME=admin123456
# Fake-IP 域名及内网 DNS。仅信任管理员配置的上游;没有严格 DNS 过滤开关。
# URL 协议、字面 IP、TLS 证书,以及隧道中继和登录 OAuth 的校验仍保留。
# 可选 Provider OAuth 客户端。Gemini CLI 授权及刷新必须配置 client secret。
# Antigravity 默认使用内置 native-app 客户端凭据;自定义 client ID 时必须同时配置
# 对应的 client secret。未配置 client ID 时使用内置的公开 native-app client ID。
# 可选 Provider OAuth 客户端。Gemini CLI 和 Antigravity 默认使用内置 native-app
# 客户端凭据;自定义 client ID 时必须同时配置对应的 client secret。
# 显式配置的 client secret 优先于默认值。
# AETHER_GEMINI_CLI_OAUTH_CLIENT_ID=
# AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET=
# AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID=
+32 -39
View File
@@ -251,18 +251,6 @@ jobs:
arch: arm64
os: ubuntu-latest
use_cross: true
- name: macos-amd64
target: x86_64-apple-darwin
platform: macos
arch: amd64
os: macos-15-intel
use_cross: false
- name: macos-arm64
target: aarch64-apple-darwin
platform: macos
arch: arm64
os: macos-15
use_cross: false
steps:
- uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5
with:
@@ -398,31 +386,29 @@ jobs:
VERSION="nightly"
mkdir -p package release-assets
for platform in linux macos; do
for arch in amd64 arm64; do
bundle="aether-${VERSION}-${platform}-${arch}"
root="package/${bundle}"
mkdir -p "${root}/bin" "${root}/frontend"
for arch in amd64 arm64; do
bundle="aether-${VERSION}-linux-${arch}"
root="package/${bundle}"
mkdir -p "${root}/bin" "${root}/frontend"
install -m 0755 \
"artifacts/nightly-gateway-${platform}-${arch}/aether-gateway" \
"${root}/bin/aether-gateway"
cp -R artifacts/nightly-frontend-dist/. "${root}/frontend/"
sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > "${root}/install.sh"
chmod 0755 "${root}/install.sh"
install -m 0755 update.sh "${root}/update.sh"
install -m 0644 docker-compose.yml "${root}/docker-compose.yml"
install -m 0644 docker-compose.single-node.yml "${root}/docker-compose.single-node.yml"
install -m 0644 .env.example "${root}/.env.example"
install -m 0755 generate_keys.sh "${root}/generate_keys.sh"
install -m 0644 README.md "${root}/README.md"
install -m 0644 LICENSE "${root}/LICENSE"
install -m 0755 \
"artifacts/nightly-gateway-linux-${arch}/aether-gateway" \
"${root}/bin/aether-gateway"
cp -R artifacts/nightly-frontend-dist/. "${root}/frontend/"
sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > "${root}/install.sh"
chmod 0755 "${root}/install.sh"
install -m 0755 update.sh "${root}/update.sh"
install -m 0644 docker-compose.yml "${root}/docker-compose.yml"
install -m 0644 docker-compose.single-node.yml "${root}/docker-compose.single-node.yml"
install -m 0644 .env.example "${root}/.env.example"
install -m 0755 generate_keys.sh "${root}/generate_keys.sh"
install -m 0644 README.md "${root}/README.md"
install -m 0644 LICENSE "${root}/LICENSE"
tar -C package -czf "release-assets/${bundle}.tar.gz" "${bundle}"
done
tar -C package -czf "release-assets/${bundle}.tar.gz" "${bundle}"
done
sed \
@@ -432,8 +418,8 @@ jobs:
chmod 0755 release-assets/install.sh
(cd release-assets && sha256sum *.tar.gz > SHA256SUMS)
test "$(find release-assets -maxdepth 1 -name '*.tar.gz' | wc -l)" -eq 4
test "$(wc -l < release-assets/SHA256SUMS)" -eq 4
test "$(find release-assets -maxdepth 1 -name '*.tar.gz' | wc -l)" -eq 2
test "$(wc -l < release-assets/SHA256SUMS)" -eq 2
(cd release-assets && sha256sum -c SHA256SUMS)
for archive in release-assets/*.tar.gz; do
tar -tzf "${archive}" >/dev/null
@@ -517,6 +503,15 @@ jobs:
--repo "${REPOSITORY}" \
--clobber
published_assets="$(gh release view "${RELEASE_TAG}" --repo "${REPOSITORY}" --json assets --jq '.assets[].name')"
while IFS= read -r asset_name; do
if [[ "${asset_name}" == aether-nightly-*.tar.gz && ! -f "release-assets/${asset_name}" ]]; then
gh release delete-asset "${RELEASE_TAG}" "${asset_name}" \
--repo "${REPOSITORY}" \
--yes
fi
done <<<"${published_assets}"
# target_commitish does not move an existing git tag. Move the ref
# only after the complete asset set is available.
if gh api "repos/${REPOSITORY}/git/ref/tags/${RELEASE_TAG}" >/dev/null 2>&1; then
@@ -541,8 +536,6 @@ jobs:
expected_assets=(
aether-nightly-linux-amd64.tar.gz
aether-nightly-linux-arm64.tar.gz
aether-nightly-macos-amd64.tar.gz
aether-nightly-macos-arm64.tar.gz
SHA256SUMS
install.sh
)
+21 -35
View File
@@ -186,18 +186,6 @@ jobs:
arch: arm64
os: ubuntu-latest
use_cross: true
- name: macos-amd64
target: x86_64-apple-darwin
platform: macos
arch: amd64
os: macos-15-intel
use_cross: false
- name: macos-arm64
target: aarch64-apple-darwin
platform: macos
arch: arm64
os: macos-15
use_cross: false
steps:
- uses: actions/checkout@fbc6f3992d24b796d5a048ff273f7fcc4a7b6c09 # v5
@@ -356,31 +344,29 @@ jobs:
fi
mkdir -p package release-assets
for platform in linux macos; do
for arch in amd64 arm64; do
bundle="aether-${VERSION}-${platform}-${arch}"
root="package/${bundle}"
mkdir -p \
"${root}/bin" \
"${root}/frontend"
for arch in amd64 arm64; do
bundle="aether-${VERSION}-linux-${arch}"
root="package/${bundle}"
mkdir -p \
"${root}/bin" \
"${root}/frontend"
install -m 0755 "artifacts/aether-gateway-${platform}-${arch}/aether-gateway" "${root}/bin/aether-gateway"
cp -R artifacts/frontend-dist/. "${root}/frontend/"
sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > "${root}/install.sh"
chmod 0755 "${root}/install.sh"
install -m 0755 update.sh "${root}/update.sh"
install -m 0644 docker-compose.yml "${root}/docker-compose.yml"
install -m 0644 docker-compose.single-node.yml "${root}/docker-compose.single-node.yml"
install -m 0644 .env.example "${root}/.env.example"
install -m 0755 generate_keys.sh "${root}/generate_keys.sh"
install -m 0644 README.md "${root}/README.md"
install -m 0644 LICENSE "${root}/LICENSE"
install -m 0755 "artifacts/aether-gateway-linux-${arch}/aether-gateway" "${root}/bin/aether-gateway"
cp -R artifacts/frontend-dist/. "${root}/frontend/"
sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > "${root}/install.sh"
chmod 0755 "${root}/install.sh"
install -m 0755 update.sh "${root}/update.sh"
install -m 0644 docker-compose.yml "${root}/docker-compose.yml"
install -m 0644 docker-compose.single-node.yml "${root}/docker-compose.single-node.yml"
install -m 0644 .env.example "${root}/.env.example"
install -m 0755 generate_keys.sh "${root}/generate_keys.sh"
install -m 0644 README.md "${root}/README.md"
install -m 0644 LICENSE "${root}/LICENSE"
tar -C package -czf "release-assets/${bundle}.tar.gz" "${bundle}"
done
tar -C package -czf "release-assets/${bundle}.tar.gz" "${bundle}"
done
sed \
Generated
+8
View File
@@ -116,6 +116,7 @@ name = "aether-billing"
version = "0.1.0"
dependencies = [
"aether-data-contracts",
"aether-runtime-state",
"aether-usage-runtime",
"async-trait",
"serde",
@@ -305,6 +306,7 @@ dependencies = [
"futures-util",
"hmac",
"http",
"http-body",
"http-body-util",
"hyper",
"hyper-util",
@@ -369,8 +371,12 @@ dependencies = [
"bytes",
"futures-util",
"http",
"http-body-util",
"hyper",
"hyper-util",
"serde_json",
"tokio",
"tokio-util",
"tower",
"tracing",
"tracing-subscriber",
@@ -420,6 +426,7 @@ dependencies = [
"aether-data",
"aether-data-contracts",
"aether-gateway",
"aether-runtime",
"aether-runtime-state",
"aether-testkit",
"async-stream",
@@ -5273,6 +5280,7 @@ dependencies = [
"bytes",
"futures-core",
"futures-sink",
"futures-util",
"pin-project-lite",
"tokio",
]
+1 -1
View File
@@ -44,5 +44,5 @@ EXPOSE 8084
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
USER 65532:65532
USER 0:0
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
+1
View File
@@ -157,4 +157,5 @@ EXPOSE 8084
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/usr/local/bin/aether-gateway", "--healthcheck"]
USER 0:0
ENTRYPOINT ["/usr/local/bin/aether-gateway"]
+1
View File
@@ -156,4 +156,5 @@ EXPOSE 8084
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
USER 0:0
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
+16 -76
View File
@@ -50,82 +50,10 @@ chmod 600 .env
./generate_keys.sh
# 编辑 .env 设置 ADMIN_PASSWORD
# 3. 首次部署 / 更新 (从以下部署形态任选其一)
# Postgres + Redis (推荐)
# 3. Docker 部署 / 更新(PostgreSQL + Redis)
docker compose pull && docker compose up -d
# Single Node:同样使用 PostgreSQL + Redis,无需挂载本地数据库文件
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`。如需使用其他身份,可在 `.env` 中设置非零的 `AETHER_CONTAINER_UID` / `AETHER_CONTAINER_GID`。数据库使用独立 PostgreSQL 容器和 named volume,不再需要调整应用数据库目录的权限。
### 一键更新
Docker Compose 部署后,可在部署目录直接执行:
```bash
./update.sh
```
`update.sh` 会拉取最新 `app` 镜像并重建 `app` 容器,Docker named volumes、`./data` 和 `./logs` 不会被删除。Single Node 部署也可显式指定:
```bash
./update.sh --mode single-node
```
现在仅支持 PostgreSQL。标准和单节点 Docker Compose 均部署 PostgreSQL + Redis;原生 systemd / launchd 安装需要显式提供 PostgreSQL `DATABASE_URL`,例如 `DATABASE_URL=postgresql://user:password@host:5432/aether`。旧数据库不会自动迁移或清空。升级时保留原有 PostgreSQL 密码、`JWT_SECRET_KEY` 和 `ENCRYPTION_KEY`,不要重新生成整个 `.env`。
仓库自带的 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/<version>`,再切换 `/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 镜像均固定到多架构 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` 调整。
标准和 Single Node Docker Compose 均使用 Docker named volume 存放 PostgreSQL 数据。
如果是本地源码构建镜像的部署,继续使用:
```bash
./deploy.sh
```
如果要在本机联调“管理后台在线更新”本身,可启动仓库内置的 release-layout 测试环境:
```bash
docker compose -f docker-compose.release-local.yml up -d --build
```
这套环境会用当前源码构建一个本地测试镜像,但编译为 `release` 类型,并默认伪装成 `v0.7.0`,这样后台会按正式发布版逻辑开放“立即更新”。默认监听 `http://127.0.0.1:18085`,数据目录使用 `./data-release-local`;日志默认走 `docker logs`,不会影响你正在跑的源码构建容器。
如果这套容器在 `prepare-update` 时访问 GitHub 失败,而你本机是通过代理出网,请在 `.env` 里把 `AETHER_UPDATE_PROXY_URL` 写成宿主机地址,例如 `http://host.docker.internal:7890`;容器内的 `127.0.0.1` 指向容器自身,不是宿主机。
如果想重置这套联调环境(包括 `/opt/aether/current` 和已下载的历史版本),执行:
```bash
docker compose -f docker-compose.release-local.yml down -v
```
可选变量:
- `AETHER_RELEASE_LOCAL_VERSION`:本地联调镜像对外声明的当前版本,默认 `v0.7.0`
- `AETHER_RELEASE_LOCAL_PORT`:本地联调端口,默认 `18085`
- `LOCAL_RELEASE_APP_IMAGE`:本地联调镜像名,默认 `aether-app:release-local`
### 一键安装(PostgreSQL + Redis)
```bash
@@ -134,7 +62,9 @@ cd Aether
curl -fsSL https://raw.githubusercontent.com/fawney19/Aether/main/install.sh | sudo bash -s -- --mode compose
```
原生 Linux systemd / macOS launchd 安装需先准备 PostgreSQL,将连接串通过 `DATABASE_URL` 传给安装进程,并选择 `--mode single-node`;不再自动创建本地数据库文件。
正式版和 Nightly 自动构建仅提供 Linux `amd64` / `arm64` 二进制包,Docker 镜像同样支持这两种架构。macOS 用户可使用 Docker 或自行从源码构建;安装脚本保留对历史 macOS 制品的兼容。独立 Aether Tunnel 的多平台发行不受此调整影响。
原生 Linux systemd 安装需先准备 PostgreSQL,将连接串通过 `DATABASE_URL` 传给安装进程,并选择 `--mode single-node`;不再自动创建本地数据库文件。
### Nightly(每日 main 构建)
@@ -193,9 +123,17 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙
- `APP_PORT`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}`
- `DATABASE_URL`:PostgreSQL 连接串,例如 `postgresql://USER:PASSWORD@HOST:5432/aether`
- `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时 PostgreSQL 按每核 `4` 条自动推导,总池范围为 `32-100`。该预算按进程计算,多实例部署应按数据库连接上限显式分配
- `AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_TIMEOUT_MS` / `AETHER_GATEWAY_DATA_POSTGRES_LOCK_TIMEOUT_MS`:普通数据库连接的单条 SQL / 锁等待期限,默认 `30000` / `3000` 毫秒,显式 `0` 关闭;不是整个事务总期限。迁移与历史 backfill 使用独立连接放宽,事务可通过局部设置覆盖
- `AETHER_USAGE_EVENT_CAPTURE_MEMORY_BUDGET_BYTES`:usage 诊断正文共享预算,默认 `134217728`(128 MiB),按 JSON 堆内存估算,覆盖进入终态队列的 seed、Redis 解码后的事件、数据库写入 DTO 及其正文副本。额度不足或显式 `0` 时先保留计费事实,再舍弃诊断正文;已有清空或禁用状态保持不变,其余标记截断。预算随正文保留到释放,后台构建或压缩不会因调用方取消而提前归还额度。该额度不覆盖原始 Redis 批次、解码临时分配、序列化及压缩结果、协议观察缓冲或进程总内存;可通过 `usage_runtime_event_capture_memory_*` 指标观察
- `AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES`:新增 usage 队列消息的完整 JSON payload 上限,默认 `1048576`(1 MiB),按序列化后的 UTF-8 字节计算,显式 `0` 非法。超限先保留计费事实并舍弃诊断字段;仍超限或无法保留计费语义时拒绝入队,终态消息尝试受限数据库落库,失败则明确失败,不继续 Redis 重试。该限制不覆盖存量 Redis 消息、整个读取批次、DLQ 或进程总内存。`usage_runtime_queue_payload_*` 导出上限及进程级降级、拒绝编码尝试次数,包含入队和重试预校验,不代表唯一事件数;`usage_runtime_enqueue_retry_permanent_failure_total` 记录永久输入错误导致的重试拒绝或终止
- `AETHER_USAGE_QUEUE_READ_PAYLOAD_BUDGET_BYTES` / `AETHER_USAGE_QUEUE_READ_BATCH_PAYLOAD_BYTES`:usage worker 读取和重领共用的进程级逻辑 payload 预留,默认总额 `134217728`(128 MiB)、单批目标 `8388608`(8 MiB)。按当前 `QUEUE_PAYLOAD_MAX_BYTES` 推导实际 COUNT,默认最多读取 8 条,自动扩容使用实际 COUNT 判断批次是否读满。预留覆盖读取、整批处理和确认,额度不足等待;取消/失败释放。单批目标至少允许一条,当前 payload 上限大于总额时读取报配置错误。`0` 或非法值回退默认,过大值收敛到约 4 GiB 的有效总额。收到消息后按全部字段值长度缩减多余预留;历史消息、其他生产者使用更高上限或额外字段可能超出估算,仍继续原计费流程并记录 `usage_runtime_queue_read_oversized_*`。`usage_runtime_queue_read_*` 同时导出预留、等待与累计字段字节;该预留不是 RESP 解码、连接缓冲容量、字段结构、诊断 JSON、DLQ 或进程 RSS 的硬上限,旧公开 Vec 读取接口不携带处理阶段预留
- `AETHER_USAGE_DLQ_ENCODING_BUDGET_BYTES` / `AETHER_USAGE_DLQ_ENCODING_MAX_JOBS`:死信原文和 JSON 编码独立共享预留,默认 `67108864`(64 MiB)、最多 `4` 个后台编码及写入任务。根据原始字段、ID、错误字符串及 JSON 最坏 6 倍转义一次预留;预算占满或单条超总额时立即失败,worker 保留原消息等待重领,不截断账务原文。编码失败会继续处理同批其他消息,只确认成功项,批次末尾仍报告失败;存储转移失败则停止该批后续处理。取消编码等待不会提前归还仍在后台使用的额度。`0`/非法值回退默认,bytes 最大约 4 GiB,jobs 最大 128;超大存量消息可能需要调高总额后恢复。`usage_runtime_dlq_encoding_*` 导出额度、在途任务、拒绝和编码尝试次数;不包含字段结构、字符串额外容量、Redis 命令/连接副本或进程 RSS。内置 Redis/Memory worker 将死信追加、源 ACK 和删除作为一次原子转移,同一源 stream、消费组及 pending ID 的并发或重试只追加一次;Redis 要求 7+ 及 `EVAL/TYPE/XPENDING/XADD/XACK/XDEL` 权限,Cluster 两键须同 slot(当前默认键不自动迁移)。源和 DLQ 不能同名。源已不在 PEL 时不宣称已归档;外部 ACK/trim/delete 及多消费组仍有原来的删除语义。公开 `push_dead_letter` 仍为追加接口,未实现新原子 trait 方法的外部后端沿用追加后 ACK,仍可能重复归档
- `AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS`:单实例请求并发上限;未配置时按 CPU 自动推导(基础范围 `512-65536`),低文件描述符预算时会进一步下调
- `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB`
- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:可选的请求体完整读取超时;默认或显式设为 `0` 时关闭,非零值限制在 `1000-600000ms`
- `AETHER_GATEWAY_MAX_HTTP_CONNECTIONS`:二进制入口全部监听分片共用的入站 TCP 连接上限,包含握手、空闲 keep-alive 和 HTTP 升级后仍存活的 socket。未设置或 `0` 时使用请求上限与 WebSocket 上限之和;自动及显式值均最多 `65536`,已知 FD soft limit 时进一步限制为 `max(1, (FD - 256) / 2)`。接入后立即尝试取得额度,满额时关闭新连接,不创建 HTTP 处理任务、不等待额度,不返回 HTTP 状态码;取消、解析失败和连接释放归还,WebSocket 升级不会提前归还。HTTP/2 多流共用一个 TCP 许可,原请求和 WebSocket 准入仍独立有效。`gateway_http_connections_*` 导出配置上限、当前数、高水位、拒绝数及 accept 错误数。该限制不包含 kernel backlog、上游、Redis 或数据库连接,也不是整个进程 FD/内存硬上限。临时 accept 错误重试,资源类错误退避一秒后重试,避免单次错误停止监听
- `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB`;压缩和未知长度上传按实际缓冲增长申请额度,解压时计入同时存活的输入和输出。额度不足返回 `503`;接近单请求上限的压缩上传需要为输入和解压输出预留额外预算
- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:请求体完整读取超时,默认 `120000ms`;显式设为 `0` 时关闭,非零值限制在 `1000-600000ms`
- `AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS`:上游流首包后的空闲超时,默认 `300000ms`;请求执行配置中的 `read_ms` 优先,显式 `0` 关闭对应超时。网关生成的 keepalive 不会重置计时
- `AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES`:进程内流式响应诊断捕获的共享字节预算,默认 `134217728`(128 MiB);包含 provider/client 捕获容量和扩容时的新旧分配。额度不足时仅截断审计副本,显式 `0` 关闭此类捕获;协议解析、客户端传输和计费观察继续执行。该预算不包含协议解析缓冲、终态编码及 usage 队列副本,不是进程总内存上限
- `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`;满载时拒绝新事件,避免控制面故障导致无界内存增长
@@ -216,6 +154,8 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙
- `RUST_LOG`:Rust 日志过滤,例如 `aether_gateway=info`、`aether_gateway=debug,sqlx=warn`
- `DB_PASSWORD` / `REDIS_PASSWORD`:Docker Compose 后端密码,首次安装时分别随机生成;手工部署必须替换示例占位值,不要互相复用
运行日志由独立后台线程写入 stdout 和文件,每个输出队列最多 4096 条、保留正文最多 8 MiB(包含正在写入的记录),单条最多 256 KiB。队列满、正文预算不足或单条超限时整条丢弃,不等待日志设备;`Both` 两个输出独立降级。`logging_stdout_*` 和 `logging_file_*` 指标记录丢弃和写入错误,网关指标沿用其命名空间前缀。正常退出时日志最多等待 2 秒排空;这不是请求优雅排空或整个进程退出期限。日志格式化仍在调用线程执行,日志预算不包含格式化临时内存,运行日志也不能作为可靠计费账本。
### S3 备份离线恢复
先从 S3 下载完整的 `.json.zst.aes256gcm` 对象,再使用原始的完整 S3 object key 做认证解密。恢复工具只验证并输出本地 JSON,不会直接写数据库;数据库导入仍应在维护窗口通过管理端完成。
+1
View File
@@ -62,6 +62,7 @@ flate2.workspace = true
futures-util.workspace = true
hmac.workspace = true
http.workspace = true
http-body = "1"
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"] }
@@ -71,8 +71,13 @@ struct Args {
distributed_request_command_timeout_ms: u64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
let _ = rustls::crypto::ring::default_provider().install_default();
init_service_runtime(ServiceRuntimeConfig::new(
@@ -88,8 +88,13 @@ struct Args {
distributed_request_command_timeout_ms: u64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_service_runtime(ServiceRuntimeConfig::new(
"aether-tunnel-standalone",
"aether_gateway=info",
@@ -996,7 +996,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
);
if page_is_exact_auth_api_key_concurrency_limited(&page) {
if self.wait_for_auth_api_key_concurrency_retry().await {
if self.wait_for_auth_api_key_concurrency_retry().await? {
continue;
}
self.persist_final_auth_api_key_concurrency_skips(page.skipped_candidates)
@@ -1087,20 +1087,23 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
}
}
async fn wait_for_auth_api_key_concurrency_retry(&mut self) -> bool {
async fn wait_for_auth_api_key_concurrency_retry(&mut self) -> Result<bool, GatewayError> {
let now = Instant::now();
let deadline = *self
.auth_api_key_concurrency_wait_deadline
.get_or_insert(now + AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET);
if now >= deadline {
return false;
if !crate::scheduler::candidate::wait_for_auth_api_key_concurrency_retry(
self.state.app(),
Some(&self.auth_snapshot),
deadline,
AUTH_API_KEY_CONCURRENCY_RETRY_DELAY,
)
.await?
{
return Ok(false);
}
let sleep_duration =
AUTH_API_KEY_CONCURRENCY_RETRY_DELAY.min(deadline.saturating_duration_since(now));
tokio::time::sleep(sleep_duration).await;
self.page_cursor.restart_scan();
true
Ok(true)
}
async fn persist_final_auth_api_key_concurrency_skips(
@@ -2289,6 +2292,103 @@ mod tests {
candidate
}
#[tokio::test]
async fn auth_concurrency_wait_paged_scan_retries_once_at_original_deadline() {
let now = current_unix_ms();
let active = serde_json::from_value(json!({
"id": "active-candidate",
"request_id": "active-request",
"api_key_id": "api-key-1",
"candidate_index": 0,
"retry_index": 0,
"status": "pending",
"is_cached": false,
"created_at_unix_ms": now,
"started_at_unix_ms": now
}))
.expect("active candidate should build");
let repository = Arc::new(InMemoryRequestCandidateRepository::seed([active]));
let app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
GatewayDataState::with_request_candidate_repository_for_tests(repository),
);
let mut auth_snapshot = sample_auth_snapshot();
auth_snapshot.api_key_concurrent_limit = Some(1);
let page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
&crate::system_features::ModelDirectivePolicySnapshot::default(),
"openai:chat",
"gpt-5",
None,
true,
None,
&auth_snapshot,
None,
None,
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
true,
Some("trace-auth-wait"),
)
.await;
let mut cursor = RequestedModelAttemptPageCursor {
state: PlannerAppState::new(&app),
trace_id: "trace-auth-wait".to_string(),
client_api_format: "openai:chat".to_string(),
requested_model: "gpt-5".to_string(),
auth_snapshot,
client_session_affinity: None,
required_capabilities: None,
routing_policy: None,
sticky_session_token: None,
request_auth_channel: None,
skipped_user_id: "user-1".to_string(),
skipped_api_key_id: "api-key-1".to_string(),
skipped_required_capabilities: None,
skipped_error_context: "test auth wait",
record_runtime_miss_diagnostic: false,
resolution_mode: LocalCandidateResolutionMode::Standard,
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::new(),
skipped_provider_ids: BTreeSet::new(),
skipped_endpoint_ids: BTreeSet::new(),
skipped_credential_ids: BTreeSet::new(),
candidate_count: 0,
next_candidate_index: 0,
remembered_affinity: false,
scheduler_cache_affinity_enabled: false,
auth_api_key_concurrency_wait_deadline: None,
deferred_error: None,
};
let started = Instant::now();
let mut scan_restarts = 0;
while cursor
.wait_for_auth_api_key_concurrency_retry()
.await
.expect("auth wait should succeed")
{
scan_restarts += 1;
}
assert_eq!(
scan_restarts, 1,
"blocked polls must not restart page scans"
);
assert!(started.elapsed() >= AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET);
let original_deadline = cursor.auth_api_key_concurrency_wait_deadline;
assert!(!cursor
.wait_for_auth_api_key_concurrency_retry()
.await
.expect("expired auth wait should succeed"));
assert_eq!(
cursor.auth_api_key_concurrency_wait_deadline,
original_deadline
);
}
#[tokio::test]
async fn pool_group_keys_are_not_persisted_as_available_before_attempt() {
let repository = Arc::new(InMemoryRequestCandidateRepository::default());
@@ -194,7 +194,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncA
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -246,7 +246,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt>
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -118,7 +118,7 @@ pub(crate) fn build_local_execution_report_context(
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) {
if let Ok(value) = serde_json::to_value(&policy.execution_policy) {
extra_fields.insert(ROUTING_EXECUTION_POLICY_REPORT_FIELD.to_string(), value);
}
}
@@ -179,7 +179,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -224,7 +224,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -257,7 +257,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -302,7 +302,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -109,7 +109,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -182,7 +182,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -232,7 +232,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttempt
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -124,7 +124,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttem
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -97,7 +97,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSo
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -729,117 +729,6 @@ fn update_normalization_codex_capabilities_digest(
update_normalization_string_vec_digest(digest, &capabilities.supported_service_tiers);
}
#[cfg(test)]
mod continuation_fingerprint_tests {
use http::HeaderValue;
use serde_json::json;
use super::ResponsesWebSocketBodyNormalization;
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
#[test]
fn normalization_fingerprint_is_stable_for_json_object_key_order() {
let first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}, "x": 1}));
let second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"x": 1, "reasoning": {"effort": "high"}}));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_changes_with_effective_contract() {
let base = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let changed_policy = base.clone().with_reasoning_replay_policy_for_tests(
OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
assert_ne!(
base.continuation_fingerprint(),
changed_policy.continuation_fingerprint()
);
let changed_patch = base
.clone()
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "low"}}));
assert_ne!(
base.continuation_fingerprint(),
changed_patch.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_ignores_unrelated_volatile_request_headers() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "x-contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
first
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-1"));
first
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-1"));
let mut second = first.clone();
second
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-2"));
second
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-2"));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"headers that no body-rule condition reads must not invalidate a persisted continuation"
);
}
#[test]
fn normalization_fingerprint_tracks_headers_used_by_body_rule_conditions() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "X-Contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules.clone());
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
let mut second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
second
.request_headers
.insert("x-contract", HeaderValue::from_static("disabled"));
assert_ne!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"a header that controls an effective body-rule condition remains part of the contract"
);
}
}
/// Builds one upstream decision for a Responses WebSocket turn. The session
/// reuses this decision for same-model turns and invokes the planner again when
/// a later `response.create` changes the public model.
@@ -1058,3 +947,114 @@ async fn release_responses_websocket_planning_lease(
}
}
}
#[cfg(test)]
mod continuation_fingerprint_tests {
use http::HeaderValue;
use serde_json::json;
use super::ResponsesWebSocketBodyNormalization;
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
#[test]
fn normalization_fingerprint_is_stable_for_json_object_key_order() {
let first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}, "x": 1}));
let second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"x": 1, "reasoning": {"effort": "high"}}));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_changes_with_effective_contract() {
let base = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let changed_policy = base.clone().with_reasoning_replay_policy_for_tests(
OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
assert_ne!(
base.continuation_fingerprint(),
changed_policy.continuation_fingerprint()
);
let changed_patch = base
.clone()
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "low"}}));
assert_ne!(
base.continuation_fingerprint(),
changed_patch.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_ignores_unrelated_volatile_request_headers() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "x-contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
first
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-1"));
first
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-1"));
let mut second = first.clone();
second
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-2"));
second
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-2"));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"headers that no body-rule condition reads must not invalidate a persisted continuation"
);
}
#[test]
fn normalization_fingerprint_tracks_headers_used_by_body_rule_conditions() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "X-Contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules.clone());
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
let mut second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
second
.request_headers
.insert("x-contract", HeaderValue::from_static("disabled"));
assert_ne!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"a header that controls an effective body-rule condition remains part of the contract"
);
}
}
@@ -166,7 +166,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -216,7 +216,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStream
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -1,9 +1,7 @@
use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate};
use std::time::Duration;
use tokio::time::Instant;
use super::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_secs;
use crate::constants::{
API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS, API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS,
};
@@ -97,11 +95,13 @@ impl<'a> PlannerAppState<'a> {
),
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));
let wait_deadline = Instant::now() + wait_timeout;
let mut attempt_now_unix_secs = now_unix_secs;
loop {
crate::scheduler::candidate::select_with_auth_concurrency_wait(
self.app(),
auth_snapshot,
now_unix_secs,
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS),
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS),
|attempt_now_unix_secs| async move {
let result = crate::scheduler::candidate::list_selectable_candidates_with_skip_reasons_for_request_operation(
self.app().data.as_ref(),
self.app(),
@@ -118,21 +118,13 @@ impl<'a> PlannerAppState<'a> {
)
.await?;
if !crate::scheduler::candidate::is_exact_all_skipped_by_auth_limit(
let auth_limit_blocked = crate::scheduler::candidate::is_exact_all_skipped_by_auth_limit(
&result.0, &result.1,
) {
return Ok(result);
}
let now = Instant::now();
if now >= wait_deadline {
return Ok(result);
}
let remaining = wait_deadline.duration_since(now);
tokio::time::sleep(wait_interval.min(remaining)).await;
attempt_now_unix_secs = current_unix_secs();
}
);
Ok((result, auth_limit_blocked))
},
)
.await
}
#[allow(clippy::too_many_arguments)]
@@ -178,13 +170,14 @@ impl<'a> PlannerAppState<'a> {
now_unix_secs: u64,
ordering_config: SchedulerOrderingConfig,
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, 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));
let wait_deadline = Instant::now() + wait_timeout;
let mut attempt_now_unix_secs = now_unix_secs;
loop {
let (result, auth_limit_blocked) = crate::scheduler::candidate::list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal(
crate::scheduler::candidate::select_with_auth_concurrency_wait(
self.app(),
auth_snapshot,
now_unix_secs,
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS),
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS),
|attempt_now_unix_secs| {
crate::scheduler::candidate::list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal(
self.app().data.as_ref(),
self.app(),
candidate_api_format,
@@ -195,20 +188,8 @@ impl<'a> PlannerAppState<'a> {
attempt_now_unix_secs,
ordering_config,
)
.await?;
if !auth_limit_blocked {
return Ok(result);
}
let now = Instant::now();
if now >= wait_deadline {
return Ok(result);
}
let remaining = wait_deadline.duration_since(now);
tokio::time::sleep(wait_interval.min(remaining)).await;
attempt_now_unix_secs = current_unix_secs();
}
},
)
.await
}
}
+14 -14
View File
@@ -23,7 +23,6 @@ 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(Clone)]
pub(crate) struct BarkPushConfig {
@@ -208,19 +207,20 @@ async fn build_bark_push_client_and_url(
let port = push_url
.port_or_known_default()
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少端口".to_string()))?;
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
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::<Vec<_>>()
};
let addresses = aether_http::lookup_host_with_limits(
host.as_str(),
port,
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
)
.await
.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "Bark 服务器 DNS 解析超时",
std::io::ErrorKind::InvalidData => "Bark 服务器 DNS 解析返回过多地址",
_ => "Bark 服务器 DNS 解析失败",
};
GatewayError::Internal(message.to_string())
})?;
let allow_benchmarking_ip = push_url.scheme() == "https"
&& push_url.port_or_known_default() == Some(443)
&& host.eq_ignore_ascii_case("api.day.app");
@@ -496,6 +496,7 @@ fn access_for_route(method: &http::Method, decision: &GatewayControlDecision) ->
Some("admin:endpoints_manage"),
Some(
"reveal_key"
| "reveal_endpoint_rules"
| "export_key"
| "create_provider_key"
| "update_key"
@@ -1490,6 +1491,12 @@ mod tests {
fn plaintext_credential_reads_require_admin_permission() {
let read_only_permissions = read_only_management_token_permissions();
let cases = [
(
"admin:endpoints_manage",
"reveal_endpoint_rules",
None,
"admin:endpoints_manage:admin",
),
(
"admin:endpoints_manage",
"reveal_key",
@@ -302,6 +302,19 @@ pub(super) fn classify_admin_endpoints_family_route(
"admin:endpoints_manage",
false,
))
} else if method == http::Method::GET
&& normalized_path
.strip_prefix("/api/admin/endpoints/")
.and_then(|path| path.strip_suffix("/rules/reveal"))
.is_some_and(|endpoint_id| !endpoint_id.is_empty() && !endpoint_id.contains('/'))
{
Some(classified(
"admin_proxy",
"endpoints_manage",
"reveal_endpoint_rules",
"admin:endpoints_manage",
false,
))
} else if method == http::Method::GET
&& normalized_path.starts_with("/api/admin/endpoints/")
&& !normalized_path.starts_with("/api/admin/endpoints/health/")
@@ -381,6 +381,28 @@ fn classifies_admin_get_endpoint_as_admin_proxy_route() {
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_admin_reveal_endpoint_rules_as_admin_proxy_route() {
let headers = headers(&[]);
let uri: Uri = "/api/admin/endpoints/endpoint-1/rules/reveal"
.parse()
.expect("uri should parse");
let decision =
classify_control_route(&http::Method::GET, &uri, &headers).expect("route should classify");
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
assert_eq!(decision.route_family.as_deref(), Some("endpoints_manage"));
assert_eq!(
decision.route_kind.as_deref(),
Some("reveal_endpoint_rules")
);
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("admin:endpoints_manage")
);
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_admin_create_endpoint_as_admin_proxy_route() {
let headers = http::HeaderMap::new();
@@ -81,6 +81,19 @@ impl GatewayDataState {
}
}
pub(crate) async fn list_recent_runtime_request_candidates(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_reader {
Some(repository) => repository
.list_recent_runtime(limit)
.await
.map(sanitize_request_candidate_rows),
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_finalized_request_candidates_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
@@ -41,16 +41,9 @@ fn usage_request_record_level_from_value(value: Option<&Value>) -> UsageRequestR
return UsageRequestRecordLevel::Basic;
};
if value.eq_ignore_ascii_case("basic")
|| value.eq_ignore_ascii_case("base")
|| value.eq_ignore_ascii_case("headers")
|| value.eq_ignore_ascii_case("minimal")
|| value.eq_ignore_ascii_case("none")
{
UsageRequestRecordLevel::Basic
if value.eq_ignore_ascii_case("full") {
UsageRequestRecordLevel::Full
} else {
// 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
}
}
@@ -501,7 +494,7 @@ mod tests {
}
#[tokio::test]
async fn usage_runtime_access_disables_full_http_capture() {
async fn usage_runtime_access_honors_explicit_full_http_capture() {
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
"request_record_level".to_string(),
json!("full"),
@@ -511,7 +504,32 @@ mod tests {
.await
.expect("request record level should read");
assert_eq!(level, UsageRequestRecordLevel::Basic);
assert_eq!(level, UsageRequestRecordLevel::Full);
}
#[tokio::test]
async fn usage_runtime_access_honors_legacy_full_without_overriding_current_config() {
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
"request_log_level".to_string(),
json!(" FULL "),
)]);
assert_eq!(
UsageRuntimeAccess::request_record_level(&state)
.await
.unwrap(),
UsageRequestRecordLevel::Full
);
let state = state.with_system_config_values_for_tests([
("request_log_level".to_string(), json!("full")),
("request_record_level".to_string(), json!("basic")),
]);
assert_eq!(
UsageRuntimeAccess::request_record_level(&state)
.await
.unwrap(),
UsageRequestRecordLevel::Basic
);
}
#[tokio::test]
@@ -1593,6 +1593,19 @@ impl GatewayDataState {
}
}
pub(crate) async fn read_request_usage_body_payload(
&self,
body_ref: &str,
) -> Result<
Option<aether_data_contracts::repository::usage::StoredUsageBodyPayload>,
DataLayerError,
> {
match &self.usage_reader {
Some(repository) => repository.read_body_payload(body_ref).await,
None => Ok(None),
}
}
pub(crate) async fn list_usage_audits(
&self,
query: &UsageAuditListQuery,
+25 -2
View File
@@ -19,9 +19,14 @@ use aether_data::repository::management_tokens::{
};
use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode};
use aether_data::repository::proxy_nodes::{ProxyNodeReadRepository, ProxyNodeWriteRepository};
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
use aether_data::repository::users::{
InMemoryUserReadRepository, StoredUserAuthRecord, UserReadRepository,
};
use aether_data_contracts::repository::routing_profiles::{
StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion,
};
use aether_routing_core::RoutingGroupConfig;
use sha2::{Digest, Sha256};
use super::{GatewayDataConfig, GatewayDataState};
@@ -142,6 +147,24 @@ impl GatewayDataState {
provider_catalog_repository;
let usage_reader: Arc<dyn UsageReadRepository> = usage_repository.clone();
let usage_writer: Arc<dyn UsageWriteRepository> = usage_repository;
let routing_groups = Arc::new(InMemoryRoutingGroupRepository::seed(
[StoredRoutingGroup {
id: "system-default".to_string(),
name: "system-default".to_string(),
description: Some("pressure harness 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: 1,
updated_at: 1,
published_at: Some(1),
}],
std::iter::empty::<StoredRoutingGroupBinding>(),
std::iter::empty::<StoredRoutingGroupVersion>(),
));
Self {
config: GatewayDataConfig::disabled().with_encryption_key(encryption_key),
@@ -174,8 +197,8 @@ impl GatewayDataState {
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
routing_group_reader: Some(routing_groups.clone()),
routing_group_writer: Some(routing_groups),
usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer),
user_reader: None,
@@ -35,7 +35,6 @@ use crate::ai_serving::{
SkippedLocalExecutionCandidate,
};
use crate::clock::current_unix_ms;
use crate::handlers::shared::provider_pool::read_admin_provider_pool_runtime_state;
use crate::handlers::shared::provider_pool::{
admin_provider_pool_cache_affinity_enabled, admin_provider_pool_config_from_config_value,
};
@@ -44,6 +43,9 @@ use crate::handlers::shared::provider_pool::{
read_admin_provider_pool_key_cooldown_reason, AdminProviderPoolConfig,
AdminProviderPoolRuntimeState, AdminProviderPoolSchedulingPreset,
};
use crate::handlers::shared::provider_pool::{
read_provider_pool_scheduling_runtime_state, read_provider_pool_sticky_bound_key_id,
};
use crate::handlers::shared::{parse_catalog_auth_config_json, provider_key_health_summary};
use crate::maintenance::spawn_pool_quota_probe_replenish_for_request;
use crate::orchestration::LocalExecutionCandidateMetadata;
@@ -141,7 +143,7 @@ async fn schedule_pool_page_candidates(
AdminProviderPoolRuntimeState::default()
} else {
let runtime_started_at = std::time::Instant::now();
let runtime = read_admin_provider_pool_runtime_state(
let runtime = read_provider_pool_scheduling_runtime_state(
state.app().runtime_state.as_ref(),
provider_id.as_str(),
&key_ids,
@@ -982,15 +984,13 @@ impl<'a> PoolKeyCursor<'a> {
if !admin_provider_pool_cache_affinity_enabled(&pool_config) {
return None;
}
let runtime = read_admin_provider_pool_runtime_state(
let sticky_key_id = read_provider_pool_sticky_bound_key_id(
self.state.app().runtime_state.as_ref(),
self.group.candidate.provider_id.as_str(),
&[],
&pool_config,
self.sticky_session_token.as_deref(),
)
.await;
let sticky_key_id = runtime.sticky_bound_key_id?;
.await?;
if self
.routing_overlay
.as_ref()
@@ -2028,7 +2028,9 @@ mod tests {
};
use crate::data::GatewayDataState;
use crate::handlers::shared::provider_pool::{
admin_provider_pool_cache_affinity_enabled, record_admin_provider_pool_error,
admin_provider_pool_cache_affinity_enabled, admin_provider_pool_config_from_config_value,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
record_admin_provider_pool_error, record_admin_provider_pool_success,
AdminProviderPoolRuntimeState,
};
use crate::orchestration::LocalExecutionCandidateMetadata;
@@ -2060,6 +2062,111 @@ mod tests {
use std::collections::{BTreeMap, BTreeSet, VecDeque};
use std::sync::Arc;
#[tokio::test]
async fn scheduling_runtime_preserves_pool_ranking_and_cost_rejections() {
let runtime = aether_runtime_state::RuntimeState::memory(
aether_runtime_state::MemoryRuntimeStateConfig::default(),
);
let writer_config = admin_provider_pool_config_from_config_value(Some(&json!({
"pool_advanced": {
"cost_limit_per_key_tokens": 100,
"scheduling_presets": [
{"preset": "cache_affinity", "enabled": true},
{"preset": "latency_first", "enabled": true}
]
}
})))
.expect("writer pool config");
for (key_id, cost, latency) in [("key-a", 100, 10), ("key-b", 20, 100)] {
record_admin_provider_pool_success(
&runtime,
"provider-pool",
key_id,
&writer_config,
Some(key_id),
cost,
Some(latency),
)
.await;
}
let key_ids = vec!["key-a".to_string(), "key-b".to_string()];
for (preset, cost_limit) in [
("cache_affinity", None),
("priority_first", None),
("latency_first", None),
("cost_first", None),
("quota_balanced", None),
("latency_first", Some(100)),
] {
let provider_config = json!({
"pool_advanced": {
"cost_limit_per_key_tokens": cost_limit,
"scheduling_presets": [{"preset": preset, "enabled": true}]
}
});
let pool_config = admin_provider_pool_config_from_config_value(Some(&provider_config))
.expect("reader pool config");
let admin = read_admin_provider_pool_runtime_state(
&runtime,
"provider-pool",
&key_ids,
&pool_config,
Some("key-a"),
)
.await;
let scheduling = read_provider_pool_scheduling_runtime_state(
&runtime,
"provider-pool",
&key_ids,
&pool_config,
Some("key-a"),
)
.await;
let run = |snapshot| {
let candidates = key_ids
.iter()
.map(|key_id| {
sample_eligible_candidate(
"provider-pool",
"endpoint-1",
key_id,
10,
Some(provider_config.clone()),
)
})
.collect();
let (scheduled, skipped) = apply_local_execution_pool_scheduler_with_runtime_map(
candidates,
&BTreeMap::from([("provider-pool".to_string(), snapshot)]),
&BTreeMap::new(),
);
(
scheduled
.into_iter()
.map(|item| item.candidate.key_id)
.collect::<Vec<_>>(),
skipped
.into_iter()
.map(|item| (item.candidate.key_id, item.skip_reason))
.collect::<Vec<_>>(),
)
};
let expected = run(admin);
let actual = run(scheduling);
assert_eq!(
actual, expected,
"preset: {preset}, cost limit: {cost_limit:?}"
);
if cost_limit.is_some() {
assert_eq!(actual.0, vec!["key-b"]);
assert_eq!(
actual.1,
vec![("key-a".to_string(), "pool_cost_limit_reached")]
);
}
}
}
#[test]
fn pool_scheduler_groups_interleaved_candidates_and_reorders_internal_keys() {
let pool_first = sample_eligible_candidate(
+194 -43
View File
@@ -136,14 +136,16 @@ pub(crate) async fn send_smtp_email(
email: ComposedEmail,
) -> Result<(), GatewayError> {
validate_smtp_delivery_inputs(&config, &email)?;
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email))
let stream = connect_tcp_stream(&config).await?;
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email, stream))
.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))
let stream = connect_tcp_stream(&config).await?;
tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config, stream))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
}
@@ -328,43 +330,58 @@ fn resolve_server_name(host: &str) -> Result<rustls::pki_types::ServerName<'stat
.map_err(|err| GatewayError::Internal(err.to_string()))
}
fn connect_tcp_stream(config: &SmtpDeliveryConfig) -> Result<std::net::TcpStream, GatewayError> {
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::<Vec<_>>();
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;
async fn connect_tcp_stream(
config: &SmtpDeliveryConfig,
) -> Result<std::net::TcpStream, GatewayError> {
connect_tcp_stream_with_dns(
aether_http::lookup_host_with_limits(
&config.host,
config.port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
),
std::time::Duration::from_secs(SMTP_TIMEOUT_SECS),
)
.await
}
async fn connect_tcp_stream_with_dns(
lookup: impl std::future::Future<Output = std::io::Result<Vec<std::net::SocketAddr>>>,
timeout: std::time::Duration,
) -> Result<std::net::TcpStream, GatewayError> {
let stream = tokio::time::timeout(timeout, async {
let addresses = lookup.await.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "smtp DNS resolution timed out",
std::io::ErrorKind::InvalidData => {
"smtp DNS resolution returned too many addresses"
}
_ => "smtp DNS resolution failed",
};
GatewayError::Internal(message.to_string())
})?;
if addresses.is_empty() {
return Err(GatewayError::Internal(
"smtp host did not resolve to an address".to_string(),
));
}
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()),
)
})?;
let attempts = addresses
.into_iter()
.map(|address| Box::pin(tokio::net::TcpStream::connect(address)));
futures_util::future::select_ok(attempts)
.await
.map(|(stream, _)| stream)
.map_err(|error| {
GatewayError::Internal(format!("smtp connection failed ({})", error.kind()))
})
})
.await
.map_err(|_| GatewayError::Internal("smtp DNS or TCP connection timed out".to_string()))??;
let stream = stream
.into_std()
.map_err(|err| GatewayError::Internal(err.to_string()))?;
stream
.set_nonblocking(false)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
stream
.set_read_timeout(Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS)))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
@@ -680,16 +697,15 @@ fn smtp_probe_connection<S: std::io::Read + std::io::Write>(
fn send_smtp_email_blocking(
config: SmtpDeliveryConfig,
email: ComposedEmail,
stream: std::net::TcpStream,
) -> Result<(), GatewayError> {
if config.use_ssl {
let stream = connect_tcp_stream(&config)?;
let tls_stream = wrap_tls_stream(stream, &config.host)?;
let mut reader = std::io::BufReader::new(tls_stream);
let _ = smtp_expect(&mut reader, &[220])?;
return smtp_send_message(&mut reader, &config, &email);
}
let stream = connect_tcp_stream(&config)?;
let mut reader = std::io::BufReader::new(stream);
let _ = smtp_expect(&mut reader, &[220])?;
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
@@ -705,16 +721,17 @@ fn send_smtp_email_blocking(
smtp_deliver_message(&mut reader, &config, &email)
}
fn probe_smtp_connection_blocking(config: SmtpDeliveryConfig) -> Result<(), GatewayError> {
fn probe_smtp_connection_blocking(
config: SmtpDeliveryConfig,
stream: std::net::TcpStream,
) -> Result<(), GatewayError> {
if config.use_ssl {
let stream = connect_tcp_stream(&config)?;
let tls_stream = wrap_tls_stream(stream, &config.host)?;
let mut reader = std::io::BufReader::new(tls_stream);
let _ = smtp_expect(&mut reader, &[220])?;
return smtp_probe_connection(&mut reader, &config);
}
let stream = connect_tcp_stream(&config)?;
let mut reader = std::io::BufReader::new(stream);
let _ = smtp_expect(&mut reader, &[220])?;
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
@@ -778,6 +795,140 @@ mod tests {
assert!(validate_smtp_delivery_inputs(&config(), &email()).is_ok());
}
#[tokio::test]
async fn smtp_connection_deadline_includes_a_stalled_dns_lookup() {
let error = connect_tcp_stream_with_dns(
std::future::pending(),
std::time::Duration::from_millis(5),
)
.await
.expect_err("DNS must not outlive the connection deadline");
assert!(format!("{error:?}").contains("smtp DNS or TCP connection timed out"));
}
#[tokio::test]
async fn smtp_dns_errors_and_empty_answers_fail_without_connecting() {
for (addresses, expected) in [
(Ok(Vec::new()), "smtp host did not resolve to an address"),
(
Err(std::io::Error::other("sensitive-dns-detail")),
"smtp DNS resolution failed",
),
(
Err(std::io::Error::from(std::io::ErrorKind::InvalidData)),
"smtp DNS resolution returned too many addresses",
),
(
Err(std::io::Error::from(std::io::ErrorKind::TimedOut)),
"smtp DNS resolution timed out",
),
] {
let error = connect_tcp_stream_with_dns(
std::future::ready(addresses),
std::time::Duration::from_secs(1),
)
.await
.expect_err("invalid DNS answers must fail before TCP connect");
assert!(format!("{error:?}").contains(expected));
assert!(!format!("{error:?}").contains("sensitive-dns-detail"));
}
}
#[tokio::test]
async fn smtp_connection_tries_answers_beyond_the_old_sixteen_address_limit() {
let unavailable = tokio::net::TcpSocket::new_v4().unwrap();
unavailable.bind("127.0.0.1:0".parse().unwrap()).unwrap();
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let available = listener.local_addr().unwrap();
let mut addresses = vec![unavailable.local_addr().unwrap(); 16];
addresses.push(available);
let stream = connect_tcp_stream_with_dns(
std::future::ready(Ok(addresses)),
std::time::Duration::from_secs(5),
)
.await
.expect("later DNS answers should remain available for fallback");
assert_eq!(stream.peer_addr().unwrap(), available);
assert_eq!(
stream.read_timeout().unwrap(),
Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))
);
}
#[tokio::test]
async fn smtp_probe_and_delivery_use_the_preconnected_stream() {
use tokio::io::{AsyncBufReadExt, AsyncWriteExt};
for deliver in [false, true] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut reader = tokio::io::BufReader::new(stream);
reader
.get_mut()
.write_all(b"220 mock SMTP ready\r\n")
.await
.unwrap();
let mut delivered = false;
loop {
let mut line = String::new();
assert!(reader.read_line(&mut line).await.unwrap() > 0);
let response = if line.starts_with("EHLO ")
|| line.starts_with("MAIL FROM:")
|| line.starts_with("RCPT TO:")
{
&b"250 OK\r\n"[..]
} else if line == "DATA\r\n" {
reader
.get_mut()
.write_all(b"354 End with dot\r\n")
.await
.unwrap();
loop {
line.clear();
assert!(reader.read_line(&mut line).await.unwrap() > 0);
if line == ".\r\n" {
break;
}
}
delivered = true;
&b"250 Accepted\r\n"[..]
} else {
assert_eq!(line, "QUIT\r\n");
reader
.get_mut()
.write_all(b"221 Goodbye\r\n")
.await
.unwrap();
break;
};
reader.get_mut().write_all(response).await.unwrap();
}
assert_eq!(delivered, deliver);
});
let config = SmtpDeliveryConfig {
host: "127.0.0.1".to_string(),
port,
user: None,
password: None,
use_tls: false,
use_ssl: false,
..config()
};
tokio::time::timeout(std::time::Duration::from_secs(5), async {
if deliver {
send_smtp_email(config, email()).await.unwrap();
} else {
probe_smtp_connection(config).await.unwrap();
}
server.await.unwrap();
})
.await
.expect("local SMTP probe and delivery should complete");
}
}
#[test]
fn rejects_authentication_over_plaintext_smtp() {
let mut insecure = config();
@@ -234,7 +234,9 @@ impl Drop for AttemptCancellationGuard {
);
return;
};
let usage_producer = state.usage_runtime.track_producer();
handle.spawn(async move {
let _usage_producer = usage_producer;
settle_cancelled_attempt(state, armed, error_type, error_message).await;
});
}
@@ -426,11 +428,8 @@ mod tests {
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() {
async fn settling_a_dropped_attempt_respects_disabled_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);
@@ -444,8 +443,6 @@ mod tests {
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,
@@ -481,7 +478,10 @@ mod tests {
.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);
assert_eq!(
usage.provider_request_body_state,
Some(UsageBodyCaptureState::Disabled)
);
}
#[tokio::test]
@@ -674,8 +674,10 @@ impl ExecutionAttemptLifecycle {
let billing_void = settlement.billing.is_void();
let usage_runtime = Arc::clone(&state.usage_runtime);
let usage_data = Arc::clone(state.usage_lifecycle_data_state());
let usage_producer = usage_runtime.track_producer();
self.stage_guard
.await_detachable_stage(self.trace_id.as_str(), "usage_terminal", async move {
let _usage_producer = usage_producer;
usage_runtime
.record_stream_terminal(
usage_data.as_ref(),
@@ -68,7 +68,6 @@ 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;
@@ -1334,24 +1333,18 @@ async fn resolve_public_web_image_addrs(
"ChatGPT-Web image URL is missing a port".to_string(),
)
})?;
let resolved = if let Ok(ip) = host.parse::<IpAddr>() {
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::<Vec<_>>()
};
let resolved = aether_http::lookup_host_with_limits(host, port, lookup_timeout)
.await
.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "ChatGPT-Web image URL DNS resolution timed out",
std::io::ErrorKind::InvalidData => {
"ChatGPT-Web image URL DNS resolution returned too many addresses"
}
_ => "ChatGPT-Web image URL DNS resolution failed",
};
ExecutionRuntimeTransportError::UpstreamRequest(message.to_string())
})?;
if resolved.is_empty() {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"ChatGPT-Web image URL DNS resolution returned no addresses".to_string(),
@@ -14,7 +14,7 @@ fn sync_plan_kind_disables_local_candidate_failover(plan_kind: &str) -> bool {
)
}
fn openai_image_success_disables_local_success_failover(
pub(super) fn openai_image_success_disables_local_success_failover(
plan: &ExecutionPlan,
status_code: u16,
) -> bool {
@@ -1036,6 +1036,7 @@ mod tests {
policy,
LocalFailoverPolicy {
max_retries: Some(1),
routing_rules: Default::default(),
max_transfer_count: 0,
max_transfer_timeout_seconds: 0,
stop_status_codes: [503].into_iter().collect(),
@@ -1826,12 +1826,12 @@ async fn fetch_grok_attachment_url(
// 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 public_addrs = public_socket_addrs_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_addr])
.resolve_to_addrs(url.host_str().unwrap_or_default(), &public_addrs)
.build()
.map_err(ExecutionRuntimeTransportError::ClientBuild)?
.get(url.clone())
@@ -1897,10 +1897,10 @@ fn validate_grok_attachment_url(url: &reqwest::Url) -> Result<(), ExecutionRunti
Ok(())
}
async fn public_socket_addr_for_url(
async fn public_socket_addrs_for_url(
url: &reqwest::Url,
) -> Result<std::net::SocketAddr, ExecutionRuntimeTransportError> {
let host = url.host().ok_or_else(|| {
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
let host = url.host_str().ok_or_else(|| {
ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL is missing a host".to_string(),
)
@@ -1910,64 +1910,34 @@ 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::<IpAddr>() {
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));
}
let mut public_addr = None;
let mut resolved_any = false;
for addr in
let addresses =
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(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
public_addr.get_or_insert(addr);
}
if !resolved_any {
})?;
validate_grok_attachment_addresses(addresses)
}
fn validate_grok_attachment_addresses(
addresses: Vec<std::net::SocketAddr>,
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
if addresses.is_empty() {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL DNS resolution returned no addresses".to_string(),
));
}
public_addr.ok_or_else(|| {
ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL has no public address".to_string(),
)
})
if addresses
.iter()
.any(|address| !grok_attachment_ip_is_public(address.ip()))
{
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
Ok(addresses)
}
fn grok_attachment_ip_is_public(ip: IpAddr) -> bool {
@@ -3898,7 +3868,7 @@ mod tests {
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, maximum_base64_len_for_decoded_limit, openai_chat_body,
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addr_for_url,
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addrs_for_url,
set_grok_image_edit_config, validate_grok_attachment_url, GrokAttachmentInput,
GrokCollected, GrokImagineImage, GrokStreamAdapter,
};
@@ -4520,7 +4490,7 @@ mod tests {
] {
let url = reqwest::Url::parse(raw_url).expect("URL should parse");
assert!(
public_socket_addr_for_url(&url).await.is_err(),
public_socket_addrs_for_url(&url).await.is_err(),
"private IPv6 literal should be rejected: {raw_url}"
);
}
@@ -4528,13 +4498,31 @@ mod tests {
let url = reqwest::Url::parse("https://[2606:4700:4700::1111]/attachment")
.expect("URL should parse");
assert_eq!(
public_socket_addr_for_url(&url)
public_socket_addrs_for_url(&url)
.await
.expect("public IPv6 literal should pass"),
"[2606:4700:4700::1111]:443".parse().unwrap()
vec!["[2606:4700:4700::1111]:443".parse().unwrap()]
);
}
#[test]
fn grok_attachment_dns_keeps_all_safe_addresses_for_connection_fallback() {
let addresses = vec![
"[2606:4700:4700::1111]:443".parse().unwrap(),
"8.8.8.8:443".parse().unwrap(),
];
assert_eq!(
super::validate_grok_attachment_addresses(addresses.clone()).unwrap(),
addresses
);
assert!(super::validate_grok_attachment_addresses(Vec::new()).is_err());
for blocked in ["198.18.0.1:443", "127.0.0.1:443", "[fd00::1]:443"] {
let mut mixed = addresses.clone();
mixed.push(blocked.parse().unwrap());
assert!(super::validate_grok_attachment_addresses(mixed).is_err());
}
}
#[test]
fn grok_attachment_url_rejects_credentials_and_fragments_on_every_hop() {
for raw_url in [
@@ -20,6 +20,7 @@ mod response_header_rules;
mod server;
pub(crate) mod stream;
mod stream_pump;
mod stream_read_timeout;
pub(crate) mod submission;
pub(crate) mod sync;
pub(crate) mod transport;
@@ -0,0 +1,251 @@
use std::ops::Deref;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, LazyLock};
const DEFAULT_STREAM_CAPTURE_MEMORY_BUDGET_BYTES: usize = 128 * 1024 * 1024;
const STREAM_CAPTURE_MEMORY_BUDGET_ENV: &str = "AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES";
static STREAM_CAPTURE_BUDGET: LazyLock<Arc<StreamCaptureBudget>> = LazyLock::new(|| {
StreamCaptureBudget::new(
std::env::var(STREAM_CAPTURE_MEMORY_BUDGET_ENV)
.ok()
.and_then(|value| value.trim().parse().ok())
.unwrap_or(DEFAULT_STREAM_CAPTURE_MEMORY_BUDGET_BYTES),
)
});
#[derive(Debug)]
pub(super) struct StreamCaptureBudget {
available: AtomicUsize,
}
impl StreamCaptureBudget {
pub(super) fn new(bytes: usize) -> Arc<Self> {
Arc::new(Self {
available: AtomicUsize::new(bytes),
})
}
fn reserve_up_to(&self, wanted: usize, minimum: usize) -> usize {
let mut available = self.available.load(Ordering::Relaxed);
loop {
let reserved = wanted.min(available);
if reserved < minimum {
return 0;
}
match self.available.compare_exchange_weak(
available,
available - reserved,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return reserved,
Err(current) => available = current,
}
}
}
fn release(&self, bytes: usize) {
self.available.fetch_add(bytes, Ordering::Relaxed);
}
}
/// Only retained diagnostic bytes belong here. Protocol and billing observers
/// must consume the original chunks independently of capture admission.
#[derive(Debug)]
pub(super) struct StreamBodyCapture {
bytes: Vec<u8>,
budget: Arc<StreamCaptureBudget>,
}
impl Default for StreamBodyCapture {
fn default() -> Self {
Self::with_budget(Arc::clone(&STREAM_CAPTURE_BUDGET))
}
}
impl StreamBodyCapture {
pub(super) fn with_budget(budget: Arc<StreamCaptureBudget>) -> Self {
Self {
bytes: Vec::new(),
budget,
}
}
pub(super) fn append(&mut self, chunk: &[u8], limit: usize, truncated: &mut bool) {
if chunk.is_empty() || *truncated {
return;
}
let wanted_len = self.bytes.len().saturating_add(chunk.len()).min(limit);
if wanted_len > self.bytes.capacity() {
// Keep the old allocation charged until its replacement has been
// allocated and copied, including their overlap during growth.
let wanted_capacity = wanted_len
.max(self.bytes.capacity().saturating_mul(2))
.min(limit);
let reserved = self
.budget
.reserve_up_to(wanted_capacity, self.bytes.capacity().saturating_add(1));
if reserved > 0 {
let mut replacement = Vec::new();
if replacement.try_reserve_exact(reserved).is_ok() {
let extra = replacement.capacity().saturating_sub(reserved);
if extra == 0 || self.budget.reserve_up_to(extra, extra) == extra {
replacement.extend_from_slice(&self.bytes);
let old = std::mem::replace(&mut self.bytes, replacement);
let old_capacity = old.capacity();
drop(old);
self.budget.release(old_capacity);
} else {
drop(replacement);
self.budget.release(reserved);
}
} else {
self.budget.release(reserved);
}
}
}
let keep = wanted_len
.min(self.bytes.capacity())
.saturating_sub(self.bytes.len());
self.bytes.extend_from_slice(&chunk[..keep]);
// Once bytes are omitted, never append a later suffix to this prefix.
*truncated = keep < chunk.len();
}
}
impl Deref for StreamBodyCapture {
type Target = [u8];
fn deref(&self) -> &Self::Target {
&self.bytes
}
}
impl Drop for StreamBodyCapture {
fn drop(&mut self) {
let bytes = std::mem::take(&mut self.bytes);
let capacity = bytes.capacity();
drop(bytes);
self.budget.release(capacity);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stream_capture_budget_is_shared_and_released_on_drop() {
let budget = StreamCaptureBudget::new(12);
let mut provider = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut client = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut provider_truncated = false;
let mut client_truncated = false;
provider.append(b"12345678", 64, &mut provider_truncated);
client.append(b"abcdefgh", 64, &mut client_truncated);
assert_eq!(&*provider, b"12345678");
assert_eq!(&*client, b"abcd");
assert!(!provider_truncated);
assert!(client_truncated);
assert_eq!(budget.available.load(Ordering::Relaxed), 0);
drop(provider);
client.append(b"later", 64, &mut client_truncated);
assert_eq!(&*client, b"abcd");
drop(client);
assert_eq!(budget.available.load(Ordering::Relaxed), 12);
}
#[test]
fn stream_capture_budget_charges_capacity_and_reallocation_overlap() {
let budget = StreamCaptureBudget::new(16);
let mut capture = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut truncated = false;
capture.append(b"1234", 64, &mut truncated);
capture.append(b"5", 64, &mut truncated);
assert_eq!(capture.bytes.capacity(), 8);
assert_eq!(budget.available.load(Ordering::Relaxed), 8);
capture.append(b"6789", 64, &mut truncated);
assert_eq!(&*capture, b"12345678");
assert!(truncated);
assert_eq!(budget.available.load(Ordering::Relaxed), 8);
drop(capture);
assert_eq!(budget.available.load(Ordering::Relaxed), 16);
}
#[test]
fn stream_capture_budget_zero_disables_capture_without_allocating() {
let budget = StreamCaptureBudget::new(0);
let mut capture = StreamBodyCapture::with_budget(budget);
let mut truncated = false;
capture.append(b"data", 64, &mut truncated);
assert!(capture.is_empty());
assert_eq!(capture.bytes.capacity(), 0);
assert!(truncated);
}
#[test]
fn stream_capture_budget_exhaustion_uses_existing_spare_capacity() {
let budget = StreamCaptureBudget::new(14);
let mut capture = StreamBodyCapture::with_budget(budget);
let mut truncated = false;
capture.append(b"1234", 64, &mut truncated);
capture.append(b"5", 64, &mut truncated);
assert_eq!(capture.bytes.capacity(), 8);
capture.append(b"6789", 64, &mut truncated);
assert_eq!(&*capture, b"12345678");
assert_eq!(capture.bytes.capacity(), 8);
assert!(truncated);
}
#[test]
fn stream_capture_budget_local_limit_keeps_a_contiguous_prefix() {
let budget = StreamCaptureBudget::new(128);
let mut capture = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut truncated = false;
capture.append(b"abcdef", 3, &mut truncated);
assert_eq!(&*capture, b"abc");
assert!(truncated);
assert_eq!(budget.available.load(Ordering::Relaxed), 125);
}
#[test]
fn stream_capture_budget_concurrent_growth_and_drop_never_exceeds_capacity() {
const LIMIT: usize = 256;
const THREADS: usize = 8;
let budget = StreamCaptureBudget::new(LIMIT);
let barrier = std::sync::Barrier::new(THREADS);
let held = AtomicUsize::new(0);
std::thread::scope(|scope| {
for index in 0..THREADS {
let budget = &budget;
let barrier = &barrier;
let held = &held;
scope.spawn(move || {
for _ in 0..32 {
let mut capture = StreamBodyCapture::with_budget(Arc::clone(budget));
let mut truncated = false;
barrier.wait();
capture.append(&[1; 16], LIMIT, &mut truncated);
capture.append(&[2; 48], LIMIT, &mut truncated);
held.fetch_add(capture.bytes.capacity(), Ordering::SeqCst);
barrier.wait();
if index == 0 {
let retained = held.load(Ordering::SeqCst);
assert!(retained <= LIMIT);
assert_eq!(retained + budget.available.load(Ordering::Relaxed), LIMIT,);
}
barrier.wait();
drop(capture);
barrier.wait();
if index == 0 {
assert_eq!(budget.available.load(Ordering::Relaxed), LIMIT);
held.store(0, Ordering::SeqCst);
}
barrier.wait();
}
});
}
});
}
}
@@ -11,6 +11,10 @@ const GEMINI_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750);
pub(super) enum StreamCommitPolicy {
ResponseHeaders,
FirstClassifiedBody,
FirstSseSemanticEvent {
max_bytes: usize,
max_wait: Duration,
},
FirstAnthropicSemanticEvent {
max_bytes: usize,
max_wait: Duration,
@@ -36,16 +40,21 @@ impl StreamCommitPolicy {
return Self::FirstClassifiedBody;
}
if force_prefetch {
return Self::FirstClassifiedBody;
}
let content_type = content_type
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or_default()
.to_ascii_lowercase();
if content_type.contains("text/event-stream") {
if provider_api_format.eq_ignore_ascii_case("openai:image")
|| client_api_format.eq_ignore_ascii_case("openai:image")
{
return if force_prefetch {
Self::FirstClassifiedBody
} else {
Self::ResponseHeaders
};
}
if provider_api_format.eq_ignore_ascii_case("claude:messages")
&& provider_api_format.eq_ignore_ascii_case(client_api_format)
&& !has_private_stream_normalizer
@@ -62,7 +71,14 @@ impl StreamCommitPolicy {
max_wait: GEMINI_PRECOMMIT_MAX_WAIT,
};
}
return Self::ResponseHeaders;
return Self::FirstSseSemanticEvent {
max_bytes: MAX_STREAM_PREFETCH_BYTES,
max_wait: Duration::from_secs(30),
};
}
if force_prefetch {
return Self::FirstClassifiedBody;
}
if has_private_stream_normalizer || has_local_stream_rewriter {
@@ -91,14 +107,17 @@ impl StreamCommitPolicy {
pub(super) const fn requires_bounded_frame_wait(self) -> bool {
matches!(
self,
Self::FirstAnthropicSemanticEvent { .. } | Self::FirstGeminiSemanticEvent { .. }
Self::FirstAnthropicSemanticEvent { .. }
| Self::FirstGeminiSemanticEvent { .. }
| Self::FirstSseSemanticEvent { .. }
)
}
pub(super) const fn max_precommit_wait(self) -> Option<Duration> {
match self {
Self::FirstAnthropicSemanticEvent { max_wait, .. }
| Self::FirstGeminiSemanticEvent { max_wait, .. } => Some(max_wait),
| Self::FirstGeminiSemanticEvent { max_wait, .. }
| Self::FirstSseSemanticEvent { max_wait, .. } => Some(max_wait),
Self::ResponseHeaders | Self::FirstClassifiedBody => None,
}
}
@@ -110,6 +129,16 @@ impl StreamCommitPolicy {
pub(super) const fn is_gemini(self) -> bool {
matches!(self, Self::FirstGeminiSemanticEvent { .. })
}
pub(super) fn with_precommit_wait(mut self, wait: Duration) -> Self {
match &mut self {
Self::FirstAnthropicSemanticEvent { max_wait, .. }
| Self::FirstGeminiSemanticEvent { max_wait, .. }
| Self::FirstSseSemanticEvent { max_wait, .. } => *max_wait = wait,
_ => {}
}
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -133,6 +162,7 @@ pub(super) struct StreamCommitGate {
observed_bytes: usize,
anthropic: AnthropicSsePrecommitInspector,
gemini: GeminiSsePrecommitInspector,
generic: GenericSsePrecommitInspector,
}
impl StreamCommitGate {
@@ -148,6 +178,7 @@ impl StreamCommitGate {
observed_bytes: 0,
anthropic: AnthropicSsePrecommitInspector::default(),
gemini: GeminiSsePrecommitInspector::default(),
generic: GenericSsePrecommitInspector::default(),
}
}
@@ -171,6 +202,9 @@ impl StreamCommitGate {
StreamCommitPolicy::FirstGeminiSemanticEvent { max_bytes, .. } => {
(max_bytes, self.gemini.observe(chunk, max_bytes))
}
StreamCommitPolicy::FirstSseSemanticEvent { max_bytes, .. } => {
(max_bytes, self.generic.observe(chunk, max_bytes))
}
StreamCommitPolicy::ResponseHeaders | StreamCommitPolicy::FirstClassifiedBody => {
return StreamPrecommitObservation::Pending;
}
@@ -217,6 +251,152 @@ enum SemanticSseObservation {
Error { status_code: u16, body_json: Value },
}
#[derive(Debug, Default)]
struct GenericSsePrecommitInspector {
buffered: Vec<u8>,
}
impl GenericSsePrecommitInspector {
fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation {
let remaining = max_bytes.saturating_sub(self.buffered.len());
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_generic_sse_record(&record) {
SemanticSseObservation::Pending => {}
observation => return observation,
}
}
if chunk.len() > remaining {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
}
}
}
fn classify_generic_sse_record(record: &[u8]) -> SemanticSseObservation {
let Ok(record) = std::str::from_utf8(record) else {
return SemanticSseObservation::SemanticEvent;
};
let normalized = record.replace("\r\n", "\n").replace('\r', "\n");
let event_type = normalized
.lines()
.find_map(|line| line.strip_prefix("event:").map(str::trim));
let data = normalized
.lines()
.filter_map(|line| line.strip_prefix("data:").map(str::trim_start))
.collect::<Vec<_>>()
.join("\n");
if data.trim().is_empty() || matches!(event_type, Some("ping" | "heartbeat" | "keepalive")) {
return SemanticSseObservation::Pending;
}
if data.trim() == "[DONE]" {
return SemanticSseObservation::SemanticEvent;
}
let Ok(body_json) = serde_json::from_str::<Value>(data.trim()) else {
return SemanticSseObservation::SemanticEvent;
};
let payload_type = body_json.get("type").and_then(Value::as_str).or(event_type);
if payload_type.is_some_and(is_anthropic_semantic_event_type) {
return classify_anthropic_sse_record(record.as_bytes());
}
let error = body_json
.get("error")
.filter(|value| !value.is_null())
.or_else(|| {
body_json
.pointer("/response/error")
.filter(|value| !value.is_null())
});
if error.is_some()
|| matches!(payload_type, Some("error" | "response.failed"))
|| body_json.get("status").and_then(Value::as_str) == Some("failed")
{
let failure = error
.map(|error| serde_json::json!({ "error": error }))
.unwrap_or_else(|| body_json.clone());
return SemanticSseObservation::Error {
status_code: crate::execution_runtime::submission::resolve_local_sync_error_status_code(
200, &failure,
),
body_json: failure,
};
}
if matches!(
payload_type,
Some("ping" | "response.created" | "response.in_progress" | "response.queued")
) {
return SemanticSseObservation::Pending;
}
if payload_type == Some("response.output_item.added")
&& matches!(
body_json.pointer("/item/type").and_then(Value::as_str),
Some("message" | "reasoning")
)
&& body_json
.pointer("/item/content")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty)
&& body_json
.pointer("/item/summary")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty)
{
return SemanticSseObservation::Pending;
}
if matches!(
payload_type,
Some("response.content_part.added" | "response.reasoning_summary_part.added")
) && matches!(
body_json.pointer("/part/type").and_then(Value::as_str),
Some("output_text" | "summary_text" | "refusal")
) && !body_json
.pointer("/part/text")
.is_some_and(value_has_semantic_content)
&& !body_json
.pointer("/part/refusal")
.is_some_and(value_has_semantic_content)
{
return SemanticSseObservation::Pending;
}
if let Some(choices) = body_json.get("choices").and_then(Value::as_array) {
let semantic = choices.iter().any(|choice| {
choice
.get("finish_reason")
.is_some_and(|value| !value.is_null())
|| choice.get("text").is_some_and(value_has_semantic_content)
|| choice
.get("delta")
.or_else(|| choice.get("message"))
.and_then(Value::as_object)
.is_some_and(|delta| {
delta.iter().any(|(name, value)| {
name != "role" && value_has_semantic_content(value)
})
})
});
return if semantic {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
};
}
SemanticSseObservation::SemanticEvent
}
fn value_has_semantic_content(value: &Value) -> bool {
match value {
Value::Null => false,
Value::String(text) => !text.is_empty(),
Value::Array(values) => !values.is_empty(),
Value::Object(values) => !values.is_empty(),
_ => true,
}
}
#[derive(Debug, Default)]
struct AnthropicSsePrecommitInspector {
buffered: Vec<u8>,
@@ -355,7 +535,30 @@ fn classify_anthropic_sse_record(record: &[u8]) -> SemanticSseObservation {
(None, Some(payload_type)) => Some(payload_type),
_ => None,
};
if semantic_type.is_some_and(is_anthropic_semantic_event_type) {
let setup_only = match semantic_type {
Some("message_start") => body_json
.pointer("/message/content")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty),
Some("content_block_start") => {
let block_type = body_json
.pointer("/content_block/type")
.and_then(Value::as_str);
matches!(block_type, Some("text" | "thinking"))
&& !body_json
.pointer("/content_block/text")
.is_some_and(value_has_semantic_content)
&& !body_json
.pointer("/content_block/thinking")
.is_some_and(value_has_semantic_content)
}
Some("content_block_stop") => true,
Some("message_delta") => body_json
.pointer("/delta/stop_reason")
.is_none_or(Value::is_null),
_ => false,
};
if !setup_only && semantic_type.is_some_and(is_anthropic_semantic_event_type) {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
@@ -507,6 +710,89 @@ pub(super) fn anthropic_error_status_code(body_json: &Value) -> u16 {
#[cfg(test)]
mod tests {
#[test]
fn image_streams_only_prefetch_when_explicitly_requested() {
for force_prefetch in [false, true] {
let policy = super::StreamCommitPolicy::for_response(
true,
Some("text/event-stream"),
"openai:image",
"openai:image",
false,
false,
force_prefetch,
);
assert_eq!(policy.commits_on_response_headers(), !force_prefetch);
assert!(!policy.requires_bounded_frame_wait());
}
}
#[test]
fn generic_sse_waits_through_setup_and_classifies_fragmented_errors() {
let setup = b"event: response.created\ndata: {\"type\":\"response.created\"}\n\n";
let failure = b"event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"capacity exhausted\"}}}\n\n";
for split in 1..failure.len() {
let policy = super::StreamCommitPolicy::FirstSseSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
};
let mut gate = super::StreamCommitGate::new(policy);
assert_eq!(
gate.observe_provider_bytes(setup),
super::StreamPrecommitObservation::Pending
);
for control in [
b"event: ping\ndata: keepalive\n\n".as_slice(),
b"data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"reasoning\",\"summary\":[]}}\n\n".as_slice(),
b"data: {\"type\":\"response.reasoning_summary_part.added\",\"part\":{\"type\":\"summary_text\",\"text\":\"\"}}\n\n".as_slice(),
] {
assert_eq!(gate.observe_provider_bytes(control), super::StreamPrecommitObservation::Pending);
}
assert_eq!(
gate.observe_provider_bytes(&failure[..split]),
super::StreamPrecommitObservation::Pending
);
assert!(matches!(
gate.observe_provider_bytes(&failure[split..]),
super::StreamPrecommitObservation::UpstreamError { .. }
));
}
}
#[test]
fn generic_sse_commits_on_content_or_tool_call_but_not_role() {
for output in [
"data: {\"choices\":[{\"delta\":{\"content\":\"hello\"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"id\":\"call-1\"}]}}]}\n\n",
] {
let mut gate =
super::StreamCommitGate::new(super::StreamCommitPolicy::FirstSseSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
});
assert_eq!(gate.observe_provider_bytes(b"data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":\"\"}}]}\n\n"), super::StreamPrecommitObservation::Pending);
assert_eq!(
gate.observe_provider_bytes(output.as_bytes()),
super::StreamPrecommitObservation::Commit
);
assert_eq!(
gate.observe_provider_bytes(b"data: {\"error\":{\"message\":\"late error\"}}\n\n"),
super::StreamPrecommitObservation::Commit
);
}
}
#[test]
fn native_anthropic_setup_does_not_hide_an_early_error() {
let mut gate =
super::StreamCommitGate::new(super::StreamCommitPolicy::FirstAnthropicSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
});
assert_eq!(gate.observe_provider_bytes(b"event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"content\":[]}}\n\n"), super::StreamPrecommitObservation::Pending);
assert_eq!(gate.observe_provider_bytes(b"event: content_block_start\ndata: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n"), super::StreamPrecommitObservation::Pending);
assert!(matches!(gate.observe_provider_bytes(b"event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n"), super::StreamPrecommitObservation::UpstreamError { status_code: 529, .. }));
}
use std::time::Duration;
use super::{
@@ -553,7 +839,7 @@ mod tests {
false,
false,
)
.commits_on_response_headers());
.requires_bounded_frame_wait());
assert!(StreamCommitPolicy::for_response(
true,
Some("text/event-stream"),
@@ -563,7 +849,7 @@ mod tests {
true,
false,
)
.commits_on_response_headers());
.requires_bounded_frame_wait());
}
#[test]
@@ -745,8 +1031,8 @@ mod tests {
let mut gate = StreamCommitGate::new(native_anthropic_policy());
let observation = gate.observe_provider_bytes(
concat!(
"event: message_start\n",
"data: {\"type\":\"message_start\",\"message\":{}}\n\n",
"event: content_block_delta\n",
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n",
"event: error\n",
"data: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n",
)
File diff suppressed because it is too large Load Diff
@@ -1,7 +1,10 @@
mod capture_budget;
mod commit_policy;
mod error;
mod execution;
mod usage_fallback;
pub(crate) use execution::{
execute_execution_runtime_stream, execute_execution_runtime_stream_with_retry_scope,
ClientVisibleStreamCompletionTracker,
};
File diff suppressed because it is too large Load Diff
@@ -12,7 +12,6 @@ use async_stream::stream;
use axum::body::Bytes;
use base64::Engine as _;
use futures_util::{Stream, StreamExt};
use http_body_util::BodyExt;
use serde_json::Value;
use tracing::warn;
@@ -21,9 +20,14 @@ use crate::ai_serving::api::{
normalize_provider_private_report_context, StreamingStandardTerminalObserver,
};
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
use crate::execution_runtime::stream::ClientVisibleStreamCompletionTracker;
use crate::execution_runtime::stream_read_timeout::{
await_stream_idle_read, stream_idle_timeout_message,
};
use crate::execution_runtime::transport::{
append_upstream_response_body_chunk, decode_response_body_bytes,
stream_first_byte_timeout_message, DirectUpstreamResponse,
direct_upstream_response_byte_stream, stream_first_byte_timeout_message,
DirectUpstreamResponse,
};
use crate::execution_runtime::DirectUpstreamStreamExecution;
use crate::GatewayError;
@@ -31,6 +35,15 @@ use crate::GatewayError;
const STREAM_USAGE_OBSERVER_MAX_LINE_BYTES: usize = 1024 * 1024;
const UPSTREAM_STREAM_READ_ERROR_MESSAGE: &str = "Upstream response stream failed";
fn upstream_stream_error_category(response: &DirectUpstreamResponse) -> &'static str {
match response {
DirectUpstreamResponse::Reqwest(_) => "reqwest_body_read_failed",
DirectUpstreamResponse::HyperH2c(_) => "hyper_body_read_failed",
DirectUpstreamResponse::BrowserWreq(_) => "browser_body_read_failed",
DirectUpstreamResponse::LocalTunnel(_) => "tunnel_body_read_failed",
}
}
pub(crate) fn build_direct_execution_frame_stream(
execution: DirectUpstreamStreamExecution,
) -> impl Stream<Item = Result<Bytes, IoError>> + Send + 'static {
@@ -49,9 +62,11 @@ pub(crate) fn build_direct_execution_frame_stream(
started_at,
response_observation,
stream_first_byte_timeout,
stream_idle_timeout,
upstream_target_permit,
} = execution;
let _upstream_target_permit = upstream_target_permit;
let upstream_error_category = upstream_stream_error_category(&response);
let mut observer_context = stream_summary_report_context;
if observer_context
@@ -74,6 +89,7 @@ pub(crate) fn build_direct_execution_frame_stream(
let mut private_stream_normalizer =
maybe_build_provider_private_stream_normalizer(Some(&observer_context));
let mut stream_terminal_observer = StreamingStandardTerminalObserver::default();
let mut stream_completion = ClientVisibleStreamCompletionTracker::default();
let mut observer_buffered = Vec::new();
if should_buffer_non_stream_response(
@@ -87,6 +103,7 @@ pub(crate) fn build_direct_execution_frame_stream(
response,
started_at,
stream_first_byte_timeout,
stream_idle_timeout,
)
.await
{
@@ -166,6 +183,7 @@ pub(crate) fn build_direct_execution_frame_stream(
ttfb_ms,
upstream_bytes,
first_byte_timeout,
idle_timeout,
}) => {
match encode_headers_frame(
status_code,
@@ -180,6 +198,8 @@ pub(crate) fn build_direct_execution_frame_stream(
}
let error_frame = if let Some(timeout) = first_byte_timeout {
encode_first_byte_timeout_frame(timeout)
} else if let Some(timeout) = idle_timeout {
encode_idle_timeout_frame(timeout)
} else {
encode_error_frame(message)
};
@@ -228,7 +248,9 @@ pub(crate) fn build_direct_execution_frame_stream(
let mut prefetched_body_failed = false;
for item in prefetched_body {
match item {
Ok(chunk) if chunk.is_empty() => continue,
Ok(chunk) => {
stream_completion.observe_chunk(&chunk);
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
@@ -280,328 +302,97 @@ pub(crate) fn build_direct_execution_frame_stream(
}
}
if !prefetched_body_failed {
match response {
DirectUpstreamResponse::Reqwest(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
let mut bytes_stream = direct_upstream_response_byte_stream(VecDeque::new(), response);
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(), started_at, stream_first_byte_timeout,
).await {
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => yield Err(err),
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
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_category = "reqwest_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
}
}
}
DirectUpstreamResponse::HyperH2c(response) => {
let mut bytes_stream = response.into_body().into_data_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
} else {
match await_stream_idle_read(bytes_stream.next(), stream_idle_timeout).await {
Ok(item) => item,
Err(timeout) => {
drop(bytes_stream);
if stream_completion.successful_completion()
|| (!stream_completion.observed_terminal()
&& stream_terminal_observer.latest_summary().is_some_and(|summary| {
summary.observed_finish && summary.parser_error.is_none()
&& summary.finish_reason.as_deref() != Some("error")
}))
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
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_category = "hyper_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
if stream_terminal_observer.latest_summary().is_some_and(|summary| {
summary.observed_finish && summary.parser_error.is_some()
}) {
// The terminal summary carries the original provider failure.
break;
}
match encode_idle_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => yield Err(err),
}
break;
}
}
}
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
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_category = "browser_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
};
let Some(item) = item else { break };
match item {
Ok(chunk) => {
stream_completion.observe_chunk(&chunk);
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
response.next_chunk(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
response.next_chunk().await
};
match item {
Ok(Some(chunk)) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
Ok(None) => break,
Err(_message) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error_category = "tunnel_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
break;
}
}
Err(_) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error_category = upstream_error_category,
"upstream body stream read error"
);
match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) {
Ok(frame) => yield Ok(frame),
Err(err) => yield Err(err),
}
break;
}
}
}
}
@@ -704,6 +495,22 @@ fn encode_first_byte_timeout_frame(timeout: Duration) -> Result<Bytes, IoError>
})
}
fn encode_idle_timeout_frame(timeout: Duration) -> Result<Bytes, IoError> {
encode_stream_frame_ndjson(&StreamFrame {
frame_type: StreamFrameType::Error,
payload: StreamFramePayload::Error {
error: ExecutionError {
kind: ExecutionErrorKind::ReadTimeout,
phase: ExecutionPhase::StreamRead,
message: stream_idle_timeout_message(timeout),
upstream_status: None,
retryable: true,
failover_recommended: true,
},
},
})
}
async fn await_stream_first_byte<T, F>(
future: F,
started_at: Instant,
@@ -737,6 +544,7 @@ struct BufferedUpstreamBodyError {
ttfb_ms: Option<u64>,
upstream_bytes: u64,
first_byte_timeout: Option<Duration>,
idle_timeout: Option<Duration>,
}
fn append_buffered_upstream_body_chunk(
@@ -752,6 +560,7 @@ fn append_buffered_upstream_body_chunk(
ttfb_ms,
upstream_bytes: *upstream_bytes,
first_byte_timeout: None,
idle_timeout: None,
}
})
}
@@ -816,16 +625,52 @@ fn should_buffer_non_stream_response(
}
async fn buffer_non_sse_upstream_body(
mut prefetched_body: VecDeque<Result<Bytes, String>>,
prefetched_body: VecDeque<Result<Bytes, String>>,
response: DirectUpstreamResponse,
started_at: Instant,
stream_first_byte_timeout: Option<Duration>,
stream_idle_timeout: Option<Duration>,
) -> Result<BufferedUpstreamBody, BufferedUpstreamBodyError> {
let mut body_bytes = Vec::new();
let mut upstream_bytes = 0u64;
let mut ttfb_ms = None;
while let Some(item) = prefetched_body.pop_front() {
let upstream_error_category = upstream_stream_error_category(&response);
let mut bytes_stream = direct_upstream_response_byte_stream(prefetched_body, response);
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
idle_timeout: None,
})
}
}
} else {
match await_stream_idle_read(bytes_stream.next(), stream_idle_timeout).await {
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_idle_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
idle_timeout: Some(timeout),
})
}
}
};
let Some(item) = item else { break };
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
@@ -838,246 +683,24 @@ async fn buffer_non_sse_upstream_body(
&mut upstream_bytes,
)?;
}
Err(_message) => {
Err(_) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error_category = upstream_error_category,
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(),
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
idle_timeout: None,
});
}
}
}
match response {
DirectUpstreamResponse::Reqwest(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
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_category = "reqwest_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
}
}
DirectUpstreamResponse::HyperH2c(response) => {
let mut bytes_stream = response.into_body().into_data_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
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_category = "hyper_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
}
}
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
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_category = "browser_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
response.next_chunk(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
response.next_chunk().await
};
match item {
Ok(Some(chunk)) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
Ok(None) => break,
Err(_message) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error_category = "tunnel_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(),
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
},
}
Ok(BufferedUpstreamBody {
body_bytes,
ttfb_ms,
@@ -1605,6 +1228,89 @@ mod tests {
assert_eq!(error.get("failover_recommended"), Some(&Value::Bool(true)));
}
#[tokio::test]
async fn direct_execution_frame_stream_enforces_idle_timeout_after_first_byte() {
for (content_type, first_chunk, expect_timeout, provider_format) in [
("text/event-stream", "data: hello\n\n", true, "openai:chat"),
("application/json", "{\"message\":", true, "openai:chat"),
("text/event-stream", "data: [DONE]\n\n", false, "openai:chat"),
("text/event-stream", "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\"}}\n\n", false, "openai:responses"),
("text/event-stream", "event: response.incomplete\ndata: {\"type\":\"response.incomplete\",\"response\":{}}\n\n", true, "openai:responses"),
] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0_u8; 4096];
assert!(socket.read(&mut request).await.unwrap() > 0);
let response = if content_type == "application/json" {
format!("HTTP/1.1 200 OK\r\ncontent-type: {content_type}\r\ncontent-length: 1024\r\n\r\n{first_chunk}")
} else {
format!(
"HTTP/1.1 200 OK\r\ncontent-type: {content_type}\r\ntransfer-encoding: chunked\r\n\r\n{:x}\r\n{first_chunk}\r\n",
first_chunk.len(),
)
};
socket.write_all(response.as_bytes()).await.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_secs(5)).await;
});
let execution = DirectSyncExecutionRuntime::new()
.execute_stream(&ExecutionPlan {
request_id: "req-stream-idle-timeout".into(),
candidate_id: Some("cand-stream-idle-timeout".into()),
provider_name: Some("openai".into()),
provider_id: "prov-1".into(),
endpoint_id: "ep-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: format!("http://{addr}/chat"),
headers: BTreeMap::from([("content-type".into(), "application/json".into())]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(serde_json::json!({"stream": true})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: provider_format.into(),
model_name: Some("gpt-5".into()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts {
first_byte_ms: Some(1_000),
read_ms: Some(10),
..ExecutionTimeouts::default()
}),
})
.await
.expect("stream response headers");
let frames = tokio::time::timeout(
Duration::from_secs(1),
build_direct_execution_frame_stream(execution).collect::<Vec<_>>(),
)
.await;
server.abort();
let frames = frames
.expect("idle timeout must terminate both SSE and buffered JSON")
.into_iter()
.map(|line| serde_json::from_slice::<Value>(&line.unwrap()).unwrap())
.collect::<Vec<_>>();
let errors = frames
.iter()
.filter(|frame| frame["type"] == "error")
.collect::<Vec<_>>();
assert_eq!(errors.len(), usize::from(expect_timeout));
if expect_timeout {
assert_eq!(errors[0]["payload"]["error"]["kind"], "read_timeout");
assert_eq!(errors[0]["payload"]["error"]["phase"], "stream_read");
}
assert!(frames.iter().any(|frame| frame["type"] == "eof"));
assert_eq!(
frames.iter().any(|frame| frame["type"] == "data"),
content_type == "text/event-stream"
);
}
}
#[tokio::test]
async fn direct_execution_frame_stream_emits_telemetry_before_first_data_frame() {
let listener = crate::test_support::bind_loopback_listener()
@@ -0,0 +1,187 @@
use std::future::Future;
use std::time::Duration;
use aether_contracts::ExecutionPlan;
use axum::body::Bytes;
use futures_util::{Stream, StreamExt};
const STREAM_IDLE_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS";
const DEFAULT_STREAM_IDLE_TIMEOUT_MS: u64 = 300_000;
pub(crate) fn resolve_stream_idle_timeout(plan: &ExecutionPlan) -> Option<Duration> {
if !plan.stream {
return None;
}
let configured = std::env::var(STREAM_IDLE_TIMEOUT_MS_ENV).ok();
stream_idle_timeout_from_config(
plan.timeouts.as_ref().and_then(|timeouts| timeouts.read_ms),
configured.as_deref(),
)
}
fn stream_idle_timeout_from_config(
read_ms: Option<u64>,
configured: Option<&str>,
) -> Option<Duration> {
let timeout_ms = read_ms
.or_else(|| configured.and_then(|value| value.trim().parse::<u64>().ok()))
.unwrap_or(DEFAULT_STREAM_IDLE_TIMEOUT_MS);
// Zero explicitly disables the idle limit for providers with long silent reasoning phases.
(timeout_ms > 0).then(|| Duration::from_millis(timeout_ms))
}
pub(crate) fn stream_idle_timeout_message(timeout: Duration) -> String {
format!(
"provider stream idle read timeout after {} ms",
timeout.as_millis()
)
}
pub(crate) async fn await_stream_idle_read<T>(
future: impl Future<Output = T>,
timeout: Option<Duration>,
) -> Result<T, Duration> {
match timeout {
Some(timeout) => tokio::time::timeout(timeout, future)
.await
.map_err(|_| timeout),
None => Ok(future.await),
}
}
pub(crate) fn skip_empty_upstream_chunks<E: Send + 'static>(
upstream: impl Stream<Item = Result<Bytes, E>> + Send + 'static,
) -> impl Stream<Item = Result<Bytes, E>> + Send {
async_stream::stream! {
tokio::pin!(upstream);
while let Some(item) = upstream.next().await {
match item {
Ok(chunk) if chunk.is_empty() => {
// Empty frames are not progress; yield so an always-ready source cannot
// monopolize the executor or prevent its enclosing timeout from firing.
tokio::task::yield_now().await;
}
item => yield item,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
#[test]
fn idle_timeout_configuration_preserves_provider_override_and_explicit_disable() {
assert_eq!(
stream_idle_timeout_from_config(None, None),
Some(Duration::from_secs(300))
);
assert_eq!(
stream_idle_timeout_from_config(None, Some(" invalid ")),
Some(Duration::from_secs(300))
);
assert_eq!(
stream_idle_timeout_from_config(None, Some(" 600000 ")),
Some(Duration::from_secs(600))
);
assert_eq!(
stream_idle_timeout_from_config(Some(120_000), Some("600000")),
Some(Duration::from_secs(120))
);
assert_eq!(
stream_idle_timeout_from_config(Some(0), Some("600000")),
None
);
assert_eq!(stream_idle_timeout_from_config(None, Some("0")), None);
}
#[tokio::test]
async fn idle_timeout_cancels_the_pending_upstream_read() {
struct DropMarker(Arc<AtomicBool>);
impl Drop for DropMarker {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let dropped = Arc::new(AtomicBool::new(false));
let marker = DropMarker(Arc::clone(&dropped));
let outcome = await_stream_idle_read(
async move {
let _marker = marker;
std::future::pending::<()>().await;
},
Some(Duration::from_millis(5)),
)
.await;
assert_eq!(outcome, Err(Duration::from_millis(5)));
assert!(dropped.load(Ordering::SeqCst));
}
#[tokio::test]
async fn idle_timeout_allows_progressing_stream_to_outlive_one_timeout() {
for _ in 0..3 {
assert_eq!(
await_stream_idle_read(
async {
tokio::time::sleep(Duration::from_millis(10)).await;
1
},
Some(Duration::from_millis(25))
)
.await,
Ok(1)
);
}
assert_eq!(
await_stream_idle_read(
async {
tokio::time::sleep(Duration::from_millis(30)).await;
2
},
None
)
.await,
Ok(2)
);
}
#[tokio::test]
async fn empty_upstream_chunks_do_not_reset_idle_timeout() {
let upstream = futures_util::stream::repeat(Ok::<_, ()>(Bytes::new()));
let filtered = skip_empty_upstream_chunks(upstream);
tokio::pin!(filtered);
let outcome = tokio::time::timeout(
Duration::from_secs(1),
await_stream_idle_read(filtered.next(), Some(Duration::from_millis(5))),
)
.await
.expect("empty ready chunks must yield to the idle timer");
assert_eq!(outcome, Err(Duration::from_millis(5)));
}
#[tokio::test]
async fn downstream_keepalive_ticks_do_not_reset_pending_upstream_idle_timeout() {
let read = await_stream_idle_read(
std::future::pending::<()>(),
Some(Duration::from_millis(30)),
);
tokio::pin!(read);
let mut keepalive = tokio::time::interval(Duration::from_millis(2));
let mut ticks = 0;
loop {
tokio::select! {
result = &mut read => {
assert_eq!(result, Err(Duration::from_millis(30)));
assert!(ticks > 0);
break;
}
_ = keepalive.tick() => { ticks += 1; }
}
}
}
}
@@ -529,7 +529,13 @@ fn classify_local_sync_error_kind(
{
return LocalCoreSyncErrorKind::Overloaded;
}
if (500..600).contains(&status_code) {
if (500..600).contains(&status_code)
|| raw_type.is_some_and(|value| {
["server_error", "internal_error", "api_error"]
.iter()
.any(|kind| value.trim().eq_ignore_ascii_case(kind))
})
{
return LocalCoreSyncErrorKind::ServerError;
}
LocalCoreSyncErrorKind::InvalidRequest
@@ -676,6 +682,13 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize(
#[cfg(test)]
mod tests {
#[test]
fn success_http_status_does_not_misclassify_explicit_server_errors_as_bad_requests() {
for error_type in ["server_error", "internal_error", "api_error"] {
let body = serde_json::json!({ "error": { "type": error_type, "message": "failed" } });
assert_eq!(super::resolve_local_sync_error_status_code(200, &body), 500);
}
}
use axum::body::to_bytes;
use serde_json::json;
@@ -270,7 +270,9 @@ impl Drop for SyncAttemptTerminalGuard {
let candidate_started_unix_ms = self.candidate_started_unix_ms;
let candidate_started_at = self.candidate_started_at;
if let Ok(handle) = tokio::runtime::Handle::try_current() {
let usage_producer = state.usage_runtime.track_producer();
handle.spawn(async move {
let _usage_producer = usage_producer;
record_sync_attempt_forced_terminal_state(
state,
plan,
@@ -21,10 +21,7 @@ use aether_contracts::{
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, is_https_or_loopback_http_url, is_private_or_reserved_ip,
HttpClientConfig,
};
use aether_http::{apply_http_client_config, is_private_or_reserved_ip, HttpClientConfig};
use aether_runtime::{MetricKind, MetricSample};
use axum::body::Bytes;
use base64::Engine as _;
@@ -39,7 +36,7 @@ use hyper::body::Incoming as HyperIncomingBody;
use hyper::client::conn::http2::SendRequest as HyperH2cSendRequest;
use hyper_util::client::legacy::connect::HttpConnector;
use hyper_util::client::legacy::Client as HyperLegacyClient;
use hyper_util::rt::{TokioExecutor, TokioIo};
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use reqwest::redirect::Policy;
use serde::Serialize;
@@ -53,6 +50,7 @@ use tokio::sync::OnceCell as TokioOnceCell;
use crate::ai_serving::api::extract_provider_private_stream_error_body;
#[cfg(test)]
use crate::execution_runtime::remote_compat::execute_sync_plan_via_remote_execution_runtime;
use crate::execution_runtime::stream_read_timeout::resolve_stream_idle_timeout;
use crate::execution_runtime::windsurf::maybe_execute_windsurf_sync;
use crate::frontdoor_loop_guard::{
configured_gateway_frontdoor_base_url, gateway_frontdoor_self_loop_guard_error,
@@ -112,7 +110,7 @@ const DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV: &str =
"AETHER_GATEWAY_DIRECT_REQWEST_PREWARM_SYNC_CLIENTS";
const DEFAULT_H2_TARGET_STREAMS_PER_CLIENT: usize = 8;
const DEFAULT_HTTP1_TARGET_STREAMS_PER_CLIENT: usize = 512;
const DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: usize = 512;
const DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: usize = 32;
const DEFAULT_DIRECT_H2C_TARGET_STREAMS_PER_CLIENT: usize = 128;
const DEFAULT_DIRECT_H2C_SENDER_SELECT_WINDOW: usize = 4;
const MAX_DIRECT_H2C_DRIVER_RUNTIME_THREADS: usize = 16;
@@ -385,6 +383,7 @@ static DIRECT_H2C_SENDER_CACHE: LazyLock<
static DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: LazyLock<usize> = LazyLock::new(|| {
env_positive_usize(DIRECT_H2C_POOL_MAX_IDLE_PER_HOST_ENV)
.unwrap_or(DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST)
.min(1024)
});
static DIRECT_H2C_SENDER_SELECT_WINDOW: LazyLock<usize> = LazyLock::new(|| {
@@ -441,7 +440,7 @@ static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock<DirectHyperH2cSenderCacheMetric
LazyLock::new(DirectHyperH2cSenderCacheMetrics::default);
#[derive(Debug, Clone, Copy, Default)]
struct ExecutionSafeDnsResolver;
pub(crate) struct ExecutionSafeDnsResolver;
#[derive(Debug, Clone, Copy, Default)]
struct ExecutionSafeHyperDnsResolver;
@@ -449,10 +448,7 @@ struct ExecutionSafeHyperDnsResolver;
fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
let host = host.trim_end_matches('.');
host.eq_ignore_ascii_case("localhost")
|| host
.parse::<IpAddr>()
.map(|ip| ip.is_loopback())
.unwrap_or(false)
|| aether_http::parse_ip_literal_host(host).is_some_and(|ip| ip.is_loopback())
}
fn validate_resolved_execution_addresses(
@@ -494,12 +490,9 @@ async fn resolve_execution_target_addresses_with_policy(
port: u16,
provider_execution: bool,
) -> Result<Vec<SocketAddr>, std::io::Error> {
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
let addresses =
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await?
};
.await?;
validate_resolved_execution_addresses(host, addresses, provider_execution)
}
@@ -1207,6 +1200,42 @@ pub(crate) enum DirectUpstreamResponse {
LocalTunnel(tunnel::DirectRelayResponse),
}
pub(crate) fn direct_upstream_response_byte_stream(
prefetched_body: VecDeque<Result<Bytes, String>>,
response: DirectUpstreamResponse,
) -> futures_util::stream::BoxStream<'static, Result<Bytes, String>> {
let response_stream = match response {
DirectUpstreamResponse::Reqwest(response) => response
.bytes_stream()
.map(|item| item.map_err(|err| format_upstream_request_error(&err)))
.boxed(),
DirectUpstreamResponse::HyperH2c(response) => response
.into_body()
.into_data_stream()
.map(|item| item.map_err(|err| format_hyper_error_chain(&err)))
.boxed(),
DirectUpstreamResponse::BrowserWreq(response) => response
.bytes_stream()
.map(|item| item.map_err(|err| format_wreq_upstream_request_error(&err)))
.boxed(),
DirectUpstreamResponse::LocalTunnel(mut response) => async_stream::stream! {
loop {
match response.next_chunk().await {
Ok(Some(chunk)) => yield Ok(chunk),
Ok(None) => break,
Err(err) => {
yield Err(err);
break;
}
}
}
}
.boxed(),
};
let upstream = futures_util::stream::iter(prefetched_body).chain(response_stream);
crate::execution_runtime::stream_read_timeout::skip_empty_upstream_chunks(upstream).boxed()
}
pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) request_id: String,
pub(crate) candidate_id: Option<String>,
@@ -1223,6 +1252,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) started_at: Instant,
pub(crate) response_observation: ExecutionResponseObservation,
pub(crate) stream_first_byte_timeout: Option<Duration>,
pub(crate) stream_idle_timeout: Option<Duration>,
pub(crate) upstream_target_permit: Option<UpstreamTargetAdmissionPermit>,
}
@@ -1356,6 +1386,7 @@ impl DirectSyncExecutionRuntime {
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
stream_idle_timeout: resolve_stream_idle_timeout(plan),
upstream_target_permit: None,
})
}
@@ -1503,6 +1534,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
stream_idle_timeout: resolve_stream_idle_timeout(plan),
upstream_target_permit: None,
}))
}
@@ -2602,6 +2634,8 @@ fn build_direct_h2c_client_from_cache_key(
builder.http2_only(true);
builder.http2_adaptive_window(true);
builder.pool_max_idle_per_host(cache_key.pool_max_idle_per_host);
builder.pool_timer(TokioTimer::new());
builder.pool_idle_timeout(Duration::from_millis(upstream_pool_idle_timeout_ms()));
builder.build(connector)
}
@@ -2863,8 +2897,18 @@ async fn send_via_browser_wreq_transport(
let profile = plan.transport_profile.as_ref().ok_or_else(|| {
ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new())
})?;
let mut client_timeouts = plan.timeouts.clone();
if plan.stream {
if let Some(timeouts) = client_timeouts.as_mut() {
// Streamed responses use the shared idle reader; sync collectors retain
// their existing client read timeout. Zero explicitly disables either.
if apply_request_total_timeout || timeouts.read_ms == Some(0) {
timeouts.read_ms = None;
}
}
}
let client = build_browser_wreq_client(
plan.timeouts.as_ref(),
client_timeouts.as_ref(),
plan.proxy.as_ref(),
profile,
transport_controls,
@@ -4223,6 +4267,7 @@ fn build_direct_reqwest_client_from_cache_key(
&HttpClientConfig {
connect_timeout_ms: cache_key.connect_timeout_ms,
pool_max_idle_per_host: Some(direct_reqwest_pool_max_idle_per_host()),
pool_idle_timeout_ms: Some(upstream_pool_idle_timeout_ms()),
..HttpClientConfig::default()
},
);
@@ -4242,12 +4287,22 @@ fn build_direct_reqwest_client_from_cache_key(
}
fn direct_reqwest_pool_max_idle_per_host() -> usize {
const DEFAULT_MAX_IDLE_PER_HOST: usize = 1024;
const DEFAULT_MAX_IDLE_PER_HOST: usize = 32;
std::env::var("AETHER_GATEWAY_UPSTREAM_POOL_MAX_IDLE_PER_HOST")
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_MAX_IDLE_PER_HOST)
.min(1024)
}
fn upstream_pool_idle_timeout_ms() -> u64 {
std::env::var("AETHER_GATEWAY_UPSTREAM_POOL_IDLE_TIMEOUT_MS")
.ok()
.and_then(|value| value.trim().parse::<u64>().ok())
.filter(|value| *value > 0)
.unwrap_or(15_000)
.min(300_000)
}
pub(crate) fn direct_reqwest_client_cache_metric_samples() -> Vec<MetricSample> {
@@ -4594,7 +4649,11 @@ pub(crate) fn build_browser_wreq_client(
) -> Result<wreq::Client, ExecutionRuntimeTransportError> {
let emulation = browser_wreq_emulation_from_profile(transport_profile)?;
let proxy_url = resolve_proxy_url(proxy)?;
let mut builder = wreq::Client::builder().no_proxy().emulation(emulation);
let mut builder = wreq::Client::builder()
.no_proxy()
.emulation(emulation)
.pool_max_idle_per_host(direct_reqwest_pool_max_idle_per_host())
.pool_idle_timeout(Duration::from_millis(upstream_pool_idle_timeout_ms()));
if proxy_url.is_none() {
builder = builder.dns_resolver(ExecutionSafeDnsResolver);
}
@@ -5152,7 +5211,7 @@ fn execution_log_url_host(url: &str) -> String {
.unwrap_or_else(|| "-".to_string())
}
fn validate_execution_upstream_url(
pub(crate) fn validate_execution_upstream_url(
raw_url: &str,
) -> Result<url::Url, ExecutionRuntimeTransportError> {
let url = url::Url::parse(raw_url).map_err(|_| {
@@ -5173,11 +5232,6 @@ fn validate_execution_upstream_url(
"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)),
@@ -5324,7 +5378,7 @@ pub(crate) fn build_execution_response_body(
mod tests {
use std::collections::BTreeMap;
use std::io::{Read, Write};
use std::sync::{Arc, Mutex, MutexGuard, OnceLock};
use std::sync::{Arc, Mutex};
use aether_contracts::tunnel::{
TUNNEL_RELAY_AUTH_NONCE_HEADER, TUNNEL_RELAY_AUTH_PAYLOAD_HEADER,
@@ -5389,9 +5443,13 @@ mod tests {
const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes";
#[test]
fn execution_upstream_url_requires_https_or_literal_loopback_http() {
fn execution_upstream_url_accepts_http_and_https_with_safe_targets() {
for allowed in [
"https://api.example.test/v1/responses?api-version=1",
"http://api.example.test:8080/v1/responses?api-version=1",
"http://8.8.8.8:8080/v1/responses",
"https://8.8.8.8/v1/responses",
"http://[2606:4700:4700::1111]:8080/v1/responses",
"http://localhost:8080/v1/responses",
"http://127.42.0.1:8080/v1/responses",
"http://[::1]:8080/v1/responses",
@@ -5403,7 +5461,6 @@ mod tests {
}
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",
@@ -5411,6 +5468,8 @@ mod tests {
"https://10.0.0.1:8443/v1/responses",
"https://[email protected]/v1/responses",
"https://example.test/v1/responses#secret",
"http://[email protected]/v1/responses",
"http://example.test/v1/responses#secret",
"ftp://localhost/resource",
] {
assert!(
@@ -5443,6 +5502,8 @@ mod tests {
"93.184.216.34:443".parse().unwrap(),
];
for host in [
"chatgpt.com",
"api.openai.com",
"oauth2.googleapis.com",
"www.googleapis.com",
"custom.example.test",
@@ -5455,6 +5516,46 @@ mod tests {
}
}
#[tokio::test]
async fn execution_dns_handles_url_ipv6_without_weakening_relay_filtering() {
for provider_execution in [false, true] {
let addresses = super::resolve_execution_target_addresses_with_policy(
"[::1]",
8443,
provider_execution,
)
.await
.expect("literal IPv6 loopback should resolve without DNS");
assert_eq!(addresses, vec!["[::1]:8443".parse().unwrap()]);
}
let error = super::resolve_execution_target_addresses_with_policy("[fd00::1]", 443, false)
.await
.expect_err("private IPv6 must remain blocked for relay traffic");
assert_eq!(error.kind(), std::io::ErrorKind::PermissionDenied);
}
#[tokio::test]
async fn execution_dns_resolvers_preserve_provider_fake_ip_answers() {
for host in ["198.18.78.41", "198.19.1.2"] {
let expected = vec![format!("{host}:0").parse::<std::net::SocketAddr>().unwrap()];
let reqwest_addresses = reqwest::dns::Resolve::resolve(
&super::ExecutionSafeDnsResolver,
host.parse().unwrap(),
)
.await
.expect("HTTP provider DNS must accept Fake-IP answers")
.collect::<Vec<_>>();
let wreq_addresses =
wreq::dns::Resolve::resolve(&super::ExecutionSafeDnsResolver, host.into())
.await
.expect("WebSocket provider DNS must accept Fake-IP answers")
.collect::<Vec<_>>();
assert_eq!(reqwest_addresses, expected);
assert_eq!(wreq_addresses, expected);
}
}
#[test]
fn execution_dns_answers_keep_relay_address_filtering() {
let public = "93.184.216.34:443".parse().unwrap();
@@ -6231,16 +6332,14 @@ mod tests {
TestEnvVarGuard { key, previous }
}
fn direct_reqwest_env_lock() -> MutexGuard<'static, ()> {
static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
LOCK.get_or_init(|| Mutex::new(()))
.lock()
.expect("direct reqwest env lock")
fn direct_reqwest_env_lock() -> &'static tokio::sync::Mutex<()> {
static LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
&LOCK
}
#[test]
fn direct_reqwest_client_cache_key_includes_transport_profile() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let timeouts = ExecutionTimeouts {
connect_ms: Some(5_000),
..ExecutionTimeouts::default()
@@ -6344,7 +6443,7 @@ mod tests {
#[test]
fn direct_reqwest_client_cache_evicts_least_recently_used_entry_at_capacity() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_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(
@@ -6442,7 +6541,7 @@ mod tests {
#[test]
fn direct_reqwest_client_cache_key_splits_origin_only_when_enabled() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-origin".into(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(),
@@ -6629,14 +6728,14 @@ mod tests {
#[test]
fn direct_h2c_client_shards_respect_explicit_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_H2C_CLIENT_SHARDS_ENV, "7");
assert_eq!(super::direct_h2c_client_shard_count(), 7);
}
#[test]
fn direct_h2c_adaptive_window_respects_explicit_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
{
let _adaptive = set_test_env_var(super::DIRECT_H2C_ADAPTIVE_WINDOW_ENV, "0");
assert!(!super::direct_h2c_adaptive_window_enabled());
@@ -6704,7 +6803,7 @@ mod tests {
#[test]
fn direct_h2c_prewarm_urls_parse_env_list() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _urls = set_test_env_var(
super::DIRECT_H2C_PREWARM_URLS_ENV,
" http://127.0.0.1:18184/v1/chat/completions,;http://127.0.0.1:18185/v1/chat/completions\nhttp://127.0.0.1:18186/v1/chat/completions ",
@@ -6722,7 +6821,7 @@ mod tests {
#[test]
fn direct_h2c_prewarm_cache_keys_dedup_by_origin() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let urls = vec![
"http://127.0.0.1:18184/v1/chat/completions".to_string(),
"http://127.0.0.1:18184/v1/responses".to_string(),
@@ -6748,7 +6847,7 @@ mod tests {
#[test]
fn direct_h2c_client_cache_splits_by_origin_and_shards() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_H2C_CLIENT_SHARDS_ENV, "3");
super::DIRECT_H2C_CLIENT_CACHE
.lock()
@@ -6773,7 +6872,7 @@ mod tests {
#[test]
fn direct_reqwest_initial_client_shards_are_bounded_by_target() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
assert_eq!(super::direct_reqwest_initial_client_shard_count(1), 1);
assert_eq!(super::direct_reqwest_initial_client_shard_count(2), 2);
assert_eq!(
@@ -6784,7 +6883,7 @@ mod tests {
#[test]
fn direct_reqwest_initial_client_shards_cap_large_sync_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "128");
assert_eq!(
super::direct_reqwest_initial_client_shard_count(128),
@@ -6794,7 +6893,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_client_shards_default_to_initial() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
assert_eq!(super::direct_reqwest_prewarm_client_shard_count(1), 1);
assert_eq!(
super::direct_reqwest_prewarm_client_shard_count(96),
@@ -6804,7 +6903,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_client_shards_do_not_exceed_request_path_cap() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "4");
let _prewarm = set_test_env_var(super::DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV, "128");
@@ -6813,7 +6912,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_populates_cache_for_plan() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "4");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-prewarm".into(),
@@ -6875,7 +6974,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_plan_keeps_large_sync_env_off_request_path() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "128");
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "4");
let _prewarm = set_test_env_var(super::DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV, "128");
@@ -6935,7 +7034,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_skips_h2c_fast_path() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _fast_path = set_test_env_var(super::DIRECT_H2C_FAST_PATH_ENV, "1");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-fast-path-prewarm-skip".into(),
@@ -6988,7 +7087,7 @@ mod tests {
#[test]
fn direct_reqwest_cache_metrics_expose_ready_state() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "1");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-ready-metrics".into(),
@@ -8572,7 +8671,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_supports_tunnel_relay() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET);
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -8739,7 +8838,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_rejects_short_tunnel_relay_secret_before_send() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", &"x".repeat(31));
let execution_runtime = DirectSyncExecutionRuntime::new();
let error = execution_runtime
@@ -8780,7 +8879,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_requires_tunnel_relay_secret_before_send() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = unset_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET");
let execution_runtime = DirectSyncExecutionRuntime::new();
let error = execution_runtime
@@ -9107,7 +9206,7 @@ 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 _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET);
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -9323,7 +9422,7 @@ mod tests {
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn direct_sync_execution_runtime_uses_h2c_prior_knowledge_on_wire() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().lock().await;
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "1");
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -522,7 +522,6 @@ where
decision,
plan_kind,
transfer_tracker,
request_first_byte_started_at: Instant::now(),
};
let loop_result = run_ai_attempt_loop(&port, plan_and_reports).await;
if loop_result.is_err() {
@@ -603,7 +602,6 @@ where
decision,
plan_kind,
transfer_tracker,
request_first_byte_started_at: Instant::now(),
};
let loop_result = run_dynamic_attempt_loop(
&port,
@@ -657,6 +655,73 @@ struct ProviderTransferState {
struct ProviderTransferStateTracker {
by_provider: BTreeMap<String, ProviderTransferState>,
exhausted_provider_ids: BTreeSet<String>,
global: GlobalTransferState,
}
#[derive(Debug, Default)]
struct GlobalTransferState {
first_attempt_started_at: Option<Instant>,
last_candidate: Option<(String, String, String)>,
transfer_count: u64,
limits: Option<ProviderTransferLimits>,
exhausted: bool,
}
impl GlobalTransferState {
fn load_policy(&mut self, report_context: Option<&serde_json::Value>) {
if self.limits.is_none() {
if let Some(policy) =
crate::orchestration::routing_execution_policy_from_report_context(report_context)
{
self.limits = Some(ProviderTransferLimits {
max_transfer_count: policy.max_transfer_count,
max_transfer_timeout_seconds: policy.max_transfer_timeout_seconds,
});
}
}
}
fn changes_candidate(&self, plan: &aether_contracts::ExecutionPlan) -> bool {
self.last_candidate
.as_ref()
.is_some_and(|(provider, endpoint, key)| {
provider != &plan.provider_id
|| endpoint != &plan.endpoint_id
|| key != &plan.key_id
})
}
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.first_attempt_started_at.get_or_insert(now);
if self.changes_candidate(plan) {
self.transfer_count = self.transfer_count.saturating_add(1);
}
self.last_candidate = Some((
plan.provider_id.clone(),
plan.endpoint_id.clone(),
plan.key_id.clone(),
));
}
fn check_before_attempt(
&mut self,
plan: &aether_contracts::ExecutionPlan,
now: Instant,
) -> Option<(bool, bool)> {
let limits = self.limits?;
let started_at = self.first_attempt_started_at?;
let count_reached = self.changes_candidate(plan)
&& limits.max_transfer_count > 0
&& self.transfer_count >= limits.max_transfer_count;
let timeout_reached = limits.max_transfer_timeout_seconds > 0
&& now.saturating_duration_since(started_at)
>= Duration::from_secs(limits.max_transfer_timeout_seconds);
if !count_reached && !timeout_reached {
return None;
}
self.exhausted = true;
Some((count_reached, timeout_reached))
}
}
#[derive(Clone, Debug, Default)]
@@ -719,6 +784,7 @@ struct ProviderTransferLimitReached {
impl ProviderTransferStateTracker {
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.global.record_attempt_started(plan, now);
match self.by_provider.entry(plan.provider_id.clone()) {
std::collections::btree_map::Entry::Vacant(entry) => {
entry.insert(ProviderTransferState {
@@ -905,11 +971,42 @@ async fn should_skip_provider_transfer_attempt<Attempt>(
where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let reached = tracker
.state
.lock()
.await
.check_before_attempt(attempt.execution_plan(), Instant::now());
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 mut tracker = tracker.state.lock().await;
tracker.global.load_policy(report_context);
if tracker.global.exhausted {
return true;
}
let now = Instant::now();
if let Some((count_reached, timeout_reached)) = tracker
.global
.check_before_attempt(attempt.execution_plan(), now)
{
warn!(
event_name = "routing_transfer_limit_reached",
log_type = "event",
trace_id,
plan_kind,
transfer_count = tracker.global.transfer_count,
elapsed_ms = tracker
.global
.first_attempt_started_at
.map(|started| now.saturating_duration_since(started).as_millis() as u64)
.unwrap_or(0),
count_reached,
timeout_reached,
"gateway exhausted the routing strategy transfer budget"
);
return true;
}
let reached = tracker.check_before_attempt(attempt.execution_plan(), now);
let Some(reached) = reached else {
return false;
};
@@ -1121,10 +1218,6 @@ 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]
@@ -1254,7 +1347,6 @@ 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(
@@ -1308,7 +1400,7 @@ where
http::StatusCode::GATEWAY_TIMEOUT.as_u16(),
"local_stream_candidate_watchdog_timeout",
stream_candidate_watchdog_timeout_message(),
self.request_first_byte_started_at.elapsed().as_millis() as u64,
watchdog_started_at.elapsed().as_millis() as u64,
)
.await?,
)
@@ -1758,7 +1850,6 @@ async fn execute_stream_candidate_with_watchdog<Fut>(
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<StreamCandidateWatchdogOutcome, GatewayError>
@@ -1768,7 +1859,6 @@ 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 {
@@ -1794,14 +1884,7 @@ where
let watchdog_progress = StreamCandidateWatchdogProgress::shared();
let execution = watchdog_progress.clone().scope(execute());
tokio::pin!(execution);
// 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);
let deadline = tokio::time::sleep(timeout_duration);
tokio::pin!(deadline);
let execution_result = tokio::select! {
biased;
@@ -1830,10 +1913,6 @@ 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,
@@ -1862,8 +1941,6 @@ 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 {
@@ -2487,6 +2564,130 @@ mod tests {
assert_eq!(port.unused.lock().unwrap().as_slice(), ["a-key3-retry0"]);
}
#[tokio::test]
async fn routing_transfer_budget_counts_switches_across_providers_not_same_key_retries() {
for (limit, succeeds) in [(1, false), (2, true)] {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] =
json!({ "max_transfer_count": limit });
}
let outcome = run_ai_attempt_loop(&port, attempts).await.unwrap();
assert_eq!(
matches!(outcome, AiAttemptLoopOutcome::Responded(_)),
succeeds
);
{
let executed = port.executed.lock().unwrap();
assert_eq!(
&executed[..3],
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(executed.len(), if succeeds { 4 } else { 3 });
}
assert_eq!(port.tracker.state.lock().await.global.transfer_count, limit);
}
}
#[tokio::test]
async fn dynamic_loop_honors_global_transfer_budget_across_providers() {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let mut source = TransferTestAttemptSource {
attempts: attempts.into(),
skipped_providers: Vec::new(),
};
let outcome = run_dynamic_attempt_loop(
&port,
&mut source,
"global-budget",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
outcome,
LocalExecutionRequestOutcome::Exhausted(_)
));
assert_eq!(
port.executed.lock().unwrap().as_slice(),
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
}
#[test]
fn routing_time_budget_is_cumulative_and_zero_is_unlimited() {
let mut global = super::GlobalTransferState::default();
global.load_policy(Some(
&json!({ "routing_execution_policy": { "max_transfer_timeout_seconds": 60 } }),
));
let now = tokio::time::Instant::now();
let plan = test_plan(None);
global.record_attempt_started(&plan, now);
global.record_attempt_started(&plan, now + Duration::from_secs(40));
assert_eq!(global.transfer_count, 0);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(59)),
None
);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(60)),
Some((false, true))
);
let mut unlimited = super::GlobalTransferState::default();
unlimited.load_policy(Some(&json!({ "routing_execution_policy": {} })));
unlimited.record_attempt_started(&plan, now);
assert_eq!(
unlimited.check_before_attempt(&plan, now + Duration::from_secs(86_400)),
None
);
}
#[tokio::test]
async fn cloned_tracker_preserves_global_budget_across_candidate_loops() {
let state = AppState::new().unwrap();
let tracker = ProviderTransferTracker::default();
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let remaining = attempts.split_off(3);
let first_port = TransferTestPort::with_tracker(&state, tracker.clone());
let first_outcome = run_ai_attempt_loop(&first_port, attempts).await.unwrap();
assert!(matches!(first_outcome, AiAttemptLoopOutcome::Exhausted(_)));
assert_eq!(tracker.state.lock().await.global.transfer_count, 1);
let second_port = TransferTestPort::with_tracker(&state, tracker.clone());
let mut source = TransferTestAttemptSource {
attempts: remaining.into(),
skipped_providers: Vec::new(),
};
let second_outcome = run_dynamic_attempt_loop(
&second_port,
&mut source,
"global-budget-across-loops",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
second_outcome,
LocalExecutionRequestOutcome::NoPath
));
assert!(second_port.executed.lock().unwrap().is_empty());
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
assert!(tracker.state.lock().await.global.exhausted);
}
#[tokio::test]
async fn cloned_tracker_preserves_transfer_budget_across_candidate_loops() {
let state = AppState::new().expect("state should build");
@@ -3150,7 +3351,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
false,
|| {
std::future::pending::<
@@ -3189,39 +3389,32 @@ 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());
async fn assert_stream_candidate_retry_gets_fresh_first_byte_budget(
provider_id: &str,
key_id: &str,
first_byte_ms: u64,
) {
let writer = TestRequestCandidateWriter::default();
let plan = test_plan(Some(ExecutionTimeouts {
first_byte_ms: Some(250),
first_byte_ms: Some(100),
..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<AiAttemptExecutionOutcome<Response<Body>>, GatewayError>,
>()
},
),
let result = execute_stream_candidate_with_watchdog(
&writer,
"trace_watchdog_retry_budget",
"claude_cli_stream",
&plan,
Some(&report_context),
false,
|| {
std::future::pending::<
Result<AiAttemptExecutionOutcome<Response<Body>>, GatewayError>,
>()
},
)
.await
.expect("an expired request-level first-byte budget must not restart per candidate");
.await;
assert!(matches!(
result,
Ok(StreamCandidateWatchdogOutcome::Executed(
@@ -3231,14 +3424,119 @@ mod tests {
}
))
));
let mut next_plan = plan.clone();
next_plan.candidate_id = Some("cand_watchdog_retry".to_string());
next_plan.provider_id = provider_id.to_string();
next_plan.key_id = key_id.to_string();
next_plan.timeouts = Some(ExecutionTimeouts {
first_byte_ms: Some(first_byte_ms),
..ExecutionTimeouts::default()
});
let mut next_report_context = report_context.clone();
next_report_context["candidate_id"] = json!("cand_watchdog_retry");
next_report_context["candidate_index"] = json!(3);
let result = execute_stream_candidate_with_watchdog(
&writer,
"trace_watchdog_retry_budget",
"claude_cli_stream",
&next_plan,
Some(&next_report_context),
false,
|| async {
tokio::time::sleep(Duration::from_millis(60)).await;
Ok(AiAttemptExecutionOutcome::Responded(Response::new(
Body::from("retry succeeded"),
)))
},
)
.await;
assert!(
matches!(
result,
Ok(StreamCandidateWatchdogOutcome::Executed(
AiAttemptExecutionOutcome::Responded(_)
))
),
"candidate {provider_id}/{key_id} must receive its own {first_byte_ms} ms budget"
);
let records = writer.records.lock().await;
assert_eq!(records.len(), 1);
assert_eq!(records[0].id, plan.candidate_id.as_deref().unwrap());
assert_eq!(records[0].status, RequestCandidateStatus::Failed);
assert_eq!(
records[0].error_type.as_deref(),
Some("local_stream_candidate_watchdog_timeout")
);
}
#[tokio::test]
async fn stream_candidate_watchdog_failover_gets_fresh_first_byte_budget() {
for first_byte_ms in [100, 75, 150] {
assert_stream_candidate_retry_gets_fresh_first_byte_budget(
"provider_next",
"key_next",
first_byte_ms,
)
.await;
}
}
#[tokio::test]
async fn stream_candidate_watchdog_same_provider_retries_get_fresh_first_byte_budget() {
for key_id in ["key_next", "key_id"] {
assert_stream_candidate_retry_gets_fresh_first_byte_budget("provider_id", key_id, 100)
.await;
}
}
#[tokio::test]
async fn stream_candidate_watchdog_starts_first_byte_budget_after_admission() {
let writer = TestRequestCandidateWriter::with_upstream_gate(1, Duration::from_secs(1));
let held_permit = writer
.upstream_gate
.as_ref()
.expect("test gate should exist")
.try_acquire()
.expect("test gate permit should acquire");
let plan = test_plan(Some(ExecutionTimeouts {
first_byte_ms: Some(50),
..ExecutionTimeouts::default()
}));
let report_context = test_report_context();
let (result, ()) = tokio::join!(
execute_stream_candidate_with_watchdog(
&writer,
"trace_watchdog_admission_budget",
"claude_cli_stream",
&plan,
Some(&report_context),
false,
|| async {
tokio::time::sleep(Duration::from_millis(20)).await;
Ok(AiAttemptExecutionOutcome::Responded(Response::new(
Body::from("admitted candidate succeeded"),
)))
},
),
async move {
tokio::time::sleep(Duration::from_millis(100)).await;
drop(held_permit);
},
);
assert!(matches!(
result,
Ok(StreamCandidateWatchdogOutcome::Executed(
AiAttemptExecutionOutcome::Responded(_)
))
));
assert!(writer.records.lock().await.is_empty());
}
#[tokio::test]
async fn stream_candidate_watchdog_can_stop_on_transport_error() {
let writer = Arc::new(TestRequestCandidateWriter::default());
@@ -3254,7 +3552,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
true,
|| {
std::future::pending::<
@@ -3292,7 +3589,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
true,
|| async {
mark_stream_candidate_watchdog_terminal_started();
@@ -3325,7 +3621,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
true,
|| async {
Err(GatewayError::UpstreamUnavailable {
@@ -3365,7 +3660,6 @@ 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")
@@ -3413,7 +3707,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
false,
|| async {
Err(GatewayError::AdmissionTimeout {
@@ -822,15 +822,20 @@ where
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = standard_text_sync_heartbeat_final_bytes(
let completion = standard_text_sync_heartbeat_final_bytes(
client_api_format.as_str(),
redaction_slot.as_ref(),
execute(state, parts, trace_id, decision, plan_kind, started_at).await,
)
.await;
tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execute(state, parts, trace_id, decision, plan_kind, started_at) => result,
},
);
let bytes = completion.await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -1097,23 +1102,26 @@ fn build_openai_image_sync_heartbeat_shell_response(
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = openai_image_sync_heartbeat_final_bytes(
execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
)
.await,
)
.await;
let execution = execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
);
let outcome = tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execution => result,
};
let bytes = openai_image_sync_heartbeat_final_bytes(outcome).await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -2331,6 +2339,45 @@ mod tests {
.expect("background completion should release admission");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_cancels_when_routing_policy_enables_it() {
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (mut release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let response = crate::request_lifecycle::run_request(async move {
crate::request_lifecycle::configure_client_disconnect(
aether_routing_core::RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
},
);
let (parts, _) = http::Request::builder()
.method("POST")
.uri("/v1/responses")
.body(())
.unwrap()
.into_parts();
build_standard_text_sync_heartbeat_shell_response(
AppState::new().unwrap(),
parts,
"trace-heartbeat-disconnect".to_string(),
test_standard_text_heartbeat_decision(),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND.to_string(),
move |_, _, _, _, _, _| async move {
started_tx.send(()).unwrap();
release_rx.await.unwrap();
Ok(LocalExecutionRequestOutcome::NoPath)
},
)
})
.await
.unwrap();
started_rx.await.unwrap();
drop(response);
tokio::time::timeout(Duration::from_secs(1), release_tx.closed())
.await
.expect("heartbeat must drop upstream execution immediately");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_propagates_request_diagnostics_to_terminal_usage() {
let (state, usage_repository) = heartbeat_usage_test_state(json!({
+29 -43
View File
@@ -756,13 +756,8 @@ fn runtime_miss_client_error_body(api_format: Option<&str>, message: &str) -> Va
}
fn runtime_miss_original_headers_json(headers: &HeaderMap) -> Value {
let mut headers = crate::headers::collect_control_headers(headers);
for (name, value) in headers.iter_mut() {
if runtime_miss_sensitive_header(name) {
*value = runtime_miss_mask_header_value(value);
}
}
serde_json::to_value(headers).unwrap_or_else(|_| json!({}))
serde_json::to_value(crate::headers::collect_control_headers(headers))
.unwrap_or_else(|_| json!({}))
}
fn runtime_miss_original_request_body_json(
@@ -784,40 +779,6 @@ fn runtime_miss_original_request_body_json(
})
}
fn runtime_miss_sensitive_header(name: &str) -> bool {
const SENSITIVE_HEADERS: &[&str] = &[
"authorization",
"x-api-key",
"api-key",
"x-goog-api-key",
"cookie",
"proxy-authorization",
];
SENSITIVE_HEADERS
.iter()
.any(|candidate| name.eq_ignore_ascii_case(candidate))
}
fn runtime_miss_mask_header_value(value: &str) -> String {
let value = value.trim();
let char_count = value.chars().count();
if char_count <= 8 {
return "****".to_string();
}
let prefix: String = value.chars().take(4).collect();
let suffix: String = value
.chars()
.rev()
.take(4)
.collect::<Vec<_>>()
.into_iter()
.rev()
.collect();
format!("{prefix}****{suffix}")
}
async fn load_runtime_miss_candidate_contexts(
state: &AppState,
request_id: &str,
@@ -1233,8 +1194,9 @@ mod tests {
apply_runtime_miss_usage_routing, beautify_local_execution_client_error_message,
insert_runtime_miss_candidate_usage_metadata,
request_candidate_represents_provider_execution, runtime_miss_client_error_body,
select_last_runtime_miss_executed_candidate, select_last_runtime_miss_routing_candidate,
LocalExecutionRuntimeMissContext, RuntimeMissCandidateContext,
runtime_miss_original_headers_json, select_last_runtime_miss_executed_candidate,
select_last_runtime_miss_routing_candidate, LocalExecutionRuntimeMissContext,
RuntimeMissCandidateContext,
};
use crate::constants::EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS;
use crate::state::LocalExecutionRuntimeMissDiagnostic;
@@ -1266,6 +1228,30 @@ mod tests {
);
}
#[test]
fn runtime_miss_usage_preserves_original_request_headers() {
let expected = json!({
"authorization": "Bearer original-client-token",
"x-api-key": "short",
"api-key": "original-api-key",
"x-goog-api-key": "original-google-key",
"cookie": "session=original-client",
"proxy-authorization": "Basic original-proxy-token",
"originator": "codex-cli",
"session-id": "original-session",
"x-codex-turn-metadata": "{\"turn_id\":\"original-turn\"}"
});
let mut headers = http::HeaderMap::new();
for (name, value) in expected.as_object().unwrap() {
headers.insert(
http::HeaderName::from_bytes(name.as_bytes()).unwrap(),
http::HeaderValue::from_str(value.as_str().unwrap()).unwrap(),
);
}
assert_eq!(runtime_miss_original_headers_json(&headers), expected);
}
#[test]
fn runtime_miss_usage_body_matches_claude_client_envelope() {
let claude = runtime_miss_client_error_body(Some("claude:messages"), "busy");
@@ -339,57 +339,6 @@ 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<'_>,
@@ -689,3 +638,54 @@ pub(crate) async fn maybe_build_local_admin_oauth_response(
Ok(None)
}
#[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());
}
}
}
@@ -278,46 +278,6 @@ 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<'_>,
@@ -357,3 +317,43 @@ async fn build_assign_to_providers_response(
&global_model_id,
))
}
#[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());
}
}
@@ -58,6 +58,7 @@ async fn admin_monitoring_trace_request_returns_local_payload() {
.expect("route should be handled locally");
assert_eq!(response.status(), http::StatusCode::OK);
assert!(response.headers().contains_key("x-aether-build-version"));
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
@@ -93,6 +94,15 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
Some(33),
Some(200),
),
sample_candidate(
"cand-other-attempt",
"trace-1",
1,
RequestCandidateStatus::Failed,
Some(100),
Some(20),
Some(502),
),
]));
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
@@ -110,6 +120,8 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
100,
);
usage.id = "usage-row-1".to_string();
usage.request_body_state = Some(UsageBodyCaptureState::Reference);
usage.response_body_state = Some(UsageBodyCaptureState::Reference);
usage.candidate_id = Some("cand-used".to_string());
usage.request_headers = Some(json!({
"x-trace-id": "trace-1"
@@ -140,6 +152,17 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["request_id"], json!("trace-1"));
assert_eq!(payload["diagnostic_request"]["usage_id"], "usage-row-1");
assert_eq!(
payload["candidates"][0]["extra_data"]["diagnostic_context"]["usage_id"],
"usage-row-1"
);
assert_eq!(
payload["candidates"][0]["extra_data"]["diagnostic_context"]["body_states"]
["response_body"],
"reference"
);
assert!(payload["candidates"][1]["extra_data"]["diagnostic_context"].is_null());
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
assert_eq!(
payload["candidates"][0]["extra_data"]["first_byte_time_ms"],
@@ -67,13 +67,19 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
let key_accounts =
build_admin_monitoring_key_account_display_map(admin_state, &resolved.trace).await?;
Ok(
build_admin_monitoring_trace_request_payload_response_with_key_accounts(
&resolved.trace,
resolved.usage.as_ref(),
&key_accounts,
),
)
let mut response = build_admin_monitoring_trace_request_payload_response_with_key_accounts(
&resolved.trace,
resolved.usage.as_ref(),
&key_accounts,
);
if let Ok(version) = axum::http::HeaderValue::from_str(
option_env!("AETHER_BUILD_VERSION").unwrap_or(env!("CARGO_PKG_VERSION")),
) {
response
.headers_mut()
.insert("x-aether-build-version", version);
}
Ok(response)
}
async fn resolve_admin_monitoring_trace(
@@ -17,7 +17,10 @@ use aether_admin::observability::usage::{
admin_usage_bad_request_response, admin_usage_data_unavailable_response,
admin_usage_provider_key_name, ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
};
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageBodyField};
use aether_data_contracts::repository::usage::{
canonical_usage_body_ref_for, StoredRequestUsageAudit, StoredUsageBodyPayload,
UsageBodyCaptureState, UsageBodyField, MAX_DECOMPRESSED_USAGE_JSON_BYTES,
};
use axum::{
body::Body,
http,
@@ -28,9 +31,59 @@ use serde_json::{json, Value};
use std::collections::BTreeMap;
use tokio::try_join;
#[derive(Default)]
struct AdminUsageDetailBodyValue {
value: Option<Value>,
load_failed: bool,
error_code: Option<&'static str>,
}
impl AdminUsageDetailBodyValue {
fn resolved(
item: &StoredRequestUsageAudit,
field: UsageBodyField,
value: Option<Value>,
) -> Self {
let missing = value.is_none()
&& item
.body_capture_result(field, item.body_value(field))
.available;
Self {
value,
error_code: missing.then_some("missing"),
}
}
}
fn admin_usage_body_load_error_code(error: &GatewayError) -> &'static str {
if let GatewayError::Internal(message) = error {
if message.contains("decompressed usage json exceeds ")
|| message.contains("encoded usage json exceeds ")
{
return "too_large";
}
if message.contains("failed to decompress usage json:")
|| message.contains("failed to parse decompressed usage json:")
{
return "decode_failed";
}
}
"storage_unavailable"
}
async fn resolve_admin_usage_detail_field(
state: &AdminAppState<'_>,
item: &StoredRequestUsageAudit,
field: UsageBodyField,
selected_field: Option<UsageBodyField>,
) -> AdminUsageDetailBodyValue {
if selected_field.is_some_and(|selected| selected != field) {
return AdminUsageDetailBodyValue::default();
}
if field == UsageBodyField::RequestBody {
resolve_admin_usage_detail_request_body(state, item).await
} else {
resolve_admin_usage_detail_body_value(state, item, field).await
}
}
async fn resolve_admin_usage_detail_request_body(
@@ -38,10 +91,7 @@ async fn resolve_admin_usage_detail_request_body(
item: &StoredRequestUsageAudit,
) -> AdminUsageDetailBodyValue {
match admin_usage_resolve_request_capture_body_for_item(state, item, None).await {
Ok(body) => AdminUsageDetailBodyValue {
value: body,
load_failed: false,
},
Ok(body) => AdminUsageDetailBodyValue::resolved(item, UsageBodyField::RequestBody, body),
Err(err) => {
tracing::warn!(
error = ?err,
@@ -52,7 +102,9 @@ async fn resolve_admin_usage_detail_request_body(
);
let value = admin_usage_resolve_request_capture_body(item, None);
AdminUsageDetailBodyValue {
load_failed: value.is_none(),
error_code: value
.is_none()
.then(|| admin_usage_body_load_error_code(&err)),
value,
}
}
@@ -66,10 +118,7 @@ async fn resolve_admin_usage_detail_body_value(
) -> AdminUsageDetailBodyValue {
let inline_body = item.body_value(field);
match admin_usage_resolve_body_value(state, item, inline_body, field).await {
Ok(body) => AdminUsageDetailBodyValue {
value: body,
load_failed: false,
},
Ok(body) => AdminUsageDetailBodyValue::resolved(item, field, body),
Err(err) => {
tracing::warn!(
error = ?err,
@@ -80,13 +129,139 @@ async fn resolve_admin_usage_detail_body_value(
);
let value = inline_body.cloned();
AdminUsageDetailBodyValue {
load_failed: value.is_none(),
error_code: value
.is_none()
.then(|| admin_usage_body_load_error_code(&err)),
value,
}
}
}
}
async fn read_admin_usage_raw_body(
state: &AdminAppState<'_>,
item: &StoredRequestUsageAudit,
field: UsageBodyField,
) -> Result<Option<StoredUsageBodyPayload>, GatewayError> {
if matches!(
item.body_state(field),
Some(
UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
| UsageBodyCaptureState::None
)
) {
return Ok(None);
}
let inline_body = item.body_value(field);
let prefer_inline = matches!(
item.body_state(field),
Some(UsageBodyCaptureState::Inline | UsageBodyCaptureState::Truncated)
) && inline_body.is_some();
if !prefer_inline {
if let Some(body_ref) = item
.body_ref(field)
.and_then(|reference| canonical_usage_body_ref_for(reference, &item.request_id, field))
{
if let Some(payload) = state.read_request_usage_body_payload(&body_ref).await? {
return Ok(Some(payload));
}
}
}
let fallback = inline_body.cloned().or_else(|| {
(field == UsageBodyField::RequestBody)
.then(|| admin_usage_resolve_request_capture_body(item, None))
.flatten()
});
fallback
.map(|value| {
serde_json::to_vec(&value)
.map(StoredUsageBodyPayload::Json)
.map_err(|error| GatewayError::Internal(error.to_string()))
})
.transpose()
}
async fn build_admin_usage_raw_body_response(
state: &AdminAppState<'_>,
item: &StoredRequestUsageAudit,
field: UsageBodyField,
) -> Response<Body> {
let result = read_admin_usage_raw_body(state, item, field).await;
let mut response = match result {
Ok(Some(payload)) => admin_usage_raw_payload_response(payload),
Ok(None) => admin_usage_raw_body_error(http::StatusCode::NOT_FOUND, "missing"),
Err(error) => {
tracing::warn!(error = ?error, usage_id = %item.id, field = field.as_storage_field(), "failed to read admin usage raw body");
let code = admin_usage_body_load_error_code(&error);
admin_usage_raw_body_error(
if code == "too_large" {
http::StatusCode::PAYLOAD_TOO_LARGE
} else {
http::StatusCode::SERVICE_UNAVAILABLE
},
code,
)
}
};
let headers = response.headers_mut();
headers.insert(
http::header::CACHE_CONTROL,
http::HeaderValue::from_static("no-store, no-transform"),
);
headers.insert(
"x-content-type-options",
http::HeaderValue::from_static("nosniff"),
);
headers.insert(
"x-aether-body-field",
http::HeaderValue::from_static(field.as_storage_field()),
);
if let Ok(value) = http::HeaderValue::from_str(&item.id) {
headers.insert("x-aether-usage-id", value);
}
attach_admin_audit_response(
response,
"admin_usage_detail_viewed",
"view_usage_detail",
"usage_record",
&item.id,
)
}
fn admin_usage_raw_payload_response(payload: StoredUsageBodyPayload) -> Response<Body> {
let (encoding, bytes, limit) = match payload {
StoredUsageBodyPayload::Gzip(bytes) => (
"gzip",
bytes,
MAX_DECOMPRESSED_USAGE_JSON_BYTES + 1024 * 1024,
),
StoredUsageBodyPayload::Json(bytes) => ("json", bytes, MAX_DECOMPRESSED_USAGE_JSON_BYTES),
};
if bytes.len() > limit {
admin_usage_raw_body_error(http::StatusCode::PAYLOAD_TOO_LARGE, "too_large")
} else {
(
[
("content-type", "application/octet-stream"),
("content-encoding", "identity"),
("x-aether-body-encoding", encoding),
],
bytes,
)
.into_response()
}
}
fn admin_usage_raw_body_error(status: http::StatusCode, code: &'static str) -> Response<Body> {
(
status,
[("x-aether-body-error", code)],
Json(json!({ "body_load_error_code": code })),
)
.into_response()
}
pub(super) async fn maybe_build_local_admin_usage_detail_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -212,6 +387,25 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
"include_bodies",
true,
);
let body_field =
match request_context
.request_query_string
.as_deref()
.and_then(|query| {
url::form_urlencoded::parse(query.as_bytes())
.find(|(key, _)| key == "body_field")
.map(|(_, value)| value.into_owned())
}) {
Some(value) => {
match UsageBodyField::from_storage_field(value.trim()) {
Some(field) if include_bodies => Some(field),
_ => return Ok(Some(admin_usage_bad_request_response(
"body_field 必须是有效的正文字段,且 include_bodies 必须为 true",
))),
}
}
None => None,
};
let Some(item) = state.find_request_usage_by_id(&usage_id).await? else {
return Ok(Some(
@@ -223,6 +417,27 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
));
};
let body_format = request_context
.request_query_string
.as_deref()
.and_then(|query| {
url::form_urlencoded::parse(query.as_bytes())
.find(|(key, _)| key == "body_format")
.map(|(_, value)| value.into_owned())
});
if let Some(format) = body_format {
if format != "raw" || body_field.is_none() {
return Ok(Some(admin_usage_bad_request_response(
"body_format=raw 必须指定 body_field",
)));
}
if let Some(field) = body_field {
return Ok(Some(
build_admin_usage_raw_body_response(state, &item, field).await,
));
}
}
let user_ids = item.user_id.clone().into_iter().collect::<Vec<_>>();
let (users_by_id, provider_key_names, api_key_names): (
BTreeMap<String, aether_data::repository::users::StoredUserSummary>,
@@ -250,23 +465,32 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
}
}
let mut body_load_errors = serde_json::Map::new();
let request_body = if include_bodies {
let mut body_load_error_codes = serde_json::Map::new();
let mut request_body = if include_bodies {
let (request_body, provider_request_body, response_body, client_response_body) = tokio::join!(
resolve_admin_usage_detail_request_body(state, &item),
resolve_admin_usage_detail_body_value(
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::RequestBody,
body_field
),
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::ProviderRequestBody,
body_field,
),
resolve_admin_usage_detail_body_value(
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::ResponseBody,
body_field,
),
resolve_admin_usage_detail_body_value(
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::ClientResponseBody,
body_field,
),
);
for (field, resolved) in [
@@ -275,20 +499,25 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
(UsageBodyField::ResponseBody, &response_body),
(UsageBodyField::ClientResponseBody, &client_response_body),
] {
if resolved.load_failed {
if let Some(error_code) = resolved.error_code {
body_load_errors.insert(field.as_storage_field().to_string(), json!(true));
body_load_error_codes
.insert(field.as_storage_field().to_string(), json!(error_code));
}
}
detail_item.provider_request_body = provider_request_body.value;
detail_item.response_body = response_body.value;
detail_item.client_response_body = client_response_body.value;
if body_field.is_none_or(|field| field == UsageBodyField::ProviderRequestBody) {
detail_item.provider_request_body = provider_request_body.value;
}
if body_field.is_none_or(|field| field == UsageBodyField::ResponseBody) {
detail_item.response_body = response_body.value;
}
if body_field.is_none_or(|field| field == UsageBodyField::ClientResponseBody) {
detail_item.client_response_body = client_response_body.value;
}
request_body.value
} else {
None
};
if include_bodies {
// request_body 已通过 request capture 解析;其余 detached body 在上方并行加载。
}
let default_headers = admin_usage_curl_headers();
let mut payload = build_admin_usage_detail_payload(
&detail_item,
@@ -297,15 +526,33 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
state.has_auth_user_data_reader(),
state.has_auth_api_key_data_reader(),
provider_key_name.as_deref(),
include_bodies,
request_body,
include_bodies && body_field.is_none(),
if body_field.is_none() {
request_body.take()
} else {
None
},
&default_headers,
);
if let Some(field) = body_field {
payload[field.as_storage_field()] = match field {
UsageBodyField::RequestBody => request_body,
UsageBodyField::ProviderRequestBody => detail_item.provider_request_body.take(),
UsageBodyField::ResponseBody => detail_item.response_body.take(),
UsageBodyField::ClientResponseBody => detail_item.client_response_body.take(),
}
.unwrap_or(Value::Null);
}
payload["body_load_errors"] = if include_bodies && !body_load_errors.is_empty() {
Value::Object(body_load_errors)
} else {
Value::Null
};
payload["body_load_error_codes"] = if body_load_error_codes.is_empty() {
Value::Null
} else {
Value::Object(body_load_error_codes)
};
return Ok(Some(attach_admin_audit_response(
Json(payload).into_response(),
@@ -320,3 +567,61 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
Ok(None)
}
#[cfg(test)]
mod tests {
use super::admin_usage_body_load_error_code;
use crate::GatewayError;
#[tokio::test]
async fn admin_usage_raw_body_does_not_decode_or_reencode_stored_bytes() {
use super::{admin_usage_raw_payload_response, StoredUsageBodyPayload};
for (payload, encoding, expected) in [
(
StoredUsageBodyPayload::Gzip(vec![31, 139, 8, 0, 1]),
"gzip",
vec![31, 139, 8, 0, 1],
),
(
StoredUsageBodyPayload::Json(b"{ \"untouched\" : true }".to_vec()),
"json",
b"{ \"untouched\" : true }".to_vec(),
),
] {
let response = admin_usage_raw_payload_response(payload);
assert_eq!(response.headers()["content-encoding"], "identity");
assert_eq!(response.headers()["x-aether-body-encoding"], encoding);
let bytes = axum::body::to_bytes(response.into_body(), 1024)
.await
.unwrap();
assert_eq!(bytes.as_ref(), expected.as_slice());
}
}
#[test]
fn body_load_errors_expose_safe_codes_instead_of_internal_messages() {
for (message, expected) in [
(
"unexpected database value: decompressed usage json exceeds 67108864 bytes",
"too_large",
),
(
"failed to decompress usage json: invalid gzip header",
"decode_failed",
),
(
"failed to parse decompressed usage json: invalid JSON",
"decode_failed",
),
(
"postgres error: private connection details",
"storage_unavailable",
),
] {
assert_eq!(
admin_usage_body_load_error_code(&GatewayError::Internal(message.to_string())),
expected
);
}
}
}
@@ -103,10 +103,11 @@ pub(crate) async fn maybe_build_local_admin_provider_reads_response(
.build_admin_provider_summary_payload(&provider_id)
.await
{
Some(payload) => Json(payload).into_response(),
None => build_admin_provider_not_found_response(format!(
Ok(Some(payload)) => Json(payload).into_response(),
Ok(None) => build_admin_provider_not_found_response(format!(
"Provider {provider_id} 不存在"
)),
Err(_) => build_admin_providers_data_unavailable_response(),
},
));
}
@@ -173,16 +173,17 @@ pub(crate) async fn maybe_build_local_admin_provider_writes_response(
.build_admin_provider_summary_payload(&provider_id)
.await
{
Some(payload) => attach_admin_audit_response(
Ok(Some(payload)) => attach_admin_audit_response(
Json(payload).into_response(),
"admin_provider_updated",
"update_provider",
"provider",
&provider_id,
),
None => build_admin_provider_not_found_response(format!(
Ok(None) => build_admin_provider_not_found_response(format!(
"Provider {provider_id} 不存在"
)),
Err(_) => build_admin_providers_data_unavailable_response(),
},
));
}
@@ -6,6 +6,7 @@ mod extractors;
mod list;
pub(crate) mod payloads;
mod reads;
mod reveal;
mod support;
mod update;
@@ -41,6 +42,10 @@ pub(crate) async fn maybe_build_local_admin_endpoints_routes_response(
return Ok(Some(response));
}
if let Some(response) = reveal::maybe_handle(state, request_context).await? {
return Ok(Some(response));
}
if let Some(response) = defaults::maybe_handle(state, request_context, request_body).await? {
return Ok(Some(response));
}
@@ -0,0 +1,72 @@
use super::extractors::admin_endpoint_id;
use super::support::build_admin_endpoints_data_unavailable_response;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{
attach_admin_audit_response, mark_sensitive_admin_response_no_store,
};
use crate::GatewayError;
use axum::{
body::Body,
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.decision() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("endpoints_manage")
|| decision.route_kind.as_deref() != Some("reveal_endpoint_rules")
{
return Ok(None);
}
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
let Some(endpoint_id) = request_context
.path()
.strip_suffix("/rules/reveal")
.and_then(admin_endpoint_id)
else {
return Ok(Some(
(
StatusCode::NOT_FOUND,
Json(json!({ "detail": "Endpoint 不存在" })),
)
.into_response(),
));
};
let Some(endpoint) = state
.read_provider_catalog_endpoints_by_ids(std::slice::from_ref(&endpoint_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
StatusCode::NOT_FOUND,
Json(json!({ "detail": "Endpoint 不存在" })),
)
.into_response(),
));
};
let payload = json!({
"header_rules": endpoint.header_rules.as_ref().and_then(|value| value.as_array()).cloned().unwrap_or_default(),
"body_rules": endpoint.body_rules.as_ref().and_then(|value| value.as_array()).cloned().unwrap_or_default(),
"response_header_rules": endpoint.config.as_ref().and_then(|config| config.get("response_header_rules")).and_then(|value| value.as_array()).cloned().unwrap_or_default(),
});
Ok(Some(mark_sensitive_admin_response_no_store(
attach_admin_audit_response(
Json(payload).into_response(),
"admin_endpoint_rules_revealed",
"reveal_endpoint_rules",
"provider_endpoint",
&endpoint_id,
),
)))
}
@@ -1463,7 +1463,7 @@ mod tests {
&auth_config,
Some(0),
),
"antigravity_[email protected]"
"[email protected]"
);
}
@@ -1,3 +1,4 @@
use super::super::helpers::admin_provider_oauth_key_name_from_auth_config;
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,
@@ -79,7 +80,7 @@ fn kiro_social_key_name(
.collect::<String>()
})
.unwrap_or_else(|| "unknown".to_string());
format!("kiro_{fallback} ({provider})")
format!("账号_{fallback} ({provider})")
}
fn kiro_social_poll_error_response(error: impl Into<String>) -> Response<Body> {
@@ -1004,10 +1005,11 @@ async fn handle_admin_provider_oauth_windsurf_browser_device_poll(
}
}
} else {
let key_name = email
.as_deref()
.map(|email| format!("windsurf_{email}"))
.unwrap_or_else(|| format!("windsurf_{}", current_unix_secs()));
let key_name = admin_provider_oauth_key_name_from_auth_config(
&provider.provider_type,
&auth_config,
None,
);
match state
.create_provider_oauth_catalog_key(
&provider.id,
@@ -1356,6 +1358,32 @@ mod tests {
use crate::control::GatewayAdminPrincipalContext;
use aether_data::repository::provider_oauth::StoredAdminProviderOAuthDeviceSession;
#[test]
fn kiro_social_key_name_preserves_email_and_auth_method() {
assert_eq!(
super::kiro_social_key_name(
Some(" [email protected] "),
Some("Github"),
Some("refresh-token-1"),
),
"[email protected] (Github)"
);
}
#[test]
fn kiro_social_key_name_without_email_uses_generic_account_prefix() {
for email in [None, Some(""), Some(" ")] {
assert_eq!(
super::kiro_social_key_name(email, Some("Google"), Some("refresh-token-1")),
"账号_154f43 (Google)"
);
assert_eq!(
super::kiro_social_key_name(email, None, None),
"账号_unknown (social)"
);
}
}
fn device_session() -> StoredAdminProviderOAuthDeviceSession {
StoredAdminProviderOAuthDeviceSession {
session_id: "device-session-1".to_string(),
@@ -52,13 +52,12 @@ pub(super) fn admin_provider_oauth_key_name_from_auth_config(
auth_config: &Map<String, Value>,
batch_index: Option<usize>,
) -> String {
let provider_type = provider_type.trim();
if let Some(email) = trimmed_auth_config_string(auth_config, "email") {
return format!("{provider_type}_{email}");
return email;
}
if provider_type.eq_ignore_ascii_case("grok") {
if provider_type.trim().eq_ignore_ascii_case("grok") {
if let Some(user_id) = trimmed_auth_config_string(auth_config, "user_id") {
return format!("grok_{user_id}");
return user_id;
}
}
@@ -68,7 +67,7 @@ pub(super) fn admin_provider_oauth_key_name_from_auth_config(
.map(|duration| duration.as_secs())
.unwrap_or(0);
match batch_index {
Some(index) => format!("{provider_type}_{timestamp}_{index}"),
Some(index) => format!("账号_{timestamp}_{index}"),
None => format!("账号_{timestamp}"),
}
}
@@ -87,6 +86,106 @@ mod tests {
use super::*;
use serde_json::{json, Map};
const PROVIDER_TYPES: &[&str] = &[
"codex",
" Codex ",
"claude_code",
"chatgpt_web",
"gemini_cli",
"antigravity",
"grok",
" Grok ",
"kiro",
"windsurf",
];
#[test]
fn default_key_name_uses_email_without_provider_prefix() {
let mut auth_config = Map::new();
auth_config.insert("email".to_string(), json!(" [email protected] "));
for provider_type in PROVIDER_TYPES {
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
"[email protected]"
);
}
}
}
#[test]
fn antigravity_default_key_name_uses_email_without_provider_prefix() {
for email in [" [email protected] ", "[email protected]"] {
let mut auth_config = Map::new();
auth_config.insert("email".to_string(), json!(email));
for provider_type in ["antigravity", " Antigravity "] {
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
email.trim()
);
}
}
}
}
#[test]
fn default_key_name_preserves_email_with_provider_prefix() {
for provider_type in PROVIDER_TYPES {
let email = format!("{}[email protected]", provider_type.trim());
let mut auth_config = Map::new();
auth_config.insert("email".to_string(), json!(email));
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
email
);
}
}
}
#[test]
fn default_key_name_without_email_uses_generic_account_name() {
for email in [None, Some(""), Some(" ")] {
let mut auth_config = Map::new();
if let Some(email) = email {
auth_config.insert("email".to_string(), json!(email));
}
for provider_type in PROVIDER_TYPES {
for batch_index in [None, Some(3)] {
let name = admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
);
let suffix = name.strip_prefix("账号_").expect("generic account prefix");
let timestamp = if batch_index.is_some() {
suffix.strip_suffix("_3").expect("batch index suffix")
} else {
suffix
};
assert!(timestamp.parse::<u64>().is_ok());
}
}
}
}
#[test]
fn grok_default_key_name_uses_full_user_id() {
let mut auth_config = Map::new();
@@ -95,10 +194,18 @@ mod tests {
json!("1619039a-0191-4e0a-a490-8f4ad21262c9"),
);
assert_eq!(
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
"grok_1619039a-0191-4e0a-a490-8f4ad21262c9"
);
for provider_type in ["grok", " Grok "] {
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
"1619039a-0191-4e0a-a490-8f4ad21262c9"
);
}
}
}
#[test]
@@ -109,17 +216,22 @@ mod tests {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
"grok_grok@example.com"
"[email protected]"
);
}
#[test]
fn batch_default_key_name_keeps_existing_timestamp_shape() {
fn batch_default_key_name_keeps_distinct_indexes_without_provider_prefix() {
let auth_config = Map::new();
let name = admin_provider_oauth_key_name_from_auth_config("codex", &auth_config, Some(3));
let name = admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, Some(3));
let other_name =
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, Some(4));
assert!(name.starts_with("codex_"));
assert!(name.starts_with("账号_"));
assert!(name.ends_with("_3"));
assert!(other_name.starts_with("账号_"));
assert!(other_name.ends_with("_4"));
assert_ne!(name, other_name);
}
#[test]
@@ -66,43 +66,6 @@ fn admin_provider_oauth_kiro_refresh_error(
}
}
#[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"));
}
}
pub(super) async fn refresh_admin_provider_oauth_kiro_auth_config(
state: &AdminAppState<'_>,
auth_config: &AdminKiroAuthConfig,
@@ -240,3 +203,40 @@ pub(super) async fn fetch_admin_provider_oauth_kiro_email(
aether_admin::provider::quota::parse_kiro_usage_response(&payload, current_unix_secs())?;
json_non_empty_string(metadata.get("email"))
}
#[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"));
}
}
@@ -5,7 +5,6 @@ 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::{
@@ -24,63 +23,6 @@ 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<String> {
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<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
@@ -380,10 +322,6 @@ 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 {
@@ -40,20 +40,6 @@ pub(super) fn pool_stream_timeout_key(provider_id: &str, key_id: &str) -> String
format!("ap:{provider_id}:stream_timeout:{key_id}")
}
pub(super) fn parse_pool_cost_member(member: &str) -> u64 {
member
.rsplit_once(':')
.and_then(|(_, suffix)| suffix.parse::<u64>().ok())
.unwrap_or(0)
}
pub(super) fn parse_pool_latency_member(member: &str) -> u64 {
member
.rsplit_once(':')
.and_then(|(_, suffix)| suffix.parse::<u64>().ok())
.unwrap_or(0)
}
pub(super) fn pool_cooldown_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
key_ids
.iter()
@@ -12,7 +12,8 @@ pub(crate) use self::mutations::{
pub(crate) use self::reads::{
read_admin_provider_pool_cooldown_count, read_admin_provider_pool_cooldown_counts,
read_admin_provider_pool_cooldown_key_ids, read_admin_provider_pool_key_cooldown_reason,
read_admin_provider_pool_runtime_state,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
read_provider_pool_sticky_bound_key_id,
};
pub(crate) use self::status::build_admin_provider_pool_status_payload;
pub(crate) use self::writes::{
@@ -1,7 +1,6 @@
use super::keys::{
parse_pool_cost_member, parse_pool_latency_member, pool_cooldown_index_key, pool_cooldown_key,
pool_cooldown_keys, pool_cost_keys, pool_latency_keys, pool_lru_key, pool_sticky_key,
pool_sticky_pattern,
pool_cooldown_index_key, pool_cooldown_key, pool_cooldown_keys, pool_cost_keys,
pool_latency_keys, pool_lru_key, pool_sticky_key, pool_sticky_pattern,
};
use crate::handlers::admin::provider::pool::config::admin_provider_pool_cache_affinity_enabled;
use crate::handlers::admin::provider::shared::support::{
@@ -12,7 +11,8 @@ use crate::maintenance::PoolQuotaProbeWorkerConfig;
use crate::provider_pool_demand::{
provider_pool_burst_pending, read_provider_pool_demand_snapshot,
};
use aether_runtime_state::{DataLayerError, RuntimeState};
use aether_pool_core::{normalize_enabled_pool_presets, PoolSchedulingPreset};
use aether_runtime_state::{DataLayerError, RuntimeState, ScoreWindowU64Stats};
use futures_util::future::join_all;
use std::collections::{BTreeMap, BTreeSet};
use std::time::{SystemTime, UNIX_EPOCH};
@@ -48,6 +48,103 @@ fn bounded_runtime_window_metric_key_ids(key_ids: &[String], limit: usize) -> &[
&key_ids[..end]
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum PoolRuntimeReadPurpose {
Admin,
Scheduling,
}
fn scheduling_window_metrics(pool_config: &AdminProviderPoolConfig) -> (bool, bool) {
let presets = pool_config
.scheduling_presets
.iter()
.map(|preset| PoolSchedulingPreset {
preset: preset.preset.clone(),
enabled: preset.enabled,
mode: preset.mode.clone(),
})
.collect::<Vec<_>>();
let active = normalize_enabled_pool_presets(&presets);
let cost = pool_config.cost_limit_per_key_tokens.is_some()
|| active
.iter()
.any(|preset| matches!(preset.as_str(), "cost_first" | "quota_balanced"));
let latency = active.iter().any(|preset| preset == "latency_first");
(cost, latency)
}
async fn read_window_stats(
runtime: &RuntimeState,
keys: &[String],
min_score: f64,
) -> Vec<ScoreWindowU64Stats> {
let aggregates = match runtime.score_window_u64_stats_by_min(keys, min_score).await {
Ok(values) => values,
Err(err) => {
warn!(
"gateway provider pool: bounded window aggregation failed, using exact range reads: {err:?}"
);
vec![None; keys.len()]
}
};
join_all(keys.iter().zip(aggregates).map(|(key, stats)| async move {
match stats {
Some(stats) => stats,
// Large windows and failed aggregation retain the original exact
// read. A missing aggregate must never be treated as zero cost.
None => {
let members = runtime
.score_range_by_min(key, min_score)
.await
.unwrap_or_default();
ScoreWindowU64Stats::from_members(members.iter().map(String::as_str))
}
}
}))
.await
}
pub(crate) async fn read_provider_pool_sticky_bound_key_id(
runtime: &RuntimeState,
provider_id: &str,
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
) -> Option<String> {
if pool_config.sticky_session_ttl_seconds == 0
|| !admin_provider_pool_cache_affinity_enabled(pool_config)
{
return None;
}
let sticky_session_token = sticky_session_token
.map(str::trim)
.filter(|value| !value.is_empty())?;
let sticky_key = pool_sticky_key(provider_id, sticky_session_token);
let bound_key_id = runtime.kv_get(&sticky_key).await.ok().flatten()?;
let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id);
match runtime.kv_exists(&cooldown_key).await {
Ok(false) => {
let _ = runtime
.key_expire(
&sticky_key,
std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds),
)
.await;
Some(bound_key_id)
}
Ok(true) => {
let _ = runtime.kv_delete(&sticky_key).await;
None
}
Err(err) => {
warn!(
"gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}",
err
);
Some(bound_key_id)
}
}
}
pub(crate) async fn read_admin_provider_pool_cooldown_counts(
runtime: &RuntimeState,
provider_ids: &[String],
@@ -71,9 +168,52 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
) -> AdminProviderPoolRuntimeState {
read_provider_pool_runtime_state(
runtime,
provider_id,
key_ids,
pool_config,
sticky_session_token,
PoolRuntimeReadPurpose::Admin,
)
.await
}
pub(crate) async fn read_provider_pool_scheduling_runtime_state(
runtime: &RuntimeState,
provider_id: &str,
key_ids: &[String],
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
) -> AdminProviderPoolRuntimeState {
read_provider_pool_runtime_state(
runtime,
provider_id,
key_ids,
pool_config,
sticky_session_token,
PoolRuntimeReadPurpose::Scheduling,
)
.await
}
async fn read_provider_pool_runtime_state(
runtime: &RuntimeState,
provider_id: &str,
key_ids: &[String],
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
purpose: PoolRuntimeReadPurpose,
) -> AdminProviderPoolRuntimeState {
let include_admin_metrics = purpose == PoolRuntimeReadPurpose::Admin;
let mut state = AdminProviderPoolRuntimeState::default();
let cooldown_keys = pool_cooldown_keys(provider_id, key_ids);
let metric_key_limit = pool_runtime_window_metric_key_limit();
let metric_key_limit = if include_admin_metrics {
pool_runtime_window_metric_key_limit()
} else {
key_ids.len()
};
// The admin display cap must not hide a candidate's strict cost limit.
let metric_key_ids = bounded_runtime_window_metric_key_ids(key_ids, metric_key_limit);
if metric_key_ids.len() < key_ids.len() {
info!(
@@ -86,44 +226,33 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
"gateway limited admin pool runtime cost/latency window reads"
);
}
let cost_keys = pool_cost_keys(provider_id, metric_key_ids);
let latency_keys = pool_latency_keys(provider_id, metric_key_ids);
let (load_cost, load_latency) = if include_admin_metrics {
(true, true)
} else {
scheduling_window_metrics(pool_config)
};
let cost_keys = if load_cost {
pool_cost_keys(provider_id, metric_key_ids)
} else {
Vec::new()
};
let latency_keys = if load_latency {
pool_latency_keys(provider_id, metric_key_ids)
} else {
Vec::new()
};
let sticky_sessions_enabled = pool_config.sticky_session_ttl_seconds > 0
&& admin_provider_pool_cache_affinity_enabled(pool_config);
if let Some(sticky_session_token) = sticky_session_token
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|_| sticky_sessions_enabled)
{
let sticky_key = pool_sticky_key(provider_id, sticky_session_token);
if let Ok(Some(bound_key_id)) = runtime.kv_get(&sticky_key).await {
let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id);
match runtime.kv_exists(&cooldown_key).await {
Ok(false) => {
let _ = runtime
.key_expire(
&sticky_key,
std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds),
)
.await;
state.sticky_bound_key_id = Some(bound_key_id);
}
Ok(true) => {
let _ = runtime.kv_delete(&sticky_key).await;
}
Err(err) => {
warn!(
"gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}",
err
);
state.sticky_bound_key_id = Some(bound_key_id);
}
}
}
}
state.sticky_bound_key_id = read_provider_pool_sticky_bound_key_id(
runtime,
provider_id,
pool_config,
sticky_session_token,
)
.await;
if sticky_sessions_enabled {
if include_admin_metrics && sticky_sessions_enabled {
let sticky_keys = runtime
.scan_keys(&pool_sticky_pattern(provider_id), 200)
.await
@@ -161,23 +290,25 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
.unwrap_or_default();
}
let probe_config = PoolQuotaProbeWorkerConfig::from_env();
let demand_snapshot = read_provider_pool_demand_snapshot(
runtime,
provider_id,
key_ids.len(),
probe_config.max_keys_per_provider,
)
.await;
state.provider_in_flight = demand_snapshot.in_flight;
state.provider_ema_in_flight = demand_snapshot.ema_in_flight;
state.provider_desired_hot = if pool_config.probing_enabled {
demand_snapshot.desired_hot
} else {
0
};
state.provider_burst_pending =
pool_config.probing_enabled && provider_pool_burst_pending(runtime, provider_id).await;
if include_admin_metrics || pool_config.probing_enabled {
let probe_config = PoolQuotaProbeWorkerConfig::from_env();
let demand_snapshot = read_provider_pool_demand_snapshot(
runtime,
provider_id,
key_ids.len(),
probe_config.max_keys_per_provider,
)
.await;
state.provider_in_flight = demand_snapshot.in_flight;
state.provider_ema_in_flight = demand_snapshot.ema_in_flight;
state.provider_desired_hot = if pool_config.probing_enabled {
demand_snapshot.desired_hot
} else {
0
};
state.provider_burst_pending =
pool_config.probing_enabled && provider_pool_burst_pending(runtime, provider_id).await;
}
if !cooldown_keys.is_empty() {
let cooldown_reasons = runtime
@@ -190,12 +321,14 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
{
if let Some(reason) = reason {
state.cooldown_reason_by_key.insert(key_id.clone(), reason);
if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await {
if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 {
state
.cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds);
if include_admin_metrics {
if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await {
if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 {
state
.cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds);
}
}
}
}
@@ -205,42 +338,24 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
let now = current_unix_secs();
let cost_window_start = now.saturating_sub(pool_config.cost_window_seconds) as f64;
let cost_results = join_all(
cost_keys
.iter()
.map(|cost_key| runtime.score_range_by_min(cost_key, cost_window_start)),
)
.await;
for (key_id, members) in metric_key_ids.iter().zip(cost_results) {
let total = members
.unwrap_or_default()
.iter()
.map(|member| parse_pool_cost_member(member))
.sum::<u64>();
if total > 0 {
state.cost_window_usage_by_key.insert(key_id.clone(), total);
let latency_window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
let (cost_results, latency_results) = tokio::join!(
read_window_stats(runtime, &cost_keys, cost_window_start),
read_window_stats(runtime, &latency_keys, latency_window_start),
);
for (key_id, stats) in metric_key_ids.iter().zip(cost_results) {
if stats.sum > 0 {
state
.cost_window_usage_by_key
.insert(key_id.clone(), stats.sum);
}
}
let latency_window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
let latency_results = join_all(
latency_keys
.iter()
.map(|latency_key| runtime.score_range_by_min(latency_key, latency_window_start)),
)
.await;
for (key_id, members) in metric_key_ids.iter().zip(latency_results) {
let samples = members
.unwrap_or_default()
.iter()
.map(|member| parse_pool_latency_member(member))
.filter(|value| *value > 0)
.collect::<Vec<_>>();
if samples.is_empty() {
for (key_id, stats) in metric_key_ids.iter().zip(latency_results) {
if stats.positive_count == 0 {
continue;
}
let total = samples.iter().sum::<u64>() as f64;
let average = total / samples.len() as f64;
let average = stats.sum as f64 / stats.positive_count as f64;
if average.is_finite() && average >= 0.0 {
state.latency_avg_ms_by_key.insert(key_id.clone(), average);
}
@@ -300,7 +415,358 @@ pub(crate) async fn read_admin_provider_pool_key_cooldown_reason(
#[cfg(test)]
mod tests {
use super::bounded_runtime_window_metric_key_ids;
use super::super::keys::{pool_cooldown_key, pool_cost_key, pool_latency_key, pool_sticky_key};
use super::{
bounded_runtime_window_metric_key_ids, current_unix_secs,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
read_provider_pool_sticky_bound_key_id,
};
use crate::handlers::admin::provider::pool::config::admin_provider_pool_config_from_config_value;
use crate::handlers::admin::provider::shared::support::AdminProviderPoolConfig;
use aether_runtime_state::{MemoryRuntimeStateConfig, RedisClientConfig, RuntimeState};
use aether_test_support::ManagedRedisServer;
use serde_json::json;
use std::time::Duration;
fn config(value: serde_json::Value) -> AdminProviderPoolConfig {
admin_provider_pool_config_from_config_value(Some(&json!({ "pool_advanced": value })))
.expect("pool config")
}
async fn seed_window_metrics(runtime: &RuntimeState, provider_id: &str, key_id: &str) {
let now = current_unix_secs() as f64;
for (key, member, timestamp) in [
(pool_cost_key(provider_id, key_id), "current:70", now),
(pool_cost_key(provider_id, key_id), "earlier:30", now - 1.0),
(
pool_cost_key(provider_id, key_id),
"expired:999",
now - 20_000.0,
),
(pool_latency_key(provider_id, key_id), "first:10", now),
(
pool_latency_key(provider_id, key_id),
"second:30",
now - 1.0,
),
] {
runtime
.score_set(&key, member, timestamp)
.await
.expect("seed window");
}
}
async fn admin_command_count(runtime: &RuntimeState) -> u64 {
runtime
.redis_diagnostics()
.await
.expect("diagnostics")
.expect("Redis runtime")
.lanes
.into_iter()
.find(|lane| lane.lane == "admin")
.expect("admin lane")
.command_count
}
#[tokio::test]
async fn scheduling_runtime_aggregates_bounded_windows_and_falls_back_for_large_windows() {
let redis = match ManagedRedisServer::start().await {
Ok(server) => server,
Err(err) if err.to_string().contains("No such file or directory") => {
eprintln!("skipping redis-backed scheduling runtime test: {err}");
return;
}
Err(err) => panic!("start Redis: {err}"),
};
let runtime = RuntimeState::redis(
RedisClientConfig {
url: redis.redis_url().to_string(),
key_prefix: Some("pool-window-aggregation-test".to_string()),
},
Some(1_000),
)
.await
.expect("runtime Redis");
let keys = vec!["bounded".to_string(), "large".to_string()];
let now = current_unix_secs() as f64;
for (key_id, count) in [(&keys[0], 512), (&keys[1], 2048)] {
let cost_key = pool_cost_key("pool", key_id);
for index in 0..count {
runtime
.score_set(&cost_key, &format!("{index}:100"), now)
.await
.expect("seed cost window");
}
runtime
.score_set(&cost_key, "expired:9999999", now - 20_000.0)
.await
.expect("expired cost");
for (member, score) in [("first:10", now), ("second:30", now), ("zero:0", now)] {
runtime
.score_set(&pool_latency_key("pool", key_id), member, score)
.await
.expect("seed latency");
}
}
let pool_config = config(json!({
"cost_limit_per_key_tokens": 50_000,
"cost_window_seconds": 600,
"latency_window_seconds": 600,
"scheduling_presets": [{"preset": "latency_first", "enabled": true}]
}));
let scheduled = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
scheduled.cost_window_usage_by_key.get("bounded"),
Some(&51_200)
);
assert_eq!(
scheduled.cost_window_usage_by_key.get("large"),
Some(&204_800)
);
assert_eq!(scheduled.latency_avg_ms_by_key.get("bounded"), Some(&20.0));
assert_eq!(scheduled.latency_avg_ms_by_key.get("large"), Some(&20.0));
runtime
.score_remove_by_score(&pool_cost_key("pool", "large"), f64::INFINITY)
.await
.expect("reset window");
runtime
.score_set(&pool_cost_key("pool", "large"), "after-reset:75", now)
.await
.expect("post-reset cost");
let reset = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(reset.cost_window_usage_by_key.get("large"), Some(&75));
runtime
.kv_set(&pool_cost_key("pool", "large"), "wrong-type", None)
.await
.expect("simulate invalid metric key");
let partial = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
partial.cost_window_usage_by_key.get("bounded"),
Some(&51_200),
"one failed aggregate must not discard another key's strict cost check"
);
}
#[tokio::test]
async fn scheduling_runtime_skips_admin_scan_and_unused_window_queries() {
let redis = match ManagedRedisServer::start().await {
Ok(server) => server,
Err(err) if err.to_string().contains("No such file or directory") => {
eprintln!("skipping redis-backed scheduling runtime test: {err}");
return;
}
Err(err) => panic!("Redis server should start: {err}"),
};
let runtime = RuntimeState::redis(
RedisClientConfig {
url: redis.redis_url().to_string(),
key_prefix: Some("scheduling-runtime-reads".to_string()),
},
Some(2_000),
)
.await
.expect("Redis runtime");
let pool_config = config(json!({}));
let keys = vec!["ready".to_string(), "cooling".to_string()];
seed_window_metrics(&runtime, "pool", "ready").await;
for session in ["current", "other"] {
runtime
.kv_set(
&pool_sticky_key("pool", session),
"ready".to_string(),
Some(Duration::from_secs(60)),
)
.await
.expect("seed sticky session");
}
runtime
.kv_set(
&pool_cooldown_key("pool", "cooling"),
"rate_limit".to_string(),
Some(Duration::from_secs(60)),
)
.await
.expect("seed cooldown");
let before = admin_command_count(&runtime).await;
let scheduled = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
Some("current"),
)
.await;
let after = admin_command_count(&runtime).await;
assert_eq!(
after - before,
1,
"only the diagnostics INFO may use the admin lane"
);
assert_eq!(scheduled.sticky_bound_key_id.as_deref(), Some("ready"));
assert_eq!(
scheduled
.cooldown_reason_by_key
.get("cooling")
.map(String::as_str),
Some("rate_limit")
);
assert!(scheduled.cooldown_ttl_by_key.is_empty());
assert!(scheduled.cost_window_usage_by_key.is_empty());
assert!(scheduled.latency_avg_ms_by_key.is_empty());
assert_eq!(scheduled.total_sticky_sessions, 0);
let admin = read_admin_provider_pool_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
Some("current"),
)
.await;
assert_eq!(admin.total_sticky_sessions, 2);
assert_eq!(admin.sticky_sessions_by_key.get("ready"), Some(&2));
assert_eq!(admin.cost_window_usage_by_key.get("ready"), Some(&100));
assert_eq!(admin.latency_avg_ms_by_key.get("ready"), Some(&20.0));
assert!(admin
.cooldown_ttl_by_key
.get("cooling")
.is_some_and(|ttl| *ttl > 0));
}
#[tokio::test]
async fn scheduling_runtime_loads_only_metrics_used_by_enabled_strategies() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let keys = vec!["key".to_string()];
seed_window_metrics(&runtime, "pool", "key").await;
for (value, expected_cost, expected_latency) in [
(json!({}), false, false),
(json!({"cost_limit_per_key_tokens": 100}), true, false),
(json!({"cost_limit_per_key_tokens": 0}), true, false),
(
json!({"scheduling_presets": [{"preset": "cost_first", "enabled": true}]}),
true,
false,
),
(
json!({"scheduling_presets": [{"preset": "quota_balanced", "enabled": true}]}),
true,
false,
),
(
json!({"scheduling_presets": [{"preset": "latency_first", "enabled": true}]}),
false,
true,
),
(
json!({"scheduling_presets": [
{"preset": "cost_first", "enabled": false},
{"preset": "latency_first", "enabled": false}
]}),
false,
false,
),
] {
let pool_config = config(value.clone());
let snapshot = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
snapshot.cost_window_usage_by_key.get("key").copied(),
expected_cost.then_some(100),
"config: {value}"
);
assert_eq!(
snapshot.latency_avg_ms_by_key.get("key").copied(),
expected_latency.then_some(20.0),
"config: {value}"
);
}
}
#[tokio::test]
async fn scheduling_runtime_checks_cost_for_candidates_beyond_admin_display_limit() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let keys = (0..513)
.map(|index| format!("key-{index}"))
.collect::<Vec<_>>();
seed_window_metrics(&runtime, "pool", &keys[512]).await;
let pool_config = config(json!({ "cost_limit_per_key_tokens": 100 }));
let snapshot = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
snapshot.cost_window_usage_by_key.get(&keys[512]),
Some(&100)
);
}
#[tokio::test]
async fn scheduling_sticky_lookup_invalidates_a_cooled_down_binding() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let pool_config = config(json!({}));
let sticky_key = pool_sticky_key("pool", "session");
runtime
.kv_set(&sticky_key, "key".to_string(), None)
.await
.expect("sticky session");
runtime
.kv_set(
&pool_cooldown_key("pool", "key"),
"rate_limit".to_string(),
None,
)
.await
.expect("cooldown");
assert!(read_provider_pool_sticky_bound_key_id(
&runtime,
"pool",
&pool_config,
Some("session")
)
.await
.is_none());
assert!(!runtime
.kv_exists(&sticky_key)
.await
.expect("sticky existence"));
}
#[test]
fn runtime_window_metric_key_ids_are_bounded() {
@@ -1,5 +1,6 @@
use super::value::build_admin_provider_summary_value;
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
@@ -10,19 +11,23 @@ use std::time::{SystemTime, UNIX_EPOCH};
pub(crate) async fn build_admin_provider_summary_payload(
state: &AdminAppState<'_>,
provider_id: &str,
) -> Option<serde_json::Value> {
) -> Result<Option<serde_json::Value>, GatewayError> {
let state = state.as_ref();
if !state.has_provider_catalog_data_reader() {
return None;
return Err(GatewayError::Internal(
"Admin provider catalog data unavailable".to_string(),
));
}
let provider_ids = vec![provider_id.to_string()];
let provider = state
let Some(provider) = state
.read_provider_catalog_providers_by_ids(&provider_ids)
.await
.ok()?
.await?
.into_iter()
.next()?;
.next()
else {
return Ok(None);
};
let (
endpoints_result,
keys_result,
@@ -36,8 +41,8 @@ pub(crate) async fn build_admin_provider_summary_payload(
state.list_provider_model_stats(&provider_ids),
state.list_active_global_model_ids_by_provider_ids(&provider_ids),
);
let endpoints = endpoints_result.ok().unwrap_or_default();
let keys = keys_result.ok().unwrap_or_default();
let endpoints = endpoints_result?;
let keys = keys_result?;
let quota_snapshot = quota_snapshot_result.ok().flatten();
let model_stats = model_stats_result
.ok()
@@ -57,7 +62,7 @@ pub(crate) async fn build_admin_provider_summary_payload(
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
Some(build_admin_provider_summary_value(
Ok(Some(build_admin_provider_summary_value(
&provider,
&endpoints,
&keys,
@@ -65,7 +70,7 @@ pub(crate) async fn build_admin_provider_summary_payload(
model_stats.as_ref(),
active_global_model_ids,
now_unix_secs,
))
)))
}
pub(crate) async fn build_admin_providers_summary_payload(
@@ -94,11 +99,7 @@ pub(crate) async fn build_admin_providers_summary_payload(
normalized_api_format != "all" && !normalized_api_format.is_empty();
let requires_model_filter = normalized_model_id != "all" && !normalized_model_id.is_empty();
let mut providers = state
.list_provider_catalog_providers(false)
.await
.ok()
.unwrap_or_default();
let mut providers = state.list_provider_catalog_providers(false).await.ok()?;
let all_provider_ids = providers
.iter()
.map(|provider| provider.id.clone())
@@ -109,8 +110,7 @@ pub(crate) async fn build_admin_providers_summary_payload(
state
.list_provider_catalog_endpoints_by_provider_ids(&all_provider_ids)
.await
.ok()
.unwrap_or_default()
.ok()?
};
let active_global_model_refs = if !requires_model_filter || all_provider_ids.is_empty() {
Vec::new()
@@ -202,8 +202,8 @@ pub(crate) async fn build_admin_providers_summary_payload(
state.list_active_global_model_ids_by_provider_ids(&provider_ids),
);
(
endpoints_result.ok().unwrap_or_default(),
keys_result.ok().unwrap_or_default(),
endpoints_result.ok()?,
keys_result.ok()?,
model_stats_result.ok().unwrap_or_default(),
active_global_model_refs_result.ok().unwrap_or_default(),
)
@@ -95,8 +95,11 @@ pub(crate) fn build_admin_provider_summary_value(
let scores = endpoint_keys
.iter()
.filter(|key| endpoint.is_active && key.is_active)
.filter_map(|key| provider_key_health_score(key, &endpoint.api_format))
.filter(|score| score.is_finite())
.map(|key| {
provider_key_health_score(key, &endpoint.api_format)
.filter(|score| score.is_finite())
.unwrap_or(1.0)
})
.collect::<Vec<_>>();
let health_score =
(!scores.is_empty()).then(|| scores.iter().sum::<f64>() / scores.len() as f64);
@@ -234,6 +234,20 @@ impl<'a> AdminAppState<'a> {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn read_request_usage_body_payload(
&self,
body_ref: &str,
) -> Result<
Option<aether_data_contracts::repository::usage::StoredUsageBodyPayload>,
GatewayError,
> {
self.app
.data
.read_request_usage_body_payload(body_ref)
.await
.map_err(|error| GatewayError::Internal(error.to_string()))
}
pub(crate) async fn build_api_format_health_monitor_payload(
&self,
lookback_hours: u64,
@@ -141,7 +141,7 @@ impl<'a> AdminAppState<'a> {
pub(crate) async fn build_admin_provider_summary_payload(
&self,
provider_id: &str,
) -> Option<serde_json::Value> {
) -> Result<Option<serde_json::Value>, GatewayError> {
crate::handlers::admin::provider::summary::build_admin_provider_summary_payload(
self,
provider_id,
@@ -1937,6 +1937,7 @@ pub(crate) async fn start_admin_system_rollback_task(
}
fn request_process_restart() -> ! {
let _ = aether_runtime::shutdown_logging(std::time::Duration::from_secs(2));
std::process::exit(RESTART_EXIT_CODE);
}
@@ -225,6 +225,9 @@ impl RequestBodyBufferError {
RequestBodyNormalizationError::RequestBodyTooLarge { .. } => {
"request_body_too_large"
}
RequestBodyNormalizationError::BodyBufferOverloaded { .. } => {
"request_body_buffer_overloaded"
}
},
Self::TooLarge { .. } => "request_body_too_large",
Self::Overloaded { .. } => "request_body_buffer_overloaded",
@@ -292,15 +295,37 @@ pub(super) async fn buffer_and_normalize_request_body(
.await
.map_err(RequestBodyBufferError::from)?;
let elapsed_ms = buffered.elapsed().as_millis() as u64;
let retained_input_capacity = buffered
.requested_bytes()
.saturating_sub(buffered.bytes().len());
let normalized = buffered
.try_map(|body| {
crate::headers::normalize_request_body_headers_and_bytes_with_limit(
.try_map_with_budget(|body, memory| {
crate::headers::normalize_request_body_headers_and_bytes_with_budget(
headers,
body,
policy.effective_max_bytes(),
&mut |requested_bytes| {
let requested_bytes = requested_bytes.saturating_add(retained_input_capacity);
memory.try_reserve_bytes(requested_bytes).map_err(|_| {
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes: policy.budget_bytes(),
}
})
},
)
})
.map_err(RequestBodyBufferError::Normalization)?;
.map_err(|error| match error {
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes,
} => RequestBodyBufferError::Overloaded {
requested_bytes,
budget_bytes,
timeout_ms: 0,
},
error => RequestBodyBufferError::Normalization(error),
})?;
info!(
event_name = "frontdoor_request_body_buffer_completed",
log_type = "event",
@@ -211,9 +211,19 @@ fn apply_sensitive_route_cache_policy(
return;
}
let preserve_no_transform = headers
.get_all(http::header::CACHE_CONTROL)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.any(|directive| directive.trim().eq_ignore_ascii_case("no-transform"));
headers.insert(
http::header::CACHE_CONTROL,
HeaderValue::from_static("no-store"),
HeaderValue::from_static(if preserve_no_transform {
"no-store, no-transform"
} else {
"no-store"
}),
);
headers.insert(http::header::PRAGMA, HeaderValue::from_static("no-cache"));
}
@@ -354,6 +364,29 @@ mod tests {
);
}
#[test]
fn raw_body_no_transform_survives_sensitive_cache_policy() {
let mut headers = HeaderMap::new();
headers.append(
http::header::CACHE_CONTROL,
HeaderValue::from_static("public, max-age=3600"),
);
headers.append(
http::header::CACHE_CONTROL,
HeaderValue::from_static(" No-Transform "),
);
apply_sensitive_route_cache_policy(
&mut headers,
"/api/admin/usage/usage-1?body_format=raw",
None,
);
assert_eq!(
headers[http::header::CACHE_CONTROL],
"no-store, no-transform"
);
assert_eq!(headers[http::header::PRAGMA], "no-cache");
}
#[test]
fn authenticated_user_data_responses_are_never_cacheable() {
let mut headers = HeaderMap::new();
+112 -6
View File
@@ -1035,11 +1035,10 @@ pub(crate) async fn proxy_request(
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
request: Request,
) -> Result<Response<Body>, GatewayError> {
crate::request_diagnostics::scope_request_diagnostics(Box::pin(proxy_request_inner(
state,
remote_addr,
request,
)))
crate::request_lifecycle::run_request_with_usage(
state.usage_runtime.clone(),
Box::pin(proxy_request_inner(state, remote_addr, request)),
)
.await
}
@@ -3228,7 +3227,7 @@ mod tests {
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])
.write_all(&[b'a'; 128])
.expect("test gzip body should encode");
let encoded = encoder.finish().expect("test gzip body should finish");
assert!(
@@ -3269,6 +3268,113 @@ mod tests {
));
}
#[tokio::test]
async fn request_body_buffer_allows_parallel_compressed_uploads() {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(br#"{"model":"test"}"#).unwrap();
let encoded = Bytes::from(encoder.finish().unwrap());
let budget_bytes = 2 * crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES;
let budget = Arc::new(Semaphore::new(2));
let policy = RequestBodyBufferPolicy::for_tests_with_budget(
budget_bytes as u64,
Duration::from_secs(1),
Duration::from_millis(50),
budget_bytes,
Arc::clone(&budget),
);
let mut headers = HeaderMap::new();
headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.insert(header::CONTENT_LENGTH, HeaderValue::from(encoded.len()));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (finish_tx, finish_rx) = tokio::sync::oneshot::channel();
let first_policy = policy.clone();
let mut first_headers = headers.clone();
let first_encoded = encoded.clone();
let first = async move {
let stream = async_stream::stream! {
let middle = first_encoded.len() / 2;
yield Ok::<_, std::io::Error>(first_encoded.slice(..middle));
let _ = started_tx.send(());
let _ = finish_rx.await;
yield Ok(first_encoded.slice(middle..));
};
buffer_and_normalize_request_body(
&mut Some(Body::from_stream(stream)),
&mut first_headers,
"test owns body",
"trace-compressed-first",
&Method::POST,
"/v1/responses",
"test",
first_policy,
)
.await
};
let second = async move {
started_rx.await.unwrap();
let result = buffer_and_normalize_request_body(
&mut Some(Body::from(encoded)),
&mut headers,
"test owns body",
"trace-compressed-second",
&Method::POST,
"/v1/responses",
"test",
policy,
)
.await;
let _ = finish_tx.send(());
result
};
let (first, second) = tokio::time::timeout(Duration::from_secs(2), async {
tokio::join!(first, second)
})
.await
.expect("concurrent compressed requests should finish");
assert_eq!(first.unwrap().as_ref(), br#"{"model":"test"}"#);
assert_eq!(second.unwrap().as_ref(), br#"{"model":"test"}"#);
assert_eq!(budget.available_permits(), 2);
}
#[tokio::test]
async fn request_body_buffer_rejects_decompression_growth_when_budget_is_busy() {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&vec![b'a'; 100_000]).unwrap();
let encoded = encoder.finish().unwrap();
let budget_bytes = 2 * crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES;
let budget = Arc::new(Semaphore::new(2));
let held = Arc::clone(&budget).acquire_owned().await.unwrap();
let mut headers = HeaderMap::new();
headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.insert(header::CONTENT_LENGTH, HeaderValue::from(encoded.len()));
let result = buffer_and_normalize_request_body(
&mut Some(Body::from(encoded)),
&mut headers,
"test owns body",
"trace-decompression-overload",
&Method::POST,
"/v1/responses",
"test",
RequestBodyBufferPolicy::for_tests_with_budget(
budget_bytes as u64,
Duration::from_secs(1),
Duration::from_secs(1),
budget_bytes,
Arc::clone(&budget),
),
)
.await
.unwrap_err();
assert!(matches!(
result,
RequestBodyBufferError::Overloaded { timeout_ms: 0, .. }
));
assert_eq!(result.http_status(), http::StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(budget.available_permits(), 1);
drop(held);
assert_eq!(budget.available_permits(), 2);
}
#[tokio::test]
async fn request_body_buffer_times_out_instead_of_waiting_forever() {
let stream = async_stream::stream! {
@@ -508,7 +508,9 @@ async fn persist_live_audit_event(
let usage_runtime = std::sync::Arc::clone(&state.usage_runtime);
let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state());
let write_request_id = request_id.clone();
let usage_producer = usage_runtime.track_producer();
let task = tokio::spawn(async move {
let _usage_producer = usage_producer;
if tokio::time::timeout(
LIVE_AUDIT_WRITE_HARD_TIMEOUT,
usage_runtime.record_terminal_event_direct(usage_data.as_ref(), event),
@@ -572,7 +574,9 @@ fn spawn_live_audit_event_detached(state: &AppState, event: UsageEvent, audit_sc
};
let usage_runtime = std::sync::Arc::clone(&state.usage_runtime);
let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state());
let usage_producer = usage_runtime.track_producer();
runtime.spawn(async move {
let _usage_producer = usage_producer;
if tokio::time::timeout(
LIVE_AUDIT_WRITE_HARD_TIMEOUT,
usage_runtime.record_terminal_event_direct(usage_data.as_ref(), event),
@@ -68,7 +68,12 @@ pub(super) async fn relay_bound_connection(
state: &AppState,
context: &WebSocketRequestContext,
) {
let mut client_connected = true;
loop {
if !client_connected && !bound.turn_state.response_in_flight() {
close_bound_upstream(bound).await;
break;
}
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
tokio::select! {
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
@@ -100,8 +105,12 @@ pub(super) async fn relay_bound_connection(
).await;
break;
}
client_message = client_socket.next() => {
client_message = client_socket.next(), if client_connected => {
let Some(client_message) = client_message else {
if retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
finalize_active_turn(
bound,
state,
@@ -111,6 +120,10 @@ pub(super) async fn relay_bound_connection(
break;
};
let Ok(client_message) = client_message else {
if retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
warn!(
event_name = "responses_websocket_client_receive_failed",
log_type = "ops",
@@ -127,6 +140,12 @@ pub(super) async fn relay_bound_connection(
close_bound_upstream(bound).await;
break;
};
if matches!(client_message, AxumWsMessage::Close(_))
&& retain_disconnected_turn(bound)
{
client_connected = false;
continue;
}
match Box::pin(forward_client_message(
client_message,
bound,
@@ -559,6 +578,7 @@ pub(super) async fn relay_bound_connection(
let mut relay_send_error = None;
let mut relay_serialization_failed = false;
match relay_directive {
_ if !client_connected => {}
Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => {
let client_frame = match parsed_upstream_frame.as_ref().map(|frame| {
bound
@@ -673,6 +693,10 @@ pub(super) async fn relay_bound_connection(
break;
}
if let Some(error) = relay_send_error {
if terminal_outcome.is_none() && retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
warn!(
event_name = "responses_websocket_client_send_failed",
log_type = "ops",
@@ -737,6 +761,20 @@ pub(super) async fn relay_bound_connection(
}
}
fn retain_disconnected_turn(bound: &mut BoundResponsesConnection) -> bool {
if bound
.turn_state
.attempt()
.is_none_or(|attempt| attempt.cancel_on_client_disconnect())
{
return false;
}
bound
.turn_state
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
true
}
struct PendingContinuationRegistration {
user_id: String,
api_key_id: String,
@@ -845,6 +845,13 @@ fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> Gatew
}
impl ResponsesProviderAttempt {
pub(super) fn cancel_on_client_disconnect(&self) -> bool {
crate::orchestration::routing_execution_policy_from_report_context(
self.lifecycle.report_context(),
)
.is_some_and(|policy| policy.cancel_on_client_disconnect)
}
/// Releases all per-turn capacity before terminal persistence starts.
/// Provider-pool runtime tokens normally use an awaited removal. The
/// bounded wait prevents a broken runtime backend from stalling the relay;
@@ -24,7 +24,7 @@ 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, normalize_execution_proxy_url,
ExecutionTransportControls,
validate_execution_upstream_url, ExecutionSafeDnsResolver, ExecutionTransportControls,
};
use crate::frontdoor_loop_guard::gateway_frontdoor_self_loop_guard_error;
use crate::handlers::proxy::websocket::session::{
@@ -66,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, &upstream_url, errors).await?;
let client = build_websocket_client(decision, errors)?;
let response = client
.websocket(upstream_url.as_str())
.headers(headers)
@@ -149,25 +149,14 @@ pub(crate) fn websocket_upstream_url(
invalid_code: &'static str,
) -> Result<Url, &'static str> {
let mut url = Url::parse(raw).map_err(|_| invalid_code)?;
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" => return Ok(url),
"ws" if aether_http::url_has_literal_loopback_host(&url) => return Ok(url),
"ws" => return Err(invalid_code),
let (http_scheme, websocket_scheme) = match url.scheme() {
"https" | "wss" => ("https", "wss"),
"http" | "ws" => ("http", "ws"),
_ => return Err(invalid_code),
};
url.set_scheme(http_scheme).map_err(|_| invalid_code)?;
let mut url = validate_execution_upstream_url(url.as_str()).map_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)
}
@@ -223,9 +212,8 @@ pub(crate) fn websocket_handshake_headers(
Ok(headers)
}
async fn build_websocket_client(
fn build_websocket_client(
decision: &AiExecutionDecision,
upstream_url: &Url,
errors: UpstreamWebSocketErrorCodes,
) -> Result<wreq::Client, &'static str> {
let timeouts = websocket_timeouts(decision);
@@ -249,41 +237,7 @@ async fn build_websocket_client(
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::<std::net::IpAddr>() {
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::<std::net::IpAddr>()
.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 = builder.dns_resolver(ExecutionSafeDnsResolver);
}
builder.build().map_err(|_| errors.client_build_failed)
}
@@ -678,15 +632,17 @@ 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, 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,
bounded_send, build_websocket_client, 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::ai_serving::AiExecutionDecision;
use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url;
use aether_contracts::ProxySnapshot;
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
use axum::http::HeaderMap;
use std::collections::BTreeMap;
use std::time::Duration;
@@ -844,15 +800,17 @@ mod tests {
#[test]
fn maps_http_url_to_websocket_url_without_losing_path_or_query() {
let url = websocket_upstream_url(
"https://example.test/backend-api/codex/responses?x=1",
"invalid",
)
.expect("URL should be converted");
assert_eq!(
url.as_str(),
"wss://example.test/backend-api/codex/responses?x=1"
);
for (http_scheme, websocket_scheme) in [("https", "wss"), ("http", "ws")] {
let url = websocket_upstream_url(
&format!("{http_scheme}://example.test:8080/backend-api/codex/responses?x=1"),
"invalid",
)
.expect("URL should be converted");
assert_eq!(
url.as_str(),
format!("{websocket_scheme}://example.test:8080/backend-api/codex/responses?x=1")
);
}
}
#[test]
@@ -861,10 +819,16 @@ mod tests {
}
#[test]
fn remote_websocket_requires_wss_but_loopback_ws_is_allowed() {
fn websocket_upstream_url_accepts_ws_and_wss_with_safe_targets() {
for allowed in [
"wss://example.test/v1/responses",
"https://example.test/v1/responses",
"ws://example.test:8080/v1/responses",
"http://example.test:8080/v1/responses",
"http://8.8.8.8:8080/v1/responses",
"wss://8.8.8.8/v1/responses",
"ws://[2606:4700:4700::1111]:8080/v1/responses",
"wss://[2606:4700:4700::1111]/v1/responses",
"ws://localhost:8080/v1/responses",
"http://127.42.0.1:8080/v1/responses",
"ws://[::1]:8080/v1/responses",
@@ -875,11 +839,22 @@ mod tests {
);
}
for rejected in [
"ws://example.test/v1/responses",
"http://10.0.0.1/v1/responses",
"wss://10.0.0.1/v1/responses",
"wss://127.0.0.1/v1/responses",
"wss://[::1]/v1/responses",
"wss://[fd00::1]/v1/responses",
"wss://[::ffff:127.0.0.1]/v1/responses",
"wss://169.254.169.254/v1/responses",
"wss://198.18.78.41/v1/responses",
"wss://198.19.1.2/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",
"ws://example.test/v1/responses#secret",
"http://[email protected]/v1/responses",
"ws://[email protected]/v1/responses",
"ftp://example.test/v1/responses",
] {
assert!(
websocket_upstream_url(rejected, "invalid").is_err(),
@@ -888,6 +863,60 @@ mod tests {
}
}
#[tokio::test]
async fn websocket_client_build_defers_provider_dns_for_all_transport_profiles() {
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 profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for proxy in [
None,
Some(ProxySnapshot {
enabled: Some(false),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("socks5h://proxy.invalid:1080".to_string()),
..Default::default()
}),
] {
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "wss://upstream.invalid/v1/responses"
}))
.expect("minimal provider decision should deserialize");
decision.transport_profile = profile.clone();
decision.proxy = proxy;
build_websocket_client(&decision, errors)
.expect("building a client must not resolve the provider or proxy hostname");
}
}
}
#[test]
fn active_websocket_proxy_without_a_target_fails_closed() {
let errors = UpstreamWebSocketErrorCodes {
@@ -922,6 +951,94 @@ mod tests {
);
}
#[tokio::test]
async fn websocket_handshake_keeps_provider_dns_remote_for_http_and_socks_proxies() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
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 profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for scheme in ["http", "socks5", "socks5h"] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let proxy_addr = listener.local_addr().unwrap();
let (release, released) = tokio::sync::oneshot::channel::<()>();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
if scheme != "http" {
let mut greeting = [0; 2];
stream.read_exact(&mut greeting).await.unwrap();
assert_eq!(greeting[0], 5);
let mut methods = vec![0; greeting[1] as usize];
stream.read_exact(&mut methods).await.unwrap();
assert!(methods.contains(&0));
stream.write_all(&[5, 0]).await.unwrap();
let mut request = [0; 4];
stream.read_exact(&mut request).await.unwrap();
assert_eq!(
request,
[5, 1, 0, 3],
"proxy must receive a domain, not an IP"
);
let host_len = stream.read_u8().await.unwrap();
let mut host = vec![0; host_len as usize];
stream.read_exact(&mut host).await.unwrap();
assert_eq!(host, b"provider-dns.invalid");
assert_eq!(stream.read_u16().await.unwrap(), 80);
stream
.write_all(&[5, 0, 0, 1, 127, 0, 0, 1, 0, 80])
.await
.unwrap();
}
let socket = tokio_tungstenite::accept_async(stream).await.unwrap();
let _ = released.await;
drop(socket);
});
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "ws://provider-dns.invalid/v1/responses",
"proxy": {"enabled": true, "url": format!("{scheme}://{proxy_addr}")}
}))
.unwrap();
decision.transport_profile = profile.clone();
let connection = tokio::time::timeout(
Duration::from_secs(5),
super::connect_upstream_websocket(
&decision,
crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS,
errors,
),
)
.await
.expect("proxied handshake must not wait for local provider DNS")
.unwrap_or_else(|error| panic!("{scheme} handshake failed: {error}"));
release.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(5), server)
.await
.unwrap()
.unwrap();
drop(connection);
}
}
}
#[test]
fn rejects_responses_websocket_frontdoor_self_loop_before_connecting() {
let base_url = configured_gateway_frontdoor_base_url();
@@ -129,9 +129,6 @@ pub(crate) fn normalize_admin_base_url(base_url: &str) -> Result<String, 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());
}
@@ -154,16 +151,42 @@ mod normalize_admin_base_url_tests {
"https://user:[email protected]/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",
"http://user:password@api.example.test/v1",
"http://api.example.test/v1?key=secret",
"http://api.example.test/v1#secret",
"ftp://api.example.test/v1",
"file:///v1",
"api.example.test/v1",
"",
"https://",
"http://",
"https://api.example.test:invalid/v1",
] {
assert!(normalize_admin_base_url(value).is_err(), "accepted {value}");
}
}
#[test]
fn endpoint_base_url_accepts_remote_http_hosts() {
for (raw_url, expected) in [
(
" HTTP://API.EXAMPLE.TEST:8080/v1/ ",
"http://api.example.test:8080/v1",
),
("http://8.8.8.8:8080/v1/", "http://8.8.8.8:8080/v1"),
("http://10.0.0.1:8080/v1/", "http://10.0.0.1:8080/v1"),
(
"http://[2606:4700:4700::1111]:8080/v1/",
"http://[2606:4700:4700::1111]:8080/v1",
),
] {
assert_eq!(
normalize_admin_base_url(raw_url).expect("HTTP base URL should be accepted"),
expected,
);
}
}
#[test]
fn endpoint_base_url_is_parsed_and_normalized() {
assert_eq!(
@@ -30,6 +30,8 @@ use std::time::{SystemTime, UNIX_EPOCH};
mod support_announcements;
#[path = "support/auth.rs"]
mod support_auth;
#[path = "support/auth_cookie_policy.rs"]
mod support_auth_cookie_policy;
#[path = "support/billing.rs"]
mod support_billing;
#[path = "support/ccswitch.rs"]
@@ -133,6 +135,31 @@ pub(crate) async fn maybe_build_local_public_support_response(
remote_addr: &std::net::SocketAddr,
client_ip: std::net::IpAddr,
request_body: Option<&Bytes>,
) -> Option<Response<Body>> {
let response = build_local_public_support_response(
state,
request_context,
headers,
remote_addr,
client_ip,
request_body,
)
.await?;
Some(support_auth_cookie_policy::finalize_refresh_cookie(
response,
headers,
request_context.host_header.as_deref(),
remote_addr,
))
}
async fn build_local_public_support_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
remote_addr: &std::net::SocketAddr,
client_ip: std::net::IpAddr,
request_body: Option<&Bytes>,
) -> Option<Response<Body>> {
let decision = request_context.control_decision.as_ref()?;
if decision.route_class.as_deref() != Some("public_support") {
@@ -0,0 +1,424 @@
use super::support_auth::{auth_refresh_cookie_name, auth_refresh_cookie_secure};
use axum::body::Body;
use axum::http::{header, HeaderMap, HeaderValue, Response};
use std::net::SocketAddr;
use url::Url;
pub(super) fn finalize_refresh_cookie(
mut response: Response<Body>,
headers: &HeaderMap,
host_header: Option<&str>,
remote_addr: &SocketAddr,
) -> Response<Body> {
if !response.headers().contains_key(header::SET_COOKIE) {
return response;
}
let cookie_name = auth_refresh_cookie_name();
let explicit_secure = std::env::var("AUTH_REFRESH_COOKIE_SECURE").ok();
let public_base_url = std::env::var("AETHER_PUBLIC_BASE_URL")
.ok()
.or_else(|| std::env::var("PUBLIC_BASE_URL").ok());
let secure = refresh_cookie_secure_for_request(
headers,
host_header,
crate::headers::trusted_proxy_ip(remote_addr.ip()),
explicit_secure.as_deref(),
public_base_url.as_deref(),
auth_refresh_cookie_secure(),
);
let cookies = response
.headers()
.get_all(header::SET_COOKIE)
.iter()
.map(|cookie| rewrite_refresh_cookie(cookie, &cookie_name, secure))
.collect::<Vec<_>>();
response.headers_mut().remove(header::SET_COOKIE);
for cookie in cookies {
response.headers_mut().append(header::SET_COOKIE, cookie);
}
response
}
fn refresh_cookie_secure_for_request(
headers: &HeaderMap,
host_header: Option<&str>,
trusted_proxy: bool,
explicit_secure: Option<&str>,
public_base_url: Option<&str>,
fallback_secure: bool,
) -> bool {
if let Some(value) = explicit_secure {
return !value.trim().eq_ignore_ascii_case("false");
}
let origin = single_header(headers, header::ORIGIN.as_str()).and_then(parse_origin);
let public_url = public_base_url.and_then(parse_http_url);
let forwarded_proto = trusted_proxy.then(|| forwarded_proto(headers)).flatten();
if origin.as_ref().is_some_and(|url| url.scheme() == "https")
|| public_url
.as_ref()
.is_some_and(|url| url.scheme() == "https")
|| forwarded_proto == Some("https")
{
return true;
}
if trusted_proxy && headers.contains_key("x-forwarded-proto") {
return forwarded_proto != Some("http");
}
if public_url
.as_ref()
.is_some_and(|url| url.scheme() == "http")
{
return false;
}
if let (Some(origin), Some(host)) = (origin, host_header) {
let request_origin = parse_origin(&format!("{}://{host}", origin.scheme()));
if request_origin.is_some_and(|url| url.origin() == origin.origin()) {
return false;
}
}
fallback_secure
}
fn single_header<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
let mut values = headers.get_all(name).iter();
let value = values.next()?.to_str().ok()?.trim();
(values.next().is_none() && !value.is_empty()).then_some(value)
}
fn parse_origin(value: &str) -> Option<Url> {
let url = parse_http_url(value)?;
(url.path() == "/").then_some(url)
}
fn parse_http_url(value: &str) -> Option<Url> {
let url = Url::parse(value.trim()).ok()?;
(matches!(url.scheme(), "http" | "https")
&& url.host_str().is_some()
&& url.username().is_empty()
&& url.password().is_none()
&& url.query().is_none()
&& url.fragment().is_none())
.then_some(url)
}
fn forwarded_proto(headers: &HeaderMap) -> Option<&'static str> {
let value = headers
.get_all("x-forwarded-proto")
.iter()
.next_back()?
.to_str()
.ok()?
.rsplit(',')
.next()?
.trim();
if value.eq_ignore_ascii_case("https") {
Some("https")
} else if value.eq_ignore_ascii_case("http") {
Some("http")
} else {
None
}
}
fn rewrite_refresh_cookie(cookie: &HeaderValue, cookie_name: &str, secure: bool) -> HeaderValue {
let Ok(value) = cookie.to_str() else {
return cookie.clone();
};
let mut attributes = value.split(';').map(str::trim);
let Some(pair) = attributes.next() else {
return cookie.clone();
};
if pair.split_once('=').map(|(name, _)| name) != Some(cookie_name) {
return cookie.clone();
}
let secure =
secure || cookie_name.starts_with("__Secure-") || cookie_name.starts_with("__Host-");
let mut parts = vec![pair.to_string()];
for attribute in attributes {
if attribute.eq_ignore_ascii_case("Secure") {
continue;
}
if !secure
&& attribute.split_once('=').is_some_and(|(name, value)| {
name.trim().eq_ignore_ascii_case("SameSite")
&& value.trim().eq_ignore_ascii_case("None")
})
{
parts.push("SameSite=Lax".to_string());
} else {
parts.push(attribute.to_string());
}
}
if secure {
parts.push("Secure".to_string());
}
let Ok(mut rewritten) = HeaderValue::from_str(&parts.join("; ")) else {
return cookie.clone();
};
rewritten.set_sensitive(cookie.is_sensitive());
rewritten
}
#[cfg(test)]
mod tests {
use super::{refresh_cookie_secure_for_request, rewrite_refresh_cookie};
use axum::http::{header, HeaderMap, HeaderValue};
fn headers(origin: Option<&str>, forwarded_proto: Option<&str>) -> HeaderMap {
let mut headers = HeaderMap::new();
if let Some(origin) = origin {
headers.insert(header::ORIGIN, HeaderValue::from_str(origin).unwrap());
}
if let Some(proto) = forwarded_proto {
headers.insert("x-forwarded-proto", HeaderValue::from_str(proto).unwrap());
}
headers
}
#[test]
fn refresh_cookie_auto_detects_same_origin_http_and_https() {
for (origin, host, secure) in [
("http://aether.test:8084", "aether.test:8084", false),
("http://aether.test", "aether.test:80", false),
("http://[2001:db8::1]:8084", "[2001:db8::1]:8084", false),
("https://aether.test", "aether.test", true),
] {
assert_eq!(
refresh_cookie_secure_for_request(
&headers(Some(origin), None),
Some(host),
false,
None,
None,
true,
),
secure,
"{origin}",
);
}
assert!(refresh_cookie_secure_for_request(
&headers(Some("https://aether.test"), None),
Some("aether.test"),
false,
None,
None,
false,
));
}
#[test]
fn refresh_cookie_does_not_infer_http_from_other_or_invalid_origins() {
for origin in [
"http://other.test",
"http://aether.test:8085",
"null",
"http://[email protected]:8084",
"http://aether.test:8084/path",
"http://aether.test:8084?query",
"http://aether.test:8084#fragment",
"http://aether.test:8084, https://aether.test:8084",
"file:///tmp/test",
] {
assert!(
refresh_cookie_secure_for_request(
&headers(Some(origin), None),
Some("aether.test:8084"),
false,
None,
None,
true,
),
"{origin}"
);
}
let mut duplicate = headers(Some("http://aether.test:8084"), None);
duplicate.append(
header::ORIGIN,
HeaderValue::from_static("https://aether.test:8084"),
);
assert!(refresh_cookie_secure_for_request(
&duplicate,
Some("aether.test:8084"),
false,
None,
None,
true,
));
}
#[test]
fn refresh_cookie_only_trusts_forwarded_protocol_from_trusted_peers() {
for (proto, trusted, secure) in [
("http", true, false),
("https", true, true),
("http", false, true),
("https", false, true),
("https, http", true, false),
("http, https", true, true),
("ftp", true, true),
("http,", true, true),
] {
assert_eq!(
refresh_cookie_secure_for_request(
&headers(None, Some(proto)),
Some("aether.test"),
trusted,
None,
None,
true,
),
secure,
"{proto}, trusted={trusted}"
);
}
let mut chained = headers(None, Some("http, http"));
chained.append("x-forwarded-proto", HeaderValue::from_static("https"));
assert!(refresh_cookie_secure_for_request(
&chained,
Some("aether.test"),
true,
None,
None,
true,
));
}
#[test]
fn refresh_cookie_https_evidence_prevents_automatic_downgrade() {
for (origin, proto, public_url) in [
("https://aether.test", "http", None),
("http://aether.test", "https", None),
("http://aether.test", "http", Some("https://aether.test")),
] {
assert!(refresh_cookie_secure_for_request(
&headers(Some(origin), Some(proto)),
Some("aether.test"),
true,
None,
public_url,
true,
));
}
}
#[test]
fn refresh_cookie_preserves_explicit_overrides_and_unknown_defaults() {
for (explicit, secure) in [
("true", true),
("FALSE", false),
("invalid", true),
("", true),
] {
assert_eq!(
refresh_cookie_secure_for_request(
&headers(Some("http://aether.test"), None),
Some("aether.test"),
false,
Some(explicit),
None,
true,
),
secure
);
}
assert!(!refresh_cookie_secure_for_request(
&headers(Some("https://aether.test"), None),
Some("aether.test"),
false,
Some("false"),
None,
true,
));
for fallback in [false, true] {
assert_eq!(
refresh_cookie_secure_for_request(
&HeaderMap::new(),
Some("aether.test"),
false,
None,
None,
fallback,
),
fallback
);
}
}
#[test]
fn refresh_cookie_accepts_an_explicit_public_http_origin() {
assert!(!refresh_cookie_secure_for_request(
&HeaderMap::new(),
Some("internal:8084"),
false,
None,
Some("http://aether.test"),
true,
));
assert!(refresh_cookie_secure_for_request(
&headers(Some("https://aether.test"), None),
Some("internal:8084"),
false,
None,
Some("http://aether.test"),
true,
));
}
#[test]
fn refresh_cookie_rewrite_preserves_secret_path_expiry_and_httponly() {
let mut cookie = HeaderValue::from_static(
"aether_refresh_token=secret; Path=/api/auth; HttpOnly; SameSite=None; Max-Age=604800; Secure",
);
cookie.set_sensitive(true);
let rewritten = rewrite_refresh_cookie(&cookie, "aether_refresh_token", false);
assert_eq!(
rewritten.to_str().unwrap(),
"aether_refresh_token=secret; Path=/api/auth; HttpOnly; SameSite=Lax; Max-Age=604800"
);
assert!(rewritten.is_sensitive());
assert_eq!(
rewrite_refresh_cookie(&cookie, "aether_refresh_token", true),
cookie
);
}
#[test]
fn refresh_cookie_rewrite_also_clears_http_cookies() {
let cookie = HeaderValue::from_static(
"aether_refresh_token=; Path=/api/auth; HttpOnly; SameSite=None; Max-Age=0; Secure",
);
assert_eq!(
rewrite_refresh_cookie(&cookie, "aether_refresh_token", false)
.to_str()
.unwrap(),
"aether_refresh_token=; Path=/api/auth; HttpOnly; SameSite=Lax; Max-Age=0"
);
}
#[test]
fn refresh_cookie_rewrite_preserves_other_cookies_and_strict_policy() {
let unrelated = HeaderValue::from_static("oauth_binding=secret; Path=/; Secure; HttpOnly");
assert_eq!(
rewrite_refresh_cookie(&unrelated, "aether_refresh_token", false),
unrelated
);
let strict = HeaderValue::from_static(
"custom_refresh=secret; Path=/api/auth; HttpOnly; SameSite=Strict",
);
assert_eq!(
rewrite_refresh_cookie(&strict, "custom_refresh", false),
strict
);
assert!(rewrite_refresh_cookie(&strict, "custom_refresh", true)
.to_str()
.unwrap()
.ends_with("; Secure"));
let prefixed =
HeaderValue::from_static("__Secure-refresh=secret; HttpOnly; SameSite=None; Secure");
assert_eq!(
rewrite_refresh_cookie(&prefixed, "__Secure-refresh", false),
prefixed
);
}
}
@@ -248,7 +248,7 @@ pub(super) fn auth_verification_send_cooldown_seconds() -> i64 {
.unwrap_or(60)
}
pub(super) fn auth_refresh_cookie_name() -> String {
pub(crate) fn auth_refresh_cookie_name() -> String {
std::env::var("AUTH_REFRESH_COOKIE_NAME")
.ok()
.map(|value| value.trim().to_string())
@@ -1,5 +1,5 @@
use std::collections::BTreeMap;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use axum::{
@@ -10,6 +10,10 @@ use axum::{
};
use serde_json::json;
use crate::execution_runtime::transport::{
validate_execution_upstream_url, ExecutionSafeDnsResolver,
};
use super::test_connection_shared::select_test_connection_provider;
use super::{
provider_catalog_key_supports_format, query_param_value, AppState, GatewayPublicRequestContext,
@@ -18,101 +22,16 @@ 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, reqwest::Error> {
reqwest::Client::builder()
.no_proxy()
.dns_resolver(Arc::new(ExecutionSafeDnsResolver))
.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<SocketAddr>,
}
/// 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<ResolvedTestConnectionTarget, &'static str> {
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::<IpAddr>().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<reqwest::Client, reqwest::Error> {
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::<IpAddr>().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,
request_context: &GatewayPublicRequestContext,
@@ -387,18 +306,14 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
// 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,
let upstream_url = match validate_execution_upstream_url(&upstream_url) {
Ok(url) => url,
Err(reason) => {
tracing::warn!(
event_name = "provider_test_connection_target_rejected",
provider_id = %provider.id,
endpoint_id = %endpoint.id,
reason,
reason = %reason,
"provider connection test target was rejected"
);
return Some(
@@ -410,7 +325,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let test_client = match build_pinned_test_connection_client(&target) {
let test_client = match build_test_connection_client() {
Ok(client) => client,
Err(_) => {
return Some(
@@ -422,7 +337,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let mut upstream_request = test_client.post(target.url);
let mut upstream_request = test_client.post(upstream_url);
for (name, value) in &provider_request_headers {
upstream_request = upstream_request.header(name, value);
}
@@ -498,7 +413,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
#[cfg(test)]
mod tests {
use super::{build_test_connection_client, resolve_test_connection_target};
use super::{build_test_connection_client, validate_execution_upstream_url};
use axum::{
body::Body,
http::{header, Request, StatusCode},
@@ -570,61 +485,68 @@ mod tests {
redirected_server.abort();
}
#[tokio::test]
async fn test_connection_target_rejects_private_addresses_in_production_mode() {
#[test]
fn test_connection_target_rejects_private_literals_like_provider_requests() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat/completions",
"http://10.0.0.1/v1/chat/completions",
"http://169.254.169.254/v1/chat/completions",
"https://10.0.0.1/v1/chat/completions",
"https://127.0.0.1/v1/chat/completions",
"https://[::1]/v1/chat/completions",
"https://localhost/v1/chat/completions",
"http://8.8.8.8/v1/chat/completions",
"https://198.18.78.41/v1/chat/completions",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).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"
);
#[test]
fn test_connection_target_accepts_public_http_and_https_addresses() {
for (raw_url, expected_port) in [
("http://8.8.8.8/v1/chat", 80),
("http://8.8.8.8:8080/v1/chat", 8080),
("https://8.8.8.8/v1/chat", 443),
("https://[2606:4700:4700::1111]/v1/chat", 443),
] {
let url = validate_execution_upstream_url(raw_url)
.expect("public HTTP(S) provider target should be valid");
assert_eq!(url.as_str(), raw_url);
assert_eq!(url.port_or_known_default(), Some(expected_port));
}
}
#[tokio::test]
async fn test_connection_target_rejects_url_credentials_and_fragments() {
async fn test_connection_target_defers_dns_and_accepts_provider_loopback_urls() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat",
"http://[::1]:8080/v1/chat",
"http://localhost:8080/v1/chat",
"https://provider-dns.invalid/v1/chat",
] {
let url = validate_execution_upstream_url(raw_url)
.expect("target validation must not depend on the current DNS answer");
let request = build_test_connection_client()
.expect("client should build without DNS")
.post(url)
.build()
.expect("provider request should build without DNS");
assert_eq!(request.url().as_str(), raw_url);
}
}
#[test]
fn test_connection_target_rejects_url_credentials_and_fragments() {
for raw_url in [
"https://user:[email protected]/v1/chat",
"https://example.com/v1/chat#fragment",
"http://user:[email protected]/v1/chat",
"http://example.com/v1/chat#fragment",
"ftp://example.com/v1/chat",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).is_err(),
"unsafe provider target should be rejected: {raw_url}"
);
}
@@ -168,52 +168,6 @@ fn wallet_public_refund_payload(mut payload: serde_json::Value) -> serde_json::V
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,
@@ -659,3 +613,49 @@ pub(super) async fn handle_wallet_create_refund(
}
}
}
#[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");
}
}
@@ -3074,6 +3074,45 @@ mod tests {
.expect("key transport should build")
}
#[test]
fn admin_provider_key_health_response_preserves_v0_7_13_defaults() {
let state = AppState::new().expect("gateway should build");
for (health, expected_score) in [
(None, json!(1.0)),
(Some(json!({})), json!(1.0)),
(
Some(json!({"openai:chat": {"consecutive_failures": 0}})),
json!(1.0),
),
(
Some(json!({"openai:chat": {"health_score": 0.0}})),
json!(0.0),
),
(
Some(json!({"openai:chat": {"health_score": 1.0}})),
json!(1.0),
),
(
Some(json!({
"openai:chat": {"health_score": 0.25},
"openai:responses": {"health_score": 0.75},
})),
json!(0.25),
),
] {
let mut key = sample_catalog_key();
key.health_by_format = health;
let payload = build_admin_provider_key_response(
&state,
&key,
"openai",
&["openai:chat".to_string()],
1_000,
);
assert_eq!(payload["health_score"], expected_score);
}
}
#[test]
fn responses_key_scope_covers_search_in_one_direction() {
let mut responses_key = sample_catalog_key();
@@ -292,7 +292,7 @@ mod tests {
assert!(controls.is_err());
let template = "{{value}}".repeat(100_000);
let variables = BTreeMap::from([(String::from("value"), String::from("x".repeat(64)))]);
let variables = BTreeMap::from([(String::from("value"), "x".repeat(64))]);
let error = render_admin_email_template_html(&template, &variables)
.expect_err("rendered output must remain bounded");
assert!(format!("{error:?}").contains("exceeds"));
@@ -3,7 +3,8 @@ pub(crate) use super::super::admin::provider::pool::config::{
};
pub(crate) use super::super::admin::provider::pool::runtime::{
admin_provider_pool_key_terminal_error_reason, read_admin_provider_pool_key_cooldown_reason,
read_admin_provider_pool_runtime_state, record_admin_provider_pool_error,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
read_provider_pool_sticky_bound_key_id, record_admin_provider_pool_error,
record_admin_provider_pool_stream_timeout, record_admin_provider_pool_success,
release_admin_provider_pool_key_lease,
};
@@ -199,10 +199,7 @@ pub(crate) fn normalize_ldap_transport_server_url(raw: &str, use_starttls: bool)
// 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,
);
aether_admin::system::normalize_ldap_transport_server_url_for_tests(raw, use_starttls)
}
#[cfg(not(test))]
{
+401 -30
View File
@@ -402,9 +402,21 @@ pub(crate) enum RequestBodyNormalizationError {
InvalidBodyFraming,
AmbiguousBodyFraming,
UnsupportedContentEncoding(String),
DecodeFailed { encoding: String, reason: String },
DecompressedBodyTooLarge { encoding: String, limit_bytes: u64 },
RequestBodyTooLarge { limit_bytes: u64 },
DecodeFailed {
encoding: String,
reason: String,
},
DecompressedBodyTooLarge {
encoding: String,
limit_bytes: u64,
},
RequestBodyTooLarge {
limit_bytes: u64,
},
BodyBufferOverloaded {
requested_bytes: usize,
budget_bytes: usize,
},
}
impl RequestBodyNormalizationError {
@@ -428,6 +440,9 @@ impl RequestBodyNormalizationError {
Self::RequestBodyTooLarge { limit_bytes } => {
format!("Request body exceeds {limit_bytes} bytes")
}
Self::BodyBufferOverloaded { .. } => {
"Request body buffering capacity is temporarily exhausted".to_string()
}
}
}
@@ -440,6 +455,7 @@ impl RequestBodyNormalizationError {
Self::UnsupportedContentEncoding(_) | Self::DecodeFailed { .. } => {
http::StatusCode::BAD_REQUEST
}
Self::BodyBufferOverloaded { .. } => http::StatusCode::SERVICE_UNAVAILABLE,
}
}
}
@@ -468,6 +484,10 @@ impl fmt::Display for RequestBodyNormalizationError {
Self::RequestBodyTooLarge { limit_bytes } => {
write!(f, "request body exceeds {limit_bytes} bytes")
}
Self::BodyBufferOverloaded { requested_bytes, budget_bytes } => write!(
f,
"request body buffering needs {requested_bytes} bytes of a {budget_bytes} byte budget"
),
}
}
}
@@ -489,9 +509,24 @@ pub(crate) fn normalize_request_body_headers_and_bytes_with_limit(
headers: &mut http::HeaderMap,
body_bytes: Bytes,
limit_bytes: u64,
) -> Result<Bytes, RequestBodyNormalizationError> {
normalize_request_body_headers_and_bytes_with_budget(
headers,
body_bytes,
limit_bytes,
&mut |_| Ok(()),
)
}
pub(crate) fn normalize_request_body_headers_and_bytes_with_budget(
headers: &mut http::HeaderMap,
body_bytes: Bytes,
limit_bytes: u64,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Bytes, RequestBodyNormalizationError> {
let body_was_encoded = !request_content_encodings(headers).is_empty();
let decoded = decoded_request_body_bytes_with_limit(headers, body_bytes.as_ref(), limit_bytes)?;
let decoded =
decoded_request_body_bytes_with_budget(headers, body_bytes.as_ref(), limit_bytes, budget)?;
if !body_was_encoded {
return Ok(body_bytes);
}
@@ -533,6 +568,15 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>(
headers: &http::HeaderMap,
body_bytes: &'a [u8],
limit: u64,
) -> Result<Cow<'a, [u8]>, RequestBodyNormalizationError> {
decoded_request_body_bytes_with_budget(headers, body_bytes, limit, &mut |_| Ok(()))
}
fn decoded_request_body_bytes_with_budget<'a>(
headers: &http::HeaderMap,
body_bytes: &'a [u8],
limit: u64,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Cow<'a, [u8]>, RequestBodyNormalizationError> {
validate_request_body_framing(headers)?;
let encodings = request_content_encodings(headers);
@@ -543,11 +587,21 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>(
return Ok(Cow::Borrowed(body_bytes));
}
let mut decoded = body_bytes.to_vec();
let mut decoded = Cow::Borrowed(body_bytes);
for encoding in encodings.iter().rev() {
decoded = decode_single_request_body_with_limit(encoding, decoded.as_slice(), limit)?;
let retained_input_bytes = body_bytes.len().saturating_add(match &decoded {
Cow::Borrowed(_) => 0,
Cow::Owned(bytes) => bytes.capacity(),
});
decoded = Cow::Owned(decode_single_request_body_with_budget(
encoding,
decoded.as_ref(),
limit,
retained_input_bytes,
budget,
)?);
}
Ok(Cow::Owned(decoded))
Ok(decoded)
}
fn request_content_encodings(headers: &http::HeaderMap) -> Vec<String> {
@@ -631,11 +685,56 @@ fn decode_single_request_body_with_limit(
encoding: &str,
body_bytes: &[u8],
limit: u64,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
decode_single_request_body_with_budget(
encoding,
body_bytes,
limit,
body_bytes.len(),
&mut |_| Ok(()),
)
}
fn decode_single_request_body_with_budget(
encoding: &str,
body_bytes: &[u8],
limit: u64,
retained_input_bytes: usize,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
match encoding {
"gzip" | "x-gzip" => decode_gzip_body_with_limit(encoding, body_bytes, limit),
"deflate" => decode_deflate_body_with_limit(encoding, body_bytes, limit),
"zstd" => decode_zstd_body_with_limit(encoding, body_bytes, limit),
"gzip" | "x-gzip" => {
let mut decoder = GzDecoder::new(body_bytes);
read_request_decoder_to_end_with_budget(
encoding,
&mut decoder,
limit,
retained_input_bytes,
budget,
)
}
"deflate" => decode_deflate_body_with_budget(
encoding,
body_bytes,
limit,
retained_input_bytes,
budget,
),
"zstd" => {
let mut decoder = zstd::stream::read::Decoder::new(body_bytes).map_err(|err| {
RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
}
})?;
read_request_decoder_to_end_with_budget(
encoding,
&mut decoder,
limit,
retained_input_bytes,
budget,
)
}
_ => Err(RequestBodyNormalizationError::UnsupportedContentEncoding(
encoding.to_string(),
)),
@@ -669,20 +768,48 @@ fn decode_deflate_body_with_limit(
encoding: &str,
body_bytes: &[u8],
limit: u64,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
decode_deflate_body_with_budget(encoding, body_bytes, limit, body_bytes.len(), &mut |_| {
Ok(())
})
}
fn decode_deflate_body_with_budget(
encoding: &str,
body_bytes: &[u8],
limit: u64,
retained_input_bytes: usize,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
let mut zlib_decoder = ZlibDecoder::new(body_bytes);
match read_request_decoder_to_end_with_limit(encoding, &mut zlib_decoder, limit) {
match read_request_decoder_to_end_with_budget(
encoding,
&mut zlib_decoder,
limit,
retained_input_bytes,
budget,
) {
Ok(decoded) => Ok(decoded),
Err(err @ RequestBodyNormalizationError::DecompressedBodyTooLarge { .. }) => Err(err),
Err(zlib_error) => {
Err(zlib_error @ RequestBodyNormalizationError::DecodeFailed { .. }) => {
let mut raw_decoder = DeflateDecoder::new(body_bytes);
read_request_decoder_to_end_with_limit(encoding, &mut raw_decoder, limit).map_err(
|raw_error| RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: format!("{zlib_error}; raw deflate fallback failed: {raw_error}"),
},
read_request_decoder_to_end_with_budget(
encoding,
&mut raw_decoder,
limit,
retained_input_bytes,
budget,
)
.map_err(|raw_error| match raw_error {
RequestBodyNormalizationError::DecodeFailed { .. } => {
RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: format!("{zlib_error}; raw deflate fallback failed: {raw_error}"),
}
}
error => error,
})
}
Err(error) => Err(error),
}
}
@@ -719,21 +846,70 @@ fn read_request_decoder_to_end_with_limit(
decoder: &mut impl Read,
limit: u64,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
let mut limited = decoder.take(limit.saturating_add(1));
read_request_decoder_to_end_with_budget(encoding, decoder, limit, 0, &mut |_| Ok(()))
}
fn read_request_decoder_to_end_with_budget(
encoding: &str,
decoder: &mut impl Read,
limit: u64,
retained_input_bytes: usize,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
let capacity_limit = usize::try_from(limit).unwrap_or(usize::MAX);
let mut scratch = [0_u8; 8 * 1024];
let mut out = Vec::new();
limited
.read_to_end(&mut out)
.map_err(|err| RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
loop {
let remaining = limit.saturating_sub(out.len() as u64).saturating_add(1);
let read_limit = scratch
.len()
.min(usize::try_from(remaining).unwrap_or(usize::MAX));
let read = decoder.read(&mut scratch[..read_limit]).map_err(|err| {
RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
}
})?;
if out.len() as u64 > limit {
return Err(RequestBodyNormalizationError::DecompressedBodyTooLarge {
encoding: encoding.to_string(),
limit_bytes: limit,
});
if read == 0 {
return Ok(out);
}
let next_len = out.len().saturating_add(read);
if next_len as u64 > limit {
return Err(RequestBodyNormalizationError::DecompressedBodyTooLarge {
encoding: encoding.to_string(),
limit_bytes: limit,
});
}
if next_len > out.capacity() {
let mut capacity = out
.capacity()
.saturating_mul(2)
.max(next_len)
.min(capacity_limit);
// The encoded body and previous decoding layer remain alive during growth.
match budget(retained_input_bytes.saturating_add(capacity)) {
Ok(()) => {}
Err(RequestBodyNormalizationError::BodyBufferOverloaded { .. })
if capacity > next_len =>
{
// Rejected reservations leave the budget unchanged. Spare capacity
// must not reject a body whose actual bytes still fit.
capacity = next_len;
budget(retained_input_bytes.saturating_add(capacity))?;
}
Err(error) => return Err(error),
}
out.try_reserve_exact(capacity.saturating_sub(out.len()))
.map_err(|err| RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
})?;
if out.capacity() > capacity {
budget(retained_input_bytes.saturating_add(out.capacity()))?;
}
}
out.extend_from_slice(&scratch[..read]);
}
Ok(out)
}
pub(crate) fn header_equals(
@@ -1223,6 +1399,14 @@ mod tests {
#[test]
fn request_body_normalization_error_maps_http_status() {
assert_eq!(
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes: 2,
budget_bytes: 1,
}
.http_status(),
http::StatusCode::SERVICE_UNAVAILABLE
);
assert_eq!(
RequestBodyNormalizationError::RequestBodyTooLarge { limit_bytes: 1 }.http_status(),
http::StatusCode::PAYLOAD_TOO_LARGE
@@ -1250,6 +1434,193 @@ mod tests {
);
}
#[test]
fn budgeted_normalization_accounts_for_encoded_input_and_output_capacity() {
let payload = vec![b'a'; 150_000];
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&payload).expect("gzip payload");
let encoded = encoder.finish().expect("gzip finish");
let encoded_len = encoded.len();
let mut headers = HeaderMap::new();
headers.insert(
http::header::CONTENT_ENCODING,
HeaderValue::from_static("gzip"),
);
let mut reservations = Vec::new();
let decoded = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.into(),
256 * 1024,
&mut |bytes| {
reservations.push(bytes);
Ok(())
},
)
.expect("budgeted gzip should decode");
assert_eq!(decoded.as_ref(), payload.as_slice());
assert!(!headers.contains_key(http::header::CONTENT_ENCODING));
assert_eq!(reservations[0], encoded_len + 8 * 1024);
assert!(reservations.windows(2).all(|pair| pair[0] <= pair[1]));
assert!(reservations.last().copied().unwrap() >= encoded_len + payload.len());
assert!(reservations.last().copied().unwrap() <= encoded_len + payload.len() * 2);
assert!(
reservations.len() <= 8,
"output growth should remain geometric"
);
}
#[test]
fn budgeted_normalization_accounts_for_retained_encoding_layers() {
let payload = b"small chained request body";
let mut inner = GzEncoder::new(Vec::new(), Compression::default());
inner.write_all(payload).expect("inner gzip payload");
let inner = inner.finish().expect("inner gzip finish");
let mut outer = GzEncoder::new(Vec::new(), Compression::default());
outer.write_all(&inner).expect("outer gzip payload");
let encoded = outer.finish().expect("outer gzip finish");
let encoded_len = encoded.len();
let mut headers = HeaderMap::new();
headers.insert(
http::header::CONTENT_ENCODING,
HeaderValue::from_static("gzip, gzip"),
);
let mut reservations = Vec::new();
let decoded = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.into(),
256,
&mut |bytes| {
reservations.push(bytes);
Ok(())
},
)
.expect("chained gzip should decode");
assert_eq!(decoded.as_ref(), payload);
assert_eq!(
reservations,
vec![
encoded_len + inner.len(),
encoded_len + inner.len() + payload.len(),
]
);
}
#[test]
fn budgeted_decoder_stops_before_collecting_when_capacity_is_exhausted() {
let source = vec![b'a'; 150_000];
let mut decoder = std::io::Cursor::new(source);
let error = super::read_request_decoder_to_end_with_budget(
"test",
&mut decoder,
256 * 1024,
100,
&mut |requested_bytes| {
Err(RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes: 100,
})
},
)
.expect_err("budget rejection must stop output growth");
assert_eq!(decoder.position(), 8 * 1024);
assert_eq!(
error,
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes: 100 + 8 * 1024,
budget_bytes: 100,
}
);
}
#[test]
fn budgeted_decoder_accepts_actual_output_when_geometric_growth_does_not_fit() {
let source = vec![b'a'; 100_000];
let mut decoder = std::io::Cursor::new(&source);
let retained_input_bytes = 100;
let budget_bytes = retained_input_bytes + source.len();
let mut reserved_bytes = retained_input_bytes;
let mut rejected_growth = false;
let decoded = super::read_request_decoder_to_end_with_budget(
"test",
&mut decoder,
256 * 1024,
retained_input_bytes,
&mut |requested_bytes| {
if requested_bytes > budget_bytes {
rejected_growth = true;
return Err(RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes,
});
}
reserved_bytes = reserved_bytes.max(requested_bytes);
Ok(())
},
)
.expect("actual output within the budget should finish decoding");
assert!(
rejected_growth,
"test must exercise oversized spare capacity"
);
assert_eq!(decoded, source);
assert_eq!(reserved_bytes, budget_bytes);
assert_eq!(decoded.capacity() + retained_input_bytes, budget_bytes);
}
#[test]
fn budgeted_deflate_preserves_capacity_and_size_rejections() {
let payload = [b'a'; 128];
let mut wrapped = ZlibEncoder::new(Vec::new(), Compression::default());
wrapped.write_all(&payload).expect("zlib payload");
let wrapped = wrapped.finish().expect("zlib finish");
let mut raw = DeflateEncoder::new(Vec::new(), Compression::default());
raw.write_all(&payload).expect("raw deflate payload");
let raw = raw.finish().expect("raw deflate finish");
for encoded in [wrapped, raw] {
let mut headers = HeaderMap::new();
headers.insert(
http::header::CONTENT_ENCODING,
HeaderValue::from_static("deflate"),
);
let rejection = RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes: 300,
budget_bytes: 200,
};
let error = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.clone().into(),
256,
&mut |_| Err(rejection.clone()),
)
.expect_err("capacity rejection must remain an overload error");
assert_eq!(error, rejection);
assert!(headers.contains_key(http::header::CONTENT_ENCODING));
let error = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.into(),
64,
&mut |_| Ok(()),
)
.expect_err("size rejection must remain a size error");
assert!(matches!(
error,
RequestBodyNormalizationError::DecompressedBodyTooLarge {
limit_bytes: 64,
..
}
));
}
}
#[test]
fn check_request_content_length_allows_missing_or_within_limit() {
let empty = HeaderMap::new();
+1
View File
@@ -71,6 +71,7 @@ mod rate_limit;
mod request_candidate_queue;
mod request_candidate_runtime;
mod request_diagnostics;
mod request_lifecycle;
mod roles;
mod router;
mod routing;
+1 -1
View File
@@ -44,7 +44,7 @@ pub(crate) fn local_auth_jwt_secret() -> Result<String, String> {
Err(std::env::VarError::NotPresent) => {
#[cfg(test)]
{
return Ok(TEST_JWT_SECRET.to_string());
Ok(TEST_JWT_SECRET.to_string())
}
#[cfg(not(test))]
+293 -51
View File
@@ -18,6 +18,7 @@ use hyper_util::{
server::conn::auto::Builder as HyperServerBuilder,
service::TowerToHyperService,
};
use tokio_util::sync::CancellationToken;
use tower::{Service as _, ServiceExt as _};
use tracing::{debug, info, warn};
@@ -117,8 +118,8 @@ where
use aether_crypto::warm_python_fernet_secret;
use aether_data::lifecycle::export::{
copy_database_records, export_database_jsonl, import_database_jsonl, DataCopyOptions,
ExportDomain, MAX_JSONL_INPUT_BYTES,
copy_database_records, export_database_jsonl, import_database_jsonl_with_options,
DataCopyOptions, DataImportOptions, ExportDomain, MAX_JSONL_INPUT_BYTES,
};
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use aether_gateway::{
@@ -127,6 +128,7 @@ use aether_gateway::{
FrontdoorCorsConfig, FrontdoorUserRpmConfig, GatewayDataConfig, UsageRuntimeConfig,
VideoTaskTruthSourceMode,
};
use aether_gateway_frontdoor::{http_connection_limit, HttpConnectionBudget};
use aether_runtime::{
init_service_runtime, FileLoggingConfig, LogDestination, LogFormat, LogRotation,
ServiceRuntimeConfig,
@@ -1007,6 +1009,13 @@ struct GatewayUsageArgs {
)]
queue_stream_maxlen: usize,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES",
default_value_t = 1024 * 1024
)]
queue_payload_max_bytes: usize,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_BATCH_SIZE",
@@ -1220,6 +1229,7 @@ impl GatewayUsageArgs {
consumer_group: self.queue_group.trim().to_string(),
dlq_stream_key: self.queue_dlq_stream_key.trim().to_string(),
stream_maxlen: self.queue_stream_maxlen.max(1),
queue_payload_max_bytes: self.queue_payload_max_bytes,
consumer_batch_size: self.queue_batch_size.max(1),
consumer_block_ms: self.queue_block_ms.max(1),
reclaim_idle_ms: self.queue_reclaim_idle_ms.max(1),
@@ -1351,6 +1361,11 @@ struct DataExportArgs {
struct DataImportArgs {
#[arg(long)]
input: PathBuf,
#[arg(
long,
help = "Preserve passwords and API/management credentials from a trusted import; imported sessions remain revoked. Without this flag identity credentials are revoked."
)]
preserve_credentials: bool,
}
#[derive(ClapArgs, Debug, Clone)]
@@ -1382,6 +1397,11 @@ struct DataCopyArgs {
#[arg(long)]
omit_request_body_details: bool,
#[arg(
long,
help = "Preserve passwords and API/management credentials from the trusted source; imported sessions remain revoked. The target must use the source encryption key."
)]
preserve_credentials: bool,
}
impl GatewayLoggingArgs {
@@ -1476,6 +1496,22 @@ struct Args {
/// Maximum number of HTTP/1 request header fields.
http_max_headers: usize,
#[arg(
long,
env = "AETHER_GATEWAY_HTTP_SHUTDOWN_TIMEOUT_MS",
default_value_t = 30_000
)]
/// Grace period for HTTP requests and upgraded connections before forced close.
http_shutdown_timeout_ms: u64,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_SHUTDOWN_TIMEOUT_MS",
default_value_t = 30_000
)]
/// Additional time for request finalizers and local usage buffers to persist.
usage_shutdown_timeout_ms: u64,
/// 容器内健康检查入口:根据当前 bind 端口探测本地 /health。
#[arg(long, hide = true, default_value_t = false)]
healthcheck: bool,
@@ -1557,6 +1593,11 @@ struct Args {
#[arg(long, env = "AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS")]
max_in_flight_requests: Option<usize>,
/// Maximum accepted HTTP TCP connections across all listener shards, including upgrades.
/// Unset or 0 follows request plus WebSocket capacity, bounded by the FD allowance.
#[arg(long, env = "AETHER_GATEWAY_MAX_HTTP_CONNECTIONS")]
max_http_connections: Option<usize>,
/// Maximum number of long-lived public WebSocket connections. When unset,
/// this follows `max_in_flight_requests` while remaining an independent
/// gate. Set `AETHER_GATEWAY_MAX_WEBSOCKET_CONNECTIONS` to override it.
@@ -1838,40 +1879,59 @@ fn gateway_listeners(
Ok(listeners)
}
async fn serve_gateway_router(
listeners: Vec<tokio::net::TcpListener>,
router: axum::Router,
#[derive(Clone, Copy)]
struct GatewayHttpLimits {
http2_max_concurrent_streams: u32,
http_header_read_timeout_ms: u64,
http_header_max_bytes: usize,
http_max_headers: usize,
}
async fn serve_gateway_router(
listeners: Vec<tokio::net::TcpListener>,
router: axum::Router,
connection_budget: Arc<HttpConnectionBudget>,
limits: GatewayHttpLimits,
shutdown: CancellationToken,
) -> Result<(), Box<dyn std::error::Error>> {
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 limits = GatewayHttpLimits {
http2_max_concurrent_streams: gateway_http2_max_concurrent_streams(
limits.http2_max_concurrent_streams,
),
http_header_read_timeout_ms: gateway_http_header_read_timeout_ms(
limits.http_header_read_timeout_ms,
),
http_header_max_bytes: gateway_http_header_max_bytes(limits.http_header_max_bytes),
http_max_headers: gateway_http_max_headers(limits.http_max_headers),
};
let mut servers = tokio::task::JoinSet::new();
for listener in listeners {
let router = router.clone();
let connection_budget = Arc::clone(&connection_budget);
let shutdown = shutdown.clone();
servers.spawn(async move {
serve_gateway_listener(
listener,
router,
http2_max_concurrent_streams,
http_header_read_timeout_ms,
http_header_max_bytes,
http_max_headers,
)
.await
serve_gateway_listener(listener, router, connection_budget, limits, shutdown).await
});
}
if let Some(result) = servers.join_next().await {
servers.abort_all();
let serve_result = result
.map_err(|err| std::io::Error::other(format!("gateway listener task failed: {err}")))?;
serve_result?;
let mut failure = None;
while let Some(result) = servers.join_next().await {
let result = result.unwrap_or_else(|err| {
Err(std::io::Error::other(format!(
"gateway listener task failed: {err}"
)))
});
if let Err(error) = result {
failure.get_or_insert(error);
shutdown.cancel();
connection_budget.force_close();
}
}
if let Some(error) = failure {
return Err(error.into());
}
// Hyper hands upgrades to application tasks; their IO still owns this budget.
while connection_budget.snapshot().in_flight != 0 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
Ok(())
}
@@ -1879,14 +1939,29 @@ async fn serve_gateway_router(
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,
connection_budget: Arc<HttpConnectionBudget>,
limits: GatewayHttpLimits,
shutdown: CancellationToken,
) -> Result<(), std::io::Error> {
let GatewayHttpLimits {
http2_max_concurrent_streams,
http_header_read_timeout_ms,
http_header_max_bytes,
http_max_headers,
} = limits;
let mut make_service = router.into_make_service_with_connect_info::<std::net::SocketAddr>();
let mut connections = tokio::task::JoinSet::new();
loop {
let (io, remote_addr) = listener.accept().await?;
let (io, remote_addr) = tokio::select! {
biased;
_ = shutdown.cancelled() => break,
_ = connections.join_next(), if !connections.is_empty() => continue,
accepted = connection_budget.accept(&listener) => accepted,
};
let Ok(io) = connection_budget.try_admit(io) else {
tokio::task::yield_now().await;
continue;
};
let tower_service = make_service
.call(remote_addr)
.await
@@ -1899,7 +1974,9 @@ async fn serve_gateway_listener(
});
let io = TokioIo::new(io);
tokio::spawn(async move {
let shutdown = shutdown.clone();
let connection_budget = Arc::clone(&connection_budget);
connections.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:
@@ -1928,17 +2005,34 @@ async fn serve_gateway_listener(
// 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;
let connection = builder.serve_connection_with_upgrades(io, hyper_service);
tokio::pin!(connection);
let draining_connection = async {
tokio::select! {
result = &mut connection => result,
_ = shutdown.cancelled() => {
connection.as_mut().graceful_shutdown();
connection.await
}
}
};
let connection_result = tokio::select! {
biased;
_ = connection_budget.wait_for_forced_close() => Ok(()),
result = drive_gateway_connection(
draining_connection,
first_request_gate,
std::time::Duration::from_millis(http_header_read_timeout_ms),
) => result,
};
if let Err(err) = connection_result {
tracing::trace!(error = ?err, "gateway connection closed with error");
}
});
}
drop(listener);
while connections.join_next().await.is_some() {}
Ok(())
}
fn resolve_local_http_base_url(app_port: u16) -> Result<String, std::io::Error> {
@@ -2055,11 +2149,14 @@ fn validate_deployment_topology(
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
tokio::runtime::Builder::new_multi_thread()
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.thread_stack_size(GATEWAY_TOKIO_WORKER_STACK_SIZE_BYTES)
.build()?
.block_on(run())
.build()?;
let result = runtime.block_on(run());
aether_usage_runtime::shutdown_usage_background_runtime(std::time::Duration::from_secs(5));
result
}
async fn run() -> Result<(), Box<dyn std::error::Error>> {
@@ -2123,6 +2220,13 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.max_websocket_connections
.filter(|limit| *limit > 0)
.unwrap_or(request_concurrency_limit);
let http_connection_limit = http_connection_limit(
args.max_http_connections,
request_concurrency_limit,
websocket_connection_limit,
soft_fd_limit(),
);
let http_connection_budget = Arc::new(HttpConnectionBudget::new(http_connection_limit));
let distributed_websocket_connection_limit = match args.distributed_websocket_connection_limit {
Some(limit) if limit > 0 => Some(limit),
Some(_) => None,
@@ -2313,7 +2417,8 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
}
state = state
.with_request_concurrency_limit(request_concurrency_limit)
.with_websocket_connection_limit(websocket_connection_limit);
.with_websocket_connection_limit(websocket_connection_limit)
.with_http_connection_budget(Arc::clone(&http_connection_budget));
if let Some(limit) = args.distributed_request_limit.filter(|limit| *limit > 0) {
let distributed_gate = state
.runtime_state()
@@ -2457,6 +2562,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
let listeners = gateway_listeners(bind_addr, listen_backlog, listener_shards)?;
let public_base_url = resolve_local_http_base_url(app_port)?;
let frontdoor_health_url = format!("{public_base_url}/_gateway/health");
let shutdown_state = state.clone();
let api_router = build_router_with_state(state);
// Compose the final router: API routes + optional static file serving.
@@ -2476,6 +2582,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
app_port,
listen_backlog,
listener_shards,
max_http_connections = http_connection_limit,
http2_max_concurrent_streams = gateway_http2_max_concurrent_streams(args.http2_max_concurrent_streams),
public_url = %public_base_url,
healthcheck_url = %frontdoor_health_url,
@@ -2483,18 +2590,63 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
"aether-gateway ready"
);
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?;
let shutdown = CancellationToken::new();
let serve_result = {
let server = serve_gateway_router(
listeners,
router,
Arc::clone(&http_connection_budget),
GatewayHttpLimits {
http2_max_concurrent_streams: args.http2_max_concurrent_streams,
http_header_read_timeout_ms: args.http_header_read_timeout_ms,
http_header_max_bytes: args.http_header_max_bytes,
http_max_headers: args.http_max_headers,
},
shutdown.clone(),
);
tokio::pin!(server);
tokio::select! {
result = &mut server => result,
signal = aether_runtime::wait_for_shutdown_signal() => {
signal?;
info!("shutdown signal received, draining gateway requests");
shutdown.cancel();
match tokio::time::timeout(
std::time::Duration::from_millis(args.http_shutdown_timeout_ms),
&mut server,
).await {
Ok(result) => result,
Err(_) => {
warn!(
event_name = "gateway_http_shutdown_deadline",
connections = http_connection_budget.snapshot().in_flight,
"HTTP drain deadline reached; closing remaining sockets"
);
http_connection_budget.force_close();
match tokio::time::timeout(std::time::Duration::from_secs(5), &mut server).await {
Ok(result) => result,
Err(_) => Err(std::io::Error::new(std::io::ErrorKind::TimedOut,
"gateway connection tasks did not stop after forced close").into()),
}
}
}
}
}
};
let usage_result = shutdown_state
.shutdown_usage_runtime(std::time::Duration::from_millis(
args.usage_shutdown_timeout_ms,
))
.await;
if let Some(background_tasks) = background_tasks {
background_tasks.shutdown().await;
}
serve_result?;
usage_result?;
info!(
event_name = "gateway_shutdown_complete",
"gateway local persistence drained"
);
Ok(())
}
@@ -2905,12 +3057,23 @@ async fn run_data_import(
let driver = database.driver;
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?;
if !args.preserve_credentials {
warn!("identity credentials will be revoked; use --preserve-credentials only for trusted recovery or migration");
}
let imported = import_database_jsonl_with_options(
database,
&input,
DataImportOptions {
preserve_credentials: args.preserve_credentials,
},
)
.await?;
info!(
driver = %driver,
input = %args.input.display(),
imported,
preserve_credentials = args.preserve_credentials,
"database import complete"
);
println!(
@@ -3171,6 +3334,9 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
let target_driver = target.driver;
let domains = requested_domains(&args.domains);
let created_at_unix_secs = current_unix_secs()?;
if !args.preserve_credentials {
warn!("identity credentials will be revoked; use --preserve-credentials only for trusted recovery or migration");
}
let imported = copy_database_records(
source,
target,
@@ -3178,6 +3344,7 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
created_at_unix_secs,
DataCopyOptions {
omit_request_body_details: args.omit_request_body_details,
preserve_credentials: args.preserve_credentials,
},
)
.await?;
@@ -3186,6 +3353,7 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
source_driver = %source_driver,
target_driver = %target_driver,
imported,
preserve_credentials = args.preserve_credentials,
"database copy complete"
);
println!(
@@ -3411,6 +3579,10 @@ fn pending_backfills_error(
#[cfg(test)]
mod tests {
mod shutdown {
include!("shutdown_tests.rs");
}
use super::{
automatic_gateway_request_concurrency_for_capacity,
automatic_gateway_request_concurrency_for_parallelism, automatic_sql_pool_config,
@@ -3452,6 +3624,8 @@ mod tests {
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,
http_shutdown_timeout_ms: 30_000,
usage_shutdown_timeout_ms: 30_000,
healthcheck: false,
healthcheck_timeout_ms: 3_000,
deployment_topology: DeploymentTopologyArg::SingleNode,
@@ -3466,6 +3640,7 @@ mod tests {
video_task_poller_batch_size: 32,
video_task_store_path: None,
max_in_flight_requests: None,
max_http_connections: None,
max_websocket_connections: None,
distributed_request_limit: None,
distributed_websocket_connection_limit: None,
@@ -3506,6 +3681,7 @@ mod tests {
queue_group: "usage_consumers".to_string(),
queue_dlq_stream_key: "usage:events:dlq".to_string(),
queue_stream_maxlen: 200_000,
queue_payload_max_bytes: 1024 * 1024,
queue_batch_size: 128,
queue_block_ms: 500,
queue_reclaim_idle_ms: 60_000,
@@ -3899,6 +4075,37 @@ mod tests {
}
}
#[test]
fn gateway_usage_queue_payload_limit_preserves_cli_override_and_rejects_zero() {
let command = <Args as clap::CommandFactory>::command();
let argument = command
.get_arguments()
.find(|argument| argument.get_id() == "queue_payload_max_bytes")
.expect("usage payload argument must be registered");
assert_eq!(
argument.get_env(),
Some(std::ffi::OsStr::new(
"AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES"
))
);
assert_eq!(argument.get_default_values()[0].to_str(), Some("1048576"));
let args = Args::try_parse_from(["aether-gateway", "--queue-payload-max-bytes", "32768"])
.expect("explicit usage payload limit should parse");
let config = args.usage.to_config(4, 8, Some(4));
assert_eq!(config.queue_payload_max_bytes, 32_768);
assert!(config.validate().is_ok());
let mut args = test_args();
assert_eq!(
args.usage.to_config(4, 8, Some(4)).queue_payload_max_bytes,
1024 * 1024
);
args.usage.queue_payload_max_bytes = 0;
let config = args.usage.to_config(4, 8, Some(4));
assert_eq!(config.queue_payload_max_bytes, 0);
assert!(config.validate().is_err());
}
#[test]
fn gateway_usage_queue_workers_manual_override_wins_and_is_capped() {
let mut args = test_args();
@@ -4357,6 +4564,41 @@ mod tests {
};
assert!(copy.source_allow_insecure);
assert!(!copy.target_allow_insecure);
assert!(!copy.preserve_credentials);
}
#[test]
fn data_import_and_copy_require_explicit_credential_preservation() {
for preserve in [false, true] {
let mut import_args = vec!["aether-gateway", "import", "--input", "trusted.jsonl"];
let mut copy_args = vec![
"aether-gateway",
"copy",
"--source-driver",
"postgres",
"--source-url",
"postgres://localhost/source",
"--target-driver",
"postgres",
"--target-url",
"postgres://localhost/target",
];
if preserve {
import_args.push("--preserve-credentials");
copy_args.push("--preserve-credentials");
}
let Some(DataCommand::Import(import)) =
Args::try_parse_from(import_args).unwrap().command
else {
panic!("expected import command");
};
let Some(DataCommand::Copy(copy)) = Args::try_parse_from(copy_args).unwrap().command
else {
panic!("expected copy command");
};
assert_eq!(import.preserve_credentials, preserve);
assert_eq!(copy.preserve_credentials, preserve);
}
}
#[cfg(unix)]
+7 -7
View File
@@ -26,11 +26,11 @@ pub(crate) use runtime::{
start_manual_usage_cleanup_task, start_proxy_upgrade_rollout, AccountSelfCheckRunSummary,
AdminCleanupRunRecord, AdminCleanupTaskKind, AdminStatsRebuildSummary,
AdminSystemCleanupSummary, ManualUsageCleanupError, ManualUsageCleanupMode,
ManualUsageCleanupOptions, OAuthTokenRefreshRunSummary, PoolQuotaProbeRunSummary,
PoolQuotaProbeWorkerConfig, ProviderCheckinRunSummary, ProviderQuotaAlertRunSummary,
ProxyUpgradeRolloutCancelSummary, ProxyUpgradeRolloutConflictClearSummary,
ProxyUpgradeRolloutNodeActionSummary, ProxyUpgradeRolloutProbeConfig,
ProxyUpgradeRolloutSkippedRestoreSummary, ProxyUpgradeRolloutStatus,
ProxyUpgradeRolloutTrackedNodeState, UsageCounterFlushRuntimeMetrics,
UsageCounterFlushWorkerConfig,
ManualUsageCleanupOptions, OAuthTokenRefreshRunSummary, PoolQuotaProbeReplenishCoordinator,
PoolQuotaProbeRunSummary, PoolQuotaProbeWorkerConfig, ProviderCheckinRunSummary,
ProviderQuotaAlertRunSummary, ProxyUpgradeRolloutCancelSummary,
ProxyUpgradeRolloutConflictClearSummary, ProxyUpgradeRolloutNodeActionSummary,
ProxyUpgradeRolloutProbeConfig, ProxyUpgradeRolloutSkippedRestoreSummary,
ProxyUpgradeRolloutStatus, ProxyUpgradeRolloutTrackedNodeState,
UsageCounterFlushRuntimeMetrics, UsageCounterFlushWorkerConfig,
};
@@ -84,7 +84,8 @@ pub(crate) use pool_quota_probe::{
perform_pool_quota_probe_once, perform_pool_quota_probe_once_for_provider_with_config,
perform_pool_quota_probe_once_with_config, pool_quota_probe_target_count,
select_pool_quota_probe_key_ids, spawn_pool_quota_probe_replenish_for_request,
spawn_pool_quota_probe_worker, PoolQuotaProbeRunSummary, PoolQuotaProbeWorkerConfig,
spawn_pool_quota_probe_worker, PoolQuotaProbeReplenishCoordinator, PoolQuotaProbeRunSummary,
PoolQuotaProbeWorkerConfig,
};
pub(crate) use pool_score_rebuild::{
ensure_provider_key_pool_scores_for_keys, perform_pool_score_rebuild_once,
@@ -115,9 +115,6 @@ pub(super) fn usage_cleanup_window(
usage_cleanup_window_with_override(now_utc, settings, None)
}
/// Clamp is non-aggressive: each tier's cutoff becomes `max(policy_cutoff, now - override)`.
/// A later cutoff = fewer records deleted, so the override can only make cleanup more
/// conservative than the configured retention, never more destructive.
pub(super) fn usage_cleanup_window_with_override(
now_utc: DateTime<Utc>,
settings: UsageCleanupSettings,
@@ -135,9 +132,9 @@ pub(super) fn usage_cleanup_window_with_override(
};
let manual_cutoff = now_utc - override_duration;
UsageCleanupWindow {
detail_cutoff: policy.detail_cutoff.max(manual_cutoff),
compressed_cutoff: policy.compressed_cutoff.max(manual_cutoff),
header_cutoff: policy.header_cutoff.max(manual_cutoff),
log_cutoff: policy.log_cutoff.max(manual_cutoff),
detail_cutoff: policy.detail_cutoff.min(manual_cutoff),
compressed_cutoff: policy.compressed_cutoff.min(manual_cutoff),
header_cutoff: policy.header_cutoff.min(manual_cutoff),
log_cutoff: policy.log_cutoff.min(manual_cutoff),
}
}
@@ -1,4 +1,6 @@
use std::collections::{BTreeMap, BTreeSet};
use std::collections::{BTreeMap, BTreeSet, HashMap};
use std::future::Future;
use std::sync::{Arc, Mutex};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_data_contracts::repository::pool_scores::{
@@ -47,6 +49,166 @@ const POOL_QUOTA_PROBE_BURST_RETRY_GUARD_SECONDS: u64 = 15;
const POOL_QUOTA_PROBE_AUTO_MIN_INTERVAL_SECONDS: u64 = 30;
const POOL_QUOTA_PROBE_AUTO_MAX_INTERVAL_SECONDS: u64 = 10 * 60;
const POOL_QUOTA_PROBE_AUTO_MAX_PRESSURE: u64 = 64;
const POOL_QUOTA_PROBE_LOCAL_MAX_PROVIDERS: usize = 1024;
#[derive(Debug)]
pub(crate) struct PoolQuotaProbeReplenishCoordinator {
capacity: usize,
state: Mutex<PoolQuotaProbeReplenishState>,
}
#[derive(Debug, Default)]
struct PoolQuotaProbeReplenishState {
providers: HashMap<String, PoolQuotaProbeReplenishEntry>,
started_total: u64,
coalesced_total: u64,
capacity_rejected_total: u64,
}
#[derive(Debug)]
struct PoolQuotaProbeReplenishEntry {
identity: Arc<()>,
pending: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct PoolQuotaProbeReplenishSnapshot {
pub(crate) capacity: usize,
pub(crate) active: usize,
pub(crate) started_total: u64,
pub(crate) coalesced_total: u64,
pub(crate) capacity_rejected_total: u64,
}
impl Default for PoolQuotaProbeReplenishCoordinator {
fn default() -> Self {
Self::new(POOL_QUOTA_PROBE_LOCAL_MAX_PROVIDERS)
}
}
impl PoolQuotaProbeReplenishCoordinator {
fn new(capacity: usize) -> Self {
Self {
capacity,
state: Mutex::new(PoolQuotaProbeReplenishState::default()),
}
}
pub(crate) fn snapshot(&self) -> PoolQuotaProbeReplenishSnapshot {
let state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
PoolQuotaProbeReplenishSnapshot {
capacity: self.capacity,
active: state.providers.len(),
started_total: state.started_total,
coalesced_total: state.coalesced_total,
capacity_rejected_total: state.capacity_rejected_total,
}
}
fn request(self: &Arc<Self>, provider_id: String) -> Option<PoolQuotaProbeReplenishGuard> {
let mut state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(entry) = state.providers.get_mut(&provider_id) {
entry.pending = true;
state.coalesced_total = state.coalesced_total.saturating_add(1);
return None;
}
if state.providers.len() >= self.capacity {
// Replenishment is best effort; the periodic base scan remains available.
state.capacity_rejected_total = state.capacity_rejected_total.saturating_add(1);
return None;
}
let identity = Arc::new(());
state.providers.insert(
provider_id.clone(),
PoolQuotaProbeReplenishEntry {
identity: Arc::clone(&identity),
pending: true,
},
);
state.started_total = state.started_total.saturating_add(1);
Some(PoolQuotaProbeReplenishGuard {
coordinator: Arc::clone(self),
provider_id,
identity,
finished: false,
})
}
fn spawn<F, Fut>(
self: &Arc<Self>,
provider_id: String,
mut replenish: F,
) -> Option<tokio::task::JoinHandle<()>>
where
F: FnMut() -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
// Own the guard before spawn so cancellation before the first poll also cleans up.
let mut guard = self.request(provider_id)?;
Some(tokio::spawn(async move {
while guard.next_pass() {
replenish().await;
}
}))
}
}
struct PoolQuotaProbeReplenishGuard {
coordinator: Arc<PoolQuotaProbeReplenishCoordinator>,
provider_id: String,
identity: Arc<()>,
finished: bool,
}
impl PoolQuotaProbeReplenishGuard {
fn next_pass(&mut self) -> bool {
let mut state = self
.coordinator
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(entry) = state
.providers
.get_mut(&self.provider_id)
.filter(|entry| Arc::ptr_eq(&entry.identity, &self.identity))
{
if entry.pending {
entry.pending = false;
return true;
}
// Check for a follow-up and release ownership in one critical section.
state.providers.remove(&self.provider_id);
}
self.finished = true;
false
}
}
impl Drop for PoolQuotaProbeReplenishGuard {
fn drop(&mut self) {
if self.finished {
return;
}
let mut state = self
.coordinator
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if state
.providers
.get(&self.provider_id)
.is_some_and(|entry| Arc::ptr_eq(&entry.identity, &self.identity))
{
state.providers.remove(&self.provider_id);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PoolQuotaProbeMode {
@@ -1672,38 +1834,62 @@ pub(crate) fn spawn_pool_quota_probe_replenish_for_request(
return None;
}
Some(tokio::spawn(async move {
let runtime = state.runtime_state.clone();
mark_probe_burst_pending(runtime.as_ref(), &provider_id).await;
let lease =
acquire_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), &provider_id).await;
if lease.is_none() {
return;
}
let coordinator = Arc::clone(&state.pool_quota_probe_replenish);
coordinator.spawn(provider_id.clone(), move || {
run_pool_quota_probe_replenish(state.clone(), provider_id.clone())
})
}
let config = PoolQuotaProbeWorkerConfig::from_env();
loop {
let pending = runtime
.kv_take(&probe_burst_pending_key(&provider_id))
.await
.ok()
.flatten()
.is_some();
if !pending {
break;
}
match perform_pool_quota_probe_once_for_provider_with_mode(
async fn run_pool_quota_probe_replenish(state: AppState, provider_id: String) {
let runtime = state.runtime_state.clone();
let runtime = runtime.as_ref();
let config = PoolQuotaProbeWorkerConfig::from_env();
run_pool_quota_probe_replenish_with(
runtime,
&provider_id,
|| {
perform_pool_quota_probe_once_for_provider_with_mode(
&state,
&provider_id,
config,
PoolQuotaProbeMode::Burst,
)
.await
{
},
|lease| release_pool_quota_probe_burst_trigger_lock(runtime, Some(lease)),
)
.await;
}
async fn run_pool_quota_probe_replenish_with<Probe, ProbeFuture, Release, ReleaseFuture>(
runtime: &RuntimeState,
provider_id: &str,
mut probe: Probe,
mut release: Release,
) where
Probe: FnMut() -> ProbeFuture,
ProbeFuture: Future<Output = Result<PoolQuotaProbeRunSummary, GatewayError>>,
Release: FnMut(RuntimeLockLease) -> ReleaseFuture,
ReleaseFuture: Future<Output = ()>,
{
mark_probe_burst_pending(runtime, provider_id).await;
loop {
let Some(lease) = acquire_pool_quota_probe_burst_trigger_lock(runtime, provider_id).await
else {
return;
};
let recheck_after_release = loop {
let pending = match runtime.kv_take(&probe_burst_pending_key(provider_id)).await {
Ok(pending) => pending.is_some(),
Err(_) => break false,
};
if !pending {
break true;
}
match probe().await {
Ok(summary) => {
if summary.providers_busy > 0 {
mark_probe_burst_pending(runtime.as_ref(), &provider_id).await;
mark_probe_burst_pending(runtime, provider_id).await;
tokio::time::sleep(Duration::from_millis(250)).await;
continue;
}
@@ -1717,17 +1903,30 @@ pub(crate) fn spawn_pool_quota_probe_replenish_for_request(
}
}
let still_pending = runtime
.kv_exists(&probe_burst_pending_key(&provider_id))
match runtime
.kv_exists(&probe_burst_pending_key(provider_id))
.await
.unwrap_or(false);
if !still_pending {
break;
{
Ok(true) => {}
Ok(false) => break true,
Err(_) => break false,
}
}
};
release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), lease).await;
}))
release(lease).await;
// A different instance may publish after the final pending check and fail
// to acquire our old lease. Recheck after release, then acquire a fresh token
// before consuming that signal. Read failures terminate instead of spinning.
if !recheck_after_release
|| !runtime
.kv_exists(&probe_burst_pending_key(provider_id))
.await
.unwrap_or(false)
{
return;
}
tokio::task::yield_now().await;
}
}
pub(crate) fn spawn_pool_quota_probe_worker(
@@ -1764,7 +1963,430 @@ pub(crate) fn spawn_pool_quota_probe_worker(
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use aether_runtime_state::MemoryRuntimeStateConfig;
use serde_json::json;
use tokio::sync::Notify;
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn pool_quota_probe_local_coalesces_before_spawn_and_keeps_one_follow_up() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(8));
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let probe_calls = Arc::new(AtomicUsize::new(0));
let release_calls = Arc::new(AtomicUsize::new(0));
let started = Arc::new(Notify::new());
let finish_first = Arc::new(Notify::new());
let leader = coordinator
.spawn("provider".to_string(), {
let runtime = Arc::clone(&runtime);
let probe_calls = Arc::clone(&probe_calls);
let release_calls = Arc::clone(&release_calls);
let started = Arc::clone(&started);
let finish_first = Arc::clone(&finish_first);
move || {
let runtime = Arc::clone(&runtime);
let probe_calls = Arc::clone(&probe_calls);
let release_calls = Arc::clone(&release_calls);
let started = Arc::clone(&started);
let finish_first = Arc::clone(&finish_first);
async move {
run_pool_quota_probe_replenish_with(
runtime.as_ref(),
"provider",
|| async {
if probe_calls.fetch_add(1, Ordering::AcqRel) == 0 {
started.notify_one();
finish_first.notified().await;
}
Ok(PoolQuotaProbeRunSummary::empty())
},
|lease| {
release_calls.fetch_add(1, Ordering::AcqRel);
release_pool_quota_probe_burst_trigger_lock(
runtime.as_ref(),
Some(lease),
)
},
)
.await;
}
}
})
.expect("one leader");
tokio::time::timeout(Duration::from_secs(2), started.notified())
.await
.expect("first probe starts");
let barrier = Arc::new(tokio::sync::Barrier::new(65));
let mut triggers = tokio::task::JoinSet::new();
for _ in 0..64 {
let coordinator = Arc::clone(&coordinator);
let barrier = Arc::clone(&barrier);
triggers.spawn(async move {
barrier.wait().await;
assert!(coordinator
.spawn("provider".to_string(), || async {
panic!("a duplicate trigger must not spawn work")
})
.is_none());
});
}
barrier.wait().await;
while let Some(result) = triggers.join_next().await {
result.expect("concurrent trigger");
}
assert_eq!(coordinator.snapshot().started_total, 1);
assert_eq!(coordinator.snapshot().coalesced_total, 64);
assert_eq!(probe_calls.load(Ordering::Acquire), 1);
assert!(
!runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("pending read"),
"local duplicates must not each write Redis pending while the leader is running"
);
finish_first.notify_one();
tokio::time::timeout(Duration::from_secs(2), leader)
.await
.expect("leader finishes")
.expect("leader task");
assert_eq!(probe_calls.load(Ordering::Acquire), 2);
assert_eq!(
release_calls.load(Ordering::Acquire),
2,
"64 retriggers produce one additional Redis lock/drain cycle"
);
assert_eq!(coordinator.snapshot().active, 0);
}
#[tokio::test]
async fn pool_quota_probe_local_different_providers_run_independently() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(2));
let started = Arc::new(tokio::sync::Semaphore::new(0));
let finish = Arc::new(tokio::sync::Semaphore::new(0));
let mut tasks = Vec::new();
for provider_id in ["provider-a", "provider-b"] {
tasks.push(
coordinator
.spawn(provider_id.to_string(), {
let started = Arc::clone(&started);
let finish = Arc::clone(&finish);
move || {
let started = Arc::clone(&started);
let finish = Arc::clone(&finish);
async move {
started.add_permits(1);
finish.acquire().await.expect("finish signal").forget();
}
}
})
.expect("independent provider leader"),
);
}
tokio::time::timeout(Duration::from_secs(2), started.acquire_many(2))
.await
.expect("both providers start")
.expect("started permits")
.forget();
assert_eq!(coordinator.snapshot().active, 2);
finish.add_permits(2);
for task in tasks {
task.await.expect("provider finishes");
}
assert_eq!(coordinator.snapshot().active, 0);
}
#[test]
fn pool_quota_probe_local_exit_handoff_keeps_exactly_one_owner() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
for _ in 0..100 {
let mut first = coordinator
.request("provider".to_string())
.expect("first owner");
assert!(first.next_pass());
let barrier = Arc::new(std::sync::Barrier::new(2));
let (mut first, continues, replacement) = std::thread::scope(|scope| {
let first_barrier = Arc::clone(&barrier);
let exit = scope.spawn(move || {
first_barrier.wait();
let continues = first.next_pass();
(first, continues)
});
let trigger = scope.spawn(|| {
barrier.wait();
coordinator.request("provider".to_string())
});
let (first, continues) = exit.join().expect("exit thread");
(first, continues, trigger.join().expect("trigger thread"))
});
assert_ne!(
continues,
replacement.is_some(),
"the signal is consumed by exactly one owner"
);
assert_eq!(coordinator.snapshot().active, 1);
if continues {
assert!(!first.next_pass());
}
drop(first);
if replacement.is_some() {
assert_eq!(
coordinator.snapshot().active,
1,
"old cleanup must not delete the replacement"
);
}
drop(replacement);
assert_eq!(coordinator.snapshot().active, 0);
}
}
#[tokio::test]
async fn pool_quota_probe_local_abort_and_panic_release_admission() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
let unpolled = coordinator
.spawn("provider".to_string(), || async {
std::future::pending::<()>().await;
})
.expect("unpolled owner");
unpolled.abort();
assert!(unpolled.await.expect_err("aborted").is_cancelled());
assert_eq!(coordinator.snapshot().active, 0);
let started = Arc::new(Notify::new());
let running = coordinator
.spawn("provider".to_string(), {
let started = Arc::clone(&started);
move || {
let started = Arc::clone(&started);
async move {
started.notify_one();
std::future::pending::<()>().await;
}
}
})
.expect("running owner");
started.notified().await;
assert!(coordinator.request("provider".to_string()).is_none());
running.abort();
assert!(running.await.expect_err("aborted").is_cancelled());
assert_eq!(coordinator.snapshot().active, 0);
let panicked = coordinator
.spawn("provider".to_string(), || async {
panic!("probe panicked")
})
.expect("panic owner");
assert!(panicked.await.expect_err("probe panic").is_panic());
assert_eq!(coordinator.snapshot().active, 0);
coordinator
.spawn("provider".to_string(), || std::future::ready(()))
.expect("later trigger can run")
.await
.expect("recovered probe");
assert_eq!(coordinator.snapshot().active, 0);
}
#[test]
fn pool_quota_probe_local_capacity_is_bounded_and_completed_keys_are_removed() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(2));
let first = coordinator.request("a".to_string()).expect("first");
let second = coordinator.request("b".to_string()).expect("second");
assert!(coordinator.request("c".to_string()).is_none());
assert!(coordinator.request("a".to_string()).is_none());
assert_eq!(coordinator.snapshot().capacity_rejected_total, 1);
assert_eq!(coordinator.snapshot().coalesced_total, 1);
drop(first);
let replacement = coordinator
.request("c".to_string())
.expect("freed capacity");
drop((second, replacement));
for index in 0..1000 {
drop(
coordinator
.request(format!("provider-{index}"))
.expect("new provider"),
);
assert_eq!(coordinator.snapshot().active, 0);
}
assert_eq!(coordinator.snapshot().started_total, 1003);
}
#[tokio::test]
async fn pool_quota_probe_local_state_clones_share_only_the_same_runtime_binding() {
let state = AppState::new().expect("state");
let cloned = state.clone();
assert!(Arc::ptr_eq(
&state.pool_quota_probe_replenish,
&cloned.pool_quota_probe_replenish
));
let rebound_same = cloned.with_runtime_state(Arc::clone(&state.runtime_state));
assert!(Arc::ptr_eq(
&state.pool_quota_probe_replenish,
&rebound_same.pool_quota_probe_replenish
));
let rebound = state
.clone()
.with_runtime_state(Arc::new(RuntimeState::memory(
MemoryRuntimeStateConfig::default(),
)));
assert!(!Arc::ptr_eq(
&state.pool_quota_probe_replenish,
&rebound.pool_quota_probe_replenish
));
let first = state
.pool_quota_probe_replenish
.request("provider".to_string())
.expect("first runtime");
assert!(rebound_same
.pool_quota_probe_replenish
.request("provider".to_string())
.is_none());
let second = rebound
.pool_quota_probe_replenish
.request("provider".to_string())
.expect("other runtime is independent");
drop((first, second));
}
#[tokio::test]
async fn pool_quota_probe_replenish_rechecks_remote_pending_after_unlock() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let first = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
let second = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
let before_unlock = Arc::new(Notify::new());
let finish_unlock = Arc::new(Notify::new());
let first_probes = Arc::new(AtomicUsize::new(0));
let first_releases = Arc::new(AtomicUsize::new(0));
let first_task = first
.spawn("provider".to_string(), {
let runtime = Arc::clone(&runtime);
let before_unlock = Arc::clone(&before_unlock);
let finish_unlock = Arc::clone(&finish_unlock);
let first_probes = Arc::clone(&first_probes);
let first_releases = Arc::clone(&first_releases);
move || {
let runtime = Arc::clone(&runtime);
let before_unlock = Arc::clone(&before_unlock);
let finish_unlock = Arc::clone(&finish_unlock);
let first_probes = Arc::clone(&first_probes);
let first_releases = Arc::clone(&first_releases);
async move {
run_pool_quota_probe_replenish_with(
runtime.as_ref(),
"provider",
|| {
first_probes.fetch_add(1, Ordering::AcqRel);
std::future::ready(Ok(PoolQuotaProbeRunSummary::empty()))
},
|lease| {
let runtime = Arc::clone(&runtime);
let before_unlock = Arc::clone(&before_unlock);
let finish_unlock = Arc::clone(&finish_unlock);
let first_release =
first_releases.fetch_add(1, Ordering::AcqRel) == 0;
async move {
if first_release {
before_unlock.notify_one();
finish_unlock.notified().await;
}
release_pool_quota_probe_burst_trigger_lock(
runtime.as_ref(),
Some(lease),
)
.await;
}
},
)
.await;
}
}
})
.expect("first instance leader");
tokio::time::timeout(Duration::from_secs(2), before_unlock.notified())
.await
.expect("first instance drained but still owns Redis lease");
assert_eq!(first_probes.load(Ordering::Acquire), 1);
assert!(!runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("drained pending"));
let second_task = second.spawn("provider".to_string(), {
let runtime = Arc::clone(&runtime);
move || {
let runtime = Arc::clone(&runtime);
async move {
run_pool_quota_probe_replenish_with(
runtime.as_ref(), "provider",
|| async { panic!("second instance must not consume pending without the Redis lease") },
|lease| release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), Some(lease)),
).await;
}
}
}).expect("independent second instance leader");
second_task
.await
.expect("second instance leaves pending for the current owner");
assert_eq!(second.snapshot().active, 0);
assert!(runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("new pending signal"));
finish_unlock.notify_one();
tokio::time::timeout(Duration::from_secs(2), first_task)
.await
.expect("handoff drains")
.expect("first instance finishes");
assert_eq!(
first_probes.load(Ordering::Acquire),
2,
"the cross-instance exit signal must trigger a second probe"
);
assert_eq!(first_releases.load(Ordering::Acquire), 2);
assert_eq!(first.snapshot().active, 0);
assert!(!runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("all pending consumed"));
let lease = acquire_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), "provider").await;
assert!(lease.is_some(), "the replacement Redis lease is released");
release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), lease).await;
}
#[tokio::test]
async fn pool_quota_probe_replenish_public_spawn_returns_none_for_merged_triggers() {
let repository = Arc::new(
aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository::seed(
Vec::new(),
Vec::new(),
Vec::new(),
),
);
let state = AppState::new().expect("state").with_data_state_for_tests(
crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(repository),
);
let leader =
spawn_pool_quota_probe_replenish_for_request(state.clone(), "provider".to_string())
.expect("leader handle");
for _ in 0..64 {
assert!(spawn_pool_quota_probe_replenish_for_request(
state.clone(),
"provider".to_string()
)
.is_none());
}
leader.await.expect("leader completes");
assert_eq!(state.pool_quota_probe_replenish.snapshot().started_total, 1);
assert_eq!(
state.pool_quota_probe_replenish.snapshot().coalesced_total,
64
);
assert_eq!(state.pool_quota_probe_replenish.snapshot().active, 0);
spawn_pool_quota_probe_replenish_for_request(state.clone(), "provider".to_string())
.expect("later leader handle")
.await
.expect("later leader finishes");
}
#[test]
fn worker_error_score_reason_drops_runtime_error_details() {
@@ -1140,19 +1140,40 @@ fn usage_cleanup_window_with_override_is_always_non_aggressive() {
let override_duration = chrono::Duration::days(180);
let clamped = usage_cleanup_window_with_override(now_utc, settings, Some(override_duration));
assert_eq!(clamped.detail_cutoff, policy.detail_cutoff);
assert_eq!(clamped.compressed_cutoff, policy.compressed_cutoff);
assert_eq!(clamped.header_cutoff, policy.header_cutoff);
assert_eq!(clamped.log_cutoff, now_utc - override_duration);
assert!(clamped.log_cutoff > policy.log_cutoff);
assert_eq!(clamped.detail_cutoff, now_utc - override_duration);
assert_eq!(clamped.compressed_cutoff, now_utc - override_duration);
assert_eq!(clamped.header_cutoff, now_utc - override_duration);
assert_eq!(clamped.log_cutoff, policy.log_cutoff);
assert!(clamped.log_cutoff <= policy.log_cutoff);
let far_override = chrono::Duration::days(5);
let far = usage_cleanup_window_with_override(now_utc, settings, Some(far_override));
assert_eq!(far.detail_cutoff, now_utc - far_override);
assert_eq!(far.compressed_cutoff, now_utc - far_override);
assert_eq!(far.header_cutoff, now_utc - far_override);
assert_eq!(far.log_cutoff, now_utc - far_override);
assert!(far.log_cutoff > policy.log_cutoff);
assert_eq!(far, policy);
for days in [0, 5, 30, 180, 400] {
let cutoff = now_utc - chrono::Duration::days(days);
let window = usage_cleanup_window_with_override(
now_utc,
settings,
Some(chrono::Duration::days(days)),
);
for (actual, configured) in [
(window.detail_cutoff, policy.detail_cutoff),
(window.compressed_cutoff, policy.compressed_cutoff),
(window.header_cutoff, policy.header_cutoff),
(window.log_cutoff, policy.log_cutoff),
] {
assert!(actual <= configured);
assert!(actual <= cutoff);
for age in [1, 7, 15, 30, 90, 180, 365, 401] {
let created_at = now_utc - chrono::Duration::days(age);
if created_at < actual {
assert!(created_at < configured);
assert!(created_at < cutoff);
}
}
}
}
let passthrough = usage_cleanup_window_with_override(now_utc, settings, None);
assert_eq!(passthrough, policy);

Some files were not shown because too many files have changed in this diff Show More