mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 03:09:50 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8b766930b0 | ||
|
|
c7e403b410 | ||
|
|
cf8ea19856 | ||
|
|
7113d04f8a | ||
|
|
099b810a2f | ||
|
|
7aa0c89244 | ||
|
|
7847ae98c6 | ||
|
|
a90d564931 | ||
|
|
a5c3699ae9 | ||
|
|
7b8048c6ae | ||
|
|
ec95f2ca1f | ||
|
|
aa7dbe67d3 | ||
|
|
a26680f460 | ||
|
|
522b979052 | ||
|
|
808946312a | ||
|
|
741107bf71 | ||
|
|
6962731220 | ||
|
|
062e111c03 | ||
|
|
470c59e197 | ||
|
|
2f929e74c7 | ||
|
|
fc0417ceb9 | ||
|
|
44174a31e0 | ||
|
|
b599fb7354 | ||
|
|
14f96c9fa0 | ||
|
|
6948852992 |
+12
-8
@@ -15,11 +15,6 @@ APP_PORT=8084
|
|||||||
# APP_IMAGE=ghcr.io/fawney19/aether:beta
|
# APP_IMAGE=ghcr.io/fawney19/aether:beta
|
||||||
# APP_IMAGE=ghcr.io/fawney19/aether:0.7.0-rc.1
|
# 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 前缀(默认 sk)
|
||||||
API_KEY_PREFIX=sk
|
API_KEY_PREFIX=sk
|
||||||
|
|
||||||
@@ -31,7 +26,11 @@ RUST_LOG=aether_gateway=info
|
|||||||
# 示例: http://localhost:5173,https://app.example.com
|
# 示例: http://localhost:5173,https://app.example.com
|
||||||
# CORS_ORIGINS=http://localhost:5173
|
# CORS_ORIGINS=http://localhost:5173
|
||||||
# CORS_ALLOW_CREDENTIALS=true
|
# 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_SAMESITE=None
|
||||||
# AUTH_REFRESH_COOKIE_SECURE=true
|
# AUTH_REFRESH_COOKIE_SECURE=true
|
||||||
|
|
||||||
@@ -111,8 +110,13 @@ ADMIN_USERNAME=admin123456
|
|||||||
# AETHER_BARK_ALLOW_HTTP=false
|
# AETHER_BARK_ALLOW_HTTP=false
|
||||||
# AETHER_BARK_ALLOW_PRIVATE_TARGETS=false
|
# AETHER_BARK_ALLOW_PRIVATE_TARGETS=false
|
||||||
|
|
||||||
# 可选 Provider OAuth 客户端。使用 Gemini CLI / Antigravity 浏览器授权时必须配置
|
# 普通 Provider 反代(包括 Provider OAuth)不按 DNS 地址过滤上游,兼容任意
|
||||||
# 对应的 client secret;client ID 未配置时使用内置的公开 native-app client ID。
|
# Fake-IP 域名及内网 DNS。仅信任管理员配置的上游;没有严格 DNS 过滤开关。
|
||||||
|
# URL 协议、字面 IP、TLS 证书,以及隧道中继和登录 OAuth 的校验仍保留。
|
||||||
|
|
||||||
|
# 可选 Provider OAuth 客户端。Gemini CLI 和 Antigravity 默认使用内置 native-app
|
||||||
|
# 客户端凭据;自定义 client ID 时必须同时配置对应的 client secret。
|
||||||
|
# 显式配置的 client secret 优先于默认值。
|
||||||
# AETHER_GEMINI_CLI_OAUTH_CLIENT_ID=
|
# AETHER_GEMINI_CLI_OAUTH_CLIENT_ID=
|
||||||
# AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET=
|
# AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET=
|
||||||
# AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID=
|
# AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID=
|
||||||
|
|||||||
@@ -13,6 +13,10 @@
|
|||||||
.plans
|
.plans
|
||||||
.playwright-mcp/
|
.playwright-mcp/
|
||||||
|
|
||||||
|
docs/architecture
|
||||||
|
!docs/architecture/architecture-dark.svg
|
||||||
|
!docs/architecture/architecture-light.svg
|
||||||
|
|
||||||
### Python ###
|
### Python ###
|
||||||
*.db
|
*.db
|
||||||
*.db-*
|
*.db-*
|
||||||
|
|||||||
Generated
+1
-1
@@ -661,7 +661,7 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "aether-tunnel"
|
name = "aether-tunnel"
|
||||||
version = "0.3.16"
|
version = "0.3.17"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"aether-contracts",
|
"aether-contracts",
|
||||||
"aether-gateway",
|
"aether-gateway",
|
||||||
|
|||||||
+1
-1
@@ -44,5 +44,5 @@ EXPOSE 8084
|
|||||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||||
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
||||||
|
|
||||||
USER 65532:65532
|
USER 0:0
|
||||||
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
||||||
|
|||||||
@@ -157,4 +157,5 @@ EXPOSE 8084
|
|||||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||||
CMD ["/usr/local/bin/aether-gateway", "--healthcheck"]
|
CMD ["/usr/local/bin/aether-gateway", "--healthcheck"]
|
||||||
|
|
||||||
|
USER 0:0
|
||||||
ENTRYPOINT ["/usr/local/bin/aether-gateway"]
|
ENTRYPOINT ["/usr/local/bin/aether-gateway"]
|
||||||
|
|||||||
@@ -156,4 +156,5 @@ EXPOSE 8084
|
|||||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||||
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
CMD ["/opt/aether/current/bin/aether-gateway", "--healthcheck"]
|
||||||
|
|
||||||
|
USER 0:0
|
||||||
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
ENTRYPOINT ["/opt/aether/current/bin/aether-gateway"]
|
||||||
|
|||||||
@@ -50,82 +50,10 @@ chmod 600 .env
|
|||||||
./generate_keys.sh
|
./generate_keys.sh
|
||||||
# 编辑 .env 设置 ADMIN_PASSWORD
|
# 编辑 .env 设置 ADMIN_PASSWORD
|
||||||
|
|
||||||
# 3. 首次部署 / 更新 (从以下部署形态任选其一)
|
# 3. Docker 部署 / 更新(PostgreSQL + Redis)
|
||||||
# Postgres + Redis (推荐)
|
|
||||||
docker compose pull && docker compose up -d
|
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)
|
### 一键安装(PostgreSQL + Redis)
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|||||||
@@ -496,6 +496,7 @@ fn access_for_route(method: &http::Method, decision: &GatewayControlDecision) ->
|
|||||||
Some("admin:endpoints_manage"),
|
Some("admin:endpoints_manage"),
|
||||||
Some(
|
Some(
|
||||||
"reveal_key"
|
"reveal_key"
|
||||||
|
| "reveal_endpoint_rules"
|
||||||
| "export_key"
|
| "export_key"
|
||||||
| "create_provider_key"
|
| "create_provider_key"
|
||||||
| "update_key"
|
| "update_key"
|
||||||
@@ -1490,6 +1491,12 @@ mod tests {
|
|||||||
fn plaintext_credential_reads_require_admin_permission() {
|
fn plaintext_credential_reads_require_admin_permission() {
|
||||||
let read_only_permissions = read_only_management_token_permissions();
|
let read_only_permissions = read_only_management_token_permissions();
|
||||||
let cases = [
|
let cases = [
|
||||||
|
(
|
||||||
|
"admin:endpoints_manage",
|
||||||
|
"reveal_endpoint_rules",
|
||||||
|
None,
|
||||||
|
"admin:endpoints_manage:admin",
|
||||||
|
),
|
||||||
(
|
(
|
||||||
"admin:endpoints_manage",
|
"admin:endpoints_manage",
|
||||||
"reveal_key",
|
"reveal_key",
|
||||||
|
|||||||
@@ -302,6 +302,19 @@ pub(super) fn classify_admin_endpoints_family_route(
|
|||||||
"admin:endpoints_manage",
|
"admin:endpoints_manage",
|
||||||
false,
|
false,
|
||||||
))
|
))
|
||||||
|
} else if method == http::Method::GET
|
||||||
|
&& normalized_path
|
||||||
|
.strip_prefix("/api/admin/endpoints/")
|
||||||
|
.and_then(|path| path.strip_suffix("/rules/reveal"))
|
||||||
|
.is_some_and(|endpoint_id| !endpoint_id.is_empty() && !endpoint_id.contains('/'))
|
||||||
|
{
|
||||||
|
Some(classified(
|
||||||
|
"admin_proxy",
|
||||||
|
"endpoints_manage",
|
||||||
|
"reveal_endpoint_rules",
|
||||||
|
"admin:endpoints_manage",
|
||||||
|
false,
|
||||||
|
))
|
||||||
} else if method == http::Method::GET
|
} else if method == http::Method::GET
|
||||||
&& normalized_path.starts_with("/api/admin/endpoints/")
|
&& normalized_path.starts_with("/api/admin/endpoints/")
|
||||||
&& !normalized_path.starts_with("/api/admin/endpoints/health/")
|
&& !normalized_path.starts_with("/api/admin/endpoints/health/")
|
||||||
|
|||||||
@@ -381,6 +381,28 @@ fn classifies_admin_get_endpoint_as_admin_proxy_route() {
|
|||||||
assert!(!decision.is_execution_runtime_candidate());
|
assert!(!decision.is_execution_runtime_candidate());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn classifies_admin_reveal_endpoint_rules_as_admin_proxy_route() {
|
||||||
|
let headers = headers(&[]);
|
||||||
|
let uri: Uri = "/api/admin/endpoints/endpoint-1/rules/reveal"
|
||||||
|
.parse()
|
||||||
|
.expect("uri should parse");
|
||||||
|
let decision =
|
||||||
|
classify_control_route(&http::Method::GET, &uri, &headers).expect("route should classify");
|
||||||
|
|
||||||
|
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
|
||||||
|
assert_eq!(decision.route_family.as_deref(), Some("endpoints_manage"));
|
||||||
|
assert_eq!(
|
||||||
|
decision.route_kind.as_deref(),
|
||||||
|
Some("reveal_endpoint_rules")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
decision.auth_endpoint_signature.as_deref(),
|
||||||
|
Some("admin:endpoints_manage")
|
||||||
|
);
|
||||||
|
assert!(!decision.is_execution_runtime_candidate());
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn classifies_admin_create_endpoint_as_admin_proxy_route() {
|
fn classifies_admin_create_endpoint_as_admin_proxy_route() {
|
||||||
let headers = http::HeaderMap::new();
|
let headers = http::HeaderMap::new();
|
||||||
|
|||||||
@@ -17,13 +17,13 @@ fn sanitize_request_candidate_rows(
|
|||||||
mut candidates: Vec<StoredRequestCandidate>,
|
mut candidates: Vec<StoredRequestCandidate>,
|
||||||
) -> Vec<StoredRequestCandidate> {
|
) -> Vec<StoredRequestCandidate> {
|
||||||
for candidate in &mut candidates {
|
for candidate in &mut candidates {
|
||||||
candidate.sanitize_sensitive_diagnostics();
|
candidate.sanitize_for_persistence();
|
||||||
}
|
}
|
||||||
candidates
|
candidates
|
||||||
}
|
}
|
||||||
|
|
||||||
fn sanitize_request_candidate_row(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
|
fn sanitize_request_candidate_row(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
|
||||||
candidate.sanitize_sensitive_diagnostics();
|
candidate.sanitize_for_persistence();
|
||||||
candidate
|
candidate
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1048,7 +1048,10 @@ mod request_candidate_security_tests {
|
|||||||
fn assert_candidate_is_sanitized(candidate: &StoredRequestCandidate) {
|
fn assert_candidate_is_sanitized(candidate: &StoredRequestCandidate) {
|
||||||
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
||||||
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
||||||
assert!(candidate.error_message.is_none());
|
assert_eq!(
|
||||||
|
candidate.error_message.as_deref(),
|
||||||
|
Some("Bearer candidate-secret")
|
||||||
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
candidate.extra_data,
|
candidate.extra_data,
|
||||||
Some(json!({"gateway_execution_runtime": true}))
|
Some(json!({"gateway_execution_runtime": true}))
|
||||||
@@ -1057,13 +1060,15 @@ mod request_candidate_security_tests {
|
|||||||
candidate.required_capabilities,
|
candidate.required_capabilities,
|
||||||
Some(json!({"vision": true}))
|
Some(json!({"vision": true}))
|
||||||
);
|
);
|
||||||
assert!(!serde_json::to_string(candidate)
|
let mut public_candidate = candidate.clone();
|
||||||
|
public_candidate.sanitize_sensitive_diagnostics();
|
||||||
|
assert!(!serde_json::to_string(&public_candidate)
|
||||||
.expect("candidate should serialize")
|
.expect("candidate should serialize")
|
||||||
.contains("candidate-secret"));
|
.contains("candidate-secret"));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn gateway_candidate_boundary_sanitizes_repository_rows_and_write_results() {
|
fn gateway_candidate_boundary_preserves_admin_errors_and_removes_request_payloads() {
|
||||||
let candidate = sanitize_request_candidate_row(untrusted_candidate());
|
let candidate = sanitize_request_candidate_row(untrusted_candidate());
|
||||||
assert_candidate_is_sanitized(&candidate);
|
assert_candidate_is_sanitized(&candidate);
|
||||||
|
|
||||||
|
|||||||
@@ -41,16 +41,9 @@ fn usage_request_record_level_from_value(value: Option<&Value>) -> UsageRequestR
|
|||||||
return UsageRequestRecordLevel::Basic;
|
return UsageRequestRecordLevel::Basic;
|
||||||
};
|
};
|
||||||
|
|
||||||
if value.eq_ignore_ascii_case("basic")
|
if value.eq_ignore_ascii_case("full") {
|
||||||
|| value.eq_ignore_ascii_case("base")
|
UsageRequestRecordLevel::Full
|
||||||
|| value.eq_ignore_ascii_case("headers")
|
|
||||||
|| value.eq_ignore_ascii_case("minimal")
|
|
||||||
|| value.eq_ignore_ascii_case("none")
|
|
||||||
{
|
|
||||||
UsageRequestRecordLevel::Basic
|
|
||||||
} else {
|
} else {
|
||||||
// Raw HTTP payload capture is disabled at the runtime boundary. The setting remains
|
|
||||||
// accepted for compatibility, but no longer authorizes collecting request/response data.
|
|
||||||
UsageRequestRecordLevel::Basic
|
UsageRequestRecordLevel::Basic
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -501,7 +494,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn usage_runtime_access_disables_full_http_capture() {
|
async fn usage_runtime_access_honors_explicit_full_http_capture() {
|
||||||
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
|
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
|
||||||
"request_record_level".to_string(),
|
"request_record_level".to_string(),
|
||||||
json!("full"),
|
json!("full"),
|
||||||
@@ -511,7 +504,32 @@ mod tests {
|
|||||||
.await
|
.await
|
||||||
.expect("request record level should read");
|
.expect("request record level should read");
|
||||||
|
|
||||||
assert_eq!(level, UsageRequestRecordLevel::Basic);
|
assert_eq!(level, UsageRequestRecordLevel::Full);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn usage_runtime_access_honors_legacy_full_without_overriding_current_config() {
|
||||||
|
let state = GatewayDataState::disabled().with_system_config_values_for_tests([(
|
||||||
|
"request_log_level".to_string(),
|
||||||
|
json!(" FULL "),
|
||||||
|
)]);
|
||||||
|
assert_eq!(
|
||||||
|
UsageRuntimeAccess::request_record_level(&state)
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
UsageRequestRecordLevel::Full
|
||||||
|
);
|
||||||
|
|
||||||
|
let state = state.with_system_config_values_for_tests([
|
||||||
|
("request_log_level".to_string(), json!("full")),
|
||||||
|
("request_record_level".to_string(), json!("basic")),
|
||||||
|
]);
|
||||||
|
assert_eq!(
|
||||||
|
UsageRuntimeAccess::request_record_level(&state)
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
UsageRequestRecordLevel::Basic
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -1593,6 +1593,19 @@ impl GatewayDataState {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(crate) async fn read_request_usage_body_payload(
|
||||||
|
&self,
|
||||||
|
body_ref: &str,
|
||||||
|
) -> Result<
|
||||||
|
Option<aether_data_contracts::repository::usage::StoredUsageBodyPayload>,
|
||||||
|
DataLayerError,
|
||||||
|
> {
|
||||||
|
match &self.usage_reader {
|
||||||
|
Some(repository) => repository.read_body_payload(body_ref).await,
|
||||||
|
None => Ok(None),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) async fn list_usage_audits(
|
pub(crate) async fn list_usage_audits(
|
||||||
&self,
|
&self,
|
||||||
query: &UsageAuditListQuery,
|
query: &UsageAuditListQuery,
|
||||||
|
|||||||
@@ -426,11 +426,8 @@ mod tests {
|
|||||||
assert!(candidate.finished_at_unix_ms.is_some());
|
assert!(candidate.finished_at_unix_ms.is_some());
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The guard holds no request body, and the persistence boundary intentionally
|
|
||||||
/// rejects request/response capture material. A dropped-attempt settlement
|
|
||||||
/// must not re-introduce an inline body or a caller-controlled body reference.
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn settling_a_dropped_attempt_does_not_reintroduce_request_body_capture() {
|
async fn settling_a_dropped_attempt_respects_disabled_request_body_capture() {
|
||||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
||||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||||
let state = test_state(&usage_repository, &request_candidate_repository);
|
let state = test_state(&usage_repository, &request_candidate_repository);
|
||||||
@@ -444,8 +441,6 @@ mod tests {
|
|||||||
candidate_started_unix_ms,
|
candidate_started_unix_ms,
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
// This deliberately supplies capture material to prove that the usage
|
|
||||||
// persistence boundary strips it before either lifecycle write stores it.
|
|
||||||
let captured_body = json!({"stream": true, "service_tier": "priority"});
|
let captured_body = json!({"stream": true, "service_tier": "priority"});
|
||||||
let mut capture = build_pending_usage_record(
|
let mut capture = build_pending_usage_record(
|
||||||
&plan,
|
&plan,
|
||||||
@@ -481,7 +476,10 @@ mod tests {
|
|||||||
.expect("cancelled usage should be recorded");
|
.expect("cancelled usage should be recorded");
|
||||||
assert_eq!(usage.provider_request_body, None);
|
assert_eq!(usage.provider_request_body, None);
|
||||||
assert_eq!(usage.provider_request_body_ref, None);
|
assert_eq!(usage.provider_request_body_ref, None);
|
||||||
assert_eq!(usage.provider_request_body_state, None);
|
assert_eq!(
|
||||||
|
usage.provider_request_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Disabled)
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -14661,15 +14661,19 @@ mod tests {
|
|||||||
assert_eq!(usage.status_code, Some(302));
|
assert_eq!(usage.status_code, Some(302));
|
||||||
assert_eq!(usage.error_category.as_deref(), Some("redirect"));
|
assert_eq!(usage.error_category.as_deref(), Some("redirect"));
|
||||||
assert!(usage.error_message.is_none());
|
assert!(usage.error_message.is_none());
|
||||||
// HTTP capture is intentionally disabled at the persistence boundary. Keep the
|
assert_eq!(
|
||||||
// protocol facts above, but do not turn provider/client headers into an audit store.
|
usage.client_response_headers.as_ref().unwrap()["content-type"],
|
||||||
assert!(usage.client_response_headers.is_none());
|
json!("application/json")
|
||||||
assert!(usage.response_headers.is_none());
|
);
|
||||||
|
assert_eq!(
|
||||||
|
usage.response_headers.as_ref().unwrap()["content-type"],
|
||||||
|
json!("text/html")
|
||||||
|
);
|
||||||
assert!(
|
assert!(
|
||||||
usage.response_body.is_none(),
|
usage.response_body.is_none(),
|
||||||
"upstream redirect did not include a body"
|
"upstream redirect did not include a body"
|
||||||
);
|
);
|
||||||
assert!(usage.client_response_body.is_none());
|
assert_eq!(usage.client_response_body.as_ref(), Some(&body_json));
|
||||||
let candidates = request_candidate_repository
|
let candidates = request_candidate_repository
|
||||||
.list_by_request_id("req-remote-runtime-stream-redirect")
|
.list_by_request_id("req-remote-runtime-stream-redirect")
|
||||||
.await
|
.await
|
||||||
@@ -14682,9 +14686,10 @@ mod tests {
|
|||||||
candidate_extra["upstream_response"]["status_code"],
|
candidate_extra["upstream_response"]["status_code"],
|
||||||
json!(302)
|
json!(302)
|
||||||
);
|
);
|
||||||
assert!(candidate_extra["upstream_response"]
|
assert_eq!(
|
||||||
.get("headers")
|
candidate_extra["upstream_response"]["headers"]["location"],
|
||||||
.is_none());
|
"/"
|
||||||
|
);
|
||||||
assert!(candidate_extra["upstream_response"].get("body").is_none());
|
assert!(candidate_extra["upstream_response"].get("body").is_none());
|
||||||
assert!(candidate_extra.get("client_response").is_none());
|
assert!(candidate_extra.get("client_response").is_none());
|
||||||
|
|
||||||
|
|||||||
@@ -21,10 +21,7 @@ use aether_contracts::{
|
|||||||
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
|
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
|
||||||
};
|
};
|
||||||
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
|
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
|
||||||
use aether_http::{
|
use aether_http::{apply_http_client_config, is_private_or_reserved_ip, HttpClientConfig};
|
||||||
apply_http_client_config, is_https_or_loopback_http_url, is_ipv4_benchmarking_fake_ip,
|
|
||||||
is_private_or_reserved_ip, HttpClientConfig,
|
|
||||||
};
|
|
||||||
use aether_runtime::{MetricKind, MetricSample};
|
use aether_runtime::{MetricKind, MetricSample};
|
||||||
use axum::body::Bytes;
|
use axum::body::Bytes;
|
||||||
use base64::Engine as _;
|
use base64::Engine as _;
|
||||||
@@ -63,8 +60,6 @@ use crate::upstream_admission::UpstreamTargetAdmissionPermit;
|
|||||||
use crate::{AppState, GatewayError};
|
use crate::{AppState, GatewayError};
|
||||||
|
|
||||||
const HUB_RELAY_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope";
|
const HUB_RELAY_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope";
|
||||||
pub(crate) const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY: &str =
|
|
||||||
aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY;
|
|
||||||
const HUB_RELAY_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
const HUB_RELAY_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
||||||
const MAX_SAFE_REDIRECTS: usize = 10;
|
const MAX_SAFE_REDIRECTS: usize = 10;
|
||||||
const MAX_UPSTREAM_ERROR_DETAIL_BYTES: usize = 2_048;
|
const MAX_UPSTREAM_ERROR_DETAIL_BYTES: usize = 2_048;
|
||||||
@@ -442,209 +437,12 @@ struct DirectHyperH2cSenderCacheMetrics {
|
|||||||
static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock<DirectHyperH2cSenderCacheMetrics> =
|
static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock<DirectHyperH2cSenderCacheMetrics> =
|
||||||
LazyLock::new(DirectHyperH2cSenderCacheMetrics::default);
|
LazyLock::new(DirectHyperH2cSenderCacheMetrics::default);
|
||||||
|
|
||||||
/// DNS resolver used for direct provider connections.
|
|
||||||
///
|
|
||||||
/// Provider endpoint URLs are frequently user/configuration supplied. The
|
|
||||||
/// platform resolver may return a different answer on every lookup, so merely
|
|
||||||
/// checking a URL's host (or resolving it once before constructing a client)
|
|
||||||
/// is not sufficient to prevent DNS rebinding. This resolver validates every
|
|
||||||
/// answer at the point reqwest/wreq asks for it. Explicit loopback targets are
|
|
||||||
/// retained for the supported local-provider workflow, but a hostname that is
|
|
||||||
/// not itself `localhost` can never resolve to a loopback/private address.
|
|
||||||
#[derive(Debug, Clone, Copy, Default)]
|
#[derive(Debug, Clone, Copy, Default)]
|
||||||
struct ExecutionSafeDnsResolver;
|
struct ExecutionSafeDnsResolver;
|
||||||
|
|
||||||
/// Resolver adapter for the legacy Hyper client retained for compatibility
|
|
||||||
/// with the non-fast-path H2C cache. Keep this path subject to the same
|
|
||||||
/// private-address and rebinding checks as reqwest/wreq clients.
|
|
||||||
#[derive(Debug, Clone, Copy, Default)]
|
#[derive(Debug, Clone, Copy, Default)]
|
||||||
struct ExecutionSafeHyperDnsResolver;
|
struct ExecutionSafeHyperDnsResolver;
|
||||||
|
|
||||||
// Local DNS interception tools may use RFC 2544's 198.18.0.0/15 range for
|
|
||||||
// synthetic answers. This exception is deliberately an allowlist rather
|
|
||||||
// than a property of the address range itself: a custom provider hostname
|
|
||||||
// must not be able to turn a local synthetic mapping into an SSRF primitive.
|
|
||||||
// Keep this list limited to origins that Aether constructs as built-in
|
|
||||||
// provider/model-fetch targets. In particular, do not use a
|
|
||||||
// suffix match for ordinary hosts (for example, `evil.chatgpt.com`).
|
|
||||||
const TRUSTED_EXECUTION_BENCHMARKING_DNS_EXACT_HOSTS: &[&str] = &[
|
|
||||||
"aiplatform.googleapis.com",
|
|
||||||
"antigravity.googleapis.com",
|
|
||||||
"api.openai.com",
|
|
||||||
"api.anthropic.com",
|
|
||||||
"api.deepseek.com",
|
|
||||||
"chatgpt.com",
|
|
||||||
"cloudcode-pa.googleapis.com",
|
|
||||||
"daily-cloudcode-pa.googleapis.com",
|
|
||||||
"daily-cloudcode-pa.sandbox.googleapis.com",
|
|
||||||
"dashscope.aliyuncs.com",
|
|
||||||
"generativelanguage.googleapis.com",
|
|
||||||
"grok.com",
|
|
||||||
"open.bigmodel.cn",
|
|
||||||
"q.us-iso-east-1.c2s.ic.gov",
|
|
||||||
"q.us-isob-east-1.sc2s.sgov.gov",
|
|
||||||
"q.us-isof-east-1.csp.hci.ic.gov",
|
|
||||||
"q.us-isof-south-1.csp.hci.ic.gov",
|
|
||||||
"server.codeium.com",
|
|
||||||
];
|
|
||||||
|
|
||||||
const TRUSTED_EXECUTION_VERTEX_DNS_REGIONS: &[&str] = &[
|
|
||||||
"africa-south1",
|
|
||||||
"asia-east1",
|
|
||||||
"asia-east2",
|
|
||||||
"asia-northeast1",
|
|
||||||
"asia-northeast2",
|
|
||||||
"asia-northeast3",
|
|
||||||
"asia-south1",
|
|
||||||
"asia-south2",
|
|
||||||
"asia-southeast1",
|
|
||||||
"asia-southeast2",
|
|
||||||
"australia-southeast1",
|
|
||||||
"australia-southeast2",
|
|
||||||
"europe-central2",
|
|
||||||
"europe-north1",
|
|
||||||
"europe-southwest1",
|
|
||||||
"europe-west1",
|
|
||||||
"europe-west2",
|
|
||||||
"europe-west3",
|
|
||||||
"europe-west4",
|
|
||||||
"europe-west6",
|
|
||||||
"europe-west8",
|
|
||||||
"europe-west9",
|
|
||||||
"europe-west10",
|
|
||||||
"europe-west12",
|
|
||||||
"me-central1",
|
|
||||||
"me-central2",
|
|
||||||
"me-west1",
|
|
||||||
"northamerica-northeast1",
|
|
||||||
"northamerica-northeast2",
|
|
||||||
"southamerica-east1",
|
|
||||||
"southamerica-west1",
|
|
||||||
"us-central1",
|
|
||||||
"us-east1",
|
|
||||||
"us-east4",
|
|
||||||
"us-east5",
|
|
||||||
"us-south1",
|
|
||||||
"us-west1",
|
|
||||||
"us-west2",
|
|
||||||
"us-west3",
|
|
||||||
"us-west4",
|
|
||||||
];
|
|
||||||
|
|
||||||
const TRUSTED_EXECUTION_AWS_DNS_REGIONS: &[&str] = &[
|
|
||||||
"af-south-1",
|
|
||||||
"ap-east-1",
|
|
||||||
"ap-northeast-1",
|
|
||||||
"ap-northeast-2",
|
|
||||||
"ap-northeast-3",
|
|
||||||
"ap-south-1",
|
|
||||||
"ap-south-2",
|
|
||||||
"ap-southeast-1",
|
|
||||||
"ap-southeast-2",
|
|
||||||
"ap-southeast-3",
|
|
||||||
"ap-southeast-4",
|
|
||||||
"ca-central-1",
|
|
||||||
"ca-west-1",
|
|
||||||
"eu-central-1",
|
|
||||||
"eu-central-2",
|
|
||||||
"eu-north-1",
|
|
||||||
"eu-south-1",
|
|
||||||
"eu-south-2",
|
|
||||||
"eu-west-1",
|
|
||||||
"eu-west-2",
|
|
||||||
"eu-west-3",
|
|
||||||
"il-central-1",
|
|
||||||
"me-central-1",
|
|
||||||
"me-south-1",
|
|
||||||
"mx-central-1",
|
|
||||||
"sa-east-1",
|
|
||||||
"us-east-1",
|
|
||||||
"us-east-2",
|
|
||||||
"us-gov-east-1",
|
|
||||||
"us-gov-west-1",
|
|
||||||
"us-west-1",
|
|
||||||
"us-west-2",
|
|
||||||
];
|
|
||||||
|
|
||||||
static EXECUTION_EXTRA_TRUSTED_DNS_HOSTS: LazyLock<StdRwLock<BTreeSet<String>>> =
|
|
||||||
LazyLock::new(|| StdRwLock::new(BTreeSet::new()));
|
|
||||||
|
|
||||||
pub(crate) fn refresh_execution_extra_trusted_dns_hosts(value: Option<&Value>) {
|
|
||||||
let hosts = value
|
|
||||||
.cloned()
|
|
||||||
.and_then(|value| {
|
|
||||||
aether_admin::system::normalize_execution_extra_trusted_dns_hosts_config_value(value)
|
|
||||||
.ok()
|
|
||||||
})
|
|
||||||
.and_then(|value| {
|
|
||||||
value.as_array().map(|hosts| {
|
|
||||||
hosts
|
|
||||||
.iter()
|
|
||||||
.filter_map(Value::as_str)
|
|
||||||
.map(ToOwned::to_owned)
|
|
||||||
.collect::<BTreeSet<_>>()
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
if let Ok(mut current) = EXECUTION_EXTRA_TRUSTED_DNS_HOSTS.write() {
|
|
||||||
*current = hosts;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Return whether `host` is one of the fixed provider origins for which a
|
|
||||||
/// local RFC-2544 synthetic answer can be accepted. The resolver receives only
|
|
||||||
/// a hostname (not the URL scheme/path), so all policy that can be expressed
|
|
||||||
/// here is intentionally host based. URL validation still requires HTTPS for
|
|
||||||
/// non-loopback upstreams before this resolver is used.
|
|
||||||
fn execution_host_allows_benchmarking_dns_answer(host: &str) -> bool {
|
|
||||||
let extra_hosts = EXECUTION_EXTRA_TRUSTED_DNS_HOSTS
|
|
||||||
.read()
|
|
||||||
.map(|hosts| hosts.clone())
|
|
||||||
.unwrap_or_default();
|
|
||||||
execution_host_allows_benchmarking_dns_answer_with_extra_hosts(host, &extra_hosts)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn execution_host_allows_benchmarking_dns_answer_with_extra_hosts(
|
|
||||||
host: &str,
|
|
||||||
extra_hosts: &BTreeSet<String>,
|
|
||||||
) -> bool {
|
|
||||||
let host = host.trim().trim_end_matches('.').to_ascii_lowercase();
|
|
||||||
if extra_hosts.contains(&host)
|
|
||||||
|| TRUSTED_EXECUTION_BENCHMARKING_DNS_EXACT_HOSTS
|
|
||||||
.iter()
|
|
||||||
.any(|trusted| *trusted == host)
|
|
||||||
{
|
|
||||||
return true;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Vertex service-account requests use `<region>-aiplatform.googleapis.com`.
|
|
||||||
// Keep this compatibility exception limited to known provider regions.
|
|
||||||
if let Some(region) = host.strip_suffix("-aiplatform.googleapis.com") {
|
|
||||||
return TRUSTED_EXECUTION_VERTEX_DNS_REGIONS.contains(®ion);
|
|
||||||
}
|
|
||||||
|
|
||||||
// Kiro uses a small, fixed set of regional service origins. Match each
|
|
||||||
// supported AWS partition explicitly; never use a broad suffix check that
|
|
||||||
// could accept an attacker-controlled subdomain.
|
|
||||||
matches_regional_service_host(&host, "q", ".amazonaws.com")
|
|
||||||
|| matches_regional_service_host(&host, "q-fips", ".amazonaws.com")
|
|
||||||
|| matches_regional_service_host(&host, "codewhisperer", ".amazonaws.com")
|
|
||||||
|| matches_regional_service_host(&host, "oidc", ".amazonaws.com")
|
|
||||||
|| matches_regional_service_host(&host, "prod", ".auth.desktop.kiro.dev")
|
|
||||||
}
|
|
||||||
|
|
||||||
fn matches_regional_service_host(host: &str, service: &str, suffix: &str) -> bool {
|
|
||||||
let Some(region) = host
|
|
||||||
.strip_prefix(service)
|
|
||||||
.and_then(|value| value.strip_prefix('.'))
|
|
||||||
.and_then(|value| value.strip_suffix(suffix))
|
|
||||||
else {
|
|
||||||
return false;
|
|
||||||
};
|
|
||||||
TRUSTED_EXECUTION_AWS_DNS_REGIONS.contains(®ion)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
|
fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
|
||||||
let host = host.trim_end_matches('.');
|
let host = host.trim_end_matches('.');
|
||||||
host.eq_ignore_ascii_case("localhost")
|
host.eq_ignore_ascii_case("localhost")
|
||||||
@@ -654,17 +452,10 @@ fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
|
|||||||
.unwrap_or(false)
|
.unwrap_or(false)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn validate_execution_dns_answers(
|
fn validate_resolved_execution_addresses(
|
||||||
host: &str,
|
host: &str,
|
||||||
addresses: Vec<SocketAddr>,
|
addresses: Vec<SocketAddr>,
|
||||||
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
provider_execution: bool,
|
||||||
validate_execution_dns_answers_with_policy(host, addresses, true)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn validate_execution_dns_answers_with_policy(
|
|
||||||
host: &str,
|
|
||||||
addresses: Vec<SocketAddr>,
|
|
||||||
allow_trusted_benchmarking_dns_answer: bool,
|
|
||||||
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
||||||
if addresses.is_empty() {
|
if addresses.is_empty() {
|
||||||
return Err(std::io::Error::new(
|
return Err(std::io::Error::new(
|
||||||
@@ -672,25 +463,22 @@ fn validate_execution_dns_answers_with_policy(
|
|||||||
"upstream DNS resolution returned no addresses",
|
"upstream DNS resolution returned no addresses",
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
if provider_execution {
|
||||||
|
return Ok(addresses);
|
||||||
|
}
|
||||||
let allows_loopback = dns_host_explicitly_allows_loopback(host);
|
let allows_loopback = dns_host_explicitly_allows_loopback(host);
|
||||||
let allows_benchmarking_dns_answer = allow_trusted_benchmarking_dns_answer
|
if addresses.iter().any(|address| {
|
||||||
&& execution_host_allows_benchmarking_dns_answer(host);
|
|
||||||
let unsafe_answer = addresses.iter().any(|address| {
|
|
||||||
if allows_loopback {
|
if allows_loopback {
|
||||||
!address.ip().is_loopback()
|
!address.ip().is_loopback()
|
||||||
} else {
|
} else {
|
||||||
is_private_or_reserved_ip(address.ip())
|
is_private_or_reserved_ip(address.ip())
|
||||||
&& !(allows_benchmarking_dns_answer && is_ipv4_benchmarking_fake_ip(address.ip()))
|
|
||||||
}
|
}
|
||||||
});
|
}) {
|
||||||
if unsafe_answer {
|
|
||||||
return Err(std::io::Error::new(
|
return Err(std::io::Error::new(
|
||||||
std::io::ErrorKind::PermissionDenied,
|
std::io::ErrorKind::PermissionDenied,
|
||||||
"upstream DNS resolution returned a private or reserved address",
|
"tunnel relay DNS resolution returned a private or reserved address",
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(addresses)
|
Ok(addresses)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -701,7 +489,7 @@ async fn resolve_execution_dns_addresses(host: &str) -> Result<Vec<SocketAddr>,
|
|||||||
async fn resolve_execution_target_addresses_with_policy(
|
async fn resolve_execution_target_addresses_with_policy(
|
||||||
host: &str,
|
host: &str,
|
||||||
port: u16,
|
port: u16,
|
||||||
allow_trusted_benchmarking_dns_answer: bool,
|
provider_execution: bool,
|
||||||
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
||||||
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
|
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
|
||||||
vec![SocketAddr::new(ip, port)]
|
vec![SocketAddr::new(ip, port)]
|
||||||
@@ -709,11 +497,7 @@ async fn resolve_execution_target_addresses_with_policy(
|
|||||||
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
|
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
|
||||||
.await?
|
.await?
|
||||||
};
|
};
|
||||||
validate_execution_dns_answers_with_policy(
|
validate_resolved_execution_addresses(host, addresses, provider_execution)
|
||||||
host,
|
|
||||||
addresses,
|
|
||||||
allow_trusted_benchmarking_dns_answer,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl reqwest::dns::Resolve for ExecutionSafeDnsResolver {
|
impl reqwest::dns::Resolve for ExecutionSafeDnsResolver {
|
||||||
@@ -3439,10 +3223,6 @@ async fn resolve_relay_target_addresses(
|
|||||||
let port = url.port_or_known_default().ok_or_else(|| {
|
let port = url.port_or_known_default().ok_or_else(|| {
|
||||||
ExecutionRuntimeTransportError::RelayError("tunnel relay URL has no port".to_string())
|
ExecutionRuntimeTransportError::RelayError("tunnel relay URL has no port".to_string())
|
||||||
})?;
|
})?;
|
||||||
// Relay destinations remain strict even when their hostname happens to be
|
|
||||||
// an official provider origin. The RFC-2544 compatibility exception is
|
|
||||||
// only for direct provider execution; allowing it here would weaken the
|
|
||||||
// relay SSRF guard.
|
|
||||||
let addresses = resolve_execution_target_addresses_with_policy(host, port, false)
|
let addresses = resolve_execution_target_addresses_with_policy(host, port, false)
|
||||||
.await
|
.await
|
||||||
.map_err(|error| match error.kind() {
|
.map_err(|error| match error.kind() {
|
||||||
@@ -5390,11 +5170,6 @@ fn validate_execution_upstream_url(
|
|||||||
"upstream URL must not include a fragment".to_string(),
|
"upstream URL must not include a fragment".to_string(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
if !is_https_or_loopback_http_url(&url) {
|
|
||||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
|
||||||
"remote upstream URL must use HTTPS".to_string(),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
let literal_ip = match url.host() {
|
let literal_ip = match url.host() {
|
||||||
Some(url::Host::Ipv4(address)) => Some(IpAddr::V4(address)),
|
Some(url::Host::Ipv4(address)) => Some(IpAddr::V4(address)),
|
||||||
Some(url::Host::Ipv6(address)) => Some(IpAddr::V6(address)),
|
Some(url::Host::Ipv6(address)) => Some(IpAddr::V6(address)),
|
||||||
@@ -5606,9 +5381,13 @@ mod tests {
|
|||||||
const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes";
|
const RELAY_TEST_SECRET: &str = "relay-test-secret-at-least-32-bytes";
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn execution_upstream_url_requires_https_or_literal_loopback_http() {
|
fn execution_upstream_url_accepts_http_and_https_with_safe_targets() {
|
||||||
for allowed in [
|
for allowed in [
|
||||||
"https://api.example.test/v1/responses?api-version=1",
|
"https://api.example.test/v1/responses?api-version=1",
|
||||||
|
"http://api.example.test:8080/v1/responses?api-version=1",
|
||||||
|
"http://8.8.8.8:8080/v1/responses",
|
||||||
|
"https://8.8.8.8/v1/responses",
|
||||||
|
"http://[2606:4700:4700::1111]:8080/v1/responses",
|
||||||
"http://localhost:8080/v1/responses",
|
"http://localhost:8080/v1/responses",
|
||||||
"http://127.42.0.1:8080/v1/responses",
|
"http://127.42.0.1:8080/v1/responses",
|
||||||
"http://[::1]:8080/v1/responses",
|
"http://[::1]:8080/v1/responses",
|
||||||
@@ -5620,7 +5399,6 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for rejected in [
|
for rejected in [
|
||||||
"http://api.example.test/v1/responses",
|
|
||||||
"http://10.0.0.1/v1/responses",
|
"http://10.0.0.1/v1/responses",
|
||||||
"http://0.0.0.0:8080/v1/responses",
|
"http://0.0.0.0:8080/v1/responses",
|
||||||
"http://[::ffff:127.0.0.1]:8080/v1/responses",
|
"http://[::ffff:127.0.0.1]:8080/v1/responses",
|
||||||
@@ -5628,6 +5406,8 @@ mod tests {
|
|||||||
"https://10.0.0.1:8443/v1/responses",
|
"https://10.0.0.1:8443/v1/responses",
|
||||||
"https://[email protected]/v1/responses",
|
"https://[email protected]/v1/responses",
|
||||||
"https://example.test/v1/responses#secret",
|
"https://example.test/v1/responses#secret",
|
||||||
|
"http://[email protected]/v1/responses",
|
||||||
|
"http://example.test/v1/responses#secret",
|
||||||
"ftp://localhost/resource",
|
"ftp://localhost/resource",
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
@@ -5650,115 +5430,75 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn execution_dns_answers_reject_private_addresses_and_allow_explicit_loopback() {
|
fn execution_dns_answers_allow_all_provider_hosts_without_address_filtering() {
|
||||||
let public = "93.184.216.34:443".parse().unwrap();
|
let addresses = vec![
|
||||||
let private = "10.0.0.8:443".parse().unwrap();
|
"198.18.78.41:443".parse().unwrap(),
|
||||||
let loopback_v4 = "127.0.0.1:8080".parse().unwrap();
|
"10.0.0.8:443".parse().unwrap(),
|
||||||
let loopback_v6 = "[::1]:8080".parse().unwrap();
|
"127.0.0.1:443".parse().unwrap(),
|
||||||
|
"169.254.169.254:443".parse().unwrap(),
|
||||||
assert!(super::validate_execution_dns_answers("api.example.test", vec![public]).is_ok());
|
"[fd00::1]:443".parse().unwrap(),
|
||||||
assert!(super::validate_execution_dns_answers("api.example.test", vec![private]).is_err());
|
"93.184.216.34:443".parse().unwrap(),
|
||||||
assert!(
|
];
|
||||||
super::validate_execution_dns_answers("localhost", vec![loopback_v4, loopback_v6])
|
|
||||||
.is_ok()
|
|
||||||
);
|
|
||||||
assert!(super::validate_execution_dns_answers("localhost", vec![private]).is_err());
|
|
||||||
assert!(super::validate_execution_dns_answers("api.example.test", Vec::new()).is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn execution_dns_answers_allow_benchmarking_range_only_for_fixed_provider_hosts() {
|
|
||||||
let fake = "198.18.75.234:443".parse().unwrap();
|
|
||||||
for host in [
|
for host in [
|
||||||
"api.openai.com",
|
"oauth2.googleapis.com",
|
||||||
"CHATGPT.COM.",
|
"www.googleapis.com",
|
||||||
"us-central1-aiplatform.googleapis.com",
|
"custom.example.test",
|
||||||
"me-central2-aiplatform.googleapis.com",
|
|
||||||
"q.us-east-1.amazonaws.com",
|
|
||||||
"q-fips.us-gov-west-1.amazonaws.com",
|
|
||||||
"codewhisperer.us-west-2.amazonaws.com",
|
|
||||||
"oidc.us-east-1.amazonaws.com",
|
|
||||||
"prod.us-east-1.auth.desktop.kiro.dev",
|
|
||||||
"q.us-iso-east-1.c2s.ic.gov",
|
|
||||||
"q.us-isob-east-1.sc2s.sgov.gov",
|
|
||||||
"q.us-isof-east-1.csp.hci.ic.gov",
|
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert_eq!(
|
||||||
super::validate_execution_dns_answers(host, vec![fake]).is_ok(),
|
super::validate_resolved_execution_addresses(host, addresses.clone(), true)
|
||||||
"fixed provider host should accept a benchmarking DNS answer: {host}"
|
.expect("provider DNS answers should pass through"),
|
||||||
);
|
addresses
|
||||||
}
|
|
||||||
|
|
||||||
for host in [
|
|
||||||
"api.example.test",
|
|
||||||
"evil.chatgpt.com",
|
|
||||||
"api.openai.com.evil.test",
|
|
||||||
"q.us-east-1.evil.amazonaws.com",
|
|
||||||
"q.us-east-1.amazonaws.com.attacker.test",
|
|
||||||
"q.localhost.amazonaws.com",
|
|
||||||
"evil-1-aiplatform.googleapis.com",
|
|
||||||
"q.evil-1.amazonaws.com",
|
|
||||||
"q-fips.evil-1.amazonaws.com",
|
|
||||||
"codewhisperer.evil-1.amazonaws.com",
|
|
||||||
"prod.evil-1.auth.desktop.kiro.dev",
|
|
||||||
"oidc.evil-1.amazonaws.com",
|
|
||||||
"q.us-central1.amazonaws.com",
|
|
||||||
"us-east-1-aiplatform.googleapis.com",
|
|
||||||
"q.us-east-1.c2s.ic.gov",
|
|
||||||
"q.us-iso-east-1.sc2s.sgov.gov",
|
|
||||||
"q-fips.us-gov-west-1.evil.amazonaws.com",
|
|
||||||
"codewhisperer.us-west-2.evil.amazonaws.com",
|
|
||||||
"oidc.us-east-1.evil.amazonaws.com",
|
|
||||||
"prod.us-east-1.auth.desktop.kiro.dev.attacker.test",
|
|
||||||
"prod.us-east-1.evil.auth.desktop.kiro.dev",
|
|
||||||
"q.us-iso-east-1.evil.c2s.ic.gov",
|
|
||||||
"q.us-iso-east-1.c2s.ic.gov.attacker.test",
|
|
||||||
"q.us-iso-east-1.c2s.ic.gov.evil",
|
|
||||||
"198.18.75.234",
|
|
||||||
] {
|
|
||||||
assert!(
|
|
||||||
super::validate_execution_dns_answers(host, vec![fake]).is_err(),
|
|
||||||
"untrusted or lookalike host must reject a benchmarking DNS answer: {host}"
|
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn execution_dns_answers_allow_benchmarking_range_for_configured_exact_hosts() {
|
fn execution_dns_answers_keep_relay_address_filtering() {
|
||||||
let fake = "198.18.75.234:443".parse().unwrap();
|
|
||||||
super::refresh_execution_extra_trusted_dns_hosts(Some(&json!(["custom.example.com",])));
|
|
||||||
|
|
||||||
assert!(super::validate_execution_dns_answers("custom.example.com", vec![fake]).is_ok());
|
|
||||||
assert!(
|
|
||||||
super::validate_execution_dns_answers("api.custom.example.com", vec![fake]).is_err()
|
|
||||||
);
|
|
||||||
|
|
||||||
super::refresh_execution_extra_trusted_dns_hosts(None);
|
|
||||||
assert!(super::validate_execution_dns_answers("custom.example.com", vec![fake]).is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn execution_dns_answers_reject_mixed_private_results_and_strict_relay_policy() {
|
|
||||||
let fake = "198.18.75.234:443".parse().unwrap();
|
|
||||||
let public = "93.184.216.34:443".parse().unwrap();
|
let public = "93.184.216.34:443".parse().unwrap();
|
||||||
let private = "10.0.0.8:443".parse().unwrap();
|
for host in ["oauth2.googleapis.com", "custom.example.test"] {
|
||||||
|
assert!(
|
||||||
// A trusted host may have a synthetic answer alongside a genuine public
|
super::validate_resolved_execution_addresses(host, vec![public], false).is_ok()
|
||||||
// answer, but any real private answer still fails closed.
|
);
|
||||||
|
for blocked in [
|
||||||
|
"198.18.78.41:443",
|
||||||
|
"10.0.0.8:443",
|
||||||
|
"127.0.0.1:443",
|
||||||
|
"169.254.169.254:443",
|
||||||
|
"[fd00::1]:443",
|
||||||
|
] {
|
||||||
|
let blocked = blocked.parse().unwrap();
|
||||||
|
assert!(
|
||||||
|
super::validate_resolved_execution_addresses(host, vec![blocked], false)
|
||||||
|
.is_err()
|
||||||
|
);
|
||||||
|
assert!(super::validate_resolved_execution_addresses(
|
||||||
|
host,
|
||||||
|
vec![public, blocked],
|
||||||
|
false
|
||||||
|
)
|
||||||
|
.is_err());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let loopback = vec![
|
||||||
|
"127.0.0.1:443".parse().unwrap(),
|
||||||
|
"[::1]:443".parse().unwrap(),
|
||||||
|
];
|
||||||
|
assert!(super::validate_resolved_execution_addresses("localhost", loopback, false).is_ok());
|
||||||
assert!(
|
assert!(
|
||||||
super::validate_execution_dns_answers("api.openai.com", vec![fake, public]).is_ok()
|
super::validate_resolved_execution_addresses("localhost", vec![public], false).is_err()
|
||||||
);
|
);
|
||||||
assert!(
|
for provider_execution in [false, true] {
|
||||||
super::validate_execution_dns_answers("api.openai.com", vec![fake, private]).is_err()
|
assert_eq!(
|
||||||
);
|
super::validate_resolved_execution_addresses(
|
||||||
|
"custom.example.test",
|
||||||
// Tunnel relay resolution opts out of the compatibility exception.
|
Vec::new(),
|
||||||
assert!(super::validate_execution_dns_answers_with_policy(
|
provider_execution
|
||||||
"api.openai.com",
|
)
|
||||||
vec![fake],
|
.expect_err("empty DNS answers must fail")
|
||||||
false,
|
.kind(),
|
||||||
)
|
std::io::ErrorKind::NotFound
|
||||||
.is_err());
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -522,7 +522,6 @@ where
|
|||||||
decision,
|
decision,
|
||||||
plan_kind,
|
plan_kind,
|
||||||
transfer_tracker,
|
transfer_tracker,
|
||||||
request_first_byte_started_at: Instant::now(),
|
|
||||||
};
|
};
|
||||||
let loop_result = run_ai_attempt_loop(&port, plan_and_reports).await;
|
let loop_result = run_ai_attempt_loop(&port, plan_and_reports).await;
|
||||||
if loop_result.is_err() {
|
if loop_result.is_err() {
|
||||||
@@ -603,7 +602,6 @@ where
|
|||||||
decision,
|
decision,
|
||||||
plan_kind,
|
plan_kind,
|
||||||
transfer_tracker,
|
transfer_tracker,
|
||||||
request_first_byte_started_at: Instant::now(),
|
|
||||||
};
|
};
|
||||||
let loop_result = run_dynamic_attempt_loop(
|
let loop_result = run_dynamic_attempt_loop(
|
||||||
&port,
|
&port,
|
||||||
@@ -1121,10 +1119,6 @@ struct StreamAttemptLoopPort<'a> {
|
|||||||
decision: &'a GatewayControlDecision,
|
decision: &'a GatewayControlDecision,
|
||||||
plan_kind: &'a str,
|
plan_kind: &'a str,
|
||||||
transfer_tracker: &'a ProviderTransferTracker,
|
transfer_tracker: &'a ProviderTransferTracker,
|
||||||
/// All candidates in one downstream stream request share this origin.
|
|
||||||
/// Without it every retry receives a fresh full first-byte timeout and a
|
|
||||||
/// 30-second provider timeout can accumulate into a 60-120 second stall.
|
|
||||||
request_first_byte_started_at: Instant,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
@@ -1254,7 +1248,6 @@ where
|
|||||||
self.plan_kind,
|
self.plan_kind,
|
||||||
plan,
|
plan,
|
||||||
watchdog_report_context,
|
watchdog_report_context,
|
||||||
self.request_first_byte_started_at,
|
|
||||||
stop_on_transport_errors,
|
stop_on_transport_errors,
|
||||||
move || async move {
|
move || async move {
|
||||||
if let Some(response) = execution_plan_cost_capacity_response(
|
if let Some(response) = execution_plan_cost_capacity_response(
|
||||||
@@ -1308,7 +1301,7 @@ where
|
|||||||
http::StatusCode::GATEWAY_TIMEOUT.as_u16(),
|
http::StatusCode::GATEWAY_TIMEOUT.as_u16(),
|
||||||
"local_stream_candidate_watchdog_timeout",
|
"local_stream_candidate_watchdog_timeout",
|
||||||
stream_candidate_watchdog_timeout_message(),
|
stream_candidate_watchdog_timeout_message(),
|
||||||
self.request_first_byte_started_at.elapsed().as_millis() as u64,
|
watchdog_started_at.elapsed().as_millis() as u64,
|
||||||
)
|
)
|
||||||
.await?,
|
.await?,
|
||||||
)
|
)
|
||||||
@@ -1758,7 +1751,6 @@ async fn execute_stream_candidate_with_watchdog<Fut>(
|
|||||||
plan_kind: &str,
|
plan_kind: &str,
|
||||||
plan: &aether_contracts::ExecutionPlan,
|
plan: &aether_contracts::ExecutionPlan,
|
||||||
report_context: Option<&serde_json::Value>,
|
report_context: Option<&serde_json::Value>,
|
||||||
request_first_byte_started_at: Instant,
|
|
||||||
stop_on_transport_errors: bool,
|
stop_on_transport_errors: bool,
|
||||||
execute: impl FnOnce() -> Fut,
|
execute: impl FnOnce() -> Fut,
|
||||||
) -> Result<StreamCandidateWatchdogOutcome, GatewayError>
|
) -> Result<StreamCandidateWatchdogOutcome, GatewayError>
|
||||||
@@ -1768,7 +1760,6 @@ where
|
|||||||
> + Send,
|
> + Send,
|
||||||
{
|
{
|
||||||
let timeout_duration = resolve_stream_candidate_watchdog_timeout(plan, report_context);
|
let timeout_duration = resolve_stream_candidate_watchdog_timeout(plan, report_context);
|
||||||
let request_first_byte_deadline = request_first_byte_started_at + timeout_duration;
|
|
||||||
let candidate_started_at = std::time::Instant::now();
|
let candidate_started_at = std::time::Instant::now();
|
||||||
let candidate_started_unix_ms = current_unix_ms();
|
let candidate_started_unix_ms = current_unix_ms();
|
||||||
let permit = match acquire_upstream_execution_gate(state, trace_id).await {
|
let permit = match acquire_upstream_execution_gate(state, trace_id).await {
|
||||||
@@ -1794,14 +1785,7 @@ where
|
|||||||
let watchdog_progress = StreamCandidateWatchdogProgress::shared();
|
let watchdog_progress = StreamCandidateWatchdogProgress::shared();
|
||||||
let execution = watchdog_progress.clone().scope(execute());
|
let execution = watchdog_progress.clone().scope(execute());
|
||||||
tokio::pin!(execution);
|
tokio::pin!(execution);
|
||||||
// This is an absolute request-level deadline, not a new timeout for this
|
let deadline = tokio::time::sleep(timeout_duration);
|
||||||
// candidate. Retries therefore consume only the budget left by earlier
|
|
||||||
// candidates instead of resetting the full provider timeout.
|
|
||||||
let candidate_budget_ms = request_first_byte_deadline
|
|
||||||
.saturating_duration_since(Instant::now())
|
|
||||||
.as_millis()
|
|
||||||
.min(u128::from(u64::MAX)) as u64;
|
|
||||||
let deadline = tokio::time::sleep_until(request_first_byte_deadline);
|
|
||||||
tokio::pin!(deadline);
|
tokio::pin!(deadline);
|
||||||
let execution_result = tokio::select! {
|
let execution_result = tokio::select! {
|
||||||
biased;
|
biased;
|
||||||
@@ -1830,10 +1814,6 @@ where
|
|||||||
.map(|value| value.to_string())
|
.map(|value| value.to_string())
|
||||||
.unwrap_or_else(|| "-".to_string());
|
.unwrap_or_else(|| "-".to_string());
|
||||||
let timeout_ms = u64::try_from(timeout_duration.as_millis()).unwrap_or(u64::MAX);
|
let timeout_ms = u64::try_from(timeout_duration.as_millis()).unwrap_or(u64::MAX);
|
||||||
let request_elapsed_ms = request_first_byte_started_at
|
|
||||||
.elapsed()
|
|
||||||
.as_millis()
|
|
||||||
.min(u128::from(u64::MAX)) as u64;
|
|
||||||
record_local_request_candidate_status(
|
record_local_request_candidate_status(
|
||||||
state,
|
state,
|
||||||
plan,
|
plan,
|
||||||
@@ -1862,8 +1842,6 @@ where
|
|||||||
model_name,
|
model_name,
|
||||||
candidate_index = candidate_index.as_str(),
|
candidate_index = candidate_index.as_str(),
|
||||||
timeout_ms,
|
timeout_ms,
|
||||||
candidate_budget_ms,
|
|
||||||
request_elapsed_ms,
|
|
||||||
"gateway local stream candidate watchdog timed out"
|
"gateway local stream candidate watchdog timed out"
|
||||||
);
|
);
|
||||||
if stop_on_transport_errors {
|
if stop_on_transport_errors {
|
||||||
@@ -3150,7 +3128,6 @@ mod tests {
|
|||||||
"claude_cli_stream",
|
"claude_cli_stream",
|
||||||
&plan,
|
&plan,
|
||||||
Some(&report_context),
|
Some(&report_context),
|
||||||
Instant::now(),
|
|
||||||
false,
|
false,
|
||||||
|| {
|
|| {
|
||||||
std::future::pending::<
|
std::future::pending::<
|
||||||
@@ -3189,39 +3166,32 @@ mod tests {
|
|||||||
assert_eq!(record.candidate_index, 2);
|
assert_eq!(record.candidate_index, 2);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
async fn assert_stream_candidate_retry_gets_fresh_first_byte_budget(
|
||||||
async fn stream_candidate_retry_does_not_reset_an_expired_request_first_byte_budget() {
|
provider_id: &str,
|
||||||
let writer = Arc::new(TestRequestCandidateWriter::default());
|
key_id: &str,
|
||||||
|
first_byte_ms: u64,
|
||||||
|
) {
|
||||||
|
let writer = TestRequestCandidateWriter::default();
|
||||||
let plan = test_plan(Some(ExecutionTimeouts {
|
let plan = test_plan(Some(ExecutionTimeouts {
|
||||||
first_byte_ms: Some(250),
|
first_byte_ms: Some(100),
|
||||||
..ExecutionTimeouts::default()
|
..ExecutionTimeouts::default()
|
||||||
}));
|
}));
|
||||||
let report_context = test_report_context();
|
let report_context = test_report_context();
|
||||||
// Stand in for earlier candidates having already consumed the request's
|
|
||||||
// complete first-byte budget. A per-candidate watchdog would wait a new
|
|
||||||
// 250 ms here; the shared absolute deadline must settle immediately.
|
|
||||||
let request_first_byte_started_at = Instant::now() - Duration::from_millis(300);
|
|
||||||
|
|
||||||
let result = tokio::time::timeout(
|
let result = execute_stream_candidate_with_watchdog(
|
||||||
Duration::from_millis(100),
|
&writer,
|
||||||
execute_stream_candidate_with_watchdog(
|
"trace_watchdog_retry_budget",
|
||||||
writer.as_ref(),
|
"claude_cli_stream",
|
||||||
"trace_watchdog_shared_budget",
|
&plan,
|
||||||
"claude_cli_stream",
|
Some(&report_context),
|
||||||
&plan,
|
false,
|
||||||
Some(&report_context),
|
|| {
|
||||||
request_first_byte_started_at,
|
std::future::pending::<
|
||||||
false,
|
Result<AiAttemptExecutionOutcome<Response<Body>>, GatewayError>,
|
||||||
|| {
|
>()
|
||||||
std::future::pending::<
|
},
|
||||||
Result<AiAttemptExecutionOutcome<Response<Body>>, GatewayError>,
|
|
||||||
>()
|
|
||||||
},
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
.await
|
.await;
|
||||||
.expect("an expired request-level first-byte budget must not restart per candidate");
|
|
||||||
|
|
||||||
assert!(matches!(
|
assert!(matches!(
|
||||||
result,
|
result,
|
||||||
Ok(StreamCandidateWatchdogOutcome::Executed(
|
Ok(StreamCandidateWatchdogOutcome::Executed(
|
||||||
@@ -3231,14 +3201,119 @@ mod tests {
|
|||||||
}
|
}
|
||||||
))
|
))
|
||||||
));
|
));
|
||||||
|
|
||||||
|
let mut next_plan = plan.clone();
|
||||||
|
next_plan.candidate_id = Some("cand_watchdog_retry".to_string());
|
||||||
|
next_plan.provider_id = provider_id.to_string();
|
||||||
|
next_plan.key_id = key_id.to_string();
|
||||||
|
next_plan.timeouts = Some(ExecutionTimeouts {
|
||||||
|
first_byte_ms: Some(first_byte_ms),
|
||||||
|
..ExecutionTimeouts::default()
|
||||||
|
});
|
||||||
|
let mut next_report_context = report_context.clone();
|
||||||
|
next_report_context["candidate_id"] = json!("cand_watchdog_retry");
|
||||||
|
next_report_context["candidate_index"] = json!(3);
|
||||||
|
|
||||||
|
let result = execute_stream_candidate_with_watchdog(
|
||||||
|
&writer,
|
||||||
|
"trace_watchdog_retry_budget",
|
||||||
|
"claude_cli_stream",
|
||||||
|
&next_plan,
|
||||||
|
Some(&next_report_context),
|
||||||
|
false,
|
||||||
|
|| async {
|
||||||
|
tokio::time::sleep(Duration::from_millis(60)).await;
|
||||||
|
Ok(AiAttemptExecutionOutcome::Responded(Response::new(
|
||||||
|
Body::from("retry succeeded"),
|
||||||
|
)))
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
assert!(
|
||||||
|
matches!(
|
||||||
|
result,
|
||||||
|
Ok(StreamCandidateWatchdogOutcome::Executed(
|
||||||
|
AiAttemptExecutionOutcome::Responded(_)
|
||||||
|
))
|
||||||
|
),
|
||||||
|
"candidate {provider_id}/{key_id} must receive its own {first_byte_ms} ms budget"
|
||||||
|
);
|
||||||
|
|
||||||
let records = writer.records.lock().await;
|
let records = writer.records.lock().await;
|
||||||
assert_eq!(records.len(), 1);
|
assert_eq!(records.len(), 1);
|
||||||
|
assert_eq!(records[0].id, plan.candidate_id.as_deref().unwrap());
|
||||||
|
assert_eq!(records[0].status, RequestCandidateStatus::Failed);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
records[0].error_type.as_deref(),
|
records[0].error_type.as_deref(),
|
||||||
Some("local_stream_candidate_watchdog_timeout")
|
Some("local_stream_candidate_watchdog_timeout")
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn stream_candidate_watchdog_failover_gets_fresh_first_byte_budget() {
|
||||||
|
for first_byte_ms in [100, 75, 150] {
|
||||||
|
assert_stream_candidate_retry_gets_fresh_first_byte_budget(
|
||||||
|
"provider_next",
|
||||||
|
"key_next",
|
||||||
|
first_byte_ms,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn stream_candidate_watchdog_same_provider_retries_get_fresh_first_byte_budget() {
|
||||||
|
for key_id in ["key_next", "key_id"] {
|
||||||
|
assert_stream_candidate_retry_gets_fresh_first_byte_budget("provider_id", key_id, 100)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn stream_candidate_watchdog_starts_first_byte_budget_after_admission() {
|
||||||
|
let writer = TestRequestCandidateWriter::with_upstream_gate(1, Duration::from_secs(1));
|
||||||
|
let held_permit = writer
|
||||||
|
.upstream_gate
|
||||||
|
.as_ref()
|
||||||
|
.expect("test gate should exist")
|
||||||
|
.try_acquire()
|
||||||
|
.expect("test gate permit should acquire");
|
||||||
|
let plan = test_plan(Some(ExecutionTimeouts {
|
||||||
|
first_byte_ms: Some(50),
|
||||||
|
..ExecutionTimeouts::default()
|
||||||
|
}));
|
||||||
|
let report_context = test_report_context();
|
||||||
|
|
||||||
|
let (result, ()) = tokio::join!(
|
||||||
|
execute_stream_candidate_with_watchdog(
|
||||||
|
&writer,
|
||||||
|
"trace_watchdog_admission_budget",
|
||||||
|
"claude_cli_stream",
|
||||||
|
&plan,
|
||||||
|
Some(&report_context),
|
||||||
|
false,
|
||||||
|
|| async {
|
||||||
|
tokio::time::sleep(Duration::from_millis(20)).await;
|
||||||
|
Ok(AiAttemptExecutionOutcome::Responded(Response::new(
|
||||||
|
Body::from("admitted candidate succeeded"),
|
||||||
|
)))
|
||||||
|
},
|
||||||
|
),
|
||||||
|
async move {
|
||||||
|
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||||
|
drop(held_permit);
|
||||||
|
},
|
||||||
|
);
|
||||||
|
|
||||||
|
assert!(matches!(
|
||||||
|
result,
|
||||||
|
Ok(StreamCandidateWatchdogOutcome::Executed(
|
||||||
|
AiAttemptExecutionOutcome::Responded(_)
|
||||||
|
))
|
||||||
|
));
|
||||||
|
assert!(writer.records.lock().await.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn stream_candidate_watchdog_can_stop_on_transport_error() {
|
async fn stream_candidate_watchdog_can_stop_on_transport_error() {
|
||||||
let writer = Arc::new(TestRequestCandidateWriter::default());
|
let writer = Arc::new(TestRequestCandidateWriter::default());
|
||||||
@@ -3254,7 +3329,6 @@ mod tests {
|
|||||||
"claude_cli_stream",
|
"claude_cli_stream",
|
||||||
&plan,
|
&plan,
|
||||||
Some(&report_context),
|
Some(&report_context),
|
||||||
Instant::now(),
|
|
||||||
true,
|
true,
|
||||||
|| {
|
|| {
|
||||||
std::future::pending::<
|
std::future::pending::<
|
||||||
@@ -3292,7 +3366,6 @@ mod tests {
|
|||||||
"claude_cli_stream",
|
"claude_cli_stream",
|
||||||
&plan,
|
&plan,
|
||||||
Some(&report_context),
|
Some(&report_context),
|
||||||
Instant::now(),
|
|
||||||
true,
|
true,
|
||||||
|| async {
|
|| async {
|
||||||
mark_stream_candidate_watchdog_terminal_started();
|
mark_stream_candidate_watchdog_terminal_started();
|
||||||
@@ -3325,7 +3398,6 @@ mod tests {
|
|||||||
"claude_cli_stream",
|
"claude_cli_stream",
|
||||||
&plan,
|
&plan,
|
||||||
Some(&report_context),
|
Some(&report_context),
|
||||||
Instant::now(),
|
|
||||||
true,
|
true,
|
||||||
|| async {
|
|| async {
|
||||||
Err(GatewayError::UpstreamUnavailable {
|
Err(GatewayError::UpstreamUnavailable {
|
||||||
@@ -3365,7 +3437,6 @@ mod tests {
|
|||||||
"claude_cli_stream",
|
"claude_cli_stream",
|
||||||
&plan,
|
&plan,
|
||||||
Some(&report_context),
|
Some(&report_context),
|
||||||
Instant::now(),
|
|
||||||
false,
|
false,
|
||||||
|| async {
|
|| async {
|
||||||
panic!("execute future should not run while upstream execution gate is saturated")
|
panic!("execute future should not run while upstream execution gate is saturated")
|
||||||
@@ -3413,7 +3484,6 @@ mod tests {
|
|||||||
"claude_cli_stream",
|
"claude_cli_stream",
|
||||||
&plan,
|
&plan,
|
||||||
Some(&report_context),
|
Some(&report_context),
|
||||||
Instant::now(),
|
|
||||||
false,
|
false,
|
||||||
|| async {
|
|| async {
|
||||||
Err(GatewayError::AdmissionTimeout {
|
Err(GatewayError::AdmissionTimeout {
|
||||||
|
|||||||
@@ -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 {
|
fn runtime_miss_original_headers_json(headers: &HeaderMap) -> Value {
|
||||||
let mut headers = crate::headers::collect_control_headers(headers);
|
serde_json::to_value(crate::headers::collect_control_headers(headers))
|
||||||
for (name, value) in headers.iter_mut() {
|
.unwrap_or_else(|_| json!({}))
|
||||||
if runtime_miss_sensitive_header(name) {
|
|
||||||
*value = runtime_miss_mask_header_value(value);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
serde_json::to_value(headers).unwrap_or_else(|_| json!({}))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn runtime_miss_original_request_body_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(
|
async fn load_runtime_miss_candidate_contexts(
|
||||||
state: &AppState,
|
state: &AppState,
|
||||||
request_id: &str,
|
request_id: &str,
|
||||||
@@ -1233,8 +1194,9 @@ mod tests {
|
|||||||
apply_runtime_miss_usage_routing, beautify_local_execution_client_error_message,
|
apply_runtime_miss_usage_routing, beautify_local_execution_client_error_message,
|
||||||
insert_runtime_miss_candidate_usage_metadata,
|
insert_runtime_miss_candidate_usage_metadata,
|
||||||
request_candidate_represents_provider_execution, runtime_miss_client_error_body,
|
request_candidate_represents_provider_execution, runtime_miss_client_error_body,
|
||||||
select_last_runtime_miss_executed_candidate, select_last_runtime_miss_routing_candidate,
|
runtime_miss_original_headers_json, select_last_runtime_miss_executed_candidate,
|
||||||
LocalExecutionRuntimeMissContext, RuntimeMissCandidateContext,
|
select_last_runtime_miss_routing_candidate, LocalExecutionRuntimeMissContext,
|
||||||
|
RuntimeMissCandidateContext,
|
||||||
};
|
};
|
||||||
use crate::constants::EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS;
|
use crate::constants::EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS;
|
||||||
use crate::state::LocalExecutionRuntimeMissDiagnostic;
|
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]
|
#[test]
|
||||||
fn runtime_miss_usage_body_matches_claude_client_envelope() {
|
fn runtime_miss_usage_body_matches_claude_client_envelope() {
|
||||||
let claude = runtime_miss_client_error_body(Some("claude:messages"), "busy");
|
let claude = runtime_miss_client_error_body(Some("claude:messages"), "busy");
|
||||||
|
|||||||
@@ -659,7 +659,7 @@ async fn admin_monitoring_trace_request_exposes_request_path_from_usage_audit()
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn admin_monitoring_trace_request_redacts_failed_candidate_response_payloads() {
|
async fn admin_monitoring_trace_request_exposes_failed_candidate_response_payloads() {
|
||||||
let mut candidate = sample_candidate(
|
let mut candidate = sample_candidate(
|
||||||
"cand-used",
|
"cand-used",
|
||||||
"request-1",
|
"request-1",
|
||||||
@@ -746,14 +746,96 @@ async fn admin_monitoring_trace_request_redacts_failed_candidate_response_payloa
|
|||||||
json!("upstream_response")
|
json!("upstream_response")
|
||||||
);
|
);
|
||||||
assert!(extra["upstream_response"].get("headers").is_none());
|
assert!(extra["upstream_response"].get("headers").is_none());
|
||||||
assert!(extra["upstream_response"].get("body").is_none());
|
assert_eq!(
|
||||||
|
extra["upstream_response"]["body"]["error"]["message"],
|
||||||
|
"redirect blocked"
|
||||||
|
);
|
||||||
assert!(extra["upstream_response"].get("body_ref").is_none());
|
assert!(extra["upstream_response"].get("body_ref").is_none());
|
||||||
assert!(extra.get("client_response").is_none());
|
assert!(extra.get("client_response").is_none());
|
||||||
assert!(extra.get("provider_response").is_none());
|
assert!(extra.get("provider_response").is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_response_body() {
|
async fn admin_monitoring_trace_request_does_not_replace_attempt_status_with_usage_status() {
|
||||||
|
for (candidate_status, upstream_status, expected_status) in [
|
||||||
|
(Some(400), Some(400), Some(400)),
|
||||||
|
(Some(400), None, Some(400)),
|
||||||
|
(Some(502), Some(200), Some(200)),
|
||||||
|
(None, None, None),
|
||||||
|
] {
|
||||||
|
let mut candidate = sample_candidate(
|
||||||
|
"cand-used",
|
||||||
|
"request-failover-status",
|
||||||
|
0,
|
||||||
|
RequestCandidateStatus::Failed,
|
||||||
|
Some(101),
|
||||||
|
Some(33),
|
||||||
|
candidate_status,
|
||||||
|
);
|
||||||
|
if let Some(status_code) = upstream_status {
|
||||||
|
candidate.extra_data = Some(json!({
|
||||||
|
"upstream_response": {
|
||||||
|
"status_code": status_code,
|
||||||
|
"body": {"error": {"message": "sensitive upstream error"}}
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
let request_candidates =
|
||||||
|
Arc::new(InMemoryRequestCandidateRepository::seed(vec![candidate]));
|
||||||
|
let mut usage = sample_usage(
|
||||||
|
"request-failover-status",
|
||||||
|
"provider-1",
|
||||||
|
"OpenAI",
|
||||||
|
0,
|
||||||
|
0.0,
|
||||||
|
"failed",
|
||||||
|
Some(503),
|
||||||
|
100,
|
||||||
|
);
|
||||||
|
usage.candidate_id = Some("cand-used".to_string());
|
||||||
|
usage.response_body_state = Some(UsageBodyCaptureState::Reference);
|
||||||
|
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
|
||||||
|
let data_state =
|
||||||
|
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
|
||||||
|
request_candidates,
|
||||||
|
usage_repository,
|
||||||
|
);
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("state should build")
|
||||||
|
.with_data_state_for_tests(data_state);
|
||||||
|
let context = request_context(
|
||||||
|
http::Method::GET,
|
||||||
|
"/api/admin/monitoring/trace/request-failover-status",
|
||||||
|
);
|
||||||
|
|
||||||
|
let response = local_monitoring_response(&state, &context)
|
||||||
|
.await
|
||||||
|
.expect("handler should not error")
|
||||||
|
.expect("route should be handled locally");
|
||||||
|
assert_eq!(response.status(), http::StatusCode::OK);
|
||||||
|
let body = to_bytes(response.into_body(), usize::MAX)
|
||||||
|
.await
|
||||||
|
.expect("body should read");
|
||||||
|
let payload: serde_json::Value =
|
||||||
|
serde_json::from_slice(&body).expect("json body should parse");
|
||||||
|
let candidate = &payload["candidates"][0];
|
||||||
|
assert_eq!(candidate["status_code"], json!(candidate_status));
|
||||||
|
assert_eq!(
|
||||||
|
candidate["extra_data"]["upstream_response"]["status_code"],
|
||||||
|
json!(expected_status),
|
||||||
|
);
|
||||||
|
if upstream_status.is_some() {
|
||||||
|
assert_eq!(
|
||||||
|
candidate["extra_data"]["upstream_response"]["body"]["error"]["message"],
|
||||||
|
"sensitive upstream error"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
assert!(candidate["error_message"].is_null());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn admin_monitoring_trace_request_keeps_candidate_errors_without_hydrating_usage_bodies() {
|
||||||
let mut candidate = sample_candidate(
|
let mut candidate = sample_candidate(
|
||||||
"cand-used",
|
"cand-used",
|
||||||
"request-ref-body",
|
"request-ref-body",
|
||||||
@@ -771,8 +853,7 @@ async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_respon
|
|||||||
"x-request-id": "stale-request-like-body"
|
"x-request-id": "stale-request-like-body"
|
||||||
},
|
},
|
||||||
"body": {
|
"body": {
|
||||||
"model": "gpt-5.6-sol",
|
"error": {"message": "candidate-specific upstream failure"}
|
||||||
"input": [{"role": "user", "content": "request prompt"}]
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}));
|
}));
|
||||||
@@ -837,8 +918,14 @@ async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_respon
|
|||||||
assert_eq!(upstream_response["status_code"], json!(400));
|
assert_eq!(upstream_response["status_code"], json!(400));
|
||||||
assert_eq!(upstream_response["source"], json!("upstream_response"));
|
assert_eq!(upstream_response["source"], json!("upstream_response"));
|
||||||
assert_eq!(upstream_response["body_state"], json!("reference"));
|
assert_eq!(upstream_response["body_state"], json!("reference"));
|
||||||
assert!(upstream_response.get("headers").is_none());
|
assert_eq!(
|
||||||
assert!(upstream_response.get("body").is_none());
|
upstream_response["headers"]["content-type"],
|
||||||
|
"text/event-stream"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
upstream_response["body"]["error"]["message"],
|
||||||
|
"candidate-specific upstream failure"
|
||||||
|
);
|
||||||
assert!(upstream_response.get("body_ref").is_none());
|
assert!(upstream_response.get("body_ref").is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -17,7 +17,10 @@ use aether_admin::observability::usage::{
|
|||||||
admin_usage_bad_request_response, admin_usage_data_unavailable_response,
|
admin_usage_bad_request_response, admin_usage_data_unavailable_response,
|
||||||
admin_usage_provider_key_name, ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
|
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::{
|
use axum::{
|
||||||
body::Body,
|
body::Body,
|
||||||
http,
|
http,
|
||||||
@@ -28,9 +31,59 @@ use serde_json::{json, Value};
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
use tokio::try_join;
|
use tokio::try_join;
|
||||||
|
|
||||||
|
#[derive(Default)]
|
||||||
struct AdminUsageDetailBodyValue {
|
struct AdminUsageDetailBodyValue {
|
||||||
value: Option<Value>,
|
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(
|
async fn resolve_admin_usage_detail_request_body(
|
||||||
@@ -38,10 +91,7 @@ async fn resolve_admin_usage_detail_request_body(
|
|||||||
item: &StoredRequestUsageAudit,
|
item: &StoredRequestUsageAudit,
|
||||||
) -> AdminUsageDetailBodyValue {
|
) -> AdminUsageDetailBodyValue {
|
||||||
match admin_usage_resolve_request_capture_body_for_item(state, item, None).await {
|
match admin_usage_resolve_request_capture_body_for_item(state, item, None).await {
|
||||||
Ok(body) => AdminUsageDetailBodyValue {
|
Ok(body) => AdminUsageDetailBodyValue::resolved(item, UsageBodyField::RequestBody, body),
|
||||||
value: body,
|
|
||||||
load_failed: false,
|
|
||||||
},
|
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
tracing::warn!(
|
tracing::warn!(
|
||||||
error = ?err,
|
error = ?err,
|
||||||
@@ -52,7 +102,9 @@ async fn resolve_admin_usage_detail_request_body(
|
|||||||
);
|
);
|
||||||
let value = admin_usage_resolve_request_capture_body(item, None);
|
let value = admin_usage_resolve_request_capture_body(item, None);
|
||||||
AdminUsageDetailBodyValue {
|
AdminUsageDetailBodyValue {
|
||||||
load_failed: value.is_none(),
|
error_code: value
|
||||||
|
.is_none()
|
||||||
|
.then(|| admin_usage_body_load_error_code(&err)),
|
||||||
value,
|
value,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -66,10 +118,7 @@ async fn resolve_admin_usage_detail_body_value(
|
|||||||
) -> AdminUsageDetailBodyValue {
|
) -> AdminUsageDetailBodyValue {
|
||||||
let inline_body = item.body_value(field);
|
let inline_body = item.body_value(field);
|
||||||
match admin_usage_resolve_body_value(state, item, inline_body, field).await {
|
match admin_usage_resolve_body_value(state, item, inline_body, field).await {
|
||||||
Ok(body) => AdminUsageDetailBodyValue {
|
Ok(body) => AdminUsageDetailBodyValue::resolved(item, field, body),
|
||||||
value: body,
|
|
||||||
load_failed: false,
|
|
||||||
},
|
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
tracing::warn!(
|
tracing::warn!(
|
||||||
error = ?err,
|
error = ?err,
|
||||||
@@ -80,13 +129,139 @@ async fn resolve_admin_usage_detail_body_value(
|
|||||||
);
|
);
|
||||||
let value = inline_body.cloned();
|
let value = inline_body.cloned();
|
||||||
AdminUsageDetailBodyValue {
|
AdminUsageDetailBodyValue {
|
||||||
load_failed: value.is_none(),
|
error_code: value
|
||||||
|
.is_none()
|
||||||
|
.then(|| admin_usage_body_load_error_code(&err)),
|
||||||
value,
|
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(
|
pub(super) async fn maybe_build_local_admin_usage_detail_response(
|
||||||
state: &AdminAppState<'_>,
|
state: &AdminAppState<'_>,
|
||||||
request_context: &AdminRequestContext<'_>,
|
request_context: &AdminRequestContext<'_>,
|
||||||
@@ -212,6 +387,25 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
|
|||||||
"include_bodies",
|
"include_bodies",
|
||||||
true,
|
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 {
|
let Some(item) = state.find_request_usage_by_id(&usage_id).await? else {
|
||||||
return Ok(Some(
|
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 user_ids = item.user_id.clone().into_iter().collect::<Vec<_>>();
|
||||||
let (users_by_id, provider_key_names, api_key_names): (
|
let (users_by_id, provider_key_names, api_key_names): (
|
||||||
BTreeMap<String, aether_data::repository::users::StoredUserSummary>,
|
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 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!(
|
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_field(
|
||||||
resolve_admin_usage_detail_body_value(
|
state,
|
||||||
|
&item,
|
||||||
|
UsageBodyField::RequestBody,
|
||||||
|
body_field
|
||||||
|
),
|
||||||
|
resolve_admin_usage_detail_field(
|
||||||
state,
|
state,
|
||||||
&item,
|
&item,
|
||||||
UsageBodyField::ProviderRequestBody,
|
UsageBodyField::ProviderRequestBody,
|
||||||
|
body_field,
|
||||||
),
|
),
|
||||||
resolve_admin_usage_detail_body_value(
|
resolve_admin_usage_detail_field(
|
||||||
state,
|
state,
|
||||||
&item,
|
&item,
|
||||||
UsageBodyField::ResponseBody,
|
UsageBodyField::ResponseBody,
|
||||||
|
body_field,
|
||||||
),
|
),
|
||||||
resolve_admin_usage_detail_body_value(
|
resolve_admin_usage_detail_field(
|
||||||
state,
|
state,
|
||||||
&item,
|
&item,
|
||||||
UsageBodyField::ClientResponseBody,
|
UsageBodyField::ClientResponseBody,
|
||||||
|
body_field,
|
||||||
),
|
),
|
||||||
);
|
);
|
||||||
for (field, resolved) in [
|
for (field, resolved) in [
|
||||||
@@ -275,20 +499,25 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
|
|||||||
(UsageBodyField::ResponseBody, &response_body),
|
(UsageBodyField::ResponseBody, &response_body),
|
||||||
(UsageBodyField::ClientResponseBody, &client_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_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;
|
if body_field.is_none_or(|field| field == UsageBodyField::ProviderRequestBody) {
|
||||||
detail_item.response_body = response_body.value;
|
detail_item.provider_request_body = provider_request_body.value;
|
||||||
detail_item.client_response_body = client_response_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
|
request_body.value
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
if include_bodies {
|
|
||||||
// request_body 已通过 request capture 解析;其余 detached body 在上方并行加载。
|
|
||||||
}
|
|
||||||
let default_headers = admin_usage_curl_headers();
|
let default_headers = admin_usage_curl_headers();
|
||||||
let mut payload = build_admin_usage_detail_payload(
|
let mut payload = build_admin_usage_detail_payload(
|
||||||
&detail_item,
|
&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_user_data_reader(),
|
||||||
state.has_auth_api_key_data_reader(),
|
state.has_auth_api_key_data_reader(),
|
||||||
provider_key_name.as_deref(),
|
provider_key_name.as_deref(),
|
||||||
include_bodies,
|
include_bodies && body_field.is_none(),
|
||||||
request_body,
|
if body_field.is_none() {
|
||||||
|
request_body.take()
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
},
|
||||||
&default_headers,
|
&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() {
|
payload["body_load_errors"] = if include_bodies && !body_load_errors.is_empty() {
|
||||||
Value::Object(body_load_errors)
|
Value::Object(body_load_errors)
|
||||||
} else {
|
} else {
|
||||||
Value::Null
|
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(
|
return Ok(Some(attach_admin_audit_response(
|
||||||
Json(payload).into_response(),
|
Json(payload).into_response(),
|
||||||
@@ -320,3 +567,61 @@ pub(super) async fn maybe_build_local_admin_usage_detail_response(
|
|||||||
|
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::admin_usage_body_load_error_code;
|
||||||
|
use crate::GatewayError;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn admin_usage_raw_body_does_not_decode_or_reencode_stored_bytes() {
|
||||||
|
use super::{admin_usage_raw_payload_response, StoredUsageBodyPayload};
|
||||||
|
for (payload, encoding, expected) in [
|
||||||
|
(
|
||||||
|
StoredUsageBodyPayload::Gzip(vec![31, 139, 8, 0, 1]),
|
||||||
|
"gzip",
|
||||||
|
vec![31, 139, 8, 0, 1],
|
||||||
|
),
|
||||||
|
(
|
||||||
|
StoredUsageBodyPayload::Json(b"{ \"untouched\" : true }".to_vec()),
|
||||||
|
"json",
|
||||||
|
b"{ \"untouched\" : true }".to_vec(),
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
let response = admin_usage_raw_payload_response(payload);
|
||||||
|
assert_eq!(response.headers()["content-encoding"], "identity");
|
||||||
|
assert_eq!(response.headers()["x-aether-body-encoding"], encoding);
|
||||||
|
let bytes = axum::body::to_bytes(response.into_body(), 1024)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(bytes.as_ref(), expected.as_slice());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn body_load_errors_expose_safe_codes_instead_of_internal_messages() {
|
||||||
|
for (message, expected) in [
|
||||||
|
(
|
||||||
|
"unexpected database value: decompressed usage json exceeds 67108864 bytes",
|
||||||
|
"too_large",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"failed to decompress usage json: invalid gzip header",
|
||||||
|
"decode_failed",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"failed to parse decompressed usage json: invalid JSON",
|
||||||
|
"decode_failed",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"postgres error: private connection details",
|
||||||
|
"storage_unavailable",
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
assert_eq!(
|
||||||
|
admin_usage_body_load_error_code(&GatewayError::Internal(message.to_string())),
|
||||||
|
expected
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ mod extractors;
|
|||||||
mod list;
|
mod list;
|
||||||
pub(crate) mod payloads;
|
pub(crate) mod payloads;
|
||||||
mod reads;
|
mod reads;
|
||||||
|
mod reveal;
|
||||||
mod support;
|
mod support;
|
||||||
mod update;
|
mod update;
|
||||||
|
|
||||||
@@ -41,6 +42,10 @@ pub(crate) async fn maybe_build_local_admin_endpoints_routes_response(
|
|||||||
return Ok(Some(response));
|
return Ok(Some(response));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if let Some(response) = reveal::maybe_handle(state, request_context).await? {
|
||||||
|
return Ok(Some(response));
|
||||||
|
}
|
||||||
|
|
||||||
if let Some(response) = defaults::maybe_handle(state, request_context, request_body).await? {
|
if let Some(response) = defaults::maybe_handle(state, request_context, request_body).await? {
|
||||||
return Ok(Some(response));
|
return Ok(Some(response));
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,72 @@
|
|||||||
|
use super::extractors::admin_endpoint_id;
|
||||||
|
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||||
|
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||||
|
use crate::handlers::admin::shared::{
|
||||||
|
attach_admin_audit_response, mark_sensitive_admin_response_no_store,
|
||||||
|
};
|
||||||
|
use crate::GatewayError;
|
||||||
|
use axum::{
|
||||||
|
body::Body,
|
||||||
|
http::StatusCode,
|
||||||
|
response::{IntoResponse, Response},
|
||||||
|
Json,
|
||||||
|
};
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
pub(super) async fn maybe_handle(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
request_context: &AdminRequestContext<'_>,
|
||||||
|
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||||
|
let Some(decision) = request_context.decision() else {
|
||||||
|
return Ok(None);
|
||||||
|
};
|
||||||
|
if decision.route_family.as_deref() != Some("endpoints_manage")
|
||||||
|
|| decision.route_kind.as_deref() != Some("reveal_endpoint_rules")
|
||||||
|
{
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
|
if !state.has_provider_catalog_data_reader() {
|
||||||
|
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
|
||||||
|
}
|
||||||
|
let Some(endpoint_id) = request_context
|
||||||
|
.path()
|
||||||
|
.strip_suffix("/rules/reveal")
|
||||||
|
.and_then(admin_endpoint_id)
|
||||||
|
else {
|
||||||
|
return Ok(Some(
|
||||||
|
(
|
||||||
|
StatusCode::NOT_FOUND,
|
||||||
|
Json(json!({ "detail": "Endpoint 不存在" })),
|
||||||
|
)
|
||||||
|
.into_response(),
|
||||||
|
));
|
||||||
|
};
|
||||||
|
let Some(endpoint) = state
|
||||||
|
.read_provider_catalog_endpoints_by_ids(std::slice::from_ref(&endpoint_id))
|
||||||
|
.await?
|
||||||
|
.into_iter()
|
||||||
|
.next()
|
||||||
|
else {
|
||||||
|
return Ok(Some(
|
||||||
|
(
|
||||||
|
StatusCode::NOT_FOUND,
|
||||||
|
Json(json!({ "detail": "Endpoint 不存在" })),
|
||||||
|
)
|
||||||
|
.into_response(),
|
||||||
|
));
|
||||||
|
};
|
||||||
|
let payload = json!({
|
||||||
|
"header_rules": endpoint.header_rules.as_ref().and_then(|value| value.as_array()).cloned().unwrap_or_default(),
|
||||||
|
"body_rules": endpoint.body_rules.as_ref().and_then(|value| value.as_array()).cloned().unwrap_or_default(),
|
||||||
|
"response_header_rules": endpoint.config.as_ref().and_then(|config| config.get("response_header_rules")).and_then(|value| value.as_array()).cloned().unwrap_or_default(),
|
||||||
|
});
|
||||||
|
Ok(Some(mark_sensitive_admin_response_no_store(
|
||||||
|
attach_admin_audit_response(
|
||||||
|
Json(payload).into_response(),
|
||||||
|
"admin_endpoint_rules_revealed",
|
||||||
|
"reveal_endpoint_rules",
|
||||||
|
"provider_endpoint",
|
||||||
|
&endpoint_id,
|
||||||
|
),
|
||||||
|
)))
|
||||||
|
}
|
||||||
@@ -1463,7 +1463,7 @@ mod tests {
|
|||||||
&auth_config,
|
&auth_config,
|
||||||
Some(0),
|
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::{
|
use super::super::kiro::{
|
||||||
admin_provider_oauth_kiro_refresh_base_url_override, fetch_admin_provider_oauth_kiro_email,
|
admin_provider_oauth_kiro_refresh_base_url_override, fetch_admin_provider_oauth_kiro_email,
|
||||||
refresh_admin_provider_oauth_kiro_auth_config,
|
refresh_admin_provider_oauth_kiro_auth_config,
|
||||||
@@ -79,7 +80,7 @@ fn kiro_social_key_name(
|
|||||||
.collect::<String>()
|
.collect::<String>()
|
||||||
})
|
})
|
||||||
.unwrap_or_else(|| "unknown".to_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> {
|
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 {
|
} else {
|
||||||
let key_name = email
|
let key_name = admin_provider_oauth_key_name_from_auth_config(
|
||||||
.as_deref()
|
&provider.provider_type,
|
||||||
.map(|email| format!("windsurf_{email}"))
|
&auth_config,
|
||||||
.unwrap_or_else(|| format!("windsurf_{}", current_unix_secs()));
|
None,
|
||||||
|
);
|
||||||
match state
|
match state
|
||||||
.create_provider_oauth_catalog_key(
|
.create_provider_oauth_catalog_key(
|
||||||
&provider.id,
|
&provider.id,
|
||||||
@@ -1356,6 +1358,32 @@ mod tests {
|
|||||||
use crate::control::GatewayAdminPrincipalContext;
|
use crate::control::GatewayAdminPrincipalContext;
|
||||||
use aether_data::repository::provider_oauth::StoredAdminProviderOAuthDeviceSession;
|
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 {
|
fn device_session() -> StoredAdminProviderOAuthDeviceSession {
|
||||||
StoredAdminProviderOAuthDeviceSession {
|
StoredAdminProviderOAuthDeviceSession {
|
||||||
session_id: "device-session-1".to_string(),
|
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>,
|
auth_config: &Map<String, Value>,
|
||||||
batch_index: Option<usize>,
|
batch_index: Option<usize>,
|
||||||
) -> String {
|
) -> String {
|
||||||
let provider_type = provider_type.trim();
|
|
||||||
if let Some(email) = trimmed_auth_config_string(auth_config, "email") {
|
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") {
|
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())
|
.map(|duration| duration.as_secs())
|
||||||
.unwrap_or(0);
|
.unwrap_or(0);
|
||||||
match batch_index {
|
match batch_index {
|
||||||
Some(index) => format!("{provider_type}_{timestamp}_{index}"),
|
Some(index) => format!("账号_{timestamp}_{index}"),
|
||||||
None => format!("账号_{timestamp}"),
|
None => format!("账号_{timestamp}"),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -87,6 +86,106 @@ mod tests {
|
|||||||
use super::*;
|
use super::*;
|
||||||
use serde_json::{json, Map};
|
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]
|
#[test]
|
||||||
fn grok_default_key_name_uses_full_user_id() {
|
fn grok_default_key_name_uses_full_user_id() {
|
||||||
let mut auth_config = Map::new();
|
let mut auth_config = Map::new();
|
||||||
@@ -95,10 +194,18 @@ mod tests {
|
|||||||
json!("1619039a-0191-4e0a-a490-8f4ad21262c9"),
|
json!("1619039a-0191-4e0a-a490-8f4ad21262c9"),
|
||||||
);
|
);
|
||||||
|
|
||||||
assert_eq!(
|
for provider_type in ["grok", " Grok "] {
|
||||||
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
|
for batch_index in [None, Some(3)] {
|
||||||
"grok_1619039a-0191-4e0a-a490-8f4ad21262c9"
|
assert_eq!(
|
||||||
);
|
admin_provider_oauth_key_name_from_auth_config(
|
||||||
|
provider_type,
|
||||||
|
&auth_config,
|
||||||
|
batch_index,
|
||||||
|
),
|
||||||
|
"1619039a-0191-4e0a-a490-8f4ad21262c9"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -109,17 +216,22 @@ mod tests {
|
|||||||
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
|
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
|
||||||
"grok_grok@example.com"
|
"[email protected]"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[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 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!(name.ends_with("_3"));
|
||||||
|
assert!(other_name.starts_with("账号_"));
|
||||||
|
assert!(other_name.ends_with("_4"));
|
||||||
|
assert_ne!(name, other_name);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ use super::shared::{
|
|||||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||||
};
|
};
|
||||||
use crate::handlers::admin::provider::shared::payloads::AdminImportProviderModelsRequest;
|
|
||||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||||
use crate::GatewayError;
|
use crate::GatewayError;
|
||||||
use aether_admin::provider::quota::{
|
use aether_admin::provider::quota::{
|
||||||
@@ -24,63 +23,6 @@ use std::collections::BTreeMap;
|
|||||||
use std::time::{SystemTime, UNIX_EPOCH};
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
use tracing::warn;
|
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(
|
async fn execute_antigravity_quota_plan(
|
||||||
state: &AdminAppState<'_>,
|
state: &AdminAppState<'_>,
|
||||||
transport: &AdminGatewayProviderTransportSnapshot,
|
transport: &AdminGatewayProviderTransportSnapshot,
|
||||||
@@ -380,10 +322,6 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
if status == "success" {
|
|
||||||
sync_antigravity_discovered_models(state, &provider.id, metadata_update.as_ref()).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
if status == "success" {
|
if status == "success" {
|
||||||
success_count += 1;
|
success_count += 1;
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -234,6 +234,20 @@ impl<'a> AdminAppState<'a> {
|
|||||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
.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(
|
pub(crate) async fn build_api_format_health_monitor_payload(
|
||||||
&self,
|
&self,
|
||||||
lookback_hours: u64,
|
lookback_hours: u64,
|
||||||
|
|||||||
@@ -1399,7 +1399,8 @@ fn rename_release_dir_noreplace(source: &Path, destination: &Path) -> Result<(),
|
|||||||
#[cfg(target_os = "linux")]
|
#[cfg(target_os = "linux")]
|
||||||
// SAFETY: both paths are valid NUL-terminated strings and remain alive for the call.
|
// SAFETY: both paths are valid NUL-terminated strings and remain alive for the call.
|
||||||
let status = unsafe {
|
let status = unsafe {
|
||||||
libc::renameat2(
|
libc::syscall(
|
||||||
|
libc::SYS_renameat2,
|
||||||
libc::AT_FDCWD,
|
libc::AT_FDCWD,
|
||||||
source.as_ptr(),
|
source.as_ptr(),
|
||||||
libc::AT_FDCWD,
|
libc::AT_FDCWD,
|
||||||
@@ -2625,6 +2626,27 @@ mod tests {
|
|||||||
std::fs::remove_dir_all(dir).ok();
|
std::fs::remove_dir_all(dir).ok();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(any(target_os = "linux", target_os = "macos"))]
|
||||||
|
#[test]
|
||||||
|
fn update_preserves_existing_empty_release_destination() {
|
||||||
|
let dir = temp_test_dir("existing-empty-release");
|
||||||
|
let source = dir.join("prepared-v1.2.3");
|
||||||
|
let release = dir.join("releases/v1.2.3");
|
||||||
|
std::fs::create_dir_all(&source).expect("prepared release should be created");
|
||||||
|
std::fs::create_dir_all(&release).expect("existing release should be created");
|
||||||
|
|
||||||
|
let err = rename_release_dir_noreplace(&source, &release)
|
||||||
|
.expect_err("existing empty release must not be replaceable");
|
||||||
|
|
||||||
|
assert!(err.contains("拒绝覆盖"));
|
||||||
|
assert!(
|
||||||
|
release.is_dir(),
|
||||||
|
"failed install must retain its destination"
|
||||||
|
);
|
||||||
|
assert!(source.is_dir(), "failed install must retain its source");
|
||||||
|
std::fs::remove_dir_all(dir).ok();
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(any(target_os = "linux", target_os = "macos"))]
|
#[cfg(any(target_os = "linux", target_os = "macos"))]
|
||||||
#[test]
|
#[test]
|
||||||
fn update_atomically_installs_absent_release_destination() {
|
fn update_atomically_installs_absent_release_destination() {
|
||||||
|
|||||||
@@ -211,9 +211,19 @@ fn apply_sensitive_route_cache_policy(
|
|||||||
return;
|
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(
|
headers.insert(
|
||||||
http::header::CACHE_CONTROL,
|
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"));
|
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]
|
#[test]
|
||||||
fn authenticated_user_data_responses_are_never_cacheable() {
|
fn authenticated_user_data_responses_are_never_cacheable() {
|
||||||
let mut headers = HeaderMap::new();
|
let mut headers = HeaderMap::new();
|
||||||
|
|||||||
@@ -157,16 +157,22 @@ pub(crate) fn websocket_upstream_url(
|
|||||||
return Err(invalid_code);
|
return Err(invalid_code);
|
||||||
}
|
}
|
||||||
let websocket_scheme = match url.scheme() {
|
let websocket_scheme = match url.scheme() {
|
||||||
"https" => "wss",
|
"https" | "wss" => "wss",
|
||||||
"http" => "ws",
|
"http" | "ws" => "ws",
|
||||||
"wss" => return Ok(url),
|
|
||||||
"ws" if aether_http::url_has_literal_loopback_host(&url) => return Ok(url),
|
|
||||||
"ws" => return Err(invalid_code),
|
|
||||||
_ => return Err(invalid_code),
|
_ => return Err(invalid_code),
|
||||||
};
|
};
|
||||||
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
|
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
|
||||||
if url.scheme() == "ws" && !aether_http::url_has_literal_loopback_host(&url) {
|
if url.scheme() == "ws" {
|
||||||
return Err(invalid_code);
|
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)
|
Ok(url)
|
||||||
}
|
}
|
||||||
@@ -844,15 +850,17 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn maps_http_url_to_websocket_url_without_losing_path_or_query() {
|
fn maps_http_url_to_websocket_url_without_losing_path_or_query() {
|
||||||
let url = websocket_upstream_url(
|
for (http_scheme, websocket_scheme) in [("https", "wss"), ("http", "ws")] {
|
||||||
"https://example.test/backend-api/codex/responses?x=1",
|
let url = websocket_upstream_url(
|
||||||
"invalid",
|
&format!("{http_scheme}://example.test:8080/backend-api/codex/responses?x=1"),
|
||||||
)
|
"invalid",
|
||||||
.expect("URL should be converted");
|
)
|
||||||
assert_eq!(
|
.expect("URL should be converted");
|
||||||
url.as_str(),
|
assert_eq!(
|
||||||
"wss://example.test/backend-api/codex/responses?x=1"
|
url.as_str(),
|
||||||
);
|
format!("{websocket_scheme}://example.test:8080/backend-api/codex/responses?x=1")
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -861,10 +869,14 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn remote_websocket_requires_wss_but_loopback_ws_is_allowed() {
|
fn websocket_upstream_url_accepts_ws_and_wss_with_safe_targets() {
|
||||||
for allowed in [
|
for allowed in [
|
||||||
"wss://example.test/v1/responses",
|
"wss://example.test/v1/responses",
|
||||||
"https://example.test/v1/responses",
|
"https://example.test/v1/responses",
|
||||||
|
"ws://example.test:8080/v1/responses",
|
||||||
|
"http://example.test:8080/v1/responses",
|
||||||
|
"http://8.8.8.8:8080/v1/responses",
|
||||||
|
"ws://[2606:4700:4700::1111]:8080/v1/responses",
|
||||||
"ws://localhost:8080/v1/responses",
|
"ws://localhost:8080/v1/responses",
|
||||||
"http://127.42.0.1:8080/v1/responses",
|
"http://127.42.0.1:8080/v1/responses",
|
||||||
"ws://[::1]:8080/v1/responses",
|
"ws://[::1]:8080/v1/responses",
|
||||||
@@ -875,11 +887,14 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
for rejected in [
|
for rejected in [
|
||||||
"ws://example.test/v1/responses",
|
|
||||||
"http://10.0.0.1/v1/responses",
|
"http://10.0.0.1/v1/responses",
|
||||||
"ws://0.0.0.0:8080/v1/responses",
|
"ws://0.0.0.0:8080/v1/responses",
|
||||||
"ws://[::ffff:127.0.0.1]:8080/v1/responses",
|
"ws://[::ffff:127.0.0.1]:8080/v1/responses",
|
||||||
"wss://example.test/v1/responses#secret",
|
"wss://example.test/v1/responses#secret",
|
||||||
|
"ws://example.test/v1/responses#secret",
|
||||||
|
"http://[email protected]/v1/responses",
|
||||||
|
"ws://[email protected]/v1/responses",
|
||||||
|
"ftp://example.test/v1/responses",
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
websocket_upstream_url(rejected, "invalid").is_err(),
|
websocket_upstream_url(rejected, "invalid").is_err(),
|
||||||
|
|||||||
@@ -129,9 +129,6 @@ pub(crate) fn normalize_admin_base_url(base_url: &str) -> Result<String, String>
|
|||||||
if parsed.host_str().is_none() {
|
if parsed.host_str().is_none() {
|
||||||
return Err("base_url 必须包含有效主机".to_string());
|
return Err("base_url 必须包含有效主机".to_string());
|
||||||
}
|
}
|
||||||
if !aether_http::is_https_or_loopback_http_url(&parsed) {
|
|
||||||
return Err("base_url 必须使用 HTTPS;HTTP 仅允许字面量 loopback 主机".to_string());
|
|
||||||
}
|
|
||||||
if !parsed.username().is_empty() || parsed.password().is_some() {
|
if !parsed.username().is_empty() || parsed.password().is_some() {
|
||||||
return Err("base_url 不允许包含用户名或密码".to_string());
|
return Err("base_url 不允许包含用户名或密码".to_string());
|
||||||
}
|
}
|
||||||
@@ -154,16 +151,42 @@ mod normalize_admin_base_url_tests {
|
|||||||
"https://user:[email protected]/v1",
|
"https://user:[email protected]/v1",
|
||||||
"https://api.example.test/v1?key=secret",
|
"https://api.example.test/v1?key=secret",
|
||||||
"https://api.example.test/v1#secret",
|
"https://api.example.test/v1#secret",
|
||||||
"http://api.example.test/v1",
|
"http://user:password@api.example.test/v1",
|
||||||
"http://10.0.0.1/v1",
|
"http://api.example.test/v1?key=secret",
|
||||||
"http://[::ffff:127.0.0.1]/v1",
|
"http://api.example.test/v1#secret",
|
||||||
|
"ftp://api.example.test/v1",
|
||||||
|
"file:///v1",
|
||||||
|
"api.example.test/v1",
|
||||||
|
"",
|
||||||
"https://",
|
"https://",
|
||||||
|
"http://",
|
||||||
"https://api.example.test:invalid/v1",
|
"https://api.example.test:invalid/v1",
|
||||||
] {
|
] {
|
||||||
assert!(normalize_admin_base_url(value).is_err(), "accepted {value}");
|
assert!(normalize_admin_base_url(value).is_err(), "accepted {value}");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn endpoint_base_url_accepts_remote_http_hosts() {
|
||||||
|
for (raw_url, expected) in [
|
||||||
|
(
|
||||||
|
" HTTP://API.EXAMPLE.TEST:8080/v1/ ",
|
||||||
|
"http://api.example.test:8080/v1",
|
||||||
|
),
|
||||||
|
("http://8.8.8.8:8080/v1/", "http://8.8.8.8:8080/v1"),
|
||||||
|
("http://10.0.0.1:8080/v1/", "http://10.0.0.1:8080/v1"),
|
||||||
|
(
|
||||||
|
"http://[2606:4700:4700::1111]:8080/v1/",
|
||||||
|
"http://[2606:4700:4700::1111]:8080/v1",
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
assert_eq!(
|
||||||
|
normalize_admin_base_url(raw_url).expect("HTTP base URL should be accepted"),
|
||||||
|
expected,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn endpoint_base_url_is_parsed_and_normalized() {
|
fn endpoint_base_url_is_parsed_and_normalized() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
|
|||||||
@@ -30,6 +30,8 @@ use std::time::{SystemTime, UNIX_EPOCH};
|
|||||||
mod support_announcements;
|
mod support_announcements;
|
||||||
#[path = "support/auth.rs"]
|
#[path = "support/auth.rs"]
|
||||||
mod support_auth;
|
mod support_auth;
|
||||||
|
#[path = "support/auth_cookie_policy.rs"]
|
||||||
|
mod support_auth_cookie_policy;
|
||||||
#[path = "support/billing.rs"]
|
#[path = "support/billing.rs"]
|
||||||
mod support_billing;
|
mod support_billing;
|
||||||
#[path = "support/ccswitch.rs"]
|
#[path = "support/ccswitch.rs"]
|
||||||
@@ -133,6 +135,31 @@ pub(crate) async fn maybe_build_local_public_support_response(
|
|||||||
remote_addr: &std::net::SocketAddr,
|
remote_addr: &std::net::SocketAddr,
|
||||||
client_ip: std::net::IpAddr,
|
client_ip: std::net::IpAddr,
|
||||||
request_body: Option<&Bytes>,
|
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>> {
|
) -> Option<Response<Body>> {
|
||||||
let decision = request_context.control_decision.as_ref()?;
|
let decision = request_context.control_decision.as_ref()?;
|
||||||
if decision.route_class.as_deref() != Some("public_support") {
|
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)
|
.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")
|
std::env::var("AUTH_REFRESH_COOKIE_NAME")
|
||||||
.ok()
|
.ok()
|
||||||
.map(|value| value.trim().to_string())
|
.map(|value| value.trim().to_string())
|
||||||
|
|||||||
@@ -53,9 +53,6 @@ async fn resolve_test_connection_target(
|
|||||||
{
|
{
|
||||||
return Err("provider endpoint must be an HTTP(S) URL without credentials or fragment");
|
return Err("provider endpoint must be an HTTP(S) URL without credentials or fragment");
|
||||||
}
|
}
|
||||||
if url.scheme() == "http" && !(allow_private_targets && literal_loopback) {
|
|
||||||
return Err("provider endpoint must use HTTPS");
|
|
||||||
}
|
|
||||||
let host = url
|
let host = url
|
||||||
.host_str()
|
.host_str()
|
||||||
.ok_or("provider endpoint is missing a host")?
|
.ok_or("provider endpoint is missing a host")?
|
||||||
@@ -574,10 +571,11 @@ mod tests {
|
|||||||
async fn test_connection_target_rejects_private_addresses_in_production_mode() {
|
async fn test_connection_target_rejects_private_addresses_in_production_mode() {
|
||||||
for raw_url in [
|
for raw_url in [
|
||||||
"http://127.0.0.1:8080/v1/chat/completions",
|
"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://10.0.0.1/v1/chat/completions",
|
||||||
"https://[::1]/v1/chat/completions",
|
"https://[::1]/v1/chat/completions",
|
||||||
"https://localhost/v1/chat/completions",
|
"https://localhost/v1/chat/completions",
|
||||||
"http://8.8.8.8/v1/chat/completions",
|
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
resolve_test_connection_target(raw_url, false)
|
resolve_test_connection_target(raw_url, false)
|
||||||
@@ -588,6 +586,26 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[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);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_connection_target_allows_loopback_only_for_test_fixtures() {
|
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)
|
let target = resolve_test_connection_target("http://127.0.0.1:8080/v1/chat", true)
|
||||||
@@ -596,10 +614,10 @@ mod tests {
|
|||||||
assert_eq!(target.host, "127.0.0.1");
|
assert_eq!(target.host, "127.0.0.1");
|
||||||
assert_eq!(target.addresses.len(), 1);
|
assert_eq!(target.addresses.len(), 1);
|
||||||
assert!(
|
assert!(
|
||||||
resolve_test_connection_target("http://8.8.8.8/v1/chat", true)
|
resolve_test_connection_target("http://10.0.0.1/v1/chat", true)
|
||||||
.await
|
.await
|
||||||
.is_err(),
|
.is_err(),
|
||||||
"test mode must not make cleartext public endpoints acceptable"
|
"test mode must not make private non-loopback HTTP endpoints acceptable"
|
||||||
);
|
);
|
||||||
assert!(
|
assert!(
|
||||||
resolve_test_connection_target("https://10.0.0.1/v1/chat", true)
|
resolve_test_connection_target("https://10.0.0.1/v1/chat", true)
|
||||||
@@ -620,6 +638,9 @@ mod tests {
|
|||||||
for raw_url in [
|
for raw_url in [
|
||||||
"https://user:[email protected]/v1/chat",
|
"https://user:[email protected]/v1/chat",
|
||||||
"https://example.com/v1/chat#fragment",
|
"https://example.com/v1/chat#fragment",
|
||||||
|
"http://user:[email protected]/v1/chat",
|
||||||
|
"http://example.com/v1/chat#fragment",
|
||||||
|
"ftp://example.com/v1/chat",
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
resolve_test_connection_target(raw_url, false)
|
resolve_test_connection_target(raw_url, false)
|
||||||
|
|||||||
@@ -1940,12 +1940,16 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn user_usage_active_override_uses_terminal_candidate_latency() {
|
fn user_usage_active_override_uses_terminal_candidate_latency() {
|
||||||
let candidate = sample_candidate(
|
let mut candidate = sample_candidate(
|
||||||
RequestCandidateStatus::Success,
|
RequestCandidateStatus::Success,
|
||||||
Some(200),
|
Some(200),
|
||||||
Some(9_210),
|
Some(9_210),
|
||||||
None,
|
None,
|
||||||
);
|
);
|
||||||
|
candidate.error_message = Some("private upstream diagnostic".to_string());
|
||||||
|
candidate.extra_data = Some(json!({
|
||||||
|
"upstream_response": {"body": {"error": {"message": "private upstream diagnostic"}}}
|
||||||
|
}));
|
||||||
|
|
||||||
let payload =
|
let payload =
|
||||||
users_me_usage_terminal_candidate_state_override(&[candidate]).expect("override");
|
users_me_usage_terminal_candidate_state_override(&[candidate]).expect("override");
|
||||||
@@ -1953,6 +1957,7 @@ mod tests {
|
|||||||
assert_eq!(payload["status"], "completed");
|
assert_eq!(payload["status"], "completed");
|
||||||
assert_eq!(payload["response_time_ms"], 9_210);
|
assert_eq!(payload["response_time_ms"], 9_210);
|
||||||
assert_eq!(payload["status_code"], 200);
|
assert_eq!(payload["status_code"], 200);
|
||||||
|
assert!(!payload.to_string().contains("private upstream diagnostic"));
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
payload["response_time_updated_at"],
|
payload["response_time_updated_at"],
|
||||||
"1970-01-01T00:00:10.210+00:00"
|
"1970-01-01T00:00:10.210+00:00"
|
||||||
|
|||||||
@@ -120,8 +120,16 @@ pub(crate) async fn decrypt_or_migrate_smtp_password(
|
|||||||
}
|
}
|
||||||
let plaintext = decrypt_system_config_secret(state, "smtp_password", stored.trim())
|
let plaintext = decrypt_system_config_secret(state, "smtp_password", stored.trim())
|
||||||
.or_else(|| {
|
.or_else(|| {
|
||||||
(!stored.trim().is_empty() && !looks_like_python_fernet_ciphertext(stored.trim()))
|
if stored_secret_uses_known_envelope_family(stored.trim()) {
|
||||||
.then(|| stored.trim().to_string())
|
return None;
|
||||||
|
}
|
||||||
|
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored.trim()).or_else(
|
||||||
|
|| {
|
||||||
|
(!stored.trim().is_empty()
|
||||||
|
&& !looks_like_python_fernet_ciphertext(stored.trim()))
|
||||||
|
.then(|| stored.trim().to_string())
|
||||||
|
},
|
||||||
|
)
|
||||||
})
|
})
|
||||||
.ok_or_else(|| system_config_secret_error("stored SMTP password cannot be decrypted"))?;
|
.ok_or_else(|| system_config_secret_error("stored SMTP password cannot be decrypted"))?;
|
||||||
if plaintext.contains('\0') {
|
if plaintext.contains('\0') {
|
||||||
@@ -700,11 +708,13 @@ mod tests {
|
|||||||
use super::{
|
use super::{
|
||||||
bark_device_key_binding, decrypt_bark_device_key_v2, decrypt_ldap_bind_password_v2,
|
bark_device_key_binding, decrypt_bark_device_key_v2, decrypt_ldap_bind_password_v2,
|
||||||
decrypt_ldap_bind_password_v3, decrypt_or_migrate_bark_device_key,
|
decrypt_ldap_bind_password_v3, decrypt_or_migrate_bark_device_key,
|
||||||
decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_system_config_secret,
|
decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_smtp_password,
|
||||||
|
decrypt_or_migrate_system_config_secret,
|
||||||
decrypt_or_migrate_system_config_secret_with_before_compare, decrypt_system_config_secret,
|
decrypt_or_migrate_system_config_secret_with_before_compare, decrypt_system_config_secret,
|
||||||
encrypt_bark_device_key, encrypt_ldap_bind_password, encrypt_system_config_secret,
|
encrypt_bark_device_key, encrypt_ldap_bind_password, encrypt_smtp_password,
|
||||||
ldap_module_config_is_valid, normalize_ldap_transport_server_url,
|
encrypt_system_config_secret, ldap_module_config_is_valid,
|
||||||
LDAP_BIND_PASSWORD_V2_PREFIX, LDAP_BIND_PASSWORD_V3_PREFIX, SYSTEM_CONFIG_SECRET_V2_PREFIX,
|
normalize_ldap_transport_server_url, smtp_password_binding, LDAP_BIND_PASSWORD_V2_PREFIX,
|
||||||
|
LDAP_BIND_PASSWORD_V3_PREFIX, SMTP_PASSWORD_V3_PREFIX, SYSTEM_CONFIG_SECRET_V2_PREFIX,
|
||||||
};
|
};
|
||||||
use crate::data::GatewayDataState;
|
use crate::data::GatewayDataState;
|
||||||
use crate::AppState;
|
use crate::AppState;
|
||||||
@@ -740,6 +750,165 @@ mod tests {
|
|||||||
state
|
state
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn smtp_password_migrates_legacy_formats_to_bound_v3() {
|
||||||
|
let binding = smtp_password_binding(
|
||||||
|
"smtp.example.com",
|
||||||
|
587,
|
||||||
|
Some("[email protected]"),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
)
|
||||||
|
.expect("SMTP binding should build");
|
||||||
|
let fixture_state = state_with_stored_secret(TEST_SECRET);
|
||||||
|
let legacy_values = [
|
||||||
|
TEST_SECRET.to_string(),
|
||||||
|
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, TEST_SECRET)
|
||||||
|
.expect("legacy SMTP password should encrypt"),
|
||||||
|
encrypt_system_config_secret(&fixture_state, TEST_KEY, TEST_SECRET)
|
||||||
|
.expect("v2 SMTP password should encrypt"),
|
||||||
|
];
|
||||||
|
|
||||||
|
for legacy in legacy_values {
|
||||||
|
let state = state_with_stored_secret(&legacy);
|
||||||
|
let plaintext = decrypt_or_migrate_smtp_password(&state, &binding, legacy.clone())
|
||||||
|
.await
|
||||||
|
.expect("legacy SMTP password should migrate");
|
||||||
|
assert_eq!(plaintext, TEST_SECRET);
|
||||||
|
let migrated = state
|
||||||
|
.read_system_config_json_value_strong(TEST_KEY)
|
||||||
|
.await
|
||||||
|
.expect("SMTP password should read")
|
||||||
|
.and_then(|value| value.as_str().map(ToOwned::to_owned))
|
||||||
|
.expect("SMTP password should be a string");
|
||||||
|
assert!(migrated.starts_with(SMTP_PASSWORD_V3_PREFIX));
|
||||||
|
assert_ne!(migrated, legacy);
|
||||||
|
assert_eq!(
|
||||||
|
decrypt_or_migrate_smtp_password(&state, &binding, migrated.clone())
|
||||||
|
.await
|
||||||
|
.expect("migrated SMTP password should decrypt"),
|
||||||
|
TEST_SECRET
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
state
|
||||||
|
.read_system_config_json_value_strong(TEST_KEY)
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
Some(json!(migrated))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn smtp_password_rejects_invalid_ciphertext_without_rewriting() {
|
||||||
|
let binding = smtp_password_binding(
|
||||||
|
"smtp.example.com",
|
||||||
|
587,
|
||||||
|
Some("[email protected]"),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
)
|
||||||
|
.expect("SMTP binding should build");
|
||||||
|
let fixture_state = state_with_stored_secret(TEST_SECRET);
|
||||||
|
let mut tampered = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, TEST_SECRET)
|
||||||
|
.expect("legacy SMTP password should encrypt");
|
||||||
|
tampered.replace_range(tampered.len() - 2.., "AA");
|
||||||
|
let invalid_values = [
|
||||||
|
tampered,
|
||||||
|
encrypt_python_fernet_plaintext("unavailable-historical-key", TEST_SECRET)
|
||||||
|
.expect("wrong-key SMTP password should encrypt"),
|
||||||
|
encrypt_system_config_secret(&fixture_state, "other_secret", TEST_SECRET)
|
||||||
|
.expect("wrong-purpose secret should encrypt"),
|
||||||
|
"aether-system-config-secret-v2:invalid".to_string(),
|
||||||
|
"aether-smtp-password-v3:invalid".to_string(),
|
||||||
|
"aether-runtime-secret-v1:invalid".to_string(),
|
||||||
|
"aether-unknown-secret-v4:invalid".to_string(),
|
||||||
|
];
|
||||||
|
|
||||||
|
for stored in invalid_values {
|
||||||
|
let state = state_with_stored_secret(&stored);
|
||||||
|
let error = decrypt_or_migrate_smtp_password(&state, &binding, stored.clone())
|
||||||
|
.await
|
||||||
|
.expect_err("invalid ciphertext must not become an SMTP password");
|
||||||
|
assert_eq!(
|
||||||
|
error.into_message(),
|
||||||
|
"stored SMTP password cannot be decrypted"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
state
|
||||||
|
.read_system_config_json_value_strong(TEST_KEY)
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
Some(json!(stored))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn smtp_password_v3_rejects_changed_transport_binding() {
|
||||||
|
let binding = smtp_password_binding(
|
||||||
|
"smtp.example.com",
|
||||||
|
587,
|
||||||
|
Some("[email protected]"),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
)
|
||||||
|
.expect("SMTP binding should build");
|
||||||
|
let stored = encrypt_smtp_password(
|
||||||
|
&state_with_stored_secret(TEST_SECRET),
|
||||||
|
&binding,
|
||||||
|
TEST_SECRET,
|
||||||
|
)
|
||||||
|
.expect("SMTP password should encrypt");
|
||||||
|
let state = state_with_stored_secret(&stored);
|
||||||
|
for changed_binding in [
|
||||||
|
smtp_password_binding(
|
||||||
|
"other.example.com",
|
||||||
|
587,
|
||||||
|
Some("[email protected]"),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
),
|
||||||
|
smtp_password_binding(
|
||||||
|
"smtp.example.com",
|
||||||
|
465,
|
||||||
|
Some("[email protected]"),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
),
|
||||||
|
smtp_password_binding(
|
||||||
|
"smtp.example.com",
|
||||||
|
587,
|
||||||
|
Some("[email protected]"),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
),
|
||||||
|
smtp_password_binding(
|
||||||
|
"smtp.example.com",
|
||||||
|
587,
|
||||||
|
Some("[email protected]"),
|
||||||
|
false,
|
||||||
|
false,
|
||||||
|
),
|
||||||
|
smtp_password_binding("smtp.example.com", 587, Some("[email protected]"), true, true),
|
||||||
|
] {
|
||||||
|
assert!(decrypt_or_migrate_smtp_password(
|
||||||
|
&state,
|
||||||
|
&changed_binding.expect("changed binding should build"),
|
||||||
|
stored.clone(),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.is_err());
|
||||||
|
}
|
||||||
|
assert_eq!(
|
||||||
|
state
|
||||||
|
.read_system_config_json_value_strong(TEST_KEY)
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
Some(json!(stored))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
fn ldap_config(bind_password: &str) -> StoredLdapModuleConfig {
|
fn ldap_config(bind_password: &str) -> StoredLdapModuleConfig {
|
||||||
StoredLdapModuleConfig {
|
StoredLdapModuleConfig {
|
||||||
server_url: "ldaps://ldap.example.com".to_string(),
|
server_url: "ldaps://ldap.example.com".to_string(),
|
||||||
|
|||||||
@@ -238,9 +238,11 @@ async fn read_notification_channel_readiness(
|
|||||||
state: &AppState,
|
state: &AppState,
|
||||||
config: &ImportantNotificationConfig,
|
config: &ImportantNotificationConfig,
|
||||||
) -> Result<NotificationChannelReadiness, GatewayError> {
|
) -> Result<NotificationChannelReadiness, GatewayError> {
|
||||||
let smtp_config = read_smtp_delivery_config(state).await?;
|
let email = config.email_enabled
|
||||||
|
&& !config.email_recipients.is_empty()
|
||||||
|
&& matches!(read_smtp_delivery_config(state).await, Ok(Some(_)));
|
||||||
Ok(NotificationChannelReadiness {
|
Ok(NotificationChannelReadiness {
|
||||||
email: config.email_enabled && !config.email_recipients.is_empty() && smtp_config.is_some(),
|
email,
|
||||||
server_chan: config.server_chan.enabled && config.server_chan.send_key.is_some(),
|
server_chan: config.server_chan.enabled && config.server_chan.send_key.is_some(),
|
||||||
bark: config.bark.enabled && config.bark.device_key.is_some(),
|
bark: config.bark.enabled && config.bark.device_key.is_some(),
|
||||||
})
|
})
|
||||||
@@ -840,13 +842,50 @@ fn escape_html(value: &str) -> String {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::{
|
use super::{
|
||||||
apply_notification_item_template, parse_channel_filter, parse_notification_items,
|
apply_notification_item_template, important_notification_configured, parse_channel_filter,
|
||||||
parse_recipient_list, ImportantNotification, ImportantNotificationChannelFilter,
|
parse_notification_items, parse_recipient_list, ImportantNotification,
|
||||||
MAX_NOTIFICATION_ITEMS, MAX_NOTIFICATION_RECIPIENTS, MAX_NOTIFICATION_RECIPIENT_BYTES,
|
ImportantNotificationChannelFilter, IMPORTANT_NOTIFICATION_EMAIL_ENABLED_KEY,
|
||||||
|
IMPORTANT_NOTIFICATION_EMAIL_RECIPIENTS_KEY, MAX_NOTIFICATION_ITEMS,
|
||||||
|
MAX_NOTIFICATION_RECIPIENTS, MAX_NOTIFICATION_RECIPIENT_BYTES,
|
||||||
MAX_NOTIFICATION_TEMPLATE_BYTES,
|
MAX_NOTIFICATION_TEMPLATE_BYTES,
|
||||||
};
|
};
|
||||||
|
use crate::{data::GatewayDataState, AppState};
|
||||||
|
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn unused_email_channel_does_not_load_or_migrate_smtp_password() {
|
||||||
|
for (email_enabled, recipients) in [(false, "[email protected]"), (true, "")] {
|
||||||
|
let data = GatewayDataState::disabled()
|
||||||
|
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||||
|
.with_system_config_values_for_tests(vec![
|
||||||
|
(
|
||||||
|
IMPORTANT_NOTIFICATION_EMAIL_ENABLED_KEY.to_string(),
|
||||||
|
json!(email_enabled),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
IMPORTANT_NOTIFICATION_EMAIL_RECIPIENTS_KEY.to_string(),
|
||||||
|
json!(recipients),
|
||||||
|
),
|
||||||
|
("smtp_host".to_string(), json!("smtp.example.com")),
|
||||||
|
("smtp_user".to_string(), json!("[email protected]")),
|
||||||
|
("smtp_password".to_string(), json!("unused-smtp-password")),
|
||||||
|
("smtp_from_email".to_string(), json!("[email protected]")),
|
||||||
|
]);
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("gateway state should build")
|
||||||
|
.with_data_state_for_tests(data);
|
||||||
|
assert!(!important_notification_configured(&state).await.unwrap());
|
||||||
|
assert_eq!(
|
||||||
|
state
|
||||||
|
.read_system_config_json_value_strong("smtp_password")
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
Some(json!("unused-smtp-password"))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn parse_recipient_list_accepts_arrays_and_delimiters() {
|
fn parse_recipient_list_accepts_arrays_and_delimiters() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
|
|||||||
@@ -117,8 +117,8 @@ where
|
|||||||
|
|
||||||
use aether_crypto::warm_python_fernet_secret;
|
use aether_crypto::warm_python_fernet_secret;
|
||||||
use aether_data::lifecycle::export::{
|
use aether_data::lifecycle::export::{
|
||||||
copy_database_records, export_database_jsonl, import_database_jsonl, DataCopyOptions,
|
copy_database_records, export_database_jsonl, import_database_jsonl_with_options,
|
||||||
ExportDomain, MAX_JSONL_INPUT_BYTES,
|
DataCopyOptions, DataImportOptions, ExportDomain, MAX_JSONL_INPUT_BYTES,
|
||||||
};
|
};
|
||||||
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
||||||
use aether_gateway::{
|
use aether_gateway::{
|
||||||
@@ -1351,6 +1351,11 @@ struct DataExportArgs {
|
|||||||
struct DataImportArgs {
|
struct DataImportArgs {
|
||||||
#[arg(long)]
|
#[arg(long)]
|
||||||
input: PathBuf,
|
input: PathBuf,
|
||||||
|
#[arg(
|
||||||
|
long,
|
||||||
|
help = "Preserve passwords and API/management credentials from a trusted import; imported sessions remain revoked. Without this flag identity credentials are revoked."
|
||||||
|
)]
|
||||||
|
preserve_credentials: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(ClapArgs, Debug, Clone)]
|
#[derive(ClapArgs, Debug, Clone)]
|
||||||
@@ -1382,6 +1387,11 @@ struct DataCopyArgs {
|
|||||||
|
|
||||||
#[arg(long)]
|
#[arg(long)]
|
||||||
omit_request_body_details: bool,
|
omit_request_body_details: bool,
|
||||||
|
#[arg(
|
||||||
|
long,
|
||||||
|
help = "Preserve passwords and API/management credentials from the trusted source; imported sessions remain revoked. The target must use the source encryption key."
|
||||||
|
)]
|
||||||
|
preserve_credentials: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl GatewayLoggingArgs {
|
impl GatewayLoggingArgs {
|
||||||
@@ -2408,10 +2418,6 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
match state.prewarm_execution_extra_trusted_dns_hosts().await {
|
|
||||||
Ok(_) => info!("prewarmed execution Fake-IP DNS allowlist"),
|
|
||||||
Err(err) => warn!(error = %err, "failed to prewarm execution Fake-IP DNS allowlist"),
|
|
||||||
}
|
|
||||||
match prewarm_direct_h2c_sender_cache_from_env_for_startup().await {
|
match prewarm_direct_h2c_sender_cache_from_env_for_startup().await {
|
||||||
Ok(Some(report)) => {
|
Ok(Some(report)) => {
|
||||||
if report.failed_targets > 0 {
|
if report.failed_targets > 0 {
|
||||||
@@ -2909,12 +2915,23 @@ async fn run_data_import(
|
|||||||
let driver = database.driver;
|
let driver = database.driver;
|
||||||
let input_path = args.input.clone();
|
let input_path = args.input.clone();
|
||||||
let input = tokio::task::spawn_blocking(move || read_data_import_input(&input_path)).await??;
|
let input = tokio::task::spawn_blocking(move || read_data_import_input(&input_path)).await??;
|
||||||
let imported = import_database_jsonl(database, &input).await?;
|
if !args.preserve_credentials {
|
||||||
|
warn!("identity credentials will be revoked; use --preserve-credentials only for trusted recovery or migration");
|
||||||
|
}
|
||||||
|
let imported = import_database_jsonl_with_options(
|
||||||
|
database,
|
||||||
|
&input,
|
||||||
|
DataImportOptions {
|
||||||
|
preserve_credentials: args.preserve_credentials,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
|
|
||||||
info!(
|
info!(
|
||||||
driver = %driver,
|
driver = %driver,
|
||||||
input = %args.input.display(),
|
input = %args.input.display(),
|
||||||
imported,
|
imported,
|
||||||
|
preserve_credentials = args.preserve_credentials,
|
||||||
"database import complete"
|
"database import complete"
|
||||||
);
|
);
|
||||||
println!(
|
println!(
|
||||||
@@ -3175,6 +3192,9 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
|
|||||||
let target_driver = target.driver;
|
let target_driver = target.driver;
|
||||||
let domains = requested_domains(&args.domains);
|
let domains = requested_domains(&args.domains);
|
||||||
let created_at_unix_secs = current_unix_secs()?;
|
let created_at_unix_secs = current_unix_secs()?;
|
||||||
|
if !args.preserve_credentials {
|
||||||
|
warn!("identity credentials will be revoked; use --preserve-credentials only for trusted recovery or migration");
|
||||||
|
}
|
||||||
let imported = copy_database_records(
|
let imported = copy_database_records(
|
||||||
source,
|
source,
|
||||||
target,
|
target,
|
||||||
@@ -3182,6 +3202,7 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
|
|||||||
created_at_unix_secs,
|
created_at_unix_secs,
|
||||||
DataCopyOptions {
|
DataCopyOptions {
|
||||||
omit_request_body_details: args.omit_request_body_details,
|
omit_request_body_details: args.omit_request_body_details,
|
||||||
|
preserve_credentials: args.preserve_credentials,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.await?;
|
.await?;
|
||||||
@@ -3190,6 +3211,7 @@ async fn run_data_copy(args: &DataCopyArgs) -> Result<(), Box<dyn std::error::Er
|
|||||||
source_driver = %source_driver,
|
source_driver = %source_driver,
|
||||||
target_driver = %target_driver,
|
target_driver = %target_driver,
|
||||||
imported,
|
imported,
|
||||||
|
preserve_credentials = args.preserve_credentials,
|
||||||
"database copy complete"
|
"database copy complete"
|
||||||
);
|
);
|
||||||
println!(
|
println!(
|
||||||
@@ -4361,6 +4383,41 @@ mod tests {
|
|||||||
};
|
};
|
||||||
assert!(copy.source_allow_insecure);
|
assert!(copy.source_allow_insecure);
|
||||||
assert!(!copy.target_allow_insecure);
|
assert!(!copy.target_allow_insecure);
|
||||||
|
assert!(!copy.preserve_credentials);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn data_import_and_copy_require_explicit_credential_preservation() {
|
||||||
|
for preserve in [false, true] {
|
||||||
|
let mut import_args = vec!["aether-gateway", "import", "--input", "trusted.jsonl"];
|
||||||
|
let mut copy_args = vec![
|
||||||
|
"aether-gateway",
|
||||||
|
"copy",
|
||||||
|
"--source-driver",
|
||||||
|
"postgres",
|
||||||
|
"--source-url",
|
||||||
|
"postgres://localhost/source",
|
||||||
|
"--target-driver",
|
||||||
|
"postgres",
|
||||||
|
"--target-url",
|
||||||
|
"postgres://localhost/target",
|
||||||
|
];
|
||||||
|
if preserve {
|
||||||
|
import_args.push("--preserve-credentials");
|
||||||
|
copy_args.push("--preserve-credentials");
|
||||||
|
}
|
||||||
|
let Some(DataCommand::Import(import)) =
|
||||||
|
Args::try_parse_from(import_args).unwrap().command
|
||||||
|
else {
|
||||||
|
panic!("expected import command");
|
||||||
|
};
|
||||||
|
let Some(DataCommand::Copy(copy)) = Args::try_parse_from(copy_args).unwrap().command
|
||||||
|
else {
|
||||||
|
panic!("expected copy command");
|
||||||
|
};
|
||||||
|
assert_eq!(import.preserve_credentials, preserve);
|
||||||
|
assert_eq!(copy.preserve_credentials, preserve);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(unix)]
|
#[cfg(unix)]
|
||||||
|
|||||||
@@ -115,9 +115,6 @@ pub(super) fn usage_cleanup_window(
|
|||||||
usage_cleanup_window_with_override(now_utc, settings, None)
|
usage_cleanup_window_with_override(now_utc, settings, None)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Clamp is non-aggressive: each tier's cutoff becomes `max(policy_cutoff, now - override)`.
|
|
||||||
/// A later cutoff = fewer records deleted, so the override can only make cleanup more
|
|
||||||
/// conservative than the configured retention, never more destructive.
|
|
||||||
pub(super) fn usage_cleanup_window_with_override(
|
pub(super) fn usage_cleanup_window_with_override(
|
||||||
now_utc: DateTime<Utc>,
|
now_utc: DateTime<Utc>,
|
||||||
settings: UsageCleanupSettings,
|
settings: UsageCleanupSettings,
|
||||||
@@ -135,9 +132,9 @@ pub(super) fn usage_cleanup_window_with_override(
|
|||||||
};
|
};
|
||||||
let manual_cutoff = now_utc - override_duration;
|
let manual_cutoff = now_utc - override_duration;
|
||||||
UsageCleanupWindow {
|
UsageCleanupWindow {
|
||||||
detail_cutoff: policy.detail_cutoff.max(manual_cutoff),
|
detail_cutoff: policy.detail_cutoff.min(manual_cutoff),
|
||||||
compressed_cutoff: policy.compressed_cutoff.max(manual_cutoff),
|
compressed_cutoff: policy.compressed_cutoff.min(manual_cutoff),
|
||||||
header_cutoff: policy.header_cutoff.max(manual_cutoff),
|
header_cutoff: policy.header_cutoff.min(manual_cutoff),
|
||||||
log_cutoff: policy.log_cutoff.max(manual_cutoff),
|
log_cutoff: policy.log_cutoff.min(manual_cutoff),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1140,19 +1140,40 @@ fn usage_cleanup_window_with_override_is_always_non_aggressive() {
|
|||||||
let override_duration = chrono::Duration::days(180);
|
let override_duration = chrono::Duration::days(180);
|
||||||
let clamped = usage_cleanup_window_with_override(now_utc, settings, Some(override_duration));
|
let clamped = usage_cleanup_window_with_override(now_utc, settings, Some(override_duration));
|
||||||
|
|
||||||
assert_eq!(clamped.detail_cutoff, policy.detail_cutoff);
|
assert_eq!(clamped.detail_cutoff, now_utc - override_duration);
|
||||||
assert_eq!(clamped.compressed_cutoff, policy.compressed_cutoff);
|
assert_eq!(clamped.compressed_cutoff, now_utc - override_duration);
|
||||||
assert_eq!(clamped.header_cutoff, policy.header_cutoff);
|
assert_eq!(clamped.header_cutoff, now_utc - override_duration);
|
||||||
assert_eq!(clamped.log_cutoff, now_utc - override_duration);
|
assert_eq!(clamped.log_cutoff, policy.log_cutoff);
|
||||||
assert!(clamped.log_cutoff > policy.log_cutoff);
|
assert!(clamped.log_cutoff <= policy.log_cutoff);
|
||||||
|
|
||||||
let far_override = chrono::Duration::days(5);
|
let far_override = chrono::Duration::days(5);
|
||||||
let far = usage_cleanup_window_with_override(now_utc, settings, Some(far_override));
|
let far = usage_cleanup_window_with_override(now_utc, settings, Some(far_override));
|
||||||
assert_eq!(far.detail_cutoff, now_utc - far_override);
|
assert_eq!(far, policy);
|
||||||
assert_eq!(far.compressed_cutoff, now_utc - far_override);
|
|
||||||
assert_eq!(far.header_cutoff, now_utc - far_override);
|
for days in [0, 5, 30, 180, 400] {
|
||||||
assert_eq!(far.log_cutoff, now_utc - far_override);
|
let cutoff = now_utc - chrono::Duration::days(days);
|
||||||
assert!(far.log_cutoff > policy.log_cutoff);
|
let window = usage_cleanup_window_with_override(
|
||||||
|
now_utc,
|
||||||
|
settings,
|
||||||
|
Some(chrono::Duration::days(days)),
|
||||||
|
);
|
||||||
|
for (actual, configured) in [
|
||||||
|
(window.detail_cutoff, policy.detail_cutoff),
|
||||||
|
(window.compressed_cutoff, policy.compressed_cutoff),
|
||||||
|
(window.header_cutoff, policy.header_cutoff),
|
||||||
|
(window.log_cutoff, policy.log_cutoff),
|
||||||
|
] {
|
||||||
|
assert!(actual <= configured);
|
||||||
|
assert!(actual <= cutoff);
|
||||||
|
for age in [1, 7, 15, 30, 90, 180, 365, 401] {
|
||||||
|
let created_at = now_utc - chrono::Duration::days(age);
|
||||||
|
if created_at < actual {
|
||||||
|
assert!(created_at < configured);
|
||||||
|
assert!(created_at < cutoff);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
let passthrough = usage_cleanup_window_with_override(now_utc, settings, None);
|
let passthrough = usage_cleanup_window_with_override(now_utc, settings, None);
|
||||||
assert_eq!(passthrough, policy);
|
assert_eq!(passthrough, policy);
|
||||||
|
|||||||
@@ -31,7 +31,11 @@ fn apply_frontdoor_cors_headers(
|
|||||||
);
|
);
|
||||||
headers.insert(
|
headers.insert(
|
||||||
http::header::ACCESS_CONTROL_EXPOSE_HEADERS,
|
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 Some(value) = requested_headers {
|
||||||
if let Ok(value) = HeaderValue::from_str(value) {
|
if let Ok(value) = HeaderValue::from_str(value) {
|
||||||
@@ -109,3 +113,36 @@ pub(crate) async fn frontdoor_cors_middleware(
|
|||||||
);
|
);
|
||||||
response
|
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"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -533,12 +533,10 @@ impl AppState {
|
|||||||
&self,
|
&self,
|
||||||
provider_ids: &[String],
|
provider_ids: &[String],
|
||||||
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
|
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
|
||||||
let keys = self
|
self.data
|
||||||
.data
|
|
||||||
.list_provider_catalog_key_summaries_by_provider_ids(provider_ids)
|
.list_provider_catalog_key_summaries_by_provider_ids(provider_ids)
|
||||||
.await
|
.await
|
||||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||||
self.open_provider_catalog_keys(keys).await
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn list_provider_catalog_key_maintenance_summaries_by_provider_ids(
|
pub(crate) async fn list_provider_catalog_key_maintenance_summaries_by_provider_ids(
|
||||||
|
|||||||
@@ -272,6 +272,50 @@ mod tests {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn app_state_reads_redacted_key_summaries_without_opening_credentials() {
|
||||||
|
for (api_key, auth_config) in [
|
||||||
|
(Some("summary"), None),
|
||||||
|
(None, Some("{}")),
|
||||||
|
(Some("summary"), Some("{}")),
|
||||||
|
] {
|
||||||
|
let health = serde_json::json!({"openai:chat": {"health_score": 0.75}});
|
||||||
|
let key = sample_key(
|
||||||
|
"key-1",
|
||||||
|
"provider-1",
|
||||||
|
api_key.map(ToOwned::to_owned),
|
||||||
|
auth_config.map(ToOwned::to_owned),
|
||||||
|
)
|
||||||
|
.with_health_fields(Some(health.clone()), None);
|
||||||
|
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![sample_provider("provider-1")],
|
||||||
|
Vec::new(),
|
||||||
|
vec![key],
|
||||||
|
));
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("test state should build")
|
||||||
|
.with_data_state_for_tests(
|
||||||
|
GatewayDataState::with_provider_catalog_reader_for_tests(repository)
|
||||||
|
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||||
|
);
|
||||||
|
let provider_ids = ["provider-1".to_string()];
|
||||||
|
|
||||||
|
let summaries = state
|
||||||
|
.list_provider_catalog_key_summaries_by_provider_ids(&provider_ids)
|
||||||
|
.await
|
||||||
|
.expect("redacted summaries should not require credential authentication");
|
||||||
|
|
||||||
|
assert_eq!(summaries.len(), 1);
|
||||||
|
assert_eq!(summaries[0].health_by_format.as_ref(), Some(&health));
|
||||||
|
assert_eq!(summaries[0].encrypted_api_key.as_deref(), api_key);
|
||||||
|
assert_eq!(summaries[0].encrypted_auth_config.as_deref(), auth_config);
|
||||||
|
assert!(state
|
||||||
|
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
|
||||||
|
.await
|
||||||
|
.is_err());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn app_state_migrates_both_legacy_fields_with_one_exact_cas() {
|
async fn app_state_migrates_both_legacy_fields_with_one_exact_cas() {
|
||||||
let legacy_api =
|
let legacy_api =
|
||||||
|
|||||||
@@ -154,15 +154,6 @@ impl AppState {
|
|||||||
.map_err(|err| format!("{err:?}"))
|
.map_err(|err| format!("{err:?}"))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn prewarm_execution_extra_trusted_dns_hosts(&self) -> Result<(), String> {
|
|
||||||
self.read_system_config_json_value(
|
|
||||||
aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map(|_| ())
|
|
||||||
.map_err(|err| format!("{err:?}"))
|
|
||||||
}
|
|
||||||
|
|
||||||
fn usage_worker_queue_for(
|
fn usage_worker_queue_for(
|
||||||
runtime_state: &Arc<RuntimeState>,
|
runtime_state: &Arc<RuntimeState>,
|
||||||
) -> Option<Arc<dyn RuntimeQueueStore>> {
|
) -> Option<Arc<dyn RuntimeQueueStore>> {
|
||||||
@@ -778,18 +769,6 @@ impl AppState {
|
|||||||
.expect("admin monitoring error stats reset cache should lock")
|
.expect("admin monitoring error stats reset cache should lock")
|
||||||
}
|
}
|
||||||
|
|
||||||
fn refresh_execution_extra_trusted_dns_hosts(
|
|
||||||
&self,
|
|
||||||
key: &str,
|
|
||||||
value: Option<&serde_json::Value>,
|
|
||||||
) {
|
|
||||||
if key.eq_ignore_ascii_case(
|
|
||||||
aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY,
|
|
||||||
) {
|
|
||||||
crate::execution_runtime::transport::refresh_execution_extra_trusted_dns_hosts(value);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
pub(crate) fn mark_admin_monitoring_error_stats_reset(&self, now_unix_secs: u64) {
|
pub(crate) fn mark_admin_monitoring_error_stats_reset(&self, now_unix_secs: u64) {
|
||||||
let mut reset_at = self
|
let mut reset_at = self
|
||||||
.admin_monitoring_error_stats_reset_at
|
.admin_monitoring_error_stats_reset_at
|
||||||
@@ -809,7 +788,6 @@ impl AppState {
|
|||||||
SYSTEM_CONFIG_CACHE_MAX_STALENESS,
|
SYSTEM_CONFIG_CACHE_MAX_STALENESS,
|
||||||
)
|
)
|
||||||
.await?;
|
.await?;
|
||||||
self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref());
|
|
||||||
Ok(value)
|
Ok(value)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -822,7 +800,6 @@ impl AppState {
|
|||||||
.find_system_config_value_strong(key)
|
.find_system_config_value_strong(key)
|
||||||
.await
|
.await
|
||||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||||
self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref());
|
|
||||||
Ok(value)
|
Ok(value)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -961,7 +938,6 @@ impl AppState {
|
|||||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||||
self.system_config_cache
|
self.system_config_cache
|
||||||
.insert(key.to_string(), None, SYSTEM_CONFIG_CACHE_MAX_STALENESS);
|
.insert(key.to_string(), None, SYSTEM_CONFIG_CACHE_MAX_STALENESS);
|
||||||
self.refresh_execution_extra_trusted_dns_hosts(key, None);
|
|
||||||
if deleted && system_config_key_affects_scheduler(key) {
|
if deleted && system_config_key_affects_scheduler(key) {
|
||||||
self.invalidate_scheduler_affinity_cache();
|
self.invalidate_scheduler_affinity_cache();
|
||||||
}
|
}
|
||||||
@@ -1043,7 +1019,6 @@ impl AppState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn remember_system_config_write(&self, key: &str, value: Option<serde_json::Value>) {
|
fn remember_system_config_write(&self, key: &str, value: Option<serde_json::Value>) {
|
||||||
self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref());
|
|
||||||
self.system_config_cache
|
self.system_config_cache
|
||||||
.insert(key.to_string(), value, SYSTEM_CONFIG_CACHE_MAX_STALENESS);
|
.insert(key.to_string(), value, SYSTEM_CONFIG_CACHE_MAX_STALENESS);
|
||||||
if system_config_key_affects_scheduler(key) {
|
if system_config_key_affects_scheduler(key) {
|
||||||
@@ -1091,7 +1066,6 @@ impl AppState {
|
|||||||
| aether_data::repository::system::AdminSystemPurgeTarget::Stats
|
| aether_data::repository::system::AdminSystemPurgeTarget::Stats
|
||||||
) {
|
) {
|
||||||
self.system_config_cache.clear();
|
self.system_config_cache.clear();
|
||||||
crate::execution_runtime::transport::refresh_execution_extra_trusted_dns_hosts(None);
|
|
||||||
self.invalidate_provider_routing_caches();
|
self.invalidate_provider_routing_caches();
|
||||||
}
|
}
|
||||||
Ok(summary)
|
Ok(summary)
|
||||||
|
|||||||
@@ -2460,15 +2460,21 @@ async fn gateway_retries_next_local_openai_chat_stream_candidate_after_retryable
|
|||||||
failed_candidate.error_type.as_deref(),
|
failed_candidate.error_type.as_deref(),
|
||||||
Some("retryable_upstream_status")
|
Some("retryable_upstream_status")
|
||||||
);
|
);
|
||||||
assert!(failed_candidate.error_message.is_none());
|
assert!(failed_candidate.error_message.is_some());
|
||||||
let failed_upstream_response = failed_candidate
|
let failed_upstream_response = failed_candidate
|
||||||
.extra_data
|
.extra_data
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|value| value.get("upstream_response"))
|
.and_then(|value| value.get("upstream_response"))
|
||||||
.expect("failed stream candidate should keep its upstream response");
|
.expect("failed stream candidate should keep its upstream response");
|
||||||
assert_eq!(failed_upstream_response["status_code"], json!(429));
|
assert_eq!(failed_upstream_response["status_code"], json!(429));
|
||||||
assert!(failed_upstream_response.get("headers").is_none());
|
assert_eq!(
|
||||||
assert!(failed_upstream_response.get("body").is_none());
|
failed_upstream_response["headers"]["content-type"],
|
||||||
|
"application/json"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
failed_upstream_response["body"]["error"]["message"],
|
||||||
|
"rate limited"
|
||||||
|
);
|
||||||
assert_eq!(success_candidate.status, RequestCandidateStatus::Success);
|
assert_eq!(success_candidate.status, RequestCandidateStatus::Success);
|
||||||
assert_eq!(success_candidate.status_code, Some(200));
|
assert_eq!(success_candidate.status_code, Some(200));
|
||||||
assert!(success_candidate.started_at_unix_ms.is_some());
|
assert!(success_candidate.started_at_unix_ms.is_some());
|
||||||
|
|||||||
@@ -1166,15 +1166,21 @@ async fn gateway_retries_next_local_openai_chat_sync_candidate_after_auth_failur
|
|||||||
assert_eq!(failed_candidate.retry_index, retry_index as u32);
|
assert_eq!(failed_candidate.retry_index, retry_index as u32);
|
||||||
assert_eq!(failed_candidate.status, RequestCandidateStatus::Failed);
|
assert_eq!(failed_candidate.status, RequestCandidateStatus::Failed);
|
||||||
assert_eq!(failed_candidate.status_code, Some(401));
|
assert_eq!(failed_candidate.status_code, Some(401));
|
||||||
assert!(failed_candidate.error_message.is_none());
|
assert!(failed_candidate.error_message.is_some());
|
||||||
let failed_upstream_response = failed_candidate
|
let failed_upstream_response = failed_candidate
|
||||||
.extra_data
|
.extra_data
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|value| value.get("upstream_response"))
|
.and_then(|value| value.get("upstream_response"))
|
||||||
.expect("failed candidate should keep its upstream response");
|
.expect("failed candidate should keep its upstream response");
|
||||||
assert_eq!(failed_upstream_response["status_code"], json!(401));
|
assert_eq!(failed_upstream_response["status_code"], json!(401));
|
||||||
assert!(failed_upstream_response.get("headers").is_none());
|
assert_eq!(
|
||||||
assert!(failed_upstream_response.get("body").is_none());
|
failed_upstream_response["headers"]["content-type"],
|
||||||
|
"application/json"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
failed_upstream_response["body"]["error"]["message"],
|
||||||
|
"invalid auth token"
|
||||||
|
);
|
||||||
}
|
}
|
||||||
assert_eq!(stored_candidates[2].candidate_index, 1);
|
assert_eq!(stored_candidates[2].candidate_index, 1);
|
||||||
assert_eq!(stored_candidates[2].status, RequestCandidateStatus::Success);
|
assert_eq!(stored_candidates[2].status, RequestCandidateStatus::Success);
|
||||||
|
|||||||
@@ -5010,6 +5010,7 @@ fn retired_api_format_occurrences_are_whitelisted() {
|
|||||||
"crates/aether-ai/formats/src/formats/registry.rs",
|
"crates/aether-ai/formats/src/formats/registry.rs",
|
||||||
"crates/aether-data/runtime/src/migrate.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.rs",
|
||||||
|
"crates/aether-data/runtime/src/lifecycle/migrate/tests/policy_nulls.rs",
|
||||||
"crates/aether-usage/runtime/src/report.rs",
|
"crates/aether-usage/runtime/src/report.rs",
|
||||||
"frontend/src/api/endpoints/types/__tests__/api-format.spec.ts",
|
"frontend/src/api/endpoints/types/__tests__/api-format.spec.ts",
|
||||||
"frontend/src/views/admin/module-management/modelDirectivesConfig.ts",
|
"frontend/src/views/admin/module-management/modelDirectivesConfig.ts",
|
||||||
|
|||||||
@@ -326,7 +326,7 @@ async fn gateway_reads_video_task_detail_via_internal_async_task_endpoint() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endpoint() {
|
async fn gateway_redirects_persisted_openai_video_url_from_authenticated_internal_endpoint() {
|
||||||
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
||||||
let mut task = sample_video_task(
|
let mut task = sample_video_task(
|
||||||
"task-redirect",
|
"task-redirect",
|
||||||
@@ -341,7 +341,10 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endp
|
|||||||
.upsert(task)
|
.upsert(task)
|
||||||
.await
|
.await
|
||||||
.expect("upsert should succeed");
|
.expect("upsert should succeed");
|
||||||
assert_eq!(stored.video_url, None);
|
assert_eq!(
|
||||||
|
stored.video_url.as_deref(),
|
||||||
|
Some("https://8.8.8.8/video-task-redirect.mp4")
|
||||||
|
);
|
||||||
|
|
||||||
let state = AppState::new()
|
let state = AppState::new()
|
||||||
.expect("gateway state should build")
|
.expect("gateway state should build")
|
||||||
@@ -349,7 +352,10 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endp
|
|||||||
let (gateway_url, gateway_handle, access_token) =
|
let (gateway_url, gateway_handle, access_token) =
|
||||||
start_authenticated_operational_server(state).await;
|
start_authenticated_operational_server(state).await;
|
||||||
|
|
||||||
let client = authenticated_operational_client(&access_token);
|
let client = super::authenticated_operational_client_with_builder(
|
||||||
|
reqwest::Client::builder().redirect(reqwest::redirect::Policy::none()),
|
||||||
|
&access_token,
|
||||||
|
);
|
||||||
let response = client
|
let response = client
|
||||||
.get(format!(
|
.get(format!(
|
||||||
"{gateway_url}/_gateway/async-tasks/video-tasks/task-redirect/video"
|
"{gateway_url}/_gateway/async-tasks/video-tasks/task-redirect/video"
|
||||||
@@ -358,7 +364,14 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_from_internal_endp
|
|||||||
.await
|
.await
|
||||||
.expect("request should succeed");
|
.expect("request should succeed");
|
||||||
|
|
||||||
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT);
|
||||||
|
assert_eq!(
|
||||||
|
response
|
||||||
|
.headers()
|
||||||
|
.get("location")
|
||||||
|
.and_then(|value| value.to_str().ok()),
|
||||||
|
stored.video_url.as_deref()
|
||||||
|
);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -534,6 +534,24 @@ async fn gateway_exposes_request_audit_bundle_via_internal_audit_endpoint() {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() {
|
async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() {
|
||||||
|
let mut failed_candidate = sample_request_candidate(
|
||||||
|
"cand-2",
|
||||||
|
"req-trace-1",
|
||||||
|
1,
|
||||||
|
RequestCandidateStatus::Failed,
|
||||||
|
Some(101),
|
||||||
|
Some(37),
|
||||||
|
Some(502),
|
||||||
|
);
|
||||||
|
failed_candidate.error_message = Some("private upstream diagnostic".to_string());
|
||||||
|
failed_candidate.extra_data = Some(json!({
|
||||||
|
"upstream_response": {
|
||||||
|
"status_code": 502,
|
||||||
|
"headers": {"x-request-id": "private-upstream-id"},
|
||||||
|
"body": {"error": {"message": "private upstream diagnostic"}}
|
||||||
|
},
|
||||||
|
"error_flow": {"status_code": 502, "message": "private upstream diagnostic"}
|
||||||
|
}));
|
||||||
let repository = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
let repository = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||||
sample_request_candidate(
|
sample_request_candidate(
|
||||||
"cand-1",
|
"cand-1",
|
||||||
@@ -544,15 +562,7 @@ async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() {
|
|||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
),
|
),
|
||||||
sample_request_candidate(
|
failed_candidate,
|
||||||
"cand-2",
|
|
||||||
"req-trace-1",
|
|
||||||
1,
|
|
||||||
RequestCandidateStatus::Failed,
|
|
||||||
Some(101),
|
|
||||||
Some(37),
|
|
||||||
Some(502),
|
|
||||||
),
|
|
||||||
]));
|
]));
|
||||||
|
|
||||||
let gateway_state = AppState::new()
|
let gateway_state = AppState::new()
|
||||||
@@ -582,6 +592,14 @@ async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() {
|
|||||||
);
|
);
|
||||||
assert_eq!(payload["candidates"][0]["id"], "cand-2");
|
assert_eq!(payload["candidates"][0]["id"], "cand-2");
|
||||||
assert_eq!(payload["candidates"][0]["status"], "failed");
|
assert_eq!(payload["candidates"][0]["status"], "failed");
|
||||||
|
assert!(payload["candidates"][0]["error_message"].is_null());
|
||||||
|
assert_eq!(
|
||||||
|
payload["candidates"][0]["extra_data"]["upstream_response"]["status_code"],
|
||||||
|
502
|
||||||
|
);
|
||||||
|
let serialized = payload.to_string();
|
||||||
|
assert!(!serialized.contains("private upstream diagnostic"));
|
||||||
|
assert!(!serialized.contains("private-upstream-id"));
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
mod keys;
|
mod keys;
|
||||||
mod quota;
|
mod quota;
|
||||||
mod routes;
|
mod routes;
|
||||||
|
mod rules_reveal;
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ use aether_data::repository::global_models::InMemoryGlobalModelReadRepository;
|
|||||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||||
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
|
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
|
||||||
use aether_data_contracts::repository::global_models::{
|
use aether_data_contracts::repository::global_models::{
|
||||||
AdminProviderModelListQuery, GlobalModelReadRepository,
|
AdminGlobalModelListQuery, AdminProviderModelListQuery, GlobalModelReadRepository,
|
||||||
};
|
};
|
||||||
use aether_data_contracts::repository::provider_catalog::{
|
use aether_data_contracts::repository::provider_catalog::{
|
||||||
ProviderCatalogReadRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
ProviderCatalogReadRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||||
@@ -20,8 +20,9 @@ use http::StatusCode;
|
|||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
|
||||||
use super::super::super::{
|
use super::super::super::{
|
||||||
build_router_with_state, build_state_with_execution_runtime_override, sample_bound_auth_config,
|
build_router_with_state, build_state_with_execution_runtime_override,
|
||||||
sample_bound_key, sample_endpoint, sample_key, sample_proxy_node, start_server, AppState,
|
sample_admin_global_model, sample_bound_auth_config, sample_bound_key, sample_endpoint,
|
||||||
|
sample_key, sample_proxy_node, start_server, AppState,
|
||||||
};
|
};
|
||||||
use crate::constants::{
|
use crate::constants::{
|
||||||
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
|
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],
|
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 (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||||
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).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")),
|
.and_then(|value| value.get("remaining_fraction")),
|
||||||
Some(&json!(0.25))
|
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 {
|
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||||
provider_id: "provider-antigravity".to_string(),
|
provider_id: "provider-antigravity".to_string(),
|
||||||
is_active: None,
|
is_active: None,
|
||||||
@@ -2541,15 +2559,8 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru
|
|||||||
limit: 100,
|
limit: 100,
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
.expect("imported Antigravity provider models should read");
|
.expect("Antigravity provider models should read after quota refresh");
|
||||||
let imported_model_names = imported_provider_models
|
assert!(provider_models.is_empty());
|
||||||
.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"));
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
reloaded[0]
|
reloaded[0]
|
||||||
.upstream_metadata
|
.upstream_metadata
|
||||||
|
|||||||
@@ -479,7 +479,7 @@ async fn gateway_returns_service_unavailable_for_admin_provider_endpoint_create_
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_principal() {
|
async fn gateway_creates_admin_http_provider_endpoint_locally_with_trusted_admin_principal() {
|
||||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||||
let upstream = Router::new().route(
|
let upstream = Router::new().route(
|
||||||
@@ -522,7 +522,7 @@ async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_prin
|
|||||||
.json(&json!({
|
.json(&json!({
|
||||||
"provider_id": "provider-openai",
|
"provider_id": "provider-openai",
|
||||||
"api_format": "openai:chat",
|
"api_format": "openai:chat",
|
||||||
"base_url": "https://api.openai.example/",
|
"base_url": "http://api.openai.example:8080/",
|
||||||
"custom_path": "/v1/chat/completions",
|
"custom_path": "/v1/chat/completions",
|
||||||
"max_retries": 5,
|
"max_retries": 5,
|
||||||
"config": {"foo": "bar"},
|
"config": {"foo": "bar"},
|
||||||
@@ -537,7 +537,7 @@ async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_prin
|
|||||||
assert_eq!(payload["provider_id"], "provider-openai");
|
assert_eq!(payload["provider_id"], "provider-openai");
|
||||||
assert_eq!(payload["provider_name"], "openai");
|
assert_eq!(payload["provider_name"], "openai");
|
||||||
assert_eq!(payload["api_format"], "openai:chat");
|
assert_eq!(payload["api_format"], "openai:chat");
|
||||||
assert_eq!(payload["base_url"], "https://api.openai.example");
|
assert_eq!(payload["base_url"], "http://api.openai.example:8080");
|
||||||
assert_eq!(payload["custom_path"], "/v1/chat/completions");
|
assert_eq!(payload["custom_path"], "/v1/chat/completions");
|
||||||
assert_eq!(payload["max_retries"], 5);
|
assert_eq!(payload["max_retries"], 5);
|
||||||
assert_eq!(payload["total_keys"], 0);
|
assert_eq!(payload["total_keys"], 0);
|
||||||
@@ -553,7 +553,7 @@ async fn gateway_creates_admin_provider_endpoint_locally_with_trusted_admin_prin
|
|||||||
assert_eq!(endpoints.len(), 1);
|
assert_eq!(endpoints.len(), 1);
|
||||||
assert_eq!(endpoints[0].provider_id, "provider-openai");
|
assert_eq!(endpoints[0].provider_id, "provider-openai");
|
||||||
assert_eq!(endpoints[0].api_format, "openai:chat");
|
assert_eq!(endpoints[0].api_format, "openai:chat");
|
||||||
assert_eq!(endpoints[0].base_url, "https://api.openai.example");
|
assert_eq!(endpoints[0].base_url, "http://api.openai.example:8080");
|
||||||
assert_eq!(endpoints[0].max_retries, Some(5));
|
assert_eq!(endpoints[0].max_retries, Some(5));
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
@@ -658,7 +658,7 @@ async fn gateway_rejects_streaming_policy_for_search_endpoint_before_catalog_wri
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_principal() {
|
async fn gateway_updates_admin_http_provider_endpoint_locally_with_trusted_admin_principal() {
|
||||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||||
let upstream = Router::new().route(
|
let upstream = Router::new().route(
|
||||||
@@ -720,7 +720,7 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin
|
|||||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||||
.json(&json!({
|
.json(&json!({
|
||||||
"base_url": "https://updated.openai.example/",
|
"base_url": "http://updated.openai.example:8080/",
|
||||||
"custom_path": "/v1/responses",
|
"custom_path": "/v1/responses",
|
||||||
"max_retries": 5,
|
"max_retries": 5,
|
||||||
"is_active": false,
|
"is_active": false,
|
||||||
@@ -736,7 +736,7 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin
|
|||||||
assert_eq!(payload["id"], "endpoint-openai-chat");
|
assert_eq!(payload["id"], "endpoint-openai-chat");
|
||||||
assert_eq!(payload["provider_id"], "provider-openai");
|
assert_eq!(payload["provider_id"], "provider-openai");
|
||||||
assert_eq!(payload["api_format"], "openai:chat");
|
assert_eq!(payload["api_format"], "openai:chat");
|
||||||
assert_eq!(payload["base_url"], "https://updated.openai.example");
|
assert_eq!(payload["base_url"], "http://updated.openai.example:8080");
|
||||||
assert_eq!(payload["custom_path"], "/v1/responses");
|
assert_eq!(payload["custom_path"], "/v1/responses");
|
||||||
assert_eq!(payload["max_retries"], 5);
|
assert_eq!(payload["max_retries"], 5);
|
||||||
assert_eq!(payload["is_active"], false);
|
assert_eq!(payload["is_active"], false);
|
||||||
@@ -751,7 +751,7 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin
|
|||||||
.await
|
.await
|
||||||
.expect("endpoints should read");
|
.expect("endpoints should read");
|
||||||
assert_eq!(endpoints.len(), 1);
|
assert_eq!(endpoints.len(), 1);
|
||||||
assert_eq!(endpoints[0].base_url, "https://updated.openai.example");
|
assert_eq!(endpoints[0].base_url, "http://updated.openai.example:8080");
|
||||||
assert_eq!(endpoints[0].custom_path.as_deref(), Some("/v1/responses"));
|
assert_eq!(endpoints[0].custom_path.as_deref(), Some("/v1/responses"));
|
||||||
assert_eq!(endpoints[0].max_retries, Some(5));
|
assert_eq!(endpoints[0].max_retries, Some(5));
|
||||||
assert!(!endpoints[0].is_active);
|
assert!(!endpoints[0].is_active);
|
||||||
|
|||||||
@@ -0,0 +1,151 @@
|
|||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||||
|
use axum::body::Body;
|
||||||
|
use http::{HeaderMap, HeaderValue, Method, Request, StatusCode};
|
||||||
|
use http_body_util::BodyExt;
|
||||||
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
|
use super::super::super::{build_router_with_state, sample_endpoint, sample_provider, AppState};
|
||||||
|
use crate::admin_api::{maybe_build_local_admin_response, AdminRouteRequest};
|
||||||
|
use crate::audit::AdminAuditEvent;
|
||||||
|
use crate::constants::{
|
||||||
|
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
|
||||||
|
TRUSTED_ADMIN_USER_ROLE_HEADER,
|
||||||
|
};
|
||||||
|
use crate::control::resolve_public_request_context;
|
||||||
|
use crate::data::GatewayDataState;
|
||||||
|
use crate::tests::send_request;
|
||||||
|
|
||||||
|
fn seeded_state() -> AppState {
|
||||||
|
let mut endpoint = sample_endpoint(
|
||||||
|
"endpoint-rules",
|
||||||
|
"provider-rules",
|
||||||
|
"openai:chat",
|
||||||
|
"https://example.test",
|
||||||
|
);
|
||||||
|
endpoint.header_rules =
|
||||||
|
Some(json!([{"action": "set", "key": "x-auth", "value": "request-secret"}]));
|
||||||
|
endpoint.body_rules =
|
||||||
|
Some(json!([{"action": "set", "path": "auth.token", "value": "body-secret"}]));
|
||||||
|
endpoint.config = Some(json!({
|
||||||
|
"private_token": "unrelated-secret",
|
||||||
|
"response_header_rules": [{"action": "set", "key": "x-auth", "value": "response-secret"}]
|
||||||
|
}));
|
||||||
|
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![sample_provider("provider-rules", "custom", 10)],
|
||||||
|
vec![endpoint],
|
||||||
|
vec![],
|
||||||
|
));
|
||||||
|
AppState::new().unwrap().with_data_state_for_tests(
|
||||||
|
GatewayDataState::with_provider_catalog_reader_for_tests(repository),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn admin_headers() -> HeaderMap {
|
||||||
|
let mut headers = HeaderMap::new();
|
||||||
|
for (name, value) in [
|
||||||
|
(GATEWAY_HEADER, "rust-phase3b"),
|
||||||
|
(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user"),
|
||||||
|
(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin"),
|
||||||
|
(TRUSTED_ADMIN_SESSION_ID_HEADER, "admin-session"),
|
||||||
|
] {
|
||||||
|
headers.insert(name, HeaderValue::from_static(value));
|
||||||
|
}
|
||||||
|
headers
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn endpoint_rules_reveal_is_scoped_audited_and_not_cached() {
|
||||||
|
let state = seeded_state();
|
||||||
|
let context = resolve_public_request_context(
|
||||||
|
&state,
|
||||||
|
&Method::GET,
|
||||||
|
&"/api/admin/endpoints/endpoint-rules/rules/reveal"
|
||||||
|
.parse()
|
||||||
|
.unwrap(),
|
||||||
|
&admin_headers(),
|
||||||
|
"reveal-test",
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let response = maybe_build_local_admin_response(AdminRouteRequest::new(
|
||||||
|
&state,
|
||||||
|
&context,
|
||||||
|
&"127.0.0.1:12345".parse().unwrap(),
|
||||||
|
&admin_headers(),
|
||||||
|
None,
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
assert_eq!(response.headers()[http::header::CACHE_CONTROL], "no-store");
|
||||||
|
assert_eq!(response.headers()[http::header::PRAGMA], "no-cache");
|
||||||
|
let audit = response.extensions().get::<AdminAuditEvent>().unwrap();
|
||||||
|
assert_eq!(audit.event_name, "admin_endpoint_rules_revealed");
|
||||||
|
assert_eq!(audit.action, "reveal_endpoint_rules");
|
||||||
|
assert_eq!(audit.target_id, "endpoint-rules");
|
||||||
|
let body = response.into_body().collect().await.unwrap().to_bytes();
|
||||||
|
let payload: Value = serde_json::from_slice(&body).unwrap();
|
||||||
|
assert_eq!(payload["header_rules"][0]["value"], "request-secret");
|
||||||
|
assert_eq!(payload["body_rules"][0]["value"], "body-secret");
|
||||||
|
assert_eq!(
|
||||||
|
payload["response_header_rules"][0]["value"],
|
||||||
|
"response-secret"
|
||||||
|
);
|
||||||
|
assert_eq!(payload.as_object().unwrap().len(), 3);
|
||||||
|
assert!(!payload.to_string().contains("unrelated-secret"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn endpoint_rules_reveal_denies_anonymous_and_non_admin_requests() {
|
||||||
|
let router = build_router_with_state(seeded_state());
|
||||||
|
for role in [None, Some("user")] {
|
||||||
|
let mut request =
|
||||||
|
Request::builder().uri("/api/admin/endpoints/endpoint-rules/rules/reveal");
|
||||||
|
if let Some(role) = role {
|
||||||
|
request = request
|
||||||
|
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||||
|
.header(TRUSTED_ADMIN_USER_ID_HEADER, "normal-user")
|
||||||
|
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, role)
|
||||||
|
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "user-session");
|
||||||
|
}
|
||||||
|
let response = send_request(router.clone(), request.body(Body::empty()).unwrap()).await;
|
||||||
|
assert!(matches!(
|
||||||
|
response.status(),
|
||||||
|
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN
|
||||||
|
));
|
||||||
|
let body = response.into_body().collect().await.unwrap().to_bytes();
|
||||||
|
assert!(!String::from_utf8_lossy(&body).contains("request-secret"));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn endpoint_rules_reveal_returns_not_found_and_data_unavailable_without_fallback() {
|
||||||
|
for (state, expected) in [
|
||||||
|
(seeded_state(), StatusCode::NOT_FOUND),
|
||||||
|
(AppState::new().unwrap(), StatusCode::SERVICE_UNAVAILABLE),
|
||||||
|
] {
|
||||||
|
let context = resolve_public_request_context(
|
||||||
|
&state,
|
||||||
|
&Method::GET,
|
||||||
|
&"/api/admin/endpoints/missing/rules/reveal".parse().unwrap(),
|
||||||
|
&admin_headers(),
|
||||||
|
"reveal-missing-test",
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let response = maybe_build_local_admin_response(AdminRouteRequest::new(
|
||||||
|
&state,
|
||||||
|
&context,
|
||||||
|
&"127.0.0.1:12345".parse().unwrap(),
|
||||||
|
&admin_headers(),
|
||||||
|
None,
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(response.status(), expected);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,7 +1,7 @@
|
|||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{Arc, Mutex};
|
||||||
use std::time::{SystemTime, UNIX_EPOCH};
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||||
use aether_data::repository::auth_modules::InMemoryAuthModuleReadRepository;
|
use aether_data::repository::auth_modules::InMemoryAuthModuleReadRepository;
|
||||||
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
|
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
|
||||||
use aether_data::repository::management_tokens::{
|
use aether_data::repository::management_tokens::{
|
||||||
@@ -32,6 +32,155 @@ use crate::data::GatewayDataState;
|
|||||||
const ADMIN_ENDPOINT_HEALTH_DATA_UNAVAILABLE_DETAIL: &str =
|
const ADMIN_ENDPOINT_HEALTH_DATA_UNAVAILABLE_DETAIL: &str =
|
||||||
"Admin endpoint health data unavailable";
|
"Admin endpoint health data unavailable";
|
||||||
|
|
||||||
|
async fn assert_admin_modules_status_with_smtp_password(
|
||||||
|
stored_password: &str,
|
||||||
|
notification_ready: bool,
|
||||||
|
server_chan_enabled: bool,
|
||||||
|
) -> AppState {
|
||||||
|
let data = GatewayDataState::with_auth_module_reader_for_tests(Arc::new(
|
||||||
|
InMemoryAuthModuleReadRepository::seed(Vec::new(), None),
|
||||||
|
))
|
||||||
|
.with_provider_catalog_reader(Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
Vec::new(),
|
||||||
|
Vec::new(),
|
||||||
|
Vec::new(),
|
||||||
|
)))
|
||||||
|
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||||
|
.with_system_config_values_for_tests(vec![
|
||||||
|
("module.management_tokens.enabled".to_string(), json!(true)),
|
||||||
|
(
|
||||||
|
"module.important_notification.enabled".to_string(),
|
||||||
|
json!(true),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"module.important_notification.email_enabled".to_string(),
|
||||||
|
json!(true),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"module.important_notification.email_recipients".to_string(),
|
||||||
|
json!("[email protected]"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"module.server_chan_push.enabled".to_string(),
|
||||||
|
json!(server_chan_enabled),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"module.server_chan_push.send_key".to_string(),
|
||||||
|
json!(if server_chan_enabled {
|
||||||
|
"SCT-test-send-key"
|
||||||
|
} else {
|
||||||
|
""
|
||||||
|
}),
|
||||||
|
),
|
||||||
|
("smtp_host".to_string(), json!("smtp.example.com")),
|
||||||
|
("smtp_port".to_string(), json!(587)),
|
||||||
|
("smtp_user".to_string(), json!("[email protected]")),
|
||||||
|
("smtp_password".to_string(), json!(stored_password)),
|
||||||
|
("smtp_use_tls".to_string(), json!(true)),
|
||||||
|
("smtp_from_email".to_string(), json!("[email protected]")),
|
||||||
|
]);
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("gateway should build")
|
||||||
|
.with_data_state_for_tests(data);
|
||||||
|
let (gateway_url, gateway_handle) = start_server(build_router_with_state(state.clone())).await;
|
||||||
|
let client = reqwest::Client::new();
|
||||||
|
|
||||||
|
for path in [
|
||||||
|
"/api/admin/modules/status",
|
||||||
|
"/api/admin/modules/status/important_notification",
|
||||||
|
] {
|
||||||
|
let response = client
|
||||||
|
.get(format!("{gateway_url}{path}"))
|
||||||
|
.header(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")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.expect("module status request should succeed");
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
let payload: serde_json::Value = response.json().await.expect("module status should parse");
|
||||||
|
assert!(!payload.to_string().contains(stored_password));
|
||||||
|
let notification = if path == "/api/admin/modules/status" {
|
||||||
|
assert_eq!(
|
||||||
|
payload
|
||||||
|
.as_object()
|
||||||
|
.expect("module list should be an object")
|
||||||
|
.len(),
|
||||||
|
14
|
||||||
|
);
|
||||||
|
assert_eq!(payload["management_tokens"]["active"], json!(true));
|
||||||
|
&payload["important_notification"]
|
||||||
|
} else {
|
||||||
|
&payload
|
||||||
|
};
|
||||||
|
assert_eq!(notification["enabled"], json!(true));
|
||||||
|
assert_eq!(notification["config_validated"], json!(notification_ready));
|
||||||
|
assert_eq!(notification["active"], json!(notification_ready));
|
||||||
|
assert_eq!(notification["config_error"].is_null(), notification_ready);
|
||||||
|
}
|
||||||
|
gateway_handle.abort();
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
crate::important_notification::important_notification_dispatch_ready_for_item(
|
||||||
|
&state,
|
||||||
|
crate::important_notification::PROVIDER_QUOTA_ALERT_ITEM_KEY,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("SMTP errors should not abort notification readiness"),
|
||||||
|
notification_ready
|
||||||
|
);
|
||||||
|
let summary = crate::maintenance::perform_provider_quota_alert_once(&state)
|
||||||
|
.await
|
||||||
|
.expect("SMTP errors should not abort the quota alert worker");
|
||||||
|
assert_eq!(summary.failed, 0);
|
||||||
|
assert_eq!(summary.alerted, 0);
|
||||||
|
state
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn gateway_handles_admin_modules_status_with_legacy_smtp_password() {
|
||||||
|
let ciphertext =
|
||||||
|
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-smtp-password")
|
||||||
|
.expect("legacy SMTP password should encrypt");
|
||||||
|
let state = assert_admin_modules_status_with_smtp_password(&ciphertext, true, false).await;
|
||||||
|
let stored = state
|
||||||
|
.read_system_config_json_value_strong("smtp_password")
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert!(stored
|
||||||
|
.as_str()
|
||||||
|
.unwrap()
|
||||||
|
.starts_with("aether-smtp-password-v3:"));
|
||||||
|
let smtp = crate::email_delivery::read_smtp_delivery_config(&state)
|
||||||
|
.await
|
||||||
|
.expect("migrated SMTP config should load")
|
||||||
|
.expect("SMTP should be configured");
|
||||||
|
assert_eq!(smtp.password.as_deref(), Some("legacy-smtp-password"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn gateway_handles_admin_modules_status_with_invalid_smtp_password() {
|
||||||
|
let ciphertext =
|
||||||
|
encrypt_python_fernet_plaintext("unavailable-historical-key", "legacy-smtp-password")
|
||||||
|
.expect("unknown-key SMTP password should encrypt");
|
||||||
|
let state = assert_admin_modules_status_with_smtp_password(&ciphertext, false, false).await;
|
||||||
|
assert_eq!(
|
||||||
|
state
|
||||||
|
.read_system_config_json_value_strong("smtp_password")
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
Some(json!(ciphertext))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn gateway_handles_admin_modules_status_with_invalid_smtp_and_working_push() {
|
||||||
|
assert_admin_modules_status_with_smtp_password("aether-smtp-password-v3:invalid", true, true)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_returns_service_unavailable_for_admin_health_api_formats_when_readers_unavailable()
|
async fn gateway_returns_service_unavailable_for_admin_health_api_formats_when_readers_unavailable()
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -473,6 +473,23 @@ async fn gateway_handles_admin_monitoring_trace_request_locally_with_trusted_adm
|
|||||||
}),
|
}),
|
||||||
);
|
);
|
||||||
|
|
||||||
|
let mut failed_candidate = sample_candidate(
|
||||||
|
"cand-used",
|
||||||
|
"request-1",
|
||||||
|
1,
|
||||||
|
RequestCandidateStatus::Failed,
|
||||||
|
Some(101),
|
||||||
|
Some(33),
|
||||||
|
Some(502),
|
||||||
|
);
|
||||||
|
failed_candidate.error_message = Some("private upstream diagnostic".to_string());
|
||||||
|
failed_candidate.extra_data = Some(json!({
|
||||||
|
"upstream_response": {
|
||||||
|
"status_code": 502,
|
||||||
|
"headers": {"x-request-id": "upstream-diagnostic-id"},
|
||||||
|
"body": {"error": {"message": "private upstream diagnostic"}}
|
||||||
|
}
|
||||||
|
}));
|
||||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||||
sample_candidate(
|
sample_candidate(
|
||||||
"cand-unused",
|
"cand-unused",
|
||||||
@@ -483,15 +500,7 @@ async fn gateway_handles_admin_monitoring_trace_request_locally_with_trusted_adm
|
|||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
),
|
),
|
||||||
sample_candidate(
|
failed_candidate,
|
||||||
"cand-used",
|
|
||||||
"request-1",
|
|
||||||
1,
|
|
||||||
RequestCandidateStatus::Failed,
|
|
||||||
Some(101),
|
|
||||||
Some(33),
|
|
||||||
Some(502),
|
|
||||||
),
|
|
||||||
]));
|
]));
|
||||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
vec![sample_provider()],
|
vec![sample_provider()],
|
||||||
@@ -538,6 +547,37 @@ async fn gateway_handles_admin_monitoring_trace_request_locally_with_trusted_adm
|
|||||||
assert_eq!(payload["candidates"][0]["key_auth_type"], json!("api_key"));
|
assert_eq!(payload["candidates"][0]["key_auth_type"], json!("api_key"));
|
||||||
assert_eq!(payload["candidates"][0]["latency_ms"], json!(33));
|
assert_eq!(payload["candidates"][0]["latency_ms"], json!(33));
|
||||||
assert_eq!(payload["candidates"][0]["status_code"], json!(502));
|
assert_eq!(payload["candidates"][0]["status_code"], json!(502));
|
||||||
|
assert_eq!(
|
||||||
|
payload["candidates"][0]["error_message"],
|
||||||
|
"private upstream diagnostic"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
payload["candidates"][0]["extra_data"]["upstream_response"]["body"]["error"]["message"],
|
||||||
|
"private upstream diagnostic"
|
||||||
|
);
|
||||||
|
for role in [None, Some("user")] {
|
||||||
|
let mut request = reqwest::Client::new().get(format!(
|
||||||
|
"{gateway_url}/api/admin/monitoring/trace/request-1"
|
||||||
|
));
|
||||||
|
if let Some(role) = role {
|
||||||
|
request = request
|
||||||
|
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||||
|
.header(TRUSTED_ADMIN_USER_ID_HEADER, "regular-user")
|
||||||
|
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, role)
|
||||||
|
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "regular-session");
|
||||||
|
}
|
||||||
|
let response = request
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.expect("unauthorized probe should complete");
|
||||||
|
assert!(matches!(
|
||||||
|
response.status(),
|
||||||
|
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN
|
||||||
|
));
|
||||||
|
let body = response.text().await.expect("denial body should read");
|
||||||
|
assert!(!body.contains("private upstream diagnostic"));
|
||||||
|
assert!(!body.contains("upstream-diagnostic-id"));
|
||||||
|
}
|
||||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
|
|||||||
@@ -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() {
|
fn gateway_names_new_antigravity_oauth_account_from_google_userinfo_email() {
|
||||||
run_admin_oauth_test(
|
run_admin_oauth_test(
|
||||||
"gateway_names_new_antigravity_oauth_account_from_google_userinfo_email",
|
"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 = Arc::new(Mutex::new(0usize));
|
||||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||||
let upstream = Router::new().fallback(any(move |_request: Request| {
|
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()
|
let response = reqwest::Client::new()
|
||||||
.post(format!(
|
.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(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||||
.json(&json!({
|
.json(&request_body)
|
||||||
"callback_url": "http://localhost:51121/oauth2callback?code=antigravity-code-123&state=cccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccccc"
|
|
||||||
}))
|
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.expect("request should succeed");
|
.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 status = response.status();
|
||||||
let payload: Value = response.json().await.expect("json body should parse");
|
let payload: Value = response.json().await.expect("json body should parse");
|
||||||
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||||
assert_eq!(payload["provider_type"], "antigravity");
|
let account_result = if operation == "batch-import" {
|
||||||
assert_eq!(payload["email"], "[email protected]");
|
assert_eq!(payload["total"], 1);
|
||||||
assert_eq!(payload["replaced"], false);
|
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!(*token_hits.lock().expect("mutex should lock"), 1);
|
||||||
assert_eq!(*user_info_hits.lock().expect("mutex should lock"), 1);
|
assert_eq!(*user_info_hits.lock().expect("mutex should lock"), 1);
|
||||||
assert_eq!(
|
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);
|
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()
|
.as_str()
|
||||||
.expect("created key id should be returned")
|
.expect("created key id should be returned")
|
||||||
.to_string();
|
.to_string();
|
||||||
|
|||||||
@@ -70,6 +70,66 @@ async fn provider_health_summary(
|
|||||||
payload["items"][0].clone()
|
payload["items"][0].clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn admin_provider_summary_health_preserves_redacted_key_summaries() {
|
||||||
|
let endpoint = sample_endpoint(
|
||||||
|
"endpoint-chat",
|
||||||
|
"provider-openai",
|
||||||
|
"openai:chat",
|
||||||
|
"https://api.openai.example",
|
||||||
|
);
|
||||||
|
let keys = [
|
||||||
|
("key-api", "api_key", None, 0.25),
|
||||||
|
("key-oauth", "oauth", Some("{}"), 0.75),
|
||||||
|
]
|
||||||
|
.into_iter()
|
||||||
|
.map(|(key_id, auth_type, auth_config, score)| {
|
||||||
|
let mut key = sample_key(key_id, "provider-openai", "openai:chat", "test")
|
||||||
|
.with_health_fields(Some(json!({"openai:chat": {"health_score": score}})), None);
|
||||||
|
key.auth_type = auth_type.to_string();
|
||||||
|
key.encrypted_api_key = Some("summary".to_string());
|
||||||
|
key.encrypted_auth_config = auth_config.map(ToOwned::to_owned);
|
||||||
|
key
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![sample_provider("provider-openai", "openai", 10)],
|
||||||
|
vec![endpoint],
|
||||||
|
keys,
|
||||||
|
));
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("gateway should build")
|
||||||
|
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||||
|
repository,
|
||||||
|
));
|
||||||
|
|
||||||
|
for uri in [
|
||||||
|
"/api/admin/providers/summary",
|
||||||
|
"/api/admin/providers/provider-openai/summary",
|
||||||
|
] {
|
||||||
|
let response = local_admin_providers_response(&state, http::Method::GET, uri, None).await;
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
let body = axum::body::to_bytes(response.into_body(), 1024 * 1024)
|
||||||
|
.await
|
||||||
|
.expect("summary body should read");
|
||||||
|
let payload: serde_json::Value =
|
||||||
|
serde_json::from_slice(&body).expect("summary should parse");
|
||||||
|
let summary = if uri == "/api/admin/providers/summary" {
|
||||||
|
&payload["items"][0]
|
||||||
|
} else {
|
||||||
|
&payload
|
||||||
|
};
|
||||||
|
|
||||||
|
assert_eq!(summary["total_keys"], 2);
|
||||||
|
assert_eq!(summary["active_keys"], 2);
|
||||||
|
assert_eq!(summary["endpoint_health_details"][0]["total_keys"], 2);
|
||||||
|
assert_eq!(summary["endpoint_health_details"][0]["active_keys"], 2);
|
||||||
|
assert_eq!(summary["endpoint_health_details"][0]["health_score"], 0.5);
|
||||||
|
assert_eq!(summary["avg_health_score"], 0.5);
|
||||||
|
assert_eq!(summary["unhealthy_endpoints"], 0);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn admin_provider_summary_health_ignores_disabled_keys() {
|
async fn admin_provider_summary_health_ignores_disabled_keys() {
|
||||||
let endpoint = sample_endpoint(
|
let endpoint = sample_endpoint(
|
||||||
|
|||||||
@@ -2366,6 +2366,265 @@ async fn gateway_handles_admin_usage_detail_with_ref_backed_bodies() {
|
|||||||
upstream_handle.abort();
|
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]
|
#[tokio::test]
|
||||||
async fn gateway_resolves_admin_usage_detail_when_inline_state_has_body_ref() {
|
async fn gateway_resolves_admin_usage_detail_when_inline_state_has_body_ref() {
|
||||||
let (_upstream_url, upstream_hits, upstream_handle) =
|
let (_upstream_url, upstream_hits, upstream_handle) =
|
||||||
|
|||||||
@@ -221,13 +221,17 @@ async fn gateway_handles_admin_video_tasks_list_locally_with_trusted_admin_princ
|
|||||||
assert_eq!(payload["pages"], json!(1));
|
assert_eq!(payload["pages"], json!(1));
|
||||||
assert_eq!(payload["items"].as_array().map(Vec::len), Some(1));
|
assert_eq!(payload["items"].as_array().map(Vec::len), Some(1));
|
||||||
assert_eq!(payload["items"][0]["id"], "task-completed");
|
assert_eq!(payload["items"][0]["id"], "task-completed");
|
||||||
// Video-task persistence intentionally drops user-facing PII. The admin
|
assert_eq!(payload["items"][0]["username"], "alice");
|
||||||
// projection must therefore use the privacy-safe fallback when no separate
|
|
||||||
// user snapshot is joined.
|
|
||||||
assert_eq!(payload["items"][0]["username"], "Unknown");
|
|
||||||
assert_eq!(payload["items"][0]["provider_name"], "OpenAI");
|
assert_eq!(payload["items"][0]["provider_name"], "OpenAI");
|
||||||
assert_eq!(payload["items"][0]["status"], "completed");
|
assert_eq!(payload["items"][0]["status"], "completed");
|
||||||
assert!(payload["items"][0]["prompt"].is_null());
|
assert_eq!(
|
||||||
|
payload["items"][0]["prompt"],
|
||||||
|
format!("{}...", "x".repeat(100))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
payload["items"][0]["video_url"],
|
||||||
|
"https://8.8.8.8/task-completed.mp4"
|
||||||
|
);
|
||||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
@@ -393,7 +397,9 @@ async fn gateway_handles_admin_video_task_detail_locally_with_trusted_admin_prin
|
|||||||
assert_eq!(response.status(), StatusCode::OK);
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||||
assert_eq!(payload["id"], "task-detail");
|
assert_eq!(payload["id"], "task-detail");
|
||||||
assert_eq!(payload["username"], "Unknown");
|
assert_eq!(payload["prompt"], "detail prompt");
|
||||||
|
assert_eq!(payload["video_url"], "https://8.8.8.8/task-detail.mp4");
|
||||||
|
assert_eq!(payload["username"], "charlie");
|
||||||
assert_eq!(payload["provider_name"], "OpenAI");
|
assert_eq!(payload["provider_name"], "OpenAI");
|
||||||
assert_eq!(payload["endpoint"]["id"], "endpoint-1");
|
assert_eq!(payload["endpoint"]["id"], "endpoint-1");
|
||||||
assert_eq!(payload["endpoint"]["api_format"], "openai:video");
|
assert_eq!(payload["endpoint"]["api_format"], "openai:video");
|
||||||
@@ -734,7 +740,7 @@ async fn local_admin_video_task_cancel_attaches_explicit_audit() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstream() {
|
async fn gateway_redirects_persisted_openai_video_url_without_forwarding_admin_request() {
|
||||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||||
let upstream = Router::new().route(
|
let upstream = Router::new().route(
|
||||||
@@ -762,7 +768,10 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstrea
|
|||||||
))
|
))
|
||||||
.await
|
.await
|
||||||
.expect("task should upsert");
|
.expect("task should upsert");
|
||||||
assert_eq!(stored.video_url, None);
|
assert_eq!(
|
||||||
|
stored.video_url.as_deref(),
|
||||||
|
Some("https://8.8.8.8/task-redirect.mp4")
|
||||||
|
);
|
||||||
|
|
||||||
let (_upstream_url, upstream_handle) = start_server(upstream).await;
|
let (_upstream_url, upstream_handle) = start_server(upstream).await;
|
||||||
let gateway = build_router_with_state(
|
let gateway = build_router_with_state(
|
||||||
@@ -788,7 +797,14 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstrea
|
|||||||
.await
|
.await
|
||||||
.expect("request should succeed");
|
.expect("request should succeed");
|
||||||
|
|
||||||
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT);
|
||||||
|
assert_eq!(
|
||||||
|
response
|
||||||
|
.headers()
|
||||||
|
.get(http::header::LOCATION)
|
||||||
|
.and_then(|value| value.to_str().ok()),
|
||||||
|
stored.video_url.as_deref()
|
||||||
|
);
|
||||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
@@ -796,22 +812,22 @@ async fn gateway_does_not_redirect_sanitized_openai_video_url_or_forward_upstrea
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn local_admin_video_task_video_is_unavailable_after_openai_url_sanitization() {
|
async fn local_admin_video_task_download_preserves_signed_url_and_attaches_audit() {
|
||||||
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
||||||
let stored = repository
|
let mut task = sample_admin_video_task(
|
||||||
.upsert(sample_admin_video_task(
|
"task-video-audit",
|
||||||
"task-video-audit",
|
VideoTaskStatus::Completed,
|
||||||
VideoTaskStatus::Completed,
|
1_710_000_550,
|
||||||
1_710_000_550,
|
"user-5",
|
||||||
"user-5",
|
"frank",
|
||||||
"frank",
|
"provider-openai",
|
||||||
"provider-openai",
|
"gpt-video",
|
||||||
"gpt-video",
|
"video audit prompt",
|
||||||
"video audit prompt",
|
);
|
||||||
))
|
task.video_url =
|
||||||
.await
|
Some("https://8.8.8.8/video.mp4?signature=a%2Fb%2Bc%3D&part=2&part=1".to_string());
|
||||||
.expect("task should upsert");
|
let stored = repository.upsert(task).await.expect("task should upsert");
|
||||||
assert_eq!(stored.video_url, None);
|
assert_eq!(stored.prompt.as_deref(), Some("video audit prompt"));
|
||||||
|
|
||||||
let state = AppState::new()
|
let state = AppState::new()
|
||||||
.expect("gateway state should build")
|
.expect("gateway state should build")
|
||||||
@@ -825,8 +841,15 @@ async fn local_admin_video_task_video_is_unavailable_after_openai_url_sanitizati
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT);
|
||||||
assert!(response.extensions().get::<AdminAuditEvent>().is_none());
|
assert_eq!(
|
||||||
|
response
|
||||||
|
.headers()
|
||||||
|
.get(http::header::LOCATION)
|
||||||
|
.and_then(|value| value.to_str().ok()),
|
||||||
|
stored.video_url.as_deref()
|
||||||
|
);
|
||||||
|
assert!(response.extensions().get::<AdminAuditEvent>().is_some());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -50,6 +50,8 @@ use chrono::{TimeZone, Utc};
|
|||||||
const TEST_EMAIL_VERIFICATION_TOKEN: &str =
|
const TEST_EMAIL_VERIFICATION_TOKEN: &str =
|
||||||
"test-email-verification-token-00000000000000000000000000000000";
|
"test-email-verification-token-00000000000000000000000000000000";
|
||||||
|
|
||||||
|
#[path = "public_support/auth_cookie.rs"]
|
||||||
|
mod auth_cookie;
|
||||||
#[path = "public_support/dashboard.rs"]
|
#[path = "public_support/dashboard.rs"]
|
||||||
mod dashboard;
|
mod dashboard;
|
||||||
#[path = "public_support/vscodex.rs"]
|
#[path = "public_support/vscodex.rs"]
|
||||||
|
|||||||
@@ -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();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,6 +12,7 @@ use super::{
|
|||||||
UsageReadRepository, UsageRuntimeConfig, DEVELOPMENT_ENCRYPTION_KEY, TRACE_ID_HEADER,
|
UsageReadRepository, UsageRuntimeConfig, DEVELOPMENT_ENCRYPTION_KEY, TRACE_ID_HEADER,
|
||||||
};
|
};
|
||||||
use crate::constants::LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER;
|
use crate::constants::LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER;
|
||||||
|
use aether_data_contracts::repository::usage::UsageBodyCaptureState;
|
||||||
|
|
||||||
fn deep_nested_metadata(levels: usize) -> serde_json::Value {
|
fn deep_nested_metadata(levels: usize) -> serde_json::Value {
|
||||||
let mut current = json!({"leaf": "value"});
|
let mut current = json!({"leaf": "value"});
|
||||||
@@ -84,6 +85,58 @@ where
|
|||||||
stored.expect("usage should be present once the expected status is observed")
|
stored.expect("usage should be present once the expected status is observed")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn load_admin_usage_capture_detail(
|
||||||
|
state: &crate::AppState,
|
||||||
|
usage_id: &str,
|
||||||
|
include_bodies: bool,
|
||||||
|
) -> serde_json::Value {
|
||||||
|
use crate::admin_api::{maybe_build_local_admin_response, AdminRouteRequest};
|
||||||
|
use crate::constants::{
|
||||||
|
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
|
||||||
|
TRUSTED_ADMIN_USER_ROLE_HEADER,
|
||||||
|
};
|
||||||
|
use crate::control::resolve_public_request_context;
|
||||||
|
use http_body_util::BodyExt;
|
||||||
|
|
||||||
|
let mut headers = http::HeaderMap::new();
|
||||||
|
for (name, value) in [
|
||||||
|
(GATEWAY_HEADER, "rust-phase3b"),
|
||||||
|
(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user"),
|
||||||
|
(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin"),
|
||||||
|
(TRUSTED_ADMIN_SESSION_ID_HEADER, "admin-session"),
|
||||||
|
] {
|
||||||
|
headers.insert(name, HeaderValue::from_static(value));
|
||||||
|
}
|
||||||
|
let uri = format!("/api/admin/usage/{usage_id}?include_bodies={include_bodies}")
|
||||||
|
.parse()
|
||||||
|
.unwrap();
|
||||||
|
let context = resolve_public_request_context(
|
||||||
|
state,
|
||||||
|
&http::Method::GET,
|
||||||
|
&uri,
|
||||||
|
&headers,
|
||||||
|
"usage-full-detail",
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let response = maybe_build_local_admin_response(AdminRouteRequest::new(
|
||||||
|
state,
|
||||||
|
&context,
|
||||||
|
&"127.0.0.1:12345".parse().unwrap(),
|
||||||
|
&headers,
|
||||||
|
None,
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
assert!(response
|
||||||
|
.extensions()
|
||||||
|
.get::<crate::audit::AdminAuditEvent>()
|
||||||
|
.is_some());
|
||||||
|
serde_json::from_slice(&response.into_body().collect().await.unwrap().to_bytes()).unwrap()
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn gateway_handles_local_openai_chat_sync_report_with_local_reporting_when_usage_runtime_enabled() {
|
fn gateway_handles_local_openai_chat_sync_report_with_local_reporting_when_usage_runtime_enabled() {
|
||||||
run_async_test_on_large_stack(
|
run_async_test_on_large_stack(
|
||||||
@@ -348,7 +401,7 @@ async fn gateway_truncates_deep_request_echo_for_local_openai_chat_sync_usage_im
|
|||||||
Arc::clone(&request_candidate_repository),
|
Arc::clone(&request_candidate_repository),
|
||||||
Arc::clone(&usage_repository),
|
Arc::clone(&usage_repository),
|
||||||
DEVELOPMENT_ENCRYPTION_KEY,
|
DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
),
|
).with_system_config_values_for_tests([("request_record_level".to_string(), json!("full"))]),
|
||||||
)
|
)
|
||||||
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
||||||
enabled: true,
|
enabled: true,
|
||||||
@@ -402,10 +455,28 @@ async fn gateway_truncates_deep_request_echo_for_local_openai_chat_sync_usage_im
|
|||||||
let stored_usage = stored_usage.expect("usage should be recorded");
|
let stored_usage = stored_usage.expect("usage should be recorded");
|
||||||
assert_eq!(stored_usage.status, "completed");
|
assert_eq!(stored_usage.status, "completed");
|
||||||
assert_eq!(stored_usage.total_tokens, 5);
|
assert_eq!(stored_usage.total_tokens, 5);
|
||||||
assert!(stored_usage.request_body.is_none());
|
let request_body = stored_usage.request_body.as_ref().unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
request_body["messages"][0]["content"]
|
||||||
|
.as_str()
|
||||||
|
.unwrap()
|
||||||
|
.len(),
|
||||||
|
128 * 1024
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
request_body["metadata"]["child"]["child"]["child"]["child"]["child"]
|
||||||
|
.get("depth")
|
||||||
|
.is_some()
|
||||||
|
);
|
||||||
assert!(stored_usage.request_body_ref.is_none());
|
assert!(stored_usage.request_body_ref.is_none());
|
||||||
assert!(stored_usage.request_body_state.is_none());
|
assert_eq!(
|
||||||
assert!(stored_usage.request_headers.is_none());
|
stored_usage.request_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Inline)
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
stored_usage.request_headers.as_ref().unwrap()["authorization"],
|
||||||
|
"Bearer sk-client-openai-local-report-sync-deep"
|
||||||
|
);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
execution_runtime_handle.abort();
|
execution_runtime_handle.abort();
|
||||||
@@ -489,10 +560,10 @@ async fn gateway_ignores_legacy_max_request_body_size_for_local_openai_chat_sync
|
|||||||
Arc::clone(&usage_repository),
|
Arc::clone(&usage_repository),
|
||||||
DEVELOPMENT_ENCRYPTION_KEY,
|
DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
)
|
)
|
||||||
.with_system_config_values_for_tests([(
|
.with_system_config_values_for_tests([
|
||||||
"max_request_body_size".to_string(),
|
("max_request_body_size".to_string(), json!(128)),
|
||||||
json!(128),
|
("request_record_level".to_string(), json!("full")),
|
||||||
)]),
|
]),
|
||||||
)
|
)
|
||||||
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
||||||
enabled: true,
|
enabled: true,
|
||||||
@@ -535,12 +606,30 @@ async fn gateway_ignores_legacy_max_request_body_size_for_local_openai_chat_sync
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
assert_eq!(stored_usage.total_tokens, 5);
|
assert_eq!(stored_usage.total_tokens, 5);
|
||||||
assert!(stored_usage.request_body.is_none());
|
assert!(
|
||||||
|
stored_usage.request_body.as_ref().unwrap()["messages"][0]["content"]
|
||||||
|
.as_str()
|
||||||
|
.unwrap()
|
||||||
|
.len()
|
||||||
|
> 128
|
||||||
|
);
|
||||||
assert!(stored_usage.request_body_ref.is_none());
|
assert!(stored_usage.request_body_ref.is_none());
|
||||||
assert!(stored_usage.request_body_state.is_none());
|
assert_eq!(
|
||||||
assert!(stored_usage.provider_request_body.is_none());
|
stored_usage.request_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Inline)
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
stored_usage.provider_request_body.as_ref().unwrap()["messages"][0]["content"]
|
||||||
|
.as_str()
|
||||||
|
.unwrap()
|
||||||
|
.len()
|
||||||
|
> 128
|
||||||
|
);
|
||||||
assert!(stored_usage.provider_request_body_ref.is_none());
|
assert!(stored_usage.provider_request_body_ref.is_none());
|
||||||
assert!(stored_usage.provider_request_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.provider_request_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Inline)
|
||||||
|
);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
execution_runtime_handle.abort();
|
execution_runtime_handle.abort();
|
||||||
@@ -551,11 +640,19 @@ async fn gateway_ignores_legacy_max_request_body_size_for_local_openai_chat_sync
|
|||||||
fn gateway_strips_request_and_response_bodies_when_request_record_level_is_base() {
|
fn gateway_strips_request_and_response_bodies_when_request_record_level_is_base() {
|
||||||
run_async_test_on_large_stack(
|
run_async_test_on_large_stack(
|
||||||
"gateway_strips_request_and_response_bodies_when_request_record_level_is_base",
|
"gateway_strips_request_and_response_bodies_when_request_record_level_is_base",
|
||||||
gateway_strips_request_and_response_bodies_when_request_record_level_is_base_impl(),
|
gateway_honors_request_record_level_impl("base"),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn gateway_strips_request_and_response_bodies_when_request_record_level_is_base_impl() {
|
#[test]
|
||||||
|
fn gateway_full_request_record_level_preserves_sync_bodies_in_admin_detail() {
|
||||||
|
run_async_test_on_large_stack(
|
||||||
|
"gateway_full_request_record_level_preserves_sync_bodies_in_admin_detail",
|
||||||
|
gateway_honors_request_record_level_impl("full"),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn gateway_honors_request_record_level_impl(record_level: &str) {
|
||||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
||||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||||
|
|
||||||
@@ -630,14 +727,14 @@ async fn gateway_strips_request_and_response_bodies_when_request_record_level_is
|
|||||||
)
|
)
|
||||||
.with_system_config_values_for_tests([(
|
.with_system_config_values_for_tests([(
|
||||||
"request_record_level".to_string(),
|
"request_record_level".to_string(),
|
||||||
json!("base"),
|
json!(record_level),
|
||||||
)]),
|
)]),
|
||||||
)
|
)
|
||||||
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
||||||
enabled: true,
|
enabled: true,
|
||||||
..UsageRuntimeConfig::default()
|
..UsageRuntimeConfig::default()
|
||||||
});
|
});
|
||||||
let gateway = build_router_with_state(gateway_state);
|
let gateway = build_router_with_state(gateway_state.clone());
|
||||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||||
|
|
||||||
let response = reqwest::Client::new()
|
let response = reqwest::Client::new()
|
||||||
@@ -678,14 +775,39 @@ async fn gateway_strips_request_and_response_bodies_when_request_record_level_is
|
|||||||
assert_eq!(stored_usage.status, "completed");
|
assert_eq!(stored_usage.status, "completed");
|
||||||
assert_eq!(stored_usage.total_tokens, 5);
|
assert_eq!(stored_usage.total_tokens, 5);
|
||||||
assert_eq!(stored_usage.response_time_ms, Some(25));
|
assert_eq!(stored_usage.response_time_ms, Some(25));
|
||||||
assert!(stored_usage.request_body.is_none());
|
let detail = load_admin_usage_capture_detail(&gateway_state, &stored_usage.id, true).await;
|
||||||
assert!(stored_usage.request_body_ref.is_none());
|
let shallow = load_admin_usage_capture_detail(&gateway_state, &stored_usage.id, false).await;
|
||||||
assert!(stored_usage.provider_request_body.is_none());
|
for field in [
|
||||||
assert!(stored_usage.provider_request_body_ref.is_none());
|
"request_body",
|
||||||
assert!(stored_usage.response_body.is_none());
|
"provider_request_body",
|
||||||
assert!(stored_usage.response_body_ref.is_none());
|
"response_body",
|
||||||
assert!(stored_usage.client_response_body.is_none());
|
"client_response_body",
|
||||||
assert!(stored_usage.client_response_body_ref.is_none());
|
] {
|
||||||
|
assert!(shallow[field].is_null());
|
||||||
|
let expected_captured = record_level == "full" && field != "client_response_body";
|
||||||
|
assert_eq!(
|
||||||
|
shallow[format!("has_{field}")],
|
||||||
|
expected_captured,
|
||||||
|
"availability for {field}"
|
||||||
|
);
|
||||||
|
if expected_captured {
|
||||||
|
assert!(!detail[field].is_null(), "full should expose {field}");
|
||||||
|
} else {
|
||||||
|
assert!(
|
||||||
|
detail[field].is_null(),
|
||||||
|
"uncaptured {field} must remain absent"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if record_level == "full" {
|
||||||
|
assert_eq!(
|
||||||
|
detail["request_body"]["messages"][0]["content"],
|
||||||
|
"request body should not be persisted"
|
||||||
|
);
|
||||||
|
assert_eq!(detail["provider_request_body"]["model"], "gpt-5-upstream");
|
||||||
|
assert_eq!(detail["response_body"], body_json);
|
||||||
|
assert!(detail["client_response_body"].is_null());
|
||||||
|
}
|
||||||
|
|
||||||
let stored_candidates = request_candidate_repository
|
let stored_candidates = request_candidate_repository
|
||||||
.list_by_request_id("trace-openai-chat-local-report-sync-base-123")
|
.list_by_request_id("trace-openai-chat-local-report-sync-base-123")
|
||||||
@@ -825,10 +947,16 @@ async fn gateway_records_failed_usage_when_all_local_openai_chat_candidates_exha
|
|||||||
);
|
);
|
||||||
assert!(stored_usage.response_body.is_none());
|
assert!(stored_usage.response_body.is_none());
|
||||||
assert!(stored_usage.response_body_ref.is_none());
|
assert!(stored_usage.response_body_ref.is_none());
|
||||||
assert!(stored_usage.response_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.response_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Disabled)
|
||||||
|
);
|
||||||
assert!(stored_usage.client_response_body.is_none());
|
assert!(stored_usage.client_response_body.is_none());
|
||||||
assert!(stored_usage.client_response_body_ref.is_none());
|
assert!(stored_usage.client_response_body_ref.is_none());
|
||||||
assert!(stored_usage.client_response_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.client_response_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Disabled)
|
||||||
|
);
|
||||||
|
|
||||||
let stored_candidates = request_candidate_repository
|
let stored_candidates = request_candidate_repository
|
||||||
.list_by_request_id("trace-openai-chat-local-report-sync-failure-123")
|
.list_by_request_id("trace-openai-chat-local-report-sync-failure-123")
|
||||||
@@ -930,7 +1058,10 @@ async fn gateway_records_failed_usage_when_sync_runtime_transport_is_unavailable
|
|||||||
assert_eq!(stored_usage.status_code, Some(503));
|
assert_eq!(stored_usage.status_code, Some(503));
|
||||||
assert!(stored_usage.response_body.is_none());
|
assert!(stored_usage.response_body.is_none());
|
||||||
assert!(stored_usage.response_body_ref.is_none());
|
assert!(stored_usage.response_body_ref.is_none());
|
||||||
assert!(stored_usage.response_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.response_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Disabled)
|
||||||
|
);
|
||||||
|
|
||||||
let stored_candidates = request_candidate_repository
|
let stored_candidates = request_candidate_repository
|
||||||
.list_by_request_id("trace-openai-chat-local-transport-unavailable-123")
|
.list_by_request_id("trace-openai-chat-local-transport-unavailable-123")
|
||||||
@@ -1272,7 +1403,10 @@ async fn gateway_records_failed_usage_for_claude_runtime_miss_without_execution_
|
|||||||
);
|
);
|
||||||
assert!(stored_usage.client_response_body.is_none());
|
assert!(stored_usage.client_response_body.is_none());
|
||||||
assert!(stored_usage.client_response_body_ref.is_none());
|
assert!(stored_usage.client_response_body_ref.is_none());
|
||||||
assert!(stored_usage.client_response_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.client_response_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Disabled)
|
||||||
|
);
|
||||||
assert!(stored_usage.error_message.is_none());
|
assert!(stored_usage.error_message.is_none());
|
||||||
|
|
||||||
let stored_candidates = request_candidate_repository
|
let stored_candidates = request_candidate_repository
|
||||||
@@ -1296,11 +1430,20 @@ fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usa
|
|||||||
{
|
{
|
||||||
run_async_test_on_large_stack(
|
run_async_test_on_large_stack(
|
||||||
"gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled",
|
"gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled",
|
||||||
gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl(),
|
gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl("basic"),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn gateway_full_request_record_level_preserves_stream_bodies_in_admin_detail() {
|
||||||
|
run_async_test_on_large_stack(
|
||||||
|
"gateway_full_request_record_level_preserves_stream_bodies_in_admin_detail",
|
||||||
|
gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl("full"),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl(
|
async fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_when_usage_runtime_enabled_impl(
|
||||||
|
record_level: &str,
|
||||||
) {
|
) {
|
||||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
||||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||||
@@ -1406,13 +1549,13 @@ async fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_wh
|
|||||||
Arc::clone(&request_candidate_repository),
|
Arc::clone(&request_candidate_repository),
|
||||||
Arc::clone(&usage_repository),
|
Arc::clone(&usage_repository),
|
||||||
DEVELOPMENT_ENCRYPTION_KEY,
|
DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
),
|
).with_system_config_values_for_tests([("request_record_level".to_string(), json!(record_level))]),
|
||||||
)
|
)
|
||||||
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
||||||
enabled: true,
|
enabled: true,
|
||||||
..UsageRuntimeConfig::default()
|
..UsageRuntimeConfig::default()
|
||||||
});
|
});
|
||||||
let gateway = build_router_with_state(gateway_state);
|
let gateway = build_router_with_state(gateway_state.clone());
|
||||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||||
|
|
||||||
let response = reqwest::Client::new()
|
let response = reqwest::Client::new()
|
||||||
@@ -1448,6 +1591,30 @@ async fn gateway_handles_local_openai_chat_stream_report_with_local_reporting_wh
|
|||||||
assert!(stored_usage.response_time_ms >= stored_usage.first_byte_time_ms);
|
assert!(stored_usage.response_time_ms >= stored_usage.first_byte_time_ms);
|
||||||
assert!(stored_usage.is_stream);
|
assert!(stored_usage.is_stream);
|
||||||
|
|
||||||
|
let detail = load_admin_usage_capture_detail(&gateway_state, &stored_usage.id, true).await;
|
||||||
|
for field in [
|
||||||
|
"request_body",
|
||||||
|
"provider_request_body",
|
||||||
|
"response_body",
|
||||||
|
"client_response_body",
|
||||||
|
] {
|
||||||
|
if record_level == "full" {
|
||||||
|
assert!(
|
||||||
|
!detail[field].is_null(),
|
||||||
|
"full stream should expose {field}"
|
||||||
|
);
|
||||||
|
} else {
|
||||||
|
assert!(
|
||||||
|
detail[field].is_null(),
|
||||||
|
"basic stream must not persist {field}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if record_level == "full" {
|
||||||
|
assert!(detail["response_body"].to_string().contains("hello"));
|
||||||
|
assert!(detail["client_response_body"].to_string().contains("hello"));
|
||||||
|
}
|
||||||
|
|
||||||
let stored_candidates = request_candidate_repository
|
let stored_candidates = request_candidate_repository
|
||||||
.list_by_request_id("trace-openai-chat-local-report-stream-123")
|
.list_by_request_id("trace-openai-chat-local-report-stream-123")
|
||||||
.await
|
.await
|
||||||
@@ -1585,10 +1752,10 @@ async fn gateway_ignores_legacy_max_response_body_size_for_stream_usage_impl() {
|
|||||||
Arc::clone(&usage_repository),
|
Arc::clone(&usage_repository),
|
||||||
DEVELOPMENT_ENCRYPTION_KEY,
|
DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
)
|
)
|
||||||
.with_system_config_values_for_tests([(
|
.with_system_config_values_for_tests([
|
||||||
"max_response_body_size".to_string(),
|
("max_response_body_size".to_string(), json!(128)),
|
||||||
json!(128),
|
("request_record_level".to_string(), json!("full")),
|
||||||
)]),
|
]),
|
||||||
)
|
)
|
||||||
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
||||||
enabled: true,
|
enabled: true,
|
||||||
@@ -1624,12 +1791,34 @@ async fn gateway_ignores_legacy_max_response_body_size_for_stream_usage_impl() {
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
assert_eq!(stored_usage.total_tokens, 6);
|
assert_eq!(stored_usage.total_tokens, 6);
|
||||||
assert!(stored_usage.response_body.is_none());
|
assert!(
|
||||||
|
stored_usage
|
||||||
|
.response_body
|
||||||
|
.as_ref()
|
||||||
|
.unwrap()
|
||||||
|
.to_string()
|
||||||
|
.len()
|
||||||
|
> 128
|
||||||
|
);
|
||||||
assert!(stored_usage.response_body_ref.is_none());
|
assert!(stored_usage.response_body_ref.is_none());
|
||||||
assert!(stored_usage.response_body_state.is_none());
|
assert_eq!(
|
||||||
assert!(stored_usage.client_response_body.is_none());
|
stored_usage.response_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Inline)
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
stored_usage
|
||||||
|
.client_response_body
|
||||||
|
.as_ref()
|
||||||
|
.unwrap()
|
||||||
|
.to_string()
|
||||||
|
.len()
|
||||||
|
> 128
|
||||||
|
);
|
||||||
assert!(stored_usage.client_response_body_ref.is_none());
|
assert!(stored_usage.client_response_body_ref.is_none());
|
||||||
assert!(stored_usage.client_response_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.client_response_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Inline)
|
||||||
|
);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
execution_runtime_handle.abort();
|
execution_runtime_handle.abort();
|
||||||
@@ -1903,10 +2092,16 @@ async fn gateway_records_failed_usage_when_all_local_claude_cli_candidates_are_s
|
|||||||
Some("all_candidates_skipped")
|
Some("all_candidates_skipped")
|
||||||
);
|
);
|
||||||
assert!(stored_usage.error_message.is_none());
|
assert!(stored_usage.error_message.is_none());
|
||||||
assert!(stored_usage.request_headers.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.request_headers.as_ref().unwrap()["authorization"],
|
||||||
|
"Bearer sk-client-claude-cli-usage-local-miss"
|
||||||
|
);
|
||||||
assert!(stored_usage.request_body.is_none());
|
assert!(stored_usage.request_body.is_none());
|
||||||
assert!(stored_usage.request_body_ref.is_none());
|
assert!(stored_usage.request_body_ref.is_none());
|
||||||
assert!(stored_usage.request_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.request_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Disabled)
|
||||||
|
);
|
||||||
assert!(stored_usage.provider_request_body.is_none());
|
assert!(stored_usage.provider_request_body.is_none());
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
stored_usage
|
stored_usage
|
||||||
@@ -2169,7 +2364,10 @@ fn gateway_keeps_failed_usage_request_capture_lightweight_for_large_local_claude
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
assert_eq!(stored_usage.status, "failed");
|
assert_eq!(stored_usage.status, "failed");
|
||||||
assert!(stored_usage.request_body_state.is_none());
|
assert_eq!(
|
||||||
|
stored_usage.request_body_state,
|
||||||
|
Some(UsageBodyCaptureState::Disabled)
|
||||||
|
);
|
||||||
assert!(stored_usage.request_body.is_none());
|
assert!(stored_usage.request_body.is_none());
|
||||||
assert!(stored_usage.request_body_ref.is_none());
|
assert!(stored_usage.request_body_ref.is_none());
|
||||||
assert!(stored_usage.provider_request_body.is_none());
|
assert!(stored_usage.provider_request_body.is_none());
|
||||||
|
|||||||
@@ -213,6 +213,13 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
|
|||||||
};
|
};
|
||||||
|
|
||||||
assert_eq!(stored.status, VideoTaskStatus::Processing);
|
assert_eq!(stored.status, VideoTaskStatus::Processing);
|
||||||
|
assert_eq!(stored.prompt.as_deref(), Some("hello"));
|
||||||
|
assert_eq!(stored.username.as_deref(), Some("video-user"));
|
||||||
|
assert_eq!(stored.api_key_name.as_deref(), Some("video-key"));
|
||||||
|
assert_eq!(stored.duration_seconds, Some(4));
|
||||||
|
assert_eq!(stored.resolution.as_deref(), Some("720p"));
|
||||||
|
assert_eq!(stored.aspect_ratio.as_deref(), Some("16:9"));
|
||||||
|
assert_eq!(stored.size.as_deref(), Some("1280x720"));
|
||||||
assert_eq!(stored.progress_percent, 37);
|
assert_eq!(stored.progress_percent, 37);
|
||||||
assert_eq!(stored.poll_count, 1);
|
assert_eq!(stored.poll_count, 1);
|
||||||
assert!(
|
assert!(
|
||||||
|
|||||||
@@ -32,6 +32,7 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
|
|||||||
struct SeenExecutionRuntimeStreamRequest {
|
struct SeenExecutionRuntimeStreamRequest {
|
||||||
method: String,
|
method: String,
|
||||||
url: String,
|
url: String,
|
||||||
|
headers: serde_json::Value,
|
||||||
}
|
}
|
||||||
|
|
||||||
fn hash_api_key(value: &str) -> String {
|
fn hash_api_key(value: &str) -> String {
|
||||||
@@ -159,6 +160,7 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
|
|||||||
.and_then(|value| value.as_str())
|
.and_then(|value| value.as_str())
|
||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
.to_string(),
|
.to_string(),
|
||||||
|
headers: payload.get("headers").cloned().unwrap_or_else(|| json!({})),
|
||||||
});
|
});
|
||||||
|
|
||||||
let frames = [
|
let frames = [
|
||||||
@@ -252,7 +254,10 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
|
|||||||
updated_at_unix_secs: 456,
|
updated_at_unix_secs: 456,
|
||||||
error_code: None,
|
error_code: None,
|
||||||
error_message: None,
|
error_message: None,
|
||||||
video_url: Some("https://cdn.example.com/video-content.mp4".to_string()),
|
video_url: Some(
|
||||||
|
"https://cdn.example.com/video-content.mp4?signature=a%2Fb%2Bc%3D&part=2&part=1"
|
||||||
|
.to_string(),
|
||||||
|
),
|
||||||
request_metadata: None,
|
request_metadata: None,
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
@@ -358,8 +363,9 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
|
|||||||
assert_eq!(seen_stream_request.method, "GET");
|
assert_eq!(seen_stream_request.method, "GET");
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
seen_stream_request.url,
|
seen_stream_request.url,
|
||||||
"https://api.openai.example/v1/videos/ext-video-content-followup-123/content"
|
"https://cdn.example.com/video-content.mp4?signature=a%2Fb%2Bc%3D&part=2&part=1"
|
||||||
);
|
);
|
||||||
|
assert!(seen_stream_request.headers.get("authorization").is_none());
|
||||||
assert_eq!(*decision_stream_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*decision_stream_hits.lock().expect("mutex should lock"), 0);
|
||||||
assert_eq!(*execute_stream_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*execute_stream_hits.lock().expect("mutex should lock"), 0);
|
||||||
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
||||||
|
|||||||
@@ -0,0 +1,186 @@
|
|||||||
|
use std::collections::VecDeque;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use bytes::{Bytes, BytesMut};
|
||||||
|
use parking_lot::Mutex;
|
||||||
|
use tokio::sync::Notify;
|
||||||
|
|
||||||
|
const CHUNK_BYTES: usize = 32 * 1024;
|
||||||
|
|
||||||
|
#[derive(Debug)]
|
||||||
|
pub enum LocalBodyEvent {
|
||||||
|
Chunk(Bytes),
|
||||||
|
End,
|
||||||
|
Error(String),
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Default)]
|
||||||
|
struct BufferState {
|
||||||
|
chunks: VecDeque<BytesMut>,
|
||||||
|
bytes: usize,
|
||||||
|
terminal: Option<Result<(), String>>,
|
||||||
|
receiver_taken: bool,
|
||||||
|
receiver_closed: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) struct ResponseBuffer {
|
||||||
|
state: Mutex<BufferState>,
|
||||||
|
notify: Notify,
|
||||||
|
capacity: usize,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ResponseBuffer {
|
||||||
|
pub(super) fn new(capacity: usize) -> Arc<Self> {
|
||||||
|
Arc::new(Self {
|
||||||
|
state: Mutex::new(BufferState::default()),
|
||||||
|
notify: Notify::new(),
|
||||||
|
capacity,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn take_receiver(self: &Arc<Self>) -> Option<BodyReceiver> {
|
||||||
|
let mut state = self.state.lock();
|
||||||
|
if state.receiver_taken {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
state.receiver_taken = true;
|
||||||
|
Some(BodyReceiver {
|
||||||
|
buffer: Arc::clone(self),
|
||||||
|
finished: false,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn push(&self, mut payload: Bytes) -> bool {
|
||||||
|
let mut state = self.state.lock();
|
||||||
|
if state.terminal.is_some()
|
||||||
|
|| state.receiver_closed
|
||||||
|
|| payload.len() > self.capacity.saturating_sub(state.bytes)
|
||||||
|
{
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
state.bytes += payload.len();
|
||||||
|
while !payload.is_empty() {
|
||||||
|
if let Some(tail) = state
|
||||||
|
.chunks
|
||||||
|
.back_mut()
|
||||||
|
.filter(|chunk| chunk.len() < CHUNK_BYTES)
|
||||||
|
{
|
||||||
|
let count = payload.len().min(CHUNK_BYTES - tail.len());
|
||||||
|
tail.extend_from_slice(&payload.split_to(count));
|
||||||
|
} else {
|
||||||
|
let count = payload.len().min(CHUNK_BYTES);
|
||||||
|
let chunk = payload.split_to(count);
|
||||||
|
state.chunks.push_back(
|
||||||
|
chunk
|
||||||
|
.try_into_mut()
|
||||||
|
.unwrap_or_else(|chunk| BytesMut::from(chunk.as_ref())),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
drop(state);
|
||||||
|
self.notify.notify_waiters();
|
||||||
|
true
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn finish(&self, result: Result<(), String>) {
|
||||||
|
let mut state = self.state.lock();
|
||||||
|
if state.terminal.is_none() {
|
||||||
|
state.terminal = Some(result);
|
||||||
|
}
|
||||||
|
drop(state);
|
||||||
|
self.notify.notify_waiters();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) struct BodyReceiver {
|
||||||
|
buffer: Arc<ResponseBuffer>,
|
||||||
|
finished: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl BodyReceiver {
|
||||||
|
pub(super) async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||||
|
if self.finished {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
loop {
|
||||||
|
let notified = self.buffer.notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
|
{
|
||||||
|
let mut state = self.buffer.state.lock();
|
||||||
|
if let Some(chunk) = state.chunks.pop_front() {
|
||||||
|
state.bytes -= chunk.len();
|
||||||
|
return Some(LocalBodyEvent::Chunk(chunk.freeze()));
|
||||||
|
}
|
||||||
|
if let Some(terminal) = state.terminal.take() {
|
||||||
|
self.finished = true;
|
||||||
|
state.receiver_closed = true;
|
||||||
|
return Some(match terminal {
|
||||||
|
Ok(()) => LocalBodyEvent::End,
|
||||||
|
Err(error) => LocalBodyEvent::Error(error),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
notified.await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for BodyReceiver {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
let mut state = self.buffer.state.lock();
|
||||||
|
state.receiver_closed = true;
|
||||||
|
state.chunks.clear();
|
||||||
|
state.bytes = 0;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn error_survives_a_full_buffer() {
|
||||||
|
let buffer = ResponseBuffer::new(CHUNK_BYTES);
|
||||||
|
let mut receiver = buffer.take_receiver().unwrap();
|
||||||
|
assert!(buffer.push(Bytes::from(vec![b'x'; CHUNK_BYTES])));
|
||||||
|
buffer.finish(Err("proxy disconnected".into()));
|
||||||
|
assert!(matches!(
|
||||||
|
receiver.recv().await,
|
||||||
|
Some(LocalBodyEvent::Chunk(_))
|
||||||
|
));
|
||||||
|
assert!(
|
||||||
|
matches!(receiver.recv().await, Some(LocalBodyEvent::Error(error)) if error == "proxy disconnected")
|
||||||
|
);
|
||||||
|
assert!(receiver.recv().await.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn small_frames_are_coalesced_within_the_byte_budget() {
|
||||||
|
let buffer = ResponseBuffer::new(4096);
|
||||||
|
let mut receiver = buffer.take_receiver().unwrap();
|
||||||
|
for _ in 0..4096 {
|
||||||
|
assert!(buffer.push(Bytes::from_static(b"x")));
|
||||||
|
}
|
||||||
|
assert!(!buffer.push(Bytes::from_static(b"x")));
|
||||||
|
assert_eq!(buffer.state.lock().chunks.len(), 1);
|
||||||
|
buffer.finish(Ok(()));
|
||||||
|
assert!(
|
||||||
|
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 4096)
|
||||||
|
);
|
||||||
|
assert!(matches!(receiver.recv().await, Some(LocalBodyEvent::End)));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn terminal_wakes_an_empty_receiver_and_is_not_overwritten() {
|
||||||
|
let buffer = ResponseBuffer::new(1024);
|
||||||
|
let mut receiver = buffer.take_receiver().unwrap();
|
||||||
|
let task = tokio::spawn(async move { receiver.recv().await });
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
buffer.finish(Err("cancelled".into()));
|
||||||
|
buffer.finish(Ok(()));
|
||||||
|
assert!(
|
||||||
|
matches!(task.await.unwrap(), Some(LocalBodyEvent::Error(error)) if error == "cancelled")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,250 @@
|
|||||||
|
use super::*;
|
||||||
|
|
||||||
|
async fn fixture(
|
||||||
|
window: u32,
|
||||||
|
capacity: usize,
|
||||||
|
) -> (
|
||||||
|
Arc<HubRouter>,
|
||||||
|
Arc<ProxyConn>,
|
||||||
|
Arc<LocalStream>,
|
||||||
|
aether_runtime::BoundedQueueReceiver<Message>,
|
||||||
|
) {
|
||||||
|
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||||
|
let (sender, receiver) = bounded_queue(capacity);
|
||||||
|
let (close_tx, _) = watch::channel(false);
|
||||||
|
let connection = Arc::new(
|
||||||
|
ProxyConn::new(
|
||||||
|
99,
|
||||||
|
"flow-test".into(),
|
||||||
|
"flow-test".into(),
|
||||||
|
sender,
|
||||||
|
close_tx,
|
||||||
|
16,
|
||||||
|
3,
|
||||||
|
)
|
||||||
|
.with_settings(protocol::SettingsPayload {
|
||||||
|
initial_stream_window_bytes: window,
|
||||||
|
min_window_update_bytes: (window / 4).max(1),
|
||||||
|
drain_deadline_ms: 1000,
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
hub.register_proxy(Arc::clone(&connection));
|
||||||
|
let stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||||
|
(hub, connection, stream, receiver)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn meta() -> protocol::RequestMeta {
|
||||||
|
protocol::RequestMeta {
|
||||||
|
provider_id: None,
|
||||||
|
endpoint_id: None,
|
||||||
|
key_id: None,
|
||||||
|
method: "GET".into(),
|
||||||
|
url: "https://example.com".into(),
|
||||||
|
headers: HashMap::new(),
|
||||||
|
stream: true,
|
||||||
|
request_timeout_ms: None,
|
||||||
|
stream_first_byte_timeout_ms: None,
|
||||||
|
timeout: 30,
|
||||||
|
follow_redirects: None,
|
||||||
|
http1_only: false,
|
||||||
|
transport_profile: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn headers(hub: &Arc<HubRouter>, stream: &LocalStream) {
|
||||||
|
let payload = serde_json::to_vec(&protocol::ResponseMeta {
|
||||||
|
status: 200,
|
||||||
|
headers: vec![],
|
||||||
|
})
|
||||||
|
.unwrap();
|
||||||
|
let mut frame = protocol::encode_frame(
|
||||||
|
stream.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_HEADERS,
|
||||||
|
0,
|
||||||
|
&payload,
|
||||||
|
);
|
||||||
|
hub.handle_proxy_frame(stream.proxy_conn_id, &mut frame)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn window_credit_is_retried_after_queue_pressure_and_cancelled_receive() {
|
||||||
|
let (hub, _, stream, mut outbound) = fixture(128, 1).await;
|
||||||
|
headers(&hub, &stream).await;
|
||||||
|
assert!(stream.push_body_chunk(Bytes::from(vec![b'x'; 64])));
|
||||||
|
let mut receiver = stream.take_body_receiver().unwrap();
|
||||||
|
assert!(
|
||||||
|
tokio::time::timeout(Duration::from_millis(10), receiver.recv())
|
||||||
|
.await
|
||||||
|
.is_err()
|
||||||
|
);
|
||||||
|
assert_eq!(*stream.response_consumed_since_update.lock(), 64);
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
let event = tokio::time::timeout(Duration::from_secs(1), receiver.recv())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert!(matches!(event, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 64));
|
||||||
|
assert_eq!(*stream.response_consumed_since_update.lock(), 0);
|
||||||
|
let Message::Binary(data) = outbound.recv().await.unwrap() else {
|
||||||
|
panic!("expected binary update")
|
||||||
|
};
|
||||||
|
let frame = aether_contracts::tunnel::Frame::decode(data).unwrap();
|
||||||
|
let update: protocol::WindowUpdatePayload = serde_json::from_slice(&frame.payload).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
frame.msg_type,
|
||||||
|
aether_contracts::tunnel::MsgType::WindowUpdate
|
||||||
|
);
|
||||||
|
assert_eq!(update.delta_bytes, 64);
|
||||||
|
hub.cancel_local_stream(stream.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn response_credit_is_not_returned_until_consumed() {
|
||||||
|
let (hub, _, stream, mut outbound) = fixture(128, 4).await;
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
let mut body = protocol::encode_frame(
|
||||||
|
stream.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_BODY,
|
||||||
|
0,
|
||||||
|
&[b'x'; 128],
|
||||||
|
);
|
||||||
|
hub.handle_proxy_frame(99, &mut body).await;
|
||||||
|
assert!(outbound.try_recv().is_err());
|
||||||
|
let mut receiver = stream.take_body_receiver().unwrap();
|
||||||
|
assert!(
|
||||||
|
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 128)
|
||||||
|
);
|
||||||
|
assert!(outbound.try_recv().is_ok());
|
||||||
|
hub.cancel_local_stream(stream.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cancelled_stream_open_releases_slot_without_resetting_connection() {
|
||||||
|
let (hub, connection, first_stream, mut outbound) = fixture(128, 1).await;
|
||||||
|
let opening_hub = Arc::clone(&hub);
|
||||||
|
let opening =
|
||||||
|
tokio::spawn(async move { opening_hub.open_local_stream("flow-test", &meta()).await });
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while hub.local_streams.len() != 2 {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
opening.abort();
|
||||||
|
assert!(matches!(opening.await, Err(error) if error.is_cancelled()));
|
||||||
|
assert_eq!(connection.stream_count.load(Ordering::Relaxed), 1);
|
||||||
|
assert_eq!(hub.local_streams.len(), 1);
|
||||||
|
assert_eq!(hub.proxy_to_local.len(), 1);
|
||||||
|
assert!(connection.is_available());
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
assert!(outbound.try_recv().is_err());
|
||||||
|
let next_stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
hub.cancel_local_stream(first_stream.id, "test complete");
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
hub.cancel_local_stream(next_stream.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn full_response_buffer_preserves_disconnect_error() {
|
||||||
|
let (hub, connection, stream, mut outbound) = fixture(4 * 1024 * 1024, 512).await;
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
headers(&hub, &stream).await;
|
||||||
|
let mut receiver = stream.take_body_receiver().unwrap();
|
||||||
|
for _ in 0..128 {
|
||||||
|
let mut frame = protocol::encode_frame(
|
||||||
|
stream.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_BODY,
|
||||||
|
0,
|
||||||
|
&vec![b'x'; 32 * 1024],
|
||||||
|
);
|
||||||
|
hub.handle_proxy_frame(99, &mut frame).await;
|
||||||
|
}
|
||||||
|
hub.unregister_proxy(connection.id, &connection.node_id);
|
||||||
|
let mut bytes = 0;
|
||||||
|
loop {
|
||||||
|
match receiver.recv().await {
|
||||||
|
Some(LocalBodyEvent::Chunk(chunk)) => bytes += chunk.len(),
|
||||||
|
Some(LocalBodyEvent::Error(error)) => {
|
||||||
|
assert!(error.contains("disconnected"));
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
event => panic!("disconnect must not become normal EOF: {event:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
assert_eq!(bytes, 4 * 1024 * 1024);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn slow_stream_does_not_block_another_stream_on_the_same_connection() {
|
||||||
|
let (hub, _, slow, mut outbound) = fixture(128, 512).await;
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
let fast = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||||
|
assert!(slow.push_body_chunk(Bytes::from(vec![b'x'; 128])));
|
||||||
|
let mut overflowing = protocol::encode_frame(
|
||||||
|
slow.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_BODY,
|
||||||
|
0,
|
||||||
|
b"overflow",
|
||||||
|
);
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
hub.handle_proxy_frame(99, &mut overflowing).await;
|
||||||
|
headers(&hub, &fast).await;
|
||||||
|
assert_eq!(
|
||||||
|
fast.wait_headers(Duration::from_secs(1))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.status,
|
||||||
|
200
|
||||||
|
);
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.expect("slow stream must not block connection reader");
|
||||||
|
assert!(!hub.local_streams.contains_key(&slow.id));
|
||||||
|
assert!(hub.local_streams.contains_key(&fast.id));
|
||||||
|
hub.cancel_local_stream(fast.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cancelling_a_stream_wakes_request_window_waiters() {
|
||||||
|
let (_, _, stream, _) = fixture(128, 512).await;
|
||||||
|
*stream.request_window.available.lock() = 0;
|
||||||
|
let waiter = tokio::spawn({
|
||||||
|
let stream = Arc::clone(&stream);
|
||||||
|
async move {
|
||||||
|
stream
|
||||||
|
.acquire_request_window(1, Duration::from_secs(30))
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
});
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
stream.fail("cancelled");
|
||||||
|
assert!(tokio::time::timeout(Duration::from_secs(1), waiter)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap()
|
||||||
|
.is_err());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||||
|
async fn concurrent_headers_and_credit_updates_do_not_lose_notifications() {
|
||||||
|
for index in 0..256 {
|
||||||
|
let stream = Arc::new(LocalStream::new(index, "test".into(), 1, 1, 1));
|
||||||
|
let window = Arc::new(StreamFlowWindow::new(0));
|
||||||
|
let waiter = tokio::spawn({
|
||||||
|
let stream = Arc::clone(&stream);
|
||||||
|
let window = Arc::clone(&window);
|
||||||
|
async move {
|
||||||
|
stream.wait_headers(Duration::from_secs(1)).await.unwrap();
|
||||||
|
window.acquire(1, Duration::from_secs(1)).await.unwrap();
|
||||||
|
}
|
||||||
|
});
|
||||||
|
stream.set_response_headers(protocol::ResponseMeta {
|
||||||
|
status: 200,
|
||||||
|
headers: vec![],
|
||||||
|
});
|
||||||
|
window.add(1);
|
||||||
|
waiter.await.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,10 +12,11 @@ use axum::extract::ws::Message;
|
|||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use dashmap::DashMap;
|
use dashmap::DashMap;
|
||||||
use parking_lot::{Mutex, RwLock};
|
use parking_lot::{Mutex, RwLock};
|
||||||
use tokio::sync::mpsc;
|
|
||||||
use tokio::sync::{watch, Notify};
|
use tokio::sync::{watch, Notify};
|
||||||
use tracing::{debug, info, warn};
|
use tracing::{debug, info, warn};
|
||||||
|
|
||||||
|
pub use super::body::LocalBodyEvent;
|
||||||
|
use super::body::{BodyReceiver, ResponseBuffer};
|
||||||
use super::control_plane::ControlPlaneClient;
|
use super::control_plane::ControlPlaneClient;
|
||||||
use super::protocol;
|
use super::protocol;
|
||||||
|
|
||||||
@@ -29,6 +30,10 @@ const DEFAULT_DRAIN_DEADLINE_MS: u64 = 30_000;
|
|||||||
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
|
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
|
||||||
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
|
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
#[path = "flow_control_tests.rs"]
|
||||||
|
mod flow_control_tests;
|
||||||
|
|
||||||
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
||||||
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
|
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
|
||||||
.ok()
|
.ok()
|
||||||
@@ -53,11 +58,16 @@ static NODE_STATUS_QUEUE_CAPACITY: LazyLock<usize> = LazyLock::new(|| {
|
|||||||
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
|
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
|
||||||
});
|
});
|
||||||
|
|
||||||
static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
pub(super) fn local_settings() -> protocol::SettingsPayload {
|
||||||
STREAM_INITIAL_WINDOW_BYTES
|
protocol::SettingsPayload {
|
||||||
.saturating_div(4)
|
initial_stream_window_bytes: (*STREAM_INITIAL_WINDOW_BYTES)
|
||||||
.clamp(1, 1024 * 1024)
|
.min(aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u32),
|
||||||
});
|
min_window_update_bytes: STREAM_INITIAL_WINDOW_BYTES
|
||||||
|
.saturating_div(4)
|
||||||
|
.clamp(1, 1024 * 1024),
|
||||||
|
drain_deadline_ms: *DRAIN_DEADLINE_MS,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
pub enum SendStatus {
|
pub enum SendStatus {
|
||||||
@@ -91,6 +101,7 @@ impl ConnHealthState {
|
|||||||
struct StreamFlowWindow {
|
struct StreamFlowWindow {
|
||||||
available: Mutex<u64>,
|
available: Mutex<u64>,
|
||||||
notify: Notify,
|
notify: Notify,
|
||||||
|
closed: AtomicBool,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl StreamFlowWindow {
|
impl StreamFlowWindow {
|
||||||
@@ -98,6 +109,7 @@ impl StreamFlowWindow {
|
|||||||
Self {
|
Self {
|
||||||
available: Mutex::new(u64::from(initial)),
|
available: Mutex::new(u64::from(initial)),
|
||||||
notify: Notify::new(),
|
notify: Notify::new(),
|
||||||
|
closed: AtomicBool::new(false),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -109,6 +121,12 @@ impl StreamFlowWindow {
|
|||||||
let requested = bytes as u64;
|
let requested = bytes as u64;
|
||||||
let started_at = Instant::now();
|
let started_at = Instant::now();
|
||||||
loop {
|
loop {
|
||||||
|
let notified = self.notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
|
if self.closed.load(Ordering::Acquire) {
|
||||||
|
return Err(());
|
||||||
|
}
|
||||||
{
|
{
|
||||||
let mut available = self.available.lock();
|
let mut available = self.available.lock();
|
||||||
if *available >= requested {
|
if *available >= requested {
|
||||||
@@ -120,10 +138,7 @@ impl StreamFlowWindow {
|
|||||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||||
return Err(());
|
return Err(());
|
||||||
};
|
};
|
||||||
if tokio::time::timeout(remaining, self.notify.notified())
|
if tokio::time::timeout(remaining, notified).await.is_err() {
|
||||||
.await
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
return Err(());
|
return Err(());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -138,6 +153,11 @@ impl StreamFlowWindow {
|
|||||||
drop(available);
|
drop(available);
|
||||||
self.notify.notify_waiters();
|
self.notify.notify_waiters();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn close(&self) {
|
||||||
|
self.closed.store(true, Ordering::Release);
|
||||||
|
self.notify.notify_waiters();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
@@ -209,6 +229,10 @@ impl BoundedOutbound {
|
|||||||
pub fn snapshot(&self) -> QueueSnapshot {
|
pub fn snapshot(&self) -> QueueSnapshot {
|
||||||
self.tx.snapshot()
|
self.tx.snapshot()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||||
|
self.close_tx.subscribe()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub struct ProxyConn {
|
pub struct ProxyConn {
|
||||||
@@ -231,6 +255,7 @@ pub struct ProxyConn {
|
|||||||
flow_window_blocked_ms: AtomicU64,
|
flow_window_blocked_ms: AtomicU64,
|
||||||
write_latency_last_us: AtomicU64,
|
write_latency_last_us: AtomicU64,
|
||||||
write_latency_ewma_us: AtomicU64,
|
write_latency_ewma_us: AtomicU64,
|
||||||
|
settings: Mutex<protocol::SettingsPayload>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ProxyConn {
|
impl ProxyConn {
|
||||||
@@ -244,6 +269,7 @@ impl ProxyConn {
|
|||||||
protocol_version: u8,
|
protocol_version: u8,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self {
|
Self {
|
||||||
|
settings: Mutex::new(local_settings()),
|
||||||
id,
|
id,
|
||||||
node_id,
|
node_id,
|
||||||
node_name,
|
node_name,
|
||||||
@@ -271,6 +297,11 @@ impl ProxyConn {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(super) fn with_settings(mut self, settings: protocol::SettingsPayload) -> Self {
|
||||||
|
*self.settings.get_mut() = settings;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
||||||
self.node_generation = tunnel_generation;
|
self.node_generation = tunnel_generation;
|
||||||
self
|
self
|
||||||
@@ -565,13 +596,6 @@ pub struct LocalResponseHead {
|
|||||||
pub headers: Vec<(String, String)>,
|
pub headers: Vec<(String, String)>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
|
||||||
pub enum LocalBodyEvent {
|
|
||||||
Chunk(Bytes),
|
|
||||||
End,
|
|
||||||
Error(String),
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Default)]
|
#[derive(Debug, Default)]
|
||||||
struct LocalWaitState {
|
struct LocalWaitState {
|
||||||
response: Option<LocalResponseHead>,
|
response: Option<LocalResponseHead>,
|
||||||
@@ -585,10 +609,11 @@ pub struct LocalStream {
|
|||||||
proxy_stream_id: u32,
|
proxy_stream_id: u32,
|
||||||
request_window: StreamFlowWindow,
|
request_window: StreamFlowWindow,
|
||||||
response_consumed_since_update: Mutex<u64>,
|
response_consumed_since_update: Mutex<u64>,
|
||||||
|
min_window_update_bytes: u32,
|
||||||
|
response_connection: Mutex<Option<std::sync::Weak<ProxyConn>>>,
|
||||||
wait_state: Mutex<LocalWaitState>,
|
wait_state: Mutex<LocalWaitState>,
|
||||||
headers_notify: Notify,
|
headers_notify: Notify,
|
||||||
body_tx: mpsc::Sender<LocalBodyEvent>,
|
body: Arc<ResponseBuffer>,
|
||||||
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
|
|
||||||
terminal: AtomicBool,
|
terminal: AtomicBool,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -600,7 +625,6 @@ impl LocalStream {
|
|||||||
proxy_stream_id: u32,
|
proxy_stream_id: u32,
|
||||||
initial_window_bytes: u32,
|
initial_window_bytes: u32,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
let (body_tx, body_rx) = mpsc::channel(128);
|
|
||||||
Self {
|
Self {
|
||||||
id,
|
id,
|
||||||
tunnel_generation,
|
tunnel_generation,
|
||||||
@@ -608,10 +632,11 @@ impl LocalStream {
|
|||||||
proxy_stream_id,
|
proxy_stream_id,
|
||||||
request_window: StreamFlowWindow::new(initial_window_bytes),
|
request_window: StreamFlowWindow::new(initial_window_bytes),
|
||||||
response_consumed_since_update: Mutex::new(0),
|
response_consumed_since_update: Mutex::new(0),
|
||||||
|
min_window_update_bytes: (initial_window_bytes / 4).clamp(1, 1024 * 1024),
|
||||||
|
response_connection: Mutex::new(None),
|
||||||
wait_state: Mutex::new(LocalWaitState::default()),
|
wait_state: Mutex::new(LocalWaitState::default()),
|
||||||
headers_notify: Notify::new(),
|
headers_notify: Notify::new(),
|
||||||
body_tx,
|
body: ResponseBuffer::new(initial_window_bytes as usize),
|
||||||
body_rx: Mutex::new(Some(body_rx)),
|
|
||||||
terminal: AtomicBool::new(false),
|
terminal: AtomicBool::new(false),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -632,26 +657,51 @@ impl LocalStream {
|
|||||||
self.request_window.add(delta);
|
self.request_window.add(delta);
|
||||||
}
|
}
|
||||||
|
|
||||||
fn response_window_update_delta(&self, bytes: usize) -> Option<u32> {
|
async fn flush_response_credit(&self) -> Result<(), String> {
|
||||||
if bytes == 0 {
|
if self.terminal.load(Ordering::Acquire) {
|
||||||
return None;
|
return Ok(());
|
||||||
}
|
}
|
||||||
|
let connection = self
|
||||||
let mut consumed = self.response_consumed_since_update.lock();
|
.response_connection
|
||||||
*consumed = consumed.saturating_add(bytes as u64);
|
.lock()
|
||||||
let threshold = u64::from(*STREAM_MIN_WINDOW_UPDATE_BYTES);
|
.as_ref()
|
||||||
if *consumed < threshold {
|
.and_then(std::sync::Weak::upgrade);
|
||||||
return None;
|
let Some(connection) = connection else {
|
||||||
|
return Ok(());
|
||||||
|
};
|
||||||
|
if connection.protocol_version() < 3 {
|
||||||
|
return Ok(());
|
||||||
}
|
}
|
||||||
|
let delta = {
|
||||||
let delta = (*consumed).min(u64::from(u32::MAX)) as u32;
|
let consumed = self.response_consumed_since_update.lock();
|
||||||
*consumed = consumed.saturating_sub(u64::from(delta));
|
if *consumed < u64::from(self.min_window_update_bytes) {
|
||||||
Some(delta)
|
return Ok(());
|
||||||
|
}
|
||||||
|
(*consumed).min(u64::from(u32::MAX)) as u32
|
||||||
|
};
|
||||||
|
let frame = protocol::encode_window_update(self.proxy_stream_id, delta);
|
||||||
|
if connection
|
||||||
|
.send_wait(Message::Binary(frame.into()), OUTBOUND_BACKPRESSURE_TIMEOUT)
|
||||||
|
.await
|
||||||
|
== SendStatus::Queued
|
||||||
|
{
|
||||||
|
let mut consumed = self.response_consumed_since_update.lock();
|
||||||
|
*consumed = consumed.saturating_sub(u64::from(delta));
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
if self.terminal.load(Ordering::Acquire) {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
connection.request_close();
|
||||||
|
Err("proxy flow-control update failed".to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
|
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
|
||||||
tokio::time::timeout(timeout, async {
|
tokio::time::timeout(timeout, async {
|
||||||
loop {
|
loop {
|
||||||
|
let notified = self.headers_notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
let outcome = {
|
let outcome = {
|
||||||
let state = self.wait_state.lock();
|
let state = self.wait_state.lock();
|
||||||
if let Some(response) = &state.response {
|
if let Some(response) = &state.response {
|
||||||
@@ -662,15 +712,20 @@ impl LocalStream {
|
|||||||
if let Some(error) = outcome {
|
if let Some(error) = outcome {
|
||||||
return Err(error);
|
return Err(error);
|
||||||
}
|
}
|
||||||
self.headers_notify.notified().await;
|
notified.await;
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
.map_err(|_| "timed out waiting for response headers".to_string())?
|
.map_err(|_| "timed out waiting for response headers".to_string())?
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
|
pub fn take_body_receiver(self: &Arc<Self>) -> Option<LocalBodyReceiver> {
|
||||||
self.body_rx.lock().take()
|
self.body.take_receiver().map(|receiver| LocalBodyReceiver {
|
||||||
|
receiver,
|
||||||
|
stream: Arc::clone(self),
|
||||||
|
failed: false,
|
||||||
|
pending: None,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
|
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
|
||||||
@@ -690,21 +745,11 @@ impl LocalStream {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn push_body_chunk(&self, payload: Bytes) -> bool {
|
fn push_body_chunk(&self, payload: Bytes) -> bool {
|
||||||
if self.terminal.load(Ordering::Acquire) {
|
if self.terminal.load(Ordering::Acquire) {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
// Use a timeout to prevent a slow consumer from blocking the shared
|
self.body.push(payload)
|
||||||
// proxy-connection reader (head-of-line blocking across streams).
|
|
||||||
match tokio::time::timeout(
|
|
||||||
Duration::from_secs(5),
|
|
||||||
self.body_tx.send(LocalBodyEvent::Chunk(payload)),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(Ok(())) => true,
|
|
||||||
_ => false,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn finish(&self) {
|
fn finish(&self) {
|
||||||
@@ -722,7 +767,8 @@ impl LocalStream {
|
|||||||
if notify {
|
if notify {
|
||||||
self.headers_notify.notify_waiters();
|
self.headers_notify.notify_waiters();
|
||||||
}
|
}
|
||||||
let _ = self.body_tx.try_send(LocalBodyEvent::End);
|
self.request_window.close();
|
||||||
|
self.body.finish(Ok(()));
|
||||||
}
|
}
|
||||||
|
|
||||||
fn fail(&self, error: impl Into<String>) {
|
fn fail(&self, error: impl Into<String>) {
|
||||||
@@ -742,7 +788,38 @@ impl LocalStream {
|
|||||||
if notify {
|
if notify {
|
||||||
self.headers_notify.notify_waiters();
|
self.headers_notify.notify_waiters();
|
||||||
}
|
}
|
||||||
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
|
self.request_window.close();
|
||||||
|
self.body.finish(Err(error));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub struct LocalBodyReceiver {
|
||||||
|
receiver: BodyReceiver,
|
||||||
|
stream: Arc<LocalStream>,
|
||||||
|
failed: bool,
|
||||||
|
pending: Option<LocalBodyEvent>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl LocalBodyReceiver {
|
||||||
|
pub async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||||
|
if self.failed {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
if self.pending.is_none() {
|
||||||
|
let event = self.receiver.recv().await?;
|
||||||
|
if let LocalBodyEvent::Chunk(chunk) = &event {
|
||||||
|
let mut consumed = self.stream.response_consumed_since_update.lock();
|
||||||
|
*consumed = consumed.saturating_add(chunk.len() as u64);
|
||||||
|
}
|
||||||
|
self.pending = Some(event);
|
||||||
|
}
|
||||||
|
if matches!(self.pending, Some(LocalBodyEvent::Chunk(_))) {
|
||||||
|
if let Err(error) = self.stream.flush_response_credit().await {
|
||||||
|
self.failed = true;
|
||||||
|
return Some(LocalBodyEvent::Error(error));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
self.pending.take()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -765,6 +842,21 @@ pub struct HubRouter {
|
|||||||
drain_reasons: Mutex<HashMap<String, u64>>,
|
drain_reasons: Mutex<HashMap<String, u64>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
struct PendingStreamGuard<'router> {
|
||||||
|
hub: &'router HubRouter,
|
||||||
|
connection: &'router ProxyConn,
|
||||||
|
stream_id: u64,
|
||||||
|
committed: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for PendingStreamGuard<'_> {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
if !self.committed && self.hub.cleanup_local_stream(self.stream_id) {
|
||||||
|
self.connection.release_stream();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
struct NodeStatusEvent {
|
struct NodeStatusEvent {
|
||||||
node_id: String,
|
node_id: String,
|
||||||
authenticated_key: Option<String>,
|
authenticated_key: Option<String>,
|
||||||
@@ -1166,17 +1258,27 @@ impl HubRouter {
|
|||||||
|
|
||||||
// Frames encoded successfully -- now register the stream.
|
// Frames encoded successfully -- now register the stream.
|
||||||
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
||||||
let local_stream = Arc::new(LocalStream::new(
|
let settings = proxy_conn.settings.lock().clone();
|
||||||
|
let mut local_stream = LocalStream::new(
|
||||||
local_stream_id,
|
local_stream_id,
|
||||||
proxy_conn.node_generation.clone(),
|
proxy_conn.node_generation.clone(),
|
||||||
proxy_conn.id,
|
proxy_conn.id,
|
||||||
proxy_stream_id,
|
proxy_stream_id,
|
||||||
*STREAM_INITIAL_WINDOW_BYTES,
|
settings.initial_stream_window_bytes,
|
||||||
));
|
);
|
||||||
|
local_stream.min_window_update_bytes = settings.min_window_update_bytes;
|
||||||
|
*local_stream.response_connection.get_mut() = Some(Arc::downgrade(&proxy_conn));
|
||||||
|
let local_stream = Arc::new(local_stream);
|
||||||
self.local_streams
|
self.local_streams
|
||||||
.insert(local_stream_id, local_stream.clone());
|
.insert(local_stream_id, local_stream.clone());
|
||||||
self.proxy_to_local
|
self.proxy_to_local
|
||||||
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
||||||
|
let mut pending_stream = PendingStreamGuard {
|
||||||
|
hub: self,
|
||||||
|
connection: &proxy_conn,
|
||||||
|
stream_id: local_stream_id,
|
||||||
|
committed: false,
|
||||||
|
};
|
||||||
|
|
||||||
let send_status = proxy_conn
|
let send_status = proxy_conn
|
||||||
.send_wait(
|
.send_wait(
|
||||||
@@ -1195,10 +1297,11 @@ impl HubRouter {
|
|||||||
"open_local_stream dispatched"
|
"open_local_stream dispatched"
|
||||||
);
|
);
|
||||||
match send_status {
|
match send_status {
|
||||||
SendStatus::Queued => Ok(local_stream),
|
SendStatus::Queued => {
|
||||||
|
pending_stream.committed = true;
|
||||||
|
Ok(local_stream)
|
||||||
|
}
|
||||||
SendStatus::Closed | SendStatus::Congested => {
|
SendStatus::Closed | SendStatus::Congested => {
|
||||||
self.cleanup_local_stream(local_stream_id);
|
|
||||||
proxy_conn.release_stream();
|
|
||||||
Err("proxy connection congested".to_string())
|
Err("proxy connection congested".to_string())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1243,7 +1346,9 @@ impl HubRouter {
|
|||||||
.map(|entry| entry.value().clone())
|
.map(|entry| entry.value().clone())
|
||||||
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
||||||
|
|
||||||
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
|
let chunk_size = MAX_REQUEST_BODY_FRAME_SIZE
|
||||||
|
.min(proxy_conn.settings.lock().initial_stream_window_bytes as usize);
|
||||||
|
let total_chunks = payload.len().div_ceil(chunk_size);
|
||||||
let result = if total_chunks == 0 {
|
let result = if total_chunks == 0 {
|
||||||
if end_stream {
|
if end_stream {
|
||||||
self.send_request_body_frame(&proxy_conn, &stream, &[], true)
|
self.send_request_body_frame(&proxy_conn, &stream, &[], true)
|
||||||
@@ -1252,7 +1357,7 @@ impl HubRouter {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
|
for (index, chunk) in payload.chunks(chunk_size).enumerate() {
|
||||||
let is_last_chunk = index + 1 == total_chunks;
|
let is_last_chunk = index + 1 == total_chunks;
|
||||||
if let Err(error) = self
|
if let Err(error) = self
|
||||||
.send_request_body_frame(
|
.send_request_body_frame(
|
||||||
@@ -1352,17 +1457,20 @@ impl HubRouter {
|
|||||||
} else {
|
} else {
|
||||||
protocol::encode_stream_error(stream.proxy_stream_id, reason)
|
protocol::encode_stream_error(stream.proxy_stream_id, reason)
|
||||||
};
|
};
|
||||||
let _ = pc.send(Message::Binary(frame.into()));
|
if pc.send(Message::Binary(frame.into())) != SendStatus::Queued {
|
||||||
|
pc.request_close();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
stream.fail(reason.to_string());
|
stream.fail(reason.to_string());
|
||||||
}
|
}
|
||||||
|
|
||||||
fn cleanup_local_stream(&self, local_stream_id: u64) {
|
fn cleanup_local_stream(&self, local_stream_id: u64) -> bool {
|
||||||
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
||||||
return;
|
return false;
|
||||||
};
|
};
|
||||||
self.proxy_to_local
|
self.proxy_to_local
|
||||||
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
||||||
|
true
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn handle_proxy_frame(self: &Arc<Self>, proxy_conn_id: u64, data: &mut [u8]) {
|
pub async fn handle_proxy_frame(self: &Arc<Self>, proxy_conn_id: u64, data: &mut [u8]) {
|
||||||
@@ -1433,9 +1541,7 @@ impl HubRouter {
|
|||||||
.get(&proxy_conn_id)
|
.get(&proxy_conn_id)
|
||||||
.map(|entry| entry.value().clone());
|
.map(|entry| entry.value().clone());
|
||||||
if let Some(pc) = pc {
|
if let Some(pc) = pc {
|
||||||
let _ = pc
|
let _ = pc.send(Message::Binary(pong.into()));
|
||||||
.send_wait(Message::Binary(pong.into()), Duration::from_millis(250))
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
protocol::PONG => {}
|
protocol::PONG => {}
|
||||||
@@ -1509,11 +1615,36 @@ impl HubRouter {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
protocol::SETTINGS => {
|
protocol::SETTINGS => {
|
||||||
debug!(
|
let settings = protocol::decode_payload_with_limit(
|
||||||
msg_type = header.msg_type,
|
data,
|
||||||
proxy_conn_id = proxy_conn_id,
|
&header,
|
||||||
"received tunnel protocol v3 SETTINGS from proxy"
|
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||||
);
|
)
|
||||||
|
.ok()
|
||||||
|
.and_then(|payload| {
|
||||||
|
serde_json::from_slice::<protocol::SettingsPayload>(&payload).ok()
|
||||||
|
})
|
||||||
|
.filter(|settings| settings.is_valid());
|
||||||
|
if let Some(connection) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||||
|
if header.stream_id != 0 || header.flags != 0 {
|
||||||
|
connection.request_close();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
let Some(settings) = settings else {
|
||||||
|
connection.request_close();
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let local = local_settings();
|
||||||
|
let settings = settings
|
||||||
|
.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||||
|
let mut current = connection.settings.lock();
|
||||||
|
if connection.stream_count.load(Ordering::Acquire) > 0 && *current != settings {
|
||||||
|
drop(current);
|
||||||
|
connection.request_close();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
*current = settings;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
protocol::WINDOW_UPDATE => {
|
protocol::WINDOW_UPDATE => {
|
||||||
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
|
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
|
||||||
@@ -1739,19 +1870,8 @@ impl HubRouter {
|
|||||||
None => return,
|
None => return,
|
||||||
};
|
};
|
||||||
|
|
||||||
let payload_len = payload.len();
|
if !stream.push_body_chunk(Bytes::from(payload)) {
|
||||||
if !stream.push_body_chunk(Bytes::from(payload)).await {
|
|
||||||
self.cancel_local_stream(local_id, "local relay response congested");
|
self.cancel_local_stream(local_id, "local relay response congested");
|
||||||
return;
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
|
||||||
if pc.protocol_version() >= 3 {
|
|
||||||
if let Some(delta) = stream.response_window_update_delta(payload_len) {
|
|
||||||
let frame = protocol::encode_window_update(header.stream_id, delta);
|
|
||||||
let _ = pc.send(Message::Binary(frame.into()));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -9,14 +9,13 @@ use axum::body::{Body, Bytes};
|
|||||||
use axum::extract::{ConnectInfo, Path, Request, State};
|
use axum::extract::{ConnectInfo, Path, Request, State};
|
||||||
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||||
use axum::response::IntoResponse;
|
use axum::response::IntoResponse;
|
||||||
use tokio::sync::mpsc;
|
|
||||||
use tracing::warn;
|
use tracing::warn;
|
||||||
|
|
||||||
use crate::api::response::apply_streaming_response_headers;
|
use crate::api::response::apply_streaming_response_headers;
|
||||||
use crate::headers::should_skip_response_header;
|
use crate::headers::should_skip_response_header;
|
||||||
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
|
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
|
||||||
|
|
||||||
use super::hub::{LocalBodyEvent, LocalStream};
|
use super::hub::{LocalBodyEvent, LocalBodyReceiver, LocalStream};
|
||||||
use super::protocol;
|
use super::protocol;
|
||||||
use super::{AppState, RelayRequestAuthenticated};
|
use super::{AppState, RelayRequestAuthenticated};
|
||||||
|
|
||||||
@@ -40,7 +39,7 @@ impl Drop for StreamGuard {
|
|||||||
pub(crate) struct DirectRelayResponse {
|
pub(crate) struct DirectRelayResponse {
|
||||||
status: u16,
|
status: u16,
|
||||||
headers: Vec<(String, String)>,
|
headers: Vec<(String, String)>,
|
||||||
body_rx: mpsc::Receiver<LocalBodyEvent>,
|
body_rx: LocalBodyReceiver,
|
||||||
request_guard: StreamGuard,
|
request_guard: StreamGuard,
|
||||||
_request_permit: Option<AdmissionPermit>,
|
_request_permit: Option<AdmissionPermit>,
|
||||||
}
|
}
|
||||||
@@ -55,10 +54,13 @@ impl DirectRelayResponse {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, String> {
|
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, String> {
|
||||||
|
if self.request_guard.finished {
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
let event = self.body_rx.recv().await;
|
let event = self.body_rx.recv().await;
|
||||||
match event {
|
match event {
|
||||||
Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)),
|
Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)),
|
||||||
Some(LocalBodyEvent::End) | None => {
|
Some(LocalBodyEvent::End) => {
|
||||||
self.request_guard.finished = true;
|
self.request_guard.finished = true;
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
@@ -66,6 +68,7 @@ impl DirectRelayResponse {
|
|||||||
self.request_guard.finished = true;
|
self.request_guard.finished = true;
|
||||||
Err(error)
|
Err(error)
|
||||||
}
|
}
|
||||||
|
None => Err("tunnel response ended without a terminal frame".to_string()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -84,6 +87,11 @@ pub(crate) async fn open_direct_relay_stream(
|
|||||||
.open_authorized_local_stream(node_id, &meta)
|
.open_authorized_local_stream(node_id, &meta)
|
||||||
.await
|
.await
|
||||||
.map_err(|error| format!("connect: {error}"))?;
|
.map_err(|error| format!("connect: {error}"))?;
|
||||||
|
let request_guard = StreamGuard {
|
||||||
|
hub: state.hub.clone(),
|
||||||
|
stream_id: stream.id,
|
||||||
|
finished: false,
|
||||||
|
};
|
||||||
if let Err(error) = state
|
if let Err(error) = state
|
||||||
.hub
|
.hub
|
||||||
.push_local_request_body(stream.id, body, true)
|
.push_local_request_body(stream.id, body, true)
|
||||||
@@ -126,11 +134,7 @@ pub(crate) async fn open_direct_relay_stream(
|
|||||||
status: response_head.status,
|
status: response_head.status,
|
||||||
headers: response_head.headers,
|
headers: response_head.headers,
|
||||||
body_rx,
|
body_rx,
|
||||||
request_guard: StreamGuard {
|
request_guard,
|
||||||
hub: state.hub.clone(),
|
|
||||||
stream_id: stream.id,
|
|
||||||
finished: false,
|
|
||||||
},
|
|
||||||
_request_permit: request_permit,
|
_request_permit: request_permit,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -259,6 +263,11 @@ pub async fn relay_request(
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
let request_guard = StreamGuard {
|
||||||
|
hub: state.hub.clone(),
|
||||||
|
stream_id: stream.id,
|
||||||
|
finished: false,
|
||||||
|
};
|
||||||
let body_stream = match spool.body_stream().await {
|
let body_stream = match spool.body_stream().await {
|
||||||
Ok(stream) => stream,
|
Ok(stream) => stream,
|
||||||
Err(error) => {
|
Err(error) => {
|
||||||
@@ -306,12 +315,6 @@ pub async fn relay_request(
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let request_guard = StreamGuard {
|
|
||||||
hub: state.hub.clone(),
|
|
||||||
stream_id: stream.id,
|
|
||||||
finished: false,
|
|
||||||
};
|
|
||||||
|
|
||||||
let wait_timeout = relay_header_timeout(&meta);
|
let wait_timeout = relay_header_timeout(&meta);
|
||||||
let response_head = match stream.wait_headers(wait_timeout).await {
|
let response_head = match stream.wait_headers(wait_timeout).await {
|
||||||
Ok(response) => response,
|
Ok(response) => response,
|
||||||
@@ -373,6 +376,9 @@ pub async fn relay_request(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if !guard.finished {
|
||||||
|
yield Err(io::Error::other("tunnel response ended without a terminal frame"));
|
||||||
|
}
|
||||||
guard.finished = true;
|
guard.finished = true;
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -563,6 +569,100 @@ mod tests {
|
|||||||
request
|
request
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cancelled_relays_reset_streams_during_upload_and_header_wait() {
|
||||||
|
for direct in [true, false] {
|
||||||
|
for during_upload in [true, false] {
|
||||||
|
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||||
|
sample_connected_proxy_node("node-123"),
|
||||||
|
]));
|
||||||
|
let data = Arc::new(
|
||||||
|
GatewayDataState::with_proxy_node_repository_for_tests(repository)
|
||||||
|
.with_system_config_values_for_tests(
|
||||||
|
Vec::<(String, serde_json::Value)>::new(),
|
||||||
|
)
|
||||||
|
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||||
|
);
|
||||||
|
let state = test_app_state().with_data(data);
|
||||||
|
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
|
||||||
|
let (proxy_close_tx, _) = watch::channel(false);
|
||||||
|
let connection = Arc::new(
|
||||||
|
ProxyConn::new(
|
||||||
|
500,
|
||||||
|
"node-123".into(),
|
||||||
|
"Node 123".into(),
|
||||||
|
proxy_tx,
|
||||||
|
proxy_close_tx,
|
||||||
|
16,
|
||||||
|
3,
|
||||||
|
)
|
||||||
|
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
|
||||||
|
.with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string())
|
||||||
|
.with_settings(protocol::SettingsPayload {
|
||||||
|
initial_stream_window_bytes: 128,
|
||||||
|
min_window_update_bytes: 32,
|
||||||
|
drain_deadline_ms: 1000,
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
state.hub.register_proxy(Arc::clone(&connection));
|
||||||
|
let meta = protocol::RequestMeta {
|
||||||
|
provider_id: None,
|
||||||
|
endpoint_id: None,
|
||||||
|
key_id: None,
|
||||||
|
method: "POST".into(),
|
||||||
|
url: "https://example.com/".into(),
|
||||||
|
headers: HashMap::new(),
|
||||||
|
stream: true,
|
||||||
|
request_timeout_ms: None,
|
||||||
|
stream_first_byte_timeout_ms: None,
|
||||||
|
timeout: 30,
|
||||||
|
follow_redirects: None,
|
||||||
|
http1_only: false,
|
||||||
|
transport_profile: None,
|
||||||
|
};
|
||||||
|
let body = Bytes::from(vec![b'x'; if during_upload { 256 } else { 0 }]);
|
||||||
|
let relay = tokio::spawn(async move {
|
||||||
|
if direct {
|
||||||
|
let _response =
|
||||||
|
super::open_direct_relay_stream(&state, "node-123", meta, body)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
} else {
|
||||||
|
let request =
|
||||||
|
authenticated_request(encode_relay_envelope(&meta, &body)).await;
|
||||||
|
let _response = relay_request(
|
||||||
|
Path("node-123".into()),
|
||||||
|
State(state),
|
||||||
|
ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4242))),
|
||||||
|
request,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
});
|
||||||
|
recv_tunnel_test_frame(&mut proxy_rx, "request headers").await;
|
||||||
|
recv_tunnel_test_frame(&mut proxy_rx, "request body").await;
|
||||||
|
relay.abort();
|
||||||
|
assert!(relay.await.unwrap_err().is_cancelled());
|
||||||
|
let Message::Binary(frame) = recv_tunnel_test_frame(&mut proxy_rx, "reset").await
|
||||||
|
else {
|
||||||
|
panic!("expected binary reset frame")
|
||||||
|
};
|
||||||
|
let frame = aether_contracts::tunnel::Frame::decode(frame).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
frame.msg_type,
|
||||||
|
aether_contracts::tunnel::MsgType::ResetStream
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
connection
|
||||||
|
.stream_count
|
||||||
|
.load(std::sync::atomic::Ordering::Relaxed),
|
||||||
|
0
|
||||||
|
);
|
||||||
|
assert!(connection.is_available());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
|
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
|
||||||
let meta = protocol::RequestMeta {
|
let meta = protocol::RequestMeta {
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
mod body;
|
||||||
mod control_plane;
|
mod control_plane;
|
||||||
mod hub;
|
mod hub;
|
||||||
mod local_relay;
|
mod local_relay;
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ use aether_runtime::bounded_queue;
|
|||||||
use axum::extract::ws::{Message, WebSocket};
|
use axum::extract::ws::{Message, WebSocket};
|
||||||
use futures_util::{SinkExt, StreamExt};
|
use futures_util::{SinkExt, StreamExt};
|
||||||
use tokio::sync::watch;
|
use tokio::sync::watch;
|
||||||
|
use tokio::task::JoinSet;
|
||||||
use tracing::{debug, info, warn};
|
use tracing::{debug, info, warn};
|
||||||
|
|
||||||
use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus};
|
use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus};
|
||||||
@@ -84,6 +85,35 @@ pub async fn handle_proxy_connection(
|
|||||||
|
|
||||||
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
|
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
|
||||||
let (close_tx, mut close_rx) = watch::channel(false);
|
let (close_tx, mut close_rx) = watch::channel(false);
|
||||||
|
let settings = if protocol_version >= 3 {
|
||||||
|
let Some(settings) = read_proxy_settings(
|
||||||
|
&mut ws_tx,
|
||||||
|
&mut ws_rx,
|
||||||
|
security.as_deref(),
|
||||||
|
protocol_version,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
else {
|
||||||
|
warn!(conn_id, "proxy SETTINGS negotiation failed");
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let local = super::hub::local_settings();
|
||||||
|
let negotiated =
|
||||||
|
settings.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||||
|
let message = Message::Binary(protocol::encode_settings(&negotiated).into());
|
||||||
|
let Ok(message) = encrypt_message(message, security.as_deref()) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
if !matches!(
|
||||||
|
tokio::time::timeout(PROXY_HELLO_TIMEOUT, ws_tx.send(message)).await,
|
||||||
|
Ok(Ok(()))
|
||||||
|
) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
negotiated
|
||||||
|
} else {
|
||||||
|
super::hub::local_settings()
|
||||||
|
};
|
||||||
let conn = ProxyConn::new(
|
let conn = ProxyConn::new(
|
||||||
conn_id,
|
conn_id,
|
||||||
node_id.clone(),
|
node_id.clone(),
|
||||||
@@ -93,7 +123,8 @@ pub async fn handle_proxy_connection(
|
|||||||
max_streams,
|
max_streams,
|
||||||
protocol_version,
|
protocol_version,
|
||||||
)
|
)
|
||||||
.with_tunnel_generation(node_generation);
|
.with_tunnel_generation(node_generation)
|
||||||
|
.with_settings(settings);
|
||||||
let conn = match (security_key.clone(), management_token_credential) {
|
let conn = match (security_key.clone(), management_token_credential) {
|
||||||
(Some(key), None) => Arc::new(conn.with_authenticated_key(key)),
|
(Some(key), None) => Arc::new(conn.with_authenticated_key(key)),
|
||||||
(None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)),
|
(None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)),
|
||||||
@@ -416,19 +447,22 @@ async fn run_proxy_reader(
|
|||||||
let idle_enabled = !idle_timeout.is_zero();
|
let idle_enabled = !idle_timeout.is_zero();
|
||||||
let mut oversized_count = 0u32;
|
let mut oversized_count = 0u32;
|
||||||
let mut frames_received: u64 = 0;
|
let mut frames_received: u64 = 0;
|
||||||
|
let mut close_rx = conn.outbound.subscribe_close();
|
||||||
|
let mut heartbeats = JoinSet::new();
|
||||||
loop {
|
loop {
|
||||||
let msg = if idle_enabled {
|
if conn.outbound.is_closing() {
|
||||||
tokio::select! {
|
break;
|
||||||
msg = ws_rx.next() => msg,
|
}
|
||||||
_ = tokio::time::sleep(idle_timeout) => {
|
while heartbeats.try_join_next().is_some() {}
|
||||||
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
let msg = tokio::select! {
|
||||||
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
|
biased;
|
||||||
conn.request_close();
|
_ = close_rx.changed() => break,
|
||||||
break;
|
msg = ws_rx.next() => msg,
|
||||||
}
|
_ = tokio::time::sleep(idle_timeout), if idle_enabled => {
|
||||||
|
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
||||||
|
conn.request_close();
|
||||||
|
break;
|
||||||
}
|
}
|
||||||
} else {
|
|
||||||
ws_rx.next().await
|
|
||||||
};
|
};
|
||||||
|
|
||||||
match msg {
|
match msg {
|
||||||
@@ -463,7 +497,27 @@ async fn run_proxy_reader(
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
let is_heartbeat = protocol::FrameHeader::parse(&data)
|
||||||
|
.is_some_and(|header| header.msg_type == protocol::HEARTBEAT_DATA);
|
||||||
|
if is_heartbeat {
|
||||||
|
if heartbeats.is_empty() {
|
||||||
|
let heartbeat_hub = Arc::clone(&hub);
|
||||||
|
let conn_id = conn.id;
|
||||||
|
heartbeats.spawn(async move {
|
||||||
|
if tokio::time::timeout(
|
||||||
|
Duration::from_secs(10),
|
||||||
|
heartbeat_hub.handle_proxy_frame(conn_id, &mut data),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.is_err()
|
||||||
|
{
|
||||||
|
warn!(conn_id, "proxy heartbeat processing timed out");
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
Some(Ok(Message::Close(_))) | None => {
|
Some(Ok(Message::Close(_))) | None => {
|
||||||
info!(
|
info!(
|
||||||
@@ -489,6 +543,56 @@ async fn run_proxy_reader(
|
|||||||
_ => {}
|
_ => {}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
heartbeats.shutdown().await;
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn read_proxy_settings(
|
||||||
|
ws_tx: &mut futures_util::stream::SplitSink<WebSocket, Message>,
|
||||||
|
ws_rx: &mut futures_util::stream::SplitStream<WebSocket>,
|
||||||
|
security: Option<&SecureFrameCodec>,
|
||||||
|
protocol_version: u8,
|
||||||
|
) -> Option<protocol::SettingsPayload> {
|
||||||
|
tokio::time::timeout(PROXY_HELLO_TIMEOUT, async {
|
||||||
|
let mut hello_received = security.is_some();
|
||||||
|
for _ in 0..MAX_PREAUTH_PINGS {
|
||||||
|
match ws_rx.next().await? {
|
||||||
|
Ok(Message::Binary(data)) => {
|
||||||
|
if data.len() > 256 * 1024 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let data = decrypt_message(data, security).ok()?;
|
||||||
|
let frame = Frame::decode(data.into()).ok()?;
|
||||||
|
if frame.stream_id != 0 || frame.flags != 0 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
match frame.msg_type {
|
||||||
|
MsgType::Hello if !hello_received => {
|
||||||
|
let hello =
|
||||||
|
serde_json::from_slice::<HelloPayload>(&frame.payload).ok()?;
|
||||||
|
if hello.protocol_version != protocol_version {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
hello_received = true;
|
||||||
|
}
|
||||||
|
MsgType::Settings if hello_received => {
|
||||||
|
let settings =
|
||||||
|
serde_json::from_slice::<protocol::SettingsPayload>(&frame.payload)
|
||||||
|
.ok()?;
|
||||||
|
return settings.is_valid().then_some(settings);
|
||||||
|
}
|
||||||
|
_ => return None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(Message::Ping(payload)) => ws_tx.send(Message::Pong(payload)).await.ok()?,
|
||||||
|
Ok(Message::Pong(_)) => {}
|
||||||
|
_ => return None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
None
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.ok()
|
||||||
|
.flatten()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn encrypt_message(
|
fn encrypt_message(
|
||||||
@@ -523,6 +627,158 @@ fn decrypt_message(
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
|
#[cfg(feature = "testkit")]
|
||||||
|
#[tokio::test]
|
||||||
|
async fn slow_heartbeat_does_not_block_response_frames() {
|
||||||
|
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||||
|
use tokio_tungstenite::tungstenite::{client::IntoClientRequest, Message as ClientMessage};
|
||||||
|
|
||||||
|
let called = Arc::new(AtomicUsize::new(0));
|
||||||
|
let callback_called = Arc::clone(&called);
|
||||||
|
let control_plane = super::super::control_plane::ControlPlaneClient::local(
|
||||||
|
move |_, _| {
|
||||||
|
callback_called.fetch_add(1, Ordering::SeqCst);
|
||||||
|
Box::pin(std::future::pending())
|
||||||
|
},
|
||||||
|
|_, _, _, _| Box::pin(async { Ok(()) }),
|
||||||
|
);
|
||||||
|
let data = crate::data::GatewayDataState::with_tunnel_management_auth_for_testkit(
|
||||||
|
"heartbeat-test",
|
||||||
|
"heartbeat-generation",
|
||||||
|
"ae-tunnel-harness-management-token",
|
||||||
|
aether_crypto::DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
let state = super::super::AppState::new(
|
||||||
|
control_plane,
|
||||||
|
ConnConfig {
|
||||||
|
ping_interval: Duration::from_secs(60),
|
||||||
|
idle_timeout: Duration::ZERO,
|
||||||
|
outbound_queue_capacity: 128,
|
||||||
|
},
|
||||||
|
16,
|
||||||
|
)
|
||||||
|
.with_data(Arc::new(data));
|
||||||
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||||
|
let address = listener.local_addr().unwrap();
|
||||||
|
let router = super::super::build_router_with_state(state.clone());
|
||||||
|
let server = tokio::spawn(async move {
|
||||||
|
axum::serve(
|
||||||
|
listener,
|
||||||
|
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
});
|
||||||
|
let mut request = format!("ws://{address}/api/internal/proxy-tunnel")
|
||||||
|
.into_client_request()
|
||||||
|
.unwrap();
|
||||||
|
let headers = request.headers_mut();
|
||||||
|
headers.insert("x-node-id", "heartbeat-test".parse().unwrap());
|
||||||
|
headers.insert(
|
||||||
|
aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER,
|
||||||
|
"heartbeat-generation".parse().unwrap(),
|
||||||
|
);
|
||||||
|
headers.insert(
|
||||||
|
aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER,
|
||||||
|
"3".parse().unwrap(),
|
||||||
|
);
|
||||||
|
headers.insert(
|
||||||
|
"authorization",
|
||||||
|
"Bearer ae-tunnel-harness-management-token".parse().unwrap(),
|
||||||
|
);
|
||||||
|
let (mut websocket, _) = tokio_tungstenite::connect_async(request).await.unwrap();
|
||||||
|
let hello = HelloPayload {
|
||||||
|
protocol_version: 3,
|
||||||
|
capabilities: vec![],
|
||||||
|
session_id: None,
|
||||||
|
replica_id: None,
|
||||||
|
};
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(protocol::encode_hello(&hello).into()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(
|
||||||
|
protocol::encode_settings(&super::super::hub::local_settings()).into(),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let ClientMessage::Binary(settings) = websocket.next().await.unwrap().unwrap() else {
|
||||||
|
panic!("expected SETTINGS")
|
||||||
|
};
|
||||||
|
assert_eq!(Frame::decode(settings).unwrap().msg_type, MsgType::Settings);
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while !state.hub.has_local_proxy("heartbeat-test") {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let meta: protocol::RequestMeta = serde_json::from_value(serde_json::json!({
|
||||||
|
"method": "GET", "url": "https://example.com", "headers": {}, "stream": true, "timeout": 10
|
||||||
|
})).unwrap();
|
||||||
|
let stream = state
|
||||||
|
.hub
|
||||||
|
.open_local_stream("heartbeat-test", &meta)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let ClientMessage::Binary(request) = websocket.next().await.unwrap().unwrap() else {
|
||||||
|
panic!("expected request headers")
|
||||||
|
};
|
||||||
|
let stream_id = Frame::decode(request).unwrap().stream_id;
|
||||||
|
let heartbeat = Frame::control(
|
||||||
|
MsgType::HeartbeatData,
|
||||||
|
serde_json::to_vec(&serde_json::json!({"node_id": "heartbeat-test"})).unwrap(),
|
||||||
|
);
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while called.load(Ordering::SeqCst) == 0 {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
for _ in 0..8 {
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
let response = Frame::new(
|
||||||
|
stream_id,
|
||||||
|
MsgType::ResponseHeaders,
|
||||||
|
0,
|
||||||
|
serde_json::to_vec(&serde_json::json!({"status": 200, "headers": []})).unwrap(),
|
||||||
|
);
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(response.encode()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
stream
|
||||||
|
.wait_headers(Duration::from_secs(1))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.status,
|
||||||
|
200
|
||||||
|
);
|
||||||
|
assert_eq!(called.load(Ordering::SeqCst), 1);
|
||||||
|
state.hub.request_close_all_proxies();
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while state.hub.has_local_proxy("heartbeat-test") {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
server.abort();
|
||||||
|
let _ = server.await;
|
||||||
|
}
|
||||||
|
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
||||||
const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
||||||
|
|||||||
@@ -1308,7 +1308,13 @@ mod tests {
|
|||||||
stored[0].error_type.as_deref(),
|
stored[0].error_type.as_deref(),
|
||||||
Some("stream_missing_terminal_event")
|
Some("stream_missing_terminal_event")
|
||||||
);
|
);
|
||||||
assert!(stored[0].error_message.is_none());
|
assert_eq!(
|
||||||
|
stored[0].error_message.as_deref(),
|
||||||
|
Some(super::STREAM_MISSING_TERMINAL_EVENT_MESSAGE)
|
||||||
|
);
|
||||||
|
let mut public_candidate = stored[0].clone();
|
||||||
|
public_candidate.sanitize_sensitive_diagnostics();
|
||||||
|
assert!(public_candidate.error_message.is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "aether-tunnel"
|
name = "aether-tunnel"
|
||||||
version = "0.3.16"
|
version = "0.3.17"
|
||||||
edition = "2021"
|
edition = "2021"
|
||||||
description = "Tunnel agent for Aether"
|
description = "Tunnel agent for Aether"
|
||||||
|
|
||||||
@@ -47,3 +47,4 @@ uuid.workspace = true
|
|||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
aether-gateway = { workspace = true, features = ["testkit"] }
|
aether-gateway = { workspace = true, features = ["testkit"] }
|
||||||
|
tokio = { version = "1", features = ["test-util"] }
|
||||||
|
|||||||
@@ -4,6 +4,15 @@ Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道
|
|||||||
|
|
||||||
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
|
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
|
||||||
|
|
||||||
|
## 流式传输与升级注意事项
|
||||||
|
|
||||||
|
- 协议 v3 连接在 `HELLO` / `SETTINGS` 协商后才接收业务请求。实际双向流窗口取 gateway 与 agent 配置的较小值,信用更新阈值不超过该窗口的四分之一;单帧也不会超过协商窗口。
|
||||||
|
- 响应缓冲按字节限额并合并小帧,结束和错误状态独立保存。慢消费者不会阻塞同一隧道其他流的读取;超出窗口或缓冲预算的流会被明确终止,不会静默截断。
|
||||||
|
- 信用更新在消费数据后可靠入队;启用重定向重放时,进入有界重放缓存也视为请求体消费。持续无法投递关键控制帧时会关闭连接并向在途请求报告错误。
|
||||||
|
- 客户端取消会终止对应上游请求,断连会回收 session 的 writer、heartbeat 和请求任务。正常 drain 在配置期限内继续处理已有流,期限到达后终止残留任务。
|
||||||
|
- 建议先升级 gateway,再升级 agent。既有 v3 agent 已发送 `HELLO` / `SETTINGS`,可连接新 gateway;自定义 v3 节点必须完成这两步握手。协议 v1/v2 保留旧握手。与旧 gateway 混用时应保持默认窗口配置,不能依赖旧 gateway 应用新的窗口协商。
|
||||||
|
- 自动重连恢复后续请求,不会自动续传已经输出的 SSE,也不会无条件重放已经发送的请求。
|
||||||
|
|
||||||
## 安装
|
## 安装
|
||||||
|
|
||||||
`aether-tunnel` 会根据宿主机自动选择服务管理器:
|
`aether-tunnel` 会根据宿主机自动选择服务管理器:
|
||||||
@@ -15,13 +24,13 @@ Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到
|
|||||||
<!-- DOWNLOAD_TABLE_START -->
|
<!-- DOWNLOAD_TABLE_START -->
|
||||||
| Platform | Download |
|
| Platform | Download |
|
||||||
|----------|----------|
|
|----------|----------|
|
||||||
| Linux x86_64 (GNU) | [aether-tunnel-linux-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.16/aether-tunnel-linux-amd64.tar.gz) |
|
| Linux x86_64 (GNU) | [aether-tunnel-linux-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.17/aether-tunnel-linux-amd64.tar.gz) |
|
||||||
| Linux ARM64 (GNU) | [aether-tunnel-linux-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.16/aether-tunnel-linux-arm64.tar.gz) |
|
| Linux ARM64 (GNU) | [aether-tunnel-linux-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.17/aether-tunnel-linux-arm64.tar.gz) |
|
||||||
| Linux x86_64 (musl) | [aether-tunnel-linux-musl-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.16/aether-tunnel-linux-musl-amd64.tar.gz) |
|
| Linux x86_64 (musl) | [aether-tunnel-linux-musl-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.17/aether-tunnel-linux-musl-amd64.tar.gz) |
|
||||||
| Linux ARM64 (musl) | [aether-tunnel-linux-musl-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.16/aether-tunnel-linux-musl-arm64.tar.gz) |
|
| Linux ARM64 (musl) | [aether-tunnel-linux-musl-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.17/aether-tunnel-linux-musl-arm64.tar.gz) |
|
||||||
| macOS x86_64 | [aether-tunnel-macos-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.16/aether-tunnel-macos-amd64.tar.gz) |
|
| macOS x86_64 | [aether-tunnel-macos-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.17/aether-tunnel-macos-amd64.tar.gz) |
|
||||||
| macOS ARM64 | [aether-tunnel-macos-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.16/aether-tunnel-macos-arm64.tar.gz) |
|
| macOS ARM64 | [aether-tunnel-macos-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.17/aether-tunnel-macos-arm64.tar.gz) |
|
||||||
| Windows x86_64 | [aether-tunnel-windows-amd64.zip](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.16/aether-tunnel-windows-amd64.zip) |
|
| Windows x86_64 | [aether-tunnel-windows-amd64.zip](https://github.com/fawney19/Aether/releases/download/tunnel-v0.3.17/aether-tunnel-windows-amd64.zip) |
|
||||||
<!-- DOWNLOAD_TABLE_END -->
|
<!-- DOWNLOAD_TABLE_END -->
|
||||||
|
|
||||||
上表展示的是最新已发布版本的下载链接。从下一次 `tunnel-v*` 发布开始,表格会自动补上 `Linux x86_64 (musl)` / `Linux ARM64 (musl)` 包,供 Alpine 等 musl 系统直接使用。
|
上表展示的是最新已发布版本的下载链接。从下一次 `tunnel-v*` 发布开始,表格会自动补上 `Linux x86_64 (musl)` / `Linux ARM64 (musl)` 包,供 Alpine 等 musl 系统直接使用。
|
||||||
|
|||||||
@@ -764,6 +764,13 @@ impl Config {
|
|||||||
if self.tunnel_stream_initial_window_bytes == 0 {
|
if self.tunnel_stream_initial_window_bytes == 0 {
|
||||||
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
|
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
|
||||||
}
|
}
|
||||||
|
if u64::from(self.tunnel_stream_initial_window_bytes)
|
||||||
|
> aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64
|
||||||
|
{
|
||||||
|
anyhow::bail!(
|
||||||
|
"tunnel_stream_initial_window_bytes exceeds the maximum tunnel payload size"
|
||||||
|
);
|
||||||
|
}
|
||||||
if self.tunnel_drain_deadline_ms == 0 {
|
if self.tunnel_drain_deadline_ms == 0 {
|
||||||
anyhow::bail!("tunnel_drain_deadline_ms must be > 0");
|
anyhow::bail!("tunnel_drain_deadline_ms must be > 0");
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -203,19 +203,34 @@ pub async fn connect_and_run(
|
|||||||
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
|
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
|
||||||
|
|
||||||
// Spawn writer task (with WebSocket ping keepalive)
|
// Spawn writer task (with WebSocket ping keepalive)
|
||||||
let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
let (frame_tx, writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
||||||
ws_sink,
|
ws_sink,
|
||||||
ping_interval,
|
ping_interval,
|
||||||
Some(Arc::clone(&server.tunnel_metrics)),
|
Some(Arc::clone(&server.tunnel_metrics)),
|
||||||
security.clone(),
|
security.clone(),
|
||||||
);
|
);
|
||||||
|
let mut writer_handle = super::task::SessionTask::new(writer_handle);
|
||||||
send_protocol_v3_hello(&frame_tx, &security_session, state).await;
|
send_protocol_v3_hello(&frame_tx, &security_session, state).await;
|
||||||
let drain_signal = spawn_drain_signal(
|
let (session_drain_tx, session_drain_rx) = watch::channel(*drain.borrow());
|
||||||
|
let forward_drain_tx = session_drain_tx.clone();
|
||||||
|
let mut external_drain = drain;
|
||||||
|
let forward_drain = super::task::SessionTask::new(tokio::spawn(async move {
|
||||||
|
loop {
|
||||||
|
if *external_drain.borrow() {
|
||||||
|
let _ = forward_drain_tx.send(true);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
if external_drain.changed().await.is_err() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
let drain_signal = super::task::SessionTask::new(spawn_drain_signal(
|
||||||
conn_idx,
|
conn_idx,
|
||||||
frame_tx.clone(),
|
frame_tx.clone(),
|
||||||
drain.clone(),
|
session_drain_rx.clone(),
|
||||||
state.config.tunnel_drain_deadline_ms,
|
state.config.tunnel_drain_deadline_ms,
|
||||||
);
|
));
|
||||||
|
|
||||||
// Spawn heartbeat task (only for primary connection to avoid
|
// Spawn heartbeat task (only for primary connection to avoid
|
||||||
// resetting shared atomic metrics via swap(0))
|
// resetting shared atomic metrics via swap(0))
|
||||||
@@ -237,16 +252,19 @@ pub async fn connect_and_run(
|
|||||||
// ensures we detect this and trigger a reconnect promptly.
|
// ensures we detect this and trigger a reconnect promptly.
|
||||||
let state_clone = Arc::clone(state);
|
let state_clone = Arc::clone(state);
|
||||||
let server_clone = Arc::clone(server);
|
let server_clone = Arc::clone(server);
|
||||||
let outcome = tokio::select! {
|
let outcome = {
|
||||||
result = dispatcher::run_with_security(
|
let dispatch = dispatcher::run_with_security(
|
||||||
state_clone,
|
state_clone,
|
||||||
server_clone,
|
server_clone,
|
||||||
ws_read,
|
ws_read,
|
||||||
frame_tx.clone(),
|
frame_tx.clone(),
|
||||||
hb_handle,
|
hb_handle,
|
||||||
drain.clone(),
|
session_drain_rx,
|
||||||
security.clone(),
|
security.clone(),
|
||||||
) => {
|
);
|
||||||
|
tokio::pin!(dispatch);
|
||||||
|
tokio::select! {
|
||||||
|
result = &mut dispatch => {
|
||||||
match result {
|
match result {
|
||||||
Ok(()) => Ok(TunnelOutcome::Disconnected),
|
Ok(()) => Ok(TunnelOutcome::Disconnected),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -258,6 +276,8 @@ pub async fn connect_and_run(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
writer_result = &mut writer_handle => {
|
writer_result = &mut writer_handle => {
|
||||||
|
frame_tx.close();
|
||||||
|
let _ = tokio::time::timeout(Duration::from_secs(1), &mut dispatch).await;
|
||||||
match writer_result {
|
match writer_result {
|
||||||
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
|
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -278,24 +298,37 @@ pub async fn connect_and_run(
|
|||||||
}
|
}
|
||||||
_ = shutdown.changed() => {
|
_ = shutdown.changed() => {
|
||||||
debug!("shutdown during tunnel dispatch");
|
debug!("shutdown during tunnel dispatch");
|
||||||
|
let _ = session_drain_tx.send(true);
|
||||||
|
let deadline = Duration::from_millis(state.config.tunnel_drain_deadline_ms).saturating_add(Duration::from_secs(1));
|
||||||
|
let _ = tokio::time::timeout(deadline, &mut dispatch).await;
|
||||||
Ok(TunnelOutcome::Shutdown)
|
Ok(TunnelOutcome::Shutdown)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
// Drop our sender; the writer will exit once all stream handler clones
|
// Drop our sender; the writer will exit once all stream handler clones
|
||||||
// are also dropped (i.e. after they finish their in-flight work).
|
// are also dropped (i.e. after they finish their in-flight work).
|
||||||
drop(frame_tx);
|
drop(frame_tx);
|
||||||
|
forward_drain.abort();
|
||||||
|
let _ = forward_drain.await;
|
||||||
if !drain_signal.is_finished() {
|
if !drain_signal.is_finished() {
|
||||||
drain_signal.abort();
|
drain_signal.abort();
|
||||||
let _ = drain_signal.await;
|
let _ = drain_signal.await;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wait for the writer task to finish with a generous timeout — the
|
|
||||||
// dispatcher already waits up to 30s for stream handlers, so 35s here
|
|
||||||
// covers that plus a small margin.
|
|
||||||
// Skip if the writer already exited (the select branch that fired).
|
|
||||||
if !writer_handle.is_finished() {
|
if !writer_handle.is_finished() {
|
||||||
let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await;
|
let flush_timeout = if *session_drain_tx.borrow() {
|
||||||
|
Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||||
|
} else {
|
||||||
|
Duration::from_secs(1)
|
||||||
|
};
|
||||||
|
if tokio::time::timeout(flush_timeout, &mut writer_handle)
|
||||||
|
.await
|
||||||
|
.is_err()
|
||||||
|
{
|
||||||
|
writer_handle.abort();
|
||||||
|
let _ = writer_handle.await;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let connected_for = connected_at.elapsed();
|
let connected_for = connected_at.elapsed();
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ use std::time::Duration;
|
|||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use futures_util::StreamExt;
|
use futures_util::StreamExt;
|
||||||
use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore};
|
use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore};
|
||||||
use tokio::task::JoinHandle;
|
use tokio::task::{AbortHandle, JoinSet};
|
||||||
use tokio_tungstenite::tungstenite::Message;
|
use tokio_tungstenite::tungstenite::Message;
|
||||||
use tracing::{debug, error, info, warn};
|
use tracing::{debug, error, info, warn};
|
||||||
|
|
||||||
@@ -41,13 +41,25 @@ impl AsRef<[u8]> for BudgetedFramePayload {
|
|||||||
enum StreamDispatchStatus {
|
enum StreamDispatchStatus {
|
||||||
Delivered,
|
Delivered,
|
||||||
Closed,
|
Closed,
|
||||||
TimedOut,
|
Congested,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
struct StreamDispatchTarget {
|
struct StreamDispatchTarget {
|
||||||
body_tx: mpsc::Sender<Frame>,
|
body_tx: mpsc::Sender<Frame>,
|
||||||
response_window: Arc<StreamSendWindow>,
|
response_window: Arc<StreamSendWindow>,
|
||||||
|
handler: Option<AbortHandle>,
|
||||||
|
}
|
||||||
|
|
||||||
|
struct StreamCompletion {
|
||||||
|
stream_id: u32,
|
||||||
|
finished_tx: mpsc::UnboundedSender<u32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for StreamCompletion {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
let _ = self.finished_tx.send(self.stream_id);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// A request stream is identified by a non-zero id and may only be opened
|
/// A request stream is identified by a non-zero id and may only be opened
|
||||||
@@ -109,7 +121,7 @@ where
|
|||||||
// reopen the same id and bypass the stream admission limit.
|
// reopen the same id and bypass the stream admission limit.
|
||||||
let mut active_handler_ids: HashSet<u32> = HashSet::new();
|
let mut active_handler_ids: HashSet<u32> = HashSet::new();
|
||||||
// Track spawned stream handlers so we can wait for them on shutdown
|
// Track spawned stream handlers so we can wait for them on shutdown
|
||||||
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
|
let mut handler_handles = JoinSet::new();
|
||||||
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
|
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
|
||||||
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
|
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
|
||||||
let mut frames_since_cleanup: u32 = 0;
|
let mut frames_since_cleanup: u32 = 0;
|
||||||
@@ -121,30 +133,48 @@ where
|
|||||||
// Track last time we received any data to detect stale connections
|
// Track last time we received any data to detect stale connections
|
||||||
let mut last_data_at = tokio::time::Instant::now();
|
let mut last_data_at = tokio::time::Instant::now();
|
||||||
let mut draining = *drain.borrow();
|
let mut draining = *drain.borrow();
|
||||||
|
let mut drain_open = true;
|
||||||
|
let mut drain_deadline = draining.then(|| {
|
||||||
|
tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||||
|
});
|
||||||
|
let mut initial_window_bytes = state.config.tunnel_stream_initial_window_bytes;
|
||||||
|
let mut close_rx = frame_tx.subscribe_close();
|
||||||
|
|
||||||
let read_err = loop {
|
let read_err = loop {
|
||||||
|
if *close_rx.borrow() {
|
||||||
|
break None;
|
||||||
|
}
|
||||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||||
info!("tunnel drained after in-flight streams completed");
|
info!("tunnel drained after in-flight streams completed");
|
||||||
break None;
|
break None;
|
||||||
}
|
}
|
||||||
|
|
||||||
let msg_result = tokio::select! {
|
let msg_result = tokio::select! {
|
||||||
|
_ = close_rx.changed() => break None,
|
||||||
msg = ws_stream.next() => {
|
msg = ws_stream.next() => {
|
||||||
match msg {
|
match msg {
|
||||||
Some(r) => r,
|
Some(r) => r,
|
||||||
None => break None,
|
None => break None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
changed = drain.changed() => {
|
changed = drain.changed(), if drain_open => {
|
||||||
if changed.is_err() {
|
if changed.is_err() {
|
||||||
|
drain_open = false;
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
if *drain.borrow() {
|
if *drain.borrow() {
|
||||||
info!("tunnel drain requested, waiting for in-flight streams");
|
info!("tunnel drain requested, waiting for in-flight streams");
|
||||||
draining = true;
|
draining = true;
|
||||||
|
drain_deadline.get_or_insert_with(|| tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms));
|
||||||
}
|
}
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
_ = async {
|
||||||
|
match drain_deadline {
|
||||||
|
Some(deadline) => tokio::time::sleep_until(deadline).await,
|
||||||
|
None => std::future::pending().await,
|
||||||
|
}
|
||||||
|
} => break None,
|
||||||
finished = handler_finished_rx.recv() => {
|
finished = handler_finished_rx.recv() => {
|
||||||
if let Some(stream_id) = finished {
|
if let Some(stream_id) = finished {
|
||||||
active_handler_ids.remove(&stream_id);
|
active_handler_ids.remove(&stream_id);
|
||||||
@@ -238,20 +268,7 @@ where
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
if draining {
|
if draining {
|
||||||
if frame_tx
|
try_send_stream_error(&frame_tx, frame.stream_id, "tunnel draining");
|
||||||
.try_send(Frame::new(
|
|
||||||
frame.stream_id,
|
|
||||||
MsgType::StreamError,
|
|
||||||
0,
|
|
||||||
Bytes::from("tunnel draining"),
|
|
||||||
))
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
warn!(
|
|
||||||
stream_id = frame.stream_id,
|
|
||||||
"writer channel full, StreamError dropped during drain"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -263,6 +280,11 @@ where
|
|||||||
Ok(p) => p,
|
Ok(p) => p,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
|
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
|
||||||
|
try_send_stream_error(
|
||||||
|
&frame_tx,
|
||||||
|
frame.stream_id,
|
||||||
|
"invalid request metadata",
|
||||||
|
);
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
@@ -270,21 +292,11 @@ where
|
|||||||
Ok(m) => m,
|
Ok(m) => m,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
|
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
|
||||||
// Use try_send to avoid blocking the read loop
|
try_send_stream_error(
|
||||||
if frame_tx
|
&frame_tx,
|
||||||
.try_send(Frame::new(
|
frame.stream_id,
|
||||||
frame.stream_id,
|
"invalid request metadata",
|
||||||
MsgType::StreamError,
|
);
|
||||||
0,
|
|
||||||
Bytes::from(format!("invalid request metadata: {e}")),
|
|
||||||
))
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
warn!(
|
|
||||||
stream_id = frame.stream_id,
|
|
||||||
"writer channel full, StreamError dropped"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
@@ -294,33 +306,27 @@ where
|
|||||||
stream_id = frame.stream_id,
|
stream_id = frame.stream_id,
|
||||||
"max concurrent streams reached"
|
"max concurrent streams reached"
|
||||||
);
|
);
|
||||||
if frame_tx
|
try_send_stream_error(
|
||||||
.try_send(Frame::new(
|
&frame_tx,
|
||||||
frame.stream_id,
|
frame.stream_id,
|
||||||
MsgType::StreamError,
|
"max concurrent streams reached",
|
||||||
0,
|
);
|
||||||
Bytes::from("max concurrent streams reached"),
|
|
||||||
))
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
warn!(
|
|
||||||
stream_id = frame.stream_id,
|
|
||||||
"writer channel full, StreamError dropped"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create body channel and spawn handler
|
// Create body channel and spawn handler
|
||||||
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
|
let body_capacity = (initial_window_bytes as usize)
|
||||||
let response_window = Arc::new(StreamSendWindow::new(
|
.div_ceil(32 * 1024)
|
||||||
state.config.tunnel_stream_initial_window_bytes,
|
.saturating_add(1)
|
||||||
));
|
.max(64);
|
||||||
|
let (body_tx, body_rx) = mpsc::channel::<Frame>(body_capacity);
|
||||||
|
let response_window = Arc::new(StreamSendWindow::new(initial_window_bytes));
|
||||||
streams.insert(
|
streams.insert(
|
||||||
frame.stream_id,
|
frame.stream_id,
|
||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx,
|
body_tx,
|
||||||
response_window: Arc::clone(&response_window),
|
response_window: Arc::clone(&response_window),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
active_handler_ids.insert(frame.stream_id);
|
active_handler_ids.insert(frame.stream_id);
|
||||||
@@ -329,9 +335,13 @@ where
|
|||||||
let state_clone = Arc::clone(&state);
|
let state_clone = Arc::clone(&state);
|
||||||
let server_clone = Arc::clone(&server);
|
let server_clone = Arc::clone(&server);
|
||||||
let tx_clone = frame_tx.clone();
|
let tx_clone = frame_tx.clone();
|
||||||
let finished_tx = handler_finished_tx.clone();
|
|
||||||
let sid = frame.stream_id;
|
let sid = frame.stream_id;
|
||||||
let handle = tokio::spawn(async move {
|
let completion = StreamCompletion {
|
||||||
|
stream_id: sid,
|
||||||
|
finished_tx: handler_finished_tx.clone(),
|
||||||
|
};
|
||||||
|
let handle = handler_handles.spawn(async move {
|
||||||
|
let _completion = completion;
|
||||||
stream_handler::handle_stream(
|
stream_handler::handle_stream(
|
||||||
state_clone,
|
state_clone,
|
||||||
server_clone,
|
server_clone,
|
||||||
@@ -342,9 +352,8 @@ where
|
|||||||
response_window,
|
response_window,
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
let _ = finished_tx.send(sid);
|
|
||||||
});
|
});
|
||||||
handler_handles.push(handle);
|
streams.get_mut(&sid).expect("new stream exists").handler = Some(handle);
|
||||||
|
|
||||||
if request_headers_end_stream {
|
if request_headers_end_stream {
|
||||||
if let Some(target) = streams.get(&sid) {
|
if let Some(target) = streams.get(&sid) {
|
||||||
@@ -365,19 +374,21 @@ where
|
|||||||
let is_end = frame.is_end_stream();
|
let is_end = frame.is_end_stream();
|
||||||
let sid = frame.stream_id;
|
let sid = frame.stream_id;
|
||||||
let dispatch = dispatch_stream_frame(&target.body_tx, frame).await;
|
let dispatch = dispatch_stream_frame(&target.body_tx, frame).await;
|
||||||
if dispatch != StreamDispatchStatus::Delivered {
|
if dispatch == StreamDispatchStatus::Congested {
|
||||||
streams.remove(&sid);
|
if let Some(target) = streams.remove(&sid) {
|
||||||
if dispatch == StreamDispatchStatus::TimedOut {
|
if let Some(handler) = target.handler {
|
||||||
server.tunnel_metrics.record_error(
|
handler.abort();
|
||||||
"stream_dispatch_timeout",
|
}
|
||||||
&format!("request body dispatch timed out for stream {}", sid),
|
|
||||||
);
|
|
||||||
try_send_stream_error(
|
|
||||||
&frame_tx,
|
|
||||||
sid,
|
|
||||||
"tunnel request body dispatch stalled",
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
server.tunnel_metrics.record_error(
|
||||||
|
"stream_dispatch_timeout",
|
||||||
|
&format!("request body dispatch congested for stream {}", sid),
|
||||||
|
);
|
||||||
|
try_send_stream_error(
|
||||||
|
&frame_tx,
|
||||||
|
sid,
|
||||||
|
"tunnel request body dispatch stalled",
|
||||||
|
);
|
||||||
if is_end && draining && streams.is_empty() && active_handler_ids.is_empty()
|
if is_end && draining && streams.is_empty() && active_handler_ids.is_empty()
|
||||||
{
|
{
|
||||||
info!("tunnel drained after request body completion");
|
info!("tunnel drained after request body completion");
|
||||||
@@ -387,10 +398,29 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::StreamEnd | MsgType::StreamError | MsgType::ResetStream => {
|
MsgType::StreamEnd => {
|
||||||
|
if let Some(target) = streams.get(&frame.stream_id) {
|
||||||
|
if dispatch_stream_frame(&target.body_tx, frame.clone()).await
|
||||||
|
== StreamDispatchStatus::Congested
|
||||||
|
{
|
||||||
|
if let Some(handler) = &target.handler {
|
||||||
|
handler.abort();
|
||||||
|
}
|
||||||
|
try_send_stream_error(
|
||||||
|
&frame_tx,
|
||||||
|
frame.stream_id,
|
||||||
|
"tunnel request body dispatch stalled",
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
MsgType::StreamError | MsgType::ResetStream => {
|
||||||
// Client-side cancellation or end
|
// Client-side cancellation or end
|
||||||
if let Some(target) = streams.remove(&frame.stream_id) {
|
if let Some(target) = streams.remove(&frame.stream_id) {
|
||||||
let _ = dispatch_stream_frame(&target.body_tx, frame).await;
|
if let Some(handler) = target.handler {
|
||||||
|
handler.abort();
|
||||||
|
}
|
||||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||||
info!("tunnel drained after stream termination");
|
info!("tunnel drained after stream termination");
|
||||||
break None;
|
break None;
|
||||||
@@ -409,12 +439,22 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
MsgType::HeartbeatAck => {
|
MsgType::HeartbeatAck => {
|
||||||
heartbeat.on_ack(frame.payload).await;
|
heartbeat.on_ack(frame.payload);
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::GoAway => {
|
MsgType::GoAway => {
|
||||||
info!("received GOAWAY");
|
info!("received GOAWAY");
|
||||||
break None;
|
draining = true;
|
||||||
|
let deadline_ms =
|
||||||
|
serde_json::from_slice::<aether_contracts::tunnel::GoAwayPayload>(
|
||||||
|
&frame.payload,
|
||||||
|
)
|
||||||
|
.map(|payload| payload.drain_deadline_ms)
|
||||||
|
.unwrap_or(state.config.tunnel_drain_deadline_ms)
|
||||||
|
.min(state.config.tunnel_drain_deadline_ms);
|
||||||
|
drain_deadline.get_or_insert_with(|| {
|
||||||
|
tokio::time::Instant::now() + Duration::from_millis(deadline_ms)
|
||||||
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::WindowUpdate => {
|
MsgType::WindowUpdate => {
|
||||||
@@ -433,7 +473,31 @@ where
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::Hello | MsgType::Settings | MsgType::LoadReport => {
|
MsgType::Settings => {
|
||||||
|
if frame.stream_id != 0 || frame.flags != 0 {
|
||||||
|
break None;
|
||||||
|
}
|
||||||
|
let settings = serde_json::from_slice::<aether_contracts::tunnel::SettingsPayload>(
|
||||||
|
&frame.payload,
|
||||||
|
)
|
||||||
|
.ok()
|
||||||
|
.filter(|settings| settings.is_valid());
|
||||||
|
let Some(settings) = settings else {
|
||||||
|
warn!("invalid tunnel SETTINGS");
|
||||||
|
break None;
|
||||||
|
};
|
||||||
|
if !streams.is_empty()
|
||||||
|
&& settings.initial_stream_window_bytes != initial_window_bytes
|
||||||
|
{
|
||||||
|
warn!("tunnel SETTINGS changed with active streams");
|
||||||
|
break None;
|
||||||
|
}
|
||||||
|
initial_window_bytes = settings
|
||||||
|
.initial_stream_window_bytes
|
||||||
|
.min(state.config.tunnel_stream_initial_window_bytes);
|
||||||
|
}
|
||||||
|
|
||||||
|
MsgType::Hello | MsgType::LoadReport => {
|
||||||
debug!(
|
debug!(
|
||||||
msg_type = ?frame.msg_type,
|
msg_type = ?frame.msg_type,
|
||||||
stream_id = frame.stream_id,
|
stream_id = frame.stream_id,
|
||||||
@@ -455,7 +519,7 @@ where
|
|||||||
// Trigger every 64 frames OR when the count exceeds max_streams.
|
// Trigger every 64 frames OR when the count exceeds max_streams.
|
||||||
frames_since_cleanup += 1;
|
frames_since_cleanup += 1;
|
||||||
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
|
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
|
||||||
handler_handles.retain(|h| !h.is_finished());
|
while handler_handles.try_join_next().is_some() {}
|
||||||
frames_since_cleanup = 0;
|
frames_since_cleanup = 0;
|
||||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||||
info!("tunnel drained after cleanup");
|
info!("tunnel drained after cleanup");
|
||||||
@@ -467,9 +531,7 @@ where
|
|||||||
// Drop body senders so stream handlers waiting on body_rx will unblock
|
// Drop body senders so stream handlers waiting on body_rx will unblock
|
||||||
streams.clear();
|
streams.clear();
|
||||||
|
|
||||||
// Wait for active stream handlers to finish so their frame_tx clones
|
handler_handles.shutdown().await;
|
||||||
// are dropped before the writer closes the sink.
|
|
||||||
drain_handlers(handler_handles).await;
|
|
||||||
|
|
||||||
match read_err {
|
match read_err {
|
||||||
Some(e) => Err(e.into()),
|
Some(e) => Err(e.into()),
|
||||||
@@ -478,30 +540,13 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
|
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
|
||||||
let stream_id = frame.stream_id;
|
let Some(frame) = attach_request_body_queue_budget(frame).await else {
|
||||||
let dispatched = tokio::time::timeout(stream_frame_dispatch_timeout(), async {
|
return StreamDispatchStatus::Congested;
|
||||||
let frame = attach_request_body_queue_budget(frame).await?;
|
};
|
||||||
tx.send(frame).await.ok()?;
|
match tx.try_send(frame) {
|
||||||
Some(())
|
Ok(()) => StreamDispatchStatus::Delivered,
|
||||||
})
|
Err(mpsc::error::TrySendError::Closed(_)) => StreamDispatchStatus::Closed,
|
||||||
.await;
|
Err(mpsc::error::TrySendError::Full(_)) => StreamDispatchStatus::Congested,
|
||||||
match dispatched {
|
|
||||||
Ok(Some(())) => StreamDispatchStatus::Delivered,
|
|
||||||
Ok(None) => {
|
|
||||||
warn!(
|
|
||||||
stream_id,
|
|
||||||
"stream handler channel or request body budget closed while dispatching tunnel frame"
|
|
||||||
);
|
|
||||||
StreamDispatchStatus::Closed
|
|
||||||
}
|
|
||||||
Err(_) => {
|
|
||||||
warn!(
|
|
||||||
stream_id,
|
|
||||||
timeout_ms = stream_frame_dispatch_timeout().as_millis(),
|
|
||||||
"stream handler channel blocked while dispatching tunnel frame"
|
|
||||||
);
|
|
||||||
StreamDispatchStatus::TimedOut
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -523,7 +568,7 @@ async fn attach_request_body_queue_budget_with(
|
|||||||
return Some(frame);
|
return Some(frame);
|
||||||
}
|
}
|
||||||
let permits = request_body_queue_permits(&frame, budget_bytes)?;
|
let permits = request_body_queue_permits(&frame, budget_bytes)?;
|
||||||
let permit = budget.acquire_many_owned(permits).await.ok()?;
|
let permit = budget.try_acquire_many_owned(permits).ok()?;
|
||||||
frame.payload = Bytes::from_owner(BudgetedFramePayload {
|
frame.payload = Bytes::from_owner(BudgetedFramePayload {
|
||||||
bytes: frame.payload,
|
bytes: frame.payload,
|
||||||
_permit: permit,
|
_permit: permit,
|
||||||
@@ -549,20 +594,6 @@ fn request_body_queue_permits(frame: &Frame, budget_bytes: usize) -> Option<u32>
|
|||||||
u32::try_from(retained_bytes).ok()
|
u32::try_from(retained_bytes).ok()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Bound how long a single stream handler is allowed to block the shared
|
|
||||||
/// WebSocket read loop while receiving request-body frames.
|
|
||||||
fn stream_frame_dispatch_timeout() -> Duration {
|
|
||||||
#[cfg(test)]
|
|
||||||
{
|
|
||||||
Duration::from_millis(25)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(not(test))]
|
|
||||||
{
|
|
||||||
Duration::from_millis(500)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) {
|
fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) {
|
||||||
if frame_tx
|
if frame_tx
|
||||||
.try_send(Frame::new(
|
.try_send(Frame::new(
|
||||||
@@ -573,6 +604,7 @@ fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'stat
|
|||||||
))
|
))
|
||||||
.is_err()
|
.is_err()
|
||||||
{
|
{
|
||||||
|
frame_tx.close();
|
||||||
warn!(
|
warn!(
|
||||||
stream_id,
|
stream_id,
|
||||||
"writer channel full, StreamError dropped while aborting stalled stream"
|
"writer channel full, StreamError dropped while aborting stalled stream"
|
||||||
@@ -587,21 +619,6 @@ fn prune_closed_stream_senders(streams: &mut HashMap<u32, StreamDispatchTarget>)
|
|||||||
before.saturating_sub(streams.len())
|
before.saturating_sub(streams.len())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Wait for all active stream handlers to finish (with a timeout).
|
|
||||||
async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
|
|
||||||
if handles.is_empty() {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
let count = handles.len();
|
|
||||||
debug!(count, "waiting for active stream handlers to finish");
|
|
||||||
let _ = tokio::time::timeout(Duration::from_secs(30), async {
|
|
||||||
for h in handles {
|
|
||||||
let _ = h.await;
|
|
||||||
}
|
|
||||||
})
|
|
||||||
.await;
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
@@ -633,7 +650,7 @@ mod tests {
|
|||||||
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
stalled_send.await.expect("dispatch task should join"),
|
stalled_send.await.expect("dispatch task should join"),
|
||||||
StreamDispatchStatus::TimedOut
|
StreamDispatchStatus::Congested
|
||||||
);
|
);
|
||||||
|
|
||||||
let retained = rx
|
let retained = rx
|
||||||
@@ -737,6 +754,7 @@ mod tests {
|
|||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx: closed_tx,
|
body_tx: closed_tx,
|
||||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
@@ -744,6 +762,7 @@ mod tests {
|
|||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx: open_tx,
|
body_tx: open_tx,
|
||||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
]);
|
]);
|
||||||
@@ -763,6 +782,7 @@ mod tests {
|
|||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx: tx,
|
body_tx: tx,
|
||||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
)]);
|
)]);
|
||||||
let mut active_handler_ids = HashSet::from([7]);
|
let mut active_handler_ids = HashSet::from([7]);
|
||||||
|
|||||||
@@ -31,14 +31,22 @@ enum AckDecision {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
||||||
#[derive(Clone)]
|
|
||||||
pub struct HeartbeatHandle {
|
pub struct HeartbeatHandle {
|
||||||
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
|
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
|
||||||
|
task: Option<tokio::task::JoinHandle<()>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl HeartbeatHandle {
|
impl HeartbeatHandle {
|
||||||
pub async fn on_ack(&self, payload: Bytes) {
|
pub fn on_ack(&self, payload: Bytes) {
|
||||||
let _ = self.ack_tx.send(payload).await;
|
let _ = self.ack_tx.try_send(payload);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for HeartbeatHandle {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
if let Some(task) = self.task.take() {
|
||||||
|
task.abort();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -48,7 +56,7 @@ impl HeartbeatHandle {
|
|||||||
pub fn spawn_noop() -> HeartbeatHandle {
|
pub fn spawn_noop() -> HeartbeatHandle {
|
||||||
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
|
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
|
||||||
// receiver is immediately dropped; on_ack() calls will silently fail
|
// receiver is immediately dropped; on_ack() calls will silently fail
|
||||||
HeartbeatHandle { ack_tx }
|
HeartbeatHandle { ack_tx, task: None }
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, Default)]
|
#[derive(Debug, Clone, Copy, Default)]
|
||||||
@@ -74,7 +82,7 @@ pub fn spawn(
|
|||||||
) -> HeartbeatHandle {
|
) -> HeartbeatHandle {
|
||||||
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
|
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
|
||||||
|
|
||||||
tokio::spawn(async move {
|
let task = tokio::spawn(async move {
|
||||||
// Read initial interval from dynamic config (may be updated by remote config).
|
// Read initial interval from dynamic config (may be updated by remote config).
|
||||||
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
||||||
let mut current_interval = initial_interval;
|
let mut current_interval = initial_interval;
|
||||||
@@ -151,7 +159,8 @@ pub fn spawn(
|
|||||||
current_interval = new_interval;
|
current_interval = new_interval;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(ack_payload) = ack_rx.recv() => {
|
ack_payload = ack_rx.recv() => {
|
||||||
|
let Some(ack_payload) = ack_payload else { break; };
|
||||||
match handle_ack(&server, &ack_payload) {
|
match handle_ack(&server, &ack_payload) {
|
||||||
AckDecision::Accept {
|
AckDecision::Accept {
|
||||||
heartbeat_id: ack_id,
|
heartbeat_id: ack_id,
|
||||||
@@ -179,7 +188,10 @@ pub fn spawn(
|
|||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
|
||||||
HeartbeatHandle { ack_tx }
|
HeartbeatHandle {
|
||||||
|
ack_tx,
|
||||||
|
task: Some(task),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn build_heartbeat_payload(
|
async fn build_heartbeat_payload(
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ pub mod dispatcher;
|
|||||||
pub mod heartbeat;
|
pub mod heartbeat;
|
||||||
pub mod protocol;
|
pub mod protocol;
|
||||||
pub mod stream_handler;
|
pub mod stream_handler;
|
||||||
|
mod task;
|
||||||
pub mod writer;
|
pub mod writer;
|
||||||
|
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
@@ -332,9 +333,9 @@ mod tests {
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
|
||||||
tokio::time::sleep(Duration::from_millis(200)).await;
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
|
let _ = (&mut gateway_handle).await;
|
||||||
|
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
||||||
|
|
||||||
let (_restarted_gateway_state, restarted_gateway_handle) =
|
let (_restarted_gateway_state, restarted_gateway_handle) =
|
||||||
start_gateway_on_port_retry(gateway_port)
|
start_gateway_on_port_retry(gateway_port)
|
||||||
@@ -349,6 +350,7 @@ mod tests {
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
|
assert!(server.tunnel_metrics.snapshot().connect_successes >= 2);
|
||||||
let _ = shutdown_tx.send(true);
|
let _ = shutdown_tx.send(true);
|
||||||
tokio::time::timeout(Duration::from_secs(5), tunnel_task)
|
tokio::time::timeout(Duration::from_secs(5), tunnel_task)
|
||||||
.await
|
.await
|
||||||
@@ -380,7 +382,17 @@ mod tests {
|
|||||||
gateway_base_url: &str,
|
gateway_base_url: &str,
|
||||||
node_id: &str,
|
node_id: &str,
|
||||||
) -> Option<(StatusCode, String)> {
|
) -> Option<(StatusCode, String)> {
|
||||||
let payload = relay_probe_envelope();
|
let response = relay_response(gateway_base_url, node_id, relay_probe_envelope()).await?;
|
||||||
|
let status = response.status();
|
||||||
|
let body = response.text().await.unwrap_or_default();
|
||||||
|
Some((status, body))
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn relay_response(
|
||||||
|
gateway_base_url: &str,
|
||||||
|
node_id: &str,
|
||||||
|
payload: Vec<u8>,
|
||||||
|
) -> Option<reqwest::Response> {
|
||||||
let timestamp = SystemTime::now()
|
let timestamp = SystemTime::now()
|
||||||
.duration_since(UNIX_EPOCH)
|
.duration_since(UNIX_EPOCH)
|
||||||
.expect("test clock should be after epoch")
|
.expect("test clock should be after epoch")
|
||||||
@@ -398,7 +410,7 @@ mod tests {
|
|||||||
&nonce,
|
&nonce,
|
||||||
&digest,
|
&digest,
|
||||||
);
|
);
|
||||||
let response = reqwest::Client::new()
|
reqwest::Client::new()
|
||||||
.post(format!(
|
.post(format!(
|
||||||
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
|
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
|
||||||
))
|
))
|
||||||
@@ -421,10 +433,7 @@ mod tests {
|
|||||||
.body(payload)
|
.body(payload)
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.ok()?;
|
.ok()
|
||||||
let status = response.status();
|
|
||||||
let body = response.text().await.unwrap_or_default();
|
|
||||||
Some((status, body))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn relay_probe_envelope() -> Vec<u8> {
|
fn relay_probe_envelope() -> Vec<u8> {
|
||||||
@@ -456,30 +465,150 @@ mod tests {
|
|||||||
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
|
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
|
||||||
// The embedded gateway now fails closed when relay authentication is
|
// The embedded gateway now fails closed when relay authentication is
|
||||||
// not configured. Keep this integration fixture explicitly authenticated.
|
// not configured. Keep this integration fixture explicitly authenticated.
|
||||||
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
|
||||||
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
let state = {
|
||||||
std::env::set_var(
|
let _guard = ENV_LOCK.lock().unwrap();
|
||||||
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
||||||
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
||||||
);
|
std::env::set_var(
|
||||||
std::env::set_var(
|
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
||||||
"AETHER_GATEWAY_INSTANCE_ID",
|
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
||||||
"tunnel-reconnect-test-gateway",
|
);
|
||||||
);
|
std::env::set_var(
|
||||||
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
"AETHER_GATEWAY_INSTANCE_ID",
|
||||||
aether_gateway::configure_test_tunnel_security(
|
"tunnel-reconnect-test-gateway",
|
||||||
&mut state,
|
);
|
||||||
"node-recovery",
|
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
||||||
"test-generation-1",
|
aether_gateway::configure_test_tunnel_security(
|
||||||
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
&mut state,
|
||||||
);
|
"node-recovery",
|
||||||
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
"test-generation-1",
|
||||||
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
||||||
|
);
|
||||||
|
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
||||||
|
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
||||||
|
state
|
||||||
|
};
|
||||||
let router = build_router_with_state(state.clone());
|
let router = build_router_with_state(state.clone());
|
||||||
let handle = spawn_router_on_port(port, router).await?;
|
let handle = spawn_router_on_port(port, router).await?;
|
||||||
Ok((state, handle))
|
Ok((state, handle))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn negotiated_small_window_streams_large_responses_and_cancels_idle_upstream() {
|
||||||
|
use axum::body::{Body, Bytes};
|
||||||
|
use axum::routing::get;
|
||||||
|
use futures_util::StreamExt;
|
||||||
|
|
||||||
|
ensure_rustls_provider();
|
||||||
|
let upstream_port = reserve_local_port().unwrap();
|
||||||
|
let upstream = Router::new()
|
||||||
|
.route(
|
||||||
|
"/large",
|
||||||
|
get(|| async { Body::from(vec![b'x'; 2 * 1024 * 1024]) }),
|
||||||
|
)
|
||||||
|
.route(
|
||||||
|
"/idle",
|
||||||
|
get(|| async {
|
||||||
|
let first = futures_util::stream::once(async {
|
||||||
|
Ok::<_, std::io::Error>(Bytes::from_static(b"data: started\n\n"))
|
||||||
|
});
|
||||||
|
(
|
||||||
|
[("content-type", "text/event-stream")],
|
||||||
|
Body::from_stream(first.chain(futures_util::stream::pending())),
|
||||||
|
)
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
let upstream_task = super::task::SessionTask::new(
|
||||||
|
spawn_router_on_port(upstream_port, upstream).await.unwrap(),
|
||||||
|
);
|
||||||
|
let gateway_port = reserve_local_port().unwrap();
|
||||||
|
let gateway_url = format!("http://127.0.0.1:{gateway_port}");
|
||||||
|
let (_, gateway_task) = start_gateway_on_port(gateway_port).await.unwrap();
|
||||||
|
let gateway_task = super::task::SessionTask::new(gateway_task);
|
||||||
|
let mut config = sample_config(&gateway_url);
|
||||||
|
config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired;
|
||||||
|
config.tunnel_encryption_key = Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".into());
|
||||||
|
config.tunnel_stream_initial_window_bytes = 512 * 1024;
|
||||||
|
config.tunnel_drain_deadline_ms = 100;
|
||||||
|
config.allow_private_targets = true;
|
||||||
|
config.allowed_ports.push(upstream_port);
|
||||||
|
let state = sample_state(config);
|
||||||
|
let server = sample_server(&state, "node-recovery");
|
||||||
|
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||||
|
let (_drain_tx, drain_rx) = watch::channel(false);
|
||||||
|
let tunnel_task = super::task::SessionTask::new(tokio::spawn({
|
||||||
|
let state = Arc::clone(&state);
|
||||||
|
let server = Arc::clone(&server);
|
||||||
|
async move {
|
||||||
|
run(&state, &server, 0, shutdown_rx, drain_rx).await;
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
wait_until_relay_status(&gateway_url, "node-recovery", StatusCode::GATEWAY_TIMEOUT).await;
|
||||||
|
|
||||||
|
let envelope = |path: &str| {
|
||||||
|
let mut meta: protocol::RequestMeta =
|
||||||
|
serde_json::from_slice(&relay_probe_envelope()[4..]).unwrap();
|
||||||
|
meta.url = format!("http://127.0.0.1:{upstream_port}/{path}");
|
||||||
|
meta.stream = true;
|
||||||
|
meta.timeout = 10;
|
||||||
|
meta.stream_first_byte_timeout_ms = Some(10_000);
|
||||||
|
let encoded = serde_json::to_vec(&meta).unwrap();
|
||||||
|
let mut result = (encoded.len() as u32).to_be_bytes().to_vec();
|
||||||
|
result.extend(encoded);
|
||||||
|
result
|
||||||
|
};
|
||||||
|
let response = relay_response(&gateway_url, "node-recovery", envelope("large"))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
let body = tokio::time::timeout(Duration::from_secs(10), response.bytes())
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(body.len(), 2 * 1024 * 1024);
|
||||||
|
assert!(body.iter().all(|byte| *byte == b'x'));
|
||||||
|
|
||||||
|
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
response.chunk().await.unwrap().unwrap(),
|
||||||
|
"data: started\n\n"
|
||||||
|
);
|
||||||
|
drop(response);
|
||||||
|
tokio::time::timeout(Duration::from_secs(3), async {
|
||||||
|
while server
|
||||||
|
.active_connections
|
||||||
|
.load(std::sync::atomic::Ordering::Acquire)
|
||||||
|
!= 0
|
||||||
|
{
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.expect("cancelled SSE must release the upstream handler");
|
||||||
|
|
||||||
|
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert!(response.chunk().await.unwrap().is_some());
|
||||||
|
shutdown_tx.send(true).unwrap();
|
||||||
|
tokio::time::timeout(Duration::from_secs(3), tunnel_task)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
server
|
||||||
|
.active_connections
|
||||||
|
.load(std::sync::atomic::Ordering::Acquire),
|
||||||
|
0
|
||||||
|
);
|
||||||
|
drop(response);
|
||||||
|
drop(gateway_task);
|
||||||
|
drop(upstream_task);
|
||||||
|
}
|
||||||
|
|
||||||
fn restore_test_env(key: &str, value: Option<std::ffi::OsString>) {
|
fn restore_test_env(key: &str, value: Option<std::ffi::OsString>) {
|
||||||
if let Some(value) = value {
|
if let Some(value) = value {
|
||||||
std::env::set_var(key, value);
|
std::env::set_var(key, value);
|
||||||
|
|||||||
@@ -52,6 +52,7 @@ static REDIRECT_REPLAY_BUFFERED_BYTES: AtomicUsize = AtomicUsize::new(0);
|
|||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub(crate) struct StreamSendWindow {
|
pub(crate) struct StreamSendWindow {
|
||||||
|
initial_window_bytes: u32,
|
||||||
available: Mutex<u64>,
|
available: Mutex<u64>,
|
||||||
notify: Notify,
|
notify: Notify,
|
||||||
}
|
}
|
||||||
@@ -59,6 +60,7 @@ pub(crate) struct StreamSendWindow {
|
|||||||
impl StreamSendWindow {
|
impl StreamSendWindow {
|
||||||
pub(crate) fn new(initial_window_bytes: u32) -> Self {
|
pub(crate) fn new(initial_window_bytes: u32) -> Self {
|
||||||
Self {
|
Self {
|
||||||
|
initial_window_bytes: initial_window_bytes.max(1),
|
||||||
available: Mutex::new(u64::from(initial_window_bytes.max(1))),
|
available: Mutex::new(u64::from(initial_window_bytes.max(1))),
|
||||||
notify: Notify::new(),
|
notify: Notify::new(),
|
||||||
}
|
}
|
||||||
@@ -82,6 +84,9 @@ impl StreamSendWindow {
|
|||||||
let requested = bytes as u64;
|
let requested = bytes as u64;
|
||||||
let started_at = Instant::now();
|
let started_at = Instant::now();
|
||||||
loop {
|
loop {
|
||||||
|
let notified = self.notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
{
|
{
|
||||||
let mut available = self.available.lock().expect("stream window lock poisoned");
|
let mut available = self.available.lock().expect("stream window lock poisoned");
|
||||||
if *available >= requested {
|
if *available >= requested {
|
||||||
@@ -93,10 +98,7 @@ impl StreamSendWindow {
|
|||||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||||
return Err(());
|
return Err(());
|
||||||
};
|
};
|
||||||
if tokio::time::timeout(remaining, self.notify.notified())
|
if tokio::time::timeout(remaining, notified).await.is_err() {
|
||||||
.await
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
return Err(());
|
return Err(());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -173,31 +175,33 @@ fn safe_stream_error_message(message: &str) -> &'static str {
|
|||||||
"upstream request failed"
|
"upstream request failed"
|
||||||
}
|
}
|
||||||
|
|
||||||
fn try_send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) {
|
async fn send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) -> bool {
|
||||||
if bytes == 0 {
|
if bytes == 0 {
|
||||||
return;
|
return true;
|
||||||
}
|
}
|
||||||
let delta = bytes.min(u32::MAX as usize) as u32;
|
let delta = bytes.min(u32::MAX as usize) as u32;
|
||||||
if frame_tx
|
if matches!(
|
||||||
.try_send(TunnelFrame::new(
|
tokio::time::timeout(
|
||||||
stream_id,
|
FLOW_CONTROL_WAIT_TIMEOUT,
|
||||||
MsgType::WindowUpdate,
|
frame_tx.send(TunnelFrame::new(
|
||||||
0,
|
stream_id,
|
||||||
Bytes::from(
|
MsgType::WindowUpdate,
|
||||||
serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload {
|
0,
|
||||||
delta_bytes: delta,
|
Bytes::from(
|
||||||
})
|
serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload {
|
||||||
.expect("window update payload should serialize"),
|
delta_bytes: delta,
|
||||||
),
|
})
|
||||||
))
|
.expect("window update payload should serialize"),
|
||||||
.is_err()
|
),
|
||||||
{
|
))
|
||||||
warn!(
|
)
|
||||||
stream_id,
|
.await,
|
||||||
delta_bytes = delta,
|
Ok(Ok(()))
|
||||||
"writer channel full, WINDOW_UPDATE dropped"
|
) {
|
||||||
);
|
return true;
|
||||||
}
|
}
|
||||||
|
frame_tx.close();
|
||||||
|
false
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Match reqwest's default redirect budget so direct execution and tunnel relay
|
/// Match reqwest's default redirect budget so direct execution and tunnel relay
|
||||||
@@ -242,6 +246,23 @@ enum ReplayableRequestBody {
|
|||||||
struct PreparedRequestBody {
|
struct PreparedRequestBody {
|
||||||
first_request_body: Option<upstream_client::UpstreamRequestBody>,
|
first_request_body: Option<upstream_client::UpstreamRequestBody>,
|
||||||
replay_body: ReplayableRequestBody,
|
replay_body: ReplayableRequestBody,
|
||||||
|
spool_task: Option<tokio::task::JoinHandle<()>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for PreparedRequestBody {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
if let Some(task) = self.spool_task.take() {
|
||||||
|
task.abort();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
struct ActiveStreamGuard(Arc<ServerContext>);
|
||||||
|
|
||||||
|
impl Drop for ActiveStreamGuard {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
self.0.active_connections.fetch_sub(1, Ordering::Release);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
@@ -331,7 +352,10 @@ impl hyper::body::Body for ReplayRequestBody {
|
|||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
enum SpoolBodyEvent {
|
enum SpoolBodyEvent {
|
||||||
Data(Bytes),
|
Data {
|
||||||
|
payload: Bytes,
|
||||||
|
credit_returned: bool,
|
||||||
|
},
|
||||||
Error(String),
|
Error(String),
|
||||||
End,
|
End,
|
||||||
}
|
}
|
||||||
@@ -563,8 +587,9 @@ impl RequestBodyReplayState {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn push_chunk(&self, payload: Bytes) {
|
fn push_chunk(&self, payload: Bytes) -> bool {
|
||||||
let mut disable_replay = false;
|
let mut disable_replay = false;
|
||||||
|
let mut retained = false;
|
||||||
let mut state = self.state.lock().expect("request body replay state lock");
|
let mut state = self.state.lock().expect("request body replay state lock");
|
||||||
if let RequestBodyReplayStatus::Collecting {
|
if let RequestBodyReplayStatus::Collecting {
|
||||||
chunks,
|
chunks,
|
||||||
@@ -577,7 +602,7 @@ impl RequestBodyReplayState {
|
|||||||
drop(state);
|
drop(state);
|
||||||
self.release_reserved_bytes();
|
self.release_reserved_bytes();
|
||||||
self.ready.notify_waiters();
|
self.ready.notify_waiters();
|
||||||
return;
|
return false;
|
||||||
};
|
};
|
||||||
let accounted_bytes = payload.len().checked_add(std::mem::size_of::<Bytes>());
|
let accounted_bytes = payload.len().checked_add(std::mem::size_of::<Bytes>());
|
||||||
if next_len > self.budget_bytes
|
if next_len > self.budget_bytes
|
||||||
@@ -590,6 +615,7 @@ impl RequestBodyReplayState {
|
|||||||
} else {
|
} else {
|
||||||
*buffered_len = next_len;
|
*buffered_len = next_len;
|
||||||
chunks.push(payload);
|
chunks.push(payload);
|
||||||
|
retained = true;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
drop(state);
|
drop(state);
|
||||||
@@ -597,6 +623,7 @@ impl RequestBodyReplayState {
|
|||||||
self.release_reserved_bytes();
|
self.release_reserved_bytes();
|
||||||
self.ready.notify_waiters();
|
self.ready.notify_waiters();
|
||||||
}
|
}
|
||||||
|
retained
|
||||||
}
|
}
|
||||||
|
|
||||||
fn try_reserve_bytes(&self, bytes: usize) -> bool {
|
fn try_reserve_bytes(&self, bytes: usize) -> bool {
|
||||||
@@ -951,9 +978,6 @@ pub(super) fn decode_request_body_frame(frame: TunnelFrame) -> Result<Bytes, std
|
|||||||
Ok(frame.payload)
|
Ok(frame.payload)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Drain tunnel body frames on a detached task so the shared dispatcher is no
|
|
||||||
// longer coupled to upstream body polling. Redirect replay retains a bounded
|
|
||||||
// copy; crossing either replay budget only disables replay for this request.
|
|
||||||
fn prepare_request_body(
|
fn prepare_request_body(
|
||||||
stream_id: u32,
|
stream_id: u32,
|
||||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||||
@@ -973,19 +997,20 @@ fn prepare_request_body(
|
|||||||
None => ReplayableRequestBody::NonReplayable,
|
None => ReplayableRequestBody::NonReplayable,
|
||||||
};
|
};
|
||||||
|
|
||||||
tokio::spawn(spool_request_body(
|
let spool_task = tokio::spawn(spool_request_body(
|
||||||
stream_id,
|
stream_id,
|
||||||
body_rx,
|
body_rx,
|
||||||
spool_tx,
|
spool_tx,
|
||||||
replay_state,
|
replay_state,
|
||||||
body_size,
|
body_size,
|
||||||
deadline,
|
deadline,
|
||||||
frame_tx,
|
frame_tx.clone(),
|
||||||
));
|
));
|
||||||
|
|
||||||
PreparedRequestBody {
|
PreparedRequestBody {
|
||||||
first_request_body: Some(build_spooled_request_body(spool_rx)),
|
first_request_body: Some(build_spooled_request_body(spool_rx, stream_id, frame_tx)),
|
||||||
replay_body,
|
replay_body,
|
||||||
|
spool_task: Some(spool_task),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1001,6 +1026,7 @@ fn prepare_bodyless_request_body(
|
|||||||
} else {
|
} else {
|
||||||
ReplayableRequestBody::NonReplayable
|
ReplayableRequestBody::NonReplayable
|
||||||
},
|
},
|
||||||
|
spool_task: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1056,11 +1082,16 @@ async fn spool_request_body(
|
|||||||
};
|
};
|
||||||
|
|
||||||
let Some(frame) = frame else {
|
let Some(frame) = frame else {
|
||||||
|
let message = "tunnel request body closed before stream end".to_string();
|
||||||
if let Some(state) = &replay_state {
|
if let Some(state) = &replay_state {
|
||||||
state.finish();
|
state.fail(message.clone());
|
||||||
}
|
}
|
||||||
let _ =
|
let _ = send_spool_event(
|
||||||
send_spool_event(&mut spool_tx, SpoolBodyEvent::End, replay_state.as_ref()).await;
|
&mut spool_tx,
|
||||||
|
SpoolBodyEvent::Error(message),
|
||||||
|
replay_state.as_ref(),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
return;
|
return;
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -1086,13 +1117,23 @@ async fn spool_request_body(
|
|||||||
|
|
||||||
if !payload.is_empty() {
|
if !payload.is_empty() {
|
||||||
body_size.fetch_add(payload.len(), Ordering::Relaxed);
|
body_size.fetch_add(payload.len(), Ordering::Relaxed);
|
||||||
try_send_window_update(&frame_tx, stream_id, payload.len());
|
let credit_returned = replay_state
|
||||||
if let Some(state) = &replay_state {
|
.as_ref()
|
||||||
state.push_chunk(payload.clone());
|
.is_some_and(|state| state.push_chunk(payload.clone()));
|
||||||
|
if credit_returned
|
||||||
|
&& !send_window_update(&frame_tx, stream_id, payload.len()).await
|
||||||
|
{
|
||||||
|
if let Some(state) = &replay_state {
|
||||||
|
state.fail("tunnel flow-control update failed".to_string());
|
||||||
|
}
|
||||||
|
return;
|
||||||
}
|
}
|
||||||
if send_spool_event(
|
if send_spool_event(
|
||||||
&mut spool_tx,
|
&mut spool_tx,
|
||||||
SpoolBodyEvent::Data(payload),
|
SpoolBodyEvent::Data {
|
||||||
|
payload,
|
||||||
|
credit_returned,
|
||||||
|
},
|
||||||
replay_state.as_ref(),
|
replay_state.as_ref(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -1479,6 +1520,7 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
let mut stream = response.into_body().into_data_stream();
|
let mut stream = response.into_body().into_data_stream();
|
||||||
|
let chunk_size = MAX_CHUNK_SIZE.min(response_window.initial_window_bytes as usize);
|
||||||
loop {
|
loop {
|
||||||
let chunk_result = if let Some(deadline) = response_body_deadline {
|
let chunk_result = if let Some(deadline) = response_body_deadline {
|
||||||
let Some(remaining) = remaining_timeout(deadline) else {
|
let Some(remaining) = remaining_timeout(deadline) else {
|
||||||
@@ -1531,7 +1573,7 @@ where
|
|||||||
|
|
||||||
match chunk_result {
|
match chunk_result {
|
||||||
Ok(chunk) => {
|
Ok(chunk) => {
|
||||||
if chunk.len() <= MAX_CHUNK_SIZE {
|
if chunk.len() <= chunk_size {
|
||||||
let (payload, extra_flags) = raw_payload(chunk);
|
let (payload, extra_flags) = raw_payload(chunk);
|
||||||
if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len())
|
if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len())
|
||||||
.await
|
.await
|
||||||
@@ -1561,7 +1603,7 @@ where
|
|||||||
} else {
|
} else {
|
||||||
let mut offset = 0;
|
let mut offset = 0;
|
||||||
while offset < chunk.len() {
|
while offset < chunk.len() {
|
||||||
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
|
let end = (offset + chunk_size).min(chunk.len());
|
||||||
let slice = chunk.slice(offset..end);
|
let slice = chunk.slice(offset..end);
|
||||||
let (payload, extra_flags) = raw_payload(slice);
|
let (payload, extra_flags) = raw_payload(slice);
|
||||||
if !acquire_response_credit(
|
if !acquire_response_credit(
|
||||||
@@ -1735,6 +1777,7 @@ pub async fn handle_stream(
|
|||||||
};
|
};
|
||||||
|
|
||||||
server.active_connections.fetch_add(1, Ordering::Release);
|
server.active_connections.fetch_add(1, Ordering::Release);
|
||||||
|
let _active_stream = ActiveStreamGuard(Arc::clone(&server));
|
||||||
|
|
||||||
let stream_io = StreamIo {
|
let stream_io = StreamIo {
|
||||||
body_rx,
|
body_rx,
|
||||||
@@ -1745,7 +1788,6 @@ pub async fn handle_stream(
|
|||||||
|
|
||||||
let connect_elapsed = handle_stream_inner(&state, &server, stream_id, meta, stream_io).await;
|
let connect_elapsed = handle_stream_inner(&state, &server, stream_id, meta, stream_io).await;
|
||||||
|
|
||||||
server.active_connections.fetch_sub(1, Ordering::Release);
|
|
||||||
if let Some(d) = connect_elapsed {
|
if let Some(d) = connect_elapsed {
|
||||||
server.metrics.record_request(d);
|
server.metrics.record_request(d);
|
||||||
}
|
}
|
||||||
@@ -1772,6 +1814,18 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
|||||||
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
|
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
|
||||||
"writer channel stalled for body frame, abandoning stream"
|
"writer channel stalled for body frame, abandoning stream"
|
||||||
);
|
);
|
||||||
|
let reset = TunnelFrame::new(
|
||||||
|
stream_id,
|
||||||
|
MsgType::ResetStream,
|
||||||
|
0,
|
||||||
|
Bytes::from_static(b"{\"reason\":\"tunnel writer stalled\"}"),
|
||||||
|
);
|
||||||
|
if !matches!(
|
||||||
|
tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(reset)).await,
|
||||||
|
Ok(Ok(()))
|
||||||
|
) {
|
||||||
|
tx.close();
|
||||||
|
}
|
||||||
false
|
false
|
||||||
}
|
}
|
||||||
Ok(Err(QueueSendError::Full(_))) => {
|
Ok(Err(QueueSendError::Full(_))) => {
|
||||||
@@ -1781,7 +1835,10 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
|||||||
} else {
|
} else {
|
||||||
match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
||||||
Ok(Ok(())) => true,
|
Ok(Ok(())) => true,
|
||||||
Ok(Err(_)) => false,
|
Ok(Err(_)) => {
|
||||||
|
tx.close();
|
||||||
|
false
|
||||||
|
}
|
||||||
Err(_) => {
|
Err(_) => {
|
||||||
warn!(
|
warn!(
|
||||||
stream_id,
|
stream_id,
|
||||||
@@ -1789,6 +1846,7 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
|||||||
flags = flags,
|
flags = flags,
|
||||||
"control frame send timeout (writer congested), abandoning stream"
|
"control frame send timeout (writer congested), abandoning stream"
|
||||||
);
|
);
|
||||||
|
tx.close();
|
||||||
false
|
false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -2126,7 +2184,6 @@ async fn handle_stream_inner(
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
||||||
// Error frames use best-effort delivery — don't block if writer is congested
|
|
||||||
let safe_message = safe_stream_error_message(msg);
|
let safe_message = safe_stream_error_message(msg);
|
||||||
let _ = send_frame(
|
let _ = send_frame(
|
||||||
tx,
|
tx,
|
||||||
@@ -2163,22 +2220,42 @@ fn build_streaming_request_body(
|
|||||||
|
|
||||||
fn build_spooled_request_body(
|
fn build_spooled_request_body(
|
||||||
spool_rx: mpsc::Receiver<SpoolBodyEvent>,
|
spool_rx: mpsc::Receiver<SpoolBodyEvent>,
|
||||||
|
stream_id: u32,
|
||||||
|
frame_tx: FrameSender,
|
||||||
) -> upstream_client::UpstreamRequestBody {
|
) -> upstream_client::UpstreamRequestBody {
|
||||||
let body_stream = stream::unfold((spool_rx, false), |(mut spool_rx, finished)| async move {
|
let body_stream = stream::unfold(
|
||||||
if finished {
|
(spool_rx, frame_tx, false),
|
||||||
return None;
|
move |(mut spool_rx, frame_tx, finished)| async move {
|
||||||
}
|
if finished {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
match spool_rx.recv().await {
|
match spool_rx.recv().await {
|
||||||
Some(SpoolBodyEvent::Data(payload)) => {
|
Some(SpoolBodyEvent::Data {
|
||||||
Some((Ok(BodyFrame::data(payload)), (spool_rx, false)))
|
payload,
|
||||||
|
credit_returned,
|
||||||
|
}) => {
|
||||||
|
if !credit_returned
|
||||||
|
&& !send_window_update(&frame_tx, stream_id, payload.len()).await
|
||||||
|
{
|
||||||
|
return Some((
|
||||||
|
Err(io::Error::other("tunnel flow-control update failed")),
|
||||||
|
(spool_rx, frame_tx, true),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
Some((Ok(BodyFrame::data(payload)), (spool_rx, frame_tx, false)))
|
||||||
|
}
|
||||||
|
Some(SpoolBodyEvent::Error(message)) => {
|
||||||
|
Some((Err(io::Error::other(message)), (spool_rx, frame_tx, true)))
|
||||||
|
}
|
||||||
|
Some(SpoolBodyEvent::End) => None,
|
||||||
|
None => Some((
|
||||||
|
Err(io::Error::other("tunnel request body ended unexpectedly")),
|
||||||
|
(spool_rx, frame_tx, true),
|
||||||
|
)),
|
||||||
}
|
}
|
||||||
Some(SpoolBodyEvent::Error(message)) => {
|
},
|
||||||
Some((Err(io::Error::other(message)), (spool_rx, true)))
|
);
|
||||||
}
|
|
||||||
Some(SpoolBodyEvent::End) | None => None,
|
|
||||||
}
|
|
||||||
});
|
|
||||||
|
|
||||||
upstream_client::stream_request_body(body_stream)
|
upstream_client::stream_request_body(body_stream)
|
||||||
}
|
}
|
||||||
@@ -2249,6 +2326,105 @@ fn build_prefixed_request_body(
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
|
#[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);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
sender
|
||||||
|
.try_send(TunnelFrame::control(MsgType::Ping, Bytes::new()))
|
||||||
|
.unwrap();
|
||||||
|
let task = tokio::spawn(async move { send_window_update(&sender, 7, 1024).await });
|
||||||
|
tokio::time::sleep(Duration::from_secs(1)).await;
|
||||||
|
assert!(!task.is_finished());
|
||||||
|
high_rx.recv().await.unwrap();
|
||||||
|
assert!(task.await.unwrap());
|
||||||
|
let update = high_rx.recv().await.unwrap();
|
||||||
|
assert_eq!(update.msg_type, MsgType::WindowUpdate);
|
||||||
|
let payload: aether_contracts::tunnel::WindowUpdatePayload =
|
||||||
|
serde_json::from_slice(&update.payload).unwrap();
|
||||||
|
assert_eq!(payload.delta_bytes, 1024);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test(start_paused = true)]
|
||||||
|
async fn stalled_body_delivery_emits_a_reset() {
|
||||||
|
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
sender
|
||||||
|
.try_send(TunnelFrame::new(
|
||||||
|
7,
|
||||||
|
MsgType::ResponseBody,
|
||||||
|
0,
|
||||||
|
Bytes::from_static(b"first"),
|
||||||
|
))
|
||||||
|
.unwrap();
|
||||||
|
assert!(
|
||||||
|
!send_frame(
|
||||||
|
&sender,
|
||||||
|
TunnelFrame::new(7, MsgType::ResponseBody, 0, Bytes::from_static(b"second"))
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
);
|
||||||
|
let reset = high_rx.recv().await.unwrap();
|
||||||
|
assert_eq!(reset.msg_type, MsgType::ResetStream);
|
||||||
|
assert_eq!(reset.stream_id, 7);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn request_credit_follows_consumption_without_redirect_replay() {
|
||||||
|
let (body_tx, body_rx) = mpsc::channel(4);
|
||||||
|
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
let mut prepared = prepare_request_body(
|
||||||
|
7,
|
||||||
|
body_rx,
|
||||||
|
Arc::new(AtomicUsize::new(0)),
|
||||||
|
Instant::now() + Duration::from_secs(10),
|
||||||
|
false,
|
||||||
|
sender,
|
||||||
|
);
|
||||||
|
body_tx
|
||||||
|
.send(TunnelFrame::new(
|
||||||
|
7,
|
||||||
|
MsgType::RequestBody,
|
||||||
|
flags::END_STREAM,
|
||||||
|
Bytes::from_static(b"body"),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
assert!(high_rx.try_recv().is_err());
|
||||||
|
let mut body = prepared.take_first_request_body();
|
||||||
|
assert!(body.frame().await.unwrap().is_ok());
|
||||||
|
assert_eq!(
|
||||||
|
high_rx.recv().await.unwrap().msg_type,
|
||||||
|
MsgType::WindowUpdate
|
||||||
|
);
|
||||||
|
assert!(body.frame().await.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn dropping_prepared_body_cancels_its_spooler() {
|
||||||
|
let (body_tx, body_rx) = mpsc::channel(4);
|
||||||
|
let (high_tx, _high_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
let prepared = prepare_request_body(
|
||||||
|
7,
|
||||||
|
body_rx,
|
||||||
|
Arc::new(AtomicUsize::new(0)),
|
||||||
|
Instant::now() + Duration::from_secs(3600),
|
||||||
|
false,
|
||||||
|
sender,
|
||||||
|
);
|
||||||
|
drop(prepared);
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), body_tx.closed())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::net::SocketAddr;
|
use std::net::SocketAddr;
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
@@ -2379,7 +2555,7 @@ mod tests {
|
|||||||
let (tx, rx) = mpsc::channel(4);
|
let (tx, rx) = mpsc::channel(4);
|
||||||
let (frame_tx, sent, writer_handle) = spawn_test_writer();
|
let (frame_tx, sent, writer_handle) = spawn_test_writer();
|
||||||
let body_size = Arc::new(AtomicUsize::new(0));
|
let body_size = Arc::new(AtomicUsize::new(0));
|
||||||
let prepared = prepare_request_body(
|
let mut prepared = prepare_request_body(
|
||||||
1,
|
1,
|
||||||
rx,
|
rx,
|
||||||
Arc::clone(&body_size),
|
Arc::clone(&body_size),
|
||||||
@@ -2389,6 +2565,7 @@ mod tests {
|
|||||||
);
|
);
|
||||||
let mut body = prepared
|
let mut body = prepared
|
||||||
.first_request_body
|
.first_request_body
|
||||||
|
.take()
|
||||||
.expect("first request body should be present");
|
.expect("first request body should be present");
|
||||||
|
|
||||||
tx.send(TunnelFrame::new(
|
tx.send(TunnelFrame::new(
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
use std::future::Future;
|
||||||
|
use std::pin::Pin;
|
||||||
|
use std::task::{Context, Poll};
|
||||||
|
|
||||||
|
use tokio::task::{JoinError, JoinHandle};
|
||||||
|
|
||||||
|
pub(super) struct SessionTask<T>(JoinHandle<T>);
|
||||||
|
|
||||||
|
impl<T> SessionTask<T> {
|
||||||
|
pub(super) fn new(handle: JoinHandle<T>) -> Self {
|
||||||
|
Self(handle)
|
||||||
|
}
|
||||||
|
pub(super) fn abort(&self) {
|
||||||
|
self.0.abort();
|
||||||
|
}
|
||||||
|
pub(super) fn is_finished(&self) -> bool {
|
||||||
|
self.0.is_finished()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T> Future for SessionTask<T> {
|
||||||
|
type Output = Result<T, JoinError>;
|
||||||
|
|
||||||
|
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
|
||||||
|
Pin::new(&mut self.0).poll(context)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T> Drop for SessionTask<T> {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
self.0.abort();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn dropping_a_session_task_aborts_its_child() {
|
||||||
|
let child = tokio::spawn(std::future::pending::<()>());
|
||||||
|
let abort = child.abort_handle();
|
||||||
|
drop(SessionTask::new(child));
|
||||||
|
tokio::time::timeout(std::time::Duration::from_secs(1), async {
|
||||||
|
while !abort.is_finished() {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -13,6 +13,7 @@ use aether_contracts::tunnel::{MsgType, HEADER_SIZE};
|
|||||||
use aether_runtime::QueueSnapshot;
|
use aether_runtime::QueueSnapshot;
|
||||||
use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError};
|
use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError};
|
||||||
use futures_util::SinkExt;
|
use futures_util::SinkExt;
|
||||||
|
use tokio::sync::watch;
|
||||||
use tokio::task::JoinHandle;
|
use tokio::task::JoinHandle;
|
||||||
use tokio_tungstenite::tungstenite::Message;
|
use tokio_tungstenite::tungstenite::Message;
|
||||||
use tracing::{debug, error, trace};
|
use tracing::{debug, error, trace};
|
||||||
@@ -24,6 +25,8 @@ use aether_contracts::tunnel_security::SecureFrameCodec;
|
|||||||
|
|
||||||
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
|
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
|
||||||
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
|
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
|
||||||
|
const WRITE_TIMEOUT: Duration = Duration::from_secs(15);
|
||||||
|
const CLOSE_TIMEOUT: Duration = Duration::from_secs(1);
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
enum FramePriority {
|
enum FramePriority {
|
||||||
@@ -43,9 +46,18 @@ pub struct FrameQueueSnapshots {
|
|||||||
pub struct FrameSender {
|
pub struct FrameSender {
|
||||||
high_tx: BoundedQueueSender<Frame>,
|
high_tx: BoundedQueueSender<Frame>,
|
||||||
normal_tx: BoundedQueueSender<Frame>,
|
normal_tx: BoundedQueueSender<Frame>,
|
||||||
|
close_tx: watch::Sender<bool>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl FrameSender {
|
impl FrameSender {
|
||||||
|
pub fn close(&self) {
|
||||||
|
let _ = self.close_tx.send(true);
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||||
|
self.close_tx.subscribe()
|
||||||
|
}
|
||||||
|
|
||||||
pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
|
pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
|
||||||
match classify_frame_priority(&frame) {
|
match classify_frame_priority(&frame) {
|
||||||
FramePriority::High => self.high_tx.send(frame).await,
|
FramePriority::High => self.high_tx.send(frame).await,
|
||||||
@@ -73,7 +85,12 @@ impl FrameSender {
|
|||||||
high_tx: BoundedQueueSender<Frame>,
|
high_tx: BoundedQueueSender<Frame>,
|
||||||
normal_tx: BoundedQueueSender<Frame>,
|
normal_tx: BoundedQueueSender<Frame>,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self { high_tx, normal_tx }
|
let (close_tx, _) = watch::channel(false);
|
||||||
|
Self {
|
||||||
|
high_tx,
|
||||||
|
normal_tx,
|
||||||
|
close_tx,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -113,15 +130,24 @@ where
|
|||||||
{
|
{
|
||||||
let (high_tx, mut high_rx) = bounded_queue::<Frame>(HIGH_PRIORITY_QUEUE_CAPACITY);
|
let (high_tx, mut high_rx) = bounded_queue::<Frame>(HIGH_PRIORITY_QUEUE_CAPACITY);
|
||||||
let (normal_tx, mut normal_rx) = bounded_queue::<Frame>(NORMAL_PRIORITY_QUEUE_CAPACITY);
|
let (normal_tx, mut normal_rx) = bounded_queue::<Frame>(NORMAL_PRIORITY_QUEUE_CAPACITY);
|
||||||
let tx = FrameSender { high_tx, normal_tx };
|
let (close_tx, mut close_rx) = watch::channel(false);
|
||||||
|
let tx = FrameSender {
|
||||||
|
high_tx,
|
||||||
|
normal_tx,
|
||||||
|
close_tx,
|
||||||
|
};
|
||||||
|
|
||||||
let handle = tokio::spawn(async move {
|
let handle = tokio::spawn(async move {
|
||||||
let mut ping_ticker = tokio::time::interval(ping_interval);
|
let mut ping_ticker = tokio::time::interval(ping_interval);
|
||||||
let mut high_open = true;
|
let mut high_open = true;
|
||||||
let mut normal_open = true;
|
let mut normal_open = true;
|
||||||
|
let mut close_open = true;
|
||||||
ping_ticker.tick().await; // skip first immediate tick
|
ping_ticker.tick().await; // skip first immediate tick
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
|
if *close_rx.borrow() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
if let Ok(frame) = high_rx.try_recv() {
|
if let Ok(frame) = high_rx.try_recv() {
|
||||||
if !write_frame(
|
if !write_frame(
|
||||||
&mut sink,
|
&mut sink,
|
||||||
@@ -141,6 +167,10 @@ where
|
|||||||
|
|
||||||
tokio::select! {
|
tokio::select! {
|
||||||
biased;
|
biased;
|
||||||
|
changed = close_rx.changed(), if close_open => {
|
||||||
|
if changed.is_err() { close_open = false; }
|
||||||
|
if *close_rx.borrow() { break; }
|
||||||
|
},
|
||||||
frame = high_rx.recv(), if high_open => {
|
frame = high_rx.recv(), if high_open => {
|
||||||
match frame {
|
match frame {
|
||||||
Some(frame) => {
|
Some(frame) => {
|
||||||
@@ -152,7 +182,7 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
_ = ping_ticker.tick(), if high_open || normal_open => {
|
_ = ping_ticker.tick(), if high_open || normal_open => {
|
||||||
if let Err(e) = sink.send(Message::Ping(vec![])).await {
|
if let Err(e) = send_message(&mut sink, Message::Ping(vec![])).await {
|
||||||
error!(error = %e, "failed to send WebSocket ping");
|
error!(error = %e, "failed to send WebSocket ping");
|
||||||
if let Some(metrics) = tunnel_metrics.as_deref() {
|
if let Some(metrics) = tunnel_metrics.as_deref() {
|
||||||
metrics.record_error("ws_ping_error", &e.to_string());
|
metrics.record_error("ws_ping_error", &e.to_string());
|
||||||
@@ -174,7 +204,7 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
debug!("writer task exiting");
|
debug!("writer task exiting");
|
||||||
let _ = sink.close().await;
|
let _ = tokio::time::timeout(CLOSE_TIMEOUT, sink.close()).await;
|
||||||
});
|
});
|
||||||
|
|
||||||
(tx, handle)
|
(tx, handle)
|
||||||
@@ -228,7 +258,7 @@ where
|
|||||||
None => frame.encode(),
|
None => frame.encode(),
|
||||||
};
|
};
|
||||||
let wire_len = data.len().max(HEADER_SIZE);
|
let wire_len = data.len().max(HEADER_SIZE);
|
||||||
if let Err(e) = sink.send(Message::Binary(data.into())).await {
|
if let Err(e) = send_message(sink, Message::Binary(data.into())).await {
|
||||||
error!(
|
error!(
|
||||||
stream_id = stream_id,
|
stream_id = stream_id,
|
||||||
msg_type = ?msg_type,
|
msg_type = ?msg_type,
|
||||||
@@ -248,8 +278,92 @@ where
|
|||||||
true
|
true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn send_message<S>(
|
||||||
|
sink: &mut S,
|
||||||
|
message: Message,
|
||||||
|
) -> Result<(), tokio_tungstenite::tungstenite::Error>
|
||||||
|
where
|
||||||
|
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin,
|
||||||
|
{
|
||||||
|
tokio::time::timeout(WRITE_TIMEOUT, sink.send(message))
|
||||||
|
.await
|
||||||
|
.map_err(|_| {
|
||||||
|
tokio_tungstenite::tungstenite::Error::Io(std::io::Error::new(
|
||||||
|
std::io::ErrorKind::TimedOut,
|
||||||
|
"tunnel WebSocket write timed out",
|
||||||
|
))
|
||||||
|
})?
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
|
#[tokio::test]
|
||||||
|
async fn dropping_last_sender_flushes_queued_body_and_end_frames() {
|
||||||
|
let sink = VecSink::default();
|
||||||
|
let sent = Arc::clone(&sink.sent);
|
||||||
|
let (sender, task) = spawn_writer(sink, Duration::from_secs(60));
|
||||||
|
sender
|
||||||
|
.send(Frame::new(
|
||||||
|
7,
|
||||||
|
MsgType::ResponseBody,
|
||||||
|
0,
|
||||||
|
bytes::Bytes::from_static(b"late"),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
sender
|
||||||
|
.send(Frame::new(7, MsgType::StreamEnd, 0, bytes::Bytes::new()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
drop(sender);
|
||||||
|
task.await.unwrap();
|
||||||
|
let frames = sent.lock().unwrap();
|
||||||
|
assert_eq!(frames.len(), 2);
|
||||||
|
let Message::Binary(body) = &frames[0] else {
|
||||||
|
panic!("expected body")
|
||||||
|
};
|
||||||
|
assert_eq!(
|
||||||
|
Frame::decode(body.clone().into()).unwrap().payload,
|
||||||
|
b"late".as_slice()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
struct StalledSink;
|
||||||
|
|
||||||
|
impl futures_util::Sink<Message> for StalledSink {
|
||||||
|
type Error = Error;
|
||||||
|
fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||||
|
Poll::Pending
|
||||||
|
}
|
||||||
|
fn start_send(self: Pin<&mut Self>, _: Message) -> Result<(), Error> {
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||||
|
Poll::Pending
|
||||||
|
}
|
||||||
|
fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||||
|
Poll::Pending
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test(start_paused = true)]
|
||||||
|
async fn stalled_socket_write_and_close_are_bounded() {
|
||||||
|
let (sender, task) = spawn_writer(StalledSink, Duration::from_secs(60));
|
||||||
|
sender
|
||||||
|
.send(Frame::new(
|
||||||
|
1,
|
||||||
|
MsgType::ResponseBody,
|
||||||
|
0,
|
||||||
|
bytes::Bytes::from_static(b"data"),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
tokio::time::timeout(Duration::from_secs(20), task)
|
||||||
|
.await
|
||||||
|
.expect("writer should time out")
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{Arc, Mutex};
|
||||||
use std::task::{Context, Poll};
|
use std::task::{Context, Poll};
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
use aether_data_contracts::repository::{
|
use aether_data_contracts::repository::{
|
||||||
candidates::{
|
candidates::{
|
||||||
sanitize_request_candidate_api_formats, sanitize_request_candidate_error_type,
|
sanitize_request_candidate_api_formats, sanitize_request_candidate_error_type,
|
||||||
sanitize_request_candidate_extra_data, sanitize_request_candidate_required_capabilities,
|
sanitize_request_candidate_extra_data_for_persistence,
|
||||||
sanitize_request_candidate_skip_reason, DecisionTrace, DecisionTraceCandidate,
|
sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason,
|
||||||
RequestCandidateStatus,
|
DecisionTrace, DecisionTraceCandidate, RequestCandidateStatus,
|
||||||
},
|
},
|
||||||
provider_catalog::StoredProviderCatalogKey,
|
provider_catalog::StoredProviderCatalogKey,
|
||||||
usage::StoredRequestUsageAudit,
|
usage::StoredRequestUsageAudit,
|
||||||
@@ -334,10 +334,13 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts(
|
|||||||
key_accounts: &BTreeMap<String, AdminMonitoringKeyAccountDisplay>,
|
key_accounts: &BTreeMap<String, AdminMonitoringKeyAccountDisplay>,
|
||||||
) -> Value {
|
) -> Value {
|
||||||
let mut item = item.clone();
|
let mut item = item.clone();
|
||||||
item.sanitize_sensitive_diagnostics();
|
item.sanitize_for_admin();
|
||||||
let candidate = &item.candidate;
|
let candidate = &item.candidate;
|
||||||
let sanitized_extra_data =
|
let sanitized_extra_data = build_admin_monitoring_trace_candidate_extra_data(
|
||||||
build_admin_monitoring_trace_candidate_extra_data(candidate.extra_data.as_ref(), usage);
|
candidate.extra_data.as_ref(),
|
||||||
|
candidate.status_code,
|
||||||
|
usage,
|
||||||
|
);
|
||||||
let sanitized_extra_data_ref =
|
let sanitized_extra_data_ref =
|
||||||
(!sanitized_extra_data.is_null()).then_some(&sanitized_extra_data);
|
(!sanitized_extra_data.is_null()).then_some(&sanitized_extra_data);
|
||||||
let sanitized_key_api_formats =
|
let sanitized_key_api_formats =
|
||||||
@@ -382,7 +385,7 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts(
|
|||||||
"is_cached": candidate.is_cached,
|
"is_cached": candidate.is_cached,
|
||||||
"status_code": candidate.status_code,
|
"status_code": candidate.status_code,
|
||||||
"error_type": sanitize_request_candidate_error_type(candidate.error_type.clone()),
|
"error_type": sanitize_request_candidate_error_type(candidate.error_type.clone()),
|
||||||
"error_message": serde_json::Value::Null,
|
"error_message": candidate.error_message,
|
||||||
"latency_ms": candidate.latency_ms,
|
"latency_ms": candidate.latency_ms,
|
||||||
"concurrent_requests": candidate.concurrent_requests,
|
"concurrent_requests": candidate.concurrent_requests,
|
||||||
"ranking": build_admin_monitoring_trace_candidate_ranking(sanitized_extra_data_ref),
|
"ranking": build_admin_monitoring_trace_candidate_ranking(sanitized_extra_data_ref),
|
||||||
@@ -516,9 +519,10 @@ fn build_admin_monitoring_trace_candidate_ranking(existing: Option<&Value>) -> V
|
|||||||
|
|
||||||
fn build_admin_monitoring_trace_candidate_extra_data(
|
fn build_admin_monitoring_trace_candidate_extra_data(
|
||||||
existing: Option<&Value>,
|
existing: Option<&Value>,
|
||||||
|
candidate_status_code: Option<u16>,
|
||||||
usage: Option<&StoredRequestUsageAudit>,
|
usage: Option<&StoredRequestUsageAudit>,
|
||||||
) -> Value {
|
) -> Value {
|
||||||
let mut extra_data = sanitize_request_candidate_extra_data(existing.cloned())
|
let mut extra_data = sanitize_request_candidate_extra_data_for_persistence(existing.cloned())
|
||||||
.and_then(|value| value.as_object().cloned());
|
.and_then(|value| value.as_object().cloned());
|
||||||
|
|
||||||
if let Some(usage) = usage {
|
if let Some(usage) = usage {
|
||||||
@@ -546,7 +550,7 @@ fn build_admin_monitoring_trace_candidate_extra_data(
|
|||||||
if admin_monitoring_usage_is_error_node(usage) {
|
if admin_monitoring_usage_is_error_node(usage) {
|
||||||
if let Some(response) = admin_monitoring_trace_response_data(
|
if let Some(response) = admin_monitoring_trace_response_data(
|
||||||
"upstream_response",
|
"upstream_response",
|
||||||
usage.status_code,
|
candidate_status_code,
|
||||||
usage.response_body_state,
|
usage.response_body_state,
|
||||||
) {
|
) {
|
||||||
merge_admin_monitoring_trace_response(extra_object, "upstream_response", response);
|
merge_admin_monitoring_trace_response(extra_object, "upstream_response", response);
|
||||||
@@ -573,7 +577,8 @@ fn build_admin_monitoring_trace_candidate_extra_data(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
sanitize_request_candidate_extra_data(extra_data.map(Value::Object)).unwrap_or(Value::Null)
|
sanitize_request_candidate_extra_data_for_persistence(extra_data.map(Value::Object))
|
||||||
|
.unwrap_or(Value::Null)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn admin_monitoring_trace_response_data(
|
fn admin_monitoring_trace_response_data(
|
||||||
@@ -607,6 +612,9 @@ fn merge_admin_monitoring_trace_response(
|
|||||||
};
|
};
|
||||||
|
|
||||||
for (field, value) in response_object {
|
for (field, value) in response_object {
|
||||||
|
if field == "status_code" && existing_object.get(field).is_some_and(Value::is_number) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
if admin_monitoring_trace_response_value_empty(value)
|
if admin_monitoring_trace_response_value_empty(value)
|
||||||
&& existing_object
|
&& existing_object
|
||||||
.get(field)
|
.get(field)
|
||||||
|
|||||||
@@ -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]
|
#[test]
|
||||||
fn detail_payload_marks_reference_backed_bodies_as_available() {
|
fn detail_payload_marks_reference_backed_bodies_as_available() {
|
||||||
let item = StoredRequestUsageAudit {
|
let item = StoredRequestUsageAudit {
|
||||||
|
|||||||
@@ -10,7 +10,8 @@ use super::redaction::{
|
|||||||
admin_restore_secret_safe_body_rules, admin_restore_secret_safe_header_rules,
|
admin_restore_secret_safe_body_rules, admin_restore_secret_safe_header_rules,
|
||||||
admin_restore_secret_safe_json, admin_restore_secret_safe_proxy, admin_restore_secret_safe_url,
|
admin_restore_secret_safe_json, admin_restore_secret_safe_proxy, admin_restore_secret_safe_url,
|
||||||
admin_secret_safe_body_rules, admin_secret_safe_header_rules, admin_secret_safe_json,
|
admin_secret_safe_body_rules, admin_secret_safe_header_rules, admin_secret_safe_json,
|
||||||
admin_secret_safe_proxy, admin_secret_safe_url,
|
admin_secret_safe_proxy, admin_secret_safe_url, admin_validate_retained_body_rule_secrets,
|
||||||
|
admin_validate_retained_header_rule_secrets,
|
||||||
};
|
};
|
||||||
|
|
||||||
pub fn normalize_endpoint_api_format(api_format: &str) -> String {
|
pub fn normalize_endpoint_api_format(api_format: &str) -> String {
|
||||||
@@ -200,6 +201,55 @@ mod endpoint_key_count_tests {
|
|||||||
assert_eq!(active, total);
|
assert_eq!(active, total);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn endpoint_updates_reject_unresolved_rule_masks_before_persistence() {
|
||||||
|
let header_rules = json!([{"action": "set", "key": "x-auth", "value": "header-secret"}]);
|
||||||
|
let body_rules = json!([{"action": "set", "path": "auth.token", "value": "body-secret"}]);
|
||||||
|
let mut endpoint = sample_endpoint("chat", "openai:chat");
|
||||||
|
endpoint.header_rules = Some(header_rules.clone());
|
||||||
|
endpoint.body_rules = Some(body_rules.clone());
|
||||||
|
endpoint.config = Some(json!({"response_header_rules": header_rules}));
|
||||||
|
let moved_header =
|
||||||
|
json!([{"action": "set", "key": "x-other-auth", "value": "***", "has_value": true}]);
|
||||||
|
let cases = [
|
||||||
|
(
|
||||||
|
"header_rules",
|
||||||
|
super::AdminProviderEndpointUpdateFields {
|
||||||
|
header_rules: Some(moved_header.clone()),
|
||||||
|
..Default::default()
|
||||||
|
},
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"body_rules",
|
||||||
|
super::AdminProviderEndpointUpdateFields {
|
||||||
|
body_rules: Some(
|
||||||
|
json!([{"action": "set", "path": "auth.api_key", "value": "***", "has_value": true}]),
|
||||||
|
),
|
||||||
|
..Default::default()
|
||||||
|
},
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"config",
|
||||||
|
super::AdminProviderEndpointUpdateFields {
|
||||||
|
config: Some(json!({"response_header_rules": moved_header})),
|
||||||
|
..Default::default()
|
||||||
|
},
|
||||||
|
),
|
||||||
|
];
|
||||||
|
for (field, payload) in cases {
|
||||||
|
let error = super::apply_admin_provider_endpoint_update_fields(
|
||||||
|
&endpoint,
|
||||||
|
|key| key == field,
|
||||||
|
|_| false,
|
||||||
|
&payload,
|
||||||
|
)
|
||||||
|
.expect_err("unresolved masks must not overwrite saved secrets");
|
||||||
|
assert!(error.contains("无法匹配"));
|
||||||
|
assert!(!error.contains("header-secret"));
|
||||||
|
assert!(!error.contains("body-secret"));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn inherited_endpoint_counts_only_include_active_formats() {
|
fn inherited_endpoint_counts_only_include_active_formats() {
|
||||||
let responses_endpoint = sample_endpoint("responses", "openai:responses");
|
let responses_endpoint = sample_endpoint("responses", "openai:responses");
|
||||||
@@ -352,6 +402,10 @@ where
|
|||||||
if !header_rules.is_array() {
|
if !header_rules.is_array() {
|
||||||
return Err("header_rules 必须是数组或 null".to_string());
|
return Err("header_rules 必须是数组或 null".to_string());
|
||||||
}
|
}
|
||||||
|
admin_validate_retained_header_rule_secrets(
|
||||||
|
existing_endpoint.header_rules.as_ref(),
|
||||||
|
header_rules,
|
||||||
|
)?;
|
||||||
Some(admin_restore_secret_safe_header_rules(
|
Some(admin_restore_secret_safe_header_rules(
|
||||||
existing_endpoint.header_rules.as_ref(),
|
existing_endpoint.header_rules.as_ref(),
|
||||||
header_rules,
|
header_rules,
|
||||||
@@ -369,6 +423,10 @@ where
|
|||||||
if !body_rules.is_array() {
|
if !body_rules.is_array() {
|
||||||
return Err("body_rules 必须是数组或 null".to_string());
|
return Err("body_rules 必须是数组或 null".to_string());
|
||||||
}
|
}
|
||||||
|
admin_validate_retained_body_rule_secrets(
|
||||||
|
existing_endpoint.body_rules.as_ref(),
|
||||||
|
body_rules,
|
||||||
|
)?;
|
||||||
Some(admin_restore_secret_safe_body_rules(
|
Some(admin_restore_secret_safe_body_rules(
|
||||||
existing_endpoint.body_rules.as_ref(),
|
existing_endpoint.body_rules.as_ref(),
|
||||||
body_rules,
|
body_rules,
|
||||||
@@ -407,6 +465,15 @@ where
|
|||||||
if !config.is_object() {
|
if !config.is_object() {
|
||||||
return Err("config 必须是对象或 null".to_string());
|
return Err("config 必须是对象或 null".to_string());
|
||||||
}
|
}
|
||||||
|
if let Some(rules) = config.get("response_header_rules") {
|
||||||
|
admin_validate_retained_header_rule_secrets(
|
||||||
|
existing_endpoint
|
||||||
|
.config
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|config| config.get("response_header_rules")),
|
||||||
|
rules,
|
||||||
|
)?;
|
||||||
|
}
|
||||||
Some(admin_restore_secret_safe_json(
|
Some(admin_restore_secret_safe_json(
|
||||||
existing_endpoint.config.as_ref(),
|
existing_endpoint.config.as_ref(),
|
||||||
config,
|
config,
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ pub fn build_kiro_batch_import_key_name(
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|byte| format!("{byte:02x}"))
|
.map(|byte| format!("{byte:02x}"))
|
||||||
.collect::<String>();
|
.collect::<String>();
|
||||||
format!("kiro_{}", &hex[..6])
|
format!("账号_{}", &hex[..6])
|
||||||
});
|
});
|
||||||
format!("{base} ({method})")
|
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 }))
|
.map(|refresh_token| json!({ "refreshToken": refresh_token }))
|
||||||
.collect()
|
.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)"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,5 +1,4 @@
|
|||||||
use serde_json::{json, Map, Value};
|
use serde_json::{json, Map, Value};
|
||||||
use std::collections::BTreeMap;
|
|
||||||
|
|
||||||
const REDACTED_VALUE: &str = "***";
|
const REDACTED_VALUE: &str = "***";
|
||||||
const REDACTED_UPSTREAM_DIAGNOSTIC: &str = "[REDACTED upstream diagnostic]";
|
const REDACTED_UPSTREAM_DIAGNOSTIC: &str = "[REDACTED upstream diagnostic]";
|
||||||
@@ -190,6 +189,20 @@ pub fn admin_restore_secret_safe_body_rules(existing: Option<&Value>, incoming:
|
|||||||
restore_rule_array(existing, incoming, RuleKind::Body)
|
restore_rule_array(existing, incoming, RuleKind::Body)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn admin_validate_retained_header_rule_secrets(
|
||||||
|
existing: Option<&Value>,
|
||||||
|
incoming: &Value,
|
||||||
|
) -> Result<(), String> {
|
||||||
|
validate_retained_rule_secrets(existing, incoming, RuleKind::Header)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn admin_validate_retained_body_rule_secrets(
|
||||||
|
existing: Option<&Value>,
|
||||||
|
incoming: &Value,
|
||||||
|
) -> Result<(), String> {
|
||||||
|
validate_retained_rule_secrets(existing, incoming, RuleKind::Body)
|
||||||
|
}
|
||||||
|
|
||||||
pub fn admin_secret_safe_url(value: Option<&str>) -> Value {
|
pub fn admin_secret_safe_url(value: Option<&str>) -> Value {
|
||||||
value
|
value
|
||||||
.and_then(sanitize_network_url)
|
.and_then(sanitize_network_url)
|
||||||
@@ -1234,11 +1247,7 @@ fn redact_proxy_json_value_for_key(key: &str, value: &Value) -> Value {
|
|||||||
return redact_proxy_secret_value(value);
|
return redact_proxy_secret_value(value);
|
||||||
}
|
}
|
||||||
if json_url_key(&compact_key) {
|
if json_url_key(&compact_key) {
|
||||||
return value
|
return redact_json_url_value(value);
|
||||||
.as_str()
|
|
||||||
.and_then(sanitize_network_url)
|
|
||||||
.map(Value::String)
|
|
||||||
.unwrap_or(Value::Null);
|
|
||||||
}
|
}
|
||||||
if compact_key == "proxy" {
|
if compact_key == "proxy" {
|
||||||
return admin_secret_safe_proxy(Some(value));
|
return admin_secret_safe_proxy(Some(value));
|
||||||
@@ -1261,11 +1270,7 @@ fn redact_json_value_for_key(key: &str, value: &Value) -> Value {
|
|||||||
return redact_secret_value(value);
|
return redact_secret_value(value);
|
||||||
}
|
}
|
||||||
if json_url_key(&compact_key) {
|
if json_url_key(&compact_key) {
|
||||||
return value
|
return redact_json_url_value(value);
|
||||||
.as_str()
|
|
||||||
.and_then(sanitize_network_url)
|
|
||||||
.map(Value::String)
|
|
||||||
.unwrap_or(Value::Null);
|
|
||||||
}
|
}
|
||||||
if compact_key == "proxy" {
|
if compact_key == "proxy" {
|
||||||
return admin_secret_safe_proxy(Some(value));
|
return admin_secret_safe_proxy(Some(value));
|
||||||
@@ -1282,6 +1287,17 @@ fn redact_json_value_for_key(key: &str, value: &Value) -> Value {
|
|||||||
redact_json_value(value)
|
redact_json_value(value)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn redact_json_url_value(value: &Value) -> Value {
|
||||||
|
match value {
|
||||||
|
Value::String(raw) if url::Url::parse(raw).is_ok_and(|url| url.scheme() == "data") => {
|
||||||
|
value.clone()
|
||||||
|
}
|
||||||
|
Value::String(raw) => admin_secret_safe_url(Some(raw)),
|
||||||
|
Value::Array(values) => Value::Array(values.iter().map(redact_json_url_value).collect()),
|
||||||
|
_ => redact_json_value(value),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn redact_body_rule(rule: &Value) -> Value {
|
fn redact_body_rule(rule: &Value) -> Value {
|
||||||
let Some(rule) = rule.as_object() else {
|
let Some(rule) = rule.as_object() else {
|
||||||
return redact_json_value(rule);
|
return redact_json_value(rule);
|
||||||
@@ -1320,7 +1336,12 @@ fn redact_header_rule(rule: &Value) -> Value {
|
|||||||
.map(|(key, value)| (key.clone(), redact_json_value_for_key(key, value)))
|
.map(|(key, value)| (key.clone(), redact_json_value_for_key(key, value)))
|
||||||
.collect::<Map<_, _>>();
|
.collect::<Map<_, _>>();
|
||||||
let is_set = normalized_string_field(rule, "action").as_deref() == Some("set");
|
let is_set = normalized_string_field(rule, "action").as_deref() == Some("set");
|
||||||
if is_set {
|
if is_set
|
||||||
|
&& rule
|
||||||
|
.get("key")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.is_none_or(header_value_is_secret)
|
||||||
|
{
|
||||||
redact_rule_secret_field(rule, &mut projected, "value", "has_value");
|
redact_rule_secret_field(rule, &mut projected, "value", "has_value");
|
||||||
}
|
}
|
||||||
if let Some(condition) = rule.get("condition") {
|
if let Some(condition) = rule.get("condition") {
|
||||||
@@ -1374,7 +1395,14 @@ fn redact_header_values(value: &Value) -> Value {
|
|||||||
Value::Object(
|
Value::Object(
|
||||||
headers
|
headers
|
||||||
.iter()
|
.iter()
|
||||||
.map(|(key, value)| (key.clone(), redact_secret_value(value)))
|
.map(|(key, value)| {
|
||||||
|
let value = if header_value_is_secret(key) {
|
||||||
|
redact_secret_value(value)
|
||||||
|
} else {
|
||||||
|
redact_json_value(value)
|
||||||
|
};
|
||||||
|
(key.clone(), value)
|
||||||
|
})
|
||||||
.collect(),
|
.collect(),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -1395,10 +1423,25 @@ fn restore_json_value(existing: Option<&Value>, incoming: &Value, key: Option<&s
|
|||||||
return restore_header_values(existing, incoming);
|
return restore_header_values(existing, incoming);
|
||||||
}
|
}
|
||||||
if json_url_key(&compact_key) {
|
if json_url_key(&compact_key) {
|
||||||
return incoming
|
if let Some(incoming_url) = incoming.as_str() {
|
||||||
.as_str()
|
return restore_url_value(existing, incoming_url);
|
||||||
.map(|incoming_url| restore_url_value(existing, incoming_url))
|
}
|
||||||
.unwrap_or_else(|| incoming.clone());
|
if let Some(values) = incoming.as_array() {
|
||||||
|
let existing_values = existing.and_then(Value::as_array);
|
||||||
|
return Value::Array(
|
||||||
|
values
|
||||||
|
.iter()
|
||||||
|
.enumerate()
|
||||||
|
.map(|(index, value)| {
|
||||||
|
restore_json_value(
|
||||||
|
existing_values.and_then(|values| values.get(index)),
|
||||||
|
value,
|
||||||
|
Some(key),
|
||||||
|
)
|
||||||
|
})
|
||||||
|
.collect(),
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if json_secret_key(&compact_key, incoming) {
|
if json_secret_key(&compact_key, incoming) {
|
||||||
return restore_masked_secret(existing, incoming, true);
|
return restore_masked_secret(existing, incoming, true);
|
||||||
@@ -1448,29 +1491,20 @@ fn restore_rule_array(existing: Option<&Value>, incoming: &Value, kind: RuleKind
|
|||||||
let Some(incoming_values) = incoming.as_array() else {
|
let Some(incoming_values) = incoming.as_array() else {
|
||||||
return incoming.clone();
|
return incoming.clone();
|
||||||
};
|
};
|
||||||
let existing_values = existing.and_then(Value::as_array);
|
if unchanged_projected_rules(existing, incoming, kind) {
|
||||||
let incoming_identities = identity_counts(incoming_values, kind);
|
return existing.cloned().unwrap_or_else(|| incoming.clone());
|
||||||
let existing_identities = existing_values
|
}
|
||||||
.map(|values| identity_counts(values, kind))
|
let existing_values = existing
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.map(Vec::as_slice)
|
||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
|
|
||||||
Value::Array(
|
Value::Array(
|
||||||
incoming_values
|
incoming_values
|
||||||
.iter()
|
.iter()
|
||||||
.map(|incoming_rule| {
|
.map(|incoming_rule| {
|
||||||
let identity = rule_identity(incoming_rule, kind);
|
let existing_rule =
|
||||||
let existing_rule = identity.as_ref().and_then(|identity| {
|
match_existing_rule(existing_values, incoming_values, incoming_rule, kind);
|
||||||
(incoming_identities.get(identity) == Some(&1)
|
|
||||||
&& existing_identities.get(identity) == Some(&1))
|
|
||||||
.then(|| {
|
|
||||||
existing_values.and_then(|values| {
|
|
||||||
values.iter().find(|candidate| {
|
|
||||||
rule_identity(candidate, kind).as_ref() == Some(identity)
|
|
||||||
})
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.flatten()
|
|
||||||
});
|
|
||||||
restore_rule(existing_rule, incoming_rule, kind)
|
restore_rule(existing_rule, incoming_rule, kind)
|
||||||
})
|
})
|
||||||
.collect(),
|
.collect(),
|
||||||
@@ -1546,7 +1580,10 @@ fn restore_condition(existing: Option<&Value>, incoming: &Value) -> Value {
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
let existing_value = existing_object.and_then(|object| object.get(key));
|
let existing_value = existing_object.and_then(|object| object.get(key));
|
||||||
let value = if key == "value" && condition_value_is_secret(incoming_object) {
|
let value = if key == "value"
|
||||||
|
&& (condition_value_is_secret(incoming_object)
|
||||||
|
|| incoming_object.get("has_value").and_then(Value::as_bool) == Some(true))
|
||||||
|
{
|
||||||
let marker_set =
|
let marker_set =
|
||||||
incoming_object.get("has_value").and_then(Value::as_bool) == Some(true);
|
incoming_object.get("has_value").and_then(Value::as_bool) == Some(true);
|
||||||
restore_masked_secret(existing_value, incoming_value, marker_set)
|
restore_masked_secret(existing_value, incoming_value, marker_set)
|
||||||
@@ -1562,29 +1599,25 @@ fn restore_condition_array(existing: Option<&Value>, incoming: &Value) -> Value
|
|||||||
let Some(incoming_values) = incoming.as_array() else {
|
let Some(incoming_values) = incoming.as_array() else {
|
||||||
return incoming.clone();
|
return incoming.clone();
|
||||||
};
|
};
|
||||||
let existing_values = existing.and_then(Value::as_array);
|
let existing_values = existing
|
||||||
let incoming_counts = condition_identity_counts(incoming_values);
|
.and_then(Value::as_array)
|
||||||
let existing_counts = existing_values
|
.map(Vec::as_slice)
|
||||||
.map(|values| condition_identity_counts(values))
|
|
||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
|
if projected_conditions_unchanged(existing_values, incoming_values) {
|
||||||
|
return Value::Array(existing_values.to_vec());
|
||||||
|
}
|
||||||
|
|
||||||
Value::Array(
|
Value::Array(
|
||||||
incoming_values
|
incoming_values
|
||||||
.iter()
|
.iter()
|
||||||
.map(|incoming_condition| {
|
.map(|incoming_condition| {
|
||||||
let identity = condition_identity(incoming_condition);
|
let existing_condition = match_existing_entry(
|
||||||
let existing_condition = identity.as_ref().and_then(|identity| {
|
existing_values,
|
||||||
(incoming_counts.get(identity) == Some(&1)
|
incoming_values,
|
||||||
&& existing_counts.get(identity) == Some(&1))
|
incoming_condition,
|
||||||
.then(|| {
|
condition_identity,
|
||||||
existing_values.and_then(|values| {
|
redact_condition,
|
||||||
values.iter().find(|candidate| {
|
);
|
||||||
condition_identity(candidate).as_ref() == Some(identity)
|
|
||||||
})
|
|
||||||
})
|
|
||||||
})
|
|
||||||
.flatten()
|
|
||||||
});
|
|
||||||
restore_condition(existing_condition, incoming_condition)
|
restore_condition(existing_condition, incoming_condition)
|
||||||
})
|
})
|
||||||
.collect(),
|
.collect(),
|
||||||
@@ -1633,20 +1666,174 @@ fn restore_url_value(existing: Option<&Value>, incoming_url: &str) -> Value {
|
|||||||
Value::String(incoming_url.to_string())
|
Value::String(incoming_url.to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
fn identity_counts(values: &[Value], kind: RuleKind) -> BTreeMap<String, usize> {
|
fn project_rule(rule: &Value, kind: RuleKind) -> Value {
|
||||||
let mut counts = BTreeMap::new();
|
match kind {
|
||||||
for identity in values.iter().filter_map(|value| rule_identity(value, kind)) {
|
RuleKind::Header => redact_header_rule(rule),
|
||||||
*counts.entry(identity).or_insert(0) += 1;
|
RuleKind::Body => redact_body_rule(rule),
|
||||||
}
|
}
|
||||||
counts
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn condition_identity_counts(values: &[Value]) -> BTreeMap<String, usize> {
|
fn unchanged_projected_rules(existing: Option<&Value>, incoming: &Value, kind: RuleKind) -> bool {
|
||||||
let mut counts = BTreeMap::new();
|
existing.and_then(Value::as_array).is_some_and(|rules| {
|
||||||
for identity in values.iter().filter_map(condition_identity) {
|
Value::Array(rules.iter().map(|rule| project_rule(rule, kind)).collect()) == *incoming
|
||||||
*counts.entry(identity).or_insert(0) += 1;
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn rule_match_shape(rule: &Value, kind: RuleKind) -> Value {
|
||||||
|
let mut projected = project_rule(rule, kind);
|
||||||
|
if let Some(object) = projected.as_object_mut() {
|
||||||
|
for field in [
|
||||||
|
"enabled",
|
||||||
|
"value",
|
||||||
|
"has_value",
|
||||||
|
"pattern",
|
||||||
|
"has_pattern",
|
||||||
|
"replacement",
|
||||||
|
"has_replacement",
|
||||||
|
] {
|
||||||
|
object.remove(field);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
counts
|
projected
|
||||||
|
}
|
||||||
|
|
||||||
|
fn match_existing_rule<'a>(
|
||||||
|
existing: &'a [Value],
|
||||||
|
incoming: &[Value],
|
||||||
|
rule: &Value,
|
||||||
|
kind: RuleKind,
|
||||||
|
) -> Option<&'a Value> {
|
||||||
|
match_existing_entry(
|
||||||
|
existing,
|
||||||
|
incoming,
|
||||||
|
rule,
|
||||||
|
|value| rule_identity(value, kind),
|
||||||
|
|value| rule_match_shape(value, kind),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn match_existing_entry<'a>(
|
||||||
|
existing: &'a [Value],
|
||||||
|
incoming: &[Value],
|
||||||
|
entry: &Value,
|
||||||
|
identity: impl Fn(&Value) -> Option<String>,
|
||||||
|
project: impl Fn(&Value) -> Value,
|
||||||
|
) -> Option<&'a Value> {
|
||||||
|
let entry_identity = identity(entry)?;
|
||||||
|
let same_identity = |value: &&Value| identity(value).as_ref() == Some(&entry_identity);
|
||||||
|
let candidates = existing.iter().filter(same_identity).collect::<Vec<_>>();
|
||||||
|
if candidates.len() == 1 && incoming.iter().filter(same_identity).count() == 1 {
|
||||||
|
return candidates.first().copied();
|
||||||
|
}
|
||||||
|
let projected = project(entry);
|
||||||
|
let mut matching = candidates
|
||||||
|
.into_iter()
|
||||||
|
.filter(|candidate| project(candidate) == projected);
|
||||||
|
let matched = matching.next()?;
|
||||||
|
if matching.next().is_some()
|
||||||
|
|| incoming
|
||||||
|
.iter()
|
||||||
|
.filter(same_identity)
|
||||||
|
.filter(|value| project(value) == projected)
|
||||||
|
.count()
|
||||||
|
!= 1
|
||||||
|
{
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
Some(matched)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn projected_conditions_unchanged(existing: &[Value], incoming: &[Value]) -> bool {
|
||||||
|
existing.len() == incoming.len()
|
||||||
|
&& existing
|
||||||
|
.iter()
|
||||||
|
.zip(incoming)
|
||||||
|
.all(|(existing, incoming)| redact_condition(existing) == *incoming)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn validate_retained_rule_secrets(
|
||||||
|
existing: Option<&Value>,
|
||||||
|
incoming: &Value,
|
||||||
|
kind: RuleKind,
|
||||||
|
) -> Result<(), String> {
|
||||||
|
let Some(incoming_values) = incoming.as_array() else {
|
||||||
|
return Ok(());
|
||||||
|
};
|
||||||
|
if unchanged_projected_rules(existing, incoming, kind) {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
let existing_values = existing
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.map(Vec::as_slice)
|
||||||
|
.unwrap_or_default();
|
||||||
|
for rule in incoming_values {
|
||||||
|
let existing_rule = match_existing_rule(existing_values, incoming_values, rule, kind);
|
||||||
|
validate_retained_secret_fields(existing_rule, rule)?;
|
||||||
|
if let Some(condition) = rule.get("condition") {
|
||||||
|
validate_retained_condition_secrets(
|
||||||
|
existing_rule.and_then(|rule| rule.get("condition")),
|
||||||
|
condition,
|
||||||
|
)?;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn validate_retained_secret_fields(
|
||||||
|
existing: Option<&Value>,
|
||||||
|
incoming: &Value,
|
||||||
|
) -> Result<(), String> {
|
||||||
|
for (field, marker) in [
|
||||||
|
("value", "has_value"),
|
||||||
|
("pattern", "has_pattern"),
|
||||||
|
("replacement", "has_replacement"),
|
||||||
|
] {
|
||||||
|
if incoming.get(marker).and_then(Value::as_bool) == Some(true)
|
||||||
|
&& incoming.get(field).and_then(Value::as_str) == Some(REDACTED_VALUE)
|
||||||
|
&& existing
|
||||||
|
.and_then(|value| value.get(field))
|
||||||
|
.filter(|value| secret_value_is_set(value))
|
||||||
|
.is_none()
|
||||||
|
{
|
||||||
|
return Err(
|
||||||
|
"无法匹配脱敏规则的原值,请查看原值后重新填写,避免将占位符保存为实际配置"
|
||||||
|
.to_string(),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn validate_retained_condition_secrets(
|
||||||
|
existing: Option<&Value>,
|
||||||
|
incoming: &Value,
|
||||||
|
) -> Result<(), String> {
|
||||||
|
for group_key in ["all", "any"] {
|
||||||
|
if let Some(children) = incoming.get(group_key).and_then(Value::as_array) {
|
||||||
|
let existing_children = existing
|
||||||
|
.and_then(|value| value.get(group_key))
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.map(Vec::as_slice)
|
||||||
|
.unwrap_or_default();
|
||||||
|
if projected_conditions_unchanged(existing_children, children) {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
for child in children {
|
||||||
|
let existing_child = match_existing_entry(
|
||||||
|
existing_children,
|
||||||
|
children,
|
||||||
|
child,
|
||||||
|
condition_identity,
|
||||||
|
redact_condition,
|
||||||
|
);
|
||||||
|
validate_retained_condition_secrets(existing_child, child)?;
|
||||||
|
}
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let existing =
|
||||||
|
existing.filter(|value| condition_identity(value) == condition_identity(incoming));
|
||||||
|
validate_retained_secret_fields(existing, incoming)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn rule_identity(value: &Value, kind: RuleKind) -> Option<String> {
|
fn rule_identity(value: &Value, kind: RuleKind) -> Option<String> {
|
||||||
@@ -1676,8 +1863,18 @@ fn rule_identity(value: &Value, kind: RuleKind) -> Option<String> {
|
|||||||
|
|
||||||
fn condition_identity(value: &Value) -> Option<String> {
|
fn condition_identity(value: &Value) -> Option<String> {
|
||||||
let value = value.as_object()?;
|
let value = value.as_object()?;
|
||||||
if value.contains_key("all") || value.contains_key("any") {
|
for group_key in ["all", "any"] {
|
||||||
return None;
|
if let Some(children) = value.get(group_key).and_then(Value::as_array) {
|
||||||
|
let mut identities = children
|
||||||
|
.iter()
|
||||||
|
.map(condition_identity)
|
||||||
|
.collect::<Option<Vec<_>>>()?;
|
||||||
|
identities.sort();
|
||||||
|
return Some(format!(
|
||||||
|
"condition:{group_key}:{}",
|
||||||
|
serde_json::to_string(&identities).ok()?
|
||||||
|
));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
let path = trimmed_string_field(value, "path")?;
|
let path = trimmed_string_field(value, "path")?;
|
||||||
let op = normalized_string_field(value, "op")?;
|
let op = normalized_string_field(value, "op")?;
|
||||||
@@ -1720,11 +1917,36 @@ fn is_rule_marker(key: &str) -> bool {
|
|||||||
|
|
||||||
fn condition_value_is_secret(condition: &Map<String, Value>) -> bool {
|
fn condition_value_is_secret(condition: &Map<String, Value>) -> bool {
|
||||||
let source = normalized_condition_source(condition.get("source").and_then(Value::as_str));
|
let source = normalized_condition_source(condition.get("source").and_then(Value::as_str));
|
||||||
source == "request_headers"
|
let path = condition.get("path").and_then(Value::as_str);
|
||||||
|| condition
|
if source == "request_headers" {
|
||||||
.get("path")
|
path.is_none_or(header_value_is_secret)
|
||||||
.and_then(Value::as_str)
|
} else {
|
||||||
.is_some_and(json_path_targets_secret)
|
path.is_some_and(json_path_targets_secret)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn header_value_is_secret(name: &str) -> bool {
|
||||||
|
!matches!(
|
||||||
|
name.trim().to_ascii_lowercase().as_str(),
|
||||||
|
"accept"
|
||||||
|
| "accept-encoding"
|
||||||
|
| "accept-language"
|
||||||
|
| "cache-control"
|
||||||
|
| "content-encoding"
|
||||||
|
| "content-type"
|
||||||
|
| "user-agent"
|
||||||
|
| "anthropic-version"
|
||||||
|
| "anthropic-beta"
|
||||||
|
| "openai-beta"
|
||||||
|
| "x-stainless-lang"
|
||||||
|
| "x-stainless-package-version"
|
||||||
|
| "x-stainless-os"
|
||||||
|
| "x-stainless-arch"
|
||||||
|
| "x-stainless-runtime"
|
||||||
|
| "x-stainless-runtime-version"
|
||||||
|
| "x-stainless-retry-count"
|
||||||
|
| "x-stainless-timeout"
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn normalized_condition_source(source: Option<&str>) -> String {
|
fn normalized_condition_source(source: Option<&str>) -> String {
|
||||||
@@ -2542,4 +2764,167 @@ mod tests {
|
|||||||
assert_eq!(projected, "https://api.example/v1");
|
assert_eq!(projected, "https://api.example/v1");
|
||||||
assert_eq!(admin_secret_safe_url(Some("not a url")), json!(null));
|
assert_eq!(admin_secret_safe_url(Some("not a url")), json!(null));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn public_protocol_headers_and_conditions_remain_editable() {
|
||||||
|
let rules = json!([
|
||||||
|
{"action": "set", "key": "Content-Type", "value": "application/json"},
|
||||||
|
{"action": "set", "key": "User-Agent", "value": "client/1.0"},
|
||||||
|
{"action": "set", "key": "anthropic-version", "value": "2023-06-01"},
|
||||||
|
{"action": "set", "key": "OpenAI-Beta", "value": "responses=experimental", "condition": {
|
||||||
|
"source": "request_headers", "path": "Accept", "op": "eq", "value": "text/event-stream"
|
||||||
|
}}
|
||||||
|
]);
|
||||||
|
assert_eq!(admin_secret_safe_header_rules(Some(&rules)), rules);
|
||||||
|
let mut legacy_projection = rules.clone();
|
||||||
|
legacy_projection[0]["value"] = json!("***");
|
||||||
|
legacy_projection[0]["has_value"] = json!(true);
|
||||||
|
legacy_projection[3]["condition"]["value"] = json!("***");
|
||||||
|
legacy_projection[3]["condition"]["has_value"] = json!(true);
|
||||||
|
assert_eq!(
|
||||||
|
admin_restore_secret_safe_header_rules(Some(&rules), &legacy_projection),
|
||||||
|
rules
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn header_maps_keep_protocol_values_but_hide_credentials_and_unknown_headers() {
|
||||||
|
let projected = admin_secret_safe_json(Some(&json!({"headers": {
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
"User-Agent": "client/1.0",
|
||||||
|
"Authorization": "Bearer secret",
|
||||||
|
"Cookie": "session=secret",
|
||||||
|
"x-custom-auth": "custom-secret"
|
||||||
|
}})));
|
||||||
|
assert_eq!(projected["headers"]["Content-Type"], "application/json");
|
||||||
|
assert_eq!(projected["headers"]["User-Agent"], "client/1.0");
|
||||||
|
for header in ["Authorization", "Cookie", "x-custom-auth"] {
|
||||||
|
assert_eq!(projected["headers"][header], "***");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn duplicate_conditional_rules_retain_secrets_when_reordered_or_disabled() {
|
||||||
|
let existing = json!([
|
||||||
|
{"action": "set", "key": "x-auth", "value": "first-secret", "condition": {
|
||||||
|
"path": "model", "op": "eq", "value": "first-model"
|
||||||
|
}},
|
||||||
|
{"action": "set", "key": "x-auth", "value": "second-secret", "condition": {
|
||||||
|
"path": "model", "op": "eq", "value": "second-model"
|
||||||
|
}}
|
||||||
|
]);
|
||||||
|
let mut incoming = admin_secret_safe_header_rules(Some(&existing));
|
||||||
|
incoming.as_array_mut().unwrap().reverse();
|
||||||
|
incoming[0]["enabled"] = json!(false);
|
||||||
|
super::admin_validate_retained_header_rule_secrets(Some(&existing), &incoming).unwrap();
|
||||||
|
let restored = admin_restore_secret_safe_header_rules(Some(&existing), &incoming);
|
||||||
|
assert_eq!(restored[0]["value"], "second-secret");
|
||||||
|
assert_eq!(restored[0]["enabled"], false);
|
||||||
|
assert_eq!(restored[1]["value"], "first-secret");
|
||||||
|
assert!(restored[0].get("has_value").is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn unchanged_duplicate_body_rules_preserve_their_original_values() {
|
||||||
|
let existing = json!([
|
||||||
|
{"action": "append", "path": "auth.cookies", "value": "first-secret"},
|
||||||
|
{"action": "append", "path": "auth.cookies", "value": "second-secret"}
|
||||||
|
]);
|
||||||
|
let incoming = admin_secret_safe_body_rules(Some(&existing));
|
||||||
|
super::admin_validate_retained_body_rule_secrets(Some(&existing), &incoming).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
admin_restore_secret_safe_body_rules(Some(&existing), &incoming),
|
||||||
|
existing
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn nested_condition_groups_preserve_secrets_after_sibling_edits_and_reordering() {
|
||||||
|
let existing = json!([{
|
||||||
|
"action": "set", "key": "x-output", "value": "header-secret",
|
||||||
|
"condition": {"all": [
|
||||||
|
{"any": [
|
||||||
|
{"path": "auth.token", "op": "eq", "value": "condition-secret"},
|
||||||
|
{"path": "model", "op": "eq", "value": "old-model"}
|
||||||
|
]},
|
||||||
|
{"path": "metadata.enabled", "op": "eq", "value": true}
|
||||||
|
]}
|
||||||
|
}]);
|
||||||
|
let mut incoming = admin_secret_safe_header_rules(Some(&existing));
|
||||||
|
incoming[0]["condition"]["all"][0]["any"][1]["value"] = json!("new-model");
|
||||||
|
incoming[0]["condition"]["all"]
|
||||||
|
.as_array_mut()
|
||||||
|
.unwrap()
|
||||||
|
.reverse();
|
||||||
|
super::admin_validate_retained_header_rule_secrets(Some(&existing), &incoming).unwrap();
|
||||||
|
let restored = admin_restore_secret_safe_header_rules(Some(&existing), &incoming);
|
||||||
|
assert_eq!(restored[0]["value"], "header-secret");
|
||||||
|
assert_eq!(
|
||||||
|
restored[0]["condition"]["all"][1]["any"][0]["value"],
|
||||||
|
"condition-secret"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
restored[0]["condition"]["all"][1]["any"][1]["value"],
|
||||||
|
"new-model"
|
||||||
|
);
|
||||||
|
assert!(!restored.to_string().contains("has_value"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn retained_masks_cannot_silently_overwrite_changed_or_ambiguous_rules() {
|
||||||
|
let existing = json!([
|
||||||
|
{"action": "set", "key": "x-auth", "value": "first-secret"},
|
||||||
|
{"action": "set", "key": "x-auth", "value": "second-secret"}
|
||||||
|
]);
|
||||||
|
let mut incoming = admin_secret_safe_header_rules(Some(&existing));
|
||||||
|
incoming[0]["enabled"] = json!(false);
|
||||||
|
assert!(
|
||||||
|
super::admin_validate_retained_header_rule_secrets(Some(&existing), &incoming).is_err()
|
||||||
|
);
|
||||||
|
|
||||||
|
let existing = json!([{"action": "set", "path": "auth.token", "value": "secret"}]);
|
||||||
|
let mut incoming = admin_secret_safe_body_rules(Some(&existing));
|
||||||
|
incoming[0]["path"] = json!("auth.api_key");
|
||||||
|
assert!(
|
||||||
|
super::admin_validate_retained_body_rule_secrets(Some(&existing), &incoming).is_err()
|
||||||
|
);
|
||||||
|
incoming[0]["value"] = json!("replacement-secret");
|
||||||
|
incoming[0].as_object_mut().unwrap().remove("has_value");
|
||||||
|
assert!(
|
||||||
|
super::admin_validate_retained_body_rule_secrets(Some(&existing), &incoming).is_ok()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn body_rule_projection_preserves_structured_image_urls_and_inline_data() {
|
||||||
|
let existing = json!([{
|
||||||
|
"action": "append", "path": "messages[0].content",
|
||||||
|
"value": {"type": "image_url", "image_url": {"url": "data:image/png;base64,aW1hZ2U=", "detail": "high"}}
|
||||||
|
}]);
|
||||||
|
let projected = admin_secret_safe_body_rules(Some(&existing));
|
||||||
|
assert_eq!(projected, existing);
|
||||||
|
assert_eq!(
|
||||||
|
admin_restore_secret_safe_body_rules(Some(&existing), &projected),
|
||||||
|
existing
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn structured_url_values_still_hide_and_restore_network_credentials() {
|
||||||
|
let existing = json!({
|
||||||
|
"image_url": {"url": "https://user:[email protected]/image?token=secret", "detail": "auto"},
|
||||||
|
"url": ["https://example.test/file?token=secret", "data:image/png;base64,aW1hZ2U="]
|
||||||
|
});
|
||||||
|
let projected = admin_secret_safe_json(Some(&existing));
|
||||||
|
assert_eq!(projected["image_url"]["url"], "https://example.test/image");
|
||||||
|
assert_eq!(projected["image_url"]["detail"], "auto");
|
||||||
|
assert_eq!(projected["url"][0], "https://example.test/file");
|
||||||
|
assert_eq!(projected["url"][1], existing["url"][1]);
|
||||||
|
assert!(!projected.to_string().contains("secret"));
|
||||||
|
assert!(!projected.to_string().contains("password"));
|
||||||
|
assert_eq!(
|
||||||
|
admin_restore_secret_safe_json(Some(&existing), &projected),
|
||||||
|
existing
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -405,18 +405,40 @@ pub fn build_kiro_device_key_name(email: Option<&str>, refresh_token: Option<&st
|
|||||||
.collect::<String>()
|
.collect::<String>()
|
||||||
})
|
})
|
||||||
.unwrap_or_else(|| "unknown".to_string());
|
.unwrap_or_else(|| "unknown".to_string());
|
||||||
format!("kiro_{fallback} (idc)")
|
format!("账号_{fallback} (idc)")
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::{
|
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,
|
parse_provider_oauth_callback_params, MAX_UNVERIFIED_JWT_CLAIMS_BYTES,
|
||||||
};
|
};
|
||||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||||
use serde_json::json;
|
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 {
|
fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
|
||||||
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
|
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
|
||||||
let payload = URL_SAFE_NO_PAD.encode(payload.to_string());
|
let payload = URL_SAFE_NO_PAD.encode(payload.to_string());
|
||||||
|
|||||||
@@ -67,17 +67,6 @@ pub fn admin_email_template_html_is_valid(value: &str) -> bool {
|
|||||||
pub const ADMIN_SYSTEM_CONFIG_EXPORT_VERSION: &str = "2.3";
|
pub const ADMIN_SYSTEM_CONFIG_EXPORT_VERSION: &str = "2.3";
|
||||||
pub const ADMIN_SYSTEM_CONFIG_SUPPORTED_VERSIONS: &[&str] =
|
pub const ADMIN_SYSTEM_CONFIG_SUPPORTED_VERSIONS: &[&str] =
|
||||||
&["2.0", "2.1", "2.2", ADMIN_SYSTEM_CONFIG_EXPORT_VERSION];
|
&["2.0", "2.1", "2.2", ADMIN_SYSTEM_CONFIG_EXPORT_VERSION];
|
||||||
pub const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY: &str = "execution_extra_trusted_dns_hosts";
|
|
||||||
pub const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_MAX_ENTRIES: usize = 128;
|
|
||||||
pub const EXECUTION_EXTRA_TRUSTED_DNS_HOST_MAX_BYTES: usize = 253;
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
||||||
pub enum ExecutionExtraTrustedDnsHostsConfigError {
|
|
||||||
InvalidValue,
|
|
||||||
TooManyEntries,
|
|
||||||
InvalidHost,
|
|
||||||
}
|
|
||||||
|
|
||||||
pub const ADMIN_SYSTEM_USERS_EXPORT_VERSION: &str = "1.6";
|
pub const ADMIN_SYSTEM_USERS_EXPORT_VERSION: &str = "1.6";
|
||||||
pub const ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS: &[&str] =
|
pub const ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS: &[&str] =
|
||||||
&["1.3", "1.4", "1.5", ADMIN_SYSTEM_USERS_EXPORT_VERSION];
|
&["1.3", "1.4", "1.5", ADMIN_SYSTEM_USERS_EXPORT_VERSION];
|
||||||
@@ -2289,7 +2278,6 @@ pub fn admin_system_config_default_value(key: &str) -> Option<serde_json::Value>
|
|||||||
"email_suffix_mode" => Some(json!("none")),
|
"email_suffix_mode" => Some(json!("none")),
|
||||||
"email_suffix_list" => Some(json!([])),
|
"email_suffix_list" => Some(json!([])),
|
||||||
"enable_format_conversion" => Some(json!(false)),
|
"enable_format_conversion" => Some(json!(false)),
|
||||||
EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY => Some(json!([])),
|
|
||||||
"enable_model_directives" => Some(json!(false)),
|
"enable_model_directives" => Some(json!(false)),
|
||||||
// Failover after a provider-side Cyber policy refusal is an explicit
|
// Failover after a provider-side Cyber policy refusal is an explicit
|
||||||
// opt-in. Keep the system-config fallback aligned with the routing
|
// opt-in. Keep the system-config fallback aligned with the routing
|
||||||
@@ -2329,58 +2317,6 @@ pub fn admin_system_config_default_value(key: &str) -> Option<serde_json::Value>
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn normalize_execution_extra_trusted_dns_hosts_config_value(
|
|
||||||
value: serde_json::Value,
|
|
||||||
) -> Result<serde_json::Value, ExecutionExtraTrustedDnsHostsConfigError> {
|
|
||||||
let values = match value {
|
|
||||||
Value::Null => Vec::new(),
|
|
||||||
Value::Array(values) => values,
|
|
||||||
_ => return Err(ExecutionExtraTrustedDnsHostsConfigError::InvalidValue),
|
|
||||||
};
|
|
||||||
if values.len() > EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_MAX_ENTRIES {
|
|
||||||
return Err(ExecutionExtraTrustedDnsHostsConfigError::TooManyEntries);
|
|
||||||
}
|
|
||||||
|
|
||||||
let mut hosts = BTreeSet::new();
|
|
||||||
for value in values {
|
|
||||||
let host = value
|
|
||||||
.as_str()
|
|
||||||
.map(str::trim)
|
|
||||||
.ok_or(ExecutionExtraTrustedDnsHostsConfigError::InvalidHost)?;
|
|
||||||
let host = host.trim_end_matches('.').to_ascii_lowercase();
|
|
||||||
if !execution_extra_trusted_dns_host_is_valid(&host) {
|
|
||||||
return Err(ExecutionExtraTrustedDnsHostsConfigError::InvalidHost);
|
|
||||||
}
|
|
||||||
hosts.insert(host);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(Value::Array(hosts.into_iter().map(Value::String).collect()))
|
|
||||||
}
|
|
||||||
|
|
||||||
fn execution_extra_trusted_dns_host_is_valid(host: &str) -> bool {
|
|
||||||
if host.is_empty()
|
|
||||||
|| host.len() > EXECUTION_EXTRA_TRUSTED_DNS_HOST_MAX_BYTES
|
|
||||||
|| !host.is_ascii()
|
|
||||||
|| host.parse::<std::net::IpAddr>().is_ok()
|
|
||||||
{
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
let labels = host.split('.').collect::<Vec<_>>();
|
|
||||||
if labels.len() < 2 {
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
labels.iter().all(|label| {
|
|
||||||
!label.is_empty()
|
|
||||||
&& label.len() <= 63
|
|
||||||
&& !label.starts_with('-')
|
|
||||||
&& !label.ends_with('-')
|
|
||||||
&& label
|
|
||||||
.bytes()
|
|
||||||
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-')
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn build_admin_system_configs_payload(
|
pub fn build_admin_system_configs_payload(
|
||||||
entries: &[StoredSystemConfigEntry],
|
entries: &[StoredSystemConfigEntry],
|
||||||
) -> serde_json::Value {
|
) -> serde_json::Value {
|
||||||
@@ -2872,15 +2808,6 @@ pub fn parse_admin_system_config_update(
|
|||||||
)
|
)
|
||||||
})?;
|
})?;
|
||||||
}
|
}
|
||||||
EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY => {
|
|
||||||
value =
|
|
||||||
normalize_execution_extra_trusted_dns_hosts_config_value(value).map_err(|_| {
|
|
||||||
(
|
|
||||||
http::StatusCode::BAD_REQUEST,
|
|
||||||
json!({ "detail": "额外可信 Fake-IP 域名配置格式无效" }),
|
|
||||||
)
|
|
||||||
})?;
|
|
||||||
}
|
|
||||||
"module.important_notification.default_channel" => {
|
"module.important_notification.default_channel" => {
|
||||||
value = normalize_notification_channel_value(value).map_err(|_| {
|
value = normalize_notification_channel_value(value).map_err(|_| {
|
||||||
(
|
(
|
||||||
@@ -4621,36 +4548,6 @@ mod tests {
|
|||||||
.is_err());
|
.is_err());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn extra_trusted_dns_hosts_update_normalizes_exact_hostnames() {
|
|
||||||
let update = parse_admin_system_config_update(
|
|
||||||
"execution_extra_trusted_dns_hosts",
|
|
||||||
br#"{"value":[" API.Example.COM. ","api.example.com"]}"#,
|
|
||||||
)
|
|
||||||
.expect("valid extra trusted DNS hosts should parse");
|
|
||||||
assert_eq!(update.value, json!(["api.example.com"]));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn extra_trusted_dns_hosts_update_rejects_non_exact_hostnames() {
|
|
||||||
for value in [
|
|
||||||
r#"["*.example.com"]"#,
|
|
||||||
r#"["example.com:443"]"#,
|
|
||||||
r#"["https://example.com/path"]"#,
|
|
||||||
r#"["10.0.0.1"]"#,
|
|
||||||
r#"["example..com"]"#,
|
|
||||||
r#"["localhost"]"#,
|
|
||||||
] {
|
|
||||||
let body = format!(r#"{{"value":{value}}}"#);
|
|
||||||
let error = parse_admin_system_config_update(
|
|
||||||
"execution_extra_trusted_dns_hosts",
|
|
||||||
body.as_bytes(),
|
|
||||||
)
|
|
||||||
.expect_err("non-exact hostname should be rejected");
|
|
||||||
assert_eq!(error.0, http::StatusCode::BAD_REQUEST);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn legacy_notification_email_config_key_normalizes_to_important_notification() {
|
fn legacy_notification_email_config_key_normalizes_to_important_notification() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
|
|||||||
@@ -588,6 +588,68 @@ pub struct SettingsPayload {
|
|||||||
pub drain_deadline_ms: u64,
|
pub drain_deadline_ms: u64,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl SettingsPayload {
|
||||||
|
pub fn is_valid(&self) -> bool {
|
||||||
|
self.initial_stream_window_bytes > 0
|
||||||
|
&& u64::from(self.initial_stream_window_bytes)
|
||||||
|
<= MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64
|
||||||
|
&& self.min_window_update_bytes > 0
|
||||||
|
&& self.min_window_update_bytes <= self.initial_stream_window_bytes
|
||||||
|
&& self.drain_deadline_ms > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn negotiate(&self, initial_window_bytes: u32, drain_deadline_ms: u64) -> Self {
|
||||||
|
let window = self
|
||||||
|
.initial_stream_window_bytes
|
||||||
|
.min(initial_window_bytes)
|
||||||
|
.max(1);
|
||||||
|
Self {
|
||||||
|
initial_stream_window_bytes: window,
|
||||||
|
min_window_update_bytes: self.min_window_update_bytes.min((window / 4).max(1)),
|
||||||
|
drain_deadline_ms: self.drain_deadline_ms.min(drain_deadline_ms),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod settings_tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn negotiation_bounds_window_updates_by_the_smaller_window() {
|
||||||
|
let settings = SettingsPayload {
|
||||||
|
initial_stream_window_bytes: 512 * 1024,
|
||||||
|
min_window_update_bytes: 128 * 1024,
|
||||||
|
drain_deadline_ms: 30_000,
|
||||||
|
};
|
||||||
|
let negotiated = settings.negotiate(4 * 1024 * 1024, 1000);
|
||||||
|
assert_eq!(negotiated.initial_stream_window_bytes, 512 * 1024);
|
||||||
|
assert_eq!(negotiated.min_window_update_bytes, 128 * 1024);
|
||||||
|
assert_eq!(negotiated.drain_deadline_ms, 1000);
|
||||||
|
assert!(negotiated.is_valid());
|
||||||
|
let tiny = settings.negotiate(1, 1);
|
||||||
|
assert_eq!(tiny.initial_stream_window_bytes, 1);
|
||||||
|
assert_eq!(tiny.min_window_update_bytes, 1);
|
||||||
|
assert!(tiny.is_valid());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn invalid_window_settings_are_rejected() {
|
||||||
|
let mut settings = SettingsPayload {
|
||||||
|
initial_stream_window_bytes: 1024,
|
||||||
|
min_window_update_bytes: 256,
|
||||||
|
drain_deadline_ms: 1,
|
||||||
|
};
|
||||||
|
settings.min_window_update_bytes = 1025;
|
||||||
|
assert!(!settings.is_valid());
|
||||||
|
settings.min_window_update_bytes = 0;
|
||||||
|
assert!(!settings.is_valid());
|
||||||
|
settings.initial_stream_window_bytes = u32::MAX;
|
||||||
|
settings.min_window_update_bytes = 1;
|
||||||
|
assert!(!settings.is_valid());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||||
pub struct WindowUpdatePayload {
|
pub struct WindowUpdatePayload {
|
||||||
pub delta_bytes: u32,
|
pub delta_bytes: u32,
|
||||||
|
|||||||
@@ -18,8 +18,6 @@ futures-util.workspace = true
|
|||||||
sqlx = { workspace = true, features = ["postgres", "runtime-tokio-rustls", "chrono", "migrate", "macros"] }
|
sqlx = { workspace = true, features = ["postgres", "runtime-tokio-rustls", "chrono", "migrate", "macros"] }
|
||||||
serde_json.workspace = true
|
serde_json.workspace = true
|
||||||
sha2.workspace = true
|
sha2.workspace = true
|
||||||
|
tokio.workspace = true
|
||||||
tracing.workspace = true
|
tracing.workspace = true
|
||||||
uuid.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';
|
||||||
@@ -172,7 +172,7 @@ DO UPDATE SET
|
|||||||
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
||||||
ELSE EXCLUDED.error_type
|
ELSE EXCLUDED.error_type
|
||||||
END,
|
END,
|
||||||
error_message = NULL,
|
error_message = __AETHER_CANDIDATE_ERROR_MESSAGE__,
|
||||||
latency_ms = CASE
|
latency_ms = CASE
|
||||||
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||||
AND EXCLUDED.status <> request_candidates.status
|
AND EXCLUDED.status <> request_candidates.status
|
||||||
@@ -184,7 +184,7 @@ DO UPDATE SET
|
|||||||
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
||||||
END,
|
END,
|
||||||
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
||||||
extra_data = EXCLUDED.extra_data,
|
extra_data = __AETHER_CANDIDATE_EXTRA_DATA__,
|
||||||
required_capabilities = EXCLUDED.required_capabilities,
|
required_capabilities = EXCLUDED.required_capabilities,
|
||||||
created_at = CASE
|
created_at = CASE
|
||||||
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
||||||
@@ -275,7 +275,7 @@ DO UPDATE SET
|
|||||||
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
||||||
ELSE EXCLUDED.error_type
|
ELSE EXCLUDED.error_type
|
||||||
END,
|
END,
|
||||||
error_message = NULL,
|
error_message = __AETHER_CANDIDATE_ERROR_MESSAGE__,
|
||||||
latency_ms = CASE
|
latency_ms = CASE
|
||||||
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||||
AND EXCLUDED.status <> request_candidates.status
|
AND EXCLUDED.status <> request_candidates.status
|
||||||
@@ -287,7 +287,7 @@ DO UPDATE SET
|
|||||||
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
||||||
END,
|
END,
|
||||||
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
||||||
extra_data = EXCLUDED.extra_data,
|
extra_data = __AETHER_CANDIDATE_EXTRA_DATA__,
|
||||||
required_capabilities = EXCLUDED.required_capabilities,
|
required_capabilities = EXCLUDED.required_capabilities,
|
||||||
created_at = CASE
|
created_at = CASE
|
||||||
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
||||||
@@ -353,7 +353,7 @@ DO UPDATE SET
|
|||||||
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
||||||
ELSE EXCLUDED.error_type
|
ELSE EXCLUDED.error_type
|
||||||
END,
|
END,
|
||||||
error_message = NULL,
|
error_message = __AETHER_CANDIDATE_ERROR_MESSAGE__,
|
||||||
latency_ms = CASE
|
latency_ms = CASE
|
||||||
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||||
AND EXCLUDED.status <> request_candidates.status
|
AND EXCLUDED.status <> request_candidates.status
|
||||||
@@ -365,7 +365,7 @@ DO UPDATE SET
|
|||||||
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
||||||
END,
|
END,
|
||||||
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
||||||
extra_data = EXCLUDED.extra_data,
|
extra_data = __AETHER_CANDIDATE_EXTRA_DATA__,
|
||||||
required_capabilities = EXCLUDED.required_capabilities,
|
required_capabilities = EXCLUDED.required_capabilities,
|
||||||
created_at = CASE
|
created_at = CASE
|
||||||
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
||||||
@@ -446,6 +446,23 @@ fn postgres_candidate_upsert_sql(template: &str) -> String {
|
|||||||
)
|
)
|
||||||
.as_str(),
|
.as_str(),
|
||||||
)
|
)
|
||||||
|
.replace(
|
||||||
|
"__AETHER_CANDIDATE_ERROR_MESSAGE__",
|
||||||
|
"CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') \
|
||||||
|
AND EXCLUDED.status <> request_candidates.status \
|
||||||
|
THEN request_candidates.error_message \
|
||||||
|
WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') \
|
||||||
|
THEN request_candidates.error_message \
|
||||||
|
WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') \
|
||||||
|
THEN request_candidates.error_message \
|
||||||
|
ELSE COALESCE(EXCLUDED.error_message, request_candidates.error_message) END",
|
||||||
|
)
|
||||||
|
.replace(
|
||||||
|
"__AETHER_CANDIDATE_EXTRA_DATA__",
|
||||||
|
"CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') \
|
||||||
|
AND (EXCLUDED.status <> request_candidates.status OR EXCLUDED.extra_data IS NULL) \
|
||||||
|
THEN request_candidates.extra_data ELSE EXCLUDED.extra_data END",
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn postgres_sanitized_legacy_diagnostic_sql(
|
fn postgres_sanitized_legacy_diagnostic_sql(
|
||||||
@@ -1381,33 +1398,35 @@ mod tests {
|
|||||||
finished_at_unix_ms: Some(2),
|
finished_at_unix_ms: Some(2),
|
||||||
};
|
};
|
||||||
|
|
||||||
assert_eq!(sanitize_request_candidate_for_postgres(&mut candidate), 0);
|
assert_eq!(sanitize_request_candidate_for_postgres(&mut candidate), 1);
|
||||||
assert!(candidate.username.is_none());
|
assert!(candidate.username.is_none());
|
||||||
assert!(candidate.api_key_name.is_none());
|
assert!(candidate.api_key_name.is_none());
|
||||||
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
||||||
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
||||||
assert!(candidate.error_message.is_none());
|
assert_eq!(candidate.error_message.as_deref(), Some("bad�message"));
|
||||||
assert!(candidate.extra_data.is_none());
|
assert!(candidate.extra_data.is_none());
|
||||||
assert!(candidate.required_capabilities.is_none());
|
assert!(candidate.required_capabilities.is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn every_postgres_candidate_conflict_path_discards_legacy_diagnostics() {
|
fn every_postgres_candidate_conflict_path_preserves_errors_without_unrelated_legacy_data() {
|
||||||
for sql in [
|
for sql in [
|
||||||
UPSERT_SQL.as_str(),
|
UPSERT_SQL.as_str(),
|
||||||
UPSERT_CONFLICT_SQL.as_str(),
|
UPSERT_CONFLICT_SQL.as_str(),
|
||||||
UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL.as_str(),
|
UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL.as_str(),
|
||||||
] {
|
] {
|
||||||
assert!(sql.contains("error_message = NULL"));
|
assert!(
|
||||||
assert!(sql.contains("extra_data = EXCLUDED.extra_data"));
|
sql.contains("COALESCE(EXCLUDED.error_message, request_candidates.error_message)")
|
||||||
|
);
|
||||||
|
assert!(sql.contains("THEN request_candidates.extra_data ELSE EXCLUDED.extra_data END"));
|
||||||
assert!(sql.contains("required_capabilities = EXCLUDED.required_capabilities"));
|
assert!(sql.contains("required_capabilities = EXCLUDED.required_capabilities"));
|
||||||
assert!(!sql.contains("request_candidates.error_message"));
|
assert!(!sql.contains("COALESCE(request_candidates.extra_data"));
|
||||||
assert!(!sql.contains("request_candidates.extra_data"));
|
|
||||||
assert!(!sql.contains("request_candidates.required_capabilities"));
|
assert!(!sql.contains("request_candidates.required_capabilities"));
|
||||||
assert!(sql.contains("ELSE 'unclassified_skip' END"));
|
assert!(sql.contains("ELSE 'unclassified_skip' END"));
|
||||||
assert!(sql.contains("ELSE 'unclassified_error' END"));
|
assert!(sql.contains("ELSE 'unclassified_error' END"));
|
||||||
assert!(sql.contains("THEN 'first_byte_timeout'"));
|
assert!(sql.contains("THEN 'first_byte_timeout'"));
|
||||||
assert!(!sql.contains("__AETHER_SANITIZED_LEGACY_"));
|
assert!(!sql.contains("__AETHER_SANITIZED_LEGACY_"));
|
||||||
|
assert!(!sql.contains("__AETHER_CANDIDATE_"));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1541,10 +1560,11 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW())
|
|||||||
.fetch_one(repository.pool())
|
.fetch_one(repository.pool())
|
||||||
.await
|
.await
|
||||||
.expect("raw candidate diagnostics should load");
|
.expect("raw candidate diagnostics should load");
|
||||||
assert!(
|
assert_eq!(
|
||||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "error_message")
|
sqlx::Row::try_get::<Option<String>, _>(&raw, "error_message")
|
||||||
.expect("error_message should decode")
|
.expect("error_message should decode")
|
||||||
.is_none()
|
.as_deref(),
|
||||||
|
Some("bad�message")
|
||||||
);
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "skip_reason")
|
sqlx::Row::try_get::<Option<String>, _>(&raw, "skip_reason")
|
||||||
@@ -1576,7 +1596,7 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW())
|
|||||||
.expect("sanitized candidate should be readable");
|
.expect("sanitized candidate should be readable");
|
||||||
assert_eq!(rows.len(), 1);
|
assert_eq!(rows.len(), 1);
|
||||||
assert_eq!(rows[0].status, RequestCandidateStatus::Success);
|
assert_eq!(rows[0].status, RequestCandidateStatus::Success);
|
||||||
assert!(rows[0].error_message.is_none());
|
assert_eq!(rows[0].error_message.as_deref(), Some("bad�message"));
|
||||||
assert!(rows[0].extra_data.is_none());
|
assert!(rows[0].extra_data.is_none());
|
||||||
assert!(rows[0].required_capabilities.is_none());
|
assert!(rows[0].required_capabilities.is_none());
|
||||||
}
|
}
|
||||||
@@ -1589,12 +1609,92 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW())
|
|||||||
1
|
1
|
||||||
);
|
);
|
||||||
|
|
||||||
|
let mut cleanup_request_ids = vec![single_request_id, batch_request_id, healthy_request_id];
|
||||||
|
for write_path in 0..3 {
|
||||||
|
let request_id = uuid::Uuid::new_v4().to_string();
|
||||||
|
cleanup_request_ids.push(request_id.clone());
|
||||||
|
let mut failed = candidate(&request_id, uuid::Uuid::new_v4().to_string());
|
||||||
|
failed.status = RequestCandidateStatus::Failed;
|
||||||
|
failed.status_code = Some(400);
|
||||||
|
failed.error_message = Some("original upstream failure".to_string());
|
||||||
|
failed.is_cached = (write_path != 2).then_some(false);
|
||||||
|
failed.extra_data = Some(json!({
|
||||||
|
"upstream_response": {
|
||||||
|
"status_code": 400,
|
||||||
|
"headers": {"x-request-id": "original-upstream-id"},
|
||||||
|
"body": {"error": {"message": "original upstream failure", "param": "model"}}
|
||||||
|
},
|
||||||
|
"error_flow": {"status_code": 400, "message": "original upstream failure"}
|
||||||
|
}));
|
||||||
|
let mut pending = failed.clone();
|
||||||
|
pending.status = RequestCandidateStatus::Pending;
|
||||||
|
pending.status_code = None;
|
||||||
|
pending.error_message = None;
|
||||||
|
pending.extra_data = None;
|
||||||
|
repository
|
||||||
|
.upsert(pending.clone())
|
||||||
|
.await
|
||||||
|
.expect("pending seed should persist");
|
||||||
|
if write_path == 0 {
|
||||||
|
repository
|
||||||
|
.upsert(failed)
|
||||||
|
.await
|
||||||
|
.expect("single failure should persist");
|
||||||
|
} else {
|
||||||
|
repository
|
||||||
|
.upsert_many(vec![failed])
|
||||||
|
.await
|
||||||
|
.expect("batch failure should persist");
|
||||||
|
}
|
||||||
|
pending.status_code = Some(200);
|
||||||
|
pending.error_message = Some("late unrelated error".to_string());
|
||||||
|
pending.extra_data = Some(json!({
|
||||||
|
"upstream_response": {"status_code": 200, "body": "late unrelated response"}
|
||||||
|
}));
|
||||||
|
if write_path == 0 {
|
||||||
|
repository
|
||||||
|
.upsert(pending)
|
||||||
|
.await
|
||||||
|
.expect("late update should persist");
|
||||||
|
} else {
|
||||||
|
repository
|
||||||
|
.upsert_many(vec![pending])
|
||||||
|
.await
|
||||||
|
.expect("late batch should persist");
|
||||||
|
}
|
||||||
|
let stored = repository
|
||||||
|
.list_by_request_id(&request_id)
|
||||||
|
.await
|
||||||
|
.expect("failure should read");
|
||||||
|
assert_eq!(stored[0].status, RequestCandidateStatus::Failed);
|
||||||
|
assert_eq!(stored[0].status_code, Some(400));
|
||||||
|
assert_eq!(
|
||||||
|
stored[0].error_message.as_deref(),
|
||||||
|
Some("original upstream failure")
|
||||||
|
);
|
||||||
|
let extra = stored[0]
|
||||||
|
.extra_data
|
||||||
|
.as_ref()
|
||||||
|
.expect("failure details should remain");
|
||||||
|
assert_eq!(extra["upstream_response"]["status_code"], 400);
|
||||||
|
assert_eq!(
|
||||||
|
extra["upstream_response"]["headers"]["x-request-id"],
|
||||||
|
"original-upstream-id"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
extra["upstream_response"]["body"]["error"]["param"],
|
||||||
|
"model"
|
||||||
|
);
|
||||||
|
assert_eq!(extra["error_flow"]["message"], "original upstream failure");
|
||||||
|
let mut public = stored[0].clone();
|
||||||
|
public.sanitize_sensitive_diagnostics();
|
||||||
|
assert!(!serde_json::to_string(&public)
|
||||||
|
.expect("public record should serialize")
|
||||||
|
.contains("original upstream failure"));
|
||||||
|
}
|
||||||
|
|
||||||
sqlx::query("DELETE FROM request_candidates WHERE request_id = ANY($1)")
|
sqlx::query("DELETE FROM request_candidates WHERE request_id = ANY($1)")
|
||||||
.bind(vec![
|
.bind(cleanup_request_ids)
|
||||||
single_request_id,
|
|
||||||
batch_request_id,
|
|
||||||
healthy_request_id,
|
|
||||||
])
|
|
||||||
.execute(repository.pool())
|
.execute(repository.pool())
|
||||||
.await
|
.await
|
||||||
.expect("candidate NUL test rows should clean up");
|
.expect("candidate NUL test rows should clean up");
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -1,9 +1,9 @@
|
|||||||
use aether_data_contracts::repository::usage::{
|
use aether_data_contracts::repository::usage::{
|
||||||
canonical_usage_body_ref_for, parse_usage_body_ref, read_decompressed_usage_json,
|
canonical_usage_body_ref_for, parse_usage_body_ref, read_decompressed_usage_json,
|
||||||
usage_body_ref, ApiKeyLastUsedDelta, ManagementTokenCounterDelta, ProxyNodeCounterDelta,
|
usage_body_ref, ApiKeyLastUsedDelta, ManagementTokenCounterDelta, ProxyNodeCounterDelta,
|
||||||
StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow,
|
StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBodyPayload,
|
||||||
StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow,
|
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
|
||||||
StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
|
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
|
||||||
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
|
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
|
||||||
StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow,
|
StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow,
|
||||||
StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow,
|
StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow,
|
||||||
@@ -63,8 +63,25 @@ pub mod cleanup;
|
|||||||
// newly captured bodies always spill to usage_body_blobs and resolve through usage_http_audits.
|
// 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_INLINE_USAGE_BODY_BYTES: usize = 0;
|
||||||
const MAX_SUPPORTED_UNIX_SECS: u64 = 253_402_300_799;
|
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");
|
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)]
|
#[derive(Debug, Clone, PartialEq, Eq, Default)]
|
||||||
struct AggregateRangeSplit {
|
struct AggregateRangeSplit {
|
||||||
@@ -1374,8 +1391,8 @@ WHERE u.request_id = ANY($1)
|
|||||||
"#;
|
"#;
|
||||||
const UPSERT_USAGE_ROUTING_SNAPSHOT_SQL: &str =
|
const UPSERT_USAGE_ROUTING_SNAPSHOT_SQL: &str =
|
||||||
include_str!("queries/upsert_usage_routing_snapshot_sql.sql");
|
include_str!("queries/upsert_usage_routing_snapshot_sql.sql");
|
||||||
#[cfg(test)]
|
|
||||||
const UPSERT_USAGE_HTTP_AUDIT_SQL: &str = include_str!("queries/upsert_usage_http_audit_sql.sql");
|
const UPSERT_USAGE_HTTP_AUDIT_SQL: &str = include_str!("queries/upsert_usage_http_audit_sql.sql");
|
||||||
|
const UPSERT_USAGE_BODY_BLOB_SQL: &str = include_str!("queries/upsert_usage_body_blob_sql.sql");
|
||||||
const UPSERT_USAGE_SETTLEMENT_PRICING_SNAPSHOT_SQL: &str =
|
const UPSERT_USAGE_SETTLEMENT_PRICING_SNAPSHOT_SQL: &str =
|
||||||
include_str!("queries/upsert_usage_settlement_pricing_snapshot_sql.sql");
|
include_str!("queries/upsert_usage_settlement_pricing_snapshot_sql.sql");
|
||||||
|
|
||||||
@@ -2862,36 +2879,82 @@ ORDER BY request_count DESC, "usage".provider_name ASC
|
|||||||
Ok(items)
|
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 {
|
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let canonical_ref = usage_body_ref(&request_id, field);
|
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(&canonical_ref)
|
||||||
.bind(&request_id)
|
.bind(&request_id)
|
||||||
.bind(field.as_storage_field())
|
.bind(field.as_storage_field())
|
||||||
|
.bind(encoded_limit)
|
||||||
.fetch_optional(&self.pool)
|
.fetch_optional(&self.pool)
|
||||||
.await
|
.await
|
||||||
.map_postgres_err()?;
|
.map_postgres_err()?;
|
||||||
if let Some(row) = blob_row.as_ref() {
|
if let Some(row) = row {
|
||||||
let payload_gzip = row
|
return row
|
||||||
.try_get::<Vec<u8>, _>("payload_gzip")
|
.try_get::<Option<Vec<u8>>, _>("payload_gzip")
|
||||||
.map_postgres_err()?;
|
.map_postgres_err()?
|
||||||
return inflate_usage_json_value(&payload_gzip).map(Some);
|
.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 (inline_column, compressed_column) = usage_body_sql_columns(field);
|
||||||
let row = sqlx::query(&format!(
|
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(request_id)
|
||||||
|
.bind(json_limit)
|
||||||
|
.bind(encoded_limit)
|
||||||
.fetch_optional(&self.pool)
|
.fetch_optional(&self.pool)
|
||||||
.await
|
.await
|
||||||
.map_postgres_err()?;
|
.map_postgres_err()?;
|
||||||
row.as_ref()
|
let Some(row) = row else {
|
||||||
.map(|row| usage_json_column(row, "inline_body", "compressed_body", true))
|
return Ok(None);
|
||||||
.transpose()
|
};
|
||||||
.map(|value| value.and_then(|column| column.value))
|
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(
|
async fn hydrate_usage_body_refs(
|
||||||
@@ -10306,6 +10369,13 @@ impl UsageReadRepository for SqlxUsageReadRepository {
|
|||||||
Self::resolve_body_ref(self, body_ref).await
|
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(
|
async fn list_usage_audits(
|
||||||
&self,
|
&self,
|
||||||
query: &UsageAuditListQuery,
|
query: &UsageAuditListQuery,
|
||||||
@@ -13362,19 +13432,36 @@ async fn sync_usage_body_blob_storage<'e, E>(
|
|||||||
executor: E,
|
executor: E,
|
||||||
request_id: &str,
|
request_id: &str,
|
||||||
field: UsageBodyField,
|
field: UsageBodyField,
|
||||||
_value: Option<&Value>,
|
value: Option<&Value>,
|
||||||
_storage: &UsageBodyStorage,
|
storage: &UsageBodyStorage,
|
||||||
_clear_existing: bool,
|
clear_existing: bool,
|
||||||
) -> Result<(), DataLayerError>
|
) -> Result<(), DataLayerError>
|
||||||
where
|
where
|
||||||
E: sqlx::Executor<'e, Database = Postgres>,
|
E: sqlx::Executor<'e, Database = Postgres>,
|
||||||
{
|
{
|
||||||
let body_ref = usage_body_ref(request_id, field);
|
let body_ref = usage_body_ref(request_id, field);
|
||||||
sqlx::query(DELETE_USAGE_BODY_BLOB_SQL)
|
if clear_existing {
|
||||||
.bind(&body_ref)
|
sqlx::query(DELETE_USAGE_BODY_BLOB_SQL)
|
||||||
.execute(executor)
|
.bind(&body_ref)
|
||||||
.await
|
.execute(executor)
|
||||||
.map_postgres_err()?;
|
.await
|
||||||
|
.map_postgres_err()?;
|
||||||
|
} else if let Some(payload_gzip) = storage.detached_blob_bytes.as_ref() {
|
||||||
|
sqlx::query(UPSERT_USAGE_BODY_BLOB_SQL)
|
||||||
|
.bind(&body_ref)
|
||||||
|
.bind(request_id)
|
||||||
|
.bind(field.as_storage_field())
|
||||||
|
.bind(payload_gzip)
|
||||||
|
.execute(executor)
|
||||||
|
.await
|
||||||
|
.map_postgres_err()?;
|
||||||
|
} else if value.is_some() {
|
||||||
|
sqlx::query(DELETE_USAGE_BODY_BLOB_SQL)
|
||||||
|
.bind(&body_ref)
|
||||||
|
.execute(executor)
|
||||||
|
.await
|
||||||
|
.map_postgres_err()?;
|
||||||
|
}
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -13383,43 +13470,46 @@ async fn sync_usage_http_audit_storage<'e, E>(
|
|||||||
request_id: &str,
|
request_id: &str,
|
||||||
headers: &UsageHttpAuditHeaders<'_>,
|
headers: &UsageHttpAuditHeaders<'_>,
|
||||||
refs: &UsageHttpAuditRefs,
|
refs: &UsageHttpAuditRefs,
|
||||||
_states: &UsageHttpAuditStates,
|
states: &UsageHttpAuditStates,
|
||||||
body_capture_mode: &str,
|
body_capture_mode: &str,
|
||||||
) -> Result<(), DataLayerError>
|
) -> Result<(), DataLayerError>
|
||||||
where
|
where
|
||||||
E: sqlx::Executor<'e, Database = Postgres>,
|
E: sqlx::Executor<'e, Database = Postgres>,
|
||||||
{
|
{
|
||||||
if headers.any_present() || refs.any_present() || body_capture_mode != "none" {
|
if !headers.any_present()
|
||||||
return Err(DataLayerError::InvalidInput(
|
&& !refs.any_present()
|
||||||
"usage HTTP capture persistence is disabled".to_string(),
|
&& !states.any_present()
|
||||||
));
|
&& body_capture_mode == "none"
|
||||||
|
{
|
||||||
|
return Ok(());
|
||||||
}
|
}
|
||||||
|
|
||||||
sqlx::query(
|
sqlx::query(UPSERT_USAGE_HTTP_AUDIT_SQL)
|
||||||
r#"
|
.bind(request_id)
|
||||||
WITH deleted_audit AS (
|
.bind(headers.request_headers_json)
|
||||||
DELETE FROM usage_http_audits WHERE request_id = $1
|
.bind(headers.provider_request_headers_json)
|
||||||
)
|
.bind(headers.response_headers_json)
|
||||||
UPDATE usage
|
.bind(headers.client_response_headers_json)
|
||||||
SET request_headers = NULL,
|
.bind(refs.request_body_ref.as_deref())
|
||||||
request_body = NULL,
|
.bind(refs.provider_request_body_ref.as_deref())
|
||||||
provider_request_headers = NULL,
|
.bind(refs.response_body_ref.as_deref())
|
||||||
provider_request_body = NULL,
|
.bind(refs.client_response_body_ref.as_deref())
|
||||||
response_headers = NULL,
|
.bind(usage_body_capture_state_bind_text(
|
||||||
response_body = NULL,
|
states.request_body_state,
|
||||||
client_response_headers = NULL,
|
))
|
||||||
client_response_body = NULL,
|
.bind(usage_body_capture_state_bind_text(
|
||||||
request_body_compressed = NULL,
|
states.provider_request_body_state,
|
||||||
provider_request_body_compressed = NULL,
|
))
|
||||||
response_body_compressed = NULL,
|
.bind(usage_body_capture_state_bind_text(
|
||||||
client_response_body_compressed = NULL
|
states.response_body_state,
|
||||||
WHERE request_id = $1
|
))
|
||||||
"#,
|
.bind(usage_body_capture_state_bind_text(
|
||||||
)
|
states.client_response_body_state,
|
||||||
.bind(request_id)
|
))
|
||||||
.execute(executor)
|
.bind(body_capture_mode)
|
||||||
.await
|
.execute(executor)
|
||||||
.map_postgres_err()?;
|
.await
|
||||||
|
.map_postgres_err()?;
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -187,6 +187,152 @@ async fn pending_batch_is_opt_in_and_rejects_non_pending_before_connecting() {
|
|||||||
.contains("pending usage batch requires pending status"));
|
.contains("pending usage batch requires pending status"));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"]
|
||||||
|
async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() {
|
||||||
|
let factory = PostgresPoolFactory::new(PostgresPoolConfig {
|
||||||
|
database_url: std::env::var("AETHER_TEST_DATABASE_URL").unwrap(),
|
||||||
|
min_connections: 1,
|
||||||
|
max_connections: 2,
|
||||||
|
acquire_timeout_ms: 10_000,
|
||||||
|
idle_timeout_ms: 30_000,
|
||||||
|
max_lifetime_ms: 60_000,
|
||||||
|
statement_cache_capacity: 64,
|
||||||
|
require_ssl: false,
|
||||||
|
})
|
||||||
|
.unwrap();
|
||||||
|
let repository = SqlxUsageReadRepository::new(factory.connect_lazy().unwrap());
|
||||||
|
crate::run_migrations(repository.pool()).await.unwrap();
|
||||||
|
|
||||||
|
for batch in [false, true] {
|
||||||
|
let request_id = format!("req-full-capture-{}", uuid::Uuid::new_v4().simple());
|
||||||
|
let now_unix_secs = Utc::now().timestamp() as u64;
|
||||||
|
let mut pending = fast_clear_usage_record(
|
||||||
|
&request_id,
|
||||||
|
"full-capture-test",
|
||||||
|
now_unix_secs,
|
||||||
|
false,
|
||||||
|
UsageBodyCaptureState::Inline,
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
pending.request_headers =
|
||||||
|
Some(json!({"content-type": "application/json", "authorization": "Bearer private"}));
|
||||||
|
pending.request_body =
|
||||||
|
Some(json!({"messages": [{"role": "user", "content": "original request"}]}));
|
||||||
|
pending.request_body_state = Some(UsageBodyCaptureState::Inline);
|
||||||
|
pending.provider_request_body = Some(json!({"input": "provider request"}));
|
||||||
|
pending.response_body = Some(json!("pending response"));
|
||||||
|
pending.response_body_state = Some(UsageBodyCaptureState::Inline);
|
||||||
|
pending.client_response_body = Some(json!("pending client response"));
|
||||||
|
pending.client_response_body_state = Some(UsageBodyCaptureState::Inline);
|
||||||
|
if batch {
|
||||||
|
repository
|
||||||
|
.upsert_pending_many(vec![pending.clone()])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
} else {
|
||||||
|
repository.upsert(pending.clone()).await.unwrap();
|
||||||
|
}
|
||||||
|
for (field, expected) in [
|
||||||
|
(UsageBodyField::RequestBody, pending.request_body.as_ref()),
|
||||||
|
(
|
||||||
|
UsageBodyField::ProviderRequestBody,
|
||||||
|
pending.provider_request_body.as_ref(),
|
||||||
|
),
|
||||||
|
(UsageBodyField::ResponseBody, pending.response_body.as_ref()),
|
||||||
|
(
|
||||||
|
UsageBodyField::ClientResponseBody,
|
||||||
|
pending.client_response_body.as_ref(),
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
assert_eq!(
|
||||||
|
repository
|
||||||
|
.resolve_body_ref(&usage_body_ref(&request_id, field))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.as_ref(),
|
||||||
|
expected,
|
||||||
|
"batch={batch}, field={field:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut terminal = fast_clear_usage_record(
|
||||||
|
&request_id,
|
||||||
|
"full-capture-test",
|
||||||
|
now_unix_secs,
|
||||||
|
true,
|
||||||
|
UsageBodyCaptureState::None,
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
terminal.provider_request_body_state = None;
|
||||||
|
terminal.response_headers =
|
||||||
|
Some(json!({"content-type": "text/event-stream", "set-cookie": "private"}));
|
||||||
|
terminal.response_body = Some(json!(format!(
|
||||||
|
"data: {}\n\ndata: [DONE]\n\n",
|
||||||
|
"streamed text".repeat(8192)
|
||||||
|
)));
|
||||||
|
terminal.response_body_state = Some(UsageBodyCaptureState::Inline);
|
||||||
|
terminal.client_response_body = Some(json!({"output": "final response"}));
|
||||||
|
terminal.client_response_body_state = Some(UsageBodyCaptureState::Inline);
|
||||||
|
repository.upsert(terminal.clone()).await.unwrap();
|
||||||
|
|
||||||
|
let stored = repository
|
||||||
|
.find_by_request_id_shallow(&request_id)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
stored.request_headers,
|
||||||
|
Some(json!({"content-type": "application/json", "authorization": "Bearer private"}))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
stored.response_headers,
|
||||||
|
Some(json!({"content-type": "text/event-stream", "set-cookie": "private"}))
|
||||||
|
);
|
||||||
|
for (field, expected) in [
|
||||||
|
(UsageBodyField::RequestBody, pending.request_body.as_ref()),
|
||||||
|
(
|
||||||
|
UsageBodyField::ProviderRequestBody,
|
||||||
|
pending.provider_request_body.as_ref(),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
UsageBodyField::ResponseBody,
|
||||||
|
terminal.response_body.as_ref(),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
UsageBodyField::ClientResponseBody,
|
||||||
|
terminal.client_response_body.as_ref(),
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
assert_eq!(
|
||||||
|
stored.body_state(field),
|
||||||
|
Some(UsageBodyCaptureState::Reference)
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
stored.body_ref(field),
|
||||||
|
Some(usage_body_ref(&request_id, field).as_str())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
repository
|
||||||
|
.resolve_body_ref(stored.body_ref(field).unwrap())
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.as_ref(),
|
||||||
|
expected,
|
||||||
|
"batch={batch}, field={field:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
let legacy_content_present: bool = sqlx::query_scalar("SELECT request_body IS NOT NULL OR request_headers IS NOT NULL OR response_body IS NOT NULL FROM usage WHERE request_id = $1")
|
||||||
|
.bind(&request_id).fetch_one(repository.pool()).await.unwrap();
|
||||||
|
assert!(!legacy_content_present);
|
||||||
|
sqlx::query("DELETE FROM usage WHERE request_id = $1")
|
||||||
|
.bind(&request_id)
|
||||||
|
.execute(repository.pool())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"]
|
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"]
|
||||||
async fn live_stale_terminal_event_is_a_full_transaction_noop() {
|
async fn live_stale_terminal_event_is_a_full_transaction_noop() {
|
||||||
@@ -253,7 +399,7 @@ async fn live_stale_terminal_event_is_a_full_transaction_noop() {
|
|||||||
.unwrap(),
|
.unwrap(),
|
||||||
);
|
);
|
||||||
let settlement_before = sqlx::query(
|
let settlement_before = sqlx::query(
|
||||||
"SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = $1",
|
"SELECT billing_status, billing_total_cost_usd::DOUBLE PRECISION AS billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = $1",
|
||||||
)
|
)
|
||||||
.bind(&request_id)
|
.bind(&request_id)
|
||||||
.fetch_one(repository.pool())
|
.fetch_one(repository.pool())
|
||||||
@@ -318,7 +464,7 @@ async fn live_stale_terminal_event_is_a_full_transaction_noop() {
|
|||||||
.unwrap(),
|
.unwrap(),
|
||||||
);
|
);
|
||||||
let settlement_after = sqlx::query(
|
let settlement_after = sqlx::query(
|
||||||
"SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = $1",
|
"SELECT billing_status, billing_total_cost_usd::DOUBLE PRECISION AS billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = $1",
|
||||||
)
|
)
|
||||||
.bind(&request_id)
|
.bind(&request_id)
|
||||||
.fetch_one(repository.pool())
|
.fetch_one(repository.pool())
|
||||||
@@ -542,7 +688,7 @@ async fn live_pending_batch_persists_auxiliary_state_and_preserves_terminal_conf
|
|||||||
.fetch_one(repository.pool())
|
.fetch_one(repository.pool())
|
||||||
.await
|
.await
|
||||||
.expect("HTTP audit count should be readable");
|
.expect("HTTP audit count should be readable");
|
||||||
assert_eq!(http_count, 0);
|
assert_eq!(http_count, 1);
|
||||||
let blob_count = sqlx::query_scalar::<_, i64>(
|
let blob_count = sqlx::query_scalar::<_, i64>(
|
||||||
"SELECT COUNT(*)::BIGINT FROM usage_body_blobs WHERE request_id = $1",
|
"SELECT COUNT(*)::BIGINT FROM usage_body_blobs WHERE request_id = $1",
|
||||||
)
|
)
|
||||||
@@ -550,7 +696,23 @@ async fn live_pending_batch_persists_auxiliary_state_and_preserves_terminal_conf
|
|||||||
.fetch_one(repository.pool())
|
.fetch_one(repository.pool())
|
||||||
.await
|
.await
|
||||||
.expect("body blob count should be readable");
|
.expect("body blob count should be readable");
|
||||||
assert_eq!(blob_count, 0);
|
assert_eq!(blob_count, 4);
|
||||||
|
let captured = repository
|
||||||
|
.find_by_request_id_shallow(&rich_request_id)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
captured.request_headers,
|
||||||
|
Some(json!({"x-request": "request-value"}))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
repository
|
||||||
|
.resolve_body_ref(captured.body_ref(UsageBodyField::RequestBody).unwrap())
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
Some(json!({"messages": [{"role": "user", "content": "hello"}]}))
|
||||||
|
);
|
||||||
|
|
||||||
let routing = sqlx::query(
|
let routing = sqlx::query(
|
||||||
"SELECT candidate_id, candidate_index, selected_provider_api_key_id FROM usage_routing_snapshots WHERE request_id = $1",
|
"SELECT candidate_id, candidate_index, selected_provider_api_key_id FROM usage_routing_snapshots WHERE request_id = $1",
|
||||||
@@ -667,7 +829,7 @@ async fn live_pending_batch_and_terminal_upserts_count_each_provider_request_onc
|
|||||||
|
|
||||||
let suffix = uuid::Uuid::new_v4().simple().to_string();
|
let suffix = uuid::Uuid::new_v4().simple().to_string();
|
||||||
let provider_name = format!("pending-terminal-race-provider-{suffix}");
|
let provider_name = format!("pending-terminal-race-provider-{suffix}");
|
||||||
let provider_key_id = format!("pending-terminal-race-key-{suffix}");
|
let provider_key_id = uuid::Uuid::new_v4().to_string();
|
||||||
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
||||||
let request_ids = (0..REQUESTS)
|
let request_ids = (0..REQUESTS)
|
||||||
.map(|index| format!("req-pending-terminal-race-{index}-{suffix}"))
|
.map(|index| format!("req-pending-terminal-race-{index}-{suffix}"))
|
||||||
@@ -793,7 +955,7 @@ async fn live_first_byte_fast_path_is_atomic_and_preserves_terminal_state() {
|
|||||||
let existing_request_id = format!("req-first-byte-existing-{suffix}");
|
let existing_request_id = format!("req-first-byte-existing-{suffix}");
|
||||||
let metadata_fill_request_id = format!("req-first-byte-metadata-fill-{suffix}");
|
let metadata_fill_request_id = format!("req-first-byte-metadata-fill-{suffix}");
|
||||||
let provider_name = format!("first-byte-fast-{suffix}");
|
let provider_name = format!("first-byte-fast-{suffix}");
|
||||||
let missing_provider_key_id = format!("key-first-byte-missing-{suffix}");
|
let missing_provider_key_id = uuid::Uuid::new_v4().to_string();
|
||||||
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
||||||
|
|
||||||
let mut missing_first_byte = first_byte_usage_record(
|
let mut missing_first_byte = first_byte_usage_record(
|
||||||
@@ -1064,7 +1226,7 @@ async fn live_first_byte_reads_provider_contribution_after_waiting_for_canonical
|
|||||||
let suffix = uuid::Uuid::new_v4().simple().to_string();
|
let suffix = uuid::Uuid::new_v4().simple().to_string();
|
||||||
let request_id = format!("req-first-byte-lock-snapshot-{suffix}");
|
let request_id = format!("req-first-byte-lock-snapshot-{suffix}");
|
||||||
let provider_name = format!("first-byte-lock-snapshot-{suffix}");
|
let provider_name = format!("first-byte-lock-snapshot-{suffix}");
|
||||||
let provider_key_id = format!("key-first-byte-lock-snapshot-{suffix}");
|
let provider_key_id = uuid::Uuid::new_v4().to_string();
|
||||||
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
||||||
let mut pending = first_byte_usage_record(
|
let mut pending = first_byte_usage_record(
|
||||||
&request_id,
|
&request_id,
|
||||||
@@ -1209,14 +1371,14 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
let request_b = format!("req-first-byte-batch-b-{suffix}");
|
let request_b = format!("req-first-byte-batch-b-{suffix}");
|
||||||
let request_missing = format!("req-first-byte-batch-missing-{suffix}");
|
let request_missing = format!("req-first-byte-batch-missing-{suffix}");
|
||||||
let request_terminal = format!("req-first-byte-batch-terminal-{suffix}");
|
let request_terminal = format!("req-first-byte-batch-terminal-{suffix}");
|
||||||
let missing_provider_key_id = format!("key-first-byte-batch-missing-{suffix}");
|
let missing_provider_key_id = uuid::Uuid::new_v4().to_string();
|
||||||
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
|
||||||
|
|
||||||
let mut pending_a = first_byte_usage_record(
|
let mut pending_a = first_byte_usage_record(
|
||||||
&request_a,
|
&request_a,
|
||||||
&provider_name,
|
&provider_name,
|
||||||
now_unix_secs,
|
now_unix_secs,
|
||||||
Some(json!({"seed": "a"})),
|
Some(json!({"trace_id": "seed-a"})),
|
||||||
);
|
);
|
||||||
pending_a.status = "pending".to_string();
|
pending_a.status = "pending".to_string();
|
||||||
pending_a.first_byte_time_ms = None;
|
pending_a.first_byte_time_ms = None;
|
||||||
@@ -1241,7 +1403,7 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
);
|
);
|
||||||
terminal.is_stream = Some(true);
|
terminal.is_stream = Some(true);
|
||||||
terminal.first_byte_time_ms = Some(44);
|
terminal.first_byte_time_ms = Some(44);
|
||||||
terminal.request_metadata = Some(json!({"terminal": true}));
|
terminal.request_metadata = Some(json!({"trace_id": "terminal"}));
|
||||||
|
|
||||||
repository
|
repository
|
||||||
.upsert(pending_a)
|
.upsert(pending_a)
|
||||||
@@ -1260,7 +1422,7 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
&request_a,
|
&request_a,
|
||||||
&provider_name,
|
&provider_name,
|
||||||
now_unix_secs + 1,
|
now_unix_secs + 1,
|
||||||
Some(json!({"incoming": "a"})),
|
Some(json!({"trace_id": "incoming-a"})),
|
||||||
);
|
);
|
||||||
first_a.first_byte_time_ms = Some(30);
|
first_a.first_byte_time_ms = Some(30);
|
||||||
first_a.has_format_conversion = None;
|
first_a.has_format_conversion = None;
|
||||||
@@ -1272,7 +1434,7 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
&request_b,
|
&request_b,
|
||||||
&provider_name,
|
&provider_name,
|
||||||
now_unix_secs + 1,
|
now_unix_secs + 1,
|
||||||
Some(json!({"incoming": "b"})),
|
Some(json!({"trace_id": "incoming-b"})),
|
||||||
);
|
);
|
||||||
first_b.has_format_conversion = Some(true);
|
first_b.has_format_conversion = Some(true);
|
||||||
|
|
||||||
@@ -1280,7 +1442,7 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
&request_terminal,
|
&request_terminal,
|
||||||
&provider_name,
|
&provider_name,
|
||||||
now_unix_secs + 2,
|
now_unix_secs + 2,
|
||||||
Some(json!({"late": true})),
|
Some(json!({"trace_id": "late"})),
|
||||||
);
|
);
|
||||||
late_terminal.first_byte_time_ms = Some(3);
|
late_terminal.first_byte_time_ms = Some(3);
|
||||||
late_terminal.has_format_conversion = None;
|
late_terminal.has_format_conversion = None;
|
||||||
@@ -1288,7 +1450,7 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
&request_missing,
|
&request_missing,
|
||||||
&provider_name,
|
&provider_name,
|
||||||
now_unix_secs + 1,
|
now_unix_secs + 1,
|
||||||
Some(json!({"incoming": "missing"})),
|
Some(json!({"trace_id": "incoming-missing"})),
|
||||||
);
|
);
|
||||||
first_missing.provider_api_key_id = Some(missing_provider_key_id.clone());
|
first_missing.provider_api_key_id = Some(missing_provider_key_id.clone());
|
||||||
|
|
||||||
@@ -1345,7 +1507,7 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
row_a
|
row_a
|
||||||
.try_get::<Option<serde_json::Value>, _>("request_metadata")
|
.try_get::<Option<serde_json::Value>, _>("request_metadata")
|
||||||
.unwrap(),
|
.unwrap(),
|
||||||
Some(json!({"seed": "a"})),
|
Some(json!({"trace_id": "seed-a"})),
|
||||||
"existing metadata remains authoritative"
|
"existing metadata remains authoritative"
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -1364,7 +1526,7 @@ async fn live_first_byte_batch_preserves_duplicate_order_and_terminal_guards() {
|
|||||||
row_b
|
row_b
|
||||||
.try_get::<Option<serde_json::Value>, _>("request_metadata")
|
.try_get::<Option<serde_json::Value>, _>("request_metadata")
|
||||||
.unwrap(),
|
.unwrap(),
|
||||||
Some(json!({"incoming": "b"}))
|
Some(json!({"trace_id": "incoming-b"}))
|
||||||
);
|
);
|
||||||
|
|
||||||
let row_terminal = rows
|
let row_terminal = rows
|
||||||
@@ -2044,7 +2206,7 @@ async fn live_provider_performance_grouping_sets_matches_separate_queries() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
#[ignore = "requires AETHER_TEST_DATABASE_URL and a populated PostgreSQL database"]
|
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"]
|
||||||
async fn live_dashboard_daily_breakdown_uses_canonical_covering_read_path() {
|
async fn live_dashboard_daily_breakdown_uses_canonical_covering_read_path() {
|
||||||
let database_url = std::env::var("AETHER_TEST_DATABASE_URL")
|
let database_url = std::env::var("AETHER_TEST_DATABASE_URL")
|
||||||
.expect("AETHER_TEST_DATABASE_URL must point at the test database");
|
.expect("AETHER_TEST_DATABASE_URL must point at the test database");
|
||||||
@@ -2062,6 +2224,20 @@ async fn live_dashboard_daily_breakdown_uses_canonical_covering_read_path() {
|
|||||||
let repository =
|
let repository =
|
||||||
SqlxUsageReadRepository::new(factory.connect_lazy().expect("lazy pool should build"));
|
SqlxUsageReadRepository::new(factory.connect_lazy().expect("lazy pool should build"));
|
||||||
let until = Utc::now().timestamp().max(0) as u64;
|
let until = Utc::now().timestamp().max(0) as u64;
|
||||||
|
crate::run_migrations(repository.pool()).await.unwrap();
|
||||||
|
let request_id = format!("daily-breakdown-{}", uuid::Uuid::new_v4().simple());
|
||||||
|
let provider_name = format!("daily-provider-{}", uuid::Uuid::new_v4().simple());
|
||||||
|
repository
|
||||||
|
.upsert(fast_clear_usage_record(
|
||||||
|
&request_id,
|
||||||
|
&provider_name,
|
||||||
|
until.saturating_sub(60),
|
||||||
|
true,
|
||||||
|
UsageBodyCaptureState::None,
|
||||||
|
None,
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
let started = std::time::Instant::now();
|
let started = std::time::Instant::now();
|
||||||
let rows = repository
|
let rows = repository
|
||||||
.list_dashboard_daily_breakdown(&UsageDashboardDailyBreakdownQuery {
|
.list_dashboard_daily_breakdown(&UsageDashboardDailyBreakdownQuery {
|
||||||
@@ -2077,7 +2253,17 @@ async fn live_dashboard_daily_breakdown_uses_canonical_covering_read_path() {
|
|||||||
started.elapsed(),
|
started.elapsed(),
|
||||||
rows.len()
|
rows.len()
|
||||||
);
|
);
|
||||||
assert!(!rows.is_empty());
|
let seeded = rows
|
||||||
|
.iter()
|
||||||
|
.find(|row| row.provider == provider_name)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(seeded.requests, 1);
|
||||||
|
assert_eq!(seeded.total_tokens, 2);
|
||||||
|
sqlx::query("DELETE FROM \"usage\" WHERE request_id = $1")
|
||||||
|
.bind(&request_id)
|
||||||
|
.execute(repository.pool())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
@@ -4088,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]
|
#[test]
|
||||||
fn prepare_usage_body_storage_compresses_large_payloads() {
|
fn prepare_usage_body_storage_compresses_large_payloads() {
|
||||||
let payload = json!({
|
let payload = json!({
|
||||||
|
|||||||
@@ -81,12 +81,12 @@ fn select_video_task_full_columns() -> String {
|
|||||||
|
|
||||||
fn select_video_task_claim_columns() -> String {
|
fn select_video_task_claim_columns() -> String {
|
||||||
select_video_task_columns(
|
select_video_task_columns(
|
||||||
"NULL::TEXT",
|
"prompt",
|
||||||
"NULL::jsonb",
|
"NULL::jsonb",
|
||||||
"NULL::INTEGER",
|
"duration_seconds",
|
||||||
"NULL::TEXT",
|
"resolution",
|
||||||
"NULL::TEXT",
|
"aspect_ratio",
|
||||||
"NULL::TEXT",
|
"size",
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1176,7 +1176,7 @@ fn map_video_task_row(row: &PgRow) -> Result<StoredVideoTask, DataLayerError> {
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::{update_if_active_sql, upsert_sql, SqlxVideoTaskRepository};
|
use super::{claim_due_sql, update_if_active_sql, upsert_sql, SqlxVideoTaskRepository};
|
||||||
use crate::{PostgresPoolConfig, PostgresPoolFactory};
|
use crate::{PostgresPoolConfig, PostgresPoolFactory};
|
||||||
use aether_data_contracts::repository::video_tasks::{
|
use aether_data_contracts::repository::video_tasks::{
|
||||||
UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository,
|
UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository,
|
||||||
@@ -1240,6 +1240,142 @@ mod tests {
|
|||||||
assert!(update.contains("created_at = COALESCE(created_at, TO_TIMESTAMP($34))"));
|
assert!(update.contains("created_at = COALESCE(created_at, TO_TIMESTAMP($34))"));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn poll_claim_returns_business_fields_required_by_identity_guards() {
|
||||||
|
let sql = claim_due_sql();
|
||||||
|
for field in [
|
||||||
|
"prompt",
|
||||||
|
"duration_seconds",
|
||||||
|
"resolution",
|
||||||
|
"aspect_ratio",
|
||||||
|
"size",
|
||||||
|
] {
|
||||||
|
assert!(
|
||||||
|
sql.contains(&format!("{field} AS {field}")),
|
||||||
|
"claim must retain {field}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
assert!(sql.contains("NULL::jsonb AS original_request_body"));
|
||||||
|
assert!(sql.contains("FOR UPDATE SKIP LOCKED"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"]
|
||||||
|
async fn live_video_task_capture_claim_and_completion_preserve_business_fields() {
|
||||||
|
let database_url = std::env::var("AETHER_TEST_DATABASE_URL")
|
||||||
|
.expect("AETHER_TEST_DATABASE_URL must point at the test database");
|
||||||
|
let pool = sqlx::postgres::PgPoolOptions::new()
|
||||||
|
.max_connections(1)
|
||||||
|
.connect(&database_url)
|
||||||
|
.await
|
||||||
|
.expect("test database should connect");
|
||||||
|
crate::run_migrations(&pool)
|
||||||
|
.await
|
||||||
|
.expect("test database should migrate");
|
||||||
|
sqlx::query("CREATE TEMP TABLE video_tasks (LIKE public.video_tasks INCLUDING ALL)")
|
||||||
|
.execute(&pool)
|
||||||
|
.await
|
||||||
|
.expect("isolated task table should be created");
|
||||||
|
let repository = SqlxVideoTaskRepository::new(pool);
|
||||||
|
for api_format in ["openai:video", "gemini:video"] {
|
||||||
|
let task_id = uuid::Uuid::new_v4().to_string();
|
||||||
|
let original = UpsertVideoTask {
|
||||||
|
id: task_id.clone(),
|
||||||
|
short_id: Some(uuid::Uuid::new_v4().simple().to_string()[..16].to_string()),
|
||||||
|
request_id: format!("request-{task_id}"),
|
||||||
|
user_id: None,
|
||||||
|
api_key_id: None,
|
||||||
|
username: Some("alice".to_string()),
|
||||||
|
api_key_name: Some("video-client".to_string()),
|
||||||
|
external_task_id: Some("upstream-task-1".to_string()),
|
||||||
|
provider_id: None,
|
||||||
|
endpoint_id: None,
|
||||||
|
key_id: None,
|
||||||
|
client_api_format: Some(api_format.to_string()),
|
||||||
|
provider_api_format: Some(api_format.to_string()),
|
||||||
|
format_converted: false,
|
||||||
|
model: Some("video-model".to_string()),
|
||||||
|
prompt: Some("business prompt".to_string()),
|
||||||
|
original_request_body: Some(serde_json::json!({"token": "private"})),
|
||||||
|
duration_seconds: Some(8),
|
||||||
|
resolution: Some("1080p".to_string()),
|
||||||
|
aspect_ratio: Some("16:9".to_string()),
|
||||||
|
size: Some("1920x1080".to_string()),
|
||||||
|
status: VideoTaskStatus::Submitted,
|
||||||
|
progress_percent: 0,
|
||||||
|
progress_message: None,
|
||||||
|
retry_count: 0,
|
||||||
|
poll_interval_seconds: 10,
|
||||||
|
next_poll_at_unix_secs: Some(10),
|
||||||
|
poll_count: 0,
|
||||||
|
max_poll_count: 360,
|
||||||
|
created_at_unix_ms: 1,
|
||||||
|
submitted_at_unix_secs: Some(1),
|
||||||
|
completed_at_unix_secs: None,
|
||||||
|
updated_at_unix_secs: 1,
|
||||||
|
error_code: None,
|
||||||
|
error_message: None,
|
||||||
|
video_url: None,
|
||||||
|
request_metadata: Some(serde_json::json!({"authorization": "private"})),
|
||||||
|
};
|
||||||
|
let stored = repository
|
||||||
|
.upsert(original.clone())
|
||||||
|
.await
|
||||||
|
.expect("task should persist");
|
||||||
|
assert_eq!(stored.prompt, original.prompt);
|
||||||
|
assert_eq!(stored.username, original.username);
|
||||||
|
assert_eq!(stored.api_key_name, original.api_key_name);
|
||||||
|
assert!(stored.original_request_body.is_none());
|
||||||
|
assert!(stored.request_metadata.is_none());
|
||||||
|
|
||||||
|
let mut claimed = repository
|
||||||
|
.claim_due(20, 50, 10)
|
||||||
|
.await
|
||||||
|
.expect("task should be claimed");
|
||||||
|
assert_eq!(claimed.len(), 1);
|
||||||
|
let mut completion: UpsertVideoTask = claimed.pop().expect("claimed task").into();
|
||||||
|
stored
|
||||||
|
.ensure_immutable_identity_matches(&completion)
|
||||||
|
.expect("claim must preserve task identity");
|
||||||
|
assert_eq!(completion.prompt, original.prompt);
|
||||||
|
let mut mismatched = completion.clone();
|
||||||
|
mismatched.duration_seconds = Some(99);
|
||||||
|
assert!(repository
|
||||||
|
.update_if_active(mismatched)
|
||||||
|
.await
|
||||||
|
.expect("guarded update should execute")
|
||||||
|
.is_none());
|
||||||
|
completion.status = VideoTaskStatus::Completed;
|
||||||
|
completion.progress_percent = 100;
|
||||||
|
completion.next_poll_at_unix_secs = None;
|
||||||
|
completion.completed_at_unix_secs = Some(21);
|
||||||
|
completion.updated_at_unix_secs = 21;
|
||||||
|
completion.video_url = Some(
|
||||||
|
"https://cdn.example.test/video.mp4?alt=media&signature=a%2Fb%2Bc%3D&part=2&part=1"
|
||||||
|
.to_string(),
|
||||||
|
);
|
||||||
|
let completed = repository
|
||||||
|
.update_if_active(completion.clone())
|
||||||
|
.await
|
||||||
|
.expect("completion should execute")
|
||||||
|
.expect("matching active task should complete");
|
||||||
|
assert_eq!(completed.video_url, completion.video_url);
|
||||||
|
let reloaded = repository
|
||||||
|
.find(VideoTaskLookupKey::Id(&task_id))
|
||||||
|
.await
|
||||||
|
.expect("task should reload")
|
||||||
|
.expect("task should exist");
|
||||||
|
assert_eq!(reloaded.status, VideoTaskStatus::Completed);
|
||||||
|
assert_eq!(reloaded.prompt, original.prompt);
|
||||||
|
assert_eq!(reloaded.video_url, completion.video_url);
|
||||||
|
assert_eq!(reloaded.duration_seconds, original.duration_seconds);
|
||||||
|
assert_eq!(reloaded.size, original.size);
|
||||||
|
assert_eq!(reloaded.username, original.username);
|
||||||
|
assert!(reloaded.request_metadata.is_none());
|
||||||
|
}
|
||||||
|
repository.pool().close().await;
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn repository_constructs_from_lazy_pool() {
|
async fn repository_constructs_from_lazy_pool() {
|
||||||
let repository = SqlxVideoTaskRepository::new(build_pool());
|
let repository = SqlxVideoTaskRepository::new(build_pool());
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ pub use types::{
|
|||||||
build_decision_trace, derive_request_candidate_final_status,
|
build_decision_trace, derive_request_candidate_final_status,
|
||||||
request_candidate_lifecycle_would_regress, sanitize_request_candidate_api_formats,
|
request_candidate_lifecycle_would_regress, sanitize_request_candidate_api_formats,
|
||||||
sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data,
|
sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data,
|
||||||
|
sanitize_request_candidate_extra_data_for_persistence,
|
||||||
sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason,
|
sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason,
|
||||||
DecisionTrace, DecisionTraceCandidate, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
DecisionTrace, DecisionTraceCandidate, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||||
RequestCandidateFinalStatus, RequestCandidateReadRepository, RequestCandidateRepository,
|
RequestCandidateFinalStatus, RequestCandidateReadRepository, RequestCandidateRepository,
|
||||||
|
|||||||
@@ -224,6 +224,21 @@ pub struct StoredRequestCandidate {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl StoredRequestCandidate {
|
impl StoredRequestCandidate {
|
||||||
|
pub fn sanitize_for_persistence(&mut self) {
|
||||||
|
self.username = None;
|
||||||
|
self.api_key_name = None;
|
||||||
|
self.skip_reason = sanitize_request_candidate_skip_reason(self.skip_reason.take());
|
||||||
|
self.error_type = sanitize_request_candidate_error_type(self.error_type.take());
|
||||||
|
self.error_message = self
|
||||||
|
.error_message
|
||||||
|
.take()
|
||||||
|
.map(limit_candidate_diagnostic_text);
|
||||||
|
self.extra_data =
|
||||||
|
sanitize_request_candidate_extra_data_for_persistence(self.extra_data.take());
|
||||||
|
self.required_capabilities =
|
||||||
|
sanitize_request_candidate_required_capabilities(self.required_capabilities.take());
|
||||||
|
}
|
||||||
|
|
||||||
pub fn sanitize_sensitive_diagnostics(&mut self) {
|
pub fn sanitize_sensitive_diagnostics(&mut self) {
|
||||||
self.username = None;
|
self.username = None;
|
||||||
self.api_key_name = None;
|
self.api_key_name = None;
|
||||||
@@ -349,7 +364,7 @@ impl StoredRequestCandidate {
|
|||||||
started_at_unix_ms,
|
started_at_unix_ms,
|
||||||
finished_at_unix_ms,
|
finished_at_unix_ms,
|
||||||
};
|
};
|
||||||
candidate.sanitize_sensitive_diagnostics();
|
candidate.sanitize_for_persistence();
|
||||||
Ok(candidate)
|
Ok(candidate)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -386,7 +401,7 @@ impl RequestCandidateTrace {
|
|||||||
attempted_only: bool,
|
attempted_only: bool,
|
||||||
) -> Option<Self> {
|
) -> Option<Self> {
|
||||||
for candidate in &mut all_candidates {
|
for candidate in &mut all_candidates {
|
||||||
candidate.sanitize_sensitive_diagnostics();
|
candidate.sanitize_for_persistence();
|
||||||
}
|
}
|
||||||
if all_candidates.is_empty() {
|
if all_candidates.is_empty() {
|
||||||
return None;
|
return None;
|
||||||
@@ -510,8 +525,17 @@ pub struct DecisionTraceCandidate {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl DecisionTraceCandidate {
|
impl DecisionTraceCandidate {
|
||||||
|
pub fn sanitize_for_admin(&mut self) {
|
||||||
|
self.candidate.sanitize_for_persistence();
|
||||||
|
self.sanitize_catalog_metadata();
|
||||||
|
}
|
||||||
|
|
||||||
pub fn sanitize_sensitive_diagnostics(&mut self) {
|
pub fn sanitize_sensitive_diagnostics(&mut self) {
|
||||||
self.candidate.sanitize_sensitive_diagnostics();
|
self.candidate.sanitize_sensitive_diagnostics();
|
||||||
|
self.sanitize_catalog_metadata();
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sanitize_catalog_metadata(&mut self) {
|
||||||
self.provider_website = self
|
self.provider_website = self
|
||||||
.provider_website
|
.provider_website
|
||||||
.take()
|
.take()
|
||||||
@@ -574,7 +598,9 @@ pub fn build_decision_trace(
|
|||||||
})
|
})
|
||||||
.collect(),
|
.collect(),
|
||||||
};
|
};
|
||||||
trace.sanitize_sensitive_diagnostics();
|
for item in &mut trace.candidates {
|
||||||
|
item.sanitize_for_admin();
|
||||||
|
}
|
||||||
trace
|
trace
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -727,8 +753,12 @@ impl UpsertRequestCandidateRecord {
|
|||||||
self.api_key_name = None;
|
self.api_key_name = None;
|
||||||
self.skip_reason = sanitize_request_candidate_skip_reason(self.skip_reason.take());
|
self.skip_reason = sanitize_request_candidate_skip_reason(self.skip_reason.take());
|
||||||
self.error_type = sanitize_request_candidate_error_type(self.error_type.take());
|
self.error_type = sanitize_request_candidate_error_type(self.error_type.take());
|
||||||
self.error_message = None;
|
self.error_message = self
|
||||||
self.extra_data = sanitize_request_candidate_extra_data(self.extra_data.take());
|
.error_message
|
||||||
|
.take()
|
||||||
|
.map(limit_candidate_diagnostic_text);
|
||||||
|
self.extra_data =
|
||||||
|
sanitize_request_candidate_extra_data_for_persistence(self.extra_data.take());
|
||||||
self.required_capabilities =
|
self.required_capabilities =
|
||||||
sanitize_request_candidate_required_capabilities(self.required_capabilities.take());
|
sanitize_request_candidate_required_capabilities(self.required_capabilities.take());
|
||||||
}
|
}
|
||||||
@@ -788,6 +818,78 @@ pub fn sanitize_request_candidate_error_type(value: Option<String>) -> Option<St
|
|||||||
Some(safe)
|
Some(safe)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const MAX_CANDIDATE_DIAGNOSTIC_BYTES: usize = 65_536;
|
||||||
|
|
||||||
|
fn limit_candidate_diagnostic_text(mut text: String) -> String {
|
||||||
|
const SUFFIX: &str = "...[truncated]";
|
||||||
|
if text.len() > MAX_CANDIDATE_DIAGNOSTIC_BYTES {
|
||||||
|
let mut boundary = MAX_CANDIDATE_DIAGNOSTIC_BYTES - SUFFIX.len();
|
||||||
|
while !text.is_char_boundary(boundary) {
|
||||||
|
boundary -= 1;
|
||||||
|
}
|
||||||
|
text.truncate(boundary);
|
||||||
|
text.push_str(SUFFIX);
|
||||||
|
}
|
||||||
|
text
|
||||||
|
}
|
||||||
|
|
||||||
|
fn limit_candidate_diagnostic_value(value: &serde_json::Value) -> serde_json::Value {
|
||||||
|
if let Some(text) = value.as_str() {
|
||||||
|
return serde_json::Value::String(limit_candidate_diagnostic_text(text.to_string()));
|
||||||
|
}
|
||||||
|
let serialized = value.to_string();
|
||||||
|
if serialized.len() > MAX_CANDIDATE_DIAGNOSTIC_BYTES {
|
||||||
|
serde_json::Value::String(limit_candidate_diagnostic_text(serialized))
|
||||||
|
} else {
|
||||||
|
value.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn sanitize_request_candidate_extra_data_for_persistence(
|
||||||
|
extra_data: Option<serde_json::Value>,
|
||||||
|
) -> Option<serde_json::Value> {
|
||||||
|
let object = extra_data.as_ref()?.as_object()?;
|
||||||
|
let mut sanitized = sanitize_request_candidate_extra_data(extra_data.clone())
|
||||||
|
.and_then(|value| value.as_object().cloned())
|
||||||
|
.unwrap_or_default();
|
||||||
|
for (key, fields) in [
|
||||||
|
("upstream_response", &["headers", "body"][..]),
|
||||||
|
("error_flow", &["message"][..]),
|
||||||
|
(
|
||||||
|
"failure_diagnostic",
|
||||||
|
&["path", "field_path", "message", "type", "reason"][..],
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"request_conversion_error",
|
||||||
|
&["path", "field_path", "message", "type", "reason"][..],
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"request_body_build_error",
|
||||||
|
&["path", "field_path", "message", "type", "reason"][..],
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
let Some(diagnostic) = object.get(key).and_then(serde_json::Value::as_object) else {
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
let mut summary = sanitized
|
||||||
|
.remove(key)
|
||||||
|
.and_then(|value| value.as_object().cloned())
|
||||||
|
.unwrap_or_default();
|
||||||
|
for field in fields {
|
||||||
|
if let Some(value) = diagnostic.get(*field).filter(|value| !value.is_null()) {
|
||||||
|
summary.insert(
|
||||||
|
(*field).to_string(),
|
||||||
|
limit_candidate_diagnostic_value(value),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !summary.is_empty() {
|
||||||
|
sanitized.insert(key.to_string(), serde_json::Value::Object(summary));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
(!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized))
|
||||||
|
}
|
||||||
|
|
||||||
pub fn sanitize_request_candidate_extra_data(
|
pub fn sanitize_request_candidate_extra_data(
|
||||||
extra_data: Option<serde_json::Value>,
|
extra_data: Option<serde_json::Value>,
|
||||||
) -> Option<serde_json::Value> {
|
) -> Option<serde_json::Value> {
|
||||||
@@ -1949,7 +2051,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn candidate_persistence_removes_credentials_and_raw_payloads() {
|
fn candidate_persistence_keeps_admin_errors_but_removes_request_credentials() {
|
||||||
let mut record = UpsertRequestCandidateRecord {
|
let mut record = UpsertRequestCandidateRecord {
|
||||||
id: "candidate-1".to_string(),
|
id: "candidate-1".to_string(),
|
||||||
request_id: "request-1".to_string(),
|
request_id: "request-1".to_string(),
|
||||||
@@ -2079,7 +2181,7 @@ mod tests {
|
|||||||
);
|
);
|
||||||
assert!(record.username.is_none());
|
assert!(record.username.is_none());
|
||||||
assert!(record.api_key_name.is_none());
|
assert!(record.api_key_name.is_none());
|
||||||
assert!(record.error_message.is_none());
|
assert_eq!(record.error_message.as_deref(), Some("unauthorized"));
|
||||||
let extra = record
|
let extra = record
|
||||||
.extra_data
|
.extra_data
|
||||||
.as_ref()
|
.as_ref()
|
||||||
@@ -2112,7 +2214,10 @@ mod tests {
|
|||||||
assert_eq!(extra["error_flow"]["stage"], "upstream");
|
assert_eq!(extra["error_flow"]["stage"], "upstream");
|
||||||
assert_eq!(extra["error_flow"]["retryable"], true);
|
assert_eq!(extra["error_flow"]["retryable"], true);
|
||||||
assert_eq!(extra["error_flow"]["status_code"], 401);
|
assert_eq!(extra["error_flow"]["status_code"], 401);
|
||||||
assert!(extra["error_flow"].get("message").is_none());
|
assert_eq!(
|
||||||
|
extra["error_flow"]["message"],
|
||||||
|
"token vertex-secret rejected"
|
||||||
|
);
|
||||||
assert_eq!(extra["gateway_execution_runtime"], true);
|
assert_eq!(extra["gateway_execution_runtime"], true);
|
||||||
assert_eq!(extra["client_api_format"], "openai:responses");
|
assert_eq!(extra["client_api_format"], "openai:responses");
|
||||||
assert_eq!(extra["provider_api_format"], "claude:messages");
|
assert_eq!(extra["provider_api_format"], "claude:messages");
|
||||||
@@ -2127,8 +2232,14 @@ mod tests {
|
|||||||
assert_eq!(extra["upstream_response"]["source"], "upstream_response");
|
assert_eq!(extra["upstream_response"]["source"], "upstream_response");
|
||||||
assert_eq!(extra["upstream_response"]["status_code"], 401);
|
assert_eq!(extra["upstream_response"]["status_code"], 401);
|
||||||
assert_eq!(extra["upstream_response"]["body_state"], "inline");
|
assert_eq!(extra["upstream_response"]["body_state"], "inline");
|
||||||
assert!(extra["upstream_response"].get("headers").is_none());
|
assert_eq!(
|
||||||
assert!(extra["upstream_response"].get("body").is_none());
|
extra["upstream_response"]["headers"]["set-cookie"],
|
||||||
|
"session=secret"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
extra["upstream_response"]["body"]["error"]["message"],
|
||||||
|
"token vertex-secret rejected"
|
||||||
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
extra["image_progress"]["last_client_visible_event"],
|
extra["image_progress"]["last_client_visible_event"],
|
||||||
"image_generation.partial_image"
|
"image_generation.partial_image"
|
||||||
@@ -2178,7 +2289,6 @@ mod tests {
|
|||||||
|
|
||||||
let serialized = serde_json::to_string(&record).expect("candidate should serialize");
|
let serialized = serde_json::to_string(&record).expect("candidate should serialize");
|
||||||
for sensitive in [
|
for sensitive in [
|
||||||
"vertex-secret",
|
|
||||||
"client-secret",
|
"client-secret",
|
||||||
"credential-label-secret",
|
"credential-label-secret",
|
||||||
"header-rule-secret",
|
"header-rule-secret",
|
||||||
@@ -2199,6 +2309,11 @@ mod tests {
|
|||||||
"candidate must not retain {sensitive}"
|
"candidate must not retain {sensitive}"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
let public_extra = super::sanitize_request_candidate_extra_data(record.extra_data);
|
||||||
|
let serialized =
|
||||||
|
serde_json::to_string(&public_extra).expect("public data should serialize");
|
||||||
|
assert!(!serialized.contains("vertex-secret"));
|
||||||
|
assert!(!serialized.contains("session=secret"));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -2278,8 +2393,8 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn candidate_database_read_sanitizes_legacy_diagnostic_text() {
|
fn candidate_database_read_preserves_errors_until_public_projection() {
|
||||||
let candidate = StoredRequestCandidate::new(
|
let mut candidate = StoredRequestCandidate::new(
|
||||||
"candidate-1".to_string(),
|
"candidate-1".to_string(),
|
||||||
"request-1".to_string(),
|
"request-1".to_string(),
|
||||||
None,
|
None,
|
||||||
@@ -2315,9 +2430,42 @@ mod tests {
|
|||||||
candidate.error_type.as_deref(),
|
candidate.error_type.as_deref(),
|
||||||
Some(UNCLASSIFIED_CANDIDATE_ERROR_TYPE)
|
Some(UNCLASSIFIED_CANDIDATE_ERROR_TYPE)
|
||||||
);
|
);
|
||||||
assert!(candidate.error_message.is_none());
|
assert_eq!(
|
||||||
|
candidate.error_message.as_deref(),
|
||||||
|
Some("legacy secret message")
|
||||||
|
);
|
||||||
assert!(candidate.username.is_none());
|
assert!(candidate.username.is_none());
|
||||||
assert!(candidate.api_key_name.is_none());
|
assert!(candidate.api_key_name.is_none());
|
||||||
|
candidate.sanitize_sensitive_diagnostics();
|
||||||
|
assert!(candidate.error_message.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn admin_diagnostics_are_bounded_and_public_projection_removes_them() {
|
||||||
|
let raw = json!({
|
||||||
|
"upstream_response": {"status_code": 400, "body": "错误内容".repeat(20_000)},
|
||||||
|
"error_flow": {"status_code": 400, "message": "private upstream failure"},
|
||||||
|
"failure_diagnostic": {"path": "$.input", "message": "private conversion failure"},
|
||||||
|
"request_body": {"input": "private prompt"}
|
||||||
|
});
|
||||||
|
let admin = super::sanitize_request_candidate_extra_data_for_persistence(Some(raw))
|
||||||
|
.expect("admin diagnostics should remain");
|
||||||
|
let body = admin["upstream_response"]["body"]
|
||||||
|
.as_str()
|
||||||
|
.expect("body should be text");
|
||||||
|
assert!(body.len() <= 65_536);
|
||||||
|
assert!(body.ends_with("...[truncated]"));
|
||||||
|
assert!(admin.get("request_body").is_none());
|
||||||
|
assert_eq!(
|
||||||
|
super::sanitize_request_candidate_extra_data_for_persistence(Some(admin.clone())),
|
||||||
|
Some(admin.clone()),
|
||||||
|
);
|
||||||
|
let public = super::sanitize_request_candidate_extra_data(Some(admin))
|
||||||
|
.expect("public status should remain");
|
||||||
|
assert_eq!(public["upstream_response"]["status_code"], 400);
|
||||||
|
assert!(public["upstream_response"].get("body").is_none());
|
||||||
|
assert!(public["error_flow"].get("message").is_none());
|
||||||
|
assert!(public.get("failure_diagnostic").is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ pub use types::{
|
|||||||
ProxyNodeCounterDelta, StoredProviderApiKeyUsageSummary,
|
ProxyNodeCounterDelta, StoredProviderApiKeyUsageSummary,
|
||||||
StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredProviderUsageWindow,
|
StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredProviderUsageWindow,
|
||||||
StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary,
|
StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary,
|
||||||
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
|
StoredUsageBodyPayload, StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
|
||||||
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
|
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
|
||||||
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow,
|
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow,
|
||||||
StoredUsageDashboardProviderCount, StoredUsageDashboardStatsSummary,
|
StoredUsageDashboardProviderCount, StoredUsageDashboardStatsSummary,
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user