Compare commits

...
19 Commits
Author SHA1 Message Date
elky 361952ada9 fix: resolve workspace lint and regression test failures 2026-09-09 11:34:45 +08:00
elky 6630856061 fix: harden routing failover, model testing, and wallet queries 2026-09-09 10:38:25 +08:00
elky a893bd0557 refactor(data): reuse payment order query 2026-09-09 09:21:09 +08:00
elky f2839ae6a7 feat(routing): add strategy failover controls 2026-09-09 09:12:09 +08:00
elky e58570d79d feat(routing): make client disconnect behavior strategy-scoped 2026-09-08 23:11:37 +08:00
elky 99f6499b2b fix(conversion): improve stream failures and diagnostic exports 2026-09-08 21:04:22 +08:00
elky 17d01d7fe0 fix(dns): unify provider resolution and bound SMTP and tunnel egress
Share provider DNS policy across WebSocket and connection probes, handle bracketed IPv6 literals, and preserve bounded address sets for outbound clients.

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

Document DNS policy boundaries and verify 809 gateway, tunnel, and HTTP regression tests.
2026-09-08 17:44:59 +08:00
elky 8b766930b0 fix(ci): resolve formatting, clippy and migration alias checks 2026-09-08 12:41:19 +08:00
elky c7e403b410 fix: restore container logging compatibility and normalize legacy policies 2026-09-08 11:43:51 +08:00
elky cf8ea19856 fix: harden OAuth identity and cookies and correct quota and JSON display 2026-09-08 10:51:25 +08:00
elky 7113d04f8a fix(usage): preserve original captured HTTP headers 2026-09-08 08:49:35 +08:00
elky 099b810a2f feat: optimize usage body viewing and provider card layout 2026-09-08 02:49:06 +08:00
elky 7aa0c89244 fix(gateway): restore HTTP and WS upstream support 2026-09-07 22:15:05 +08:00
elky 7847ae98c6 fix(gateway): reset stream first-byte timeout per candidate 2026-09-07 21:56:16 +08:00
elky a90d564931 fix: restore security hardening compatibility and validation
Restore authorized rule reveal, explicit full HTTP capture and retention, video task business fields, and valid payment URLs. Add opt-in credential preservation for trusted recovery, fix frontend type contracts and async races, and eliminate PostgreSQL test fixture resource leaks. Document audit coverage and successful fmt and CI-scoped Clippy checks.
2026-09-07 21:14:27 +08:00
github-actions[bot] a5c3699ae9 chore(tunnel): update download links for tunnel-v0.3.17 2026-09-07 08:06:59 +00:00
elky 7b8048c6ae chore(tunnel): release v0.3.17 2026-09-07 15:57:39 +08:00
elky ec95f2ca1f fix(tunnel): prevent stream stalls and harden session cleanup
Reliably deliver flow-control credits and terminal states, isolate slow streams and heartbeats, negotiate stream windows, and clean up cancelled streams and session tasks.

