mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-09 20:50:20 +08:00
Merge origin/main into fix/gemini-cli-v1internal
This commit is contained in:
@@ -57,6 +57,38 @@ ADMIN_USERNAME=admin123456
|
||||
# docker compose 下 app 启动前自动执行 pending migration/backfill(默认 true)
|
||||
# AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true
|
||||
|
||||
# 管理后台更新策略:
|
||||
# - systemd/二进制部署使用 self:下载 GitHub Release 包,校验 SHA256 后切换 current 并重启。
|
||||
# - Docker Compose 使用 docker:后台只提示版本,实际更新请在 compose 目录执行 ./update.sh。
|
||||
# - 源码/本地构建使用 manual:手动拉取源码或下载 release。
|
||||
# Compose 默认把持久化文件放在 ./datas/{postgres,mysql,sqlite,redis},日志放在 ./logs。
|
||||
# 分布式/多节点部署不要使用 ./datas 作为共享数据目录;应使用外部共享 Postgres/MySQL 和 Redis。
|
||||
# 多节点不要从管理后台一键更新单个节点,应使用镜像滚动更新、systemd 分批发布或外部编排。
|
||||
# AETHER_BASE_DIR=/opt/aether
|
||||
# AETHER_UPDATE_STRATEGY=docker
|
||||
# AETHER_DOCKER_UPDATE_COMMAND=./update.sh
|
||||
# AETHER_GATEWAY_DEPLOYMENT_TOPOLOGY=single-node
|
||||
# AETHER_GATEWAY_NODE_ROLE=all
|
||||
# Docker Compose 默认强制把应用日志输出到 stdout/stderr,并由 Docker 轮转日志。
|
||||
# 如需文件日志,需要在 compose 里把 AETHER_LOG_DESTINATION 改成 file 或 both,
|
||||
# 并把容器用户可写目录挂载到 /opt/aether/logs。
|
||||
# AETHER_LOG_DESTINATION=stdout
|
||||
# AETHER_LOG_FORMAT=pretty
|
||||
# AETHER_LOG_DIR=/opt/aether/logs
|
||||
# 服务器访问 GitHub 需要代理时可配置;也兼容 UPDATE_PROXY_URL / HTTPS_PROXY / ALL_PROXY / HTTP_PROXY。
|
||||
# 如果 Aether 跑在 Docker 容器里,想走宿主机代理时请写 host.docker.internal,不要写 127.0.0.1。
|
||||
# AETHER_UPDATE_PROXY_URL=http://host.docker.internal:7890
|
||||
# 共享出口触发 GitHub API 限流时可配置只读 token;也兼容 GITHUB_TOKEN / GH_TOKEN。
|
||||
# AETHER_UPDATE_GITHUB_TOKEN=
|
||||
# 下载超时控制:总超时默认 600 秒;连续无响应/无数据默认 30 秒。
|
||||
# AETHER_UPDATE_DOWNLOAD_TIMEOUT_SECS=600
|
||||
# AETHER_UPDATE_DOWNLOAD_IDLE_TIMEOUT_SECS=30
|
||||
# 本地联调后台在线更新(配合 docker-compose.release-local.yml):
|
||||
# 会用当前源码构建 release-layout 测试镜像,并伪装成较旧版本以触发升级入口。
|
||||
# AETHER_RELEASE_LOCAL_VERSION=v0.7.0
|
||||
# AETHER_RELEASE_LOCAL_PORT=18085
|
||||
# LOCAL_RELEASE_APP_IMAGE=aether-app:release-local
|
||||
|
||||
# PostgreSQL 连接池配置(默认适合单实例/小型部署;高并发可按需调大)
|
||||
# 推荐计算方式(单实例):
|
||||
# MAX = CPU 核数 × 10(AI 网关偏 IO 等待,可激进些;纯 OLTP 用 × 4)
|
||||
@@ -69,6 +101,7 @@ ADMIN_USERNAME=admin123456
|
||||
|
||||
# PostgreSQL 性能调优(默认值适合 2核4GB 机器,按实际配置覆盖)
|
||||
# 参考:shared_buffers ≈ 可用内存 25%,effective_cache_size ≈ 可用内存 50-75%
|
||||
# POSTGRES_SHM_SIZE 控制 Docker 容器 /dev/shm;仪表盘统计等并行查询会使用它。
|
||||
# work_mem 是每个连接每个排序操作的内存,不要设太大(并发数 × work_mem 是实际占用)
|
||||
# | 系统内存 | shared_buffers | effective_cache_size | work_mem |
|
||||
# | 2GB | 256MB | 768MB | 4MB |
|
||||
@@ -78,5 +111,6 @@ ADMIN_USERNAME=admin123456
|
||||
# | 32GB+ | 8GB | 24GB | 32MB |
|
||||
# POSTGRES_SHARED_BUFFERS=1GB
|
||||
# POSTGRES_EFFECTIVE_CACHE_SIZE=3GB
|
||||
# POSTGRES_SHM_SIZE=512mb
|
||||
# POSTGRES_WORK_MEM=16MB
|
||||
# POSTGRES_MAINTENANCE_WORK_MEM=256MB
|
||||
|
||||
@@ -146,6 +146,7 @@ jobs:
|
||||
- name: Build
|
||||
env:
|
||||
AETHER_VERSION: ${{ needs.preflight.outputs.version_tag }}
|
||||
AETHER_BUILD_TYPE: release
|
||||
CARGO_TERM_COLOR: always
|
||||
shell: bash
|
||||
run: |
|
||||
@@ -260,8 +261,7 @@ jobs:
|
||||
root="package/${bundle}"
|
||||
mkdir -p \
|
||||
"${root}/bin" \
|
||||
"${root}/frontend" \
|
||||
"${root}/scripts"
|
||||
"${root}/frontend"
|
||||
|
||||
install -m 0755 "artifacts/aether-gateway-${platform}-${arch}/aether-gateway" "${root}/bin/aether-gateway"
|
||||
cp -R artifacts/frontend-dist/. "${root}/frontend/"
|
||||
@@ -270,10 +270,9 @@ jobs:
|
||||
-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 0755 scripts/migrate-pg-compose-to-single-node.sh "${root}/scripts/migrate-pg-compose-to-single-node.sh"
|
||||
install -m 0755 scripts/migrate-pg-to-single-node.sh "${root}/scripts/migrate-pg-to-single-node.sh"
|
||||
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"
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
# Created by https://www.toptal.com/developers/gitignore/api/python
|
||||
# Edit at https://www.toptal.com/developers/gitignore?templates=python
|
||||
|
||||
*.rsa
|
||||
|
||||
# AI Assistant Configuration
|
||||
.codex/
|
||||
.claude/
|
||||
|
||||
Generated
+107
-5
@@ -123,10 +123,14 @@ version = "0.1.0"
|
||||
name = "aether-contracts"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"flate2",
|
||||
"hmac",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2",
|
||||
"thiserror 2.0.18",
|
||||
]
|
||||
|
||||
@@ -256,6 +260,7 @@ dependencies = [
|
||||
"http",
|
||||
"ldap3",
|
||||
"md-5",
|
||||
"object_store",
|
||||
"parking_lot",
|
||||
"regex",
|
||||
"reqwest",
|
||||
@@ -266,6 +271,7 @@ dependencies = [
|
||||
"sha1",
|
||||
"sha2",
|
||||
"sqlx",
|
||||
"tar",
|
||||
"thiserror 2.0.18",
|
||||
"tikv-jemallocator",
|
||||
"tokio",
|
||||
@@ -427,6 +433,7 @@ dependencies = [
|
||||
"aether-contracts",
|
||||
"aether-data-contracts",
|
||||
"aether-wallet",
|
||||
"chrono",
|
||||
"regex",
|
||||
"serde",
|
||||
"serde_json",
|
||||
@@ -473,7 +480,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "aether-tunnel"
|
||||
version = "0.3.12"
|
||||
version = "0.3.13"
|
||||
dependencies = [
|
||||
"aether-contracts",
|
||||
"aether-gateway",
|
||||
@@ -511,6 +518,7 @@ dependencies = [
|
||||
"tower-service",
|
||||
"tracing",
|
||||
"url",
|
||||
"uuid",
|
||||
"webpki-roots 0.26.11",
|
||||
]
|
||||
|
||||
@@ -1274,6 +1282,16 @@ dependencies = [
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "core-foundation"
|
||||
version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b2a6cd9ae233e7f62ba4e9353e81a88df7fc8a5987b8d445b4d90c879bd156f6"
|
||||
dependencies = [
|
||||
"core-foundation-sys",
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "core-foundation-sys"
|
||||
version = "0.8.7"
|
||||
@@ -2120,6 +2138,12 @@ version = "1.0.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9"
|
||||
|
||||
[[package]]
|
||||
name = "humantime"
|
||||
version = "2.3.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "135b12329e5e3ce057a9f972339ea52bc954fe1e9358ef27f95e89716fbc5424"
|
||||
|
||||
[[package]]
|
||||
name = "hyper"
|
||||
version = "1.8.1"
|
||||
@@ -2153,6 +2177,7 @@ dependencies = [
|
||||
"hyper",
|
||||
"hyper-util",
|
||||
"rustls 0.23.37",
|
||||
"rustls-native-certs 0.8.3",
|
||||
"rustls-pki-types",
|
||||
"tokio",
|
||||
"tokio-rustls 0.26.4",
|
||||
@@ -2484,7 +2509,7 @@ dependencies = [
|
||||
"percent-encoding",
|
||||
"ring 0.16.20",
|
||||
"rustls 0.21.12",
|
||||
"rustls-native-certs",
|
||||
"rustls-native-certs 0.6.3",
|
||||
"thiserror 1.0.69",
|
||||
"tokio",
|
||||
"tokio-rustls 0.24.1",
|
||||
@@ -2831,6 +2856,41 @@ dependencies = [
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "object_store"
|
||||
version = "0.12.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fbfbfff40aeccab00ec8a910b57ca8ecf4319b335c542f2edcd19dd25a1e2a00"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"chrono",
|
||||
"form_urlencoded",
|
||||
"futures",
|
||||
"http",
|
||||
"http-body-util",
|
||||
"humantime",
|
||||
"hyper",
|
||||
"itertools 0.14.0",
|
||||
"md-5",
|
||||
"parking_lot",
|
||||
"percent-encoding",
|
||||
"quick-xml",
|
||||
"rand 0.9.2",
|
||||
"reqwest",
|
||||
"ring 0.17.14",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serde_urlencoded",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"url",
|
||||
"wasm-bindgen-futures",
|
||||
"web-time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "oid-registry"
|
||||
version = "0.6.1"
|
||||
@@ -2875,6 +2935,12 @@ version = "0.1.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d05e27ee213611ffe7d6348b942e8f942b37114c00cc03cec254295a4a17852e"
|
||||
|
||||
[[package]]
|
||||
name = "openssl-probe"
|
||||
version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe"
|
||||
|
||||
[[package]]
|
||||
name = "ordered-float"
|
||||
version = "4.6.0"
|
||||
@@ -3157,6 +3223,16 @@ dependencies = [
|
||||
"unicode-ident",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "quick-xml"
|
||||
version = "0.38.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b66c2058c55a409d601666cffe35f04333cf1013010882cec174a7467cd4e21c"
|
||||
dependencies = [
|
||||
"memchr",
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "quinn"
|
||||
version = "0.11.9"
|
||||
@@ -3490,6 +3566,7 @@ dependencies = [
|
||||
"pin-project-lite",
|
||||
"quinn",
|
||||
"rustls 0.23.37",
|
||||
"rustls-native-certs 0.8.3",
|
||||
"rustls-pki-types",
|
||||
"serde",
|
||||
"serde_json",
|
||||
@@ -3642,10 +3719,22 @@ version = "0.6.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a9aace74cb666635c918e9c12bc0d348266037aa8eb599b5cba565709a8dff00"
|
||||
dependencies = [
|
||||
"openssl-probe",
|
||||
"openssl-probe 0.1.6",
|
||||
"rustls-pemfile",
|
||||
"schannel",
|
||||
"security-framework",
|
||||
"security-framework 2.11.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustls-native-certs"
|
||||
version = "0.8.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "612460d5f7bea540c490b2b6395d8e34a953e52b491accd6c86c8164c5932a63"
|
||||
dependencies = [
|
||||
"openssl-probe 0.2.1",
|
||||
"rustls-pki-types",
|
||||
"schannel",
|
||||
"security-framework 3.7.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -3744,7 +3833,20 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "897b2245f0b511c87893af39b033e5ca9cce68824c4d7e7630b5a1d339658d02"
|
||||
dependencies = [
|
||||
"bitflags 2.11.0",
|
||||
"core-foundation",
|
||||
"core-foundation 0.9.4",
|
||||
"core-foundation-sys",
|
||||
"libc",
|
||||
"security-framework-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "security-framework"
|
||||
version = "3.7.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d"
|
||||
dependencies = [
|
||||
"bitflags 2.11.0",
|
||||
"core-foundation 0.10.1",
|
||||
"core-foundation-sys",
|
||||
"libc",
|
||||
"security-framework-sys",
|
||||
|
||||
@@ -67,6 +67,7 @@ aether-http = { path = "crates/aether-http" }
|
||||
aether-runtime = { path = "crates/aether-runtime" }
|
||||
aether-testkit = { path = "crates/aether-testkit" }
|
||||
aes = "0.8"
|
||||
aes-gcm = "0.10"
|
||||
async-stream = "0.3"
|
||||
async-trait = "0.1"
|
||||
axum = "0.8"
|
||||
@@ -80,6 +81,7 @@ flate2 = "1"
|
||||
futures-util = "0.3"
|
||||
hmac = "0.12"
|
||||
http = "1"
|
||||
object_store = { version = "0.12", default-features = false, features = ["aws"] }
|
||||
pbkdf2 = { version = "0.12", default-features = false, features = ["hmac"] }
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls", "http2", "socks"] }
|
||||
redis = { version = "0.28", default-features = false, features = ["tokio-comp", "script", "streams", "connection-manager"] }
|
||||
@@ -90,6 +92,7 @@ serde = { version = "1", features = ["derive"] }
|
||||
serde_json = { version = "1", features = ["preserve_order"] }
|
||||
serde_path_to_error = "0.1"
|
||||
sha2 = "0.10"
|
||||
tar = "0.4"
|
||||
sqlx = { version = "0.8", default-features = false, features = ["postgres", "mysql", "sqlite", "runtime-tokio-rustls", "chrono"] }
|
||||
thiserror = "2"
|
||||
tokio = { version = "1", features = ["macros", "net", "rt-multi-thread", "signal", "sync", "time"] }
|
||||
|
||||
+27
-15
@@ -1,31 +1,43 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
# Aether Gateway 运行时镜像(交叉编译方案)
|
||||
# 二进制和前端产物均由 CI 预先构建,此 Dockerfile 仅做打包
|
||||
# 用法: docker buildx build --platform linux/amd64,linux/arm64 -f Dockerfile.app .
|
||||
# Aether Gateway runtime image (cross-compilation)
|
||||
# Binary and frontend assets are pre-built by CI; this Dockerfile only packages them.
|
||||
# Usage: docker buildx build --platform linux/amd64,linux/arm64 -f Dockerfile.app .
|
||||
#
|
||||
# 构建上下文中须包含:
|
||||
# dist/aether-gateway-amd64 (x86_64-unknown-linux-musl 交叉编译产物)
|
||||
# dist/aether-gateway-arm64 (aarch64-unknown-linux-musl 交叉编译产物)
|
||||
# dist/frontend/ (npm run build 产物)
|
||||
# Build context must contain:
|
||||
# dist/aether-gateway-amd64 (x86_64-unknown-linux-musl cross-compiled binary)
|
||||
# dist/aether-gateway-arm64 (aarch64-unknown-linux-musl cross-compiled binary)
|
||||
# dist/frontend/ (npm run build output)
|
||||
|
||||
FROM gcr.io/distroless/static-debian12
|
||||
# --- layout stage: create /opt/aether directory structure with symlink ---
|
||||
# distroless has no shell, so we use busybox to set up the symlink.
|
||||
FROM busybox:1.37-musl AS layout
|
||||
|
||||
# TARGETARCH 由 buildx 自动注入: amd64 或 arm64
|
||||
ARG TARGETARCH
|
||||
|
||||
COPY dist/aether-gateway-${TARGETARCH} /usr/local/bin/aether-gateway
|
||||
COPY dist/frontend/ /srv/frontend
|
||||
RUN mkdir -p /opt/aether/releases/image/bin /opt/aether/releases/image/frontend /opt/aether/logs
|
||||
|
||||
WORKDIR /app
|
||||
COPY dist/aether-gateway-${TARGETARCH} /opt/aether/releases/image/bin/aether-gateway
|
||||
RUN chmod 0755 /opt/aether/releases/image/bin/aether-gateway
|
||||
COPY dist/frontend/ /opt/aether/releases/image/frontend/
|
||||
|
||||
RUN ln -s /opt/aether/releases/image /opt/aether/current
|
||||
|
||||
# --- final stage: distroless runtime ---
|
||||
FROM gcr.io/distroless/static-debian12
|
||||
|
||||
COPY --from=layout /opt/aether /opt/aether
|
||||
|
||||
WORKDIR /opt/aether
|
||||
|
||||
ENV RUST_LOG=aether_gateway=info \
|
||||
APP_PORT=8084 \
|
||||
AETHER_GATEWAY_STATIC_DIR=/srv/frontend
|
||||
AETHER_UPDATE_STRATEGY=docker \
|
||||
AETHER_GATEWAY_STATIC_DIR=/opt/aether/current/frontend
|
||||
|
||||
EXPOSE 8084
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD ["/usr/local/bin/aether-gateway", "--healthcheck"]
|
||||
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
||||
|
||||
USER root
|
||||
ENTRYPOINT ["/usr/local/bin/aether-gateway"]
|
||||
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
# syntax=docker.m.daocloud.io/docker/dockerfile:1
|
||||
# Aether 运行镜像:Rust gateway 直接服务 API + 前端静态文件(国内镜像源版本)
|
||||
# 构建命令: docker build --build-arg AETHER_BUILD_VERSION=v0.7.2 -f Dockerfile.app.local -t aether-app:latest .
|
||||
|
||||
ARG RUST_VERSION=1.95.0
|
||||
ARG NODE_BASE_IMAGE=docker.m.daocloud.io/library/node:22-slim
|
||||
ARG RUST_BASE_IMAGE=docker.m.daocloud.io/library/rust:${RUST_VERSION}-slim
|
||||
|
||||
# ==================== 前端构建 ====================
|
||||
FROM node:22-slim AS frontend-builder
|
||||
FROM ${NODE_BASE_IMAGE} AS frontend-builder
|
||||
ARG AETHER_BUILD_VERSION
|
||||
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
|
||||
AETHER_VERSION=${AETHER_BUILD_VERSION}
|
||||
@@ -18,7 +20,7 @@ COPY frontend/ ./
|
||||
RUN npm run build
|
||||
|
||||
# ==================== Rust gateway 构建 ====================
|
||||
FROM rust:${RUST_VERSION}-slim AS gateway-base
|
||||
FROM ${RUST_BASE_IMAGE} AS gateway-base
|
||||
WORKDIR /build
|
||||
|
||||
# 本地镜像优先缩短构建时间,保留 release 语义,但改用更快的 thin LTO。
|
||||
@@ -134,6 +136,7 @@ ENV LANG=C.UTF-8 \
|
||||
LC_ALL=C.UTF-8 \
|
||||
RUST_LOG=aether_gateway=info \
|
||||
APP_PORT=8084 \
|
||||
AETHER_UPDATE_STRATEGY=manual \
|
||||
AETHER_GATEWAY_STATIC_DIR=/srv/frontend
|
||||
|
||||
EXPOSE 8084
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
# syntax=docker.m.daocloud.io/docker/dockerfile:1
|
||||
# Aether 本地发布版联调镜像
|
||||
# 作用:用当前源码构建一个 release-layout 容器,专门测试管理后台在线更新流程。
|
||||
|
||||
ARG RUST_VERSION=1.95.0
|
||||
ARG NODE_BASE_IMAGE=docker.m.daocloud.io/library/node:22-slim
|
||||
ARG RUST_BASE_IMAGE=docker.m.daocloud.io/library/rust:${RUST_VERSION}-slim
|
||||
|
||||
# ==================== 前端构建 ====================
|
||||
FROM ${NODE_BASE_IMAGE} AS frontend-builder
|
||||
ARG AETHER_BUILD_VERSION
|
||||
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
|
||||
AETHER_VERSION=${AETHER_BUILD_VERSION}
|
||||
WORKDIR /app/frontend
|
||||
COPY frontend/package*.json ./
|
||||
RUN --mount=type=cache,id=aether-npm-cache,target=/root/.npm,sharing=locked \
|
||||
npm config set registry https://registry.npmmirror.com && \
|
||||
npm ci --no-audit --no-fund
|
||||
COPY frontend/ ./
|
||||
RUN npm run build
|
||||
|
||||
# ==================== Rust gateway 构建 ====================
|
||||
FROM ${RUST_BASE_IMAGE} AS gateway-base
|
||||
WORKDIR /build
|
||||
|
||||
ENV CARGO_REGISTRIES_CRATES_IO_PROTOCOL=sparse \
|
||||
CARGO_PROFILE_RELEASE_LTO=thin \
|
||||
CARGO_PROFILE_RELEASE_CODEGEN_UNITS=16
|
||||
|
||||
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
||||
--mount=type=cache,target=/var/lib/apt,sharing=locked \
|
||||
sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
|
||||
apt-get update && apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
ca-certificates \
|
||||
cmake \
|
||||
git \
|
||||
libclang-dev \
|
||||
libssl-dev \
|
||||
pkg-config \
|
||||
perl
|
||||
|
||||
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
|
||||
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
|
||||
cargo install cargo-chef --locked
|
||||
|
||||
FROM gateway-base AS gateway-planner
|
||||
COPY Cargo.toml Cargo.lock ./
|
||||
COPY apps/ ./apps/
|
||||
COPY crates/ ./crates/
|
||||
RUN cargo chef prepare --recipe-path recipe.json
|
||||
|
||||
FROM gateway-base AS gateway-builder
|
||||
ARG AETHER_BUILD_VERSION
|
||||
ARG AETHER_BUILD_TYPE=release
|
||||
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
|
||||
AETHER_VERSION=${AETHER_BUILD_VERSION} \
|
||||
AETHER_BUILD_TYPE=${AETHER_BUILD_TYPE}
|
||||
COPY --from=gateway-planner /build/recipe.json ./recipe.json
|
||||
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
|
||||
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
|
||||
--mount=type=cache,id=aether-cargo-target-release-local,target=/build/target,sharing=locked \
|
||||
cargo chef cook --release --locked --package aether-gateway --bin aether-gateway --recipe-path recipe.json
|
||||
|
||||
COPY Cargo.toml Cargo.lock ./
|
||||
COPY apps/ ./apps/
|
||||
COPY crates/ ./crates/
|
||||
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
|
||||
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \
|
||||
--mount=type=cache,id=aether-cargo-target-release-local,target=/build/target,sharing=locked \
|
||||
cargo build --release --locked -p aether-gateway && \
|
||||
cp target/release/aether-gateway /tmp/aether-gateway
|
||||
|
||||
# ==================== 最小运行时打包 ====================
|
||||
FROM gateway-builder AS runtime-prep
|
||||
RUN set -eux; \
|
||||
mkdir -p \
|
||||
/runtime-root/app/data \
|
||||
/runtime-root/etc \
|
||||
/runtime-root/etc/ssl \
|
||||
/runtime-root/lib \
|
||||
/runtime-root/lib64 \
|
||||
/runtime-root/usr/lib \
|
||||
/runtime-root/opt/aether/logs \
|
||||
/runtime-root/opt/aether/releases/image/bin \
|
||||
/runtime-root/opt/aether/releases/image/frontend; \
|
||||
cp /tmp/aether-gateway /runtime-root/opt/aether/releases/image/bin/aether-gateway; \
|
||||
ln -s /opt/aether/releases/image /runtime-root/opt/aether/current; \
|
||||
: > /tmp/runtime-libs.txt; \
|
||||
: > /tmp/runtime-scan-queue.txt; \
|
||||
printf '%s\n' /tmp/aether-gateway >> /tmp/runtime-scan-queue.txt; \
|
||||
while [ -s /tmp/runtime-scan-queue.txt ]; do \
|
||||
current="$(head -n1 /tmp/runtime-scan-queue.txt)"; \
|
||||
sed -i '1d' /tmp/runtime-scan-queue.txt; \
|
||||
ldd "$current" | awk '/=>/ { print $3 } $1 ~ /^\// { print $1 }' | while read -r lib; do \
|
||||
[ -n "$lib" ]; \
|
||||
if ! grep -Fxq "$lib" /tmp/runtime-libs.txt; then \
|
||||
printf '%s\n' "$lib" >> /tmp/runtime-libs.txt; \
|
||||
printf '%s\n' "$lib" >> /tmp/runtime-scan-queue.txt; \
|
||||
fi; \
|
||||
done; \
|
||||
done; \
|
||||
sort -u /tmp/runtime-libs.txt -o /tmp/runtime-libs.txt; \
|
||||
while read -r lib; do \
|
||||
[ -n "$lib" ]; \
|
||||
dest="/runtime-root$(dirname "$lib")"; \
|
||||
mkdir -p "$dest"; \
|
||||
cp -L "$lib" "$dest/"; \
|
||||
done < /tmp/runtime-libs.txt; \
|
||||
for lib in \
|
||||
/lib/x86_64-linux-gnu/libnss_dns.so.2 \
|
||||
/lib/x86_64-linux-gnu/libnss_files.so.2 \
|
||||
/lib/x86_64-linux-gnu/libresolv.so.2; do \
|
||||
if [ -f "$lib" ]; then \
|
||||
dest="/runtime-root$(dirname "$lib")"; \
|
||||
mkdir -p "$dest"; \
|
||||
cp -L "$lib" "$dest/"; \
|
||||
fi; \
|
||||
done; \
|
||||
cp -a /usr/lib/ssl /runtime-root/usr/lib/; \
|
||||
cp -a /etc/ssl/certs /runtime-root/etc/ssl/; \
|
||||
if [ -f /etc/ssl/openssl.cnf ]; then \
|
||||
cp /etc/ssl/openssl.cnf /runtime-root/etc/ssl/openssl.cnf; \
|
||||
fi; \
|
||||
if [ -f /etc/nsswitch.conf ]; then \
|
||||
cp /etc/nsswitch.conf /runtime-root/etc/nsswitch.conf; \
|
||||
fi
|
||||
COPY --from=frontend-builder /app/frontend/dist /runtime-root/opt/aether/releases/image/frontend
|
||||
|
||||
# ==================== 运行时镜像 ====================
|
||||
FROM scratch
|
||||
|
||||
COPY --from=runtime-prep /runtime-root/ /
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
ENV LANG=C.UTF-8 \
|
||||
LC_ALL=C.UTF-8 \
|
||||
RUST_LOG=aether_gateway=info \
|
||||
APP_PORT=8084 \
|
||||
AETHER_BASE_DIR=/opt/aether \
|
||||
AETHER_UPDATE_STRATEGY=self \
|
||||
AETHER_GATEWAY_STATIC_DIR=/opt/aether/current/frontend
|
||||
|
||||
EXPOSE 8084
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
||||
|
||||
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
||||
@@ -55,10 +55,61 @@ docker compose pull && docker compose up -d
|
||||
docker compose -f docker-compose.single-node.yml pull && docker compose -f docker-compose.single-node.yml up -d
|
||||
```
|
||||
|
||||
### 一键更新
|
||||
|
||||
Docker Compose 部署后,可在部署目录直接执行:
|
||||
|
||||
```bash
|
||||
./update.sh
|
||||
```
|
||||
|
||||
`update.sh` 会拉取最新 `app` 镜像并重建 `app` 容器,Docker named volumes、`./data` 和 `./logs` 不会被删除。Single Node 部署也可显式指定:
|
||||
|
||||
```bash
|
||||
./update.sh --mode single-node
|
||||
```
|
||||
|
||||
仓库自带的 Docker Compose 默认把应用日志输出到容器 `stdout/stderr`,直接用 `docker compose logs -f app` 查看,并由 Docker 轮转日志,避免正式发布镜像切换到非 root 用户后再被宿主机挂载日志目录的权限问题拖垮启动。如果你确实需要文件日志,需要在 compose 里把 `AETHER_LOG_DESTINATION` 改成 `file|both`,并额外挂载一个容器用户可写的目录到 `/opt/aether/logs`。
|
||||
|
||||
管理后台右上角“版本信息”会检测新版本。Docker Compose 部署只提示版本,实际更新继续执行 `./update.sh`;systemd / launchd / 二进制部署才使用后台自更新,流程是下载对应平台的 GitHub Release 包、强制校验 `SHA256SUMS`、解压到 `/opt/aether/releases/<version>`,再切换 `/opt/aether/current` 并退出进程,交给 systemd / launchd 拉起新版本。
|
||||
|
||||
源码或本地构建版本不会启用后台在线更新,请继续使用源码更新流程。Docker Compose 用户如果希望“容器重建后也保持镜像层面的新版本”,仍建议定期运行 `./update.sh` 拉取并重建 app 镜像。服务器访问 GitHub 需要代理时,可设置 `AETHER_UPDATE_PROXY_URL`,也兼容 `UPDATE_PROXY_URL`、`HTTPS_PROXY`、`ALL_PROXY`、`HTTP_PROXY` 以及 `NO_PROXY`。共享出口触发 GitHub API 限流时,可设置只读 `AETHER_UPDATE_GITHUB_TOKEN`,也兼容 `GITHUB_TOKEN` / `GH_TOKEN`。下载总超时默认 600 秒,连续无响应/无数据默认 30 秒,可通过 `AETHER_UPDATE_DOWNLOAD_TIMEOUT_SECS` 和 `AETHER_UPDATE_DOWNLOAD_IDLE_TIMEOUT_SECS` 调整。
|
||||
|
||||
标准 Docker Compose 使用 Docker named volumes 存放 Postgres/Redis/MySQL 数据;Single Node 使用部署目录下的 `./data` 存放 SQLite 数据。
|
||||
|
||||
如果是本地源码构建镜像的部署,继续使用:
|
||||
|
||||
```bash
|
||||
./deploy.sh
|
||||
```
|
||||
|
||||
如果要在本机联调“管理后台在线更新”本身,可启动仓库内置的 release-layout 测试环境:
|
||||
|
||||
```bash
|
||||
docker compose -f docker-compose.release-local.yml up -d --build
|
||||
```
|
||||
|
||||
这套环境会用当前源码构建一个本地测试镜像,但编译为 `release` 类型,并默认伪装成 `v0.7.0`,这样后台会按正式发布版逻辑开放“立即更新”。默认监听 `http://127.0.0.1:18085`,数据目录使用 `./data-release-local`;日志默认走 `docker logs`,不会影响你正在跑的源码构建容器。
|
||||
|
||||
如果这套容器在 `prepare-update` 时访问 GitHub 失败,而你本机是通过代理出网,请在 `.env` 里把 `AETHER_UPDATE_PROXY_URL` 写成宿主机地址,例如 `http://host.docker.internal:7890`;容器内的 `127.0.0.1` 指向容器自身,不是宿主机。
|
||||
|
||||
如果想重置这套联调环境(包括 `/opt/aether/current` 和已下载的历史版本),执行:
|
||||
|
||||
```bash
|
||||
docker compose -f docker-compose.release-local.yml down -v
|
||||
```
|
||||
|
||||
可选变量:
|
||||
|
||||
- `AETHER_RELEASE_LOCAL_VERSION`:本地联调镜像对外声明的当前版本,默认 `v0.7.0`
|
||||
- `AETHER_RELEASE_LOCAL_PORT`:本地联调端口,默认 `18085`
|
||||
- `LOCAL_RELEASE_APP_IMAGE`:本地联调镜像名,默认 `aether-app:release-local`
|
||||
|
||||
### 一键安装(默认 Single Node:Linux systemd / macOS launchd + SQLite)
|
||||
|
||||
```bash
|
||||
cd Aether && cd Aether
|
||||
git clone https://github.com/fawney19/Aether.git
|
||||
cd Aether
|
||||
curl -fsSL https://raw.githubusercontent.com/fawney19/Aether/main/install.sh | sudo bash
|
||||
```
|
||||
|
||||
|
||||
@@ -31,7 +31,7 @@ aether-task-runtime.workspace = true
|
||||
aether-usage-runtime.workspace = true
|
||||
aether-video-tasks-core.workspace = true
|
||||
aether-wallet.workspace = true
|
||||
aes-gcm = "0.10"
|
||||
aes-gcm.workspace = true
|
||||
async-stream.workspace = true
|
||||
async-trait.workspace = true
|
||||
axum = { version = "0.8", features = ["ws"] }
|
||||
@@ -48,6 +48,7 @@ hmac.workspace = true
|
||||
http.workspace = true
|
||||
ldap3 = { version = "0.11", default-features = false, features = ["sync", "tls-rustls"] }
|
||||
md-5 = "0.10"
|
||||
object_store.workspace = true
|
||||
parking_lot = "0.12"
|
||||
regex.workspace = true
|
||||
reqwest.workspace = true
|
||||
@@ -57,6 +58,7 @@ serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
sha1 = "0.10"
|
||||
sha2 = { workspace = true, features = ["oid"] }
|
||||
tar.workspace = true
|
||||
sqlx.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio.workspace = true
|
||||
|
||||
@@ -3,6 +3,7 @@ use std::process::Command;
|
||||
|
||||
fn main() {
|
||||
println!("cargo:rerun-if-env-changed=AETHER_BUILD_VERSION");
|
||||
println!("cargo:rerun-if-env-changed=AETHER_BUILD_TYPE");
|
||||
println!("cargo:rerun-if-env-changed=AETHER_VERSION");
|
||||
println!("cargo:rerun-if-env-changed=GITHUB_REF_NAME");
|
||||
println!("cargo:rerun-if-changed=../../.git/HEAD");
|
||||
@@ -27,6 +28,12 @@ fn main() {
|
||||
.unwrap_or(package_version);
|
||||
|
||||
println!("cargo:rustc-env=AETHER_BUILD_VERSION={version}");
|
||||
|
||||
let build_type = env::var("AETHER_BUILD_TYPE")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| "source".to_string());
|
||||
println!("cargo:rustc-env=AETHER_BUILD_TYPE={build_type}");
|
||||
}
|
||||
|
||||
fn git_describe_version() -> Option<String> {
|
||||
|
||||
@@ -8,6 +8,44 @@ fn utf8(bytes: Vec<u8>) -> String {
|
||||
String::from_utf8(bytes).expect("utf8 should decode")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn same_format_claude_local_stream_rewriter_sanitizes_read_input_json_delta() {
|
||||
let report_context = json!({
|
||||
"provider_api_format": "claude:messages",
|
||||
"client_api_format": "claude:messages",
|
||||
"needs_conversion": false,
|
||||
});
|
||||
let mut rewriter =
|
||||
maybe_build_local_stream_rewriter(Some(&report_context)).expect("rewriter should exist");
|
||||
let mut output = rewriter
|
||||
.push_chunk(
|
||||
b"event: content_block_start\n\
|
||||
data: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"tool_use\",\"id\":\"call_read_1\",\"name\":\"Read\",\"input\":{}}}\n\n",
|
||||
)
|
||||
.expect("start should be accepted");
|
||||
output.extend(
|
||||
rewriter
|
||||
.push_chunk(
|
||||
b"event: content_block_delta\n\
|
||||
data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"input_json_delta\",\"partial_json\":\"{\\\"file_path\\\":\\\"/tmp/a.txt\\\",\\\"pages\\\":\\\"\\\"}\"}}\n\n",
|
||||
)
|
||||
.expect("delta should be accepted"),
|
||||
);
|
||||
output.extend(
|
||||
rewriter
|
||||
.push_chunk(
|
||||
b"event: content_block_stop\n\
|
||||
data: {\"type\":\"content_block_stop\",\"index\":0}\n\n",
|
||||
)
|
||||
.expect("stop should flush sanitized delta"),
|
||||
);
|
||||
|
||||
let output_text = utf8(output);
|
||||
assert!(output_text.contains("\"name\":\"Read\""));
|
||||
assert!(output_text.contains("\\\"file_path\\\":\\\"/tmp/a.txt\\\""));
|
||||
assert!(!output_text.contains("\\\"pages\\\":\\\"\\\""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_sync_bridge_converts_openai_chat_sync_json_to_openai_chat_sse() {
|
||||
let outcome = maybe_bridge_standard_sync_json_to_stream(
|
||||
|
||||
@@ -42,7 +42,7 @@ use crate::clock::current_unix_ms;
|
||||
use crate::dispatch::refs::dispatch_ref_for_local_candidate;
|
||||
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
|
||||
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
|
||||
use crate::scheduler::candidate::API_KEY_CONCURRENCY_LIMIT_SKIP_REASON;
|
||||
use crate::scheduler::candidate::is_auth_api_key_concurrency_limit_skip_reason;
|
||||
use crate::scheduler::config::SchedulerSchedulingMode;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
@@ -988,7 +988,7 @@ fn page_is_exact_auth_api_key_concurrency_limited(
|
||||
&& page
|
||||
.skipped_candidates
|
||||
.iter()
|
||||
.all(|skipped| skipped.skip_reason == API_KEY_CONCURRENCY_LIMIT_SKIP_REASON)
|
||||
.all(|skipped| is_auth_api_key_concurrency_limit_skip_reason(skipped.skip_reason))
|
||||
}
|
||||
|
||||
async fn pop_attempt_from_items(
|
||||
|
||||
@@ -18,6 +18,7 @@ mod passthrough;
|
||||
mod plan_builders;
|
||||
mod pool_scheduler;
|
||||
pub(crate) mod pool_scores;
|
||||
mod redaction;
|
||||
mod report_context;
|
||||
mod route;
|
||||
mod runtime_miss;
|
||||
|
||||
@@ -11,7 +11,8 @@ use crate::ai_serving::planner::materialization_policy::{
|
||||
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
|
||||
};
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, LocalExecutionReportContextParts,
|
||||
build_local_execution_report_context, insert_native_client_envelope_name,
|
||||
LocalExecutionReportContextParts,
|
||||
};
|
||||
use crate::ai_serving::planner::spec_metadata::local_same_format_provider_spec_metadata;
|
||||
use crate::ai_serving::planner::CandidateFailureDiagnostic;
|
||||
@@ -55,10 +56,15 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
let Some(resolved) = resolve_local_same_format_provider_candidate_payload_parts(
|
||||
state, parts, trace_id, body_json, input, &attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let original_request_body_json = if resolved.request_redacted {
|
||||
Some(&resolved.provider_request_body)
|
||||
} else {
|
||||
Some(body_json)
|
||||
};
|
||||
|
||||
let prompt_cache_key = resolved
|
||||
.provider_request_body
|
||||
@@ -90,6 +96,11 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
"envelope_name".to_string(),
|
||||
json!(super::super::ANTIGRAVITY_ENVELOPE_NAME),
|
||||
);
|
||||
insert_native_client_envelope_name(
|
||||
&mut extra_fields,
|
||||
super::super::ANTIGRAVITY_ENVELOPE_NAME,
|
||||
parts.uri.path(),
|
||||
);
|
||||
} else if resolved.is_gemini_cli {
|
||||
extra_fields.insert(
|
||||
"envelope_name".to_string(),
|
||||
@@ -129,7 +140,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
original_request_body_json: Some(body_json),
|
||||
original_request_body_json,
|
||||
original_request_body_base64: None,
|
||||
client_session_affinity: input.client_session_affinity.as_ref(),
|
||||
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
|
||||
@@ -164,6 +175,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
provider_request_headers,
|
||||
provider_request_body,
|
||||
transport_profile: _,
|
||||
request_redacted: _,
|
||||
} = resolved;
|
||||
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
|
||||
@@ -7,6 +7,9 @@ use serde_json::Value;
|
||||
use crate::ai_serving::planner::common::{
|
||||
enforce_provider_body_stream_policy, request_requires_body_stream_field,
|
||||
};
|
||||
use crate::ai_serving::planner::redaction::{
|
||||
request_identity_response_encoding_when_redacted, resolve_provider_chat_pii_redaction,
|
||||
};
|
||||
use crate::ai_serving::transport::antigravity::{
|
||||
build_antigravity_safe_v1internal_request, build_antigravity_static_identity_headers,
|
||||
classify_local_antigravity_request_support, AntigravityEnvelopeRequestType,
|
||||
@@ -16,11 +19,10 @@ use crate::ai_serving::transport::gemini_cli::resolve_gemini_cli_project_id;
|
||||
use crate::ai_serving::transport::{
|
||||
build_gemini_cli_v1internal_request, build_grok_browser_headers, build_grok_upstream_url,
|
||||
build_same_format_provider_headers, GeminiCliRequestEnvelopeSupport, GrokHeaderInput,
|
||||
SameFormatProviderHeadersInput, GEMINI_CLI_USER_AGENT, GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME,
|
||||
GROK_CHAT_PATH,
|
||||
SameFormatProviderHeadersInput, GEMINI_CLI_USER_AGENT, GROK_CHAT_PATH,
|
||||
};
|
||||
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
|
||||
use crate::AppState;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
mod policy;
|
||||
mod prepare;
|
||||
@@ -103,6 +105,7 @@ pub(crate) struct LocalSameFormatProviderCandidatePayloadParts {
|
||||
pub(super) provider_request_headers: BTreeMap<String, String>,
|
||||
pub(super) provider_request_body: Value,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
pub(super) request_redacted: bool,
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
@@ -113,9 +116,9 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
input: &LocalSameFormatProviderDecisionInput,
|
||||
attempt: &LocalSameFormatProviderCandidateAttempt,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
) -> Option<LocalSameFormatProviderCandidatePayloadParts> {
|
||||
) -> Result<Option<LocalSameFormatProviderCandidatePayloadParts>, GatewayError> {
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let prepared = prepare_local_same_format_provider_candidate(
|
||||
let Some(prepared) = prepare_local_same_format_provider_candidate(
|
||||
state,
|
||||
trace_id,
|
||||
input,
|
||||
@@ -124,7 +127,10 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
&attempt.candidate_id,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let enable_model_directives =
|
||||
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
|
||||
state,
|
||||
@@ -133,6 +139,16 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
)
|
||||
.await;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let redaction = resolve_provider_chat_pii_redaction(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
&input.auth_context,
|
||||
spec.api_format,
|
||||
&attempt.candidate_id,
|
||||
)
|
||||
.await?;
|
||||
let body_json = redaction.body_json.as_ref();
|
||||
let mut transport = Arc::clone(&prepared.transport);
|
||||
|
||||
let Some(mut base_provider_request_body) =
|
||||
@@ -170,7 +186,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
if let Some(mapping) =
|
||||
crate::system_features::reasoning_model_directive_mapping_for_api_format_and_model(
|
||||
@@ -216,7 +232,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
"transport_unsupported",
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -246,7 +262,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
"transport_auth_unavailable",
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -280,7 +296,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else if let Some(project_id) = gemini_cli_project_id.as_deref() {
|
||||
@@ -308,7 +324,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -352,7 +368,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let mut extra_headers = antigravity_auth
|
||||
@@ -362,7 +378,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
if prepared.behavior.is_gemini_cli {
|
||||
extra_headers.insert("user-agent".to_string(), GEMINI_CLI_USER_AGENT.to_string());
|
||||
}
|
||||
let Some(provider_request_headers) = (if is_grok {
|
||||
let Some(mut provider_request_headers) = (if is_grok {
|
||||
build_grok_browser_headers(GrokHeaderInput {
|
||||
transport: &transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
@@ -406,10 +422,14 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
redaction.redacted,
|
||||
);
|
||||
|
||||
Some(LocalSameFormatProviderCandidatePayloadParts {
|
||||
Ok(Some(LocalSameFormatProviderCandidatePayloadParts {
|
||||
transport,
|
||||
is_antigravity: prepared.is_antigravity,
|
||||
is_gemini_cli: prepared.behavior.is_gemini_cli,
|
||||
@@ -424,5 +444,6 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
provider_request_headers,
|
||||
provider_request_body,
|
||||
transport_profile,
|
||||
})
|
||||
request_redacted: redaction.redacted,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKe
|
||||
use aether_pool_core::{
|
||||
score_pool_member_with_rules, PoolMemberScoreInput, PoolMemberScoreRules, POOL_SCORE_VERSION,
|
||||
};
|
||||
use aether_scheduler_core::any_provider_key_circuit_open_at;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::handlers::shared::{provider_key_health_summary, provider_key_status_snapshot_payload};
|
||||
@@ -98,7 +99,8 @@ fn provider_key_score_input(
|
||||
.as_object()
|
||||
.and_then(|snapshot| snapshot.get("account"))
|
||||
.and_then(Value::as_object);
|
||||
let (health_score, _, _, any_circuit_open, _) = provider_key_health_summary(key);
|
||||
let (health_score, _, _, _, _) = provider_key_health_summary(key);
|
||||
let active_circuit_open = any_provider_key_circuit_open_at(key, now_unix_secs);
|
||||
let health_score = key
|
||||
.health_by_format
|
||||
.as_ref()
|
||||
@@ -125,7 +127,7 @@ fn provider_key_score_input(
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false),
|
||||
oauth_invalid_reason: key.oauth_invalid_reason.clone(),
|
||||
circuit_open: any_circuit_open,
|
||||
circuit_open: active_circuit_open,
|
||||
success_count: key.success_count.unwrap_or(0).into(),
|
||||
error_count: key.error_count.unwrap_or(0).into(),
|
||||
total_response_time_ms: key.total_response_time_ms.unwrap_or(0).into(),
|
||||
@@ -159,3 +161,70 @@ fn stable_hash(bytes: &[u8]) -> u64 {
|
||||
}
|
||||
hash
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aether_data_contracts::repository::pool_scores::PoolMemberHardState;
|
||||
use serde_json::json;
|
||||
|
||||
fn sample_key_with_circuit_next_probe(
|
||||
next_probe_at_unix_secs: u64,
|
||||
) -> StoredProviderCatalogKey {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key-gemini-5".to_string(),
|
||||
"provider-google-api".to_string(),
|
||||
"5".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("sample key should be valid");
|
||||
key.health_by_format = Some(json!({
|
||||
"gemini:generate_content": {
|
||||
"health_score": 0.2,
|
||||
"consecutive_failures": 8
|
||||
}
|
||||
}));
|
||||
key.circuit_breaker_by_format = Some(json!({
|
||||
"gemini:generate_content": {
|
||||
"open": true,
|
||||
"reason": "consecutive_failures_8",
|
||||
"next_probe_at_unix_secs": next_probe_at_unix_secs
|
||||
}
|
||||
}));
|
||||
key
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn expired_circuit_probe_deadline_does_not_leave_pool_score_in_cooldown() {
|
||||
let now_unix_secs = 1_000;
|
||||
let key = sample_key_with_circuit_next_probe(900);
|
||||
|
||||
let score = build_provider_key_pool_score_upsert(
|
||||
&key,
|
||||
"custom",
|
||||
None,
|
||||
now_unix_secs,
|
||||
PoolMemberScoreRules::default(),
|
||||
);
|
||||
|
||||
assert_eq!(score.hard_state, PoolMemberHardState::Available);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn future_circuit_probe_deadline_keeps_pool_score_in_cooldown() {
|
||||
let now_unix_secs = 1_000;
|
||||
let key = sample_key_with_circuit_next_probe(1_100);
|
||||
|
||||
let score = build_provider_key_pool_score_upsert(
|
||||
&key,
|
||||
"custom",
|
||||
None,
|
||||
now_unix_secs,
|
||||
PoolMemberScoreRules::default(),
|
||||
);
|
||||
|
||||
assert_eq!(score.hard_state, PoolMemberHardState::Cooldown);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,192 @@
|
||||
use std::borrow::Cow;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use serde_json::Value;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::ExecutionRuntimeAuthContext;
|
||||
use crate::privacy::{
|
||||
build_redaction_session_config, read_chat_pii_redaction_runtime_config,
|
||||
try_mask_chat_pii_request_json_with_cache_options, ChatPiiRedactionRequestFormat,
|
||||
MaskChatRequestOptions, RedactionMaskError, RedactionSessionSlot, RedisRedactionMappingCache,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(crate) struct ProviderRequestRedaction<'a> {
|
||||
pub(crate) body_json: Cow<'a, Value>,
|
||||
pub(crate) redacted: bool,
|
||||
}
|
||||
|
||||
impl<'a> ProviderRequestRedaction<'a> {
|
||||
fn disabled(body_json: &'a Value) -> Self {
|
||||
Self {
|
||||
body_json: Cow::Borrowed(body_json),
|
||||
redacted: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
struct ChatPiiRedactionFeatureSettings {
|
||||
enabled: Option<bool>,
|
||||
inject_model_instruction: Option<bool>,
|
||||
}
|
||||
|
||||
impl ChatPiiRedactionFeatureSettings {
|
||||
fn merge_from_value(&mut self, value: Option<&Value>) {
|
||||
let Some(settings) = value
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|features| features.get("chat_pii_redaction"))
|
||||
.and_then(Value::as_object)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if let Some(enabled) = settings.get("enabled").and_then(Value::as_bool) {
|
||||
self.enabled = Some(enabled);
|
||||
}
|
||||
if let Some(inject_model_instruction) = settings
|
||||
.get("inject_model_instruction")
|
||||
.and_then(Value::as_bool)
|
||||
{
|
||||
self.inject_model_instruction = Some(inject_model_instruction);
|
||||
}
|
||||
}
|
||||
|
||||
fn effective_enabled(self) -> bool {
|
||||
self.enabled.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn effective_inject_model_instruction(self) -> bool {
|
||||
self.inject_model_instruction.unwrap_or(true)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn request_identity_response_encoding_when_redacted(
|
||||
headers: &mut std::collections::BTreeMap<String, String>,
|
||||
redacted: bool,
|
||||
) {
|
||||
if redacted {
|
||||
headers.insert("accept-encoding".to_string(), "identity".to_string());
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
body_json: &'a Value,
|
||||
auth_context: &ExecutionRuntimeAuthContext,
|
||||
client_api_format: &str,
|
||||
candidate_id: &str,
|
||||
) -> Result<ProviderRequestRedaction<'a>, GatewayError> {
|
||||
let Some(format) = ChatPiiRedactionRequestFormat::from_api_format(client_api_format) else {
|
||||
return Ok(ProviderRequestRedaction::disabled(body_json));
|
||||
};
|
||||
let Some(slot) = parts.extensions.get::<RedactionSessionSlot>() else {
|
||||
return Ok(ProviderRequestRedaction::disabled(body_json));
|
||||
};
|
||||
let runtime_config = read_chat_pii_redaction_runtime_config(state)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to read chat pii redaction runtime config"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
if !runtime_config.enabled {
|
||||
return Ok(ProviderRequestRedaction::disabled(body_json));
|
||||
}
|
||||
let feature_settings = resolve_chat_pii_redaction_feature_settings(state, auth_context).await?;
|
||||
if !feature_settings.effective_enabled() {
|
||||
return Ok(ProviderRequestRedaction::disabled(body_json));
|
||||
}
|
||||
let Some(hmac_key) = state.encryption_key().map(str::as_bytes).map(Vec::from) else {
|
||||
warn!("gateway chat pii redaction is enabled but encryption key is unavailable");
|
||||
return Err(GatewayError::Internal(
|
||||
"chat pii redaction setup failed".to_string(),
|
||||
));
|
||||
};
|
||||
let body_bytes = serde_json::to_vec(body_json).map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to serialize provider chat pii redaction body"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs();
|
||||
let cache = RedisRedactionMappingCache::new(state.runtime_state.as_ref());
|
||||
let masked = try_mask_chat_pii_request_json_with_cache_options(
|
||||
&body_bytes,
|
||||
format,
|
||||
build_redaction_session_config(hmac_key, &runtime_config, now_unix_secs),
|
||||
MaskChatRequestOptions::runtime(feature_settings.effective_inject_model_instruction()),
|
||||
Some(&cache),
|
||||
)
|
||||
.await
|
||||
.map_err(redaction_mask_error_to_gateway_error)?;
|
||||
if !masked.redacted {
|
||||
return Ok(ProviderRequestRedaction {
|
||||
body_json: Cow::Borrowed(body_json),
|
||||
redacted: false,
|
||||
});
|
||||
}
|
||||
let masked_body_json = serde_json::from_slice::<Value>(&masked.body).map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to decode redacted provider chat pii body"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
slot.put_for_candidate(candidate_id, masked.session);
|
||||
Ok(ProviderRequestRedaction {
|
||||
body_json: Cow::Owned(masked_body_json),
|
||||
redacted: true,
|
||||
})
|
||||
}
|
||||
|
||||
async fn resolve_chat_pii_redaction_feature_settings(
|
||||
state: &AppState,
|
||||
auth_context: &ExecutionRuntimeAuthContext,
|
||||
) -> Result<ChatPiiRedactionFeatureSettings, GatewayError> {
|
||||
let user_settings = state
|
||||
.read_user_feature_settings(&auth_context.user_id)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to read user chat pii redaction feature settings"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
let key_settings = state
|
||||
.read_auth_api_key_feature_settings(
|
||||
&auth_context.user_id,
|
||||
&auth_context.api_key_id,
|
||||
auth_context.api_key_is_standalone,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to read api key chat pii redaction feature settings"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
|
||||
let mut settings = ChatPiiRedactionFeatureSettings::default();
|
||||
settings.merge_from_value(user_settings.as_ref());
|
||||
settings.merge_from_value(key_settings.as_ref());
|
||||
Ok(settings)
|
||||
}
|
||||
|
||||
fn redaction_mask_error_to_gateway_error(error: RedactionMaskError) -> GatewayError {
|
||||
match error {
|
||||
RedactionMaskError::Limit(limit) => GatewayError::Client {
|
||||
status: limit.client_status(),
|
||||
message: limit.safe_message().to_string(),
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -198,6 +198,21 @@ pub(crate) fn insert_provider_stream_event_api_format(
|
||||
insert_ai_provider_stream_event_api_format(extra_fields, provider_type);
|
||||
}
|
||||
|
||||
pub(crate) fn insert_native_client_envelope_name(
|
||||
extra_fields: &mut Map<String, Value>,
|
||||
envelope_name: &str,
|
||||
request_path: &str,
|
||||
) {
|
||||
if envelope_name.eq_ignore_ascii_case("antigravity:v1internal")
|
||||
&& request_path == "/v1internal:streamGenerateContent"
|
||||
{
|
||||
extra_fields.insert(
|
||||
"client_envelope_name".to_string(),
|
||||
Value::String(envelope_name.to_string()),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_incoming_tls_fingerprint(extra_fields: &mut Map<String, Value>, incoming_tls: Value) {
|
||||
let entry = extra_fields
|
||||
.entry("tls_fingerprint".to_string())
|
||||
|
||||
@@ -0,0 +1,386 @@
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub(crate) fn is_deepseek_provider(provider_type: &str, base_url: &str) -> bool {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
if matches!(
|
||||
provider_type.as_str(),
|
||||
"deepseek" | "deepseek_openai" | "deepseek_anthropic" | "deepseek_compatible"
|
||||
) {
|
||||
return true;
|
||||
}
|
||||
|
||||
let host = base_url_host(base_url);
|
||||
host == "deepseek.com" || host.ends_with(".deepseek.com")
|
||||
}
|
||||
|
||||
pub(crate) fn apply_deepseek_tool_call_thinking_compat(
|
||||
provider_request_body: &mut Value,
|
||||
provider_type: &str,
|
||||
base_url: &str,
|
||||
provider_api_format: &str,
|
||||
original_request_body: Option<&Value>,
|
||||
) {
|
||||
if !is_deepseek_provider(provider_type, base_url) {
|
||||
return;
|
||||
}
|
||||
|
||||
match crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str() {
|
||||
"openai:chat" => {
|
||||
apply_deepseek_openai_chat_thinking_compat(provider_request_body, original_request_body)
|
||||
}
|
||||
"claude:messages" => apply_deepseek_claude_messages_thinking_compat(
|
||||
provider_request_body,
|
||||
original_request_body,
|
||||
),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn base_url_host(base_url: &str) -> String {
|
||||
let lower = base_url.trim().to_ascii_lowercase();
|
||||
let without_scheme = lower
|
||||
.split_once("://")
|
||||
.map(|(_, rest)| rest)
|
||||
.unwrap_or(lower.as_str());
|
||||
let without_userinfo = without_scheme
|
||||
.rsplit_once('@')
|
||||
.map(|(_, host)| host)
|
||||
.unwrap_or(without_scheme);
|
||||
without_userinfo
|
||||
.split(['/', '?', '#'])
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.split(':')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
fn source_disables_thinking(
|
||||
original_request_body: Option<&Value>,
|
||||
provider_request_body: &Value,
|
||||
) -> bool {
|
||||
request_explicitly_disables_thinking(provider_request_body)
|
||||
|| original_request_body.is_some_and(request_explicitly_disables_thinking)
|
||||
}
|
||||
|
||||
fn request_explicitly_disables_thinking(body: &Value) -> bool {
|
||||
thinking_type(body).is_some_and(|value| value.eq_ignore_ascii_case("disabled"))
|
||||
|| reasoning_effort(body).is_some_and(|value| value.eq_ignore_ascii_case("none"))
|
||||
}
|
||||
|
||||
fn thinking_type(body: &Value) -> Option<&str> {
|
||||
body.get("thinking")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|thinking| thinking.get("type"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn reasoning_effort(body: &Value) -> Option<&str> {
|
||||
body.get("reasoning_effort")
|
||||
.and_then(Value::as_str)
|
||||
.or_else(|| {
|
||||
body.get("reasoning")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|reasoning| reasoning.get("effort"))
|
||||
.and_then(Value::as_str)
|
||||
})
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn set_deepseek_thinking_type(body: &mut Value, thinking_type: &str) {
|
||||
let Some(object) = body.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
|
||||
match object.get_mut("thinking") {
|
||||
Some(Value::Object(thinking)) => {
|
||||
thinking.insert("type".to_string(), Value::String(thinking_type.to_string()));
|
||||
}
|
||||
_ => {
|
||||
object.insert(
|
||||
"thinking".to_string(),
|
||||
json!({
|
||||
"type": thinking_type,
|
||||
}),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_deepseek_openai_chat_thinking_compat(
|
||||
provider_request_body: &mut Value,
|
||||
original_request_body: Option<&Value>,
|
||||
) {
|
||||
let disabled = source_disables_thinking(original_request_body, provider_request_body);
|
||||
set_deepseek_thinking_type(
|
||||
provider_request_body,
|
||||
if disabled { "disabled" } else { "enabled" },
|
||||
);
|
||||
|
||||
let Some(object) = provider_request_body.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
if disabled {
|
||||
if reasoning_effort(&Value::Object(object.clone()))
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case("none"))
|
||||
{
|
||||
object.remove("reasoning_effort");
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
let Some(messages) = object.get_mut("messages").and_then(Value::as_array_mut) else {
|
||||
return;
|
||||
};
|
||||
for message in messages {
|
||||
let Some(message_object) = message.as_object_mut() else {
|
||||
continue;
|
||||
};
|
||||
let is_assistant = message_object
|
||||
.get("role")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|role| role.trim().eq_ignore_ascii_case("assistant"));
|
||||
if !is_assistant {
|
||||
continue;
|
||||
}
|
||||
if message_object
|
||||
.get("reasoning_content")
|
||||
.is_some_and(|value| !value.is_null())
|
||||
{
|
||||
continue;
|
||||
}
|
||||
message_object.insert(
|
||||
"reasoning_content".to_string(),
|
||||
Value::String(String::new()),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_deepseek_claude_messages_thinking_compat(
|
||||
provider_request_body: &mut Value,
|
||||
original_request_body: Option<&Value>,
|
||||
) {
|
||||
if source_disables_thinking(original_request_body, provider_request_body) {
|
||||
set_deepseek_thinking_type(provider_request_body, "disabled");
|
||||
return;
|
||||
}
|
||||
|
||||
let Some(messages) = provider_request_body
|
||||
.get_mut("messages")
|
||||
.and_then(Value::as_array_mut)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
for message in messages {
|
||||
let Some(message_object) = message.as_object_mut() else {
|
||||
continue;
|
||||
};
|
||||
let is_assistant = message_object
|
||||
.get("role")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|role| role.trim().eq_ignore_ascii_case("assistant"));
|
||||
if !is_assistant {
|
||||
continue;
|
||||
}
|
||||
ensure_claude_assistant_message_has_thinking_block(message_object);
|
||||
}
|
||||
}
|
||||
|
||||
fn ensure_claude_assistant_message_has_thinking_block(
|
||||
message: &mut serde_json::Map<String, Value>,
|
||||
) {
|
||||
let thinking_block = json!({
|
||||
"type": "thinking",
|
||||
"thinking": "",
|
||||
});
|
||||
match message.get_mut("content") {
|
||||
Some(Value::Array(blocks)) => {
|
||||
if blocks.iter().any(is_claude_thinking_block) {
|
||||
return;
|
||||
}
|
||||
blocks.insert(0, thinking_block);
|
||||
}
|
||||
Some(Value::String(text)) => {
|
||||
let text = std::mem::take(text);
|
||||
message.insert(
|
||||
"content".to_string(),
|
||||
Value::Array(vec![
|
||||
thinking_block,
|
||||
json!({
|
||||
"type": "text",
|
||||
"text": text,
|
||||
}),
|
||||
]),
|
||||
);
|
||||
}
|
||||
Some(Value::Null) | None => {
|
||||
message.insert("content".to_string(), Value::Array(vec![thinking_block]));
|
||||
}
|
||||
Some(other) => {
|
||||
let existing = std::mem::take(other);
|
||||
message.insert(
|
||||
"content".to_string(),
|
||||
Value::Array(vec![thinking_block, existing]),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn is_claude_thinking_block(block: &Value) -> bool {
|
||||
block
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|block_type| block_type.trim().eq_ignore_ascii_case("thinking"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{apply_deepseek_tool_call_thinking_compat, is_deepseek_provider};
|
||||
|
||||
#[test]
|
||||
fn detects_deepseek_provider_by_type_or_host() {
|
||||
assert!(is_deepseek_provider(
|
||||
"deepseek",
|
||||
"https://relay.example.com"
|
||||
));
|
||||
assert!(is_deepseek_provider(
|
||||
"custom",
|
||||
"https://api.deepseek.com/v1"
|
||||
));
|
||||
assert!(!is_deepseek_provider(
|
||||
"custom",
|
||||
"https://example.com/deepseek"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_chat_deepseek_adds_thinking_and_empty_reasoning_content() {
|
||||
let mut body = json!({
|
||||
"model": "deepseek-chat",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": null, "tool_calls": [{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "lookup", "arguments": "{}"}
|
||||
}]},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "{}"}
|
||||
]
|
||||
});
|
||||
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut body,
|
||||
"deepseek",
|
||||
"https://api.deepseek.com/v1",
|
||||
"openai:chat",
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(body["thinking"]["type"], "enabled");
|
||||
assert_eq!(body["messages"][1]["reasoning_content"], "");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_chat_deepseek_honors_disabled_thinking() {
|
||||
let original = json!({"reasoning_effort": "none"});
|
||||
let mut body = json!({
|
||||
"model": "deepseek-chat",
|
||||
"reasoning_effort": "none",
|
||||
"messages": [
|
||||
{"role": "assistant", "content": "hi"}
|
||||
]
|
||||
});
|
||||
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut body,
|
||||
"deepseek",
|
||||
"https://api.deepseek.com/v1",
|
||||
"openai:chat",
|
||||
Some(&original),
|
||||
);
|
||||
|
||||
assert_eq!(body["thinking"]["type"], "disabled");
|
||||
assert!(body.get("reasoning_effort").is_none());
|
||||
assert!(body["messages"][0].get("reasoning_content").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_messages_deepseek_prepends_empty_thinking_block() {
|
||||
let mut body = json!({
|
||||
"model": "deepseek-3.2",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "tool_use", "id": "call_1", "name": "lookup", "input": {}}
|
||||
]}
|
||||
]
|
||||
});
|
||||
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut body,
|
||||
"deepseek",
|
||||
"https://api.deepseek.com",
|
||||
"claude:messages",
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(body["messages"][1]["content"][0]["type"], "thinking");
|
||||
assert_eq!(body["messages"][1]["content"][0]["thinking"], "");
|
||||
assert_eq!(body["messages"][1]["content"][1]["type"], "tool_use");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_messages_deepseek_converts_string_assistant_content_to_blocks() {
|
||||
let mut body = json!({
|
||||
"model": "deepseek-3.2",
|
||||
"messages": [{
|
||||
"role": "assistant",
|
||||
"content": "done"
|
||||
}]
|
||||
});
|
||||
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut body,
|
||||
"deepseek",
|
||||
"https://api.deepseek.com",
|
||||
"claude:messages",
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(body["messages"][0]["content"][0]["type"], "thinking");
|
||||
assert_eq!(body["messages"][0]["content"][1]["type"], "text");
|
||||
assert_eq!(body["messages"][0]["content"][1]["text"], "done");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_messages_deepseek_preserves_existing_thinking_block() {
|
||||
let mut body = json!({
|
||||
"model": "deepseek-3.2",
|
||||
"messages": [{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": "sig"},
|
||||
{"type": "text", "text": "answer"}
|
||||
]
|
||||
}]
|
||||
});
|
||||
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut body,
|
||||
"deepseek",
|
||||
"https://api.deepseek.com",
|
||||
"claude:messages",
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(body["messages"][0]["content"].as_array().unwrap().len(), 2);
|
||||
assert_eq!(body["messages"][0]["content"][0]["thinking"], "plan");
|
||||
assert_eq!(body["messages"][0]["content"][0]["signature"], "sig");
|
||||
}
|
||||
}
|
||||
@@ -9,7 +9,8 @@ use crate::ai_serving::planner::materialization_policy::{
|
||||
};
|
||||
use crate::ai_serving::planner::passthrough::maybe_build_local_same_format_provider_decision_payload_for_candidate;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, LocalExecutionReportContextParts,
|
||||
build_local_execution_report_context, insert_native_client_envelope_name,
|
||||
LocalExecutionReportContextParts,
|
||||
};
|
||||
use crate::ai_serving::planner::spec_metadata::local_standard_spec_metadata;
|
||||
use crate::ai_serving::planner::CandidateFailureDiagnostic;
|
||||
@@ -74,10 +75,15 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
let Some(resolved) = resolve_local_standard_candidate_payload_parts(
|
||||
state, parts, trace_id, body_json, input, &attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let original_request_body_json = if resolved.request_redacted {
|
||||
Some(&resolved.provider_request_body)
|
||||
} else {
|
||||
Some(body_json)
|
||||
};
|
||||
let proxy = state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
|
||||
.await;
|
||||
@@ -92,6 +98,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
"envelope_name".to_string(),
|
||||
serde_json::Value::String(envelope_name.to_string()),
|
||||
);
|
||||
insert_native_client_envelope_name(&mut extra_fields, envelope_name, parts.uri.path());
|
||||
}
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
@@ -129,7 +136,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
original_request_body_json: Some(body_json),
|
||||
original_request_body_json,
|
||||
original_request_body_base64: None,
|
||||
client_session_affinity: input.client_session_affinity.as_ref(),
|
||||
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
|
||||
@@ -166,6 +173,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
envelope_name: _,
|
||||
transport,
|
||||
transport_profile: _,
|
||||
request_redacted: _,
|
||||
} = resolved;
|
||||
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
|
||||
@@ -16,9 +16,13 @@ use crate::ai_serving::planner::gemini_cli::{
|
||||
build_gemini_cli_v1internal_provider_request, GeminiCliV1InternalRequestError,
|
||||
GeminiCliV1InternalRequestInput,
|
||||
};
|
||||
use crate::ai_serving::planner::redaction::{
|
||||
request_identity_response_encoding_when_redacted, resolve_provider_chat_pii_redaction,
|
||||
};
|
||||
use crate::ai_serving::planner::spec_metadata::local_standard_spec_metadata;
|
||||
use crate::ai_serving::planner::standard::{
|
||||
apply_codex_openai_responses_special_headers, request_body_build_failure_extra_data,
|
||||
apply_codex_openai_responses_special_headers, apply_deepseek_tool_call_thinking_compat,
|
||||
is_deepseek_provider, request_body_build_failure_extra_data,
|
||||
};
|
||||
use crate::ai_serving::transport::kiro::{
|
||||
build_kiro_provider_headers, build_kiro_provider_request_body,
|
||||
@@ -41,7 +45,7 @@ use crate::ai_serving::{
|
||||
build_openai_image_request_body_from_gemini_image_request, gemini_request_is_image_generation,
|
||||
CandidateFailureDiagnostic, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use crate::AppState;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
use super::payload::{
|
||||
mark_skipped_local_standard_candidate, mark_skipped_local_standard_candidate_with_extra_data,
|
||||
@@ -49,6 +53,8 @@ use super::payload::{
|
||||
};
|
||||
use super::{LocalStandardCandidateAttempt, LocalStandardDecisionInput, LocalStandardSpec};
|
||||
|
||||
const OMITTED_THINKING_TEXT: &str = "Previous thinking omitted.";
|
||||
|
||||
pub(crate) struct LocalStandardCandidatePayloadParts {
|
||||
pub(super) auth_header: String,
|
||||
pub(super) auth_value: String,
|
||||
@@ -61,6 +67,7 @@ pub(crate) struct LocalStandardCandidatePayloadParts {
|
||||
pub(super) envelope_name: Option<&'static str>,
|
||||
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
pub(super) request_redacted: bool,
|
||||
}
|
||||
|
||||
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
|
||||
@@ -119,8 +126,6 @@ fn sanitize_claude_thinking_block(block: Value) -> (Option<Value>, bool) {
|
||||
}
|
||||
|
||||
fn sanitize_claude_message_content_for_non_native_thinking(content: &mut Value) -> bool {
|
||||
const OMITTED_THINKING_TEXT: &str = "Previous thinking omitted.";
|
||||
|
||||
if content.is_object() {
|
||||
let original = std::mem::take(content);
|
||||
let (sanitized, changed) = sanitize_claude_thinking_block(original);
|
||||
@@ -180,6 +185,81 @@ fn sanitize_claude_request_thinking_signatures_for_non_native(body_json: &mut Va
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn remove_claude_redacted_thinking_block(block: Value) -> (Option<Value>, bool) {
|
||||
let Some(object) = block.as_object() else {
|
||||
return (Some(block), false);
|
||||
};
|
||||
let block_type = object
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default();
|
||||
if block_type == "redacted_thinking" {
|
||||
return (None, true);
|
||||
}
|
||||
(Some(block), false)
|
||||
}
|
||||
|
||||
fn sanitize_claude_message_content_for_deepseek_thinking(content: &mut Value) -> bool {
|
||||
if content.is_object() {
|
||||
let original = std::mem::take(content);
|
||||
let (sanitized, changed) = remove_claude_redacted_thinking_block(original);
|
||||
if changed {
|
||||
*content = sanitized.unwrap_or_else(|| {
|
||||
serde_json::json!({
|
||||
"type": "text",
|
||||
"text": OMITTED_THINKING_TEXT,
|
||||
})
|
||||
});
|
||||
}
|
||||
return changed;
|
||||
}
|
||||
|
||||
let Some(blocks) = content.as_array_mut() else {
|
||||
return false;
|
||||
};
|
||||
let original_blocks = std::mem::take(blocks);
|
||||
let mut changed = false;
|
||||
let mut sanitized_blocks = Vec::with_capacity(original_blocks.len());
|
||||
for block in original_blocks {
|
||||
let (sanitized, block_changed) = remove_claude_redacted_thinking_block(block);
|
||||
changed |= block_changed;
|
||||
if let Some(sanitized) = sanitized {
|
||||
sanitized_blocks.push(sanitized);
|
||||
}
|
||||
}
|
||||
if changed && sanitized_blocks.is_empty() {
|
||||
sanitized_blocks.push(serde_json::json!({
|
||||
"type": "text",
|
||||
"text": OMITTED_THINKING_TEXT,
|
||||
}));
|
||||
}
|
||||
*blocks = sanitized_blocks;
|
||||
changed
|
||||
}
|
||||
|
||||
fn sanitize_claude_request_redacted_thinking_for_deepseek(body_json: &mut Value) -> bool {
|
||||
body_json
|
||||
.get_mut("messages")
|
||||
.and_then(Value::as_array_mut)
|
||||
.map(|messages| {
|
||||
messages.iter_mut().fold(false, |changed, message| {
|
||||
let is_assistant = message
|
||||
.get("role")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|role| role.trim().eq_ignore_ascii_case("assistant"));
|
||||
if !is_assistant {
|
||||
return changed;
|
||||
}
|
||||
let content_changed = message
|
||||
.get_mut("content")
|
||||
.is_some_and(sanitize_claude_message_content_for_deepseek_thinking);
|
||||
changed || content_changed
|
||||
})
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn apply_non_native_claude_thinking_signature_compat(
|
||||
provider_request_body: &mut Value,
|
||||
provider_api_format: &str,
|
||||
@@ -188,6 +268,13 @@ fn apply_non_native_claude_thinking_signature_compat(
|
||||
if crate::ai_serving::normalize_api_format_alias(provider_api_format) != "claude:messages" {
|
||||
return;
|
||||
}
|
||||
if is_deepseek_provider(
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.base_url.as_str(),
|
||||
) {
|
||||
let _ = sanitize_claude_request_redacted_thinking_for_deepseek(provider_request_body);
|
||||
return;
|
||||
}
|
||||
if provider_preserves_claude_thinking_signatures(
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.base_url.as_str(),
|
||||
@@ -206,7 +293,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
input: &LocalStandardDecisionInput,
|
||||
attempt: &LocalStandardCandidateAttempt,
|
||||
spec: LocalStandardSpec,
|
||||
) -> Option<LocalStandardCandidatePayloadParts> {
|
||||
) -> Result<Option<LocalStandardCandidatePayloadParts>, GatewayError> {
|
||||
let spec_metadata = local_standard_spec_metadata(spec);
|
||||
let planner_state = crate::ai_serving::PlannerAppState::new(state);
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
@@ -223,10 +310,12 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
&& provider_api_format == "openai:image"
|
||||
&& gemini_request_is_image_generation(body_json)
|
||||
{
|
||||
return resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
|
||||
state, parts, trace_id, body_json, input, attempt,
|
||||
)
|
||||
.await;
|
||||
return Ok(
|
||||
resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
|
||||
state, parts, trace_id, body_json, input, attempt,
|
||||
)
|
||||
.await,
|
||||
);
|
||||
}
|
||||
let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format);
|
||||
if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
@@ -255,10 +344,21 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
let redaction = resolve_provider_chat_pii_redaction(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
&input.auth_context,
|
||||
spec_metadata.api_format,
|
||||
&attempt.candidate_id,
|
||||
)
|
||||
.await?;
|
||||
let body_json = redaction.body_json.as_ref();
|
||||
|
||||
let mut provider_request_body = body_json.clone();
|
||||
if let Some(object) = provider_request_body.as_object_mut() {
|
||||
object.insert(
|
||||
@@ -284,7 +384,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
);
|
||||
|
||||
let upstream_url = build_grok_upstream_url(transport, GROK_CHAT_PATH);
|
||||
let Some(provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
|
||||
let Some(mut provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
|
||||
transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(effective_headers),
|
||||
@@ -309,10 +409,14 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
redaction.redacted,
|
||||
);
|
||||
|
||||
return Some(LocalStandardCandidatePayloadParts {
|
||||
return Ok(Some(LocalStandardCandidatePayloadParts {
|
||||
auth_header: prepared_candidate.auth_header,
|
||||
auth_value: prepared_candidate.auth_value,
|
||||
mapped_model: prepared_candidate.mapped_model,
|
||||
@@ -324,7 +428,8 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile,
|
||||
});
|
||||
request_redacted: redaction.redacted,
|
||||
}));
|
||||
}
|
||||
|
||||
if !crate::ai_serving::request_pair_allowed_for_transport(
|
||||
@@ -332,7 +437,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
) {
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let is_windsurf_cascade =
|
||||
@@ -357,7 +462,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let oauth_context = OauthPreparationContext {
|
||||
@@ -385,7 +490,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
"transport_auth_unavailable",
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -410,7 +515,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -435,7 +540,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -456,6 +561,16 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
Some(&input.requested_model),
|
||||
)
|
||||
.await;
|
||||
let redaction = resolve_provider_chat_pii_redaction(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
&input.auth_context,
|
||||
spec_metadata.api_format,
|
||||
&attempt.candidate_id,
|
||||
)
|
||||
.await?;
|
||||
let body_json = redaction.body_json.as_ref();
|
||||
let mut provider_request_body =
|
||||
match crate::ai_serving::planner::standard::build_standard_request_body_with_model_directives_and_request_headers(
|
||||
body_json,
|
||||
@@ -491,7 +606,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
enforce_provider_body_stream_policy(
|
||||
@@ -521,13 +636,20 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
apply_non_native_claude_thinking_signature_compat(
|
||||
&mut provider_request_body,
|
||||
provider_api_format,
|
||||
transport,
|
||||
);
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.base_url.as_str(),
|
||||
provider_api_format,
|
||||
Some(body_json),
|
||||
);
|
||||
if let Some(mapping) =
|
||||
crate::system_features::reasoning_model_directive_mapping_for_api_format_and_model(
|
||||
state,
|
||||
@@ -569,17 +691,24 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
apply_non_native_claude_thinking_signature_compat(
|
||||
&mut provider_request_body,
|
||||
provider_api_format,
|
||||
transport,
|
||||
);
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.base_url.as_str(),
|
||||
provider_api_format,
|
||||
Some(body_json),
|
||||
);
|
||||
}
|
||||
|
||||
if let Some(kiro_auth) = kiro_auth.as_ref() {
|
||||
return build_kiro_cross_format_payload_parts(
|
||||
return Ok(build_kiro_cross_format_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
@@ -594,11 +723,12 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
provider_request_body,
|
||||
upstream_is_stream,
|
||||
kiro_auth,
|
||||
redaction.redacted,
|
||||
)
|
||||
.await;
|
||||
.await);
|
||||
}
|
||||
if is_windsurf_cascade {
|
||||
return build_windsurf_cross_format_payload_parts(
|
||||
return Ok(build_windsurf_cross_format_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
@@ -612,8 +742,9 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
prepared_candidate.auth_value,
|
||||
provider_request_body,
|
||||
upstream_is_stream,
|
||||
redaction.redacted,
|
||||
)
|
||||
.await;
|
||||
.await);
|
||||
}
|
||||
|
||||
let normalized_provider_api_format =
|
||||
@@ -621,7 +752,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
if normalized_provider_api_format == "gemini:generate_content"
|
||||
&& is_gemini_cli_provider_transport(transport)
|
||||
{
|
||||
return build_gemini_cli_cross_format_payload_parts(
|
||||
return Ok(build_gemini_cli_cross_format_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
@@ -636,8 +767,9 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
prepared_candidate.auth_value,
|
||||
provider_request_body,
|
||||
upstream_is_stream,
|
||||
redaction.redacted,
|
||||
)
|
||||
.await;
|
||||
.await);
|
||||
}
|
||||
|
||||
let upstream_url = match crate::ai_serving::planner::standard::build_standard_upstream_url(
|
||||
@@ -665,7 +797,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let Some(resolved_headers) =
|
||||
@@ -698,7 +830,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
let mut provider_request_headers = resolved_headers.headers;
|
||||
apply_codex_openai_responses_special_headers(
|
||||
@@ -710,8 +842,12 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
Some(trace_id),
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
redaction.redacted,
|
||||
);
|
||||
|
||||
Some(LocalStandardCandidatePayloadParts {
|
||||
Ok(Some(LocalStandardCandidatePayloadParts {
|
||||
auth_header: resolved_headers.auth_header,
|
||||
auth_value: resolved_headers.auth_value,
|
||||
mapped_model: prepared_candidate.mapped_model,
|
||||
@@ -723,7 +859,8 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
})
|
||||
request_redacted: redaction.redacted,
|
||||
}))
|
||||
}
|
||||
|
||||
fn apply_transport_request_body_semantics(
|
||||
@@ -754,6 +891,7 @@ async fn build_gemini_cli_cross_format_payload_parts(
|
||||
auth_value: String,
|
||||
gemini_request_body: Value,
|
||||
upstream_is_stream: bool,
|
||||
request_redacted: bool,
|
||||
) -> Option<LocalStandardCandidatePayloadParts> {
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
@@ -854,6 +992,10 @@ async fn build_gemini_cli_cross_format_payload_parts(
|
||||
Some(trace_id),
|
||||
resolved.transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
request_redacted,
|
||||
);
|
||||
|
||||
Some(LocalStandardCandidatePayloadParts {
|
||||
auth_header: resolved.headers.auth_header,
|
||||
@@ -867,6 +1009,7 @@ async fn build_gemini_cli_cross_format_payload_parts(
|
||||
envelope_name: Some(GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME),
|
||||
transport: resolved.transport,
|
||||
transport_profile: None,
|
||||
request_redacted,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -885,6 +1028,7 @@ async fn build_windsurf_cross_format_payload_parts(
|
||||
auth_value: String,
|
||||
openai_chat_request_body: Value,
|
||||
upstream_is_stream: bool,
|
||||
request_redacted: bool,
|
||||
) -> Option<LocalStandardCandidatePayloadParts> {
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
@@ -940,7 +1084,7 @@ async fn build_windsurf_cross_format_payload_parts(
|
||||
return None;
|
||||
}
|
||||
};
|
||||
let provider_request_headers = match build_windsurf_cascade_headers(
|
||||
let mut provider_request_headers = match build_windsurf_cascade_headers(
|
||||
effective_headers,
|
||||
&provider_request_body,
|
||||
original_body_json,
|
||||
@@ -969,6 +1113,10 @@ async fn build_windsurf_cross_format_payload_parts(
|
||||
return None;
|
||||
}
|
||||
};
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
request_redacted,
|
||||
);
|
||||
|
||||
Some(LocalStandardCandidatePayloadParts {
|
||||
auth_header,
|
||||
@@ -982,6 +1130,7 @@ async fn build_windsurf_cross_format_payload_parts(
|
||||
envelope_name: Some(WINDSURF_ENVELOPE_NAME),
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
request_redacted,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1120,6 +1269,7 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
request_redacted: false,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1139,6 +1289,7 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
claude_request_body: Value,
|
||||
upstream_is_stream: bool,
|
||||
kiro_auth: &KiroRequestAuth,
|
||||
request_redacted: bool,
|
||||
) -> Option<LocalStandardCandidatePayloadParts> {
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
@@ -1197,7 +1348,7 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
return None;
|
||||
}
|
||||
};
|
||||
let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
let mut provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
headers: effective_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: original_body_json,
|
||||
@@ -1227,6 +1378,10 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
return None;
|
||||
}
|
||||
};
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
request_redacted,
|
||||
);
|
||||
|
||||
Some(LocalStandardCandidatePayloadParts {
|
||||
auth_header,
|
||||
@@ -1240,6 +1395,7 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
envelope_name: Some(KIRO_ENVELOPE_NAME),
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
request_redacted,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1247,6 +1403,7 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
mod tests {
|
||||
use super::{
|
||||
provider_preserves_claude_thinking_signatures,
|
||||
sanitize_claude_request_redacted_thinking_for_deepseek,
|
||||
sanitize_claude_request_thinking_signatures_for_non_native,
|
||||
};
|
||||
use serde_json::json;
|
||||
@@ -1310,6 +1467,46 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn deepseek_sanitizer_preserves_plain_thinking_but_removes_redacted() {
|
||||
let mut body = json!({
|
||||
"model": "claude-opus-4-1",
|
||||
"messages": [{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "thinking",
|
||||
"thinking": "I should keep this short.",
|
||||
"signature": "sig_123"
|
||||
},
|
||||
{
|
||||
"type": "redacted_thinking",
|
||||
"data": "opaque"
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Done."
|
||||
}
|
||||
]
|
||||
}]
|
||||
});
|
||||
|
||||
assert!(sanitize_claude_request_redacted_thinking_for_deepseek(
|
||||
&mut body
|
||||
));
|
||||
assert_eq!(body["messages"][0]["content"].as_array().unwrap().len(), 2);
|
||||
assert_eq!(body["messages"][0]["content"][0]["type"], json!("thinking"));
|
||||
assert_eq!(
|
||||
body["messages"][0]["content"][0]["thinking"],
|
||||
json!("I should keep this short.")
|
||||
);
|
||||
assert_eq!(
|
||||
body["messages"][0]["content"][0]["signature"],
|
||||
json!("sig_123")
|
||||
);
|
||||
assert_eq!(body["messages"][0]["content"][1]["text"], json!("Done."));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn official_claude_providers_preserve_thinking_signatures() {
|
||||
assert!(provider_preserves_claude_thinking_signatures(
|
||||
@@ -1328,6 +1525,14 @@ mod tests {
|
||||
"amazon_bedrock",
|
||||
"https://relay.example.com"
|
||||
));
|
||||
assert!(!provider_preserves_claude_thinking_signatures(
|
||||
"deepseek",
|
||||
"https://relay.example.com"
|
||||
));
|
||||
assert!(!provider_preserves_claude_thinking_signatures(
|
||||
"custom",
|
||||
"https://api.deepseek.com"
|
||||
));
|
||||
assert!(!provider_preserves_claude_thinking_signatures(
|
||||
"openai",
|
||||
"https://relay.example.com"
|
||||
|
||||
@@ -8,6 +8,7 @@ use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
mod claude;
|
||||
mod codex;
|
||||
mod deepseek;
|
||||
mod family;
|
||||
mod gemini;
|
||||
mod normalize;
|
||||
@@ -16,6 +17,7 @@ mod openai;
|
||||
pub(crate) use self::codex::{
|
||||
apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers,
|
||||
};
|
||||
pub(crate) use self::deepseek::{apply_deepseek_tool_call_thinking_compat, is_deepseek_provider};
|
||||
pub(crate) use self::family::{
|
||||
build_local_stream_attempt_source, build_local_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source, build_local_sync_plan_and_reports,
|
||||
|
||||
@@ -292,6 +292,47 @@ fn strips_metadata_for_codex_openai_responses_requests() {
|
||||
assert!(provider_request_body.get("metadata").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn openai_chat_to_codex_responses_preserves_json_mode_chat_messages() {
|
||||
let body_json = json!({
|
||||
"model": "gpt-5.5",
|
||||
"messages": [
|
||||
{"role": "system", "content": "Return a JSON object."},
|
||||
{"role": "user", "content": "Why did this JSON request fail?"}
|
||||
],
|
||||
"response_format": {"type": "json_object"}
|
||||
});
|
||||
|
||||
let provider_request_body = build_cross_format_openai_responses_request_body(
|
||||
&body_json,
|
||||
"gpt-5.5-upstream",
|
||||
"openai:chat",
|
||||
"openai:responses",
|
||||
false,
|
||||
false,
|
||||
"codex",
|
||||
None,
|
||||
None,
|
||||
&http::HeaderMap::new(),
|
||||
false,
|
||||
)
|
||||
.expect("openai chat to codex responses request should build");
|
||||
|
||||
assert_eq!(
|
||||
provider_request_body["text"]["format"]["type"],
|
||||
"json_object"
|
||||
);
|
||||
assert_eq!(provider_request_body["input"][0]["role"], "user");
|
||||
assert_eq!(
|
||||
provider_request_body["input"][0]["content"][0]["text"],
|
||||
"Why did this JSON request fail?"
|
||||
);
|
||||
assert_eq!(
|
||||
provider_request_body["instructions"],
|
||||
"Return a JSON object."
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn applies_codex_defaults_unless_body_rules_handle_the_field() {
|
||||
let body_json = json!({
|
||||
@@ -356,7 +397,7 @@ fn injects_codex_prompt_cache_key_for_openai_responses_cross_format_requests() {
|
||||
|
||||
assert_eq!(
|
||||
provider_request_body["prompt_cache_key"],
|
||||
"b4dfeb75-b105-544c-a706-39b92f0bddb0"
|
||||
"4ee6ea6e-3ac6-5a18-8cb8-1f8b956419e5"
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -2,8 +2,8 @@ use crate::ai_serving::build_request_trace_proxy_value;
|
||||
use crate::ai_serving::planner::common::OPENAI_CHAT_STREAM_PLAN_KIND;
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, insert_provider_stream_event_api_format,
|
||||
LocalExecutionReportContextParts,
|
||||
build_local_execution_report_context, insert_native_client_envelope_name,
|
||||
insert_provider_stream_event_api_format, LocalExecutionReportContextParts,
|
||||
};
|
||||
use crate::ai_serving::planner::{
|
||||
build_ai_execution_decision_response, AiExecutionDecisionResponseParts,
|
||||
@@ -84,6 +84,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
|
||||
"envelope_name".to_string(),
|
||||
serde_json::Value::String(envelope_name.to_string()),
|
||||
);
|
||||
insert_native_client_envelope_name(&mut extra_fields, envelope_name, parts.uri.path());
|
||||
}
|
||||
insert_provider_stream_event_api_format(
|
||||
&mut extra_fields,
|
||||
|
||||
+30
-193
@@ -1,7 +1,5 @@
|
||||
use std::borrow::Cow;
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_contracts::ResolvedTransportProfile;
|
||||
use serde_json::{json, Value};
|
||||
@@ -19,11 +17,14 @@ use crate::ai_serving::planner::gemini_cli::{
|
||||
build_gemini_cli_v1internal_provider_request, GeminiCliV1InternalRequestError,
|
||||
GeminiCliV1InternalRequestInput,
|
||||
};
|
||||
use crate::ai_serving::planner::redaction::{
|
||||
request_identity_response_encoding_when_redacted, resolve_provider_chat_pii_redaction,
|
||||
};
|
||||
use crate::ai_serving::planner::standard::{
|
||||
apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers,
|
||||
build_cross_format_openai_chat_request_body, build_cross_format_openai_chat_upstream_url,
|
||||
build_local_openai_chat_request_body, build_local_openai_chat_upstream_url,
|
||||
request_body_build_failure_extra_data,
|
||||
apply_deepseek_tool_call_thinking_compat, build_cross_format_openai_chat_request_body,
|
||||
build_cross_format_openai_chat_upstream_url, build_local_openai_chat_request_body,
|
||||
build_local_openai_chat_upstream_url, request_body_build_failure_extra_data,
|
||||
};
|
||||
use crate::ai_serving::transport::auth::resolve_local_openai_bearer_auth;
|
||||
use crate::ai_serving::transport::kiro::{
|
||||
@@ -52,13 +53,7 @@ use crate::ai_serving::{
|
||||
LocalResolvedOAuthRequestAuth,
|
||||
};
|
||||
use crate::ai_serving::{ConversionMode, ExecutionStrategy};
|
||||
use crate::privacy::{
|
||||
build_redaction_session_config, read_chat_pii_redaction_runtime_config,
|
||||
try_mask_chat_request_json_with_cache_options, MaskChatRequestOptions, RedactionMaskError,
|
||||
RedactionSessionSlot, RedisRedactionMappingCache,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
use tracing::warn;
|
||||
|
||||
use super::support::{
|
||||
mark_skipped_local_openai_chat_candidate,
|
||||
@@ -92,100 +87,6 @@ fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
|
||||
)
|
||||
}
|
||||
|
||||
fn request_identity_response_encoding_when_redacted(
|
||||
headers: &mut BTreeMap<String, String>,
|
||||
redacted: bool,
|
||||
) {
|
||||
if redacted {
|
||||
headers.insert("accept-encoding".to_string(), "identity".to_string());
|
||||
}
|
||||
}
|
||||
|
||||
struct ProviderChatRequestRedaction<'a> {
|
||||
body_json: Cow<'a, Value>,
|
||||
redacted: bool,
|
||||
}
|
||||
|
||||
impl<'a> ProviderChatRequestRedaction<'a> {
|
||||
fn disabled(body_json: &'a Value, _parts: &http::request::Parts) -> Self {
|
||||
Self {
|
||||
body_json: Cow::Borrowed(body_json),
|
||||
redacted: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
struct ChatPiiRedactionFeatureSettings {
|
||||
enabled: Option<bool>,
|
||||
inject_model_instruction: Option<bool>,
|
||||
}
|
||||
|
||||
impl ChatPiiRedactionFeatureSettings {
|
||||
fn merge_from_value(&mut self, value: Option<&Value>) {
|
||||
let Some(settings) = value
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|features| features.get("chat_pii_redaction"))
|
||||
.and_then(Value::as_object)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if let Some(enabled) = settings.get("enabled").and_then(Value::as_bool) {
|
||||
self.enabled = Some(enabled);
|
||||
}
|
||||
if let Some(inject_model_instruction) = settings
|
||||
.get("inject_model_instruction")
|
||||
.and_then(Value::as_bool)
|
||||
{
|
||||
self.inject_model_instruction = Some(inject_model_instruction);
|
||||
}
|
||||
}
|
||||
|
||||
fn effective_enabled(self) -> bool {
|
||||
self.enabled.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn effective_inject_model_instruction(self) -> bool {
|
||||
self.inject_model_instruction.unwrap_or(true)
|
||||
}
|
||||
}
|
||||
|
||||
async fn resolve_chat_pii_redaction_feature_settings(
|
||||
state: &AppState,
|
||||
input: &LocalOpenAiChatDecisionInput,
|
||||
) -> Result<ChatPiiRedactionFeatureSettings, GatewayError> {
|
||||
let user_settings = state
|
||||
.read_user_feature_settings(&input.auth_context.user_id)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to read user chat pii redaction feature settings"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
let key_settings = state
|
||||
.read_auth_api_key_feature_settings(
|
||||
&input.auth_context.user_id,
|
||||
&input.auth_context.api_key_id,
|
||||
input.auth_context.api_key_is_standalone,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
warn!(
|
||||
|
||||
error = ?err,
|
||||
"gateway failed to read api key chat pii redaction feature settings"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
|
||||
let mut settings = ChatPiiRedactionFeatureSettings::default();
|
||||
settings.merge_from_value(user_settings.as_ref());
|
||||
settings.merge_from_value(key_settings.as_ref());
|
||||
Ok(settings)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
state: &AppState,
|
||||
@@ -214,9 +115,15 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
Some(&input.requested_model),
|
||||
)
|
||||
.await;
|
||||
let redaction =
|
||||
resolve_provider_chat_request_redaction(state, parts, body_json, input, candidate_id)
|
||||
.await?;
|
||||
let redaction = resolve_provider_chat_pii_redaction(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
&input.auth_context,
|
||||
"openai:chat",
|
||||
candidate_id,
|
||||
)
|
||||
.await?;
|
||||
let body_json = redaction.body_json.as_ref();
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let is_grok = transport
|
||||
@@ -408,7 +315,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
}
|
||||
};
|
||||
|
||||
let Some(provider_request_body) = build_local_openai_chat_request_body(
|
||||
let Some(mut provider_request_body) = build_local_openai_chat_request_body(
|
||||
body_json,
|
||||
&prepared_candidate.mapped_model,
|
||||
upstream_is_stream,
|
||||
@@ -434,6 +341,13 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
.await;
|
||||
return Ok(None);
|
||||
};
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.base_url.as_str(),
|
||||
"openai:chat",
|
||||
Some(body_json),
|
||||
);
|
||||
|
||||
let Some(upstream_url) = build_local_openai_chat_upstream_url(parts, transport) else {
|
||||
mark_skipped_local_openai_chat_candidate_with_failure_diagnostic(
|
||||
@@ -712,6 +626,13 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
request_requires_body_stream_field(body_json, force_body_stream_field),
|
||||
);
|
||||
}
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.base_url.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
Some(body_json),
|
||||
);
|
||||
|
||||
if let Some(kiro_auth) = kiro_auth.as_ref() {
|
||||
return Ok(build_kiro_openai_chat_cross_format_payload_parts(
|
||||
@@ -1773,90 +1694,6 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
|
||||
})
|
||||
}
|
||||
|
||||
async fn resolve_provider_chat_request_redaction<'a>(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
body_json: &'a Value,
|
||||
input: &LocalOpenAiChatDecisionInput,
|
||||
candidate_id: &str,
|
||||
) -> Result<ProviderChatRequestRedaction<'a>, GatewayError> {
|
||||
if parts.uri.path() != "/v1/chat/completions" {
|
||||
return Ok(ProviderChatRequestRedaction::disabled(body_json, parts));
|
||||
}
|
||||
let Some(slot) = parts.extensions.get::<RedactionSessionSlot>() else {
|
||||
return Ok(ProviderChatRequestRedaction::disabled(body_json, parts));
|
||||
};
|
||||
let runtime_config = read_chat_pii_redaction_runtime_config(state)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to read chat pii redaction runtime config"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
if !runtime_config.enabled {
|
||||
return Ok(ProviderChatRequestRedaction::disabled(body_json, parts));
|
||||
}
|
||||
let feature_settings = resolve_chat_pii_redaction_feature_settings(state, input).await?;
|
||||
if !feature_settings.effective_enabled() {
|
||||
return Ok(ProviderChatRequestRedaction::disabled(body_json, parts));
|
||||
}
|
||||
let Some(hmac_key) = state.encryption_key().map(str::as_bytes).map(Vec::from) else {
|
||||
warn!("gateway chat pii redaction is enabled but encryption key is unavailable");
|
||||
return Err(GatewayError::Internal(
|
||||
"chat pii redaction setup failed".to_string(),
|
||||
));
|
||||
};
|
||||
let body_bytes = serde_json::to_vec(body_json).map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to serialize provider chat pii redaction body"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs();
|
||||
let cache = RedisRedactionMappingCache::new(state.runtime_state.as_ref());
|
||||
let masked = try_mask_chat_request_json_with_cache_options(
|
||||
&body_bytes,
|
||||
build_redaction_session_config(hmac_key, &runtime_config, now_unix_secs),
|
||||
MaskChatRequestOptions::runtime(feature_settings.effective_inject_model_instruction()),
|
||||
Some(&cache),
|
||||
)
|
||||
.await
|
||||
.map_err(redaction_mask_error_to_gateway_error)?;
|
||||
if !masked.redacted {
|
||||
return Ok(ProviderChatRequestRedaction {
|
||||
body_json: Cow::Borrowed(body_json),
|
||||
redacted: false,
|
||||
});
|
||||
}
|
||||
let masked_body_json = serde_json::from_slice::<Value>(&masked.body).map_err(|err| {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway failed to decode redacted provider chat pii body"
|
||||
);
|
||||
GatewayError::Internal("chat pii redaction setup failed".to_string())
|
||||
})?;
|
||||
slot.put_for_candidate(candidate_id, masked.session);
|
||||
Ok(ProviderChatRequestRedaction {
|
||||
body_json: Cow::Owned(masked_body_json),
|
||||
redacted: true,
|
||||
})
|
||||
}
|
||||
|
||||
fn redaction_mask_error_to_gateway_error(error: RedactionMaskError) -> GatewayError {
|
||||
match error {
|
||||
RedactionMaskError::Limit(limit) => GatewayError::Client {
|
||||
status: limit.client_status(),
|
||||
message: limit.safe_message().to_string(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
+11
-4
@@ -4,8 +4,8 @@ use tracing::debug;
|
||||
use crate::ai_serving::build_request_trace_proxy_value;
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, insert_provider_stream_event_api_format,
|
||||
LocalExecutionReportContextParts,
|
||||
build_local_execution_report_context, insert_native_client_envelope_name,
|
||||
insert_provider_stream_event_api_format, LocalExecutionReportContextParts,
|
||||
};
|
||||
use crate::ai_serving::planner::spec_metadata::local_openai_responses_spec_metadata;
|
||||
use crate::ai_serving::planner::{
|
||||
@@ -51,11 +51,16 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
&candidate_id,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let candidate = &eligible.candidate;
|
||||
let original_request_body_json = if resolved.request_redacted {
|
||||
Some(&resolved.provider_request_body)
|
||||
} else {
|
||||
Some(body_json)
|
||||
};
|
||||
|
||||
let prompt_cache_key = resolved
|
||||
.provider_request_body
|
||||
@@ -80,6 +85,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
}
|
||||
if let Some(envelope_name) = resolved.envelope_name {
|
||||
extra_fields.insert("envelope_name".to_string(), json!(envelope_name));
|
||||
insert_native_client_envelope_name(&mut extra_fields, envelope_name, parts.uri.path());
|
||||
}
|
||||
if let Some(image_request_summary) = resolved.image_request_summary.as_ref() {
|
||||
extra_fields.insert("image_request".to_string(), image_request_summary.clone());
|
||||
@@ -141,7 +147,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
original_request_body_json: Some(body_json),
|
||||
original_request_body_json,
|
||||
original_request_body_base64: None,
|
||||
client_session_affinity: input.client_session_affinity.as_ref(),
|
||||
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
|
||||
@@ -204,6 +210,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
transport,
|
||||
transport_profile: _,
|
||||
image_request_summary: _,
|
||||
request_redacted: _,
|
||||
} = resolved;
|
||||
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
|
||||
+132
-31
@@ -18,10 +18,13 @@ use crate::ai_serving::planner::gemini_cli::{
|
||||
build_gemini_cli_v1internal_provider_request, GeminiCliV1InternalRequestError,
|
||||
GeminiCliV1InternalRequestInput,
|
||||
};
|
||||
use crate::ai_serving::planner::redaction::{
|
||||
request_identity_response_encoding_when_redacted, resolve_provider_chat_pii_redaction,
|
||||
};
|
||||
use crate::ai_serving::planner::spec_metadata::local_openai_responses_spec_metadata;
|
||||
use crate::ai_serving::planner::standard::{
|
||||
apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers,
|
||||
build_cross_format_openai_responses_request_body,
|
||||
apply_deepseek_tool_call_thinking_compat, build_cross_format_openai_responses_request_body,
|
||||
build_cross_format_openai_responses_upstream_url, build_local_openai_responses_request_body,
|
||||
build_local_openai_responses_upstream_url, request_body_build_failure_extra_data,
|
||||
};
|
||||
@@ -58,7 +61,7 @@ use crate::ai_serving::{
|
||||
LocalResolvedOAuthRequestAuth, PlannerAppState,
|
||||
};
|
||||
use crate::ai_serving::{ConversionMode, ExecutionStrategy};
|
||||
use crate::AppState;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
use super::support::{
|
||||
mark_skipped_local_openai_responses_candidate,
|
||||
@@ -93,6 +96,7 @@ pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts {
|
||||
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
pub(super) image_request_summary: Option<Value>,
|
||||
pub(super) request_redacted: bool,
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
@@ -106,7 +110,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
candidate_index: u32,
|
||||
candidate_id: &str,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
) -> Option<LocalOpenAiResponsesCandidatePayloadParts> {
|
||||
) -> Result<Option<LocalOpenAiResponsesCandidatePayloadParts>, GatewayError> {
|
||||
let spec_metadata = local_openai_responses_spec_metadata(spec);
|
||||
let client_api_format = spec_metadata.api_format.trim().to_ascii_lowercase();
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
@@ -123,7 +127,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
.eq_ignore_ascii_case("grok");
|
||||
|
||||
if !is_grok && provider_api_format.eq_ignore_ascii_case("openai:image") {
|
||||
return resolve_openai_responses_to_openai_image_payload_parts(
|
||||
return Ok(resolve_openai_responses_to_openai_image_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
@@ -134,7 +138,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
candidate_id,
|
||||
spec,
|
||||
)
|
||||
.await;
|
||||
.await);
|
||||
}
|
||||
let is_windsurf_cascade =
|
||||
provider_api_format == "openai:chat" && is_windsurf_provider_transport(transport);
|
||||
@@ -171,7 +175,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let oauth_context = OauthPreparationContext {
|
||||
@@ -199,7 +203,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
"transport_auth_unavailable",
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -240,7 +244,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -265,7 +269,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
};
|
||||
@@ -279,6 +283,16 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
Some(&input.requested_model),
|
||||
)
|
||||
.await;
|
||||
let redaction = resolve_provider_chat_pii_redaction(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
&input.auth_context,
|
||||
spec_metadata.api_format,
|
||||
candidate_id,
|
||||
)
|
||||
.await?;
|
||||
let body_json = redaction.body_json.as_ref();
|
||||
|
||||
let needs_bidirectional_conversion = !same_format && conversion_kind.is_some();
|
||||
let upstream_is_stream = resolve_upstream_is_stream_for_provider(
|
||||
@@ -357,7 +371,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
if let Some(mapping) =
|
||||
crate::system_features::reasoning_model_directive_mapping_for_api_format_and_model(
|
||||
@@ -380,6 +394,13 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
request_requires_body_stream_field(body_json, force_body_stream_field),
|
||||
);
|
||||
}
|
||||
apply_deepseek_tool_call_thinking_compat(
|
||||
&mut base_provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.base_url.as_str(),
|
||||
provider_api_format,
|
||||
Some(body_json),
|
||||
);
|
||||
let antigravity_auth = if is_antigravity {
|
||||
match classify_local_antigravity_request_support(
|
||||
transport,
|
||||
@@ -398,7 +419,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
"transport_unsupported",
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -429,7 +450,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
@@ -456,11 +477,12 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
upstream_is_stream,
|
||||
needs_bidirectional_conversion,
|
||||
kiro_auth,
|
||||
redaction.redacted,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
if is_windsurf_cascade {
|
||||
return build_windsurf_openai_responses_payload_parts(
|
||||
return Ok(build_windsurf_openai_responses_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
@@ -477,13 +499,14 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
auth_value,
|
||||
provider_request_body,
|
||||
upstream_is_stream,
|
||||
redaction.redacted,
|
||||
)
|
||||
.await;
|
||||
.await);
|
||||
}
|
||||
if provider_api_format == "gemini:generate_content"
|
||||
&& is_gemini_cli_provider_transport(transport)
|
||||
{
|
||||
return build_gemini_cli_openai_responses_payload_parts(
|
||||
return Ok(build_gemini_cli_openai_responses_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
@@ -500,8 +523,9 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
auth_value,
|
||||
provider_request_body,
|
||||
upstream_is_stream,
|
||||
redaction.redacted,
|
||||
)
|
||||
.await;
|
||||
.await);
|
||||
}
|
||||
|
||||
let Some(upstream_url) = (if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
@@ -537,7 +561,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
let extra_headers = antigravity_auth
|
||||
.as_ref()
|
||||
@@ -569,7 +593,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
crate::ai_serving::transport::StandardProviderRequestHeaders {
|
||||
headers,
|
||||
@@ -607,7 +631,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
resolved_headers
|
||||
};
|
||||
@@ -623,6 +647,10 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
}
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
redaction.redacted,
|
||||
);
|
||||
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats(spec_metadata.api_format, provider_api_format);
|
||||
@@ -651,7 +679,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
"gateway resolved local openai responses upstream url"
|
||||
);
|
||||
|
||||
Some(LocalOpenAiResponsesCandidatePayloadParts {
|
||||
Ok(Some(LocalOpenAiResponsesCandidatePayloadParts {
|
||||
auth_header: resolved_headers.auth_header,
|
||||
auth_value: resolved_headers.auth_value,
|
||||
mapped_model,
|
||||
@@ -672,7 +700,8 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile,
|
||||
image_request_summary: None,
|
||||
})
|
||||
request_redacted: redaction.redacted,
|
||||
}))
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
@@ -693,6 +722,7 @@ async fn build_gemini_cli_openai_responses_payload_parts(
|
||||
auth_value: String,
|
||||
gemini_request_body: Value,
|
||||
upstream_is_stream: bool,
|
||||
request_redacted: bool,
|
||||
) -> Option<LocalOpenAiResponsesCandidatePayloadParts> {
|
||||
let candidate = &eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
@@ -792,6 +822,10 @@ async fn build_gemini_cli_openai_responses_payload_parts(
|
||||
Some(trace_id),
|
||||
resolved.transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
request_redacted,
|
||||
);
|
||||
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats(client_api_format, provider_api_format);
|
||||
@@ -812,6 +846,7 @@ async fn build_gemini_cli_openai_responses_payload_parts(
|
||||
transport: resolved.transport,
|
||||
transport_profile: None,
|
||||
image_request_summary: None,
|
||||
request_redacted,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -833,6 +868,7 @@ async fn build_windsurf_openai_responses_payload_parts(
|
||||
auth_value: String,
|
||||
openai_chat_request_body: Value,
|
||||
upstream_is_stream: bool,
|
||||
request_redacted: bool,
|
||||
) -> Option<LocalOpenAiResponsesCandidatePayloadParts> {
|
||||
let candidate = &eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
@@ -888,7 +924,7 @@ async fn build_windsurf_openai_responses_payload_parts(
|
||||
return None;
|
||||
}
|
||||
};
|
||||
let provider_request_headers = match build_windsurf_cascade_headers(
|
||||
let mut provider_request_headers = match build_windsurf_cascade_headers(
|
||||
effective_headers,
|
||||
&provider_request_body,
|
||||
original_body_json,
|
||||
@@ -917,6 +953,10 @@ async fn build_windsurf_openai_responses_payload_parts(
|
||||
return None;
|
||||
}
|
||||
};
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
request_redacted,
|
||||
);
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats(client_api_format, provider_api_format);
|
||||
|
||||
@@ -936,6 +976,7 @@ async fn build_windsurf_openai_responses_payload_parts(
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
image_request_summary: None,
|
||||
request_redacted,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1125,6 +1166,7 @@ async fn resolve_openai_responses_to_openai_image_payload_parts(
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
image_request_summary: Some(image_request_summary),
|
||||
request_redacted: false,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1241,22 +1283,42 @@ fn build_chatgpt_web_image_provider_body_from_openai_responses_body(
|
||||
.unwrap_or("gpt-5-5-thinking");
|
||||
let image_urls = openai_image_inputs_as_urls(&images);
|
||||
|
||||
let body = json!({
|
||||
let mut body = json!({
|
||||
"operation": operation,
|
||||
"model": if model.is_empty() { "gpt-image-2" } else { model },
|
||||
"web_model": web_model,
|
||||
"prompt": prompt,
|
||||
"size": size,
|
||||
"ratio": chatgpt_web_ratio_for_size(size),
|
||||
"quality": quality,
|
||||
"output_format": output_format,
|
||||
"images": image_urls,
|
||||
});
|
||||
let summary = json!({
|
||||
if let Some(partial_images) = tool
|
||||
.as_ref()
|
||||
.and_then(|tool| tool.get("partial_images"))
|
||||
.or_else(|| object.get("partial_images"))
|
||||
.cloned()
|
||||
{
|
||||
body.as_object_mut()?
|
||||
.insert("partial_images".to_string(), partial_images);
|
||||
}
|
||||
let mut summary = json!({
|
||||
"operation": operation,
|
||||
"output_format": output_format,
|
||||
"size": size,
|
||||
"quality": quality,
|
||||
});
|
||||
if let Some(partial_images) = tool
|
||||
.as_ref()
|
||||
.and_then(|tool| tool.get("partial_images"))
|
||||
.or_else(|| object.get("partial_images"))
|
||||
.cloned()
|
||||
{
|
||||
summary
|
||||
.as_object_mut()?
|
||||
.insert("partial_images".to_string(), partial_images);
|
||||
}
|
||||
Some((body, summary))
|
||||
}
|
||||
|
||||
@@ -1428,7 +1490,8 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
upstream_is_stream: bool,
|
||||
needs_bidirectional_conversion: bool,
|
||||
kiro_auth: &KiroRequestAuth,
|
||||
) -> Option<LocalOpenAiResponsesCandidatePayloadParts> {
|
||||
request_redacted: bool,
|
||||
) -> Result<Option<LocalOpenAiResponsesCandidatePayloadParts>, GatewayError> {
|
||||
let candidate = &eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let provider_request_body = match build_kiro_provider_request_body(
|
||||
@@ -1455,7 +1518,7 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let upstream_url = match build_kiro_cross_format_upstream_url(
|
||||
@@ -1483,10 +1546,10 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
let mut provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
headers: effective_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: original_body_json,
|
||||
@@ -1513,7 +1576,7 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let (execution_strategy, conversion_mode) =
|
||||
@@ -1538,7 +1601,12 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
"gateway resolved local openai responses kiro upstream url"
|
||||
);
|
||||
|
||||
Some(LocalOpenAiResponsesCandidatePayloadParts {
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
request_redacted,
|
||||
);
|
||||
|
||||
Ok(Some(LocalOpenAiResponsesCandidatePayloadParts {
|
||||
auth_header,
|
||||
auth_value,
|
||||
mapped_model,
|
||||
@@ -1554,7 +1622,8 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
image_request_summary: None,
|
||||
})
|
||||
request_redacted,
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -1594,4 +1663,36 @@ mod tests {
|
||||
assert_eq!(summary["operation"], "generate");
|
||||
assert_eq!(summary["output_format"], "png");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chatgpt_web_responses_image_body_preserves_usage_options() {
|
||||
let body_json = json!({
|
||||
"model": "gpt-image-2",
|
||||
"input": "Draw a glass city",
|
||||
"tools": [
|
||||
{
|
||||
"type": "image_generation",
|
||||
"size": "1024x1024",
|
||||
"quality": "high",
|
||||
"output_format": "png",
|
||||
"partial_images": 2
|
||||
}
|
||||
],
|
||||
"tool_choice": {
|
||||
"type": "image_generation"
|
||||
}
|
||||
});
|
||||
|
||||
let (provider_body, summary) =
|
||||
build_chatgpt_web_image_provider_body_from_openai_responses_body(
|
||||
&body_json,
|
||||
"gpt-image-2",
|
||||
)
|
||||
.expect("responses image body should convert");
|
||||
|
||||
assert_eq!(provider_body["quality"], "high");
|
||||
assert_eq!(provider_body["partial_images"], 2);
|
||||
assert_eq!(summary["quality"], "high");
|
||||
assert_eq!(summary["partial_images"], 2);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,6 +17,14 @@ const AI_POST_ROUTE_PATTERNS: &[&str] = &[
|
||||
"/v1/responses/compact",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1internal:loadCodeAssist",
|
||||
"/v1internal:fetchAvailableModels",
|
||||
"/v1internal:fetchUserInfo",
|
||||
"/v1internal:fetchAdminControls",
|
||||
"/v1internal:setUserSettings",
|
||||
"/v1internal:listExperiments",
|
||||
"/v1internal:recordCodeAssistMetrics",
|
||||
"/v1internal:streamGenerateContent",
|
||||
];
|
||||
|
||||
const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
|
||||
|
||||
@@ -26,7 +26,6 @@ pub(crate) fn mount_public_support_routes(router: Router<AppState>) -> Router<Ap
|
||||
.route("/api/capabilities/model/{*model_path}", get(proxy_request))
|
||||
.route("/install/{*install_path}", get(proxy_request))
|
||||
.route("/install-tunnel/{*install_path}", get(proxy_request))
|
||||
.route("/install-proxy/{*install_path}", get(proxy_request))
|
||||
.route("/i/{*install_path}", get(proxy_request))
|
||||
.route("/", get(proxy_request))
|
||||
}
|
||||
|
||||
@@ -20,6 +20,12 @@ pub(crate) fn mount_core_routes(router: Router<AppState>) -> Router<AppState> {
|
||||
.route("/_gateway/health", get(health))
|
||||
}
|
||||
|
||||
fn current_gateway_version() -> &'static str {
|
||||
option_env!("AETHER_BUILD_VERSION")
|
||||
.filter(|version| !version.is_empty())
|
||||
.unwrap_or(env!("CARGO_PKG_VERSION"))
|
||||
}
|
||||
|
||||
pub(crate) async fn health(State(state): State<AppState>) -> impl IntoResponse {
|
||||
let request_concurrency = state.request_concurrency_snapshot().map(|snapshot| {
|
||||
json!({
|
||||
@@ -78,7 +84,7 @@ pub(crate) async fn frontdoor_manifest(State(state): State<AppState>) -> impl In
|
||||
Json(json!({
|
||||
"component": "aether-gateway",
|
||||
"manifest_version": FRONTDOOR_MANIFEST_VERSION,
|
||||
"version": env!("CARGO_PKG_VERSION"),
|
||||
"version": current_gateway_version(),
|
||||
"mode": "compatibility_frontdoor",
|
||||
"entrypoints": {
|
||||
"public_manifest": FRONTDOOR_MANIFEST_PATH,
|
||||
|
||||
@@ -0,0 +1,398 @@
|
||||
use std::fmt;
|
||||
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::schedule::{BackupSchedule, BackupScheduleUnit};
|
||||
use super::scopes::BackupScope;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct S3BackupConfig {
|
||||
pub(crate) enabled: bool,
|
||||
pub(crate) scope: BackupScope,
|
||||
pub(crate) endpoint: String,
|
||||
pub(crate) region: String,
|
||||
pub(crate) bucket: String,
|
||||
pub(crate) prefix: String,
|
||||
pub(crate) access_key_id: String,
|
||||
pub(crate) secret_access_key: String,
|
||||
pub(crate) path_style: bool,
|
||||
pub(crate) compression: String,
|
||||
pub(crate) schedule: BackupSchedule,
|
||||
pub(crate) retention_count: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct BackupConfigError {
|
||||
message: String,
|
||||
}
|
||||
|
||||
impl BackupConfigError {
|
||||
fn new(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for BackupConfigError {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(&self.message)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for BackupConfigError {}
|
||||
|
||||
impl S3BackupConfig {
|
||||
pub(crate) fn from_json_map(entries: &Map<String, Value>) -> Result<Self, BackupConfigError> {
|
||||
let enabled = optional_bool(entries, "backup_s3_enabled")?.unwrap_or(false);
|
||||
let mut schedule = BackupSchedule::default();
|
||||
schedule.unit = optional_string(entries, "backup_s3_schedule_unit")?
|
||||
.map(|value| {
|
||||
BackupScheduleUnit::from_config_value(&value)
|
||||
.ok_or_else(|| BackupConfigError::new("Schedule Unit(计划单位)配置值无效"))
|
||||
})
|
||||
.transpose()?
|
||||
.unwrap_or(schedule.unit);
|
||||
schedule.interval =
|
||||
optional_u32(entries, "backup_s3_schedule_interval")?.unwrap_or(schedule.interval);
|
||||
schedule.minute =
|
||||
optional_u32(entries, "backup_s3_schedule_minute")?.unwrap_or(schedule.minute);
|
||||
schedule.hour = optional_u32(entries, "backup_s3_schedule_hour")?.unwrap_or(schedule.hour);
|
||||
schedule.weekday =
|
||||
optional_u32(entries, "backup_s3_schedule_weekday")?.unwrap_or(schedule.weekday);
|
||||
schedule.month_day =
|
||||
optional_u32(entries, "backup_s3_schedule_month_day")?.unwrap_or(schedule.month_day);
|
||||
validate_range("Interval(计划间隔)", schedule.interval, 1, u32::MAX)?;
|
||||
validate_range("Minute(计划分钟)", schedule.minute, 0, 59)?;
|
||||
validate_range("Hour(计划小时)", schedule.hour, 0, 23)?;
|
||||
validate_range("Weekday(计划星期)", schedule.weekday, 1, 7)?;
|
||||
validate_range("Month Day(计划月日)", schedule.month_day, 1, 31)?;
|
||||
|
||||
let scope = optional_string(entries, "backup_s3_scope")?
|
||||
.map(|value| {
|
||||
BackupScope::from_config_value(&value)
|
||||
.ok_or_else(|| BackupConfigError::new("Scope(备份范围)配置值无效"))
|
||||
})
|
||||
.transpose()?
|
||||
.unwrap_or(BackupScope::Data);
|
||||
let retention_count = optional_u32(entries, "backup_s3_retention_count")?.unwrap_or(7);
|
||||
validate_range("Retention(保留份数)", retention_count, 1, u32::MAX)?;
|
||||
let endpoint = required_or_disabled_string(
|
||||
entries,
|
||||
"backup_s3_endpoint",
|
||||
"Endpoint(S3 地址)",
|
||||
enabled,
|
||||
)?;
|
||||
let bucket =
|
||||
required_or_disabled_string(entries, "backup_s3_bucket", "Bucket(存储桶)", enabled)?;
|
||||
let access_key_id = required_or_disabled_string(
|
||||
entries,
|
||||
"backup_s3_access_key_id",
|
||||
"Access Key ID(访问密钥 ID)",
|
||||
enabled,
|
||||
)?;
|
||||
let secret_access_key = required_or_disabled_string(
|
||||
entries,
|
||||
"backup_s3_secret_access_key",
|
||||
"Secret Access Key(访问密钥)",
|
||||
enabled,
|
||||
)?;
|
||||
|
||||
Ok(Self {
|
||||
enabled,
|
||||
scope,
|
||||
endpoint,
|
||||
region: optional_string(entries, "backup_s3_region")?
|
||||
.unwrap_or_else(|| "auto".to_string()),
|
||||
bucket,
|
||||
prefix: optional_string(entries, "backup_s3_prefix")?
|
||||
.unwrap_or_else(|| "aether/backups/".to_string()),
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
path_style: optional_bool(entries, "backup_s3_path_style")?.unwrap_or(true),
|
||||
compression: optional_string(entries, "backup_s3_compression")?
|
||||
.unwrap_or_else(|| "zstd".to_string()),
|
||||
schedule,
|
||||
retention_count,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_range(label: &str, value: u32, min: u32, max: u32) -> Result<(), BackupConfigError> {
|
||||
if (min..=max).contains(&value) {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(BackupConfigError::new(format!(
|
||||
"{label}配置值无效,应在 {min}..={max} 范围内"
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
fn required_or_disabled_string(
|
||||
entries: &Map<String, Value>,
|
||||
key: &str,
|
||||
label: &str,
|
||||
enabled: bool,
|
||||
) -> Result<String, BackupConfigError> {
|
||||
if enabled {
|
||||
required_string(entries, key, label)
|
||||
} else {
|
||||
Ok(optional_string(entries, key)?.unwrap_or_default())
|
||||
}
|
||||
}
|
||||
|
||||
fn required_string(
|
||||
entries: &Map<String, Value>,
|
||||
key: &str,
|
||||
label: &str,
|
||||
) -> Result<String, BackupConfigError> {
|
||||
optional_string(entries, key)?
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| BackupConfigError::new(format!("{label}为必填配置")))
|
||||
}
|
||||
|
||||
fn optional_string(
|
||||
entries: &Map<String, Value>,
|
||||
key: &str,
|
||||
) -> Result<Option<String>, BackupConfigError> {
|
||||
let Some(value) = entries.get(key) else {
|
||||
return Ok(None);
|
||||
};
|
||||
match value {
|
||||
Value::String(value) => {
|
||||
let trimmed = value.trim();
|
||||
if trimmed.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(Some(trimmed.to_string()))
|
||||
}
|
||||
}
|
||||
Value::Null => Ok(None),
|
||||
_ => Err(BackupConfigError::new(format!(
|
||||
"{} 字符串值无效",
|
||||
config_label(key)
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn config_label(key: &str) -> &str {
|
||||
match key {
|
||||
"backup_s3_endpoint" => "Endpoint(S3 地址)",
|
||||
"backup_s3_region" => "Region(S3 区域)",
|
||||
"backup_s3_bucket" => "Bucket(存储桶)",
|
||||
"backup_s3_prefix" => "Prefix(备份前缀)",
|
||||
"backup_s3_access_key_id" => "Access Key ID(访问密钥 ID)",
|
||||
"backup_s3_secret_access_key" => "Secret Access Key(访问密钥)",
|
||||
"backup_s3_compression" => "Compression(压缩格式)",
|
||||
"backup_s3_scope" => "Scope(备份范围)",
|
||||
"backup_s3_schedule_unit" => "Schedule Unit(计划单位)",
|
||||
_ => key,
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_bool(
|
||||
entries: &Map<String, Value>,
|
||||
key: &str,
|
||||
) -> Result<Option<bool>, BackupConfigError> {
|
||||
let Some(value) = entries.get(key) else {
|
||||
return Ok(None);
|
||||
};
|
||||
match value {
|
||||
Value::Bool(value) => Ok(Some(*value)),
|
||||
Value::String(value) => match value.trim() {
|
||||
"true" => Ok(Some(true)),
|
||||
"false" => Ok(Some(false)),
|
||||
"" => Ok(None),
|
||||
_ => Err(BackupConfigError::new(format!("{key} 布尔值无效"))),
|
||||
},
|
||||
Value::Null => Ok(None),
|
||||
_ => Err(BackupConfigError::new(format!("{key} 布尔值无效"))),
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_u32(entries: &Map<String, Value>, key: &str) -> Result<Option<u32>, BackupConfigError> {
|
||||
let Some(value) = entries.get(key) else {
|
||||
return Ok(None);
|
||||
};
|
||||
match value {
|
||||
Value::Number(value) => value
|
||||
.as_u64()
|
||||
.and_then(|value| u32::try_from(value).ok())
|
||||
.map(Some)
|
||||
.ok_or_else(|| BackupConfigError::new(format!("{key} 数值无效"))),
|
||||
Value::String(value) => {
|
||||
let trimmed = value.trim();
|
||||
if trimmed.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
trimmed
|
||||
.parse::<u32>()
|
||||
.map(Some)
|
||||
.map_err(|_| BackupConfigError::new(format!("{key} 数值无效")))
|
||||
}
|
||||
}
|
||||
Value::Null => Ok(None),
|
||||
_ => Err(BackupConfigError::new(format!("{key} 数值无效"))),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::schedule::BackupScheduleUnit;
|
||||
use super::super::scopes::BackupScope;
|
||||
use super::S3BackupConfig;
|
||||
|
||||
#[test]
|
||||
fn parses_minimal_valid_s3_backup_config() {
|
||||
let entries = serde_json::json!({
|
||||
"backup_s3_enabled": true,
|
||||
"backup_s3_scope": "data",
|
||||
"backup_s3_endpoint": "https://s3.example.com",
|
||||
"backup_s3_region": "auto",
|
||||
"backup_s3_bucket": "aether-backups",
|
||||
"backup_s3_prefix": "prod/",
|
||||
"backup_s3_access_key_id": "access",
|
||||
"backup_s3_secret_access_key": "secret",
|
||||
"backup_s3_path_style": true,
|
||||
"backup_s3_compression": "zstd",
|
||||
"backup_s3_schedule_unit": "days",
|
||||
"backup_s3_schedule_interval": 1,
|
||||
"backup_s3_schedule_hour": 3,
|
||||
"backup_s3_schedule_minute": 15,
|
||||
"backup_s3_retention_count": 7
|
||||
});
|
||||
|
||||
let config = S3BackupConfig::from_json_map(entries.as_object().unwrap())
|
||||
.expect("config should parse");
|
||||
|
||||
assert_eq!(config.scope, BackupScope::Data);
|
||||
assert_eq!(config.bucket, "aether-backups");
|
||||
assert_eq!(config.prefix, "prod/");
|
||||
assert_eq!(config.schedule.unit, BackupScheduleUnit::Days);
|
||||
assert_eq!(config.retention_count, 7);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_missing_bucket_for_backup() {
|
||||
let entries = serde_json::json!({
|
||||
"backup_s3_enabled": true,
|
||||
"backup_s3_endpoint": "https://s3.example.com",
|
||||
"backup_s3_access_key_id": "access",
|
||||
"backup_s3_secret_access_key": "secret"
|
||||
});
|
||||
|
||||
let err = S3BackupConfig::from_json_map(entries.as_object().unwrap())
|
||||
.expect_err("bucket is required");
|
||||
|
||||
assert!(err.to_string().contains("Bucket"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_disabled_default_s3_backup_config_with_null_credentials() {
|
||||
let entries = serde_json::json!({
|
||||
"backup_s3_enabled": false,
|
||||
"backup_s3_scope": "data",
|
||||
"backup_s3_endpoint": null,
|
||||
"backup_s3_region": "auto",
|
||||
"backup_s3_bucket": null,
|
||||
"backup_s3_prefix": "aether/backups/",
|
||||
"backup_s3_access_key_id": null,
|
||||
"backup_s3_secret_access_key": null,
|
||||
"backup_s3_path_style": true,
|
||||
"backup_s3_compression": "zstd",
|
||||
"backup_s3_schedule_unit": "days",
|
||||
"backup_s3_schedule_interval": 1,
|
||||
"backup_s3_schedule_hour": 3,
|
||||
"backup_s3_schedule_minute": 0,
|
||||
"backup_s3_schedule_weekday": 1,
|
||||
"backup_s3_schedule_month_day": 1,
|
||||
"backup_s3_retention_count": 7
|
||||
});
|
||||
|
||||
let config = S3BackupConfig::from_json_map(entries.as_object().unwrap())
|
||||
.expect("disabled default config should parse");
|
||||
|
||||
assert_eq!(config.enabled, false);
|
||||
assert_eq!(config.endpoint, "");
|
||||
assert_eq!(config.bucket, "");
|
||||
assert_eq!(config.access_key_id, "");
|
||||
assert_eq!(config.secret_access_key, "");
|
||||
assert_eq!(config.schedule.unit, BackupScheduleUnit::Days);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_schedule_numbers() {
|
||||
let cases = [
|
||||
("backup_s3_schedule_interval", 0, "Interval"),
|
||||
("backup_s3_schedule_minute", 60, "Minute"),
|
||||
("backup_s3_schedule_hour", 24, "Hour"),
|
||||
("backup_s3_schedule_weekday", 0, "Weekday"),
|
||||
("backup_s3_schedule_month_day", 32, "Month Day"),
|
||||
("backup_s3_retention_count", 0, "Retention"),
|
||||
];
|
||||
|
||||
for (key, value, label) in cases {
|
||||
let mut entries = serde_json::json!({
|
||||
"backup_s3_enabled": true,
|
||||
"backup_s3_endpoint": "https://s3.example.com",
|
||||
"backup_s3_bucket": "aether-backups",
|
||||
"backup_s3_access_key_id": "access",
|
||||
"backup_s3_secret_access_key": "secret"
|
||||
});
|
||||
entries.as_object_mut().unwrap().insert(
|
||||
key.to_string(),
|
||||
serde_json::Value::Number(serde_json::Number::from(value)),
|
||||
);
|
||||
|
||||
let err = S3BackupConfig::from_json_map(entries.as_object().unwrap())
|
||||
.expect_err("invalid numeric config should fail");
|
||||
|
||||
assert!(
|
||||
err.to_string().contains(label),
|
||||
"{key} should mention {label}, got {err}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_string_endpoint_config() {
|
||||
let entries = serde_json::json!({
|
||||
"backup_s3_enabled": true,
|
||||
"backup_s3_endpoint": {"url": "https://s3.example.com"},
|
||||
"backup_s3_bucket": "aether-backups",
|
||||
"backup_s3_access_key_id": "access",
|
||||
"backup_s3_secret_access_key": "secret"
|
||||
});
|
||||
|
||||
let err = S3BackupConfig::from_json_map(entries.as_object().unwrap())
|
||||
.expect_err("endpoint object should fail");
|
||||
|
||||
assert!(err.to_string().contains("Endpoint"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn applies_default_values_from_system_config_contract() {
|
||||
let entries = serde_json::json!({
|
||||
"backup_s3_endpoint": "https://s3.example.com",
|
||||
"backup_s3_bucket": "aether-backups",
|
||||
"backup_s3_access_key_id": "access",
|
||||
"backup_s3_secret_access_key": "secret"
|
||||
});
|
||||
|
||||
let config = S3BackupConfig::from_json_map(entries.as_object().unwrap())
|
||||
.expect("config should parse with defaults");
|
||||
|
||||
assert_eq!(config.scope, BackupScope::Data);
|
||||
assert_eq!(config.region, "auto");
|
||||
assert_eq!(config.prefix, "aether/backups/");
|
||||
assert_eq!(config.path_style, true);
|
||||
assert_eq!(config.compression, "zstd");
|
||||
assert_eq!(config.schedule.unit, BackupScheduleUnit::Days);
|
||||
assert_eq!(config.schedule.interval, 1);
|
||||
assert_eq!(config.schedule.hour, 3);
|
||||
assert_eq!(config.schedule.minute, 0);
|
||||
assert_eq!(config.schedule.weekday, 1);
|
||||
assert_eq!(config.schedule.month_day, 1);
|
||||
assert_eq!(config.retention_count, 7);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,224 @@
|
||||
use bytes::Bytes;
|
||||
use chrono::{DateTime, SecondsFormat, Utc};
|
||||
use serde_json::Value;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::config::S3BackupConfig;
|
||||
use super::scopes::BackupScope;
|
||||
use super::store::{BackupObjectStore, BackupStoreError};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct BackupRunResult {
|
||||
pub(crate) scope: BackupScope,
|
||||
pub(crate) bucket: String,
|
||||
pub(crate) object_key: String,
|
||||
pub(crate) bytes: usize,
|
||||
pub(crate) sha256: String,
|
||||
pub(crate) export_version: String,
|
||||
pub(crate) exported_at: String,
|
||||
pub(crate) compression: String,
|
||||
pub(crate) deleted_old_objects: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub(crate) enum BackupExecutionError {
|
||||
#[error("S3 backup JSON serialization failed: {0}")]
|
||||
Json(#[from] serde_json::Error),
|
||||
|
||||
#[error("S3 backup compression failed: {0}")]
|
||||
Compression(#[from] std::io::Error),
|
||||
|
||||
#[error("{0}")]
|
||||
Store(#[from] BackupStoreError),
|
||||
|
||||
#[error("S3 backup compression `{0}` is not supported; expected `zstd`")]
|
||||
UnknownCompression(String),
|
||||
}
|
||||
|
||||
pub(crate) async fn run_backup_with_store<S>(
|
||||
config: &S3BackupConfig,
|
||||
store: &S,
|
||||
payload: Value,
|
||||
now_utc: DateTime<Utc>,
|
||||
) -> Result<BackupRunResult, BackupExecutionError>
|
||||
where
|
||||
S: BackupObjectStore + ?Sized,
|
||||
{
|
||||
let export_version = payload_string_field(&payload, "version").unwrap_or_default();
|
||||
let exported_at = payload_string_field(&payload, "exported_at")
|
||||
.unwrap_or_else(|| now_utc.to_rfc3339_opts(SecondsFormat::Secs, true));
|
||||
let json_bytes = serde_json::to_vec(&payload)?;
|
||||
let compression = config.compression.trim().to_string();
|
||||
let upload_bytes = match compression.as_str() {
|
||||
"zstd" => zstd::stream::encode_all(json_bytes.as_slice(), 0)?,
|
||||
other => return Err(BackupExecutionError::UnknownCompression(other.to_string())),
|
||||
};
|
||||
let bytes = upload_bytes.len();
|
||||
let sha256 = format!("{:x}", Sha256::digest(&upload_bytes));
|
||||
let timestamp = now_utc.format("%Y%m%d-%H%M%S").to_string();
|
||||
let object_key = config.scope.object_key(&config.prefix, ×tamp);
|
||||
|
||||
store
|
||||
.put_object(&object_key, Bytes::from(upload_bytes))
|
||||
.await?;
|
||||
let deleted_old_objects = prune_old_backups(config, store, &object_key).await?;
|
||||
|
||||
Ok(BackupRunResult {
|
||||
scope: config.scope,
|
||||
bucket: config.bucket.clone(),
|
||||
object_key,
|
||||
bytes,
|
||||
sha256,
|
||||
export_version,
|
||||
exported_at,
|
||||
compression,
|
||||
deleted_old_objects,
|
||||
})
|
||||
}
|
||||
|
||||
async fn prune_old_backups<S>(
|
||||
config: &S3BackupConfig,
|
||||
store: &S,
|
||||
current_object_key: &str,
|
||||
) -> Result<usize, BackupExecutionError>
|
||||
where
|
||||
S: BackupObjectStore + ?Sized,
|
||||
{
|
||||
let keys = store.list_keys(&config.prefix).await?;
|
||||
let mut matching_keys = config.scope.matching_backup_keys(&config.prefix, keys);
|
||||
matching_keys.sort_by(|left, right| right.cmp(left));
|
||||
|
||||
let mut deleted = 0;
|
||||
let mut retained = usize::from(
|
||||
config.retention_count > 0 && matching_keys.iter().any(|key| key == current_object_key),
|
||||
);
|
||||
for key in matching_keys {
|
||||
if config.retention_count > 0 && key == current_object_key {
|
||||
continue;
|
||||
}
|
||||
if retained < config.retention_count as usize {
|
||||
retained += 1;
|
||||
continue;
|
||||
}
|
||||
|
||||
store.delete_object(&key).await?;
|
||||
deleted += 1;
|
||||
}
|
||||
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
fn payload_string_field(payload: &Value, field: &str) -> Option<String> {
|
||||
payload
|
||||
.get(field)
|
||||
.and_then(Value::as_str)
|
||||
.map(ToString::to_string)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::config::S3BackupConfig;
|
||||
use super::super::schedule::BackupSchedule;
|
||||
use super::super::scopes::BackupScope;
|
||||
use super::super::store::{BackupObjectStore, FakeBackupObjectStore};
|
||||
use super::run_backup_with_store;
|
||||
use bytes::Bytes;
|
||||
use chrono::{DateTime, Utc};
|
||||
use serde_json::json;
|
||||
|
||||
fn sample_backup_config(scope: BackupScope, retention_count: u32) -> S3BackupConfig {
|
||||
S3BackupConfig {
|
||||
enabled: true,
|
||||
scope,
|
||||
endpoint: "https://example.com".to_string(),
|
||||
region: "auto".to_string(),
|
||||
bucket: "aether-backups".to_string(),
|
||||
prefix: "prod/".to_string(),
|
||||
access_key_id: "test-access-key".to_string(),
|
||||
secret_access_key: "test-secret-key".to_string(),
|
||||
path_style: true,
|
||||
compression: "zstd".to_string(),
|
||||
schedule: BackupSchedule::default(),
|
||||
retention_count,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn backup_executor_uploads_payload_and_prunes_same_scope_only() {
|
||||
let store = FakeBackupObjectStore::default();
|
||||
store
|
||||
.put_object(
|
||||
"prod/aether-data-backup-20260524-010000.json.zst",
|
||||
Bytes::from_static(b"old"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.put_object(
|
||||
"prod/aether-config-backup-20260524-010000.json.zst",
|
||||
Bytes::from_static(b"keep-config"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let config = sample_backup_config(BackupScope::Data, 1);
|
||||
let payload = json!({
|
||||
"version": "1.0",
|
||||
"exported_at": "2026-05-24T03:15:00Z",
|
||||
"config_data": {},
|
||||
"user_data": {}
|
||||
});
|
||||
let now_utc = DateTime::parse_from_rfc3339("2026-05-24T03:15:00+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&Utc);
|
||||
|
||||
let result = run_backup_with_store(&config, &store, payload, now_utc)
|
||||
.await
|
||||
.expect("backup should succeed");
|
||||
|
||||
assert_eq!(result.scope, BackupScope::Data);
|
||||
assert_eq!(result.bucket, "aether-backups");
|
||||
assert_eq!(
|
||||
result.object_key,
|
||||
"prod/aether-data-backup-20260523-191500.json.zst"
|
||||
);
|
||||
assert!(result.bytes > 0);
|
||||
assert_eq!(result.sha256.len(), 64);
|
||||
assert_eq!(result.export_version, "1.0");
|
||||
assert_eq!(result.exported_at, "2026-05-24T03:15:00Z");
|
||||
assert_eq!(result.compression, "zstd");
|
||||
assert_eq!(result.deleted_old_objects, 1);
|
||||
|
||||
let keys = store.list_keys("prod/").await.unwrap();
|
||||
assert!(keys
|
||||
.iter()
|
||||
.any(|key| key == "prod/aether-config-backup-20260524-010000.json.zst"));
|
||||
assert!(keys
|
||||
.iter()
|
||||
.any(|key| key == "prod/aether-data-backup-20260523-191500.json.zst"));
|
||||
assert!(!keys
|
||||
.iter()
|
||||
.any(|key| key == "prod/aether-data-backup-20260524-010000.json.zst"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn backup_executor_rejects_unknown_compression() {
|
||||
let store = FakeBackupObjectStore::default();
|
||||
let mut config = sample_backup_config(BackupScope::Data, 1);
|
||||
config.compression = "brotli".to_string();
|
||||
|
||||
let payload = json!({
|
||||
"version": "1.0",
|
||||
"exported_at": "2026-05-24T03:15:00Z"
|
||||
});
|
||||
let now_utc = DateTime::parse_from_rfc3339("2026-05-24T03:15:00+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&Utc);
|
||||
|
||||
let error = run_backup_with_store(&config, &store, payload, now_utc)
|
||||
.await
|
||||
.expect_err("unknown compression should fail");
|
||||
|
||||
assert!(error.to_string().contains("brotli"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
pub(crate) mod config;
|
||||
pub(crate) mod executor;
|
||||
pub(crate) mod schedule;
|
||||
pub(crate) mod scopes;
|
||||
pub(crate) mod store;
|
||||
pub(crate) mod task;
|
||||
pub(crate) mod worker;
|
||||
|
||||
pub(crate) const S3_BACKUP_ENABLED_KEY: &str = "backup_s3_enabled";
|
||||
pub(crate) const S3_BACKUP_LAST_SLOT_KEY: &str = "backup_s3_last_slot";
|
||||
@@ -0,0 +1,366 @@
|
||||
use chrono::{DateTime, Datelike, NaiveDate, TimeZone, Timelike, Utc};
|
||||
use chrono_tz::Tz;
|
||||
|
||||
const BACKUP_SCHEDULE_DEFAULT_TIMEZONE: &str = "Asia/Shanghai";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) enum BackupScheduleUnit {
|
||||
Hours,
|
||||
Days,
|
||||
Weeks,
|
||||
Months,
|
||||
}
|
||||
|
||||
impl BackupScheduleUnit {
|
||||
pub(crate) fn from_config_value(value: &str) -> Option<Self> {
|
||||
match value.trim() {
|
||||
"hours" => Some(Self::Hours),
|
||||
"days" => Some(Self::Days),
|
||||
"weeks" => Some(Self::Weeks),
|
||||
"months" => Some(Self::Months),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn slot_prefix(self) -> &'static str {
|
||||
match self {
|
||||
Self::Hours => "hours",
|
||||
Self::Days => "days",
|
||||
Self::Weeks => "weeks",
|
||||
Self::Months => "months",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) struct BackupSchedule {
|
||||
pub(crate) unit: BackupScheduleUnit,
|
||||
pub(crate) interval: u32,
|
||||
pub(crate) minute: u32,
|
||||
pub(crate) hour: u32,
|
||||
pub(crate) weekday: u32,
|
||||
pub(crate) month_day: u32,
|
||||
}
|
||||
|
||||
impl Default for BackupSchedule {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
unit: BackupScheduleUnit::Days,
|
||||
interval: 1,
|
||||
minute: 0,
|
||||
hour: 3,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl BackupSchedule {
|
||||
pub(crate) fn due_slot(&self, now_utc: DateTime<Utc>) -> Option<String> {
|
||||
let timezone = backup_schedule_timezone();
|
||||
let local_now = now_utc.with_timezone(&timezone);
|
||||
let interval = self.interval.max(1);
|
||||
if local_now.minute() != self.minute {
|
||||
return None;
|
||||
}
|
||||
|
||||
let due = match self.unit {
|
||||
BackupScheduleUnit::Hours => {
|
||||
(local_epoch_hour(local_now.date_naive(), local_now.hour()) - i64::from(self.hour))
|
||||
.rem_euclid(i64::from(interval))
|
||||
== 0
|
||||
}
|
||||
BackupScheduleUnit::Days => {
|
||||
local_now.hour() == self.hour
|
||||
&& local_epoch_day(local_now.date_naive()) % i64::from(interval) == 0
|
||||
}
|
||||
BackupScheduleUnit::Weeks => {
|
||||
local_now.hour() == self.hour
|
||||
&& local_now.weekday().number_from_monday() == self.weekday
|
||||
&& local_epoch_week(local_now.date_naive()) % i64::from(interval) == 0
|
||||
}
|
||||
BackupScheduleUnit::Months => {
|
||||
local_now.hour() == self.hour
|
||||
&& local_now.day() == self.month_day
|
||||
&& month_ordinal(local_now.year(), local_now.month0()) % i64::from(interval)
|
||||
== 0
|
||||
}
|
||||
};
|
||||
if !due {
|
||||
return None;
|
||||
}
|
||||
|
||||
let slot = Utc.from_utc_datetime(&now_utc.date_naive().and_hms_opt(
|
||||
now_utc.hour(),
|
||||
now_utc.minute(),
|
||||
0,
|
||||
)?);
|
||||
Some(format!(
|
||||
"{}:{}",
|
||||
self.unit.slot_prefix(),
|
||||
slot.to_rfc3339_opts(chrono::SecondsFormat::Secs, true)
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn backup_schedule_timezone() -> Tz {
|
||||
std::env::var("APP_TIMEZONE")
|
||||
.ok()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.as_deref()
|
||||
.unwrap_or(BACKUP_SCHEDULE_DEFAULT_TIMEZONE)
|
||||
.parse()
|
||||
.unwrap_or(chrono_tz::Asia::Shanghai)
|
||||
}
|
||||
|
||||
fn local_epoch_day(date: NaiveDate) -> i64 {
|
||||
let epoch = NaiveDate::from_ymd_opt(1970, 1, 1).expect("unix epoch date should be valid");
|
||||
date.signed_duration_since(epoch).num_days()
|
||||
}
|
||||
|
||||
fn local_epoch_week(date: NaiveDate) -> i64 {
|
||||
local_epoch_day(date).div_euclid(7)
|
||||
}
|
||||
|
||||
fn local_epoch_hour(date: NaiveDate, hour: u32) -> i64 {
|
||||
local_epoch_day(date) * 24 + i64::from(hour)
|
||||
}
|
||||
|
||||
fn month_ordinal(year: i32, month0: u32) -> i64 {
|
||||
i64::from(year) * 12 + i64::from(month0)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{BackupSchedule, BackupScheduleUnit};
|
||||
use std::ffi::OsString;
|
||||
use std::sync::{Mutex, OnceLock};
|
||||
|
||||
struct TimezoneEnvGuard {
|
||||
previous: Option<OsString>,
|
||||
}
|
||||
|
||||
impl TimezoneEnvGuard {
|
||||
fn set(value: Option<&str>) -> Self {
|
||||
let previous = std::env::var_os("APP_TIMEZONE");
|
||||
match value {
|
||||
Some(value) => std::env::set_var("APP_TIMEZONE", value),
|
||||
None => std::env::remove_var("APP_TIMEZONE"),
|
||||
}
|
||||
Self { previous }
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TimezoneEnvGuard {
|
||||
fn drop(&mut self) {
|
||||
match &self.previous {
|
||||
Some(value) => std::env::set_var("APP_TIMEZONE", value),
|
||||
None => std::env::remove_var("APP_TIMEZONE"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn timezone_env_lock() -> std::sync::MutexGuard<'static, ()> {
|
||||
static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
|
||||
LOCK.get_or_init(|| Mutex::new(()))
|
||||
.lock()
|
||||
.unwrap_or_else(|err| err.into_inner())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn hourly_schedule_returns_stable_slot_once_per_due_hour() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(None);
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Hours,
|
||||
interval: 6,
|
||||
minute: 10,
|
||||
hour: 0,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let now = chrono::DateTime::parse_from_rfc3339("2026-05-24T12:10:30+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(
|
||||
schedule.due_slot(now).as_deref(),
|
||||
Some("hours:2026-05-24T04:10:00Z")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn hourly_schedule_interval_does_not_reset_at_midnight() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(None);
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Hours,
|
||||
interval: 5,
|
||||
minute: 10,
|
||||
hour: 0,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let midnight = chrono::DateTime::parse_from_rfc3339("2026-05-25T00:10:30+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
let due_after_midnight = chrono::DateTime::parse_from_rfc3339("2026-05-25T03:10:30+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(schedule.due_slot(midnight), None);
|
||||
assert_eq!(
|
||||
schedule.due_slot(due_after_midnight).as_deref(),
|
||||
Some("hours:2026-05-24T19:10:00Z")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn hourly_schedule_supports_intervals_longer_than_one_day() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(None);
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Hours,
|
||||
interval: 25,
|
||||
minute: 10,
|
||||
hour: 0,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let daily_midnight = chrono::DateTime::parse_from_rfc3339("2026-05-25T00:10:30+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
let due_after_25_hours = chrono::DateTime::parse_from_rfc3339("2026-05-25T23:10:30+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(schedule.due_slot(daily_midnight), None);
|
||||
assert_eq!(
|
||||
schedule.due_slot(due_after_25_hours).as_deref(),
|
||||
Some("hours:2026-05-25T15:10:00Z")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn daily_schedule_uses_maintenance_timezone_for_hour_and_interval() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(None);
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Days,
|
||||
interval: 2,
|
||||
minute: 15,
|
||||
hour: 3,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let due = chrono::DateTime::parse_from_rfc3339("2026-05-23T03:15:45+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
let not_due = chrono::DateTime::parse_from_rfc3339("2026-05-24T03:15:45+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(
|
||||
schedule.due_slot(due).as_deref(),
|
||||
Some("days:2026-05-22T19:15:00Z")
|
||||
);
|
||||
assert_eq!(schedule.due_slot(not_due), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn weekly_schedule_uses_maintenance_timezone_for_weekday_and_interval() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(None);
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Weeks,
|
||||
interval: 2,
|
||||
minute: 30,
|
||||
hour: 5,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let due = chrono::DateTime::parse_from_rfc3339("2026-05-25T05:30:59+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
let wrong_week = chrono::DateTime::parse_from_rfc3339("2026-05-18T05:30:59+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(
|
||||
schedule.due_slot(due).as_deref(),
|
||||
Some("weeks:2026-05-24T21:30:00Z")
|
||||
);
|
||||
assert_eq!(schedule.due_slot(wrong_week), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn monthly_schedule_uses_maintenance_timezone_for_month_day_and_interval() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(None);
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Months,
|
||||
interval: 3,
|
||||
minute: 45,
|
||||
hour: 2,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let due = chrono::DateTime::parse_from_rfc3339("2026-04-01T02:45:01+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
let wrong_month = chrono::DateTime::parse_from_rfc3339("2026-05-01T02:45:01+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(
|
||||
schedule.due_slot(due).as_deref(),
|
||||
Some("months:2026-03-31T18:45:00Z")
|
||||
);
|
||||
assert_eq!(schedule.due_slot(wrong_month), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn hourly_schedule_uses_utc_when_app_timezone_is_utc() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(Some("UTC"));
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Hours,
|
||||
interval: 6,
|
||||
minute: 10,
|
||||
hour: 4,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let now = chrono::DateTime::parse_from_rfc3339("2026-05-24T04:10:30Z")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(
|
||||
schedule.due_slot(now).as_deref(),
|
||||
Some("hours:2026-05-24T04:10:00Z")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn daily_schedule_slot_uses_actual_instant_for_non_whole_hour_timezone() {
|
||||
let _guard = timezone_env_lock();
|
||||
let _env = TimezoneEnvGuard::set(Some("Asia/Kolkata"));
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Days,
|
||||
interval: 1,
|
||||
minute: 30,
|
||||
hour: 3,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let now = chrono::DateTime::parse_from_rfc3339("2026-05-24T03:30:45+05:30")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
|
||||
assert_eq!(
|
||||
schedule.due_slot(now).as_deref(),
|
||||
Some("days:2026-05-23T22:00:00Z")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,211 @@
|
||||
use std::fmt;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) enum BackupScope {
|
||||
Config,
|
||||
Users,
|
||||
Data,
|
||||
}
|
||||
|
||||
impl BackupScope {
|
||||
pub(crate) fn from_config_value(value: &str) -> Option<Self> {
|
||||
match value.trim() {
|
||||
"config" => Some(Self::Config),
|
||||
"users" => Some(Self::Users),
|
||||
"data" => Some(Self::Data),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn as_config_value(self) -> &'static str {
|
||||
match self {
|
||||
Self::Config => "config",
|
||||
Self::Users => "users",
|
||||
Self::Data => "data",
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn route_kind(self) -> &'static str {
|
||||
match self {
|
||||
Self::Config => "config_export",
|
||||
Self::Users => "users_export",
|
||||
Self::Data => "data_export",
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn file_stem(self) -> &'static str {
|
||||
match self {
|
||||
Self::Config => "aether-config-backup",
|
||||
Self::Users => "aether-users-backup",
|
||||
Self::Data => "aether-data-backup",
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn object_key(self, prefix: &str, timestamp: &str) -> String {
|
||||
let file_name = self.file_name(timestamp);
|
||||
let prefix = normalized_prefix(prefix);
|
||||
|
||||
if prefix.is_empty() {
|
||||
file_name
|
||||
} else {
|
||||
format!("{prefix}/{file_name}")
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn matching_backup_keys(
|
||||
self,
|
||||
prefix: &str,
|
||||
keys: impl IntoIterator<Item = String>,
|
||||
) -> Vec<String> {
|
||||
let normalized_prefix = normalized_prefix(prefix);
|
||||
let expected_prefix = if normalized_prefix.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!("{normalized_prefix}/")
|
||||
};
|
||||
let file_prefix = format!("{}-", self.file_stem());
|
||||
let file_suffix = ".json.zst";
|
||||
|
||||
keys.into_iter()
|
||||
.filter(|key| {
|
||||
let Some(file_name) = key.strip_prefix(&expected_prefix) else {
|
||||
return false;
|
||||
};
|
||||
if file_name.contains('/') {
|
||||
return false;
|
||||
}
|
||||
let Some(timestamp) = file_name
|
||||
.strip_prefix(&file_prefix)
|
||||
.and_then(|rest| rest.strip_suffix(file_suffix))
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
|
||||
is_aether_backup_timestamp(timestamp)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn file_name(self, timestamp: &str) -> String {
|
||||
format!("{}-{timestamp}.json.zst", self.file_stem())
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for BackupScope {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(self.as_config_value())
|
||||
}
|
||||
}
|
||||
|
||||
fn normalized_prefix(prefix: &str) -> &str {
|
||||
prefix.trim_end_matches('/')
|
||||
}
|
||||
|
||||
fn is_aether_backup_timestamp(timestamp: &str) -> bool {
|
||||
let bytes = timestamp.as_bytes();
|
||||
|
||||
bytes.len() == 15
|
||||
&& bytes[8] == b'-'
|
||||
&& bytes[..8].iter().all(|byte| byte.is_ascii_digit())
|
||||
&& bytes[9..].iter().all(|byte| byte.is_ascii_digit())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::BackupScope;
|
||||
|
||||
#[test]
|
||||
fn backup_scope_matches_export_routes_and_object_prefixes() {
|
||||
assert_eq!(BackupScope::Config.as_config_value(), "config");
|
||||
assert_eq!(BackupScope::Users.as_config_value(), "users");
|
||||
assert_eq!(BackupScope::Data.as_config_value(), "data");
|
||||
|
||||
assert_eq!(BackupScope::Config.route_kind(), "config_export");
|
||||
assert_eq!(BackupScope::Users.route_kind(), "users_export");
|
||||
assert_eq!(BackupScope::Data.route_kind(), "data_export");
|
||||
|
||||
assert_eq!(BackupScope::Config.file_stem(), "aether-config-backup");
|
||||
assert_eq!(BackupScope::Users.file_stem(), "aether-users-backup");
|
||||
assert_eq!(BackupScope::Data.file_stem(), "aether-data-backup");
|
||||
|
||||
assert_eq!(
|
||||
BackupScope::Config.object_key("prod/", "20260524-031500"),
|
||||
"prod/aether-config-backup-20260524-031500.json.zst"
|
||||
);
|
||||
assert_eq!(
|
||||
BackupScope::Users.object_key("prod/", "20260524-031500"),
|
||||
"prod/aether-users-backup-20260524-031500.json.zst"
|
||||
);
|
||||
assert_eq!(
|
||||
BackupScope::Data.object_key("prod/", "20260524-031500"),
|
||||
"prod/aether-data-backup-20260524-031500.json.zst"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn retention_filter_only_matches_same_scope() {
|
||||
let keys = vec![
|
||||
"prod/aether-config-backup-20260524-010000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-20260524-010000.json.zst".to_string(),
|
||||
"prod/aether-data-backup-20260524-010000.json.zst".to_string(),
|
||||
"prod/random.json.zst".to_string(),
|
||||
];
|
||||
|
||||
let matched = BackupScope::Users.matching_backup_keys("prod/", keys);
|
||||
|
||||
assert_eq!(
|
||||
matched,
|
||||
vec!["prod/aether-users-backup-20260524-010000.json.zst"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn retention_filter_requires_aether_timestamp_format() {
|
||||
let keys = vec![
|
||||
"prod/aether-users-backup-20260524-010000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-foo.json.zst".to_string(),
|
||||
"prod/aether-users-backup-2026052-010000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-202605240-010000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-20260524-01000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-20260524-0100000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-20260524010000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-2026052a-010000.json.zst".to_string(),
|
||||
"prod/aether-users-backup-20260524-01000x.json.zst".to_string(),
|
||||
];
|
||||
|
||||
let matched = BackupScope::Users.matching_backup_keys("prod/", keys);
|
||||
|
||||
assert_eq!(
|
||||
matched,
|
||||
vec!["prod/aether-users-backup-20260524-010000.json.zst"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backup_key_prefix_boundaries_are_exact() {
|
||||
assert_eq!(
|
||||
BackupScope::Config.object_key("", "20260524-031500"),
|
||||
"aether-config-backup-20260524-031500.json.zst"
|
||||
);
|
||||
assert_eq!(
|
||||
BackupScope::Config.object_key("prod", "20260524-031500"),
|
||||
"prod/aether-config-backup-20260524-031500.json.zst"
|
||||
);
|
||||
|
||||
let keys = vec![
|
||||
"prod/aether-config-backup-20260524-010000.json.zst".to_string(),
|
||||
"prod//aether-config-backup-20260524-010000.json.zst".to_string(),
|
||||
"prod-backups/aether-config-backup-20260524-010000.json.zst".to_string(),
|
||||
"prod/aether-config-backup-20260524-010000.json".to_string(),
|
||||
"prod/aether-config-backup-.json.zst".to_string(),
|
||||
"aether-config-backup-20260524-010000.json.zst".to_string(),
|
||||
];
|
||||
|
||||
let matched = BackupScope::Config.matching_backup_keys("prod", keys);
|
||||
|
||||
assert_eq!(
|
||||
matched,
|
||||
vec!["prod/aether-config-backup-20260524-010000.json.zst"]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,227 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fmt;
|
||||
use std::sync::Arc;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::TryStreamExt;
|
||||
use object_store::aws::AmazonS3Builder;
|
||||
use object_store::path::Path;
|
||||
use object_store::ObjectStore;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use super::config::S3BackupConfig;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
pub(crate) trait BackupObjectStore: Send + Sync {
|
||||
async fn put_object(&self, key: &str, bytes: Bytes) -> Result<(), BackupStoreError>;
|
||||
|
||||
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, BackupStoreError>;
|
||||
|
||||
async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError>;
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct BackupStoreError {
|
||||
message: String,
|
||||
}
|
||||
|
||||
impl BackupStoreError {
|
||||
fn new(message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn object_store(operation: &str, key: &str, error: impl fmt::Display) -> Self {
|
||||
Self::new(format!(
|
||||
"S3 backup object store {operation} failed for `{key}`: {error}"
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for BackupStoreError {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(&self.message)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for BackupStoreError {}
|
||||
|
||||
#[derive(Debug, Default, Clone)]
|
||||
pub(crate) struct FakeBackupObjectStore {
|
||||
objects: Arc<RwLock<BTreeMap<String, Bytes>>>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl BackupObjectStore for FakeBackupObjectStore {
|
||||
async fn put_object(&self, key: &str, bytes: Bytes) -> Result<(), BackupStoreError> {
|
||||
self.objects.write().await.insert(key.to_string(), bytes);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, BackupStoreError> {
|
||||
let prefix = directory_list_prefix(prefix);
|
||||
Ok(self
|
||||
.objects
|
||||
.read()
|
||||
.await
|
||||
.keys()
|
||||
.filter(|key| key.starts_with(&prefix))
|
||||
.cloned()
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> {
|
||||
self.objects.write().await.remove(key);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) struct ObjectStoreS3BackupStore {
|
||||
store: object_store::aws::AmazonS3,
|
||||
}
|
||||
|
||||
impl ObjectStoreS3BackupStore {
|
||||
pub(crate) fn from_config(config: &S3BackupConfig) -> Result<Self, BackupStoreError> {
|
||||
let store = AmazonS3Builder::new()
|
||||
.with_endpoint(config.endpoint.clone())
|
||||
.with_region(config.region.clone())
|
||||
.with_bucket_name(config.bucket.clone())
|
||||
.with_access_key_id(config.access_key_id.clone())
|
||||
.with_secret_access_key(config.secret_access_key.clone())
|
||||
.with_virtual_hosted_style_request(!config.path_style)
|
||||
.build()
|
||||
.map_err(|error| {
|
||||
BackupStoreError::new(format!(
|
||||
"S3 backup object store configuration failed: {error}"
|
||||
))
|
||||
})?;
|
||||
|
||||
Ok(Self { store })
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl BackupObjectStore for ObjectStoreS3BackupStore {
|
||||
async fn put_object(&self, key: &str, bytes: Bytes) -> Result<(), BackupStoreError> {
|
||||
self.store
|
||||
.put(&Path::from(key), bytes.into())
|
||||
.await
|
||||
.map(|_| ())
|
||||
.map_err(|error| BackupStoreError::object_store("put", key, error))
|
||||
}
|
||||
|
||||
async fn list_keys(&self, prefix: &str) -> Result<Vec<String>, BackupStoreError> {
|
||||
let prefix_path = list_prefix_path(prefix);
|
||||
let mut keys = self
|
||||
.store
|
||||
.list(prefix_path.as_ref())
|
||||
.map_ok(|meta| meta.location.to_string())
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.map_err(|error| BackupStoreError::object_store("list", prefix, error))?;
|
||||
keys.sort();
|
||||
Ok(keys)
|
||||
}
|
||||
|
||||
async fn delete_object(&self, key: &str) -> Result<(), BackupStoreError> {
|
||||
self.store
|
||||
.delete(&Path::from(key))
|
||||
.await
|
||||
.map_err(|error| BackupStoreError::object_store("delete", key, error))
|
||||
}
|
||||
}
|
||||
|
||||
fn directory_list_prefix(prefix: &str) -> String {
|
||||
let prefix = prefix.trim_end_matches('/');
|
||||
if prefix.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!("{prefix}/")
|
||||
}
|
||||
}
|
||||
|
||||
fn list_prefix_path(prefix: &str) -> Option<Path> {
|
||||
let prefix = prefix.trim_end_matches('/');
|
||||
if prefix.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(Path::from(prefix))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{list_prefix_path, BackupObjectStore, FakeBackupObjectStore};
|
||||
|
||||
#[tokio::test]
|
||||
async fn fake_backup_object_store_puts_lists_and_deletes() {
|
||||
let store = FakeBackupObjectStore::default();
|
||||
store
|
||||
.put_object(
|
||||
"prod/aether-data-backup-20260524-010000.json.zst",
|
||||
bytes::Bytes::from_static(b"one"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.put_object(
|
||||
"prod/aether-data-backup-20260524-020000.json.zst",
|
||||
bytes::Bytes::from_static(b"two"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let keys = store.list_keys("prod/").await.unwrap();
|
||||
assert_eq!(keys.len(), 2);
|
||||
|
||||
store
|
||||
.delete_object("prod/aether-data-backup-20260524-010000.json.zst")
|
||||
.await
|
||||
.unwrap();
|
||||
let keys = store.list_keys("prod/").await.unwrap();
|
||||
assert_eq!(
|
||||
keys,
|
||||
vec!["prod/aether-data-backup-20260524-020000.json.zst"]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fake_backup_object_store_lists_normalized_directory_prefixes() {
|
||||
let store = FakeBackupObjectStore::default();
|
||||
store
|
||||
.put_object(
|
||||
"prod/aether-data-backup-20260524-010000.json.zst",
|
||||
bytes::Bytes::from_static(b"one"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
store
|
||||
.put_object(
|
||||
"prod-backups/aether-data-backup-20260524-010000.json.zst",
|
||||
bytes::Bytes::from_static(b"two"),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let keys = store.list_keys("prod").await.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
keys,
|
||||
vec!["prod/aether-data-backup-20260524-010000.json.zst"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn s3_list_prefix_path_lets_object_store_add_directory_delimiter() {
|
||||
assert_eq!(
|
||||
list_prefix_path("prod/")
|
||||
.as_ref()
|
||||
.map(std::string::ToString::to_string)
|
||||
.as_deref(),
|
||||
Some("prod")
|
||||
);
|
||||
assert!(list_prefix_path("").is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,821 @@
|
||||
use std::fmt;
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_admin::system::admin_system_config_default_value;
|
||||
use aether_data_contracts::repository::background_tasks::{
|
||||
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskStatus, StoredBackgroundTaskRun,
|
||||
UpsertBackgroundTaskRun,
|
||||
};
|
||||
use aether_runtime_state::RuntimeLockLease;
|
||||
use axum::http::StatusCode;
|
||||
use chrono::Utc;
|
||||
use futures_util::FutureExt;
|
||||
use serde::Serialize;
|
||||
use serde_json::{json, Map, Value};
|
||||
use tracing::warn;
|
||||
|
||||
use super::config::S3BackupConfig;
|
||||
use super::executor::{run_backup_with_store, BackupRunResult};
|
||||
use super::scopes::BackupScope;
|
||||
use super::store::ObjectStoreS3BackupStore;
|
||||
use crate::admin_api::AdminAppState;
|
||||
use crate::handlers::shared::decrypt_catalog_secret_with_fallbacks;
|
||||
use crate::task_runtime::{
|
||||
append_event_with_logging, build_task_run_id, now_unix_secs, spawn_fire_and_forget,
|
||||
task_definition, update_run_status, upsert_run_with_logging, TASK_KEY_SYSTEM_S3_BACKUP,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const S3_BACKUP_CONFIG_KEYS: &[&str] = &[
|
||||
"backup_s3_enabled",
|
||||
"backup_s3_scope",
|
||||
"backup_s3_endpoint",
|
||||
"backup_s3_region",
|
||||
"backup_s3_bucket",
|
||||
"backup_s3_prefix",
|
||||
"backup_s3_access_key_id",
|
||||
"backup_s3_secret_access_key",
|
||||
"backup_s3_path_style",
|
||||
"backup_s3_compression",
|
||||
"backup_s3_schedule_unit",
|
||||
"backup_s3_schedule_interval",
|
||||
"backup_s3_schedule_minute",
|
||||
"backup_s3_schedule_hour",
|
||||
"backup_s3_schedule_weekday",
|
||||
"backup_s3_schedule_month_day",
|
||||
"backup_s3_retention_count",
|
||||
];
|
||||
|
||||
const S3_BACKUP_QUEUED_MESSAGE: &str = "S3 备份任务已提交";
|
||||
const S3_BACKUP_TASK_LOCK_KEY: &str = "task_runtime:lock:system.s3.backup";
|
||||
const S3_BACKUP_TASK_LOCK_TTL: Duration = Duration::from_secs(60 * 60 * 6);
|
||||
const S3_BACKUP_TASK_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(60 * 5);
|
||||
const S3_BACKUP_ACTIVE_TASK_STALE_AFTER_SECS: u64 = 60 * 60 * 6;
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub(crate) struct S3BackupTaskStart {
|
||||
pub(crate) id: String,
|
||||
pub(crate) task_key: &'static str,
|
||||
pub(crate) status: &'static str,
|
||||
pub(crate) progress_message: &'static str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct S3BackupTaskError {
|
||||
status: StatusCode,
|
||||
detail: String,
|
||||
}
|
||||
|
||||
impl S3BackupTaskError {
|
||||
fn bad_request(detail: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::BAD_REQUEST,
|
||||
detail: detail.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn internal(detail: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
detail: detail.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn service_unavailable(detail: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::SERVICE_UNAVAILABLE,
|
||||
detail: detail.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn conflict(detail: impl Into<String>) -> Self {
|
||||
Self {
|
||||
status: StatusCode::CONFLICT,
|
||||
detail: detail.into(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn status(&self) -> StatusCode {
|
||||
self.status
|
||||
}
|
||||
|
||||
pub(crate) fn detail(&self) -> &str {
|
||||
&self.detail
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for S3BackupTaskError {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(&self.detail)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for S3BackupTaskError {}
|
||||
|
||||
impl From<GatewayError> for S3BackupTaskError {
|
||||
fn from(error: GatewayError) -> Self {
|
||||
Self::internal(format!("{error:?}"))
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn start_s3_backup_task(
|
||||
app: AppState,
|
||||
trigger: &str,
|
||||
created_by: Option<&str>,
|
||||
) -> Result<S3BackupTaskStart, S3BackupTaskError> {
|
||||
start_s3_backup_task_with_slot(app, trigger, created_by, None).await
|
||||
}
|
||||
|
||||
pub(crate) async fn start_s3_backup_task_for_schedule(
|
||||
app: AppState,
|
||||
scheduled_slot: String,
|
||||
) -> Result<S3BackupTaskStart, S3BackupTaskError> {
|
||||
start_s3_backup_task_with_slot(app, "scheduled", None, Some(scheduled_slot)).await
|
||||
}
|
||||
|
||||
async fn start_s3_backup_task_with_slot(
|
||||
app: AppState,
|
||||
trigger: &str,
|
||||
created_by: Option<&str>,
|
||||
scheduled_slot: Option<String>,
|
||||
) -> Result<S3BackupTaskStart, S3BackupTaskError> {
|
||||
let config = load_s3_backup_config_for_run(&app).await?;
|
||||
ensure_background_task_storage(&app)?;
|
||||
let lock = acquire_s3_backup_task_lock(&app).await?;
|
||||
let active_run_exists = match has_active_s3_backup_task(&app).await {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
release_s3_backup_task_lock(&app, lock).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
if active_run_exists {
|
||||
release_s3_backup_task_lock(&app, lock).await;
|
||||
return Err(S3BackupTaskError::conflict(
|
||||
"已有 S3 备份任务正在执行,请等待当前任务完成后再试",
|
||||
));
|
||||
}
|
||||
|
||||
let run_id = build_task_run_id();
|
||||
let created_at = now_unix_secs();
|
||||
let max_attempts = task_definition(TASK_KEY_SYSTEM_S3_BACKUP)
|
||||
.map(|item| item.retry_policy.max_attempts)
|
||||
.unwrap_or(1);
|
||||
|
||||
let run = UpsertBackgroundTaskRun {
|
||||
id: run_id.clone(),
|
||||
task_key: TASK_KEY_SYSTEM_S3_BACKUP.to_string(),
|
||||
kind: BackgroundTaskKind::Scheduled,
|
||||
trigger: trigger.to_string(),
|
||||
status: BackgroundTaskStatus::Queued,
|
||||
attempt: 1,
|
||||
max_attempts,
|
||||
owner_instance: Some(app.tunnel.local_instance_id().to_string()),
|
||||
progress_percent: 0,
|
||||
progress_message: Some(S3_BACKUP_QUEUED_MESSAGE.to_string()),
|
||||
payload_json: Some(s3_backup_task_payload_json(
|
||||
&config,
|
||||
trigger,
|
||||
scheduled_slot.as_deref(),
|
||||
)),
|
||||
result_json: None,
|
||||
error_message: None,
|
||||
cancel_requested: false,
|
||||
created_by: Some(created_by.unwrap_or("admin").to_string()),
|
||||
created_at_unix_secs: created_at,
|
||||
started_at_unix_secs: None,
|
||||
finished_at_unix_secs: None,
|
||||
updated_at_unix_secs: created_at,
|
||||
};
|
||||
if upsert_run_with_logging(&app, run).await.is_none() {
|
||||
release_s3_backup_task_lock(&app, lock).await;
|
||||
return Err(S3BackupTaskError::service_unavailable(
|
||||
"无法创建 S3 备份后台任务记录,请检查后台任务存储是否可用",
|
||||
));
|
||||
}
|
||||
append_event_with_logging(
|
||||
&app,
|
||||
&run_id,
|
||||
"queued",
|
||||
"S3 backup task queued",
|
||||
Some(json!({ "trigger": trigger })),
|
||||
)
|
||||
.await;
|
||||
|
||||
spawn_s3_backup_worker(app, run_id.clone(), config, lock, scheduled_slot);
|
||||
|
||||
Ok(S3BackupTaskStart {
|
||||
id: run_id,
|
||||
task_key: TASK_KEY_SYSTEM_S3_BACKUP,
|
||||
status: BackgroundTaskStatus::Queued.as_database(),
|
||||
progress_message: S3_BACKUP_QUEUED_MESSAGE,
|
||||
})
|
||||
}
|
||||
|
||||
fn s3_backup_task_payload_json(
|
||||
config: &S3BackupConfig,
|
||||
trigger: &str,
|
||||
scheduled_slot: Option<&str>,
|
||||
) -> Value {
|
||||
let mut payload = json!({
|
||||
"scope": config.scope.as_config_value(),
|
||||
"bucket": config.bucket.clone(),
|
||||
"prefix": config.prefix.clone(),
|
||||
"compression": config.compression.clone(),
|
||||
"trigger": trigger,
|
||||
});
|
||||
if let Some(scheduled_slot) = scheduled_slot {
|
||||
payload["scheduled_slot"] = Value::String(scheduled_slot.to_string());
|
||||
}
|
||||
payload
|
||||
}
|
||||
|
||||
fn spawn_s3_backup_worker(
|
||||
app: AppState,
|
||||
run_id: String,
|
||||
config: S3BackupConfig,
|
||||
lock: RuntimeLockLease,
|
||||
scheduled_slot: Option<String>,
|
||||
) {
|
||||
spawn_fire_and_forget("task-runtime-system-s3-backup", async move {
|
||||
let app_for_worker = app.clone();
|
||||
let run_id_for_worker = run_id.clone();
|
||||
let result = std::panic::AssertUnwindSafe(run_s3_backup_worker_inner(
|
||||
app_for_worker,
|
||||
run_id_for_worker,
|
||||
config,
|
||||
lock.clone(),
|
||||
scheduled_slot,
|
||||
))
|
||||
.catch_unwind()
|
||||
.await;
|
||||
if result.is_err() {
|
||||
warn!(run_id = %run_id, "S3 backup task panicked");
|
||||
let _ = update_run_status(
|
||||
&app,
|
||||
&run_id,
|
||||
BackgroundTaskStatus::Failed,
|
||||
Some(100),
|
||||
Some("S3 备份任务异常退出".to_string()),
|
||||
None,
|
||||
Some("S3 backup task panicked".to_string()),
|
||||
None,
|
||||
Some(now_unix_secs()),
|
||||
)
|
||||
.await;
|
||||
append_event_with_logging(&app, &run_id, "failed", "S3 backup task panicked", None)
|
||||
.await;
|
||||
}
|
||||
|
||||
release_s3_backup_task_lock(&app, lock).await;
|
||||
});
|
||||
}
|
||||
|
||||
async fn run_s3_backup_worker_inner(
|
||||
app: AppState,
|
||||
run_id: String,
|
||||
config: S3BackupConfig,
|
||||
lock: RuntimeLockLease,
|
||||
scheduled_slot: Option<String>,
|
||||
) {
|
||||
let started_at = now_unix_secs();
|
||||
let _ = update_run_status(
|
||||
&app,
|
||||
&run_id,
|
||||
BackgroundTaskStatus::Running,
|
||||
Some(5),
|
||||
Some("S3 备份任务开始执行".to_string()),
|
||||
None,
|
||||
None,
|
||||
Some(started_at),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
append_event_with_logging(&app, &run_id, "running", "S3 backup task started", None).await;
|
||||
|
||||
let heartbeat = spawn_s3_backup_task_heartbeat(app.clone(), run_id.clone(), lock);
|
||||
let result = run_s3_backup_once(&app, &config).await;
|
||||
heartbeat.abort();
|
||||
let _ = heartbeat.await;
|
||||
|
||||
match result {
|
||||
Ok(result) => {
|
||||
if let Some(slot) = scheduled_backup_slot_to_record(scheduled_slot.as_deref(), true) {
|
||||
if let Err(error) = record_scheduled_backup_slot(&app, &slot).await {
|
||||
warn!(error = ?error, run_id = %run_id, "S3 backup slot record failed");
|
||||
let _ = update_run_status(
|
||||
&app,
|
||||
&run_id,
|
||||
BackgroundTaskStatus::Failed,
|
||||
Some(100),
|
||||
Some("S3 备份任务完成,但记录调度时间失败".to_string()),
|
||||
None,
|
||||
Some(format!("S3 backup slot record failed: {error:?}")),
|
||||
None,
|
||||
Some(now_unix_secs()),
|
||||
)
|
||||
.await;
|
||||
append_event_with_logging(
|
||||
&app,
|
||||
&run_id,
|
||||
"failed",
|
||||
"S3 backup slot record failed",
|
||||
Some(json!({ "error": format!("{error:?}") })),
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
let result_json = backup_run_result_json(&result);
|
||||
let _ = update_run_status(
|
||||
&app,
|
||||
&run_id,
|
||||
BackgroundTaskStatus::Succeeded,
|
||||
Some(100),
|
||||
Some("S3 备份任务完成".to_string()),
|
||||
Some(result_json.clone()),
|
||||
None,
|
||||
None,
|
||||
Some(now_unix_secs()),
|
||||
)
|
||||
.await;
|
||||
append_event_with_logging(
|
||||
&app,
|
||||
&run_id,
|
||||
"succeeded",
|
||||
"S3 backup task completed",
|
||||
Some(result_json),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Err(error) => {
|
||||
warn!(error = %error, run_id = %run_id, "S3 backup task failed");
|
||||
let _ = update_run_status(
|
||||
&app,
|
||||
&run_id,
|
||||
BackgroundTaskStatus::Failed,
|
||||
Some(100),
|
||||
Some("S3 备份任务失败".to_string()),
|
||||
None,
|
||||
Some(error.to_string()),
|
||||
None,
|
||||
Some(now_unix_secs()),
|
||||
)
|
||||
.await;
|
||||
append_event_with_logging(
|
||||
&app,
|
||||
&run_id,
|
||||
"failed",
|
||||
"S3 backup task failed",
|
||||
Some(json!({ "error": error.to_string() })),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn spawn_s3_backup_task_heartbeat(
|
||||
app: AppState,
|
||||
run_id: String,
|
||||
lock: RuntimeLockLease,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
spawn_fire_and_forget("task-runtime-system-s3-backup-heartbeat", async move {
|
||||
let mut interval = tokio::time::interval(S3_BACKUP_TASK_HEARTBEAT_INTERVAL);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
interval.tick().await;
|
||||
loop {
|
||||
interval.tick().await;
|
||||
let _ = app
|
||||
.runtime_state
|
||||
.lock_renew(&lock, S3_BACKUP_TASK_LOCK_TTL)
|
||||
.await;
|
||||
let _ = update_run_status(
|
||||
&app,
|
||||
&run_id,
|
||||
BackgroundTaskStatus::Running,
|
||||
Some(50),
|
||||
Some("S3 备份任务执行中".to_string()),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn scheduled_backup_slot_to_record(
|
||||
scheduled_slot: Option<&str>,
|
||||
task_succeeded: bool,
|
||||
) -> Option<String> {
|
||||
if !task_succeeded {
|
||||
return None;
|
||||
}
|
||||
scheduled_slot
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
async fn record_scheduled_backup_slot(
|
||||
app: &AppState,
|
||||
scheduled_slot: &str,
|
||||
) -> Result<(), GatewayError> {
|
||||
app.upsert_system_config_json_value(
|
||||
super::S3_BACKUP_LAST_SLOT_KEY,
|
||||
&Value::String(scheduled_slot.to_string()),
|
||||
None,
|
||||
)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn ensure_background_task_storage(app: &AppState) -> Result<(), S3BackupTaskError> {
|
||||
if app.has_background_task_data_reader() && app.has_background_task_data_writer() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
Err(S3BackupTaskError::service_unavailable(
|
||||
"当前节点未启用后台任务存储,无法提交 S3 备份任务",
|
||||
))
|
||||
}
|
||||
|
||||
async fn acquire_s3_backup_task_lock(
|
||||
app: &AppState,
|
||||
) -> Result<RuntimeLockLease, S3BackupTaskError> {
|
||||
match app
|
||||
.runtime_state
|
||||
.lock_try_acquire(
|
||||
S3_BACKUP_TASK_LOCK_KEY,
|
||||
app.tunnel.local_instance_id(),
|
||||
S3_BACKUP_TASK_LOCK_TTL,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(lock)) => Ok(lock),
|
||||
Ok(None) => Err(S3BackupTaskError::conflict(
|
||||
"已有 S3 备份任务正在执行,请等待当前任务完成后再试",
|
||||
)),
|
||||
Err(error) => Err(S3BackupTaskError::service_unavailable(format!(
|
||||
"无法获取 S3 备份任务锁:{error}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
async fn release_s3_backup_task_lock(app: &AppState, lock: RuntimeLockLease) {
|
||||
let _ = app.runtime_state.lock_release(&lock).await;
|
||||
}
|
||||
|
||||
async fn has_active_s3_backup_task(app: &AppState) -> Result<bool, S3BackupTaskError> {
|
||||
let now = now_unix_secs();
|
||||
for status in [BackgroundTaskStatus::Queued, BackgroundTaskStatus::Running] {
|
||||
let page = app
|
||||
.list_background_task_runs(&BackgroundTaskListQuery {
|
||||
task_key_substring: Some(TASK_KEY_SYSTEM_S3_BACKUP.to_string()),
|
||||
kind: Some(BackgroundTaskKind::Scheduled),
|
||||
status: Some(status),
|
||||
trigger: None,
|
||||
offset: 0,
|
||||
limit: 100,
|
||||
})
|
||||
.await?;
|
||||
if page
|
||||
.items
|
||||
.iter()
|
||||
.any(|run| is_blocking_active_s3_backup_run(run, status, now))
|
||||
{
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
fn is_blocking_active_s3_backup_run(
|
||||
run: &StoredBackgroundTaskRun,
|
||||
status: BackgroundTaskStatus,
|
||||
now_unix_secs: u64,
|
||||
) -> bool {
|
||||
run.task_key == TASK_KEY_SYSTEM_S3_BACKUP
|
||||
&& run.kind == BackgroundTaskKind::Scheduled
|
||||
&& run.status == status
|
||||
&& !run.cancel_requested
|
||||
&& now_unix_secs.saturating_sub(run.updated_at_unix_secs)
|
||||
< S3_BACKUP_ACTIVE_TASK_STALE_AFTER_SECS
|
||||
}
|
||||
|
||||
async fn run_s3_backup_once(
|
||||
app: &AppState,
|
||||
config: &S3BackupConfig,
|
||||
) -> Result<BackupRunResult, S3BackupTaskError> {
|
||||
let admin_state = AdminAppState::new(app);
|
||||
let payload = match config.scope {
|
||||
BackupScope::Config => {
|
||||
admin_state
|
||||
.build_admin_system_config_export_payload()
|
||||
.await?
|
||||
}
|
||||
BackupScope::Users => {
|
||||
admin_state
|
||||
.build_admin_system_users_export_payload()
|
||||
.await?
|
||||
}
|
||||
BackupScope::Data => admin_state.build_admin_system_data_export_payload().await?,
|
||||
};
|
||||
let store = ObjectStoreS3BackupStore::from_config(config)
|
||||
.map_err(|error| S3BackupTaskError::internal(error.to_string()))?;
|
||||
run_backup_with_store(config, &store, payload, Utc::now())
|
||||
.await
|
||||
.map_err(|error| S3BackupTaskError::internal(error.to_string()))
|
||||
}
|
||||
|
||||
async fn load_s3_backup_config_for_run(
|
||||
app: &AppState,
|
||||
) -> Result<S3BackupConfig, S3BackupTaskError> {
|
||||
let mut values = load_s3_backup_config_values(app).await?;
|
||||
values.insert("backup_s3_enabled".to_string(), Value::Bool(true));
|
||||
S3BackupConfig::from_json_map(&values)
|
||||
.map_err(|error| S3BackupTaskError::bad_request(format!("S3 备份配置无效:{error}")))
|
||||
}
|
||||
|
||||
pub(crate) async fn load_s3_backup_config_values(
|
||||
app: &AppState,
|
||||
) -> Result<Map<String, Value>, S3BackupTaskError> {
|
||||
let mut values = Map::new();
|
||||
for key in S3_BACKUP_CONFIG_KEYS {
|
||||
let value = app
|
||||
.read_system_config_json_value(key)
|
||||
.await
|
||||
.map_err(S3BackupTaskError::from)?
|
||||
.or_else(|| admin_system_config_default_value(key));
|
||||
if let Some(value) = value {
|
||||
let value = if *key == "backup_s3_secret_access_key" {
|
||||
decrypt_s3_secret_access_key(app, value)?
|
||||
} else {
|
||||
value
|
||||
};
|
||||
values.insert((*key).to_string(), value);
|
||||
}
|
||||
}
|
||||
Ok(values)
|
||||
}
|
||||
|
||||
fn decrypt_s3_secret_access_key(app: &AppState, value: Value) -> Result<Value, S3BackupTaskError> {
|
||||
let Some(ciphertext) = value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(value);
|
||||
};
|
||||
let Some(plaintext) = decrypt_catalog_secret_with_fallbacks(app.encryption_key(), ciphertext)
|
||||
else {
|
||||
return Err(S3BackupTaskError::bad_request(
|
||||
"S3 备份配置无效:Secret Access Key(访问密钥)无法解密,请重新填写",
|
||||
));
|
||||
};
|
||||
Ok(Value::String(plaintext))
|
||||
}
|
||||
|
||||
fn backup_run_result_json(result: &BackupRunResult) -> Value {
|
||||
json!({
|
||||
"scope": result.scope.as_config_value(),
|
||||
"bucket": result.bucket,
|
||||
"object_key": result.object_key,
|
||||
"bytes": result.bytes,
|
||||
"sha256": result.sha256,
|
||||
"export_version": result.export_version,
|
||||
"exported_at": result.exported_at,
|
||||
"compression": result.compression,
|
||||
"deleted_old_objects": result.deleted_old_objects,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
use aether_data::repository::background_tasks::InMemoryBackgroundTaskRepository;
|
||||
use aether_data_contracts::repository::background_tasks::{
|
||||
BackgroundTaskKind, BackgroundTaskStatus, StoredBackgroundTaskRun,
|
||||
};
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::state::AppState;
|
||||
use crate::task_runtime::{now_unix_secs, TASK_KEY_SYSTEM_S3_BACKUP};
|
||||
|
||||
fn valid_s3_backup_config_values() -> Vec<(String, serde_json::Value)> {
|
||||
vec![
|
||||
(
|
||||
"backup_s3_endpoint".to_string(),
|
||||
serde_json::json!("https://s3.example.com"),
|
||||
),
|
||||
(
|
||||
"backup_s3_bucket".to_string(),
|
||||
serde_json::json!("aether-backups"),
|
||||
),
|
||||
("backup_s3_prefix".to_string(), serde_json::json!("prod/")),
|
||||
(
|
||||
"backup_s3_access_key_id".to_string(),
|
||||
serde_json::json!("access-key-id"),
|
||||
),
|
||||
(
|
||||
"backup_s3_secret_access_key".to_string(),
|
||||
serde_json::json!(encrypt_python_fernet_plaintext(
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
"secret"
|
||||
)
|
||||
.expect("test secret should encrypt")),
|
||||
),
|
||||
]
|
||||
}
|
||||
|
||||
fn stored_s3_backup_run(status: BackgroundTaskStatus) -> StoredBackgroundTaskRun {
|
||||
let now = now_unix_secs();
|
||||
StoredBackgroundTaskRun {
|
||||
id: "existing-s3-backup-run".to_string(),
|
||||
task_key: TASK_KEY_SYSTEM_S3_BACKUP.to_string(),
|
||||
kind: BackgroundTaskKind::Scheduled,
|
||||
trigger: "manual".to_string(),
|
||||
status,
|
||||
attempt: 1,
|
||||
max_attempts: 1,
|
||||
owner_instance: Some("test-instance".to_string()),
|
||||
progress_percent: 5,
|
||||
progress_message: Some("running".to_string()),
|
||||
payload_json: None,
|
||||
result_json: None,
|
||||
error_message: None,
|
||||
cancel_requested: false,
|
||||
created_by: Some("admin".to_string()),
|
||||
created_at_unix_secs: now,
|
||||
started_at_unix_secs: Some(now),
|
||||
finished_at_unix_secs: None,
|
||||
updated_at_unix_secs: now,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn start_s3_backup_task_rejects_missing_bucket_for_manual_run() {
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::disabled().with_system_config_values_for_tests(vec![(
|
||||
"backup_s3_endpoint".to_string(),
|
||||
serde_json::json!("https://s3.example.com"),
|
||||
)]),
|
||||
);
|
||||
|
||||
let err = super::start_s3_backup_task(app, "manual", Some("admin-user-123"))
|
||||
.await
|
||||
.expect_err("missing bucket should reject the backup run");
|
||||
|
||||
assert_eq!(err.status(), axum::http::StatusCode::BAD_REQUEST);
|
||||
assert!(err.to_string().contains("Bucket"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn start_s3_backup_task_requires_background_task_storage() {
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::disabled()
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||
.with_system_config_values_for_tests(valid_s3_backup_config_values()),
|
||||
);
|
||||
|
||||
let err = super::start_s3_backup_task(app, "manual", Some("admin-user-123"))
|
||||
.await
|
||||
.expect_err("manual backup should require observable background task storage");
|
||||
|
||||
assert_eq!(err.status(), axum::http::StatusCode::SERVICE_UNAVAILABLE);
|
||||
assert!(err.to_string().contains("后台任务存储"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn start_s3_backup_task_rejects_when_same_task_is_active() {
|
||||
let repository = Arc::new(InMemoryBackgroundTaskRepository::seed_runs([
|
||||
stored_s3_backup_run(BackgroundTaskStatus::Running),
|
||||
]));
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::disabled()
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||
.with_system_config_values_for_tests(valid_s3_backup_config_values())
|
||||
.with_background_task_repository_for_tests(repository),
|
||||
);
|
||||
|
||||
let err = super::start_s3_backup_task(app, "manual", Some("admin-user-123"))
|
||||
.await
|
||||
.expect_err("manual backup should reject duplicate active runs");
|
||||
|
||||
assert_eq!(err.status(), axum::http::StatusCode::CONFLICT);
|
||||
assert!(err.to_string().contains("已有 S3 备份任务正在执行"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn active_s3_backup_detection_blocks_queued_runs() {
|
||||
let repository = Arc::new(InMemoryBackgroundTaskRepository::seed_runs([
|
||||
stored_s3_backup_run(BackgroundTaskStatus::Queued),
|
||||
]));
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::disabled().with_background_task_repository_for_tests(repository),
|
||||
);
|
||||
|
||||
assert!(super::has_active_s3_backup_task(&app)
|
||||
.await
|
||||
.expect("active task lookup should succeed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn active_s3_backup_detection_ignores_cancelled_and_stale_runs() {
|
||||
let now = now_unix_secs();
|
||||
let mut stale = stored_s3_backup_run(BackgroundTaskStatus::Running);
|
||||
stale.id = "stale-s3-backup-run".to_string();
|
||||
stale.updated_at_unix_secs = now
|
||||
.saturating_sub(super::S3_BACKUP_ACTIVE_TASK_STALE_AFTER_SECS)
|
||||
.saturating_sub(1);
|
||||
let mut cancelled = stored_s3_backup_run(BackgroundTaskStatus::Queued);
|
||||
cancelled.id = "cancelled-s3-backup-run".to_string();
|
||||
cancelled.cancel_requested = true;
|
||||
let repository = Arc::new(InMemoryBackgroundTaskRepository::seed_runs([
|
||||
stale, cancelled,
|
||||
]));
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::disabled().with_background_task_repository_for_tests(repository),
|
||||
);
|
||||
|
||||
assert!(!super::has_active_s3_backup_task(&app)
|
||||
.await
|
||||
.expect("active task lookup should succeed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn queued_s3_backup_task_payload_does_not_include_secret() {
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::disabled()
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||
.with_system_config_values_for_tests(valid_s3_backup_config_values()),
|
||||
);
|
||||
let values = super::load_s3_backup_config_values(&app)
|
||||
.await
|
||||
.expect("config should load");
|
||||
let config = super::S3BackupConfig::from_json_map(&values)
|
||||
.expect("config should parse for payload test");
|
||||
|
||||
let payload = super::s3_backup_task_payload_json(&config, "manual", None);
|
||||
|
||||
assert!(payload["bucket"].is_string());
|
||||
assert_eq!(payload["trigger"], serde_json::json!("manual"));
|
||||
assert!(!payload.to_string().contains("secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scheduled_backup_slot_records_only_successful_scheduled_runs() {
|
||||
assert_eq!(
|
||||
super::scheduled_backup_slot_to_record(Some("days:2026-05-24T19:00:00Z"), true),
|
||||
Some("days:2026-05-24T19:00:00Z".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
super::scheduled_backup_slot_to_record(Some("days:2026-05-24T19:00:00Z"), false),
|
||||
None
|
||||
);
|
||||
assert_eq!(super::scheduled_backup_slot_to_record(None, true), None);
|
||||
assert_eq!(
|
||||
super::scheduled_backup_slot_to_record(Some(" "), true),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn record_scheduled_backup_slot_updates_system_config() {
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::disabled().with_system_config_values_for_tests(Vec::<(
|
||||
String,
|
||||
serde_json::Value,
|
||||
)>::new(
|
||||
)),
|
||||
);
|
||||
|
||||
super::record_scheduled_backup_slot(&app, "days:2026-05-24T19:00:00Z")
|
||||
.await
|
||||
.expect("slot record should write system config");
|
||||
|
||||
assert_eq!(
|
||||
app.read_system_config_json_value(super::super::S3_BACKUP_LAST_SLOT_KEY)
|
||||
.await
|
||||
.expect("slot config should be readable"),
|
||||
Some(serde_json::json!("days:2026-05-24T19:00:00Z"))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use tokio::task::JoinHandle;
|
||||
use tracing::warn;
|
||||
|
||||
use super::config::S3BackupConfig;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(crate) const S3_BACKUP_WORKER_TASK_KEY: &str =
|
||||
crate::task_runtime::TASK_KEY_SYSTEM_S3_BACKUP_WORKER;
|
||||
const S3_BACKUP_WORKER_INTERVAL: Duration = Duration::from_secs(60);
|
||||
|
||||
pub(crate) fn should_start_scheduled_backup(last_slot: Option<&str>, current_slot: &str) -> bool {
|
||||
last_slot != Some(current_slot)
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_s3_backup_worker(app: AppState) -> Option<JoinHandle<()>> {
|
||||
if !app.data.has_system_config_store() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(tokio::spawn(async move {
|
||||
let mut interval = tokio::time::interval(S3_BACKUP_WORKER_INTERVAL);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
interval.tick().await;
|
||||
loop {
|
||||
interval.tick().await;
|
||||
if let Err(error) = run_s3_backup_schedule_tick(&app, Utc::now()).await {
|
||||
warn!(error = ?error, "S3 backup schedule tick failed");
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
async fn run_s3_backup_schedule_tick(
|
||||
app: &AppState,
|
||||
now: DateTime<Utc>,
|
||||
) -> Result<(), GatewayError> {
|
||||
let values = match super::task::load_s3_backup_config_values(app).await {
|
||||
Ok(values) => values,
|
||||
Err(error) => {
|
||||
warn!(error = %error, "S3 backup schedule config load failed");
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
let config = match S3BackupConfig::from_json_map(&values) {
|
||||
Ok(config) => config,
|
||||
Err(error) => {
|
||||
warn!(error = %error, "S3 backup schedule config is invalid");
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
if !config.enabled {
|
||||
return Ok(());
|
||||
}
|
||||
let Some(slot) = config.schedule.due_slot(now) else {
|
||||
return Ok(());
|
||||
};
|
||||
let last_slot = read_last_backup_slot(app).await?;
|
||||
if !should_start_scheduled_backup(last_slot.as_deref(), &slot) {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
match super::task::start_s3_backup_task_for_schedule(app.clone(), slot).await {
|
||||
Ok(_) => {}
|
||||
Err(error) => {
|
||||
warn!(error = %error, "S3 backup scheduled task submission failed");
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn read_last_backup_slot(app: &AppState) -> Result<Option<String>, GatewayError> {
|
||||
Ok(app
|
||||
.read_system_config_json_value(super::S3_BACKUP_LAST_SLOT_KEY)
|
||||
.await?
|
||||
.and_then(|value| value.as_str().map(str::trim).map(str::to_string))
|
||||
.filter(|value| !value.is_empty()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::backup::schedule::{BackupSchedule, BackupScheduleUnit};
|
||||
use crate::task_runtime::{task_definition, TASK_KEY_SYSTEM_S3_BACKUP};
|
||||
|
||||
#[test]
|
||||
fn backup_worker_skips_already_recorded_slot() {
|
||||
let schedule = BackupSchedule {
|
||||
unit: BackupScheduleUnit::Days,
|
||||
interval: 1,
|
||||
minute: 0,
|
||||
hour: 3,
|
||||
weekday: 1,
|
||||
month_day: 1,
|
||||
};
|
||||
let now = chrono::DateTime::parse_from_rfc3339("2026-05-24T03:00:30+08:00")
|
||||
.unwrap()
|
||||
.with_timezone(&chrono::Utc);
|
||||
let slot = schedule.due_slot(now).expect("slot should be due");
|
||||
|
||||
assert!(super::should_start_scheduled_backup(
|
||||
Some("days:2026-05-22T19:00:00Z"),
|
||||
&slot
|
||||
));
|
||||
assert!(!super::should_start_scheduled_backup(Some(&slot), &slot));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backup_worker_has_distinct_supervisor_task_key() {
|
||||
assert_ne!(super::S3_BACKUP_WORKER_TASK_KEY, TASK_KEY_SYSTEM_S3_BACKUP);
|
||||
assert!(task_definition(super::S3_BACKUP_WORKER_TASK_KEY).is_some());
|
||||
}
|
||||
}
|
||||
@@ -56,7 +56,7 @@ impl ClientSessionScope {
|
||||
pub(crate) fn scheduler_affinity(&self) -> Option<ClientSessionAffinity> {
|
||||
let client_family = self.client_family.trim();
|
||||
let client_family = if client_family.is_empty() {
|
||||
"generic".to_string()
|
||||
"unknown".to_string()
|
||||
} else {
|
||||
client_family.to_ascii_lowercase()
|
||||
};
|
||||
@@ -84,6 +84,15 @@ struct GenericSessionScopeAdapter;
|
||||
struct CodexSessionScopeAdapter;
|
||||
struct ClaudeCodeSessionScopeAdapter;
|
||||
struct OpenCodeSessionScopeAdapter;
|
||||
struct QwenCodeSessionScopeAdapter;
|
||||
struct RooCodeSessionScopeAdapter;
|
||||
struct KiloCodeSessionScopeAdapter;
|
||||
struct CherryStudioSessionScopeAdapter;
|
||||
struct OpenUiSessionScopeAdapter;
|
||||
struct OpenAiJsSdkSessionScopeAdapter;
|
||||
struct OpenAiPythonSdkSessionScopeAdapter;
|
||||
struct AnthropicJsSdkSessionScopeAdapter;
|
||||
struct AnthropicPythonSdkSessionScopeAdapter;
|
||||
|
||||
pub(crate) fn client_session_affinity_from_request(
|
||||
headers: &http::HeaderMap,
|
||||
@@ -171,14 +180,85 @@ fn detect_client_family(request: &ClientSessionRequest<'_>) -> String {
|
||||
return adapter.family().to_string();
|
||||
}
|
||||
}
|
||||
if let Some(client_family) = detect_fingerprint_client_family(request) {
|
||||
return client_family.to_string();
|
||||
}
|
||||
GenericSessionScopeAdapter.family().to_string()
|
||||
}
|
||||
|
||||
fn specific_client_session_scope_adapters() -> [&'static dyn ClientSessionScopeAdapter; 3] {
|
||||
fn detect_fingerprint_client_family(request: &ClientSessionRequest<'_>) -> Option<&'static str> {
|
||||
if header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"geminicli",
|
||||
) || header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"gemini-cli",
|
||||
) {
|
||||
return Some("gemini_cli");
|
||||
}
|
||||
if header_contains(request.headers, http::header::USER_AGENT.as_str(), "cursor") {
|
||||
return Some("cursor");
|
||||
}
|
||||
if header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"windsurf",
|
||||
) {
|
||||
return Some("windsurf");
|
||||
}
|
||||
if header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"continue",
|
||||
) {
|
||||
return Some("continue");
|
||||
}
|
||||
if header_contains(request.headers, http::header::USER_AGENT.as_str(), "cline") {
|
||||
return Some("cline");
|
||||
}
|
||||
if header_contains(request.headers, http::header::USER_AGENT.as_str(), "aider") {
|
||||
return Some("aider");
|
||||
}
|
||||
if header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"langchain",
|
||||
) {
|
||||
return Some("langchain");
|
||||
}
|
||||
if header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"llamaindex",
|
||||
) || header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"llama-index",
|
||||
) {
|
||||
return Some("llamaindex");
|
||||
}
|
||||
if has_header_with_prefix(request.headers, "x-stainless-") {
|
||||
return Some("sdk");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn specific_client_session_scope_adapters() -> [&'static dyn ClientSessionScopeAdapter; 12] {
|
||||
[
|
||||
&CodexSessionScopeAdapter,
|
||||
&ClaudeCodeSessionScopeAdapter,
|
||||
&OpenCodeSessionScopeAdapter,
|
||||
&QwenCodeSessionScopeAdapter,
|
||||
&RooCodeSessionScopeAdapter,
|
||||
&KiloCodeSessionScopeAdapter,
|
||||
&CherryStudioSessionScopeAdapter,
|
||||
&OpenUiSessionScopeAdapter,
|
||||
&AnthropicJsSdkSessionScopeAdapter,
|
||||
&AnthropicPythonSdkSessionScopeAdapter,
|
||||
&OpenAiJsSdkSessionScopeAdapter,
|
||||
&OpenAiPythonSdkSessionScopeAdapter,
|
||||
]
|
||||
}
|
||||
|
||||
@@ -199,6 +279,7 @@ fn extract_scope_from_other_specific_adapters(
|
||||
specific_client_session_scope_adapters()
|
||||
.into_iter()
|
||||
.filter(|adapter| adapter.family() != client_family)
|
||||
.filter(|adapter| adapter.detect(request))
|
||||
.find_map(|adapter| adapter.extract_scope(request))
|
||||
}
|
||||
|
||||
@@ -218,7 +299,7 @@ fn extract_generic_scope_for_client_family(
|
||||
|
||||
impl ClientSessionScopeAdapter for GenericSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"generic"
|
||||
"unknown"
|
||||
}
|
||||
|
||||
fn detect(&self, _request: &ClientSessionRequest<'_>) -> bool {
|
||||
@@ -226,6 +307,18 @@ impl ClientSessionScopeAdapter for GenericSessionScopeAdapter {
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
if let Some(root_session) = header_value_str(request.headers, "session_id")
|
||||
.or_else(|| header_value_str(request.headers, "conversation_id"))
|
||||
{
|
||||
return Some(ClientSessionScope::new(
|
||||
self.family(),
|
||||
root_session,
|
||||
None,
|
||||
None,
|
||||
ClientSessionSignalSource::Header,
|
||||
));
|
||||
}
|
||||
|
||||
let body = request.body_json?;
|
||||
let root_session = value_at_paths(
|
||||
body,
|
||||
@@ -378,6 +471,246 @@ impl ClientSessionScopeAdapter for OpenCodeSessionScopeAdapter {
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for QwenCodeSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"qwen_code"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"qwencode",
|
||||
) || header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"qwen-code",
|
||||
) || header_contains(request.headers, "x-dashscope-useragent", "qwencode")
|
||||
|| header_contains(request.headers, "x-dashscope-useragent", "qwen-code")
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for RooCodeSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"roo_code"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"roo-code",
|
||||
) || header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"roocode",
|
||||
) || header_contains(request.headers, "originator", "roo-code")
|
||||
|| header_contains(request.headers, "originator", "roocode")
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for KiloCodeSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"kilocode"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"kilo-code",
|
||||
) || header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"kilocode",
|
||||
) || has_header_with_prefix(request.headers, "x-kilocode-")
|
||||
|| header_value_str(request.headers, "x-kilo-directory").is_some()
|
||||
|| header_value_str(request.headers, "x-kilo-workspace").is_some()
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for CherryStudioSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"cherrystudio"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"cherrystudio",
|
||||
) || header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"cherry-studio",
|
||||
) || header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"cherry studio",
|
||||
)
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for OpenUiSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"openui"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(request.headers, http::header::USER_AGENT.as_str(), "openui")
|
||||
|| header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"openui-agent-manager",
|
||||
)
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for OpenAiJsSdkSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"openai_js_sdk"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"openai/js",
|
||||
) || (header_contains(request.headers, http::header::USER_AGENT.as_str(), "/js ")
|
||||
&& header_contains(request.headers, "x-stainless-lang", "js"))
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for OpenAiPythonSdkSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"openai_python_sdk"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"openai/python",
|
||||
) || (header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"/python ",
|
||||
) && header_contains(request.headers, "x-stainless-lang", "python")
|
||||
&& header_value_str(request.headers, "anthropic-version").is_none())
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for AnthropicJsSdkSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"anthropic_js_sdk"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"anthropic/js",
|
||||
) || (header_contains(request.headers, http::header::USER_AGENT.as_str(), "/js ")
|
||||
&& header_contains(request.headers, "x-stainless-lang", "js")
|
||||
&& header_value_str(request.headers, "anthropic-version").is_some())
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
impl ClientSessionScopeAdapter for AnthropicPythonSdkSessionScopeAdapter {
|
||||
fn family(&self) -> &'static str {
|
||||
"anthropic_python_sdk"
|
||||
}
|
||||
|
||||
fn detect(&self, request: &ClientSessionRequest<'_>) -> bool {
|
||||
header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"anthropic/python",
|
||||
) || (header_contains(
|
||||
request.headers,
|
||||
http::header::USER_AGENT.as_str(),
|
||||
"/python ",
|
||||
) && header_contains(request.headers, "x-stainless-lang", "python")
|
||||
&& header_value_str(request.headers, "anthropic-version").is_some())
|
||||
}
|
||||
|
||||
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
|
||||
scoped_from_standard_session_headers(self.family(), request)
|
||||
.or_else(|| scoped_from_generic_body(self.family(), request))
|
||||
}
|
||||
}
|
||||
|
||||
fn scoped_from_standard_session_headers(
|
||||
client_family: &str,
|
||||
request: &ClientSessionRequest<'_>,
|
||||
) -> Option<ClientSessionScope> {
|
||||
header_value_str(request.headers, "session_id")
|
||||
.or_else(|| header_value_str(request.headers, "conversation_id"))
|
||||
.map(|root_session| {
|
||||
ClientSessionScope::new(
|
||||
client_family,
|
||||
root_session,
|
||||
None,
|
||||
None,
|
||||
ClientSessionSignalSource::Header,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn scoped_from_generic_body(
|
||||
client_family: &str,
|
||||
request: &ClientSessionRequest<'_>,
|
||||
) -> Option<ClientSessionScope> {
|
||||
let body_session = GenericSessionScopeAdapter.extract_scope(request)?;
|
||||
Some(ClientSessionScope::new(
|
||||
client_family,
|
||||
body_session.session_id,
|
||||
body_session.agent_id,
|
||||
body_session.account_hint,
|
||||
body_session.source,
|
||||
))
|
||||
}
|
||||
|
||||
fn claude_code_session_id_from_body(body: &Value) -> Option<&str> {
|
||||
value_at_path(body, &["metadata", "user_id"]).and_then(|user_id| {
|
||||
user_id
|
||||
@@ -443,6 +776,12 @@ fn header_contains(headers: &http::HeaderMap, key: &str, needle: &str) -> bool {
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn has_header_with_prefix(headers: &http::HeaderMap, prefix: &str) -> bool {
|
||||
headers
|
||||
.keys()
|
||||
.any(|key| key.as_str().to_ascii_lowercase().starts_with(prefix))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
@@ -455,7 +794,7 @@ mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn generic_adapter_extracts_body_session_and_agent() {
|
||||
fn unknown_adapter_extracts_body_session_and_agent() {
|
||||
let body = json!({
|
||||
"metadata": {
|
||||
"session_id": "session-1",
|
||||
@@ -466,7 +805,7 @@ mod tests {
|
||||
let affinity = client_session_affinity_from_request(&HeaderMap::new(), Some(&body))
|
||||
.expect("affinity should build");
|
||||
|
||||
assert_eq!(affinity.client_family.as_deref(), Some("generic"));
|
||||
assert_eq!(affinity.client_family.as_deref(), Some("unknown"));
|
||||
assert_eq!(
|
||||
affinity.session_key.as_deref(),
|
||||
Some("session=session-1;agent=planner")
|
||||
@@ -669,6 +1008,84 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen_code_detection_keeps_body_session() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::USER_AGENT,
|
||||
HeaderValue::from_static("QwenCode/0.1.0 (linux; x64)"),
|
||||
);
|
||||
let body = json!({"conversation_id": "qwen-session"});
|
||||
|
||||
let scope = client_session_scope_from_request(&headers, Some(&body))
|
||||
.expect("session scope should build");
|
||||
|
||||
assert_eq!(scope.client_family, "qwen_code");
|
||||
assert_eq!(scope.session_id, "qwen-session");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn roo_code_detection_uses_originator_and_session_header() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("originator", HeaderValue::from_static("roo-code"));
|
||||
headers.insert("session_id", HeaderValue::from_static("roo-session"));
|
||||
|
||||
let scope =
|
||||
client_session_scope_from_request(&headers, None).expect("session scope should build");
|
||||
|
||||
assert_eq!(scope.client_family, "roo_code");
|
||||
assert_eq!(scope.session_id, "roo-session");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_user_agent_with_session_header_stays_unknown() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::USER_AGENT,
|
||||
HeaderValue::from_static("CustomClient/1.0"),
|
||||
);
|
||||
headers.insert("session_id", HeaderValue::from_static("custom-session"));
|
||||
|
||||
let scope =
|
||||
client_session_scope_from_request(&headers, None).expect("session scope should build");
|
||||
|
||||
assert_eq!(scope.client_family, "unknown");
|
||||
assert_eq!(scope.session_id, "custom-session");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vscode_copilot_user_agent_is_not_cherrystudio_by_itself() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::USER_AGENT,
|
||||
HeaderValue::from_static("Visual Studio Code (desktop) GithubCopilot/1.155.0"),
|
||||
);
|
||||
headers.insert("session_id", HeaderValue::from_static("vscode-session"));
|
||||
|
||||
let scope =
|
||||
client_session_scope_from_request(&headers, None).expect("session scope should build");
|
||||
|
||||
assert_eq!(scope.client_family, "unknown");
|
||||
assert_eq!(scope.session_id, "vscode-session");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sdk_detection_uses_user_agent_before_stainless_fallback() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::USER_AGENT,
|
||||
HeaderValue::from_static("OpenAI/JS 6.0.0"),
|
||||
);
|
||||
headers.insert("x-stainless-lang", HeaderValue::from_static("js"));
|
||||
let body = json!({"metadata": {"session_id": "sdk-session"}});
|
||||
|
||||
let scope = client_session_scope_from_request(&headers, Some(&body))
|
||||
.expect("session scope should build");
|
||||
|
||||
assert_eq!(scope.client_family, "openai_js_sdk");
|
||||
assert_eq!(scope.session_id, "sdk-session");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn specific_adapter_wins_over_generic_body_session() {
|
||||
let mut headers = HeaderMap::new();
|
||||
|
||||
@@ -26,6 +26,8 @@ pub(crate) const EXECUTION_PATH_EXECUTION_RUNTIME_STREAM: &str = "execution_runt
|
||||
pub(crate) const EXECUTION_PATH_CONTROL_EXECUTE_SYNC: &str = "control_execute_sync";
|
||||
pub(crate) const EXECUTION_PATH_CONTROL_EXECUTE_STREAM: &str = "control_execute_stream";
|
||||
pub(crate) const EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS: &str = "local_execution_runtime_miss";
|
||||
pub(crate) const EXECUTION_PATH_LOCAL_EXECUTION_PLANNING_TIMEOUT: &str =
|
||||
"local_execution_planning_timeout";
|
||||
pub(crate) const EXECUTION_PATH_LOCAL_API_KEY_CONCURRENCY_LIMITED: &str =
|
||||
"local_api_key_concurrency_limited";
|
||||
pub(crate) const API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS: u64 = 150;
|
||||
@@ -63,6 +65,7 @@ pub(crate) const TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER: &str =
|
||||
"x-aether-admin-management-token-id";
|
||||
pub(crate) const TRUSTED_RATE_LIMIT_PREFLIGHT_HEADER: &str = "x-aether-rate-limit-preflight";
|
||||
pub(crate) const DEFAULT_USER_GROUP_CONFIG_KEY: &str = "default_user_group_id";
|
||||
pub(crate) const ANTIGRAVITY_BEARER_BRIDGE_CONFIG_KEY: &str = "module.antigravity.bearer_bridge";
|
||||
pub(crate) const BUILTIN_DEFAULT_USER_GROUP_ID: &str = "00000000-0000-0000-0000-000000000001";
|
||||
|
||||
pub(crate) const FRONTDOOR_REPLACEABLE_ROUTE_GROUPS: &[&str] = &["frontdoor_compat_router"];
|
||||
@@ -134,6 +137,14 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
|
||||
"/upload/v1beta/files",
|
||||
"/v1beta/files",
|
||||
"/v1beta/files/{path...}",
|
||||
"/v1internal:loadCodeAssist",
|
||||
"/v1internal:fetchAvailableModels",
|
||||
"/v1internal:fetchUserInfo",
|
||||
"/v1internal:fetchAdminControls",
|
||||
"/v1internal:setUserSettings",
|
||||
"/v1internal:listExperiments",
|
||||
"/v1internal:recordCodeAssistMetrics",
|
||||
"/v1internal:streamGenerateContent",
|
||||
"/",
|
||||
"/{*path}",
|
||||
];
|
||||
|
||||
@@ -186,6 +186,9 @@ fn select_primary_credential(
|
||||
if signature.starts_with("gemini:") {
|
||||
return select_gemini_credential(bundle);
|
||||
}
|
||||
if signature.starts_with("antigravity:") {
|
||||
return select_antigravity_credential(bundle);
|
||||
}
|
||||
if signature.starts_with("claude:") {
|
||||
return select_claude_messages_credential(bundle);
|
||||
}
|
||||
@@ -196,6 +199,20 @@ fn select_primary_credential(
|
||||
select_generic_credential(bundle)
|
||||
}
|
||||
|
||||
fn select_antigravity_credential(
|
||||
bundle: &GatewayCredentialBundle,
|
||||
) -> Option<GatewayPrimaryCredential> {
|
||||
first_provider_api_key(
|
||||
bundle,
|
||||
&[
|
||||
GatewayCredentialCarrier::XApiKey,
|
||||
GatewayCredentialCarrier::ApiKey,
|
||||
],
|
||||
)
|
||||
.or_else(|| first_bearer_token(bundle))
|
||||
.or_else(|| select_cookie_credential(bundle))
|
||||
}
|
||||
|
||||
fn select_openai_credential(bundle: &GatewayCredentialBundle) -> Option<GatewayPrimaryCredential> {
|
||||
first_provider_api_key(
|
||||
bundle,
|
||||
@@ -459,6 +476,29 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prefers_antigravity_aether_api_key_over_google_bearer() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::AUTHORIZATION,
|
||||
"Bearer google-oauth-access-token".parse().unwrap(),
|
||||
);
|
||||
headers.insert("x-api-key", "sk-aether-antigravity".parse().unwrap());
|
||||
|
||||
let extracted = extract_request_credentials(
|
||||
&headers,
|
||||
&uri("/v1internal:streamGenerateContent?alt=sse"),
|
||||
"antigravity:v1internal",
|
||||
);
|
||||
assert_eq!(
|
||||
extracted.primary,
|
||||
Some(GatewayPrimaryCredential::ProviderApiKey {
|
||||
raw: "sk-aether-antigravity".to_string(),
|
||||
carrier: GatewayCredentialCarrier::XApiKey,
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prefers_gemini_query_key_over_header_key() {
|
||||
let mut headers = http::HeaderMap::new();
|
||||
|
||||
@@ -16,16 +16,48 @@ use crate::{AppState, GatewayError};
|
||||
use super::super::GatewayControlDecision;
|
||||
use super::credentials::{
|
||||
build_auth_context_cache_key, current_unix_secs, extract_request_credentials,
|
||||
extract_trusted_admin_headers,
|
||||
extract_trusted_admin_headers, hash_api_key,
|
||||
};
|
||||
use super::gate::GatewayLocalAuthRejection;
|
||||
use super::principal::derive_principal_candidate;
|
||||
use super::types::{GatewayPrincipalCandidate, GatewayTrustedAuthHeaders};
|
||||
use super::types::{
|
||||
GatewayCredentialCarrier, GatewayPrincipalCandidate, GatewayTrustedAuthHeaders,
|
||||
};
|
||||
use crate::headers::header_value_str;
|
||||
|
||||
const AUTH_CONTEXT_CACHE_TTL: Duration = Duration::from_secs(60);
|
||||
const AUTH_CONTEXT_CACHE_MAX_ENTRIES: usize = 256;
|
||||
|
||||
#[derive(Debug, Clone, Deserialize)]
|
||||
struct AntigravityBearerBridgeConfig {
|
||||
#[serde(default)]
|
||||
enabled: bool,
|
||||
#[serde(default)]
|
||||
auth_user_id: String,
|
||||
#[serde(default)]
|
||||
auth_api_key_id: String,
|
||||
#[serde(default)]
|
||||
bearer_sha256_allowlist: Vec<String>,
|
||||
#[serde(default)]
|
||||
allow_unverified_google_bearer: bool,
|
||||
}
|
||||
|
||||
impl AntigravityBearerBridgeConfig {
|
||||
fn bearer_validation_mode(&self, raw_bearer: &str) -> Option<&'static str> {
|
||||
if !self.bearer_sha256_allowlist.is_empty() {
|
||||
let bearer_hash = hash_api_key(raw_bearer);
|
||||
return self
|
||||
.bearer_sha256_allowlist
|
||||
.iter()
|
||||
.any(|allowed| allowed.trim().eq_ignore_ascii_case(&bearer_hash))
|
||||
.then_some("sha256_allowlist");
|
||||
}
|
||||
|
||||
self.allow_unverified_google_bearer
|
||||
.then_some("explicit_unverified")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, Serialize)]
|
||||
pub(crate) struct GatewayControlAuthContext {
|
||||
pub(crate) user_id: String,
|
||||
@@ -607,14 +639,117 @@ pub(super) async fn resolve_data_backed_auth_context(
|
||||
.await,
|
||||
))
|
||||
}
|
||||
Some(
|
||||
GatewayPrincipalCandidate::DeferredBearerToken { .. }
|
||||
| GatewayPrincipalCandidate::DeferredCookieHeader { .. },
|
||||
) => Ok(None),
|
||||
Some(GatewayPrincipalCandidate::DeferredBearerToken { raw, carrier }) => {
|
||||
if let Some(auth_context) = resolve_antigravity_bearer_bridge_auth_context(
|
||||
state,
|
||||
signature,
|
||||
raw.as_str(),
|
||||
carrier,
|
||||
now_unix_secs,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(auth_context));
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
Some(GatewayPrincipalCandidate::DeferredCookieHeader { .. }) => Ok(None),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
async fn resolve_antigravity_bearer_bridge_auth_context(
|
||||
state: &AppState,
|
||||
auth_endpoint_signature: &str,
|
||||
raw_bearer: &str,
|
||||
carrier: GatewayCredentialCarrier,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<Option<GatewayControlAuthContext>, GatewayError> {
|
||||
if carrier != GatewayCredentialCarrier::AuthorizationBearer
|
||||
|| !auth_endpoint_signature
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity:v1internal")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(config_value) = state
|
||||
.read_system_config_json_value(crate::constants::ANTIGRAVITY_BEARER_BRIDGE_CONFIG_KEY)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
if config_value.is_null() {
|
||||
return Ok(None);
|
||||
}
|
||||
let config: AntigravityBearerBridgeConfig =
|
||||
serde_json::from_value(config_value).map_err(|err| {
|
||||
GatewayError::Internal(format!(
|
||||
"{} invalid: {err}",
|
||||
crate::constants::ANTIGRAVITY_BEARER_BRIDGE_CONFIG_KEY
|
||||
))
|
||||
})?;
|
||||
if !config.enabled {
|
||||
return Ok(None);
|
||||
}
|
||||
let Some(validation_mode) = config.bearer_validation_mode(raw_bearer) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let user_id = config.auth_user_id.trim();
|
||||
let api_key_id = config.auth_api_key_id.trim();
|
||||
if user_id.is_empty() || api_key_id.is_empty() {
|
||||
return Err(GatewayError::Internal(format!(
|
||||
"{} requires auth_user_id and auth_api_key_id",
|
||||
crate::constants::ANTIGRAVITY_BEARER_BRIDGE_CONFIG_KEY
|
||||
)));
|
||||
}
|
||||
|
||||
let snapshot = state
|
||||
.data
|
||||
.read_auth_api_key_snapshot(user_id, api_key_id, now_unix_secs)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let Some(snapshot) = snapshot else {
|
||||
return Ok(Some(GatewayControlAuthContext {
|
||||
user_id: user_id.to_string(),
|
||||
api_key_id: api_key_id.to_string(),
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
balance_remaining: None,
|
||||
access_allowed: false,
|
||||
user_rate_limit: None,
|
||||
api_key_rate_limit: None,
|
||||
api_key_is_standalone: false,
|
||||
admin_bypass_limits: false,
|
||||
local_rejection: Some(GatewayLocalAuthRejection::InvalidApiKey),
|
||||
allowed_models: None,
|
||||
ip_rules: None,
|
||||
}));
|
||||
};
|
||||
|
||||
let wallet_access = resolve_wallet_auth_gate(state, &snapshot).await?;
|
||||
let auth_context = build_data_backed_auth_context(
|
||||
state,
|
||||
snapshot,
|
||||
auth_endpoint_signature,
|
||||
None,
|
||||
None,
|
||||
wallet_access,
|
||||
)
|
||||
.await;
|
||||
info!(
|
||||
event_name = "antigravity_bearer_bridge_auth_context_resolved",
|
||||
log_type = "event",
|
||||
validation_mode,
|
||||
user_id = auth_context.user_id.as_str(),
|
||||
api_key_id = auth_context.api_key_id.as_str(),
|
||||
access_allowed = auth_context.access_allowed,
|
||||
has_local_rejection = auth_context.local_rejection.is_some(),
|
||||
"resolved Antigravity bearer bridge auth context"
|
||||
);
|
||||
Ok(Some(auth_context))
|
||||
}
|
||||
|
||||
async fn resolve_trusted_auth_context(
|
||||
state: &AppState,
|
||||
auth_endpoint_signature: &str,
|
||||
@@ -712,7 +847,12 @@ async fn build_data_backed_auth_context(
|
||||
})
|
||||
} else if snapshot
|
||||
.effective_allowed_api_formats()
|
||||
.is_some_and(|allowed| !contains_api_format_or_alias(allowed, auth_endpoint_signature))
|
||||
.is_some_and(|allowed| {
|
||||
!contains_api_format_or_alias(
|
||||
allowed,
|
||||
auth_gate_api_format(auth_endpoint_signature).as_str(),
|
||||
)
|
||||
})
|
||||
{
|
||||
Some(GatewayLocalAuthRejection::ApiFormatNotAllowed {
|
||||
api_format: auth_endpoint_signature.to_string(),
|
||||
@@ -747,6 +887,15 @@ fn normalize_api_format_alias(value: &str) -> String {
|
||||
crate::ai_serving::normalize_api_format_alias(value)
|
||||
}
|
||||
|
||||
fn auth_gate_api_format(auth_endpoint_signature: &str) -> String {
|
||||
let normalized = normalize_api_format_alias(auth_endpoint_signature);
|
||||
if normalized == "antigravity:v1internal" {
|
||||
"gemini:generate_content".to_string()
|
||||
} else {
|
||||
normalized
|
||||
}
|
||||
}
|
||||
|
||||
fn api_format_matches(left: &str, right: &str) -> bool {
|
||||
aether_scheduler_core::api_format_matches_allowed_value(left, right)
|
||||
}
|
||||
@@ -1261,6 +1410,58 @@ mod tests {
|
||||
assert_eq!(auth_context.local_rejection, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn data_backed_auth_context_allows_antigravity_v1internal_for_gemini_generate_content_keys(
|
||||
) {
|
||||
let api_key = "sk-test-antigravity-v1internal";
|
||||
let mut snapshot = sample_snapshot("key-ant-v1internal", "user-ant-v1internal");
|
||||
snapshot.user_allowed_providers = Some(vec!["antigravity".to_string()]);
|
||||
snapshot.api_key_allowed_providers = Some(vec!["antigravity".to_string()]);
|
||||
snapshot.user_allowed_api_formats = Some(vec!["gemini:generate_content".to_string()]);
|
||||
snapshot.api_key_allowed_api_formats = Some(vec!["gemini:generate_content".to_string()]);
|
||||
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key(api_key)),
|
||||
snapshot,
|
||||
)]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider(
|
||||
"provider-antigravity-1",
|
||||
"Antigravity",
|
||||
"antigravity",
|
||||
)],
|
||||
vec![sample_endpoint(
|
||||
"endpoint-antigravity-1",
|
||||
"provider-antigravity-1",
|
||||
"gemini:generate_content",
|
||||
)],
|
||||
Vec::new(),
|
||||
));
|
||||
let data = GatewayDataState::with_auth_api_key_reader_for_tests(repository)
|
||||
.with_provider_catalog_reader(provider_catalog);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data);
|
||||
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-api-key", api_key.parse().unwrap());
|
||||
headers.insert(
|
||||
http::header::AUTHORIZATION,
|
||||
"Bearer google-oauth-access-token".parse().unwrap(),
|
||||
);
|
||||
|
||||
let auth_context = resolve_data_backed_auth_context(
|
||||
&state,
|
||||
&headers,
|
||||
&uri("/v1internal:streamGenerateContent?alt=sse"),
|
||||
Some("antigravity:v1internal"),
|
||||
)
|
||||
.await
|
||||
.expect("resolution should succeed")
|
||||
.expect("auth context should exist");
|
||||
|
||||
assert_eq!(auth_context.local_rejection, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn data_backed_auth_context_allows_provider_id_for_convertible_endpoint_format() {
|
||||
let api_key = "sk-test-provider-convertible-endpoint";
|
||||
|
||||
@@ -23,6 +23,65 @@ pub(super) fn classify_admin_system_family_route(
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET && normalized_path == "/api/admin/system/releases" {
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"releases",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET
|
||||
&& normalized_path == "/api/admin/system/update-capability"
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"update_capability",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST && normalized_path == "/api/admin/system/prepare-update"
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"prepare_update",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST && normalized_path == "/api/admin/system/apply-update" {
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"apply_update",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST && normalized_path == "/api/admin/system/rollback" {
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"rollback",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET && normalized_path == "/api/admin/system/update-status" {
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"update_status",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET && normalized_path == "/api/admin/system/update-history" {
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"update_history",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET && normalized_path == "/api/admin/system/aws-regions" {
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
@@ -71,6 +130,15 @@ pub(super) fn classify_admin_system_family_route(
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST && normalized_path == "/api/admin/system/backups/s3/run"
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"system_manage",
|
||||
"s3_backup_run",
|
||||
"admin:system",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST && normalized_path == "/api/admin/system/config/import" {
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
|
||||
@@ -8,7 +8,9 @@ pub(super) fn classify_ai_public_route(
|
||||
normalized_path: &str,
|
||||
headers: &http::HeaderMap,
|
||||
) -> Option<ClassifiedRoute> {
|
||||
if method == http::Method::POST && normalized_path == "/v1/chat/completions" {
|
||||
if let Some(route) = classify_antigravity_v1internal_route(method, normalized_path) {
|
||||
Some(route)
|
||||
} else if method == http::Method::POST && normalized_path == "/v1/chat/completions" {
|
||||
Some(classified(
|
||||
"ai_public",
|
||||
"openai",
|
||||
@@ -167,3 +169,34 @@ fn is_gemini_files_method(method: &http::Method, normalized_path: &str) -> bool
|
||||
|| ((method == http::Method::GET || method == http::Method::DELETE)
|
||||
&& normalized_path.starts_with("/v1beta/files"))
|
||||
}
|
||||
|
||||
fn classify_antigravity_v1internal_route(
|
||||
method: &http::Method,
|
||||
normalized_path: &str,
|
||||
) -> Option<ClassifiedRoute> {
|
||||
if method != http::Method::POST {
|
||||
return None;
|
||||
}
|
||||
|
||||
let action = normalized_path.strip_prefix("/v1internal:")?;
|
||||
let (route_kind, execution_runtime_candidate) = match action {
|
||||
"loadCodeAssist" => ("load_code_assist", false),
|
||||
"fetchAvailableModels" => ("fetch_available_models", false),
|
||||
"fetchUserInfo" => ("fetch_user_info", false),
|
||||
"fetchAdminControls" => ("fetch_admin_controls", false),
|
||||
"setUserSettings" => ("set_user_settings", false),
|
||||
"listExperiments" => ("list_experiments", false),
|
||||
"recordCodeAssistMetrics" => ("record_code_assist_metrics", false),
|
||||
"streamGenerateContent" => ("stream_generate_content", true),
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
Some(classified_with_request_auth_channel(
|
||||
"ai_public",
|
||||
"antigravity",
|
||||
route_kind,
|
||||
"bearer_like",
|
||||
"antigravity:v1internal",
|
||||
execution_runtime_candidate,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -797,7 +797,6 @@ pub(super) fn classify_public_support_route(
|
||||
} else if method == http::Method::GET
|
||||
&& (has_single_segment_after_prefix(normalized_path, "/install/")
|
||||
|| has_single_segment_after_prefix(normalized_path, "/install-tunnel/")
|
||||
|| has_single_segment_after_prefix(normalized_path, "/install-proxy/")
|
||||
|| has_single_segment_after_prefix(normalized_path, "/i/"))
|
||||
{
|
||||
Some(classified(
|
||||
|
||||
@@ -175,6 +175,25 @@ fn classifies_admin_system_data_export_as_admin_proxy_route() {
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_system_s3_backup_start_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
let uri: Uri = "/api/admin/system/backups/s3/run"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("system_manage"));
|
||||
assert_eq!(decision.route_kind.as_deref(), Some("s3_backup_run"));
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:system")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_system_maintenance_write_routes_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
@@ -238,6 +257,55 @@ fn classifies_admin_system_check_update_as_admin_proxy_route() {
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_system_update_routes_as_admin_proxy_routes() {
|
||||
let headers = headers(&[]);
|
||||
let cases = [
|
||||
(
|
||||
http::Method::GET,
|
||||
"/api/admin/system/update-capability",
|
||||
"update_capability",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/system/prepare-update",
|
||||
"prepare_update",
|
||||
),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/system/apply-update",
|
||||
"apply_update",
|
||||
),
|
||||
(http::Method::POST, "/api/admin/system/rollback", "rollback"),
|
||||
(http::Method::GET, "/api/admin/system/releases", "releases"),
|
||||
(
|
||||
http::Method::GET,
|
||||
"/api/admin/system/update-history",
|
||||
"update_history",
|
||||
),
|
||||
(
|
||||
http::Method::GET,
|
||||
"/api/admin/system/update-status",
|
||||
"update_status",
|
||||
),
|
||||
];
|
||||
|
||||
for (method, path, expected_kind) in cases {
|
||||
let uri: Uri = path.parse().expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&method, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("system_manage"));
|
||||
assert_eq!(decision.route_kind.as_deref(), Some(expected_kind));
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:system")
|
||||
);
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_system_aws_regions_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
|
||||
@@ -291,3 +291,86 @@ fn classifies_gemini_predict_long_running_as_video_route() {
|
||||
);
|
||||
assert!(decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_antigravity_v1internal_control_plane_routes() {
|
||||
let headers = headers(&[
|
||||
("authorization", "Bearer ant-access-token"),
|
||||
("user-agent", "antigravity/cli/1.0.2 linux/arm64"),
|
||||
]);
|
||||
|
||||
for (path, route_kind) in [
|
||||
("/v1internal:loadCodeAssist", "load_code_assist"),
|
||||
("/v1internal:fetchAvailableModels", "fetch_available_models"),
|
||||
("/v1internal:fetchUserInfo", "fetch_user_info"),
|
||||
("/v1internal:fetchAdminControls", "fetch_admin_controls"),
|
||||
("/v1internal:setUserSettings", "set_user_settings"),
|
||||
("/v1internal:listExperiments", "list_experiments"),
|
||||
(
|
||||
"/v1internal:recordCodeAssistMetrics",
|
||||
"record_code_assist_metrics",
|
||||
),
|
||||
] {
|
||||
let uri: Uri = path.parse().expect("uri should parse");
|
||||
let decision = classify_control_route(&http::Method::POST, &uri, &headers)
|
||||
.expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("ai_public"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("antigravity"));
|
||||
assert_eq!(decision.route_kind.as_deref(), Some(route_kind));
|
||||
assert_eq!(
|
||||
decision.request_auth_channel.as_deref(),
|
||||
Some("bearer_like")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("antigravity:v1internal")
|
||||
);
|
||||
assert!(
|
||||
!decision.is_execution_runtime_candidate(),
|
||||
"control-plane route {path} must be handled by local facade before execution runtime"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_antigravity_stream_generate_content_as_execution_route() {
|
||||
let headers = headers(&[
|
||||
("authorization", "Bearer ant-access-token"),
|
||||
("user-agent", "antigravity/cli/1.0.2 linux/arm64"),
|
||||
]);
|
||||
let uri: Uri = "/v1internal:streamGenerateContent?alt=sse"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_class.as_deref(), Some("ai_public"));
|
||||
assert_eq!(decision.route_family.as_deref(), Some("antigravity"));
|
||||
assert_eq!(
|
||||
decision.route_kind.as_deref(),
|
||||
Some("stream_generate_content")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.request_auth_channel.as_deref(),
|
||||
Some("bearer_like")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("antigravity:v1internal")
|
||||
);
|
||||
assert!(decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unknown_antigravity_v1internal_route() {
|
||||
let headers = headers(&[
|
||||
("authorization", "Bearer ant-access-token"),
|
||||
("user-agent", "antigravity/cli/1.0.2 linux/arm64"),
|
||||
]);
|
||||
let uri: Uri = "/v1internal:deleteEverything"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
|
||||
assert!(classify_control_route(&http::Method::POST, &uri, &headers).is_none());
|
||||
}
|
||||
|
||||
@@ -1684,6 +1684,28 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn set_api_key_usage_totals(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
total_requests: u64,
|
||||
total_tokens: u64,
|
||||
total_cost_usd: f64,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
match &self.auth_api_key_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.set_api_key_usage_totals(
|
||||
api_key_id,
|
||||
total_requests,
|
||||
total_tokens,
|
||||
total_cost_usd,
|
||||
)
|
||||
.await
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn set_standalone_api_key_feature_settings(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::collections::HashSet;
|
||||
use std::collections::HashMap;
|
||||
use std::future::Future;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
@@ -13,15 +13,26 @@ use aether_data_contracts::repository::candidate_selection::{
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::Notify;
|
||||
use tokio::time::timeout;
|
||||
use tracing::warn;
|
||||
|
||||
const CANDIDATE_SELECTION_CACHE_TTL: Duration = Duration::from_secs(5);
|
||||
const CANDIDATE_SELECTION_CACHE_MAX_ENTRIES: usize = 4096;
|
||||
#[cfg(not(test))]
|
||||
const CANDIDATE_SELECTION_CACHE_LOAD_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
#[cfg(test)]
|
||||
const CANDIDATE_SELECTION_CACHE_LOAD_TIMEOUT: Duration = Duration::from_millis(50);
|
||||
#[cfg(not(test))]
|
||||
const CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
#[cfg(test)]
|
||||
const CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT: Duration = Duration::from_millis(50);
|
||||
|
||||
pub(super) struct CachedMinimalCandidateSelectionReadRepository {
|
||||
inner: Arc<dyn MinimalCandidateSelectionReadRepository>,
|
||||
entries: ExpiringMap<CandidateSelectionCacheKey, Vec<StoredMinimalCandidateSelectionRow>>,
|
||||
inflight: Mutex<HashSet<CandidateSelectionCacheKey>>,
|
||||
inflight: Mutex<HashMap<CandidateSelectionCacheKey, u64>>,
|
||||
inflight_notify: Notify,
|
||||
next_inflight_token: AtomicU64,
|
||||
epoch: AtomicU64,
|
||||
}
|
||||
|
||||
@@ -30,8 +41,9 @@ impl CachedMinimalCandidateSelectionReadRepository {
|
||||
Self {
|
||||
inner,
|
||||
entries: ExpiringMap::new(),
|
||||
inflight: Mutex::new(HashSet::new()),
|
||||
inflight: Mutex::new(HashMap::new()),
|
||||
inflight_notify: Notify::new(),
|
||||
next_inflight_token: AtomicU64::new(1),
|
||||
epoch: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
@@ -52,67 +64,169 @@ impl CachedMinimalCandidateSelectionReadRepository {
|
||||
loop {
|
||||
let notified = self.inflight_notify.notified();
|
||||
match self.register_inflight(&key) {
|
||||
InflightRegistration::Bypass => return load().await,
|
||||
InflightRegistration::Bypass => {
|
||||
return load_candidate_selection_rows_with_timeout(&key, load()).await;
|
||||
}
|
||||
InflightRegistration::Follower => {
|
||||
notified.await;
|
||||
if timeout(CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT, notified)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
self.expire_inflight(&key);
|
||||
}
|
||||
if let Some(rows) = self.entries.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
|
||||
{
|
||||
return Ok(rows);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
InflightRegistration::Leader => {}
|
||||
}
|
||||
|
||||
let load_epoch = self.epoch.load(Ordering::Acquire);
|
||||
let result = load().await;
|
||||
if let Ok(rows) = &result {
|
||||
if load_epoch == self.epoch.load(Ordering::Acquire) {
|
||||
self.entries.insert(
|
||||
key.clone(),
|
||||
rows.clone(),
|
||||
CANDIDATE_SELECTION_CACHE_TTL,
|
||||
CANDIDATE_SELECTION_CACHE_MAX_ENTRIES,
|
||||
);
|
||||
InflightRegistration::Leader(token) => {
|
||||
let mut guard = InflightGuard::new(self, key.clone(), token);
|
||||
let load_epoch = self.epoch.load(Ordering::Acquire);
|
||||
let result = load_candidate_selection_rows_with_timeout(&key, load()).await;
|
||||
if let Ok(rows) = &result {
|
||||
if load_epoch == self.epoch.load(Ordering::Acquire) {
|
||||
self.entries.insert(
|
||||
key.clone(),
|
||||
rows.clone(),
|
||||
CANDIDATE_SELECTION_CACHE_TTL,
|
||||
CANDIDATE_SELECTION_CACHE_MAX_ENTRIES,
|
||||
);
|
||||
}
|
||||
}
|
||||
guard.finish();
|
||||
return result;
|
||||
}
|
||||
}
|
||||
self.finish_inflight(&key);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
fn register_inflight(&self, key: &CandidateSelectionCacheKey) -> InflightRegistration {
|
||||
match self.inflight.lock() {
|
||||
Ok(mut inflight) => {
|
||||
if inflight.insert(key.clone()) {
|
||||
InflightRegistration::Leader
|
||||
} else {
|
||||
InflightRegistration::Follower
|
||||
if inflight.contains_key(key) {
|
||||
return InflightRegistration::Follower;
|
||||
}
|
||||
let token = self.next_inflight_token.fetch_add(1, Ordering::AcqRel);
|
||||
inflight.insert(key.clone(), token);
|
||||
InflightRegistration::Leader(token)
|
||||
}
|
||||
Err(_) => InflightRegistration::Bypass,
|
||||
}
|
||||
}
|
||||
|
||||
fn finish_inflight(&self, key: &CandidateSelectionCacheKey) {
|
||||
fn finish_inflight(&self, key: &CandidateSelectionCacheKey, token: u64) {
|
||||
let mut removed = false;
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
inflight.remove(key);
|
||||
if inflight.get(key).copied() == Some(token) {
|
||||
inflight.remove(key);
|
||||
removed = true;
|
||||
}
|
||||
}
|
||||
if removed {
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
fn expire_inflight(&self, key: &CandidateSelectionCacheKey) {
|
||||
let mut removed = false;
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
removed = inflight.remove(key).is_some();
|
||||
}
|
||||
if removed {
|
||||
warn!(
|
||||
event_name = "candidate_selection_cache_inflight_expired",
|
||||
log_type = "ops",
|
||||
cache_key = ?key,
|
||||
wait_timeout_ms = CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT.as_millis() as u64,
|
||||
"gateway candidate selection cache expired stale inflight load"
|
||||
);
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
|
||||
fn clear(&self) {
|
||||
self.epoch.fetch_add(1, Ordering::AcqRel);
|
||||
self.entries.clear();
|
||||
let mut cleared_inflight = false;
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
cleared_inflight = !inflight.is_empty();
|
||||
inflight.clear();
|
||||
}
|
||||
if cleared_inflight {
|
||||
warn!(
|
||||
event_name = "candidate_selection_cache_inflight_cleared",
|
||||
log_type = "ops",
|
||||
"gateway candidate selection cache cleared in-flight loads"
|
||||
);
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
enum InflightRegistration {
|
||||
Leader,
|
||||
Leader(u64),
|
||||
Follower,
|
||||
Bypass,
|
||||
}
|
||||
|
||||
struct InflightGuard<'a> {
|
||||
cache: &'a CachedMinimalCandidateSelectionReadRepository,
|
||||
key: Option<CandidateSelectionCacheKey>,
|
||||
token: u64,
|
||||
}
|
||||
|
||||
impl<'a> InflightGuard<'a> {
|
||||
fn new(
|
||||
cache: &'a CachedMinimalCandidateSelectionReadRepository,
|
||||
key: CandidateSelectionCacheKey,
|
||||
token: u64,
|
||||
) -> Self {
|
||||
Self {
|
||||
cache,
|
||||
key: Some(key),
|
||||
token,
|
||||
}
|
||||
}
|
||||
|
||||
fn finish(&mut self) {
|
||||
if let Some(key) = self.key.take() {
|
||||
self.cache.finish_inflight(&key, self.token);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for InflightGuard<'_> {
|
||||
fn drop(&mut self) {
|
||||
self.finish();
|
||||
}
|
||||
}
|
||||
|
||||
async fn load_candidate_selection_rows_with_timeout<Fut>(
|
||||
key: &CandidateSelectionCacheKey,
|
||||
load: Fut,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>
|
||||
where
|
||||
Fut: Future<Output = Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>>,
|
||||
{
|
||||
match timeout(CANDIDATE_SELECTION_CACHE_LOAD_TIMEOUT, load).await {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
warn!(
|
||||
event_name = "candidate_selection_cache_load_timeout",
|
||||
log_type = "ops",
|
||||
cache_key = ?key,
|
||||
timeout_ms = CANDIDATE_SELECTION_CACHE_LOAD_TIMEOUT.as_millis() as u64,
|
||||
"gateway candidate selection cache load timed out"
|
||||
);
|
||||
Err(DataLayerError::TimedOut(format!(
|
||||
"candidate selection cache load exceeded {}ms for {key:?}",
|
||||
CANDIDATE_SELECTION_CACHE_LOAD_TIMEOUT.as_millis()
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for CachedMinimalCandidateSelectionReadRepository {
|
||||
fn clear_local_cache(&self) {
|
||||
@@ -286,6 +400,7 @@ fn normalize_api_format_key(api_format: &str) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::future::pending;
|
||||
use std::sync::atomic::AtomicUsize;
|
||||
|
||||
struct StubCandidateSelectionRepository {
|
||||
@@ -361,6 +476,77 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
struct FirstLoadPendingThenFastRepository {
|
||||
calls: AtomicUsize,
|
||||
}
|
||||
|
||||
impl FirstLoadPendingThenFastRepository {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
calls: AtomicUsize::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
fn calls(&self) -> usize {
|
||||
self.calls.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
async fn load(&self) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let call = self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
if call == 0 {
|
||||
pending::<()>().await;
|
||||
}
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for FirstLoadPendingThenFastRepository {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
_query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group_key_ids(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn candidate_selection_cache_coalesces_concurrent_loads() {
|
||||
let inner = Arc::new(StubCandidateSelectionRepository::new(
|
||||
@@ -398,4 +584,117 @@ mod tests {
|
||||
cache.list_for_exact_api_format("openai").await.unwrap();
|
||||
assert_eq!(inner.calls(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn candidate_selection_cache_releases_inflight_when_leader_is_cancelled() {
|
||||
let inner = Arc::new(FirstLoadPendingThenFastRepository::new());
|
||||
let cache = Arc::new(CachedMinimalCandidateSelectionReadRepository::new(
|
||||
inner.clone(),
|
||||
));
|
||||
let leader_cache = cache.clone();
|
||||
let leader = tokio::spawn(async move {
|
||||
leader_cache
|
||||
.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "gpt-5.5".to_string(),
|
||||
offset: 0,
|
||||
limit: 64,
|
||||
},
|
||||
)
|
||||
.await
|
||||
});
|
||||
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
leader.abort();
|
||||
let _ = leader.await;
|
||||
|
||||
tokio::time::timeout(
|
||||
Duration::from_millis(200),
|
||||
cache.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "gpt-5.5".to_string(),
|
||||
offset: 0,
|
||||
limit: 64,
|
||||
},
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("cancelled leader must not leave a permanent inflight wait")
|
||||
.unwrap();
|
||||
assert_eq!(inner.calls(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn candidate_selection_cache_times_out_and_clears_stuck_load() {
|
||||
let inner = Arc::new(FirstLoadPendingThenFastRepository::new());
|
||||
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner.clone());
|
||||
let err = cache
|
||||
.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "gpt-5.5".to_string(),
|
||||
offset: 0,
|
||||
limit: 64,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, DataLayerError::TimedOut(_)));
|
||||
|
||||
cache
|
||||
.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "gpt-5.5".to_string(),
|
||||
offset: 0,
|
||||
limit: 64,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(inner.calls(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn candidate_selection_cache_clear_releases_inflight_waiters() {
|
||||
let inner = Arc::new(FirstLoadPendingThenFastRepository::new());
|
||||
let cache = Arc::new(CachedMinimalCandidateSelectionReadRepository::new(
|
||||
inner.clone(),
|
||||
));
|
||||
let leader_cache = cache.clone();
|
||||
let leader = tokio::spawn(async move {
|
||||
leader_cache
|
||||
.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "gpt-5.5".to_string(),
|
||||
offset: 0,
|
||||
limit: 64,
|
||||
},
|
||||
)
|
||||
.await
|
||||
});
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
|
||||
cache.clear_local_cache();
|
||||
tokio::time::timeout(
|
||||
Duration::from_millis(200),
|
||||
cache.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "gpt-5.5".to_string(),
|
||||
offset: 0,
|
||||
limit: 64,
|
||||
},
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("cache clear must release stale inflight waiters")
|
||||
.unwrap();
|
||||
assert_eq!(inner.calls(), 2);
|
||||
leader.abort();
|
||||
let _ = leader.await;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,7 +2,8 @@ use super::{
|
||||
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
|
||||
GeminiFileMappingStats, ProviderCatalogKeyListQuery, PublicHealthStatusCount,
|
||||
PublicHealthTimelineBucket, StoredGeminiFileMapping, StoredGeminiFileMappingListPage,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
|
||||
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
|
||||
};
|
||||
@@ -285,6 +286,20 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_provider_catalog_key_maintenance_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKeyMaintenanceSummary>, DataLayerError> {
|
||||
match &self.provider_catalog_reader {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.list_key_maintenance_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
}
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_provider_catalog_key_page(
|
||||
&self,
|
||||
query: &ProviderCatalogKeyListQuery,
|
||||
@@ -319,7 +334,7 @@ impl GatewayDataState {
|
||||
encrypted_auth_config: Option<&str>,
|
||||
expires_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.update_key_oauth_credentials(
|
||||
@@ -331,17 +346,25 @@ impl GatewayDataState {
|
||||
.await
|
||||
}
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_key(
|
||||
&self,
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<Option<StoredProviderCatalogKey>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let created = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.create_key(key).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if created.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(created)
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_provider(
|
||||
@@ -349,33 +372,45 @@ impl GatewayDataState {
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
shift_existing_priorities_from: Option<i32>,
|
||||
) -> Result<Option<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let created = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository
|
||||
.create_provider(provider, shift_existing_priorities_from)
|
||||
.await
|
||||
.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if created.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(created)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_provider(
|
||||
&self,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Result<Option<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.update_provider(provider).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if updated.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_provider_catalog_provider(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let deleted = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.delete_provider(provider_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if deleted {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_deleted_provider_catalog_refs(
|
||||
@@ -384,54 +419,74 @@ impl GatewayDataState {
|
||||
endpoint_ids: &[String],
|
||||
key_ids: &[String],
|
||||
) -> Result<(), DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let cleaned = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.cleanup_deleted_provider_refs(provider_id, endpoint_ids, key_ids)
|
||||
.await
|
||||
}
|
||||
None => Ok(()),
|
||||
};
|
||||
if !endpoint_ids.is_empty() || !key_ids.is_empty() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
cleaned
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_endpoint(
|
||||
&self,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
) -> Result<Option<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let created = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.create_endpoint(endpoint).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if created.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(created)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_endpoint(
|
||||
&self,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
) -> Result<Option<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.update_endpoint(endpoint).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if updated.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_provider_catalog_endpoint(
|
||||
&self,
|
||||
endpoint_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let deleted = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.delete_endpoint(endpoint_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if deleted {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key(
|
||||
&self,
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<Option<StoredProviderCatalogKey>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.update_key(key).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if updated.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_upstream_metadata(
|
||||
@@ -440,34 +495,46 @@ impl GatewayDataState {
|
||||
upstream_metadata: Option<&serde_json::Value>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.update_key_upstream_metadata(key_id, upstream_metadata, updated_at_unix_secs)
|
||||
.await
|
||||
}
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_provider_catalog_key(
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let deleted = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.delete_key(key_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if deleted {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn clear_provider_catalog_key_oauth_invalid_marker(
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.clear_key_oauth_invalid_marker(key_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_health_state(
|
||||
@@ -477,7 +544,7 @@ impl GatewayDataState {
|
||||
health_by_format: Option<&serde_json::Value>,
|
||||
circuit_breaker_by_format: Option<&serde_json::Value>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.update_key_health_state(
|
||||
@@ -489,6 +556,10 @@ impl GatewayDataState {
|
||||
.await
|
||||
}
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use aether_data::{DataBackends, DataLayerError, DatabaseDriver};
|
||||
use aether_data_contracts::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
|
||||
use aether_runtime_state::RuntimeQueueStore;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -99,7 +100,11 @@ impl GatewayDataState {
|
||||
let request_candidate_reader = backends.read().request_candidates();
|
||||
let request_candidate_writer = backends.write().request_candidates();
|
||||
let gemini_file_mapping_writer = backends.write().gemini_file_mappings();
|
||||
let provider_catalog_reader = backends.read().provider_catalog();
|
||||
let provider_catalog_reader = backends.read().provider_catalog().map(|repository| {
|
||||
Arc::new(
|
||||
super::provider_catalog_cache::CachedProviderCatalogReadRepository::new(repository),
|
||||
) as Arc<dyn ProviderCatalogReadRepository>
|
||||
});
|
||||
let provider_catalog_writer = backends.write().provider_catalog();
|
||||
let pool_score_reader = backends.read().pool_scores();
|
||||
let pool_score_writer = backends.write().pool_scores();
|
||||
@@ -276,6 +281,12 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn clear_provider_catalog_cache(&self) {
|
||||
if let Some(repository) = &self.provider_catalog_reader {
|
||||
repository.clear_local_cache();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn has_request_candidate_reader(&self) -> bool {
|
||||
self.request_candidate_reader.is_some()
|
||||
}
|
||||
@@ -511,6 +522,45 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn export_admin_system_usage_aggregates(
|
||||
&self,
|
||||
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateSnapshot, DataLayerError>
|
||||
{
|
||||
match self.backends.as_ref() {
|
||||
Some(backends) => backends.export_admin_system_usage_aggregates().await,
|
||||
None => {
|
||||
Ok(aether_data::repository::system::AdminSystemUsageAggregateSnapshot::default())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn import_admin_system_usage_aggregates(
|
||||
&self,
|
||||
snapshot: &aether_data::repository::system::AdminSystemUsageAggregateSnapshot,
|
||||
user_id_map: &std::collections::BTreeMap<String, String>,
|
||||
api_key_id_map: &std::collections::BTreeMap<String, String>,
|
||||
mode: aether_data::repository::system::AdminSystemUsageAggregateImportMode,
|
||||
) -> Result<
|
||||
aether_data::repository::system::AdminSystemUsageAggregateImportSummary,
|
||||
DataLayerError,
|
||||
> {
|
||||
match self.backends.as_ref() {
|
||||
Some(backends) => {
|
||||
backends
|
||||
.import_admin_system_usage_aggregates(
|
||||
snapshot,
|
||||
user_id_map,
|
||||
api_key_id_map,
|
||||
mode,
|
||||
)
|
||||
.await
|
||||
}
|
||||
None => Ok(
|
||||
aether_data::repository::system::AdminSystemUsageAggregateImportSummary::default(),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn purge_admin_request_bodies_batch(
|
||||
&self,
|
||||
batch_size: usize,
|
||||
|
||||
@@ -119,7 +119,8 @@ use aether_data_contracts::repository::pool_scores::{
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_data_contracts::repository::quota::{
|
||||
@@ -323,6 +324,7 @@ mod core;
|
||||
mod integrations;
|
||||
mod models;
|
||||
mod pool_scores;
|
||||
mod provider_catalog_cache;
|
||||
mod referrals;
|
||||
mod routing_profiles;
|
||||
mod runtime;
|
||||
|
||||
@@ -0,0 +1,360 @@
|
||||
use std::collections::HashMap;
|
||||
use std::future::Future;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_cache::ExpiringMap;
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
|
||||
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::Notify;
|
||||
|
||||
const PROVIDER_CATALOG_CACHE_TTL: Duration = Duration::from_secs(5);
|
||||
const PROVIDER_CATALOG_CACHE_MAX_ENTRIES: usize = 1024;
|
||||
|
||||
pub(super) struct CachedProviderCatalogReadRepository {
|
||||
inner: Arc<dyn ProviderCatalogReadRepository>,
|
||||
entries: ExpiringMap<ProviderCatalogCacheKey, ProviderCatalogCacheValue>,
|
||||
inflight: Mutex<HashMap<ProviderCatalogCacheKey, u64>>,
|
||||
inflight_notify: Notify,
|
||||
next_inflight_token: AtomicU64,
|
||||
epoch: AtomicU64,
|
||||
}
|
||||
|
||||
impl CachedProviderCatalogReadRepository {
|
||||
pub(super) fn new(inner: Arc<dyn ProviderCatalogReadRepository>) -> Self {
|
||||
Self {
|
||||
inner,
|
||||
entries: ExpiringMap::new(),
|
||||
inflight: Mutex::new(HashMap::new()),
|
||||
inflight_notify: Notify::new(),
|
||||
next_inflight_token: AtomicU64::new(1),
|
||||
epoch: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_or_load<F, Fut>(
|
||||
&self,
|
||||
key: ProviderCatalogCacheKey,
|
||||
load: F,
|
||||
) -> Result<ProviderCatalogCacheValue, DataLayerError>
|
||||
where
|
||||
F: Fn() -> Fut,
|
||||
Fut: Future<Output = Result<ProviderCatalogCacheValue, DataLayerError>>,
|
||||
{
|
||||
if let Some(value) = self.entries.get_fresh(&key, PROVIDER_CATALOG_CACHE_TTL) {
|
||||
return Ok(value);
|
||||
}
|
||||
|
||||
loop {
|
||||
let notified = self.inflight_notify.notified();
|
||||
match self.register_inflight(&key) {
|
||||
InflightRegistration::Bypass => return load().await,
|
||||
InflightRegistration::Follower => {
|
||||
notified.await;
|
||||
if let Some(value) = self.entries.get_fresh(&key, PROVIDER_CATALOG_CACHE_TTL) {
|
||||
return Ok(value);
|
||||
}
|
||||
}
|
||||
InflightRegistration::Leader(token) => {
|
||||
let mut guard = InflightGuard::new(self, key.clone(), token);
|
||||
let load_epoch = self.epoch.load(Ordering::Acquire);
|
||||
let result = load().await;
|
||||
if let Ok(value) = &result {
|
||||
if load_epoch == self.epoch.load(Ordering::Acquire) {
|
||||
self.entries.insert(
|
||||
key.clone(),
|
||||
value.clone(),
|
||||
PROVIDER_CATALOG_CACHE_TTL,
|
||||
PROVIDER_CATALOG_CACHE_MAX_ENTRIES,
|
||||
);
|
||||
}
|
||||
}
|
||||
guard.finish();
|
||||
return result;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn register_inflight(&self, key: &ProviderCatalogCacheKey) -> InflightRegistration {
|
||||
match self.inflight.lock() {
|
||||
Ok(mut inflight) => {
|
||||
if inflight.contains_key(key) {
|
||||
return InflightRegistration::Follower;
|
||||
}
|
||||
let token = self.next_inflight_token.fetch_add(1, Ordering::AcqRel);
|
||||
inflight.insert(key.clone(), token);
|
||||
InflightRegistration::Leader(token)
|
||||
}
|
||||
Err(_) => InflightRegistration::Bypass,
|
||||
}
|
||||
}
|
||||
|
||||
fn finish_inflight(&self, key: &ProviderCatalogCacheKey, token: u64) {
|
||||
let mut removed = false;
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
if inflight.get(key).copied() == Some(token) {
|
||||
inflight.remove(key);
|
||||
removed = true;
|
||||
}
|
||||
}
|
||||
if removed {
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
fn clear(&self) {
|
||||
self.epoch.fetch_add(1, Ordering::AcqRel);
|
||||
self.entries.clear();
|
||||
let mut cleared_inflight = false;
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
cleared_inflight = !inflight.is_empty();
|
||||
inflight.clear();
|
||||
}
|
||||
if cleared_inflight {
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ProviderCatalogReadRepository for CachedProviderCatalogReadRepository {
|
||||
fn clear_local_cache(&self) {
|
||||
self.clear();
|
||||
self.inner.clear_local_cache();
|
||||
}
|
||||
|
||||
async fn list_providers(
|
||||
&self,
|
||||
active_only: bool,
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
match self
|
||||
.get_or_load(
|
||||
ProviderCatalogCacheKey::Providers { active_only },
|
||||
|| async move {
|
||||
self.inner
|
||||
.list_providers(active_only)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Providers)
|
||||
},
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Providers(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_providers_by_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::ProvidersByIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_providers_by_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Providers)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Providers(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_endpoints_by_ids(
|
||||
&self,
|
||||
endpoint_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
self.inner.list_endpoints_by_ids(endpoint_ids).await
|
||||
}
|
||||
|
||||
async fn list_endpoints_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::EndpointsByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_endpoints_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Endpoints)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Endpoints(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_keys_by_ids(
|
||||
&self,
|
||||
key_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.inner.list_keys_by_ids(key_ids).await
|
||||
}
|
||||
|
||||
async fn list_keys_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeysByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_keys_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Keys)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Keys(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_key_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeySummariesByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_key_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Keys)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Keys(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_key_maintenance_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKeyMaintenanceSummary>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeyMaintenanceSummariesByProviderIds(normalize_ids(
|
||||
provider_ids,
|
||||
));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_key_maintenance_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::KeyMaintenanceSummaries)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::KeyMaintenanceSummaries(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_keys_page(
|
||||
&self,
|
||||
query: &ProviderCatalogKeyListQuery,
|
||||
) -> Result<StoredProviderCatalogKeyPage, DataLayerError> {
|
||||
self.inner.list_keys_page(query).await
|
||||
}
|
||||
|
||||
async fn list_key_stats_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKeyStats>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeyStatsByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_key_stats_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::KeyStats)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::KeyStats(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
|
||||
enum ProviderCatalogCacheKey {
|
||||
Providers { active_only: bool },
|
||||
ProvidersByIds(Vec<String>),
|
||||
EndpointsByProviderIds(Vec<String>),
|
||||
KeysByProviderIds(Vec<String>),
|
||||
KeySummariesByProviderIds(Vec<String>),
|
||||
KeyMaintenanceSummariesByProviderIds(Vec<String>),
|
||||
KeyStatsByProviderIds(Vec<String>),
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
enum ProviderCatalogCacheValue {
|
||||
Providers(Vec<StoredProviderCatalogProvider>),
|
||||
Endpoints(Vec<StoredProviderCatalogEndpoint>),
|
||||
Keys(Vec<StoredProviderCatalogKey>),
|
||||
KeyMaintenanceSummaries(Vec<StoredProviderCatalogKeyMaintenanceSummary>),
|
||||
KeyStats(Vec<StoredProviderCatalogKeyStats>),
|
||||
}
|
||||
|
||||
enum InflightRegistration {
|
||||
Leader(u64),
|
||||
Follower,
|
||||
Bypass,
|
||||
}
|
||||
|
||||
struct InflightGuard<'a> {
|
||||
cache: &'a CachedProviderCatalogReadRepository,
|
||||
key: Option<ProviderCatalogCacheKey>,
|
||||
token: u64,
|
||||
}
|
||||
|
||||
impl<'a> InflightGuard<'a> {
|
||||
fn new(
|
||||
cache: &'a CachedProviderCatalogReadRepository,
|
||||
key: ProviderCatalogCacheKey,
|
||||
token: u64,
|
||||
) -> Self {
|
||||
Self {
|
||||
cache,
|
||||
key: Some(key),
|
||||
token,
|
||||
}
|
||||
}
|
||||
|
||||
fn finish(&mut self) {
|
||||
if let Some(key) = self.key.take() {
|
||||
self.cache.finish_inflight(&key, self.token);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for InflightGuard<'_> {
|
||||
fn drop(&mut self) {
|
||||
self.finish();
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_ids(ids: &[String]) -> Vec<String> {
|
||||
let mut normalized = ids
|
||||
.iter()
|
||||
.map(|id| id.trim())
|
||||
.filter(|id| !id.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<Vec<_>>();
|
||||
normalized.sort();
|
||||
normalized.dedup();
|
||||
normalized
|
||||
}
|
||||
@@ -9,7 +9,8 @@ use aether_data_contracts::repository::usage::UsageRepository;
|
||||
use super::{
|
||||
AnnouncementReadRepository, AnnouncementWriteRepository, AuthApiKeyReadRepository,
|
||||
AuthApiKeyWriteRepository, AuthModuleReadRepository, AuthModuleWriteRepository,
|
||||
BillingReadRepository, GatewayDataConfig, GatewayDataState, GeminiFileMappingReadRepository,
|
||||
BackgroundTaskReadRepository, BackgroundTaskWriteRepository, BillingReadRepository,
|
||||
GatewayDataConfig, GatewayDataState, GeminiFileMappingReadRepository,
|
||||
GeminiFileMappingWriteRepository, GlobalModelReadRepository, GlobalModelWriteRepository,
|
||||
ManagementTokenReadRepository, ManagementTokenWriteRepository,
|
||||
MinimalCandidateSelectionReadRepository, OAuthProviderReadRepository,
|
||||
@@ -261,6 +262,18 @@ impl GatewayDataState {
|
||||
self
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn with_background_task_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
|
||||
where
|
||||
T: BackgroundTaskReadRepository + BackgroundTaskWriteRepository + 'static,
|
||||
{
|
||||
let background_task_reader: Arc<dyn BackgroundTaskReadRepository> = repository.clone();
|
||||
let background_task_writer: Arc<dyn BackgroundTaskWriteRepository> = repository;
|
||||
self.background_task_reader = Some(background_task_reader);
|
||||
self.background_task_writer = Some(background_task_writer);
|
||||
self
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn with_global_model_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
|
||||
where
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
use std::collections::{btree_map::Entry, BTreeMap, BTreeSet, VecDeque};
|
||||
use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
|
||||
use std::sync::{
|
||||
atomic::{AtomicU64, Ordering as AtomicOrdering},
|
||||
Arc, LazyLock,
|
||||
};
|
||||
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
@@ -19,7 +22,8 @@ use aether_pool_core::{
|
||||
};
|
||||
use aether_provider_pool::ProviderPoolService;
|
||||
use aether_routing_core::{RankingOverlay, ResolvedRoutingPolicy};
|
||||
use tracing::warn;
|
||||
use tokio::sync::Semaphore;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
use crate::ai_serving::{
|
||||
candidate_auth_channel_skip_reason, candidate_common_transport_skip_reason,
|
||||
@@ -43,8 +47,13 @@ use crate::maintenance::spawn_pool_quota_probe_replenish_for_request;
|
||||
use crate::orchestration::LocalExecutionCandidateMetadata;
|
||||
|
||||
static LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
static POOL_SCORE_SCHEDULE_INTEREST_SEMAPHORE: LazyLock<Arc<Semaphore>> =
|
||||
LazyLock::new(|| Arc::new(Semaphore::new(POOL_SCORE_SCHEDULE_INTEREST_CONCURRENCY)));
|
||||
const POOL_ACTIVE_PROBE_SEALED_SKIP_REASON: &str = "pool_active_probe_sealed";
|
||||
const ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON: &str = "routing_profile_disallowed_key";
|
||||
const POOL_SCORE_SCHEDULE_INTEREST_CONCURRENCY: usize = 4;
|
||||
const POOL_SCORE_SCHEDULE_INTEREST_MAX_PER_BATCH: usize = 16;
|
||||
const POOL_SCORE_SCHEDULE_INTEREST_MIN_INTERVAL_SECS: u64 = 60;
|
||||
|
||||
type PoolCatalogKeyContext = PoolMemberSignals;
|
||||
|
||||
@@ -334,9 +343,11 @@ pub(crate) struct PoolKeyCursor<'a> {
|
||||
pool_key_order: StoredPoolKeyCandidateOrder,
|
||||
next_offset: u32,
|
||||
scanned_keys: u32,
|
||||
budget_scanned_keys: u32,
|
||||
window_size: u32,
|
||||
page_size: u32,
|
||||
max_scanned_keys: u32,
|
||||
absolute_max_scanned_keys: u32,
|
||||
score_top_n: u32,
|
||||
score_phase_loaded: bool,
|
||||
skip_reason_counts: BTreeMap<&'static str, u32>,
|
||||
@@ -384,13 +395,14 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
.map(|config| config.score_top_n)
|
||||
.unwrap_or(u64::from(aether_dispatch_core::DEFAULT_POOL_PAGE_SIZE))
|
||||
.clamp(1, u64::from(u32::MAX)) as u32;
|
||||
let max_scanned_keys = pool_config
|
||||
let configured_max_scanned_keys = pool_config
|
||||
.as_ref()
|
||||
.map(|config| config.score_fallback_scan_limit)
|
||||
.unwrap_or(u64::from(aether_dispatch_core::DEFAULT_POOL_MAX_SCAN))
|
||||
.clamp(1, u64::from(u32::MAX)) as u32;
|
||||
let window_config = crate::dispatch::pool::default_pool_window_config().normalized();
|
||||
let max_scanned_keys = max_scanned_keys.min(window_config.max_scan);
|
||||
let max_scanned_keys = configured_max_scanned_keys.min(window_config.max_scan);
|
||||
let absolute_max_scanned_keys = configured_max_scanned_keys.max(max_scanned_keys);
|
||||
Self {
|
||||
state,
|
||||
group,
|
||||
@@ -403,9 +415,11 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
pool_key_order,
|
||||
next_offset: 0,
|
||||
scanned_keys: 0,
|
||||
budget_scanned_keys: 0,
|
||||
window_size: window_config.window_size,
|
||||
page_size: window_config.page_size,
|
||||
max_scanned_keys: max_scanned_keys.max(window_config.window_size),
|
||||
absolute_max_scanned_keys: absolute_max_scanned_keys.max(window_config.window_size),
|
||||
score_top_n,
|
||||
score_phase_loaded: false,
|
||||
skip_reason_counts: BTreeMap::new(),
|
||||
@@ -477,6 +491,7 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
extra_data: Some(serde_json::json!({
|
||||
"pool_group_exhaustion": {
|
||||
"scanned_keys": self.scanned_keys,
|
||||
"budget_scanned_keys": self.budget_scanned_keys,
|
||||
"skip_reason_counts": skip_reason_counts,
|
||||
}
|
||||
})),
|
||||
@@ -495,6 +510,9 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
endpoint_id = %self.group.candidate.endpoint_id,
|
||||
model_id = %self.group.candidate.model_id,
|
||||
scanned_keys = self.scanned_keys,
|
||||
budget_scanned_keys = self.budget_scanned_keys,
|
||||
max_scanned_keys = self.max_scanned_keys,
|
||||
absolute_max_scanned_keys = self.absolute_max_scanned_keys,
|
||||
skip_reason_counts = ?self.skip_reason_counts,
|
||||
"gateway pool scheduler exhausted pool group without a schedulable key"
|
||||
);
|
||||
@@ -539,13 +557,16 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
}
|
||||
}
|
||||
|
||||
if self.scanned_keys >= self.max_scanned_keys {
|
||||
if self.budget_scanned_keys >= self.max_scanned_keys
|
||||
|| self.scanned_keys >= self.absolute_max_scanned_keys
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let limit = self
|
||||
.page_size
|
||||
.min(self.max_scanned_keys - self.scanned_keys);
|
||||
.min(self.max_scanned_keys - self.budget_scanned_keys)
|
||||
.min(self.absolute_max_scanned_keys - self.scanned_keys);
|
||||
let query = StoredPoolKeyCandidateRowsQuery {
|
||||
api_format: self.group.candidate.endpoint_api_format.clone(),
|
||||
provider_id: self.group.candidate.provider_id.clone(),
|
||||
@@ -584,11 +605,21 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
}
|
||||
|
||||
self.scanned_keys += rows.len() as u32;
|
||||
self.budget_scanned_keys += rows.len() as u32;
|
||||
self.next_offset = self.next_offset.saturating_add(rows.len() as u32);
|
||||
Some(self.build_page_eligible_candidates(rows).await)
|
||||
}
|
||||
|
||||
async fn next_score_candidates(&mut self) -> Option<Vec<EligibleLocalExecutionCandidate>> {
|
||||
if self.scanned_keys >= self.absolute_max_scanned_keys {
|
||||
return None;
|
||||
}
|
||||
let limit = self
|
||||
.score_top_n
|
||||
.min(self.absolute_max_scanned_keys - self.scanned_keys);
|
||||
if limit == 0 {
|
||||
return None;
|
||||
}
|
||||
let scope = provider_key_pool_score_scope();
|
||||
let query = ListRankedPoolMembersQuery {
|
||||
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
|
||||
@@ -599,7 +630,7 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
hard_states: vec![PoolMemberHardState::Available, PoolMemberHardState::Unknown],
|
||||
probe_statuses: None,
|
||||
offset: 0,
|
||||
limit: self.score_top_n as usize,
|
||||
limit: limit as usize,
|
||||
};
|
||||
let scores = match self.state.app().data.list_ranked_pool_members(&query).await {
|
||||
Ok(scores) => scores,
|
||||
@@ -621,7 +652,7 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
return None;
|
||||
}
|
||||
|
||||
self.record_score_schedule_interest(&scores).await;
|
||||
self.spawn_score_schedule_interest_recording(&scores);
|
||||
|
||||
let key_ids = scores
|
||||
.iter()
|
||||
@@ -658,6 +689,7 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
}
|
||||
};
|
||||
self.scanned_keys = self.scanned_keys.saturating_add(scores.len() as u32);
|
||||
self.budget_scanned_keys = self.budget_scanned_keys.saturating_add(scores.len() as u32);
|
||||
Some(self.build_page_eligible_candidates(rows).await)
|
||||
}
|
||||
|
||||
@@ -849,60 +881,97 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
);
|
||||
}
|
||||
|
||||
async fn record_score_schedule_interest(&self, scores: &[StoredPoolMemberScore]) {
|
||||
if scores.is_empty() {
|
||||
fn spawn_score_schedule_interest_recording(&self, scores: &[StoredPoolMemberScore]) {
|
||||
if scores.is_empty() || !self.state.app().data.has_pool_score_writer() {
|
||||
return;
|
||||
}
|
||||
|
||||
let scheduled_at = current_unix_ms() / 1000;
|
||||
let mut failed = 0usize;
|
||||
for score in scores {
|
||||
let identity = PoolMemberIdentity {
|
||||
pool_kind: score.pool_kind.clone(),
|
||||
pool_id: score.pool_id.clone(),
|
||||
member_kind: score.member_kind.clone(),
|
||||
member_id: score.member_id.clone(),
|
||||
};
|
||||
let scope = PoolScoreScope {
|
||||
capability: score.capability.clone(),
|
||||
scope_kind: score.scope_kind.clone(),
|
||||
scope_id: score.scope_id.clone(),
|
||||
};
|
||||
let result = self
|
||||
.state
|
||||
.app()
|
||||
.data
|
||||
.record_pool_member_schedule_feedback(PoolMemberScheduleFeedback {
|
||||
identity,
|
||||
scope: Some(scope),
|
||||
scheduled_at,
|
||||
succeeded: None,
|
||||
hard_state: None,
|
||||
score_delta: None,
|
||||
score_reason_patch: Some(serde_json::json!({
|
||||
"last_schedule_interest": {
|
||||
"provider_id": self.group.candidate.provider_id.as_str(),
|
||||
"endpoint_id": self.group.candidate.endpoint_id.as_str(),
|
||||
"model_id": self.group.candidate.model_id.as_str()
|
||||
}
|
||||
})),
|
||||
let provider_id = self.group.candidate.provider_id.clone();
|
||||
let endpoint_id = self.group.candidate.endpoint_id.clone();
|
||||
let model_id = self.group.candidate.model_id.clone();
|
||||
let feedback = scores
|
||||
.iter()
|
||||
.filter(|score| {
|
||||
score.last_scheduled_at.is_none_or(|last_scheduled_at| {
|
||||
scheduled_at.saturating_sub(last_scheduled_at)
|
||||
>= POOL_SCORE_SCHEDULE_INTEREST_MIN_INTERVAL_SECS
|
||||
})
|
||||
.await;
|
||||
if result.is_err() {
|
||||
failed += 1;
|
||||
}
|
||||
})
|
||||
.take(POOL_SCORE_SCHEDULE_INTEREST_MAX_PER_BATCH)
|
||||
.map(|score| PoolMemberScheduleFeedback {
|
||||
identity: PoolMemberIdentity {
|
||||
pool_kind: score.pool_kind.clone(),
|
||||
pool_id: score.pool_id.clone(),
|
||||
member_kind: score.member_kind.clone(),
|
||||
member_id: score.member_id.clone(),
|
||||
},
|
||||
scope: Some(PoolScoreScope {
|
||||
capability: score.capability.clone(),
|
||||
scope_kind: score.scope_kind.clone(),
|
||||
scope_id: score.scope_id.clone(),
|
||||
}),
|
||||
scheduled_at,
|
||||
succeeded: None,
|
||||
hard_state: None,
|
||||
score_delta: None,
|
||||
score_reason_patch: Some(serde_json::json!({
|
||||
"last_schedule_interest": {
|
||||
"provider_id": provider_id.as_str(),
|
||||
"endpoint_id": endpoint_id.as_str(),
|
||||
"model_id": model_id.as_str()
|
||||
}
|
||||
})),
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let score_count = feedback.len();
|
||||
if feedback.is_empty() {
|
||||
return;
|
||||
}
|
||||
if failed > 0 {
|
||||
warn!(
|
||||
event_name = "pool_group_score_interest_update_failed",
|
||||
|
||||
let Ok(permit) = POOL_SCORE_SCHEDULE_INTEREST_SEMAPHORE
|
||||
.clone()
|
||||
.try_acquire_owned()
|
||||
else {
|
||||
debug!(
|
||||
event_name = "pool_group_score_interest_dropped",
|
||||
log_type = "event",
|
||||
provider_id = %self.group.candidate.provider_id,
|
||||
endpoint_id = %self.group.candidate.endpoint_id,
|
||||
model_id = %self.group.candidate.model_id,
|
||||
failed_count = failed,
|
||||
score_count = scores.len(),
|
||||
"gateway pool scheduler failed to record some pool score schedule interests"
|
||||
"gateway pool scheduler dropped score schedule interest because the background writer is saturated"
|
||||
);
|
||||
}
|
||||
return;
|
||||
};
|
||||
|
||||
let app = self.state.app().clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let _permit = permit;
|
||||
let mut failed = 0usize;
|
||||
for feedback in feedback {
|
||||
let result = app
|
||||
.data
|
||||
.record_pool_member_schedule_feedback(feedback)
|
||||
.await;
|
||||
if result.is_err() {
|
||||
failed += 1;
|
||||
}
|
||||
}
|
||||
if failed > 0 {
|
||||
warn!(
|
||||
event_name = "pool_group_score_interest_update_failed",
|
||||
log_type = "event",
|
||||
provider_id = %provider_id,
|
||||
endpoint_id = %endpoint_id,
|
||||
model_id = %model_id,
|
||||
failed_count = failed,
|
||||
score_count,
|
||||
"gateway pool scheduler failed to record some pool score schedule interests"
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
async fn build_page_eligible_candidates(
|
||||
@@ -964,9 +1033,20 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
for skipped_candidate in skipped_candidates {
|
||||
self.record_skip_reason(skipped_candidate.skip_reason);
|
||||
}
|
||||
let prefiltered_count = skipped_candidates
|
||||
.iter()
|
||||
.filter(|candidate| pool_skip_reason_releases_scan_budget(candidate.skip_reason))
|
||||
.count();
|
||||
self.budget_scanned_keys = self
|
||||
.budget_scanned_keys
|
||||
.saturating_sub(u32::try_from(prefiltered_count).unwrap_or(u32::MAX));
|
||||
}
|
||||
}
|
||||
|
||||
fn pool_skip_reason_releases_scan_budget(skip_reason: &str) -> bool {
|
||||
skip_reason == POOL_ACCOUNT_EXHAUSTED_SKIP_REASON
|
||||
}
|
||||
|
||||
fn pool_candidate_transport_policy_facts(
|
||||
candidate: &aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate,
|
||||
) -> CandidateTransportPolicyFacts<'_> {
|
||||
@@ -1098,12 +1178,54 @@ fn build_pool_catalog_key_context(
|
||||
let mut signals =
|
||||
provider_pool_service.member_signals(provider_type, key, auth_config.as_ref());
|
||||
signals.account_blocked |= admin_provider_pool_pure::admin_pool_key_is_known_banned(key);
|
||||
signals.account_blocked |=
|
||||
pool_key_requires_reauth_for_scheduling(key, current_unix_ms().saturating_div(1000));
|
||||
signals.health_score = health_score;
|
||||
signals.latency_avg_ms = latency_avg_ms;
|
||||
signals.catalog_lru_score = Some(key.last_used_at_unix_secs.unwrap_or(0) as f64);
|
||||
signals
|
||||
}
|
||||
|
||||
fn pool_key_requires_reauth_for_scheduling(
|
||||
key: &StoredProviderCatalogKey,
|
||||
now_unix_secs: u64,
|
||||
) -> bool {
|
||||
if !key.auth_type.trim().eq_ignore_ascii_case("oauth") {
|
||||
return false;
|
||||
}
|
||||
|
||||
let invalid_reason = key
|
||||
.oauth_invalid_reason
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.unwrap_or_default();
|
||||
if !invalid_reason.is_empty() {
|
||||
if pool_oauth_reason_has_tag(invalid_reason, "[OAUTH_EXPIRED]")
|
||||
|| pool_oauth_reason_has_tag(invalid_reason, "[ACCOUNT_BLOCK]")
|
||||
{
|
||||
return true;
|
||||
}
|
||||
if pool_oauth_reason_has_tag(invalid_reason, "[REQUEST_FAILED]") {
|
||||
return false;
|
||||
}
|
||||
if pool_oauth_reason_has_tag(invalid_reason, "[REFRESH_FAILED]") {
|
||||
return key
|
||||
.expires_at_unix_secs
|
||||
.is_none_or(|expires_at| expires_at == 0 || expires_at <= now_unix_secs);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
key.oauth_invalid_at_unix_secs.is_some()
|
||||
}
|
||||
|
||||
fn pool_oauth_reason_has_tag(reason: &str, tag: &str) -> bool {
|
||||
reason
|
||||
.lines()
|
||||
.map(str::trim)
|
||||
.any(|line| line.starts_with(tag))
|
||||
}
|
||||
|
||||
fn apply_local_execution_pool_scheduler_with_runtime_map(
|
||||
candidates: Vec<EligibleLocalExecutionCandidate>,
|
||||
runtime_by_provider: &BTreeMap<String, AdminProviderPoolRuntimeState>,
|
||||
@@ -1454,10 +1576,11 @@ fn apply_pool_orchestration(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
admin_provider_pool_quota_probe_active_members_key,
|
||||
admin_provider_pool_quota_probe_active_members_key, apply_local_execution_pool_scheduler,
|
||||
apply_local_execution_pool_scheduler_with_runtime_map,
|
||||
apply_local_execution_pool_scheduler_with_runtime_map_outcome,
|
||||
build_pool_catalog_key_context, pool_config_for_candidate,
|
||||
pool_key_requires_reauth_for_scheduling,
|
||||
prune_unschedulable_active_probe_members_for_request,
|
||||
remove_active_probe_members_for_request, should_trigger_active_probe_burst_for_request,
|
||||
PoolCatalogKeyContext, PoolKeyCursor, POOL_ACTIVE_PROBE_SEALED_SKIP_REASON,
|
||||
@@ -2878,6 +3001,324 @@ mod tests {
|
||||
}));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_key_cursor_does_not_spend_effective_scan_budget_on_exhausted_accounts() {
|
||||
let provider_config = Some(json!({
|
||||
"pool_advanced": {
|
||||
"skip_exhausted_accounts": true
|
||||
}
|
||||
}));
|
||||
let (provider, endpoint, mut keys, rows) = large_pool_fixture(700, provider_config.clone());
|
||||
for key in keys.iter_mut().take(600) {
|
||||
key.status_snapshot = Some(json!({
|
||||
"quota": {
|
||||
"provider_type": "openai",
|
||||
"exhausted": true,
|
||||
"usage_ratio": 1.0,
|
||||
"windows": [
|
||||
{
|
||||
"code": "daily",
|
||||
"used_ratio": 1.0,
|
||||
"remaining_ratio": 0.0
|
||||
}
|
||||
]
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
let data_state =
|
||||
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
|
||||
Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![endpoint],
|
||||
keys,
|
||||
)),
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)),
|
||||
)
|
||||
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY);
|
||||
let app = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let group = sample_eligible_candidate(
|
||||
"provider-pool",
|
||||
"endpoint-1",
|
||||
"pool-group",
|
||||
10,
|
||||
provider_config,
|
||||
);
|
||||
|
||||
let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None);
|
||||
assert_eq!(
|
||||
cursor.max_scanned_keys,
|
||||
aether_dispatch_core::DEFAULT_POOL_MAX_SCAN
|
||||
);
|
||||
assert!(
|
||||
cursor.absolute_max_scanned_keys > cursor.max_scanned_keys,
|
||||
"pool config scan limit should be retained as the absolute cap"
|
||||
);
|
||||
|
||||
let candidate = cursor
|
||||
.next_key()
|
||||
.await
|
||||
.expect("cursor should scan past exhausted accounts within the absolute cap");
|
||||
let key_index = candidate
|
||||
.candidate
|
||||
.key_id
|
||||
.strip_prefix("key-")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.expect("fixture key id should contain a numeric suffix");
|
||||
assert!(
|
||||
key_index >= 600,
|
||||
"cursor should not return one of the exhausted leading keys"
|
||||
);
|
||||
assert_eq!(candidate.orchestration.pool_key_index, Some(0));
|
||||
assert_eq!(cursor.scanned_keys, 640);
|
||||
assert_eq!(cursor.budget_scanned_keys, 40);
|
||||
assert_eq!(
|
||||
cursor
|
||||
.skip_reason_counts
|
||||
.get(aether_pool_core::POOL_ACCOUNT_EXHAUSTED_SKIP_REASON),
|
||||
Some(&600)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_scheduler_skips_invalid_and_exhausted_high_priority_hot_pool_before_fallback_provider(
|
||||
) {
|
||||
let provider_config = Some(json!({
|
||||
"pool_advanced": {
|
||||
"probing_enabled": true,
|
||||
"skip_exhausted_accounts": true,
|
||||
"scheduling_presets": [
|
||||
{"preset": "single_account", "enabled": true}
|
||||
]
|
||||
}
|
||||
}));
|
||||
let provider_a = sample_codex_pool_provider("provider-a", 0, provider_config.clone());
|
||||
let provider_b = sample_codex_pool_provider("provider-b", 10, provider_config.clone());
|
||||
let endpoint_a = sample_codex_pool_endpoint("provider-a", "endpoint-a");
|
||||
let endpoint_b = sample_codex_pool_endpoint("provider-b", "endpoint-b");
|
||||
|
||||
let mut key_a_invalid = sample_codex_pool_key("provider-a", "key-a-invalid");
|
||||
key_a_invalid.oauth_invalid_at_unix_secs = Some(1_710_000_000);
|
||||
key_a_invalid.oauth_invalid_reason =
|
||||
Some("[OAUTH_EXPIRED] Codex Token 无效或已过期 (401)".to_string());
|
||||
let exhausted_status_snapshot = json!({
|
||||
"quota": {
|
||||
"provider_type": "codex",
|
||||
"exhausted": true,
|
||||
"usage_ratio": 1.0,
|
||||
"windows": [
|
||||
{
|
||||
"code": "daily",
|
||||
"used_ratio": 1.0,
|
||||
"remaining_ratio": 0.0
|
||||
}
|
||||
]
|
||||
}
|
||||
});
|
||||
key_a_invalid.status_snapshot = Some(exhausted_status_snapshot.clone());
|
||||
let mut key_a_exhausted = sample_codex_pool_key("provider-a", "key-a-exhausted");
|
||||
key_a_exhausted.status_snapshot = Some(exhausted_status_snapshot);
|
||||
let key_b_ready = sample_codex_pool_key("provider-b", "key-b-ready");
|
||||
|
||||
let rows = vec![
|
||||
sample_codex_pool_row("provider-a", "endpoint-a", "key-a-invalid", 0),
|
||||
sample_codex_pool_row("provider-a", "endpoint-a", "key-a-exhausted", 0),
|
||||
sample_codex_pool_row("provider-b", "endpoint-b", "key-b-ready", 10),
|
||||
];
|
||||
|
||||
let data_state =
|
||||
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
|
||||
Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider_a, provider_b],
|
||||
vec![endpoint_a, endpoint_b],
|
||||
vec![key_a_invalid, key_a_exhausted, key_b_ready],
|
||||
)),
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)),
|
||||
)
|
||||
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY);
|
||||
let app = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
app.runtime_state
|
||||
.set_add(
|
||||
&admin_provider_pool_quota_probe_active_members_key("provider-a"),
|
||||
"key-a-invalid",
|
||||
)
|
||||
.await
|
||||
.expect("provider-a hot member should insert");
|
||||
app.runtime_state
|
||||
.set_add(
|
||||
&admin_provider_pool_quota_probe_active_members_key("provider-b"),
|
||||
"key-b-ready",
|
||||
)
|
||||
.await
|
||||
.expect("provider-b hot member should insert");
|
||||
|
||||
let group_a =
|
||||
sample_codex_pool_group("provider-a", "endpoint-a", 0, provider_config.clone());
|
||||
let group_b = sample_codex_pool_group("provider-b", "endpoint-b", 10, provider_config);
|
||||
|
||||
let (scheduled, skipped) = apply_local_execution_pool_scheduler(
|
||||
PlannerAppState::new(&app),
|
||||
vec![group_a, group_b],
|
||||
None,
|
||||
Some("gpt-5"),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
scheduled
|
||||
.iter()
|
||||
.map(|item| item.candidate.key_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["key-b-ready"]
|
||||
);
|
||||
let skipped_pairs = skipped
|
||||
.iter()
|
||||
.map(|item| (item.candidate.key_id.as_str(), item.skip_reason))
|
||||
.collect::<Vec<_>>();
|
||||
assert!(skipped_pairs.contains(&("key-a-invalid", "pool_account_blocked")));
|
||||
assert!(skipped_pairs.contains(&(
|
||||
"key-a-exhausted",
|
||||
aether_pool_core::POOL_ACCOUNT_EXHAUSTED_SKIP_REASON
|
||||
)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_scheduler_skips_invalid_high_priority_hot_pool_account_even_with_remaining_quota()
|
||||
{
|
||||
let provider_config = Some(json!({
|
||||
"pool_advanced": {
|
||||
"probing_enabled": true,
|
||||
"skip_exhausted_accounts": true,
|
||||
"scheduling_presets": [
|
||||
{"preset": "single_account", "enabled": true}
|
||||
]
|
||||
}
|
||||
}));
|
||||
let provider_a = sample_codex_pool_provider("provider-a", 0, provider_config.clone());
|
||||
let provider_b = sample_codex_pool_provider("provider-b", 10, provider_config.clone());
|
||||
let endpoint_a = sample_codex_pool_endpoint("provider-a", "endpoint-a");
|
||||
let endpoint_b = sample_codex_pool_endpoint("provider-b", "endpoint-b");
|
||||
|
||||
let mut key_a_invalid = sample_codex_pool_key("provider-a", "key-a-invalid");
|
||||
key_a_invalid.oauth_invalid_at_unix_secs = Some(1_710_000_000);
|
||||
key_a_invalid.oauth_invalid_reason =
|
||||
Some("[OAUTH_EXPIRED] Codex Token 无效或已过期 (401)".to_string());
|
||||
key_a_invalid.status_snapshot = Some(json!({
|
||||
"quota": {
|
||||
"provider_type": "codex",
|
||||
"exhausted": false,
|
||||
"usage_ratio": 0.25,
|
||||
"windows": [
|
||||
{
|
||||
"code": "daily",
|
||||
"used_ratio": 0.25,
|
||||
"remaining_ratio": 0.75
|
||||
}
|
||||
]
|
||||
}
|
||||
}));
|
||||
let key_b_ready = sample_codex_pool_key("provider-b", "key-b-ready");
|
||||
|
||||
let rows = vec![
|
||||
sample_codex_pool_row("provider-a", "endpoint-a", "key-a-invalid", 0),
|
||||
sample_codex_pool_row("provider-b", "endpoint-b", "key-b-ready", 10),
|
||||
];
|
||||
|
||||
let data_state =
|
||||
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
|
||||
Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider_a, provider_b],
|
||||
vec![endpoint_a, endpoint_b],
|
||||
vec![key_a_invalid, key_b_ready],
|
||||
)),
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)),
|
||||
)
|
||||
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY);
|
||||
let app = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
app.runtime_state
|
||||
.set_add(
|
||||
&admin_provider_pool_quota_probe_active_members_key("provider-a"),
|
||||
"key-a-invalid",
|
||||
)
|
||||
.await
|
||||
.expect("provider-a hot member should insert");
|
||||
app.runtime_state
|
||||
.set_add(
|
||||
&admin_provider_pool_quota_probe_active_members_key("provider-b"),
|
||||
"key-b-ready",
|
||||
)
|
||||
.await
|
||||
.expect("provider-b hot member should insert");
|
||||
|
||||
let group_a =
|
||||
sample_codex_pool_group("provider-a", "endpoint-a", 0, provider_config.clone());
|
||||
let group_b = sample_codex_pool_group("provider-b", "endpoint-b", 10, provider_config);
|
||||
|
||||
let (scheduled, skipped) = apply_local_execution_pool_scheduler(
|
||||
PlannerAppState::new(&app),
|
||||
vec![group_a, group_b],
|
||||
None,
|
||||
Some("gpt-5"),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
scheduled
|
||||
.iter()
|
||||
.map(|item| item.candidate.key_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["key-b-ready"]
|
||||
);
|
||||
let skipped_pairs = skipped
|
||||
.iter()
|
||||
.map(|item| (item.candidate.key_id.as_str(), item.skip_reason))
|
||||
.collect::<Vec<_>>();
|
||||
assert!(skipped_pairs.contains(&("key-a-invalid", "pool_account_blocked")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pool_key_reauth_scheduling_keeps_recoverable_oauth_markers_usable() {
|
||||
let mut key = sample_codex_pool_key("provider-a", "key-refresh-failed");
|
||||
key.expires_at_unix_secs = Some(200);
|
||||
key.oauth_invalid_reason = Some(
|
||||
"[REFRESH_FAILED] Token 续期失败 (401): refresh_token 已被使用并轮换,请重新登录授权"
|
||||
.to_string(),
|
||||
);
|
||||
|
||||
assert!(!pool_key_requires_reauth_for_scheduling(&key, 100));
|
||||
assert!(pool_key_requires_reauth_for_scheduling(&key, 200));
|
||||
|
||||
key.oauth_invalid_reason = Some("[REQUEST_FAILED] 账号状态检查失败".to_string());
|
||||
key.oauth_invalid_at_unix_secs = Some(100);
|
||||
assert!(!pool_key_requires_reauth_for_scheduling(&key, 300));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pool_key_reauth_scheduling_blocks_invalid_oauth_markers_without_affecting_non_oauth_keys() {
|
||||
let mut key = sample_codex_pool_key("provider-a", "key-invalid");
|
||||
key.oauth_invalid_reason = Some("[ACCOUNT_BLOCK] account has been deactivated".to_string());
|
||||
assert!(pool_key_requires_reauth_for_scheduling(&key, 100));
|
||||
|
||||
key.oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string());
|
||||
key.oauth_invalid_at_unix_secs = None;
|
||||
assert!(pool_key_requires_reauth_for_scheduling(&key, 100));
|
||||
|
||||
key.oauth_invalid_reason = None;
|
||||
key.oauth_invalid_at_unix_secs = Some(100);
|
||||
assert!(pool_key_requires_reauth_for_scheduling(&key, 100));
|
||||
|
||||
key.auth_type = "api_key".to_string();
|
||||
assert!(!pool_key_requires_reauth_for_scheduling(&key, 100));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_key_cursor_simulates_large_lru_pool_with_lazy_pages_and_dynamic_skips() {
|
||||
const KEY_COUNT: usize = 2048;
|
||||
@@ -3325,6 +3766,210 @@ mod tests {
|
||||
(provider, endpoint, keys, rows)
|
||||
}
|
||||
|
||||
fn sample_codex_pool_provider(
|
||||
provider_id: &str,
|
||||
provider_priority: i32,
|
||||
provider_config: Option<serde_json::Value>,
|
||||
) -> StoredProviderCatalogProvider {
|
||||
StoredProviderCatalogProvider::new(
|
||||
provider_id.to_string(),
|
||||
provider_id.to_string(),
|
||||
Some("https://example.com".to_string()),
|
||||
"codex".to_string(),
|
||||
)
|
||||
.expect("provider should build")
|
||||
.with_routing_fields(provider_priority)
|
||||
.with_transport_fields(
|
||||
true,
|
||||
false,
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
provider_config,
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_codex_pool_endpoint(
|
||||
provider_id: &str,
|
||||
endpoint_id: &str,
|
||||
) -> StoredProviderCatalogEndpoint {
|
||||
StoredProviderCatalogEndpoint::new(
|
||||
endpoint_id.to_string(),
|
||||
provider_id.to_string(),
|
||||
"openai:responses".to_string(),
|
||||
Some("openai".to_string()),
|
||||
Some("responses".to_string()),
|
||||
true,
|
||||
)
|
||||
.expect("endpoint should build")
|
||||
.with_health_score(1.0)
|
||||
.with_transport_fields(
|
||||
"https://example.com/v1/responses".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("endpoint transport should build")
|
||||
}
|
||||
|
||||
fn sample_codex_pool_key(provider_id: &str, key_id: &str) -> StoredProviderCatalogKey {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
key_id.to_string(),
|
||||
provider_id.to_string(),
|
||||
key_id.to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
.with_transport_fields(
|
||||
Some(json!(["openai:responses"])),
|
||||
Some(format!("secret-{key_id}")),
|
||||
None,
|
||||
None,
|
||||
Some(json!({"openai:responses": 1})),
|
||||
None,
|
||||
Some(4_102_444_800),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("key transport should build");
|
||||
key.internal_priority = 10;
|
||||
key
|
||||
}
|
||||
|
||||
fn sample_codex_pool_row(
|
||||
provider_id: &str,
|
||||
endpoint_id: &str,
|
||||
key_id: &str,
|
||||
provider_priority: i32,
|
||||
) -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: provider_id.to_string(),
|
||||
provider_name: provider_id.to_string(),
|
||||
provider_type: "codex".to_string(),
|
||||
provider_priority,
|
||||
provider_is_active: true,
|
||||
endpoint_id: endpoint_id.to_string(),
|
||||
endpoint_api_format: "openai:responses".to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("responses".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: key_id.to_string(),
|
||||
key_name: key_id.to_string(),
|
||||
key_auth_type: "oauth".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec!["openai:responses".to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 10,
|
||||
key_global_priority_by_format: Some(json!({"openai:responses": 1})),
|
||||
model_id: "model-1".to_string(),
|
||||
global_model_id: "global-model-1".to_string(),
|
||||
global_model_name: "gpt-5".to_string(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: Some(true),
|
||||
model_provider_model_name: "gpt-5".to_string(),
|
||||
model_provider_model_mappings: None,
|
||||
model_supports_streaming: Some(true),
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_codex_pool_group(
|
||||
provider_id: &str,
|
||||
endpoint_id: &str,
|
||||
provider_priority: i32,
|
||||
provider_config: Option<serde_json::Value>,
|
||||
) -> EligibleLocalExecutionCandidate {
|
||||
EligibleLocalExecutionCandidate {
|
||||
kind: LocalExecutionCandidateKind::PoolGroup,
|
||||
candidate: SchedulerMinimalCandidateSelectionCandidate {
|
||||
provider_id: provider_id.to_string(),
|
||||
provider_name: provider_id.to_string(),
|
||||
provider_type: "codex".to_string(),
|
||||
provider_priority,
|
||||
endpoint_id: endpoint_id.to_string(),
|
||||
endpoint_api_format: "openai:responses".to_string(),
|
||||
key_id: format!("{provider_id}-pool-group"),
|
||||
key_name: format!("{provider_id}-pool-group"),
|
||||
key_auth_type: "oauth".to_string(),
|
||||
key_internal_priority: 10,
|
||||
key_global_priority_for_format: Some(1),
|
||||
key_capabilities: None,
|
||||
model_id: "model-1".to_string(),
|
||||
global_model_id: "global-model-1".to_string(),
|
||||
global_model_name: "gpt-5".to_string(),
|
||||
selected_provider_model_name: "gpt-5".to_string(),
|
||||
mapping_matched_model: None,
|
||||
},
|
||||
provider_api_format: "openai:responses".to_string(),
|
||||
orchestration: LocalExecutionCandidateMetadata::default(),
|
||||
ranking: None,
|
||||
transport: Arc::new(crate::ai_serving::GatewayProviderTransportSnapshot {
|
||||
provider: GatewayProviderTransportProvider {
|
||||
id: provider_id.to_string(),
|
||||
name: provider_id.to_string(),
|
||||
provider_type: "codex".to_string(),
|
||||
website: None,
|
||||
is_active: true,
|
||||
keep_priority_on_conversion: false,
|
||||
enable_format_conversion: false,
|
||||
concurrent_limit: None,
|
||||
max_retries: None,
|
||||
proxy: None,
|
||||
request_timeout_secs: None,
|
||||
stream_first_byte_timeout_secs: None,
|
||||
config: provider_config,
|
||||
},
|
||||
endpoint: GatewayProviderTransportEndpoint {
|
||||
id: endpoint_id.to_string(),
|
||||
provider_id: provider_id.to_string(),
|
||||
api_format: "openai:responses".to_string(),
|
||||
api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("responses".to_string()),
|
||||
is_active: true,
|
||||
base_url: "https://example.com/v1/responses".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: GatewayProviderTransportKey {
|
||||
id: format!("{provider_id}-pool-group"),
|
||||
provider_id: provider_id.to_string(),
|
||||
name: format!("{provider_id}-pool-group"),
|
||||
auth_type: "oauth".to_string(),
|
||||
is_active: true,
|
||||
api_formats: Some(vec!["openai:responses".to_string()]),
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
allowed_models: None,
|
||||
capabilities: None,
|
||||
rate_multipliers: None,
|
||||
global_priority_by_format: None,
|
||||
expires_at_unix_secs: None,
|
||||
proxy: None,
|
||||
fingerprint: None,
|
||||
upstream_metadata: None,
|
||||
decrypted_api_key: "secret".to_string(),
|
||||
decrypted_auth_config: None,
|
||||
},
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn routing_policy_with_allowed_keys<const N: usize>(
|
||||
key_ids: [&str; N],
|
||||
) -> ResolvedRoutingPolicy {
|
||||
|
||||
@@ -11,12 +11,42 @@ use crate::insert_header_if_missing;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) enum GatewayError {
|
||||
UpstreamUnavailable { trace_id: String, message: String },
|
||||
ControlUnavailable { trace_id: String, message: String },
|
||||
Client { status: StatusCode, message: String },
|
||||
UpstreamUnavailable {
|
||||
trace_id: String,
|
||||
message: String,
|
||||
},
|
||||
ControlUnavailable {
|
||||
trace_id: String,
|
||||
message: String,
|
||||
},
|
||||
LocalExecutionPlanningTimeout {
|
||||
trace_id: String,
|
||||
phase: &'static str,
|
||||
timeout_ms: u64,
|
||||
},
|
||||
Client {
|
||||
status: StatusCode,
|
||||
message: String,
|
||||
},
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
impl GatewayError {
|
||||
pub(crate) fn into_message(self) -> String {
|
||||
match self {
|
||||
Self::UpstreamUnavailable { message, .. }
|
||||
| Self::ControlUnavailable { message, .. }
|
||||
| Self::Client { message, .. }
|
||||
| Self::Internal(message) => message,
|
||||
Self::LocalExecutionPlanningTimeout {
|
||||
phase, timeout_ms, ..
|
||||
} => {
|
||||
format!("local execution planning timed out in {phase} after {timeout_ms}ms")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoResponse for GatewayError {
|
||||
fn into_response(self) -> Response<Body> {
|
||||
match self {
|
||||
@@ -56,6 +86,33 @@ impl IntoResponse for GatewayError {
|
||||
);
|
||||
response
|
||||
}
|
||||
Self::LocalExecutionPlanningTimeout {
|
||||
trace_id,
|
||||
phase,
|
||||
timeout_ms,
|
||||
} => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
phase,
|
||||
timeout_ms,
|
||||
"gateway local execution planning timed out"
|
||||
);
|
||||
let body = Json(json!({
|
||||
"error": {
|
||||
"message": "gateway local execution planning timed out",
|
||||
"trace_id": trace_id,
|
||||
}
|
||||
}));
|
||||
let mut response = (StatusCode::GATEWAY_TIMEOUT, body).into_response();
|
||||
let _ =
|
||||
insert_header_if_missing(response.headers_mut(), TRACE_ID_HEADER, &trace_id);
|
||||
let _ = insert_header_if_missing(
|
||||
response.headers_mut(),
|
||||
GATEWAY_HEADER,
|
||||
"rust-phase3b",
|
||||
);
|
||||
response
|
||||
}
|
||||
Self::Client { status, message } => (
|
||||
status,
|
||||
Json(json!({
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,5 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::future::Future;
|
||||
use std::io::Error as IoError;
|
||||
use std::net::IpAddr;
|
||||
use std::sync::OnceLock;
|
||||
@@ -29,7 +30,9 @@ use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
|
||||
use crate::execution_runtime::transport::{
|
||||
build_browser_wreq_client, build_request_body, build_request_headers,
|
||||
decode_response_body_bytes, format_upstream_request_error, format_wreq_upstream_request_error,
|
||||
send_request, DirectHttpResponse, ExecutionRuntimeTransportError, ExecutionTransportControls,
|
||||
resolve_stream_first_byte_timeout, send_request, stream_first_byte_timeout_message,
|
||||
with_non_stream_total_timeout, DirectHttpResponse, ExecutionRuntimeTransportError,
|
||||
ExecutionTransportControls,
|
||||
};
|
||||
|
||||
const GROK_INTERNAL_HEADER: &str = "x-aether-grok-runtime";
|
||||
@@ -148,9 +151,12 @@ pub(crate) async fn maybe_execute_grok_sync(
|
||||
if !is_grok_plan(plan, report_context) {
|
||||
return Ok(None);
|
||||
}
|
||||
let mut collected = execute_grok_app_chat(plan, report_context).await?;
|
||||
materialize_grok_image_assets(plan, &mut collected).await;
|
||||
Ok(Some(grok_execution_result(plan, collected, report_context)))
|
||||
with_non_stream_total_timeout(plan, async move {
|
||||
let mut collected = execute_grok_app_chat(plan, report_context).await?;
|
||||
materialize_grok_image_assets(plan, &mut collected).await;
|
||||
Ok(Some(grok_execution_result(plan, collected, report_context)))
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_execute_grok_stream(
|
||||
@@ -387,6 +393,7 @@ async fn grok_imagine_websocket_images(
|
||||
plan.proxy.as_ref(),
|
||||
profile,
|
||||
ExecutionTransportControls::default(),
|
||||
true,
|
||||
)?;
|
||||
let response = client
|
||||
.websocket(GROK_IMAGINE_WS_URL)
|
||||
@@ -561,6 +568,7 @@ fn grok_success_frame_stream(
|
||||
started_at: Instant,
|
||||
mut body_stream: GrokUpstreamBodyStream,
|
||||
) -> BoxStream<'static, Result<Bytes, IoError>> {
|
||||
let stream_first_byte_timeout = resolve_stream_first_byte_timeout(&plan);
|
||||
async_stream::stream! {
|
||||
match encode_grok_headers_frame(
|
||||
status_code,
|
||||
@@ -581,8 +589,36 @@ fn grok_success_frame_stream(
|
||||
let mut text_len = 0usize;
|
||||
let mut thinking_len = 0usize;
|
||||
let mut image_len = 0usize;
|
||||
let mut terminal_error_emitted = false;
|
||||
|
||||
while let Some(item) = body_stream.next().await {
|
||||
loop {
|
||||
let item = if ttfb_ms.is_none() {
|
||||
match await_grok_stream_first_byte(
|
||||
body_stream.next(),
|
||||
started_at,
|
||||
stream_first_byte_timeout,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(item) => item,
|
||||
Err(timeout) => {
|
||||
match encode_grok_first_byte_timeout_frame(timeout) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
terminal_error_emitted = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
body_stream.next().await
|
||||
};
|
||||
let Some(item) = item else {
|
||||
break;
|
||||
};
|
||||
let chunk = match item {
|
||||
Ok(chunk) => chunk,
|
||||
Err(message) => {
|
||||
@@ -593,6 +629,7 @@ fn grok_success_frame_stream(
|
||||
return;
|
||||
}
|
||||
}
|
||||
terminal_error_emitted = true;
|
||||
break;
|
||||
}
|
||||
};
|
||||
@@ -630,6 +667,25 @@ fn grok_success_frame_stream(
|
||||
}
|
||||
}
|
||||
|
||||
if terminal_error_emitted {
|
||||
match encode_grok_telemetry_frame(
|
||||
ttfb_ms,
|
||||
Some(started_at.elapsed().as_millis() as u64),
|
||||
upstream_bytes,
|
||||
) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
match encode_stream_frame_ndjson(&StreamFrame::eof_with_summary(None)) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => yield Err(err),
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
adapter.finish();
|
||||
match emit_grok_adapter_deltas(
|
||||
&mut client_emitter,
|
||||
@@ -691,6 +747,28 @@ fn grok_success_frame_stream(
|
||||
.boxed()
|
||||
}
|
||||
|
||||
async fn await_grok_stream_first_byte<T, F>(
|
||||
future: F,
|
||||
started_at: Instant,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<T, Duration>
|
||||
where
|
||||
F: Future<Output = T>,
|
||||
{
|
||||
let Some(timeout) = timeout else {
|
||||
return Ok(future.await);
|
||||
};
|
||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||
return Err(timeout);
|
||||
};
|
||||
if remaining.is_zero() {
|
||||
return Err(timeout);
|
||||
}
|
||||
tokio::time::timeout(remaining, future)
|
||||
.await
|
||||
.map_err(|_| timeout)
|
||||
}
|
||||
|
||||
fn emit_grok_adapter_deltas(
|
||||
client_emitter: &mut GrokClientStreamEmitter,
|
||||
adapter: &GrokStreamAdapter,
|
||||
@@ -790,6 +868,22 @@ fn encode_grok_error_frame(status_code: u16, message: String) -> Result<Bytes, I
|
||||
})
|
||||
}
|
||||
|
||||
fn encode_grok_first_byte_timeout_frame(timeout: Duration) -> Result<Bytes, IoError> {
|
||||
encode_stream_frame_ndjson(&StreamFrame {
|
||||
frame_type: StreamFrameType::Error,
|
||||
payload: StreamFramePayload::Error {
|
||||
error: aether_contracts::ExecutionError {
|
||||
kind: aether_contracts::ExecutionErrorKind::FirstByteTimeout,
|
||||
phase: aether_contracts::ExecutionPhase::FirstByte,
|
||||
message: stream_first_byte_timeout_message(timeout),
|
||||
upstream_status: Some(504),
|
||||
retryable: true,
|
||||
failover_recommended: true,
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
enum GrokClientStreamEmitter {
|
||||
OpenAiChat {
|
||||
id: String,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,6 +1,7 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::future::Future;
|
||||
use std::io::Error as IoError;
|
||||
use std::time::Instant;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use aether_contracts::{
|
||||
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionStreamTerminalSummary,
|
||||
@@ -19,7 +20,7 @@ use crate::ai_serving::api::{
|
||||
};
|
||||
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
|
||||
use crate::execution_runtime::transport::{
|
||||
format_wreq_upstream_request_error, DirectUpstreamResponse,
|
||||
format_wreq_upstream_request_error, stream_first_byte_timeout_message, DirectUpstreamResponse,
|
||||
};
|
||||
use crate::execution_runtime::DirectUpstreamStreamExecution;
|
||||
use crate::GatewayError;
|
||||
@@ -37,6 +38,7 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
stream_summary_report_context,
|
||||
response,
|
||||
started_at,
|
||||
stream_first_byte_timeout,
|
||||
} = execution;
|
||||
|
||||
let mut observer_context = stream_summary_report_context;
|
||||
@@ -64,7 +66,7 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
|
||||
if should_buffer_non_stream_response(&headers, &observer_context) {
|
||||
let original_headers = headers.clone();
|
||||
match buffer_non_sse_upstream_body(response, started_at).await {
|
||||
match buffer_non_sse_upstream_body(response, started_at, stream_first_byte_timeout).await {
|
||||
Ok(buffered) => {
|
||||
let mut response_headers = original_headers;
|
||||
let mut response_body = Bytes::from(buffered.body_bytes);
|
||||
@@ -134,6 +136,7 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
message,
|
||||
ttfb_ms,
|
||||
upstream_bytes,
|
||||
first_byte_timeout,
|
||||
}) => {
|
||||
match encode_headers_frame(status_code, original_headers) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
@@ -142,7 +145,12 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
return;
|
||||
}
|
||||
}
|
||||
match encode_error_frame(status_code, message) {
|
||||
let error_frame = if let Some(timeout) = first_byte_timeout {
|
||||
encode_first_byte_timeout_frame(timeout)
|
||||
} else {
|
||||
encode_error_frame(status_code, message)
|
||||
};
|
||||
match error_frame {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
@@ -183,7 +191,33 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
match response {
|
||||
DirectUpstreamResponse::Reqwest(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
while let Some(item) = bytes_stream.next().await {
|
||||
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() {
|
||||
@@ -239,7 +273,33 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
}
|
||||
DirectUpstreamResponse::BrowserWreq(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
while let Some(item) = bytes_stream.next().await {
|
||||
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() {
|
||||
@@ -294,7 +354,30 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
}
|
||||
}
|
||||
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
|
||||
match response.next_chunk().await {
|
||||
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);
|
||||
@@ -428,6 +511,44 @@ fn encode_error_frame(status_code: u16, message: String) -> Result<Bytes, IoErro
|
||||
})
|
||||
}
|
||||
|
||||
fn encode_first_byte_timeout_frame(timeout: Duration) -> Result<Bytes, IoError> {
|
||||
encode_stream_frame_ndjson(&StreamFrame {
|
||||
frame_type: StreamFrameType::Error,
|
||||
payload: StreamFramePayload::Error {
|
||||
error: ExecutionError {
|
||||
kind: ExecutionErrorKind::FirstByteTimeout,
|
||||
phase: ExecutionPhase::FirstByte,
|
||||
message: stream_first_byte_timeout_message(timeout),
|
||||
upstream_status: Some(504),
|
||||
retryable: true,
|
||||
failover_recommended: true,
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
async fn await_stream_first_byte<T, F>(
|
||||
future: F,
|
||||
started_at: Instant,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<T, Duration>
|
||||
where
|
||||
F: Future<Output = T>,
|
||||
{
|
||||
let Some(timeout) = timeout else {
|
||||
return Ok(future.await);
|
||||
};
|
||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||
return Err(timeout);
|
||||
};
|
||||
if remaining.is_zero() {
|
||||
return Err(timeout);
|
||||
}
|
||||
tokio::time::timeout(remaining, future)
|
||||
.await
|
||||
.map_err(|_| timeout)
|
||||
}
|
||||
|
||||
struct BufferedUpstreamBody {
|
||||
body_bytes: Vec<u8>,
|
||||
ttfb_ms: Option<u64>,
|
||||
@@ -438,6 +559,7 @@ struct BufferedUpstreamBodyError {
|
||||
message: String,
|
||||
ttfb_ms: Option<u64>,
|
||||
upstream_bytes: u64,
|
||||
first_byte_timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
fn response_headers_indicate_sse(headers: &BTreeMap<String, String>) -> bool {
|
||||
@@ -477,6 +599,7 @@ fn should_buffer_non_stream_response(
|
||||
async fn buffer_non_sse_upstream_body(
|
||||
response: DirectUpstreamResponse,
|
||||
started_at: Instant,
|
||||
stream_first_byte_timeout: Option<Duration>,
|
||||
) -> Result<BufferedUpstreamBody, BufferedUpstreamBodyError> {
|
||||
let mut body_bytes = Vec::new();
|
||||
let mut upstream_bytes = 0u64;
|
||||
@@ -485,7 +608,31 @@ async fn buffer_non_sse_upstream_body(
|
||||
match response {
|
||||
DirectUpstreamResponse::Reqwest(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
while let Some(item) = bytes_stream.next().await {
|
||||
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() {
|
||||
@@ -507,6 +654,7 @@ async fn buffer_non_sse_upstream_body(
|
||||
message,
|
||||
ttfb_ms,
|
||||
upstream_bytes,
|
||||
first_byte_timeout: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -514,7 +662,31 @@ async fn buffer_non_sse_upstream_body(
|
||||
}
|
||||
DirectUpstreamResponse::BrowserWreq(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
while let Some(item) = bytes_stream.next().await {
|
||||
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() {
|
||||
@@ -536,13 +708,35 @@ async fn buffer_non_sse_upstream_body(
|
||||
message,
|
||||
ttfb_ms,
|
||||
upstream_bytes,
|
||||
first_byte_timeout: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
|
||||
match response.next_chunk().await {
|
||||
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);
|
||||
@@ -563,6 +757,7 @@ async fn buffer_non_sse_upstream_body(
|
||||
message,
|
||||
ttfb_ms,
|
||||
upstream_bytes,
|
||||
first_byte_timeout: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -761,6 +956,7 @@ mod tests {
|
||||
use base64::Engine as _;
|
||||
use futures_util::StreamExt;
|
||||
use serde_json::Value;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::sync::watch;
|
||||
|
||||
use super::{
|
||||
@@ -903,6 +1099,93 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_execution_frame_stream_applies_first_byte_timeout_after_headers() {
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("local addr should resolve");
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("client should connect");
|
||||
let mut request = [0_u8; 1024];
|
||||
let _ = socket
|
||||
.read(&mut request)
|
||||
.await
|
||||
.expect("request should read");
|
||||
socket
|
||||
.write_all(
|
||||
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n",
|
||||
)
|
||||
.await
|
||||
.expect("headers should write");
|
||||
socket.flush().await.expect("headers should flush");
|
||||
tokio::time::sleep(Duration::from_millis(200)).await;
|
||||
let _ = socket.write_all(b"d\r\ndata: hello\n\n\r\n0\r\n\r\n").await;
|
||||
});
|
||||
|
||||
let execution = DirectSyncExecutionRuntime::new()
|
||||
.execute_stream(&ExecutionPlan {
|
||||
request_id: "req-stream-first-byte-timeout".into(),
|
||||
candidate_id: Some("cand-stream-first-byte-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: "openai:chat".into(),
|
||||
model_name: Some("gpt-5".into()),
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
first_byte_ms: Some(50),
|
||||
total_ms: Some(5_000),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.expect("stream execution should receive response headers");
|
||||
|
||||
let frames = build_direct_execution_frame_stream(execution)
|
||||
.map(|item| item.expect("frame should encode"))
|
||||
.collect::<Vec<_>>()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|bytes| String::from_utf8(bytes.to_vec()).expect("frame should be utf8"))
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
server.abort();
|
||||
|
||||
let error_frame = frames
|
||||
.iter()
|
||||
.map(|line| serde_json::from_str::<Value>(line).expect("frame should parse"))
|
||||
.find(|frame| frame.get("type").and_then(Value::as_str) == Some("error"))
|
||||
.expect("timeout should emit an error frame");
|
||||
|
||||
assert_eq!(
|
||||
error_frame
|
||||
.get("payload")
|
||||
.and_then(|payload| payload.get("error"))
|
||||
.and_then(|error| error.get("kind"))
|
||||
.and_then(Value::as_str),
|
||||
Some("first_byte_timeout")
|
||||
);
|
||||
assert!(error_frame
|
||||
.get("payload")
|
||||
.and_then(|payload| payload.get("error"))
|
||||
.and_then(|error| error.get("message"))
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(
|
||||
|message| message.contains("provider stream first byte timeout after 50 ms")
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_execution_frame_stream_emits_telemetry_before_first_data_frame() {
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,6 +1,6 @@
|
||||
use std::collections::{BTreeMap, HashMap, HashSet};
|
||||
use std::fs;
|
||||
use std::io::Error as IoError;
|
||||
use std::io::{Error as IoError, Read, Seek, SeekFrom};
|
||||
use std::net::{SocketAddr, TcpListener, TcpStream};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::process::{Child, Command, Stdio};
|
||||
@@ -18,9 +18,9 @@ use aether_provider_transport::windsurf::cascade::{
|
||||
build_get_trajectory_steps_request, build_get_user_status_request, build_heartbeat_request,
|
||||
build_initialize_panel_state_request, build_send_cascade_message_request_with_options,
|
||||
build_start_cascade_request, build_update_panel_state_with_user_status_request,
|
||||
build_update_workspace_trust_request, extract_grpc_frames, extract_user_status_bytes,
|
||||
grpc_frame, parse_generator_metadata, parse_start_cascade_response, parse_trajectory_status,
|
||||
parse_trajectory_steps, CascadeImage, CascadeUsage, SendCascadeMessageOptions,
|
||||
build_update_workspace_trust_request, extract_user_status_bytes, parse_generator_metadata,
|
||||
parse_start_cascade_response, parse_trajectory_status, parse_trajectory_steps, CascadeImage,
|
||||
CascadeUsage, SendCascadeMessageOptions,
|
||||
};
|
||||
use aether_provider_transport::windsurf::models::resolve_windsurf_model;
|
||||
use aether_provider_transport::windsurf::{GET_CHAT_MESSAGE_PATH, WINDSURF_ENVELOPE_NAME};
|
||||
@@ -32,11 +32,11 @@ use regex::Regex;
|
||||
use serde_json::{json, Value};
|
||||
use sha2::{Digest, Sha256};
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{debug, error, info, warn};
|
||||
use tracing::{debug, info, warn};
|
||||
use uuid::Uuid;
|
||||
|
||||
use super::ndjson::encode_stream_frame_ndjson;
|
||||
use super::transport::ExecutionRuntimeTransportError;
|
||||
use super::transport::{with_non_stream_total_timeout, ExecutionRuntimeTransportError};
|
||||
use crate::AppState;
|
||||
|
||||
const LS_SERVICE: &str = "/exa.language_server_pb.LanguageServerService";
|
||||
@@ -51,6 +51,7 @@ const CASCADE_TEXT_STALL: Duration = Duration::from_secs(45);
|
||||
const CASCADE_THINKING_STALL: Duration = Duration::from_secs(120);
|
||||
const SSE_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(15);
|
||||
const LS_READY_TIMEOUT: Duration = Duration::from_secs(25);
|
||||
const WINDOWS_LS_READY_TIMEOUT: Duration = Duration::from_secs(90);
|
||||
const GRPC_SHORT_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
const GRPC_STATUS_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
const GRPC_REQUEST_TIMEOUT: Duration = Duration::from_secs(45);
|
||||
@@ -215,66 +216,69 @@ pub(crate) async fn maybe_execute_windsurf_sync(
|
||||
let Some(input) = detect_windsurf_request(plan, report_context) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let key_upstream_metadata = read_windsurf_key_upstream_metadata(state, plan).await;
|
||||
let prepared = prepare_windsurf_cascade(plan, input, key_upstream_metadata).await?;
|
||||
let started_at = Instant::now();
|
||||
let mut deltas = Vec::new();
|
||||
let poll_result = poll_windsurf_cascade_with_transport_recovery(&prepared, |event| {
|
||||
if let WindsurfPollEvent::TextDelta(delta) = event {
|
||||
deltas.push(sanitize_windsurf_text(&delta));
|
||||
with_non_stream_total_timeout(plan, async move {
|
||||
let key_upstream_metadata = read_windsurf_key_upstream_metadata(state, plan).await;
|
||||
let prepared = prepare_windsurf_cascade(plan, input, key_upstream_metadata).await?;
|
||||
let started_at = Instant::now();
|
||||
let mut deltas = Vec::new();
|
||||
let poll_result = poll_windsurf_cascade_with_transport_recovery(&prepared, |event| {
|
||||
if let WindsurfPollEvent::TextDelta(delta) = event {
|
||||
deltas.push(sanitize_windsurf_text(&delta));
|
||||
}
|
||||
Ok(())
|
||||
})
|
||||
.await?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
let content = deltas.concat();
|
||||
let parsed_tool_calls = parse_and_filter_windsurf_tool_calls(&content, &prepared.input);
|
||||
let mut tool_calls = poll_result.native_tool_calls;
|
||||
tool_calls.extend(parsed_tool_calls.tool_calls);
|
||||
let has_tool_calls = !tool_calls.is_empty();
|
||||
let message = if has_tool_calls {
|
||||
json!({
|
||||
"role": "assistant",
|
||||
"content": Value::Null,
|
||||
"tool_calls": openai_tool_call_values(&tool_calls),
|
||||
})
|
||||
} else {
|
||||
json!({
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
})
|
||||
};
|
||||
let mut body_json = json!({
|
||||
"id": format!("chatcmpl-{}", prepared.request_id),
|
||||
"object": "chat.completion",
|
||||
"created": current_unix_secs(),
|
||||
"model": prepared.model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": message,
|
||||
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" },
|
||||
}],
|
||||
});
|
||||
if let Some(usage) = poll_result.usage {
|
||||
body_json["usage"] = windsurf_openai_usage_json(&usage);
|
||||
}
|
||||
Ok(())
|
||||
})
|
||||
.await?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
let content = deltas.concat();
|
||||
let parsed_tool_calls = parse_and_filter_windsurf_tool_calls(&content, &prepared.input);
|
||||
let mut tool_calls = poll_result.native_tool_calls;
|
||||
tool_calls.extend(parsed_tool_calls.tool_calls);
|
||||
let has_tool_calls = !tool_calls.is_empty();
|
||||
let message = if has_tool_calls {
|
||||
json!({
|
||||
"role": "assistant",
|
||||
"content": Value::Null,
|
||||
"tool_calls": openai_tool_call_values(&tool_calls),
|
||||
})
|
||||
} else {
|
||||
json!({
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
})
|
||||
};
|
||||
let mut body_json = json!({
|
||||
"id": format!("chatcmpl-{}", prepared.request_id),
|
||||
"object": "chat.completion",
|
||||
"created": current_unix_secs(),
|
||||
"model": prepared.model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": message,
|
||||
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" },
|
||||
}],
|
||||
});
|
||||
if let Some(usage) = poll_result.usage {
|
||||
body_json["usage"] = windsurf_openai_usage_json(&usage);
|
||||
}
|
||||
|
||||
Ok(Some(ExecutionResult {
|
||||
request_id: prepared.request_id,
|
||||
candidate_id: prepared.candidate_id,
|
||||
status_code: 200,
|
||||
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
|
||||
body: Some(ResponseBody {
|
||||
json_body: Some(body_json),
|
||||
body_bytes_b64: None,
|
||||
}),
|
||||
telemetry: Some(ExecutionTelemetry {
|
||||
ttfb_ms: None,
|
||||
elapsed_ms: Some(elapsed_ms),
|
||||
upstream_bytes: None,
|
||||
}),
|
||||
error: None,
|
||||
}))
|
||||
Ok(Some(ExecutionResult {
|
||||
request_id: prepared.request_id,
|
||||
candidate_id: prepared.candidate_id,
|
||||
status_code: 200,
|
||||
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
|
||||
body: Some(ResponseBody {
|
||||
json_body: Some(body_json),
|
||||
body_bytes_b64: None,
|
||||
}),
|
||||
telemetry: Some(ExecutionTelemetry {
|
||||
ttfb_ms: None,
|
||||
elapsed_ms: Some(elapsed_ms),
|
||||
upstream_bytes: None,
|
||||
}),
|
||||
error: None,
|
||||
}))
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn prepare_windsurf_cascade(
|
||||
@@ -1172,13 +1176,19 @@ async fn ensure_windsurf_language_server(
|
||||
let mut command = Command::new(&binary_path);
|
||||
command
|
||||
.arg(format!("--api_server_url={}", codeium_api_url()))
|
||||
.arg("--run_child")
|
||||
.arg(format!("--server_port={port}"))
|
||||
.arg(format!("--csrf_token={DEFAULT_CSRF_TOKEN}"))
|
||||
.arg(format!("--register_user_url={DEFAULT_REGISTER_USER_URL}"))
|
||||
.arg(format!("--codeium_dir={}", data_dir.display()))
|
||||
.arg(format!("--database_dir={}", data_dir.join("db").display()))
|
||||
.arg("--detect_proxy=false")
|
||||
.env_clear()
|
||||
.arg("--detect_proxy=false");
|
||||
|
||||
if !cfg!(target_os = "windows") {
|
||||
command.env_clear();
|
||||
}
|
||||
|
||||
command
|
||||
.envs(language_server_env(proxy_url.as_deref()))
|
||||
.stdin(Stdio::null())
|
||||
.stdout(Stdio::null())
|
||||
@@ -1191,7 +1201,8 @@ async fn ensure_windsurf_language_server(
|
||||
))
|
||||
})?;
|
||||
|
||||
if let Err(err) = wait_language_server_ready(port).await {
|
||||
if let Err(err) = wait_language_server_ready(port, &mut child, stderr_log_path.as_deref()).await
|
||||
{
|
||||
let _ = child.kill();
|
||||
return Err(err);
|
||||
}
|
||||
@@ -1379,7 +1390,7 @@ async fn windsurf_warmup_unary(
|
||||
Err(err) if is_windsurf_cascade_transport_error(&err) => Err(err),
|
||||
Err(err) => {
|
||||
if stage == "UpdateWorkspaceTrust" {
|
||||
error!(
|
||||
warn!(
|
||||
event_name = "windsurf_workspace_trust_update_failed",
|
||||
log_type = "ops",
|
||||
port,
|
||||
@@ -1512,44 +1523,39 @@ async fn windsurf_grpc_unary(
|
||||
) -> Result<Vec<u8>, ExecutionRuntimeTransportError> {
|
||||
let url = format!("http://127.0.0.1:{port}{LS_SERVICE}/{method}");
|
||||
let client = reqwest::Client::builder()
|
||||
.http2_prior_knowledge()
|
||||
.http1_only()
|
||||
.timeout(timeout)
|
||||
.build()
|
||||
.map_err(ExecutionRuntimeTransportError::ClientBuild)?;
|
||||
let response = client
|
||||
.post(url)
|
||||
.header("content-type", "application/grpc")
|
||||
.header("te", "trailers")
|
||||
.header("user-agent", "grpc-node/1.108.2")
|
||||
.header("content-type", "application/proto")
|
||||
.header("connect-protocol-version", "1")
|
||||
.header("user-agent", "connect-es/1.5.0")
|
||||
.header("x-codeium-csrf-token", csrf_token)
|
||||
.body(grpc_frame(&payload))
|
||||
.body(payload)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format!(
|
||||
"Windsurf gRPC {method} request failed: {}",
|
||||
"Windsurf Connect {method} request failed: {}",
|
||||
super::transport::format_upstream_request_error(&err)
|
||||
))
|
||||
})?;
|
||||
let status = response.status();
|
||||
let body = response.bytes().await.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format!(
|
||||
"Windsurf gRPC {method} response read failed: {}",
|
||||
"Windsurf Connect {method} response read failed: {}",
|
||||
super::transport::format_upstream_request_error(&err)
|
||||
))
|
||||
})?;
|
||||
if !status.is_success() {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!(
|
||||
"Windsurf gRPC {method} returned HTTP {status}: {}",
|
||||
"Windsurf Connect {method} returned HTTP {status}: {}",
|
||||
String::from_utf8_lossy(&body)
|
||||
)));
|
||||
}
|
||||
let frames = extract_grpc_frames(&body);
|
||||
if frames.is_empty() {
|
||||
Ok(body.to_vec())
|
||||
} else {
|
||||
Ok(frames.concat())
|
||||
}
|
||||
Ok(body.to_vec())
|
||||
}
|
||||
|
||||
fn detect_windsurf_request(
|
||||
@@ -3897,10 +3903,27 @@ fn port_is_free(port: u16) -> bool {
|
||||
TcpListener::bind(("127.0.0.1", port)).is_ok()
|
||||
}
|
||||
|
||||
async fn wait_language_server_ready(port: u16) -> Result<(), ExecutionRuntimeTransportError> {
|
||||
async fn wait_language_server_ready(
|
||||
port: u16,
|
||||
child: &mut Child,
|
||||
stderr_log_path: Option<&Path>,
|
||||
) -> Result<(), ExecutionRuntimeTransportError> {
|
||||
let timeout = if cfg!(target_os = "windows") {
|
||||
WINDOWS_LS_READY_TIMEOUT
|
||||
} else {
|
||||
LS_READY_TIMEOUT
|
||||
};
|
||||
let started = Instant::now();
|
||||
let addr = SocketAddr::from(([127, 0, 0, 1], port));
|
||||
while started.elapsed() < LS_READY_TIMEOUT {
|
||||
while started.elapsed() < timeout {
|
||||
if let Ok(Some(status)) = child.try_wait() {
|
||||
let stderr_tail = stderr_log_path
|
||||
.and_then(|path| read_log_tail(path, 8 * 1024))
|
||||
.unwrap_or_else(|| "<stderr log unavailable>".to_string());
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(format!(
|
||||
"Windsurf language server exited before port {port} became ready with status {status}; stderr tail: {stderr_tail}"
|
||||
)));
|
||||
}
|
||||
if TcpStream::connect_timeout(&addr, Duration::from_millis(200)).is_ok() {
|
||||
debug!(
|
||||
event_name = "windsurf_language_server_port_ready",
|
||||
@@ -3912,12 +3935,30 @@ async fn wait_language_server_ready(port: u16) -> Result<(), ExecutionRuntimeTra
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(250)).await;
|
||||
}
|
||||
let child_status = child.try_wait().ok().flatten();
|
||||
let stderr_tail = stderr_log_path
|
||||
.and_then(|path| read_log_tail(path, 8 * 1024))
|
||||
.unwrap_or_else(|| "<stderr log unavailable>".to_string());
|
||||
Err(ExecutionRuntimeTransportError::UpstreamRequest(format!(
|
||||
"Windsurf language server port {port} was not ready after {}ms",
|
||||
LS_READY_TIMEOUT.as_millis()
|
||||
"Windsurf language server port {port} was not ready after {}ms{}; stderr tail: {stderr_tail}",
|
||||
timeout.as_millis(),
|
||||
child_status
|
||||
.map(|status| format!(" (child status: {status})"))
|
||||
.unwrap_or_default()
|
||||
)))
|
||||
}
|
||||
|
||||
fn read_log_tail(path: &Path, max_bytes: usize) -> Option<String> {
|
||||
let mut file = fs::File::open(path).ok()?;
|
||||
let len = file.metadata().ok()?.len() as i64;
|
||||
let max_bytes = max_bytes.max(1) as i64;
|
||||
let start = len.saturating_sub(max_bytes);
|
||||
file.seek(SeekFrom::Start(start as u64)).ok()?;
|
||||
let mut buf = Vec::with_capacity((len - start) as usize);
|
||||
file.read_to_end(&mut buf).ok()?;
|
||||
Some(String::from_utf8_lossy(&buf).to_string())
|
||||
}
|
||||
|
||||
fn language_server_data_dir(key: &str) -> PathBuf {
|
||||
for env_key in ["WINDSURF_LS_DATA_DIR", "LS_DATA_DIR"] {
|
||||
if let Some(path) = std::env::var_os(env_key).filter(|value| !value.is_empty()) {
|
||||
@@ -3979,9 +4020,36 @@ fn language_server_env(proxy_url: Option<&str>) -> BTreeMap<String, String> {
|
||||
}
|
||||
}
|
||||
}
|
||||
if cfg!(target_os = "windows") {
|
||||
for key in [
|
||||
"USERPROFILE",
|
||||
"APPDATA",
|
||||
"LOCALAPPDATA",
|
||||
"SystemRoot",
|
||||
"WINDIR",
|
||||
"ComSpec",
|
||||
] {
|
||||
if let Ok(value) = std::env::var(key) {
|
||||
if !value.trim().is_empty() {
|
||||
env.insert(key.to_string(), value);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if !env.contains_key("HOME") {
|
||||
env.insert("HOME".to_string(), home_dir().display().to_string());
|
||||
}
|
||||
if cfg!(target_os = "windows") && !env.contains_key("USERPROFILE") {
|
||||
if let Some(home) = std::env::var_os("USERPROFILE")
|
||||
.or_else(|| std::env::var_os("HOME"))
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
env.insert(
|
||||
"USERPROFILE".to_string(),
|
||||
PathBuf::from(home).display().to_string(),
|
||||
);
|
||||
}
|
||||
}
|
||||
if let Some(proxy_url) = proxy_url.filter(|value| !value.trim().is_empty()) {
|
||||
for key in ["HTTP_PROXY", "HTTPS_PROXY", "http_proxy", "https_proxy"] {
|
||||
env.insert(key.to_string(), proxy_url.to_string());
|
||||
@@ -3998,9 +4066,15 @@ fn codeium_api_url() -> String {
|
||||
}
|
||||
|
||||
fn home_dir() -> PathBuf {
|
||||
std::env::var_os("HOME")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| PathBuf::from("."))
|
||||
if let Some(home) = std::env::var_os("HOME").filter(|value| !value.is_empty()) {
|
||||
return PathBuf::from(home);
|
||||
}
|
||||
if cfg!(target_os = "windows") {
|
||||
if let Some(home) = std::env::var_os("USERPROFILE").filter(|value| !value.is_empty()) {
|
||||
return PathBuf::from(home);
|
||||
}
|
||||
}
|
||||
PathBuf::from(".")
|
||||
}
|
||||
|
||||
fn first_string(body: &Value, keys: &[&str]) -> Option<String> {
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
use aether_ai_serving::{
|
||||
run_ai_attempt_loop, AiAttemptLoopOutcome, AiAttemptLoopPort, AiExecutionAttempt,
|
||||
UPSTREAM_IS_STREAM_KEY,
|
||||
};
|
||||
use aether_data_contracts::repository::candidates::RequestCandidateStatus;
|
||||
use aether_scheduler_core::{
|
||||
@@ -26,7 +25,7 @@ use crate::request_candidate_runtime::{
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const DEFAULT_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS: u64 = 300_000;
|
||||
const DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS: u64 = 30_000;
|
||||
|
||||
fn attach_redaction_execution_candidate(response: &mut Response<Body>, candidate_id: Option<&str>) {
|
||||
if let Some(candidate_id) = candidate_id
|
||||
@@ -125,7 +124,16 @@ where
|
||||
decision,
|
||||
plan_kind,
|
||||
};
|
||||
run_dynamic_attempt_loop(&port, &mut source).await
|
||||
run_dynamic_attempt_loop(
|
||||
&port,
|
||||
&mut source,
|
||||
trace_id,
|
||||
plan_kind,
|
||||
state
|
||||
.frontdoor_runtime_guards
|
||||
.local_execution_planning_timeout,
|
||||
)
|
||||
.await
|
||||
}
|
||||
.instrument(span)
|
||||
.await
|
||||
@@ -281,7 +289,16 @@ where
|
||||
decision,
|
||||
plan_kind,
|
||||
};
|
||||
run_dynamic_attempt_loop(&port, &mut source).await
|
||||
run_dynamic_attempt_loop(
|
||||
&port,
|
||||
&mut source,
|
||||
trace_id,
|
||||
plan_kind,
|
||||
state
|
||||
.frontdoor_runtime_guards
|
||||
.local_execution_planning_timeout,
|
||||
)
|
||||
.await
|
||||
}
|
||||
.instrument(span)
|
||||
.await
|
||||
@@ -290,6 +307,9 @@ where
|
||||
async fn run_dynamic_attempt_loop<Port, Source, Attempt>(
|
||||
port: &Port,
|
||||
source: &mut Source,
|
||||
trace_id: &str,
|
||||
plan_kind: &str,
|
||||
planning_timeout: Duration,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError>
|
||||
where
|
||||
Port: AiAttemptLoopPort<
|
||||
@@ -303,7 +323,9 @@ where
|
||||
{
|
||||
let mut last_attempted = None;
|
||||
|
||||
while let Some(attempt) = source.next_execution_attempt().await? {
|
||||
while let Some(attempt) =
|
||||
next_execution_attempt_with_timeout(source, trace_id, plan_kind, planning_timeout).await?
|
||||
{
|
||||
last_attempted = Some((attempt.execution_plan().clone(), attempt.report_context()));
|
||||
if let Some(response) = port.execute_attempt(&attempt).await? {
|
||||
let remaining = source.drain_execution_attempts().await?;
|
||||
@@ -322,6 +344,37 @@ where
|
||||
))
|
||||
}
|
||||
|
||||
async fn next_execution_attempt_with_timeout<Source, Attempt>(
|
||||
source: &mut Source,
|
||||
trace_id: &str,
|
||||
plan_kind: &str,
|
||||
planning_timeout: Duration,
|
||||
) -> Result<Option<Attempt>, GatewayError>
|
||||
where
|
||||
Source: LocalExecutionAttemptSource<Attempt>,
|
||||
{
|
||||
match timeout(planning_timeout, source.next_execution_attempt()).await {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
let timeout_ms = planning_timeout.as_millis() as u64;
|
||||
warn!(
|
||||
event_name = "local_execution_candidate_planning_timeout",
|
||||
log_type = "ops",
|
||||
trace_id,
|
||||
plan_kind,
|
||||
timeout_ms,
|
||||
phase = "next_execution_attempt",
|
||||
"gateway timed out while planning the next local execution candidate"
|
||||
);
|
||||
Err(GatewayError::LocalExecutionPlanningTimeout {
|
||||
trace_id: trace_id.to_string(),
|
||||
phase: "next_execution_attempt",
|
||||
timeout_ms,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct StreamAttemptLoopPort<'a> {
|
||||
state: &'a AppState,
|
||||
trace_id: &'a str,
|
||||
@@ -474,27 +527,21 @@ fn should_skip_unused_persistence_from_metadata(
|
||||
|
||||
fn resolve_stream_candidate_watchdog_timeout(
|
||||
plan: &aether_contracts::ExecutionPlan,
|
||||
report_context: Option<&serde_json::Value>,
|
||||
_report_context: Option<&serde_json::Value>,
|
||||
) -> Duration {
|
||||
let upstream_is_stream = report_context
|
||||
.and_then(|context| context.get(UPSTREAM_IS_STREAM_KEY))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(true);
|
||||
let timeout_ms = plan
|
||||
.timeouts
|
||||
.as_ref()
|
||||
.and_then(|timeouts| {
|
||||
if upstream_is_stream {
|
||||
timeouts.first_byte_ms.or(timeouts.total_ms)
|
||||
} else {
|
||||
timeouts.total_ms.or(timeouts.first_byte_ms)
|
||||
}
|
||||
})
|
||||
.unwrap_or(DEFAULT_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS)
|
||||
.and_then(|timeouts| timeouts.first_byte_ms)
|
||||
.unwrap_or(DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS)
|
||||
.max(1);
|
||||
Duration::from_millis(timeout_ms)
|
||||
}
|
||||
|
||||
fn stream_candidate_watchdog_timeout_message() -> &'static str {
|
||||
"Stream first byte timeout"
|
||||
}
|
||||
|
||||
async fn execute_stream_candidate_with_watchdog<Fut>(
|
||||
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
|
||||
trace_id: &str,
|
||||
@@ -534,9 +581,7 @@ where
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: Some(http::StatusCode::GATEWAY_TIMEOUT.as_u16()),
|
||||
error_type: Some("local_stream_candidate_watchdog_timeout".to_string()),
|
||||
error_message: Some(format!(
|
||||
"local stream candidate attempt exceeded watchdog timeout of {timeout_ms}ms"
|
||||
)),
|
||||
error_message: Some(stream_candidate_watchdog_timeout_message().to_string()),
|
||||
latency_ms: None,
|
||||
started_at_unix_ms: Some(candidate_started_unix_ms),
|
||||
finished_at_unix_ms: Some(finished_at_unix_ms),
|
||||
@@ -632,6 +677,20 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
struct PendingAttemptSource;
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<()> for PendingAttemptSource {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<()>, GatewayError> {
|
||||
std::future::pending::<()>().await;
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<()>, GatewayError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
fn test_plan(timeouts: Option<ExecutionTimeouts>) -> ExecutionPlan {
|
||||
ExecutionPlan {
|
||||
request_id: "req_watchdog".to_string(),
|
||||
@@ -656,6 +715,33 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn next_execution_attempt_times_out_instead_of_waiting_forever() {
|
||||
let mut source = PendingAttemptSource;
|
||||
|
||||
let err = next_execution_attempt_with_timeout(
|
||||
&mut source,
|
||||
"trace-planning-timeout",
|
||||
"openai_responses_sync",
|
||||
Duration::from_millis(5),
|
||||
)
|
||||
.await
|
||||
.expect_err("pending candidate planning should time out");
|
||||
|
||||
match err {
|
||||
GatewayError::LocalExecutionPlanningTimeout {
|
||||
trace_id,
|
||||
phase,
|
||||
timeout_ms,
|
||||
} => {
|
||||
assert_eq!(trace_id, "trace-planning-timeout");
|
||||
assert_eq!(phase, "next_execution_attempt");
|
||||
assert_eq!(timeout_ms, 5);
|
||||
}
|
||||
other => panic!("unexpected error: {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn test_report_context() -> serde_json::Value {
|
||||
json!({
|
||||
"request_id": "req_watchdog",
|
||||
@@ -688,37 +774,57 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
timeout,
|
||||
Duration::from_millis(DEFAULT_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS)
|
||||
Duration::from_millis(DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_candidate_watchdog_prefers_total_timeout_when_upstream_non_stream() {
|
||||
fn stream_candidate_watchdog_ignores_total_timeout_for_stream_upstream() {
|
||||
let report_context = json!({"upstream_is_stream": true});
|
||||
let timeout = resolve_stream_candidate_watchdog_timeout(
|
||||
&test_plan(Some(ExecutionTimeouts {
|
||||
total_ms: Some(90_000),
|
||||
..ExecutionTimeouts::default()
|
||||
})),
|
||||
Some(&report_context),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
timeout,
|
||||
Duration::from_millis(DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_candidate_watchdog_prefers_first_byte_timeout_when_upstream_non_stream() {
|
||||
let report_context = json!({"upstream_is_stream": false});
|
||||
let timeout = resolve_stream_candidate_watchdog_timeout(
|
||||
&test_plan(Some(ExecutionTimeouts {
|
||||
first_byte_ms: Some(300_000),
|
||||
first_byte_ms: Some(12_345),
|
||||
total_ms: Some(599_000),
|
||||
..ExecutionTimeouts::default()
|
||||
})),
|
||||
Some(&report_context),
|
||||
);
|
||||
|
||||
assert_eq!(timeout, Duration::from_millis(599_000));
|
||||
assert_eq!(timeout, Duration::from_millis(12_345));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_candidate_watchdog_falls_back_to_first_byte_when_upstream_non_stream_lacks_total() {
|
||||
fn stream_candidate_watchdog_ignores_total_timeout_when_upstream_non_stream() {
|
||||
let report_context = json!({"upstream_is_stream": false});
|
||||
let timeout = resolve_stream_candidate_watchdog_timeout(
|
||||
&test_plan(Some(ExecutionTimeouts {
|
||||
first_byte_ms: Some(300_000),
|
||||
total_ms: Some(599_000),
|
||||
..ExecutionTimeouts::default()
|
||||
})),
|
||||
Some(&report_context),
|
||||
);
|
||||
|
||||
assert_eq!(timeout, Duration::from_millis(300_000));
|
||||
assert_eq!(
|
||||
timeout,
|
||||
Duration::from_millis(DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -795,7 +901,7 @@ mod tests {
|
||||
assert!(record
|
||||
.error_message
|
||||
.as_deref()
|
||||
.is_some_and(|message| message.contains("25ms")));
|
||||
.is_some_and(|message| message == "Stream first byte timeout"));
|
||||
assert_eq!(record.candidate_index, 2);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -48,6 +48,7 @@ pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool {
|
||||
) || path.starts_with("/v1/videos/")
|
||||
|| path.starts_with("/v1beta/files/")
|
||||
|| path.starts_with("/v1beta/operations/")
|
||||
|| path.starts_with("/v1internal:")
|
||||
|| is_gemini_generation_path(path)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
use aether_scheduler_core::count_recent_rpm_requests_for_provider_key_since;
|
||||
use aether_scheduler_core::{
|
||||
count_recent_rpm_requests_for_provider_key_since,
|
||||
provider_key_circuit_payload_is_active_open_at,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
@@ -18,6 +21,10 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.await
|
||||
.ok()
|
||||
.and_then(|mut keys| keys.drain(..).next())?;
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or_default();
|
||||
let provider = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
|
||||
.await
|
||||
@@ -81,10 +88,10 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.and_then(|value| value.get("last_failure_at"))
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
payload["circuit_breaker_open"] = json!(circuit_data
|
||||
.and_then(|value| value.get("open"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false));
|
||||
payload["circuit_breaker_open"] =
|
||||
json!(circuit_data.is_some_and(
|
||||
|value| provider_key_circuit_payload_is_active_open_at(value, now_unix_secs)
|
||||
));
|
||||
payload["circuit_breaker_open_at"] = circuit_data
|
||||
.and_then(|value| value.get("open_at"))
|
||||
.cloned()
|
||||
@@ -107,15 +114,16 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.unwrap_or(0));
|
||||
} else {
|
||||
let mut formats_payload = serde_json::Map::new();
|
||||
let mut any_circuit_open = false;
|
||||
for format_name in
|
||||
provider_key_effective_api_formats(&key, &provider.provider_type, &endpoints)
|
||||
{
|
||||
let health_data = health_by_format.and_then(|formats| formats.get(&format_name));
|
||||
let circuit_data = circuit_by_format.and_then(|formats| formats.get(&format_name));
|
||||
let is_open = circuit_data
|
||||
.and_then(|value| value.get("open"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let active_open = circuit_data.is_some_and(|value| {
|
||||
provider_key_circuit_payload_is_active_open_at(value, now_unix_secs)
|
||||
});
|
||||
any_circuit_open |= active_open;
|
||||
formats_payload.insert(
|
||||
format_name.clone(),
|
||||
json!({
|
||||
@@ -134,8 +142,8 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null),
|
||||
"circuit_breaker": {
|
||||
"state": if is_open { "open" } else { "closed" },
|
||||
"open": is_open,
|
||||
"state": if active_open { "open" } else { "closed" },
|
||||
"open": active_open,
|
||||
"open_at": circuit_data
|
||||
.and_then(|value| value.get("open_at"))
|
||||
.cloned()
|
||||
@@ -167,13 +175,6 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.filter_map(serde_json::Value::as_f64)
|
||||
.reduce(f64::min)
|
||||
.unwrap_or(1.0);
|
||||
let any_circuit_open = formats_payload.values().any(|value| {
|
||||
value
|
||||
.get("circuit_breaker")
|
||||
.and_then(|circuit| circuit.get("open"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
});
|
||||
|
||||
payload["key_health_score"] = json!(key_health_score);
|
||||
payload["any_circuit_open"] = json!(any_circuit_open);
|
||||
|
||||
@@ -4,7 +4,9 @@ use crate::handlers::public::{api_format_display_name, build_public_health_timel
|
||||
use crate::handlers::shared::unix_ms_to_rfc3339;
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
use aether_data_contracts::repository::candidates::PublicHealthTimelineBucket;
|
||||
use aether_scheduler_core::{is_provider_key_circuit_open, provider_key_health_score};
|
||||
use aether_scheduler_core::{
|
||||
any_provider_key_circuit_open_at, is_provider_key_circuit_open_at, provider_key_health_score,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -99,7 +101,9 @@ pub(crate) async fn build_admin_endpoint_health_status_payload(
|
||||
.entry(api_format.clone())
|
||||
.or_default()
|
||||
.insert(key.id.clone());
|
||||
if key.is_active && !is_provider_key_circuit_open(&key, &api_format) {
|
||||
if key.is_active
|
||||
&& !is_provider_key_circuit_open_at(&key, &api_format, now_unix_secs)
|
||||
{
|
||||
let key_health_score =
|
||||
provider_key_health_score(&key, &api_format).unwrap_or(1.0);
|
||||
active_keys_by_format
|
||||
@@ -229,6 +233,10 @@ pub(crate) async fn build_admin_health_summary_payload(
|
||||
return None;
|
||||
}
|
||||
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or_default();
|
||||
let providers = state
|
||||
.list_provider_catalog_providers(false)
|
||||
.await
|
||||
@@ -286,20 +294,7 @@ pub(crate) async fn build_admin_health_summary_payload(
|
||||
.count();
|
||||
let circuit_open_keys = keys
|
||||
.iter()
|
||||
.filter(|key| {
|
||||
key.circuit_breaker_by_format
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.map(|formats| {
|
||||
formats.values().any(|circuit| {
|
||||
circuit
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
})
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.filter(|key| any_provider_key_circuit_open_at(key, now_unix_secs))
|
||||
.count();
|
||||
|
||||
Some(json!({
|
||||
|
||||
@@ -8,10 +8,12 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_scheduler_core::{
|
||||
is_provider_key_circuit_open, matches_model_mapping, provider_key_health_score,
|
||||
is_provider_key_circuit_open_at, matches_model_mapping,
|
||||
provider_key_circuit_payload_is_active_open_at, provider_key_health_score,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
@@ -86,6 +88,10 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
.flatten()
|
||||
.and_then(|value| value.as_bool())
|
||||
.unwrap_or(false);
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or_default();
|
||||
|
||||
let global_model_mappings = global_model
|
||||
.config
|
||||
@@ -163,9 +169,11 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
entries
|
||||
.iter()
|
||||
.filter_map(|(api_format, value)| {
|
||||
value.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.filter(|is_open| *is_open)
|
||||
provider_key_circuit_payload_is_active_open_at(
|
||||
value,
|
||||
now_unix_secs,
|
||||
)
|
||||
.then_some(())
|
||||
.map(|_| api_format.clone())
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
@@ -188,7 +196,7 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
"effective_rpm": effective_rpm,
|
||||
"allowed_models": allowed_models,
|
||||
"health_score": provider_key_health_score(key, &endpoint.api_format),
|
||||
"circuit_breaker_open": is_provider_key_circuit_open(key, &endpoint.api_format),
|
||||
"circuit_breaker_open": is_provider_key_circuit_open_at(key, &endpoint.api_format, now_unix_secs),
|
||||
"circuit_breaker_formats": circuit_breaker_formats,
|
||||
"next_probe_at": next_probe_at,
|
||||
});
|
||||
|
||||
+8
-7
@@ -1,6 +1,6 @@
|
||||
use super::super::usage_helpers::admin_monitoring_usage_is_error;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::{provider_key_health_summary, unix_secs_to_rfc3339};
|
||||
use crate::handlers::admin::shared::{provider_key_health_summary_at, unix_secs_to_rfc3339};
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::{
|
||||
provider_catalog::StoredProviderCatalogKey, usage::UsageMonitoringErrorListQuery,
|
||||
@@ -99,7 +99,7 @@ pub(super) async fn build_admin_monitoring_resilience_snapshot(
|
||||
last_failure_at,
|
||||
circuit_breaker_open,
|
||||
circuit_by_format,
|
||||
) = provider_key_health_summary(key);
|
||||
) = provider_key_health_summary_at(key, now.timestamp().max(0) as u64);
|
||||
if health_score < 0.8 {
|
||||
degraded_keys += 1;
|
||||
}
|
||||
@@ -110,11 +110,12 @@ pub(super) async fn build_admin_monitoring_resilience_snapshot(
|
||||
let open_formats = circuit_by_format
|
||||
.iter()
|
||||
.filter_map(|(api_format, value)| {
|
||||
value
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.filter(|open| *open)
|
||||
.map(|_| api_format.clone())
|
||||
aether_scheduler_core::provider_key_circuit_payload_is_active_open_at(
|
||||
value,
|
||||
now.timestamp().max(0) as u64,
|
||||
)
|
||||
.then_some(())
|
||||
.map(|_| api_format.clone())
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
|
||||
@@ -1237,7 +1237,8 @@ async fn admin_monitoring_circuit_history_returns_local_payload() {
|
||||
"openai:chat": {
|
||||
"open": true,
|
||||
"open_at": "2026-03-30T12:00:00+00:00",
|
||||
"next_probe_at": "2026-03-30T12:05:00+00:00",
|
||||
"next_probe_at": "2099-03-30T12:05:00+00:00",
|
||||
"recovery_seconds": 300,
|
||||
"reason": "错误率过高"
|
||||
}
|
||||
})),
|
||||
|
||||
@@ -81,6 +81,133 @@ async fn admin_monitoring_trace_request_returns_local_payload() {
|
||||
assert_eq!(payload["candidates"][0]["status_code"], json!(502));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
sample_candidate(
|
||||
"cand-used",
|
||||
"trace-1",
|
||||
0,
|
||||
RequestCandidateStatus::Success,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(200),
|
||||
),
|
||||
]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
vec![sample_endpoint()],
|
||||
vec![sample_key()],
|
||||
));
|
||||
let mut usage = sample_usage(
|
||||
"usage-request-1",
|
||||
"provider-1",
|
||||
"OpenAI",
|
||||
40,
|
||||
0.02,
|
||||
"completed",
|
||||
Some(200),
|
||||
100,
|
||||
);
|
||||
usage.id = "usage-row-1".to_string();
|
||||
usage.candidate_id = Some("cand-used".to_string());
|
||||
usage.request_headers = Some(json!({
|
||||
"x-trace-id": "trace-1"
|
||||
}));
|
||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
|
||||
request_candidates,
|
||||
usage_repository,
|
||||
)
|
||||
.with_provider_catalog_reader(provider_catalog);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let context = request_context(
|
||||
http::Method::GET,
|
||||
"/api/admin/monitoring/trace/usage-row-1?attempted_only=true",
|
||||
);
|
||||
|
||||
let response = local_monitoring_response(&state, &context)
|
||||
.await
|
||||
.expect("handler should not error")
|
||||
.expect("route should be handled locally");
|
||||
|
||||
assert_eq!(response.status(), http::StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||
assert_eq!(payload["request_id"], json!("trace-1"));
|
||||
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
|
||||
assert_eq!(
|
||||
payload["candidates"][0]["extra_data"]["first_byte_time_ms"],
|
||||
json!(30)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_resolves_usage_request_id_to_metadata_trace_id() {
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
sample_candidate(
|
||||
"cand-used",
|
||||
"trace-2",
|
||||
0,
|
||||
RequestCandidateStatus::Success,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(200),
|
||||
),
|
||||
]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
vec![sample_endpoint()],
|
||||
vec![sample_key()],
|
||||
));
|
||||
let mut usage = sample_usage(
|
||||
"usage-request-2",
|
||||
"provider-1",
|
||||
"OpenAI",
|
||||
40,
|
||||
0.02,
|
||||
"completed",
|
||||
Some(200),
|
||||
100,
|
||||
);
|
||||
usage.candidate_id = Some("cand-used".to_string());
|
||||
usage.request_metadata = Some(json!({
|
||||
"trace_id": "trace-2"
|
||||
}));
|
||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
|
||||
request_candidates,
|
||||
usage_repository,
|
||||
)
|
||||
.with_provider_catalog_reader(provider_catalog);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let context = request_context(
|
||||
http::Method::GET,
|
||||
"/api/admin/monitoring/trace/usage-request-2",
|
||||
);
|
||||
|
||||
let response = local_monitoring_response(&state, &context)
|
||||
.await
|
||||
.expect("handler should not error")
|
||||
.expect("route should be handled locally");
|
||||
|
||||
assert_eq!(response.status(), http::StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||
assert_eq!(payload["request_id"], json!("trace-2"));
|
||||
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_returns_oauth_account_label_from_auth_config() {
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
|
||||
@@ -12,6 +12,7 @@ use aether_admin::observability::monitoring::{
|
||||
use aether_data_contracts::repository::{
|
||||
candidates::{DecisionTrace, RequestCandidateStatus},
|
||||
provider_catalog::StoredProviderCatalogKey,
|
||||
usage::StoredRequestUsageAudit,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
@@ -21,12 +22,16 @@ use serde_json::{Map, Value};
|
||||
use std::collections::BTreeMap;
|
||||
use tracing::debug;
|
||||
|
||||
struct ResolvedAdminMonitoringTrace {
|
||||
trace: DecisionTrace,
|
||||
usage: Option<StoredRequestUsageAudit>,
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_monitoring_trace_request_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let admin_state = state;
|
||||
let state = state.as_ref();
|
||||
let Some(request_id) =
|
||||
admin_monitoring_trace_request_id_from_path(&request_context.request_path)
|
||||
else {
|
||||
@@ -39,11 +44,8 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
|
||||
Err(detail) => return Ok(admin_monitoring_bad_request_response(detail)),
|
||||
};
|
||||
|
||||
let Some(trace) = state
|
||||
.data
|
||||
.read_decision_trace(&request_id, attempted_only)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
let Some(resolved) =
|
||||
resolve_admin_monitoring_trace(admin_state, &request_id, attempted_only).await?
|
||||
else {
|
||||
debug!(
|
||||
event_name = "admin_monitoring_request_trace_not_found",
|
||||
@@ -58,22 +60,113 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
|
||||
attempted_only,
|
||||
));
|
||||
};
|
||||
let usage = state
|
||||
.data
|
||||
.read_request_usage_audit(&request_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let key_accounts = build_admin_monitoring_key_account_display_map(admin_state, &trace).await?;
|
||||
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(
|
||||
&trace,
|
||||
usage.as_ref(),
|
||||
&resolved.trace,
|
||||
resolved.usage.as_ref(),
|
||||
&key_accounts,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
async fn resolve_admin_monitoring_trace(
|
||||
state: &AdminAppState<'_>,
|
||||
request_id: &str,
|
||||
attempted_only: bool,
|
||||
) -> Result<Option<ResolvedAdminMonitoringTrace>, GatewayError> {
|
||||
let app = state.as_ref();
|
||||
if let Some(trace) = app
|
||||
.data
|
||||
.read_decision_trace(request_id, attempted_only)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
{
|
||||
let usage = app
|
||||
.data
|
||||
.read_request_usage_audit(request_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
return Ok(Some(ResolvedAdminMonitoringTrace { trace, usage }));
|
||||
}
|
||||
|
||||
let mut usage_candidates = Vec::new();
|
||||
if let Some(usage) = app
|
||||
.data
|
||||
.read_request_usage_audit(request_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
{
|
||||
usage_candidates.push(usage);
|
||||
}
|
||||
if let Some(usage) = state.find_request_usage_by_id(request_id).await? {
|
||||
if !usage_candidates.iter().any(|item| item.id == usage.id) {
|
||||
usage_candidates.push(usage);
|
||||
}
|
||||
}
|
||||
|
||||
for usage in usage_candidates {
|
||||
for trace_request_id in admin_monitoring_usage_trace_request_ids(&usage) {
|
||||
if trace_request_id == request_id {
|
||||
continue;
|
||||
}
|
||||
if let Some(trace) = app
|
||||
.data
|
||||
.read_decision_trace(&trace_request_id, attempted_only)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
{
|
||||
return Ok(Some(ResolvedAdminMonitoringTrace {
|
||||
trace,
|
||||
usage: Some(usage),
|
||||
}));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn admin_monitoring_usage_trace_request_ids(usage: &StoredRequestUsageAudit) -> Vec<String> {
|
||||
let mut ids = Vec::new();
|
||||
push_non_empty_unique(&mut ids, usage.request_id.as_str());
|
||||
if let Some(trace_id) = usage.trace_id() {
|
||||
push_non_empty_unique(&mut ids, trace_id);
|
||||
}
|
||||
if let Some(trace_id) = usage_trace_id_from_headers(usage.request_headers.as_ref()) {
|
||||
push_non_empty_unique(&mut ids, trace_id.as_str());
|
||||
}
|
||||
if let Some(trace_id) = usage_trace_id_from_headers(usage.provider_request_headers.as_ref()) {
|
||||
push_non_empty_unique(&mut ids, trace_id.as_str());
|
||||
}
|
||||
ids
|
||||
}
|
||||
|
||||
fn usage_trace_id_from_headers(headers: Option<&Value>) -> Option<String> {
|
||||
let object = headers?.as_object()?;
|
||||
object.iter().find_map(|(key, value)| {
|
||||
key.eq_ignore_ascii_case(crate::constants::TRACE_ID_HEADER)
|
||||
.then(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
.flatten()
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
fn push_non_empty_unique(values: &mut Vec<String>, value: &str) {
|
||||
let value = value.trim();
|
||||
if value.is_empty() || values.iter().any(|existing| existing == value) {
|
||||
return;
|
||||
}
|
||||
values.push(value.to_string());
|
||||
}
|
||||
|
||||
async fn build_admin_monitoring_key_account_display_map(
|
||||
state: &AdminAppState<'_>,
|
||||
trace: &DecisionTrace,
|
||||
|
||||
@@ -91,6 +91,26 @@ fn coerce_admin_provider_oauth_import_project_id(
|
||||
}
|
||||
}
|
||||
|
||||
fn json_import_expiry_value(value: Option<&serde_json::Value>) -> Option<u64> {
|
||||
let value = value?;
|
||||
json_u64_value(Some(value)).or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok())
|
||||
.and_then(|value| u64::try_from(value.timestamp()).ok())
|
||||
})
|
||||
}
|
||||
|
||||
fn json_import_expiry_from_keys(
|
||||
object: &serde_json::Map<String, serde_json::Value>,
|
||||
keys: &[&str],
|
||||
) -> Option<u64> {
|
||||
keys.iter()
|
||||
.find_map(|key| json_import_expiry_value(object.get(*key)))
|
||||
}
|
||||
|
||||
fn grok_cookie_value(raw: &str, name: &str) -> Option<String> {
|
||||
raw.trim()
|
||||
.strip_prefix("Cookie:")
|
||||
@@ -204,6 +224,8 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
object
|
||||
.get("sso_token")
|
||||
.or_else(|| object.get("ssoToken"))
|
||||
.or_else(|| object.get("session_token"))
|
||||
.or_else(|| object.get("sessionToken"))
|
||||
.or(grok_token_alias),
|
||||
)
|
||||
.or_else(|| {
|
||||
@@ -254,7 +276,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
refresh_token
|
||||
};
|
||||
let expires_at =
|
||||
json_u64_value(object.get("expires_at").or_else(|| object.get("expiresAt")));
|
||||
json_import_expiry_from_keys(object, &["expires_at", "expiresAt", "expired"]);
|
||||
let account_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_id")
|
||||
@@ -660,6 +682,21 @@ mod tests {
|
||||
assert_eq!(entries[0].email.as_deref(), Some("u@example.com"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_common_chatgpt_web_json_aliases() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"chatgpt_web",
|
||||
r#"[{"session_token":"session-1","expired":"2030-01-01T00:00:00Z","chatgpt_account_id":"acc-1","chatgpt_plan_type":"plus"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("session-1"));
|
||||
assert_eq!(entries[0].expires_at, Some(1_893_456_000));
|
||||
assert_eq!(entries[0].account_id.as_deref(), Some("acc-1"));
|
||||
assert_eq!(entries[0].plan_type.as_deref(), Some("plus"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_plain_jwt_line_as_access_token() {
|
||||
let token = unsigned_jwt(json!({
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::provider_oauth_token_payload_expires_at_unix_secs;
|
||||
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::super::super::runtime::resolve_provider_oauth_runtime_endpoints;
|
||||
use super::super::super::runtime::{
|
||||
resolve_provider_oauth_runtime_endpoints,
|
||||
spawn_provider_oauth_account_state_refresh_after_update,
|
||||
};
|
||||
use super::super::super::state::{
|
||||
admin_provider_oauth_template, enrich_admin_provider_oauth_auth_config,
|
||||
is_fixed_provider_type_for_provider_oauth, json_non_empty_string,
|
||||
@@ -219,50 +221,20 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
));
|
||||
}
|
||||
|
||||
let mut account_state_recheck_attempted = false;
|
||||
let mut account_state_recheck_error = None::<String>;
|
||||
if provider_type == "codex" {
|
||||
if let Some(endpoint) = runtime_endpoint {
|
||||
let refreshed_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
.unwrap_or_else(|| key.clone());
|
||||
if let Some(result) = refresh_codex_provider_quota_locally(
|
||||
state,
|
||||
&provider,
|
||||
&endpoint,
|
||||
vec![refreshed_key],
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
account_state_recheck_attempted = true;
|
||||
let success = result
|
||||
.get("success")
|
||||
.and_then(serde_json::Value::as_u64)
|
||||
.unwrap_or(0);
|
||||
if success == 0 {
|
||||
account_state_recheck_error = result
|
||||
.get("results")
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.and_then(|results| results.first())
|
||||
.and_then(|value| value.get("message"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
provider.clone(),
|
||||
key_id.clone(),
|
||||
request_proxy.clone(),
|
||||
);
|
||||
|
||||
Ok(Json(json!({
|
||||
"provider_type": provider_type,
|
||||
"expires_at": expires_at,
|
||||
"has_refresh_token": refresh_token.is_some(),
|
||||
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
|
||||
"account_state_recheck_attempted": account_state_recheck_attempted,
|
||||
"account_state_recheck_error": account_state_recheck_error,
|
||||
"account_state_recheck_attempted": false,
|
||||
"account_state_recheck_error": serde_json::Value::Null,
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
@@ -828,13 +828,7 @@ async fn handle_admin_provider_oauth_windsurf_browser_device_poll(
|
||||
};
|
||||
callback_token.to_string()
|
||||
} else {
|
||||
let token = token.unwrap_or_default();
|
||||
if windsurf_raw_api_key(token).is_none() {
|
||||
return Ok(windsurf_browser_poll_error_response(
|
||||
"浏览器授权请提交包含 state 的回调 URL;纯 token 请使用导入授权",
|
||||
));
|
||||
}
|
||||
token.to_string()
|
||||
token.unwrap_or_default().to_string()
|
||||
};
|
||||
|
||||
let mut raw_credentials = serde_json::Map::new();
|
||||
|
||||
@@ -106,12 +106,21 @@ fn import_payload_string_any(
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn import_payload_u64(
|
||||
fn import_payload_u64_any(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
snake_case: &str,
|
||||
camel_case: &str,
|
||||
keys: &[&str],
|
||||
) -> Option<u64> {
|
||||
json_u64_value(payload.get(snake_case).or_else(|| payload.get(camel_case)))
|
||||
keys.iter().find_map(|key| {
|
||||
let value = payload.get(*key)?;
|
||||
json_u64_value(Some(value)).or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok())
|
||||
.and_then(|value| u64::try_from(value.timestamp()).ok())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn apply_single_import_hints(
|
||||
@@ -409,9 +418,17 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
||||
let access_token_input = import_payload_string_any(
|
||||
&raw_payload,
|
||||
&["access_token", "accessToken", "sso_token", "ssoToken"],
|
||||
&[
|
||||
"access_token",
|
||||
"accessToken",
|
||||
"sso_token",
|
||||
"ssoToken",
|
||||
"session_token",
|
||||
"sessionToken",
|
||||
],
|
||||
);
|
||||
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
|
||||
let imported_expires_at =
|
||||
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
|
||||
let name = raw_payload
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
@@ -622,8 +639,41 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::sanitize_windsurf_import_error;
|
||||
use super::{
|
||||
import_payload_string_any, import_payload_u64_any, sanitize_windsurf_import_error,
|
||||
};
|
||||
use aether_oauth::core::OAuthError;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn single_import_accepts_session_token_alias() {
|
||||
let payload = json!({
|
||||
"session_token": "session-1",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(
|
||||
import_payload_string_any(&payload, &["access_token", "session_token"]).as_deref(),
|
||||
Some("session-1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_import_accepts_iso_expired_alias() {
|
||||
let payload = json!({
|
||||
"expired": "2030-01-01T00:00:00Z",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(
|
||||
import_payload_u64_any(&payload, &["expires_at", "expiresAt", "expired"]),
|
||||
Some(1_893_456_000)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_import_error_redacts_http_body() {
|
||||
|
||||
+50
-40
@@ -81,48 +81,42 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
.await?;
|
||||
if provider_auto_remove_banned_keys(provider.config.as_ref()) {
|
||||
let now_unix_secs = helpers::unix_now_secs();
|
||||
let latest_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next();
|
||||
if latest_key.as_ref().is_some_and(|latest_key| {
|
||||
should_auto_remove_oauth_invalid_key(
|
||||
latest_key,
|
||||
None,
|
||||
false,
|
||||
now_unix_secs,
|
||||
)
|
||||
}) {
|
||||
state
|
||||
.clear_admin_provider_pool_cooldown(&provider.id, &key_id)
|
||||
.await;
|
||||
state
|
||||
.reset_admin_provider_pool_cost(&provider.id, &key_id)
|
||||
.await;
|
||||
if state.delete_provider_catalog_key(&key_id).await? {
|
||||
let deleted_key_ids = [key_id.clone()];
|
||||
state
|
||||
.cleanup_deleted_provider_catalog_refs(
|
||||
&provider.id,
|
||||
&[],
|
||||
&deleted_key_ids,
|
||||
let auto_removed = state
|
||||
.cleanup_provider_catalog_key_if_current(
|
||||
&provider,
|
||||
&key_id,
|
||||
|latest_key| {
|
||||
should_auto_remove_oauth_invalid_key(
|
||||
latest_key,
|
||||
Some(&failure_reason),
|
||||
false,
|
||||
now_unix_secs,
|
||||
)
|
||||
.await?;
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if auto_removed {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
}
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_failed_retained",
|
||||
"gateway manual provider oauth refresh failure retained key"
|
||||
);
|
||||
}
|
||||
}
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
@@ -171,9 +165,25 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
};
|
||||
|
||||
if !helpers::key_is_account_blocked(&key, OAUTH_ACCOUNT_BLOCK_PREFIX) {
|
||||
let _ = state
|
||||
let previous_oauth_refresh_issue =
|
||||
key.oauth_invalid_reason.as_deref().is_some_and(|reason| {
|
||||
reason.lines().map(str::trim).any(|line| {
|
||||
line.starts_with("[OAUTH_EXPIRED]") || line.starts_with("[REFRESH_FAILED]")
|
||||
})
|
||||
});
|
||||
let cleared = state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
|
||||
.await?;
|
||||
if cleared && previous_oauth_refresh_issue {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_fixed",
|
||||
"gateway manual provider oauth refresh cleared oauth invalid marker"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let refreshed_key = state
|
||||
|
||||
@@ -201,10 +201,10 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let mut updated = existing_key.clone();
|
||||
updated.is_active = true;
|
||||
updated.encrypted_api_key = Some(encrypted_api_key);
|
||||
updated.encrypted_auth_config = Some(encrypted_auth_config);
|
||||
updated.api_formats = provider_oauth_catalog_key_api_formats(provider_type, api_formats);
|
||||
updated.is_active = true;
|
||||
updated.expires_at_unix_secs = expires_at_unix_secs;
|
||||
updated.oauth_invalid_at_unix_secs = None;
|
||||
updated.oauth_invalid_reason = None;
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload, coerce_json_f64,
|
||||
coerce_json_string, default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
extract_execution_error_message, oauth_refresh_auto_removed_result,
|
||||
persist_provider_quota_refresh_state, quota_key_auto_removed,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
coerce_json_string, execute_provider_quota_plan, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -33,11 +33,10 @@ async fn execute_antigravity_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec = build_antigravity_pool_quota_request(
|
||||
&transport.key.id,
|
||||
&transport.endpoint.base_url,
|
||||
|
||||
@@ -1,16 +1,19 @@
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, extract_execution_error_message,
|
||||
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::quota::parse_chatgpt_web_conversation_init_response;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_contracts::{
|
||||
ExecutionResult, ProxySnapshot, ResolvedTransportProfile, TRANSPORT_BACKEND_BROWSER_WREQ,
|
||||
TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
@@ -18,10 +21,12 @@ use aether_provider_pool::{
|
||||
build_chatgpt_web_pool_quota_request, enrich_chatgpt_web_quota_metadata,
|
||||
normalize_chatgpt_web_image_quota_limit,
|
||||
};
|
||||
use base64::Engine as _;
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
|
||||
const CHATGPT_WEB_BROWSER_PROFILE: &str = "chrome143";
|
||||
|
||||
fn chatgpt_web_auth_config(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
@@ -67,26 +72,105 @@ async fn execute_chatgpt_web_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec =
|
||||
build_chatgpt_web_pool_quota_request(&transport.key.id, &endpoint.base_url, authorization);
|
||||
let resolved_transport_profile = state.resolve_transport_profile(transport);
|
||||
let plan = super::shared::build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
proxy,
|
||||
state.resolve_transport_profile(transport),
|
||||
chatgpt_web_quota_transport_profile(resolved_transport_profile.as_ref()),
|
||||
timeouts,
|
||||
);
|
||||
|
||||
execute_provider_quota_plan(state, transport, plan, "chatgpt_web").await
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_transport_profile(
|
||||
transport_profile: Option<&ResolvedTransportProfile>,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
match transport_profile {
|
||||
Some(profile)
|
||||
if profile
|
||||
.backend
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ) =>
|
||||
{
|
||||
Some(profile.clone())
|
||||
}
|
||||
_ => Some(default_chatgpt_web_quota_transport_profile()),
|
||||
}
|
||||
}
|
||||
|
||||
fn default_chatgpt_web_quota_transport_profile() -> ResolvedTransportProfile {
|
||||
ResolvedTransportProfile {
|
||||
profile_id: CHATGPT_WEB_BROWSER_PROFILE.to_string(),
|
||||
backend: TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: Some(json!({
|
||||
"browser_profile": CHATGPT_WEB_BROWSER_PROFILE,
|
||||
"source": "chatgpt_web_quota_default",
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_error_detail(result: &ExecutionResult) -> Option<String> {
|
||||
extract_execution_error_message(result).or_else(|| {
|
||||
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
|
||||
let decoded = base64::engine::general_purpose::STANDARD
|
||||
.decode(body)
|
||||
.ok()?;
|
||||
let text = String::from_utf8_lossy(&decoded).trim().to_string();
|
||||
(!text.is_empty()).then_some(text)
|
||||
})
|
||||
}
|
||||
|
||||
fn chatgpt_web_is_structured_account_block(message: &str) -> bool {
|
||||
let lowered = message.to_ascii_lowercase();
|
||||
[
|
||||
"account has been disabled",
|
||||
"account disabled",
|
||||
"account has been deactivated",
|
||||
"account_deactivated",
|
||||
"account deactivated",
|
||||
"organization has been disabled",
|
||||
"organization_disabled",
|
||||
"deactivated_workspace",
|
||||
"account suspended",
|
||||
"account banned",
|
||||
"account_block",
|
||||
"account blocked",
|
||||
"访问被禁止",
|
||||
"账户访问被禁止",
|
||||
"账户已封禁",
|
||||
"封禁",
|
||||
"封号",
|
||||
"被封",
|
||||
]
|
||||
.iter()
|
||||
.any(|keyword| lowered.contains(keyword))
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_403_refresh_failed_reason(message: Option<&str>) -> String {
|
||||
let detail = message
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.filter(|value| !value.contains('<'))
|
||||
.unwrap_or("ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制");
|
||||
format!("{OAUTH_REFRESH_FAILED_PREFIX}{detail}")
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<&str>) -> String {
|
||||
let message = upstream_message.unwrap_or_default().trim();
|
||||
if status_code == 403 && !chatgpt_web_is_structured_account_block(message) {
|
||||
return chatgpt_web_quota_403_refresh_failed_reason(upstream_message);
|
||||
}
|
||||
let detail = if message.is_empty() {
|
||||
match status_code {
|
||||
401 => "ChatGPT Web Token 无效或已过期",
|
||||
@@ -103,6 +187,19 @@ fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<&
|
||||
}
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_result_message(reason: &str) -> String {
|
||||
for prefix in [
|
||||
OAUTH_REFRESH_FAILED_PREFIX,
|
||||
OAUTH_EXPIRED_PREFIX,
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX,
|
||||
] {
|
||||
if let Some(message) = reason.strip_prefix(prefix) {
|
||||
return message.trim().to_string();
|
||||
}
|
||||
}
|
||||
reason.trim().to_string()
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -216,8 +313,20 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
message = Some("响应中未包含 ChatGPT Web 生图限额信息".to_string());
|
||||
}
|
||||
} else {
|
||||
let err_msg = extract_execution_error_message(&result);
|
||||
message = Some(match err_msg.as_deref() {
|
||||
let err_msg = chatgpt_web_quota_error_detail(&result);
|
||||
let invalid_reason = if matches!(result.status_code, 401 | 403) {
|
||||
Some(chatgpt_web_quota_invalid_reason(
|
||||
result.status_code,
|
||||
err_msg.as_deref(),
|
||||
))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let display_detail = invalid_reason
|
||||
.as_deref()
|
||||
.map(chatgpt_web_quota_result_message)
|
||||
.or_else(|| err_msg.clone());
|
||||
message = Some(match display_detail.as_deref() {
|
||||
Some(detail) if !detail.is_empty() => {
|
||||
format!(
|
||||
"conversation/init 返回状态码 {}: {}",
|
||||
@@ -229,12 +338,14 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
|
||||
if matches!(result.status_code, 401 | 403) {
|
||||
oauth_invalid_at_unix_secs = Some(now_unix_secs);
|
||||
oauth_invalid_reason = Some(chatgpt_web_quota_invalid_reason(
|
||||
result.status_code,
|
||||
err_msg.as_deref(),
|
||||
));
|
||||
oauth_invalid_reason = invalid_reason;
|
||||
status = if result.status_code == 401 {
|
||||
"auth_invalid".to_string()
|
||||
} else if oauth_invalid_reason
|
||||
.as_deref()
|
||||
.is_some_and(|reason| reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX))
|
||||
{
|
||||
"refresh_failed".to_string()
|
||||
} else {
|
||||
"forbidden".to_string()
|
||||
};
|
||||
@@ -303,3 +414,81 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
"auto_removed": auto_removed_count,
|
||||
})))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aether_contracts::{ResponseBody, TRANSPORT_BACKEND_REQWEST_RUSTLS};
|
||||
use base64::Engine as _;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn quota_refresh_defaults_to_browser_wreq_transport() {
|
||||
let profile = chatgpt_web_quota_transport_profile(None).expect("transport profile");
|
||||
|
||||
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
|
||||
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
|
||||
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_AUTO);
|
||||
assert_eq!(profile.pool_scope, TRANSPORT_POOL_SCOPE_KEY);
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("browser_profile"))
|
||||
.and_then(serde_json::Value::as_str),
|
||||
Some(CHATGPT_WEB_BROWSER_PROFILE)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_refresh_overrides_non_browser_transport() {
|
||||
let reqwest_profile = ResolvedTransportProfile {
|
||||
profile_id: "chrome_136".to_string(),
|
||||
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: None,
|
||||
};
|
||||
|
||||
let profile =
|
||||
chatgpt_web_quota_transport_profile(Some(&reqwest_profile)).expect("transport profile");
|
||||
|
||||
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
|
||||
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn browser_challenge_403_is_not_account_block() {
|
||||
let body = "<!DOCTYPE html><html><head><title>Just a moment...</title></head><body>Cloudflare</body></html>";
|
||||
let result = ExecutionResult {
|
||||
request_id: "chatgpt-web-quota:test".to_string(),
|
||||
candidate_id: None,
|
||||
status_code: 403,
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
};
|
||||
|
||||
let detail = chatgpt_web_quota_error_detail(&result).expect("html body should decode");
|
||||
let reason = chatgpt_web_quota_invalid_reason(result.status_code, Some(&detail));
|
||||
|
||||
assert!(reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX));
|
||||
assert!(!reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX));
|
||||
assert_eq!(
|
||||
chatgpt_web_quota_result_message(&reason),
|
||||
"ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_account_block_403_remains_account_block() {
|
||||
let reason = chatgpt_web_quota_invalid_reason(403, Some("account has been deactivated"));
|
||||
|
||||
assert!(reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -44,6 +44,15 @@ fn merge_codex_quota_metadata(
|
||||
serde_json::Value::Object(merged)
|
||||
}
|
||||
|
||||
fn codex_oauth_refresh_issue_reason(reason: Option<&str>) -> bool {
|
||||
reason.is_some_and(|reason| {
|
||||
reason
|
||||
.lines()
|
||||
.map(str::trim)
|
||||
.any(|line| line.starts_with("[OAUTH_EXPIRED]") || line.starts_with("[REFRESH_FAILED]"))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -56,8 +65,13 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
let mut success_count = 0usize;
|
||||
let mut failed_count = 0usize;
|
||||
let mut auto_removed_count = 0usize;
|
||||
let mut refresh_fixed_count = 0usize;
|
||||
let mut refresh_failed_retained_count = 0usize;
|
||||
let mut auto_removed_hard_banned_count = 0usize;
|
||||
|
||||
for key in keys {
|
||||
let had_oauth_refresh_issue =
|
||||
codex_oauth_refresh_issue_reason(key.oauth_invalid_reason.as_deref());
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
@@ -276,13 +290,9 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
}
|
||||
|
||||
let auto_removed = auto_remove_abnormal_keys
|
||||
let auto_remove_candidate = auto_remove_abnormal_keys
|
||||
&& should_auto_remove_structured_reason(oauth_invalid_reason.as_deref());
|
||||
if auto_removed {
|
||||
if state.delete_provider_catalog_key(&key.id).await? {
|
||||
auto_removed_count += 1;
|
||||
}
|
||||
} else if !persist_provider_quota_refresh_state(
|
||||
let persisted = persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
@@ -290,8 +300,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
oauth_invalid_reason.clone(),
|
||||
None,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
.await?;
|
||||
if !persisted {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
@@ -301,6 +311,29 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
let auto_removed = if auto_remove_candidate {
|
||||
state
|
||||
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
|
||||
should_auto_remove_structured_reason(latest_key.oauth_invalid_reason.as_deref())
|
||||
})
|
||||
.await?
|
||||
} else {
|
||||
false
|
||||
};
|
||||
if auto_removed {
|
||||
auto_removed_count += 1;
|
||||
auto_removed_hard_banned_count += 1;
|
||||
}
|
||||
let refresh_fixed =
|
||||
status == "success" && had_oauth_refresh_issue && oauth_invalid_reason.is_none();
|
||||
if refresh_fixed {
|
||||
refresh_fixed_count += 1;
|
||||
}
|
||||
let refresh_failed_retained =
|
||||
status != "success" && oauth_invalid_reason.is_some() && !auto_removed;
|
||||
if refresh_failed_retained {
|
||||
refresh_failed_retained_count += 1;
|
||||
}
|
||||
|
||||
if status == "success" {
|
||||
success_count += 1;
|
||||
@@ -336,6 +369,13 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
if auto_removed {
|
||||
payload.insert("auto_removed".to_string(), json!(true));
|
||||
payload.insert("auto_removed_hard_banned".to_string(), json!(true));
|
||||
}
|
||||
if refresh_fixed {
|
||||
payload.insert("refresh_fixed".to_string(), json!(true));
|
||||
}
|
||||
if refresh_failed_retained {
|
||||
payload.insert("refresh_failed_retained".to_string(), json!(true));
|
||||
}
|
||||
results.push(serde_json::Value::Object(payload));
|
||||
}
|
||||
@@ -346,5 +386,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
"total": results.len(),
|
||||
"results": results,
|
||||
"auto_removed": auto_removed_count,
|
||||
"refresh_fixed": refresh_fixed_count,
|
||||
"refresh_failed_retained": refresh_failed_retained_count,
|
||||
"auto_removed_hard_banned": auto_removed_hard_banned_count,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::super::shared::{
|
||||
build_provider_quota_execution_plan, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, ProviderQuotaExecutionOutcome,
|
||||
build_provider_quota_execution_plan, execute_provider_quota_plan,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -38,11 +38,10 @@ pub(super) async fn execute_codex_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, extract_execution_error_message,
|
||||
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
|
||||
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
@@ -245,11 +244,10 @@ async fn execute_grok_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let transport_profile = state.resolve_transport_profile(transport);
|
||||
let base_url = grok_base_url(endpoint);
|
||||
let headers = build_grok_quota_headers(
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::super::shared::{
|
||||
build_provider_quota_execution_plan, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, ProviderQuotaExecutionOutcome,
|
||||
build_provider_quota_execution_plan, execute_provider_quota_plan,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{
|
||||
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminKiroRequestAuth,
|
||||
@@ -23,11 +23,10 @@ pub(super) async fn execute_kiro_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec = build_kiro_pool_quota_request(
|
||||
&transport.key.id,
|
||||
&KiroPoolQuotaAuthInput {
|
||||
|
||||
@@ -46,6 +46,23 @@ pub(super) fn default_provider_quota_execution_timeouts(
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider_quota_execution_timeouts(
|
||||
configured: Option<ExecutionTimeouts>,
|
||||
proxy: Option<&ProxySnapshot>,
|
||||
) -> ExecutionTimeouts {
|
||||
let defaults = default_provider_quota_execution_timeouts(proxy);
|
||||
let Some(mut timeouts) = configured else {
|
||||
return defaults;
|
||||
};
|
||||
timeouts.connect_ms = timeouts.connect_ms.or(defaults.connect_ms);
|
||||
timeouts.read_ms = timeouts.read_ms.or(defaults.read_ms);
|
||||
timeouts.write_ms = timeouts.write_ms.or(defaults.write_ms);
|
||||
timeouts.pool_ms = timeouts.pool_ms.or(defaults.pool_ms);
|
||||
timeouts.total_ms = timeouts.total_ms.or(defaults.total_ms);
|
||||
timeouts.first_byte_ms = timeouts.first_byte_ms.or(defaults.first_byte_ms);
|
||||
timeouts
|
||||
}
|
||||
|
||||
pub(crate) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
|
||||
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
|
||||
}
|
||||
@@ -317,12 +334,7 @@ pub(super) async fn execute_provider_quota_plan(
|
||||
match state.execute_execution_runtime_sync_plan(None, &plan).await {
|
||||
Ok(result) => Ok(ProviderQuotaExecutionOutcome::Response(result)),
|
||||
Err(err) => {
|
||||
let error = match err {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
};
|
||||
let error = err.into_message();
|
||||
let proxy_node_id = plan
|
||||
.proxy
|
||||
.as_ref()
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload,
|
||||
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload, execute_provider_quota_plan,
|
||||
extract_execution_error_message, persist_provider_quota_refresh_state,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
quota_refresh_success_invalid_state, resolve_provider_quota_execution_timeouts,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -33,11 +33,10 @@ async fn execute_windsurf_probe_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
|
||||
@@ -301,12 +301,7 @@ fn admin_provider_ops_decode_response_bytes(
|
||||
}
|
||||
|
||||
fn admin_provider_ops_gateway_error_message(error: GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
error.into_message()
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_verify_execution_error_message(error: &str) -> String {
|
||||
|
||||
@@ -16,7 +16,12 @@ use aether_runtime_state::{DataLayerError, RuntimeState};
|
||||
use futures_util::future::join_all;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use tracing::warn;
|
||||
use tracing::{info, warn};
|
||||
|
||||
const DEFAULT_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT: usize = 512;
|
||||
const MAX_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT: usize = 10_000;
|
||||
const POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT_ENV: &str =
|
||||
"AETHER_GATEWAY_ADMIN_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT";
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
SystemTime::now()
|
||||
@@ -29,6 +34,20 @@ fn should_load_active_probe_members(pool_config: &AdminProviderPoolConfig) -> bo
|
||||
pool_config.probing_enabled
|
||||
}
|
||||
|
||||
fn pool_runtime_window_metric_key_limit() -> usize {
|
||||
std::env::var(POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT_ENV)
|
||||
.ok()
|
||||
.and_then(|value| value.trim().parse::<usize>().ok())
|
||||
.filter(|value| *value > 0)
|
||||
.unwrap_or(DEFAULT_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT)
|
||||
.clamp(1, MAX_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT)
|
||||
}
|
||||
|
||||
fn bounded_runtime_window_metric_key_ids(key_ids: &[String], limit: usize) -> &[String] {
|
||||
let end = key_ids.len().min(limit.max(1));
|
||||
&key_ids[..end]
|
||||
}
|
||||
|
||||
pub(crate) async fn read_admin_provider_pool_cooldown_counts(
|
||||
runtime: &RuntimeState,
|
||||
provider_ids: &[String],
|
||||
@@ -54,8 +73,21 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
) -> AdminProviderPoolRuntimeState {
|
||||
let mut state = AdminProviderPoolRuntimeState::default();
|
||||
let cooldown_keys = pool_cooldown_keys(provider_id, key_ids);
|
||||
let cost_keys = pool_cost_keys(provider_id, key_ids);
|
||||
let latency_keys = pool_latency_keys(provider_id, key_ids);
|
||||
let metric_key_limit = pool_runtime_window_metric_key_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!(
|
||||
event_name = "admin_pool_runtime_window_metrics_truncated",
|
||||
log_type = "event",
|
||||
provider_id,
|
||||
total_key_count = key_ids.len(),
|
||||
scanned_key_count = metric_key_ids.len(),
|
||||
metric_key_limit,
|
||||
"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 sticky_sessions_enabled = pool_config.sticky_session_ttl_seconds > 0
|
||||
&& admin_provider_pool_cache_affinity_enabled(pool_config);
|
||||
|
||||
@@ -179,7 +211,7 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
.map(|cost_key| runtime.score_range_by_min(cost_key, cost_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(cost_results) {
|
||||
for (key_id, members) in metric_key_ids.iter().zip(cost_results) {
|
||||
let total = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
@@ -197,7 +229,7 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
.map(|latency_key| runtime.score_range_by_min(latency_key, latency_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(latency_results) {
|
||||
for (key_id, members) in metric_key_ids.iter().zip(latency_results) {
|
||||
let samples = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
@@ -265,3 +297,30 @@ pub(crate) async fn read_admin_provider_pool_key_cooldown_reason(
|
||||
.kv_get(&pool_cooldown_key(provider_id, key_id))
|
||||
.await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::bounded_runtime_window_metric_key_ids;
|
||||
|
||||
#[test]
|
||||
fn runtime_window_metric_key_ids_are_bounded() {
|
||||
let key_ids = vec![
|
||||
"key-1".to_string(),
|
||||
"key-2".to_string(),
|
||||
"key-3".to_string(),
|
||||
];
|
||||
|
||||
let bounded = bounded_runtime_window_metric_key_ids(&key_ids, 2);
|
||||
|
||||
assert_eq!(bounded, &key_ids[..2]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_window_metric_key_ids_keep_at_least_one_key() {
|
||||
let key_ids = vec!["key-1".to_string(), "key-2".to_string()];
|
||||
|
||||
let bounded = bounded_runtime_window_metric_key_ids(&key_ids, 0);
|
||||
|
||||
assert_eq!(bounded, &key_ids[..1]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_data_contracts::repository::usage::StoredProviderApiKeyWindowUsageSummary;
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -933,19 +934,14 @@ fn admin_pool_health_score(key: &StoredProviderCatalogKey) -> f64 {
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_pool_circuit_breaker_open(key: &StoredProviderCatalogKey) -> bool {
|
||||
fn admin_pool_circuit_breaker_open(key: &StoredProviderCatalogKey, now_unix_secs: u64) -> bool {
|
||||
key.circuit_breaker_by_format
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.map(|formats| {
|
||||
formats
|
||||
.values()
|
||||
.filter_map(serde_json::Value::as_object)
|
||||
.any(|item| {
|
||||
item.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.any(|item| provider_key_circuit_payload_is_active_open_at(item, now_unix_secs))
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
@@ -1044,7 +1040,7 @@ pub(super) fn build_admin_pool_key_payload(
|
||||
.as_ref()
|
||||
.and_then(|_| runtime.cooldown_ttl_by_key.get(&key.id).copied());
|
||||
let health_score = admin_pool_health_score(key);
|
||||
let circuit_breaker_open = admin_pool_circuit_breaker_open(key);
|
||||
let circuit_breaker_open = admin_pool_circuit_breaker_open(key, now_unix_secs);
|
||||
let auth_semantics = provider_key_auth_semantics(key, provider_type);
|
||||
let account_quota_exhausted = pool_config
|
||||
.as_ref()
|
||||
|
||||
@@ -99,6 +99,7 @@ static PROVIDER_QUERY_POOL_LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
struct ProviderQueryKeyFetchResult {
|
||||
models: Vec<Value>,
|
||||
error: Option<String>,
|
||||
warning: Option<String>,
|
||||
from_cache: bool,
|
||||
has_success: bool,
|
||||
}
|
||||
@@ -288,6 +289,7 @@ fn provider_query_codex_preset_fallback(
|
||||
Some(ProviderQueryKeyFetchResult {
|
||||
models: aggregate_models_for_cache(&models),
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
})
|
||||
@@ -427,6 +429,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models,
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: true,
|
||||
has_success: true,
|
||||
});
|
||||
@@ -444,6 +447,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models,
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
});
|
||||
@@ -451,6 +455,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -477,6 +482,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -492,6 +498,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -528,18 +535,25 @@ async fn provider_query_fetch_models_for_key(
|
||||
}
|
||||
}
|
||||
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
} else {
|
||||
let has_models = !unique_models.is_empty();
|
||||
let mut error = if !has_models && !all_errors.is_empty() {
|
||||
Some(all_errors.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if unique_models.is_empty() && error.is_none() {
|
||||
let warning = if has_models && !all_errors.is_empty() {
|
||||
Some(all_errors.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if !has_models && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL.to_string());
|
||||
}
|
||||
|
||||
Ok(ProviderQueryKeyFetchResult {
|
||||
models: provider_query_filter_models_for_key(provider, key, unique_models),
|
||||
error,
|
||||
warning,
|
||||
from_cache: false,
|
||||
has_success: outcome.has_success,
|
||||
})
|
||||
@@ -600,6 +614,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": result.error,
|
||||
"warning": result.warning,
|
||||
"from_cache": result.from_cache,
|
||||
},
|
||||
"provider": provider_query_provider_payload(&provider),
|
||||
@@ -632,6 +647,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": serde_json::Value::Null,
|
||||
"warning": serde_json::Value::Null,
|
||||
"from_cache": true,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": active_key_count,
|
||||
@@ -655,6 +671,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
|
||||
let mut all_models = Vec::new();
|
||||
let mut all_errors = Vec::new();
|
||||
let mut all_warnings = Vec::new();
|
||||
let mut cache_hit_count = 0usize;
|
||||
let mut fetch_count = 0usize;
|
||||
for key in &ordered_keys {
|
||||
@@ -669,6 +686,13 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
error
|
||||
));
|
||||
}
|
||||
if let Some(warning) = result.warning {
|
||||
all_warnings.push(format!(
|
||||
"Key {}: {}",
|
||||
provider_query_key_display_name(key),
|
||||
warning
|
||||
));
|
||||
}
|
||||
if result.from_cache {
|
||||
cache_hit_count += 1;
|
||||
} else {
|
||||
@@ -694,10 +718,17 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
provider_query_write_provider_cached_models(state, &provider.id, &models).await;
|
||||
}
|
||||
let success = !models.is_empty();
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
let mut all_issues = all_errors;
|
||||
all_issues.extend(all_warnings);
|
||||
let mut error = if !success && !all_issues.is_empty() {
|
||||
Some(all_issues.join("; "))
|
||||
} else {
|
||||
Some(all_errors.join("; "))
|
||||
None
|
||||
};
|
||||
let warning = if success && !all_issues.is_empty() {
|
||||
Some(all_issues.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if !success && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
|
||||
@@ -709,6 +740,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": error,
|
||||
"warning": warning,
|
||||
"from_cache": fetch_count == 0 && cache_hit_count > 0,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": cache_hit_count,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use super::super::payload::{
|
||||
provider_query_extract_api_key_id, provider_query_extract_force_refresh,
|
||||
provider_query_extract_api_key_ids, provider_query_extract_force_refresh,
|
||||
provider_query_extract_model, provider_query_extract_provider_id,
|
||||
provider_query_extract_request_id,
|
||||
};
|
||||
@@ -68,6 +68,7 @@ use aether_model_fetch::{
|
||||
aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
|
||||
preset_models_for_provider, selected_models_fetch_endpoints,
|
||||
};
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
http::{self, HeaderMap, HeaderName, HeaderValue},
|
||||
@@ -122,6 +123,7 @@ const ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL: &str =
|
||||
"No active endpoint or API key found";
|
||||
const ADMIN_PROVIDER_QUERY_INVALID_MAPPED_MODEL_DETAIL: &str =
|
||||
"mapped_model_name is not valid for the selected model and endpoint";
|
||||
const PROVIDER_QUERY_KEY_MODEL_NOT_ALLOWED_SKIP_REASON: &str = "key_model_not_allowed";
|
||||
const ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX: &str = "upstream_models_provider:";
|
||||
const DEFAULT_PROVIDER_QUERY_TEST_MESSAGE: &str = "Hello! This is a test message.";
|
||||
static PROVIDER_QUERY_POOL_LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
@@ -859,7 +861,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
selected_key_id: Option<&str>,
|
||||
selected_key_ids: Option<&BTreeSet<String>>,
|
||||
) -> Option<StoredProviderCatalogEndpoint> {
|
||||
for priority in 0..=2 {
|
||||
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
|
||||
@@ -872,7 +874,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
}
|
||||
for key in keys {
|
||||
if !key.is_active
|
||||
|| selected_key_id.is_some_and(|value| value != key.id.as_str())
|
||||
|| !provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
|
||||
|| !provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
@@ -904,7 +906,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
endpoint.is_active
|
||||
&& keys.iter().any(|key| {
|
||||
key.is_active
|
||||
&& selected_key_id.is_none_or(|value| value == key.id.as_str())
|
||||
&& provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
|
||||
&& provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
@@ -916,10 +918,56 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn provider_query_selected_key_ids_allow_key(
|
||||
selected_key_ids: Option<&BTreeSet<String>>,
|
||||
key_id: &str,
|
||||
) -> bool {
|
||||
selected_key_ids.is_none_or(|ids| ids.contains(key_id))
|
||||
}
|
||||
|
||||
fn provider_query_selected_key_ids_all_exist(
|
||||
selected_key_ids: &BTreeSet<String>,
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
) -> bool {
|
||||
selected_key_ids
|
||||
.iter()
|
||||
.all(|id| keys.iter().any(|key| key.id == *id))
|
||||
}
|
||||
|
||||
fn provider_query_model_name_matches(left: &str, right: &str) -> bool {
|
||||
let left = left.trim();
|
||||
let right = right.trim();
|
||||
!left.is_empty() && !right.is_empty() && left.eq_ignore_ascii_case(right)
|
||||
}
|
||||
|
||||
fn provider_query_key_allows_effective_test_model(
|
||||
key: &StoredProviderCatalogKey,
|
||||
requested_model: &str,
|
||||
effective_model: &str,
|
||||
) -> bool {
|
||||
let allowed_models = json_string_list(key.allowed_models.as_ref());
|
||||
if key.allowed_models.is_none() || allowed_models.is_empty() {
|
||||
return true;
|
||||
}
|
||||
|
||||
let requested_base_model = crate::ai_serving::model_directive_base_model(requested_model);
|
||||
allowed_models
|
||||
.iter()
|
||||
.map(String::as_str)
|
||||
.any(|allowed_model| {
|
||||
provider_query_model_name_matches(allowed_model, requested_model)
|
||||
|| provider_query_model_name_matches(allowed_model, effective_model)
|
||||
|| requested_base_model.as_deref().is_some_and(|base_model| {
|
||||
provider_query_model_name_matches(allowed_model, base_model)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_test_key_sort_key(
|
||||
provider_type: &str,
|
||||
key: &StoredProviderCatalogKey,
|
||||
endpoint_api_format: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> (u8, u8, i32, u64, i32) {
|
||||
let quota_exhausted =
|
||||
admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type);
|
||||
@@ -928,10 +976,7 @@ fn provider_query_test_key_sort_key(
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get(endpoint_api_format))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("open"))
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
.is_some_and(|value| provider_key_circuit_payload_is_active_open_at(value, now_unix_secs));
|
||||
let health_score = key
|
||||
.health_by_format
|
||||
.as_ref()
|
||||
@@ -1263,7 +1308,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
)
|
||||
})?;
|
||||
let selected_key_id = provider_query_extract_api_key_id(payload);
|
||||
let selected_key_ids = provider_query_extract_api_key_ids(payload);
|
||||
let requested_endpoint_id = provider_query_extract_endpoint_id(payload);
|
||||
let requested_api_format = provider_query_extract_api_format(payload);
|
||||
let endpoint = if requested_endpoint_id.is_none()
|
||||
@@ -1275,7 +1320,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider,
|
||||
&endpoints,
|
||||
&all_keys,
|
||||
selected_key_id.as_deref(),
|
||||
selected_key_ids.as_ref(),
|
||||
)
|
||||
.await
|
||||
.ok_or_else(|| {
|
||||
@@ -1308,22 +1353,11 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(api_key_id) = selected_key_id.as_deref() {
|
||||
let Some(key) = all_keys.iter().find(|key| key.id == api_key_id) else {
|
||||
if let Some(selected_key_ids) = selected_key_ids.as_ref() {
|
||||
if !provider_query_selected_key_ids_all_exist(selected_key_ids, &all_keys) {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
|
||||
));
|
||||
};
|
||||
if !key.is_active
|
||||
|| !provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
&endpoint.api_format,
|
||||
)
|
||||
{
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1378,20 +1412,43 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
.unwrap_or(requested_model.clone())
|
||||
};
|
||||
|
||||
let mut keys = all_keys
|
||||
let now_unix_secs = current_unix_ms() / 1000;
|
||||
let mut keys = Vec::new();
|
||||
let mut model_skipped_candidates = Vec::new();
|
||||
|
||||
for key in all_keys
|
||||
.into_iter()
|
||||
.filter(|key| key.is_active)
|
||||
.filter(|key| {
|
||||
selected_key_id
|
||||
.as_deref()
|
||||
.is_none_or(|value| value == key.id.as_str())
|
||||
})
|
||||
.filter(|key| provider_query_selected_key_ids_allow_key(selected_key_ids.as_ref(), &key.id))
|
||||
.filter(|key| {
|
||||
provider_query_key_supports_endpoint(key, &provider.provider_type, &endpoint.api_format)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
{
|
||||
if provider_query_key_allows_effective_test_model(&key, &requested_model, &effective_model)
|
||||
{
|
||||
keys.push(key);
|
||||
} else {
|
||||
model_skipped_candidates.push(ProviderQueryTestCandidate {
|
||||
endpoint: endpoint.clone(),
|
||||
key,
|
||||
effective_model: effective_model.clone(),
|
||||
scheduler_skip_reason: Some(
|
||||
PROVIDER_QUERY_KEY_MODEL_NOT_ALLOWED_SKIP_REASON.to_string(),
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let candidates = if test_mode.eq_ignore_ascii_case("pool") {
|
||||
model_skipped_candidates.sort_by_key(|candidate| {
|
||||
provider_query_test_key_sort_key(
|
||||
provider.provider_type.as_str(),
|
||||
&candidate.key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
|
||||
let scheduled_candidates = if test_mode.eq_ignore_ascii_case("pool") {
|
||||
if let Some(pool_config) =
|
||||
admin_provider_pool_config_from_config_value(provider.config.as_ref())
|
||||
{
|
||||
@@ -1411,6 +1468,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider.provider_type.as_str(),
|
||||
key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
keys.into_iter()
|
||||
@@ -1428,6 +1486,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider.provider_type.as_str(),
|
||||
key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
keys.into_iter()
|
||||
@@ -1439,6 +1498,8 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
let mut candidates = model_skipped_candidates;
|
||||
candidates.extend(scheduled_candidates);
|
||||
|
||||
if candidates.is_empty() {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
|
||||
@@ -64,6 +64,68 @@ fn sample_openai_image_transport(provider_type: &str) -> AdminGatewayProviderTra
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_catalog_key_with_allowed_models(
|
||||
allowed_models: Option<serde_json::Value>,
|
||||
) -> aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey {
|
||||
let mut key =
|
||||
aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"key".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("sample provider key should build");
|
||||
key.allowed_models = allowed_models;
|
||||
key
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_allows_keys_without_model_restrictions() {
|
||||
let unrestricted = sample_catalog_key_with_allowed_models(None);
|
||||
let empty = sample_catalog_key_with_allowed_models(Some(json!([])));
|
||||
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&unrestricted,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&empty,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_filters_key_disallowed_for_requested_model() {
|
||||
let key = sample_catalog_key_with_allowed_models(Some(json!(["model-a"])));
|
||||
|
||||
assert!(!provider_query_key_allows_effective_test_model(
|
||||
&key,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_allows_key_for_requested_or_mapped_model() {
|
||||
let requested_allowed = sample_catalog_key_with_allowed_models(Some(json!(["model-b"])));
|
||||
let mapped_allowed = sample_catalog_key_with_allowed_models(Some(json!(["MODEL-B-UPSTREAM"])));
|
||||
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&requested_allowed,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&mapped_allowed,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_preserves_custom_model() {
|
||||
let payload = json!({
|
||||
@@ -232,6 +294,27 @@ fn provider_query_request_body_model_uses_non_empty_string_only() {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_extracts_multiple_selected_key_ids() {
|
||||
let payload = json!({
|
||||
"api_key_ids": [" key-b ", "", "key-a", "key-b"],
|
||||
"api_key_id": "key-c"
|
||||
});
|
||||
|
||||
let ids = provider_query_extract_api_key_ids(&payload)
|
||||
.expect("non-empty key selection should be extracted")
|
||||
.into_iter()
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(ids, vec!["key-a", "key-b", "key-c"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_empty_selected_key_ids_keep_default_selection() {
|
||||
assert!(provider_query_extract_api_key_ids(&json!({})).is_none());
|
||||
assert!(provider_query_extract_api_key_ids(&json!({ "api_key_ids": [] })).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() {
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use axum::body::Bytes;
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
pub(crate) fn parse_admin_provider_query_body(
|
||||
request_body: Option<&Bytes>,
|
||||
@@ -36,6 +37,47 @@ pub(crate) fn provider_query_extract_api_key_id(payload: &serde_json::Value) ->
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn provider_query_insert_api_key_id(ids: &mut BTreeSet<String>, value: &str) {
|
||||
let value = value.trim();
|
||||
if !value.is_empty() {
|
||||
ids.insert(value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn provider_query_extract_api_key_ids(
|
||||
payload: &serde_json::Value,
|
||||
) -> Option<BTreeSet<String>> {
|
||||
let mut ids = BTreeSet::new();
|
||||
|
||||
if let Some(value) = payload
|
||||
.get("api_key_ids")
|
||||
.or_else(|| payload.get("provider_key_ids"))
|
||||
.or_else(|| payload.get("key_ids"))
|
||||
{
|
||||
match value {
|
||||
serde_json::Value::Array(items) => {
|
||||
for item in items {
|
||||
if let Some(value) = item.as_str() {
|
||||
provider_query_insert_api_key_id(&mut ids, value);
|
||||
}
|
||||
}
|
||||
}
|
||||
serde_json::Value::String(value) => {
|
||||
for item in value.split(',') {
|
||||
provider_query_insert_api_key_id(&mut ids, item);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(api_key_id) = provider_query_extract_api_key_id(payload) {
|
||||
ids.insert(api_key_id);
|
||||
}
|
||||
|
||||
(!ids.is_empty()).then_some(ids)
|
||||
}
|
||||
|
||||
pub(crate) fn provider_query_extract_force_refresh(payload: &serde_json::Value) -> bool {
|
||||
payload
|
||||
.get("force_refresh")
|
||||
|
||||
@@ -662,10 +662,5 @@ fn admin_provider_oauth_decode_response_bytes(
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_gateway_error_message(error: GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
error.into_message()
|
||||
}
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
use super::*;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -193,8 +196,6 @@ impl<'a> AdminAppState<'a> {
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
|
||||
let Some(provider) = self
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id.to_string()))
|
||||
.await?
|
||||
@@ -208,6 +209,30 @@ impl<'a> AdminAppState<'a> {
|
||||
.into_response());
|
||||
};
|
||||
|
||||
let affected = self
|
||||
.cleanup_known_banned_provider_catalog_keys(&provider)
|
||||
.await?;
|
||||
if affected == 0 {
|
||||
return Ok(Json(
|
||||
aether_admin::provider::pool::build_admin_pool_cleanup_empty_payload(
|
||||
"未发现可清理的异常账号",
|
||||
),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
|
||||
Ok(
|
||||
Json(aether_admin::provider::pool::build_admin_pool_cleanup_result_payload(affected))
|
||||
.into_response(),
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_known_banned_provider_catalog_keys(
|
||||
&self,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Result<usize, GatewayError> {
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
|
||||
let banned_keys = self
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await?
|
||||
@@ -215,12 +240,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.filter(admin_provider_pool_pure::admin_pool_key_is_known_banned)
|
||||
.collect::<Vec<_>>();
|
||||
if banned_keys.is_empty() {
|
||||
return Ok(Json(
|
||||
admin_provider_pool_pure::build_admin_pool_cleanup_empty_payload(
|
||||
"未发现可清理的异常账号",
|
||||
),
|
||||
)
|
||||
.into_response());
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
let deleted_key_ids = banned_keys
|
||||
@@ -243,10 +263,42 @@ impl<'a> AdminAppState<'a> {
|
||||
self.cleanup_deleted_provider_catalog_refs(&provider.id, &[], &deleted_key_ids)
|
||||
.await?;
|
||||
|
||||
Ok(
|
||||
Json(admin_provider_pool_pure::build_admin_pool_cleanup_result_payload(affected))
|
||||
.into_response(),
|
||||
)
|
||||
Ok(affected)
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_provider_catalog_key_if_current<F>(
|
||||
&self,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
key_id: &str,
|
||||
should_delete: F,
|
||||
) -> Result<bool, GatewayError>
|
||||
where
|
||||
F: FnOnce(&StoredProviderCatalogKey) -> bool,
|
||||
{
|
||||
let key_ids = [key_id.to_string()];
|
||||
let Some(key) = self
|
||||
.read_provider_catalog_keys_by_ids(&key_ids)
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if key.provider_id != provider.id || !should_delete(&key) {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
self.clear_admin_provider_pool_cooldown(&provider.id, &key.id)
|
||||
.await;
|
||||
self.reset_admin_provider_pool_cost(&provider.id, &key.id)
|
||||
.await;
|
||||
let deleted = self.delete_provider_catalog_key(&key.id).await?;
|
||||
if deleted {
|
||||
let deleted_key_ids = [key.id.clone()];
|
||||
self.cleanup_deleted_provider_catalog_refs(&provider.id, &[], &deleted_key_ids)
|
||||
.await?;
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_pool_batch_action_response(
|
||||
|
||||
@@ -40,6 +40,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.map(|model| AdminSystemConfigGlobalModel {
|
||||
name: model.name.clone(),
|
||||
display_name: model.display_name.clone(),
|
||||
usage_count: Some(model.usage_count),
|
||||
default_price_per_request: model.default_price_per_request,
|
||||
default_tiered_pricing: model.default_tiered_pricing.clone(),
|
||||
supported_capabilities: model.supported_capabilities.as_ref().and_then(|value| {
|
||||
@@ -169,6 +170,13 @@ impl<'a> AdminAppState<'a> {
|
||||
) -> Result<serde_json::Value, GatewayError> {
|
||||
let users = self.list_non_admin_export_users().await?;
|
||||
let user_ids = users.iter().map(|user| user.id.clone()).collect::<Vec<_>>();
|
||||
let user_usage_totals = self
|
||||
.app
|
||||
.summarize_usage_totals_by_user_ids(&user_ids)
|
||||
.await?
|
||||
.into_iter()
|
||||
.map(|totals| (totals.user_id.clone(), totals))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let user_wallets = self.list_wallet_snapshots_by_user_ids(&user_ids).await?;
|
||||
let user_api_keys = self
|
||||
.list_auth_api_key_export_records_by_user_ids(&user_ids)
|
||||
@@ -185,6 +193,7 @@ impl<'a> AdminAppState<'a> {
|
||||
let standalone_wallets = self
|
||||
.list_wallet_snapshots_by_api_key_ids(&standalone_api_key_ids)
|
||||
.await?;
|
||||
let usage_aggregates = self.export_admin_system_usage_aggregates().await?;
|
||||
|
||||
let wallets_by_user_id = user_wallets
|
||||
.into_iter()
|
||||
@@ -260,8 +269,10 @@ impl<'a> AdminAppState<'a> {
|
||||
self.build_admin_system_users_export_api_key_payload(key, None, true)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let usage_totals = user_usage_totals.get(&user.id);
|
||||
|
||||
json!({
|
||||
"id": user.id.clone(),
|
||||
"email": user.email.clone(),
|
||||
"email_verified": user.email_verified,
|
||||
"username": user.username.clone(),
|
||||
@@ -284,6 +295,12 @@ impl<'a> AdminAppState<'a> {
|
||||
.unwrap_or(false),
|
||||
"wallet": wallet_payload,
|
||||
"is_active": user.is_active,
|
||||
"request_count": usage_totals
|
||||
.map(|totals| totals.request_count)
|
||||
.unwrap_or(0),
|
||||
"total_tokens": usage_totals
|
||||
.map(|totals| totals.total_tokens)
|
||||
.unwrap_or(0),
|
||||
"api_keys": api_keys_payload,
|
||||
})
|
||||
})
|
||||
@@ -306,6 +323,7 @@ impl<'a> AdminAppState<'a> {
|
||||
"user_groups": user_groups_data,
|
||||
"users": users_data,
|
||||
"standalone_keys": standalone_keys_data,
|
||||
"usage_aggregates": usage_aggregates,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -330,6 +348,7 @@ impl<'a> AdminAppState<'a> {
|
||||
include_is_standalone: bool,
|
||||
) -> serde_json::Value {
|
||||
let mut payload = serde_json::Map::from_iter([
|
||||
("api_key_id".to_string(), json!(key.api_key_id.clone())),
|
||||
("key_hash".to_string(), json!(key.key_hash.clone())),
|
||||
("name".to_string(), json!(key.name.clone())),
|
||||
(
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -70,6 +70,26 @@ impl<'a> AdminAppState<'a> {
|
||||
self.app.purge_admin_system_data(target).await
|
||||
}
|
||||
|
||||
pub(crate) async fn export_admin_system_usage_aggregates(
|
||||
&self,
|
||||
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateSnapshot, GatewayError>
|
||||
{
|
||||
self.app.export_admin_system_usage_aggregates().await
|
||||
}
|
||||
|
||||
pub(crate) async fn import_admin_system_usage_aggregates(
|
||||
&self,
|
||||
snapshot: &aether_data::repository::system::AdminSystemUsageAggregateSnapshot,
|
||||
user_id_map: &std::collections::BTreeMap<String, String>,
|
||||
api_key_id_map: &std::collections::BTreeMap<String, String>,
|
||||
mode: aether_data::repository::system::AdminSystemUsageAggregateImportMode,
|
||||
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateImportSummary, GatewayError>
|
||||
{
|
||||
self.app
|
||||
.import_admin_system_usage_aggregates(snapshot, user_id_map, api_key_id_map, mode)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn run_admin_system_cleanup_once(
|
||||
&self,
|
||||
) -> Result<crate::maintenance::AdminSystemCleanupSummary, GatewayError> {
|
||||
|
||||
@@ -706,6 +706,19 @@ impl<'a> AdminAppState<'a> {
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn set_api_key_usage_totals(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
total_requests: u64,
|
||||
total_tokens: u64,
|
||||
total_cost_usd: f64,
|
||||
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
|
||||
{
|
||||
self.app
|
||||
.set_api_key_usage_totals(api_key_id, total_requests, total_tokens, total_cost_usd)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_user_api_key(
|
||||
&self,
|
||||
user_id: &str,
|
||||
|
||||
@@ -13,7 +13,8 @@ pub(crate) use crate::handlers::shared::{
|
||||
effective_catalog_encryption_key, encrypt_catalog_secret_with_fallbacks, json_string_list,
|
||||
masked_catalog_api_key, normalize_json_array, normalize_json_object, normalize_string_list,
|
||||
parse_catalog_auth_config_json, provider_catalog_key_supports_format,
|
||||
provider_key_health_summary, provider_key_status_snapshot_payload, query_param_bool,
|
||||
query_param_optional_bool, query_param_value, take_secret_prefix, take_secret_suffix,
|
||||
unix_secs_to_rfc3339, OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
|
||||
provider_key_health_summary, provider_key_health_summary_at,
|
||||
provider_key_status_snapshot_payload, query_param_bool, query_param_optional_bool,
|
||||
query_param_value, take_secret_prefix, take_secret_suffix, unix_secs_to_rfc3339,
|
||||
OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
|
||||
};
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user