mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
Compare commits
22
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
28f61ec45b | ||
|
|
6aeadcd1d7 | ||
|
|
3a8dadcd6b | ||
|
|
ecc16673eb | ||
|
|
d28dd89039 | ||
|
|
8260a87215 | ||
|
|
361952ada9 | ||
|
|
6630856061 | ||
|
|
a893bd0557 | ||
|
|
f2839ae6a7 | ||
|
|
e58570d79d | ||
|
|
99f6499b2b | ||
|
|
17d01d7fe0 | ||
|
|
8b766930b0 | ||
|
|
c7e403b410 | ||
|
|
cf8ea19856 | ||
|
|
7113d04f8a | ||
|
|
099b810a2f | ||
|
|
7aa0c89244 | ||
|
|
7847ae98c6 | ||
|
|
a90d564931 | ||
|
|
a5c3699ae9 |
+58
-12
@@ -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=
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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,不会直接写数据库;数据库导入仍应在维护窗口通过管理端完成。
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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!({
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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))]
|
||||
{
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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)]
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user