Compare commits

...
19 Commits
Author SHA1 Message Date
elky 7aa0c89244 fix(gateway): restore HTTP and WS upstream support 2026-09-07 22:15:05 +08:00
elky 7847ae98c6 fix(gateway): reset stream first-byte timeout per candidate 2026-09-07 21:56:16 +08:00
elky a90d564931 fix: restore security hardening compatibility and validation
Restore authorized rule reveal, explicit full HTTP capture and retention, video task business fields, and valid payment URLs. Add opt-in credential preservation for trusted recovery, fix frontend type contracts and async races, and eliminate PostgreSQL test fixture resource leaks. Document audit coverage and successful fmt and CI-scoped Clippy checks.
2026-09-07 21:14:27 +08:00
github-actions[bot] a5c3699ae9 chore(tunnel): update download links for tunnel-v0.3.17 2026-09-07 08:06:59 +00:00
elky 7b8048c6ae chore(tunnel): release v0.3.17 2026-09-07 15:57:39 +08:00
elky ec95f2ca1f fix(tunnel): prevent stream stalls and harden session cleanup
Reliably deliver flow-control credits and terminal states, isolate slow streams and heartbeats, negotiate stream windows, and clean up cancelled streams and session tasks.

Add regression coverage for queue pressure, early cancellation, small-window streaming, drain, and reconnect. Validate 185 agent tests, 88 gateway tunnel tests, and 21 protocol tests.
2026-09-07 15:39:40 +08:00
elky aa7dbe67d3 feat(providers): add persistent card view and shared drag ordering 2026-09-07 14:13:59 +08:00
elky a26680f460 fix(modules): restore legacy SMTP password migration 2026-09-07 12:07:53 +08:00
elky 522b979052 refactor(transport): remove provider DNS filtering and allowlist settings 2026-09-07 12:07:50 +08:00
elky 808946312a fix(providers): retry initial empty quota before showing feedback 2026-09-07 11:18:31 +08:00
elky 741107bf71 fix(transport): make provider DNS address filtering opt-in 2026-09-07 11:18:31 +08:00
elky 6962731220 fix(antigravity): restore default OAuth client compatibility 2026-09-07 10:51:24 +08:00
elky 062e111c03 fix(observability): preserve admin upstream error diagnostics 2026-09-07 10:39:48 +08:00
fawney19 470c59e197 Merge pull request #804 from AAEE86/fix-frontend-eslint
Fix frontend ESLint issues
2026-09-07 10:06:40 +08:00
elky 2f929e74c7 fix(frontend): preserve session cleanup and Unicode navigation 2026-09-07 09:57:53 +08:00
AAEE86 fc0417ceb9 fix(frontend): resolve ESLint issues 2026-09-07 09:03:51 +08:00
elky 44174a31e0 chore: update architecture documentation ignore rules 2026-09-07 08:57:44 +08:00
elky b599fb7354 fix(frontend): complete i18n coverage and responsive layouts 2026-09-07 08:54:19 +08:00
elky 14f96c9fa0 fix(providers): preserve health in redacted key summaries 2026-09-07 08:53:41 +08:00
315 changed files with 16868 additions and 4520 deletions
+7 -2
View File
@@ -111,8 +111,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=
+4
View File
@@ -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
View File
@@ -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",
@@ -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();
+10 -5
View File
@@ -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]
@@ -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(&region);
}
// 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(&region)
}
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 {
@@ -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());
} }
@@ -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,
),
)))
}
@@ -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!(
@@ -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!(
+64 -7
View File
@@ -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);
+2 -4
View File
@@ -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 =
-26
View File
@@ -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);
+17 -4
View File
@@ -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();
} }
+27 -9
View File
@@ -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;
@@ -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();
@@ -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(
@@ -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]
+240 -42
View File
@@ -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"],
"[redacted]"
);
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"],
"[redacted]"
);
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();
}
}
+203 -83
View File
@@ -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]
+2 -1
View File
@@ -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"] }
+16 -7
View File
@@ -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 系统直接使用。
+7
View File
@@ -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");
} }
+46 -13
View File
@@ -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();
+151 -131
View File
@@ -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]);
+19 -7
View File
@@ -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(
+156 -27
View File
@@ -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);
+237 -60
View File
@@ -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(
+52
View File
@@ -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();
}
}
+119 -5
View File
@@ -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)
+68 -1
View File
@@ -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,
+454 -69
View File
@@ -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
);
}
} }
-103
View File
@@ -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!(
+62
View File
@@ -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,
@@ -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
@@ -1374,8 +1374,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");
@@ -13362,19 +13362,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 +13400,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": "[redacted]"}))
);
assert_eq!(
stored.response_headers,
Some(json!({"content-type": "text/event-stream", "set-cookie": "[redacted]"}))
);
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": "[redacted]"}))
);
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]
@@ -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]
@@ -268,7 +268,7 @@ pub fn strip_deprecated_usage_display_fields(mut usage: UpsertUsageRecord) -> Up
usage usage
} }
pub fn sanitize_usage_for_persistence(mut usage: UpsertUsageRecord) -> UpsertUsageRecord { fn sanitize_usage_record_metadata(mut usage: UpsertUsageRecord) -> UpsertUsageRecord {
usage = strip_deprecated_usage_display_fields(usage); usage = strip_deprecated_usage_display_fields(usage);
sanitize_usage_routing_fields(&mut usage, None); sanitize_usage_routing_fields(&mut usage, None);
usage.error_message = None; usage.error_message = None;
@@ -280,6 +280,11 @@ pub fn sanitize_usage_for_persistence(mut usage: UpsertUsageRecord) -> UpsertUsa
.map(str::to_string); .map(str::to_string);
} }
usage.request_metadata = super::sanitize_usage_request_metadata(usage.request_metadata); usage.request_metadata = super::sanitize_usage_request_metadata(usage.request_metadata);
usage
}
pub fn sanitize_usage_for_persistence(usage: UpsertUsageRecord) -> UpsertUsageRecord {
let mut usage = sanitize_usage_record_metadata(usage);
usage.request_headers = None; usage.request_headers = None;
usage.request_body = None; usage.request_body = None;
usage.request_body_ref = None; usage.request_body_ref = None;
@@ -299,40 +304,96 @@ pub fn sanitize_usage_for_persistence(mut usage: UpsertUsageRecord) -> UpsertUsa
usage usage
} }
/// Project an event onto the non-content controls accepted by auxiliary usage storage.
///
/// Explicit `none` states are retained only as tombstones for removing historical captures.
/// Every header, body, reference, and non-clear capture state is discarded.
pub fn sanitize_usage_capture_controls_for_persistence( pub fn sanitize_usage_capture_controls_for_persistence(
mut usage: UpsertUsageRecord, mut usage: UpsertUsageRecord,
) -> UpsertUsageRecord { ) -> UpsertUsageRecord {
// Routing facts are allowed in the transient event metadata for compatibility with older
// writers. Project only the known scalar fields into typed slots before the general metadata
// sanitizer drops unknown keys. This keeps snapshots useful without re-persisting arbitrary
// metadata (or any body/header material).
let metadata = usage let metadata = usage
.request_metadata .request_metadata
.as_ref() .as_ref()
.and_then(Value::as_object) .and_then(Value::as_object)
.cloned(); .cloned();
sanitize_usage_routing_fields(&mut usage, metadata.as_ref()); sanitize_usage_routing_fields(&mut usage, metadata.as_ref());
let clear_request_body = usage.request_body_state == Some(super::UsageBodyCaptureState::None); let mut usage = sanitize_usage_record_metadata(usage);
let clear_provider_request_body = for headers in [
usage.provider_request_body_state == Some(super::UsageBodyCaptureState::None); &mut usage.request_headers,
let clear_response_body = usage.response_body_state == Some(super::UsageBodyCaptureState::None); &mut usage.provider_request_headers,
let clear_client_response_body = &mut usage.response_headers,
usage.client_response_body_state == Some(super::UsageBodyCaptureState::None); &mut usage.client_response_headers,
] {
let mut usage = sanitize_usage_for_persistence(usage); *headers = sanitize_usage_headers_for_persistence(headers.take());
usage.request_body_state = clear_request_body.then_some(super::UsageBodyCaptureState::None); }
usage.provider_request_body_state = for (field, body, body_ref, state) in [
clear_provider_request_body.then_some(super::UsageBodyCaptureState::None); (
usage.response_body_state = clear_response_body.then_some(super::UsageBodyCaptureState::None); super::UsageBodyField::RequestBody,
usage.client_response_body_state = &mut usage.request_body,
clear_client_response_body.then_some(super::UsageBodyCaptureState::None); &mut usage.request_body_ref,
usage.request_body_state,
),
(
super::UsageBodyField::ProviderRequestBody,
&mut usage.provider_request_body,
&mut usage.provider_request_body_ref,
usage.provider_request_body_state,
),
(
super::UsageBodyField::ResponseBody,
&mut usage.response_body,
&mut usage.response_body_ref,
usage.response_body_state,
),
(
super::UsageBodyField::ClientResponseBody,
&mut usage.client_response_body,
&mut usage.client_response_body_ref,
usage.client_response_body_state,
),
] {
if matches!(
state,
Some(
super::UsageBodyCaptureState::None
| super::UsageBodyCaptureState::Disabled
| super::UsageBodyCaptureState::Unavailable
)
) {
*body = None;
*body_ref = None;
} else {
*body_ref = body_ref.as_deref().and_then(|value| {
super::canonical_usage_body_ref_for(value, &usage.request_id, field)
});
}
}
usage usage
} }
pub fn usage_header_value_is_sensitive(name: &str) -> bool {
![
"accept",
"accept-encoding",
"content-encoding",
"content-length",
"content-type",
"transfer-encoding",
"x-request-id",
"x-trace-id",
]
.iter()
.any(|candidate| name.trim().eq_ignore_ascii_case(candidate))
}
pub fn sanitize_usage_headers_for_persistence(value: Option<Value>) -> Option<Value> {
let Value::Object(mut headers) = value? else {
return None;
};
for (name, value) in &mut headers {
if usage_header_value_is_sensitive(name) && !value.is_null() {
*value = Value::String("[redacted]".to_string());
}
}
Some(Value::Object(headers))
}
fn sanitize_usage_routing_fields( fn sanitize_usage_routing_fields(
usage: &mut UpsertUsageRecord, usage: &mut UpsertUsageRecord,
metadata: Option<&Map<String, Value>>, metadata: Option<&Map<String, Value>>,
@@ -773,20 +834,79 @@ mod tests {
} }
#[test] #[test]
fn auxiliary_capture_projection_keeps_only_explicit_clear_tombstones() { fn auxiliary_capture_projection_preserves_captures_and_honors_disabled_states() {
let mut input = usage_with_http_capture(); let mut input = usage_with_http_capture();
input.request_body_state = Some(UsageBodyCaptureState::None); input.request_body_state = Some(UsageBodyCaptureState::None);
input.response_body_state = Some(UsageBodyCaptureState::Disabled); input.response_body_state = Some(UsageBodyCaptureState::Disabled);
let usage = sanitize_usage_capture_controls_for_persistence(input); let usage = sanitize_usage_capture_controls_for_persistence(input);
assert!(usage.request_headers.is_none()); assert_eq!(
usage.request_headers,
Some(json!({"authorization": "[redacted]"}))
);
assert!(usage.request_body.is_none()); assert!(usage.request_body.is_none());
assert!(usage.request_body_ref.is_none()); assert!(usage.request_body_ref.is_none());
assert_eq!(usage.request_body_state, Some(UsageBodyCaptureState::None)); assert_eq!(usage.request_body_state, Some(UsageBodyCaptureState::None));
assert!(usage.provider_request_body_state.is_none()); assert_eq!(
assert!(usage.response_body_state.is_none()); usage.provider_request_body,
assert!(usage.client_response_body_state.is_none()); Some(json!({"prompt": "private"}))
);
assert_eq!(
usage.provider_request_body_state,
Some(UsageBodyCaptureState::Inline)
);
assert!(usage.response_body.is_none());
assert!(usage.response_body_ref.is_none());
assert_eq!(
usage.response_body_state,
Some(UsageBodyCaptureState::Disabled)
);
assert!(usage.client_response_body.is_none());
assert_eq!(
usage.client_response_body_state,
Some(UsageBodyCaptureState::Disabled)
);
}
#[test]
fn auxiliary_capture_projection_preserves_all_body_directions_and_scopes_references() {
let mut input = usage_with_http_capture();
input.request_headers =
Some(json!({"Content-Type": "application/json", "Authorization": "Bearer secret"}));
input.request_body_state = Some(UsageBodyCaptureState::Inline);
input.client_response_body_state = Some(UsageBodyCaptureState::Inline);
input.request_body_ref = Some(super::super::usage_body_ref(
&input.request_id,
super::super::UsageBodyField::RequestBody,
));
input.provider_request_body_ref = input.request_body_ref.clone();
input.response_body_ref = Some(super::super::usage_body_ref(
"another-request",
super::super::UsageBodyField::ResponseBody,
));
let captured = sanitize_usage_capture_controls_for_persistence(input.clone());
assert_eq!(captured.request_body, input.request_body);
assert_eq!(captured.provider_request_body, input.provider_request_body);
assert_eq!(captured.response_body, input.response_body);
assert_eq!(captured.client_response_body, input.client_response_body);
assert_eq!(captured.request_body_ref, input.request_body_ref);
assert!(captured.provider_request_body_ref.is_none());
assert!(captured.response_body_ref.is_none());
assert!(captured.client_response_body_ref.is_none());
assert_eq!(
captured.request_headers,
Some(json!({"Content-Type": "application/json", "Authorization": "[redacted]"}))
);
assert_eq!(
captured.provider_request_headers,
Some(json!({"x-api-key": "[redacted]"}))
);
assert_eq!(
captured.response_headers,
Some(json!({"set-cookie": "[redacted]"}))
);
assert!(captured.error_message.is_none());
} }
#[test] #[test]
@@ -1,8 +1,6 @@
use async_trait::async_trait; use async_trait::async_trait;
use serde_json::Value; use serde_json::Value;
const SAFE_VIDEO_URL_QUERY_KEYS: &[(&str, &str)] = &[("alt", "media")];
#[derive( #[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize, Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)] )]
@@ -221,16 +219,11 @@ impl StoredVideoTask {
} }
fn sanitize_persisted_diagnostics(&mut self) { fn sanitize_persisted_diagnostics(&mut self) {
self.prompt = None;
self.original_request_body = None; self.original_request_body = None;
self.progress_message = None; self.progress_message = None;
self.error_code = sanitize_video_task_error_code(self.error_code.take()); self.error_code = sanitize_video_task_error_code(self.error_code.take());
self.error_message = None; self.error_message = None;
self.video_url = sanitize_video_task_url( self.video_url = sanitize_video_task_url(self.video_url.take());
self.client_api_format.as_deref(),
self.provider_api_format.as_deref(),
self.video_url.take(),
);
self.request_metadata = None; self.request_metadata = None;
} }
@@ -339,18 +332,11 @@ pub struct UpsertVideoTask {
impl UpsertVideoTask { impl UpsertVideoTask {
pub fn sanitize_for_persistence(&mut self) { pub fn sanitize_for_persistence(&mut self) {
self.username = None;
self.api_key_name = None;
self.prompt = None;
self.original_request_body = None; self.original_request_body = None;
self.progress_message = None; self.progress_message = None;
self.error_code = sanitize_video_task_error_code(self.error_code.take()); self.error_code = sanitize_video_task_error_code(self.error_code.take());
self.error_message = None; self.error_message = None;
self.video_url = sanitize_video_task_url( self.video_url = sanitize_video_task_url(self.video_url.take());
self.client_api_format.as_deref(),
self.provider_api_format.as_deref(),
self.video_url.take(),
);
self.request_metadata = None; self.request_metadata = None;
} }
@@ -421,16 +407,7 @@ fn sanitize_video_task_error_code(value: Option<String>) -> Option<String> {
}) })
} }
fn sanitize_video_task_url( fn sanitize_video_task_url(value: Option<String>) -> Option<String> {
client_api_format: Option<&str>,
provider_api_format: Option<&str>,
value: Option<String>,
) -> Option<String> {
if effective_video_task_api_format(client_api_format, provider_api_format)
!= Some("gemini:video")
{
return None;
}
let mut url = url::Url::parse(value?.trim()).ok()?; let mut url = url::Url::parse(value?.trim()).ok()?;
if !matches!(url.scheme(), "http" | "https") if !matches!(url.scheme(), "http" | "https")
|| url.host_str().is_none() || url.host_str().is_none()
@@ -440,15 +417,6 @@ fn sanitize_video_task_url(
return None; return None;
} }
let query = url
.query_pairs()
.filter(|(key, value)| SAFE_VIDEO_URL_QUERY_KEYS.contains(&(key.as_ref(), value.as_ref())))
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect::<Vec<_>>();
url.set_query(None);
if !query.is_empty() {
url.query_pairs_mut().extend_pairs(query);
}
url.set_fragment(None); url.set_fragment(None);
Some(url.into()) Some(url.into())
} }
@@ -940,20 +908,23 @@ mod tests {
assert_eq!(task.user_id.as_deref(), Some("user-1")); assert_eq!(task.user_id.as_deref(), Some("user-1"));
assert_eq!(task.api_key_id.as_deref(), Some("api-key-1")); assert_eq!(task.api_key_id.as_deref(), Some("api-key-1"));
assert_eq!(task.username, None); assert_eq!(task.username.as_deref(), Some("private-user-name"));
assert_eq!(task.api_key_name, None); assert_eq!(task.api_key_name.as_deref(), Some("private-key-name"));
assert_eq!(task.original_request_body, None); assert_eq!(task.original_request_body, None);
assert_eq!(task.progress_message, None); assert_eq!(task.progress_message, None);
assert_eq!(task.error_message, None); assert_eq!(task.error_message, None);
assert_eq!(task.error_code.as_deref(), Some("provider_error")); assert_eq!(task.error_code.as_deref(), Some("provider_error"));
assert_eq!(task.video_url, None); assert_eq!(
task.video_url.as_deref(),
Some("https://cdn.example.test/video.mp4?token=secret")
);
assert_eq!(task.request_metadata, None); assert_eq!(task.request_metadata, None);
assert_eq!(task.prompt, None); assert_eq!(task.prompt.as_deref(), Some("prompt"));
assert_eq!(task.duration_seconds, Some(4)); assert_eq!(task.duration_seconds, Some(4));
} }
#[test] #[test]
fn upsert_sanitization_keeps_only_noncredential_video_urls() { fn stored_task_preserves_prompt_and_signed_download_url() {
let mut args = base_new_args(); let mut args = base_new_args();
args.12 = Some("gemini:video".to_string()); args.12 = Some("gemini:video".to_string());
args.15 = Some("private prompt".to_string()); args.15 = Some("private prompt".to_string());
@@ -968,15 +939,15 @@ mod tests {
args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36, args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36,
) )
.expect("stored task should build"); .expect("stored task should build");
assert_eq!(task.prompt, None); assert_eq!(task.prompt.as_deref(), Some("private prompt"));
assert_eq!( assert_eq!(
task.video_url.as_deref(), task.video_url.as_deref(),
Some("https://cdn.example.test/video.mp4?alt=media") Some("https://cdn.example.test/video.mp4?key=secret&alt=media&signature=private")
); );
} }
#[test] #[test]
fn upsert_sanitization_uses_client_format_when_legacy_provider_format_is_blank() { fn stored_task_preserves_download_url_when_legacy_provider_format_is_blank() {
let mut args = base_new_args(); let mut args = base_new_args();
args.11 = Some("gemini:video".to_string()); args.11 = Some("gemini:video".to_string());
args.12 = Some(" ".to_string()); args.12 = Some(" ".to_string());
@@ -992,7 +963,32 @@ mod tests {
assert_eq!( assert_eq!(
task.video_url.as_deref(), task.video_url.as_deref(),
Some("https://cdn.example.test/video.mp4?alt=media") Some("https://cdn.example.test/video.mp4?key=secret&alt=media")
); );
} }
#[test]
fn video_url_sanitization_preserves_signed_query_encoding_and_order() {
let video_url =
"https://cdn.example.test/video.mp4?signature=a%2Fb%2Bc%3D&part=2&part=1&name=a%20b";
assert_eq!(
super::sanitize_video_task_url(Some(format!("{video_url}#fragment"))).as_deref(),
Some(video_url)
);
}
#[test]
fn video_url_sanitization_rejects_invalid_schemes_and_embedded_credentials() {
for video_url in [
"file:///etc/passwd",
"javascript:alert(1)",
"data:video/mp4;base64,AAAA",
"https://user:[email protected]/video.mp4",
"https://[email protected]/video.mp4",
"/relative/video.mp4",
"not a url",
] {
assert!(super::sanitize_video_task_url(Some(video_url.to_string())).is_none());
}
}
} }
@@ -1,13 +1,8 @@
use std::{ use sqlx::{query, query_scalar, PgPool};
path::{Path, PathBuf},
process::{Child, Command, Stdio},
time::{Duration, Instant},
};
use sqlx::{query, query_scalar, Connection, PgConnection, PgPool};
use super::{pending_backfills, pending_backfills_from_applied, run_backfills, AppliedBackfill}; use super::{pending_backfills, pending_backfills_from_applied, run_backfills, AppliedBackfill};
use crate::lifecycle::migrate::prepare_database_for_startup; use crate::lifecycle::migrate::prepare_database_for_startup;
use crate::lifecycle::postgres_test_support::ManagedPostgresServer;
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION: i64 = 20260517012000; const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION: i64 = 20260517012000;
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_SQL: &str = const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_SQL: &str =
@@ -84,194 +79,6 @@ fn corrected_legacy_backfill_is_not_requeued_after_application() {
assert!(!pending_versions.contains(&LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION)); assert!(!pending_versions.contains(&LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION));
} }
#[derive(Debug)]
struct ManagedPostgresServer {
child: Option<Child>,
workdir: PathBuf,
database_url: String,
}
impl ManagedPostgresServer {
async fn try_start() -> Result<Option<Self>, Box<dyn std::error::Error>> {
let initdb_bin = std::env::var("AETHER_INITDB_BIN")
.ok()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "initdb".to_string());
let postgres_bin = std::env::var("AETHER_POSTGRES_BIN")
.ok()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "postgres".to_string());
if !command_exists(&initdb_bin) || !command_exists(&postgres_bin) {
eprintln!(
"skipping postgres backfill test because required binaries are unavailable: initdb={}, postgres={}",
initdb_bin, postgres_bin
);
return Ok(None);
}
match Self::start(initdb_bin, postgres_bin).await {
Ok(server) => Ok(Some(server)),
Err(err) if postgres_local_startup_unavailable(err.to_string().as_str()) => {
eprintln!(
"skipping postgres backfill test because local postgres could not start in this environment: {err}"
);
Ok(None)
}
Err(err) => Err(err),
}
}
async fn start(
initdb_bin: String,
postgres_bin: String,
) -> Result<Self, Box<dyn std::error::Error>> {
let port = reserve_local_port()?;
let workdir = std::env::temp_dir().join(format!(
"aether-backfill-tests-{}-{}",
std::process::id(),
port
));
let data_dir = workdir.join("data");
std::fs::create_dir_all(&workdir)?;
let init_output = Command::new(&initdb_bin)
.arg("-D")
.arg(&data_dir)
.arg("-U")
.arg("aether")
.arg("--auth=trust")
.arg("--encoding=UTF8")
.arg("--no-instructions")
.output()?;
if !init_output.status.success() {
return Err(std::io::Error::other(format!(
"initdb failed: {}",
String::from_utf8_lossy(&init_output.stderr)
))
.into());
}
let database_url = format!("postgres://[email protected]:{port}/postgres");
let log_path = workdir.join("postgres.log");
let stdout = std::fs::File::create(&log_path)?;
let stderr = stdout.try_clone()?;
let mut child = Command::new(&postgres_bin)
.arg("-D")
.arg(&data_dir)
.arg("-h")
.arg("127.0.0.1")
.arg("-p")
.arg(port.to_string())
.arg("-F")
.arg("-c")
.arg("fsync=off")
.arg("-c")
.arg("synchronous_commit=off")
.arg("-c")
.arg("full_page_writes=off")
.arg("-c")
.arg("shared_buffers=8MB")
.arg("-c")
.arg("max_connections=8")
.arg("-c")
.arg("dynamic_shared_memory_type=mmap")
.arg("-c")
.arg("autovacuum=off")
.stdout(Stdio::from(stdout))
.stderr(Stdio::from(stderr))
.spawn()?;
if let Err(err) = wait_for_postgres(&database_url).await {
let _ = child.kill();
let _ = child.wait();
return Err(err);
}
Ok(Self {
child: Some(child),
workdir,
database_url,
})
}
fn database_url(&self) -> &str {
&self.database_url
}
fn stop(&mut self) {
if let Some(mut child) = self.child.take() {
let _ = child.kill();
let _ = child.wait();
}
}
}
impl Drop for ManagedPostgresServer {
fn drop(&mut self) {
self.stop();
let _ = std::fs::remove_dir_all(&self.workdir);
}
}
fn command_exists(bin: &str) -> bool {
if bin.contains(std::path::MAIN_SEPARATOR) {
return Path::new(bin).exists();
}
let Some(paths) = std::env::var_os("PATH") else {
return false;
};
std::env::split_paths(&paths).any(|path| path.join(bin).exists())
}
fn reserve_local_port() -> Result<u16, std::io::Error> {
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
Ok(port)
}
fn postgres_shared_memory_unavailable(message: &str) -> bool {
let message = message.to_ascii_lowercase();
message.contains("shared memory")
&& (message.contains("could not create shared memory segment")
|| message.contains("shmget")
|| message.contains("no space left on device"))
}
fn postgres_local_startup_unavailable(message: &str) -> bool {
let message = message.to_ascii_lowercase();
postgres_shared_memory_unavailable(&message)
|| (message.contains("timed out waiting for local postgres")
&& (message.contains("connection refused")
|| message.contains("os error 61")
|| message.contains("os error 111")))
}
async fn wait_for_postgres(database_url: &str) -> Result<(), Box<dyn std::error::Error>> {
let deadline = Instant::now() + Duration::from_secs(10);
loop {
match PgConnection::connect(database_url).await {
Ok(connection) => {
connection.close().await?;
return Ok(());
}
Err(_) if Instant::now() < deadline => {
tokio::time::sleep(Duration::from_millis(50)).await
}
Err(err) => {
return Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("timed out waiting for local postgres: {err}"),
)
.into())
}
}
}
}
#[tokio::test] #[tokio::test]
async fn run_backfills_rebuilds_stats_and_records_execution() { async fn run_backfills_rebuilds_stats_and_records_execution() {
let Some(server) = ManagedPostgresServer::try_start() let Some(server) = ManagedPostgresServer::try_start()
@@ -538,9 +538,15 @@ fn imported_payload_ids(
.collect() .collect()
} }
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct DataImportOptions {
pub preserve_credentials: bool,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] #[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct DataCopyOptions { pub struct DataCopyOptions {
pub omit_request_body_details: bool, pub omit_request_body_details: bool,
pub preserve_credentials: bool,
} }
#[derive(Debug, Clone, PartialEq, Eq)] #[derive(Debug, Clone, PartialEq, Eq)]
@@ -572,6 +578,25 @@ fn set_supported_import_value(
} }
} }
fn apply_import_credential_policy(
table_name: &str,
object: &mut serde_json::Map<String, Value>,
target_has_column: impl Fn(&str) -> bool,
options: DataImportOptions,
) {
let normalized_table = table_name
.rsplit('.')
.next()
.unwrap_or(table_name)
.trim_matches(|character| matches!(character, '"' | '`'));
if options.preserve_credentials
&& matches!(normalized_table, "users" | "api_keys" | "management_tokens")
{
return;
}
deactivate_imported_credentials(table_name, object, target_has_column);
}
fn deactivate_imported_credentials( fn deactivate_imported_credentials(
table_name: &str, table_name: &str,
object: &mut serde_json::Map<String, Value>, object: &mut serde_json::Map<String, Value>,
@@ -1085,6 +1110,14 @@ pub async fn export_database_jsonl(
pub async fn import_database_jsonl( pub async fn import_database_jsonl(
database: SqlDatabaseConfig, database: SqlDatabaseConfig,
input: &str, input: &str,
) -> Result<usize, DataLayerError> {
import_database_jsonl_with_options(database, input, DataImportOptions::default()).await
}
pub async fn import_database_jsonl_with_options(
database: SqlDatabaseConfig,
input: &str,
options: DataImportOptions,
) -> Result<usize, DataLayerError> { ) -> Result<usize, DataLayerError> {
match database.driver { match database.driver {
#[cfg(feature = "postgres")] #[cfg(feature = "postgres")]
@@ -1092,7 +1125,7 @@ pub async fn import_database_jsonl(
let pool = let pool =
crate::driver::postgres::PostgresPoolFactory::new(database.to_postgres_config()?)? crate::driver::postgres::PostgresPoolFactory::new(database.to_postgres_config()?)?
.connect_lazy()?; .connect_lazy()?;
import_postgres_jsonl(&pool, input).await postgres::import_postgres_jsonl_with_options(&pool, input, options).await
} }
#[cfg(not(feature = "postgres"))] #[cfg(not(feature = "postgres"))]
DatabaseDriver::Postgres => Err(DataLayerError::InvalidInput( DatabaseDriver::Postgres => Err(DataLayerError::InvalidInput(
@@ -1113,7 +1146,14 @@ pub async fn copy_database_records(
if options.omit_request_body_details { if options.omit_request_body_details {
omit_request_body_details_from_records(&mut records); omit_request_body_details_from_records(&mut records);
} }
import_database_jsonl(target, &encode_jsonl(&records)?).await import_database_jsonl_with_options(
target,
&encode_jsonl(&records)?,
DataImportOptions {
preserve_credentials: options.preserve_credentials,
},
)
.await
} }
fn omit_request_body_details_from_records(records: &mut Vec<DataExportRecord>) { fn omit_request_body_details_from_records(records: &mut Vec<DataExportRecord>) {
@@ -58,14 +58,30 @@ pub async fn export_postgres_jsonl(
pub async fn import_postgres_jsonl( pub async fn import_postgres_jsonl(
pool: &crate::driver::postgres::PostgresPool, pool: &crate::driver::postgres::PostgresPool,
input: &str, input: &str,
) -> Result<usize, DataLayerError> {
import_postgres_jsonl_with_options(pool, input, DataImportOptions::default()).await
}
pub(super) async fn import_postgres_jsonl_with_options(
pool: &crate::driver::postgres::PostgresPool,
input: &str,
options: DataImportOptions,
) -> Result<usize, DataLayerError> { ) -> Result<usize, DataLayerError> {
let plan = build_import_plan(input)?; let plan = build_import_plan(input)?;
import_postgres_plan(pool, &plan).await import_postgres_plan_with_options(pool, &plan, options).await
} }
pub async fn import_postgres_plan( pub async fn import_postgres_plan(
pool: &crate::driver::postgres::PostgresPool, pool: &crate::driver::postgres::PostgresPool,
plan: &DataImportPlan, plan: &DataImportPlan,
) -> Result<usize, DataLayerError> {
import_postgres_plan_with_options(pool, plan, DataImportOptions::default()).await
}
async fn import_postgres_plan_with_options(
pool: &crate::driver::postgres::PostgresPool,
plan: &DataImportPlan,
options: DataImportOptions,
) -> Result<usize, DataLayerError> { ) -> Result<usize, DataLayerError> {
let identity_scope = IdentityImportScope::from_plan(plan)?; let identity_scope = IdentityImportScope::from_plan(plan)?;
let mut tx = pool.begin().await.map_sql_err()?; let mut tx = pool.begin().await.map_sql_err()?;
@@ -75,7 +91,7 @@ pub async fn import_postgres_plan(
for domain in &plan.manifest.domains { for domain in &plan.manifest.domains {
if *domain == ExportDomain::Auxiliary { if *domain == ExportDomain::Auxiliary {
for row in plan.rows(*domain) { for row in plan.rows(*domain) {
import_postgres_auxiliary_row(&mut tx, row, &mut column_cache).await?; import_postgres_auxiliary_row(&mut tx, row, &mut column_cache, options).await?;
imported = imported.saturating_add(1); imported = imported.saturating_add(1);
} }
continue; continue;
@@ -110,6 +126,7 @@ pub async fn import_postgres_plan(
*domain, *domain,
row, row,
&target_columns, &target_columns,
options,
) )
.await?; .await?;
imported = imported.saturating_add(1); imported = imported.saturating_add(1);
@@ -587,11 +604,15 @@ async fn import_postgres_row(
domain: ExportDomain, domain: ExportDomain,
row: &ExportRow, row: &ExportRow,
target_columns: &PostgresImportColumns, target_columns: &PostgresImportColumns,
options: DataImportOptions,
) -> Result<(), DataLayerError> { ) -> Result<(), DataLayerError> {
let mut object = normalize_postgres_import_payload(table_name, domain, row, target_columns)?; let mut object = normalize_postgres_import_payload(table_name, domain, row, target_columns)?;
deactivate_imported_credentials(table_name, &mut object, |column_name| { apply_import_credential_policy(
target_columns.contains_key(column_name) table_name,
}); &mut object,
|column_name| target_columns.contains_key(column_name),
options,
);
let columns = object.keys().map(String::as_str).collect::<Vec<_>>(); let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
let column_sql = columns let column_sql = columns
@@ -839,6 +860,7 @@ async fn import_postgres_billing_row(
payload, payload,
}, },
&target_columns, &target_columns,
DataImportOptions::default(),
) )
.await .await
} }
@@ -847,6 +869,7 @@ async fn import_postgres_auxiliary_row(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>, tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
row: &ExportRow, row: &ExportRow,
column_cache: &mut BTreeMap<String, PostgresImportColumns>, column_cache: &mut BTreeMap<String, PostgresImportColumns>,
options: DataImportOptions,
) -> Result<(), DataLayerError> { ) -> Result<(), DataLayerError> {
let (table_name, payload) = domain_payload_table(row, "auxiliary", None)?; let (table_name, payload) = domain_payload_table(row, "auxiliary", None)?;
let table = auxiliary_table(&table_name)?; let table = auxiliary_table(&table_name)?;
@@ -862,6 +885,7 @@ async fn import_postgres_auxiliary_row(
payload, payload,
}, },
&target_columns, &target_columns,
options,
) )
.await .await
} }
@@ -897,6 +921,7 @@ async fn import_postgres_wallet_row(
payload, payload,
}, },
&target_columns, &target_columns,
DataImportOptions::default(),
) )
.await .await
} }
@@ -3,11 +3,12 @@ use std::collections::{BTreeMap, BTreeSet};
use serde_json::{json, Value}; use serde_json::{json, Value};
use super::{ use super::{
build_import_plan, deactivate_imported_credentials, decode_jsonl, decode_jsonl_with_limits, apply_import_credential_policy, build_import_plan, deactivate_imported_credentials,
encode_jsonl, export_postgres_core_jsonl, normalize_imported_binary, decode_jsonl, decode_jsonl_with_limits, encode_jsonl, export_postgres_core_jsonl,
normalize_imported_integer_timestamp, normalize_postgres_import_payload, normalize_imported_binary, normalize_imported_integer_timestamp,
postgres_bytea_json_value, postgres_core_export_domains, DataExportManifest, DataExportRecord, normalize_postgres_import_payload, postgres_bytea_json_value, postgres_core_export_domains,
ExportDomain, ExportRow, PostgresImportColumn, DataExportManifest, DataExportRecord, DataImportOptions, ExportDomain, ExportRow,
PostgresImportColumn,
}; };
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory}; use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::lifecycle::migrate::run_migrations as run_postgres_migrations; use crate::lifecycle::migrate::run_migrations as run_postgres_migrations;
@@ -416,6 +417,159 @@ fn imported_proxy_nodes_receive_a_new_offline_tunnel_generation() {
); );
} }
#[test]
fn trusted_import_preserves_stable_credentials_only_when_explicitly_requested() {
assert!(!DataImportOptions::default().preserve_credentials);
for (table, payload) in [
("users", json!({"password_hash": "$2b$12$trusted-hash"})),
(
"public.\"api_keys\"",
json!({
"key_hash": "trusted-key-hash", "key_encrypted": "trusted-ciphertext",
"status": "active", "is_active": true, "is_locked": false,
}),
),
(
"management_tokens",
json!({"token_hash": "trusted-token-hash", "is_active": true}),
),
] {
for preserve_credentials in [false, true] {
let mut object = payload.as_object().unwrap().clone();
apply_import_credential_policy(
table,
&mut object,
|_| true,
DataImportOptions {
preserve_credentials,
},
);
if preserve_credentials {
assert_eq!(&object, payload.as_object().unwrap());
} else {
assert_ne!(&object, payload.as_object().unwrap());
}
}
}
}
#[test]
fn trusted_import_still_revokes_imported_sessions_and_live_tunnels() {
let options = DataImportOptions {
preserve_credentials: true,
};
let mut session = json!({
"refresh_token_hash": "old-session", "prev_refresh_token_hash": "older-session",
"revoked_at": null, "revoke_reason": null,
})
.as_object()
.unwrap()
.clone();
apply_import_credential_policy("public.user_sessions", &mut session, |_| true, options);
assert_ne!(session["refresh_token_hash"], json!("old-session"));
assert_eq!(session["prev_refresh_token_hash"], Value::Null);
assert_eq!(
session["revoke_reason"],
json!("imported_credentials_revoked")
);
let mut node = json!({
"tunnel_generation": "old-generation", "tunnel_connected": true,
"status": "online", "active_connections": 10,
})
.as_object()
.unwrap()
.clone();
apply_import_credential_policy("proxy_nodes", &mut node, |_| true, options);
assert_ne!(node["tunnel_generation"], json!("old-generation"));
assert_eq!(node["tunnel_connected"], json!(false));
assert_eq!(node["status"], json!("offline"));
assert_eq!(node["active_connections"], json!(0));
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_POSTGRES_URL and PostgreSQL migrations"]
async fn live_import_credential_policy_round_trips_through_postgres() {
let pool = PostgresPoolFactory::new(PostgresPoolConfig {
database_url: std::env::var("AETHER_TEST_POSTGRES_URL").unwrap(),
..Default::default()
})
.unwrap()
.connect_lazy()
.unwrap();
run_postgres_migrations(&pool).await.unwrap();
for preserve_credentials in [false, true] {
let user_id = uuid::Uuid::new_v4().to_string();
let key_id = uuid::Uuid::new_v4().to_string();
let password_hash = "$2b$12$trusted-import-hash";
let key_hash = format!("trusted-{key_id}");
sqlx::query("INSERT INTO users (id, username, password_hash, auth_source, email_verified) VALUES ($1, $2, $3, 'local', FALSE)")
.bind(&user_id).bind(format!("import-{}", &user_id[..8])).bind(password_hash)
.execute(&pool).await.unwrap();
sqlx::query("INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name) VALUES ($1, $2, $3, 'trusted-ciphertext', 'Import probe')")
.bind(&key_id).bind(&user_id).bind(&key_hash).execute(&pool).await.unwrap();
let user: Value = sqlx::query_scalar("SELECT to_jsonb(users) FROM users WHERE id = $1")
.bind(&user_id)
.fetch_one(&pool)
.await
.unwrap();
let key: Value =
sqlx::query_scalar("SELECT to_jsonb(api_keys) FROM api_keys WHERE id = $1")
.bind(&key_id)
.fetch_one(&pool)
.await
.unwrap();
let input = encode_jsonl(&[
DataExportRecord::manifest(DataExportManifest::new(
1_788_739_200,
Some(DatabaseDriver::Postgres),
vec![ExportDomain::Users, ExportDomain::ApiKeys],
)),
DataExportRecord::row(ExportDomain::Users, &user_id, user),
DataExportRecord::row(ExportDomain::ApiKeys, &key_id, key),
])
.unwrap();
super::postgres::import_postgres_jsonl_with_options(
&pool,
&input,
DataImportOptions {
preserve_credentials,
},
)
.await
.unwrap();
let imported_password: String =
sqlx::query_scalar("SELECT password_hash FROM users WHERE id = $1")
.bind(&user_id)
.fetch_one(&pool)
.await
.unwrap();
let imported_key: (String, Option<String>, bool) =
sqlx::query_as("SELECT key_hash, key_encrypted, is_active FROM api_keys WHERE id = $1")
.bind(&key_id)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(imported_password == password_hash, preserve_credentials);
assert_eq!(imported_key.0 == key_hash, preserve_credentials);
assert_eq!(
imported_key.1.as_deref(),
preserve_credentials.then_some("trusted-ciphertext")
);
assert_eq!(imported_key.2, preserve_credentials);
sqlx::query("DELETE FROM api_keys WHERE id = $1")
.bind(&key_id)
.execute(&pool)
.await
.unwrap();
sqlx::query("DELETE FROM users WHERE id = $1")
.bind(&user_id)
.execute(&pool)
.await
.unwrap();
}
}
fn postgres_column(data_type: &str, udt_name: &str) -> PostgresImportColumn { fn postgres_column(data_type: &str, udt_name: &str) -> PostgresImportColumn {
PostgresImportColumn { PostgresImportColumn {
data_type: data_type.to_ascii_lowercase(), data_type: data_type.to_ascii_lowercase(),
@@ -1,9 +1,7 @@
use std::borrow::Cow; use std::borrow::Cow;
use std::collections::BTreeSet; use std::collections::BTreeSet;
use std::fs; use std::fs;
use std::path::{Path, PathBuf}; use std::path::PathBuf;
use std::process::{Child, Command, Stdio};
use std::time::{Duration, Instant};
use sqlx::{ use sqlx::{
migrate::{AppliedMigration, Migrate}, migrate::{AppliedMigration, Migrate},
@@ -18,6 +16,8 @@ use aether_data_contracts::repository::{
}, },
}; };
use crate::lifecycle::postgres_test_support::ManagedPostgresServer;
use super::{ use super::{
postgres::{all_up_migrations, pending_migrations_from_applied, POSTGRES_MIGRATOR}, postgres::{all_up_migrations, pending_migrations_from_applied, POSTGRES_MIGRATOR},
prepare_database_for_startup, prepare_database_for_startup,
@@ -27,146 +27,6 @@ use crate::lifecycle::bootstrap::postgres::{
EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION, EMPTY_DATABASE_SNAPSHOT_SQL, EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION, EMPTY_DATABASE_SNAPSHOT_SQL,
}; };
#[derive(Debug)]
struct ManagedPostgresServer {
child: Option<Child>,
workdir: PathBuf,
database_url: String,
}
impl ManagedPostgresServer {
async fn try_start() -> Result<Option<Self>, Box<dyn std::error::Error>> {
let required = local_postgres_tests_required();
let initdb_bin = std::env::var("AETHER_INITDB_BIN")
.ok()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "initdb".to_string());
let postgres_bin = std::env::var("AETHER_POSTGRES_BIN")
.ok()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "postgres".to_string());
if !command_exists(&initdb_bin) || !command_exists(&postgres_bin) {
let message = format!(
"required postgres integration test binaries are unavailable: initdb={initdb_bin}, postgres={postgres_bin}"
);
if required {
return Err(std::io::Error::new(std::io::ErrorKind::NotFound, message).into());
}
eprintln!("skipping postgres integration test because {message}");
return Ok(None);
}
match Self::start(initdb_bin, postgres_bin).await {
Ok(server) => Ok(Some(server)),
Err(err)
if !required && postgres_local_startup_unavailable(err.to_string().as_str()) =>
{
eprintln!(
"skipping postgres integration test because local postgres could not start in this environment: {err}"
);
Ok(None)
}
Err(err) => Err(err),
}
}
async fn start(
initdb_bin: String,
postgres_bin: String,
) -> Result<Self, Box<dyn std::error::Error>> {
let port = reserve_local_port()?;
let workdir = std::env::temp_dir().join(format!(
"aether-migrate-tests-{}-{}",
std::process::id(),
port
));
let data_dir = workdir.join("data");
std::fs::create_dir_all(&workdir)?;
let init_output = Command::new(&initdb_bin)
.arg("-D")
.arg(&data_dir)
.arg("-U")
.arg("aether")
.arg("--auth=trust")
.arg("--encoding=UTF8")
.arg("--no-instructions")
.output()?;
if !init_output.status.success() {
return Err(std::io::Error::other(format!(
"initdb failed: {}",
String::from_utf8_lossy(&init_output.stderr)
))
.into());
}
let database_url = format!("postgres://[email protected]:{port}/postgres");
let log_path = workdir.join("postgres.log");
let stdout = std::fs::File::create(&log_path)?;
let stderr = stdout.try_clone()?;
let mut child = Command::new(&postgres_bin)
.arg("-D")
.arg(&data_dir)
.arg("-h")
.arg("127.0.0.1")
.arg("-p")
.arg(port.to_string())
.arg("-k")
.arg(&workdir)
.arg("-F")
.arg("-c")
.arg("fsync=off")
.arg("-c")
.arg("synchronous_commit=off")
.arg("-c")
.arg("full_page_writes=off")
.arg("-c")
.arg("shared_buffers=8MB")
.arg("-c")
.arg("max_connections=8")
.arg("-c")
.arg("dynamic_shared_memory_type=mmap")
.arg("-c")
.arg("autovacuum=off")
.stdout(Stdio::from(stdout))
.stderr(Stdio::from(stderr))
.spawn()?;
if let Err(err) = wait_for_postgres(&database_url).await {
let _ = child.kill();
let exit_status = child
.wait()
.map(|status| status.to_string())
.unwrap_or_else(|wait_err| format!("unavailable ({wait_err})"));
let logs = fs::read_to_string(&log_path)
.unwrap_or_else(|read_err| format!("<failed to read postgres log: {read_err}>"));
return Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("{err}; postgres exit status: {exit_status}; logs:\n{logs}"),
)
.into());
}
Ok(Self {
child: Some(child),
workdir,
database_url,
})
}
fn database_url(&self) -> &str {
&self.database_url
}
fn stop(&mut self) {
if let Some(mut child) = self.child.take() {
let _ = child.kill();
let _ = child.wait();
}
}
}
/// A clean PostgreSQL database is bootstrapped from the schema snapshot first; /// A clean PostgreSQL database is bootstrapped from the schema snapshot first;
/// migrations after the privacy/security frontier are intentionally left /// migrations after the privacy/security frontier are intentionally left
/// pending so their data-preserving changes still execute. Exercise the same /// pending so their data-preserving changes still execute. Exercise the same
@@ -191,83 +51,6 @@ async fn prepare_and_apply_clean_postgres_database(pool: &PgPool) {
); );
} }
fn local_postgres_tests_required() -> bool {
// CI can opt into failing when the isolated local PostgreSQL fixture is unavailable.
std::env::var("AETHER_REQUIRE_LOCAL_POSTGRES_TESTS")
.ok()
.is_some_and(|value| {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
)
})
}
impl Drop for ManagedPostgresServer {
fn drop(&mut self) {
self.stop();
let _ = std::fs::remove_dir_all(&self.workdir);
}
}
fn command_exists(bin: &str) -> bool {
if bin.contains(std::path::MAIN_SEPARATOR) {
return Path::new(bin).exists();
}
let Some(paths) = std::env::var_os("PATH") else {
return false;
};
std::env::split_paths(&paths).any(|path| path.join(bin).exists())
}
fn reserve_local_port() -> Result<u16, std::io::Error> {
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
Ok(port)
}
fn postgres_shared_memory_unavailable(message: &str) -> bool {
let message = message.to_ascii_lowercase();
message.contains("shared memory")
&& (message.contains("could not create shared memory segment")
|| message.contains("shmget")
|| message.contains("no space left on device"))
}
fn postgres_local_startup_unavailable(message: &str) -> bool {
let message = message.to_ascii_lowercase();
postgres_shared_memory_unavailable(&message)
|| (message.contains("timed out waiting for local postgres")
&& (message.contains("connection refused")
|| message.contains("os error 61")
|| message.contains("os error 111")))
}
async fn wait_for_postgres(database_url: &str) -> Result<(), Box<dyn std::error::Error>> {
let deadline = Instant::now() + Duration::from_secs(10);
loop {
match PgConnection::connect(database_url).await {
Ok(connection) => {
connection.close().await?;
return Ok(());
}
Err(_) if Instant::now() < deadline => {
tokio::time::sleep(Duration::from_millis(50)).await
}
Err(err) => {
return Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("timed out waiting for local postgres: {err}"),
)
.into())
}
}
}
}
async fn table_exists(pool: &PgPool, table_name: &str) -> Result<bool, sqlx::Error> { async fn table_exists(pool: &PgPool, table_name: &str) -> Result<bool, sqlx::Error> {
query_scalar::<_, bool>("SELECT to_regclass($1) IS NOT NULL") query_scalar::<_, bool>("SELECT to_regclass($1) IS NOT NULL")
.bind(format!("public.{table_name}")) .bind(format!("public.{table_name}"))
@@ -8,3 +8,5 @@ pub mod backfill;
pub(crate) mod bootstrap; pub(crate) mod bootstrap;
pub mod export; pub mod export;
pub mod migrate; pub mod migrate;
#[cfg(all(test, feature = "postgres"))]
mod postgres_test_support;
@@ -0,0 +1,276 @@
use std::{
path::{Path, PathBuf},
process::{Child, Command, Stdio},
time::{Duration, Instant},
};
use sqlx::{Connection, PgConnection};
#[derive(Debug)]
pub(super) struct ManagedPostgresServer {
child: Option<Child>,
pg_ctl_bin: PathBuf,
workdir: PathBuf,
data_dir: PathBuf,
database_url: String,
}
impl ManagedPostgresServer {
pub(super) async fn try_start() -> Result<Option<Self>, Box<dyn std::error::Error>> {
let required = local_postgres_tests_required();
let initdb_bin = configured_binary("AETHER_INITDB_BIN", "initdb");
let postgres_bin = configured_binary("AETHER_POSTGRES_BIN", "postgres");
let pg_ctl_bin = std::env::var("AETHER_PG_CTL_BIN")
.ok()
.filter(|value| !value.trim().is_empty())
.map(PathBuf::from)
.unwrap_or_else(|| {
PathBuf::from(&postgres_bin).with_file_name(if cfg!(windows) {
"pg_ctl.exe"
} else {
"pg_ctl"
})
});
if !command_exists(Path::new(&initdb_bin))
|| !command_exists(Path::new(&postgres_bin))
|| !command_exists(&pg_ctl_bin)
{
let message = format!(
"required postgres integration test binaries are unavailable: initdb={initdb_bin}, postgres={postgres_bin}, pg_ctl={}",
pg_ctl_bin.display()
);
if required {
return Err(std::io::Error::new(std::io::ErrorKind::NotFound, message).into());
}
eprintln!("skipping postgres integration test because {message}");
return Ok(None);
}
match Self::start(initdb_bin, postgres_bin, pg_ctl_bin).await {
Ok(server) => Ok(Some(server)),
Err(error) if !required && postgres_local_startup_unavailable(&error.to_string()) => {
eprintln!(
"skipping postgres integration test because local postgres could not start in this environment: {error}"
);
Ok(None)
}
Err(error) => Err(error),
}
}
async fn start(
initdb_bin: String,
postgres_bin: String,
pg_ctl_bin: PathBuf,
) -> Result<Self, Box<dyn std::error::Error>> {
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
let workdir = std::env::temp_dir().join(format!(
"aether-lifecycle-tests-{}-{port}",
std::process::id()
));
std::fs::create_dir(&workdir)?;
let mut server = Self {
child: None,
pg_ctl_bin,
data_dir: workdir.join("data"),
workdir,
database_url: format!("postgres://[email protected]:{port}/postgres"),
};
let init_output = Command::new(&initdb_bin)
.arg("-D")
.arg(&server.data_dir)
.args([
"-U",
"aether",
"--auth=trust",
"--encoding=UTF8",
"--no-instructions",
])
.output()?;
if !init_output.status.success() {
return Err(std::io::Error::other(format!(
"initdb failed: {}",
String::from_utf8_lossy(&init_output.stderr)
))
.into());
}
let log_path = server.workdir.join("postgres.log");
let stdout = std::fs::File::create(&log_path)?;
let stderr = stdout.try_clone()?;
server.child = Some(
Command::new(&postgres_bin)
.arg("-D")
.arg(&server.data_dir)
.args(["-h", "127.0.0.1", "-p"])
.arg(port.to_string())
.arg("-F")
.args(["-c", "unix_socket_directories="])
.args(["-c", "fsync=off"])
.args(["-c", "synchronous_commit=off"])
.args(["-c", "full_page_writes=off"])
.args(["-c", "shared_buffers=8MB"])
.args(["-c", "max_connections=8"])
.args(["-c", "dynamic_shared_memory_type=mmap"])
.args(["-c", "autovacuum=off"])
.stdout(Stdio::from(stdout))
.stderr(Stdio::from(stderr))
.spawn()?,
);
if let Err(error) = wait_for_postgres(&server.database_url).await {
let logs = std::fs::read_to_string(&log_path).unwrap_or_default();
server.stop()?;
return Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("{error}; postgres logs:\n{logs}"),
)
.into());
}
Ok(server)
}
pub(super) fn database_url(&self) -> &str {
&self.database_url
}
fn stop(&mut self) -> Result<(), std::io::Error> {
let Some(child) = self.child.as_mut() else {
return Ok(());
};
if child.try_wait()?.is_some() {
self.child = None;
return Ok(());
}
let output = Command::new(&self.pg_ctl_bin)
.arg("-D")
.arg(&self.data_dir)
.args(["stop", "-m", "fast", "-w", "-t", "10"])
.output()?;
if !output.status.success() && child.try_wait()?.is_none() {
return Err(std::io::Error::other(format!(
"pg_ctl stop failed for {}: {}{}",
self.data_dir.display(),
String::from_utf8_lossy(&output.stdout),
String::from_utf8_lossy(&output.stderr),
)));
}
child.wait()?;
self.child = None;
Ok(())
}
}
impl Drop for ManagedPostgresServer {
fn drop(&mut self) {
match self.stop() {
Ok(()) => {
let _ = std::fs::remove_dir_all(&self.workdir);
}
Err(error) => {
eprintln!(
"failed to stop managed postgres; preserving {}: {error}",
self.workdir.display(),
);
}
}
}
}
fn configured_binary(variable: &str, default: &str) -> String {
std::env::var(variable)
.ok()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| default.to_string())
}
fn command_exists(binary: &Path) -> bool {
if binary.is_absolute() || binary.components().count() > 1 {
return binary.is_file();
}
std::env::var_os("PATH")
.is_some_and(|paths| std::env::split_paths(&paths).any(|path| path.join(binary).is_file()))
}
fn local_postgres_tests_required() -> bool {
std::env::var("AETHER_REQUIRE_LOCAL_POSTGRES_TESTS")
.ok()
.is_some_and(|value| {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
)
})
}
fn postgres_local_startup_unavailable(message: &str) -> bool {
let message = message.to_ascii_lowercase();
(message.contains("shared memory")
&& (message.contains("could not create shared memory segment")
|| message.contains("shmget")
|| message.contains("no space left on device")))
|| (message.contains("timed out waiting for local postgres")
&& (message.contains("connection refused")
|| message.contains("os error 61")
|| message.contains("os error 111")))
}
async fn wait_for_postgres(database_url: &str) -> Result<(), Box<dyn std::error::Error>> {
let deadline = Instant::now() + Duration::from_secs(10);
loop {
match PgConnection::connect(database_url).await {
Ok(connection) => {
connection.close().await?;
return Ok(());
}
Err(_) if Instant::now() < deadline => {
tokio::time::sleep(Duration::from_millis(50)).await;
}
Err(error) => {
return Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("timed out waiting for local postgres: {error}"),
)
.into());
}
}
}
}
#[tokio::test]
async fn managed_postgres_stops_cleanly_with_open_connections() {
let Some(mut server) = ManagedPostgresServer::try_start().await.unwrap() else {
return;
};
let connection = PgConnection::connect(server.database_url()).await.unwrap();
let workdir = server.workdir.clone();
assert!(server.data_dir.join("postmaster.pid").exists());
server.stop().unwrap();
assert!(server.child.is_none());
assert!(!server.data_dir.join("postmaster.pid").exists());
server.stop().unwrap();
drop(connection);
drop(server);
assert!(!workdir.exists());
}
#[tokio::test]
async fn failed_postgres_stop_retains_ownership_for_retry() {
let Some(mut server) = ManagedPostgresServer::try_start().await.unwrap() else {
return;
};
let pg_ctl_bin = server.pg_ctl_bin.clone();
let workdir = server.workdir.clone();
server.pg_ctl_bin = workdir.join("missing-pg-ctl");
assert!(server.stop().is_err());
assert!(server.child.as_mut().unwrap().try_wait().unwrap().is_none());
assert!(server.data_dir.exists());
server.pg_ctl_bin = pg_ctl_bin;
server.stop().unwrap();
drop(server);
assert!(!workdir.exists());
}
@@ -10,19 +10,31 @@ use crate::DataLayerError;
use async_trait::async_trait; use async_trait::async_trait;
fn sanitize_stored_candidate(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate { fn sanitize_stored_candidate(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
candidate.sanitize_sensitive_diagnostics(); candidate.sanitize_for_persistence();
candidate candidate
} }
fn merge_extra_data( fn merge_extra_data(
existing: Option<serde_json::Value>, existing: Option<serde_json::Value>,
overlay: Option<serde_json::Value>, overlay: Option<serde_json::Value>,
preserve_error_details: bool,
) -> Option<serde_json::Value> { ) -> Option<serde_json::Value> {
match (existing, overlay) { match (existing, overlay) {
( (
Some(serde_json::Value::Object(mut existing_object)), Some(serde_json::Value::Object(mut existing_object)),
Some(serde_json::Value::Object(overlay_object)), Some(serde_json::Value::Object(mut overlay_object)),
) => { ) => {
if preserve_error_details {
for key in [
"upstream_response",
"error_flow",
"failure_diagnostic",
"request_conversion_error",
"request_body_build_error",
] {
overlay_object.remove(key);
}
}
existing_object.extend(overlay_object); existing_object.extend(overlay_object);
Some(serde_json::Value::Object(existing_object)) Some(serde_json::Value::Object(existing_object))
} }
@@ -397,7 +409,13 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
.error_type .error_type
.or_else(|| existing.as_ref().and_then(|row| row.error_type.clone())) .or_else(|| existing.as_ref().and_then(|row| row.error_type.clone()))
}, },
error_message: None, error_message: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.error_message.clone())
} else {
candidate
.error_message
.or_else(|| existing.as_ref().and_then(|row| row.error_message.clone()))
},
latency_ms: if preserve_existing_lifecycle { latency_ms: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.latency_ms) existing.as_ref().and_then(|row| row.latency_ms)
} else { } else {
@@ -411,6 +429,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
extra_data: merge_extra_data( extra_data: merge_extra_data(
existing.as_ref().and_then(|row| row.extra_data.clone()), existing.as_ref().and_then(|row| row.extra_data.clone()),
candidate.extra_data, candidate.extra_data,
preserve_existing_lifecycle,
), ),
required_capabilities: candidate.required_capabilities.or_else(|| { required_capabilities: candidate.required_capabilities.or_else(|| {
existing existing
@@ -589,7 +608,10 @@ mod tests {
let candidate = stored let candidate = stored
.get("cand-raw") .get("cand-raw")
.expect("seeded candidate should exist"); .expect("seeded candidate should exist");
assert!(candidate.error_message.is_none()); assert_eq!(
candidate.error_message.as_deref(),
Some("Bearer secret-token")
);
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_eq!( assert_eq!(
@@ -619,7 +641,10 @@ mod tests {
.iter() .iter()
.find(|candidate| candidate.id == "cand-bypassed") .find(|candidate| candidate.id == "cand-bypassed")
.expect("bypassed candidate should be returned"); .expect("bypassed candidate should be returned");
assert!(candidate.error_message.is_none()); assert_eq!(
candidate.error_message.as_deref(),
Some("Bearer secret-token")
);
assert_eq!( assert_eq!(
candidate.extra_data, candidate.extra_data,
Some(json!({"gateway_execution_runtime": true})) Some(json!({"gateway_execution_runtime": true}))
@@ -665,7 +690,7 @@ mod tests {
.await .await
.expect("candidate merge should succeed"); .expect("candidate merge should succeed");
assert_eq!(merged.id, "cand-bypassed"); assert_eq!(merged.id, "cand-bypassed");
assert!(merged.error_message.is_none()); assert_eq!(merged.error_message.as_deref(), Some("Bearer secret-token"));
assert_eq!( assert_eq!(
merged.extra_data, merged.extra_data,
Some(json!({ Some(json!({
@@ -839,7 +864,13 @@ mod tests {
Some("retryable upstream failure".to_string()), Some("retryable upstream failure".to_string()),
Some(45), Some(45),
Some(1), Some(1),
Some(json!({"stream_completed": true})), Some(json!({
"stream_completed": true,
"upstream_response": {
"status_code": 503,
"body": {"error": {"message": "original upstream failure"}}
}
})),
None, None,
100, 100,
Some(101), Some(101),
@@ -869,7 +900,10 @@ mod tests {
error_message: None, error_message: None,
latency_ms: Some(9_999), latency_ms: Some(9_999),
concurrent_requests: Some(2), concurrent_requests: Some(2),
extra_data: Some(json!({"gateway_execution_runtime": true})), extra_data: Some(json!({
"gateway_execution_runtime": true,
"upstream_response": {"status_code": 200, "body": "late unrelated response"}
})),
required_capabilities: None, required_capabilities: None,
created_at_unix_ms: None, created_at_unix_ms: None,
started_at_unix_ms: Some(102), started_at_unix_ms: Some(102),
@@ -882,7 +916,10 @@ mod tests {
assert_eq!(updated.status, RequestCandidateStatus::Failed); assert_eq!(updated.status, RequestCandidateStatus::Failed);
assert_eq!(updated.status_code, Some(503)); assert_eq!(updated.status_code, Some(503));
assert_eq!(updated.error_type.as_deref(), Some("upstream_error")); assert_eq!(updated.error_type.as_deref(), Some("upstream_error"));
assert!(updated.error_message.is_none()); assert_eq!(
updated.error_message.as_deref(),
Some("retryable upstream failure")
);
assert_eq!(updated.latency_ms, Some(45)); assert_eq!(updated.latency_ms, Some(45));
assert_eq!(updated.concurrent_requests, Some(2)); assert_eq!(updated.concurrent_requests, Some(2));
assert_eq!(updated.finished_at_unix_ms, Some(145)); assert_eq!(updated.finished_at_unix_ms, Some(145));
@@ -890,7 +927,11 @@ mod tests {
updated.extra_data, updated.extra_data,
Some(json!({ Some(json!({
"gateway_execution_runtime": true, "gateway_execution_runtime": true,
"stream_completed": true "stream_completed": true,
"upstream_response": {
"status_code": 503,
"body": {"error": {"message": "original upstream failure"}}
}
})) }))
); );
} }
@@ -2752,6 +2752,40 @@ fn hydrate_client_family(item: &mut StoredRequestUsageAudit) {
} }
} }
fn merge_usage_body_capture(
incoming_body: Option<Value>,
incoming_ref: Option<String>,
incoming_state: Option<UsageBodyCaptureState>,
existing: Option<&StoredRequestUsageAudit>,
field: UsageBodyField,
) -> (Option<Value>, Option<String>, Option<UsageBodyCaptureState>) {
if matches!(
incoming_state,
Some(
UsageBodyCaptureState::None
| UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
)
) {
return (None, None, incoming_state);
}
if incoming_body.is_some() {
return (
incoming_body,
None,
incoming_state.or(Some(UsageBodyCaptureState::Inline)),
);
}
if incoming_ref.is_some() {
return (None, incoming_ref, Some(UsageBodyCaptureState::Reference));
}
(
existing.and_then(|item| item.body_value(field).cloned()),
existing.and_then(|item| item.body_ref(field).map(ToOwned::to_owned)),
incoming_state.or_else(|| existing.and_then(|item| item.body_state(field))),
)
}
fn request_body_capture_replaces_derived_facts( fn request_body_capture_replaces_derived_facts(
request_body: Option<&Value>, request_body: Option<&Value>,
request_body_state: Option<UsageBodyCaptureState>, request_body_state: Option<UsageBodyCaptureState>,
@@ -2883,36 +2917,47 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
return Ok(existing.clone()); return Ok(existing.clone());
} }
} }
let capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage); let mut capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage);
if let Some(existing) = by_request_id.get_mut(&usage.request_id) { if let Some(existing) = by_request_id.get_mut(&usage.request_id) {
existing.request_headers = None;
existing.request_body = None;
existing.request_body_ref = None;
existing.request_body_state = None;
existing.provider_request_headers = None;
existing.provider_request_body = None;
existing.provider_request_body_ref = None;
existing.provider_request_body_state = None;
existing.response_headers = None;
existing.response_body = None;
existing.response_body_ref = None;
existing.response_body_state = None;
existing.client_response_headers = None;
existing.client_response_body = None;
existing.client_response_body_ref = None;
existing.client_response_body_state = None;
existing.request_metadata = existing.request_metadata =
sanitize_usage_request_metadata(existing.request_metadata.take()); sanitize_usage_request_metadata(existing.request_metadata.take());
} }
{ {
let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock"); let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock");
for field in [ for (field, state, body) in [
UsageBodyField::RequestBody, (
UsageBodyField::ProviderRequestBody, UsageBodyField::RequestBody,
UsageBodyField::ResponseBody, capture_usage.request_body_state,
UsageBodyField::ClientResponseBody, &capture_usage.request_body,
),
(
UsageBodyField::ProviderRequestBody,
capture_usage.provider_request_body_state,
&capture_usage.provider_request_body,
),
(
UsageBodyField::ResponseBody,
capture_usage.response_body_state,
&capture_usage.response_body,
),
(
UsageBodyField::ClientResponseBody,
capture_usage.client_response_body_state,
&capture_usage.client_response_body,
),
] { ] {
detached_bodies.remove(&usage_body_ref(&usage.request_id, field)); if body.is_some()
|| matches!(
state,
Some(
UsageBodyCaptureState::None
| UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
)
)
{
detached_bodies.remove(&usage_body_ref(&usage.request_id, field));
}
} }
} }
@@ -2977,10 +3022,36 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
} }
}); });
let request_metadata = sanitize_memory_request_metadata(request_metadata); let request_metadata = sanitize_memory_request_metadata(request_metadata);
let request_body_ref = None; let (request_body, request_body_ref, request_body_state) = merge_usage_body_capture(
let provider_request_body_ref = None; capture_usage.request_body.take(),
let response_body_ref = None; capture_usage.request_body_ref.take(),
let client_response_body_ref = None; capture_usage.request_body_state,
existing.as_ref(),
UsageBodyField::RequestBody,
);
let (provider_request_body, provider_request_body_ref, provider_request_body_state) =
merge_usage_body_capture(
capture_usage.provider_request_body.take(),
capture_usage.provider_request_body_ref.take(),
capture_usage.provider_request_body_state,
existing.as_ref(),
UsageBodyField::ProviderRequestBody,
);
let (response_body, response_body_ref, response_body_state) = merge_usage_body_capture(
capture_usage.response_body.take(),
capture_usage.response_body_ref.take(),
capture_usage.response_body_state,
existing.as_ref(),
UsageBodyField::ResponseBody,
);
let (client_response_body, client_response_body_ref, client_response_body_state) =
merge_usage_body_capture(
capture_usage.client_response_body.take(),
capture_usage.client_response_body_ref.take(),
capture_usage.client_response_body_state,
existing.as_ref(),
UsageBodyField::ClientResponseBody,
);
let stored = StoredRequestUsageAudit { let stored = StoredRequestUsageAudit {
id: existing id: existing
.as_ref() .as_ref()
@@ -3089,22 +3160,38 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
), ),
status: usage.status, status: usage.status,
billing_status: usage.billing_status, billing_status: usage.billing_status,
request_headers: None, request_headers: capture_usage.request_headers.or_else(|| {
request_body: None, existing
.as_ref()
.and_then(|item| item.request_headers.clone())
}),
request_body,
request_body_ref, request_body_ref,
request_body_state: capture_usage.request_body_state, request_body_state,
provider_request_headers: None, provider_request_headers: capture_usage.provider_request_headers.or_else(|| {
provider_request_body: None, existing
.as_ref()
.and_then(|item| item.provider_request_headers.clone())
}),
provider_request_body,
provider_request_body_ref, provider_request_body_ref,
provider_request_body_state: capture_usage.provider_request_body_state, provider_request_body_state,
response_headers: None, response_headers: capture_usage.response_headers.or_else(|| {
response_body: None, existing
.as_ref()
.and_then(|item| item.response_headers.clone())
}),
response_body,
response_body_ref, response_body_ref,
response_body_state: capture_usage.response_body_state, response_body_state,
client_response_headers: None, client_response_headers: capture_usage.client_response_headers.or_else(|| {
client_response_body: None, existing
.as_ref()
.and_then(|item| item.client_response_headers.clone())
}),
client_response_body,
client_response_body_ref, client_response_body_ref,
client_response_body_state: capture_usage.client_response_body_state, client_response_body_state,
candidate_id: if replace_routing_snapshot { candidate_id: if replace_routing_snapshot {
capture_usage.candidate_id capture_usage.candidate_id
} else { } else {
@@ -136,6 +136,76 @@ fn sample_upsert_usage_record(request_id: &str) -> UpsertUsageRecord {
} }
} }
#[tokio::test]
async fn upsert_preserves_full_http_captures_across_lifecycle_updates() {
let repository = InMemoryUsageReadRepository::default();
let mut pending = sample_upsert_usage_record("req-full-capture");
pending.request_headers =
Some(json!({"content-type": "application/json", "authorization": "Bearer private"}));
pending.request_body =
Some(json!({"messages": [{"role": "user", "content": "original request"}]}));
pending.provider_request_body = Some(json!({"input": "provider request"}));
pending.request_body_state = Some(UsageBodyCaptureState::Inline);
pending.provider_request_body_state = Some(UsageBodyCaptureState::Inline);
let stored_pending = repository.upsert(pending.clone()).await.unwrap();
assert_eq!(stored_pending.request_body, pending.request_body);
assert_eq!(
stored_pending.provider_request_body,
pending.provider_request_body
);
assert_eq!(
stored_pending.request_headers,
Some(json!({"content-type": "application/json", "authorization": "[redacted]"}))
);
let mut streaming = sample_upsert_usage_record(&pending.request_id);
streaming.status = "streaming".to_string();
streaming.updated_at_unix_secs += 1;
let stored_streaming = repository.upsert(streaming).await.unwrap();
assert_eq!(stored_streaming.request_body, pending.request_body);
assert_eq!(
stored_streaming.provider_request_body,
pending.provider_request_body
);
let mut terminal = sample_upsert_usage_record(&pending.request_id);
terminal.status = "completed".to_string();
terminal.updated_at_unix_secs += 2;
terminal.finalized_at_unix_secs = Some(terminal.updated_at_unix_secs);
terminal.response_headers =
Some(json!({"content-type": "text/event-stream", "set-cookie": "private"}));
terminal.response_body = Some(json!("data: upstream response\n\ndata: [DONE]\n\n"));
terminal.client_response_body =
Some(json!({"choices": [{"message": {"content": "client response"}}]}));
let stored_terminal = repository.upsert(terminal.clone()).await.unwrap();
assert_eq!(stored_terminal.request_body, pending.request_body);
assert_eq!(
stored_terminal.provider_request_body,
pending.provider_request_body
);
assert_eq!(stored_terminal.response_body, terminal.response_body);
assert_eq!(
stored_terminal.client_response_body,
terminal.client_response_body
);
assert_eq!(
stored_terminal.response_headers,
Some(json!({"content-type": "text/event-stream", "set-cookie": "[redacted]"}))
);
let found = repository
.find_by_request_id(&pending.request_id)
.await
.unwrap()
.unwrap();
assert_eq!(found.request_body, pending.request_body);
assert_eq!(found.response_body, terminal.response_body);
assert_eq!(
repository.upsert(pending).await.unwrap().response_body,
terminal.response_body
);
}
#[tokio::test] #[tokio::test]
async fn upsert_uses_typed_provider_capture_as_the_fast_fact_snapshot() { async fn upsert_uses_typed_provider_capture_as_the_fast_fact_snapshot() {
for (name, state, incoming_tier, expected_tier) in [ for (name, state, incoming_tier, expected_tier) in [
@@ -238,7 +238,28 @@ impl GenericProviderOAuthAdapter {
return required_client_secret(env_name, non_empty_owned(value)).map(Some); return required_client_secret(env_name, non_empty_owned(value)).map(Some);
} }
required_client_secret(env_name, non_empty_environment_value(env_name)).map(Some) let configured = non_empty_environment_value(env_name);
self.resolve_client_secret(configured).map(Some)
}
fn resolve_client_secret(&self, configured: Option<String>) -> Result<String, OAuthError> {
let configured = configured.or_else(|| {
let default_template = template_for_provider_type(self.template.provider_type)?;
if self.client_id() != default_template.client_id {
return None;
}
match self.template.provider_type {
"gemini_cli" => Some("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".to_string()),
"antigravity" => Some("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".to_string()),
_ => None,
}
});
required_client_secret(
self.template
.client_secret_env
.unwrap_or(ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV),
configured,
)
} }
async fn exchange_grant( async fn exchange_grant(
@@ -957,6 +978,81 @@ mod tests {
); );
} }
#[tokio::test]
async fn google_default_credentials_support_authorization_and_refresh() {
for provider_type in ["gemini_cli", "antigravity"] {
let mut template = template_for_provider_type(provider_type).expect("template");
template.client_id_env = None;
template.client_secret_env = Some("AETHER_TEST_UNUSED_ANTIGRAVITY_SECRET");
let adapter = GenericProviderOAuthAdapter::new(template);
let ctx = transport_context(provider_type);
adapter
.build_authorize_url(&ctx, "state", None)
.expect("authorize");
let seen_request = Arc::new(Mutex::new(None));
let executor = StaticExecutor {
seen_request: Arc::clone(&seen_request),
response_payload: json!({"access_token": "new-token", "expires_in": 3600}),
};
adapter
.refresh(&executor, &ctx, &oauth_account(provider_type))
.await
.expect("refresh");
let seen = seen_request.lock().expect("lock").clone().expect("request");
let body = seen.body_bytes.expect("body");
let fields = url::form_urlencoded::parse(&body)
.into_owned()
.collect::<BTreeMap<_, _>>();
assert_eq!(fields["client_id"], template.client_id);
assert_eq!(
fields["client_secret"],
adapter.resolve_client_secret(None).expect("default")
);
assert_eq!(fields["refresh_token"], "old-refresh-token");
assert!(!format!("{adapter:?}").contains(&fields["client_secret"]));
}
}
#[test]
fn google_custom_client_requires_its_own_secret() {
for provider_type in ["gemini_cli", "antigravity"] {
let adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)
.expect("adapter")
.with_oauth_credentials_for_tests("custom-client", "custom-secret");
assert!(adapter.resolve_client_secret(None).is_err());
assert_eq!(
adapter
.resolve_client_secret(Some("custom-secret".to_string()))
.expect("configured secret"),
"custom-secret"
);
assert_eq!(
adapter.client_secret().expect("override"),
Some("custom-secret".to_string())
);
}
}
#[test]
fn google_configured_secret_overrides_the_native_app_default() {
for provider_type in ["gemini_cli", "antigravity"] {
let adapter =
GenericProviderOAuthAdapter::for_provider_type(provider_type).expect("adapter");
assert_eq!(
adapter
.resolve_client_secret(Some("configured-secret".to_string()))
.unwrap(),
"configured-secret"
);
let mut template = template_for_provider_type(provider_type).expect("template");
template.client_id = "custom-template-client";
template.client_id_env = None;
assert!(GenericProviderOAuthAdapter::new(template)
.resolve_client_secret(None)
.is_err());
}
}
#[test] #[test]
fn generic_adapter_debug_redacts_oauth_credentials() { fn generic_adapter_debug_redacts_oauth_credentials() {
let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli") let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli")
+111 -16
View File
@@ -11,6 +11,7 @@ use crate::wait_until;
pub struct ManagedPostgresServer { pub struct ManagedPostgresServer {
child: Option<Child>, child: Option<Child>,
postgres_bin: String, postgres_bin: String,
pg_ctl_bin: PathBuf,
port: u16, port: u16,
workdir: PathBuf, workdir: PathBuf,
data_dir: PathBuf, data_dir: PathBuf,
@@ -26,7 +27,7 @@ impl ManagedPostgresServer {
port port
)); ));
let data_dir = workdir.join("data"); let data_dir = workdir.join("data");
std::fs::create_dir_all(&workdir)?; std::fs::create_dir(&workdir)?;
let initdb_bin = std::env::var("AETHER_INITDB_BIN") let initdb_bin = std::env::var("AETHER_INITDB_BIN")
.ok() .ok()
@@ -36,10 +37,31 @@ impl ManagedPostgresServer {
.ok() .ok()
.filter(|value| !value.trim().is_empty()) .filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "postgres".to_string()); .unwrap_or_else(|| "postgres".to_string());
let pg_ctl_bin = std::env::var("AETHER_PG_CTL_BIN")
.ok()
.filter(|value| !value.trim().is_empty())
.map(PathBuf::from)
.unwrap_or_else(|| {
PathBuf::from(&postgres_bin).with_file_name(if cfg!(windows) {
"pg_ctl.exe"
} else {
"pg_ctl"
})
});
let database_url = format!("postgres://[email protected]:{port}/postgres");
let mut server = Self {
child: None,
postgres_bin,
pg_ctl_bin,
port,
workdir,
data_dir,
database_url,
};
let init_output = Command::new(&initdb_bin) let init_output = Command::new(&initdb_bin)
.arg("-D") .arg("-D")
.arg(&data_dir) .arg(&server.data_dir)
.arg("-U") .arg("-U")
.arg("aether") .arg("aether")
.arg("--auth=trust") .arg("--auth=trust")
@@ -54,15 +76,6 @@ impl ManagedPostgresServer {
.into()); .into());
} }
let database_url = format!("postgres://[email protected]:{port}/postgres");
let mut server = Self {
child: None,
postgres_bin,
port,
workdir,
data_dir,
database_url,
};
server.restart().await?; server.restart().await?;
Ok(server) Ok(server)
} }
@@ -76,10 +89,29 @@ impl ManagedPostgresServer {
} }
pub fn stop(&mut self) -> Result<(), std::io::Error> { pub fn stop(&mut self) -> Result<(), std::io::Error> {
if let Some(mut child) = self.child.take() { let Some(child) = self.child.as_mut() else {
let _ = child.kill(); return Ok(());
let _ = child.wait(); };
if child.try_wait()?.is_some() {
self.child = None;
return Ok(());
} }
let output = Command::new(&self.pg_ctl_bin)
.arg("-D")
.arg(&self.data_dir)
.args(["stop", "-m", "fast", "-w", "-t", "10"])
.output()?;
if !output.status.success() && child.try_wait()?.is_none() {
return Err(std::io::Error::other(format!(
"pg_ctl stop failed for {}: {}{}",
self.data_dir.display(),
String::from_utf8_lossy(&output.stdout),
String::from_utf8_lossy(&output.stderr),
)));
}
child.wait()?;
self.child = None;
Ok(()) Ok(())
} }
@@ -139,8 +171,17 @@ impl ManagedPostgresServer {
impl Drop for ManagedPostgresServer { impl Drop for ManagedPostgresServer {
fn drop(&mut self) { fn drop(&mut self) {
let _ = self.stop(); match self.stop() {
let _ = std::fs::remove_dir_all(&self.workdir); Ok(()) => {
let _ = std::fs::remove_dir_all(&self.workdir);
}
Err(error) => {
eprintln!(
"failed to stop managed postgres; preserving {}: {error}",
self.workdir.display(),
);
}
}
} }
} }
@@ -170,3 +211,57 @@ fn reserve_local_port() -> Result<u16, std::io::Error> {
drop(listener); drop(listener);
Ok(port) Ok(port)
} }
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
#[ignore = "requires local initdb, postgres, and pg_ctl binaries"]
async fn live_managed_postgres_restarts_cleanly_with_open_connections() {
let mut server = ManagedPostgresServer::start().await.unwrap();
let workdir = server.workdir.clone();
let mut connection = PgConnection::connect(server.database_url()).await.unwrap();
sqlx::query("CREATE TABLE restart_probe (value INTEGER NOT NULL)")
.execute(&mut connection)
.await
.unwrap();
sqlx::query("INSERT INTO restart_probe VALUES (42)")
.execute(&mut connection)
.await
.unwrap();
for _iteration in 0..4 {
server.stop().unwrap();
server.stop().unwrap();
assert!(server.child.is_none());
assert!(!server.data_dir.join("postmaster.pid").exists());
assert!(server.data_dir.exists());
server.restart().await.unwrap();
connection = PgConnection::connect(server.database_url()).await.unwrap();
let value: i32 = sqlx::query_scalar("SELECT value FROM restart_probe")
.fetch_one(&mut connection)
.await
.unwrap();
assert_eq!(value, 42);
}
drop(server);
assert!(!workdir.exists());
}
#[tokio::test]
#[ignore = "requires local initdb, postgres, and pg_ctl binaries"]
async fn live_failed_postgres_stop_can_be_retried_without_losing_ownership() {
let mut server = ManagedPostgresServer::start().await.unwrap();
let pg_ctl_bin = server.pg_ctl_bin.clone();
server.pg_ctl_bin = server.workdir.join("missing-pg-ctl");
assert!(server.stop().is_err());
assert!(server.child.as_mut().unwrap().try_wait().unwrap().is_none());
assert!(server.data_dir.exists());
server.pg_ctl_bin = pg_ctl_bin;
server.stop().unwrap();
assert!(server.child.is_none());
assert!(!server.data_dir.join("postmaster.pid").exists());
}
}
+2 -37
View File
@@ -2726,25 +2726,8 @@ fn headers_to_json(headers: &BTreeMap<String, String>) -> Option<Value> {
const REDACTED_USAGE_VALUE: &str = "[redacted]"; const REDACTED_USAGE_VALUE: &str = "[redacted]";
/// Only headers whose values are protocol metadata are persisted verbatim.
/// Unknown headers are treated as credentials because providers commonly use
/// custom `X-*` names for authentication.
const SAFE_USAGE_HEADER_VALUE_NAMES: &[&str] = &[
"accept",
"accept-encoding",
"content-encoding",
"content-length",
"content-type",
"transfer-encoding",
"x-request-id",
"x-trace-id",
];
fn is_sensitive_header(name: &str) -> bool { fn is_sensitive_header(name: &str) -> bool {
let trimmed = name.trim(); aether_data_contracts::repository::usage::usage_header_value_is_sensitive(name)
!SAFE_USAGE_HEADER_VALUE_NAMES
.iter()
.any(|candidate| trimmed.eq_ignore_ascii_case(candidate))
} }
fn mask_header_value(name: &str, value: &str) -> String { fn mask_header_value(name: &str, value: &str) -> String {
@@ -2761,25 +2744,7 @@ fn mask_sensitive_header_value(_value: &str) -> String {
/// Non-object values cannot be established as a valid header map and are /// Non-object values cannot be established as a valid header map and are
/// discarded instead of being persisted verbatim. /// discarded instead of being persisted verbatim.
fn mask_sensitive_headers_in_json_value(value: Option<Value>) -> Option<Value> { fn mask_sensitive_headers_in_json_value(value: Option<Value>) -> Option<Value> {
let mut value = value?; aether_data_contracts::repository::usage::sanitize_usage_headers_for_persistence(value)
let Value::Object(map) = &mut value else {
return None;
};
for (key, val) in map.iter_mut() {
if !is_sensitive_header(key) {
continue;
}
match val {
Value::String(text) => {
*text = mask_sensitive_header_value(text);
}
Value::Null => {}
other => {
*other = Value::String(mask_sensitive_header_value(&other.to_string()));
}
}
}
Some(value)
} }
fn mask_sensitive_body_fields(mut value: Value) -> Value { fn mask_sensitive_body_fields(mut value: Value) -> Value {
+61 -4
View File
@@ -4,6 +4,7 @@ use aether_data_contracts::repository::video_tasks::{
}; };
use serde_json::{json, Map, Value}; use serde_json::{json, Map, Value};
use crate::transport::{gemini_response_video_url, gemini_video_metadata};
use crate::types::sanitize_video_task_error_code; use crate::types::sanitize_video_task_error_code;
use crate::{ use crate::{
build_video_follow_up_report_context, current_unix_timestamp_secs, gemini_metadata_video_url, build_video_follow_up_report_context, current_unix_timestamp_secs, gemini_metadata_video_url,
@@ -112,7 +113,16 @@ impl GeminiVideoTaskSeed {
self.error_code = None; self.error_code = None;
self.error_message = None; self.error_message = None;
} }
self.metadata = json!({}); self.metadata = if error.is_none() {
gemini_video_metadata(
provider_body
.get("response")
.and_then(gemini_response_video_url)
.as_deref(),
)
} else {
json!({})
};
return; return;
} }
@@ -365,8 +375,8 @@ mod tests {
use serde_json::json; use serde_json::json;
use crate::{ use crate::{
GeminiVideoTaskSeed, LocalVideoTaskPersistence, LocalVideoTaskStatus, GeminiVideoTaskSeed, LocalVideoTaskPersistence, LocalVideoTaskSnapshot,
LocalVideoTaskTransport, LocalVideoTaskStatus, LocalVideoTaskTransport,
}; };
use super::map_gemini_stored_task_to_read_response; use super::map_gemini_stored_task_to_read_response;
@@ -484,7 +494,54 @@ mod tests {
assert!(record.request_metadata.is_none()); assert!(record.request_metadata.is_none());
assert_eq!( assert_eq!(
record.video_url.as_deref(), record.video_url.as_deref(),
Some("https://files.example/video.mp4?alt=media") Some("https://files.example/video.mp4?alt=media&token=sensitive")
); );
assert_eq!(record.prompt.as_deref(), Some("business prompt"));
let transport = seed.transport.clone();
let mut snapshot = LocalVideoTaskSnapshot::Gemini(seed);
snapshot.apply_provider_body(json!({
"done": true,
"debug": "private-provider-debug",
"response": {
"provider_token": "private-provider-token",
"generateVideoResponse": {
"generatedSamples": [{
"video": {
"uri": "https://files.example/video.mp4?alt=media&token=signed%2Bvalue",
"debug": "private-video-debug"
}
}]
}
}
}).as_object().expect("provider body"));
assert!(!serde_json::to_string(&snapshot)
.expect("serialize snapshot")
.contains("private-"));
snapshot.sanitize_persisted_diagnostics();
let stored = snapshot.to_upsert_record().into_stored();
assert_eq!(stored.status, VideoTaskStatus::Completed);
assert_eq!(stored.prompt.as_deref(), Some("business prompt"));
assert_eq!(
stored.video_url.as_deref(),
Some("https://files.example/video.mp4?alt=media&token=signed%2Bvalue")
);
assert!(stored.request_metadata.is_none());
assert!(stored.original_request_body.is_none());
let mut restored =
LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, transport)
.expect("stored Gemini task should reconstruct");
restored.sanitize_persisted_diagnostics();
let restored_record = restored.to_upsert_record();
assert_eq!(restored_record.video_url, stored.video_url);
assert_eq!(restored_record.prompt, stored.prompt);
let response = restored.read_response();
assert_eq!(response.body_json["done"], true);
assert!(response
.body_json
.to_string()
.contains("/v1beta/files/aev_"));
assert!(!response.body_json.to_string().contains("signed%2Bvalue"));
} }
} }
+28 -4
View File
@@ -610,7 +610,7 @@ impl OpenAiVideoTaskSeed {
updated_at_unix_secs: self.completed_at_unix_secs.unwrap_or(now_unix_secs), updated_at_unix_secs: self.completed_at_unix_secs.unwrap_or(now_unix_secs),
error_code: self.error_code.clone(), error_code: self.error_code.clone(),
error_message: None, error_message: None,
video_url: None, video_url: self.video_url.clone(),
request_metadata: None, request_metadata: None,
}; };
record.sanitize_for_persistence(); record.sanitize_for_persistence();
@@ -626,8 +626,8 @@ mod tests {
use serde_json::json; use serde_json::json;
use crate::{ use crate::{
LocalVideoTaskPersistence, LocalVideoTaskStatus, LocalVideoTaskTransport, LocalVideoTaskContentAction, LocalVideoTaskPersistence, LocalVideoTaskSnapshot,
OpenAiVideoTaskSeed, LocalVideoTaskStatus, LocalVideoTaskTransport, OpenAiVideoTaskSeed,
}; };
use super::map_openai_stored_task_to_read_response; use super::map_openai_stored_task_to_read_response;
@@ -751,9 +751,33 @@ mod tests {
assert!(record.original_request_body.is_none()); assert!(record.original_request_body.is_none());
assert!(record.progress_message.is_none()); assert!(record.progress_message.is_none());
assert!(record.error_message.is_none()); assert!(record.error_message.is_none());
assert!(record.video_url.is_none()); assert_eq!(record.video_url, seed.video_url);
assert_eq!(record.prompt, seed.prompt);
assert!(record.request_metadata.is_none()); assert!(record.request_metadata.is_none());
assert_eq!(record.duration_seconds, Some(4)); assert_eq!(record.duration_seconds, Some(4));
assert_eq!(record.size.as_deref(), Some("1280x720")); assert_eq!(record.size.as_deref(), Some("1280x720"));
let mut stored = record.into_stored();
stored.status = VideoTaskStatus::Completed;
let snapshot =
LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, seed.transport)
.expect("stored task should reconstruct with current transport");
let LocalVideoTaskSnapshot::OpenAi(restored) = snapshot else {
panic!("expected OpenAI snapshot");
};
assert_eq!(restored.prompt, stored.prompt);
assert_eq!(restored.to_upsert_record().video_url, stored.video_url);
let Some(LocalVideoTaskContentAction::StreamPlan(plan)) =
restored.build_content_stream_action(None, "trace-download")
else {
panic!("completed stored task should stream content");
};
assert_eq!(Some(plan.url.as_str()), stored.video_url.as_deref());
assert!(plan.headers.is_empty());
let response = map_openai_stored_task_to_read_response(stored.clone());
assert_eq!(
response.body_json["video_url"].as_str(),
stored.video_url.as_deref()
);
} }
} }
@@ -1,6 +1,7 @@
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, UpsertVideoTask}; use aether_data_contracts::repository::video_tasks::{StoredVideoTask, UpsertVideoTask};
use serde_json::{json, Map, Value}; use serde_json::{json, Map, Value};
use crate::transport::{gemini_metadata_video_url, gemini_video_metadata};
use crate::types::sanitize_video_task_error_code; use crate::types::sanitize_video_task_error_code;
use crate::{ use crate::{
local_status_from_stored, non_empty_owned, request_body_string, GeminiVideoTaskSeed, local_status_from_stored, non_empty_owned, request_body_string, GeminiVideoTaskSeed,
@@ -104,7 +105,7 @@ impl LocalVideoTaskSnapshot {
progress_percent: task.progress_percent, progress_percent: task.progress_percent,
error_code: task.error_code.clone(), error_code: task.error_code.clone(),
error_message: task.error_message.clone(), error_message: task.error_message.clone(),
metadata: Value::Object(Map::new()), metadata: gemini_video_metadata(task.video_url.as_deref()),
persistence, persistence,
transport, transport,
})) }))
@@ -128,7 +129,8 @@ impl LocalVideoTaskSnapshot {
let previous_error_code = seed.error_code.clone(); let previous_error_code = seed.error_code.clone();
let error_code = let error_code =
sanitized_error_code_for_status(seed.status, seed.error_code.take()); sanitized_error_code_for_status(seed.status, seed.error_code.take());
let safe_metadata = Value::Object(Map::new()); let safe_metadata =
gemini_video_metadata(gemini_metadata_video_url(&seed.metadata).as_deref());
let changed = previous_error_code != error_code let changed = previous_error_code != error_code
|| seed.error_message.is_some() || seed.error_message.is_some()
|| seed.metadata != safe_metadata; || seed.metadata != safe_metadata;
@@ -450,6 +450,49 @@ mod tests {
}) })
} }
#[test]
fn completed_gemini_download_survives_encrypted_file_reload() {
let path = temp_store_path("completed-download");
let video_url = "https://cdn.example.test/video.mp4?signature=a%2Fb%2Bc%3D";
let mut snapshot = sensitive_gemini_snapshot();
snapshot.apply_provider_body(
json!({
"done": true,
"debug": "private-debug",
"response": {
"generateVideoResponse": {
"generatedSamples": [{"video": {"uri": video_url}}]
}
}
})
.as_object()
.expect("provider body"),
);
let store = FileVideoTaskStore::new(&path, DEVELOPMENT_ENCRYPTION_KEY).expect("store");
store.insert(snapshot);
drop(store);
let bytes = std::fs::read_to_string(&path).expect("encrypted store file");
assert!(bytes.starts_with(ENCRYPTED_VIDEO_TASK_STORE_PREFIX));
assert!(!bytes.contains(video_url));
assert!(!bytes.contains("transport-key-required-for-resume"));
let restored =
FileVideoTaskStore::new(&path, DEVELOPMENT_ENCRYPTION_KEY).expect("restored store");
let task = restored
.clone_gemini("task-sensitive")
.expect("completed Gemini task");
let record = task.to_upsert_record();
assert_eq!(
record.status,
aether_data_contracts::repository::video_tasks::VideoTaskStatus::Completed
);
assert_eq!(record.video_url.as_deref(), Some(video_url));
assert_eq!(record.prompt.as_deref(), Some("create a video"));
assert!(!task.metadata.to_string().contains("private-debug"));
assert!(record.request_metadata.is_none());
drop(restored);
cleanup_store_path(&path);
}
#[test] #[test]
fn loading_encrypted_store_rewrites_legacy_provider_diagnostics() { fn loading_encrypted_store_rewrites_legacy_provider_diagnostics() {
let path = temp_store_path("diagnostic-migration"); let path = temp_store_path("diagnostic-migration");
@@ -1,4 +1,4 @@
use serde_json::Value; use serde_json::{json, Value};
use crate::LocalVideoTaskStatus; use crate::LocalVideoTaskStatus;
@@ -20,9 +20,12 @@ pub fn parse_video_content_variant(query_string: Option<&str>) -> Option<&'stati
} }
pub fn gemini_metadata_video_url(metadata: &Value) -> Option<String> { pub fn gemini_metadata_video_url(metadata: &Value) -> Option<String> {
metadata metadata.get("response").and_then(gemini_response_video_url)
.get("response") }
.and_then(|value| value.get("generateVideoResponse"))
pub(crate) fn gemini_response_video_url(response: &Value) -> Option<String> {
response
.get("generateVideoResponse")
.and_then(|value| value.get("generatedSamples")) .and_then(|value| value.get("generatedSamples"))
.and_then(Value::as_array) .and_then(Value::as_array)
.and_then(|value| value.first()) .and_then(|value| value.first())
@@ -32,6 +35,19 @@ pub fn gemini_metadata_video_url(metadata: &Value) -> Option<String> {
.map(str::to_string) .map(str::to_string)
} }
pub(crate) fn gemini_video_metadata(video_url: Option<&str>) -> Value {
match video_url {
Some(video_url) => json!({
"response": {
"generateVideoResponse": {
"generatedSamples": [{"video": {"uri": video_url}}]
}
}
}),
None => json!({}),
}
}
pub fn map_openai_task_status(status: LocalVideoTaskStatus) -> &'static str { pub fn map_openai_task_status(status: LocalVideoTaskStatus) -> &'static str {
match status { match status {
LocalVideoTaskStatus::Submitted | LocalVideoTaskStatus::Queued => "queued", LocalVideoTaskStatus::Submitted | LocalVideoTaskStatus::Queued => "queued",
@@ -98,10 +98,15 @@ impl LocalVideoTaskPersistence {
api_key_name: task.api_key_name.clone(), api_key_name: task.api_key_name.clone(),
client_api_format, client_api_format,
provider_api_format, provider_api_format,
original_request_body: task original_request_body: task.original_request_body.clone().unwrap_or_else(|| {
.original_request_body serde_json::json!({
.clone() "prompt": task.prompt,
.unwrap_or_else(|| Value::Object(Map::new())), "seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size,
})
}),
format_converted: task.format_converted, format_converted: task.format_converted,
}) })
} }
+9
View File
@@ -17,6 +17,15 @@
`total_keys` 和 `active_keys` 仍反映密钥配置数量,不因缺少健康数据而减少。 `total_keys` 和 `active_keys` 仍反映密钥配置数量,不因缺少健康数据而减少。
## 数据读取
摘要从密钥的轻量投影读取 API 格式、启用状态和 `health_by_format`。
PostgreSQL 投影中的凭据字段使用 `summary` / `{}` 等脱敏占位值,并非真实密文,
因此摘要读取不执行凭据解密、认证或迁移。完整密钥读取仍保留原有的凭据安全校验。
若将这些占位值送入凭据校验,读取会失败,并被摘要聚合当作空密钥列表,
导致已配置密钥的端点也被错误显示为灰色;不能通过给缺失分数默认填 `100%` 来修复。
## 页面展示 ## 页面展示
桌面表格和手机卡片使用相同规则: 桌面表格和手机卡片使用相同规则:
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,98 @@
# 安全加固兼容性复核(2026-09-07)
后续扩展到全仓的检查范围、新发现、逐文件清单及限制见 `security-hardening-full-audit.md` 和 `security-hardening-audit-manifest.tsv`。
## 范围
本轮针对 `579f2c7cc`(2026-09-04)及其后续修复进行复核。该提交涉及 1019 个文件,混合了安全边界、订阅计费、数据持久化和运行时变更,不适合整体 revert。
本轮重点检查管理端脱敏与编辑往返、Provider OAuth 默认配置、已有 DNS/代理兼容性修复、管理员错误诊断、SMTP/LDAP/OAuth 配置、导入导出及请求大小限制。以下是已确认的问题及处理,不代表对全部文件逐行审计或对生产环境的完整验证。
## 新增:管理员按需查看规则原值
- 端点管理的请求/响应规则工具栏增加“查看原值”。只读弹窗展示服务端已保存的请求头、请求体和响应头规则,不覆盖未保存的草稿。
- 新接口:`GET /api/admin/endpoints/{endpoint_id}/rules/reveal`。
- 仅管理员可访问;管理令牌需要 `admin:endpoints_manage:admin`,只读和普通写权限不足。
- 返回 `Cache-Control: no-store`、`Pragma: no-cache`,记录 `admin_endpoint_rules_revealed` 审计事件。响应仅包含三组规则,不附带端点的其他配置或凭据。
- 关闭弹窗、切换端点或卸载组件时取消请求并清除明文;旧请求的迟到结果不会重新显示。
## 本轮修复的六类误伤
| 问题 | 影响 | 处理 |
| --- | --- | --- |
| 普通协议头也全部脱敏 | `Content-Type`、`User-Agent`、版本/客户端信息及相应条件无法正常查看;配置导出也受影响 | 恢复明确的常用非凭据头的显示;认证头、Cookie 和未知自定义头继续默认隐藏。旧版本返回的保留标记仍可正常保存 |
| 重复规则无法恢复原值 | 同一头或路径配置多条条件规则后,保存或重排可能将原值变成 `***` | 在相同操作和目标范围内按可区分的规则结构匹配;未修改的重复规则保留原顺序与原值。无法唯一定位的脱敏编辑明确报错,禁止静默写入占位符 |
| 嵌套条件组丢失原值 | 多层 `all`/`any` 在修改其他条件或重排后,内部脱敏值无法恢复 | 增加递归条件组匹配、未修改条件组往返保留和保留值校验 |
| URL 类型判断过宽 | `image_url` 等对象/数组被变成 `null`,内嵌 `data:` 图片 URL 也受影响,继续编辑可能破坏请求规则 | 保留结构化 URL 和内嵌数据;递归处理网络 URL,继续隐藏并在保存时恢复网络凭据 |
| 保存后的编辑状态没有同步 | 后台返回脱敏数据后,界面仍保留提交前明文,持续显示“未保存” | 使用保存响应重新建立编辑基准;保存期间的新编辑不被覆盖,关闭后清除草稿,忽略跨弹窗的旧响应 |
| Gemini CLI 默认 OAuth 被误移除 | 未额外配置环境变量时,原有默认客户端不能授权或刷新 | 恢复加固前内置 native-app 客户端的配套默认值;显式配置优先,自定义客户端 ID 仍必须提供自己的凭据,调试输出继续脱敏 |
明确可直接显示的头包括 `Accept`/编码/语言、`Content-Type`/编码、`Cache-Control`、`User-Agent`、Anthropic/OpenAI 协议版本与 beta 标记,以及明确列出的 `x-stainless-*` 客户端运行时字段;不是按前缀放行任意自定义头。
## 补充修复:显式 `full` 请求记录被禁用
原则:安全处理应约束未授权访问、未配置时的默认行为和真实凭据泄露,不应静默覆盖管理员明确启用的功能。
本次确认不是单纯的前端隐藏,而是同一链路上的多重清空:
- `request_record_level=full` 在运行时被强制解析为 `basic`。
- 独立 HTTP 审计存储的输入投影清除了所有请求头、正文、正文引用和采集状态。
- 内存仓库每次更新都丢弃正文与引用;PostgreSQL 的单条写入把正文 blob 同步改成删除,并拒绝 HTTP 审计内容。
处理:恢复显式 `full`(含旧配置名 `request_log_level`)的采集与独立存储,保留当前配置名优先;恢复请求、上游请求、上游响应和客户端响应各方向已有的采集内容,流转状态更新不再无条件删除它们。PostgreSQL 正文仍写入原有压缩 blob 表,HTTP 头与引用仍写入独立审计表,不重新塞回计费主表。
保留:缺失、无效配置和读取配置失败时默认 `basic`;`basic` 不记录正文;OAuth 令牌交换等内部凭据请求不采集正文;认证头及未知自定义头继续脱敏;正文引用仍校验所属请求和字段;管理端权限、审计及无缓存策略不变。普通同格式非流式响应仍只保存一份原有响应体,前端沿用回退显示,不制造重复的客户端副本。
新增回归覆盖配置解析、内存生命周期、流式/非流式管理端正文读取、旧正文大小配置不覆盖 `full`、PostgreSQL 单条/批量写入与正文读回、前端按需加载与无正文时的展示。真实数据库测试使用本机临时隔离 PostgreSQL,不访问现有业务数据库。
此修复只能恢复之后新采集的记录;此前未采集或已删除的正文无法由代码补回。运行时沿用原有的 30 秒采集策略缓存,刚切换记录级别时需等待缓存刷新。
扩大执行原有、默认忽略的 PostgreSQL 测试时,另发现 5 个旧用例的测试数据/类型与当前 schema 不兼容:4 个使用超过 `varchar(36)` 限制的 Provider Key ID,1 个直接按 `f64` 解码 `NUMERIC` 列。它们并非本次正文修复引入,也不作为本次已通过项;未修改相关业务逻辑或测试夹具。
## 继续复核:视频任务链路的四类回归
继续沿“业务字段被当作诊断数据清空”“读取投影与更新条件不一致”检查,另确认以下四类问题,均可定位到本次加固新增的清空或校验条件,不是用户配置错误:
| 问题 | 实际影响 | 本轮修复 |
| --- | --- | --- |
| 任务业务信息被清空 | 提示词、用户名和客户端 Key 名称丢失,管理页显示空白或 `Unknown`;提示词被写成 `NULL` 还与 PostgreSQL 的非空约束冲突,可能直接阻止任务入库 | 保留这些明确用于任务展示的业务字段,不再作为敏感诊断一律删除 |
| 下载地址被清空或破坏签名 | OpenAI 任务的 `video_url` 被无条件丢弃;Gemini 地址仅保留 `alt=media`,使下载签名、有效期等参数失效 | 保留 HTTP(S) 产物地址及完整查询串,不重排重复参数、不重编码签名;继续拒绝非法协议和 URL userinfo,移除 fragment |
| Gemini 完成结果和重载丢失 | 完成轮询直接清空元数据;数据库重建与加密文件加载再次清空结果 URI,出现任务已完成但无法取得视频 | 仅保留下载所需的最小结果 URI,不保留上游完整响应;从数据库的业务字段重建必要的提示词、尺寸信息和结果 URI,保证再次保存不会丢失 |
| PostgreSQL 轮询更新条件自相矛盾 | 领取任务时不读出时长、分辨率、宽高比和尺寸,而加固后的更新要求这些不可变字段逐项相等,导致正常轮询更新被拒绝、任务停留在处理中 | 领取时读回必要业务字段,保留原有归属/身份校验和并发领取锁,不靠放宽更新条件绕过问题 |
保留原有播放/下载接口与鉴权语义:签名产物链接是有权限用户访问任务结果所需的业务数据,不等同于可随意删除的调试凭据。OpenAI 直接下载不附带 Provider 认证头;Gemini 仍通过网关文件接口访问,并校验上游地址同源后才附加 Provider Key。私网/保留地址拦截、跨域凭据隔离、任务归属校验及管理操作审计不撤销。
任务数据库仍不保存完整请求体、原始错误消息或含认证头的传输快照;本地文件仍使用原有认证加密格式,调试输出继续脱敏。回归覆盖任务保存后重读、签名参数顺序与编码、加密文件重载、管理列表/详情、带审计的下载,以及轮询前后提示词和尺寸不丢失。
真实 PostgreSQL 回归使用临时 Unix socket 实例,并从正式迁移后的 schema 复制会话隔离的任务表。OpenAI/Gemini 两种任务均验证了入库、领取、正常完成更新、拒绝不匹配的尺寸更新和重新查询签名地址;不连接现有业务数据库。
这四类是本轮已经确认的新增回归,不代表对混合提交全部 1019 个文件给出“零遗漏”保证。原有用户策略字段停用早于本次加固,不据此回退;后台任务和缓存中的裁剪未发现足以确认本次功能回归的调用链,不做猜测性恢复。数据库仍保存的历史值可以重新读取;已经写空、未入库或已过期的资源不能凭本次源码修复重建。
## 已存在的后续修复:保留,不重复撤销
- `522b97905`:已移除普通 Provider 代理的 DNS 地址过滤和相关白名单设置。
- `696273122`:已恢复 Antigravity 默认 native-app OAuth 配套凭据;本轮将同类遗漏补到 Gemini CLI。
- `a26680f46`:已恢复旧 SMTP 密码迁移。
- `062e111c0`:已恢复管理员查看上游错误诊断的能力,同时保留普通用户侧脱敏。
- `14f96c9fa` 等:已修复脱敏后的密钥健康状态摘要,不将其退回旧实现。
## 保留的安全边界
- 管理员/普通用户与管理令牌的权限隔离;凭据查看接口的审计和无缓存要求。
- 凭据加密存储及凭据与目标/身份绑定;未知头和真实敏感字段的默认脱敏。
- 登录 OAuth、支付、隧道中继等独立网络边界,TLS 校验和隧道防重放。
- 有界请求/响应缓冲、备份恢复模式、导入校验、计费与配额一致性;本轮没有因“安全加固”标签撤回这些功能。
## 回归覆盖与使用限制
回归用例覆盖规则投影/恢复、重复与嵌套规则、占位符拒绝、结构化 URL、保存期间的并发编辑、弹窗关闭/切换时的迟到响应、接口权限/审计/无缓存,以及 Gemini CLI 和 Antigravity 的默认配置与显式覆盖。
早期兼容性修复的阶段性验证记录(最终全量结果见 `security-hardening-full-audit.md`):
- 视频数据契约 12 项、视频核心 33 项(含加密文件重载)、内存视频仓库 11 项、PostgreSQL 视频仓库 9 项、Provider 视频传输 9 项均通过。
- 新增的真实 PostgreSQL 回归单独执行通过,同时覆盖 OpenAI 和 Gemini;临时数据库实例已停止并清理。
- 网关 `video` 101 项、`async_task::` 20 项、`usage` 302 项、端点规则原值查看 3 项、管理令牌权限 28 项均通过;筛选结果可能重叠,不合计为独立用例总数。
- 前端 199 个测试文件、1411 项测试通过;补齐保存测试的类型夹具后,该文件 4 项测试及 ESLint 再次通过;Rust 格式检查和 `git diff --check` 通过。
- 阶段性检查曾发现 461 条既有类型诊断,与未修改 HEAD 的同依赖基线一致。2026-09-07 的“全部修复”续轮已修正这些诊断,并将 `npm run type-check` 接入 `vue-tsc -b --force --pretty false`;当前真实全量类型检查为 0 错误。最新全量测试、构建及剩余验证边界以 `security-hardening-full-audit.md` 为准,未关闭严格模式或排除测试。
上述修复与回归不涉及生产服务器部署或现有业务数据库更改;源码提交、推送不代表生产环境已经部署或完成验证。若历史版本已经把真实值覆盖为字面量 `***`,查看接口不能重建丢失的原值,需重新填写或从可靠备份恢复。
@@ -0,0 +1,134 @@
# 安全加固全仓兼容性复核(2026-09-07)
## 范围与方法
- 加固基线:`579f2c7cc`(2026-09-04);复核时 HEAD:`a5c3699ae`。保留已有未提交的兼容性修复,不整体回退该混合提交,也不撤销后续修复。
- 历史提交涉及 1019 个文件,其中 907 个路径仍存在,112 个已在后续重构中移除。逐路径清单见 `security-hardening-audit-manifest.tsv`;已移除的旧数据库适配器、未发布迁移等不凭历史清单重新引入。
- 对复核时 Git 跟踪的 2843 个 Rust、TypeScript、Vue、JavaScript、Python、SQL 和 Shell 源文件进行清单、读取与模式扫描,并检查加固差异、后续修订及关键生产调用链。正文清空、无条件禁用、脱敏、URL 拒绝、权限与数据归属等是重点搜索模式。
- 清单中的 `heuristic_added_risk_lines` 是启发式候选行数,含初始化代码和内联测试,**不是 Bug 数、漏洞数或逐行人工审计完成标记**。全仓扫描、模块复核、回归测试是不同的覆盖层次,不等于人工精读所有源文件,也不承诺零遗漏。
- 判断标准:恢复管理员明确启用的功能和必要业务数据;保留未授权访问阻断、默认保护、凭据隔离和有界资源使用。不因一条旧测试要求“所有内容都为空”,就把正常功能重新禁用。
之前的规则、OAuth、`full` 记录及视频四类修复详见 `security-hardening-compatibility-review.md`。本报告只把本次扩展检查的新发现单独计数。
## 本次新增确认并修复的两类加固回归
### 1. 正文的“压缩保留”被替换成删除
涉及 `crates/aether-data/adapters/postgres/src/usage/cleanup.rs`。
原本分为详细正文、压缩正文、请求头、计费记录等独立保留周期。加固后,详细正文清理不再压缩迁移,而是删除正文、独立 blob 与审计引用;选择条件还包含已经独立存储的正文。因此采用默认 7/30 天设置时,明确启用 `full` 后保存的正文也会在约 7 天后提前失去,而不是按压缩正文的 30 天保留。
修复内容:
- 详细正文到期只处理主表中的旧内联/压缩数据,将四个方向的正文迁入独立 gzip blob,并保留 HTTP 审计引用;不选择已经迁出的正文反复处理。
- 旧 metadata 中的正文引用按当前请求和字段校验后迁入审计表,删除旧 metadata 键,但不删除实际正文。无效、跨请求及跨字段引用不恢复。
- 每条迁移在事务中锁定并重读来源行。正文写入或引用更新失败时回滚,不先清空原值;其他 metadata 内容保留。
- 预览统计与实际清理选择条件一致。压缩正文到期仍删除;明确选择“立即清理正文”的操作仍执行删除,未把隐私清理功能改成永远保留。
真实 PostgreSQL 回归使用正式迁移后的 schema 与会话隔离的临时表,覆盖四个方向的内联 JSON、旧 gzip、独立 blob、旧引用迁移、外部请求引用拒绝、7/30 天区间、过期删除、重复清理幂等、显式即时删除和中途写入失败后的事务回滚。不会读取或修改业务数据库。
同时修复了该迁移路径原有的 SQL 错误:`usage` 表没有 `updated_at` 列,旧更新语句却写入它,真实 schema 下会导致迁移失败。这个错误在加固前已经存在,**不混算成第三类加固新增回归**。
### 2. 合法支付会话链接的 fragment 被误拒绝
涉及 `frontend/src/utils/paymentUrl.ts`,调用方包括钱包充值和订阅购买。
加固把包含 `#fragment` 的 HTTPS 支付链接一律拒绝。支付会话链接可能依赖该片段,不能将其等同于脚本协议或 URL 内嵌凭据。现在保留原有片段、查询参数及编码,同时继续拒绝非 HTTPS、相对地址、反斜杠和 URL 用户名/密码。
对应工具函数新增会话片段和编码保留测试,危险协议及凭据注入用例继续执行。服务器端支付 API 基址的同源、TLS 和无片段限制不变;没有放宽 webhook 验签、支付金额、订单归属或实际支付状态校验。此次不进行真实扣款测试。
## 按模块的复核边界
| 模块 | 复核内容与处理 |
| --- | --- |
| 公共入口、认证、会话、管理令牌 | 核对路由分类、认证缓存刷新、跨节点撤权和令牌权限目录;保留身份头防伪、用户/管理员隔离及敏感动作的权限要求。 |
| 端点配置及编辑往返 | 保留前序规则原值查看、重复/嵌套规则恢复、占位符拒绝和并发编辑保护,避免只恢复显示而破坏再次保存。 |
| 普通 Provider 连接与代理 | 对照后续 DNS/FakeIP/SOCKS 修复,不重新引入已移除的普通 Provider 地址过滤;凭据专用连接与普通业务代理分开判断。 |
| OAuth、SMTP、LDAP 与密钥 | 检查默认客户端、显式覆盖、凭据迁移与目标绑定;保留 Gemini CLI/Antigravity 兼容性恢复及既有 SMTP 修复,不撤销传输凭据隔离。 |
| 请求记录、候选诊断及保留策略 | 核对配置读取、采集、投影、单条/批量写入、读取和清理链路;修复本报告的提前删除,保留显式 `full` 及必要诊断,不把原始正文重新塞回计费主表。 |
| 视频与异步任务 | 保留前序提示词/展示名、签名产物 URL、Gemini 结果 URI、数据库领取字段修复;身份不可变条件、任务归属、加密文件和下载凭据隔离不撤销。 |
| 模型、调度、配额和池状态 | 对照模型获取、手动模型、健康摘要、候选选择、额度预留及状态流转;不回退混入该提交的订阅/计费功能。 |
| 钱包、订阅、兑换与支付 | 检查业务 URL 到前端跳转链路并修复 fragment 误拦;交易归属、金额、并发更新、幂等及回调校验维持原边界。 |
| 流式、WebSocket、格式转换 | 检查请求限制、认证头转发、取消/重定向、会话延续和正文捕获之间的关系;修正两处与恢复后的日志语义冲突的旧断言,不退回禁用功能的实现。 |
| 系统导入、导出、备份和恢复 | 区分交互式导入、恢复备份、回滚模式;保留加密备份、用途绑定、导入锁和凭据保留规则,不把安全导出误当作完整灾备。 |
| 数据契约、PostgreSQL 和迁移 | 检查脱敏投影与业务字段消费者是否矛盾;旧驱动不复活,使用真实迁移 schema 验证正文及视频的关键写入链路。 |
| 管理前端与外链 | 复核脱敏状态、权限显示、URL/导航校验、支付调用方,并执行全部前端测试与真实项目类型检查。 |
| 安装、更新、容器和发布流程 | 执行归档、链接、目录、写入目标、来源信任、compose 参数及发布供应链的隔离 Shell 测试;不运行实际部署。 |
## “全部修复”续轮完成项
以下问题已按根因处理,不再作为上一轮的未解决清单。历史类型问题、测试夹具与辅助器缺陷不统一归因为本次加固,功能恢复也不以撤销所有保护为代价。
### 1. 真实前端类型检查及相关运行时缺陷
- 修复原先 116 个文件中的 461 条类型诊断,将 `npm run type-check` 改为 `vue-tsc -b --force --pretty false`,确实检查引用的应用和工具项目,不再依赖顶层空项目的成功退出。
- 按现有接口契约补齐响应泛型、可空字段、配置 schema、图表时间轴、用户角色、事件参数及测试夹具。保留严格检查、ES2021、未知数据边界与动态配置扩展字段,未通过 `any`、忽略诊断、排除测试或降低配置绕过错误。
- 修复请求缓存和登录校验在 fetcher 同步抛错后不能正确清理 in-flight 状态的问题;失败后可按原有退避策略重试,仍隔离不同身份和旧请求。
- 修复 Provider 余额重试的迟到响应覆盖新加载结果,以及卸载后更新状态的问题;新增回归同时覆盖合法的零余额与签到失败值,不因真假值判断隐藏正常数据。
- 修复钱包模板中的刷新调用,并在用户密钥删除、路由配置保存等异步流程中固定操作目标,避免确认期间切换页面后作用到另一个对象。
### 2. 手动保留天数取了更激进的截止点
`usage_cleanup_window_with_override` 原来使用 `max` 合并截止时间,与界面“在策略内取更保守时间点”的设计相反:删除条件是记录时间早于截止点,选更晚的时间会扩大删除范围。
现改为每一保留层级分别取 `min`。详细正文、压缩正文、请求头、完整日志均不短于既有策略,也不短于本次手动指定天数。测试覆盖 0、5、30、180、400 天及“选中记录是原策略子集”的关系。显式立即清理模式不受此改动影响。
### 3. 可信原始数据库恢复可显式保留凭据
原始 JSONL 导入及数据库复制曾无条件替换密码哈希、API Key 和管理令牌,导致可信恢复/迁移也不能继续使用原凭据。新增 CLI `--preserve-credentials`,仅由操作者显式选择,不能由导入文件中的字段自行启用:
```sh
aether-gateway --database-driver postgres --postgres-url "$TARGET_DATABASE_URL" import --input /path/to/trusted.jsonl --preserve-credentials
aether-gateway copy --source-driver postgres --source-url "$SOURCE_DATABASE_URL" --target-driver postgres --target-url "$TARGET_DATABASE_URL" --preserve-credentials
```
- 不加该参数时仍默认撤销导入的身份凭据,并在操作前给出警告;原有库函数入口也保持该默认值。
- 显式保留时,只保留导入文件中用户密码、API Key 和管理令牌原有字段/状态;不会把原本停用的凭据重新启用。
- 导入的登录会话仍撤销,代理隧道仍重置在线状态和代际。外键、身份归属、OAuth 绑定一致性、输入大小及事务校验不变。
- 仅适用于可信且由操作者控制的备份/来源。目标实例须使用与来源兼容的加密配置;此开关不会解密、重新加密或重建已经丢失的值。复制链接的 TLS 默认保护不变。
- 这是原始数据库工具的选项,不改变 HTTP 配置导入权限,也不替代已有的加密备份恢复模式。操作示例会写入指定目标,执行前须自行核对目标并备份;本轮仅在临时隔离库中验证。
真实 PostgreSQL 回归分别验证默认撤销与显式保留后的密码哈希、密钥哈希、密文和启用状态;单元回归另覆盖管理令牌、会话及隧道边界。
### 4. 数据库回归夹具与测试资源回收
- 修正历史测试的超长 ID、NUMERIC/f64 解码、依赖预填充数据库的统计夹具,以及使用不在持久化契约内的 metadata 字段。使用真实 UUID、显式 SQL 类型转换、自建自清理数据和合法 trace 字段,未放宽生产 schema 或脱敏规则来迁就测试。
- 保留现有 usage 导出时间戳以秒计的兼容契约,未因旧字段名包含 `_ms` 就改变导入导出单位。
- 临时 PostgreSQL 辅助器改为 `pg_ctl stop -m fast` 并等待子进程退出,再清理自有目录;关闭失败时保留进程与目录的所有权,允许重试,不强杀后直接删数据。
- 新增两个真实 PostgreSQL 回归,覆盖打开连接下四次停止/重启、数据保留、重复停止、关闭失败重试及析构清理。专项运行前后共享内存段清单未新增残留;未清除其他业务或历史进程的 IPC 资源。
- 续轮全量回归进一步定位到迁移、回填测试各自复制的旧辅助器,同样在强杀后泄漏资源,导致后续网关测试无法初始化数据库。两份实现已合并到仅测试使用的 `postgres_test_support.rs`,采用相同的正常关闭与失败重试规则,另补两个回归;没有只清理环境而保留泄漏代码。
- 用 `AETHER_REQUIRE_LOCAL_POSTGRES_TESTS=1` 强制迁移、回填和共用辅助器的真实数据库用例执行,生命周期筛选结果为 63 通过、0 失败、1 忽略;该忽略项为已在隔离数据库中单独通过的导入测试。此次完整生命周期运行前后 IPC 清单一致。
- 辅助器支持 `AETHER_PG_CTL_BIN`,默认查找 `AETHER_POSTGRES_BIN` 同目录下的 `pg_ctl`;初始化失败也回收本次拥有的工作目录。
## 验证结果
| 验证层 | 结果 |
| --- | --- |
| Rust 工作区所有 library 与 binary 测试目标 | 最终 `cargo test --workspace --lib --bins --locked -- --test-threads=4`:65 个目标,8944 通过、0 失败、19 默认忽略;其中 41 个 library 目标 8645 通过,24 个 binary 目标 299 通过。设置 `AETHER_REQUIRE_LOCAL_POSTGRES_TESTS=1`,迁移/回填测试不能因环境问题静默跳过。 |
| 最终网关全量 | 随工作区 feature 合并执行:5124 通过,0 失败,含此前因数据库初始化失败的两项跨节点认证回归。使用仓库既有的 `RUST_MIN_STACK=16777216` 测试配置;不直接运行遗漏该配置的产物,也不改生产配置规避问题。 |
| PostgreSQL 适配器全量含真实数据库回归 | `AETHER_TEST_DATABASE_URL=<隔离库> cargo test -p aether-data-postgres --lib --locked -- --include-ignored --test-threads=1`:232 通过,0 失败,0 忽略,包含原先默认忽略的 16 项。 |
| 原始数据库导入/导出 | 设置独立的 `AETHER_TEST_POSTGRES_URL` 并执行 `lifecycle::export::tests --include-ignored`:18 通过,0 失败,0 忽略,包含迁移后数据库读取及两种凭据策略的真实往返。 |
| 临时 PostgreSQL 关闭/重启 | `aether-testkit --features postgres` 的两个真实数据库专项回归均通过;四次打开连接下重启、关闭失败重试、目录与共享内存回收均已验证。 |
| WebSocket 真实网关集成 | 最终重新执行 11 通过、0 失败,含跨连接继续对话、连续计费、长连接撤权、认证头隔离、未知字段透传、断开结算、额度重试和 PII 恢复。使用临时 PostgreSQL 与模拟上游。 |
| 可执行入口 | 299 项通过中包括网关主入口 61、隧道 185、备份恢复 CLI 10、两种 WebSocket 探针 3/4,以及压力测试种子/探针、模拟上游等入口测试。没有启动这些工具的实际生产操作。 |
| 额外入口与身份隔离集成 | `aether-data --test public_entrypoints` 3 项、`aether-gateway --test admin_unsigned_identity_headers` 1 项均重新执行通过。 |
| 前端全量与构建 | 最终 `npm run test:run`:200 个测试文件、1419 项测试通过;`npm run build:with-typecheck` 构建通过。 |
| 类型与格式 | `npm run type-check` 实际运行 `vue-tsc -b --force --pretty false`:0 错误,原先 461 条诊断均已修正;142 个改动/新增前端源码文件 ESLint 为 0 错误、0 警告。`cargo fmt --all -- --check` 和 `git diff --check` 通过。 |
| 提交前 Clippy | 本地按 `.github/workflows/rust-ci.yml` 的 Gateway、Data、其余工作区三组范围执行,均使用 `--locked` 和 `-D warnings`,全部通过;未关闭或放宽 lint 规则。 |
| 安装/发布脚本 | 10 份隔离 Shell 测试全部通过;`python3 tests/compose_database_config_test.py` 通过,只渲染配置,不启动容器。未实际安装、升级或部署。 |
工作区默认忽略的 19 项全部另行显式执行通过:PostgreSQL 适配器 16 项、真实导入 1 项、测试辅助器 2 项。默认命令中的忽略不计作通过;工作区 feature、单包测试及过滤测试的范围不同且存在重叠,上表不累加成独立用例总数。此前的规则、完整正文单条/批量持久化和视频真实数据库回归结果保留在兼容性复核报告中。
中间轮次曾因旧测试辅助器留下的 IPC 残留导致 2 项网关测试和随后 10 项 WebSocket 测试初始化失败;没有把这些命令记为成功。现已修正三处辅助器的关闭路径,并只清理本轮确认创建、无连接且创建进程已退出的 11 个段,未清理此前已有的 21 个段,也未修改内核共享内存限制。
最终工作区、WebSocket、公开入口和身份隔离测试顺序运行前后的 IPC 清单一致,均为此前已有的 21 个段,无新增残留。单独用于数据库适配器及导入往返的临时 PostgreSQL 已正常停止并删除本次自有目录。
脚本范围:`tests/deploy_state_safety_test.sh`、`tests/install_archive_safety_test.sh`、`tests/install_container_runtime_security_test.sh`、`tests/install_current_release_link_test.sh`、`tests/install_local_bundle_safety_test.sh`、`tests/install_privileged_write_safety_test.sh`、`tests/install_source_trust_test.sh`、`tests/release_supply_chain_test.sh`、`tests/tunnel_installer_config_security_test.sh`、`tests/update_compose_safety_test.sh`。
## 交付限制
本轮已确认、可复现且可在本仓库修复的剩余问题均已处理,未保留已确认却未修复的本轮源码问题;这一结论不等于“所有代码及所有部署环境绝无未知缺陷”。
这是本地源码、静态检查与隔离回归,不是对生产配置、第三方账户或所有部署组合的认证。上述验证未执行服务器部署或改动现有业务数据库;后续源码提交、推送不代表已经部署或完成生产环境验证。
已经被旧代码写空、删除、未采集的正文或视频字段不能由修复自动重建;仍在数据库中的旧内联/压缩正文可在修正后的保留链路中迁移。真实第三方 OAuth、支付、对象存储和远程隧道服务未使用生产凭据进行端到端验证。
+1 -1
View File
@@ -16,7 +16,7 @@
"test:ui": "node --experimental-require-module --disable-warning=ExperimentalWarning ./node_modules/vitest/vitest.mjs --ui", "test:ui": "node --experimental-require-module --disable-warning=ExperimentalWarning ./node_modules/vitest/vitest.mjs --ui",
"test:run": "node --experimental-require-module --disable-warning=ExperimentalWarning ./node_modules/vitest/vitest.mjs run", "test:run": "node --experimental-require-module --disable-warning=ExperimentalWarning ./node_modules/vitest/vitest.mjs run",
"lint": "eslint . --fix", "lint": "eslint . --fix",
"type-check": "vue-tsc --noEmit", "type-check": "vue-tsc -b --force --pretty false",
"version": "git describe --tags --always" "version": "git describe --tags --always"
}, },
"dependencies": { "dependencies": {
+17 -21
View File
@@ -1,5 +1,5 @@
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import type { AxiosAdapter, AxiosInstance, InternalAxiosRequestConfig } from 'axios' import type { AxiosAdapter, InternalAxiosRequestConfig } from 'axios'
import apiClient, { import apiClient, {
AUTH_SESSION_SIGNAL_KEY, AUTH_SESSION_SIGNAL_KEY,
@@ -8,10 +8,6 @@ import apiClient, {
} from '@/api/client' } from '@/api/client'
import { cache, cachedRequest } from '@/utils/cache' import { cache, cachedRequest } from '@/utils/cache'
type TestableApiClient = typeof apiClient & {
client: AxiosInstance
}
describe('apiClient auth state change event', () => { describe('apiClient auth state change event', () => {
beforeEach(() => { beforeEach(() => {
localStorage.clear() localStorage.clear()
@@ -55,10 +51,10 @@ describe('apiClient auth state change event', () => {
}) })
it('restores a session through the refresh cookie and stores the result in memory only', async () => { it('restores a session through the refresh cookie and stores the result in memory only', async () => {
const rawClient = apiClient as TestableApiClient const rawClient = apiClient['client']
const previousAdapter = rawClient.client.defaults.adapter const previousAdapter = rawClient.defaults.adapter
rawClient.client.defaults.adapter = (async (config: InternalAxiosRequestConfig) => ({ rawClient.defaults.adapter = (async (config: InternalAxiosRequestConfig) => ({
data: { access_token: 'restored-access-token' }, data: { access_token: 'restored-access-token' },
status: 200, status: 200,
statusText: 'OK', statusText: 'OK',
@@ -72,16 +68,16 @@ describe('apiClient auth state change event', () => {
expect(localStorage.getItem('access_token')).toBeNull() expect(localStorage.getItem('access_token')).toBeNull()
expect(sessionStorage.getItem('access_token')).toBeNull() expect(sessionStorage.getItem('access_token')).toBeNull()
} finally { } finally {
rawClient.client.defaults.adapter = previousAdapter rawClient.defaults.adapter = previousAdapter
} }
}) })
it('does not resurrect a session when logout wins an in-flight restore', async () => { it('does not resurrect a session when logout wins an in-flight restore', async () => {
const rawClient = apiClient as TestableApiClient const rawClient = apiClient['client']
const previousAdapter = rawClient.client.defaults.adapter const previousAdapter = rawClient.defaults.adapter
let resolveRefresh!: (response: Awaited<ReturnType<AxiosAdapter>>) => void let resolveRefresh!: (response: Awaited<ReturnType<AxiosAdapter>>) => void
rawClient.client.defaults.adapter = (() => new Promise((resolve) => { rawClient.defaults.adapter = (() => new Promise((resolve) => {
resolveRefresh = resolve resolveRefresh = resolve
})) as AxiosAdapter })) as AxiosAdapter
@@ -100,7 +96,7 @@ describe('apiClient auth state change event', () => {
await expect(restore).rejects.toThrow('Auth state changed') await expect(restore).rejects.toThrow('Auth state changed')
expect(apiClient.getToken()).toBeNull() expect(apiClient.getToken()).toBeNull()
} finally { } finally {
rawClient.client.defaults.adapter = previousAdapter rawClient.defaults.adapter = previousAdapter
} }
}) })
@@ -141,11 +137,11 @@ describe('apiClient auth state change event', () => {
}) })
it('sends auth refresh without a request body', async () => { it('sends auth refresh without a request body', async () => {
const rawClient = apiClient as TestableApiClient const rawClient = apiClient['client']
const previousAdapter = rawClient.client.defaults.adapter const previousAdapter = rawClient.defaults.adapter
const requests: InternalAxiosRequestConfig[] = [] const requests: InternalAxiosRequestConfig[] = []
rawClient.client.defaults.adapter = (async (config: InternalAxiosRequestConfig) => { rawClient.defaults.adapter = (async (config: InternalAxiosRequestConfig) => {
requests.push(config) requests.push(config)
return { return {
data: { access_token: 'new-access-token' }, data: { access_token: 'new-access-token' },
@@ -165,16 +161,16 @@ describe('apiClient auth state change event', () => {
expect(requests[0].method).toBe('post') expect(requests[0].method).toBe('post')
expect(requests[0].data).toBeUndefined() expect(requests[0].data).toBeUndefined()
} finally { } finally {
rawClient.client.defaults.adapter = previousAdapter rawClient.defaults.adapter = previousAdapter
} }
}) })
it('authenticates protected gateway operational requests', async () => { it('authenticates protected gateway operational requests', async () => {
const rawClient = apiClient as TestableApiClient const rawClient = apiClient['client']
const previousAdapter = rawClient.client.defaults.adapter const previousAdapter = rawClient.defaults.adapter
const requests: InternalAxiosRequestConfig[] = [] const requests: InternalAxiosRequestConfig[] = []
rawClient.client.defaults.adapter = (async (config: InternalAxiosRequestConfig) => { rawClient.defaults.adapter = (async (config: InternalAxiosRequestConfig) => {
requests.push(config) requests.push(config)
return { return {
data: '', data: '',
@@ -193,7 +189,7 @@ describe('apiClient auth state change event', () => {
expect(requests[0].headers.Authorization).toBe('Bearer operational-access-token') expect(requests[0].headers.Authorization).toBe('Bearer operational-access-token')
expect(requests[0].headers['X-Client-Device-Id']).toBeTruthy() expect(requests[0].headers['X-Client-Device-Id']).toBeTruthy()
} finally { } finally {
rawClient.client.defaults.adapter = previousAdapter rawClient.defaults.adapter = previousAdapter
} }
}) })
}) })
@@ -0,0 +1,26 @@
import { beforeEach, describe, expect, it, vi } from 'vitest'
import { revealEndpointRules } from '../endpoints/endpoints'
const client = vi.hoisted(() => ({ get: vi.fn() }))
vi.mock('../client', () => ({ default: client }))
beforeEach(() => {
client.get.mockReset().mockResolvedValue({ data: { header_rules: [], body_rules: [], response_header_rules: [] } })
})
describe('endpoint rule reveal API', () => {
it('uses the scoped route and forwards cancellation', async () => {
const controller = new AbortController()
await revealEndpointRules('endpoint/with?reserved', controller.signal)
expect(client.get).toHaveBeenCalledWith(
'/api/admin/endpoints/endpoint%2Fwith%3Freserved/rules/reveal',
{ signal: controller.signal },
)
})
it('fetches each reveal without retaining a cached response', async () => {
await revealEndpointRules('endpoint-1')
await revealEndpointRules('endpoint-1')
expect(client.get).toHaveBeenCalledTimes(2)
})
})
+27 -5
View File
@@ -215,7 +215,18 @@ export const adminWalletApi = {
credited_at: string | null credited_at: string | null
} }
}> { }> {
const response = await apiClient.post(`/api/admin/wallets/${walletId}/recharge`, payload) const response = await apiClient.post<{
wallet: AdminWallet
payment_order: {
id: string
order_no: string
amount_usd: number
payment_method: string
status: string
created_at: string
credited_at: string | null
}
}>(`/api/admin/wallets/${walletId}/recharge`, payload)
return response.data return response.data
}, },
@@ -223,7 +234,10 @@ export const adminWalletApi = {
wallet: AdminWallet wallet: AdminWallet
transaction: WalletTransaction transaction: WalletTransaction
}> { }> {
const response = await apiClient.post(`/api/admin/wallets/${walletId}/adjust`, payload) const response = await apiClient.post<{
wallet: AdminWallet
transaction: WalletTransaction
}>(`/api/admin/wallets/${walletId}/adjust`, payload)
return response.data return response.data
}, },
@@ -232,7 +246,11 @@ export const adminWalletApi = {
refund: RefundRequest refund: RefundRequest
transaction: WalletTransaction transaction: WalletTransaction
}> { }> {
const response = await apiClient.post( const response = await apiClient.post<{
wallet: AdminWallet
refund: RefundRequest
transaction: WalletTransaction
}>(
`/api/admin/wallets/${walletId}/refunds/${refundId}/process`, `/api/admin/wallets/${walletId}/refunds/${refundId}/process`,
{} {}
) )
@@ -244,7 +262,11 @@ export const adminWalletApi = {
refund: RefundRequest refund: RefundRequest
transaction: WalletTransaction | null transaction: WalletTransaction | null
}> { }> {
const response = await apiClient.post( const response = await apiClient.post<{
wallet: AdminWallet
refund: RefundRequest
transaction: WalletTransaction | null
}>(
`/api/admin/wallets/${walletId}/refunds/${refundId}/fail`, `/api/admin/wallets/${walletId}/refunds/${refundId}/fail`,
payload payload
) )
@@ -256,7 +278,7 @@ export const adminWalletApi = {
refundId: string, refundId: string,
payload: RefundCompleteRequest payload: RefundCompleteRequest
): Promise<{ refund: RefundRequest }> { ): Promise<{ refund: RefundRequest }> {
const response = await apiClient.post( const response = await apiClient.post<{ refund: RefundRequest }>(
`/api/admin/wallets/${walletId}/refunds/${refundId}/complete`, `/api/admin/wallets/${walletId}/refunds/${refundId}/complete`,
payload payload
) )
+7 -2
View File
@@ -8,6 +8,11 @@ import type { ApiKeyInstallSession, InstallSessionTargetSystem, InstallTargetCli
const SYSTEM_DATA_IMPORT_TIMEOUT_MS = 10 * 60 * 1000 const SYSTEM_DATA_IMPORT_TIMEOUT_MS = 10 * 60 * 1000
const ALL_SYSTEM_CONFIGS_CACHE_KEY = 'admin:system:configs' const ALL_SYSTEM_CONFIGS_CACHE_KEY = 'admin:system:configs'
export interface AdminTimeSeriesPoint extends Record<string, unknown> {
date: string
total_cost: number
}
export interface AdminSystemConfigItem { export interface AdminSystemConfigItem {
key: string key: string
value: unknown value: unknown
@@ -1594,12 +1599,12 @@ export const adminApi = {
provider_name?: string provider_name?: string
}, },
options?: AdminAnalyticsRequestOptions options?: AdminAnalyticsRequestOptions
): Promise<Array<Record<string, unknown>>> { ): Promise<AdminTimeSeriesPoint[]> {
const cacheKey = buildCacheKey('admin:stats:time-series', params) const cacheKey = buildCacheKey('admin:stats:time-series', params)
return cachedRequest( return cachedRequest(
cacheKey, cacheKey,
async () => { async () => {
const response = await apiClient.get<Array<Record<string, unknown>>>('/api/admin/stats/time-series', { params }) const response = await apiClient.get<AdminTimeSeriesPoint[]>('/api/admin/stats/time-series', { params })
return response.data return response.data
}, },
options?.skipCache ? 0 : 20 * 1000 options?.skipCache ? 0 : 20 * 1000

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