mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 02:47:45 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
361952ada9 | ||
|
|
6630856061 | ||
|
|
a893bd0557 | ||
|
|
f2839ae6a7 | ||
|
|
e58570d79d | ||
|
|
99f6499b2b | ||
|
|
17d01d7fe0 | ||
|
|
8b766930b0 | ||
|
|
c7e403b410 | ||
|
|
cf8ea19856 | ||
|
|
7113d04f8a | ||
|
|
099b810a2f |
+5
-6
@@ -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
@@ -305,6 +305,7 @@ dependencies = [
|
||||
"futures-util",
|
||||
"hmac",
|
||||
"http",
|
||||
"http-body",
|
||||
"http-body-util",
|
||||
"hyper",
|
||||
"hyper-util",
|
||||
|
||||
+1
-1
@@ -44,5 +44,5 @@ EXPOSE 8084
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
||||
|
||||
USER 65532:65532
|
||||
USER 0:0
|
||||
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
||||
|
||||
@@ -157,4 +157,5 @@ EXPOSE 8084
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD ["/usr/local/bin/aether-gateway", "--healthcheck"]
|
||||
|
||||
USER 0:0
|
||||
ENTRYPOINT ["/usr/local/bin/aether-gateway"]
|
||||
|
||||
@@ -156,4 +156,5 @@ EXPOSE 8084
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
||||
|
||||
USER 0:0
|
||||
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
||||
|
||||
@@ -50,82 +50,10 @@ chmod 600 .env
|
||||
./generate_keys.sh
|
||||
# 编辑 .env 设置 ADMIN_PASSWORD
|
||||
|
||||
# 3. 首次部署 / 更新 (从以下部署形态任选其一)
|
||||
# Postgres + Redis (推荐)
|
||||
# 3. Docker 部署 / 更新(PostgreSQL + Redis)
|
||||
docker compose pull && docker compose up -d
|
||||
# Single Node:同样使用 PostgreSQL + Redis,无需挂载本地数据库文件
|
||||
docker compose -f docker-compose.single-node.yml pull && docker compose -f docker-compose.single-node.yml up -d
|
||||
```
|
||||
|
||||
应用镜像默认以固定非 root 身份 `65532:65532` 运行;Compose 移除全部 Linux capabilities、禁止提权、启用只读根文件系统,并提供带 `nosuid,nodev,noexec` 的 `/tmp`。如需使用其他身份,可在 `.env` 中设置非零的 `AETHER_CONTAINER_UID` / `AETHER_CONTAINER_GID`。数据库使用独立 PostgreSQL 容器和 named volume,不再需要调整应用数据库目录的权限。
|
||||
|
||||
|
||||
### 一键更新
|
||||
|
||||
Docker Compose 部署后,可在部署目录直接执行:
|
||||
|
||||
```bash
|
||||
./update.sh
|
||||
```
|
||||
|
||||
`update.sh` 会拉取最新 `app` 镜像并重建 `app` 容器,Docker named volumes、`./data` 和 `./logs` 不会被删除。Single Node 部署也可显式指定:
|
||||
|
||||
```bash
|
||||
./update.sh --mode single-node
|
||||
```
|
||||
|
||||
现在仅支持 PostgreSQL。标准和单节点 Docker Compose 均部署 PostgreSQL + Redis;原生 systemd / launchd 安装需要显式提供 PostgreSQL `DATABASE_URL`,例如 `DATABASE_URL=postgresql://user:password@host:5432/aether`。旧数据库不会自动迁移或清空。升级时保留原有 PostgreSQL 密码、`JWT_SECRET_KEY` 和 `ENCRYPTION_KEY`,不要重新生成整个 `.env`。
|
||||
|
||||
仓库自带的 Docker Compose 默认把应用日志输出到容器 `stdout/stderr`,直接用 `docker compose logs -f app` 查看,并由 Docker 轮转日志,避免非 root 用户被宿主机日志目录权限拖垮启动。如果你确实需要文件日志,需要在 compose 里把 `AETHER_LOG_DESTINATION` 改成 `file|both`,额外挂载目录到 `/opt/aether/logs`,并让它归 `.env` 中配置的容器 UID/GID 所有;只读根文件系统不会阻止显式可写挂载。
|
||||
|
||||
管理后台右上角“版本信息”会检测新版本。Docker Compose 部署只提示版本,实际更新继续执行 `./update.sh`;systemd / launchd / 二进制部署才使用后台自更新,流程是下载对应平台的 GitHub Release 包、强制校验 `SHA256SUMS`、解压到 `/opt/aether/releases/<version>`,再切换 `/opt/aether/current` 并退出进程,交给 systemd / launchd 拉起新版本。
|
||||
|
||||
正式 Release 还会发布由 GitHub Actions OIDC / Sigstore 签发的 SLSA build provenance。需要验证发布者身份时,下载目标 tarball 和 `AETHER_RELEASE_PROVENANCE.sigstore.json`,并把 `TAG` 设置为对应 Release tag:
|
||||
|
||||
```bash
|
||||
gh attestation verify "aether-${TAG}-linux-amd64.tar.gz" \
|
||||
--repo fawney19/Aether \
|
||||
--signer-workflow fawney19/Aether/.github/workflows/release.yml \
|
||||
--source-ref "refs/tags/${TAG}" \
|
||||
--bundle AETHER_RELEASE_PROVENANCE.sigstore.json
|
||||
```
|
||||
|
||||
`docker-compose.yml` 中的官方 PostgreSQL 和 Redis 镜像均固定到多架构 OCI index digest。升级这些依赖时应在发布变更中显式更新 digest,避免同名 tag 在无人审查的情况下改变部署内容。
|
||||
|
||||
正式发布到 GHCR 和 Docker Hub 的多架构 Aether 镜像也带有同一 GitHub Actions OIDC / Sigstore provenance;生产 `Dockerfile.app` 的 BusyBox 与 Distroless 基础镜像同样固定到多架构 OCI index digest。
|
||||
|
||||
源码或本地构建版本不会启用后台在线更新,请继续使用源码更新流程。Docker Compose 用户如果希望“容器重建后也保持镜像层面的新版本”,仍建议定期运行 `./update.sh` 拉取并重建 app 镜像。服务器访问 GitHub 需要代理时,可设置 `AETHER_UPDATE_PROXY_URL`,也兼容 `UPDATE_PROXY_URL`、`HTTPS_PROXY`、`ALL_PROXY`、`HTTP_PROXY` 以及 `NO_PROXY`。共享出口触发 GitHub API 限流时,可设置只读 `AETHER_UPDATE_GITHUB_TOKEN`,也兼容 `GITHUB_TOKEN` / `GH_TOKEN`。下载总超时默认 600 秒,连续无响应/无数据默认 30 秒,可通过 `AETHER_UPDATE_DOWNLOAD_TIMEOUT_SECS` 和 `AETHER_UPDATE_DOWNLOAD_IDLE_TIMEOUT_SECS` 调整。
|
||||
|
||||
标准和 Single Node Docker Compose 均使用 Docker named volume 存放 PostgreSQL 数据。
|
||||
|
||||
如果是本地源码构建镜像的部署,继续使用:
|
||||
|
||||
```bash
|
||||
./deploy.sh
|
||||
```
|
||||
|
||||
如果要在本机联调“管理后台在线更新”本身,可启动仓库内置的 release-layout 测试环境:
|
||||
|
||||
```bash
|
||||
docker compose -f docker-compose.release-local.yml up -d --build
|
||||
```
|
||||
|
||||
这套环境会用当前源码构建一个本地测试镜像,但编译为 `release` 类型,并默认伪装成 `v0.7.0`,这样后台会按正式发布版逻辑开放“立即更新”。默认监听 `http://127.0.0.1:18085`,数据目录使用 `./data-release-local`;日志默认走 `docker logs`,不会影响你正在跑的源码构建容器。
|
||||
|
||||
如果这套容器在 `prepare-update` 时访问 GitHub 失败,而你本机是通过代理出网,请在 `.env` 里把 `AETHER_UPDATE_PROXY_URL` 写成宿主机地址,例如 `http://host.docker.internal:7890`;容器内的 `127.0.0.1` 指向容器自身,不是宿主机。
|
||||
|
||||
如果想重置这套联调环境(包括 `/opt/aether/current` 和已下载的历史版本),执行:
|
||||
|
||||
```bash
|
||||
docker compose -f docker-compose.release-local.yml down -v
|
||||
```
|
||||
|
||||
可选变量:
|
||||
|
||||
- `AETHER_RELEASE_LOCAL_VERSION`:本地联调镜像对外声明的当前版本,默认 `v0.7.0`
|
||||
- `AETHER_RELEASE_LOCAL_PORT`:本地联调端口,默认 `18085`
|
||||
- `LOCAL_RELEASE_APP_IMAGE`:本地联调镜像名,默认 `aether-app:release-local`
|
||||
|
||||
### 一键安装(PostgreSQL + Redis)
|
||||
|
||||
```bash
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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!({
|
||||
|
||||
@@ -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))]
|
||||
{
|
||||
|
||||
@@ -71,6 +71,7 @@ mod rate_limit;
|
||||
mod request_candidate_queue;
|
||||
mod request_candidate_runtime;
|
||||
mod request_diagnostics;
|
||||
mod request_lifecycle;
|
||||
mod roles;
|
||||
mod router;
|
||||
mod routing;
|
||||
|
||||
@@ -44,7 +44,7 @@ pub(crate) fn local_auth_jwt_secret() -> Result<String, String> {
|
||||
Err(std::env::VarError::NotPresent) => {
|
||||
#[cfg(test)]
|
||||
{
|
||||
return Ok(TEST_JWT_SECRET.to_string());
|
||||
Ok(TEST_JWT_SECRET.to_string())
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
@@ -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(
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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 缓存容量(条目数) |
|
||||
|
||||
#### 日志
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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}-{}",
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
+45
@@ -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
Reference in New Issue
Block a user