Add regression coverage for queue pressure, early cancellation, small-window streaming, drain, and reconnect. Validate 185 agent tests, 88 gateway tunnel tests, and 21 protocol tests.
2026-09-07 15:39:40 +08:00
elky aa7dbe67d3 feat(providers): add persistent card view and shared drag ordering 2026-09-07 14:13:59 +08:00
379 changed files with 23584 additions and 4518 deletions
+8 -9
View File
@@ -15,11 +15,6 @@ APP_PORT=8084
# APP_IMAGE=ghcr.io/fawney19/aether:beta
# APP_IMAGE=ghcr.io/fawney19/aether:0.7.0-rc.1
# Compose 应用容器的非 root 数字身份。
# install.sh 会自动写入安装用户的 UID/GID。
AETHER_CONTAINER_UID=65532
AETHER_CONTAINER_GID=65532
# API Key 前缀(默认 sk)
API_KEY_PREFIX=sk
@@ -31,7 +26,11 @@ RUST_LOG=aether_gateway=info
# 示例: http://localhost:5173,https://app.example.com
# CORS_ORIGINS=http://localhost:5173
# CORS_ALLOW_CREDENTIALS=true
# 如果前后端跨站并依赖登录刷新 Cookie,还要配合:
# 登录刷新 Cookie 对同源浏览器请求和可信反代自动适配 HTTP/HTTPS。
# HTTP 自动使用兼容的 SameSite=Lax(显式 Strict 保留);HTTPS 保留原有 SameSite 配置。
# 无法确认访问协议时保留安全默认值;HTTPS 反代请正确传递 X-Forwarded-Proto。
# AUTH_REFRESH_COOKIE_SECURE 可显式覆盖自动判断,公网部署仍建议使用 HTTPS。
# 如果前后端跨站并依赖登录刷新 Cookie,必须使用 HTTPS,并配合:
# AUTH_REFRESH_COOKIE_SAMESITE=None
# AUTH_REFRESH_COOKIE_SECURE=true
@@ -115,9 +114,9 @@ ADMIN_USERNAME=admin123456
# Fake-IP 域名及内网 DNS。仅信任管理员配置的上游;没有严格 DNS 过滤开关。
# URL 协议、字面 IP、TLS 证书,以及隧道中继和登录 OAuth 的校验仍保留。
# 可选 Provider OAuth 客户端。Gemini CLI 授权及刷新必须配置 client secret。
# Antigravity 默认使用内置 native-app 客户端凭据;自定义 client ID 时必须同时配置
# 对应的 client secret。未配置 client ID 时使用内置的公开 native-app client ID。
# 可选 Provider OAuth 客户端。Gemini CLI 和 Antigravity 默认使用内置 native-app
# 客户端凭据;自定义 client ID 时必须同时配置对应的 client secret。
# 显式配置的 client secret 优先于默认值。
# AETHER_GEMINI_CLI_OAUTH_CLIENT_ID=
# AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET=
# AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID=
Generated
+2 -1
View File
@@ -305,6 +305,7 @@ dependencies = [
"futures-util",
"hmac",
"http",
"http-body",
"http-body-util",
"hyper",
"hyper-util",
@@ -661,7 +662,7 @@ dependencies = [
[[package]]
name = "aether-tunnel"
version = "0.3.16"
version = "0.3.17"
dependencies = [
"aether-contracts",
"aether-gateway",
+1 -1
View File
@@ -44,5 +44,5 @@ EXPOSE 8084
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
USER 65532:65532
USER 0:0
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
+1
View File
@@ -157,4 +157,5 @@ EXPOSE 8084
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/usr/local/bin/aether-gateway", "--healthcheck"]
USER 0:0
ENTRYPOINT ["/usr/local/bin/aether-gateway"]
+1
View File
@@ -156,4 +156,5 @@ EXPOSE 8084
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
USER 0:0
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
+1 -73
View File
@@ -50,82 +50,10 @@ chmod 600 .env
./generate_keys.sh
# 编辑 .env 设置 ADMIN_PASSWORD
# 3. 首次部署 / 更新 (从以下部署形态任选其一)
# Postgres + Redis (推荐)
# 3. Docker 部署 / 更新(PostgreSQL + Redis)
docker compose pull && docker compose up -d
# Single Node:同样使用 PostgreSQL + Redis,无需挂载本地数据库文件
docker compose -f docker-compose.single-node.yml pull && docker compose -f docker-compose.single-node.yml up -d
```
应用镜像默认以固定非 root 身份 `65532:65532` 运行;Compose 移除全部 Linux capabilities、禁止提权、启用只读根文件系统,并提供带 `nosuid,nodev,noexec` 的 `/tmp`。如需使用其他身份,可在 `.env` 中设置非零的 `AETHER_CONTAINER_UID` / `AETHER_CONTAINER_GID`。数据库使用独立 PostgreSQL 容器和 named volume,不再需要调整应用数据库目录的权限。
### 一键更新
Docker Compose 部署后,可在部署目录直接执行:
```bash
./update.sh
```
`update.sh` 会拉取最新 `app` 镜像并重建 `app` 容器,Docker named volumes、`./data` 和 `./logs` 不会被删除。Single Node 部署也可显式指定:
```bash
./update.sh --mode single-node
```
现在仅支持 PostgreSQL。标准和单节点 Docker Compose 均部署 PostgreSQL + Redis;原生 systemd / launchd 安装需要显式提供 PostgreSQL `DATABASE_URL`,例如 `DATABASE_URL=postgresql://user:password@host:5432/aether`。旧数据库不会自动迁移或清空。升级时保留原有 PostgreSQL 密码、`JWT_SECRET_KEY` 和 `ENCRYPTION_KEY`,不要重新生成整个 `.env`。
仓库自带的 Docker Compose 默认把应用日志输出到容器 `stdout/stderr`,直接用 `docker compose logs -f app` 查看,并由 Docker 轮转日志,避免非 root 用户被宿主机日志目录权限拖垮启动。如果你确实需要文件日志,需要在 compose 里把 `AETHER_LOG_DESTINATION` 改成 `file|both`,额外挂载目录到 `/opt/aether/logs`,并让它归 `.env` 中配置的容器 UID/GID 所有;只读根文件系统不会阻止显式可写挂载。
管理后台右上角“版本信息”会检测新版本。Docker Compose 部署只提示版本,实际更新继续执行 `./update.sh`;systemd / launchd / 二进制部署才使用后台自更新,流程是下载对应平台的 GitHub Release 包、强制校验 `SHA256SUMS`、解压到 `/opt/aether/releases/<version>`,再切换 `/opt/aether/current` 并退出进程,交给 systemd / launchd 拉起新版本。
正式 Release 还会发布由 GitHub Actions OIDC / Sigstore 签发的 SLSA build provenance。需要验证发布者身份时,下载目标 tarball 和 `AETHER_RELEASE_PROVENANCE.sigstore.json`,并把 `TAG` 设置为对应 Release tag:
```bash
gh attestation verify "aether-${TAG}-linux-amd64.tar.gz" \
--repo fawney19/Aether \
--signer-workflow fawney19/Aether/.github/workflows/release.yml \
--source-ref "refs/tags/${TAG}" \
--bundle AETHER_RELEASE_PROVENANCE.sigstore.json
```
`docker-compose.yml` 中的官方 PostgreSQL 和 Redis 镜像均固定到多架构 OCI index digest。升级这些依赖时应在发布变更中显式更新 digest,避免同名 tag 在无人审查的情况下改变部署内容。
正式发布到 GHCR 和 Docker Hub 的多架构 Aether 镜像也带有同一 GitHub Actions OIDC / Sigstore provenance;生产 `Dockerfile.app` 的 BusyBox 与 Distroless 基础镜像同样固定到多架构 OCI index digest。
源码或本地构建版本不会启用后台在线更新,请继续使用源码更新流程。Docker Compose 用户如果希望“容器重建后也保持镜像层面的新版本”,仍建议定期运行 `./update.sh` 拉取并重建 app 镜像。服务器访问 GitHub 需要代理时,可设置 `AETHER_UPDATE_PROXY_URL`,也兼容 `UPDATE_PROXY_URL`、`HTTPS_PROXY`、`ALL_PROXY`、`HTTP_PROXY` 以及 `NO_PROXY`。共享出口触发 GitHub API 限流时,可设置只读 `AETHER_UPDATE_GITHUB_TOKEN`,也兼容 `GITHUB_TOKEN` / `GH_TOKEN`。下载总超时默认 600 秒,连续无响应/无数据默认 30 秒,可通过 `AETHER_UPDATE_DOWNLOAD_TIMEOUT_SECS` 和 `AETHER_UPDATE_DOWNLOAD_IDLE_TIMEOUT_SECS` 调整。
标准和 Single Node Docker Compose 均使用 Docker named volume 存放 PostgreSQL 数据。
如果是本地源码构建镜像的部署,继续使用:
```bash
./deploy.sh
```
如果要在本机联调“管理后台在线更新”本身,可启动仓库内置的 release-layout 测试环境:
```bash
docker compose -f docker-compose.release-local.yml up -d --build
```
这套环境会用当前源码构建一个本地测试镜像,但编译为 `release` 类型,并默认伪装成 `v0.7.0`,这样后台会按正式发布版逻辑开放“立即更新”。默认监听 `http://127.0.0.1:18085`,数据目录使用 `./data-release-local`;日志默认走 `docker logs`,不会影响你正在跑的源码构建容器。
如果这套容器在 `prepare-update` 时访问 GitHub 失败,而你本机是通过代理出网,请在 `.env` 里把 `AETHER_UPDATE_PROXY_URL` 写成宿主机地址,例如 `http://host.docker.internal:7890`;容器内的 `127.0.0.1` 指向容器自身,不是宿主机。
如果想重置这套联调环境(包括 `/opt/aether/current` 和已下载的历史版本),执行:
```bash
docker compose -f docker-compose.release-local.yml down -v
```
可选变量:
- `AETHER_RELEASE_LOCAL_VERSION`:本地联调镜像对外声明的当前版本,默认 `v0.7.0`
- `AETHER_RELEASE_LOCAL_PORT`:本地联调端口,默认 `18085`
- `LOCAL_RELEASE_APP_IMAGE`:本地联调镜像名,默认 `aether-app:release-local`
### 一键安装(PostgreSQL + Redis)
```bash
+1
View File
@@ -62,6 +62,7 @@ flate2.workspace = true
futures-util.workspace = true
hmac.workspace = true
http.workspace = true
http-body = "1"
http-body-util = "0.1"
hyper = { version = "1", features = ["client", "server", "http1", "http2"] }
hyper-util = { version = "0.1", features = ["client-legacy", "client-pool", "server-auto", "service", "tokio"] }
@@ -194,7 +194,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncA
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -246,7 +246,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt>
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -118,7 +118,7 @@ pub(crate) fn build_local_execution_report_context(
insert_pool_key_lease_report_context_fields(&mut extra_fields, parts.pool_key_lease);
insert_scheduler_affinity_policy_report_context_field(&mut extra_fields, parts.routing_policy);
if let Some(policy) = parts.routing_policy {
if let Ok(value) = serde_json::to_value(policy.execution_policy) {
if let Ok(value) = serde_json::to_value(&policy.execution_policy) {
extra_fields.insert(ROUTING_EXECUTION_POLICY_REPORT_FIELD.to_string(), value);
}
}
@@ -179,7 +179,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -224,7 +224,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -257,7 +257,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -302,7 +302,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -109,7 +109,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -182,7 +182,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -232,7 +232,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttempt
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -124,7 +124,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttem
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -97,7 +97,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSo
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -729,117 +729,6 @@ fn update_normalization_codex_capabilities_digest(
update_normalization_string_vec_digest(digest, &capabilities.supported_service_tiers);
}
#[cfg(test)]
mod continuation_fingerprint_tests {
use http::HeaderValue;
use serde_json::json;
use super::ResponsesWebSocketBodyNormalization;
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
#[test]
fn normalization_fingerprint_is_stable_for_json_object_key_order() {
let first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}, "x": 1}));
let second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"x": 1, "reasoning": {"effort": "high"}}));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_changes_with_effective_contract() {
let base = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let changed_policy = base.clone().with_reasoning_replay_policy_for_tests(
OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
assert_ne!(
base.continuation_fingerprint(),
changed_policy.continuation_fingerprint()
);
let changed_patch = base
.clone()
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "low"}}));
assert_ne!(
base.continuation_fingerprint(),
changed_patch.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_ignores_unrelated_volatile_request_headers() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "x-contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
first
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-1"));
first
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-1"));
let mut second = first.clone();
second
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-2"));
second
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-2"));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"headers that no body-rule condition reads must not invalidate a persisted continuation"
);
}
#[test]
fn normalization_fingerprint_tracks_headers_used_by_body_rule_conditions() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "X-Contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules.clone());
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
let mut second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
second
.request_headers
.insert("x-contract", HeaderValue::from_static("disabled"));
assert_ne!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"a header that controls an effective body-rule condition remains part of the contract"
);
}
}
/// Builds one upstream decision for a Responses WebSocket turn. The session
/// reuses this decision for same-model turns and invokes the planner again when
/// a later `response.create` changes the public model.
@@ -1058,3 +947,114 @@ async fn release_responses_websocket_planning_lease(
}
}
}
#[cfg(test)]
mod continuation_fingerprint_tests {
use http::HeaderValue;
use serde_json::json;
use super::ResponsesWebSocketBodyNormalization;
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
#[test]
fn normalization_fingerprint_is_stable_for_json_object_key_order() {
let first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}, "x": 1}));
let second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"x": 1, "reasoning": {"effort": "high"}}));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_changes_with_effective_contract() {
let base = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let changed_policy = base.clone().with_reasoning_replay_policy_for_tests(
OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
assert_ne!(
base.continuation_fingerprint(),
changed_policy.continuation_fingerprint()
);
let changed_patch = base
.clone()
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "low"}}));
assert_ne!(
base.continuation_fingerprint(),
changed_patch.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_ignores_unrelated_volatile_request_headers() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "x-contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
first
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-1"));
first
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-1"));
let mut second = first.clone();
second
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-2"));
second
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-2"));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"headers that no body-rule condition reads must not invalidate a persisted continuation"
);
}
#[test]
fn normalization_fingerprint_tracks_headers_used_by_body_rule_conditions() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "X-Contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules.clone());
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
let mut second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
second
.request_headers
.insert("x-contract", HeaderValue::from_static("disabled"));
assert_ne!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"a header that controls an effective body-rule condition remains part of the contract"
);
}
}
@@ -166,7 +166,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -216,7 +216,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStream
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
+14 -14
View File
@@ -23,7 +23,6 @@ const MAX_BARK_TEMPLATE_BYTES: usize = 256 * 1024;
const MAX_BARK_TITLE_BYTES: usize = 512;
const MAX_BARK_BODY_BYTES: usize = 2 * 1024 * 1024;
const MAX_BARK_RENDERED_BODY_BYTES: usize = 2 * 1024 * 1024;
const MAX_BARK_RESOLVED_ADDRESSES: usize = 32;
#[derive(Clone)]
pub(crate) struct BarkPushConfig {
@@ -208,19 +207,20 @@ async fn build_bark_push_client_and_url(
let port = push_url
.port_or_known_default()
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少端口".to_string()))?;
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
tokio::time::timeout(
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
tokio::net::lookup_host((host.as_str(), port)),
)
.await
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析超时".to_string()))?
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析失败".to_string()))?
.take(MAX_BARK_RESOLVED_ADDRESSES)
.collect::<Vec<_>>()
};
let addresses = aether_http::lookup_host_with_limits(
host.as_str(),
port,
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
)
.await
.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "Bark 服务器 DNS 解析超时",
std::io::ErrorKind::InvalidData => "Bark 服务器 DNS 解析返回过多地址",
_ => "Bark 服务器 DNS 解析失败",
};
GatewayError::Internal(message.to_string())
})?;
let allow_benchmarking_ip = push_url.scheme() == "https"
&& push_url.port_or_known_default() == Some(443)
&& host.eq_ignore_ascii_case("api.day.app");
@@ -496,6 +496,7 @@ fn access_for_route(method: &http::Method, decision: &GatewayControlDecision) ->
Some("admin:endpoints_manage"),
Some(
"reveal_key"
| "reveal_endpoint_rules"
| "export_key"
| "create_provider_key"
| "update_key"
@@ -1490,6 +1491,12 @@ mod tests {
fn plaintext_credential_reads_require_admin_permission() {
let read_only_permissions = read_only_management_token_permissions();
let cases = [
(
"admin:endpoints_manage",
"reveal_endpoint_rules",
None,
"admin:endpoints_manage:admin",
),
(
"admin:endpoints_manage",
"reveal_key",
@@ -302,6 +302,19 @@ pub(super) fn classify_admin_endpoints_family_route(
"admin:endpoints_manage",
false,
))
} else if method == http::Method::GET
&& normalized_path
.strip_prefix("/api/admin/endpoints/")
.and_then(|path| path.strip_suffix("/rules/reveal"))
.is_some_and(|endpoint_id| !endpoint_id.is_empty() && !endpoint_id.contains('/'))
{
Some(classified(
"admin_proxy",
"endpoints_manage",
"reveal_endpoint_rules",
"admin:endpoints_manage",
false,
))
} else if method == http::Method::GET
&& normalized_path.starts_with("/api/admin/endpoints/")
&& !normalized_path.starts_with("/api/admin/endpoints/health/")
@@ -381,6 +381,28 @@ fn classifies_admin_get_endpoint_as_admin_proxy_route() {
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_admin_reveal_endpoint_rules_as_admin_proxy_route() {
let headers = headers(&[]);
let uri: Uri = "/api/admin/endpoints/endpoint-1/rules/reveal"
.parse()
.expect("uri should parse");
let decision =
classify_control_route(&http::Method::GET, &uri, &headers).expect("route should classify");
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
assert_eq!(decision.route_family.as_deref(), Some("endpoints_manage"));
assert_eq!(
decision.route_kind.as_deref(),
Some("reveal_endpoint_rules")
);
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("admin:endpoints_manage")
);
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_admin_create_endpoint_as_admin_proxy_route() {
let headers = http::HeaderMap::new();
@@ -41,16 +41,9 @@ fn usage_request_record_level_from_value(value: Option<&Value>) -> UsageRequestR
return UsageRequestRecordLevel::Basic;
};
if value.eq_ignore_ascii_case("basic")
|| value.eq_ignore_ascii_case("base")
|| value.eq_ignore_ascii_case("headers")
|| value.eq_ignore_ascii_case("minimal")
|| value.eq_ignore_ascii_case("none")
{
UsageRequestRecordLevel::Basic
if value.eq_ignore_ascii_case("full") {
UsageRequestRecordLevel::Full
} else {
// Raw HTTP payload capture is disabled at the runtime boundary. The setting remains
// accepted for compatibility, but no longer authorizes collecting request/response data.
UsageRequestRecordLevel::Basic
}
}
@@ -501,7 +494,7 @@ mod tests {
}
#[tokio::test]
async fn usage_runtime_access_disables_full_http_capture() {
async fn usage_runtime_access_honors_explicit_full_http_capture() {
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
"request_record_level".to_string(),
json!("full"),
@@ -511,7 +504,32 @@ mod tests {
.await
.expect("request record level should read");
assert_eq!(level, UsageRequestRecordLevel::Basic);
assert_eq!(level, UsageRequestRecordLevel::Full);
}
#[tokio::test]
async fn usage_runtime_access_honors_legacy_full_without_overriding_current_config() {
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
"request_log_level".to_string(),
json!(" FULL "),
)]);
assert_eq!(
UsageRuntimeAccess::request_record_level(&state)
.await
.unwrap(),
UsageRequestRecordLevel::Full
);
let state = state.with_system_config_values_for_tests([
("request_log_level".to_string(), json!("full")),
("request_record_level".to_string(), json!("basic")),
]);
assert_eq!(
UsageRuntimeAccess::request_record_level(&state)
.await
.unwrap(),
UsageRequestRecordLevel::Basic
);
}
#[tokio::test]
@@ -1593,6 +1593,19 @@ impl GatewayDataState {
}
}
pub(crate) async fn read_request_usage_body_payload(
&self,
body_ref: &str,
) -> Result<
Option<aether_data_contracts::repository::usage::StoredUsageBodyPayload>,
DataLayerError,
> {
match &self.usage_reader {
Some(repository) => repository.read_body_payload(body_ref).await,
None => Ok(None),
}
}
pub(crate) async fn list_usage_audits(
&self,
query: &UsageAuditListQuery,
+194 -43
View File
@@ -136,14 +136,16 @@ pub(crate) async fn send_smtp_email(
email: ComposedEmail,
) -> Result<(), GatewayError> {
validate_smtp_delivery_inputs(&config, &email)?;
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email))
let stream = connect_tcp_stream(&config).await?;
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email, stream))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
}
pub(crate) async fn probe_smtp_connection(config: SmtpDeliveryConfig) -> Result<(), GatewayError> {
validate_smtp_config(&config)?;
tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config))
let stream = connect_tcp_stream(&config).await?;
tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config, stream))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
}
@@ -328,43 +330,58 @@ fn resolve_server_name(host: &str) -> Result<rustls::pki_types::ServerName<'stat
.map_err(|err| GatewayError::Internal(err.to_string()))
}
fn connect_tcp_stream(config: &SmtpDeliveryConfig) -> Result<std::net::TcpStream, GatewayError> {
use std::net::ToSocketAddrs;
let addresses = (config.host.as_str(), config.port)
.to_socket_addrs()
.map_err(|err| GatewayError::Internal(err.to_string()))?
.take(16)
.collect::<Vec<_>>();
if addresses.is_empty() {
return Err(GatewayError::Internal(
"smtp host did not resolve to an address".to_string(),
));
}
let deadline = std::time::Instant::now()
.checked_add(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))
.unwrap_or_else(std::time::Instant::now);
let mut last_error = None;
let mut stream = None;
for address in addresses {
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
if remaining.is_zero() {
break;
async fn connect_tcp_stream(
config: &SmtpDeliveryConfig,
) -> Result<std::net::TcpStream, GatewayError> {
connect_tcp_stream_with_dns(
aether_http::lookup_host_with_limits(
&config.host,
config.port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
),
std::time::Duration::from_secs(SMTP_TIMEOUT_SECS),
)
.await
}
async fn connect_tcp_stream_with_dns(
lookup: impl std::future::Future<Output = std::io::Result<Vec<std::net::SocketAddr>>>,
timeout: std::time::Duration,
) -> Result<std::net::TcpStream, GatewayError> {
let stream = tokio::time::timeout(timeout, async {
let addresses = lookup.await.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "smtp DNS resolution timed out",
std::io::ErrorKind::InvalidData => {
"smtp DNS resolution returned too many addresses"
}
_ => "smtp DNS resolution failed",
};
GatewayError::Internal(message.to_string())
})?;
if addresses.is_empty() {
return Err(GatewayError::Internal(
"smtp host did not resolve to an address".to_string(),
));
}
match std::net::TcpStream::connect_timeout(&address, remaining) {
Ok(candidate) => {
stream = Some(candidate);
break;
}
Err(err) => last_error = Some(err),
}
}
let stream = stream.ok_or_else(|| {
GatewayError::Internal(
last_error
.map(|err| err.to_string())
.unwrap_or_else(|| "smtp connection timed out".to_string()),
)
})?;
let attempts = addresses
.into_iter()
.map(|address| Box::pin(tokio::net::TcpStream::connect(address)));
futures_util::future::select_ok(attempts)
.await
.map(|(stream, _)| stream)
.map_err(|error| {
GatewayError::Internal(format!("smtp connection failed ({})", error.kind()))
})
})
.await
.map_err(|_| GatewayError::Internal("smtp DNS or TCP connection timed out".to_string()))??;
let stream = stream
.into_std()
.map_err(|err| GatewayError::Internal(err.to_string()))?;
stream
.set_nonblocking(false)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
stream
.set_read_timeout(Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS)))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
@@ -680,16 +697,15 @@ fn smtp_probe_connection<S: std::io::Read + std::io::Write>(
fn send_smtp_email_blocking(
config: SmtpDeliveryConfig,
email: ComposedEmail,
stream: std::net::TcpStream,
) -> Result<(), GatewayError> {
if config.use_ssl {
let stream = connect_tcp_stream(&config)?;
let tls_stream = wrap_tls_stream(stream, &config.host)?;
let mut reader = std::io::BufReader::new(tls_stream);
let _ = smtp_expect(&mut reader, &[220])?;
return smtp_send_message(&mut reader, &config, &email);
}
let stream = connect_tcp_stream(&config)?;
let mut reader = std::io::BufReader::new(stream);
let _ = smtp_expect(&mut reader, &[220])?;
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
@@ -705,16 +721,17 @@ fn send_smtp_email_blocking(
smtp_deliver_message(&mut reader, &config, &email)
}
fn probe_smtp_connection_blocking(config: SmtpDeliveryConfig) -> Result<(), GatewayError> {
fn probe_smtp_connection_blocking(
config: SmtpDeliveryConfig,
stream: std::net::TcpStream,
) -> Result<(), GatewayError> {
if config.use_ssl {
let stream = connect_tcp_stream(&config)?;
let tls_stream = wrap_tls_stream(stream, &config.host)?;
let mut reader = std::io::BufReader::new(tls_stream);
let _ = smtp_expect(&mut reader, &[220])?;
return smtp_probe_connection(&mut reader, &config);
}
let stream = connect_tcp_stream(&config)?;
let mut reader = std::io::BufReader::new(stream);
let _ = smtp_expect(&mut reader, &[220])?;
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
@@ -778,6 +795,140 @@ mod tests {
assert!(validate_smtp_delivery_inputs(&config(), &email()).is_ok());
}
#[tokio::test]
async fn smtp_connection_deadline_includes_a_stalled_dns_lookup() {
let error = connect_tcp_stream_with_dns(
std::future::pending(),
std::time::Duration::from_millis(5),
)
.await
.expect_err("DNS must not outlive the connection deadline");
assert!(format!("{error:?}").contains("smtp DNS or TCP connection timed out"));
}
#[tokio::test]
async fn smtp_dns_errors_and_empty_answers_fail_without_connecting() {
for (addresses, expected) in [
(Ok(Vec::new()), "smtp host did not resolve to an address"),
(
Err(std::io::Error::other("sensitive-dns-detail")),
"smtp DNS resolution failed",
),
(
Err(std::io::Error::from(std::io::ErrorKind::InvalidData)),
"smtp DNS resolution returned too many addresses",
),
(
Err(std::io::Error::from(std::io::ErrorKind::TimedOut)),
"smtp DNS resolution timed out",
),
] {
let error = connect_tcp_stream_with_dns(
std::future::ready(addresses),
std::time::Duration::from_secs(1),
)
.await
.expect_err("invalid DNS answers must fail before TCP connect");
assert!(format!("{error:?}").contains(expected));
assert!(!format!("{error:?}").contains("sensitive-dns-detail"));
}
}
#[tokio::test]
async fn smtp_connection_tries_answers_beyond_the_old_sixteen_address_limit() {
let unavailable = tokio::net::TcpSocket::new_v4().unwrap();
unavailable.bind("127.0.0.1:0".parse().unwrap()).unwrap();
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let available = listener.local_addr().unwrap();
let mut addresses = vec![unavailable.local_addr().unwrap(); 16];
addresses.push(available);
let stream = connect_tcp_stream_with_dns(
std::future::ready(Ok(addresses)),
std::time::Duration::from_secs(5),
)
.await
.expect("later DNS answers should remain available for fallback");
assert_eq!(stream.peer_addr().unwrap(), available);
assert_eq!(
stream.read_timeout().unwrap(),
Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))
);
}
#[tokio::test]
async fn smtp_probe_and_delivery_use_the_preconnected_stream() {
use tokio::io::{AsyncBufReadExt, AsyncWriteExt};
for deliver in [false, true] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut reader = tokio::io::BufReader::new(stream);
reader
.get_mut()
.write_all(b"220 mock SMTP ready\r\n")
.await
.unwrap();
let mut delivered = false;
loop {
let mut line = String::new();
assert!(reader.read_line(&mut line).await.unwrap() > 0);
let response = if line.starts_with("EHLO ")
|| line.starts_with("MAIL FROM:")
|| line.starts_with("RCPT TO:")
{
&b"250 OK\r\n"[..]
} else if line == "DATA\r\n" {
reader
.get_mut()
.write_all(b"354 End with dot\r\n")
.await
.unwrap();
loop {
line.clear();
assert!(reader.read_line(&mut line).await.unwrap() > 0);
if line == ".\r\n" {
break;
}
}
delivered = true;
&b"250 Accepted\r\n"[..]
} else {
assert_eq!(line, "QUIT\r\n");
reader
.get_mut()
.write_all(b"221 Goodbye\r\n")
.await
.unwrap();
break;
};
reader.get_mut().write_all(response).await.unwrap();
}
assert_eq!(delivered, deliver);
});
let config = SmtpDeliveryConfig {
host: "127.0.0.1".to_string(),
port,
user: None,
password: None,
use_tls: false,
use_ssl: false,
..config()
};
tokio::time::timeout(std::time::Duration::from_secs(5), async {
if deliver {
send_smtp_email(config, email()).await.unwrap();
} else {
probe_smtp_connection(config).await.unwrap();
}
server.await.unwrap();
})
.await
.expect("local SMTP probe and delivery should complete");
}
}
#[test]
fn rejects_authentication_over_plaintext_smtp() {
let mut insecure = config();
@@ -426,11 +426,8 @@ mod tests {
assert!(candidate.finished_at_unix_ms.is_some());
}
/// The guard holds no request body, and the persistence boundary intentionally
/// rejects request/response capture material. A dropped-attempt settlement
/// must not re-introduce an inline body or a caller-controlled body reference.
#[tokio::test]
async fn settling_a_dropped_attempt_does_not_reintroduce_request_body_capture() {
async fn settling_a_dropped_attempt_respects_disabled_request_body_capture() {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let state = test_state(&usage_repository, &request_candidate_repository);
@@ -444,8 +441,6 @@ mod tests {
candidate_started_unix_ms,
)
.await;
// This deliberately supplies capture material to prove that the usage
// persistence boundary strips it before either lifecycle write stores it.
let captured_body = json!({"stream": true, "service_tier": "priority"});
let mut capture = build_pending_usage_record(
&plan,
@@ -481,7 +476,10 @@ mod tests {
.expect("cancelled usage should be recorded");
assert_eq!(usage.provider_request_body, None);
assert_eq!(usage.provider_request_body_ref, None);
assert_eq!(usage.provider_request_body_state, None);
assert_eq!(
usage.provider_request_body_state,
Some(UsageBodyCaptureState::Disabled)
);
}
#[tokio::test]
@@ -68,7 +68,6 @@ const CHATGPT_WEB_IMAGE_PUBLIC_CONNECT_TIMEOUT_MS: u64 = 10_000;
const CHATGPT_WEB_IMAGE_PUBLIC_READ_TIMEOUT_MS: u64 = 30_000;
const CHATGPT_WEB_IMAGE_PUBLIC_TOTAL_TIMEOUT_MS: u64 = 300_000;
const CHATGPT_WEB_OPAQUE_ID_MAX_BYTES: usize = 256;
const CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES: usize = 32;
const CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES: usize = 64 * 1024;
const CHATGPT_WEB_IMAGE_UPLOAD_RESPONSE_LIMIT_BYTES: usize = 64 * 1024;
const CHATGPT_WEB_IMAGE_MAX_PROMPT_BYTES: usize = 32 * 1024;
@@ -1334,24 +1333,18 @@ async fn resolve_public_web_image_addrs(
"ChatGPT-Web image URL is missing a port".to_string(),
)
})?;
let resolved = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
tokio::time::timeout(lookup_timeout, tokio::net::lookup_host((host, port)))
.await
.map_err(|_| {
ExecutionRuntimeTransportError::UpstreamRequest(
"ChatGPT-Web image URL DNS resolution timed out".to_string(),
)
})?
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format!(
"ChatGPT-Web image URL DNS resolution failed: {err}"
))
})?
.take(CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES)
.collect::<Vec<_>>()
};
let resolved = aether_http::lookup_host_with_limits(host, port, lookup_timeout)
.await
.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "ChatGPT-Web image URL DNS resolution timed out",
std::io::ErrorKind::InvalidData => {
"ChatGPT-Web image URL DNS resolution returned too many addresses"
}
_ => "ChatGPT-Web image URL DNS resolution failed",
};
ExecutionRuntimeTransportError::UpstreamRequest(message.to_string())
})?;
if resolved.is_empty() {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"ChatGPT-Web image URL DNS resolution returned no addresses".to_string(),
@@ -14,7 +14,7 @@ fn sync_plan_kind_disables_local_candidate_failover(plan_kind: &str) -> bool {
)
}
fn openai_image_success_disables_local_success_failover(
pub(super) fn openai_image_success_disables_local_success_failover(
plan: &ExecutionPlan,
status_code: u16,
) -> bool {
@@ -1036,6 +1036,7 @@ mod tests {
policy,
LocalFailoverPolicy {
max_retries: Some(1),
routing_rules: Default::default(),
max_transfer_count: 0,
max_transfer_timeout_seconds: 0,
stop_status_codes: [503].into_iter().collect(),
@@ -1826,12 +1826,12 @@ async fn fetch_grok_attachment_url(
// a fragment from the previous URL, while an absolute Location can
// introduce either explicitly.
validate_grok_attachment_url(&url)?;
let public_addr = public_socket_addr_for_url(&url).await?;
let public_addrs = public_socket_addrs_for_url(&url).await?;
let response = reqwest::Client::builder()
.no_proxy()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.resolve_to_addrs(url.host_str().unwrap_or_default(), &[public_addr])
.resolve_to_addrs(url.host_str().unwrap_or_default(), &public_addrs)
.build()
.map_err(ExecutionRuntimeTransportError::ClientBuild)?
.get(url.clone())
@@ -1897,10 +1897,10 @@ fn validate_grok_attachment_url(url: &reqwest::Url) -> Result<(), ExecutionRunti
Ok(())
}
async fn public_socket_addr_for_url(
async fn public_socket_addrs_for_url(
url: &reqwest::Url,
) -> Result<std::net::SocketAddr, ExecutionRuntimeTransportError> {
let host = url.host().ok_or_else(|| {
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
let host = url.host_str().ok_or_else(|| {
ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL is missing a host".to_string(),
)
@@ -1910,64 +1910,34 @@ async fn public_socket_addr_for_url(
"Grok attachment URL is missing a port".to_string(),
)
})?;
let host = match host {
url::Host::Ipv4(ip) => {
let ip = IpAddr::V4(ip);
if !grok_attachment_ip_is_public(ip) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
return Ok(std::net::SocketAddr::new(ip, port));
}
url::Host::Ipv6(ip) => {
let ip = IpAddr::V6(ip);
if !grok_attachment_ip_is_public(ip) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
return Ok(std::net::SocketAddr::new(ip, port));
}
url::Host::Domain(host) => host,
};
if let Ok(ip) = host.parse::<IpAddr>() {
if !grok_attachment_ip_is_public(ip) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
return Ok(std::net::SocketAddr::new(ip, port));
}
let mut public_addr = None;
let mut resolved_any = false;
for addr in
let addresses =
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format!(
"Grok attachment URL DNS resolution failed: {err}"
))
})?
{
resolved_any = true;
if !grok_attachment_ip_is_public(addr.ip()) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
public_addr.get_or_insert(addr);
}
if !resolved_any {
})?;
validate_grok_attachment_addresses(addresses)
}
fn validate_grok_attachment_addresses(
addresses: Vec<std::net::SocketAddr>,
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
if addresses.is_empty() {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL DNS resolution returned no addresses".to_string(),
));
}
public_addr.ok_or_else(|| {
ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL has no public address".to_string(),
)
})
if addresses
.iter()
.any(|address| !grok_attachment_ip_is_public(address.ip()))
{
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
Ok(addresses)
}
fn grok_attachment_ip_is_public(ip: IpAddr) -> bool {
@@ -3898,7 +3868,7 @@ mod tests {
grok_should_use_imagine_websocket, grok_success_frame_stream, grok_upload_url,
grok_upstream_model_name, grok_usage_estimate, grok_user_id_from_cookie_header,
materialize_grok_image_assets, maximum_base64_len_for_decoded_limit, openai_chat_body,
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addr_for_url,
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addrs_for_url,
set_grok_image_edit_config, validate_grok_attachment_url, GrokAttachmentInput,
GrokCollected, GrokImagineImage, GrokStreamAdapter,
};
@@ -4520,7 +4490,7 @@ mod tests {
] {
let url = reqwest::Url::parse(raw_url).expect("URL should parse");
assert!(
public_socket_addr_for_url(&url).await.is_err(),
public_socket_addrs_for_url(&url).await.is_err(),
"private IPv6 literal should be rejected: {raw_url}"
);
}
@@ -4528,13 +4498,31 @@ mod tests {
let url = reqwest::Url::parse("https://[2606:4700:4700::1111]/attachment")
.expect("URL should parse");
assert_eq!(
public_socket_addr_for_url(&url)
public_socket_addrs_for_url(&url)
.await
.expect("public IPv6 literal should pass"),
"[2606:4700:4700::1111]:443".parse().unwrap()
vec!["[2606:4700:4700::1111]:443".parse().unwrap()]
);
}
#[test]
fn grok_attachment_dns_keeps_all_safe_addresses_for_connection_fallback() {
let addresses = vec![
"[2606:4700:4700::1111]:443".parse().unwrap(),
"8.8.8.8:443".parse().unwrap(),
];
assert_eq!(
super::validate_grok_attachment_addresses(addresses.clone()).unwrap(),
addresses
);
assert!(super::validate_grok_attachment_addresses(Vec::new()).is_err());
for blocked in ["198.18.0.1:443", "127.0.0.1:443", "[fd00::1]:443"] {
let mut mixed = addresses.clone();
mixed.push(blocked.parse().unwrap());
assert!(super::validate_grok_attachment_addresses(mixed).is_err());
}
}
#[test]
fn grok_attachment_url_rejects_credentials_and_fragments_on_every_hop() {
for raw_url in [
@@ -11,6 +11,10 @@ const GEMINI_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750);
pub(super) enum StreamCommitPolicy {
ResponseHeaders,
FirstClassifiedBody,
FirstSseSemanticEvent {
max_bytes: usize,
max_wait: Duration,
},
FirstAnthropicSemanticEvent {
max_bytes: usize,
max_wait: Duration,
@@ -36,16 +40,21 @@ impl StreamCommitPolicy {
return Self::FirstClassifiedBody;
}
if force_prefetch {
return Self::FirstClassifiedBody;
}
let content_type = content_type
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or_default()
.to_ascii_lowercase();
if content_type.contains("text/event-stream") {
if provider_api_format.eq_ignore_ascii_case("openai:image")
|| client_api_format.eq_ignore_ascii_case("openai:image")
{
return if force_prefetch {
Self::FirstClassifiedBody
} else {
Self::ResponseHeaders
};
}
if provider_api_format.eq_ignore_ascii_case("claude:messages")
&& provider_api_format.eq_ignore_ascii_case(client_api_format)
&& !has_private_stream_normalizer
@@ -62,7 +71,14 @@ impl StreamCommitPolicy {
max_wait: GEMINI_PRECOMMIT_MAX_WAIT,
};
}
return Self::ResponseHeaders;
return Self::FirstSseSemanticEvent {
max_bytes: MAX_STREAM_PREFETCH_BYTES,
max_wait: Duration::from_secs(30),
};
}
if force_prefetch {
return Self::FirstClassifiedBody;
}
if has_private_stream_normalizer || has_local_stream_rewriter {
@@ -91,14 +107,17 @@ impl StreamCommitPolicy {
pub(super) const fn requires_bounded_frame_wait(self) -> bool {
matches!(
self,
Self::FirstAnthropicSemanticEvent { .. } | Self::FirstGeminiSemanticEvent { .. }
Self::FirstAnthropicSemanticEvent { .. }
| Self::FirstGeminiSemanticEvent { .. }
| Self::FirstSseSemanticEvent { .. }
)
}
pub(super) const fn max_precommit_wait(self) -> Option<Duration> {
match self {
Self::FirstAnthropicSemanticEvent { max_wait, .. }
| Self::FirstGeminiSemanticEvent { max_wait, .. } => Some(max_wait),
| Self::FirstGeminiSemanticEvent { max_wait, .. }
| Self::FirstSseSemanticEvent { max_wait, .. } => Some(max_wait),
Self::ResponseHeaders | Self::FirstClassifiedBody => None,
}
}
@@ -110,6 +129,16 @@ impl StreamCommitPolicy {
pub(super) const fn is_gemini(self) -> bool {
matches!(self, Self::FirstGeminiSemanticEvent { .. })
}
pub(super) fn with_precommit_wait(mut self, wait: Duration) -> Self {
match &mut self {
Self::FirstAnthropicSemanticEvent { max_wait, .. }
| Self::FirstGeminiSemanticEvent { max_wait, .. }
| Self::FirstSseSemanticEvent { max_wait, .. } => *max_wait = wait,
_ => {}
}
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -133,6 +162,7 @@ pub(super) struct StreamCommitGate {
observed_bytes: usize,
anthropic: AnthropicSsePrecommitInspector,
gemini: GeminiSsePrecommitInspector,
generic: GenericSsePrecommitInspector,
}
impl StreamCommitGate {
@@ -148,6 +178,7 @@ impl StreamCommitGate {
observed_bytes: 0,
anthropic: AnthropicSsePrecommitInspector::default(),
gemini: GeminiSsePrecommitInspector::default(),
generic: GenericSsePrecommitInspector::default(),
}
}
@@ -171,6 +202,9 @@ impl StreamCommitGate {
StreamCommitPolicy::FirstGeminiSemanticEvent { max_bytes, .. } => {
(max_bytes, self.gemini.observe(chunk, max_bytes))
}
StreamCommitPolicy::FirstSseSemanticEvent { max_bytes, .. } => {
(max_bytes, self.generic.observe(chunk, max_bytes))
}
StreamCommitPolicy::ResponseHeaders | StreamCommitPolicy::FirstClassifiedBody => {
return StreamPrecommitObservation::Pending;
}
@@ -217,6 +251,152 @@ enum SemanticSseObservation {
Error { status_code: u16, body_json: Value },
}
#[derive(Debug, Default)]
struct GenericSsePrecommitInspector {
buffered: Vec<u8>,
}
impl GenericSsePrecommitInspector {
fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation {
let remaining = max_bytes.saturating_sub(self.buffered.len());
self.buffered
.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
while let Some((record_end, separator_len)) = find_sse_record_boundary(&self.buffered) {
let record = self.buffered[..record_end].to_vec();
self.buffered.drain(..record_end + separator_len);
match classify_generic_sse_record(&record) {
SemanticSseObservation::Pending => {}
observation => return observation,
}
}
if chunk.len() > remaining {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
}
}
}
fn classify_generic_sse_record(record: &[u8]) -> SemanticSseObservation {
let Ok(record) = std::str::from_utf8(record) else {
return SemanticSseObservation::SemanticEvent;
};
let normalized = record.replace("\r\n", "\n").replace('\r', "\n");
let event_type = normalized
.lines()
.find_map(|line| line.strip_prefix("event:").map(str::trim));
let data = normalized
.lines()
.filter_map(|line| line.strip_prefix("data:").map(str::trim_start))
.collect::<Vec<_>>()
.join("\n");
if data.trim().is_empty() || matches!(event_type, Some("ping" | "heartbeat" | "keepalive")) {
return SemanticSseObservation::Pending;
}
if data.trim() == "[DONE]" {
return SemanticSseObservation::SemanticEvent;
}
let Ok(body_json) = serde_json::from_str::<Value>(data.trim()) else {
return SemanticSseObservation::SemanticEvent;
};
let payload_type = body_json.get("type").and_then(Value::as_str).or(event_type);
if payload_type.is_some_and(is_anthropic_semantic_event_type) {
return classify_anthropic_sse_record(record.as_bytes());
}
let error = body_json
.get("error")
.filter(|value| !value.is_null())
.or_else(|| {
body_json
.pointer("/response/error")
.filter(|value| !value.is_null())
});
if error.is_some()
|| matches!(payload_type, Some("error" | "response.failed"))
|| body_json.get("status").and_then(Value::as_str) == Some("failed")
{
let failure = error
.map(|error| serde_json::json!({ "error": error }))
.unwrap_or_else(|| body_json.clone());
return SemanticSseObservation::Error {
status_code: crate::execution_runtime::submission::resolve_local_sync_error_status_code(
200, &failure,
),
body_json: failure,
};
}
if matches!(
payload_type,
Some("ping" | "response.created" | "response.in_progress" | "response.queued")
) {
return SemanticSseObservation::Pending;
}
if payload_type == Some("response.output_item.added")
&& matches!(
body_json.pointer("/item/type").and_then(Value::as_str),
Some("message" | "reasoning")
)
&& body_json
.pointer("/item/content")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty)
&& body_json
.pointer("/item/summary")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty)
{
return SemanticSseObservation::Pending;
}
if matches!(
payload_type,
Some("response.content_part.added" | "response.reasoning_summary_part.added")
) && matches!(
body_json.pointer("/part/type").and_then(Value::as_str),
Some("output_text" | "summary_text" | "refusal")
) && !body_json
.pointer("/part/text")
.is_some_and(value_has_semantic_content)
&& !body_json
.pointer("/part/refusal")
.is_some_and(value_has_semantic_content)
{
return SemanticSseObservation::Pending;
}
if let Some(choices) = body_json.get("choices").and_then(Value::as_array) {
let semantic = choices.iter().any(|choice| {
choice
.get("finish_reason")
.is_some_and(|value| !value.is_null())
|| choice.get("text").is_some_and(value_has_semantic_content)
|| choice
.get("delta")
.or_else(|| choice.get("message"))
.and_then(Value::as_object)
.is_some_and(|delta| {
delta.iter().any(|(name, value)| {
name != "role" && value_has_semantic_content(value)
})
})
});
return if semantic {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
};
}
SemanticSseObservation::SemanticEvent
}
fn value_has_semantic_content(value: &Value) -> bool {
match value {
Value::Null => false,
Value::String(text) => !text.is_empty(),
Value::Array(values) => !values.is_empty(),
Value::Object(values) => !values.is_empty(),
_ => true,
}
}
#[derive(Debug, Default)]
struct AnthropicSsePrecommitInspector {
buffered: Vec<u8>,
@@ -355,7 +535,30 @@ fn classify_anthropic_sse_record(record: &[u8]) -> SemanticSseObservation {
(None, Some(payload_type)) => Some(payload_type),
_ => None,
};
if semantic_type.is_some_and(is_anthropic_semantic_event_type) {
let setup_only = match semantic_type {
Some("message_start") => body_json
.pointer("/message/content")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty),
Some("content_block_start") => {
let block_type = body_json
.pointer("/content_block/type")
.and_then(Value::as_str);
matches!(block_type, Some("text" | "thinking"))
&& !body_json
.pointer("/content_block/text")
.is_some_and(value_has_semantic_content)
&& !body_json
.pointer("/content_block/thinking")
.is_some_and(value_has_semantic_content)
}
Some("content_block_stop") => true,
Some("message_delta") => body_json
.pointer("/delta/stop_reason")
.is_none_or(Value::is_null),
_ => false,
};
if !setup_only && semantic_type.is_some_and(is_anthropic_semantic_event_type) {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
@@ -507,6 +710,89 @@ pub(super) fn anthropic_error_status_code(body_json: &Value) -> u16 {
#[cfg(test)]
mod tests {
#[test]
fn image_streams_only_prefetch_when_explicitly_requested() {
for force_prefetch in [false, true] {
let policy = super::StreamCommitPolicy::for_response(
true,
Some("text/event-stream"),
"openai:image",
"openai:image",
false,
false,
force_prefetch,
);
assert_eq!(policy.commits_on_response_headers(), !force_prefetch);
assert!(!policy.requires_bounded_frame_wait());
}
}
#[test]
fn generic_sse_waits_through_setup_and_classifies_fragmented_errors() {
let setup = b"event: response.created\ndata: {\"type\":\"response.created\"}\n\n";
let failure = b"event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"capacity exhausted\"}}}\n\n";
for split in 1..failure.len() {
let policy = super::StreamCommitPolicy::FirstSseSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
};
let mut gate = super::StreamCommitGate::new(policy);
assert_eq!(
gate.observe_provider_bytes(setup),
super::StreamPrecommitObservation::Pending
);
for control in [
b"event: ping\ndata: keepalive\n\n".as_slice(),
b"data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"reasoning\",\"summary\":[]}}\n\n".as_slice(),
b"data: {\"type\":\"response.reasoning_summary_part.added\",\"part\":{\"type\":\"summary_text\",\"text\":\"\"}}\n\n".as_slice(),
] {
assert_eq!(gate.observe_provider_bytes(control), super::StreamPrecommitObservation::Pending);
}
assert_eq!(
gate.observe_provider_bytes(&failure[..split]),
super::StreamPrecommitObservation::Pending
);
assert!(matches!(
gate.observe_provider_bytes(&failure[split..]),
super::StreamPrecommitObservation::UpstreamError { .. }
));
}
}
#[test]
fn generic_sse_commits_on_content_or_tool_call_but_not_role() {
for output in [
"data: {\"choices\":[{\"delta\":{\"content\":\"hello\"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"id\":\"call-1\"}]}}]}\n\n",
] {
let mut gate =
super::StreamCommitGate::new(super::StreamCommitPolicy::FirstSseSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
});
assert_eq!(gate.observe_provider_bytes(b"data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":\"\"}}]}\n\n"), super::StreamPrecommitObservation::Pending);
assert_eq!(
gate.observe_provider_bytes(output.as_bytes()),
super::StreamPrecommitObservation::Commit
);
assert_eq!(
gate.observe_provider_bytes(b"data: {\"error\":{\"message\":\"late error\"}}\n\n"),
super::StreamPrecommitObservation::Commit
);
}
}
#[test]
fn native_anthropic_setup_does_not_hide_an_early_error() {
let mut gate =
super::StreamCommitGate::new(super::StreamCommitPolicy::FirstAnthropicSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
});
assert_eq!(gate.observe_provider_bytes(b"event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"content\":[]}}\n\n"), super::StreamPrecommitObservation::Pending);
assert_eq!(gate.observe_provider_bytes(b"event: content_block_start\ndata: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n"), super::StreamPrecommitObservation::Pending);
assert!(matches!(gate.observe_provider_bytes(b"event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n"), super::StreamPrecommitObservation::UpstreamError { status_code: 529, .. }));
}
use std::time::Duration;
use super::{
@@ -553,7 +839,7 @@ mod tests {
false,
false,
)
.commits_on_response_headers());
.requires_bounded_frame_wait());
assert!(StreamCommitPolicy::for_response(
true,
Some("text/event-stream"),
@@ -563,7 +849,7 @@ mod tests {
true,
false,
)
.commits_on_response_headers());
.requires_bounded_frame_wait());
}
#[test]
@@ -745,8 +1031,8 @@ mod tests {
let mut gate = StreamCommitGate::new(native_anthropic_policy());
let observation = gate.observe_provider_bytes(
concat!(
"event: message_start\n",
"data: {\"type\":\"message_start\",\"message\":{}}\n\n",
"event: content_block_delta\n",
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n",
"event: error\n",
"data: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n",
)
File diff suppressed because it is too large Load Diff
@@ -529,7 +529,13 @@ fn classify_local_sync_error_kind(
{
return LocalCoreSyncErrorKind::Overloaded;
}
if (500..600).contains(&status_code) {
if (500..600).contains(&status_code)
|| raw_type.is_some_and(|value| {
["server_error", "internal_error", "api_error"]
.iter()
.any(|kind| value.trim().eq_ignore_ascii_case(kind))
})
{
return LocalCoreSyncErrorKind::ServerError;
}
LocalCoreSyncErrorKind::InvalidRequest
@@ -676,6 +682,13 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize(
#[cfg(test)]
mod tests {
#[test]
fn success_http_status_does_not_misclassify_explicit_server_errors_as_bad_requests() {
for error_type in ["server_error", "internal_error", "api_error"] {
let body = serde_json::json!({ "error": { "type": error_type, "message": "failed" } });
assert_eq!(super::resolve_local_sync_error_status_code(200, &body), 500);
}
}
use axum::body::to_bytes;
use serde_json::json;
@@ -21,10 +21,7 @@ use aether_contracts::{
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
};
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
use aether_http::{
apply_http_client_config, is_https_or_loopback_http_url, is_private_or_reserved_ip,
HttpClientConfig,
};
use aether_http::{apply_http_client_config, is_private_or_reserved_ip, HttpClientConfig};
use aether_runtime::{MetricKind, MetricSample};
use axum::body::Bytes;
use base64::Engine as _;
@@ -441,7 +438,7 @@ static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock<DirectHyperH2cSenderCacheMetric
LazyLock::new(DirectHyperH2cSenderCacheMetrics::default);
#[derive(Debug, Clone, Copy, Default)]
struct ExecutionSafeDnsResolver;
pub(crate) struct ExecutionSafeDnsResolver;
#[derive(Debug, Clone, Copy, Default)]
struct ExecutionSafeHyperDnsResolver;
@@ -449,10 +446,7 @@ struct ExecutionSafeHyperDnsResolver;
fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
let host = host.trim_end_matches('.');
host.eq_ignore_ascii_case("localhost")
|| host
.parse::<IpAddr>()
.map(|ip| ip.is_loopback())
.unwrap_or(false)
|| aether_http::parse_ip_literal_host(host).is_some_and(|ip| ip.is_loopback())
}
fn validate_resolved_execution_addresses(
@@ -494,12 +488,9 @@ async fn resolve_execution_target_addresses_with_policy(
port: u16,
provider_execution: bool,
) -> Result<Vec<SocketAddr>, std::io::Error> {
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
let addresses =
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await?
};
.await?;
validate_resolved_execution_addresses(host, addresses, provider_execution)
}
@@ -5152,7 +5143,7 @@ fn execution_log_url_host(url: &str) -> String {
.unwrap_or_else(|| "-".to_string())
}
fn validate_execution_upstream_url(
pub(crate) fn validate_execution_upstream_url(
raw_url: &str,
) -> Result<url::Url, ExecutionRuntimeTransportError> {
let url = url::Url::parse(raw_url).map_err(|_| {
@@ -5173,11 +5164,6 @@ fn validate_execution_upstream_url(
"upstream URL must not include a fragment".to_string(),
));
}
if !is_https_or_loopback_http_url(&url) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"remote upstream URL must use HTTPS".to_string(),
));
}
let literal_ip = match url.host() {
Some(url::Host::Ipv4(address)) => Some(IpAddr::V4(address)),
Some(url::Host::Ipv6(address)) => Some(IpAddr::V6(address)),
@@ -5324,7 +5310,7 @@ pub(crate) fn build_execution_response_body(
mod tests {
use std::collections::BTreeMap;
use std::io::{Read, Write};
use std::sync::{Arc, Mutex, MutexGuard, OnceLock};
use std::sync::{Arc, Mutex};
use aether_contracts::tunnel::{
TUNNEL_RELAY_AUTH_NONCE_HEADER, TUNNEL_RELAY_AUTH_PAYLOAD_HEADER,
@@ -5389,9 +5375,13 @@ mod tests {
const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes";
#[test]
fn execution_upstream_url_requires_https_or_literal_loopback_http() {
fn execution_upstream_url_accepts_http_and_https_with_safe_targets() {
for allowed in [
"https://api.example.test/v1/responses?api-version=1",
"http://api.example.test:8080/v1/responses?api-version=1",
"http://8.8.8.8:8080/v1/responses",
"https://8.8.8.8/v1/responses",
"http://[2606:4700:4700::1111]:8080/v1/responses",
"http://localhost:8080/v1/responses",
"http://127.42.0.1:8080/v1/responses",
"http://[::1]:8080/v1/responses",
@@ -5403,7 +5393,6 @@ mod tests {
}
for rejected in [
"http://api.example.test/v1/responses",
"http://10.0.0.1/v1/responses",
"http://0.0.0.0:8080/v1/responses",
"http://[::ffff:127.0.0.1]:8080/v1/responses",
@@ -5411,6 +5400,8 @@ mod tests {
"https://10.0.0.1:8443/v1/responses",
"https://[email protected]/v1/responses",
"https://example.test/v1/responses#secret",
"http://[email protected]/v1/responses",
"http://example.test/v1/responses#secret",
"ftp://localhost/resource",
] {
assert!(
@@ -5443,6 +5434,8 @@ mod tests {
"93.184.216.34:443".parse().unwrap(),
];
for host in [
"chatgpt.com",
"api.openai.com",
"oauth2.googleapis.com",
"www.googleapis.com",
"custom.example.test",
@@ -5455,6 +5448,46 @@ mod tests {
}
}
#[tokio::test]
async fn execution_dns_handles_url_ipv6_without_weakening_relay_filtering() {
for provider_execution in [false, true] {
let addresses = super::resolve_execution_target_addresses_with_policy(
"[::1]",
8443,
provider_execution,
)
.await
.expect("literal IPv6 loopback should resolve without DNS");
assert_eq!(addresses, vec!["[::1]:8443".parse().unwrap()]);
}
let error = super::resolve_execution_target_addresses_with_policy("[fd00::1]", 443, false)
.await
.expect_err("private IPv6 must remain blocked for relay traffic");
assert_eq!(error.kind(), std::io::ErrorKind::PermissionDenied);
}
#[tokio::test]
async fn execution_dns_resolvers_preserve_provider_fake_ip_answers() {
for host in ["198.18.78.41", "198.19.1.2"] {
let expected = vec![format!("{host}:0").parse::<std::net::SocketAddr>().unwrap()];
let reqwest_addresses = reqwest::dns::Resolve::resolve(
&super::ExecutionSafeDnsResolver,
host.parse().unwrap(),
)
.await
.expect("HTTP provider DNS must accept Fake-IP answers")
.collect::<Vec<_>>();
let wreq_addresses =
wreq::dns::Resolve::resolve(&super::ExecutionSafeDnsResolver, host.into())
.await
.expect("WebSocket provider DNS must accept Fake-IP answers")
.collect::<Vec<_>>();
assert_eq!(reqwest_addresses, expected);
assert_eq!(wreq_addresses, expected);
}
}
#[test]
fn execution_dns_answers_keep_relay_address_filtering() {
let public = "93.184.216.34:443".parse().unwrap();
@@ -6231,16 +6264,14 @@ mod tests {
TestEnvVarGuard { key, previous }
}
fn direct_reqwest_env_lock() -> MutexGuard<'static, ()> {
static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
LOCK.get_or_init(|| Mutex::new(()))
.lock()
.expect("direct reqwest env lock")
fn direct_reqwest_env_lock() -> &'static tokio::sync::Mutex<()> {
static LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
&LOCK
}
#[test]
fn direct_reqwest_client_cache_key_includes_transport_profile() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let timeouts = ExecutionTimeouts {
connect_ms: Some(5_000),
..ExecutionTimeouts::default()
@@ -6344,7 +6375,7 @@ mod tests {
#[test]
fn direct_reqwest_client_cache_evicts_least_recently_used_entry_at_capacity() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _capacity = set_test_env_var(super::DIRECT_REQWEST_CACHE_MAX_ENTRIES_ENV, "2");
let cache_key = |suffix| {
super::direct_reqwest_client_cache_key(
@@ -6442,7 +6473,7 @@ mod tests {
#[test]
fn direct_reqwest_client_cache_key_splits_origin_only_when_enabled() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-origin".into(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(),
@@ -6629,14 +6660,14 @@ mod tests {
#[test]
fn direct_h2c_client_shards_respect_explicit_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_H2C_CLIENT_SHARDS_ENV, "7");
assert_eq!(super::direct_h2c_client_shard_count(), 7);
}
#[test]
fn direct_h2c_adaptive_window_respects_explicit_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
{
let _adaptive = set_test_env_var(super::DIRECT_H2C_ADAPTIVE_WINDOW_ENV, "0");
assert!(!super::direct_h2c_adaptive_window_enabled());
@@ -6704,7 +6735,7 @@ mod tests {
#[test]
fn direct_h2c_prewarm_urls_parse_env_list() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _urls = set_test_env_var(
super::DIRECT_H2C_PREWARM_URLS_ENV,
" http://127.0.0.1:18184/v1/chat/completions,;http://127.0.0.1:18185/v1/chat/completions\nhttp://127.0.0.1:18186/v1/chat/completions ",
@@ -6722,7 +6753,7 @@ mod tests {
#[test]
fn direct_h2c_prewarm_cache_keys_dedup_by_origin() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let urls = vec![
"http://127.0.0.1:18184/v1/chat/completions".to_string(),
"http://127.0.0.1:18184/v1/responses".to_string(),
@@ -6748,7 +6779,7 @@ mod tests {
#[test]
fn direct_h2c_client_cache_splits_by_origin_and_shards() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_H2C_CLIENT_SHARDS_ENV, "3");
super::DIRECT_H2C_CLIENT_CACHE
.lock()
@@ -6773,7 +6804,7 @@ mod tests {
#[test]
fn direct_reqwest_initial_client_shards_are_bounded_by_target() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
assert_eq!(super::direct_reqwest_initial_client_shard_count(1), 1);
assert_eq!(super::direct_reqwest_initial_client_shard_count(2), 2);
assert_eq!(
@@ -6784,7 +6815,7 @@ mod tests {
#[test]
fn direct_reqwest_initial_client_shards_cap_large_sync_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "128");
assert_eq!(
super::direct_reqwest_initial_client_shard_count(128),
@@ -6794,7 +6825,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_client_shards_default_to_initial() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
assert_eq!(super::direct_reqwest_prewarm_client_shard_count(1), 1);
assert_eq!(
super::direct_reqwest_prewarm_client_shard_count(96),
@@ -6804,7 +6835,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_client_shards_do_not_exceed_request_path_cap() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "4");
let _prewarm = set_test_env_var(super::DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV, "128");
@@ -6813,7 +6844,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_populates_cache_for_plan() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "4");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-prewarm".into(),
@@ -6875,7 +6906,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_plan_keeps_large_sync_env_off_request_path() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "128");
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "4");
let _prewarm = set_test_env_var(super::DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV, "128");
@@ -6935,7 +6966,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_skips_h2c_fast_path() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _fast_path = set_test_env_var(super::DIRECT_H2C_FAST_PATH_ENV, "1");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-fast-path-prewarm-skip".into(),
@@ -6988,7 +7019,7 @@ mod tests {
#[test]
fn direct_reqwest_cache_metrics_expose_ready_state() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "1");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-ready-metrics".into(),
@@ -8572,7 +8603,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_supports_tunnel_relay() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET);
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -8739,7 +8770,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_rejects_short_tunnel_relay_secret_before_send() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", &"x".repeat(31));
let execution_runtime = DirectSyncExecutionRuntime::new();
let error = execution_runtime
@@ -8780,7 +8811,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_requires_tunnel_relay_secret_before_send() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = unset_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET");
let execution_runtime = DirectSyncExecutionRuntime::new();
let error = execution_runtime
@@ -9107,7 +9138,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_forwards_http1_only_control_to_tunnel_relay() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET);
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -9323,7 +9354,7 @@ mod tests {
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn direct_sync_execution_runtime_uses_h2c_prior_knowledge_on_wire() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().lock().await;
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "1");
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -522,7 +522,6 @@ where
decision,
plan_kind,
transfer_tracker,
request_first_byte_started_at: Instant::now(),
};
let loop_result = run_ai_attempt_loop(&port, plan_and_reports).await;
if loop_result.is_err() {
@@ -603,7 +602,6 @@ where
decision,
plan_kind,
transfer_tracker,
request_first_byte_started_at: Instant::now(),
};
let loop_result = run_dynamic_attempt_loop(
&port,
@@ -657,6 +655,73 @@ struct ProviderTransferState {
struct ProviderTransferStateTracker {
by_provider: BTreeMap<String, ProviderTransferState>,
exhausted_provider_ids: BTreeSet<String>,
global: GlobalTransferState,
}
#[derive(Debug, Default)]
struct GlobalTransferState {
first_attempt_started_at: Option<Instant>,
last_candidate: Option<(String, String, String)>,
transfer_count: u64,
limits: Option<ProviderTransferLimits>,
exhausted: bool,
}
impl GlobalTransferState {
fn load_policy(&mut self, report_context: Option<&serde_json::Value>) {
if self.limits.is_none() {
if let Some(policy) =
crate::orchestration::routing_execution_policy_from_report_context(report_context)
{
self.limits = Some(ProviderTransferLimits {
max_transfer_count: policy.max_transfer_count,
max_transfer_timeout_seconds: policy.max_transfer_timeout_seconds,
});
}
}
}
fn changes_candidate(&self, plan: &aether_contracts::ExecutionPlan) -> bool {
self.last_candidate
.as_ref()
.is_some_and(|(provider, endpoint, key)| {
provider != &plan.provider_id
|| endpoint != &plan.endpoint_id
|| key != &plan.key_id
})
}
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.first_attempt_started_at.get_or_insert(now);
if self.changes_candidate(plan) {
self.transfer_count = self.transfer_count.saturating_add(1);
}
self.last_candidate = Some((
plan.provider_id.clone(),
plan.endpoint_id.clone(),
plan.key_id.clone(),
));
}
fn check_before_attempt(
&mut self,
plan: &aether_contracts::ExecutionPlan,
now: Instant,
) -> Option<(bool, bool)> {
let limits = self.limits?;
let started_at = self.first_attempt_started_at?;
let count_reached = self.changes_candidate(plan)
&& limits.max_transfer_count > 0
&& self.transfer_count >= limits.max_transfer_count;
let timeout_reached = limits.max_transfer_timeout_seconds > 0
&& now.saturating_duration_since(started_at)
>= Duration::from_secs(limits.max_transfer_timeout_seconds);
if !count_reached && !timeout_reached {
return None;
}
self.exhausted = true;
Some((count_reached, timeout_reached))
}
}
#[derive(Clone, Debug, Default)]
@@ -719,6 +784,7 @@ struct ProviderTransferLimitReached {
impl ProviderTransferStateTracker {
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.global.record_attempt_started(plan, now);
match self.by_provider.entry(plan.provider_id.clone()) {
std::collections::btree_map::Entry::Vacant(entry) => {
entry.insert(ProviderTransferState {
@@ -905,11 +971,42 @@ async fn should_skip_provider_transfer_attempt<Attempt>(
where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let reached = tracker
.state
.lock()
.await
.check_before_attempt(attempt.execution_plan(), Instant::now());
let owned_report_context = attempt
.report_context_ref()
.is_none()
.then(|| attempt.report_context())
.flatten();
let report_context = attempt
.report_context_ref()
.or(owned_report_context.as_ref());
let mut tracker = tracker.state.lock().await;
tracker.global.load_policy(report_context);
if tracker.global.exhausted {
return true;
}
let now = Instant::now();
if let Some((count_reached, timeout_reached)) = tracker
.global
.check_before_attempt(attempt.execution_plan(), now)
{
warn!(
event_name = "routing_transfer_limit_reached",
log_type = "event",
trace_id,
plan_kind,
transfer_count = tracker.global.transfer_count,
elapsed_ms = tracker
.global
.first_attempt_started_at
.map(|started| now.saturating_duration_since(started).as_millis() as u64)
.unwrap_or(0),
count_reached,
timeout_reached,
"gateway exhausted the routing strategy transfer budget"
);
return true;
}
let reached = tracker.check_before_attempt(attempt.execution_plan(), now);
let Some(reached) = reached else {
return false;
};
@@ -1121,10 +1218,6 @@ struct StreamAttemptLoopPort<'a> {
decision: &'a GatewayControlDecision,
plan_kind: &'a str,
transfer_tracker: &'a ProviderTransferTracker,
/// All candidates in one downstream stream request share this origin.
/// Without it every retry receives a fresh full first-byte timeout and a
/// 30-second provider timeout can accumulate into a 60-120 second stall.
request_first_byte_started_at: Instant,
}
#[async_trait]
@@ -1254,7 +1347,6 @@ where
self.plan_kind,
plan,
watchdog_report_context,
self.request_first_byte_started_at,
stop_on_transport_errors,
move || async move {
if let Some(response) = execution_plan_cost_capacity_response(
@@ -1308,7 +1400,7 @@ where
http::StatusCode::GATEWAY_TIMEOUT.as_u16(),
"local_stream_candidate_watchdog_timeout",
stream_candidate_watchdog_timeout_message(),
self.request_first_byte_started_at.elapsed().as_millis() as u64,
watchdog_started_at.elapsed().as_millis() as u64,
)
.await?,
)
@@ -1758,7 +1850,6 @@ async fn execute_stream_candidate_with_watchdog<Fut>(
plan_kind: &str,
plan: &aether_contracts::ExecutionPlan,
report_context: Option<&serde_json::Value>,
request_first_byte_started_at: Instant,
stop_on_transport_errors: bool,
execute: impl FnOnce() -> Fut,
) -> Result<StreamCandidateWatchdogOutcome, GatewayError>
@@ -1768,7 +1859,6 @@ where
> + Send,
{
let timeout_duration = resolve_stream_candidate_watchdog_timeout(plan, report_context);
let request_first_byte_deadline = request_first_byte_started_at + timeout_duration;
let candidate_started_at = std::time::Instant::now();
let candidate_started_unix_ms = current_unix_ms();
let permit = match acquire_upstream_execution_gate(state, trace_id).await {
@@ -1794,14 +1884,7 @@ where
let watchdog_progress = StreamCandidateWatchdogProgress::shared();
let execution = watchdog_progress.clone().scope(execute());
tokio::pin!(execution);
// This is an absolute request-level deadline, not a new timeout for this
// candidate. Retries therefore consume only the budget left by earlier
// candidates instead of resetting the full provider timeout.
let candidate_budget_ms = request_first_byte_deadline
.saturating_duration_since(Instant::now())
.as_millis()
.min(u128::from(u64::MAX)) as u64;
let deadline = tokio::time::sleep_until(request_first_byte_deadline);
let deadline = tokio::time::sleep(timeout_duration);
tokio::pin!(deadline);
let execution_result = tokio::select! {
biased;
@@ -1830,10 +1913,6 @@ where
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string());
let timeout_ms = u64::try_from(timeout_duration.as_millis()).unwrap_or(u64::MAX);
let request_elapsed_ms = request_first_byte_started_at
.elapsed()
.as_millis()
.min(u128::from(u64::MAX)) as u64;
record_local_request_candidate_status(
state,
plan,
@@ -1862,8 +1941,6 @@ where
model_name,
candidate_index = candidate_index.as_str(),
timeout_ms,
candidate_budget_ms,
request_elapsed_ms,
"gateway local stream candidate watchdog timed out"
);
if stop_on_transport_errors {
@@ -2487,6 +2564,130 @@ mod tests {
assert_eq!(port.unused.lock().unwrap().as_slice(), ["a-key3-retry0"]);
}
#[tokio::test]
async fn routing_transfer_budget_counts_switches_across_providers_not_same_key_retries() {
for (limit, succeeds) in [(1, false), (2, true)] {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] =
json!({ "max_transfer_count": limit });
}
let outcome = run_ai_attempt_loop(&port, attempts).await.unwrap();
assert_eq!(
matches!(outcome, AiAttemptLoopOutcome::Responded(_)),
succeeds
);
{
let executed = port.executed.lock().unwrap();
assert_eq!(
&executed[..3],
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(executed.len(), if succeeds { 4 } else { 3 });
}
assert_eq!(port.tracker.state.lock().await.global.transfer_count, limit);
}
}
#[tokio::test]
async fn dynamic_loop_honors_global_transfer_budget_across_providers() {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let mut source = TransferTestAttemptSource {
attempts: attempts.into(),
skipped_providers: Vec::new(),
};
let outcome = run_dynamic_attempt_loop(
&port,
&mut source,
"global-budget",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
outcome,
LocalExecutionRequestOutcome::Exhausted(_)
));
assert_eq!(
port.executed.lock().unwrap().as_slice(),
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
}
#[test]
fn routing_time_budget_is_cumulative_and_zero_is_unlimited() {
let mut global = super::GlobalTransferState::default();
global.load_policy(Some(
&json!({ "routing_execution_policy": { "max_transfer_timeout_seconds": 60 } }),
));
let now = tokio::time::Instant::now();
let plan = test_plan(None);
global.record_attempt_started(&plan, now);
global.record_attempt_started(&plan, now + Duration::from_secs(40));
assert_eq!(global.transfer_count, 0);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(59)),
None
);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(60)),
Some((false, true))
);
let mut unlimited = super::GlobalTransferState::default();
unlimited.load_policy(Some(&json!({ "routing_execution_policy": {} })));
unlimited.record_attempt_started(&plan, now);
assert_eq!(
unlimited.check_before_attempt(&plan, now + Duration::from_secs(86_400)),
None
);
}
#[tokio::test]
async fn cloned_tracker_preserves_global_budget_across_candidate_loops() {
let state = AppState::new().unwrap();
let tracker = ProviderTransferTracker::default();
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let remaining = attempts.split_off(3);
let first_port = TransferTestPort::with_tracker(&state, tracker.clone());
let first_outcome = run_ai_attempt_loop(&first_port, attempts).await.unwrap();
assert!(matches!(first_outcome, AiAttemptLoopOutcome::Exhausted(_)));
assert_eq!(tracker.state.lock().await.global.transfer_count, 1);
let second_port = TransferTestPort::with_tracker(&state, tracker.clone());
let mut source = TransferTestAttemptSource {
attempts: remaining.into(),
skipped_providers: Vec::new(),
};
let second_outcome = run_dynamic_attempt_loop(
&second_port,
&mut source,
"global-budget-across-loops",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
second_outcome,
LocalExecutionRequestOutcome::NoPath
));
assert!(second_port.executed.lock().unwrap().is_empty());
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
assert!(tracker.state.lock().await.global.exhausted);
}
#[tokio::test]
async fn cloned_tracker_preserves_transfer_budget_across_candidate_loops() {
let state = AppState::new().expect("state should build");
@@ -3150,7 +3351,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
false,
|| {
std::future::pending::<
@@ -3189,39 +3389,32 @@ mod tests {
assert_eq!(record.candidate_index, 2);
}
#[tokio::test]
async fn stream_candidate_retry_does_not_reset_an_expired_request_first_byte_budget() {
let writer = Arc::new(TestRequestCandidateWriter::default());
async fn assert_stream_candidate_retry_gets_fresh_first_byte_budget(
provider_id: &str,
key_id: &str,
first_byte_ms: u64,
) {
let writer = TestRequestCandidateWriter::default();
let plan = test_plan(Some(ExecutionTimeouts {
first_byte_ms: Some(250),
first_byte_ms: Some(100),
..ExecutionTimeouts::default()
}));
let report_context = test_report_context();
// Stand in for earlier candidates having already consumed the request's
// complete first-byte budget. A per-candidate watchdog would wait a new
// 250 ms here; the shared absolute deadline must settle immediately.
let request_first_byte_started_at = Instant::now() - Duration::from_millis(300);
let result = tokio::time::timeout(
Duration::from_millis(100),
execute_stream_candidate_with_watchdog(
writer.as_ref(),
"trace_watchdog_shared_budget",
"claude_cli_stream",
&plan,
Some(&report_context),
request_first_byte_started_at,
false,
|| {
std::future::pending::<
Result<AiAttemptExecutionOutcome<Response<Body>>, GatewayError>,
>()
},
),
let result = execute_stream_candidate_with_watchdog(
&writer,
"trace_watchdog_retry_budget",
"claude_cli_stream",
&plan,
Some(&report_context),
false,
|| {
std::future::pending::<
Result<AiAttemptExecutionOutcome<Response<Body>>, GatewayError>,
>()
},
)
.await
.expect("an expired request-level first-byte budget must not restart per candidate");
.await;
assert!(matches!(
result,
Ok(StreamCandidateWatchdogOutcome::Executed(
@@ -3231,14 +3424,119 @@ mod tests {
}
))
));
let mut next_plan = plan.clone();
next_plan.candidate_id = Some("cand_watchdog_retry".to_string());
next_plan.provider_id = provider_id.to_string();
next_plan.key_id = key_id.to_string();
next_plan.timeouts = Some(ExecutionTimeouts {
first_byte_ms: Some(first_byte_ms),
..ExecutionTimeouts::default()
});
let mut next_report_context = report_context.clone();
next_report_context["candidate_id"] = json!("cand_watchdog_retry");
next_report_context["candidate_index"] = json!(3);
let result = execute_stream_candidate_with_watchdog(
&writer,
"trace_watchdog_retry_budget",
"claude_cli_stream",
&next_plan,
Some(&next_report_context),
false,
|| async {
tokio::time::sleep(Duration::from_millis(60)).await;
Ok(AiAttemptExecutionOutcome::Responded(Response::new(
Body::from("retry succeeded"),
)))
},
)
.await;
assert!(
matches!(
result,
Ok(StreamCandidateWatchdogOutcome::Executed(
AiAttemptExecutionOutcome::Responded(_)
))
),
"candidate {provider_id}/{key_id} must receive its own {first_byte_ms} ms budget"
);
let records = writer.records.lock().await;
assert_eq!(records.len(), 1);
assert_eq!(records[0].id, plan.candidate_id.as_deref().unwrap());
assert_eq!(records[0].status, RequestCandidateStatus::Failed);
assert_eq!(
records[0].error_type.as_deref(),
Some("local_stream_candidate_watchdog_timeout")
);
}
#[tokio::test]
async fn stream_candidate_watchdog_failover_gets_fresh_first_byte_budget() {
for first_byte_ms in [100, 75, 150] {
assert_stream_candidate_retry_gets_fresh_first_byte_budget(
"provider_next",
"key_next",
first_byte_ms,
)
.await;
}
}
#[tokio::test]
async fn stream_candidate_watchdog_same_provider_retries_get_fresh_first_byte_budget() {
for key_id in ["key_next", "key_id"] {
assert_stream_candidate_retry_gets_fresh_first_byte_budget("provider_id", key_id, 100)
.await;
}
}
#[tokio::test]
async fn stream_candidate_watchdog_starts_first_byte_budget_after_admission() {
let writer = TestRequestCandidateWriter::with_upstream_gate(1, Duration::from_secs(1));
let held_permit = writer
.upstream_gate
.as_ref()
.expect("test gate should exist")
.try_acquire()
.expect("test gate permit should acquire");
let plan = test_plan(Some(ExecutionTimeouts {
first_byte_ms: Some(50),
..ExecutionTimeouts::default()
}));
let report_context = test_report_context();
let (result, ()) = tokio::join!(
execute_stream_candidate_with_watchdog(
&writer,
"trace_watchdog_admission_budget",
"claude_cli_stream",
&plan,
Some(&report_context),
false,
|| async {
tokio::time::sleep(Duration::from_millis(20)).await;
Ok(AiAttemptExecutionOutcome::Responded(Response::new(
Body::from("admitted candidate succeeded"),
)))
},
),
async move {
tokio::time::sleep(Duration::from_millis(100)).await;
drop(held_permit);
},
);
assert!(matches!(
result,
Ok(StreamCandidateWatchdogOutcome::Executed(
AiAttemptExecutionOutcome::Responded(_)
))
));
assert!(writer.records.lock().await.is_empty());
}
#[tokio::test]
async fn stream_candidate_watchdog_can_stop_on_transport_error() {
let writer = Arc::new(TestRequestCandidateWriter::default());
@@ -3254,7 +3552,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
true,
|| {
std::future::pending::<
@@ -3292,7 +3589,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
true,
|| async {
mark_stream_candidate_watchdog_terminal_started();
@@ -3325,7 +3621,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
true,
|| async {
Err(GatewayError::UpstreamUnavailable {
@@ -3365,7 +3660,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
false,
|| async {
panic!("execute future should not run while upstream execution gate is saturated")
@@ -3413,7 +3707,6 @@ mod tests {
"claude_cli_stream",
&plan,
Some(&report_context),
Instant::now(),
false,
|| async {
Err(GatewayError::AdmissionTimeout {
@@ -822,15 +822,20 @@ where
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = standard_text_sync_heartbeat_final_bytes(
let completion = standard_text_sync_heartbeat_final_bytes(
client_api_format.as_str(),
redaction_slot.as_ref(),
execute(state, parts, trace_id, decision, plan_kind, started_at).await,
)
.await;
tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execute(state, parts, trace_id, decision, plan_kind, started_at) => result,
},
);
let bytes = completion.await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -1097,23 +1102,26 @@ fn build_openai_image_sync_heartbeat_shell_response(
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = openai_image_sync_heartbeat_final_bytes(
execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
)
.await,
)
.await;
let execution = execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
);
let outcome = tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execution => result,
};
let bytes = openai_image_sync_heartbeat_final_bytes(outcome).await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -2331,6 +2339,45 @@ mod tests {
.expect("background completion should release admission");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_cancels_when_routing_policy_enables_it() {
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (mut release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let response = crate::request_lifecycle::run_request(async move {
crate::request_lifecycle::configure_client_disconnect(
aether_routing_core::RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
},
);
let (parts, _) = http::Request::builder()
.method("POST")
.uri("/v1/responses")
.body(())
.unwrap()
.into_parts();
build_standard_text_sync_heartbeat_shell_response(
AppState::new().unwrap(),
parts,
"trace-heartbeat-disconnect".to_string(),
test_standard_text_heartbeat_decision(),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND.to_string(),
move |_, _, _, _, _, _| async move {
started_tx.send(()).unwrap();
release_rx.await.unwrap();
Ok(LocalExecutionRequestOutcome::NoPath)
},
)
})
.await
.unwrap();
started_rx.await.unwrap();
drop(response);
tokio::time::timeout(Duration::from_secs(1), release_tx.closed())
.await
.expect("heartbeat must drop upstream execution immediately");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_propagates_request_diagnostics_to_terminal_usage() {
let (state, usage_repository) = heartbeat_usage_test_state(json!({
+29 -43
View File
@@ -756,13 +756,8 @@ fn runtime_miss_client_error_body(api_format: Option<&str>, message: &str) -> Va
}
fn runtime_miss_original_headers_json(headers: &HeaderMap) -> Value {
let mut headers = crate::headers::collect_control_headers(headers);
for (name, value) in headers.iter_mut() {
if runtime_miss_sensitive_header(name) {
*value = runtime_miss_mask_header_value(value);
}
}
serde_json::to_value(headers).unwrap_or_else(|_| json!({}))
serde_json::to_value(crate::headers::collect_control_headers(headers))
.unwrap_or_else(|_| json!({}))
}
fn runtime_miss_original_request_body_json(
@@ -784,40 +779,6 @@ fn runtime_miss_original_request_body_json(
})
}
fn runtime_miss_sensitive_header(name: &str) -> bool {
const SENSITIVE_HEADERS: &[&str] = &[
"authorization",
"x-api-key",
"api-key",
"x-goog-api-key",
"cookie",
"proxy-authorization",
];
SENSITIVE_HEADERS
.iter()
.any(|candidate| name.eq_ignore_ascii_case(candidate))
}
fn runtime_miss_mask_header_value(value: &str) -> String {
let value = value.trim();
let char_count = value.chars().count();
if char_count <= 8 {
return "****".to_string();
}
let prefix: String = value.chars().take(4).collect();
let suffix: String = value
.chars()
.rev()
.take(4)
.collect::<Vec<_>>()
.into_iter()
.rev()
.collect();
format!("{prefix}****{suffix}")
}
async fn load_runtime_miss_candidate_contexts(
state: &AppState,
request_id: &str,
@@ -1233,8 +1194,9 @@ mod tests {
apply_runtime_miss_usage_routing, beautify_local_execution_client_error_message,
insert_runtime_miss_candidate_usage_metadata,
request_candidate_represents_provider_execution, runtime_miss_client_error_body,
select_last_runtime_miss_executed_candidate, select_last_runtime_miss_routing_candidate,
LocalExecutionRuntimeMissContext, RuntimeMissCandidateContext,
runtime_miss_original_headers_json, select_last_runtime_miss_executed_candidate,
select_last_runtime_miss_routing_candidate, LocalExecutionRuntimeMissContext,
RuntimeMissCandidateContext,
};
use crate::constants::EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS;
use crate::state::LocalExecutionRuntimeMissDiagnostic;
@@ -1266,6 +1228,30 @@ mod tests {
);
}
#[test]
fn runtime_miss_usage_preserves_original_request_headers() {
let expected = json!({
"authorization": "Bearer original-client-token",
"x-api-key": "short",
"api-key": "original-api-key",
"x-goog-api-key": "original-google-key",
"cookie": "session=original-client",
"proxy-authorization": "Basic original-proxy-token",
"originator": "codex-cli",
"session-id": "original-session",
"x-codex-turn-metadata": "{\"turn_id\":\"original-turn\"}"
});
let mut headers = http::HeaderMap::new();
for (name, value) in expected.as_object().unwrap() {
headers.insert(
http::HeaderName::from_bytes(name.as_bytes()).unwrap(),
http::HeaderValue::from_str(value.as_str().unwrap()).unwrap(),
);
}
assert_eq!(runtime_miss_original_headers_json(&headers), expected);
}
#[test]
fn runtime_miss_usage_body_matches_claude_client_envelope() {
let claude = runtime_miss_client_error_body(Some("claude:messages"), "busy");
@@ -339,57 +339,6 @@ async fn build_admin_oauth_test_payload(
}))
}
#[cfg(test)]
mod tests {
use super::{
is_fixed_linuxdo_oauth_origin, resolve_public_admin_oauth_endpoint,
validate_public_admin_oauth_resolved_addrs,
};
use std::net::SocketAddr;
#[tokio::test]
async fn oauth_test_endpoint_rejects_loopback_https_targets_before_connecting() {
let url = reqwest::Url::parse("https://127.0.0.1/oauth/token").expect("URL");
assert!(resolve_public_admin_oauth_endpoint(&url).await.is_err());
}
#[test]
fn linuxdo_builtin_origin_allows_only_benchmarking_addresses() {
let fixed = reqwest::Url::parse("https://connect.linux.do/oauth2/token")
.expect("LinuxDo URL should parse");
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
assert!(is_fixed_linuxdo_oauth_origin(&fixed));
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], true).is_ok());
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], false).is_err());
assert!(validate_public_admin_oauth_resolved_addrs(
&fixed,
&[fake, SocketAddr::from(([127, 0, 0, 1], 443))],
true,
)
.is_err());
}
#[test]
fn custom_or_non_default_oauth_origins_reject_benchmarking_addresses() {
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
for raw_url in [
"https://oauth.example.test/token",
"https://connect.linux.do:8443/oauth2/token",
"https://connect.linuxdo.org/oauth2/token",
"https://connect.linux.do.evil.test/oauth2/token",
"https://connect.linux.do/oauth2/token?tenant=unexpected",
] {
let url = reqwest::Url::parse(raw_url).expect("test URL should parse");
assert!(
!is_fixed_linuxdo_oauth_origin(&url),
"must not trust {raw_url}"
);
assert!(validate_public_admin_oauth_resolved_addrs(&url, &[fake], true).is_err());
}
}
}
pub(crate) async fn maybe_build_local_admin_oauth_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -689,3 +638,54 @@ pub(crate) async fn maybe_build_local_admin_oauth_response(
Ok(None)
}
#[cfg(test)]
mod tests {
use super::{
is_fixed_linuxdo_oauth_origin, resolve_public_admin_oauth_endpoint,
validate_public_admin_oauth_resolved_addrs,
};
use std::net::SocketAddr;
#[tokio::test]
async fn oauth_test_endpoint_rejects_loopback_https_targets_before_connecting() {
let url = reqwest::Url::parse("https://127.0.0.1/oauth/token").expect("URL");
assert!(resolve_public_admin_oauth_endpoint(&url).await.is_err());
}
#[test]
fn linuxdo_builtin_origin_allows_only_benchmarking_addresses() {
let fixed = reqwest::Url::parse("https://connect.linux.do/oauth2/token")
.expect("LinuxDo URL should parse");
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
assert!(is_fixed_linuxdo_oauth_origin(&fixed));
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], true).is_ok());
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], false).is_err());
assert!(validate_public_admin_oauth_resolved_addrs(
&fixed,
&[fake, SocketAddr::from(([127, 0, 0, 1], 443))],
true,
)
.is_err());
}
#[test]
fn custom_or_non_default_oauth_origins_reject_benchmarking_addresses() {
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
for raw_url in [
"https://oauth.example.test/token",
"https://connect.linux.do:8443/oauth2/token",
"https://connect.linuxdo.org/oauth2/token",
"https://connect.linux.do.evil.test/oauth2/token",
"https://connect.linux.do/oauth2/token?tenant=unexpected",
] {
let url = reqwest::Url::parse(raw_url).expect("test URL should parse");
assert!(
!is_fixed_linuxdo_oauth_origin(&url),
"must not trust {raw_url}"
);
assert!(validate_public_admin_oauth_resolved_addrs(&url, &[fake], true).is_err());
}
}
}
@@ -278,46 +278,6 @@ async fn build_batch_delete_global_models_response(
))
}
#[cfg(test)]
mod batch_boundary_tests {
use super::{normalize_admin_global_model_batch_ids, MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS};
#[test]
fn global_model_batch_ids_are_bounded_and_deduplicated() {
assert_eq!(
normalize_admin_global_model_batch_ids(
vec![
"model-2".to_string(),
"model-1".to_string(),
" model-2 ".to_string(),
" ".to_string(),
],
"ids",
)
.expect("valid ids"),
vec![
"model-2".to_string(),
"model-1".to_string(),
" ".to_string(),
]
);
assert!(normalize_admin_global_model_batch_ids(
(0..=MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("model-{index}"))
.collect(),
"ids",
)
.is_err());
assert!(normalize_admin_global_model_batch_ids(
(0..MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("provider-{index}"))
.collect(),
"provider_ids",
)
.is_ok());
}
}
async fn build_assign_to_providers_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -357,3 +317,43 @@ async fn build_assign_to_providers_response(
&global_model_id,
))
}
#[cfg(test)]
mod batch_boundary_tests {
use super::{normalize_admin_global_model_batch_ids, MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS};
#[test]
fn global_model_batch_ids_are_bounded_and_deduplicated() {
assert_eq!(
normalize_admin_global_model_batch_ids(
vec![
"model-2".to_string(),
"model-1".to_string(),
" model-2 ".to_string(),
" ".to_string(),
],
"ids",
)
.expect("valid ids"),
vec![
"model-2".to_string(),
"model-1".to_string(),
" ".to_string(),
]
);
assert!(normalize_admin_global_model_batch_ids(
(0..=MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("model-{index}"))
.collect(),
"ids",
)
.is_err());
assert!(normalize_admin_global_model_batch_ids(
(0..MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("provider-{index}"))
.collect(),
"provider_ids",
)
.is_ok());
}
}
@@ -58,6 +58,7 @@ async fn admin_monitoring_trace_request_returns_local_payload() {
.expect("route should be handled locally");
assert_eq!(response.status(), http::StatusCode::OK);
assert!(response.headers().contains_key("x-aether-build-version"));
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
@@ -93,6 +94,15 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
Some(33),
Some(200),
),
sample_candidate(
"cand-other-attempt",
"trace-1",
1,
RequestCandidateStatus::Failed,
Some(100),
Some(20),
Some(502),
),
]));
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
@@ -110,6 +120,8 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
100,
);
usage.id = "usage-row-1".to_string();
usage.request_body_state = Some(UsageBodyCaptureState::Reference);
usage.response_body_state = Some(UsageBodyCaptureState::Reference);
usage.candidate_id = Some("cand-used".to_string());
usage.request_headers = Some(json!({
"x-trace-id": "trace-1"
@@ -140,6 +152,17 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["request_id"], json!("trace-1"));
assert_eq!(payload["diagnostic_request"]["usage_id"], "usage-row-1");
assert_eq!(
payload["candidates"][0]["extra_data"]["diagnostic_context"]["usage_id"],
"usage-row-1"
);
assert_eq!(
payload["candidates"][0]["extra_data"]["diagnostic_context"]["body_states"]
["response_body"],
"reference"
);
assert!(payload["candidates"][1]["extra_data"]["diagnostic_context"].is_null());
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
assert_eq!(
payload["candidates"][0]["extra_data"]["first_byte_time_ms"],
@@ -67,13 +67,19 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
let key_accounts =
build_admin_monitoring_key_account_display_map(admin_state, &resolved.trace).await?;
Ok(
build_admin_monitoring_trace_request_payload_response_with_key_accounts(
&resolved.trace,
resolved.usage.as_ref(),
&key_accounts,
),
)
let mut response = build_admin_monitoring_trace_request_payload_response_with_key_accounts(
&resolved.trace,
resolved.usage.as_ref(),
&key_accounts,
);
if let Ok(version) = axum::http::HeaderValue::from_str(
option_env!("AETHER_BUILD_VERSION").unwrap_or(env!("CARGO_PKG_VERSION")),
) {
response
.headers_mut()
.insert("x-aether-build-version", version);
}
Ok(response)
}
async fn resolve_admin_monitoring_trace(
@@ -17,7 +17,10 @@ use aether_admin::observability::usage::{
admin_usage_bad_request_response, admin_usage_data_unavailable_response,
admin_usage_provider_key_name, ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
};
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageBodyField};
use aether_data_contracts::repository::usage::{
canonical_usage_body_ref_for, StoredRequestUsageAudit, StoredUsageBodyPayload,
UsageBodyCaptureState, UsageBodyField, MAX_DECOMPRESSED_USAGE_JSON_BYTES,
};
use axum::{
body::Body,
http,
@@ -28,9 +31,59 @@ use serde_json::{json, Value};
use std::collections::BTreeMap;
use tokio::try_join;
#[derive(Default)]
struct AdminUsageDetailBodyValue {
value: Option<Value>,
load_failed: bool,
error_code: Option<&'static str>,
}
impl AdminUsageDetailBodyValue {
fn resolved(
item: &StoredRequestUsageAudit,
field: UsageBodyField,
value: Option<Value>,
) -> Self {
let missing = value.is_none()
&& item
.body_capture_result(field, item.body_value(field))
.available;
Self {
value,
error_code: missing.then_some("missing"),
}
}
}
fn admin_usage_body_load_error_code(error: &GatewayError) -> &'static str {
if let GatewayError::Internal(message) = error {
if message.contains("decompressed usage json exceeds ")
|| message.contains("encoded usage json exceeds ")
{
return "too_large";
}
if message.contains("failed to decompress usage json:")
|| message.contains("failed to parse decompressed usage json:")
{
return "decode_failed";
}
}
"storage_unavailable"
}
async fn resolve_admin_usage_detail_field(
state: &AdminAppState<'_>,
item: &StoredRequestUsageAudit,
field: UsageBodyField,
selected_field: Option<UsageBodyField>,
) -> AdminUsageDetailBodyValue {
if selected_field.is_some_and(|selected| selected != field) {
return AdminUsageDetailBodyValue::default();
}
if field == UsageBodyField::RequestBody {
resolve_admin_usage_detail_request_body(state, item).await
} else {
resolve_admin_usage_detail_body_value(state, item, field).await
}
}
async fn resolve_admin_usage_detail_request_body(
@@ -38,10 +91,7 @@ async fn resolve_admin_usage_detail_request_body(
item: &StoredRequestUsageAudit,
) -> AdminUsageDetailBodyValue {
match admin_usage_resolve_request_capture_body_for_item(state, item, None).await {
Ok(body) => AdminUsageDetailBodyValue {
value: body,
load_failed: false,
},
Ok(body) => AdminUsageDetailBodyValue::resolved(item, UsageBodyField::RequestBody, body),
Err(err) => {
tracing::warn!(
error = ?err,
@@ -52,7 +102,9 @@ async fn resolve_admin_usage_detail_request_body(
);
let value = admin_usage_resolve_request_capture_body(item, None);
AdminUsageDetailBodyValue {
load_failed: value.is_none(),
error_code: value
.is_none()
.then(|| admin_usage_body_load_error_code(&err)),
value,
}
}
@@ -66,10 +118,7 @@ async fn resolve_admin_usage_detail_body_value(
) -> AdminUsageDetailBodyValue {
let inline_body = item.body_value(field);
match admin_usage_resolve_body_value(state, item, inline_body, field).await {
Ok(body) => AdminUsageDetailBodyValue {
value: body,
load_failed: false,
},
Ok(body) => AdminUsageDetailBodyValue::resolved(item, field, body),
Err(err) => {
tracing::warn!(
error = ?err,
@@ -80,13 +129,139 @@ async fn resolve_admin_usage_detail_body_value(
);
let value = inline_body.cloned();
AdminUsageDetailBodyValue {
load_failed: value.is_none(),
error_code: value
.is_none()
.then(|| admin_usage_body_load_error_code(&err)),
value,
}
}
}
}
async fn read_admin_usage_raw_body(
state: &AdminAppState<'_>,
item: &StoredRequestUsageAudit,
field: UsageBodyField,
) -> Result<Option<StoredUsageBodyPayload>, GatewayError> {
if matches!(
item.body_state(field),
Some(
UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
| UsageBodyCaptureState::None
)
) {
return Ok(None);
}
let inline_body = item.body_value(field);
let prefer_inline = matches!(
item.body_state(field),
Some(UsageBodyCaptureState::Inline | UsageBodyCaptureState::Truncated)
) && inline_body.is_some();
if !prefer_inline {
if let Some(body_ref) = item
.body_ref(field)
.and_then(|reference| canonical_usage_body_ref_for(reference, &item.request_id, field))
{
if let Some(payload) = state.read_request_usage_body_payload(&body_ref).await? {
return Ok(Some(payload));
}
}
}
let fallback = inline_body.cloned().or_else(|| {
(field == UsageBodyField::RequestBody)
.then(|| admin_usage_resolve_request_capture_body(item, None))
.flatten()
});
fallback
.map(|value| {
serde_json::to_vec(&value)
.map(StoredUsageBodyPayload::Json)
.map_err(|error| GatewayError::Internal(error.to_string()))
})
.transpose()
}
async fn build_admin_usage_raw_body_response(
state: &AdminAppState<'_>,
item: &StoredRequestUsageAudit,
field: UsageBodyField,
) -> Response<Body> {
let result = read_admin_usage_raw_body(state, item, field).await;
let mut response = match result {
Ok(Some(payload)) => admin_usage_raw_payload_response(payload),
Ok(None) => admin_usage_raw_body_error(http::StatusCode::NOT_FOUND, "missing"),
Err(error) => {
tracing::warn!(error = ?error, usage_id = %item.id, field = field.as_storage_field(), "failed to read admin usage raw body");
let code = admin_usage_body_load_error_code(&error);
admin_usage_raw_body_error(
if code == "too_large" {
http::StatusCode::PAYLOAD_TOO_LARGE
} else {
http::StatusCode::SERVICE_UNAVAILABLE
},
code,
)
}
};
let headers = response.headers_mut();
headers.insert(
http::header::CACHE_CONTROL,
http::HeaderValue::from_static("no-store, no-transform"),
);
headers.insert(
"x-content-type-options",
http::HeaderValue::from_static("nosniff"),
);
headers.insert(
"x-aether-body-field",
http::HeaderValue::from_static(field.as_storage_field()),
);
if let Ok(value) = http::HeaderValue::from_str(&item.id) {
headers.insert("x-aether-usage-id", value);
}
attach_admin_audit_response(
response,
"admin_usage_detail_viewed",
"view_usage_detail",
"usage_record",
&item.id,
)
}
fn admin_usage_raw_payload_response(payload: StoredUsageBodyPayload) -> Response<Body> {
let (encoding, bytes, limit) = match payload {
StoredUsageBodyPayload::Gzip(bytes) => (
"gzip",
bytes,
MAX_DECOMPRESSED_USAGE_JSON_BYTES + 1024 * 1024,
),
StoredUsageBodyPayload::Json(bytes) => ("json", bytes, MAX_DECOMPRESSED_USAGE_JSON_BYTES),
};
if bytes.len() > limit {
admin_usage_raw_body_error(http::StatusCode::PAYLOAD_TOO_LARGE, "too_large")
} else {
(
[
("content-type", "application/octet-stream"),
("content-encoding", "identity"),
("x-aether-body-encoding", encoding),
],
bytes,
)
.into_response()
}
}
fn admin_usage_raw_body_error(status: http::StatusCode, code: &'static str) -> Response<Body> {
(
status,
[("x-aether-body-error", code)],
Json(json!({ "body_load_error_code": code })),
)
.into_response()
}
pub(super) async fn maybe_build_local_admin_usage_detail_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -212,6 +387,25 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
"include_bodies",
true,
);
let body_field =
match request_context
.request_query_string
.as_deref()
.and_then(|query| {
url::form_urlencoded::parse(query.as_bytes())
.find(|(key, _)| key == "body_field")
.map(|(_, value)| value.into_owned())
}) {
Some(value) => {
match UsageBodyField::from_storage_field(value.trim()) {
Some(field) if include_bodies => Some(field),
_ => return Ok(Some(admin_usage_bad_request_response(
"body_field 必须是有效的正文字段,且 include_bodies 必须为 true",
))),
}
}
None => None,
};
let Some(item) = state.find_request_usage_by_id(&usage_id).await? else {
return Ok(Some(
@@ -223,6 +417,27 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
));
};
let body_format = request_context
.request_query_string
.as_deref()
.and_then(|query| {
url::form_urlencoded::parse(query.as_bytes())
.find(|(key, _)| key == "body_format")
.map(|(_, value)| value.into_owned())
});
if let Some(format) = body_format {
if format != "raw" || body_field.is_none() {
return Ok(Some(admin_usage_bad_request_response(
"body_format=raw 必须指定 body_field",
)));
}
if let Some(field) = body_field {
return Ok(Some(
build_admin_usage_raw_body_response(state, &item, field).await,
));
}
}
let user_ids = item.user_id.clone().into_iter().collect::<Vec<_>>();
let (users_by_id, provider_key_names, api_key_names): (
BTreeMap<String, aether_data::repository::users::StoredUserSummary>,
@@ -250,23 +465,32 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
}
}
let mut body_load_errors = serde_json::Map::new();
let request_body = if include_bodies {
let mut body_load_error_codes = serde_json::Map::new();
let mut request_body = if include_bodies {
let (request_body, provider_request_body, response_body, client_response_body) = tokio::join!(
resolve_admin_usage_detail_request_body(state, &item),
resolve_admin_usage_detail_body_value(
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::RequestBody,
body_field
),
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::ProviderRequestBody,
body_field,
),
resolve_admin_usage_detail_body_value(
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::ResponseBody,
body_field,
),
resolve_admin_usage_detail_body_value(
resolve_admin_usage_detail_field(
state,
&item,
UsageBodyField::ClientResponseBody,
body_field,
),
);
for (field, resolved) in [
@@ -275,20 +499,25 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
(UsageBodyField::ResponseBody, &response_body),
(UsageBodyField::ClientResponseBody, &client_response_body),
] {
if resolved.load_failed {
if let Some(error_code) = resolved.error_code {
body_load_errors.insert(field.as_storage_field().to_string(), json!(true));
body_load_error_codes
.insert(field.as_storage_field().to_string(), json!(error_code));
}
}
detail_item.provider_request_body = provider_request_body.value;
detail_item.response_body = response_body.value;
detail_item.client_response_body = client_response_body.value;
if body_field.is_none_or(|field| field == UsageBodyField::ProviderRequestBody) {
detail_item.provider_request_body = provider_request_body.value;
}
if body_field.is_none_or(|field| field == UsageBodyField::ResponseBody) {
detail_item.response_body = response_body.value;
}
if body_field.is_none_or(|field| field == UsageBodyField::ClientResponseBody) {
detail_item.client_response_body = client_response_body.value;
}
request_body.value
} else {
None
};
if include_bodies {
// request_body 已通过 request capture 解析;其余 detached body 在上方并行加载。
}
let default_headers = admin_usage_curl_headers();
let mut payload = build_admin_usage_detail_payload(
&detail_item,
@@ -297,15 +526,33 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
state.has_auth_user_data_reader(),
state.has_auth_api_key_data_reader(),
provider_key_name.as_deref(),
include_bodies,
request_body,
include_bodies && body_field.is_none(),
if body_field.is_none() {
request_body.take()
} else {
None
},
&default_headers,
);
if let Some(field) = body_field {
payload[field.as_storage_field()] = match field {
UsageBodyField::RequestBody => request_body,
UsageBodyField::ProviderRequestBody => detail_item.provider_request_body.take(),
UsageBodyField::ResponseBody => detail_item.response_body.take(),
UsageBodyField::ClientResponseBody => detail_item.client_response_body.take(),
}
.unwrap_or(Value::Null);
}
payload["body_load_errors"] = if include_bodies && !body_load_errors.is_empty() {
Value::Object(body_load_errors)
} else {
Value::Null
};
payload["body_load_error_codes"] = if body_load_error_codes.is_empty() {
Value::Null
} else {
Value::Object(body_load_error_codes)
};
return Ok(Some(attach_admin_audit_response(
Json(payload).into_response(),
@@ -320,3 +567,61 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
Ok(None)
}
#[cfg(test)]
mod tests {
use super::admin_usage_body_load_error_code;
use crate::GatewayError;
#[tokio::test]
async fn admin_usage_raw_body_does_not_decode_or_reencode_stored_bytes() {
use super::{admin_usage_raw_payload_response, StoredUsageBodyPayload};
for (payload, encoding, expected) in [
(
StoredUsageBodyPayload::Gzip(vec![31, 139, 8, 0, 1]),
"gzip",
vec![31, 139, 8, 0, 1],
),
(
StoredUsageBodyPayload::Json(b"{ \"untouched\" : true }".to_vec()),
"json",
b"{ \"untouched\" : true }".to_vec(),
),
] {
let response = admin_usage_raw_payload_response(payload);
assert_eq!(response.headers()["content-encoding"], "identity");
assert_eq!(response.headers()["x-aether-body-encoding"], encoding);
let bytes = axum::body::to_bytes(response.into_body(), 1024)
.await
.unwrap();
assert_eq!(bytes.as_ref(), expected.as_slice());
}
}
#[test]
fn body_load_errors_expose_safe_codes_instead_of_internal_messages() {
for (message, expected) in [
(
"unexpected database value: decompressed usage json exceeds 67108864 bytes",
"too_large",
),
(
"failed to decompress usage json: invalid gzip header",
"decode_failed",
),
(
"failed to parse decompressed usage json: invalid JSON",
"decode_failed",
),
(
"postgres error: private connection details",
"storage_unavailable",
),
] {
assert_eq!(
admin_usage_body_load_error_code(&GatewayError::Internal(message.to_string())),
expected
);
}
}
}
@@ -6,6 +6,7 @@ mod extractors;
mod list;
pub(crate) mod payloads;
mod reads;
mod reveal;
mod support;
mod update;
@@ -41,6 +42,10 @@ pub(crate) async fn maybe_build_local_admin_endpoints_routes_response(
return Ok(Some(response));
}
if let Some(response) = reveal::maybe_handle(state, request_context).await? {
return Ok(Some(response));
}
if let Some(response) = defaults::maybe_handle(state, request_context, request_body).await? {
return Ok(Some(response));
}
@@ -0,0 +1,72 @@
use super::extractors::admin_endpoint_id;
use super::support::build_admin_endpoints_data_unavailable_response;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{
attach_admin_audit_response, mark_sensitive_admin_response_no_store,
};
use crate::GatewayError;
use axum::{
body::Body,
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.decision() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("endpoints_manage")
|| decision.route_kind.as_deref() != Some("reveal_endpoint_rules")
{
return Ok(None);
}
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
let Some(endpoint_id) = request_context
.path()
.strip_suffix("/rules/reveal")
.and_then(admin_endpoint_id)
else {
return Ok(Some(
(
StatusCode::NOT_FOUND,
Json(json!({ "detail": "Endpoint 不存在" })),
)
.into_response(),
));
};
let Some(endpoint) = state
.read_provider_catalog_endpoints_by_ids(std::slice::from_ref(&endpoint_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
StatusCode::NOT_FOUND,
Json(json!({ "detail": "Endpoint 不存在" })),
)
.into_response(),
));
};
let payload = json!({
"header_rules": endpoint.header_rules.as_ref().and_then(|value| value.as_array()).cloned().unwrap_or_default(),
"body_rules": endpoint.body_rules.as_ref().and_then(|value| value.as_array()).cloned().unwrap_or_default(),
"response_header_rules": endpoint.config.as_ref().and_then(|config| config.get("response_header_rules")).and_then(|value| value.as_array()).cloned().unwrap_or_default(),
});
Ok(Some(mark_sensitive_admin_response_no_store(
attach_admin_audit_response(
Json(payload).into_response(),
"admin_endpoint_rules_revealed",
"reveal_endpoint_rules",
"provider_endpoint",
&endpoint_id,
),
)))
}
@@ -1463,7 +1463,7 @@ mod tests {
&auth_config,
Some(0),
),
"antigravity_[email protected]"
"[email protected]"
);
}
@@ -1,3 +1,4 @@
use super::super::helpers::admin_provider_oauth_key_name_from_auth_config;
use super::super::kiro::{
admin_provider_oauth_kiro_refresh_base_url_override, fetch_admin_provider_oauth_kiro_email,
refresh_admin_provider_oauth_kiro_auth_config,
@@ -79,7 +80,7 @@ fn kiro_social_key_name(
.collect::<String>()
})
.unwrap_or_else(|| "unknown".to_string());
format!("kiro_{fallback} ({provider})")
format!("账号_{fallback} ({provider})")
}
fn kiro_social_poll_error_response(error: impl Into<String>) -> Response<Body> {
@@ -1004,10 +1005,11 @@ async fn handle_admin_provider_oauth_windsurf_browser_device_poll(
}
}
} else {
let key_name = email
.as_deref()
.map(|email| format!("windsurf_{email}"))
.unwrap_or_else(|| format!("windsurf_{}", current_unix_secs()));
let key_name = admin_provider_oauth_key_name_from_auth_config(
&provider.provider_type,
&auth_config,
None,
);
match state
.create_provider_oauth_catalog_key(
&provider.id,
@@ -1356,6 +1358,32 @@ mod tests {
use crate::control::GatewayAdminPrincipalContext;
use aether_data::repository::provider_oauth::StoredAdminProviderOAuthDeviceSession;
#[test]
fn kiro_social_key_name_preserves_email_and_auth_method() {
assert_eq!(
super::kiro_social_key_name(
Some(" [email protected] "),
Some("Github"),
Some("refresh-token-1"),
),
"[email protected] (Github)"
);
}
#[test]
fn kiro_social_key_name_without_email_uses_generic_account_prefix() {
for email in [None, Some(""), Some(" ")] {
assert_eq!(
super::kiro_social_key_name(email, Some("Google"), Some("refresh-token-1")),
"账号_154f43 (Google)"
);
assert_eq!(
super::kiro_social_key_name(email, None, None),
"账号_unknown (social)"
);
}
}
fn device_session() -> StoredAdminProviderOAuthDeviceSession {
StoredAdminProviderOAuthDeviceSession {
session_id: "device-session-1".to_string(),
@@ -52,13 +52,12 @@ pub(super) fn admin_provider_oauth_key_name_from_auth_config(
auth_config: &Map<String, Value>,
batch_index: Option<usize>,
) -> String {
let provider_type = provider_type.trim();
if let Some(email) = trimmed_auth_config_string(auth_config, "email") {
return format!("{provider_type}_{email}");
return email;
}
if provider_type.eq_ignore_ascii_case("grok") {
if provider_type.trim().eq_ignore_ascii_case("grok") {
if let Some(user_id) = trimmed_auth_config_string(auth_config, "user_id") {
return format!("grok_{user_id}");
return user_id;
}
}
@@ -68,7 +67,7 @@ pub(super) fn admin_provider_oauth_key_name_from_auth_config(
.map(|duration| duration.as_secs())
.unwrap_or(0);
match batch_index {
Some(index) => format!("{provider_type}_{timestamp}_{index}"),
Some(index) => format!("账号_{timestamp}_{index}"),
None => format!("账号_{timestamp}"),
}
}
@@ -87,6 +86,106 @@ mod tests {
use super::*;
use serde_json::{json, Map};
const PROVIDER_TYPES: &[&str] = &[
"codex",
" Codex ",
"claude_code",
"chatgpt_web",
"gemini_cli",
"antigravity",
"grok",
" Grok ",
"kiro",
"windsurf",
];
#[test]
fn default_key_name_uses_email_without_provider_prefix() {
let mut auth_config = Map::new();
auth_config.insert("email".to_string(), json!(" [email protected] "));
for provider_type in PROVIDER_TYPES {
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
"[email protected]"
);
}
}
}
#[test]
fn antigravity_default_key_name_uses_email_without_provider_prefix() {
for email in [" [email protected] ", "[email protected]"] {
let mut auth_config = Map::new();
auth_config.insert("email".to_string(), json!(email));
for provider_type in ["antigravity", " Antigravity "] {
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
email.trim()
);
}
}
}
}
#[test]
fn default_key_name_preserves_email_with_provider_prefix() {
for provider_type in PROVIDER_TYPES {
let email = format!("{}[email protected]", provider_type.trim());
let mut auth_config = Map::new();
auth_config.insert("email".to_string(), json!(email));
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
email
);
}
}
}
#[test]
fn default_key_name_without_email_uses_generic_account_name() {
for email in [None, Some(""), Some(" ")] {
let mut auth_config = Map::new();
if let Some(email) = email {
auth_config.insert("email".to_string(), json!(email));
}
for provider_type in PROVIDER_TYPES {
for batch_index in [None, Some(3)] {
let name = admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
);
let suffix = name.strip_prefix("账号_").expect("generic account prefix");
let timestamp = if batch_index.is_some() {
suffix.strip_suffix("_3").expect("batch index suffix")
} else {
suffix
};
assert!(timestamp.parse::<u64>().is_ok());
}
}
}
}
#[test]
fn grok_default_key_name_uses_full_user_id() {
let mut auth_config = Map::new();
@@ -95,10 +194,18 @@ mod tests {
json!("1619039a-0191-4e0a-a490-8f4ad21262c9"),
);
assert_eq!(
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
"grok_1619039a-0191-4e0a-a490-8f4ad21262c9"
);
for provider_type in ["grok", " Grok "] {
for batch_index in [None, Some(3)] {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
batch_index,
),
"1619039a-0191-4e0a-a490-8f4ad21262c9"
);
}
}
}
#[test]
@@ -109,17 +216,22 @@ mod tests {
assert_eq!(
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
"grok_grok@example.com"
"[email protected]"
);
}
#[test]
fn batch_default_key_name_keeps_existing_timestamp_shape() {
fn batch_default_key_name_keeps_distinct_indexes_without_provider_prefix() {
let auth_config = Map::new();
let name = admin_provider_oauth_key_name_from_auth_config("codex", &auth_config, Some(3));
let name = admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, Some(3));
let other_name =
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, Some(4));
assert!(name.starts_with("codex_"));
assert!(name.starts_with("账号_"));
assert!(name.ends_with("_3"));
assert!(other_name.starts_with("账号_"));
assert!(other_name.ends_with("_4"));
assert_ne!(name, other_name);
}
#[test]
@@ -66,43 +66,6 @@ fn admin_provider_oauth_kiro_refresh_error(
}
}
#[cfg(test)]
mod refresh_error_tests {
use super::admin_provider_oauth_kiro_refresh_error;
use crate::handlers::admin::request::AdminKiroAuthConfig;
use aether_oauth::core::OAuthError;
#[test]
fn kiro_refresh_error_does_not_reflect_upstream_body() {
let auth_config = AdminKiroAuthConfig {
auth_method: None,
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let detail = admin_provider_oauth_kiro_refresh_error(
&auth_config,
OAuthError::HttpStatus {
status_code: 502,
body_excerpt: "authorization=Bearer upstream-secret".to_string(),
},
);
assert_eq!(detail, "social refresh 失败: HTTP 502");
assert!(!detail.contains("upstream-secret"));
}
}
pub(super) async fn refresh_admin_provider_oauth_kiro_auth_config(
state: &AdminAppState<'_>,
auth_config: &AdminKiroAuthConfig,
@@ -240,3 +203,40 @@ pub(super) async fn fetch_admin_provider_oauth_kiro_email(
aether_admin::provider::quota::parse_kiro_usage_response(&payload, current_unix_secs())?;
json_non_empty_string(metadata.get("email"))
}
#[cfg(test)]
mod refresh_error_tests {
use super::admin_provider_oauth_kiro_refresh_error;
use crate::handlers::admin::request::AdminKiroAuthConfig;
use aether_oauth::core::OAuthError;
#[test]
fn kiro_refresh_error_does_not_reflect_upstream_body() {
let auth_config = AdminKiroAuthConfig {
auth_method: None,
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let detail = admin_provider_oauth_kiro_refresh_error(
&auth_config,
OAuthError::HttpStatus {
status_code: 502,
body_excerpt: "authorization=Bearer upstream-secret".to_string(),
},
);
assert_eq!(detail, "social refresh 失败: HTTP 502");
assert!(!detail.contains("upstream-secret"));
}
}
@@ -5,7 +5,6 @@ use super::shared::{
quota_key_auto_removed, quota_refresh_success_invalid_state,
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::payloads::AdminImportProviderModelsRequest;
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::{
@@ -24,63 +23,6 @@ use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
fn antigravity_discovered_model_ids(metadata_update: Option<&serde_json::Value>) -> Vec<String> {
metadata_update
.and_then(|value| value.pointer("/antigravity/quota_by_model"))
.and_then(serde_json::Value::as_object)
.into_iter()
.flat_map(|models| models.keys())
.map(String::as_str)
.filter(|model_id| aether_model_fetch::antigravity_model_id_is_routable(model_id))
.map(ToOwned::to_owned)
.collect()
}
async fn sync_antigravity_discovered_models(
state: &AdminAppState<'_>,
provider_id: &str,
metadata_update: Option<&serde_json::Value>,
) {
if !state.has_global_model_data_reader() || !state.has_global_model_data_writer() {
return;
}
let model_ids = antigravity_discovered_model_ids(metadata_update);
if model_ids.is_empty() {
return;
}
let result = state
.build_admin_import_provider_models_payload(
provider_id,
AdminImportProviderModelsRequest {
model_ids,
tiered_pricing: None,
price_per_request: None,
},
)
.await;
match result {
Ok(payload) => {
let errors = payload
.get("errors")
.and_then(serde_json::Value::as_array)
.map(Vec::len)
.unwrap_or(0);
if errors > 0 {
warn!(
provider_id,
errors, "Antigravity discovered-model catalog sync completed with item errors"
);
}
}
Err(error) => warn!(
provider_id,
error = %error,
"Antigravity discovered-model catalog sync failed"
),
}
}
async fn execute_antigravity_quota_plan(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
@@ -380,10 +322,6 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
continue;
}
if status == "success" {
sync_antigravity_discovered_models(state, &provider.id, metadata_update.as_ref()).await;
}
if status == "success" {
success_count += 1;
} else {
@@ -234,6 +234,20 @@ impl<'a> AdminAppState<'a> {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn read_request_usage_body_payload(
&self,
body_ref: &str,
) -> Result<
Option<aether_data_contracts::repository::usage::StoredUsageBodyPayload>,
GatewayError,
> {
self.app
.data
.read_request_usage_body_payload(body_ref)
.await
.map_err(|error| GatewayError::Internal(error.to_string()))
}
pub(crate) async fn build_api_format_health_monitor_payload(
&self,
lookback_hours: u64,
@@ -211,9 +211,19 @@ fn apply_sensitive_route_cache_policy(
return;
}
let preserve_no_transform = headers
.get_all(http::header::CACHE_CONTROL)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.any(|directive| directive.trim().eq_ignore_ascii_case("no-transform"));
headers.insert(
http::header::CACHE_CONTROL,
HeaderValue::from_static("no-store"),
HeaderValue::from_static(if preserve_no_transform {
"no-store, no-transform"
} else {
"no-store"
}),
);
headers.insert(http::header::PRAGMA, HeaderValue::from_static("no-cache"));
}
@@ -354,6 +364,29 @@ mod tests {
);
}
#[test]
fn raw_body_no_transform_survives_sensitive_cache_policy() {
let mut headers = HeaderMap::new();
headers.append(
http::header::CACHE_CONTROL,
HeaderValue::from_static("public, max-age=3600"),
);
headers.append(
http::header::CACHE_CONTROL,
HeaderValue::from_static(" No-Transform "),
);
apply_sensitive_route_cache_policy(
&mut headers,
"/api/admin/usage/usage-1?body_format=raw",
None,
);
assert_eq!(
headers[http::header::CACHE_CONTROL],
"no-store, no-transform"
);
assert_eq!(headers[http::header::PRAGMA], "no-cache");
}
#[test]
fn authenticated_user_data_responses_are_never_cacheable() {
let mut headers = HeaderMap::new();
@@ -1035,7 +1035,7 @@ pub(crate) async fn proxy_request(
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
request: Request,
) -> Result<Response<Body>, GatewayError> {
crate::request_diagnostics::scope_request_diagnostics(Box::pin(proxy_request_inner(
crate::request_lifecycle::run_request(Box::pin(proxy_request_inner(
state,
remote_addr,
request,
@@ -3228,7 +3228,7 @@ mod tests {
async fn request_body_buffer_caps_decompressed_body_at_shared_budget() {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder
.write_all(&vec![b'a'; 128])
.write_all(&[b'a'; 128])
.expect("test gzip body should encode");
let encoded = encoder.finish().expect("test gzip body should finish");
assert!(
@@ -68,7 +68,12 @@ pub(super) async fn relay_bound_connection(
state: &AppState,
context: &WebSocketRequestContext,
) {
let mut client_connected = true;
loop {
if !client_connected && !bound.turn_state.response_in_flight() {
close_bound_upstream(bound).await;
break;
}
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
tokio::select! {
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
@@ -100,8 +105,12 @@ pub(super) async fn relay_bound_connection(
).await;
break;
}
client_message = client_socket.next() => {
client_message = client_socket.next(), if client_connected => {
let Some(client_message) = client_message else {
if retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
finalize_active_turn(
bound,
state,
@@ -111,6 +120,10 @@ pub(super) async fn relay_bound_connection(
break;
};
let Ok(client_message) = client_message else {
if retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
warn!(
event_name = "responses_websocket_client_receive_failed",
log_type = "ops",
@@ -127,6 +140,12 @@ pub(super) async fn relay_bound_connection(
close_bound_upstream(bound).await;
break;
};
if matches!(client_message, AxumWsMessage::Close(_))
&& retain_disconnected_turn(bound)
{
client_connected = false;
continue;
}
match Box::pin(forward_client_message(
client_message,
bound,
@@ -559,6 +578,7 @@ pub(super) async fn relay_bound_connection(
let mut relay_send_error = None;
let mut relay_serialization_failed = false;
match relay_directive {
_ if !client_connected => {}
Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => {
let client_frame = match parsed_upstream_frame.as_ref().map(|frame| {
bound
@@ -673,6 +693,10 @@ pub(super) async fn relay_bound_connection(
break;
}
if let Some(error) = relay_send_error {
if terminal_outcome.is_none() && retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
warn!(
event_name = "responses_websocket_client_send_failed",
log_type = "ops",
@@ -737,6 +761,20 @@ pub(super) async fn relay_bound_connection(
}
}
fn retain_disconnected_turn(bound: &mut BoundResponsesConnection) -> bool {
if bound
.turn_state
.attempt()
.is_none_or(|attempt| attempt.cancel_on_client_disconnect())
{
return false;
}
bound
.turn_state
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
true
}
struct PendingContinuationRegistration {
user_id: String,
api_key_id: String,
@@ -845,6 +845,13 @@ fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> Gatew
}
impl ResponsesProviderAttempt {
pub(super) fn cancel_on_client_disconnect(&self) -> bool {
crate::orchestration::routing_execution_policy_from_report_context(
self.lifecycle.report_context(),
)
.is_some_and(|policy| policy.cancel_on_client_disconnect)
}
/// Releases all per-turn capacity before terminal persistence starts.
/// Provider-pool runtime tokens normally use an awaited removal. The
/// bounded wait prevents a broken runtime backend from stalling the relay;
@@ -24,7 +24,7 @@ use wreq::ws::message::{CloseFrame as WreqCloseFrame, Message as WreqWsMessage};
use crate::ai_serving::AiExecutionDecision;
use crate::execution_runtime::transport::{
build_browser_wreq_client, build_request_headers, normalize_execution_proxy_url,
ExecutionTransportControls,
validate_execution_upstream_url, ExecutionSafeDnsResolver, ExecutionTransportControls,
};
use crate::frontdoor_loop_guard::gateway_frontdoor_self_loop_guard_error;
use crate::handlers::proxy::websocket::session::{
@@ -66,7 +66,7 @@ pub(crate) async fn connect_upstream_websocket(
)?;
let headers =
websocket_handshake_headers(&decision.provider_request_headers, errors.headers_invalid)?;
let client = build_websocket_client(decision, &upstream_url, errors).await?;
let client = build_websocket_client(decision, errors)?;
let response = client
.websocket(upstream_url.as_str())
.headers(headers)
@@ -149,25 +149,14 @@ pub(crate) fn websocket_upstream_url(
invalid_code: &'static str,
) -> Result<Url, &'static str> {
let mut url = Url::parse(raw).map_err(|_| invalid_code)?;
if url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
|| url.fragment().is_some()
{
return Err(invalid_code);
}
let websocket_scheme = match url.scheme() {
"https" => "wss",
"http" => "ws",
"wss" => return Ok(url),
"ws" if aether_http::url_has_literal_loopback_host(&url) => return Ok(url),
"ws" => return Err(invalid_code),
let (http_scheme, websocket_scheme) = match url.scheme() {
"https" | "wss" => ("https", "wss"),
"http" | "ws" => ("http", "ws"),
_ => return Err(invalid_code),
};
url.set_scheme(http_scheme).map_err(|_| invalid_code)?;
let mut url = validate_execution_upstream_url(url.as_str()).map_err(|_| invalid_code)?;
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
if url.scheme() == "ws" && !aether_http::url_has_literal_loopback_host(&url) {
return Err(invalid_code);
}
Ok(url)
}
@@ -223,9 +212,8 @@ pub(crate) fn websocket_handshake_headers(
Ok(headers)
}
async fn build_websocket_client(
fn build_websocket_client(
decision: &AiExecutionDecision,
upstream_url: &Url,
errors: UpstreamWebSocketErrorCodes,
) -> Result<wreq::Client, &'static str> {
let timeouts = websocket_timeouts(decision);
@@ -249,41 +237,7 @@ async fn build_websocket_client(
let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?;
builder = builder.proxy(proxy);
} else {
// Pin every direct WebSocket connection to the DNS answers validated
// here. This also covers the explicitly permitted loopback `ws://`
// form; otherwise the client would perform a second lookup and a
// rebinding could escape the loopback-only policy.
let host = upstream_url.host_str().ok_or(errors.upstream_url_invalid)?;
let port = upstream_url
.port_or_known_default()
.ok_or(errors.upstream_url_invalid)?;
let addresses = if let Ok(ip) = host.parse::<std::net::IpAddr>() {
vec![std::net::SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host,
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| errors.upstream_url_invalid)?
};
let allows_loopback = host.trim_end_matches('.').eq_ignore_ascii_case("localhost")
|| host
.parse::<std::net::IpAddr>()
.map(|ip| ip.is_loopback())
.unwrap_or(false);
let unsafe_answer = if allows_loopback {
addresses.iter().any(|address| !address.ip().is_loopback())
} else {
addresses
.iter()
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()))
};
if addresses.is_empty() || unsafe_answer {
return Err(errors.upstream_url_invalid);
}
builder = builder.resolve_to_addrs(host.to_string(), addresses.iter().copied());
builder = builder.dns_resolver(ExecutionSafeDnsResolver);
}
builder.build().map_err(|_| errors.client_build_failed)
}
@@ -678,15 +632,17 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16
#[cfg(test)]
mod tests {
use super::{
bounded_send, guarded_websocket_upstream_url, resolve_websocket_proxy_url,
responses_websocket_error_event, responses_websocket_error_event_with_stream_id,
websocket_handshake_headers, websocket_relay_frame_queue, websocket_response_headers,
websocket_upstream_url, UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl,
WebSocketRelayQueueError, WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY,
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
bounded_send, build_websocket_client, guarded_websocket_upstream_url,
resolve_websocket_proxy_url, responses_websocket_error_event,
responses_websocket_error_event_with_stream_id, websocket_handshake_headers,
websocket_relay_frame_queue, websocket_response_headers, websocket_upstream_url,
UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl, WebSocketRelayQueueError,
WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY, RELAY_WRITE_TIMEOUT,
TEARDOWN_WRITE_TIMEOUT,
};
use crate::ai_serving::AiExecutionDecision;
use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url;
use aether_contracts::ProxySnapshot;
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
use axum::http::HeaderMap;
use std::collections::BTreeMap;
use std::time::Duration;
@@ -844,15 +800,17 @@ mod tests {
#[test]
fn maps_http_url_to_websocket_url_without_losing_path_or_query() {
let url = websocket_upstream_url(
"https://example.test/backend-api/codex/responses?x=1",
"invalid",
)
.expect("URL should be converted");
assert_eq!(
url.as_str(),
"wss://example.test/backend-api/codex/responses?x=1"
);
for (http_scheme, websocket_scheme) in [("https", "wss"), ("http", "ws")] {
let url = websocket_upstream_url(
&format!("{http_scheme}://example.test:8080/backend-api/codex/responses?x=1"),
"invalid",
)
.expect("URL should be converted");
assert_eq!(
url.as_str(),
format!("{websocket_scheme}://example.test:8080/backend-api/codex/responses?x=1")
);
}
}
#[test]
@@ -861,10 +819,16 @@ mod tests {
}
#[test]
fn remote_websocket_requires_wss_but_loopback_ws_is_allowed() {
fn websocket_upstream_url_accepts_ws_and_wss_with_safe_targets() {
for allowed in [
"wss://example.test/v1/responses",
"https://example.test/v1/responses",
"ws://example.test:8080/v1/responses",
"http://example.test:8080/v1/responses",
"http://8.8.8.8:8080/v1/responses",
"wss://8.8.8.8/v1/responses",
"ws://[2606:4700:4700::1111]:8080/v1/responses",
"wss://[2606:4700:4700::1111]/v1/responses",
"ws://localhost:8080/v1/responses",
"http://127.42.0.1:8080/v1/responses",
"ws://[::1]:8080/v1/responses",
@@ -875,11 +839,22 @@ mod tests {
);
}
for rejected in [
"ws://example.test/v1/responses",
"http://10.0.0.1/v1/responses",
"wss://10.0.0.1/v1/responses",
"wss://127.0.0.1/v1/responses",
"wss://[::1]/v1/responses",
"wss://[fd00::1]/v1/responses",
"wss://[::ffff:127.0.0.1]/v1/responses",
"wss://169.254.169.254/v1/responses",
"wss://198.18.78.41/v1/responses",
"wss://198.19.1.2/v1/responses",
"ws://0.0.0.0:8080/v1/responses",
"ws://[::ffff:127.0.0.1]:8080/v1/responses",
"wss://example.test/v1/responses#secret",
"ws://example.test/v1/responses#secret",
"http://[email protected]/v1/responses",
"ws://[email protected]/v1/responses",
"ftp://example.test/v1/responses",
] {
assert!(
websocket_upstream_url(rejected, "invalid").is_err(),
@@ -888,6 +863,60 @@ mod tests {
}
}
#[tokio::test]
async fn websocket_client_build_defers_provider_dns_for_all_transport_profiles() {
let errors = UpstreamWebSocketErrorCodes {
upstream_url_missing: "missing",
upstream_url_invalid: "upstream_invalid",
frontdoor_self_loop: "frontdoor_self_loop",
headers_invalid: "headers_invalid",
client_build_failed: "client_build_failed",
proxy_invalid: "proxy_invalid",
tunnel_proxy_unsupported: "tunnel_unsupported",
handshake_failed: "handshake_failed",
upgrade_rejected: "upgrade_rejected",
upgrade_failed: "upgrade_failed",
};
for profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for proxy in [
None,
Some(ProxySnapshot {
enabled: Some(false),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("socks5h://proxy.invalid:1080".to_string()),
..Default::default()
}),
] {
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "wss://upstream.invalid/v1/responses"
}))
.expect("minimal provider decision should deserialize");
decision.transport_profile = profile.clone();
decision.proxy = proxy;
build_websocket_client(&decision, errors)
.expect("building a client must not resolve the provider or proxy hostname");
}
}
}
#[test]
fn active_websocket_proxy_without_a_target_fails_closed() {
let errors = UpstreamWebSocketErrorCodes {
@@ -922,6 +951,94 @@ mod tests {
);
}
#[tokio::test]
async fn websocket_handshake_keeps_provider_dns_remote_for_http_and_socks_proxies() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let errors = UpstreamWebSocketErrorCodes {
upstream_url_missing: "missing",
upstream_url_invalid: "upstream_invalid",
frontdoor_self_loop: "frontdoor_self_loop",
headers_invalid: "headers_invalid",
client_build_failed: "client_build_failed",
proxy_invalid: "proxy_invalid",
tunnel_proxy_unsupported: "tunnel_unsupported",
handshake_failed: "handshake_failed",
upgrade_rejected: "upgrade_rejected",
upgrade_failed: "upgrade_failed",
};
for profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for scheme in ["http", "socks5", "socks5h"] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let proxy_addr = listener.local_addr().unwrap();
let (release, released) = tokio::sync::oneshot::channel::<()>();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
if scheme != "http" {
let mut greeting = [0; 2];
stream.read_exact(&mut greeting).await.unwrap();
assert_eq!(greeting[0], 5);
let mut methods = vec![0; greeting[1] as usize];
stream.read_exact(&mut methods).await.unwrap();
assert!(methods.contains(&0));
stream.write_all(&[5, 0]).await.unwrap();
let mut request = [0; 4];
stream.read_exact(&mut request).await.unwrap();
assert_eq!(
request,
[5, 1, 0, 3],
"proxy must receive a domain, not an IP"
);
let host_len = stream.read_u8().await.unwrap();
let mut host = vec![0; host_len as usize];
stream.read_exact(&mut host).await.unwrap();
assert_eq!(host, b"provider-dns.invalid");
assert_eq!(stream.read_u16().await.unwrap(), 80);
stream
.write_all(&[5, 0, 0, 1, 127, 0, 0, 1, 0, 80])
.await
.unwrap();
}
let socket = tokio_tungstenite::accept_async(stream).await.unwrap();
let _ = released.await;
drop(socket);
});
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "ws://provider-dns.invalid/v1/responses",
"proxy": {"enabled": true, "url": format!("{scheme}://{proxy_addr}")}
}))
.unwrap();
decision.transport_profile = profile.clone();
let connection = tokio::time::timeout(
Duration::from_secs(5),
super::connect_upstream_websocket(
&decision,
crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS,
errors,
),
)
.await
.expect("proxied handshake must not wait for local provider DNS")
.unwrap_or_else(|error| panic!("{scheme} handshake failed: {error}"));
release.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(5), server)
.await
.unwrap()
.unwrap();
drop(connection);
}
}
}
#[test]
fn rejects_responses_websocket_frontdoor_self_loop_before_connecting() {
let base_url = configured_gateway_frontdoor_base_url();
@@ -129,9 +129,6 @@ pub(crate) fn normalize_admin_base_url(base_url: &str) -> Result<String, String>
if parsed.host_str().is_none() {
return Err("base_url 必须包含有效主机".to_string());
}
if !aether_http::is_https_or_loopback_http_url(&parsed) {
return Err("base_url 必须使用 HTTPS;HTTP 仅允许字面量 loopback 主机".to_string());
}
if !parsed.username().is_empty() || parsed.password().is_some() {
return Err("base_url 不允许包含用户名或密码".to_string());
}
@@ -154,16 +151,42 @@ mod normalize_admin_base_url_tests {
"https://user:[email protected]/v1",
"https://api.example.test/v1?key=secret",
"https://api.example.test/v1#secret",
"http://api.example.test/v1",
"http://10.0.0.1/v1",
"http://[::ffff:127.0.0.1]/v1",
"http://user:password@api.example.test/v1",
"http://api.example.test/v1?key=secret",
"http://api.example.test/v1#secret",
"ftp://api.example.test/v1",
"file:///v1",
"api.example.test/v1",
"",
"https://",
"http://",
"https://api.example.test:invalid/v1",
] {
assert!(normalize_admin_base_url(value).is_err(), "accepted {value}");
}
}
#[test]
fn endpoint_base_url_accepts_remote_http_hosts() {
for (raw_url, expected) in [
(
" HTTP://API.EXAMPLE.TEST:8080/v1/ ",
"http://api.example.test:8080/v1",
),
("http://8.8.8.8:8080/v1/", "http://8.8.8.8:8080/v1"),
("http://10.0.0.1:8080/v1/", "http://10.0.0.1:8080/v1"),
(
"http://[2606:4700:4700::1111]:8080/v1/",
"http://[2606:4700:4700::1111]:8080/v1",
),
] {
assert_eq!(
normalize_admin_base_url(raw_url).expect("HTTP base URL should be accepted"),
expected,
);
}
}
#[test]
fn endpoint_base_url_is_parsed_and_normalized() {
assert_eq!(
@@ -30,6 +30,8 @@ use std::time::{SystemTime, UNIX_EPOCH};
mod support_announcements;
#[path = "support/auth.rs"]
mod support_auth;
#[path = "support/auth_cookie_policy.rs"]
mod support_auth_cookie_policy;
#[path = "support/billing.rs"]
mod support_billing;
#[path = "support/ccswitch.rs"]
@@ -133,6 +135,31 @@ pub(crate) async fn maybe_build_local_public_support_response(
remote_addr: &std::net::SocketAddr,
client_ip: std::net::IpAddr,
request_body: Option<&Bytes>,
) -> Option<Response<Body>> {
let response = build_local_public_support_response(
state,
request_context,
headers,
remote_addr,
client_ip,
request_body,
)
.await?;
Some(support_auth_cookie_policy::finalize_refresh_cookie(
response,
headers,
request_context.host_header.as_deref(),
remote_addr,
))
}
async fn build_local_public_support_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
remote_addr: &std::net::SocketAddr,
client_ip: std::net::IpAddr,
request_body: Option<&Bytes>,
) -> Option<Response<Body>> {
let decision = request_context.control_decision.as_ref()?;
if decision.route_class.as_deref() != Some("public_support") {
@@ -0,0 +1,424 @@
use super::support_auth::{auth_refresh_cookie_name, auth_refresh_cookie_secure};
use axum::body::Body;
use axum::http::{header, HeaderMap, HeaderValue, Response};
use std::net::SocketAddr;
use url::Url;
pub(super) fn finalize_refresh_cookie(
mut response: Response<Body>,
headers: &HeaderMap,
host_header: Option<&str>,
remote_addr: &SocketAddr,
) -> Response<Body> {
if !response.headers().contains_key(header::SET_COOKIE) {
return response;
}
let cookie_name = auth_refresh_cookie_name();
let explicit_secure = std::env::var("AUTH_REFRESH_COOKIE_SECURE").ok();
let public_base_url = std::env::var("AETHER_PUBLIC_BASE_URL")
.ok()
.or_else(|| std::env::var("PUBLIC_BASE_URL").ok());
let secure = refresh_cookie_secure_for_request(
headers,
host_header,
crate::headers::trusted_proxy_ip(remote_addr.ip()),
explicit_secure.as_deref(),
public_base_url.as_deref(),
auth_refresh_cookie_secure(),
);
let cookies = response
.headers()
.get_all(header::SET_COOKIE)
.iter()
.map(|cookie| rewrite_refresh_cookie(cookie, &cookie_name, secure))
.collect::<Vec<_>>();
response.headers_mut().remove(header::SET_COOKIE);
for cookie in cookies {
response.headers_mut().append(header::SET_COOKIE, cookie);
}
response
}
fn refresh_cookie_secure_for_request(
headers: &HeaderMap,
host_header: Option<&str>,
trusted_proxy: bool,
explicit_secure: Option<&str>,
public_base_url: Option<&str>,
fallback_secure: bool,
) -> bool {
if let Some(value) = explicit_secure {
return !value.trim().eq_ignore_ascii_case("false");
}
let origin = single_header(headers, header::ORIGIN.as_str()).and_then(parse_origin);
let public_url = public_base_url.and_then(parse_http_url);
let forwarded_proto = trusted_proxy.then(|| forwarded_proto(headers)).flatten();
if origin.as_ref().is_some_and(|url| url.scheme() == "https")
|| public_url
.as_ref()
.is_some_and(|url| url.scheme() == "https")
|| forwarded_proto == Some("https")
{
return true;
}
if trusted_proxy && headers.contains_key("x-forwarded-proto") {
return forwarded_proto != Some("http");
}
if public_url
.as_ref()
.is_some_and(|url| url.scheme() == "http")
{
return false;
}
if let (Some(origin), Some(host)) = (origin, host_header) {
let request_origin = parse_origin(&format!("{}://{host}", origin.scheme()));
if request_origin.is_some_and(|url| url.origin() == origin.origin()) {
return false;
}
}
fallback_secure
}
fn single_header<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
let mut values = headers.get_all(name).iter();
let value = values.next()?.to_str().ok()?.trim();
(values.next().is_none() && !value.is_empty()).then_some(value)
}
fn parse_origin(value: &str) -> Option<Url> {
let url = parse_http_url(value)?;
(url.path() == "/").then_some(url)
}
fn parse_http_url(value: &str) -> Option<Url> {
let url = Url::parse(value.trim()).ok()?;
(matches!(url.scheme(), "http" | "https")
&& url.host_str().is_some()
&& url.username().is_empty()
&& url.password().is_none()
&& url.query().is_none()
&& url.fragment().is_none())
.then_some(url)
}
fn forwarded_proto(headers: &HeaderMap) -> Option<&'static str> {
let value = headers
.get_all("x-forwarded-proto")
.iter()
.next_back()?
.to_str()
.ok()?
.rsplit(',')
.next()?
.trim();
if value.eq_ignore_ascii_case("https") {
Some("https")
} else if value.eq_ignore_ascii_case("http") {
Some("http")
} else {
None
}
}
fn rewrite_refresh_cookie(cookie: &HeaderValue, cookie_name: &str, secure: bool) -> HeaderValue {
let Ok(value) = cookie.to_str() else {
return cookie.clone();
};
let mut attributes = value.split(';').map(str::trim);
let Some(pair) = attributes.next() else {
return cookie.clone();
};
if pair.split_once('=').map(|(name, _)| name) != Some(cookie_name) {
return cookie.clone();
}
let secure =
secure || cookie_name.starts_with("__Secure-") || cookie_name.starts_with("__Host-");
let mut parts = vec![pair.to_string()];
for attribute in attributes {
if attribute.eq_ignore_ascii_case("Secure") {
continue;
}
if !secure
&& attribute.split_once('=').is_some_and(|(name, value)| {
name.trim().eq_ignore_ascii_case("SameSite")
&& value.trim().eq_ignore_ascii_case("None")
})
{
parts.push("SameSite=Lax".to_string());
} else {
parts.push(attribute.to_string());
}
}
if secure {
parts.push("Secure".to_string());
}
let Ok(mut rewritten) = HeaderValue::from_str(&parts.join("; ")) else {
return cookie.clone();
};
rewritten.set_sensitive(cookie.is_sensitive());
rewritten
}
#[cfg(test)]
mod tests {
use super::{refresh_cookie_secure_for_request, rewrite_refresh_cookie};
use axum::http::{header, HeaderMap, HeaderValue};
fn headers(origin: Option<&str>, forwarded_proto: Option<&str>) -> HeaderMap {
let mut headers = HeaderMap::new();
if let Some(origin) = origin {
headers.insert(header::ORIGIN, HeaderValue::from_str(origin).unwrap());
}
if let Some(proto) = forwarded_proto {
headers.insert("x-forwarded-proto", HeaderValue::from_str(proto).unwrap());
}
headers
}
#[test]
fn refresh_cookie_auto_detects_same_origin_http_and_https() {
for (origin, host, secure) in [
("http://aether.test:8084", "aether.test:8084", false),
("http://aether.test", "aether.test:80", false),
("http://[2001:db8::1]:8084", "[2001:db8::1]:8084", false),
("https://aether.test", "aether.test", true),
] {
assert_eq!(
refresh_cookie_secure_for_request(
&headers(Some(origin), None),
Some(host),
false,
None,
None,
true,
),
secure,
"{origin}",
);
}
assert!(refresh_cookie_secure_for_request(
&headers(Some("https://aether.test"), None),
Some("aether.test"),
false,
None,
None,
false,
));
}
#[test]
fn refresh_cookie_does_not_infer_http_from_other_or_invalid_origins() {
for origin in [
"http://other.test",
"http://aether.test:8085",
"null",
"http://[email protected]:8084",
"http://aether.test:8084/path",
"http://aether.test:8084?query",
"http://aether.test:8084#fragment",
"http://aether.test:8084, https://aether.test:8084",
"file:///tmp/test",
] {
assert!(
refresh_cookie_secure_for_request(
&headers(Some(origin), None),
Some("aether.test:8084"),
false,
None,
None,
true,
),
"{origin}"
);
}
let mut duplicate = headers(Some("http://aether.test:8084"), None);
duplicate.append(
header::ORIGIN,
HeaderValue::from_static("https://aether.test:8084"),
);
assert!(refresh_cookie_secure_for_request(
&duplicate,
Some("aether.test:8084"),
false,
None,
None,
true,
));
}
#[test]
fn refresh_cookie_only_trusts_forwarded_protocol_from_trusted_peers() {
for (proto, trusted, secure) in [
("http", true, false),
("https", true, true),
("http", false, true),
("https", false, true),
("https, http", true, false),
("http, https", true, true),
("ftp", true, true),
("http,", true, true),
] {
assert_eq!(
refresh_cookie_secure_for_request(
&headers(None, Some(proto)),
Some("aether.test"),
trusted,
None,
None,
true,
),
secure,
"{proto}, trusted={trusted}"
);
}
let mut chained = headers(None, Some("http, http"));
chained.append("x-forwarded-proto", HeaderValue::from_static("https"));
assert!(refresh_cookie_secure_for_request(
&chained,
Some("aether.test"),
true,
None,
None,
true,
));
}
#[test]
fn refresh_cookie_https_evidence_prevents_automatic_downgrade() {
for (origin, proto, public_url) in [
("https://aether.test", "http", None),
("http://aether.test", "https", None),
("http://aether.test", "http", Some("https://aether.test")),
] {
assert!(refresh_cookie_secure_for_request(
&headers(Some(origin), Some(proto)),
Some("aether.test"),
true,
None,
public_url,
true,
));
}
}
#[test]
fn refresh_cookie_preserves_explicit_overrides_and_unknown_defaults() {
for (explicit, secure) in [
("true", true),
("FALSE", false),
("invalid", true),
("", true),
] {
assert_eq!(
refresh_cookie_secure_for_request(
&headers(Some("http://aether.test"), None),
Some("aether.test"),
false,
Some(explicit),
None,
true,
),
secure
);
}
assert!(!refresh_cookie_secure_for_request(
&headers(Some("https://aether.test"), None),
Some("aether.test"),
false,
Some("false"),
None,
true,
));
for fallback in [false, true] {
assert_eq!(
refresh_cookie_secure_for_request(
&HeaderMap::new(),
Some("aether.test"),
false,
None,
None,
fallback,
),
fallback
);
}
}
#[test]
fn refresh_cookie_accepts_an_explicit_public_http_origin() {
assert!(!refresh_cookie_secure_for_request(
&HeaderMap::new(),
Some("internal:8084"),
false,
None,
Some("http://aether.test"),
true,
));
assert!(refresh_cookie_secure_for_request(
&headers(Some("https://aether.test"), None),
Some("internal:8084"),
false,
None,
Some("http://aether.test"),
true,
));
}
#[test]
fn refresh_cookie_rewrite_preserves_secret_path_expiry_and_httponly() {
let mut cookie = HeaderValue::from_static(
"aether_refresh_token=secret; Path=/api/auth; HttpOnly; SameSite=None; Max-Age=604800; Secure",
);
cookie.set_sensitive(true);
let rewritten = rewrite_refresh_cookie(&cookie, "aether_refresh_token", false);
assert_eq!(
rewritten.to_str().unwrap(),
"aether_refresh_token=secret; Path=/api/auth; HttpOnly; SameSite=Lax; Max-Age=604800"
);
assert!(rewritten.is_sensitive());
assert_eq!(
rewrite_refresh_cookie(&cookie, "aether_refresh_token", true),
cookie
);
}
#[test]
fn refresh_cookie_rewrite_also_clears_http_cookies() {
let cookie = HeaderValue::from_static(
"aether_refresh_token=; Path=/api/auth; HttpOnly; SameSite=None; Max-Age=0; Secure",
);
assert_eq!(
rewrite_refresh_cookie(&cookie, "aether_refresh_token", false)
.to_str()
.unwrap(),
"aether_refresh_token=; Path=/api/auth; HttpOnly; SameSite=Lax; Max-Age=0"
);
}
#[test]
fn refresh_cookie_rewrite_preserves_other_cookies_and_strict_policy() {
let unrelated = HeaderValue::from_static("oauth_binding=secret; Path=/; Secure; HttpOnly");
assert_eq!(
rewrite_refresh_cookie(&unrelated, "aether_refresh_token", false),
unrelated
);
let strict = HeaderValue::from_static(
"custom_refresh=secret; Path=/api/auth; HttpOnly; SameSite=Strict",
);
assert_eq!(
rewrite_refresh_cookie(&strict, "custom_refresh", false),
strict
);
assert!(rewrite_refresh_cookie(&strict, "custom_refresh", true)
.to_str()
.unwrap()
.ends_with("; Secure"));
let prefixed =
HeaderValue::from_static("__Secure-refresh=secret; HttpOnly; SameSite=None; Secure");
assert_eq!(
rewrite_refresh_cookie(&prefixed, "__Secure-refresh", false),
prefixed
);
}
}
@@ -248,7 +248,7 @@ pub(super) fn auth_verification_send_cooldown_seconds() -> i64 {
.unwrap_or(60)
}
pub(super) fn auth_refresh_cookie_name() -> String {
pub(crate) fn auth_refresh_cookie_name() -> String {
std::env::var("AUTH_REFRESH_COOKIE_NAME")
.ok()
.map(|value| value.trim().to_string())
@@ -1,5 +1,5 @@
use std::collections::BTreeMap;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use axum::{
@@ -10,6 +10,10 @@ use axum::{
};
use serde_json::json;
use crate::execution_runtime::transport::{
validate_execution_upstream_url, ExecutionSafeDnsResolver,
};
use super::test_connection_shared::select_test_connection_provider;
use super::{
provider_catalog_key_supports_format, query_param_value, AppState, GatewayPublicRequestContext,
@@ -18,101 +22,16 @@ use super::{
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
const MAX_TEST_CONNECTION_RESPONSE_BYTES: usize = 256 * 1024;
#[cfg(test)]
fn build_test_connection_client() -> Result<reqwest::Client, reqwest::Error> {
reqwest::Client::builder()
.no_proxy()
.dns_resolver(Arc::new(ExecutionSafeDnsResolver))
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(10))
.http2_adaptive_window(true)
.build()
}
#[derive(Debug)]
struct ResolvedTestConnectionTarget {
url: reqwest::Url,
host: String,
addresses: Vec<SocketAddr>,
}
/// Resolve the provider endpoint once and pin reqwest to that answer. The
/// test-connection route is reachable through the public front door, so it
/// must not perform an unbounded DNS lookup on every connect (which would
/// permit DNS rebinding into private/reserved networks).
async fn resolve_test_connection_target(
raw_url: &str,
allow_private_targets: bool,
) -> Result<ResolvedTestConnectionTarget, &'static str> {
let url = reqwest::Url::parse(raw_url).map_err(|_| "provider endpoint URL is invalid")?;
let literal_loopback = aether_http::url_has_literal_loopback_host(&url);
if !matches!(url.scheme(), "http" | "https")
|| url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
|| url.fragment().is_some()
{
return Err("provider endpoint must be an HTTP(S) URL without credentials or fragment");
}
if url.scheme() == "http" && !(allow_private_targets && literal_loopback) {
return Err("provider endpoint must use HTTPS");
}
let host = url
.host_str()
.ok_or("provider endpoint is missing a host")?
.to_string();
let literal_ip = host.parse::<IpAddr>().ok();
let port = url
.port_or_known_default()
.ok_or("provider endpoint is missing a port")?;
let addresses = if let Some(ip) = literal_ip {
vec![SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host.as_str(),
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| "provider endpoint DNS resolution failed")?
};
if addresses.is_empty() {
return Err("provider endpoint DNS resolution returned no addresses");
}
let has_private_answer = addresses
.iter()
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()));
// `allow_private_targets` is only enabled for in-process test fixtures.
// Keep that escape hatch narrowly scoped to literal loopback URLs whose
// every DNS answer is loopback; otherwise a test-only build (or an
// accidentally reused helper) could turn this public route into a
// private-network HTTP client.
let test_loopback_target = allow_private_targets
&& literal_loopback
&& addresses.iter().all(|address| address.ip().is_loopback());
if has_private_answer && !test_loopback_target {
return Err("provider endpoint resolves to a private or reserved address");
}
Ok(ResolvedTestConnectionTarget {
url,
host,
addresses,
})
}
fn build_pinned_test_connection_client(
target: &ResolvedTestConnectionTarget,
) -> Result<reqwest::Client, reqwest::Error> {
let mut builder = reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(10))
.http2_adaptive_window(true);
if target.host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(&target.host, &target.addresses);
}
builder.build()
}
pub(super) async fn maybe_build_local_test_connection_route_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -387,18 +306,14 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
// Resolve and pin the endpoint before constructing the request. This
// keeps the public health-check route subject to the same DNS/SSRF
// boundary as the main execution transport. Unit-test fixtures may use
// loopback listeners; production requests never opt into private targets.
let target = match resolve_test_connection_target(&upstream_url, cfg!(test)).await {
Ok(target) => target,
let upstream_url = match validate_execution_upstream_url(&upstream_url) {
Ok(url) => url,
Err(reason) => {
tracing::warn!(
event_name = "provider_test_connection_target_rejected",
provider_id = %provider.id,
endpoint_id = %endpoint.id,
reason,
reason = %reason,
"provider connection test target was rejected"
);
return Some(
@@ -410,7 +325,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let test_client = match build_pinned_test_connection_client(&target) {
let test_client = match build_test_connection_client() {
Ok(client) => client,
Err(_) => {
return Some(
@@ -422,7 +337,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let mut upstream_request = test_client.post(target.url);
let mut upstream_request = test_client.post(upstream_url);
for (name, value) in &provider_request_headers {
upstream_request = upstream_request.header(name, value);
}
@@ -498,7 +413,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
#[cfg(test)]
mod tests {
use super::{build_test_connection_client, resolve_test_connection_target};
use super::{build_test_connection_client, validate_execution_upstream_url};
use axum::{
body::Body,
http::{header, Request, StatusCode},
@@ -570,61 +485,68 @@ mod tests {
redirected_server.abort();
}
#[tokio::test]
async fn test_connection_target_rejects_private_addresses_in_production_mode() {
#[test]
fn test_connection_target_rejects_private_literals_like_provider_requests() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat/completions",
"http://10.0.0.1/v1/chat/completions",
"http://169.254.169.254/v1/chat/completions",
"https://10.0.0.1/v1/chat/completions",
"https://127.0.0.1/v1/chat/completions",
"https://[::1]/v1/chat/completions",
"https://localhost/v1/chat/completions",
"http://8.8.8.8/v1/chat/completions",
"https://198.18.78.41/v1/chat/completions",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).is_err(),
"private provider target should be rejected: {raw_url}"
);
}
}
#[tokio::test]
async fn test_connection_target_allows_loopback_only_for_test_fixtures() {
let target = resolve_test_connection_target("http://127.0.0.1:8080/v1/chat", true)
.await
.expect("test fixture target should resolve");
assert_eq!(target.host, "127.0.0.1");
assert_eq!(target.addresses.len(), 1);
assert!(
resolve_test_connection_target("http://8.8.8.8/v1/chat", true)
.await
.is_err(),
"test mode must not make cleartext public endpoints acceptable"
);
assert!(
resolve_test_connection_target("https://10.0.0.1/v1/chat", true)
.await
.is_err(),
"test mode must not make private non-loopback endpoints acceptable"
);
assert!(
resolve_test_connection_target("http://localhost:8080/v1/chat", true)
.await
.is_ok(),
"literal localhost should remain available for local fixtures"
);
#[test]
fn test_connection_target_accepts_public_http_and_https_addresses() {
for (raw_url, expected_port) in [
("http://8.8.8.8/v1/chat", 80),
("http://8.8.8.8:8080/v1/chat", 8080),
("https://8.8.8.8/v1/chat", 443),
("https://[2606:4700:4700::1111]/v1/chat", 443),
] {
let url = validate_execution_upstream_url(raw_url)
.expect("public HTTP(S) provider target should be valid");
assert_eq!(url.as_str(), raw_url);
assert_eq!(url.port_or_known_default(), Some(expected_port));
}
}
#[tokio::test]
async fn test_connection_target_rejects_url_credentials_and_fragments() {
async fn test_connection_target_defers_dns_and_accepts_provider_loopback_urls() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat",
"http://[::1]:8080/v1/chat",
"http://localhost:8080/v1/chat",
"https://provider-dns.invalid/v1/chat",
] {
let url = validate_execution_upstream_url(raw_url)
.expect("target validation must not depend on the current DNS answer");
let request = build_test_connection_client()
.expect("client should build without DNS")
.post(url)
.build()
.expect("provider request should build without DNS");
assert_eq!(request.url().as_str(), raw_url);
}
}
#[test]
fn test_connection_target_rejects_url_credentials_and_fragments() {
for raw_url in [
"https://user:[email protected]/v1/chat",
"https://example.com/v1/chat#fragment",
"http://user:[email protected]/v1/chat",
"http://example.com/v1/chat#fragment",
"ftp://example.com/v1/chat",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).is_err(),
"unsafe provider target should be rejected: {raw_url}"
);
}
@@ -168,52 +168,6 @@ fn wallet_public_refund_payload(mut payload: serde_json::Value) -> serde_json::V
payload
}
#[cfg(test)]
mod tests {
use super::wallet_refund_payload_from_record;
use aether_data::repository::wallet::StoredAdminWalletRefund;
use serde_json::json;
#[test]
fn public_refund_projection_excludes_payout_proof_and_upstream_payload() {
let record = StoredAdminWalletRefund {
id: "refund-1".to_string(),
refund_no: "rf_1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
payment_order_id: Some("order-1".to_string()),
source_type: "payment_order".to_string(),
source_id: Some("order-1".to_string()),
refund_mode: "original_channel".to_string(),
amount_usd: 10.0,
status: "processing".to_string(),
reason: Some("requested".to_string()),
failure_reason: None,
gateway_refund_id: Some("gateway-refund-1".to_string()),
payout_method: None,
payout_reference: None,
payout_proof: Some(json!({
"gateway_refund": {
"id": "gateway-refund-1",
"payload": {"payer": "sensitive", "credential": "secret"}
}
})),
requested_by: Some("user-1".to_string()),
approved_by: Some("admin-1".to_string()),
processed_by: Some("admin-1".to_string()),
created_at_unix_ms: 1,
updated_at_unix_secs: 1,
processed_at_unix_secs: Some(1),
completed_at_unix_secs: None,
};
let payload = wallet_refund_payload_from_record(&record);
assert!(payload.get("payout_proof").is_none());
assert_eq!(payload["status"], "processing");
assert_eq!(payload["gateway_refund_id"], "gateway-refund-1");
}
}
pub(super) async fn handle_wallet_refunds_list(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -659,3 +613,49 @@ pub(super) async fn handle_wallet_create_refund(
}
}
}
#[cfg(test)]
mod tests {
use super::wallet_refund_payload_from_record;
use aether_data::repository::wallet::StoredAdminWalletRefund;
use serde_json::json;
#[test]
fn public_refund_projection_excludes_payout_proof_and_upstream_payload() {
let record = StoredAdminWalletRefund {
id: "refund-1".to_string(),
refund_no: "rf_1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
payment_order_id: Some("order-1".to_string()),
source_type: "payment_order".to_string(),
source_id: Some("order-1".to_string()),
refund_mode: "original_channel".to_string(),
amount_usd: 10.0,
status: "processing".to_string(),
reason: Some("requested".to_string()),
failure_reason: None,
gateway_refund_id: Some("gateway-refund-1".to_string()),
payout_method: None,
payout_reference: None,
payout_proof: Some(json!({
"gateway_refund": {
"id": "gateway-refund-1",
"payload": {"payer": "sensitive", "credential": "secret"}
}
})),
requested_by: Some("user-1".to_string()),
approved_by: Some("admin-1".to_string()),
processed_by: Some("admin-1".to_string()),
created_at_unix_ms: 1,
updated_at_unix_secs: 1,
processed_at_unix_secs: Some(1),
completed_at_unix_secs: None,
};
let payload = wallet_refund_payload_from_record(&record);
assert!(payload.get("payout_proof").is_none());
assert_eq!(payload["status"], "processing");
assert_eq!(payload["gateway_refund_id"], "gateway-refund-1");
}
}
@@ -292,7 +292,7 @@ mod tests {
assert!(controls.is_err());
let template = "{{value}}".repeat(100_000);
let variables = BTreeMap::from([(String::from("value"), String::from("x".repeat(64)))]);
let variables = BTreeMap::from([(String::from("value"), "x".repeat(64))]);
let error = render_admin_email_template_html(&template, &variables)
.expect_err("rendered output must remain bounded");
assert!(format!("{error:?}").contains("exceeds"));
@@ -199,10 +199,7 @@ pub(crate) fn normalize_ldap_transport_server_url(raw: &str, use_starttls: bool)
// Gateway unit/integration fixtures use an in-process mock endpoint. Keep
// this exception behind the gateway test configuration; production code
// always uses the strict parser without custom schemes.
return aether_admin::system::normalize_ldap_transport_server_url_for_tests(
raw,
use_starttls,
);
aether_admin::system::normalize_ldap_transport_server_url_for_tests(raw, use_starttls)
}
#[cfg(not(test))]
{
+1
View File
@@ -71,6 +71,7 @@ mod rate_limit;
mod request_candidate_queue;
mod request_candidate_runtime;
mod request_diagnostics;
mod request_lifecycle;
mod roles;
mod router;
mod routing;
+1 -1
View File
@@ -44,7 +44,7 @@ pub(crate) fn local_auth_jwt_secret() -> Result<String, String> {
Err(std::env::VarError::NotPresent) => {
#[cfg(test)]
{
return Ok(TEST_JWT_SECRET.to_string());
Ok(TEST_JWT_SECRET.to_string())
}
#[cfg(not(test))]
+64 -3
View File
@@ -117,8 +117,8 @@ where
use aether_crypto::warm_python_fernet_secret;
use aether_data::lifecycle::export::{
copy_database_records, export_database_jsonl, import_database_jsonl, DataCopyOptions,
ExportDomain, MAX_JSONL_INPUT_BYTES,
copy_database_records, export_database_jsonl, import_database_jsonl_with_options,
DataCopyOptions, DataImportOptions, ExportDomain, MAX_JSONL_INPUT_BYTES,
};
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use aether_gateway::{
@@ -1351,6 +1351,11 @@ struct DataExportArgs {
struct DataImportArgs {
#[arg(long)]
input: PathBuf,
#[arg(
long,
help = "Preserve passwords and API/management credentials from a trusted import; imported sessions remain revoked. Without this flag identity credentials are revoked."
)]
preserve_credentials: bool,
}
#[derive(ClapArgs, Debug, Clone)]
@@ -1382,6 +1387,11 @@ struct DataCopyArgs {
#[arg(long)]
omit_request_body_details: bool,
#[arg(
long,
help = "Preserve passwords and API/management credentials from the trusted source; imported sessions remain revoked. The target must use the source encryption key."
)]
preserve_credentials: bool,
}
impl GatewayLoggingArgs {
@@ -2905,12 +2915,23 @@ async fn run_data_import(
let driver = database.driver;
let input_path = args.input.clone();
let input = tokio::task::spawn_blocking(move || read_data_import_input(&input_path)).await??;
let imported = import_database_jsonl(database, &input).await?;
if !args.preserve_credentials {
warn!("identity credentials will be revoked; use --preserve-credentials only for trusted recovery or migration");
}
let imported = import_database_jsonl_with_options(
database,
&input,
DataImportOptions {
preserve_credentials: args.preserve_credentials,
},
)
.await?;
info!(
driver = %driver,
input = %args.input.display(),
imported,
preserve_credentials = args.preserve_credentials,
"database import complete"
);
println!(
@@ -3171,6 +3192,9 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
let target_driver = target.driver;
let domains = requested_domains(&args.domains);
let created_at_unix_secs = current_unix_secs()?;
if !args.preserve_credentials {
warn!("identity credentials will be revoked; use --preserve-credentials only for trusted recovery or migration");
}
let imported = copy_database_records(
source,
target,
@@ -3178,6 +3202,7 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
created_at_unix_secs,
DataCopyOptions {
omit_request_body_details: args.omit_request_body_details,
preserve_credentials: args.preserve_credentials,
},
)
.await?;
@@ -3186,6 +3211,7 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
source_driver = %source_driver,
target_driver = %target_driver,
imported,
preserve_credentials = args.preserve_credentials,
"database copy complete"
);
println!(
@@ -4357,6 +4383,41 @@ mod tests {
};
assert!(copy.source_allow_insecure);
assert!(!copy.target_allow_insecure);
assert!(!copy.preserve_credentials);
}
#[test]
fn data_import_and_copy_require_explicit_credential_preservation() {
for preserve in [false, true] {
let mut import_args = vec!["aether-gateway", "import", "--input", "trusted.jsonl"];
let mut copy_args = vec![
"aether-gateway",
"copy",
"--source-driver",
"postgres",
"--source-url",
"postgres://localhost/source",
"--target-driver",
"postgres",
"--target-url",
"postgres://localhost/target",
];
if preserve {
import_args.push("--preserve-credentials");
copy_args.push("--preserve-credentials");
}
let Some(DataCommand::Import(import)) =
Args::try_parse_from(import_args).unwrap().command
else {
panic!("expected import command");
};
let Some(DataCommand::Copy(copy)) = Args::try_parse_from(copy_args).unwrap().command
else {
panic!("expected copy command");
};
assert_eq!(import.preserve_credentials, preserve);
assert_eq!(copy.preserve_credentials, preserve);
}
}
#[cfg(unix)]
@@ -115,9 +115,6 @@ pub(super) fn usage_cleanup_window(
usage_cleanup_window_with_override(now_utc, settings, None)
}
/// Clamp is non-aggressive: each tier's cutoff becomes `max(policy_cutoff, now - override)`.
/// A later cutoff = fewer records deleted, so the override can only make cleanup more
/// conservative than the configured retention, never more destructive.
pub(super) fn usage_cleanup_window_with_override(
now_utc: DateTime<Utc>,
settings: UsageCleanupSettings,
@@ -135,9 +132,9 @@ pub(super) fn usage_cleanup_window_with_override(
};
let manual_cutoff = now_utc - override_duration;
UsageCleanupWindow {
detail_cutoff: policy.detail_cutoff.max(manual_cutoff),
compressed_cutoff: policy.compressed_cutoff.max(manual_cutoff),
header_cutoff: policy.header_cutoff.max(manual_cutoff),
log_cutoff: policy.log_cutoff.max(manual_cutoff),
detail_cutoff: policy.detail_cutoff.min(manual_cutoff),
compressed_cutoff: policy.compressed_cutoff.min(manual_cutoff),
header_cutoff: policy.header_cutoff.min(manual_cutoff),
log_cutoff: policy.log_cutoff.min(manual_cutoff),
}
}
@@ -1140,19 +1140,40 @@ fn usage_cleanup_window_with_override_is_always_non_aggressive() {
let override_duration = chrono::Duration::days(180);
let clamped = usage_cleanup_window_with_override(now_utc, settings, Some(override_duration));
assert_eq!(clamped.detail_cutoff, policy.detail_cutoff);
assert_eq!(clamped.compressed_cutoff, policy.compressed_cutoff);
assert_eq!(clamped.header_cutoff, policy.header_cutoff);
assert_eq!(clamped.log_cutoff, now_utc - override_duration);
assert!(clamped.log_cutoff > policy.log_cutoff);
assert_eq!(clamped.detail_cutoff, now_utc - override_duration);
assert_eq!(clamped.compressed_cutoff, now_utc - override_duration);
assert_eq!(clamped.header_cutoff, now_utc - override_duration);
assert_eq!(clamped.log_cutoff, policy.log_cutoff);
assert!(clamped.log_cutoff <= policy.log_cutoff);
let far_override = chrono::Duration::days(5);
let far = usage_cleanup_window_with_override(now_utc, settings, Some(far_override));
assert_eq!(far.detail_cutoff, now_utc - far_override);
assert_eq!(far.compressed_cutoff, now_utc - far_override);
assert_eq!(far.header_cutoff, now_utc - far_override);
assert_eq!(far.log_cutoff, now_utc - far_override);
assert!(far.log_cutoff > policy.log_cutoff);
assert_eq!(far, policy);
for days in [0, 5, 30, 180, 400] {
let cutoff = now_utc - chrono::Duration::days(days);
let window = usage_cleanup_window_with_override(
now_utc,
settings,
Some(chrono::Duration::days(days)),
);
for (actual, configured) in [
(window.detail_cutoff, policy.detail_cutoff),
(window.compressed_cutoff, policy.compressed_cutoff),
(window.header_cutoff, policy.header_cutoff),
(window.log_cutoff, policy.log_cutoff),
] {
assert!(actual <= configured);
assert!(actual <= cutoff);
for age in [1, 7, 15, 30, 90, 180, 365, 401] {
let created_at = now_utc - chrono::Duration::days(age);
if created_at < actual {
assert!(created_at < configured);
assert!(created_at < cutoff);
}
}
}
}
let passthrough = usage_cleanup_window_with_override(now_utc, settings, None);
assert_eq!(passthrough, policy);
@@ -31,7 +31,11 @@ fn apply_frontdoor_cors_headers(
);
headers.insert(
http::header::ACCESS_CONTROL_EXPOSE_HEADERS,
HeaderValue::from_static("*"),
HeaderValue::from_static(if headers.contains_key("x-aether-body-field") {
"*, X-Aether-Body-Encoding, X-Aether-Body-Field, X-Aether-Body-Error, X-Aether-Usage-Id"
} else {
"*"
}),
);
if let Some(value) = requested_headers {
if let Ok(value) = HeaderValue::from_str(value) {
@@ -109,3 +113,36 @@ pub(crate) async fn frontdoor_cors_middleware(
);
response
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn raw_body_headers_are_explicitly_exposed_for_credentialed_requests() {
let cors =
FrontdoorCorsConfig::new(vec!["https://console.example".to_string()], true).unwrap();
let mut headers = http::HeaderMap::new();
headers.insert(
"x-aether-body-field",
HeaderValue::from_static("request_body"),
);
apply_frontdoor_cors_headers(&mut headers, &cors, "https://console.example", None);
let exposed = headers[http::header::ACCESS_CONTROL_EXPOSE_HEADERS]
.to_str()
.unwrap()
.to_ascii_lowercase();
for name in [
"x-aether-body-encoding",
"x-aether-body-field",
"x-aether-body-error",
"x-aether-usage-id",
] {
assert!(exposed.split(',').any(|header| header.trim() == name));
}
assert_eq!(
headers[http::header::ACCESS_CONTROL_ALLOW_CREDENTIALS],
"true"
);
}
}
@@ -3083,7 +3083,7 @@ mod tests {
.await
.expect("stale LKG read must not wait for retention lock");
assert_eq!(stale.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
assert_eq!(stale.stale_targets(), &[target.clone()]);
assert_eq!(stale.stale_targets(), std::slice::from_ref(&target));
assert_eq!(runtime.execution_count(), 1);
assert!(runtime
@@ -3111,7 +3111,7 @@ mod tests {
let load = load_one(&runtime, &client_version).await;
assert_eq!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
assert_eq!(load.stale_targets(), &[target.clone()]);
assert_eq!(load.stale_targets(), std::slice::from_ref(&target));
assert!(runtime
.state
.kv_get(&catalog_lkg_key(&target, client_version.as_str()))
@@ -3142,7 +3142,7 @@ mod tests {
let load = load_one(&runtime, &client_version).await;
assert_eq!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
assert_eq!(load.stale_targets(), &[target.clone()]);
assert_eq!(load.stale_targets(), std::slice::from_ref(&target));
assert_eq!(runtime.execution_count(), 1);
}
@@ -300,6 +300,34 @@ pub(crate) fn classify_local_failover(
policy: &LocalFailoverPolicy,
input: LocalFailoverInput<'_>,
) -> LocalFailoverClassification {
if input.status_code >= 400
&& policy.routing_rules.error_stop_patterns.iter().any(|rule| {
failover_pattern_matches(
&rule.pattern,
&rule.status_codes,
input.response_text,
input.status_code,
)
})
{
return LocalFailoverClassification::StopErrorPattern;
}
if input.status_code == 200
&& policy
.routing_rules
.success_failover_patterns
.iter()
.any(|rule| {
failover_pattern_matches(
&rule.pattern,
&rule.status_codes,
input.response_text,
input.status_code,
)
})
{
return LocalFailoverClassification::RetrySuccessPattern;
}
if policy.stop_status_codes.contains(&input.status_code) {
return LocalFailoverClassification::StopStatusCode;
}
@@ -487,13 +515,27 @@ fn local_failover_regex_rule_matches(
response_text: Option<&str>,
status_code: u16,
) -> bool {
if !rule.status_codes.is_empty() && !rule.status_codes.contains(&status_code) {
failover_pattern_matches(
&rule.pattern,
&rule.status_codes,
response_text,
status_code,
)
}
fn failover_pattern_matches(
pattern: &str,
status_codes: &std::collections::BTreeSet<u16>,
response_text: Option<&str>,
status_code: u16,
) -> bool {
if !status_codes.is_empty() && !status_codes.contains(&status_code) {
return false;
}
let pattern = rule.pattern.trim();
let pattern = pattern.trim();
if pattern.is_empty() {
return !rule.status_codes.is_empty();
return !status_codes.is_empty();
}
let Some(response_text) = response_text else {
@@ -507,6 +549,71 @@ fn local_failover_regex_rule_matches(
#[cfg(test)]
mod tests {
#[test]
fn routing_rules_precede_provider_rules_and_keep_provider_fallback() {
let policy = super::LocalFailoverPolicy {
routing_rules: aether_routing_core::RoutingFailoverRules {
success_failover_patterns: vec![aether_routing_core::RoutingFailoverRule {
pattern: "(?i)capacity.*exhausted".to_string(),
..Default::default()
}],
error_stop_patterns: vec![aether_routing_core::RoutingFailoverRule {
pattern: "invalid.*parameter".to_string(),
status_codes: [400].into_iter().collect(),
}],
},
stop_status_codes: [200, 403].into_iter().collect(),
continue_status_codes: [400].into_iter().collect(),
..Default::default()
};
for (status, body, expected) in [
(
200,
"CAPACITY exhausted",
super::LocalFailoverClassification::RetrySuccessPattern,
),
(
400,
"invalid request parameter",
super::LocalFailoverClassification::StopErrorPattern,
),
(
400,
"capacity exhausted",
super::LocalFailoverClassification::RetryStatusCode,
),
(
403,
"permission denied",
super::LocalFailoverClassification::StopStatusCode,
),
(
429,
"rate limited",
super::LocalFailoverClassification::RetryUpstreamFailure,
),
] {
assert_eq!(
super::classify_local_failover(
&policy,
super::LocalFailoverInput::new(status, Some(body))
),
expected
);
}
}
#[test]
fn provider_transport_stop_rule_is_respected() {
let policy = super::LocalFailoverPolicy {
stop_on_transport_errors: true,
..Default::default()
};
assert_eq!(
super::classify_local_transport_error(&policy),
super::LocalTransportFailoverClassification::StopTransportError
);
}
use std::collections::BTreeSet;
use super::{
@@ -4,7 +4,7 @@ use aether_contracts::ExecutionPlan;
use serde_json::{json, Value};
use tracing::debug;
use aether_routing_core::RoutingExecutionPolicy;
use aether_routing_core::{RoutingExecutionPolicy, RoutingFailoverRules};
use crate::provider_transport::GatewayProviderTransportSnapshot;
use crate::AppState;
@@ -14,6 +14,7 @@ pub(crate) const ROUTING_EXECUTION_POLICY_REPORT_FIELD: &str = "routing_executio
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct LocalFailoverPolicy {
pub(crate) routing_rules: RoutingFailoverRules,
pub(crate) max_retries: Option<u64>,
pub(crate) max_transfer_count: u64,
pub(crate) max_transfer_timeout_seconds: u64,
@@ -29,6 +30,7 @@ pub(crate) struct LocalFailoverPolicy {
impl Default for LocalFailoverPolicy {
fn default() -> Self {
Self {
routing_rules: RoutingFailoverRules::default(),
max_retries: None,
max_transfer_count: 0,
max_transfer_timeout_seconds: 0,
@@ -61,8 +63,10 @@ pub(crate) async fn resolve_local_failover_policy(
Ok(Some(transport)) => local_failover_policy_from_transport(&transport),
Ok(None) | Err(_) => LocalFailoverPolicy::default(),
};
let cyber_continue_failover = routing_execution_policy_from_report_context(report_context)
.is_some_and(|policy| policy.cyber_continue_failover);
let routing_policy =
routing_execution_policy_from_report_context(report_context).unwrap_or_default();
let cyber_continue_failover = routing_policy.cyber_continue_failover;
policy.routing_rules = routing_policy.failover_rules;
policy.stop_cyber_policy_errors = !cyber_continue_failover;
debug!(
event_name = "local_failover_policy_loaded",
@@ -80,6 +84,8 @@ pub(crate) async fn resolve_local_failover_policy(
stop_on_transport_errors = policy.stop_on_transport_errors,
success_failover_pattern_count = policy.success_failover_patterns.len(),
error_stop_pattern_count = policy.error_stop_patterns.len(),
global_success_pattern_count = policy.routing_rules.success_failover_patterns.len(),
global_stop_pattern_count = policy.routing_rules.error_stop_patterns.len(),
cyber_continue_failover,
"gateway loaded local failover policy from transport snapshot"
);
@@ -122,6 +128,7 @@ pub(crate) fn local_failover_policy_from_transport(
});
LocalFailoverPolicy {
routing_rules: RoutingFailoverRules::default(),
max_retries,
max_transfer_count: provider_config
.and_then(|value| value.get("max_transfer_count"))
@@ -184,6 +191,10 @@ pub(crate) fn local_failover_policy_from_report_context(
.as_object()?;
Some(LocalFailoverPolicy {
routing_rules: object
.get("routing_rules")
.and_then(|value| serde_json::from_value(value.clone()).ok())
.unwrap_or_default(),
max_retries: object.get("max_retries").and_then(parse_u64_value),
max_transfer_count: object
.get("max_transfer_count")
@@ -267,6 +278,7 @@ fn parse_status_code_list(value: &Value) -> BTreeSet<u16> {
fn local_failover_policy_to_value(policy: &LocalFailoverPolicy) -> Value {
json!({
"routing_rules": policy.routing_rules,
"max_retries": policy.max_retries,
"max_transfer_count": policy.max_transfer_count,
"max_transfer_timeout_seconds": policy.max_transfer_timeout_seconds,
@@ -525,6 +537,7 @@ mod tests {
assert_eq!(
local_failover_policy_from_report_context(Some(&report_context)),
Some(LocalFailoverPolicy {
routing_rules: Default::default(),
max_retries: Some(2),
max_transfer_count: 10,
max_transfer_timeout_seconds: 60,
@@ -3764,18 +3764,19 @@ mod tests {
.await;
assert_eq!(normal_batch.len(), 1);
let retry_states = metrics
.retry_states
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
assert_eq!(
retry_states
.get(&(0, RequestCandidateQueueLane::Normal))
.map(|state| state.attempt),
Some(1)
);
assert!(!retry_states.contains_key(&(0, RequestCandidateQueueLane::Active)));
drop(retry_states);
{
let retry_states = metrics
.retry_states
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
assert_eq!(
retry_states
.get(&(0, RequestCandidateQueueLane::Normal))
.map(|state| state.attempt),
Some(1)
);
assert!(!retry_states.contains_key(&(0, RequestCandidateQueueLane::Active)));
}
assert!(request_candidate_retry_is_ready(
&metrics,
0,
@@ -0,0 +1,336 @@
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll};
use aether_routing_core::RoutingExecutionPolicy;
use axum::body::{Body, Bytes, HttpBody};
use http::Response;
use http_body::{Frame, SizeHint};
use http_body_util::BodyExt;
use crate::request_diagnostics::{scope_request_diagnostics_with, RequestDiagnostics};
use crate::GatewayError;
tokio::task_local! {
static CANCEL_ON_CLIENT_DISCONNECT: Arc<AtomicBool>;
}
pub(crate) fn configure_client_disconnect(policy: RoutingExecutionPolicy) {
let _ = CANCEL_ON_CLIENT_DISCONNECT.try_with(|cancel| {
cancel.store(policy.cancel_on_client_disconnect, Ordering::Release);
});
}
pub(crate) fn cancel_on_client_disconnect() -> bool {
CANCEL_ON_CLIENT_DISCONNECT
.try_with(|cancel| cancel.load(Ordering::Acquire))
.unwrap_or(false)
}
pub(crate) async fn run_request<F>(future: F) -> Result<Response<Body>, GatewayError>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
let cancel = Arc::new(AtomicBool::new(true));
let diagnostics = Arc::new(RequestDiagnostics::default());
let cancel_for_response = Arc::clone(&cancel);
let future = CANCEL_ON_CLIENT_DISCONNECT.scope(
Arc::clone(&cancel),
scope_request_diagnostics_with(Some(Arc::clone(&diagnostics)), async move {
let response = future.await?;
if cancel_for_response.load(Ordering::Acquire) {
return Ok(response);
}
Ok(response.map(|body| {
Body::new(CompleteOnDisconnectBody {
body: Some(body),
diagnostics,
})
}))
}),
);
CompleteOnDisconnectRequest {
future: Some(Box::pin(future)),
cancel,
}
.await
}
struct CompleteOnDisconnectRequest<F>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
future: Option<Pin<Box<F>>>,
cancel: Arc<AtomicBool>,
}
impl<F> Future for CompleteOnDisconnectRequest<F>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
type Output = Result<Response<Body>, GatewayError>;
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
let result = self
.future
.as_mut()
.expect("request future")
.as_mut()
.poll(context);
if result.is_ready() {
self.future.take();
}
result
}
}
impl<F> Drop for CompleteOnDisconnectRequest<F>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
fn drop(&mut self) {
if self.cancel.load(Ordering::Acquire) {
return;
}
if let (Some(future), Ok(runtime)) =
(self.future.take(), tokio::runtime::Handle::try_current())
{
runtime.spawn(async move {
if let Ok(response) = future.await {
drain_body(response.into_body()).await;
}
});
}
}
}
struct CompleteOnDisconnectBody {
body: Option<Body>,
diagnostics: Arc<RequestDiagnostics>,
}
impl HttpBody for CompleteOnDisconnectBody {
type Data = Bytes;
type Error = axum::Error;
fn poll_frame(
mut self: Pin<&mut Self>,
context: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
let Some(body) = self.body.as_mut() else {
return Poll::Ready(None);
};
let result = Pin::new(body).poll_frame(context);
if matches!(result, Poll::Ready(None | Some(Err(_)))) {
self.body.take();
}
result
}
fn is_end_stream(&self) -> bool {
self.body.as_ref().is_none_or(HttpBody::is_end_stream)
}
fn size_hint(&self) -> SizeHint {
self.body
.as_ref()
.map(HttpBody::size_hint)
.unwrap_or_else(|| SizeHint::with_exact(0))
}
}
impl Drop for CompleteOnDisconnectBody {
fn drop(&mut self) {
let Some(body) = self.body.take().filter(|body| !body.is_end_stream()) else {
return;
};
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
runtime.spawn(scope_request_diagnostics_with(
Some(Arc::clone(&self.diagnostics)),
drain_body(body),
));
}
}
}
async fn drain_body(mut body: Body) {
while let Some(frame) = body.frame().await {
if frame.is_err() {
break;
}
}
}
#[cfg(test)]
mod tests {
use std::io;
use std::time::Duration;
use futures_util::stream;
use http::HeaderMap;
use http_body_util::StreamBody;
use tokio::sync::{mpsc, oneshot};
use super::*;
#[tokio::test]
async fn disconnected_request_finishes_and_keeps_admission_and_diagnostics() {
let gate = aether_runtime::ConcurrencyGate::new("disconnect_request", 1);
let permit = gate.try_acquire().unwrap();
let (started_tx, started_rx) = oneshot::channel();
let (release_tx, release_rx) = oneshot::channel();
let (finished_tx, finished_rx) = oneshot::channel();
let request = tokio::spawn(run_request(async move {
let _permit = permit;
configure_client_disconnect(RoutingExecutionPolicy::default());
started_tx.send(()).unwrap();
release_rx.await.unwrap();
assert!(crate::request_diagnostics::current_request_diagnostics().is_some());
finished_tx.send(()).unwrap();
Ok(Response::new(Body::empty()))
}));
started_rx.await.unwrap();
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert_eq!(gate.snapshot().in_flight, 1);
release_tx.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(1), finished_rx)
.await
.unwrap()
.unwrap();
assert_eq!(gate.snapshot().in_flight, 0);
}
#[tokio::test]
async fn enabled_cancellation_and_unresolved_requests_drop_immediately() {
for resolve_policy in [false, true] {
let (started_tx, started_rx) = oneshot::channel();
let (release_tx, release_rx) = oneshot::channel::<()>();
let request = tokio::spawn(run_request(async move {
if resolve_policy {
configure_client_disconnect(RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
});
}
started_tx.send(()).unwrap();
release_rx.await.unwrap();
Ok(Response::new(Body::empty()))
}));
started_rx.await.unwrap();
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert!(release_tx.send(()).is_err());
}
}
#[tokio::test]
async fn disconnected_body_drains_without_buffering_and_holds_admission() {
for consume_first_chunk in [false, true] {
let gate = aether_runtime::ConcurrencyGate::new("disconnect_body", 1);
let permit = gate.try_acquire().unwrap();
let (sender, receiver) = mpsc::channel(1);
let (finished_tx, finished_rx) = oneshot::channel();
let response = run_request(async move {
configure_client_disconnect(RoutingExecutionPolicy::default());
let body = Body::from_stream(stream::unfold(
(receiver, finished_tx, permit),
|(mut receiver, finished_tx, permit)| async move {
match receiver.recv().await {
Some(bytes) => {
Some((Ok::<_, io::Error>(bytes), (receiver, finished_tx, permit)))
}
None => {
assert!(crate::request_diagnostics::current_request_diagnostics()
.is_some());
finished_tx.send(()).unwrap();
None
}
}
},
));
Ok(Response::new(body))
})
.await
.unwrap();
let mut body = response.into_body();
if consume_first_chunk {
sender.send(Bytes::from_static(b"first")).await.unwrap();
assert_eq!(
body.frame().await.unwrap().unwrap().into_data().unwrap(),
"first"
);
}
drop(body);
assert_eq!(gate.snapshot().in_flight, 1);
tokio::time::timeout(Duration::from_secs(1), async {
for _ in 0..100 {
sender.send(Bytes::from_static(b"remaining")).await.unwrap();
}
drop(sender);
finished_rx.await.unwrap();
})
.await
.unwrap();
assert_eq!(gate.snapshot().in_flight, 0);
}
}
#[tokio::test]
async fn enabled_cancellation_drops_stream_receiver() {
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
let response = run_request(async move {
configure_client_disconnect(RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
});
Ok(Response::new(Body::from_stream(stream::unfold(
receiver,
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
))))
})
.await
.unwrap();
drop(response);
assert!(sender.is_closed());
}
#[tokio::test]
async fn connected_response_preserves_headers_size_hint_and_trailers() {
let response = run_request(async {
configure_client_disconnect(RoutingExecutionPolicy::default());
Ok(Response::builder()
.status(201)
.header("x-test", "unchanged")
.body(Body::from("hello"))
.unwrap())
})
.await
.unwrap();
assert_eq!(response.status(), 201);
assert_eq!(response.headers()["x-test"], "unchanged");
assert_eq!(response.body().size_hint().exact(), Some(5));
assert_eq!(
response.into_body().collect().await.unwrap().to_bytes(),
"hello"
);
let mut trailers = HeaderMap::new();
trailers.insert("x-finished", "yes".parse().unwrap());
let response = run_request(async move {
configure_client_disconnect(RoutingExecutionPolicy::default());
let frames = stream::iter([
Ok::<_, io::Error>(Frame::data(Bytes::from_static(b"hello"))),
Ok(Frame::trailers(trailers)),
]);
Ok(Response::new(Body::new(StreamBody::new(frames))))
})
.await
.unwrap();
let collected = response.into_body().collect().await.unwrap();
assert_eq!(collected.trailers().unwrap()["x-finished"], "yes");
assert_eq!(collected.to_bytes(), "hello");
}
}
+31 -32
View File
@@ -57,7 +57,7 @@ pub(crate) fn resolve_gateway_routing_policy(
let config = serde_json::from_value::<RoutingGroupConfig>(input.group_config_json.clone())
.map_err(|_| invalid_routing_group_config())?;
resolve_routing_policy(
let policy = resolve_routing_policy(
&config,
RoutingPolicyInput {
group_id: input.group_id,
@@ -73,7 +73,9 @@ pub(crate) fn resolve_gateway_routing_policy(
phase: input.phase,
},
)
.map_err(routing_policy_error)
.map_err(routing_policy_error)?;
crate::request_lifecycle::configure_client_disconnect(policy.execution_policy.clone());
Ok(policy)
}
pub(crate) fn resolve_gateway_static_default_routing_policy(
@@ -82,6 +84,7 @@ pub(crate) fn resolve_gateway_static_default_routing_policy(
let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else {
return Ok(None);
};
crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy.clone());
Ok(Some(ResolvedRoutingPolicy {
group_id: input.group_id.map(str::to_string),
@@ -142,28 +145,11 @@ fn static_default_policy_fields(
.ok_or_else(invalid_routing_group_config)?,
None => DEFAULT_STICKY_KEY_ATTEMPTS,
};
let enable_cf_heartbeat = routing_bool_field(
default_policy.get("enable_cf_heartbeat"),
"enable_cf_heartbeat",
)?;
// Older strategies stored separate image/text heartbeat flags. Treat
// either legacy flag as enabling the unified CF heartbeat setting while
// allowing newly saved strategies to use only the canonical key.
let legacy_image_heartbeat = routing_bool_field(
default_policy.get("enable_openai_image_sync_heartbeat"),
"enable_openai_image_sync_heartbeat",
)?;
let legacy_text_heartbeat = routing_bool_field(
default_policy.get("enable_standard_text_sync_heartbeat"),
"enable_standard_text_sync_heartbeat",
)?;
let execution_policy = aether_routing_core::RoutingExecutionPolicy {
enable_cf_heartbeat: enable_cf_heartbeat || legacy_image_heartbeat || legacy_text_heartbeat,
cyber_continue_failover: routing_bool_field(
default_policy.get("cyber_continue_failover"),
"cyber_continue_failover",
)?,
};
let execution_policy: aether_routing_core::RoutingExecutionPolicy =
serde_json::from_value(Value::Object(default_policy.clone()))
.map_err(|_| invalid_routing_group_config())?;
aether_routing_core::validate_routing_failover_rules(&execution_policy.failover_rules)
.map_err(|_| invalid_routing_group_config())?;
Ok(Some(RoutingDefaultPolicy {
priority_mode,
@@ -174,13 +160,6 @@ fn static_default_policy_fields(
}))
}
fn routing_bool_field(value: Option<&Value>, _field: &str) -> Result<bool, GatewayError> {
match value {
Some(value) => value.as_bool().ok_or_else(invalid_routing_group_config),
None => Ok(false),
}
}
fn routing_array_field_is_missing_or_empty(
object: &serde_json::Map<String, Value>,
key: &str,
@@ -238,7 +217,14 @@ mod tests {
"default_policy": {
"priority_mode": "global_key",
"scheduling_mode": "load_balance",
"keep_priority_on_conversion": true
"keep_priority_on_conversion": true,
"cancel_on_client_disconnect": true,
"max_transfer_count": 3,
"max_transfer_timeout_seconds": 90,
"failover_rules": {
"success_failover_patterns": [{"pattern": "(?i)capacity.*exhausted"}],
"error_stop_patterns": [{"status_codes": [400]}]
}
},
"allowed_models": ["legacy-model"],
"model_policies": [],
@@ -274,6 +260,19 @@ mod tests {
.expect("full policy should resolve");
assert_eq!(static_policy, full_policy);
assert_eq!(static_policy.execution_policy.max_transfer_count, 3);
assert_eq!(
static_policy.execution_policy.max_transfer_timeout_seconds,
90
);
assert_eq!(
static_policy
.execution_policy
.failover_rules
.error_stop_patterns
.len(),
1
);
assert_eq!(
static_policy.priority_mode,
RoutingSetPriorityMode::GlobalKey
+14 -2
View File
@@ -274,7 +274,13 @@ mod tests {
"priority_mode": "provider",
"scheduling_mode": "cache_affinity",
"keep_priority_on_conversion": false,
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS,
"max_transfer_count": 0,
"max_transfer_timeout_seconds": 0,
"failover_rules": {
"success_failover_patterns": [],
"error_stop_patterns": []
}
})
);
@@ -323,7 +329,13 @@ mod tests {
"priority_mode": "provider",
"scheduling_mode": "cache_affinity",
"keep_priority_on_conversion": false,
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS,
"max_transfer_count": 0,
"max_transfer_timeout_seconds": 0,
"failover_rules": {
"success_failover_patterns": [],
"error_stop_patterns": []
}
})
);
}
@@ -196,11 +196,12 @@ impl AppState {
{
let session = session.into();
#[cfg(test)]
if self.auth_session_store.is_some() && self.auth_user_store.is_some() {
if let (Some(session_store), Some(user_store)) = (
self.auth_session_store.as_ref(),
self.auth_user_store.as_ref(),
) {
let existing = {
self.auth_user_store
.as_ref()
.expect("checked auth user store")
user_store
.lock()
.expect("auth user store should lock")
.get(&session.user_id)
@@ -217,12 +218,7 @@ impl AppState {
let Some(existing) = existing else {
return Ok(None);
};
let mut users = self
.auth_user_store
.as_ref()
.expect("checked auth user store")
.lock()
.expect("auth user store should lock");
let mut users = user_store.lock().expect("auth user store should lock");
let user = users.entry(session.user_id.clone()).or_insert(existing);
if user.password_hash.as_deref() != Some(expected_password_hash)
|| !user.auth_source.eq_ignore_ascii_case("local")
@@ -238,10 +234,7 @@ impl AppState {
.or(session.last_seen_at)
.unwrap_or_else(chrono::Utc::now);
user.last_login_at = Some(now);
let mut sessions = self
.auth_session_store
.as_ref()
.expect("checked auth session store")
let mut sessions = session_store
.lock()
.expect("auth session store should lock");
for existing in sessions.values_mut() {
@@ -802,7 +802,7 @@ impl AppState {
}
return Ok(Some(LdapAuthProvisioningResult {
user,
owned_wallet_id: initialized.created.then(|| initialized.wallet.id),
owned_wallet_id: initialized.created.then_some(initialized.wallet.id),
}));
}
@@ -1,7 +1,7 @@
use super::{
any, build_router_with_state, build_state_with_execution_runtime_override, json, start_server,
Arc, Body, Bytes, HeaderValue, Infallible, Json, Mutex, Request, Response, Router, StatusCode,
TRACE_ID_HEADER,
LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER,
};
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
@@ -55,6 +55,24 @@ fn hash_api_key(value: &str) -> String {
format!("{:x}", hasher.finalize())
}
async fn build_cancelling_gateway(state: crate::AppState) -> Router {
state
.data
.update_routing_group(
"system-default",
aether_data_contracts::repository::routing_profiles::UpdateRoutingGroupRecord {
config_json: Some(json!({"default_policy": {"cancel_on_client_disconnect": true}})),
version: Some(2),
updated_at: 2,
..Default::default()
},
)
.await
.expect("routing policy should update")
.expect("default strategy should exist");
build_router_with_state(state)
}
fn sample_local_openai_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
StoredAuthApiKeySnapshot::new(
user_id.to_string(),
@@ -384,7 +402,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
vec![sample_local_openai_key()],
));
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let gateway = build_router_with_state(
let gateway = build_cancelling_gateway(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
@@ -395,7 +413,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
DEVELOPMENT_ENCRYPTION_KEY,
),
),
);
).await;
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
@@ -467,7 +485,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
vec![sample_local_openai_key()],
));
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let gateway = build_router_with_state(
let gateway = build_cancelling_gateway(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
@@ -478,7 +496,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
DEVELOPMENT_ENCRYPTION_KEY,
),
),
);
).await;
let (gateway_url, gateway_handle) = start_server(gateway).await;
let request = reqwest::Client::new()
@@ -608,7 +626,7 @@ async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error_
auth_repository,
candidate_selection_repository,
provider_catalog_repository,
request_candidate_repository,
Arc::clone(&request_candidate_repository),
DEVELOPMENT_ENCRYPTION_KEY,
),
),
@@ -631,17 +649,41 @@ async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error_
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
response
.headers()
.get(http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok()),
Some("text/event-stream")
Some("application/json")
);
let body_text = response.text().await.expect("response body should read");
assert!(body_text.contains("\"rate_limit_error\""));
assert!(body_text.contains("\"slow down\""));
assert_eq!(
response
.headers()
.get(LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER)
.and_then(|value| value.to_str().ok()),
Some("execution_runtime_candidates_exhausted")
);
let body_json: serde_json::Value = response.json().await.expect("response body should parse");
assert_eq!(body_json["error"]["type"], "http_error");
let stored_candidates = request_candidate_repository
.list_by_request_id("trace-openai-chat-stream-prefetch-error-123")
.await
.expect("request candidate trace should read");
let failed_candidate = stored_candidates
.iter()
.find(|candidate| candidate.status == RequestCandidateStatus::Failed)
.expect("prefetched error should mark the attempted candidate as failed");
assert!(stored_candidates
.iter()
.all(|candidate| candidate.status != RequestCandidateStatus::Success));
assert_eq!(failed_candidate.status_code, Some(429));
assert_eq!(
failed_candidate.error_type.as_deref(),
Some("rate_limit_error")
);
assert_eq!(failed_candidate.error_message.as_deref(), Some("slow down"));
assert!(failed_candidate.finished_at_unix_ms.is_some());
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
@@ -349,7 +349,7 @@ async fn gateway_executes_gemini_chat_stream_via_local_decision_gate_with_local_
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":33,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -419,7 +419,7 @@ async fn gateway_executes_gemini_chat_stream_via_local_decision_gate_with_local_
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_execution_runtime_request = seen_execution_runtime
@@ -326,7 +326,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_with_local_s
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -396,7 +396,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_with_local_s
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_execution_runtime_request = seen_execution_runtime
@@ -846,7 +846,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"response\\\":{\\\"candidates\\\":[]},\\\"remainingCredits\\\":42,\\\"consumedCredits\\\":1,\\\"traceId\\\":\\\"trace-upstream-1\\\"}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"response\\\":{\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]},\\\"remainingCredits\\\":42,\\\"consumedCredits\\\":1,\\\"traceId\\\":\\\"trace-upstream-1\\\"}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -934,7 +934,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_refresh_request = seen_refresh
@@ -1354,7 +1354,7 @@ async fn gateway_executes_vertex_ai_gemini_cli_stream_via_local_decision_gate_wi
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -1422,7 +1422,7 @@ async fn gateway_executes_vertex_ai_gemini_cli_stream_via_local_decision_gate_wi
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_execution_runtime_request = seen_execution_runtime
@@ -5010,6 +5010,7 @@ fn retired_api_format_occurrences_are_whitelisted() {
"crates/aether-ai/formats/src/formats/registry.rs",
"crates/aether-data/runtime/src/migrate.rs",
"crates/aether-data/runtime/src/lifecycle/migrate/tests.rs",
"crates/aether-data/runtime/src/lifecycle/migrate/tests/policy_nulls.rs",
"crates/aether-usage/runtime/src/report.rs",
"frontend/src/api/endpoints/types/__tests__/api-format.spec.ts",
"frontend/src/views/admin/module-management/modelDirectivesConfig.ts",
@@ -96,12 +96,11 @@ fn aether_data_backend_pool_modules_do_not_own_maintenance_sql() {
#[test]
fn wallet_maintenance_sql_is_partitioned_by_driver() {
let facade = read_workspace_file("crates/aether-data/runtime/src/backend/wallet.rs");
for module in ["mod postgres;"] {
assert!(
facade.contains(module),
"wallet facade should declare {module}"
);
}
let module = "mod postgres;";
assert!(
facade.contains(module),
"wallet facade should declare {module}"
);
for forbidden in [
"sqlx::",
"PostgresBackend",
@@ -115,24 +114,22 @@ fn wallet_maintenance_sql_is_partitioned_by_driver() {
);
}
for (driver, backend) in [("postgres", "PostgresBackend")] {
let path = format!("crates/aether-data/runtime/src/backend/wallet/{driver}.rs");
let source = read_workspace_file(&path);
assert!(source.contains(&format!("impl {backend}")));
assert!(source.contains("aggregate_wallet_daily_usage"));
assert!(source.contains("sqlx::query"));
}
let (driver, backend) = ("postgres", "PostgresBackend");
let path = format!("crates/aether-data/runtime/src/backend/wallet/{driver}.rs");
let source = read_workspace_file(&path);
assert!(source.contains(&format!("impl {backend}")));
assert!(source.contains("aggregate_wallet_daily_usage"));
assert!(source.contains("sqlx::query"));
}
#[test]
fn table_maintenance_is_partitioned_for_each_driver() {
let facade = read_workspace_file("crates/aether-data/runtime/src/backend/maintenance.rs");
for module in ["mod postgres;"] {
assert!(
facade.contains(module),
"maintenance facade should declare {module}"
);
}
let module = "mod postgres;";
assert!(
facade.contains(module),
"maintenance facade should declare {module}"
);
for forbidden in [
"impl PostgresBackend",
"VACUUM ANALYZE",
@@ -156,12 +153,11 @@ fn table_maintenance_is_partitioned_for_each_driver() {
#[test]
fn system_driver_database_operations_are_partitioned() {
let facade = read_workspace_file("crates/aether-data/runtime/src/backend/system.rs");
for module in ["mod postgres;"] {
assert!(
facade.contains(module),
"system facade should declare {module}"
);
}
let module = "mod postgres;";
assert!(
facade.contains(module),
"system facade should declare {module}"
);
for forbidden in [
"impl PostgresBackend",
"fn map_postgres_stats_daily_aggregate(",
@@ -1521,12 +1517,11 @@ fn lifecycle_migrations_are_partitioned_by_driver() {
types.contains(required),
"migrate/types.rs should own {required}"
);
for forbidden in ["PgPool"] {
assert!(
!types.contains(forbidden),
"migrate/types.rs should remain driver-independent from {forbidden}"
);
}
let forbidden = "PgPool";
assert!(
!types.contains(forbidden),
"migrate/types.rs should remain driver-independent from {forbidden}"
);
let postgres =
read_workspace_file("crates/aether-data/runtime/src/lifecycle/migrate/postgres.rs");
@@ -1905,12 +1900,11 @@ fn gateway_system_config_types_are_owned_by_aether_data() {
}
let data_backends =
read_workspace_file("crates/aether-data/runtime/src/backend/maintenance.rs");
for pattern in ["postgres.list_system_config_entries().await"] {
assert!(
data_backends.contains(pattern),
"aether-data backends should own driver-specific system config dispatch {pattern}"
);
}
let pattern = "postgres.list_system_config_entries().await";
assert!(
data_backends.contains(pattern),
"aether-data backends should own driver-specific system config dispatch {pattern}"
);
for pattern in [
"|(key, value, description, updated_at_unix_secs)|",
"Ok((0, 0, 0, 0))",
+17 -4
View File
@@ -326,7 +326,7 @@ async fn gateway_reads_video_task_detail_via_internal_async_task_endpoint() {
}
#[tokio::test]
async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endpoint() {
async fn gateway_redirects_persisted_openai_video_url_from_authenticated_internal_endpoint() {
let repository = Arc::new(InMemoryVideoTaskRepository::default());
let mut task = sample_video_task(
"task-redirect",
@@ -341,7 +341,10 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endp
.upsert(task)
.await
.expect("upsert should succeed");
assert_eq!(stored.video_url, None);
assert_eq!(
stored.video_url.as_deref(),
Some("https://8.8.8.8/video-task-redirect.mp4")
);
let state = AppState::new()
.expect("gateway state should build")
@@ -349,7 +352,10 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endp
let (gateway_url, gateway_handle, access_token) =
start_authenticated_operational_server(state).await;
let client = authenticated_operational_client(&access_token);
let client = super::authenticated_operational_client_with_builder(
reqwest::Client::builder().redirect(reqwest::redirect::Policy::none()),
&access_token,
);
let response = client
.get(format!(
"{gateway_url}/_gateway/async-tasks/video-tasks/task-redirect/video"
@@ -358,7 +364,14 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endp
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::NOT_FOUND);
assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT);
assert_eq!(
response
.headers()
.get("location")
.and_then(|value| value.to_str().ok()),
stored.video_url.as_deref()
);
gateway_handle.abort();
}
@@ -1,3 +1,4 @@
mod keys;
mod quota;
mod routes;
mod rules_reveal;
@@ -8,7 +8,7 @@ use aether_data::repository::global_models::InMemoryGlobalModelReadRepository;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
use aether_data_contracts::repository::global_models::{
AdminProviderModelListQuery, GlobalModelReadRepository,
AdminGlobalModelListQuery, AdminProviderModelListQuery, GlobalModelReadRepository,
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider,
@@ -20,8 +20,9 @@ use http::StatusCode;
use serde_json::json;
use super::super::super::{
build_router_with_state, build_state_with_execution_runtime_override, sample_bound_auth_config,
sample_bound_key, sample_endpoint, sample_key, sample_proxy_node, start_server, AppState,
build_router_with_state, build_state_with_execution_runtime_override,
sample_admin_global_model, sample_bound_auth_config, sample_bound_key, sample_endpoint,
sample_key, sample_proxy_node, start_server, AppState,
};
use crate::constants::{
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
@@ -2421,7 +2422,15 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru
)],
vec![key],
));
let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::default());
let existing_global_model = sample_admin_global_model(
"global-claude-sonnet-4",
"claude-sonnet-4",
"Claude Sonnet 4",
);
let global_model_repository = Arc::new(
InMemoryGlobalModelReadRepository::default()
.with_admin_global_models(vec![existing_global_model.clone()]),
);
let (upstream_url, upstream_handle) = start_server(upstream).await;
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
@@ -2533,7 +2542,16 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru
.and_then(|value| value.get("remaining_fraction")),
Some(&json!(0.25))
);
let imported_provider_models = global_model_repository
let global_models = global_model_repository
.list_admin_global_models(&AdminGlobalModelListQuery {
limit: 100,
..Default::default()
})
.await
.expect("global models should read after quota refresh");
assert_eq!(global_models.total, 1);
assert_eq!(global_models.items, vec![existing_global_model]);
let provider_models = global_model_repository
.list_admin_provider_models(&AdminProviderModelListQuery {
provider_id: "provider-antigravity".to_string(),
is_active: None,
@@ -2541,15 +2559,8 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru
limit: 100,
})
.await
.expect("imported Antigravity provider models should read");
let imported_model_names = imported_provider_models
.iter()
.map(|model| model.provider_model_name.as_str())
.collect::<std::collections::BTreeSet<_>>();
assert!(imported_model_names.contains("claude-sonnet-4"));
assert!(imported_model_names.contains("gemini-2.5-pro"));
assert!(imported_model_names.contains("gemini-3.7-flash-tiered"));
assert!(!imported_model_names.contains("chat_23310"));
.expect("Antigravity provider models should read after quota refresh");
assert!(provider_models.is_empty());
assert_eq!(
reloaded[0]
.upstream_metadata
@@ -479,7 +479,7 @@ async fn gateway_returns_service_unavailable_for_admin_provider_endpoint_create_
}
#[tokio::test]
async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_principal() {
async fn gateway_creates_admin_http_provider_endpoint_locally_with_trusted_admin_principal() {
let upstream_hits = Arc::new(Mutex::new(0usize));
let upstream_hits_clone = Arc::clone(&upstream_hits);
let upstream = Router::new().route(
@@ -522,7 +522,7 @@ async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_prin
.json(&json!({
"provider_id": "provider-openai",
"api_format": "openai:chat",
"base_url": "https://api.openai.example/",
"base_url": "http://api.openai.example:8080/",
"custom_path": "/v1/chat/completions",
"max_retries": 5,
"config": {"foo": "bar"},
@@ -537,7 +537,7 @@ async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_prin
assert_eq!(payload["provider_id"], "provider-openai");
assert_eq!(payload["provider_name"], "openai");
assert_eq!(payload["api_format"], "openai:chat");
assert_eq!(payload["base_url"], "https://api.openai.example");
assert_eq!(payload["base_url"], "http://api.openai.example:8080");
assert_eq!(payload["custom_path"], "/v1/chat/completions");
assert_eq!(payload["max_retries"], 5);
assert_eq!(payload["total_keys"], 0);
@@ -553,7 +553,7 @@ async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_prin
assert_eq!(endpoints.len(), 1);
assert_eq!(endpoints[0].provider_id, "provider-openai");
assert_eq!(endpoints[0].api_format, "openai:chat");
assert_eq!(endpoints[0].base_url, "https://api.openai.example");
assert_eq!(endpoints[0].base_url, "http://api.openai.example:8080");
assert_eq!(endpoints[0].max_retries, Some(5));
gateway_handle.abort();
@@ -658,7 +658,7 @@ async fn gateway_rejects_streaming_policy_for_search_endpoint_before_catalog_wri
}
#[tokio::test]
async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_principal() {
async fn gateway_updates_admin_http_provider_endpoint_locally_with_trusted_admin_principal() {
let upstream_hits = Arc::new(Mutex::new(0usize));
let upstream_hits_clone = Arc::clone(&upstream_hits);
let upstream = Router::new().route(
@@ -720,7 +720,7 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"base_url": "https://updated.openai.example/",
"base_url": "http://updated.openai.example:8080/",
"custom_path": "/v1/responses",
"max_retries": 5,
"is_active": false,
@@ -736,7 +736,7 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin
assert_eq!(payload["id"], "endpoint-openai-chat");
assert_eq!(payload["provider_id"], "provider-openai");
assert_eq!(payload["api_format"], "openai:chat");
assert_eq!(payload["base_url"], "https://updated.openai.example");
assert_eq!(payload["base_url"], "http://updated.openai.example:8080");
assert_eq!(payload["custom_path"], "/v1/responses");
assert_eq!(payload["max_retries"], 5);
assert_eq!(payload["is_active"], false);
@@ -751,7 +751,7 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin
.await
.expect("endpoints should read");
assert_eq!(endpoints.len(), 1);
assert_eq!(endpoints[0].base_url, "https://updated.openai.example");
assert_eq!(endpoints[0].base_url, "http://updated.openai.example:8080");
assert_eq!(endpoints[0].custom_path.as_deref(), Some("/v1/responses"));
assert_eq!(endpoints[0].max_retries, Some(5));
assert!(!endpoints[0].is_active);
@@ -0,0 +1,151 @@
use std::sync::Arc;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use axum::body::Body;
use http::{HeaderMap, HeaderValue, Method, Request, StatusCode};
use http_body_util::BodyExt;
use serde_json::{json, Value};
use super::super::super::{build_router_with_state, sample_endpoint, sample_provider, AppState};
use crate::admin_api::{maybe_build_local_admin_response, AdminRouteRequest};
use crate::audit::AdminAuditEvent;
use crate::constants::{
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
TRUSTED_ADMIN_USER_ROLE_HEADER,
};
use crate::control::resolve_public_request_context;
use crate::data::GatewayDataState;
use crate::tests::send_request;
fn seeded_state() -> AppState {
let mut endpoint = sample_endpoint(
"endpoint-rules",
"provider-rules",
"openai:chat",
"https://example.test",
);
endpoint.header_rules =
Some(json!([{"action": "set", "key": "x-auth", "value": "request-secret"}]));
endpoint.body_rules =
Some(json!([{"action": "set", "path": "auth.token", "value": "body-secret"}]));
endpoint.config = Some(json!({
"private_token": "unrelated-secret",
"response_header_rules": [{"action": "set", "key": "x-auth", "value": "response-secret"}]
}));
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-rules", "custom", 10)],
vec![endpoint],
vec![],
));
AppState::new().unwrap().with_data_state_for_tests(
GatewayDataState::with_provider_catalog_reader_for_tests(repository),
)
}
fn admin_headers() -> HeaderMap {
let mut headers = HeaderMap::new();
for (name, value) in [
(GATEWAY_HEADER, "rust-phase3b"),
(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user"),
(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin"),
(TRUSTED_ADMIN_SESSION_ID_HEADER, "admin-session"),
] {
headers.insert(name, HeaderValue::from_static(value));
}
headers
}
#[tokio::test]
async fn endpoint_rules_reveal_is_scoped_audited_and_not_cached() {
let state = seeded_state();
let context = resolve_public_request_context(
&state,
&Method::GET,
&"/api/admin/endpoints/endpoint-rules/rules/reveal"
.parse()
.unwrap(),
&admin_headers(),
"reveal-test",
)
.await
.unwrap();
let response = maybe_build_local_admin_response(AdminRouteRequest::new(
&state,
&context,
&"127.0.0.1:12345".parse().unwrap(),
&admin_headers(),
None,
))
.await
.unwrap()
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()[http::header::CACHE_CONTROL], "no-store");
assert_eq!(response.headers()[http::header::PRAGMA], "no-cache");
let audit = response.extensions().get::<AdminAuditEvent>().unwrap();
assert_eq!(audit.event_name, "admin_endpoint_rules_revealed");
assert_eq!(audit.action, "reveal_endpoint_rules");
assert_eq!(audit.target_id, "endpoint-rules");
let body = response.into_body().collect().await.unwrap().to_bytes();
let payload: Value = serde_json::from_slice(&body).unwrap();
assert_eq!(payload["header_rules"][0]["value"], "request-secret");
assert_eq!(payload["body_rules"][0]["value"], "body-secret");
assert_eq!(
payload["response_header_rules"][0]["value"],
"response-secret"
);
assert_eq!(payload.as_object().unwrap().len(), 3);
assert!(!payload.to_string().contains("unrelated-secret"));
}
#[tokio::test]
async fn endpoint_rules_reveal_denies_anonymous_and_non_admin_requests() {
let router = build_router_with_state(seeded_state());
for role in [None, Some("user")] {
let mut request =
Request::builder().uri("/api/admin/endpoints/endpoint-rules/rules/reveal");
if let Some(role) = role {
request = request
.header(GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "normal-user")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, role)
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "user-session");
}
let response = send_request(router.clone(), request.body(Body::empty()).unwrap()).await;
assert!(matches!(
response.status(),
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN
));
let body = response.into_body().collect().await.unwrap().to_bytes();
assert!(!String::from_utf8_lossy(&body).contains("request-secret"));
}
}
#[tokio::test]
async fn endpoint_rules_reveal_returns_not_found_and_data_unavailable_without_fallback() {
for (state, expected) in [
(seeded_state(), StatusCode::NOT_FOUND),
(AppState::new().unwrap(), StatusCode::SERVICE_UNAVAILABLE),
] {
let context = resolve_public_request_context(
&state,
&Method::GET,
&"/api/admin/endpoints/missing/rules/reveal".parse().unwrap(),
&admin_headers(),
"reveal-missing-test",
)
.await
.unwrap();
let response = maybe_build_local_admin_response(AdminRouteRequest::new(
&state,
&context,
&"127.0.0.1:12345".parse().unwrap(),
&admin_headers(),
None,
))
.await
.unwrap()
.unwrap();
assert_eq!(response.status(), expected);
}
}
@@ -3705,9 +3705,9 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p
scores[0].hard_state.schedulable(),
"OAuth completion should replace AuthInvalid with a schedulable score"
);
let decrypted_api_key = decrypt_persisted_provider_api_key(&persisted);
let decrypted_api_key = decrypt_persisted_provider_api_key(persisted);
assert_eq!(decrypted_api_key, "new-codex-access-token");
let decrypted_auth_config = decrypt_persisted_provider_auth_config(&persisted);
let decrypted_auth_config = decrypt_persisted_provider_auth_config(persisted);
let auth_config: serde_json::Value =
serde_json::from_str(&decrypted_auth_config).expect("auth config json should parse");
assert_eq!(auth_config["provider_type"], "codex");
@@ -3916,11 +3916,47 @@ async fn gateway_completes_admin_provider_oauth_provider_locally_with_trusted_ad
fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email() {
run_admin_oauth_test(
"gateway_names_new_antigravity_oauth_account_from_google_userinfo_email",
gateway_names_new_antigravity_oauth_account_from_google_userinfo_email_impl,
|| {
assert_antigravity_oauth_account_uses_google_userinfo_email(
"complete",
json!({
"callback_url": "http://localhost:51121/oauth2callback?code=antigravity-code-123&state=cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc"
}),
)
},
);
}
async fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email_impl() {
#[test]
fn gateway_names_imported_antigravity_oauth_account_from_google_userinfo_email() {
run_admin_oauth_test(
"gateway_names_imported_antigravity_oauth_account_from_google_userinfo_email",
|| {
assert_antigravity_oauth_account_uses_google_userinfo_email(
"import-refresh-token",
json!({"refresh_token": "antigravity-import-refresh-token"}),
)
},
);
}
#[test]
fn gateway_names_batch_imported_antigravity_oauth_account_from_google_userinfo_email() {
run_admin_oauth_test(
"gateway_names_batch_imported_antigravity_oauth_account_from_google_userinfo_email",
|| {
assert_antigravity_oauth_account_uses_google_userinfo_email(
"batch-import",
json!({"credentials": "antigravity-import-refresh-token"}),
)
},
);
}
async fn assert_antigravity_oauth_account_uses_google_userinfo_email(
operation: &str,
request_body: Value,
) {
let upstream_hits = Arc::new(Mutex::new(0usize));
let upstream_hits_clone = Arc::clone(&upstream_hits);
let upstream = Router::new().fallback(any(move |_request: Request| {
@@ -4025,15 +4061,13 @@ async fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email_
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/provider-oauth/providers/provider-antigravity/complete"
"{gateway_url}/api/admin/provider-oauth/providers/provider-antigravity/{operation}"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"callback_url": "http://localhost:51121/oauth2callback?code=antigravity-code-123&state=cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc"
}))
.json(&request_body)
.send()
.await
.expect("request should succeed");
@@ -4041,9 +4075,22 @@ async fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email_
let status = response.status();
let payload: Value = response.json().await.expect("json body should parse");
assert_eq!(status, StatusCode::OK, "payload={payload}");
assert_eq!(payload["provider_type"], "antigravity");
assert_eq!(payload["email"], "[email protected]");
assert_eq!(payload["replaced"], false);
let account_result = if operation == "batch-import" {
assert_eq!(payload["total"], 1);
assert_eq!(payload["success"], 1, "payload={payload}");
assert_eq!(payload["failed"], 0);
assert_eq!(payload["results"][0]["status"], "success");
assert_eq!(
payload["results"][0]["key_name"],
"[email protected]"
);
&payload["results"][0]
} else {
assert_eq!(payload["provider_type"], "antigravity");
assert_eq!(payload["email"], "[email protected]");
&payload
};
assert_eq!(account_result["replaced"], false);
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
assert_eq!(*user_info_hits.lock().expect("mutex should lock"), 1);
assert_eq!(
@@ -4055,7 +4102,7 @@ async fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email_
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
let key_id = payload["key_id"]
let key_id = account_result["key_id"]
.as_str()
.expect("created key id should be returned")
.to_string();
@@ -3834,12 +3834,12 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
Some(&Value::Null)
);
assert_eq!(
decrypt_test_provider_catalog_credential(&key, ProviderCatalogCredentialField::ApiKey,),
decrypt_test_provider_catalog_credential(key, ProviderCatalogCredentialField::ApiKey,),
"oauth-access-token-new"
);
let auth_config =
decrypt_test_provider_catalog_credential(&key, ProviderCatalogCredentialField::AuthConfig);
decrypt_test_provider_catalog_credential(key, ProviderCatalogCredentialField::AuthConfig);
let auth_config: Value =
serde_json::from_str(&auth_config).expect("oauth auth config json should parse");
assert_eq!(auth_config["provider_type"], "codex");
@@ -2366,6 +2366,265 @@ async fn gateway_handles_admin_usage_detail_with_ref_backed_bodies() {
upstream_handle.abort();
}
fn sample_selective_body_usage() -> StoredRequestUsageAudit {
let mut usage = sample_usage_row(
"usage-selected-body",
"req-selected-body",
Some("user-1"),
Some("key-1"),
Some("primary"),
"OpenAI",
"gpt-5",
"completed",
120,
30,
0.3,
0.36,
DAY_1_UNIX_SECS,
);
usage.request_body = Some(json!({ "marker": "request_body" }));
usage.provider_request_body = Some(json!({ "marker": "provider_request_body" }));
usage.response_body = Some(json!({ "marker": "response_body" }));
usage.client_response_body = Some(json!({ "marker": "client_response_body" }));
usage
}
#[tokio::test]
async fn gateway_admin_usage_detail_returns_only_the_selected_body() {
let fields = [
"request_body",
"provider_request_body",
"response_body",
"client_response_body",
];
for detached in [false, true] {
let usage = sample_selective_body_usage();
let repository = if detached {
InMemoryUsageReadRepository::seed_with_detached_bodies(vec![usage])
} else {
InMemoryUsageReadRepository::seed(vec![usage])
};
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_usage_reader_for_tests(Arc::new(
repository,
)));
for selected in fields {
let response = local_admin_usage_response(
&state, http::Method::GET,
&format!("/api/admin/usage/usage-selected-body?include_bodies=true&body_field={selected}"),
None,
).await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should collect");
let payload: serde_json::Value =
serde_json::from_slice(&bytes).expect("json should parse");
for field in fields {
assert_eq!(payload[format!("has_{field}")], true);
if field == selected {
assert_eq!(payload[field]["marker"], selected);
} else {
assert!(
payload[field].is_null(),
"unselected {field} must not be returned"
);
}
}
assert!(payload["body_load_errors"].is_null());
assert!(payload["body_load_error_codes"].is_null());
}
}
}
#[tokio::test]
async fn gateway_admin_usage_detail_validates_body_field_selection() {
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_usage_reader_for_tests(Arc::new(
InMemoryUsageReadRepository::seed(vec![sample_selective_body_usage()]),
)));
for query in [
"body_field=unknown",
"body_field=",
"include_bodies=false&body_field=request_body",
"body_format=raw",
"body_format=unknown&body_field=request_body",
"body_format=&body_field=request_body",
] {
let response = local_admin_usage_response(
&state,
http::Method::GET,
&format!("/api/admin/usage/usage-selected-body?{query}"),
None,
)
.await;
assert_eq!(
response.status(),
StatusCode::BAD_REQUEST,
"invalid query: {query}"
);
}
}
#[tokio::test]
async fn gateway_admin_usage_detail_raw_reads_only_selected_body_and_is_not_cacheable() {
for detached in [false, true] {
let repository = if detached {
InMemoryUsageReadRepository::seed_with_detached_bodies(vec![
sample_selective_body_usage(),
])
} else {
InMemoryUsageReadRepository::seed(vec![sample_selective_body_usage()])
};
let state = AppState::new().unwrap().with_data_state_for_tests(
GatewayDataState::with_usage_reader_for_tests(Arc::new(repository)),
);
for field in [
"request_body",
"provider_request_body",
"response_body",
"client_response_body",
] {
let response = local_admin_usage_response(
&state,
http::Method::GET,
&format!("/api/admin/usage/usage-selected-body?body_field={field}&body_format=raw"),
None,
)
.await;
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response.headers()["cache-control"],
"no-store, no-transform"
);
assert_eq!(
response.headers()["x-aether-usage-id"],
"usage-selected-body"
);
assert_eq!(response.headers()["x-aether-body-field"], field);
assert_eq!(response.headers()["x-aether-body-encoding"], "json");
assert!(response.extensions().get::<AdminAuditEvent>().is_some());
let bytes = to_bytes(response.into_body(), 1024).await.unwrap();
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap(),
json!({ "marker": field })
);
}
}
}
#[tokio::test]
async fn gateway_admin_usage_detail_raw_rejects_foreign_refs_and_disabled_capture() {
for body_state in [
UsageBodyCaptureState::Reference,
UsageBodyCaptureState::Disabled,
] {
let mut usage = sample_selective_body_usage();
usage.response_body = None;
usage.response_body_ref = Some("usage://request/foreign-request/response_body".to_string());
usage.response_body_state = Some(body_state);
let mut foreign = sample_selective_body_usage();
foreign.id = "foreign-usage".to_string();
foreign.request_id = "foreign-request".to_string();
foreign.response_body = Some(json!({ "secret": "must not leak" }));
let state = AppState::new().unwrap().with_data_state_for_tests(
GatewayDataState::with_usage_reader_for_tests(Arc::new(
InMemoryUsageReadRepository::seed_with_detached_bodies(vec![usage, foreign]),
)),
);
let response = local_admin_usage_response(
&state,
http::Method::GET,
"/api/admin/usage/usage-selected-body?body_field=response_body&body_format=raw",
None,
)
.await;
assert_eq!(response.status(), StatusCode::NOT_FOUND);
assert_eq!(response.headers()["x-aether-body-error"], "missing");
let bytes = to_bytes(response.into_body(), 1024).await.unwrap();
assert!(!String::from_utf8_lossy(&bytes).contains("secret"));
}
}
#[tokio::test]
async fn gateway_admin_usage_detail_raw_preserves_authorization_and_binary_headers() {
let state = AppState::new().unwrap().with_data_state_for_tests(
GatewayDataState::with_usage_reader_for_tests(Arc::new(
InMemoryUsageReadRepository::seed_with_detached_bodies(vec![
sample_selective_body_usage(),
]),
)),
);
let gateway =
build_router_with_state(state).layer(tower_http::compression::CompressionLayer::new());
let (url, server) = start_server(gateway).await;
let client = reqwest::Client::new();
let endpoint = format!(
"{url}/api/admin/usage/usage-selected-body?body_field=response_body&body_format=raw"
);
let unauthorized = client.get(&endpoint).send().await.unwrap();
assert!(matches!(
unauthorized.status(),
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN
));
let response = admin_request(client.get(&endpoint))
.header("accept-encoding", "gzip")
.send()
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()["content-encoding"], "identity");
assert_eq!(response.headers()["x-aether-body-encoding"], "json");
assert_eq!(
response.headers()["cache-control"],
"no-store, no-transform"
);
let bytes = response.bytes().await.unwrap();
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap(),
json!({ "marker": "response_body" })
);
server.abort();
}
#[tokio::test]
async fn gateway_admin_usage_detail_isolates_missing_body_errors() {
let mut usage = sample_selective_body_usage();
usage.request_body = None;
usage.request_body_ref = Some("usage://request/req-selected-body/request_body".to_string());
usage.request_body_state = Some(UsageBodyCaptureState::Reference);
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_usage_reader_for_tests(Arc::new(
InMemoryUsageReadRepository::seed_with_detached_bodies(vec![usage]),
)));
for selected in ["response_body", "request_body"] {
let response = local_admin_usage_response(
&state,
http::Method::GET,
&format!("/api/admin/usage/usage-selected-body?body_field={selected}"),
None,
)
.await;
assert_eq!(response.status(), StatusCode::OK);
let bytes = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should collect");
let payload: serde_json::Value = serde_json::from_slice(&bytes).expect("json should parse");
if selected == "request_body" {
assert_eq!(payload["body_load_errors"]["request_body"], true);
assert_eq!(payload["body_load_error_codes"]["request_body"], "missing");
assert!(payload["request_body"].is_null());
} else {
assert_eq!(payload["response_body"]["marker"], "response_body");
assert!(payload["body_load_errors"].is_null());
assert!(payload["body_load_error_codes"].is_null());
}
}
}
#[tokio::test]
async fn gateway_resolves_admin_usage_detail_when_inline_state_has_body_ref() {
let (_upstream_url, upstream_hits, upstream_handle) =
@@ -221,13 +221,17 @@ async fn gateway_handles_admin_video_tasks_list_locally_with_trusted_admin_princ
assert_eq!(payload["pages"], json!(1));
assert_eq!(payload["items"].as_array().map(Vec::len), Some(1));
assert_eq!(payload["items"][0]["id"], "task-completed");
// Video-task persistence intentionally drops user-facing PII. The admin
// projection must therefore use the privacy-safe fallback when no separate
// user snapshot is joined.
assert_eq!(payload["items"][0]["username"], "Unknown");
assert_eq!(payload["items"][0]["username"], "alice");
assert_eq!(payload["items"][0]["provider_name"], "OpenAI");
assert_eq!(payload["items"][0]["status"], "completed");
assert!(payload["items"][0]["prompt"].is_null());
assert_eq!(
payload["items"][0]["prompt"],
format!("{}...", "x".repeat(100))
);
assert_eq!(
payload["items"][0]["video_url"],
"https://8.8.8.8/task-completed.mp4"
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
@@ -393,7 +397,9 @@ async fn gateway_handles_admin_video_task_detail_locally_with_trusted_admin_prin
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["id"], "task-detail");
assert_eq!(payload["username"], "Unknown");
assert_eq!(payload["prompt"], "detail prompt");
assert_eq!(payload["video_url"], "https://8.8.8.8/task-detail.mp4");
assert_eq!(payload["username"], "charlie");
assert_eq!(payload["provider_name"], "OpenAI");
assert_eq!(payload["endpoint"]["id"], "endpoint-1");
assert_eq!(payload["endpoint"]["api_format"], "openai:video");
@@ -734,7 +740,7 @@ async fn local_admin_video_task_cancel_attaches_explicit_audit() {
}
#[tokio::test]
async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstream() {
async fn gateway_redirects_persisted_openai_video_url_without_forwarding_admin_request() {
let upstream_hits = Arc::new(Mutex::new(0usize));
let upstream_hits_clone = Arc::clone(&upstream_hits);
let upstream = Router::new().route(
@@ -762,7 +768,10 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstrea
))
.await
.expect("task should upsert");
assert_eq!(stored.video_url, None);
assert_eq!(
stored.video_url.as_deref(),
Some("https://8.8.8.8/task-redirect.mp4")
);
let (_upstream_url, upstream_handle) = start_server(upstream).await;
let gateway = build_router_with_state(
@@ -788,7 +797,14 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstrea
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::NOT_FOUND);
assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT);
assert_eq!(
response
.headers()
.get(http::header::LOCATION)
.and_then(|value| value.to_str().ok()),
stored.video_url.as_deref()
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
@@ -796,22 +812,22 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstrea
}
#[tokio::test]
async fn local_admin_video_task_video_is_unavailable_after_openai_url_sanitization() {
async fn local_admin_video_task_download_preserves_signed_url_and_attaches_audit() {
let repository = Arc::new(InMemoryVideoTaskRepository::default());
let stored = repository
.upsert(sample_admin_video_task(
"task-video-audit",
VideoTaskStatus::Completed,
1_710_000_550,
"user-5",
"frank",
"provider-openai",
"gpt-video",
"video audit prompt",
))
.await
.expect("task should upsert");
assert_eq!(stored.video_url, None);
let mut task = sample_admin_video_task(
"task-video-audit",
VideoTaskStatus::Completed,
1_710_000_550,
"user-5",
"frank",
"provider-openai",
"gpt-video",
"video audit prompt",
);
task.video_url =
Some("https://8.8.8.8/video.mp4?signature=a%2Fb%2Bc%3D&part=2&part=1".to_string());
let stored = repository.upsert(task).await.expect("task should upsert");
assert_eq!(stored.prompt.as_deref(), Some("video audit prompt"));
let state = AppState::new()
.expect("gateway state should build")
@@ -825,8 +841,15 @@ async fn local_admin_video_task_video_is_unavailable_after_openai_url_sanitizati
)
.await;
assert_eq!(response.status(), StatusCode::NOT_FOUND);
assert!(response.extensions().get::<AdminAuditEvent>().is_none());
assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT);
assert_eq!(
response
.headers()
.get(http::header::LOCATION)
.and_then(|value| value.to_str().ok()),
stored.video_url.as_deref()
);
assert!(response.extensions().get::<AdminAuditEvent>().is_some());
}
#[tokio::test]
@@ -50,6 +50,8 @@ use chrono::{TimeZone, Utc};
const TEST_EMAIL_VERIFICATION_TOKEN: &str =
"test-email-verification-token-00000000000000000000000000000000";
#[path = "public_support/auth_cookie.rs"]
mod auth_cookie;
#[path = "public_support/dashboard.rs"]
mod dashboard;
#[path = "public_support/vscodex.rs"]
@@ -2657,9 +2659,9 @@ fn set_test_env_var(key: &'static str, value: &str) -> TestEnvVarGuard {
}
#[cfg(test)]
fn payment_callback_env_lock() -> &'static std::sync::Mutex<()> {
static LOCK: std::sync::OnceLock<std::sync::Mutex<()>> = std::sync::OnceLock::new();
LOCK.get_or_init(|| std::sync::Mutex::new(()))
fn payment_callback_env_lock() -> &'static tokio::sync::Mutex<()> {
static LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
&LOCK
}
const TEST_PAYMENT_CALLBACK_SECRET: &str = "test-callback-secret-0123456789abcdef";
@@ -11873,9 +11875,7 @@ async fn gateway_does_not_report_logout_success_when_session_revoke_is_rejected(
#[tokio::test]
async fn gateway_handles_payment_callback_route_locally_without_proxying_upstream() {
let _env_lock = payment_callback_env_lock()
.lock()
.expect("payment callback test env lock should not be poisoned");
let _env_lock = payment_callback_env_lock().lock().await;
let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", TEST_PAYMENT_CALLBACK_SECRET);
let now = Utc::now();
let user = StoredUserAuthRecord::new(
@@ -12018,9 +12018,7 @@ async fn gateway_handles_payment_callback_route_locally_without_proxying_upstrea
#[tokio::test]
async fn gateway_rejects_payment_callback_with_mismatched_payment_method_locally() {
let _env_lock = payment_callback_env_lock()
.lock()
.expect("payment callback test env lock should not be poisoned");
let _env_lock = payment_callback_env_lock().lock().await;
let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", TEST_PAYMENT_CALLBACK_SECRET);
let now = Utc::now();
let user = StoredUserAuthRecord::new(
@@ -0,0 +1,124 @@
use super::{sample_auth_user, sample_auth_wallet, start_auth_gateway_with_state};
use axum::http::{header, StatusCode};
use chrono::Utc;
use serde_json::json;
fn refresh_cookie(response: &reqwest::Response, secure: bool) -> String {
let cookie = response
.headers()
.get(header::SET_COOKIE)
.unwrap()
.to_str()
.unwrap();
assert!(cookie.starts_with("aether_refresh_token="));
assert!(cookie.contains("HttpOnly"));
assert!(cookie.contains("Path=/api/auth"));
assert_eq!(
cookie
.split(';')
.any(|attribute| attribute.trim() == "Secure"),
secure
);
if !secure {
assert!(!cookie.contains("SameSite=None"));
assert!(cookie.contains("SameSite=Lax"));
}
assert_eq!(
response.headers().get(header::CACHE_CONTROL).unwrap(),
"no-store"
);
cookie.to_string()
}
#[tokio::test]
async fn gateway_auth_refresh_cookie_roundtrip_adapts_to_http_and_https() {
for (origin_scheme, forwarded_proto, secure) in [
("http", None, false),
("https", None, true),
("https", Some("https"), true),
("http", Some("https"), true),
] {
let now = Utc::now();
let (gateway_url, upstream_hits, gateway_handle, upstream_handle) =
start_auth_gateway_with_state(
sample_auth_user(now),
sample_auth_wallet("user-auth-1", now),
[],
)
.await;
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
header::ORIGIN,
gateway_url
.replacen("http:", &format!("{origin_scheme}:"), 1)
.parse()
.unwrap(),
);
headers.insert(
"x-client-device-id",
"cookie-roundtrip-device".parse().unwrap(),
);
if let Some(proto) = forwarded_proto {
headers.insert("x-forwarded-proto", proto.parse().unwrap());
}
let client = reqwest::Client::builder()
.default_headers(headers)
.build()
.unwrap();
let login = client.post(format!("{gateway_url}/api/auth/login"))
.json(&json!({ "email": "[email protected]", "password": "secret123", "auth_type": "local" }))
.send().await.unwrap();
assert_eq!(login.status(), StatusCode::OK);
let mut cookie = refresh_cookie(&login, secure);
for _ in 0..3 {
let refreshed = client
.post(format!("{gateway_url}/api/auth/refresh"))
.header(header::COOKIE, cookie.split(';').next().unwrap())
.send()
.await
.unwrap();
assert_eq!(refreshed.status(), StatusCode::OK);
let rotated = refresh_cookie(&refreshed, secure);
assert_ne!(rotated, cookie);
cookie = rotated;
let payload: serde_json::Value = refreshed.json().await.unwrap();
let current_user = client
.get(format!("{gateway_url}/api/auth/me"))
.bearer_auth(payload["access_token"].as_str().unwrap())
.send()
.await
.unwrap();
assert_eq!(current_user.status(), StatusCode::OK);
}
let logout = client
.post(format!("{gateway_url}/api/auth/logout"))
.header(header::COOKIE, cookie.split(';').next().unwrap())
.send()
.await
.unwrap();
assert_eq!(logout.status(), StatusCode::OK);
assert!(refresh_cookie(&logout, secure).contains("Max-Age=0"));
let revoked = client
.post(format!("{gateway_url}/api/auth/refresh"))
.header(header::COOKIE, cookie.split(';').next().unwrap())
.send()
.await
.unwrap();
assert_eq!(revoked.status(), StatusCode::UNAUTHORIZED);
assert!(refresh_cookie(&revoked, secure).contains("Max-Age=0"));
let missing = client
.post(format!("{gateway_url}/api/auth/refresh"))
.send()
.await
.unwrap();
assert_eq!(missing.status(), StatusCode::UNAUTHORIZED);
assert!(refresh_cookie(&missing, secure).contains("Max-Age=0"));
assert_eq!(*upstream_hits.lock().unwrap(), 0);
gateway_handle.abort();
upstream_handle.abort();
}
}
+1 -1
View File
@@ -1076,7 +1076,7 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider_with_request_timeout(
"provider-owner",
Some(0.1),
Some(2.0),
)],
vec![sample_endpoint("endpoint-owner", "provider-owner")],
vec![sample_bound_key(
+241 -43
View File
@@ -12,6 +12,7 @@ use super::{
UsageReadRepository, UsageRuntimeConfig, DEVELOPMENT_ENCRYPTION_KEY, TRACE_ID_HEADER,
};
use crate::constants::LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER;
use aether_data_contracts::repository::usage::UsageBodyCaptureState;
fn deep_nested_metadata(levels: usize) -> serde_json::Value {
let mut current = json!({"leaf": "value"});
@@ -84,6 +85,58 @@ where
stored.expect("usage should be present once the expected status is observed")
}
async fn load_admin_usage_capture_detail(
state: &crate::AppState,
usage_id: &str,
include_bodies: bool,
) -> serde_json::Value {
use crate::admin_api::{maybe_build_local_admin_response, AdminRouteRequest};
use crate::constants::{
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
TRUSTED_ADMIN_USER_ROLE_HEADER,
};
use crate::control::resolve_public_request_context;
use http_body_util::BodyExt;
let mut headers = http::HeaderMap::new();
for (name, value) in [
(GATEWAY_HEADER, "rust-phase3b"),
(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user"),
(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin"),
(TRUSTED_ADMIN_SESSION_ID_HEADER, "admin-session"),
] {
headers.insert(name, HeaderValue::from_static(value));
}
let uri = format!("/api/admin/usage/{usage_id}?include_bodies={include_bodies}")
.parse()
.unwrap();
let context = resolve_public_request_context(
state,
&http::Method::GET,
&uri,
&headers,
"usage-full-detail",
)
.await
.unwrap();
let response = maybe_build_local_admin_response(AdminRouteRequest::new(
state,
&context,
&"127.0.0.1:12345".parse().unwrap(),
&headers,
None,
))
.await
.unwrap()
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(response
.extensions()
.get::<crate::audit::AdminAuditEvent>()
.is_some());
serde_json::from_slice(&response.into_body().collect().await.unwrap().to_bytes()).unwrap()
}
#[test]
fn gateway_handles_local_openai_chat_sync_report_with_local_reporting_when_usage_runtime_enabled() {
run_async_test_on_large_stack(
@@ -348,7 +401,7 @@ async fn gateway_truncates_deep_request_echo_for_local_openai_chat_sync_usage_im
Arc::clone(&request_candidate_repository),
Arc::clone(&usage_repository),
DEVELOPMENT_ENCRYPTION_KEY,
),
).with_system_config_values_for_tests([("request_record_level".to_string(), json!("full"))]),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
@@ -402,10 +455,28 @@ async fn gateway_truncates_deep_request_echo_for_local_openai_chat_sync_usage_im
let stored_usage = stored_usage.expect("usage should be recorded");
assert_eq!(stored_usage.status, "completed");
assert_eq!(stored_usage.total_tokens, 5);
assert!(stored_usage.request_body.is_none());
let request_body = stored_usage.request_body.as_ref().unwrap();
assert_eq!(
request_body["messages"][0]["content"]
.as_str()
.unwrap()
.len(),
128 * 1024
);
assert!(
request_body["metadata"]["child"]["child"]["child"]["child"]["child"]
.get("depth")
.is_some()
);
assert!(stored_usage.request_body_ref.is_none());
assert!(stored_usage.request_body_state.is_none());
assert!(stored_usage.request_headers.is_none());
assert_eq!(
stored_usage.request_body_state,
Some(UsageBodyCaptureState::Inline)
);
assert_eq!(
stored_usage.request_headers.as_ref().unwrap()["authorization"],
"Bearer sk-client-openai-local-report-sync-deep"
);
gateway_handle.abort();
execution_runtime_handle.abort();
@@ -489,10 +560,10 @@ async fn gateway_ignores_legacy_max_request_body_size_for_local_openai_chat_sync
Arc::clone(&usage_repository),
DEVELOPMENT_ENCRYPTION_KEY,
)
.with_system_config_values_for_tests([(
"max_request_body_size".to_string(),
json!(128),
)]),
.with_system_config_values_for_tests([
("max_request_body_size".to_string(), json!(128)),
("request_record_level".to_string(), json!("full")),
]),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
@@ -535,12 +606,30 @@ async fn gateway_ignores_legacy_max_request_body_size_for_local_openai_chat_sync
)
.await;
assert_eq!(stored_usage.total_tokens, 5);
assert!(stored_usage.request_body.is_none());
assert!(
stored_usage.request_body.as_ref().unwrap()["messages"][0]["content"]
.as_str()
.unwrap()
.len()
> 128
);
assert!(stored_usage.request_body_ref.is_none());
assert!(stored_usage.request_body_state.is_none());
assert!(stored_usage.provider_request_body.is_none());
assert_eq!(
stored_usage.request_body_state,
Some(UsageBodyCaptureState::Inline)
);
assert!(
stored_usage.provider_request_body.as_ref().unwrap()["messages"][0]["content"]
.as_str()
.unwrap()
.len()
> 128
);
assert!(stored_usage.provider_request_body_ref.is_none());
assert!(stored_usage.provider_request_body_state.is_none());
assert_eq!(
stored_usage.provider_request_body_state,
Some(UsageBodyCaptureState::Inline)
);
gateway_handle.abort();
execution_runtime_handle.abort();
@@ -551,11 +640,19 @@ async fn gateway_ignores_legacy_max_request_body_size_for_local_openai_chat_sync
fn gateway_strips_request_and_response_bodies_when_request_record_level_is_base() {
run_async_test_on_large_stack(
"gateway_strips_request_and_response_bodies_when_request_record_level_is_base",
gateway_strips_request_and_response_bodies_when_request_record_level_is_base_impl(),
gateway_honors_request_record_level_impl("base"),
);
}
async fn gateway_strips_request_and_response_bodies_when_request_record_level_is_base_impl() {
#[test]
fn gateway_full_request_record_level_preserves_sync_bodies_in_admin_detail() {
run_async_test_on_large_stack(
"gateway_full_request_record_level_preserves_sync_bodies_in_admin_detail",
gateway_honors_request_record_level_impl("full"),
);
}
async fn gateway_honors_request_record_level_impl(record_level: &str) {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
@@ -630,14 +727,14 @@ async fn gateway_strips_request_and_response_bodies_when_request_record_level_is
)
.with_system_config_values_for_tests([(
"request_record_level".to_string(),
json!("base"),
json!(record_level),
)]),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
..UsageRuntimeConfig::default()
});
let gateway = build_router_with_state(gateway_state);
let gateway = build_router_with_state(gateway_state.clone());
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
@@ -678,14 +775,39 @@ async fn gateway_strips_request_and_response_bodies_when_request_record_level_is
assert_eq!(stored_usage.status, "completed");
assert_eq!(stored_usage.total_tokens, 5);
assert_eq!(stored_usage.response_time_ms, Some(25));
assert!(stored_usage.request_body.is_none());
assert!(stored_usage.request_body_ref.is_none());
assert!(stored_usage.provider_request_body.is_none());
assert!(stored_usage.provider_request_body_ref.is_none());
assert!(stored_usage.response_body.is_none());
assert!(stored_usage.response_body_ref.is_none());
assert!(stored_usage.client_response_body.is_none());
assert!(stored_usage.client_response_body_ref.is_none());
let detail = load_admin_usage_capture_detail(&gateway_state, &stored_usage.id, true).await;
let shallow = load_admin_usage_capture_detail(&gateway_state, &stored_usage.id, false).await;
for field in [
"request_body",
"provider_request_body",
"response_body",
"client_response_body",
] {
assert!(shallow[field].is_null());
let expected_captured = record_level == "full" && field != "client_response_body";
assert_eq!(
shallow[format!("has_{field}")],
expected_captured,
"availability for {field}"
);
if expected_captured {
assert!(!detail[field].is_null(), "full should expose {field}");
} else {
assert!(
detail[field].is_null(),
"uncaptured {field} must remain absent"
);
}
}
if record_level == "full" {
assert_eq!(
detail["request_body"]["messages"][0]["content"],
"request body should not be persisted"
);
assert_eq!(detail["provider_request_body"]["model"], "gpt-5-upstream");
assert_eq!(detail["response_body"], body_json);
assert!(detail["client_response_body"].is_null());
}
let stored_candidates = request_candidate_repository
.list_by_request_id("trace-openai-chat-local-report-sync-base-123")
@@ -825,10 +947,16 @@ async fn gateway_records_failed_usage_when_all_local_openai_chat_candidates_exha
);
assert!(stored_usage.response_body.is_none());
assert!(stored_usage.response_body_ref.is_none());
assert!(stored_usage.response_body_state.is_none());
assert_eq!(
stored_usage.response_body_state,
Some(UsageBodyCaptureState::Disabled)
);
assert!(stored_usage.client_response_body.is_none());
assert!(stored_usage.client_response_body_ref.is_none());
assert!(stored_usage.client_response_body_state.is_none());
assert_eq!(
stored_usage.client_response_body_state,
Some(UsageBodyCaptureState::Disabled)
);
let stored_candidates = request_candidate_repository
.list_by_request_id("trace-openai-chat-local-report-sync-failure-123")
@@ -930,7 +1058,10 @@ async fn gateway_records_failed_usage_when_sync_runtime_transport_is_unavailable
assert_eq!(stored_usage.status_code, Some(503));
assert!(stored_usage.response_body.is_none());
assert!(stored_usage.response_body_ref.is_none());
assert!(stored_usage.response_body_state.is_none());
assert_eq!(
stored_usage.response_body_state,
Some(UsageBodyCaptureState::Disabled)
);
let stored_candidates = request_candidate_repository
.list_by_request_id("trace-openai-chat-local-transport-unavailable-123")
@@ -969,7 +1100,7 @@ async fn sync_transport_error_policy_stops_or_retries_candidates_end_to_end_impl
let mut second_candidate = sample_local_openai_candidate_row();
second_candidate.key_id = "key-openai-usage-local-2".to_string();
second_candidate.key_name = "secondary".to_string();
second_candidate.key_internal_priority = second_candidate.key_internal_priority - 1;
second_candidate.key_internal_priority -= 1;
let candidate_selection_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_local_openai_candidate_row(),
@@ -1272,7 +1403,10 @@ async fn gateway_records_failed_usage_for_claude_runtime_miss_without_execution_
);
assert!(stored_usage.client_response_body.is_none());
assert!(stored_usage.client_response_body_ref.is_none());
assert!(stored_usage.client_response_body_state.is_none());
assert_eq!(
stored_usage.client_response_body_state,
Some(UsageBodyCaptureState::Disabled)
);
assert!(stored_usage.error_message.is_none());
let stored_candidates = request_candidate_repository
@@ -1296,11 +1430,20 @@ fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usa
{
run_async_test_on_large_stack(
"gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled",
gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl(),
gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl("basic"),
);
}
#[test]
fn gateway_full_request_record_level_preserves_stream_bodies_in_admin_detail() {
run_async_test_on_large_stack(
"gateway_full_request_record_level_preserves_stream_bodies_in_admin_detail",
gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl("full"),
);
}
async fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl(
record_level: &str,
) {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
@@ -1406,13 +1549,13 @@ async fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_wh
Arc::clone(&request_candidate_repository),
Arc::clone(&usage_repository),
DEVELOPMENT_ENCRYPTION_KEY,
),
).with_system_config_values_for_tests([("request_record_level".to_string(), json!(record_level))]),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
..UsageRuntimeConfig::default()
});
let gateway = build_router_with_state(gateway_state);
let gateway = build_router_with_state(gateway_state.clone());
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
@@ -1448,6 +1591,30 @@ async fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_wh
assert!(stored_usage.response_time_ms >= stored_usage.first_byte_time_ms);
assert!(stored_usage.is_stream);
let detail = load_admin_usage_capture_detail(&gateway_state, &stored_usage.id, true).await;
for field in [
"request_body",
"provider_request_body",
"response_body",
"client_response_body",
] {
if record_level == "full" {
assert!(
!detail[field].is_null(),
"full stream should expose {field}"
);
} else {
assert!(
detail[field].is_null(),
"basic stream must not persist {field}"
);
}
}
if record_level == "full" {
assert!(detail["response_body"].to_string().contains("hello"));
assert!(detail["client_response_body"].to_string().contains("hello"));
}
let stored_candidates = request_candidate_repository
.list_by_request_id("trace-openai-chat-local-report-stream-123")
.await
@@ -1585,10 +1752,10 @@ async fn gateway_ignores_legacy_max_response_body_size_for_stream_usage_impl() {
Arc::clone(&usage_repository),
DEVELOPMENT_ENCRYPTION_KEY,
)
.with_system_config_values_for_tests([(
"max_response_body_size".to_string(),
json!(128),
)]),
.with_system_config_values_for_tests([
("max_response_body_size".to_string(), json!(128)),
("request_record_level".to_string(), json!("full")),
]),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
@@ -1624,12 +1791,34 @@ async fn gateway_ignores_legacy_max_response_body_size_for_stream_usage_impl() {
)
.await;
assert_eq!(stored_usage.total_tokens, 6);
assert!(stored_usage.response_body.is_none());
assert!(
stored_usage
.response_body
.as_ref()
.unwrap()
.to_string()
.len()
> 128
);
assert!(stored_usage.response_body_ref.is_none());
assert!(stored_usage.response_body_state.is_none());
assert!(stored_usage.client_response_body.is_none());
assert_eq!(
stored_usage.response_body_state,
Some(UsageBodyCaptureState::Inline)
);
assert!(
stored_usage
.client_response_body
.as_ref()
.unwrap()
.to_string()
.len()
> 128
);
assert!(stored_usage.client_response_body_ref.is_none());
assert!(stored_usage.client_response_body_state.is_none());
assert_eq!(
stored_usage.client_response_body_state,
Some(UsageBodyCaptureState::Inline)
);
gateway_handle.abort();
execution_runtime_handle.abort();
@@ -1903,10 +2092,16 @@ async fn gateway_records_failed_usage_when_all_local_claude_cli_candidates_are_s
Some("all_candidates_skipped")
);
assert!(stored_usage.error_message.is_none());
assert!(stored_usage.request_headers.is_none());
assert_eq!(
stored_usage.request_headers.as_ref().unwrap()["authorization"],
"Bearer sk-client-claude-cli-usage-local-miss"
);
assert!(stored_usage.request_body.is_none());
assert!(stored_usage.request_body_ref.is_none());
assert!(stored_usage.request_body_state.is_none());
assert_eq!(
stored_usage.request_body_state,
Some(UsageBodyCaptureState::Disabled)
);
assert!(stored_usage.provider_request_body.is_none());
assert_eq!(
stored_usage
@@ -2169,7 +2364,10 @@ fn gateway_keeps_failed_usage_request_capture_lightweight_for_large_local_claude
)
.await;
assert_eq!(stored_usage.status, "failed");
assert!(stored_usage.request_body_state.is_none());
assert_eq!(
stored_usage.request_body_state,
Some(UsageBodyCaptureState::Disabled)
);
assert!(stored_usage.request_body.is_none());
assert!(stored_usage.request_body_ref.is_none());
assert!(stored_usage.provider_request_body.is_none());
@@ -213,6 +213,13 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
};
assert_eq!(stored.status, VideoTaskStatus::Processing);
assert_eq!(stored.prompt.as_deref(), Some("hello"));
assert_eq!(stored.username.as_deref(), Some("video-user"));
assert_eq!(stored.api_key_name.as_deref(), Some("video-key"));
assert_eq!(stored.duration_seconds, Some(4));
assert_eq!(stored.resolution.as_deref(), Some("720p"));
assert_eq!(stored.aspect_ratio.as_deref(), Some("16:9"));
assert_eq!(stored.size.as_deref(), Some("1280x720"));
assert_eq!(stored.progress_percent, 37);
assert_eq!(stored.poll_count, 1);
assert!(
@@ -32,6 +32,7 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
struct SeenExecutionRuntimeStreamRequest {
method: String,
url: String,
headers: serde_json::Value,
}
fn hash_api_key(value: &str) -> String {
@@ -159,6 +160,7 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
.and_then(|value| value.as_str())
.unwrap_or_default()
.to_string(),
headers: payload.get("headers").cloned().unwrap_or_else(|| json!({})),
});
let frames = [
@@ -252,7 +254,10 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
updated_at_unix_secs: 456,
error_code: None,
error_message: None,
video_url: Some("https://cdn.example.com/video-content.mp4".to_string()),
video_url: Some(
"https://cdn.example.com/video-content.mp4?signature=a%2Fb%2Bc%3D&part=2&part=1"
.to_string(),
),
request_metadata: None,
})
.await
@@ -358,8 +363,9 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
assert_eq!(seen_stream_request.method, "GET");
assert_eq!(
seen_stream_request.url,
"https://api.openai.example/v1/videos/ext-video-content-followup-123/content"
"https://cdn.example.com/video-content.mp4?signature=a%2Fb%2Bc%3D&part=2&part=1"
);
assert!(seen_stream_request.headers.get("authorization").is_none());
assert_eq!(*decision_stream_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*execute_stream_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
@@ -0,0 +1,186 @@
use std::collections::VecDeque;
use std::sync::Arc;
use bytes::{Bytes, BytesMut};
use parking_lot::Mutex;
use tokio::sync::Notify;
const CHUNK_BYTES: usize = 32 * 1024;
#[derive(Debug)]
pub enum LocalBodyEvent {
Chunk(Bytes),
End,
Error(String),
}
#[derive(Default)]
struct BufferState {
chunks: VecDeque<BytesMut>,
bytes: usize,
terminal: Option<Result<(), String>>,
receiver_taken: bool,
receiver_closed: bool,
}
pub(super) struct ResponseBuffer {
state: Mutex<BufferState>,
notify: Notify,
capacity: usize,
}
impl ResponseBuffer {
pub(super) fn new(capacity: usize) -> Arc<Self> {
Arc::new(Self {
state: Mutex::new(BufferState::default()),
notify: Notify::new(),
capacity,
})
}
pub(super) fn take_receiver(self: &Arc<Self>) -> Option<BodyReceiver> {
let mut state = self.state.lock();
if state.receiver_taken {
return None;
}
state.receiver_taken = true;
Some(BodyReceiver {
buffer: Arc::clone(self),
finished: false,
})
}
pub(super) fn push(&self, mut payload: Bytes) -> bool {
let mut state = self.state.lock();
if state.terminal.is_some()
|| state.receiver_closed
|| payload.len() > self.capacity.saturating_sub(state.bytes)
{
return false;
}
state.bytes += payload.len();
while !payload.is_empty() {
if let Some(tail) = state
.chunks
.back_mut()
.filter(|chunk| chunk.len() < CHUNK_BYTES)
{
let count = payload.len().min(CHUNK_BYTES - tail.len());
tail.extend_from_slice(&payload.split_to(count));
} else {
let count = payload.len().min(CHUNK_BYTES);
let chunk = payload.split_to(count);
state.chunks.push_back(
chunk
.try_into_mut()
.unwrap_or_else(|chunk| BytesMut::from(chunk.as_ref())),
);
}
}
drop(state);
self.notify.notify_waiters();
true
}
pub(super) fn finish(&self, result: Result<(), String>) {
let mut state = self.state.lock();
if state.terminal.is_none() {
state.terminal = Some(result);
}
drop(state);
self.notify.notify_waiters();
}
}
pub(super) struct BodyReceiver {
buffer: Arc<ResponseBuffer>,
finished: bool,
}
impl BodyReceiver {
pub(super) async fn recv(&mut self) -> Option<LocalBodyEvent> {
if self.finished {
return None;
}
loop {
let notified = self.buffer.notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
{
let mut state = self.buffer.state.lock();
if let Some(chunk) = state.chunks.pop_front() {
state.bytes -= chunk.len();
return Some(LocalBodyEvent::Chunk(chunk.freeze()));
}
if let Some(terminal) = state.terminal.take() {
self.finished = true;
state.receiver_closed = true;
return Some(match terminal {
Ok(()) => LocalBodyEvent::End,
Err(error) => LocalBodyEvent::Error(error),
});
}
}
notified.await;
}
}
}
impl Drop for BodyReceiver {
fn drop(&mut self) {
let mut state = self.buffer.state.lock();
state.receiver_closed = true;
state.chunks.clear();
state.bytes = 0;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn error_survives_a_full_buffer() {
let buffer = ResponseBuffer::new(CHUNK_BYTES);
let mut receiver = buffer.take_receiver().unwrap();
assert!(buffer.push(Bytes::from(vec![b'x'; CHUNK_BYTES])));
buffer.finish(Err("proxy disconnected".into()));
assert!(matches!(
receiver.recv().await,
Some(LocalBodyEvent::Chunk(_))
));
assert!(
matches!(receiver.recv().await, Some(LocalBodyEvent::Error(error)) if error == "proxy disconnected")
);
assert!(receiver.recv().await.is_none());
}
#[tokio::test]
async fn small_frames_are_coalesced_within_the_byte_budget() {
let buffer = ResponseBuffer::new(4096);
let mut receiver = buffer.take_receiver().unwrap();
for _ in 0..4096 {
assert!(buffer.push(Bytes::from_static(b"x")));
}
assert!(!buffer.push(Bytes::from_static(b"x")));
assert_eq!(buffer.state.lock().chunks.len(), 1);
buffer.finish(Ok(()));
assert!(
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 4096)
);
assert!(matches!(receiver.recv().await, Some(LocalBodyEvent::End)));
}
#[tokio::test]
async fn terminal_wakes_an_empty_receiver_and_is_not_overwritten() {
let buffer = ResponseBuffer::new(1024);
let mut receiver = buffer.take_receiver().unwrap();
let task = tokio::spawn(async move { receiver.recv().await });
tokio::task::yield_now().await;
buffer.finish(Err("cancelled".into()));
buffer.finish(Ok(()));
assert!(
matches!(task.await.unwrap(), Some(LocalBodyEvent::Error(error)) if error == "cancelled")
);
}
}
@@ -0,0 +1,250 @@
use super::*;
async fn fixture(
window: u32,
capacity: usize,
) -> (
Arc<HubRouter>,
Arc<ProxyConn>,
Arc<LocalStream>,
aether_runtime::BoundedQueueReceiver<Message>,
) {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (sender, receiver) = bounded_queue(capacity);
let (close_tx, _) = watch::channel(false);
let connection = Arc::new(
ProxyConn::new(
99,
"flow-test".into(),
"flow-test".into(),
sender,
close_tx,
16,
3,
)
.with_settings(protocol::SettingsPayload {
initial_stream_window_bytes: window,
min_window_update_bytes: (window / 4).max(1),
drain_deadline_ms: 1000,
}),
);
hub.register_proxy(Arc::clone(&connection));
let stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
(hub, connection, stream, receiver)
}
fn meta() -> protocol::RequestMeta {
protocol::RequestMeta {
provider_id: None,
endpoint_id: None,
key_id: None,
method: "GET".into(),
url: "https://example.com".into(),
headers: HashMap::new(),
stream: true,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: 30,
follow_redirects: None,
http1_only: false,
transport_profile: None,
}
}
async fn headers(hub: &Arc<HubRouter>, stream: &LocalStream) {
let payload = serde_json::to_vec(&protocol::ResponseMeta {
status: 200,
headers: vec![],
})
.unwrap();
let mut frame = protocol::encode_frame(
stream.proxy_stream_id,
protocol::RESPONSE_HEADERS,
0,
&payload,
);
hub.handle_proxy_frame(stream.proxy_conn_id, &mut frame)
.await;
}
#[tokio::test]
async fn window_credit_is_retried_after_queue_pressure_and_cancelled_receive() {
let (hub, _, stream, mut outbound) = fixture(128, 1).await;
headers(&hub, &stream).await;
assert!(stream.push_body_chunk(Bytes::from(vec![b'x'; 64])));
let mut receiver = stream.take_body_receiver().unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(10), receiver.recv())
.await
.is_err()
);
assert_eq!(*stream.response_consumed_since_update.lock(), 64);
outbound.recv().await.unwrap();
let event = tokio::time::timeout(Duration::from_secs(1), receiver.recv())
.await
.unwrap();
assert!(matches!(event, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 64));
assert_eq!(*stream.response_consumed_since_update.lock(), 0);
let Message::Binary(data) = outbound.recv().await.unwrap() else {
panic!("expected binary update")
};
let frame = aether_contracts::tunnel::Frame::decode(data).unwrap();
let update: protocol::WindowUpdatePayload = serde_json::from_slice(&frame.payload).unwrap();
assert_eq!(
frame.msg_type,
aether_contracts::tunnel::MsgType::WindowUpdate
);
assert_eq!(update.delta_bytes, 64);
hub.cancel_local_stream(stream.id, "test complete");
}
#[tokio::test]
async fn response_credit_is_not_returned_until_consumed() {
let (hub, _, stream, mut outbound) = fixture(128, 4).await;
outbound.recv().await.unwrap();
let mut body = protocol::encode_frame(
stream.proxy_stream_id,
protocol::RESPONSE_BODY,
0,
&[b'x'; 128],
);
hub.handle_proxy_frame(99, &mut body).await;
assert!(outbound.try_recv().is_err());
let mut receiver = stream.take_body_receiver().unwrap();
assert!(
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 128)
);
assert!(outbound.try_recv().is_ok());
hub.cancel_local_stream(stream.id, "test complete");
}
#[tokio::test]
async fn cancelled_stream_open_releases_slot_without_resetting_connection() {
let (hub, connection, first_stream, mut outbound) = fixture(128, 1).await;
let opening_hub = Arc::clone(&hub);
let opening =
tokio::spawn(async move { opening_hub.open_local_stream("flow-test", &meta()).await });
tokio::time::timeout(Duration::from_secs(1), async {
while hub.local_streams.len() != 2 {
tokio::task::yield_now().await;
}
})
.await
.unwrap();
opening.abort();
assert!(matches!(opening.await, Err(error) if error.is_cancelled()));
assert_eq!(connection.stream_count.load(Ordering::Relaxed), 1);
assert_eq!(hub.local_streams.len(), 1);
assert_eq!(hub.proxy_to_local.len(), 1);
assert!(connection.is_available());
outbound.recv().await.unwrap();
assert!(outbound.try_recv().is_err());
let next_stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
outbound.recv().await.unwrap();
hub.cancel_local_stream(first_stream.id, "test complete");
outbound.recv().await.unwrap();
hub.cancel_local_stream(next_stream.id, "test complete");
}
#[tokio::test]
async fn full_response_buffer_preserves_disconnect_error() {
let (hub, connection, stream, mut outbound) = fixture(4 * 1024 * 1024, 512).await;
outbound.recv().await.unwrap();
headers(&hub, &stream).await;
let mut receiver = stream.take_body_receiver().unwrap();
for _ in 0..128 {
let mut frame = protocol::encode_frame(
stream.proxy_stream_id,
protocol::RESPONSE_BODY,
0,
&vec![b'x'; 32 * 1024],
);
hub.handle_proxy_frame(99, &mut frame).await;
}
hub.unregister_proxy(connection.id, &connection.node_id);
let mut bytes = 0;
loop {
match receiver.recv().await {
Some(LocalBodyEvent::Chunk(chunk)) => bytes += chunk.len(),
Some(LocalBodyEvent::Error(error)) => {
assert!(error.contains("disconnected"));
break;
}
event => panic!("disconnect must not become normal EOF: {event:?}"),
}
}
assert_eq!(bytes, 4 * 1024 * 1024);
}
#[tokio::test]
async fn slow_stream_does_not_block_another_stream_on_the_same_connection() {
let (hub, _, slow, mut outbound) = fixture(128, 512).await;
outbound.recv().await.unwrap();
let fast = hub.open_local_stream("flow-test", &meta()).await.unwrap();
assert!(slow.push_body_chunk(Bytes::from(vec![b'x'; 128])));
let mut overflowing = protocol::encode_frame(
slow.proxy_stream_id,
protocol::RESPONSE_BODY,
0,
b"overflow",
);
tokio::time::timeout(Duration::from_secs(1), async {
hub.handle_proxy_frame(99, &mut overflowing).await;
headers(&hub, &fast).await;
assert_eq!(
fast.wait_headers(Duration::from_secs(1))
.await
.unwrap()
.status,
200
);
})
.await
.expect("slow stream must not block connection reader");
assert!(!hub.local_streams.contains_key(&slow.id));
assert!(hub.local_streams.contains_key(&fast.id));
hub.cancel_local_stream(fast.id, "test complete");
}
#[tokio::test]
async fn cancelling_a_stream_wakes_request_window_waiters() {
let (_, _, stream, _) = fixture(128, 512).await;
*stream.request_window.available.lock() = 0;
let waiter = tokio::spawn({
let stream = Arc::clone(&stream);
async move {
stream
.acquire_request_window(1, Duration::from_secs(30))
.await
}
});
tokio::task::yield_now().await;
stream.fail("cancelled");
assert!(tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.unwrap()
.unwrap()
.is_err());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn concurrent_headers_and_credit_updates_do_not_lose_notifications() {
for index in 0..256 {
let stream = Arc::new(LocalStream::new(index, "test".into(), 1, 1, 1));
let window = Arc::new(StreamFlowWindow::new(0));
let waiter = tokio::spawn({
let stream = Arc::clone(&stream);
let window = Arc::clone(&window);
async move {
stream.wait_headers(Duration::from_secs(1)).await.unwrap();
window.acquire(1, Duration::from_secs(1)).await.unwrap();
}
});
stream.set_response_headers(protocol::ResponseMeta {
status: 200,
headers: vec![],
});
window.add(1);
waiter.await.unwrap();
}
}
+203 -83
View File
@@ -12,10 +12,11 @@ use axum::extract::ws::Message;
use bytes::Bytes;
use dashmap::DashMap;
use parking_lot::{Mutex, RwLock};
use tokio::sync::mpsc;
use tokio::sync::{watch, Notify};
use tracing::{debug, info, warn};
pub use super::body::LocalBodyEvent;
use super::body::{BodyReceiver, ResponseBuffer};
use super::control_plane::ControlPlaneClient;
use super::protocol;
@@ -29,6 +30,10 @@ const DEFAULT_DRAIN_DEADLINE_MS: u64 = 30_000;
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
#[cfg(test)]
#[path = "flow_control_tests.rs"]
mod flow_control_tests;
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
.ok()
@@ -53,11 +58,16 @@ static NODE_STATUS_QUEUE_CAPACITY: LazyLock<usize> = LazyLock::new(|| {
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
});
static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock<u32> = LazyLock::new(|| {
STREAM_INITIAL_WINDOW_BYTES
.saturating_div(4)
.clamp(1, 1024 * 1024)
});
pub(super) fn local_settings() -> protocol::SettingsPayload {
protocol::SettingsPayload {
initial_stream_window_bytes: (*STREAM_INITIAL_WINDOW_BYTES)
.min(aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u32),
min_window_update_bytes: STREAM_INITIAL_WINDOW_BYTES
.saturating_div(4)
.clamp(1, 1024 * 1024),
drain_deadline_ms: *DRAIN_DEADLINE_MS,
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SendStatus {
@@ -91,6 +101,7 @@ impl ConnHealthState {
struct StreamFlowWindow {
available: Mutex<u64>,
notify: Notify,
closed: AtomicBool,
}
impl StreamFlowWindow {
@@ -98,6 +109,7 @@ impl StreamFlowWindow {
Self {
available: Mutex::new(u64::from(initial)),
notify: Notify::new(),
closed: AtomicBool::new(false),
}
}
@@ -109,6 +121,12 @@ impl StreamFlowWindow {
let requested = bytes as u64;
let started_at = Instant::now();
loop {
let notified = self.notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.closed.load(Ordering::Acquire) {
return Err(());
}
{
let mut available = self.available.lock();
if *available >= requested {
@@ -120,10 +138,7 @@ impl StreamFlowWindow {
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
return Err(());
};
if tokio::time::timeout(remaining, self.notify.notified())
.await
.is_err()
{
if tokio::time::timeout(remaining, notified).await.is_err() {
return Err(());
}
}
@@ -138,6 +153,11 @@ impl StreamFlowWindow {
drop(available);
self.notify.notify_waiters();
}
fn close(&self) {
self.closed.store(true, Ordering::Release);
self.notify.notify_waiters();
}
}
#[derive(Debug, Clone, Copy)]
@@ -209,6 +229,10 @@ impl BoundedOutbound {
pub fn snapshot(&self) -> QueueSnapshot {
self.tx.snapshot()
}
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
self.close_tx.subscribe()
}
}
pub struct ProxyConn {
@@ -231,6 +255,7 @@ pub struct ProxyConn {
flow_window_blocked_ms: AtomicU64,
write_latency_last_us: AtomicU64,
write_latency_ewma_us: AtomicU64,
settings: Mutex<protocol::SettingsPayload>,
}
impl ProxyConn {
@@ -244,6 +269,7 @@ impl ProxyConn {
protocol_version: u8,
) -> Self {
Self {
settings: Mutex::new(local_settings()),
id,
node_id,
node_name,
@@ -271,6 +297,11 @@ impl ProxyConn {
self
}
pub(super) fn with_settings(mut self, settings: protocol::SettingsPayload) -> Self {
*self.settings.get_mut() = settings;
self
}
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
self.node_generation = tunnel_generation;
self
@@ -565,13 +596,6 @@ pub struct LocalResponseHead {
pub headers: Vec<(String, String)>,
}
#[derive(Debug)]
pub enum LocalBodyEvent {
Chunk(Bytes),
End,
Error(String),
}
#[derive(Debug, Default)]
struct LocalWaitState {
response: Option<LocalResponseHead>,
@@ -585,10 +609,11 @@ pub struct LocalStream {
proxy_stream_id: u32,
request_window: StreamFlowWindow,
response_consumed_since_update: Mutex<u64>,
min_window_update_bytes: u32,
response_connection: Mutex<Option<std::sync::Weak<ProxyConn>>>,
wait_state: Mutex<LocalWaitState>,
headers_notify: Notify,
body_tx: mpsc::Sender<LocalBodyEvent>,
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
body: Arc<ResponseBuffer>,
terminal: AtomicBool,
}
@@ -600,7 +625,6 @@ impl LocalStream {
proxy_stream_id: u32,
initial_window_bytes: u32,
) -> Self {
let (body_tx, body_rx) = mpsc::channel(128);
Self {
id,
tunnel_generation,
@@ -608,10 +632,11 @@ impl LocalStream {
proxy_stream_id,
request_window: StreamFlowWindow::new(initial_window_bytes),
response_consumed_since_update: Mutex::new(0),
min_window_update_bytes: (initial_window_bytes / 4).clamp(1, 1024 * 1024),
response_connection: Mutex::new(None),
wait_state: Mutex::new(LocalWaitState::default()),
headers_notify: Notify::new(),
body_tx,
body_rx: Mutex::new(Some(body_rx)),
body: ResponseBuffer::new(initial_window_bytes as usize),
terminal: AtomicBool::new(false),
}
}
@@ -632,26 +657,51 @@ impl LocalStream {
self.request_window.add(delta);
}
fn response_window_update_delta(&self, bytes: usize) -> Option<u32> {
if bytes == 0 {
return None;
async fn flush_response_credit(&self) -> Result<(), String> {
if self.terminal.load(Ordering::Acquire) {
return Ok(());
}
let mut consumed = self.response_consumed_since_update.lock();
*consumed = consumed.saturating_add(bytes as u64);
let threshold = u64::from(*STREAM_MIN_WINDOW_UPDATE_BYTES);
if *consumed < threshold {
return None;
let connection = self
.response_connection
.lock()
.as_ref()
.and_then(std::sync::Weak::upgrade);
let Some(connection) = connection else {
return Ok(());
};
if connection.protocol_version() < 3 {
return Ok(());
}
let delta = (*consumed).min(u64::from(u32::MAX)) as u32;
*consumed = consumed.saturating_sub(u64::from(delta));
Some(delta)
let delta = {
let consumed = self.response_consumed_since_update.lock();
if *consumed < u64::from(self.min_window_update_bytes) {
return Ok(());
}
(*consumed).min(u64::from(u32::MAX)) as u32
};
let frame = protocol::encode_window_update(self.proxy_stream_id, delta);
if connection
.send_wait(Message::Binary(frame.into()), OUTBOUND_BACKPRESSURE_TIMEOUT)
.await
== SendStatus::Queued
{
let mut consumed = self.response_consumed_since_update.lock();
*consumed = consumed.saturating_sub(u64::from(delta));
return Ok(());
}
if self.terminal.load(Ordering::Acquire) {
return Ok(());
}
connection.request_close();
Err("proxy flow-control update failed".to_string())
}
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
tokio::time::timeout(timeout, async {
loop {
let notified = self.headers_notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
let outcome = {
let state = self.wait_state.lock();
if let Some(response) = &state.response {
@@ -662,15 +712,20 @@ impl LocalStream {
if let Some(error) = outcome {
return Err(error);
}
self.headers_notify.notified().await;
notified.await;
}
})
.await
.map_err(|_| "timed out waiting for response headers".to_string())?
}
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
self.body_rx.lock().take()
pub fn take_body_receiver(self: &Arc<Self>) -> Option<LocalBodyReceiver> {
self.body.take_receiver().map(|receiver| LocalBodyReceiver {
receiver,
stream: Arc::clone(self),
failed: false,
pending: None,
})
}
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
@@ -690,21 +745,11 @@ impl LocalStream {
}
}
async fn push_body_chunk(&self, payload: Bytes) -> bool {
fn push_body_chunk(&self, payload: Bytes) -> bool {
if self.terminal.load(Ordering::Acquire) {
return false;
}
// Use a timeout to prevent a slow consumer from blocking the shared
// proxy-connection reader (head-of-line blocking across streams).
match tokio::time::timeout(
Duration::from_secs(5),
self.body_tx.send(LocalBodyEvent::Chunk(payload)),
)
.await
{
Ok(Ok(())) => true,
_ => false,
}
self.body.push(payload)
}
fn finish(&self) {
@@ -722,7 +767,8 @@ impl LocalStream {
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::End);
self.request_window.close();
self.body.finish(Ok(()));
}
fn fail(&self, error: impl Into<String>) {
@@ -742,7 +788,38 @@ impl LocalStream {
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
self.request_window.close();
self.body.finish(Err(error));
}
}
pub struct LocalBodyReceiver {
receiver: BodyReceiver,
stream: Arc<LocalStream>,
failed: bool,
pending: Option<LocalBodyEvent>,
}
impl LocalBodyReceiver {
pub async fn recv(&mut self) -> Option<LocalBodyEvent> {
if self.failed {
return None;
}
if self.pending.is_none() {
let event = self.receiver.recv().await?;
if let LocalBodyEvent::Chunk(chunk) = &event {
let mut consumed = self.stream.response_consumed_since_update.lock();
*consumed = consumed.saturating_add(chunk.len() as u64);
}
self.pending = Some(event);
}
if matches!(self.pending, Some(LocalBodyEvent::Chunk(_))) {
if let Err(error) = self.stream.flush_response_credit().await {
self.failed = true;
return Some(LocalBodyEvent::Error(error));
}
}
self.pending.take()
}
}
@@ -765,6 +842,21 @@ pub struct HubRouter {
drain_reasons: Mutex<HashMap<String, u64>>,
}
struct PendingStreamGuard<'router> {
hub: &'router HubRouter,
connection: &'router ProxyConn,
stream_id: u64,
committed: bool,
}
impl Drop for PendingStreamGuard<'_> {
fn drop(&mut self) {
if !self.committed && self.hub.cleanup_local_stream(self.stream_id) {
self.connection.release_stream();
}
}
}
struct NodeStatusEvent {
node_id: String,
authenticated_key: Option<String>,
@@ -1166,17 +1258,27 @@ impl HubRouter {
// Frames encoded successfully -- now register the stream.
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
let local_stream = Arc::new(LocalStream::new(
let settings = proxy_conn.settings.lock().clone();
let mut local_stream = LocalStream::new(
local_stream_id,
proxy_conn.node_generation.clone(),
proxy_conn.id,
proxy_stream_id,
*STREAM_INITIAL_WINDOW_BYTES,
));
settings.initial_stream_window_bytes,
);
local_stream.min_window_update_bytes = settings.min_window_update_bytes;
*local_stream.response_connection.get_mut() = Some(Arc::downgrade(&proxy_conn));
let local_stream = Arc::new(local_stream);
self.local_streams
.insert(local_stream_id, local_stream.clone());
self.proxy_to_local
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
let mut pending_stream = PendingStreamGuard {
hub: self,
connection: &proxy_conn,
stream_id: local_stream_id,
committed: false,
};
let send_status = proxy_conn
.send_wait(
@@ -1195,10 +1297,11 @@ impl HubRouter {
"open_local_stream dispatched"
);
match send_status {
SendStatus::Queued => Ok(local_stream),
SendStatus::Queued => {
pending_stream.committed = true;
Ok(local_stream)
}
SendStatus::Closed | SendStatus::Congested => {
self.cleanup_local_stream(local_stream_id);
proxy_conn.release_stream();
Err("proxy connection congested".to_string())
}
}
@@ -1243,7 +1346,9 @@ impl HubRouter {
.map(|entry| entry.value().clone())
.ok_or_else(|| "proxy connection unavailable".to_string())?;
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
let chunk_size = MAX_REQUEST_BODY_FRAME_SIZE
.min(proxy_conn.settings.lock().initial_stream_window_bytes as usize);
let total_chunks = payload.len().div_ceil(chunk_size);
let result = if total_chunks == 0 {
if end_stream {
self.send_request_body_frame(&proxy_conn, &stream, &[], true)
@@ -1252,7 +1357,7 @@ impl HubRouter {
Ok(())
}
} else {
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
for (index, chunk) in payload.chunks(chunk_size).enumerate() {
let is_last_chunk = index + 1 == total_chunks;
if let Err(error) = self
.send_request_body_frame(
@@ -1352,17 +1457,20 @@ impl HubRouter {
} else {
protocol::encode_stream_error(stream.proxy_stream_id, reason)
};
let _ = pc.send(Message::Binary(frame.into()));
if pc.send(Message::Binary(frame.into())) != SendStatus::Queued {
pc.request_close();
}
}
stream.fail(reason.to_string());
}
fn cleanup_local_stream(&self, local_stream_id: u64) {
fn cleanup_local_stream(&self, local_stream_id: u64) -> bool {
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
return;
return false;
};
self.proxy_to_local
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
true
}
pub async fn handle_proxy_frame(self: &Arc<Self>, proxy_conn_id: u64, data: &mut [u8]) {
@@ -1433,9 +1541,7 @@ impl HubRouter {
.get(&proxy_conn_id)
.map(|entry| entry.value().clone());
if let Some(pc) = pc {
let _ = pc
.send_wait(Message::Binary(pong.into()), Duration::from_millis(250))
.await;
let _ = pc.send(Message::Binary(pong.into()));
}
}
protocol::PONG => {}
@@ -1509,11 +1615,36 @@ impl HubRouter {
);
}
protocol::SETTINGS => {
debug!(
msg_type = header.msg_type,
proxy_conn_id = proxy_conn_id,
"received tunnel protocol v3 SETTINGS from proxy"
);
let settings = protocol::decode_payload_with_limit(
data,
&header,
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
)
.ok()
.and_then(|payload| {
serde_json::from_slice::<protocol::SettingsPayload>(&payload).ok()
})
.filter(|settings| settings.is_valid());
if let Some(connection) = self.proxy_conns_by_id.get(&proxy_conn_id) {
if header.stream_id != 0 || header.flags != 0 {
connection.request_close();
return;
}
let Some(settings) = settings else {
connection.request_close();
return;
};
let local = local_settings();
let settings = settings
.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
let mut current = connection.settings.lock();
if connection.stream_count.load(Ordering::Acquire) > 0 && *current != settings {
drop(current);
connection.request_close();
return;
}
*current = settings;
}
}
protocol::WINDOW_UPDATE => {
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
@@ -1739,19 +1870,8 @@ impl HubRouter {
None => return,
};
let payload_len = payload.len();
if !stream.push_body_chunk(Bytes::from(payload)).await {
if !stream.push_body_chunk(Bytes::from(payload)) {
self.cancel_local_stream(local_id, "local relay response congested");
return;
}
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
if pc.protocol_version() >= 3 {
if let Some(delta) = stream.response_window_update_delta(payload_len) {
let frame = protocol::encode_window_update(header.stream_id, delta);
let _ = pc.send(Message::Binary(frame.into()));
}
}
}
}
@@ -9,14 +9,13 @@ use axum::body::{Body, Bytes};
use axum::extract::{ConnectInfo, Path, Request, State};
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
use axum::response::IntoResponse;
use tokio::sync::mpsc;
use tracing::warn;
use crate::api::response::apply_streaming_response_headers;
use crate::headers::should_skip_response_header;
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
use super::hub::{LocalBodyEvent, LocalStream};
use super::hub::{LocalBodyEvent, LocalBodyReceiver, LocalStream};
use super::protocol;
use super::{AppState, RelayRequestAuthenticated};
@@ -40,7 +39,7 @@ impl Drop for StreamGuard {
pub(crate) struct DirectRelayResponse {
status: u16,
headers: Vec<(String, String)>,
body_rx: mpsc::Receiver<LocalBodyEvent>,
body_rx: LocalBodyReceiver,
request_guard: StreamGuard,
_request_permit: Option<AdmissionPermit>,
}
@@ -55,10 +54,13 @@ impl DirectRelayResponse {
}
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, String> {
if self.request_guard.finished {
return Ok(None);
}
let event = self.body_rx.recv().await;
match event {
Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)),
Some(LocalBodyEvent::End) | None => {
Some(LocalBodyEvent::End) => {
self.request_guard.finished = true;
Ok(None)
}
@@ -66,6 +68,7 @@ impl DirectRelayResponse {
self.request_guard.finished = true;
Err(error)
}
None => Err("tunnel response ended without a terminal frame".to_string()),
}
}
}
@@ -84,6 +87,11 @@ pub(crate) async fn open_direct_relay_stream(
.open_authorized_local_stream(node_id, &meta)
.await
.map_err(|error| format!("connect: {error}"))?;
let request_guard = StreamGuard {
hub: state.hub.clone(),
stream_id: stream.id,
finished: false,
};
if let Err(error) = state
.hub
.push_local_request_body(stream.id, body, true)
@@ -126,11 +134,7 @@ pub(crate) async fn open_direct_relay_stream(
status: response_head.status,
headers: response_head.headers,
body_rx,
request_guard: StreamGuard {
hub: state.hub.clone(),
stream_id: stream.id,
finished: false,
},
request_guard,
_request_permit: request_permit,
})
}
@@ -259,6 +263,11 @@ pub async fn relay_request(
);
}
};
let request_guard = StreamGuard {
hub: state.hub.clone(),
stream_id: stream.id,
finished: false,
};
let body_stream = match spool.body_stream().await {
Ok(stream) => stream,
Err(error) => {
@@ -306,12 +315,6 @@ pub async fn relay_request(
);
}
let request_guard = StreamGuard {
hub: state.hub.clone(),
stream_id: stream.id,
finished: false,
};
let wait_timeout = relay_header_timeout(&meta);
let response_head = match stream.wait_headers(wait_timeout).await {
Ok(response) => response,
@@ -373,6 +376,9 @@ pub async fn relay_request(
}
}
}
if !guard.finished {
yield Err(io::Error::other("tunnel response ended without a terminal frame"));
}
guard.finished = true;
};
@@ -563,6 +569,100 @@ mod tests {
request
}
#[tokio::test]
async fn cancelled_relays_reset_streams_during_upload_and_header_wait() {
for direct in [true, false] {
for during_upload in [true, false] {
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
sample_connected_proxy_node("node-123"),
]));
let data = Arc::new(
GatewayDataState::with_proxy_node_repository_for_tests(repository)
.with_system_config_values_for_tests(
Vec::<(String, serde_json::Value)>::new(),
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
);
let state = test_app_state().with_data(data);
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
let (proxy_close_tx, _) = watch::channel(false);
let connection = Arc::new(
ProxyConn::new(
500,
"node-123".into(),
"Node 123".into(),
proxy_tx,
proxy_close_tx,
16,
3,
)
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
.with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string())
.with_settings(protocol::SettingsPayload {
initial_stream_window_bytes: 128,
min_window_update_bytes: 32,
drain_deadline_ms: 1000,
}),
);
state.hub.register_proxy(Arc::clone(&connection));
let meta = protocol::RequestMeta {
provider_id: None,
endpoint_id: None,
key_id: None,
method: "POST".into(),
url: "https://example.com/".into(),
headers: HashMap::new(),
stream: true,
request_timeout_ms: None,
stream_first_byte_timeout_ms: None,
timeout: 30,
follow_redirects: None,
http1_only: false,
transport_profile: None,
};
let body = Bytes::from(vec![b'x'; if during_upload { 256 } else { 0 }]);
let relay = tokio::spawn(async move {
if direct {
let _response =
super::open_direct_relay_stream(&state, "node-123", meta, body)
.await
.unwrap();
} else {
let request =
authenticated_request(encode_relay_envelope(&meta, &body)).await;
let _response = relay_request(
Path("node-123".into()),
State(state),
ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4242))),
request,
)
.await;
}
});
recv_tunnel_test_frame(&mut proxy_rx, "request headers").await;
recv_tunnel_test_frame(&mut proxy_rx, "request body").await;
relay.abort();
assert!(relay.await.unwrap_err().is_cancelled());
let Message::Binary(frame) = recv_tunnel_test_frame(&mut proxy_rx, "reset").await
else {
panic!("expected binary reset frame")
};
let frame = aether_contracts::tunnel::Frame::decode(frame).unwrap();
assert_eq!(
frame.msg_type,
aether_contracts::tunnel::MsgType::ResetStream
);
assert_eq!(
connection
.stream_count
.load(std::sync::atomic::Ordering::Relaxed),
0
);
assert!(connection.is_available());
}
}
}
#[test]
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
let meta = protocol::RequestMeta {

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