mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
28f61ec45b | ||
|
|
6aeadcd1d7 | ||
|
|
3a8dadcd6b | ||
|
|
ecc16673eb | ||
|
|
d28dd89039 | ||
|
|
8260a87215 | ||
|
|
361952ada9 | ||
|
|
6630856061 | ||
|
|
a893bd0557 | ||
|
|
f2839ae6a7 | ||
|
|
e58570d79d | ||
|
|
99f6499b2b | ||
|
|
17d01d7fe0 |
+50
-3
@@ -83,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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -62,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 构建)
|
||||
|
||||
@@ -121,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`;满载时拒绝新事件,避免控制面故障导致无界内存增长
|
||||
@@ -144,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");
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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;
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -36,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;
|
||||
@@ -50,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,
|
||||
@@ -109,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;
|
||||
@@ -382,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(|| {
|
||||
@@ -438,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;
|
||||
@@ -446,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(
|
||||
@@ -491,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)
|
||||
}
|
||||
|
||||
@@ -1204,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>,
|
||||
@@ -1220,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>,
|
||||
}
|
||||
|
||||
@@ -1353,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,
|
||||
})
|
||||
}
|
||||
@@ -1500,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,
|
||||
}))
|
||||
}
|
||||
@@ -2599,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)
|
||||
}
|
||||
|
||||
@@ -2860,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,
|
||||
@@ -4220,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()
|
||||
},
|
||||
);
|
||||
@@ -4239,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> {
|
||||
@@ -4591,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);
|
||||
}
|
||||
@@ -5149,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(|_| {
|
||||
@@ -5316,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,
|
||||
@@ -5440,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",
|
||||
@@ -5452,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();
|
||||
@@ -6228,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()
|
||||
@@ -6341,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(
|
||||
@@ -6439,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(),
|
||||
@@ -6626,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());
|
||||
@@ -6701,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 ",
|
||||
@@ -6719,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(),
|
||||
@@ -6745,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()
|
||||
@@ -6770,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!(
|
||||
@@ -6781,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),
|
||||
@@ -6791,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),
|
||||
@@ -6801,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");
|
||||
|
||||
@@ -6810,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(),
|
||||
@@ -6872,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");
|
||||
@@ -6932,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(),
|
||||
@@ -6985,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(),
|
||||
@@ -8569,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
|
||||
@@ -8736,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
|
||||
@@ -8777,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
|
||||
@@ -9104,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
|
||||
@@ -9320,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
|
||||
|
||||
@@ -655,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)]
|
||||
@@ -717,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 {
|
||||
@@ -903,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;
|
||||
};
|
||||
@@ -2465,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");
|
||||
|
||||
@@ -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!({
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,31 +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" => "wss",
|
||||
"http" | "ws" => "ws",
|
||||
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" {
|
||||
let literal_ip = match url.host() {
|
||||
Some(url::Host::Ipv4(address)) => Some(std::net::IpAddr::V4(address)),
|
||||
Some(url::Host::Ipv6(address)) => Some(std::net::IpAddr::V6(address)),
|
||||
_ => None,
|
||||
};
|
||||
if literal_ip.is_some_and(|address| {
|
||||
aether_http::is_private_or_reserved_ip(address) && !address.is_loopback()
|
||||
}) {
|
||||
return Err(invalid_code);
|
||||
}
|
||||
}
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
@@ -229,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);
|
||||
@@ -255,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)
|
||||
}
|
||||
@@ -684,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;
|
||||
@@ -876,7 +826,9 @@ mod tests {
|
||||
"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",
|
||||
@@ -888,6 +840,14 @@ mod tests {
|
||||
}
|
||||
for rejected in [
|
||||
"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",
|
||||
@@ -903,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 {
|
||||
@@ -937,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();
|
||||
|
||||
@@ -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,98 +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");
|
||||
}
|
||||
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,
|
||||
@@ -384,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(
|
||||
@@ -407,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(
|
||||
@@ -419,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);
|
||||
}
|
||||
@@ -495,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},
|
||||
@@ -567,74 +485,59 @@ 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",
|
||||
"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_accepts_public_http_and_https_addresses() {
|
||||
for allow_private_targets in [false, true] {
|
||||
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),
|
||||
] {
|
||||
let target = resolve_test_connection_target(raw_url, allow_private_targets)
|
||||
.await
|
||||
.expect("public HTTP(S) provider target should resolve");
|
||||
assert_eq!(target.url.as_str(), raw_url);
|
||||
assert_eq!(target.host, "8.8.8.8");
|
||||
assert_eq!(target.addresses.len(), 1);
|
||||
assert_eq!(target.addresses[0].ip().to_string(), "8.8.8.8");
|
||||
assert_eq!(target.addresses[0].port(), expected_port);
|
||||
}
|
||||
#[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_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://10.0.0.1/v1/chat", true)
|
||||
.await
|
||||
.is_err(),
|
||||
"test mode must not make private non-loopback HTTP 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"
|
||||
);
|
||||
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);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_connection_target_rejects_url_credentials_and_fragments() {
|
||||
#[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",
|
||||
@@ -643,9 +546,7 @@ mod tests {
|
||||
"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))]
|
||||
|
||||
+229
-48
@@ -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};
|
||||
|
||||
@@ -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),
|
||||
@@ -1486,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,
|
||||
@@ -1567,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.
|
||||
@@ -1848,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(())
|
||||
}
|
||||
@@ -1889,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
|
||||
@@ -1909,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:
|
||||
@@ -1938,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> {
|
||||
@@ -2065,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>> {
|
||||
@@ -2133,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,
|
||||
@@ -2323,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()
|
||||
@@ -2467,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.
|
||||
@@ -2486,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,
|
||||
@@ -2493,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(())
|
||||
}
|
||||
|
||||
@@ -3437,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,
|
||||
@@ -3478,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,
|
||||
@@ -3492,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,
|
||||
@@ -3532,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,
|
||||
@@ -3925,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();
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -3083,7 +3083,7 @@ mod tests {
|
||||
.await
|
||||
.expect("stale LKG read must not wait for retention lock");
|
||||
assert_eq!(stale.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
|
||||
assert_eq!(stale.stale_targets(), &[target.clone()]);
|
||||
assert_eq!(stale.stale_targets(), std::slice::from_ref(&target));
|
||||
assert_eq!(runtime.execution_count(), 1);
|
||||
|
||||
assert!(runtime
|
||||
@@ -3111,7 +3111,7 @@ mod tests {
|
||||
let load = load_one(&runtime, &client_version).await;
|
||||
|
||||
assert_eq!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
|
||||
assert_eq!(load.stale_targets(), &[target.clone()]);
|
||||
assert_eq!(load.stale_targets(), std::slice::from_ref(&target));
|
||||
assert!(runtime
|
||||
.state
|
||||
.kv_get(&catalog_lkg_key(&target, client_version.as_str()))
|
||||
@@ -3142,7 +3142,7 @@ mod tests {
|
||||
|
||||
let load = load_one(&runtime, &client_version).await;
|
||||
assert_eq!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
|
||||
assert_eq!(load.stale_targets(), &[target.clone()]);
|
||||
assert_eq!(load.stale_targets(), std::slice::from_ref(&target));
|
||||
assert_eq!(runtime.execution_count(), 1);
|
||||
}
|
||||
|
||||
|
||||
@@ -300,6 +300,34 @@ pub(crate) fn classify_local_failover(
|
||||
policy: &LocalFailoverPolicy,
|
||||
input: LocalFailoverInput<'_>,
|
||||
) -> LocalFailoverClassification {
|
||||
if input.status_code >= 400
|
||||
&& policy.routing_rules.error_stop_patterns.iter().any(|rule| {
|
||||
failover_pattern_matches(
|
||||
&rule.pattern,
|
||||
&rule.status_codes,
|
||||
input.response_text,
|
||||
input.status_code,
|
||||
)
|
||||
})
|
||||
{
|
||||
return LocalFailoverClassification::StopErrorPattern;
|
||||
}
|
||||
if input.status_code == 200
|
||||
&& policy
|
||||
.routing_rules
|
||||
.success_failover_patterns
|
||||
.iter()
|
||||
.any(|rule| {
|
||||
failover_pattern_matches(
|
||||
&rule.pattern,
|
||||
&rule.status_codes,
|
||||
input.response_text,
|
||||
input.status_code,
|
||||
)
|
||||
})
|
||||
{
|
||||
return LocalFailoverClassification::RetrySuccessPattern;
|
||||
}
|
||||
if policy.stop_status_codes.contains(&input.status_code) {
|
||||
return LocalFailoverClassification::StopStatusCode;
|
||||
}
|
||||
@@ -487,13 +515,27 @@ fn local_failover_regex_rule_matches(
|
||||
response_text: Option<&str>,
|
||||
status_code: u16,
|
||||
) -> bool {
|
||||
if !rule.status_codes.is_empty() && !rule.status_codes.contains(&status_code) {
|
||||
failover_pattern_matches(
|
||||
&rule.pattern,
|
||||
&rule.status_codes,
|
||||
response_text,
|
||||
status_code,
|
||||
)
|
||||
}
|
||||
|
||||
fn failover_pattern_matches(
|
||||
pattern: &str,
|
||||
status_codes: &std::collections::BTreeSet<u16>,
|
||||
response_text: Option<&str>,
|
||||
status_code: u16,
|
||||
) -> bool {
|
||||
if !status_codes.is_empty() && !status_codes.contains(&status_code) {
|
||||
return false;
|
||||
}
|
||||
|
||||
let pattern = rule.pattern.trim();
|
||||
let pattern = pattern.trim();
|
||||
if pattern.is_empty() {
|
||||
return !rule.status_codes.is_empty();
|
||||
return !status_codes.is_empty();
|
||||
}
|
||||
|
||||
let Some(response_text) = response_text else {
|
||||
@@ -507,6 +549,71 @@ fn local_failover_regex_rule_matches(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[test]
|
||||
fn routing_rules_precede_provider_rules_and_keep_provider_fallback() {
|
||||
let policy = super::LocalFailoverPolicy {
|
||||
routing_rules: aether_routing_core::RoutingFailoverRules {
|
||||
success_failover_patterns: vec![aether_routing_core::RoutingFailoverRule {
|
||||
pattern: "(?i)capacity.*exhausted".to_string(),
|
||||
..Default::default()
|
||||
}],
|
||||
error_stop_patterns: vec![aether_routing_core::RoutingFailoverRule {
|
||||
pattern: "invalid.*parameter".to_string(),
|
||||
status_codes: [400].into_iter().collect(),
|
||||
}],
|
||||
},
|
||||
stop_status_codes: [200, 403].into_iter().collect(),
|
||||
continue_status_codes: [400].into_iter().collect(),
|
||||
..Default::default()
|
||||
};
|
||||
for (status, body, expected) in [
|
||||
(
|
||||
200,
|
||||
"CAPACITY exhausted",
|
||||
super::LocalFailoverClassification::RetrySuccessPattern,
|
||||
),
|
||||
(
|
||||
400,
|
||||
"invalid request parameter",
|
||||
super::LocalFailoverClassification::StopErrorPattern,
|
||||
),
|
||||
(
|
||||
400,
|
||||
"capacity exhausted",
|
||||
super::LocalFailoverClassification::RetryStatusCode,
|
||||
),
|
||||
(
|
||||
403,
|
||||
"permission denied",
|
||||
super::LocalFailoverClassification::StopStatusCode,
|
||||
),
|
||||
(
|
||||
429,
|
||||
"rate limited",
|
||||
super::LocalFailoverClassification::RetryUpstreamFailure,
|
||||
),
|
||||
] {
|
||||
assert_eq!(
|
||||
super::classify_local_failover(
|
||||
&policy,
|
||||
super::LocalFailoverInput::new(status, Some(body))
|
||||
),
|
||||
expected
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_transport_stop_rule_is_respected() {
|
||||
let policy = super::LocalFailoverPolicy {
|
||||
stop_on_transport_errors: true,
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
super::classify_local_transport_error(&policy),
|
||||
super::LocalTransportFailoverClassification::StopTransportError
|
||||
);
|
||||
}
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use super::{
|
||||
|
||||
@@ -973,7 +973,7 @@ async fn record_adaptive_rate_limit_effect(
|
||||
let _effect_guard = effect_lock.lock().await;
|
||||
let observed_at_unix_secs = current_unix_secs();
|
||||
let current_rpm = state
|
||||
.read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
|
||||
.read_recent_runtime_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
|
||||
.await
|
||||
.ok()
|
||||
.map(|recent_candidates| {
|
||||
@@ -1123,7 +1123,7 @@ async fn record_adaptive_success_effect(
|
||||
return;
|
||||
}
|
||||
let Some(recent_candidates) = state
|
||||
.read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
|
||||
.read_recent_runtime_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
|
||||
.await
|
||||
.ok()
|
||||
else {
|
||||
|
||||
@@ -4,7 +4,7 @@ use aether_contracts::ExecutionPlan;
|
||||
use serde_json::{json, Value};
|
||||
use tracing::debug;
|
||||
|
||||
use aether_routing_core::RoutingExecutionPolicy;
|
||||
use aether_routing_core::{RoutingExecutionPolicy, RoutingFailoverRules};
|
||||
|
||||
use crate::provider_transport::GatewayProviderTransportSnapshot;
|
||||
use crate::AppState;
|
||||
@@ -14,6 +14,7 @@ pub(crate) const ROUTING_EXECUTION_POLICY_REPORT_FIELD: &str = "routing_executio
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct LocalFailoverPolicy {
|
||||
pub(crate) routing_rules: RoutingFailoverRules,
|
||||
pub(crate) max_retries: Option<u64>,
|
||||
pub(crate) max_transfer_count: u64,
|
||||
pub(crate) max_transfer_timeout_seconds: u64,
|
||||
@@ -29,6 +30,7 @@ pub(crate) struct LocalFailoverPolicy {
|
||||
impl Default for LocalFailoverPolicy {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
routing_rules: RoutingFailoverRules::default(),
|
||||
max_retries: None,
|
||||
max_transfer_count: 0,
|
||||
max_transfer_timeout_seconds: 0,
|
||||
@@ -61,8 +63,10 @@ pub(crate) async fn resolve_local_failover_policy(
|
||||
Ok(Some(transport)) => local_failover_policy_from_transport(&transport),
|
||||
Ok(None) | Err(_) => LocalFailoverPolicy::default(),
|
||||
};
|
||||
let cyber_continue_failover = routing_execution_policy_from_report_context(report_context)
|
||||
.is_some_and(|policy| policy.cyber_continue_failover);
|
||||
let routing_policy =
|
||||
routing_execution_policy_from_report_context(report_context).unwrap_or_default();
|
||||
let cyber_continue_failover = routing_policy.cyber_continue_failover;
|
||||
policy.routing_rules = routing_policy.failover_rules;
|
||||
policy.stop_cyber_policy_errors = !cyber_continue_failover;
|
||||
debug!(
|
||||
event_name = "local_failover_policy_loaded",
|
||||
@@ -80,6 +84,8 @@ pub(crate) async fn resolve_local_failover_policy(
|
||||
stop_on_transport_errors = policy.stop_on_transport_errors,
|
||||
success_failover_pattern_count = policy.success_failover_patterns.len(),
|
||||
error_stop_pattern_count = policy.error_stop_patterns.len(),
|
||||
global_success_pattern_count = policy.routing_rules.success_failover_patterns.len(),
|
||||
global_stop_pattern_count = policy.routing_rules.error_stop_patterns.len(),
|
||||
cyber_continue_failover,
|
||||
"gateway loaded local failover policy from transport snapshot"
|
||||
);
|
||||
@@ -122,6 +128,7 @@ pub(crate) fn local_failover_policy_from_transport(
|
||||
});
|
||||
|
||||
LocalFailoverPolicy {
|
||||
routing_rules: RoutingFailoverRules::default(),
|
||||
max_retries,
|
||||
max_transfer_count: provider_config
|
||||
.and_then(|value| value.get("max_transfer_count"))
|
||||
@@ -184,6 +191,10 @@ pub(crate) fn local_failover_policy_from_report_context(
|
||||
.as_object()?;
|
||||
|
||||
Some(LocalFailoverPolicy {
|
||||
routing_rules: object
|
||||
.get("routing_rules")
|
||||
.and_then(|value| serde_json::from_value(value.clone()).ok())
|
||||
.unwrap_or_default(),
|
||||
max_retries: object.get("max_retries").and_then(parse_u64_value),
|
||||
max_transfer_count: object
|
||||
.get("max_transfer_count")
|
||||
@@ -267,6 +278,7 @@ fn parse_status_code_list(value: &Value) -> BTreeSet<u16> {
|
||||
|
||||
fn local_failover_policy_to_value(policy: &LocalFailoverPolicy) -> Value {
|
||||
json!({
|
||||
"routing_rules": policy.routing_rules,
|
||||
"max_retries": policy.max_retries,
|
||||
"max_transfer_count": policy.max_transfer_count,
|
||||
"max_transfer_timeout_seconds": policy.max_transfer_timeout_seconds,
|
||||
@@ -525,6 +537,7 @@ mod tests {
|
||||
assert_eq!(
|
||||
local_failover_policy_from_report_context(Some(&report_context)),
|
||||
Some(LocalFailoverPolicy {
|
||||
routing_rules: Default::default(),
|
||||
max_retries: Some(2),
|
||||
max_transfer_count: 10,
|
||||
max_transfer_timeout_seconds: 60,
|
||||
|
||||
@@ -643,6 +643,24 @@ impl RequestCandidateQueueRuntime {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn pending_writes(&self) -> usize {
|
||||
[
|
||||
self.metrics.pending_current.load(Ordering::Acquire),
|
||||
self.metrics
|
||||
.priority_pending_current
|
||||
.load(Ordering::Acquire),
|
||||
self.metrics.active_pending_current.load(Ordering::Acquire),
|
||||
self.metrics
|
||||
.terminal_pending_current
|
||||
.load(Ordering::Acquire),
|
||||
self.metrics
|
||||
.terminal_barrier_pending
|
||||
.load(Ordering::Acquire),
|
||||
]
|
||||
.into_iter()
|
||||
.fold(0_usize, usize::saturating_add)
|
||||
}
|
||||
|
||||
pub(crate) fn metric_samples(&self) -> Vec<MetricSample> {
|
||||
vec![
|
||||
MetricSample::new(
|
||||
@@ -3764,18 +3782,19 @@ mod tests {
|
||||
.await;
|
||||
|
||||
assert_eq!(normal_batch.len(), 1);
|
||||
let retry_states = metrics
|
||||
.retry_states
|
||||
.lock()
|
||||
.unwrap_or_else(|poisoned| poisoned.into_inner());
|
||||
assert_eq!(
|
||||
retry_states
|
||||
.get(&(0, RequestCandidateQueueLane::Normal))
|
||||
.map(|state| state.attempt),
|
||||
Some(1)
|
||||
);
|
||||
assert!(!retry_states.contains_key(&(0, RequestCandidateQueueLane::Active)));
|
||||
drop(retry_states);
|
||||
{
|
||||
let retry_states = metrics
|
||||
.retry_states
|
||||
.lock()
|
||||
.unwrap_or_else(|poisoned| poisoned.into_inner());
|
||||
assert_eq!(
|
||||
retry_states
|
||||
.get(&(0, RequestCandidateQueueLane::Normal))
|
||||
.map(|state| state.attempt),
|
||||
Some(1)
|
||||
);
|
||||
assert!(!retry_states.contains_key(&(0, RequestCandidateQueueLane::Active)));
|
||||
}
|
||||
assert!(request_candidate_retry_is_ready(
|
||||
&metrics,
|
||||
0,
|
||||
|
||||
@@ -0,0 +1,449 @@
|
||||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use aether_routing_core::RoutingExecutionPolicy;
|
||||
use aether_usage_runtime::{UsageProducerGuard, UsageRuntime};
|
||||
use axum::body::{Body, Bytes, HttpBody};
|
||||
use http::Response;
|
||||
use http_body::{Frame, SizeHint};
|
||||
use http_body_util::BodyExt;
|
||||
|
||||
use crate::request_diagnostics::{scope_request_diagnostics_with, RequestDiagnostics};
|
||||
use crate::GatewayError;
|
||||
|
||||
tokio::task_local! {
|
||||
static CANCEL_ON_CLIENT_DISCONNECT: Arc<AtomicBool>;
|
||||
}
|
||||
|
||||
pub(crate) fn configure_client_disconnect(policy: RoutingExecutionPolicy) {
|
||||
let _ = CANCEL_ON_CLIENT_DISCONNECT.try_with(|cancel| {
|
||||
cancel.store(policy.cancel_on_client_disconnect, Ordering::Release);
|
||||
});
|
||||
}
|
||||
|
||||
pub(crate) fn cancel_on_client_disconnect() -> bool {
|
||||
CANCEL_ON_CLIENT_DISCONNECT
|
||||
.try_with(|cancel| cancel.load(Ordering::Acquire))
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) async fn run_request<F>(future: F) -> Result<Response<Body>, GatewayError>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
run_tracked_request(future, None).await
|
||||
}
|
||||
|
||||
pub(crate) async fn run_request_with_usage<F>(
|
||||
usage: Arc<UsageRuntime>,
|
||||
future: F,
|
||||
) -> Result<Response<Body>, GatewayError>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
run_tracked_request(future, Some(Arc::new(usage.track_producer()))).await
|
||||
}
|
||||
|
||||
async fn run_tracked_request<F>(
|
||||
future: F,
|
||||
producer: Option<Arc<UsageProducerGuard>>,
|
||||
) -> Result<Response<Body>, GatewayError>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
let cancel = Arc::new(AtomicBool::new(true));
|
||||
let diagnostics = Arc::new(RequestDiagnostics::default());
|
||||
let cancel_for_response = Arc::clone(&cancel);
|
||||
let producer_for_request = producer.clone();
|
||||
let future = CANCEL_ON_CLIENT_DISCONNECT.scope(
|
||||
Arc::clone(&cancel),
|
||||
scope_request_diagnostics_with(Some(Arc::clone(&diagnostics)), async move {
|
||||
let response = future.await?;
|
||||
let complete_on_disconnect = !cancel_for_response.load(Ordering::Acquire);
|
||||
if !complete_on_disconnect && producer.is_none() {
|
||||
return Ok(response);
|
||||
}
|
||||
Ok(response.map(|body| {
|
||||
Body::new(CompleteOnDisconnectBody {
|
||||
body: Some(body),
|
||||
diagnostics,
|
||||
complete_on_disconnect,
|
||||
producer,
|
||||
})
|
||||
}))
|
||||
}),
|
||||
);
|
||||
CompleteOnDisconnectRequest {
|
||||
future: Some(Box::pin(future)),
|
||||
cancel,
|
||||
producer: producer_for_request,
|
||||
}
|
||||
.await
|
||||
}
|
||||
|
||||
struct CompleteOnDisconnectRequest<F>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
future: Option<Pin<Box<F>>>,
|
||||
cancel: Arc<AtomicBool>,
|
||||
producer: Option<Arc<UsageProducerGuard>>,
|
||||
}
|
||||
|
||||
impl<F> Future for CompleteOnDisconnectRequest<F>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
type Output = Result<Response<Body>, GatewayError>;
|
||||
|
||||
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
|
||||
let result = self
|
||||
.future
|
||||
.as_mut()
|
||||
.expect("request future")
|
||||
.as_mut()
|
||||
.poll(context);
|
||||
if result.is_ready() {
|
||||
self.future.take();
|
||||
}
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
impl<F> Drop for CompleteOnDisconnectRequest<F>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
fn drop(&mut self) {
|
||||
if self.cancel.load(Ordering::Acquire) {
|
||||
return;
|
||||
}
|
||||
if let (Some(future), Ok(runtime)) =
|
||||
(self.future.take(), tokio::runtime::Handle::try_current())
|
||||
{
|
||||
let producer = self.producer.take();
|
||||
runtime.spawn(async move {
|
||||
let _producer = producer;
|
||||
if let Ok(response) = future.await {
|
||||
drain_body(response.into_body()).await;
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct CompleteOnDisconnectBody {
|
||||
body: Option<Body>,
|
||||
diagnostics: Arc<RequestDiagnostics>,
|
||||
complete_on_disconnect: bool,
|
||||
// Drop the body first so its terminal handoff registers before this guard ends.
|
||||
producer: Option<Arc<UsageProducerGuard>>,
|
||||
}
|
||||
|
||||
impl HttpBody for CompleteOnDisconnectBody {
|
||||
type Data = Bytes;
|
||||
type Error = axum::Error;
|
||||
|
||||
fn poll_frame(
|
||||
mut self: Pin<&mut Self>,
|
||||
context: &mut Context<'_>,
|
||||
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
|
||||
let Some(body) = self.body.as_mut() else {
|
||||
return Poll::Ready(None);
|
||||
};
|
||||
let result = Pin::new(body).poll_frame(context);
|
||||
if matches!(result, Poll::Ready(None | Some(Err(_)))) {
|
||||
self.body.take();
|
||||
self.producer.take();
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn is_end_stream(&self) -> bool {
|
||||
self.body.as_ref().is_none_or(HttpBody::is_end_stream)
|
||||
}
|
||||
|
||||
fn size_hint(&self) -> SizeHint {
|
||||
self.body
|
||||
.as_ref()
|
||||
.map(HttpBody::size_hint)
|
||||
.unwrap_or_else(|| SizeHint::with_exact(0))
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for CompleteOnDisconnectBody {
|
||||
fn drop(&mut self) {
|
||||
if !self.complete_on_disconnect {
|
||||
return;
|
||||
}
|
||||
let Some(body) = self.body.take().filter(|body| !body.is_end_stream()) else {
|
||||
return;
|
||||
};
|
||||
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
|
||||
let producer = self.producer.take();
|
||||
runtime.spawn(scope_request_diagnostics_with(
|
||||
Some(Arc::clone(&self.diagnostics)),
|
||||
async move {
|
||||
let _producer = producer;
|
||||
drain_body(body).await;
|
||||
},
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn drain_body(mut body: Body) {
|
||||
while let Some(frame) = body.frame().await {
|
||||
if frame.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::io;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::stream;
|
||||
use http::HeaderMap;
|
||||
use http_body_util::StreamBody;
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn disconnected_request_finishes_and_keeps_admission_and_diagnostics() {
|
||||
let gate = aether_runtime::ConcurrencyGate::new("disconnect_request", 1);
|
||||
let permit = gate.try_acquire().unwrap();
|
||||
let (started_tx, started_rx) = oneshot::channel();
|
||||
let (release_tx, release_rx) = oneshot::channel();
|
||||
let (finished_tx, finished_rx) = oneshot::channel();
|
||||
let request = tokio::spawn(run_request(async move {
|
||||
let _permit = permit;
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
started_tx.send(()).unwrap();
|
||||
release_rx.await.unwrap();
|
||||
assert!(crate::request_diagnostics::current_request_diagnostics().is_some());
|
||||
finished_tx.send(()).unwrap();
|
||||
Ok(Response::new(Body::empty()))
|
||||
}));
|
||||
started_rx.await.unwrap();
|
||||
request.abort();
|
||||
assert!(request.await.unwrap_err().is_cancelled());
|
||||
assert_eq!(gate.snapshot().in_flight, 1);
|
||||
release_tx.send(()).unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(1), finished_rx)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(gate.snapshot().in_flight, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enabled_cancellation_and_unresolved_requests_drop_immediately() {
|
||||
for resolve_policy in [false, true] {
|
||||
let (started_tx, started_rx) = oneshot::channel();
|
||||
let (release_tx, release_rx) = oneshot::channel::<()>();
|
||||
let request = tokio::spawn(run_request(async move {
|
||||
if resolve_policy {
|
||||
configure_client_disconnect(RoutingExecutionPolicy {
|
||||
cancel_on_client_disconnect: true,
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
started_tx.send(()).unwrap();
|
||||
release_rx.await.unwrap();
|
||||
Ok(Response::new(Body::empty()))
|
||||
}));
|
||||
started_rx.await.unwrap();
|
||||
request.abort();
|
||||
assert!(request.await.unwrap_err().is_cancelled());
|
||||
assert!(release_tx.send(()).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn disconnected_body_drains_without_buffering_and_holds_admission() {
|
||||
for consume_first_chunk in [false, true] {
|
||||
let gate = aether_runtime::ConcurrencyGate::new("disconnect_body", 1);
|
||||
let permit = gate.try_acquire().unwrap();
|
||||
let (sender, receiver) = mpsc::channel(1);
|
||||
let (finished_tx, finished_rx) = oneshot::channel();
|
||||
let response = run_request(async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
let body = Body::from_stream(stream::unfold(
|
||||
(receiver, finished_tx, permit),
|
||||
|(mut receiver, finished_tx, permit)| async move {
|
||||
match receiver.recv().await {
|
||||
Some(bytes) => {
|
||||
Some((Ok::<_, io::Error>(bytes), (receiver, finished_tx, permit)))
|
||||
}
|
||||
None => {
|
||||
assert!(crate::request_diagnostics::current_request_diagnostics()
|
||||
.is_some());
|
||||
finished_tx.send(()).unwrap();
|
||||
None
|
||||
}
|
||||
}
|
||||
},
|
||||
));
|
||||
Ok(Response::new(body))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut body = response.into_body();
|
||||
if consume_first_chunk {
|
||||
sender.send(Bytes::from_static(b"first")).await.unwrap();
|
||||
assert_eq!(
|
||||
body.frame().await.unwrap().unwrap().into_data().unwrap(),
|
||||
"first"
|
||||
);
|
||||
}
|
||||
drop(body);
|
||||
assert_eq!(gate.snapshot().in_flight, 1);
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
for _ in 0..100 {
|
||||
sender.send(Bytes::from_static(b"remaining")).await.unwrap();
|
||||
}
|
||||
drop(sender);
|
||||
finished_rx.await.unwrap();
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(gate.snapshot().in_flight, 0);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enabled_cancellation_drops_stream_receiver() {
|
||||
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
|
||||
let response = run_request(async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy {
|
||||
cancel_on_client_disconnect: true,
|
||||
..Default::default()
|
||||
});
|
||||
Ok(Response::new(Body::from_stream(stream::unfold(
|
||||
receiver,
|
||||
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
|
||||
))))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
drop(response);
|
||||
assert!(sender.is_closed());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn usage_shutdown_waits_for_a_disconnected_request_before_headers() {
|
||||
let usage = Arc::new(UsageRuntime::disabled());
|
||||
let (started_tx, started_rx) = oneshot::channel();
|
||||
let (release_tx, release_rx) = oneshot::channel::<()>();
|
||||
let request = tokio::spawn(run_request_with_usage(usage.clone(), async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
started_tx.send(()).unwrap();
|
||||
release_rx.await.unwrap();
|
||||
Ok(Response::new(Body::empty()))
|
||||
}));
|
||||
started_rx.await.unwrap();
|
||||
request.abort();
|
||||
assert!(request.await.unwrap_err().is_cancelled());
|
||||
assert!(usage.shutdown(Duration::from_millis(30)).await.is_err());
|
||||
assert_eq!(usage.metrics_snapshot().producers_in_flight, 1);
|
||||
release_tx.send(()).unwrap();
|
||||
usage.shutdown(Duration::from_secs(1)).await.unwrap();
|
||||
assert_eq!(usage.metrics_snapshot().producers_in_flight, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn usage_shutdown_waits_for_disconnected_body_drain() {
|
||||
let usage = Arc::new(UsageRuntime::disabled());
|
||||
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
|
||||
let response = run_request_with_usage(usage.clone(), async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
Ok(Response::new(Body::from_stream(stream::unfold(
|
||||
receiver,
|
||||
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
|
||||
))))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
drop(response);
|
||||
assert!(usage.shutdown(Duration::from_millis(30)).await.is_err());
|
||||
sender.send(Ok(Bytes::from_static(b"last"))).await.unwrap();
|
||||
drop(sender);
|
||||
usage.shutdown(Duration::from_secs(1)).await.unwrap();
|
||||
assert_eq!(usage.metrics_snapshot().producers_in_flight, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tracked_bodies_release_shutdown_on_cancellation_or_eof() {
|
||||
for cancel_on_client_disconnect in [false, true] {
|
||||
let usage = Arc::new(UsageRuntime::disabled());
|
||||
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
|
||||
let response = run_request_with_usage(usage.clone(), async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy {
|
||||
cancel_on_client_disconnect,
|
||||
..Default::default()
|
||||
});
|
||||
Ok(Response::new(Body::from_stream(stream::unfold(
|
||||
receiver,
|
||||
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
|
||||
))))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut body = response.into_body();
|
||||
if cancel_on_client_disconnect {
|
||||
drop(body);
|
||||
assert!(sender.is_closed());
|
||||
} else {
|
||||
drop(sender);
|
||||
assert!(body.frame().await.is_none());
|
||||
assert_eq!(usage.metrics_snapshot().producers_in_flight, 0);
|
||||
}
|
||||
usage.shutdown(Duration::from_secs(1)).await.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn connected_response_preserves_headers_size_hint_and_trailers() {
|
||||
let response = run_request(async {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
Ok(Response::builder()
|
||||
.status(201)
|
||||
.header("x-test", "unchanged")
|
||||
.body(Body::from("hello"))
|
||||
.unwrap())
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), 201);
|
||||
assert_eq!(response.headers()["x-test"], "unchanged");
|
||||
assert_eq!(response.body().size_hint().exact(), Some(5));
|
||||
assert_eq!(
|
||||
response.into_body().collect().await.unwrap().to_bytes(),
|
||||
"hello"
|
||||
);
|
||||
|
||||
let mut trailers = HeaderMap::new();
|
||||
trailers.insert("x-finished", "yes".parse().unwrap());
|
||||
let response = run_request(async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
let frames = stream::iter([
|
||||
Ok::<_, io::Error>(Frame::data(Bytes::from_static(b"hello"))),
|
||||
Ok(Frame::trailers(trailers)),
|
||||
]);
|
||||
Ok(Response::new(Body::new(StreamBody::new(frames))))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let collected = response.into_body().collect().await.unwrap();
|
||||
assert_eq!(collected.trailers().unwrap()["x-finished"], "yes");
|
||||
assert_eq!(collected.to_bytes(), "hello");
|
||||
}
|
||||
}
|
||||
@@ -18,6 +18,7 @@ use tower::{Service as _, ServiceExt};
|
||||
use tower_http::services::{ServeDir, ServeFile};
|
||||
use tracing::warn;
|
||||
|
||||
use aether_gateway_frontdoor::{http_connection_limit, HttpConnectionBudget};
|
||||
use aether_runtime::{prometheus_response, ConcurrencyError};
|
||||
use aether_runtime_state::RuntimeSemaphoreError;
|
||||
|
||||
@@ -244,9 +245,23 @@ pub(crate) enum RequestAdmissionError {
|
||||
pub async fn serve_tcp(bind: &str) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let listener = tokio::net::TcpListener::bind(bind).await?;
|
||||
let router = build_router()?;
|
||||
let configured_connection_limit = std::env::var("AETHER_GATEWAY_MAX_HTTP_CONNECTIONS")
|
||||
.ok()
|
||||
.and_then(|value| value.trim().parse::<usize>().ok());
|
||||
// This compatibility entry point has no configured request capacities or FD probe.
|
||||
let connection_budget = Arc::new(HttpConnectionBudget::new(http_connection_limit(
|
||||
configured_connection_limit,
|
||||
2048,
|
||||
2048,
|
||||
None,
|
||||
)));
|
||||
let mut make_service = router.into_make_service_with_connect_info::<std::net::SocketAddr>();
|
||||
loop {
|
||||
let (io, remote_addr) = listener.accept().await?;
|
||||
let (io, remote_addr) = connection_budget.accept(&listener).await;
|
||||
let Ok(io) = connection_budget.try_admit(io) else {
|
||||
tokio::task::yield_now().await;
|
||||
continue;
|
||||
};
|
||||
let tower_service = make_service
|
||||
.call(remote_addr)
|
||||
.await
|
||||
|
||||
@@ -57,7 +57,7 @@ pub(crate) fn resolve_gateway_routing_policy(
|
||||
|
||||
let config = serde_json::from_value::<RoutingGroupConfig>(input.group_config_json.clone())
|
||||
.map_err(|_| invalid_routing_group_config())?;
|
||||
resolve_routing_policy(
|
||||
let policy = resolve_routing_policy(
|
||||
&config,
|
||||
RoutingPolicyInput {
|
||||
group_id: input.group_id,
|
||||
@@ -73,7 +73,9 @@ pub(crate) fn resolve_gateway_routing_policy(
|
||||
phase: input.phase,
|
||||
},
|
||||
)
|
||||
.map_err(routing_policy_error)
|
||||
.map_err(routing_policy_error)?;
|
||||
crate::request_lifecycle::configure_client_disconnect(policy.execution_policy.clone());
|
||||
Ok(policy)
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_gateway_static_default_routing_policy(
|
||||
@@ -82,6 +84,7 @@ pub(crate) fn resolve_gateway_static_default_routing_policy(
|
||||
let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else {
|
||||
return Ok(None);
|
||||
};
|
||||
crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy.clone());
|
||||
|
||||
Ok(Some(ResolvedRoutingPolicy {
|
||||
group_id: input.group_id.map(str::to_string),
|
||||
@@ -142,28 +145,11 @@ fn static_default_policy_fields(
|
||||
.ok_or_else(invalid_routing_group_config)?,
|
||||
None => DEFAULT_STICKY_KEY_ATTEMPTS,
|
||||
};
|
||||
let enable_cf_heartbeat = routing_bool_field(
|
||||
default_policy.get("enable_cf_heartbeat"),
|
||||
"enable_cf_heartbeat",
|
||||
)?;
|
||||
// Older strategies stored separate image/text heartbeat flags. Treat
|
||||
// either legacy flag as enabling the unified CF heartbeat setting while
|
||||
// allowing newly saved strategies to use only the canonical key.
|
||||
let legacy_image_heartbeat = routing_bool_field(
|
||||
default_policy.get("enable_openai_image_sync_heartbeat"),
|
||||
"enable_openai_image_sync_heartbeat",
|
||||
)?;
|
||||
let legacy_text_heartbeat = routing_bool_field(
|
||||
default_policy.get("enable_standard_text_sync_heartbeat"),
|
||||
"enable_standard_text_sync_heartbeat",
|
||||
)?;
|
||||
let execution_policy = aether_routing_core::RoutingExecutionPolicy {
|
||||
enable_cf_heartbeat: enable_cf_heartbeat || legacy_image_heartbeat || legacy_text_heartbeat,
|
||||
cyber_continue_failover: routing_bool_field(
|
||||
default_policy.get("cyber_continue_failover"),
|
||||
"cyber_continue_failover",
|
||||
)?,
|
||||
};
|
||||
let execution_policy: aether_routing_core::RoutingExecutionPolicy =
|
||||
serde_json::from_value(Value::Object(default_policy.clone()))
|
||||
.map_err(|_| invalid_routing_group_config())?;
|
||||
aether_routing_core::validate_routing_failover_rules(&execution_policy.failover_rules)
|
||||
.map_err(|_| invalid_routing_group_config())?;
|
||||
|
||||
Ok(Some(RoutingDefaultPolicy {
|
||||
priority_mode,
|
||||
@@ -174,13 +160,6 @@ fn static_default_policy_fields(
|
||||
}))
|
||||
}
|
||||
|
||||
fn routing_bool_field(value: Option<&Value>, _field: &str) -> Result<bool, GatewayError> {
|
||||
match value {
|
||||
Some(value) => value.as_bool().ok_or_else(invalid_routing_group_config),
|
||||
None => Ok(false),
|
||||
}
|
||||
}
|
||||
|
||||
fn routing_array_field_is_missing_or_empty(
|
||||
object: &serde_json::Map<String, Value>,
|
||||
key: &str,
|
||||
@@ -238,7 +217,14 @@ mod tests {
|
||||
"default_policy": {
|
||||
"priority_mode": "global_key",
|
||||
"scheduling_mode": "load_balance",
|
||||
"keep_priority_on_conversion": true
|
||||
"keep_priority_on_conversion": true,
|
||||
"cancel_on_client_disconnect": true,
|
||||
"max_transfer_count": 3,
|
||||
"max_transfer_timeout_seconds": 90,
|
||||
"failover_rules": {
|
||||
"success_failover_patterns": [{"pattern": "(?i)capacity.*exhausted"}],
|
||||
"error_stop_patterns": [{"status_codes": [400]}]
|
||||
}
|
||||
},
|
||||
"allowed_models": ["legacy-model"],
|
||||
"model_policies": [],
|
||||
@@ -274,6 +260,19 @@ mod tests {
|
||||
.expect("full policy should resolve");
|
||||
|
||||
assert_eq!(static_policy, full_policy);
|
||||
assert_eq!(static_policy.execution_policy.max_transfer_count, 3);
|
||||
assert_eq!(
|
||||
static_policy.execution_policy.max_transfer_timeout_seconds,
|
||||
90
|
||||
);
|
||||
assert_eq!(
|
||||
static_policy
|
||||
.execution_policy
|
||||
.failover_rules
|
||||
.error_stop_patterns
|
||||
.len(),
|
||||
1
|
||||
);
|
||||
assert_eq!(
|
||||
static_policy.priority_mode,
|
||||
RoutingSetPriorityMode::GlobalKey
|
||||
|
||||
@@ -32,6 +32,9 @@ use regex::Regex;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
pub(crate) use self::runtime::{
|
||||
select_with_auth_concurrency_wait, wait_for_auth_api_key_concurrency_retry,
|
||||
};
|
||||
pub(crate) use self::selection::{
|
||||
is_auth_api_key_concurrency_limit_skip_reason, SchedulerSkippedCandidate,
|
||||
API_KEY_CONCURRENCY_LIMIT_SKIP_REASON, AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON,
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::future::Future;
|
||||
|
||||
use aether_admin::provider::{
|
||||
pool as admin_provider_pool_pure, status as admin_provider_status_pure,
|
||||
@@ -10,7 +11,8 @@ use aether_scheduler_core::{
|
||||
candidate_is_selectable_with_runtime_state, candidate_runtime_skip_reason_with_state,
|
||||
effective_provider_key_rpm_limit, CandidateRuntimeSelectabilityInput,
|
||||
};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
use tokio::time::Instant;
|
||||
|
||||
use crate::data::auth::GatewayAuthApiKeySnapshot;
|
||||
use crate::GatewayError;
|
||||
@@ -109,25 +111,92 @@ pub(super) fn auth_snapshot_concurrency_limit_reached(
|
||||
snapshot: &CandidateRuntimeSelectionSnapshot,
|
||||
now_unix_secs: u64,
|
||||
) -> bool {
|
||||
auth_snapshot
|
||||
.and_then(|snapshot| {
|
||||
usize::try_from(snapshot.api_key_concurrent_limit?)
|
||||
.ok()
|
||||
.and_then(|limit| {
|
||||
if limit == 0 {
|
||||
return None;
|
||||
}
|
||||
Some((snapshot.api_key_id.as_str(), limit))
|
||||
})
|
||||
})
|
||||
.is_some_and(|(api_key_id, limit)| {
|
||||
auth_api_key_concurrency_limit_reached(
|
||||
&snapshot.recent_candidates,
|
||||
now_unix_secs,
|
||||
api_key_id,
|
||||
limit,
|
||||
auth_snapshot_concurrency_limit(auth_snapshot).is_some_and(|(api_key_id, limit)| {
|
||||
auth_api_key_concurrency_limit_reached(
|
||||
&snapshot.recent_candidates,
|
||||
now_unix_secs,
|
||||
api_key_id,
|
||||
limit,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn auth_snapshot_concurrency_limit(
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
) -> Option<(&str, usize)> {
|
||||
let snapshot = auth_snapshot?;
|
||||
let limit = usize::try_from(snapshot.api_key_concurrent_limit?).ok()?;
|
||||
(limit > 0).then_some((snapshot.api_key_id.as_str(), limit))
|
||||
}
|
||||
|
||||
async fn read_auth_api_key_concurrency_limit_reached(
|
||||
state: &(impl SchedulerRuntimeState + ?Sized),
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
) -> Result<bool, GatewayError> {
|
||||
let Some((api_key_id, limit)) = auth_snapshot_concurrency_limit(auth_snapshot) else {
|
||||
return Ok(false);
|
||||
};
|
||||
let recent_candidates = state.read_recent_request_candidates(128).await?;
|
||||
Ok(auth_api_key_concurrency_limit_reached(
|
||||
&recent_candidates,
|
||||
crate::clock::current_unix_secs(),
|
||||
api_key_id,
|
||||
limit,
|
||||
))
|
||||
}
|
||||
|
||||
/// A retry always rebuilds candidates, including at the deadline. Only the
|
||||
/// intervening polls omit catalog, quota and ranking work while auth is blocked.
|
||||
pub(crate) async fn wait_for_auth_api_key_concurrency_retry(
|
||||
state: &(impl SchedulerRuntimeState + ?Sized),
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
deadline: Instant,
|
||||
poll_interval: Duration,
|
||||
) -> Result<bool, GatewayError> {
|
||||
if Instant::now() >= deadline {
|
||||
return Ok(false);
|
||||
}
|
||||
let poll_interval = poll_interval.max(Duration::from_millis(1));
|
||||
loop {
|
||||
let remaining = deadline.saturating_duration_since(Instant::now());
|
||||
tokio::time::sleep(poll_interval.min(remaining)).await;
|
||||
if Instant::now() >= deadline
|
||||
|| !read_auth_api_key_concurrency_limit_reached(state, auth_snapshot).await?
|
||||
{
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn select_with_auth_concurrency_wait<T, Select, Selection>(
|
||||
state: &(impl SchedulerRuntimeState + ?Sized),
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
now_unix_secs: u64,
|
||||
wait_timeout: Duration,
|
||||
poll_interval: Duration,
|
||||
mut select: Select,
|
||||
) -> Result<T, GatewayError>
|
||||
where
|
||||
Select: FnMut(u64) -> Selection,
|
||||
Selection: Future<Output = Result<(T, bool), GatewayError>>,
|
||||
{
|
||||
let deadline = Instant::now() + wait_timeout;
|
||||
let mut attempt_now_unix_secs = now_unix_secs;
|
||||
loop {
|
||||
let (result, auth_limit_blocked) = select(attempt_now_unix_secs).await?;
|
||||
if !auth_limit_blocked
|
||||
|| !wait_for_auth_api_key_concurrency_retry(
|
||||
state,
|
||||
auth_snapshot,
|
||||
deadline,
|
||||
poll_interval,
|
||||
)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
return Ok(result);
|
||||
}
|
||||
attempt_now_unix_secs = crate::clock::current_unix_secs();
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn is_candidate_selectable(
|
||||
|
||||
@@ -0,0 +1,548 @@
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::Mutex;
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use aether_data_contracts::repository::candidates::{
|
||||
RequestCandidateStatus, StoredRequestCandidate,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot;
|
||||
use aether_scheduler_core::{SchedulerAffinityTarget, SchedulerMinimalCandidateSelectionCandidate};
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::Notify;
|
||||
|
||||
use crate::data::auth::GatewayAuthApiKeySnapshot;
|
||||
use crate::data::candidate_selection::MinimalCandidateSelectionRowSource;
|
||||
use crate::scheduler::config::SchedulerOrderingConfig;
|
||||
use crate::scheduler::state::SchedulerRuntimeState;
|
||||
use crate::GatewayError;
|
||||
|
||||
use super::super::{
|
||||
is_exact_all_skipped_by_auth_limit,
|
||||
list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal,
|
||||
list_selectable_candidates_with_skip_reasons, select_with_auth_concurrency_wait,
|
||||
SchedulerSkippedCandidate,
|
||||
};
|
||||
use super::support::{sample_auth_snapshot, sample_provider, sample_row};
|
||||
|
||||
const POLL_INTERVAL: Duration = Duration::from_millis(2);
|
||||
|
||||
enum RecentReadAction {
|
||||
Keep,
|
||||
Release,
|
||||
ReleaseAndReplaceCandidate,
|
||||
ReleaseAndRemoveCandidates,
|
||||
Fail,
|
||||
}
|
||||
|
||||
struct CountingState {
|
||||
rows: Mutex<Vec<StoredMinimalCandidateSelectionRow>>,
|
||||
recent: Mutex<Vec<StoredRequestCandidate>>,
|
||||
recent_actions: Mutex<VecDeque<RecentReadAction>>,
|
||||
row_reads: AtomicUsize,
|
||||
format_reads: AtomicUsize,
|
||||
provider_reads: AtomicUsize,
|
||||
key_reads: AtomicUsize,
|
||||
quota_reads: AtomicUsize,
|
||||
recent_reads: AtomicUsize,
|
||||
row_error_at: Option<usize>,
|
||||
first_row_delay: Duration,
|
||||
poll_observed: Notify,
|
||||
}
|
||||
|
||||
impl CountingState {
|
||||
fn blocked() -> Self {
|
||||
Self {
|
||||
rows: Mutex::new(vec![sample_row()]),
|
||||
recent: Mutex::new(vec![active_candidate()]),
|
||||
recent_actions: Mutex::new(VecDeque::new()),
|
||||
row_reads: AtomicUsize::new(0),
|
||||
format_reads: AtomicUsize::new(0),
|
||||
provider_reads: AtomicUsize::new(0),
|
||||
key_reads: AtomicUsize::new(0),
|
||||
quota_reads: AtomicUsize::new(0),
|
||||
recent_reads: AtomicUsize::new(0),
|
||||
row_error_at: None,
|
||||
first_row_delay: Duration::ZERO,
|
||||
poll_observed: Notify::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn on_recent_reads(self, actions: impl IntoIterator<Item = RecentReadAction>) -> Self {
|
||||
*self.recent_actions.lock().unwrap() = actions.into_iter().collect();
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
fn active_candidate() -> StoredRequestCandidate {
|
||||
let now_ms = i64::try_from(crate::clock::current_unix_secs() * 1000).unwrap();
|
||||
StoredRequestCandidate::new(
|
||||
"active-candidate".to_string(),
|
||||
"active-request".to_string(),
|
||||
Some("user-1".to_string()),
|
||||
Some("api-key-1".to_string()),
|
||||
None,
|
||||
None,
|
||||
0,
|
||||
0,
|
||||
Some("provider-1".to_string()),
|
||||
Some("endpoint-1".to_string()),
|
||||
Some("key-1".to_string()),
|
||||
RequestCandidateStatus::Streaming,
|
||||
None,
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
now_ms,
|
||||
Some(now_ms),
|
||||
None,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn limited_auth() -> GatewayAuthApiKeySnapshot {
|
||||
let mut auth = sample_auth_snapshot("api-key-1");
|
||||
auth.api_key_concurrent_limit = Some(1);
|
||||
auth
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionRowSource for CountingState {
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_global_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
panic!("requested-model selection should use its paged query")
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
panic!("requested-model selection should use its paged query")
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let read = self.row_reads.fetch_add(1, Ordering::SeqCst) + 1;
|
||||
if read == 1 && !self.first_row_delay.is_zero() {
|
||||
tokio::time::sleep(self.first_row_delay).await;
|
||||
}
|
||||
if self.row_error_at == Some(read) {
|
||||
return Err(DataLayerError::Postgres(
|
||||
"candidate query failed".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(self
|
||||
.rows
|
||||
.lock()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter(|row| {
|
||||
row.endpoint_api_format == query.api_format
|
||||
&& row.global_model_name == query.requested_model_name
|
||||
})
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.cloned()
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.format_reads.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(self.rows.lock().unwrap().clone())
|
||||
}
|
||||
|
||||
async fn read_pool_key_candidate_rows_for_group(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
panic!("the test has no provider pool")
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SchedulerRuntimeState for CountingState {
|
||||
async fn read_provider_quota_snapshot(
|
||||
&self,
|
||||
_provider_id: &str,
|
||||
) -> Result<Option<StoredProviderQuotaSnapshot>, GatewayError> {
|
||||
self.quota_reads.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn read_provider_catalog_providers_by_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, GatewayError> {
|
||||
self.provider_reads.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(provider_ids
|
||||
.iter()
|
||||
.map(|id| sample_provider(id, None))
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn read_provider_catalog_keys_by_ids(
|
||||
&self,
|
||||
_key_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> {
|
||||
self.key_reads.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn read_recent_request_candidates(
|
||||
&self,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StoredRequestCandidate>, GatewayError> {
|
||||
assert_eq!(
|
||||
limit, 128,
|
||||
"polls must use the same sample as full selection"
|
||||
);
|
||||
let read = self.recent_reads.fetch_add(1, Ordering::SeqCst) + 1;
|
||||
let action = self
|
||||
.recent_actions
|
||||
.lock()
|
||||
.unwrap()
|
||||
.pop_front()
|
||||
.unwrap_or(RecentReadAction::Keep);
|
||||
if matches!(action, RecentReadAction::Fail) {
|
||||
return Err(GatewayError::Internal(
|
||||
"recent candidates failed".to_string(),
|
||||
));
|
||||
}
|
||||
if matches!(
|
||||
action,
|
||||
RecentReadAction::Release
|
||||
| RecentReadAction::ReleaseAndReplaceCandidate
|
||||
| RecentReadAction::ReleaseAndRemoveCandidates
|
||||
) {
|
||||
for candidate in self.recent.lock().unwrap().iter_mut() {
|
||||
candidate.status = RequestCandidateStatus::Success;
|
||||
candidate.finished_at_unix_ms = Some(crate::clock::current_unix_secs() * 1000);
|
||||
}
|
||||
}
|
||||
match action {
|
||||
RecentReadAction::ReleaseAndReplaceCandidate => {
|
||||
self.rows.lock().unwrap()[0].key_id = "replacement-key".to_string();
|
||||
}
|
||||
RecentReadAction::ReleaseAndRemoveCandidates => self.rows.lock().unwrap().clear(),
|
||||
_ => {}
|
||||
}
|
||||
if read > 1 {
|
||||
self.poll_observed.notify_one();
|
||||
}
|
||||
Ok(self.recent.lock().unwrap().clone())
|
||||
}
|
||||
|
||||
fn provider_key_rpm_reset_at(&self, _key_id: &str, _now_unix_secs: u64) -> Option<u64> {
|
||||
None
|
||||
}
|
||||
|
||||
fn read_cached_scheduler_affinity_target(
|
||||
&self,
|
||||
_cache_key: &str,
|
||||
_ttl: Duration,
|
||||
) -> Option<SchedulerAffinityTarget> {
|
||||
None
|
||||
}
|
||||
|
||||
fn scheduler_affinity_epoch(&self) -> u64 {
|
||||
0
|
||||
}
|
||||
|
||||
fn remember_scheduler_affinity_target(
|
||||
&self,
|
||||
_cache_key: &str,
|
||||
_target: SchedulerAffinityTarget,
|
||||
_ttl: Duration,
|
||||
_max_entries: usize,
|
||||
) {
|
||||
}
|
||||
|
||||
fn remember_scheduler_affinity_target_for_epoch(
|
||||
&self,
|
||||
_cache_key: &str,
|
||||
_target: SchedulerAffinityTarget,
|
||||
_ttl: Duration,
|
||||
_max_entries: usize,
|
||||
_expected_epoch: Option<u64>,
|
||||
) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
type Selection = (
|
||||
Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||
Vec<SchedulerSkippedCandidate>,
|
||||
);
|
||||
|
||||
async fn select_requested_model(
|
||||
state: &CountingState,
|
||||
auth: Option<&GatewayAuthApiKeySnapshot>,
|
||||
timeout: Duration,
|
||||
) -> Result<Selection, GatewayError> {
|
||||
select_with_auth_concurrency_wait(
|
||||
state,
|
||||
auth,
|
||||
crate::clock::current_unix_secs(),
|
||||
timeout,
|
||||
POLL_INTERVAL,
|
||||
|now| async move {
|
||||
let result = list_selectable_candidates_with_skip_reasons(
|
||||
state,
|
||||
state,
|
||||
"openai:chat",
|
||||
"gpt-4.1",
|
||||
false,
|
||||
None,
|
||||
auth,
|
||||
None,
|
||||
now,
|
||||
false,
|
||||
SchedulerOrderingConfig::default(),
|
||||
)
|
||||
.await?;
|
||||
let blocked = is_exact_all_skipped_by_auth_limit(&result.0, &result.1);
|
||||
Ok((result, blocked))
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_blocked_selectors_only_prepare_at_start_and_deadline() {
|
||||
let state = CountingState::blocked();
|
||||
let auth = limited_auth();
|
||||
let outcomes = futures_util::future::join_all(
|
||||
(0..8).map(|_| select_requested_model(&state, Some(&auth), Duration::from_millis(80))),
|
||||
)
|
||||
.await;
|
||||
|
||||
for outcome in outcomes {
|
||||
let (selected, skipped) = outcome.unwrap();
|
||||
assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped));
|
||||
}
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 16);
|
||||
assert_eq!(state.provider_reads.load(Ordering::SeqCst), 32);
|
||||
assert_eq!(state.key_reads.load(Ordering::SeqCst), 16);
|
||||
assert_eq!(state.quota_reads.load(Ordering::SeqCst), 16);
|
||||
assert!(state.recent_reads.load(Ordering::SeqCst) > 16);
|
||||
assert_eq!(state.format_reads.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn released_auth_slot_rebuilds_changed_candidates() {
|
||||
let state = CountingState::blocked().on_recent_reads([
|
||||
RecentReadAction::Keep,
|
||||
RecentReadAction::Keep,
|
||||
RecentReadAction::ReleaseAndReplaceCandidate,
|
||||
]);
|
||||
let auth = limited_auth();
|
||||
let (selected, skipped) = select_requested_model(&state, Some(&auth), Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(skipped.is_empty());
|
||||
assert_eq!(selected.len(), 1);
|
||||
assert_eq!(selected[0].key_id, "replacement-key");
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 4);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn released_auth_slot_does_not_reuse_removed_candidates() {
|
||||
let state = CountingState::blocked().on_recent_reads([
|
||||
RecentReadAction::Keep,
|
||||
RecentReadAction::ReleaseAndRemoveCandidates,
|
||||
]);
|
||||
let (selected, skipped) =
|
||||
select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(selected.is_empty());
|
||||
assert!(skipped.is_empty());
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn lightweight_poll_errors_are_propagated_without_another_full_query() {
|
||||
let state =
|
||||
CountingState::blocked().on_recent_reads([RecentReadAction::Keep, RecentReadAction::Fail]);
|
||||
let error = select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(
|
||||
matches!(error, GatewayError::Internal(message) if message == "recent candidates failed")
|
||||
);
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn recovered_selection_errors_are_propagated() {
|
||||
let mut state = CountingState::blocked()
|
||||
.on_recent_reads([RecentReadAction::Keep, RecentReadAction::Release]);
|
||||
state.row_error_at = Some(2);
|
||||
let error = select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(
|
||||
matches!(error, GatewayError::Internal(message) if message.contains("candidate query failed"))
|
||||
);
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn first_selection_time_consumes_the_wait_budget() {
|
||||
let mut state = CountingState::blocked();
|
||||
state.first_row_delay = Duration::from_millis(30);
|
||||
let (selected, skipped) =
|
||||
select_requested_model(&state, Some(&limited_auth()), Duration::from_millis(5))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped));
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_stops_polling_and_candidate_queries() {
|
||||
let state = CountingState::blocked();
|
||||
let auth = limited_auth();
|
||||
let mut selection = Box::pin(select_requested_model(
|
||||
&state,
|
||||
Some(&auth),
|
||||
Duration::from_secs(1),
|
||||
));
|
||||
tokio::select! {
|
||||
result = &mut selection => panic!("selection completed before cancellation: {result:?}"),
|
||||
_ = state.poll_observed.notified() => {}
|
||||
}
|
||||
drop(selection);
|
||||
let recent_reads = state.recent_reads.load(Ordering::SeqCst);
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
|
||||
assert!(recent_reads >= 2);
|
||||
assert_eq!(state.recent_reads.load(Ordering::SeqCst), recent_reads);
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn no_model_capability_selection_does_not_reenumerate_during_polls() {
|
||||
let state = CountingState::blocked();
|
||||
let auth = limited_auth();
|
||||
let selected = select_with_auth_concurrency_wait(
|
||||
&state,
|
||||
Some(&auth),
|
||||
crate::clock::current_unix_secs(),
|
||||
Duration::from_millis(80),
|
||||
POLL_INTERVAL,
|
||||
|now| {
|
||||
list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal(
|
||||
&state,
|
||||
&state,
|
||||
"openai:chat",
|
||||
"cache_1h",
|
||||
false,
|
||||
Some(&auth),
|
||||
None,
|
||||
now,
|
||||
SchedulerOrderingConfig::default(),
|
||||
)
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(selected.is_empty());
|
||||
assert_eq!(state.format_reads.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
|
||||
assert!(state.recent_reads.load(Ordering::SeqCst) > 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn absent_or_disabled_auth_limits_do_not_poll() {
|
||||
for limit in [None, Some(0), Some(-1)] {
|
||||
let state = CountingState::blocked();
|
||||
let mut auth = limited_auth();
|
||||
auth.api_key_concurrent_limit = limit;
|
||||
let (selected, _) = select_requested_model(&state, Some(&auth), Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(selected.len(), 1);
|
||||
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
let state = CountingState::blocked();
|
||||
assert_eq!(
|
||||
select_requested_model(&state, None, Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap()
|
||||
.0
|
||||
.len(),
|
||||
1
|
||||
);
|
||||
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn auth_wait_keeps_candidate_row_and_lifecycle_counting_semantics() {
|
||||
let state = CountingState::blocked();
|
||||
let mut duplicate = active_candidate();
|
||||
duplicate.id = "second-attempt-same-request".to_string();
|
||||
state.recent.lock().unwrap().push(duplicate);
|
||||
let mut auth = limited_auth();
|
||||
auth.api_key_concurrent_limit = Some(2);
|
||||
let (selected, skipped) = select_requested_model(&state, Some(&auth), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped));
|
||||
|
||||
let state = CountingState::blocked();
|
||||
state.recent.lock().unwrap()[0].finished_at_unix_ms =
|
||||
Some(crate::clock::current_unix_secs() * 1000);
|
||||
assert_eq!(
|
||||
select_requested_model(&state, Some(&limited_auth()), Duration::ZERO)
|
||||
.await
|
||||
.unwrap()
|
||||
.0
|
||||
.len(),
|
||||
1
|
||||
);
|
||||
|
||||
let state = CountingState::blocked();
|
||||
state.recent.lock().unwrap()[0].started_at_unix_ms =
|
||||
Some(crate::clock::current_unix_secs().saturating_sub(301) * 1000);
|
||||
assert_eq!(
|
||||
select_requested_model(&state, Some(&limited_auth()), Duration::ZERO)
|
||||
.await
|
||||
.unwrap()
|
||||
.0
|
||||
.len(),
|
||||
1
|
||||
);
|
||||
}
|
||||
@@ -1,4 +1,5 @@
|
||||
mod affinity;
|
||||
mod concurrency_wait;
|
||||
mod model;
|
||||
mod required_capability;
|
||||
mod selection;
|
||||
|
||||
@@ -274,7 +274,13 @@ mod tests {
|
||||
"priority_mode": "provider",
|
||||
"scheduling_mode": "cache_affinity",
|
||||
"keep_priority_on_conversion": false,
|
||||
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS
|
||||
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS,
|
||||
"max_transfer_count": 0,
|
||||
"max_transfer_timeout_seconds": 0,
|
||||
"failover_rules": {
|
||||
"success_failover_patterns": [],
|
||||
"error_stop_patterns": []
|
||||
}
|
||||
})
|
||||
);
|
||||
|
||||
@@ -323,7 +329,13 @@ mod tests {
|
||||
"priority_mode": "provider",
|
||||
"scheduling_mode": "cache_affinity",
|
||||
"keep_priority_on_conversion": false,
|
||||
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS
|
||||
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS,
|
||||
"max_transfer_count": 0,
|
||||
"max_transfer_timeout_seconds": 0,
|
||||
"failover_rules": {
|
||||
"success_failover_patterns": [],
|
||||
"error_stop_patterns": []
|
||||
}
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::{body::Body, routing::get, Router};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
use tokio::sync::Notify;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use super::super::{serve_gateway_router, GatewayHttpLimits, HttpConnectionBudget};
|
||||
|
||||
async fn start(
|
||||
router: Router,
|
||||
) -> (
|
||||
std::net::SocketAddr,
|
||||
CancellationToken,
|
||||
Arc<HttpConnectionBudget>,
|
||||
tokio::task::JoinHandle<Result<(), String>>,
|
||||
) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let shutdown = CancellationToken::new();
|
||||
let budget = Arc::new(HttpConnectionBudget::new(16));
|
||||
let stop = shutdown.clone();
|
||||
let shared = Arc::clone(&budget);
|
||||
let server = tokio::spawn(async move {
|
||||
serve_gateway_router(
|
||||
vec![listener],
|
||||
router,
|
||||
shared,
|
||||
GatewayHttpLimits {
|
||||
http2_max_concurrent_streams: 16,
|
||||
http_header_read_timeout_ms: 10_000,
|
||||
http_header_max_bytes: 32_768,
|
||||
http_max_headers: 100,
|
||||
},
|
||||
stop,
|
||||
)
|
||||
.await
|
||||
.map_err(|error| error.to_string())
|
||||
});
|
||||
(address, shutdown, budget, server)
|
||||
}
|
||||
|
||||
async fn within<T>(future: impl std::future::Future<Output = T>) -> T {
|
||||
tokio::time::timeout(Duration::from_secs(3), future)
|
||||
.await
|
||||
.expect("shutdown deadline")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_shutdown_drains_in_flight_http1_and_http2_responses() {
|
||||
for http2 in [false, true] {
|
||||
let started = Arc::new(Notify::new());
|
||||
let release = Arc::new(Notify::new());
|
||||
let handler_started = Arc::clone(&started);
|
||||
let handler_release = Arc::clone(&release);
|
||||
let router = Router::new().route(
|
||||
"/",
|
||||
get(move || {
|
||||
let started = Arc::clone(&handler_started);
|
||||
let release = Arc::clone(&handler_release);
|
||||
async move {
|
||||
started.notify_one();
|
||||
release.notified().await;
|
||||
"complete response"
|
||||
}
|
||||
}),
|
||||
);
|
||||
let (address, shutdown, budget, server) = start(router).await;
|
||||
let client = if http2 {
|
||||
reqwest::Client::builder().http2_prior_knowledge()
|
||||
} else {
|
||||
reqwest::Client::builder().http1_only()
|
||||
}
|
||||
.build()
|
||||
.unwrap();
|
||||
let request = tokio::spawn(async move {
|
||||
client
|
||||
.get(format!("http://{address}/"))
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.text()
|
||||
.await
|
||||
.unwrap()
|
||||
});
|
||||
within(started.notified()).await;
|
||||
shutdown.cancel();
|
||||
tokio::time::sleep(Duration::from_millis(20)).await;
|
||||
assert!(!server.is_finished());
|
||||
assert!(!request.is_finished());
|
||||
assert!(TcpStream::connect(address).await.is_err());
|
||||
release.notify_one();
|
||||
assert_eq!(within(request).await.unwrap(), "complete response");
|
||||
within(server).await.unwrap().unwrap();
|
||||
assert_eq!(budget.snapshot().in_flight, 0);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_shutdown_closes_idle_protocol_detection_connections() {
|
||||
let (address, shutdown, budget, server) = start(Router::new()).await;
|
||||
let mut peer = TcpStream::connect(address).await.unwrap();
|
||||
within(async {
|
||||
while budget.snapshot().in_flight == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await;
|
||||
shutdown.cancel();
|
||||
within(server).await.unwrap().unwrap();
|
||||
assert_eq!(within(peer.read(&mut [0_u8; 1])).await.unwrap(), 0);
|
||||
assert_eq!(budget.snapshot().in_flight, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_shutdown_force_cancels_a_handler_without_socket_io() {
|
||||
struct HandlerDrop(Arc<Notify>);
|
||||
impl Drop for HandlerDrop {
|
||||
fn drop(&mut self) {
|
||||
self.0.notify_one();
|
||||
}
|
||||
}
|
||||
for http2 in [false, true] {
|
||||
let started = Arc::new(Notify::new());
|
||||
let dropped = Arc::new(Notify::new());
|
||||
let request_started = Arc::clone(&started);
|
||||
let request_dropped = Arc::clone(&dropped);
|
||||
let router = Router::new().route(
|
||||
"/",
|
||||
get(move || {
|
||||
let started = Arc::clone(&request_started);
|
||||
let dropped = Arc::clone(&request_dropped);
|
||||
async move {
|
||||
let _guard = HandlerDrop(dropped);
|
||||
started.notify_one();
|
||||
std::future::pending::<&'static str>().await
|
||||
}
|
||||
}),
|
||||
);
|
||||
let (address, shutdown, budget, server) = start(router).await;
|
||||
let client = if http2 {
|
||||
reqwest::Client::builder().http2_prior_knowledge()
|
||||
} else {
|
||||
reqwest::Client::builder().http1_only()
|
||||
}
|
||||
.build()
|
||||
.unwrap();
|
||||
let request =
|
||||
tokio::spawn(async move { client.get(format!("http://{address}/")).send().await });
|
||||
within(started.notified()).await;
|
||||
shutdown.cancel();
|
||||
budget.force_close();
|
||||
within(server).await.unwrap().unwrap();
|
||||
within(dropped.notified()).await;
|
||||
assert!(within(request).await.unwrap().is_err());
|
||||
assert_eq!(budget.snapshot().in_flight, 0);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_shutdown_force_closes_upgraded_io_and_waits_for_release() {
|
||||
let upgraded_done = Arc::new(Notify::new());
|
||||
let done = Arc::clone(&upgraded_done);
|
||||
let router = Router::new().route(
|
||||
"/",
|
||||
get(move |mut request: axum::extract::Request| {
|
||||
let done = Arc::clone(&done);
|
||||
async move {
|
||||
let upgrade = hyper::upgrade::on(&mut request);
|
||||
tokio::spawn(async move {
|
||||
let upgraded = upgrade.await.unwrap();
|
||||
let mut io = hyper_util::rt::TokioIo::new(upgraded);
|
||||
let error = io.read_u8().await.unwrap_err();
|
||||
assert_eq!(error.kind(), std::io::ErrorKind::ConnectionAborted);
|
||||
drop(io);
|
||||
done.notify_one();
|
||||
});
|
||||
axum::http::Response::builder()
|
||||
.status(101)
|
||||
.header("connection", "upgrade")
|
||||
.header("upgrade", "echo")
|
||||
.body(Body::empty())
|
||||
.unwrap()
|
||||
}
|
||||
}),
|
||||
);
|
||||
let (address, shutdown, budget, server) = start(router).await;
|
||||
let mut peer = TcpStream::connect(address).await.unwrap();
|
||||
peer.write_all(
|
||||
b"GET / HTTP/1.1\r\nHost: localhost\r\nConnection: upgrade\r\nUpgrade: echo\r\n\r\n",
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let mut header = Vec::new();
|
||||
within(async {
|
||||
while !header.ends_with(b"\r\n\r\n") {
|
||||
header.push(peer.read_u8().await.unwrap());
|
||||
}
|
||||
})
|
||||
.await;
|
||||
assert!(header.starts_with(b"HTTP/1.1 101"));
|
||||
shutdown.cancel();
|
||||
tokio::time::sleep(Duration::from_millis(20)).await;
|
||||
assert!(!server.is_finished());
|
||||
budget.force_close();
|
||||
within(upgraded_done.notified()).await;
|
||||
within(server).await.unwrap().unwrap();
|
||||
assert_eq!(budget.snapshot().in_flight, 0);
|
||||
}
|
||||
@@ -8,6 +8,7 @@ use aether_data::repository::users::StoredUserGroup;
|
||||
use aether_data_contracts::repository::billing::UserDailyQuotaAvailabilityRecord;
|
||||
use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot;
|
||||
use aether_data_contracts::repository::usage::UsageCounterHealthSnapshot;
|
||||
use aether_gateway_frontdoor::HttpConnectionBudget;
|
||||
use aether_runtime::ConcurrencyGate;
|
||||
use aether_runtime_state::{RuntimeSemaphore, RuntimeState};
|
||||
use dashmap::DashMap;
|
||||
@@ -35,6 +36,7 @@ use super::{
|
||||
};
|
||||
|
||||
const MIN_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 1_000;
|
||||
const DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 120_000;
|
||||
const MAX_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 600_000;
|
||||
const REQUEST_BODY_READ_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS";
|
||||
const DEFAULT_REQUEST_BODY_BUFFER_BUDGET_MB: usize = 256;
|
||||
@@ -193,7 +195,9 @@ fn optional_env_duration_ms(key: &str, min_ms: u64, max_ms: u64) -> Option<Durat
|
||||
}
|
||||
|
||||
fn parse_optional_duration_ms(raw: Option<&str>, min_ms: u64, max_ms: u64) -> Option<Duration> {
|
||||
let parsed = raw?.trim().parse::<u64>().ok()?;
|
||||
let parsed = raw
|
||||
.and_then(|value| value.trim().parse::<u64>().ok())
|
||||
.unwrap_or(DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS);
|
||||
if parsed == 0 {
|
||||
return None;
|
||||
}
|
||||
@@ -388,6 +392,7 @@ pub struct AppState {
|
||||
pub(crate) video_task_poller: Option<VideoTaskPollerConfig>,
|
||||
pub(crate) frontdoor_runtime_guards: Arc<FrontdoorRuntimeGuardConfig>,
|
||||
pub(crate) request_body_buffer_budget: Arc<Semaphore>,
|
||||
pub(crate) http_connection_budget: Option<Arc<HttpConnectionBudget>>,
|
||||
pub(crate) request_gate: Option<Arc<ConcurrencyGate>>,
|
||||
pub(crate) websocket_connection_gate: Option<Arc<ConcurrencyGate>>,
|
||||
pub(crate) auth_snapshot_load_gate: Option<Arc<ConcurrencyGate>>,
|
||||
@@ -456,6 +461,8 @@ pub struct AppState {
|
||||
Arc<DashMap<String, LocalExecutionRuntimeMissDiagnostic>>,
|
||||
pub(crate) admin_monitoring_error_stats_reset_at: Arc<StdMutex<Option<u64>>>,
|
||||
pub(crate) provider_delete_tasks: Arc<StdMutex<HashMap<String, LocalProviderDeleteTaskState>>>,
|
||||
pub(crate) pool_quota_probe_replenish:
|
||||
Arc<crate::maintenance::PoolQuotaProbeReplenishCoordinator>,
|
||||
#[cfg(test)]
|
||||
pub(crate) turnstile_siteverify_url_override: Option<String>,
|
||||
#[cfg(test)]
|
||||
@@ -533,20 +540,22 @@ mod tests {
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn request_body_read_timeout_parser_defaults_to_disabled() {
|
||||
assert_eq!(
|
||||
parse_optional_duration_ms(
|
||||
None,
|
||||
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
|
||||
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
|
||||
),
|
||||
None
|
||||
);
|
||||
fn request_body_read_timeout_parser_uses_a_finite_default() {
|
||||
for value in [None, Some(""), Some("invalid"), Some("-1")] {
|
||||
assert_eq!(
|
||||
parse_optional_duration_ms(
|
||||
value,
|
||||
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
|
||||
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
|
||||
),
|
||||
Some(Duration::from_millis(DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS)),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_body_read_timeout_parser_disables_zero_and_invalid_values() {
|
||||
for value in ["", "invalid", "-1", "0", " 0 "] {
|
||||
fn request_body_read_timeout_parser_only_disables_explicit_zero() {
|
||||
for value in ["0", " 0 "] {
|
||||
assert_eq!(
|
||||
parse_optional_duration_ms(
|
||||
Some(value),
|
||||
|
||||
@@ -13,6 +13,7 @@ use aether_data::repository::proxy_nodes::{
|
||||
use aether_data_contracts::repository::usage::{
|
||||
UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot,
|
||||
};
|
||||
use aether_gateway_frontdoor::{HttpConnectionBudget, HttpConnectionBudgetSnapshot};
|
||||
use aether_http::{apply_http_client_config, HttpClientConfig};
|
||||
use aether_runtime::{
|
||||
service_up_sample, AdmissionPermit, ConcurrencyGate, ConcurrencySnapshot, MetricKind,
|
||||
@@ -358,6 +359,7 @@ impl AppState {
|
||||
request_body_buffer_budget: Arc::new(tokio::sync::Semaphore::new(
|
||||
frontdoor_runtime_guards.request_body_buffer_budget_permits,
|
||||
)),
|
||||
http_connection_budget: None,
|
||||
request_gate: None,
|
||||
websocket_connection_gate: None,
|
||||
auth_snapshot_load_gate: frontdoor_runtime_guards
|
||||
@@ -438,6 +440,9 @@ impl AppState {
|
||||
local_execution_runtime_miss_diagnostics: Arc::new(DashMap::new()),
|
||||
admin_monitoring_error_stats_reset_at: Arc::new(StdMutex::new(None)),
|
||||
provider_delete_tasks: Arc::new(StdMutex::new(HashMap::new())),
|
||||
pool_quota_probe_replenish: Arc::new(
|
||||
crate::maintenance::PoolQuotaProbeReplenishCoordinator::default(),
|
||||
),
|
||||
#[cfg(test)]
|
||||
turnstile_siteverify_url_override: None,
|
||||
#[cfg(test)]
|
||||
@@ -614,6 +619,11 @@ impl AppState {
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_http_connection_budget(mut self, budget: Arc<HttpConnectionBudget>) -> Self {
|
||||
self.http_connection_budget = Some(budget);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_request_concurrency_limit(mut self, limit: usize) -> Self {
|
||||
let limit = limit.max(1);
|
||||
self.request_gate = Some(Arc::new(ConcurrencyGate::new("gateway_requests", limit)));
|
||||
@@ -635,6 +645,10 @@ impl AppState {
|
||||
}
|
||||
|
||||
pub fn with_runtime_state(mut self, runtime_state: Arc<RuntimeState>) -> Self {
|
||||
if !Arc::ptr_eq(&self.runtime_state, &runtime_state) {
|
||||
self.pool_quota_probe_replenish =
|
||||
Arc::new(crate::maintenance::PoolQuotaProbeReplenishCoordinator::default());
|
||||
}
|
||||
self.runtime_state = runtime_state;
|
||||
self.admin_security_blacklist_cache.clear();
|
||||
self.admin_security_whitelist_cache.clear();
|
||||
@@ -1666,6 +1680,9 @@ impl AppState {
|
||||
.unwrap_or(u64::MAX),
|
||||
),
|
||||
]);
|
||||
if let Some(budget) = &self.http_connection_budget {
|
||||
samples.extend(http_connection_metric_samples(&budget.snapshot()));
|
||||
}
|
||||
if let Some(snapshot) = self.request_concurrency_snapshot() {
|
||||
samples.extend(snapshot.to_metric_samples("gateway_requests"));
|
||||
}
|
||||
@@ -1820,10 +1837,44 @@ impl AppState {
|
||||
&self.task_supervisor_metrics.snapshot(),
|
||||
));
|
||||
samples.extend(crate::tokio_metrics::gateway_tokio_runtime_metric_samples());
|
||||
samples.extend(aether_runtime::logging_metric_samples());
|
||||
samples.extend(
|
||||
crate::execution_runtime::transport::direct_reqwest_client_cache_metric_samples(),
|
||||
);
|
||||
samples.extend(self.upstream_target_admission.metric_samples());
|
||||
let probe_replenish = self.pool_quota_probe_replenish.snapshot();
|
||||
samples.extend([
|
||||
MetricSample::new(
|
||||
"pool_quota_probe_replenish_provider_capacity",
|
||||
"Maximum providers with an active local request-triggered probe task.",
|
||||
MetricKind::Gauge,
|
||||
probe_replenish.capacity as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"pool_quota_probe_replenish_active_providers",
|
||||
"Providers with an active local request-triggered probe task.",
|
||||
MetricKind::Gauge,
|
||||
probe_replenish.active as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"pool_quota_probe_replenish_started_total",
|
||||
"Local request-triggered provider probe tasks admitted.",
|
||||
MetricKind::Counter,
|
||||
probe_replenish.started_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"pool_quota_probe_replenish_coalesced_total",
|
||||
"Request triggers merged into an existing local provider probe task.",
|
||||
MetricKind::Counter,
|
||||
probe_replenish.coalesced_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"pool_quota_probe_replenish_capacity_rejected_total",
|
||||
"Request-triggered probe tasks skipped because the local provider limit was reached.",
|
||||
MetricKind::Counter,
|
||||
probe_replenish.capacity_rejected_total,
|
||||
),
|
||||
]);
|
||||
samples.extend(crate::cache::candidate_page_cache_metric_samples());
|
||||
samples.extend(crate::stage_metrics::gateway_stage_metric_samples());
|
||||
samples.extend(self.tunnel.metric_samples());
|
||||
@@ -2118,6 +2169,36 @@ impl AppState {
|
||||
state
|
||||
}
|
||||
|
||||
pub async fn shutdown_usage_runtime(
|
||||
&self,
|
||||
timeout: Duration,
|
||||
) -> Result<(), aether_data_contracts::DataLayerError> {
|
||||
tokio::time::timeout(timeout, async {
|
||||
let local_queue = self.runtime_state.is_memory().then(|| {
|
||||
let queue: Arc<dyn RuntimeQueueStore> = self.runtime_state.clone();
|
||||
queue
|
||||
});
|
||||
self.usage_runtime
|
||||
.shutdown_with_local_queue(timeout, local_queue)
|
||||
.await?;
|
||||
while self
|
||||
.request_candidate_queue
|
||||
.as_ref()
|
||||
.is_some_and(|queue| queue.pending_writes() != 0)
|
||||
{
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
}
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.map_err(|_| {
|
||||
aether_data_contracts::DataLayerError::TimedOut(
|
||||
"gateway usage or candidate persistence did not drain before shutdown deadline"
|
||||
.to_string(),
|
||||
)
|
||||
})?
|
||||
}
|
||||
|
||||
pub fn spawn_background_tasks(&self) -> crate::task_runtime::TaskSupervisor {
|
||||
let background_state = self.background_worker_state();
|
||||
let mut supervisor =
|
||||
@@ -2277,6 +2358,41 @@ fn database_bounded_auth_load_limit(
|
||||
})
|
||||
}
|
||||
|
||||
fn http_connection_metric_samples(snapshot: &HttpConnectionBudgetSnapshot) -> Vec<MetricSample> {
|
||||
vec![
|
||||
MetricSample::new(
|
||||
"gateway_http_connections_limit",
|
||||
"Maximum admitted inbound TCP connections across gateway listeners.",
|
||||
MetricKind::Gauge,
|
||||
u64::try_from(snapshot.limit).unwrap_or(u64::MAX),
|
||||
),
|
||||
MetricSample::new(
|
||||
"gateway_http_connections_in_flight",
|
||||
"Currently admitted inbound TCP connections, including upgraded connections.",
|
||||
MetricKind::Gauge,
|
||||
u64::try_from(snapshot.in_flight).unwrap_or(u64::MAX),
|
||||
),
|
||||
MetricSample::new(
|
||||
"gateway_http_connections_high_watermark",
|
||||
"Highest simultaneous admitted inbound TCP connection count.",
|
||||
MetricKind::Gauge,
|
||||
u64::try_from(snapshot.high_watermark).unwrap_or(u64::MAX),
|
||||
),
|
||||
MetricSample::new(
|
||||
"gateway_http_connections_rejected_total",
|
||||
"Inbound TCP connections closed because the connection budget was full.",
|
||||
MetricKind::Counter,
|
||||
snapshot.rejected_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"gateway_http_connections_accept_errors_total",
|
||||
"Listener accept errors retried with bounded backoff.",
|
||||
MetricKind::Counter,
|
||||
snapshot.accept_errors_total,
|
||||
),
|
||||
]
|
||||
}
|
||||
|
||||
fn task_supervisor_metric_samples(
|
||||
snapshot: &aether_task_runtime::TaskSupervisorMetricsSnapshot,
|
||||
) -> Vec<MetricSample> {
|
||||
@@ -3190,6 +3306,18 @@ fn usage_runtime_metric_samples(
|
||||
snapshot: &usage::UsageRuntimeMetricsSnapshot,
|
||||
) -> Vec<MetricSample> {
|
||||
vec![
|
||||
MetricSample::new(
|
||||
"usage_runtime_shutdown_started", "Whether local usage admission is closed for shutdown.",
|
||||
MetricKind::Gauge, u64::from(snapshot.shutdown_started),
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_producers_in_flight", "Tracked requests and finalizers that can still submit usage events.",
|
||||
MetricKind::Gauge, snapshot.producers_in_flight as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_delayed_lifecycle_pending", "Delayed lifecycle events still held in this process.",
|
||||
MetricKind::Gauge, snapshot.delayed_lifecycle_pending as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_enabled",
|
||||
"Whether the gateway usage runtime is enabled.",
|
||||
@@ -3322,6 +3450,132 @@ fn usage_runtime_metric_samples(
|
||||
MetricKind::Gauge,
|
||||
u64::from(snapshot.retry_deferred_lifecycle_events),
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_payload_max_bytes",
|
||||
"Maximum serialized JSON payload bytes for a new usage queue message.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.queue_payload_max_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_payload_downgraded_total",
|
||||
"Process-wide enqueue and retry validation encoding attempts that omitted diagnostic data after exceeding the payload limit; not unique events.",
|
||||
MetricKind::Counter,
|
||||
snapshot.queue_payload_downgraded_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_payload_rejected_total",
|
||||
"Process-wide enqueue and retry validation encoding attempts rejected because the payload exceeded its limit or billing facts could not be preserved; not unique events.",
|
||||
MetricKind::Counter,
|
||||
snapshot.queue_payload_rejected_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_payload_budget_bytes",
|
||||
"Process-wide payload reservation limit for usage worker reads and reclaims; not a wire or heap limit.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.queue_read_payload_budget_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_batch_payload_bytes",
|
||||
"Target payload reservation per usage worker batch, allowing at least one configured maximum payload.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.queue_read_batch_payload_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_payload_reserved_bytes",
|
||||
"Process-wide logical payload bytes reserved by usage worker reads, reclaims and unprocessed batches.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.queue_read_payload_reserved_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_payload_waiters",
|
||||
"Usage workers currently waiting for shared payload reservation capacity.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.queue_read_payload_waiters as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_payload_wait_total",
|
||||
"Process-wide usage worker batch reservation attempts that had to wait for capacity.",
|
||||
MetricKind::Counter,
|
||||
snapshot.queue_read_payload_wait_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_actual_field_bytes_total",
|
||||
"Cumulative key and value bytes observed in reserved usage worker batches, including repeated reclaims; excludes Redis framing and allocations.",
|
||||
MetricKind::Counter,
|
||||
snapshot.queue_read_actual_field_bytes_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_oversized_entries_total",
|
||||
"Observed usage entries whose combined field value bytes exceed the consumer payload estimate, including repeated reclaims.",
|
||||
MetricKind::Counter,
|
||||
snapshot.queue_read_oversized_entries_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_queue_read_oversized_batches_total",
|
||||
"Observed usage batches whose combined field value bytes exceed their initial reservation; entries remain eligible for processing.",
|
||||
MetricKind::Counter,
|
||||
snapshot.queue_read_oversized_batches_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_event_capture_memory_budget_bytes",
|
||||
"Process-wide diagnostic JSON heap estimate budget for retained usage events.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.event_capture_memory_budget_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_dlq_encoding_budget_bytes",
|
||||
"Process-wide logical raw string and JSON reservation limit for dead letter encoding; excludes Redis buffers.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.dlq_encoding_budget_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_dlq_encoding_max_jobs",
|
||||
"Maximum admitted dead letter encoding and write jobs per process.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.dlq_encoding_max_jobs as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_dlq_encoding_reserved_bytes",
|
||||
"Logical raw string and JSON bytes reserved by admitted dead letter jobs.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.dlq_encoding_reserved_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_dlq_encoding_active_jobs",
|
||||
"Admitted dead letter jobs awaiting or performing encoding or queue writes.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.dlq_encoding_active_jobs as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_dlq_encoding_capacity_rejected_total",
|
||||
"Dead letter attempts rejected because encoding byte or job capacity was occupied; source remains pending.",
|
||||
MetricKind::Counter,
|
||||
snapshot.dlq_encoding_capacity_rejected_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_dlq_encoding_oversized_rejected_total",
|
||||
"Dead letter attempts rejected because the conservative encoding reservation exceeded the total budget or overflowed.",
|
||||
MetricKind::Counter,
|
||||
snapshot.dlq_encoding_oversized_rejected_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_dlq_encoding_encoded_total",
|
||||
"Completed dead letter JSON encodings, including repeated attempts; not successful archives.",
|
||||
MetricKind::Counter,
|
||||
snapshot.dlq_encoding_encoded_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_event_capture_memory_retained_bytes",
|
||||
"Estimated diagnostic JSON heap retained by budgeted usage events and their clones.",
|
||||
MetricKind::Gauge,
|
||||
snapshot.event_capture_memory_retained_bytes as u64,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_event_capture_memory_downgraded_total",
|
||||
"Usage event diagnostic captures omitted after memory budget exhaustion.",
|
||||
MetricKind::Counter,
|
||||
snapshot.event_capture_memory_downgraded_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_terminal_submission_limit",
|
||||
"Maximum concurrent end-to-end terminal usage submissions.",
|
||||
@@ -3700,6 +3954,12 @@ fn usage_runtime_metric_samples(
|
||||
MetricKind::Counter,
|
||||
snapshot.enqueue_retry_failed_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_enqueue_retry_permanent_failure_total",
|
||||
"Usage enqueue retry submissions rejected or queued events terminated because of permanent input errors.",
|
||||
MetricKind::Counter,
|
||||
snapshot.enqueue_retry_permanent_failure_total,
|
||||
),
|
||||
MetricSample::new(
|
||||
"usage_runtime_enqueue_retry_closed_or_unavailable_total",
|
||||
"Total usage events rejected because the local enqueue dispatcher was full, closed, or unavailable.",
|
||||
@@ -3972,6 +4232,126 @@ mod tests {
|
||||
assert_eq!(database_bounded_auth_load_limit(Some(64), None), Some(64));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_connection_budget_is_optional_and_shared_by_app_state_clones() {
|
||||
let state = AppState::new().expect("app state should build");
|
||||
assert!(state.http_connection_budget.is_none());
|
||||
let budget = Arc::new(super::HttpConnectionBudget::new(1));
|
||||
let state = state.with_http_connection_budget(Arc::clone(&budget));
|
||||
let cloned = state.clone();
|
||||
assert!(Arc::ptr_eq(
|
||||
cloned.http_connection_budget.as_ref().unwrap(),
|
||||
&budget,
|
||||
));
|
||||
|
||||
let connection = budget.try_admit(()).expect("first connection admitted");
|
||||
assert!(cloned
|
||||
.http_connection_budget
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.try_admit(())
|
||||
.is_err());
|
||||
let samples = state.collect_metric_samples().await;
|
||||
for (name, expected) in [
|
||||
("gateway_http_connections_limit", 1),
|
||||
("gateway_http_connections_in_flight", 1),
|
||||
("gateway_http_connections_high_watermark", 1),
|
||||
("gateway_http_connections_rejected_total", 1),
|
||||
("gateway_http_connections_accept_errors_total", 0),
|
||||
] {
|
||||
assert_eq!(
|
||||
samples
|
||||
.iter()
|
||||
.find(|sample| sample.name == name)
|
||||
.unwrap()
|
||||
.value,
|
||||
expected,
|
||||
);
|
||||
}
|
||||
drop(connection);
|
||||
assert_eq!(budget.snapshot().in_flight, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_connection_metrics_export_all_budget_counters() {
|
||||
let samples = super::http_connection_metric_samples(&super::HttpConnectionBudgetSnapshot {
|
||||
limit: 4096,
|
||||
in_flight: 11,
|
||||
high_watermark: 30,
|
||||
rejected_total: 7,
|
||||
accept_errors_total: 3,
|
||||
});
|
||||
for (name, kind, value) in [
|
||||
("gateway_http_connections_limit", MetricKind::Gauge, 4096),
|
||||
("gateway_http_connections_in_flight", MetricKind::Gauge, 11),
|
||||
(
|
||||
"gateway_http_connections_high_watermark",
|
||||
MetricKind::Gauge,
|
||||
30,
|
||||
),
|
||||
(
|
||||
"gateway_http_connections_rejected_total",
|
||||
MetricKind::Counter,
|
||||
7,
|
||||
),
|
||||
(
|
||||
"gateway_http_connections_accept_errors_total",
|
||||
MetricKind::Counter,
|
||||
3,
|
||||
),
|
||||
] {
|
||||
let matching = samples
|
||||
.iter()
|
||||
.filter(|sample| sample.name == name)
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(matching.len(), 1);
|
||||
assert_eq!((matching[0].kind, matching[0].value), (kind, value));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_runtime_metrics_export_queue_payload_limit_and_attempt_counters() {
|
||||
let mut snapshot = crate::usage::UsageRuntimeMetricsSnapshot::default();
|
||||
snapshot.queue_payload_max_bytes = 1024 * 1024;
|
||||
snapshot.queue_payload_downgraded_total = 11;
|
||||
snapshot.queue_payload_rejected_total = 3;
|
||||
snapshot.enqueue_retry_permanent_failure_total = 2;
|
||||
let samples = usage_runtime_metric_samples(&snapshot);
|
||||
for (name, kind, value) in [
|
||||
(
|
||||
"usage_runtime_queue_payload_max_bytes",
|
||||
MetricKind::Gauge,
|
||||
1024 * 1024,
|
||||
),
|
||||
(
|
||||
"usage_runtime_queue_payload_downgraded_total",
|
||||
MetricKind::Counter,
|
||||
11,
|
||||
),
|
||||
(
|
||||
"usage_runtime_queue_payload_rejected_total",
|
||||
MetricKind::Counter,
|
||||
3,
|
||||
),
|
||||
(
|
||||
"usage_runtime_enqueue_retry_permanent_failure_total",
|
||||
MetricKind::Counter,
|
||||
2,
|
||||
),
|
||||
] {
|
||||
let matching = samples
|
||||
.iter()
|
||||
.filter(|sample| sample.name == name)
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
matching.len(),
|
||||
1,
|
||||
"each payload metric must be emitted once"
|
||||
);
|
||||
assert_eq!((matching[0].kind, matching[0].value), (kind, value));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_runtime_metrics_export_first_byte_batch_counters() {
|
||||
let mut snapshot = crate::usage::UsageRuntimeMetricsSnapshot::default();
|
||||
|
||||
@@ -671,7 +671,7 @@ impl SchedulerRuntimeState for AppState {
|
||||
&self,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StoredRequestCandidate>, GatewayError> {
|
||||
AppState::read_recent_request_candidates(self, limit).await
|
||||
AppState::read_recent_runtime_request_candidates(self, limit).await
|
||||
}
|
||||
|
||||
fn provider_key_rpm_reset_at(&self, key_id: &str, now_unix_secs: u64) -> Option<u64> {
|
||||
|
||||
@@ -196,11 +196,12 @@ impl AppState {
|
||||
{
|
||||
let session = session.into();
|
||||
#[cfg(test)]
|
||||
if self.auth_session_store.is_some() && self.auth_user_store.is_some() {
|
||||
if let (Some(session_store), Some(user_store)) = (
|
||||
self.auth_session_store.as_ref(),
|
||||
self.auth_user_store.as_ref(),
|
||||
) {
|
||||
let existing = {
|
||||
self.auth_user_store
|
||||
.as_ref()
|
||||
.expect("checked auth user store")
|
||||
user_store
|
||||
.lock()
|
||||
.expect("auth user store should lock")
|
||||
.get(&session.user_id)
|
||||
@@ -217,12 +218,7 @@ impl AppState {
|
||||
let Some(existing) = existing else {
|
||||
return Ok(None);
|
||||
};
|
||||
let mut users = self
|
||||
.auth_user_store
|
||||
.as_ref()
|
||||
.expect("checked auth user store")
|
||||
.lock()
|
||||
.expect("auth user store should lock");
|
||||
let mut users = user_store.lock().expect("auth user store should lock");
|
||||
let user = users.entry(session.user_id.clone()).or_insert(existing);
|
||||
if user.password_hash.as_deref() != Some(expected_password_hash)
|
||||
|| !user.auth_source.eq_ignore_ascii_case("local")
|
||||
@@ -238,10 +234,7 @@ impl AppState {
|
||||
.or(session.last_seen_at)
|
||||
.unwrap_or_else(chrono::Utc::now);
|
||||
user.last_login_at = Some(now);
|
||||
let mut sessions = self
|
||||
.auth_session_store
|
||||
.as_ref()
|
||||
.expect("checked auth session store")
|
||||
let mut sessions = session_store
|
||||
.lock()
|
||||
.expect("auth session store should lock");
|
||||
for existing in sessions.values_mut() {
|
||||
|
||||
@@ -802,7 +802,7 @@ impl AppState {
|
||||
}
|
||||
return Ok(Some(LdapAuthProvisioningResult {
|
||||
user,
|
||||
owned_wallet_id: initialized.created.then(|| initialized.wallet.id),
|
||||
owned_wallet_id: initialized.created.then_some(initialized.wallet.id),
|
||||
}));
|
||||
}
|
||||
|
||||
|
||||
@@ -109,6 +109,16 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn read_recent_runtime_request_candidates(
|
||||
&self,
|
||||
limit: usize,
|
||||
) -> Result<Vec<candidates::StoredRequestCandidate>, GatewayError> {
|
||||
self.data
|
||||
.list_recent_runtime_request_candidates(limit)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn upsert_request_candidate(
|
||||
&self,
|
||||
mut candidate: candidates::UpsertRequestCandidateRecord,
|
||||
|
||||
@@ -19,6 +19,17 @@ use sha2::{Digest, Sha256};
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::AppState;
|
||||
|
||||
pub use crate::tunnel::build_tunnel_pressure_router;
|
||||
|
||||
pub async fn gateway_metric_samples(
|
||||
state: &AppState,
|
||||
) -> Result<Vec<aether_runtime::MetricSample>, String> {
|
||||
if !state.prewarm_metric_snapshot().await {
|
||||
return Err("gateway harness metric refresh timed out".to_string());
|
||||
}
|
||||
Ok(state.metric_samples().await)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct OpenAiChatPressureTarget {
|
||||
pub base_url: String,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use super::{
|
||||
any, build_router_with_state, build_state_with_execution_runtime_override, json, start_server,
|
||||
Arc, Body, Bytes, HeaderValue, Infallible, Json, Mutex, Request, Response, Router, StatusCode,
|
||||
TRACE_ID_HEADER,
|
||||
LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER,
|
||||
};
|
||||
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
@@ -55,6 +55,24 @@ fn hash_api_key(value: &str) -> String {
|
||||
format!("{:x}", hasher.finalize())
|
||||
}
|
||||
|
||||
async fn build_cancelling_gateway(state: crate::AppState) -> Router {
|
||||
state
|
||||
.data
|
||||
.update_routing_group(
|
||||
"system-default",
|
||||
aether_data_contracts::repository::routing_profiles::UpdateRoutingGroupRecord {
|
||||
config_json: Some(json!({"default_policy": {"cancel_on_client_disconnect": true}})),
|
||||
version: Some(2),
|
||||
updated_at: 2,
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("routing policy should update")
|
||||
.expect("default strategy should exist");
|
||||
build_router_with_state(state)
|
||||
}
|
||||
|
||||
fn sample_local_openai_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
|
||||
StoredAuthApiKeySnapshot::new(
|
||||
user_id.to_string(),
|
||||
@@ -384,7 +402,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
|
||||
vec![sample_local_openai_key()],
|
||||
));
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let gateway = build_router_with_state(
|
||||
let gateway = build_cancelling_gateway(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||
@@ -395,7 +413,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
),
|
||||
);
|
||||
).await;
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
@@ -467,7 +485,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
|
||||
vec![sample_local_openai_key()],
|
||||
));
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let gateway = build_router_with_state(
|
||||
let gateway = build_cancelling_gateway(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||
@@ -478,7 +496,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
),
|
||||
);
|
||||
).await;
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let request = reqwest::Client::new()
|
||||
@@ -608,7 +626,7 @@ async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error_
|
||||
auth_repository,
|
||||
candidate_selection_repository,
|
||||
provider_catalog_repository,
|
||||
request_candidate_repository,
|
||||
Arc::clone(&request_candidate_repository),
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
),
|
||||
@@ -631,17 +649,41 @@ async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error_
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(http::header::CONTENT_TYPE)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("text/event-stream")
|
||||
Some("application/json")
|
||||
);
|
||||
let body_text = response.text().await.expect("response body should read");
|
||||
assert!(body_text.contains("\"rate_limit_error\""));
|
||||
assert!(body_text.contains("\"slow down\""));
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("execution_runtime_candidates_exhausted")
|
||||
);
|
||||
let body_json: serde_json::Value = response.json().await.expect("response body should parse");
|
||||
assert_eq!(body_json["error"]["type"], "http_error");
|
||||
let stored_candidates = request_candidate_repository
|
||||
.list_by_request_id("trace-openai-chat-stream-prefetch-error-123")
|
||||
.await
|
||||
.expect("request candidate trace should read");
|
||||
let failed_candidate = stored_candidates
|
||||
.iter()
|
||||
.find(|candidate| candidate.status == RequestCandidateStatus::Failed)
|
||||
.expect("prefetched error should mark the attempted candidate as failed");
|
||||
assert!(stored_candidates
|
||||
.iter()
|
||||
.all(|candidate| candidate.status != RequestCandidateStatus::Success));
|
||||
assert_eq!(failed_candidate.status_code, Some(429));
|
||||
assert_eq!(
|
||||
failed_candidate.error_type.as_deref(),
|
||||
Some("rate_limit_error")
|
||||
);
|
||||
assert_eq!(failed_candidate.error_message.as_deref(), Some("slow down"));
|
||||
assert!(failed_candidate.finished_at_unix_ms.is_some());
|
||||
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
|
||||
@@ -349,7 +349,7 @@ async fn gateway_executes_gemini_chat_stream_via_local_decision_gate_with_local_
|
||||
});
|
||||
let frames = concat!(
|
||||
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
|
||||
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":33,\"upstream_bytes\":26}}}\n",
|
||||
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
|
||||
);
|
||||
@@ -419,7 +419,7 @@ async fn gateway_executes_gemini_chat_stream_via_local_decision_gate_with_local_
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
|
||||
"data: {\"candidates\":[]}\n\n"
|
||||
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
|
||||
);
|
||||
|
||||
let seen_execution_runtime_request = seen_execution_runtime
|
||||
|
||||
@@ -326,7 +326,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_with_local_s
|
||||
});
|
||||
let frames = concat!(
|
||||
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
|
||||
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
|
||||
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
|
||||
);
|
||||
@@ -396,7 +396,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_with_local_s
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
|
||||
"data: {\"candidates\":[]}\n\n"
|
||||
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
|
||||
);
|
||||
|
||||
let seen_execution_runtime_request = seen_execution_runtime
|
||||
@@ -846,7 +846,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_
|
||||
});
|
||||
let frames = concat!(
|
||||
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"response\\\":{\\\"candidates\\\":[]},\\\"remainingCredits\\\":42,\\\"consumedCredits\\\":1,\\\"traceId\\\":\\\"trace-upstream-1\\\"}\\n\\n\"}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"response\\\":{\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]},\\\"remainingCredits\\\":42,\\\"consumedCredits\\\":1,\\\"traceId\\\":\\\"trace-upstream-1\\\"}\\n\\n\"}}\n",
|
||||
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
|
||||
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
|
||||
);
|
||||
@@ -934,7 +934,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
|
||||
"data: {\"candidates\":[]}\n\n"
|
||||
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
|
||||
);
|
||||
|
||||
let seen_refresh_request = seen_refresh
|
||||
@@ -1354,7 +1354,7 @@ async fn gateway_executes_vertex_ai_gemini_cli_stream_via_local_decision_gate_wi
|
||||
});
|
||||
let frames = concat!(
|
||||
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
|
||||
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
|
||||
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
|
||||
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
|
||||
);
|
||||
@@ -1422,7 +1422,7 @@ async fn gateway_executes_vertex_ai_gemini_cli_stream_via_local_decision_gate_wi
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
|
||||
"data: {\"candidates\":[]}\n\n"
|
||||
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
|
||||
);
|
||||
|
||||
let seen_execution_runtime_request = seen_execution_runtime
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user