Compare commits

...
12 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
188 changed files with 13559 additions and 2561 deletions
+5 -6
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
Generated
+1
View File
@@ -305,6 +305,7 @@ dependencies = [
"futures-util",
"hmac",
"http",
"http-body",
"http-body-util",
"hyper",
"hyper-util",
+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");
@@ -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();
@@ -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;
@@ -438,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;
@@ -446,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(
@@ -491,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)
}
@@ -5149,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(|_| {
@@ -5316,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,
@@ -5440,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",
@@ -5452,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();
@@ -6228,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()
@@ -6341,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(
@@ -6439,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(),
@@ -6626,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());
@@ -6701,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 ",
@@ -6719,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(),
@@ -6745,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()
@@ -6770,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!(
@@ -6781,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),
@@ -6791,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),
@@ -6801,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");
@@ -6810,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(),
@@ -6872,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");
@@ -6932,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(),
@@ -6985,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(),
@@ -8569,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
@@ -8736,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
@@ -8777,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
@@ -9104,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
@@ -9320,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
@@ -655,6 +655,73 @@ struct ProviderTransferState {
struct ProviderTransferStateTracker {
by_provider: BTreeMap<String, ProviderTransferState>,
exhausted_provider_ids: BTreeSet<String>,
global: GlobalTransferState,
}
#[derive(Debug, Default)]
struct GlobalTransferState {
first_attempt_started_at: Option<Instant>,
last_candidate: Option<(String, String, String)>,
transfer_count: u64,
limits: Option<ProviderTransferLimits>,
exhausted: bool,
}
impl GlobalTransferState {
fn load_policy(&mut self, report_context: Option<&serde_json::Value>) {
if self.limits.is_none() {
if let Some(policy) =
crate::orchestration::routing_execution_policy_from_report_context(report_context)
{
self.limits = Some(ProviderTransferLimits {
max_transfer_count: policy.max_transfer_count,
max_transfer_timeout_seconds: policy.max_transfer_timeout_seconds,
});
}
}
}
fn changes_candidate(&self, plan: &aether_contracts::ExecutionPlan) -> bool {
self.last_candidate
.as_ref()
.is_some_and(|(provider, endpoint, key)| {
provider != &plan.provider_id
|| endpoint != &plan.endpoint_id
|| key != &plan.key_id
})
}
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.first_attempt_started_at.get_or_insert(now);
if self.changes_candidate(plan) {
self.transfer_count = self.transfer_count.saturating_add(1);
}
self.last_candidate = Some((
plan.provider_id.clone(),
plan.endpoint_id.clone(),
plan.key_id.clone(),
));
}
fn check_before_attempt(
&mut self,
plan: &aether_contracts::ExecutionPlan,
now: Instant,
) -> Option<(bool, bool)> {
let limits = self.limits?;
let started_at = self.first_attempt_started_at?;
let count_reached = self.changes_candidate(plan)
&& limits.max_transfer_count > 0
&& self.transfer_count >= limits.max_transfer_count;
let timeout_reached = limits.max_transfer_timeout_seconds > 0
&& now.saturating_duration_since(started_at)
>= Duration::from_secs(limits.max_transfer_timeout_seconds);
if !count_reached && !timeout_reached {
return None;
}
self.exhausted = true;
Some((count_reached, timeout_reached))
}
}
#[derive(Clone, Debug, Default)]
@@ -717,6 +784,7 @@ struct ProviderTransferLimitReached {
impl ProviderTransferStateTracker {
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.global.record_attempt_started(plan, now);
match self.by_provider.entry(plan.provider_id.clone()) {
std::collections::btree_map::Entry::Vacant(entry) => {
entry.insert(ProviderTransferState {
@@ -903,11 +971,42 @@ async fn should_skip_provider_transfer_attempt<Attempt>(
where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let reached = tracker
.state
.lock()
.await
.check_before_attempt(attempt.execution_plan(), Instant::now());
let owned_report_context = attempt
.report_context_ref()
.is_none()
.then(|| attempt.report_context())
.flatten();
let report_context = attempt
.report_context_ref()
.or(owned_report_context.as_ref());
let mut tracker = tracker.state.lock().await;
tracker.global.load_policy(report_context);
if tracker.global.exhausted {
return true;
}
let now = Instant::now();
if let Some((count_reached, timeout_reached)) = tracker
.global
.check_before_attempt(attempt.execution_plan(), now)
{
warn!(
event_name = "routing_transfer_limit_reached",
log_type = "event",
trace_id,
plan_kind,
transfer_count = tracker.global.transfer_count,
elapsed_ms = tracker
.global
.first_attempt_started_at
.map(|started| now.saturating_duration_since(started).as_millis() as u64)
.unwrap_or(0),
count_reached,
timeout_reached,
"gateway exhausted the routing strategy transfer budget"
);
return true;
}
let reached = tracker.check_before_attempt(attempt.execution_plan(), now);
let Some(reached) = reached else {
return false;
};
@@ -2465,6 +2564,130 @@ mod tests {
assert_eq!(port.unused.lock().unwrap().as_slice(), ["a-key3-retry0"]);
}
#[tokio::test]
async fn routing_transfer_budget_counts_switches_across_providers_not_same_key_retries() {
for (limit, succeeds) in [(1, false), (2, true)] {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] =
json!({ "max_transfer_count": limit });
}
let outcome = run_ai_attempt_loop(&port, attempts).await.unwrap();
assert_eq!(
matches!(outcome, AiAttemptLoopOutcome::Responded(_)),
succeeds
);
{
let executed = port.executed.lock().unwrap();
assert_eq!(
&executed[..3],
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(executed.len(), if succeeds { 4 } else { 3 });
}
assert_eq!(port.tracker.state.lock().await.global.transfer_count, limit);
}
}
#[tokio::test]
async fn dynamic_loop_honors_global_transfer_budget_across_providers() {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let mut source = TransferTestAttemptSource {
attempts: attempts.into(),
skipped_providers: Vec::new(),
};
let outcome = run_dynamic_attempt_loop(
&port,
&mut source,
"global-budget",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
outcome,
LocalExecutionRequestOutcome::Exhausted(_)
));
assert_eq!(
port.executed.lock().unwrap().as_slice(),
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
}
#[test]
fn routing_time_budget_is_cumulative_and_zero_is_unlimited() {
let mut global = super::GlobalTransferState::default();
global.load_policy(Some(
&json!({ "routing_execution_policy": { "max_transfer_timeout_seconds": 60 } }),
));
let now = tokio::time::Instant::now();
let plan = test_plan(None);
global.record_attempt_started(&plan, now);
global.record_attempt_started(&plan, now + Duration::from_secs(40));
assert_eq!(global.transfer_count, 0);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(59)),
None
);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(60)),
Some((false, true))
);
let mut unlimited = super::GlobalTransferState::default();
unlimited.load_policy(Some(&json!({ "routing_execution_policy": {} })));
unlimited.record_attempt_started(&plan, now);
assert_eq!(
unlimited.check_before_attempt(&plan, now + Duration::from_secs(86_400)),
None
);
}
#[tokio::test]
async fn cloned_tracker_preserves_global_budget_across_candidate_loops() {
let state = AppState::new().unwrap();
let tracker = ProviderTransferTracker::default();
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let remaining = attempts.split_off(3);
let first_port = TransferTestPort::with_tracker(&state, tracker.clone());
let first_outcome = run_ai_attempt_loop(&first_port, attempts).await.unwrap();
assert!(matches!(first_outcome, AiAttemptLoopOutcome::Exhausted(_)));
assert_eq!(tracker.state.lock().await.global.transfer_count, 1);
let second_port = TransferTestPort::with_tracker(&state, tracker.clone());
let mut source = TransferTestAttemptSource {
attempts: remaining.into(),
skipped_providers: Vec::new(),
};
let second_outcome = run_dynamic_attempt_loop(
&second_port,
&mut source,
"global-budget-across-loops",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
second_outcome,
LocalExecutionRequestOutcome::NoPath
));
assert!(second_port.executed.lock().unwrap().is_empty());
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
assert!(tracker.state.lock().await.global.exhausted);
}
#[tokio::test]
async fn cloned_tracker_preserves_transfer_budget_across_candidate_loops() {
let state = AppState::new().expect("state should build");
@@ -822,15 +822,20 @@ where
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = standard_text_sync_heartbeat_final_bytes(
let completion = standard_text_sync_heartbeat_final_bytes(
client_api_format.as_str(),
redaction_slot.as_ref(),
execute(state, parts, trace_id, decision, plan_kind, started_at).await,
)
.await;
tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execute(state, parts, trace_id, decision, plan_kind, started_at) => result,
},
);
let bytes = completion.await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -1097,23 +1102,26 @@ fn build_openai_image_sync_heartbeat_shell_response(
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = openai_image_sync_heartbeat_final_bytes(
execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
)
.await,
)
.await;
let execution = execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
);
let outcome = tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execution => result,
};
let bytes = openai_image_sync_heartbeat_final_bytes(outcome).await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -2331,6 +2339,45 @@ mod tests {
.expect("background completion should release admission");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_cancels_when_routing_policy_enables_it() {
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (mut release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let response = crate::request_lifecycle::run_request(async move {
crate::request_lifecycle::configure_client_disconnect(
aether_routing_core::RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
},
);
let (parts, _) = http::Request::builder()
.method("POST")
.uri("/v1/responses")
.body(())
.unwrap()
.into_parts();
build_standard_text_sync_heartbeat_shell_response(
AppState::new().unwrap(),
parts,
"trace-heartbeat-disconnect".to_string(),
test_standard_text_heartbeat_decision(),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND.to_string(),
move |_, _, _, _, _, _| async move {
started_tx.send(()).unwrap();
release_rx.await.unwrap();
Ok(LocalExecutionRequestOutcome::NoPath)
},
)
})
.await
.unwrap();
started_rx.await.unwrap();
drop(response);
tokio::time::timeout(Duration::from_secs(1), release_tx.closed())
.await
.expect("heartbeat must drop upstream execution immediately");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_propagates_request_diagnostics_to_terminal_usage() {
let (state, usage_repository) = heartbeat_usage_test_state(json!({
+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
);
}
}
}
@@ -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,31 +149,14 @@ pub(crate) fn websocket_upstream_url(
invalid_code: &'static str,
) -> Result<Url, &'static str> {
let mut url = Url::parse(raw).map_err(|_| invalid_code)?;
if url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
|| url.fragment().is_some()
{
return Err(invalid_code);
}
let websocket_scheme = match url.scheme() {
"https" | "wss" => "wss",
"http" | "ws" => "ws",
let (http_scheme, websocket_scheme) = match url.scheme() {
"https" | "wss" => ("https", "wss"),
"http" | "ws" => ("http", "ws"),
_ => return Err(invalid_code),
};
url.set_scheme(http_scheme).map_err(|_| invalid_code)?;
let mut url = validate_execution_upstream_url(url.as_str()).map_err(|_| invalid_code)?;
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
if url.scheme() == "ws" {
let literal_ip = match url.host() {
Some(url::Host::Ipv4(address)) => Some(std::net::IpAddr::V4(address)),
Some(url::Host::Ipv6(address)) => Some(std::net::IpAddr::V6(address)),
_ => None,
};
if literal_ip.is_some_and(|address| {
aether_http::is_private_or_reserved_ip(address) && !address.is_loopback()
}) {
return Err(invalid_code);
}
}
Ok(url)
}
@@ -229,9 +212,8 @@ pub(crate) fn websocket_handshake_headers(
Ok(headers)
}
async fn build_websocket_client(
fn build_websocket_client(
decision: &AiExecutionDecision,
upstream_url: &Url,
errors: UpstreamWebSocketErrorCodes,
) -> Result<wreq::Client, &'static str> {
let timeouts = websocket_timeouts(decision);
@@ -255,41 +237,7 @@ async fn build_websocket_client(
let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?;
builder = builder.proxy(proxy);
} else {
// Pin every direct WebSocket connection to the DNS answers validated
// here. This also covers the explicitly permitted loopback `ws://`
// form; otherwise the client would perform a second lookup and a
// rebinding could escape the loopback-only policy.
let host = upstream_url.host_str().ok_or(errors.upstream_url_invalid)?;
let port = upstream_url
.port_or_known_default()
.ok_or(errors.upstream_url_invalid)?;
let addresses = if let Ok(ip) = host.parse::<std::net::IpAddr>() {
vec![std::net::SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host,
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| errors.upstream_url_invalid)?
};
let allows_loopback = host.trim_end_matches('.').eq_ignore_ascii_case("localhost")
|| host
.parse::<std::net::IpAddr>()
.map(|ip| ip.is_loopback())
.unwrap_or(false);
let unsafe_answer = if allows_loopback {
addresses.iter().any(|address| !address.ip().is_loopback())
} else {
addresses
.iter()
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()))
};
if addresses.is_empty() || unsafe_answer {
return Err(errors.upstream_url_invalid);
}
builder = builder.resolve_to_addrs(host.to_string(), addresses.iter().copied());
builder = builder.dns_resolver(ExecutionSafeDnsResolver);
}
builder.build().map_err(|_| errors.client_build_failed)
}
@@ -684,15 +632,17 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16
#[cfg(test)]
mod tests {
use super::{
bounded_send, guarded_websocket_upstream_url, resolve_websocket_proxy_url,
responses_websocket_error_event, responses_websocket_error_event_with_stream_id,
websocket_handshake_headers, websocket_relay_frame_queue, websocket_response_headers,
websocket_upstream_url, UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl,
WebSocketRelayQueueError, WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY,
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
bounded_send, build_websocket_client, guarded_websocket_upstream_url,
resolve_websocket_proxy_url, responses_websocket_error_event,
responses_websocket_error_event_with_stream_id, websocket_handshake_headers,
websocket_relay_frame_queue, websocket_response_headers, websocket_upstream_url,
UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl, WebSocketRelayQueueError,
WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY, RELAY_WRITE_TIMEOUT,
TEARDOWN_WRITE_TIMEOUT,
};
use crate::ai_serving::AiExecutionDecision;
use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url;
use aether_contracts::ProxySnapshot;
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
use axum::http::HeaderMap;
use std::collections::BTreeMap;
use std::time::Duration;
@@ -876,7 +826,9 @@ mod tests {
"ws://example.test:8080/v1/responses",
"http://example.test:8080/v1/responses",
"http://8.8.8.8:8080/v1/responses",
"wss://8.8.8.8/v1/responses",
"ws://[2606:4700:4700::1111]:8080/v1/responses",
"wss://[2606:4700:4700::1111]/v1/responses",
"ws://localhost:8080/v1/responses",
"http://127.42.0.1:8080/v1/responses",
"ws://[::1]:8080/v1/responses",
@@ -888,6 +840,14 @@ mod tests {
}
for rejected in [
"http://10.0.0.1/v1/responses",
"wss://10.0.0.1/v1/responses",
"wss://127.0.0.1/v1/responses",
"wss://[::1]/v1/responses",
"wss://[fd00::1]/v1/responses",
"wss://[::ffff:127.0.0.1]/v1/responses",
"wss://169.254.169.254/v1/responses",
"wss://198.18.78.41/v1/responses",
"wss://198.19.1.2/v1/responses",
"ws://0.0.0.0:8080/v1/responses",
"ws://[::ffff:127.0.0.1]:8080/v1/responses",
"wss://example.test/v1/responses#secret",
@@ -903,6 +863,60 @@ mod tests {
}
}
#[tokio::test]
async fn websocket_client_build_defers_provider_dns_for_all_transport_profiles() {
let errors = UpstreamWebSocketErrorCodes {
upstream_url_missing: "missing",
upstream_url_invalid: "upstream_invalid",
frontdoor_self_loop: "frontdoor_self_loop",
headers_invalid: "headers_invalid",
client_build_failed: "client_build_failed",
proxy_invalid: "proxy_invalid",
tunnel_proxy_unsupported: "tunnel_unsupported",
handshake_failed: "handshake_failed",
upgrade_rejected: "upgrade_rejected",
upgrade_failed: "upgrade_failed",
};
for profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for proxy in [
None,
Some(ProxySnapshot {
enabled: Some(false),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("socks5h://proxy.invalid:1080".to_string()),
..Default::default()
}),
] {
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "wss://upstream.invalid/v1/responses"
}))
.expect("minimal provider decision should deserialize");
decision.transport_profile = profile.clone();
decision.proxy = proxy;
build_websocket_client(&decision, errors)
.expect("building a client must not resolve the provider or proxy hostname");
}
}
}
#[test]
fn active_websocket_proxy_without_a_target_fails_closed() {
let errors = UpstreamWebSocketErrorCodes {
@@ -937,6 +951,94 @@ mod tests {
);
}
#[tokio::test]
async fn websocket_handshake_keeps_provider_dns_remote_for_http_and_socks_proxies() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let errors = UpstreamWebSocketErrorCodes {
upstream_url_missing: "missing",
upstream_url_invalid: "upstream_invalid",
frontdoor_self_loop: "frontdoor_self_loop",
headers_invalid: "headers_invalid",
client_build_failed: "client_build_failed",
proxy_invalid: "proxy_invalid",
tunnel_proxy_unsupported: "tunnel_unsupported",
handshake_failed: "handshake_failed",
upgrade_rejected: "upgrade_rejected",
upgrade_failed: "upgrade_failed",
};
for profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for scheme in ["http", "socks5", "socks5h"] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let proxy_addr = listener.local_addr().unwrap();
let (release, released) = tokio::sync::oneshot::channel::<()>();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
if scheme != "http" {
let mut greeting = [0; 2];
stream.read_exact(&mut greeting).await.unwrap();
assert_eq!(greeting[0], 5);
let mut methods = vec![0; greeting[1] as usize];
stream.read_exact(&mut methods).await.unwrap();
assert!(methods.contains(&0));
stream.write_all(&[5, 0]).await.unwrap();
let mut request = [0; 4];
stream.read_exact(&mut request).await.unwrap();
assert_eq!(
request,
[5, 1, 0, 3],
"proxy must receive a domain, not an IP"
);
let host_len = stream.read_u8().await.unwrap();
let mut host = vec![0; host_len as usize];
stream.read_exact(&mut host).await.unwrap();
assert_eq!(host, b"provider-dns.invalid");
assert_eq!(stream.read_u16().await.unwrap(), 80);
stream
.write_all(&[5, 0, 0, 1, 127, 0, 0, 1, 0, 80])
.await
.unwrap();
}
let socket = tokio_tungstenite::accept_async(stream).await.unwrap();
let _ = released.await;
drop(socket);
});
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "ws://provider-dns.invalid/v1/responses",
"proxy": {"enabled": true, "url": format!("{scheme}://{proxy_addr}")}
}))
.unwrap();
decision.transport_profile = profile.clone();
let connection = tokio::time::timeout(
Duration::from_secs(5),
super::connect_upstream_websocket(
&decision,
crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS,
errors,
),
)
.await
.expect("proxied handshake must not wait for local provider DNS")
.unwrap_or_else(|error| panic!("{scheme} handshake failed: {error}"));
release.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(5), server)
.await
.unwrap()
.unwrap();
drop(connection);
}
}
}
#[test]
fn rejects_responses_websocket_frontdoor_self_loop_before_connecting() {
let base_url = configured_gateway_frontdoor_base_url();
@@ -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,98 +22,16 @@ use super::{
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
const MAX_TEST_CONNECTION_RESPONSE_BYTES: usize = 256 * 1024;
#[cfg(test)]
fn build_test_connection_client() -> Result<reqwest::Client, reqwest::Error> {
reqwest::Client::builder()
.no_proxy()
.dns_resolver(Arc::new(ExecutionSafeDnsResolver))
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(10))
.http2_adaptive_window(true)
.build()
}
#[derive(Debug)]
struct ResolvedTestConnectionTarget {
url: reqwest::Url,
host: String,
addresses: Vec<SocketAddr>,
}
/// Resolve the provider endpoint once and pin reqwest to that answer. The
/// test-connection route is reachable through the public front door, so it
/// must not perform an unbounded DNS lookup on every connect (which would
/// permit DNS rebinding into private/reserved networks).
async fn resolve_test_connection_target(
raw_url: &str,
allow_private_targets: bool,
) -> Result<ResolvedTestConnectionTarget, &'static str> {
let url = reqwest::Url::parse(raw_url).map_err(|_| "provider endpoint URL is invalid")?;
let literal_loopback = aether_http::url_has_literal_loopback_host(&url);
if !matches!(url.scheme(), "http" | "https")
|| url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
|| url.fragment().is_some()
{
return Err("provider endpoint must be an HTTP(S) URL without credentials or fragment");
}
let host = url
.host_str()
.ok_or("provider endpoint is missing a host")?
.to_string();
let literal_ip = host.parse::<IpAddr>().ok();
let port = url
.port_or_known_default()
.ok_or("provider endpoint is missing a port")?;
let addresses = if let Some(ip) = literal_ip {
vec![SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host.as_str(),
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| "provider endpoint DNS resolution failed")?
};
if addresses.is_empty() {
return Err("provider endpoint DNS resolution returned no addresses");
}
let has_private_answer = addresses
.iter()
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()));
// `allow_private_targets` is only enabled for in-process test fixtures.
// Keep that escape hatch narrowly scoped to literal loopback URLs whose
// every DNS answer is loopback; otherwise a test-only build (or an
// accidentally reused helper) could turn this public route into a
// private-network HTTP client.
let test_loopback_target = allow_private_targets
&& literal_loopback
&& addresses.iter().all(|address| address.ip().is_loopback());
if has_private_answer && !test_loopback_target {
return Err("provider endpoint resolves to a private or reserved address");
}
Ok(ResolvedTestConnectionTarget {
url,
host,
addresses,
})
}
fn build_pinned_test_connection_client(
target: &ResolvedTestConnectionTarget,
) -> Result<reqwest::Client, reqwest::Error> {
let mut builder = reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(10))
.http2_adaptive_window(true);
if target.host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(&target.host, &target.addresses);
}
builder.build()
}
pub(super) async fn maybe_build_local_test_connection_route_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -384,18 +306,14 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
// Resolve and pin the endpoint before constructing the request. This
// keeps the public health-check route subject to the same DNS/SSRF
// boundary as the main execution transport. Unit-test fixtures may use
// loopback listeners; production requests never opt into private targets.
let target = match resolve_test_connection_target(&upstream_url, cfg!(test)).await {
Ok(target) => target,
let upstream_url = match validate_execution_upstream_url(&upstream_url) {
Ok(url) => url,
Err(reason) => {
tracing::warn!(
event_name = "provider_test_connection_target_rejected",
provider_id = %provider.id,
endpoint_id = %endpoint.id,
reason,
reason = %reason,
"provider connection test target was rejected"
);
return Some(
@@ -407,7 +325,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let test_client = match build_pinned_test_connection_client(&target) {
let test_client = match build_test_connection_client() {
Ok(client) => client,
Err(_) => {
return Some(
@@ -419,7 +337,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let mut upstream_request = test_client.post(target.url);
let mut upstream_request = test_client.post(upstream_url);
for (name, value) in &provider_request_headers {
upstream_request = upstream_request.header(name, value);
}
@@ -495,7 +413,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
#[cfg(test)]
mod tests {
use super::{build_test_connection_client, resolve_test_connection_target};
use super::{build_test_connection_client, validate_execution_upstream_url};
use axum::{
body::Body,
http::{header, Request, StatusCode},
@@ -567,74 +485,59 @@ mod tests {
redirected_server.abort();
}
#[tokio::test]
async fn test_connection_target_rejects_private_addresses_in_production_mode() {
#[test]
fn test_connection_target_rejects_private_literals_like_provider_requests() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat/completions",
"http://10.0.0.1/v1/chat/completions",
"http://169.254.169.254/v1/chat/completions",
"https://10.0.0.1/v1/chat/completions",
"https://127.0.0.1/v1/chat/completions",
"https://[::1]/v1/chat/completions",
"https://localhost/v1/chat/completions",
"https://198.18.78.41/v1/chat/completions",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).is_err(),
"private provider target should be rejected: {raw_url}"
);
}
}
#[tokio::test]
async fn test_connection_target_accepts_public_http_and_https_addresses() {
for allow_private_targets in [false, true] {
for (raw_url, expected_port) in [
("http://8.8.8.8/v1/chat", 80),
("http://8.8.8.8:8080/v1/chat", 8080),
("https://8.8.8.8/v1/chat", 443),
] {
let target = resolve_test_connection_target(raw_url, allow_private_targets)
.await
.expect("public HTTP(S) provider target should resolve");
assert_eq!(target.url.as_str(), raw_url);
assert_eq!(target.host, "8.8.8.8");
assert_eq!(target.addresses.len(), 1);
assert_eq!(target.addresses[0].ip().to_string(), "8.8.8.8");
assert_eq!(target.addresses[0].port(), expected_port);
}
#[test]
fn test_connection_target_accepts_public_http_and_https_addresses() {
for (raw_url, expected_port) in [
("http://8.8.8.8/v1/chat", 80),
("http://8.8.8.8:8080/v1/chat", 8080),
("https://8.8.8.8/v1/chat", 443),
("https://[2606:4700:4700::1111]/v1/chat", 443),
] {
let url = validate_execution_upstream_url(raw_url)
.expect("public HTTP(S) provider target should be valid");
assert_eq!(url.as_str(), raw_url);
assert_eq!(url.port_or_known_default(), Some(expected_port));
}
}
#[tokio::test]
async fn test_connection_target_allows_loopback_only_for_test_fixtures() {
let target = resolve_test_connection_target("http://127.0.0.1:8080/v1/chat", true)
.await
.expect("test fixture target should resolve");
assert_eq!(target.host, "127.0.0.1");
assert_eq!(target.addresses.len(), 1);
assert!(
resolve_test_connection_target("http://10.0.0.1/v1/chat", true)
.await
.is_err(),
"test mode must not make private non-loopback HTTP endpoints acceptable"
);
assert!(
resolve_test_connection_target("https://10.0.0.1/v1/chat", true)
.await
.is_err(),
"test mode must not make private non-loopback endpoints acceptable"
);
assert!(
resolve_test_connection_target("http://localhost:8080/v1/chat", true)
.await
.is_ok(),
"literal localhost should remain available for local fixtures"
);
async fn test_connection_target_defers_dns_and_accepts_provider_loopback_urls() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat",
"http://[::1]:8080/v1/chat",
"http://localhost:8080/v1/chat",
"https://provider-dns.invalid/v1/chat",
] {
let url = validate_execution_upstream_url(raw_url)
.expect("target validation must not depend on the current DNS answer");
let request = build_test_connection_client()
.expect("client should build without DNS")
.post(url)
.build()
.expect("provider request should build without DNS");
assert_eq!(request.url().as_str(), raw_url);
}
}
#[tokio::test]
async fn test_connection_target_rejects_url_credentials_and_fragments() {
#[test]
fn test_connection_target_rejects_url_credentials_and_fragments() {
for raw_url in [
"https://user:[email protected]/v1/chat",
"https://example.com/v1/chat#fragment",
@@ -643,9 +546,7 @@ mod tests {
"ftp://example.com/v1/chat",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).is_err(),
"unsafe provider target should be rejected: {raw_url}"
);
}
@@ -168,52 +168,6 @@ fn wallet_public_refund_payload(mut payload: serde_json::Value) -> serde_json::V
payload
}
#[cfg(test)]
mod tests {
use super::wallet_refund_payload_from_record;
use aether_data::repository::wallet::StoredAdminWalletRefund;
use serde_json::json;
#[test]
fn public_refund_projection_excludes_payout_proof_and_upstream_payload() {
let record = StoredAdminWalletRefund {
id: "refund-1".to_string(),
refund_no: "rf_1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
payment_order_id: Some("order-1".to_string()),
source_type: "payment_order".to_string(),
source_id: Some("order-1".to_string()),
refund_mode: "original_channel".to_string(),
amount_usd: 10.0,
status: "processing".to_string(),
reason: Some("requested".to_string()),
failure_reason: None,
gateway_refund_id: Some("gateway-refund-1".to_string()),
payout_method: None,
payout_reference: None,
payout_proof: Some(json!({
"gateway_refund": {
"id": "gateway-refund-1",
"payload": {"payer": "sensitive", "credential": "secret"}
}
})),
requested_by: Some("user-1".to_string()),
approved_by: Some("admin-1".to_string()),
processed_by: Some("admin-1".to_string()),
created_at_unix_ms: 1,
updated_at_unix_secs: 1,
processed_at_unix_secs: Some(1),
completed_at_unix_secs: None,
};
let payload = wallet_refund_payload_from_record(&record);
assert!(payload.get("payout_proof").is_none());
assert_eq!(payload["status"], "processing");
assert_eq!(payload["gateway_refund_id"], "gateway-refund-1");
}
}
pub(super) async fn handle_wallet_refunds_list(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -659,3 +613,49 @@ pub(super) async fn handle_wallet_create_refund(
}
}
}
#[cfg(test)]
mod tests {
use super::wallet_refund_payload_from_record;
use aether_data::repository::wallet::StoredAdminWalletRefund;
use serde_json::json;
#[test]
fn public_refund_projection_excludes_payout_proof_and_upstream_payload() {
let record = StoredAdminWalletRefund {
id: "refund-1".to_string(),
refund_no: "rf_1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
payment_order_id: Some("order-1".to_string()),
source_type: "payment_order".to_string(),
source_id: Some("order-1".to_string()),
refund_mode: "original_channel".to_string(),
amount_usd: 10.0,
status: "processing".to_string(),
reason: Some("requested".to_string()),
failure_reason: None,
gateway_refund_id: Some("gateway-refund-1".to_string()),
payout_method: None,
payout_reference: None,
payout_proof: Some(json!({
"gateway_refund": {
"id": "gateway-refund-1",
"payload": {"payer": "sensitive", "credential": "secret"}
}
})),
requested_by: Some("user-1".to_string()),
approved_by: Some("admin-1".to_string()),
processed_by: Some("admin-1".to_string()),
created_at_unix_ms: 1,
updated_at_unix_secs: 1,
processed_at_unix_secs: Some(1),
completed_at_unix_secs: None,
};
let payload = wallet_refund_payload_from_record(&record);
assert!(payload.get("payout_proof").is_none());
assert_eq!(payload["status"], "processing");
assert_eq!(payload["gateway_refund_id"], "gateway-refund-1");
}
}
@@ -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))]
@@ -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))",
@@ -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
@@ -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) =
@@ -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(
+3 -3
View File
@@ -475,7 +475,7 @@ async fn gateway_truncates_deep_request_echo_for_local_openai_chat_sync_usage_im
);
assert_eq!(
stored_usage.request_headers.as_ref().unwrap()["authorization"],
"[redacted]"
"Bearer sk-client-openai-local-report-sync-deep"
);
gateway_handle.abort();
@@ -1100,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(),
@@ -2094,7 +2094,7 @@ async fn gateway_records_failed_usage_when_all_local_claude_cli_candidates_are_s
assert!(stored_usage.error_message.is_none());
assert_eq!(
stored_usage.request_headers.as_ref().unwrap()["authorization"],
"[redacted]"
"Bearer sk-client-claude-cli-usage-local-miss"
);
assert!(stored_usage.request_body.is_none());
assert!(stored_usage.request_body_ref.is_none());
+31 -6
View File
@@ -170,12 +170,13 @@ Linux/macOS 可运行 `sudo aether-tunnel upgrade [version]`。自更新只接
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--upstream-connect-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_CONNECT_TIMEOUT_SECS` | `30` | 上游建连超时(秒) |
| `--upstream-connect-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_CONNECT_TIMEOUT` | `30` | 上游建连超时(秒) |
| `--upstream-pool-max-idle-per-host` | `AETHER_TUNNEL_UPSTREAM_POOL_MAX_IDLE_PER_HOST` | `64` | 每 Host 最大空闲连接数 |
| `--upstream-pool-idle-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_POOL_IDLE_TIMEOUT_SECS` | `300` | 连接池空闲超时(秒) |
| `--upstream-tcp-keepalive-secs` | `AETHER_TUNNEL_UPSTREAM_TCP_KEEPALIVE_SECS` | `60` | TCP keepalive(秒,0 关闭) |
| `--upstream-pool-idle-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_POOL_IDLE_TIMEOUT` | `300` | 连接池空闲超时(秒) |
| `--upstream-tcp-keepalive-secs` | `AETHER_TUNNEL_UPSTREAM_TCP_KEEPALIVE` | `60` | TCP keepalive(秒,0 关闭) |
| `--upstream-tcp-nodelay` | `AETHER_TUNNEL_UPSTREAM_TCP_NODELAY` | `true` | 启用 TCP_NODELAY |
| `--upstream-proxy-url` | `AETHER_TUNNEL_UPSTREAM_PROXY_URL` | 空 | 仅 provider 上游请求使用的出口代理 |
| `--upstream-proxy-remote-dns` | `AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS` | `false` | 显式信任 HTTP/SOCKS5h 代理解析供应商域名并执行目标 IP 访问控制;需重启 |
启用 `follow_redirects` 后,同源 307/308 会在请求体不超过 5 MiB 时重放。首个上游请求始终流式传输;超过重放预算时不会拒绝或截断原请求,而是将 307/308 响应原样返回给调用方。
@@ -185,14 +186,38 @@ Linux/macOS 可运行 `sudo aether-tunnel upgrade [version]`。自更新只接
upstream_proxy_url = "socks5h://microwarp:1080"
```
默认仍由隧道本机解析供应商域名、执行端口/IP ACL,再把已校验的 IP 交给代理;仅配置
`socks5h://` 不会跳过本地 DNS。这保留现有的防 DNS 重绑定及内网访问边界。
如果隧道本机 DNS 不可用、被污染或返回不可路由的 Fake-IP,可显式委托**受信任且配置了
目的地址访问控制的代理**解析域名。在 TOML 顶层(第一个 `[[servers]]` 之前)配置:
```toml
upstream_proxy_url = "socks5h://microwarp:1080"
upstream_proxy_remote_dns = true
```
也可启用环境变量 `AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS=true`、CLI 参数
`--upstream-proxy-remote-dns` 或 setup 中的 `Proxy Remote DNS` 开关,保存后重启。
该模式仅支持 `http://` 和 `socks5h://`,不支持本地解析语义的 `socks5://`;未配置代理时
启动会报错。域名原样交给 HTTP CONNECT/SOCKS5h,HTTP Host 和 TLS SNI/证书校验仍使用
原域名,不会在失败时偷偷回退到本地 DNS。
**安全边界:**普通 HTTP CONNECT/SOCKS5 不能让隧道校验代理最终解析出的目标 IP,因此
启用该模式代表把域名目标的 IP ACL 委托给代理,而不只是换一个 DNS 服务器。隧道仍检查
端口、URL 凭据/fragment、`localhost` 和 IP 字面地址;默认继续拒绝私网/保留 IP 字面地址。
这不需要打开 `allow_private_targets`。代理本身的域名仍需本地解析;如果本地 DNS 完全
不可用,使用代理 IP 地址或修复本地解析。代理 DNS、TCP、CONNECT/SOCKS 和 TLS 握手共同
受 `upstream_connect_timeout_secs` 限制。
如果需要让 Aether 管理 API 和 WebSocket tunnel 也走代理,使用 `aether_outbound_proxy_url`。
#### Aether API 客户端
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-request-timeout-secs` | `AETHER_TUNNEL_AETHER_REQUEST_TIMEOUT_SECS` | `10` | 请求总超时(秒) |
| `--aether-connect-timeout-secs` | `AETHER_TUNNEL_AETHER_CONNECT_TIMEOUT_SECS` | `10` | 建连超时(秒) |
| `--aether-request-timeout-secs` | `AETHER_TUNNEL_AETHER_REQUEST_TIMEOUT` | `10` | 请求总超时(秒) |
| `--aether-connect-timeout-secs` | `AETHER_TUNNEL_AETHER_CONNECT_TIMEOUT` | `10` | 建连超时(秒) |
| `--aether-outbound-proxy-url` | `AETHER_TUNNEL_AETHER_OUTBOUND_PROXY_URL` | 空 | Aether 注册、心跳和 WebSocket tunnel 回连使用的出口代理(默认不走代理) |
| `--aether-retry-max-attempts` | `AETHER_TUNNEL_AETHER_RETRY_MAX_ATTEMPTS` | `3` | 最大重试次数 |
@@ -201,7 +226,7 @@ upstream_proxy_url = "socks5h://microwarp:1080"
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `false` | 默认拦截 private/reserved 目标地址;仅在明确需要访问内网服务时设为 `true`,且仅影响重启后的进程 |
| `--dns-cache-ttl-secs` | `AETHER_TUNNEL_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL(秒) |
| `--dns-cache-ttl-secs` | `AETHER_TUNNEL_DNS_CACHE_TTL` | `60` | DNS 缓存 TTL(秒) |
| `--dns-cache-capacity` | `AETHER_TUNNEL_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) |
#### 日志
+6
View File
@@ -126,11 +126,16 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
if let Ok(proxy) = crate::egress_proxy::UpstreamProxyConfig::parse(proxy_url) {
info!(
upstream_proxy_url = %proxy.redacted_url(),
upstream_proxy_remote_dns = config.upstream_proxy_remote_dns,
"provider upstream egress proxy configured"
);
}
}
if config.upstream_proxy_remote_dns {
warn!("provider hostname DNS resolution and destination IP access controls are delegated to the trusted upstream proxy");
}
// Resolve public IP (best-effort for region info)
let public_ip = match &config.public_ip {
Some(ip) => ip.clone(),
@@ -1420,6 +1425,7 @@ mod tests {
upstream_tcp_keepalive_secs: 60,
upstream_tcp_nodelay: true,
upstream_proxy_url: None,
upstream_proxy_remote_dns: false,
legacy_redirect_replay_budget_bytes_ignored: None,
emit_proxy_timing_header: true,
log_level: "info".to_string(),
+75
View File
@@ -480,6 +480,14 @@ pub struct Config {
#[arg(long, env = "AETHER_TUNNEL_UPSTREAM_PROXY_URL")]
pub upstream_proxy_url: Option<String>,
#[arg(
long,
env = "AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS",
default_value_t = false,
help = "Trust an HTTP or SOCKS5h upstream proxy to resolve hostnames and enforce destination IP access controls"
)]
pub upstream_proxy_remote_dns: bool,
/// Accepted only so older launch commands and environments keep working.
/// Redirect request bodies are always replayed without a cumulative size limit.
#[arg(
@@ -820,6 +828,16 @@ impl Config {
crate::egress_proxy::UpstreamProxyConfig::parse(proxy_url)
.map_err(|err| anyhow::anyhow!("upstream_proxy_url invalid: {err}"))?;
}
if self.upstream_proxy_remote_dns {
let proxy_url = normalized_proxy_url(&self.upstream_proxy_url).ok_or_else(|| {
anyhow::anyhow!("upstream_proxy_remote_dns requires upstream_proxy_url")
})?;
let proxy = crate::egress_proxy::UpstreamProxyConfig::parse(proxy_url)
.map_err(anyhow::Error::msg)?;
if !proxy.supports_remote_target_dns() {
anyhow::bail!("upstream_proxy_remote_dns requires an http:// or socks5h:// proxy");
}
}
if matches!(self.max_in_flight_streams, Some(0)) {
anyhow::bail!("max_in_flight_streams must be > 0");
}
@@ -1087,6 +1105,8 @@ pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_proxy_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_proxy_remote_dns: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub emit_proxy_timing_header: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub log_level: Option<String>,
@@ -1306,6 +1326,10 @@ impl ConfigFile {
self.upstream_tcp_nodelay
);
set!("AETHER_TUNNEL_UPSTREAM_PROXY_URL", self.upstream_proxy_url);
set!(
"AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS",
self.upstream_proxy_remote_dns
);
set!(
"AETHER_TUNNEL_EMIT_PROXY_TIMING_HEADER",
self.emit_proxy_timing_header
@@ -1670,6 +1694,57 @@ mod tests {
use super::*;
use crate::hardware::HardwareInfo;
#[test]
fn proxy_remote_dns_requires_explicit_trust_and_a_remote_dns_proxy() {
let mut config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
]);
assert!(!config.upstream_proxy_remote_dns);
let argument = Config::command()
.get_arguments()
.find(|argument| argument.get_id() == "upstream_proxy_remote_dns")
.unwrap()
.clone();
assert_eq!(
argument.get_env().unwrap(),
"AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS"
);
config.upstream_proxy_remote_dns = true;
for proxy in [None, Some(" "), Some("socks5://127.0.0.1:1080")] {
config.upstream_proxy_url = proxy.map(str::to_string);
assert!(config
.validate()
.unwrap_err()
.to_string()
.contains("upstream_proxy_remote_dns"));
}
for proxy in ["http://127.0.0.1:8080", "socks5h://127.0.0.1:1080"] {
config.upstream_proxy_url = Some(proxy.to_string());
config
.validate()
.expect("explicit remote DNS configuration should validate");
}
}
#[test]
fn config_file_round_trips_proxy_remote_dns() {
let config = parse_config_file_content(
"upstream_proxy_url = \"socks5h://127.0.0.1:1080\"\nupstream_proxy_remote_dns = true",
)
.unwrap();
assert_eq!(config.upstream_proxy_remote_dns, Some(true));
let round_trip: ConfigFile = toml::from_str(&toml::to_string(&config).unwrap()).unwrap();
assert_eq!(round_trip.upstream_proxy_remote_dns, Some(true));
assert_eq!(ConfigFile::default().upstream_proxy_remote_dns, None);
}
fn config_save_test_dir(label: &str) -> std::path::PathBuf {
let path = std::env::temp_dir().join(format!(
"aether-tunnel-config-{label}-{}",
+21 -1
View File
@@ -142,6 +142,13 @@ impl UpstreamProxyConfig {
self.scheme == UpstreamProxyScheme::Socks5h
}
pub(crate) fn supports_remote_target_dns(&self) -> bool {
matches!(
self.scheme,
UpstreamProxyScheme::Http | UpstreamProxyScheme::Socks5h
)
}
pub(crate) fn basic_auth_header(&self) -> Option<String> {
let username = self.username()?;
let mut credentials = String::with_capacity(
@@ -445,7 +452,7 @@ pub(crate) async fn socks5_target_address(
remote_dns: bool,
) -> io::Result<Vec<u8>> {
let mut request = vec![0x05, 0x01, 0x00];
if let Ok(ip) = target_host.parse::<IpAddr>() {
if let Some(ip) = aether_http::parse_ip_literal_host(target_host) {
push_socks5_ip_address(&mut request, ip);
} else if remote_dns {
let host = target_host.as_bytes();
@@ -524,6 +531,19 @@ fn non_empty_url_part(value: &str) -> Option<String> {
mod tests {
use super::*;
#[tokio::test]
async fn socks_proxy_encodes_bracketed_ipv6_as_an_ip_for_both_dns_modes() {
for remote_dns in [false, true] {
let expected = socks5_target_address("::1", 443, remote_dns).await.unwrap();
let actual = socks5_target_address("[::1]", 443, remote_dns)
.await
.unwrap();
assert_eq!(actual, expected);
assert_eq!(&actual[..4], &[5, 1, 0, 4]);
assert_eq!(&actual[20..], &443u16.to_be_bytes());
}
}
#[test]
fn parses_http_proxy_with_default_port() {
let proxy = UpstreamProxyConfig::parse("http://proxy.example").expect("proxy should parse");
+55 -1
View File
@@ -214,6 +214,14 @@ impl App {
required: false,
help: "Heartbeat interval in seconds; default is 5",
},
Field {
label: "Proxy Remote DNS",
key: "upstream_proxy_remote_dns",
value: "false".into(),
kind: FieldKind::Bool,
required: false,
help: "Trust HTTP/SOCKS5h egress proxy to resolve provider hosts and enforce destination IP ACLs; restart required",
},
],
selected: 0,
mode: Mode::Normal,
@@ -288,6 +296,9 @@ impl App {
"allow_private_targets" => cfg.allow_private_targets.map(|v| v.to_string()),
"heartbeat_interval" => cfg.heartbeat_interval.map(|v| v.to_string()),
"upstream_proxy_url" => cfg.upstream_proxy_url.clone(),
"upstream_proxy_remote_dns" => {
cfg.upstream_proxy_remote_dns.map(|value| value.to_string())
}
_ => None,
};
if let Some(v) = val {
@@ -396,11 +407,23 @@ impl App {
let get_tab = |tab: &ServerTab, key: &str| -> Option<String> { Self::get_tab(tab, key) };
let save_logs_to_file = self.toggle_enabled("save_logs_to_file");
let upstream_proxy_url = self.parse_optional_upstream_proxy_url()?;
let upstream_proxy_remote_dns = self.toggle_enabled("upstream_proxy_remote_dns");
if upstream_proxy_remote_dns {
let proxy_url = upstream_proxy_url
.as_deref()
.ok_or_else(|| anyhow::anyhow!("Proxy Remote DNS requires an egress proxy"))?;
let proxy = UpstreamProxyConfig::parse(proxy_url).map_err(anyhow::Error::msg)?;
if !proxy.supports_remote_target_dns() {
anyhow::bail!("Proxy Remote DNS requires an http:// or socks5h:// proxy");
}
}
let mut cfg = ConfigFile {
log_level: get_global("log_level"),
allow_private_targets: Some(self.toggle_enabled("allow_private_targets")),
heartbeat_interval: self.parse_optional_heartbeat_interval()?,
upstream_proxy_url: self.parse_optional_upstream_proxy_url()?,
upstream_proxy_url,
upstream_proxy_remote_dns: Some(upstream_proxy_remote_dns),
log_destination: Some(if save_logs_to_file {
TunnelLogDestinationArg::Both
} else {
@@ -1086,6 +1109,37 @@ mod tests {
app
}
#[test]
fn proxy_remote_dns_toggle_round_trips_and_requires_a_trusted_proxy() {
let mut app = sample_app();
assert_eq!(
app.to_config().unwrap().upstream_proxy_remote_dns,
Some(false)
);
set_global_field(&mut app, "upstream_proxy_remote_dns", "true");
assert!(app
.to_config()
.unwrap_err()
.to_string()
.contains("requires an egress proxy"));
set_global_field(&mut app, "upstream_proxy_url", "socks5://127.0.0.1:1080");
assert!(app
.to_config()
.unwrap_err()
.to_string()
.contains("socks5h://"));
for proxy in ["http://127.0.0.1:8080", "socks5h://127.0.0.1:1080"] {
set_global_field(&mut app, "upstream_proxy_url", proxy);
let config = app.to_config().unwrap();
let mut restored = sample_app();
restored.apply_config(&config);
let round_trip = restored.to_config().unwrap();
assert_eq!(round_trip.upstream_proxy_remote_dns, Some(true));
assert_eq!(round_trip.upstream_proxy_url.as_deref(), Some(proxy));
}
}
fn unique_temp_config_path(name: &str) -> PathBuf {
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
+19 -6
View File
@@ -226,21 +226,34 @@ pub async fn validate_target(
allow_private: bool,
dns_cache: &DnsCache,
) -> Result<Vec<SocketAddr>, FilterError> {
// Port whitelist check
if let Some(address) = validate_target_literal(host, port, allowed_ports, allow_private)? {
return Ok(vec![address]);
}
resolve_public_addrs(host, port, allow_private, dns_cache).await
}
pub(crate) fn validate_target_literal(
host: &str,
port: u16,
allowed_ports: &HashSet<u16>,
allow_private: bool,
) -> Result<Option<SocketAddr>, FilterError> {
if !allowed_ports.contains(&port) {
return Err(FilterError::PortNotAllowed(port));
}
// Try parsing as IP directly (no DNS needed)
if let Ok(ip) = host.parse::<IpAddr>() {
if let Some(ip) = aether_http::parse_ip_literal_host(host) {
if !allow_private && is_private_ip(&ip) {
return Err(FilterError::PrivateIp(ip));
}
return Ok(vec![SocketAddr::new(ip, port)]);
return Ok(Some(SocketAddr::new(ip, port)));
}
// Resolve and return the exact addresses authorized for this request.
resolve_public_addrs(host, port, allow_private, dns_cache).await
if !allow_private && host.trim_end_matches('.').eq_ignore_ascii_case("localhost") {
return Err(FilterError::NoPublicAddrs(host.to_string()));
}
Ok(None)
}
#[cfg(test)]
+1
View File
@@ -737,6 +737,7 @@ mod tests {
upstream_tcp_keepalive_secs: 60,
upstream_tcp_nodelay: true,
upstream_proxy_url: None,
upstream_proxy_remote_dns: false,
legacy_redirect_replay_budget_bytes_ignored: None,
emit_proxy_timing_header: true,
log_level: "info".to_string(),
+123 -13
View File
@@ -1276,6 +1276,40 @@ fn resolve_redirect<B>(
}
}
async fn resolve_upstream_target(
current_url: &url::Url,
allowed_ports: &std::collections::HashSet<u16>,
allow_private_targets: bool,
proxy_remote_dns: bool,
dns_cache: &target_filter::DnsCache,
) -> Result<upstream_client::ValidatedUpstreamTarget, String> {
validate_tunnel_upstream_url(current_url, allow_private_targets).map_err(str::to_string)?;
let host = current_url
.host_str()
.ok_or_else(|| "missing host in URL".to_string())?;
let port = current_url
.port_or_known_default()
.ok_or_else(|| "missing port in URL".to_string())?;
let addresses = if proxy_remote_dns {
match target_filter::validate_target_literal(
host,
port,
allowed_ports,
allow_private_targets,
)
.map_err(|_| "upstream target blocked".to_string())?
{
Some(address) => vec![address],
None => return upstream_client::ValidatedUpstreamTarget::proxy_resolved(current_url),
}
} else {
target_filter::validate_target(host, port, allowed_ports, allow_private_targets, dns_cache)
.await
.map_err(|_| "upstream target blocked".to_string())?
};
upstream_client::ValidatedUpstreamTarget::new(current_url, addresses)
}
#[allow(clippy::too_many_arguments)]
async fn execute_upstream_request(
state: &AppState,
@@ -1288,24 +1322,19 @@ async fn execute_upstream_request(
timeout: Duration,
http1_only: bool,
) -> Result<UpstreamResponseContext, String> {
let host = current_url
.host_str()
.ok_or_else(|| "missing host in URL".to_string())?;
let port = current_url.port_or_known_default().unwrap_or(443);
let dns_start = Instant::now();
let validated_addrs = {
let validated_target = {
let allowed_ports = Arc::clone(&server.dynamic.load().allowed_ports);
match target_filter::validate_target(
host,
port,
match resolve_upstream_target(
current_url,
&allowed_ports,
state.config.allow_private_targets,
state.config.upstream_proxy_remote_dns,
&state.dns_cache,
)
.await
{
Ok(addrs) => addrs,
Ok(target) => target,
Err(_error) => {
server.metrics.dns_failures.fetch_add(1, Ordering::Release);
// Keep the detailed filter error out of the tunnel response;
@@ -1317,9 +1346,6 @@ async fn execute_upstream_request(
};
let dns_ms = dns_start.elapsed().as_millis() as u64;
let validated_target =
upstream_client::ValidatedUpstreamTarget::new(current_url, validated_addrs)?;
let client_key = upstream_client::upstream_client_pool_key(
meta.provider_id.as_deref(),
meta.endpoint_id.as_deref(),
@@ -2326,6 +2352,89 @@ fn build_prefixed_request_body(
#[cfg(test)]
mod tests {
#[tokio::test]
async fn remote_dns_target_resolution_skips_local_dns_and_keeps_literal_acl() {
let cache = target_filter::DnsCache::new(Duration::from_secs(60), 16);
let ports = [80, 443].into_iter().collect();
let url = url::Url::parse("https://remote-dns-test.invalid/path").unwrap();
let target = resolve_upstream_target(&url, &ports, false, true, &cache)
.await
.expect("trusted proxy should receive an unresolved hostname");
assert!(target.uses_proxy_dns());
assert!(cache.get("remote-dns-test.invalid", 443).await.is_none());
for address in ["https://8.8.8.8/", "https://[2606:4700:4700::1111]/"] {
let target = resolve_upstream_target(
&url::Url::parse(address).unwrap(),
&ports,
false,
true,
&cache,
)
.await
.unwrap();
assert!(!target.uses_proxy_dns(), "IP literals must remain pinned");
}
for address in [
"https://127.0.0.1/",
"https://10.0.0.1/",
"https://198.18.0.1/",
"https://[::1]/",
"https://[::ffff:127.0.0.1]/",
"https://localhost/",
"https://LOCALHOST./",
"https://remote-dns-test.invalid:25/",
"https://user:[email protected]/",
"https://remote-dns-test.invalid/#fragment",
"ftp://remote-dns-test.invalid/",
] {
assert!(
resolve_upstream_target(
&url::Url::parse(address).unwrap(),
&ports,
false,
true,
&cache,
)
.await
.is_err(),
"target should remain blocked: {address}"
);
}
}
#[tokio::test]
async fn strict_dns_targets_stay_pinned_and_separate_from_remote_dns_targets() {
let cache = target_filter::DnsCache::new(Duration::from_secs(60), 16);
let ports = [443].into_iter().collect();
let url = url::Url::parse("https://remote-dns-test.invalid/").unwrap();
cache
.insert(
"remote-dns-test.invalid",
443,
Arc::new(vec!["8.8.8.8:443".parse().unwrap()]),
)
.await;
let strict = resolve_upstream_target(&url, &ports, false, false, &cache)
.await
.unwrap();
let remote = resolve_upstream_target(&url, &ports, false, true, &cache)
.await
.unwrap();
assert!(!strict.uses_proxy_dns());
assert!(remote.uses_proxy_dns());
assert_ne!(strict, remote);
let private_url = url::Url::parse("http://[::1]/").unwrap();
let private_ports = [80].into_iter().collect();
let explicitly_allowed =
resolve_upstream_target(&private_url, &private_ports, true, true, &cache)
.await
.unwrap();
assert!(!explicitly_allowed.uses_proxy_dns());
}
#[tokio::test(start_paused = true)]
async fn window_updates_wait_for_capacity_instead_of_disappearing() {
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(1);
@@ -3820,6 +3929,7 @@ mod tests {
upstream_tcp_keepalive_secs: 60,
upstream_tcp_nodelay: true,
upstream_proxy_url: None,
upstream_proxy_remote_dns: false,
legacy_redirect_replay_budget_bytes_ignored: None,
emit_proxy_timing_header: true,
log_level: "info".to_string(),
+350 -42
View File
@@ -34,7 +34,8 @@ use tower_service::Service;
use crate::config::Config;
use crate::egress_proxy::{
connect_validated_target_via_proxy, ProxyConnectOptions, UpstreamProxyConfig,
connect_target_via_proxy, connect_validated_target_via_proxy, ProxyConnectOptions,
UpstreamProxyConfig,
};
use crate::target_filter::DnsCache;
@@ -66,11 +67,42 @@ pub struct ValidatedUpstreamTarget {
scheme: String,
host: String,
port: u16,
addrs: Vec<SocketAddr>,
resolution: UpstreamTargetResolution,
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
enum UpstreamTargetResolution {
Pinned(Vec<SocketAddr>),
ProxyDns,
}
impl ValidatedUpstreamTarget {
pub fn new(target_url: &url::Url, mut addrs: Vec<SocketAddr>) -> Result<Self, String> {
if addrs.is_empty() {
return Err("validated upstream target has no addresses".to_string());
}
addrs.sort_unstable();
addrs.dedup();
Self::with_resolution(target_url, UpstreamTargetResolution::Pinned(addrs))
}
pub(crate) fn proxy_resolved(target_url: &url::Url) -> Result<Self, String> {
if !matches!(target_url.host(), Some(url::Host::Domain(_))) {
return Err("IP literal targets must use pinned addresses".to_string());
}
Self::with_resolution(target_url, UpstreamTargetResolution::ProxyDns)
}
fn with_resolution(
target_url: &url::Url,
resolution: UpstreamTargetResolution,
) -> Result<Self, String> {
if !target_url.username().is_empty()
|| target_url.password().is_some()
|| target_url.fragment().is_some()
{
return Err("upstream target must not contain credentials or a fragment".to_string());
}
let scheme = target_url.scheme().to_ascii_lowercase();
if !matches!(scheme.as_str(), "http" | "https") {
return Err(format!("unsupported upstream scheme {scheme}"));
@@ -86,19 +118,16 @@ impl ValidatedUpstreamTarget {
let port = target_url
.port_or_known_default()
.ok_or_else(|| "missing port in upstream URL".to_string())?;
if addrs.is_empty() {
return Err("validated upstream target has no addresses".to_string());
if let UpstreamTargetResolution::Pinned(addrs) = &resolution {
if addrs.iter().any(|addr| addr.port() != port) {
return Err("validated upstream target address has the wrong port".to_string());
}
}
if addrs.iter().any(|addr| addr.port() != port) {
return Err("validated upstream target address has the wrong port".to_string());
}
addrs.sort_unstable();
addrs.dedup();
Ok(Self {
scheme,
host,
port,
addrs,
resolution,
})
}
@@ -119,8 +148,8 @@ impl ValidatedUpstreamTarget {
Ok(())
}
fn addrs(&self) -> &[SocketAddr] {
&self.addrs
pub(crate) fn uses_proxy_dns(&self) -> bool {
matches!(self.resolution, UpstreamTargetResolution::ProxyDns)
}
}
@@ -333,9 +362,14 @@ impl Service<Name> for PinnedResolver {
"DNS request does not match the validated upstream host",
));
}
Ok(ValidatedAddrs {
inner: target.addrs.into_iter(),
})
match target.resolution {
UpstreamTargetResolution::Pinned(addrs) => Ok(ValidatedAddrs {
inner: addrs.into_iter(),
}),
UpstreamTargetResolution::ProxyDns => Err(io::Error::other(
"proxy-resolved target must not fall back to local DNS",
)),
}
})
}
}
@@ -376,16 +410,25 @@ impl Service<Uri> for InstrumentedConnector {
};
let connect_start = std::time::Instant::now();
return Box::pin(async move {
connect_via_proxy(
dst,
scheme,
tls_config,
proxy,
validated_target,
options,
connect_start,
tokio::time::timeout(
options.connect_timeout,
connect_via_proxy(
dst,
scheme,
tls_config,
proxy,
validated_target,
options,
connect_start,
),
)
.await
.map_err(|_| {
Box::new(io::Error::new(
io::ErrorKind::TimedOut,
"upstream proxy connection timed out",
)) as BoxError
})?
});
}
let connecting = self.http.call(dst.clone());
@@ -441,20 +484,40 @@ async fn connect_via_proxy(
connect_start: std::time::Instant,
) -> Result<TimedConn, BoxError> {
let scheme = scheme.ok_or_else(|| io::Error::other("missing scheme"))?;
let mut last_error = None;
let mut connected = None;
for target_addr in validated_target.addrs().iter().copied() {
match connect_validated_target_via_proxy(&proxy, target_addr, options).await {
Ok(tcp) => {
connected = Some(tcp);
break;
let tcp = match &validated_target.resolution {
UpstreamTargetResolution::ProxyDns => {
if !proxy.supports_remote_target_dns() {
return Err(
io::Error::other("upstream proxy does not support remote target DNS").into(),
);
}
Err(error) => last_error = Some(error),
connect_target_via_proxy(
&proxy,
&validated_target.host,
validated_target.port,
options,
)
.await?
}
}
let tcp = connected.ok_or_else(|| {
last_error.unwrap_or_else(|| io::Error::other("validated upstream target has no addresses"))
})?;
UpstreamTargetResolution::Pinned(addrs) => {
let mut last_error = None;
let mut connected = None;
for target_addr in addrs.iter().copied() {
match connect_validated_target_via_proxy(&proxy, target_addr, options).await {
Ok(tcp) => {
connected = Some(tcp);
break;
}
Err(error) => last_error = Some(error),
}
}
connected.ok_or_else(|| {
last_error.unwrap_or_else(|| {
io::Error::other("validated upstream target has no addresses")
})
})?
}
};
let connect_ms = connect_start.elapsed().as_millis() as u64;
@@ -512,6 +575,24 @@ fn build_upstream_client_with_protocol(
http1_only: bool,
h2c_prior_knowledge: bool,
) -> Result<UpstreamClient, String> {
let proxy = config
.upstream_proxy_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(UpstreamProxyConfig::parse)
.transpose()?;
if validated_target.uses_proxy_dns()
&& (!config.upstream_proxy_remote_dns
|| !proxy
.as_ref()
.is_some_and(UpstreamProxyConfig::supports_remote_target_dns))
{
return Err(
"proxy-resolved upstream requires explicit remote DNS and an HTTP or SOCKS5h proxy"
.to_string(),
);
}
let mut http = HttpConnector::new_with_resolver(PinnedResolver::new(validated_target.clone()));
http.enforce_http(false);
http.set_connect_timeout(Some(Duration::from_secs(
@@ -530,13 +611,7 @@ fn build_upstream_client_with_protocol(
http,
tls_config: build_tls_config(http1_only),
validated_target,
proxy: config
.upstream_proxy_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(UpstreamProxyConfig::parse)
.transpose()?,
proxy,
connect_timeout: Duration::from_secs(config.upstream_connect_timeout_secs),
tcp_nodelay: config.upstream_tcp_nodelay,
tcp_keepalive: (config.upstream_tcp_keepalive_secs > 0)
@@ -1035,6 +1110,239 @@ mod tests {
);
}
#[tokio::test]
async fn trusted_http_proxy_resolves_hostname_without_local_dns() {
let (proxy_url, connect_rx, request_rx) = spawn_http_proxy().await;
let client = remote_dns_client(&proxy_url, "http://remote-dns-test.invalid/");
let request = hyper::Request::builder()
.uri("http://remote-dns-test.invalid/remote-dns")
.body(full_request_body(Bytes::new()))
.unwrap();
let response = tokio::time::timeout(Duration::from_secs(5), client.request(request))
.await
.unwrap()
.expect("proxy should resolve the target without local DNS");
assert_eq!(response.status(), hyper::StatusCode::OK);
assert_eq!(
response.into_body().collect().await.unwrap().to_bytes(),
"ok"
);
assert!(connect_rx
.await
.unwrap()
.starts_with("CONNECT remote-dns-test.invalid:80 HTTP/1.1\r\n"));
let request = request_rx.await.unwrap().to_ascii_lowercase();
assert!(request.starts_with("get /remote-dns http/1.1\r\n"));
assert!(request.contains("\r\nhost: remote-dns-test.invalid\r\n"));
}
#[tokio::test]
async fn trusted_socks5h_proxy_receives_hostname_not_a_locally_resolved_ip() {
let (proxy_url, target_rx, request_rx) = spawn_remote_dns_socks_proxy().await;
let client = remote_dns_client(&proxy_url, "http://remote-dns-test.invalid/");
let request = hyper::Request::builder()
.uri("http://remote-dns-test.invalid/remote-dns")
.body(full_request_body(Bytes::new()))
.unwrap();
let response = tokio::time::timeout(Duration::from_secs(5), client.request(request))
.await
.unwrap()
.expect("SOCKS proxy should receive the unresolved target");
assert_eq!(response.status(), hyper::StatusCode::OK);
assert_eq!(
response.into_body().collect().await.unwrap().to_bytes(),
"ok"
);
assert_eq!(
target_rx.await.unwrap(),
("remote-dns-test.invalid".to_string(), 80)
);
assert!(request_rx
.await
.unwrap()
.to_ascii_lowercase()
.contains("\r\nhost: remote-dns-test.invalid\r\n"));
}
#[tokio::test]
async fn remote_dns_https_preserves_hostname_for_connect_and_sni() {
let (proxy_url, connect_rx) = spawn_connect_only_http_proxy().await;
let client = remote_dns_client(&proxy_url, "https://remote-dns-test.invalid/");
let uri: Uri = "https://remote-dns-test.invalid/secure".parse().unwrap();
let request = hyper::Request::builder()
.uri(uri.clone())
.body(full_request_body(Bytes::new()))
.unwrap();
let _ = tokio::time::timeout(Duration::from_secs(5), client.request(request))
.await
.unwrap();
assert!(connect_rx
.await
.unwrap()
.starts_with("CONNECT remote-dns-test.invalid:443 HTTP/1.1\r\n"));
match resolve_server_name(&uri).unwrap() {
ServerName::DnsName(name) => assert_eq!(name.as_ref(), "remote-dns-test.invalid"),
other => panic!("expected hostname for TLS verification, got {other:?}"),
}
}
#[tokio::test]
async fn remote_dns_targets_cannot_fall_back_to_local_dns_or_change_origin() {
let url = url::Url::parse("https://remote-dns-test.invalid/").unwrap();
let target = ValidatedUpstreamTarget::proxy_resolved(&url).unwrap();
let mut resolver = PinnedResolver::new(target.clone());
let error = resolver
.call("remote-dns-test.invalid".parse().unwrap())
.await
.err()
.unwrap();
assert!(error
.to_string()
.contains("must not fall back to local DNS"));
for uri in [
"http://remote-dns-test.invalid/",
"https://another-target.invalid/",
"https://remote-dns-test.invalid:8443/",
] {
assert!(target.ensure_matches_uri(&uri.parse().unwrap()).is_err());
}
for url in ["http://127.0.0.1/", "https://[::1]/"] {
assert!(
ValidatedUpstreamTarget::proxy_resolved(&url::Url::parse(url).unwrap()).is_err()
);
}
}
#[test]
fn remote_dns_clients_require_opt_in_and_do_not_share_pinned_pool_entries() {
let mut config = remote_dns_config("http://127.0.0.1:8080");
let url = url::Url::parse("https://remote-dns-test.invalid/").unwrap();
let remote = ValidatedUpstreamTarget::proxy_resolved(&url).unwrap();
config.upstream_proxy_remote_dns = false;
assert!(build_upstream_client_with_protocol(&config, remote.clone(), true, false).is_err());
config.upstream_proxy_remote_dns = true;
for proxy in [None, Some("socks5://127.0.0.1:1080")] {
config.upstream_proxy_url = proxy.map(str::to_string);
assert!(
build_upstream_client_with_protocol(&config, remote.clone(), true, false).is_err()
);
}
config.upstream_proxy_url = Some("http://127.0.0.1:8080".to_string());
let pinned =
ValidatedUpstreamTarget::new(&url, vec!["8.8.8.8:443".parse().unwrap()]).unwrap();
let remote_key = upstream_client_pool_key(None, None, None, None, false, remote);
let pinned_key = upstream_client_pool_key(None, None, None, None, false, pinned);
assert_ne!(remote_key, pinned_key);
let pool = UpstreamClientPool::new(
Arc::new(config),
Arc::new(DnsCache::new(Duration::from_secs(60), 16)),
);
pool.get_or_build(remote_key).unwrap();
pool.get_or_build(pinned_key).unwrap();
assert_eq!(pool.clients.lock().unwrap().len(), 2);
}
#[tokio::test]
async fn remote_dns_proxy_connect_timeout_covers_connect_and_tls_handshakes() {
for tls in [false, true] {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_url = format!("http://{}", listener.local_addr().unwrap());
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
read_http_headers(&mut stream).await;
if tls {
stream
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await
.unwrap();
}
std::future::pending::<()>().await;
drop(stream);
});
let target_url = if tls {
"https://remote-dns-test.invalid/"
} else {
"http://remote-dns-test.invalid/"
};
let client = remote_dns_client(&proxy_url, target_url);
let request = hyper::Request::builder()
.uri(target_url)
.body(full_request_body(Bytes::new()))
.unwrap();
let result =
tokio::time::timeout(Duration::from_secs(5), client.request(request)).await;
server.abort();
let error = result
.expect("configured connect timeout must include proxy and TLS handshakes")
.unwrap_err();
assert!(error.is_connect());
}
}
fn remote_dns_config(proxy_url: &str) -> Config {
let _ = rustls::crypto::ring::default_provider().install_default();
Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
"--upstream-proxy-url",
proxy_url,
"--upstream-proxy-remote-dns",
"--upstream-connect-timeout-secs",
"1",
])
}
fn remote_dns_client(proxy_url: &str, target_url: &str) -> UpstreamClient {
let config = remote_dns_config(proxy_url);
let target =
ValidatedUpstreamTarget::proxy_resolved(&url::Url::parse(target_url).unwrap()).unwrap();
build_upstream_client_with_protocol(&config, target, true, false).unwrap()
}
async fn spawn_remote_dns_socks_proxy() -> (
String,
tokio::sync::oneshot::Receiver<(String, u16)>,
tokio::sync::oneshot::Receiver<String>,
) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_url = format!("socks5h://{}", listener.local_addr().unwrap());
let (target_tx, target_rx) = tokio::sync::oneshot::channel();
let (request_tx, request_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut greeting = [0u8; 3];
stream.read_exact(&mut greeting).await.unwrap();
assert_eq!(greeting, [0x05, 0x01, 0x00]);
stream.write_all(&[0x05, 0x00]).await.unwrap();
let mut header = [0u8; 5];
stream.read_exact(&mut header).await.unwrap();
assert_eq!(&header[..4], &[0x05, 0x01, 0x00, 0x03]);
let mut hostname = vec![0; header[4] as usize];
stream.read_exact(&mut hostname).await.unwrap();
let port = stream.read_u16().await.unwrap();
target_tx
.send((String::from_utf8(hostname).unwrap(), port))
.unwrap();
stream
.write_all(&[0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0])
.await
.unwrap();
request_tx
.send(read_http_headers(&mut stream).await)
.unwrap();
stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
.await
.unwrap();
});
(proxy_url, target_rx, request_rx)
}
fn proxied_client(
proxy_url: &str,
target_url: &str,
@@ -312,6 +312,10 @@ pub fn build_admin_monitoring_trace_request_payload_response_with_key_accounts(
"total_candidates": trace.total_candidates,
"final_status": trace.final_status,
"total_latency_ms": trace.total_latency_ms,
"diagnostic_request": usage.filter(|usage| admin_monitoring_usage_matches_trace(usage, &trace.request_id)).map(|usage| json!({
"usage_id": usage.id,
"body_state": usage.request_body_state.map(|state| state.as_str()),
})),
"candidates": candidates,
}))
.into_response()
@@ -337,6 +341,7 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts(
item.sanitize_for_admin();
let candidate = &item.candidate;
let sanitized_extra_data = build_admin_monitoring_trace_candidate_extra_data(
&candidate.id,
candidate.extra_data.as_ref(),
candidate.status_code,
usage,
@@ -518,6 +523,7 @@ fn build_admin_monitoring_trace_candidate_ranking(existing: Option<&Value>) -> V
}
fn build_admin_monitoring_trace_candidate_extra_data(
candidate_id: &str,
existing: Option<&Value>,
candidate_status_code: Option<u16>,
usage: Option<&StoredRequestUsageAudit>,
@@ -527,6 +533,19 @@ fn build_admin_monitoring_trace_candidate_extra_data(
if let Some(usage) = usage {
let extra_object = extra_data.get_or_insert_with(serde_json::Map::new);
if usage.routing_candidate_id() == Some(candidate_id) {
extra_object.insert("diagnostic_context".to_string(), json!({
"usage_id": usage.id,
"model": usage.model,
"target_model": usage.target_model,
"body_states": {
"request_body": usage.request_body_state.map(|state| state.as_str()),
"provider_request_body": usage.provider_request_body_state.map(|state| state.as_str()),
"response_body": usage.response_body_state.map(|state| state.as_str()),
"client_response_body": usage.client_response_body_state.map(|state| state.as_str()),
}
}));
}
if let Some(first_byte_time_ms) = usage.first_byte_time_ms {
extra_object
.entry("first_byte_time_ms".to_string())
@@ -577,8 +596,21 @@ fn build_admin_monitoring_trace_candidate_extra_data(
}
}
sanitize_request_candidate_extra_data_for_persistence(extra_data.map(Value::Object))
.unwrap_or(Value::Null)
let diagnostic_context = extra_data
.as_mut()
.and_then(|extra| extra.remove("diagnostic_context"));
let mut sanitized =
sanitize_request_candidate_extra_data_for_persistence(extra_data.map(Value::Object))
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
if let Some(context) = diagnostic_context {
sanitized.insert("diagnostic_context".to_string(), context);
}
if sanitized.is_empty() {
Value::Null
} else {
Value::Object(sanitized)
}
}
fn admin_monitoring_trace_response_data(
@@ -3398,6 +3398,60 @@ mod tests {
));
}
#[test]
fn detail_payload_preserves_original_headers_with_or_without_bodies() {
let item = StoredRequestUsageAudit {
request_headers: Some(json!({
"authorization": "Bearer original-client-token",
"originator": "codex-cli",
"session-id": "original-session",
"thread-id": "original-thread",
"x-codex-turn-metadata": "{\"turn_id\":\"original-turn\"}",
"x-openai-subagent": "reviewer"
})),
provider_request_headers: Some(json!({
"authorization": "Bearer original-provider-token",
"x-api-key": "original-provider-key"
})),
response_headers: Some(json!({
"set-cookie": ["session=original-upstream", "preference=original"],
"x-upstream-custom": "original-upstream-value"
})),
client_response_headers: Some(json!({
"set-cookie": "session=original-client",
"x-client-custom": "original-client-value"
})),
..sample_usage("completed", Some(200), None)
};
for include_bodies in [false, true] {
let payload = build_admin_usage_detail_payload(
&item,
&BTreeMap::new(),
&BTreeMap::new(),
false,
false,
None,
include_bodies,
None,
&BTreeMap::new(),
);
for (field, expected) in [
("request_headers", &item.request_headers),
("provider_request_headers", &item.provider_request_headers),
("response_headers", &item.response_headers),
("client_response_headers", &item.client_response_headers),
] {
assert_eq!(
&payload[field],
expected.as_ref().unwrap(),
"{field} should retain original values when include_bodies={include_bodies}"
);
}
}
}
#[test]
fn detail_payload_marks_reference_backed_bodies_as_available() {
let item = StoredRequestUsageAudit {
+34 -1
View File
@@ -20,7 +20,7 @@ pub fn build_kiro_batch_import_key_name(
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>();
format!("kiro_{}", &hex[..6])
format!("账号_{}", &hex[..6])
});
format!("{base} ({method})")
}
@@ -130,3 +130,36 @@ pub fn parse_admin_provider_oauth_kiro_batch_import_entries(raw_credentials: &st
.map(|refresh_token| json!({ "refreshToken": refresh_token }))
.collect()
}
#[cfg(test)]
mod tests {
use super::build_kiro_batch_import_key_name;
#[test]
fn kiro_batch_import_key_name_preserves_email_and_auth_method() {
for method in ["social", "idc"] {
assert_eq!(
build_kiro_batch_import_key_name(
Some(" [email protected] "),
Some(method),
Some("refresh-token-1"),
),
format!("[email protected] ({method})")
);
}
}
#[test]
fn kiro_batch_import_key_name_without_email_uses_generic_account_prefix() {
for email in [None, Some(""), Some(" ")] {
assert_eq!(
build_kiro_batch_import_key_name(email, None, Some("refresh-token-1")),
"账号_154f43 (social)"
);
assert_eq!(
build_kiro_batch_import_key_name(email, Some("idc"), Some("refresh-token-1")),
"账号_154f43 (idc)"
);
}
}
}
+24 -2
View File
@@ -405,18 +405,40 @@ pub fn build_kiro_device_key_name(email: Option<&str>, refresh_token: Option<&st
.collect::<String>()
})
.unwrap_or_else(|| "unknown".to_string());
format!("kiro_{fallback} (idc)")
format!("账号_{fallback} (idc)")
}
#[cfg(test)]
mod tests {
use super::{
decode_jwt_claims, enrich_admin_provider_oauth_auth_config,
build_kiro_device_key_name, decode_jwt_claims, enrich_admin_provider_oauth_auth_config,
parse_provider_oauth_callback_params, MAX_UNVERIFIED_JWT_CLAIMS_BYTES,
};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
#[test]
fn kiro_device_key_name_preserves_email_and_auth_method() {
assert_eq!(
build_kiro_device_key_name(Some(" [email protected] "), Some("refresh-token-1")),
"[email protected] (idc)"
);
}
#[test]
fn kiro_device_key_name_without_email_uses_generic_account_prefix() {
for email in [None, Some(""), Some(" ")] {
assert_eq!(
build_kiro_device_key_name(email, Some("refresh-token-1")),
"账号_154f43 (idc)"
);
assert_eq!(
build_kiro_device_key_name(email, None),
"账号_unknown (idc)"
);
}
}
fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(payload.to_string());
@@ -238,3 +238,124 @@ impl fmt::Display for FormatError {
}
impl Error for FormatError {}
impl FormatError {
pub fn diagnostic(&self) -> Value {
let (code, operation, field, reason) = match self {
Self::UnsupportedFormat(_) => ("unsupported_format", "select_format", None, None),
Self::RequestParseFailed { .. } => {
("request_parse_failed", "parse_request", None, None)
}
Self::RequestEmitFailed { .. } => ("request_emit_failed", "emit_request", None, None),
Self::ResponseParseFailed { .. } => {
("response_parse_failed", "parse_response", None, None)
}
Self::ResponseEmitFailed { .. } => {
("response_emit_failed", "emit_response", None, None)
}
Self::UnsupportedField { field, reason, .. } => {
("unsupported_field", "validate", Some(field), Some(reason))
}
Self::UnauditedField { field, reason, .. } => {
("unaudited_field", "validate", Some(field), Some(reason))
}
Self::InvalidEnumValue { field, .. } => {
("invalid_enum_value", "validate", Some(field), None)
}
Self::LossyConversionBlocked { field, reason, .. } => (
"lossy_conversion_blocked",
"convert",
Some(field),
Some(reason),
),
Self::InvalidTargetField { field, reason, .. } => (
"invalid_target_field",
"validate_target",
Some(field),
Some(reason),
),
};
let path = field.map(|field| {
let path = if field.starts_with('$') {
field.clone()
} else {
format!("$.{field}")
};
path.replace("[]", "[*]")
});
let mut diagnostic = json!({
"code": code,
"operation": operation,
"path": path.as_deref().unwrap_or("$"),
"path_source": if path.is_some() { "structured" } else { "unavailable" },
"reason": reason,
"expected": reason,
"actual": null,
"missing_context": if path.is_some() { vec![] } else { vec!["field_path", "underlying_cause"] }
});
if let Self::InvalidEnumValue { value, .. } = self {
diagnostic["actual"] = json!(value);
}
match self {
Self::UnsupportedFormat(format)
| Self::RequestParseFailed { format }
| Self::RequestEmitFailed { format }
| Self::ResponseParseFailed { format }
| Self::ResponseEmitFailed { format }
| Self::UnsupportedField { format, .. }
| Self::InvalidEnumValue { format, .. }
| Self::InvalidTargetField { format, .. } => diagnostic["format"] = json!(format),
_ => {}
}
if let Self::UnauditedField {
source_format,
target_format,
..
}
| Self::LossyConversionBlocked {
source_format,
target_format,
..
} = self
{
diagnostic["source_format"] = json!(source_format);
diagnostic["target_format"] = json!(target_format);
}
diagnostic
}
}
#[cfg(test)]
mod diagnostic_tests {
use super::FormatError;
use serde_json::json;
#[test]
fn enum_diagnostic_retains_full_path_and_actual_value() {
let diagnostic = FormatError::InvalidEnumValue {
format: "openai:chat".to_string(),
field: "choices[].finish_reason".to_string(),
value: "future_reason".to_string(),
}
.diagnostic();
assert_eq!(diagnostic["code"], "invalid_enum_value");
assert_eq!(diagnostic["path"], "$.choices[*].finish_reason");
assert_eq!(diagnostic["actual"], "future_reason");
assert_eq!(diagnostic["format"], "openai:chat");
}
#[test]
fn generic_parse_failure_reports_missing_cause_without_a_fake_path() {
let diagnostic = FormatError::ResponseParseFailed {
format: "claude:messages".to_string(),
}
.diagnostic();
assert_eq!(diagnostic["code"], "response_parse_failed");
assert_eq!(diagnostic["operation"], "parse_response");
assert_eq!(diagnostic["path_source"], "unavailable");
assert_eq!(
diagnostic["missing_context"],
json!(["field_path", "underlying_cause"])
);
}
}
@@ -23,6 +23,9 @@ use crate::formats::shared::stream_core::common::{
};
use crate::formats::shared::AiSurfaceFinalizeError;
const PROVIDER_STREAM_FINISH_ERROR_MESSAGE: &str =
"Upstream stream ended with finish reason: error";
#[derive(Default)]
pub struct StreamingStandardFormatMatrix {
provider: Option<ProviderStreamParser>,
@@ -129,7 +132,7 @@ impl StreamingStandardFormatMatrix {
{
if !canonical_stream_finish_reason_is_supported(finish_reason) {
self.terminated = true;
out.extend(client.emit_unsupported_finish_reason(finish_reason)?);
out.extend(client.emit_finish_reason_error(finish_reason)?);
break;
}
}
@@ -388,7 +391,13 @@ impl StreamingStandardTerminalObserver {
if let Some(parser_error) = finish_reason
.as_deref()
.filter(|reason| !canonical_stream_finish_reason_is_supported(reason))
.map(|reason| format!("unsupported provider stream finish reason: {reason}"))
.map(|reason| {
if reason.trim() == "error" {
PROVIDER_STREAM_FINISH_ERROR_MESSAGE.to_string()
} else {
format!("unsupported provider stream finish reason: {reason}")
}
})
{
summary.parser_error.get_or_insert(parser_error);
}
@@ -660,17 +669,28 @@ impl ClientStreamEmitter {
self.emit_error(error_body)
}
fn emit_unsupported_finish_reason(
fn emit_finish_reason_error(
&mut self,
finish_reason: &str,
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
let (message, code) = if finish_reason.trim() == "error" {
(
PROVIDER_STREAM_FINISH_ERROR_MESSAGE.to_string(),
"stream_terminal_error",
)
} else {
(
format!(
"Unsupported provider stream finish reason cannot be converted losslessly: field $.finish_reason = {}",
serde_json::json!(finish_reason)
),
"unsupported_finish_reason",
)
};
let Some(error_body) = build_core_error_body_for_client_format(
self.api_format(),
&format!(
"Unsupported provider stream finish reason cannot be converted losslessly: field $.finish_reason = {}",
serde_json::json!(finish_reason)
),
Some("unsupported_finish_reason"),
&message,
Some(code),
LocalCoreSyncErrorKind::ServerError,
) else {
return Ok(Vec::new());
@@ -2121,6 +2141,163 @@ mod tests {
assert!(sse.contains("\"stop_reason\":\"tool_use\""), "{sse}");
}
#[test]
fn terminal_observer_treats_claude_error_finish_reason_as_upstream_failure() {
for upstream_message in [None, Some("Provider temporarily overloaded")] {
let context = report_context("claude:messages", "claude:messages");
let mut observer = StreamingStandardTerminalObserver::default();
observer
.push_line(
&context,
data_line(json!({
"type": "message_start",
"message": {
"id": "msg_error_finish",
"model": "claude-sonnet-4-5",
"usage": {
"input_tokens": 22,
"cache_read_input_tokens": 7,
"cache_creation_input_tokens": 3,
"cache_creation": { "ephemeral_5m_input_tokens": 3 }
}
}
})),
)
.expect("message start should be observed");
if let Some(message) = upstream_message {
observer
.push_line(
&context,
data_line(json!({
"type": "error",
"error": { "type": "overloaded_error", "message": message }
})),
)
.expect("upstream error should be observed");
}
observer
.push_line(
&context,
data_line(json!({
"type": "message_delta",
"delta": { "stop_reason": "error" },
"usage": { "output_tokens": 5 }
})),
)
.expect("error finish reason should be observed");
let summary = observer
.finish(&context)
.expect("terminal observation should finish")
.expect("failed stream should have a summary");
assert!(summary.observed_finish);
assert_eq!(summary.finish_reason.as_deref(), Some("error"));
assert_eq!(
summary.parser_error.as_deref(),
Some(upstream_message.unwrap_or("Upstream stream ended with finish reason: error"))
);
let usage = summary
.standardized_usage
.expect("usage should be retained");
assert_eq!(usage.input_tokens, 22);
assert_eq!(usage.output_tokens, 5);
assert_eq!(usage.cache_read_tokens, 7);
assert_eq!(usage.cache_creation_tokens, 3);
assert_eq!(usage.cache_creation_ephemeral_5m_tokens, 3);
}
}
#[test]
fn transforms_claude_error_finish_reason_to_terminal_errors() {
let cases = [
(
"openai:chat",
"data: {\"error\":",
"\"code\":\"stream_terminal_error\"",
),
(
"openai:responses",
"event: response.failed\n",
"\"code\":\"stream_terminal_error\"",
),
(
"claude:messages",
"event: error\n",
"\"code\":\"stream_terminal_error\"",
),
(
"gemini:generate_content",
"data: {\"error\":",
"\"status\":\"INTERNAL\"",
),
];
for (client_api_format, prefix, marker) in cases {
let context = report_context("claude:messages", client_api_format);
let mut matrix = StreamingStandardFormatMatrix::default();
let mut output = matrix
.transform_line(
&context,
data_line(json!({
"type": "content_block_delta",
"index": 0,
"delta": { "type": "text_delta", "text": "Partial answer" }
})),
)
.expect("partial response should be emitted");
output.extend(
matrix
.transform_line(
&context,
data_line(json!({
"type": "message_delta",
"delta": { "stop_reason": "error" },
"usage": { "output_tokens": 5 }
})),
)
.expect("error finish reason should emit a terminal error"),
);
let sse = String::from_utf8(output).expect("sse should be utf8");
assert!(sse.contains("Partial answer"), "{client_api_format}: {sse}");
assert!(sse.contains(prefix), "{client_api_format}: {sse}");
assert!(sse.contains(marker), "{client_api_format}: {sse}");
assert!(
sse.contains("Upstream stream ended with finish reason: error"),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("unsupported_finish_reason"),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("response.completed"),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("\"stop_reason\":\"end_turn\""),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("\"finish_reason\":\"stop\""),
"{client_api_format}: {sse}"
);
assert!(matrix
.transform_line(
&context,
data_line(json!({
"type": "message_delta",
"delta": { "stop_reason": "end_turn" }
}))
)
.expect("events after the error should be ignored")
.is_empty());
assert!(matrix
.finish(&context)
.expect("failed matrix should stay terminated")
.is_empty());
}
}
#[test]
fn transforms_unknown_stream_finish_reasons_to_visible_client_errors() {
let cases = [
@@ -34,6 +34,7 @@ pub struct CandidateFailureDiagnostic {
client_api_format: Option<String>,
provider_api_format: Option<String>,
safe_to_show: bool,
details: Option<Value>,
}
impl CandidateFailureDiagnostic {
@@ -50,6 +51,7 @@ impl CandidateFailureDiagnostic {
client_api_format: None,
provider_api_format: None,
safe_to_show: true,
details: None,
}
}
@@ -58,6 +60,11 @@ impl CandidateFailureDiagnostic {
self
}
pub fn details(mut self, details: Value) -> Self {
self.details = Some(details);
self
}
pub fn formats(
mut self,
client_api_format: impl Into<String>,
@@ -219,6 +226,10 @@ impl CandidateFailureDiagnostic {
"client_api_format": self.client_api_format,
"provider_api_format": self.provider_api_format,
"safe_to_show": self.safe_to_show,
"details": self.details,
"stage": "request",
"source_format": self.client_api_format,
"target_format": self.provider_api_format,
})
}
}
@@ -224,6 +224,7 @@ fn diagnostic_from_format_error(
format_error_path(error),
format_error_message(error, client_api_format, provider_api_format),
)
.details(error.diagnostic())
}
fn format_error_path(error: &FormatError) -> String {
@@ -982,6 +983,23 @@ mod tests {
"request_conversion"
);
assert_eq!(diagnostic["failure_diagnostic"]["path"], "$.n");
assert_eq!(diagnostic["failure_diagnostic"]["stage"], "request");
assert_eq!(
diagnostic["failure_diagnostic"]["details"]["code"],
"lossy_conversion_blocked"
);
assert_eq!(
diagnostic["failure_diagnostic"]["details"]["path_source"],
"structured"
);
assert_eq!(
diagnostic["failure_diagnostic"]["source_format"],
"openai:chat"
);
assert_eq!(
diagnostic["failure_diagnostic"]["target_format"],
"openai:responses"
);
assert_eq!(diagnostic["request_conversion_error"]["path"], "$.n");
assert!(diagnostic["failure_diagnostic"]["message"]
.as_str()
+124 -62
View File
@@ -1,8 +1,8 @@
use aether_data_contracts::repository::billing::StoredBillingModelContext;
use aether_data_contracts::repository::usage::{
extract_provider_cache_ttl_minutes_from_metadata, resolve_provider_cache_ttl_minutes,
resolve_provider_service_tier_from_request_capture, USAGE_AVAILABLE_METADATA_KEY,
USAGE_PRICING_AVAILABLE_METADATA_KEY,
resolve_provider_service_tier_from_request_capture, CANCELLED_REQUEST_FEE_METADATA_KEY,
USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY,
};
use aether_data_contracts::DataLayerError;
use aether_usage_runtime::{UsageEvent, UsageEventType};
@@ -40,6 +40,18 @@ pub async fn enrich_usage_event_with_billing(
data: &dyn BillingModelContextLookup,
event: &mut UsageEvent,
) -> Result<(), DataLayerError> {
if matches!(event.event_type, UsageEventType::Cancelled) {
event.data.total_cost_usd = Some(0.0);
event.data.actual_total_cost_usd = Some(0.0);
if let Some(metadata) = event
.data
.request_metadata
.as_mut()
.and_then(Value::as_object_mut)
{
metadata.remove(CANCELLED_REQUEST_FEE_METADATA_KEY);
}
}
// Session transports such as Codex Live expose lifecycle telemetry but no
// authoritative token/cost object. Do not run request-based pricing with
// zero default tokens: that would turn "unknown" into a fabricated charge.
@@ -65,7 +77,10 @@ pub async fn enrich_usage_event_with_billing(
clear_usage_costs(event);
return Ok(());
}
if !matches!(event.event_type, UsageEventType::Completed) {
if !matches!(
event.event_type,
UsageEventType::Completed | UsageEventType::Cancelled
) {
event.data.total_cost_usd = Some(0.0);
event.data.actual_total_cost_usd = Some(0.0);
return Ok(());
@@ -189,7 +204,10 @@ fn calculate_billing_computation(
} else {
usage_event_image_count(&event.data).unwrap_or(0)
};
let request_count = if failed {
let cancelled = matches!(event.event_type, UsageEventType::Cancelled);
let request_count = if cancelled {
1
} else if failed {
0
} else if is_image_usage && image_count > 0 {
image_count
@@ -197,7 +215,7 @@ fn calculate_billing_computation(
1
};
let processing_tiers = usage_event_processing_tiers(&event.data);
let input = BillingUsageInput {
let mut input = BillingUsageInput {
task_type: if is_image_usage {
"image".to_string()
} else {
@@ -237,6 +255,16 @@ fn calculate_billing_computation(
.or(pricing.provider_api_key_cache_ttl_minutes),
};
if cancelled {
input.input_tokens = 0;
input.output_tokens = 0;
input.cache_creation_tokens = 0;
input.cache_creation_ephemeral_5m_tokens = 0;
input.cache_creation_ephemeral_1h_tokens = 0;
input.cache_read_tokens = 0;
input.image_count = 0;
}
BillingService::new()
.calculate(pricing, &input)
.map_err(|err| {
@@ -356,9 +384,32 @@ fn apply_billing_computation(
pricing: &BillingModelPricingSnapshot,
computation: BillingComputation,
) -> Result<(), DataLayerError> {
let cancelled = matches!(event.event_type, UsageEventType::Cancelled);
if cancelled
&& !computation
.pricing_resolution
.price_per_request
.is_some_and(|price| price > 0.0)
{
return Ok(());
}
event.data.total_cost_usd = Some(computation.cost_result.cost);
event.data.actual_total_cost_usd = Some(computation.actual_total_cost);
merge_billing_snapshot_metadata(&mut event.data.request_metadata, pricing, &computation)
merge_billing_snapshot_metadata(&mut event.data.request_metadata, pricing, &computation)?;
if cancelled {
if let Some(metadata) = event
.data
.request_metadata
.as_mut()
.and_then(Value::as_object_mut)
{
metadata.insert(
CANCELLED_REQUEST_FEE_METADATA_KEY.to_string(),
Value::Bool(true),
);
}
}
Ok(())
}
fn map_pricing_context(context: StoredBillingModelContext) -> BillingModelPricingSnapshot {
@@ -1272,8 +1323,11 @@ mod tests {
}
#[tokio::test]
async fn cancelled_usage_event_remains_unbilled() {
let lookup = TestLookup {
async fn cancelled_usage_bills_only_configured_request_fee() {
for (request_type, request_price) in
[("chat", None), ("chat", Some(0.02)), ("image", Some(0.02))]
{
let lookup = TestLookup {
name_context: Some(
StoredBillingModelContext::new(
"provider-1".to_string(),
@@ -1284,7 +1338,7 @@ mod tests {
"global-model-1".to_string(),
"gpt-5".to_string(),
None,
Some(0.02),
request_price,
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":3.0,"output_price_per_1m":15.0,"cache_creation_price_per_1m":3.75,"cache_read_price_per_1m":0.30}]})),
Some("model-1".to_string()),
Some("gpt-5-upstream".to_string()),
@@ -1296,61 +1350,69 @@ mod tests {
),
model_id_context: None,
};
let mut event = UsageEvent::new(
UsageEventType::Cancelled,
"req-billing-cancelled-1",
UsageEventData {
provider_name: "OpenAI".to_string(),
model: "gpt-5".to_string(),
provider_id: Some("provider-1".to_string()),
provider_api_key_id: Some("key-1".to_string()),
request_type: Some("chat".to_string()),
api_format: Some("openai:responses".to_string()),
endpoint_api_format: Some("openai:responses".to_string()),
input_tokens: Some(1_000),
output_tokens: Some(500),
cache_read_input_tokens: Some(100),
status_code: Some(499),
..UsageEventData::default()
},
);
let mut event = UsageEvent::new(
UsageEventType::Cancelled,
"req-billing-cancelled-1",
UsageEventData {
provider_name: "OpenAI".to_string(),
model: "gpt-5".to_string(),
provider_id: Some("provider-1".to_string()),
provider_api_key_id: Some("key-1".to_string()),
request_type: Some(request_type.to_string()),
api_format: Some("openai:responses".to_string()),
endpoint_api_format: Some("openai:responses".to_string()),
input_tokens: Some(1_000),
output_tokens: Some(500),
cache_read_input_tokens: Some(100),
status_code: Some(499),
request_metadata: Some(
json!({"cancelled_request_fee": true, "image_count": 3}),
),
..UsageEventData::default()
},
);
enrich_usage_event_with_billing(&lookup, &mut event)
.await
.expect("billing should succeed");
enrich_usage_event_with_billing(&lookup, &mut event)
.await
.expect("billing should succeed");
assert_eq!(event.data.total_cost_usd, Some(0.0));
assert_eq!(event.data.actual_total_cost_usd, Some(0.0));
assert_eq!(
event
.data
.request_metadata
.as_ref()
.and_then(|value| value.get("billing_snapshot"))
.and_then(|value| value.get("status"))
.and_then(Value::as_str),
None
);
assert_eq!(
event
.data
.request_metadata
.as_ref()
.and_then(|value| value.get("billing_dimensions"))
.and_then(|value| value.get("input_tokens"))
.and_then(Value::as_i64),
None
);
assert_eq!(
event
.data
.request_metadata
.as_ref()
.and_then(|value| value.get("billing_dimensions"))
.and_then(|value| value.get("cache_read_tokens"))
.and_then(Value::as_i64),
None
);
let expected_cost = request_price.unwrap_or(0.0);
assert_eq!(event.data.total_cost_usd, Some(expected_cost));
assert_eq!(event.data.actual_total_cost_usd, Some(expected_cost * 0.5));
assert_eq!(event.data.input_tokens, Some(1_000));
assert_eq!(event.data.output_tokens, Some(500));
let metadata = event.data.request_metadata.as_ref().unwrap();
assert_eq!(
aether_data_contracts::repository::usage::cancelled_request_fee_is_billable(Some(
metadata
)),
request_price.is_some()
);
if request_price.is_some() {
assert_eq!(
metadata.pointer("/billing_snapshot/cost_breakdown/request_cost"),
Some(&json!(expected_cost))
);
assert_eq!(
metadata.pointer("/billing_dimensions/input_tokens"),
Some(&json!(0))
);
assert_eq!(
metadata.pointer("/billing_dimensions/output_tokens"),
Some(&json!(0))
);
assert_eq!(
metadata.pointer("/billing_dimensions/cache_read_tokens"),
Some(&json!(0))
);
assert_eq!(
metadata.pointer("/billing_dimensions/request_count"),
Some(&json!(1))
);
} else {
assert!(metadata.get("billing_snapshot").is_none());
}
}
}
#[tokio::test]
@@ -18,8 +18,6 @@ futures-util.workspace = true
sqlx = { workspace = true, features = ["postgres", "runtime-tokio-rustls", "chrono", "migrate", "macros"] }
serde_json.workspace = true
sha2.workspace = true
tokio.workspace = true
tracing.workspace = true
uuid.workspace = true
[dev-dependencies]
tokio.workspace = true
@@ -0,0 +1,45 @@
DO $migration$
DECLARE
policy_column record;
BEGIN
FOR policy_column IN
SELECT *
FROM (VALUES
('api_keys', 'allowed_providers'),
('api_keys', 'allowed_api_formats'),
('api_keys', 'allowed_models'),
('api_keys', 'ip_rules'),
('users', 'allowed_providers'),
('users', 'allowed_api_formats'),
('users', 'allowed_models'),
('user_groups', 'allowed_providers'),
('user_groups', 'allowed_api_formats'),
('user_groups', 'allowed_models'),
('provider_api_keys', 'api_formats'),
('provider_api_keys', 'allowed_models')
) AS policy_columns(table_name, column_name)
LOOP
EXECUTE format(
$statement$
UPDATE public.%1$I
SET %2$I = NULL
WHERE json_typeof(%2$I::json) = 'null'
OR (
json_typeof(%2$I::json) = 'string'
AND (%2$I::json #>> '{}') ~* '^[[:space:]]*(null)?[[:space:]]*$'
)
$statement$,
policy_column.table_name,
policy_column.column_name
);
END LOOP;
END;
$migration$;
UPDATE public.management_tokens
SET allowed_ips = NULL
WHERE json_typeof(allowed_ips::json) = 'null';
UPDATE public.management_tokens
SET permissions = NULL
WHERE json_typeof(permissions::json) = 'null';
@@ -1,9 +1,9 @@
use aether_data_contracts::repository::usage::{
canonical_usage_body_ref_for, parse_usage_body_ref, read_decompressed_usage_json,
usage_body_ref, ApiKeyLastUsedDelta, ManagementTokenCounterDelta, ProxyNodeCounterDelta,
StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow,
StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow,
StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBodyPayload,
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow,
StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow,
@@ -63,8 +63,25 @@ pub mod cleanup;
// newly captured bodies always spill to usage_body_blobs and resolve through usage_http_audits.
const MAX_INLINE_USAGE_BODY_BYTES: usize = 0;
const MAX_SUPPORTED_UNIX_SECS: u64 = 253_402_300_799;
const FIND_USAGE_BODY_BLOB_BY_REF_SQL: &str = r#"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = $1 AND request_id = $2 AND body_field = $3 LIMIT 1"#;
const FIND_USAGE_BODY_BLOB_BY_REF_SQL: &str = r#"SELECT CASE WHEN octet_length(payload_gzip) <= $4 THEN payload_gzip END AS payload_gzip FROM usage_body_blobs WHERE body_ref = $1 AND request_id = $2 AND body_field = $3 LIMIT 1"#;
const DELETE_USAGE_BODY_BLOB_SQL: &str = include_str!("queries/delete_usage_body_blob_sql.sql");
static USAGE_BODY_DECODE_SLOTS: tokio::sync::Semaphore = tokio::sync::Semaphore::const_new(4);
async fn decode_usage_body_in_background(
decode: impl FnOnce() -> Result<Option<Value>, DataLayerError> + Send + 'static,
) -> Result<Option<Value>, DataLayerError> {
let permit = USAGE_BODY_DECODE_SLOTS.acquire().await.map_err(|error| {
DataLayerError::UnexpectedValue(format!("usage body decoder unavailable: {error}"))
})?;
tokio::task::spawn_blocking(move || {
let _permit = permit;
decode()
})
.await
.map_err(|error| {
DataLayerError::UnexpectedValue(format!("usage body decoder failed: {error}"))
})?
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
struct AggregateRangeSplit {
@@ -2862,36 +2879,82 @@ ORDER BY request_count DESC, "usage".provider_name ASC
Ok(items)
}
pub async fn resolve_body_ref(&self, body_ref: &str) -> Result<Option<Value>, DataLayerError> {
pub async fn read_body_payload(
&self,
body_ref: &str,
) -> Result<Option<StoredUsageBodyPayload>, DataLayerError> {
let json_limit =
aether_data_contracts::repository::usage::MAX_DECOMPRESSED_USAGE_JSON_BYTES as i64;
let encoded_limit = json_limit + 1024 * 1024;
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
return Ok(None);
};
let canonical_ref = usage_body_ref(&request_id, field);
let blob_row = sqlx::query(FIND_USAGE_BODY_BLOB_BY_REF_SQL)
let row = sqlx::query(FIND_USAGE_BODY_BLOB_BY_REF_SQL)
.bind(&canonical_ref)
.bind(&request_id)
.bind(field.as_storage_field())
.bind(encoded_limit)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
if let Some(row) = blob_row.as_ref() {
let payload_gzip = row
.try_get::<Vec<u8>, _>("payload_gzip")
.map_postgres_err()?;
return inflate_usage_json_value(&payload_gzip).map(Some);
if let Some(row) = row {
return row
.try_get::<Option<Vec<u8>>, _>("payload_gzip")
.map_postgres_err()?
.map(|bytes| Some(StoredUsageBodyPayload::Gzip(bytes)))
.ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"encoded usage json exceeds {encoded_limit} bytes"
))
});
}
let (inline_column, compressed_column) = usage_body_sql_columns(field);
let row = sqlx::query(&format!(
"SELECT {inline_column} AS inline_body, {compressed_column} AS compressed_body FROM \"usage\" WHERE request_id = $1 LIMIT 1"
"SELECT CASE WHEN octet_length({inline_column}::text) <= $2 THEN {inline_column}::text END AS inline_body, CASE WHEN octet_length({compressed_column}) <= $3 THEN {compressed_column} END AS compressed_body, (COALESCE(octet_length({inline_column}::text) > $2, false) OR ({inline_column} IS NULL AND COALESCE(octet_length({compressed_column}) > $3, false))) AS too_large FROM \"usage\" WHERE request_id = $1 LIMIT 1"
))
.bind(request_id)
.bind(json_limit)
.bind(encoded_limit)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref()
.map(|row| usage_json_column(row, "inline_body", "compressed_body", true))
.transpose()
.map(|value| value.and_then(|column| column.value))
let Some(row) = row else {
return Ok(None);
};
if row.try_get::<bool, _>("too_large").map_postgres_err()? {
return Err(DataLayerError::UnexpectedValue(
"encoded usage json exceeds preview limit".to_string(),
));
}
if let Some(body) = row
.try_get::<Option<String>, _>("inline_body")
.map_postgres_err()?
{
return Ok(Some(StoredUsageBodyPayload::Json(body.into_bytes())));
}
Ok(row
.try_get::<Option<Vec<u8>>, _>("compressed_body")
.map_postgres_err()?
.map(StoredUsageBodyPayload::Gzip))
}
pub async fn resolve_body_ref(&self, body_ref: &str) -> Result<Option<Value>, DataLayerError> {
let Some(payload) = self.read_body_payload(body_ref).await? else {
return Ok(None);
};
decode_usage_body_in_background(move || match payload {
StoredUsageBodyPayload::Gzip(bytes) => inflate_usage_json_value(&bytes).map(Some),
StoredUsageBodyPayload::Json(bytes) => {
let bytes = read_decompressed_usage_json(std::io::Cursor::new(bytes))?;
serde_json::from_slice(&bytes).map(Some).map_err(|error| {
DataLayerError::UnexpectedValue(format!(
"failed to parse decompressed usage json: {error}"
))
})
}
})
.await
}
async fn hydrate_usage_body_refs(
@@ -10306,6 +10369,13 @@ impl UsageReadRepository for SqlxUsageReadRepository {
Self::resolve_body_ref(self, body_ref).await
}
async fn read_body_payload(
&self,
body_ref: &str,
) -> Result<Option<StoredUsageBodyPayload>, DataLayerError> {
Self::read_body_payload(self, body_ref).await
}
async fn list_usage_audits(
&self,
query: &UsageAuditListQuery,
@@ -283,11 +283,11 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() {
.unwrap();
assert_eq!(
stored.request_headers,
Some(json!({"content-type": "application/json", "authorization": "[redacted]"}))
Some(json!({"content-type": "application/json", "authorization": "Bearer private"}))
);
assert_eq!(
stored.response_headers,
Some(json!({"content-type": "text/event-stream", "set-cookie": "[redacted]"}))
Some(json!({"content-type": "text/event-stream", "set-cookie": "private"}))
);
for (field, expected) in [
(UsageBodyField::RequestBody, pending.request_body.as_ref()),
@@ -704,7 +704,7 @@ async fn live_pending_batch_persists_auxiliary_state_and_preserves_terminal_conf
.unwrap();
assert_eq!(
captured.request_headers,
Some(json!({"x-request": "[redacted]"}))
Some(json!({"x-request": "request-value"}))
);
assert_eq!(
repository
@@ -4274,6 +4274,38 @@ fn prepare_usage_body_storage_detaches_small_payloads_into_blob_storage() {
);
}
#[tokio::test(flavor = "current_thread")]
async fn usage_body_decode_does_not_block_the_async_runtime_thread() {
let runtime_thread = std::thread::current().id();
let payload = json!({"message": "background decoding"});
let compressed = prepare_usage_body_storage(Some(&payload))
.expect("body should compress")
.detached_blob_bytes
.expect("body should be detached");
let decoded = super::decode_usage_body_in_background(move || {
assert_ne!(std::thread::current().id(), runtime_thread);
inflate_usage_json_value(&compressed).map(Some)
})
.await
.expect("body should decode");
assert_eq!(decoded, Some(payload));
}
#[tokio::test]
async fn usage_body_decode_preserves_storage_decode_errors() {
let error = super::decode_usage_body_in_background(|| {
inflate_usage_json_value(b"invalid gzip").map(Some)
})
.await
.expect_err("corrupt bodies should fail");
assert!(error
.to_string()
.contains("failed to decompress usage json:"));
}
#[test]
fn prepare_usage_body_storage_compresses_large_payloads() {
let payload = json!({

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