Compare commits

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

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

Document DNS policy boundaries and verify 809 gateway, tunnel, and HTTP regression tests.
2026-09-08 17:44:59 +08:00
114 changed files with 8082 additions and 1390 deletions
Generated
+1
View File
@@ -305,6 +305,7 @@ dependencies = [
"futures-util",
"hmac",
"http",
"http-body",
"http-body-util",
"hyper",
"hyper-util",
+1
View File
@@ -62,6 +62,7 @@ flate2.workspace = true
futures-util.workspace = true
hmac.workspace = true
http.workspace = true
http-body = "1"
http-body-util = "0.1"
hyper = { version = "1", features = ["client", "server", "http1", "http2"] }
hyper-util = { version = "0.1", features = ["client-legacy", "client-pool", "server-auto", "service", "tokio"] }
@@ -194,7 +194,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncA
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -246,7 +246,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt>
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -118,7 +118,7 @@ pub(crate) fn build_local_execution_report_context(
insert_pool_key_lease_report_context_fields(&mut extra_fields, parts.pool_key_lease);
insert_scheduler_affinity_policy_report_context_field(&mut extra_fields, parts.routing_policy);
if let Some(policy) = parts.routing_policy {
if let Ok(value) = serde_json::to_value(policy.execution_policy) {
if let Ok(value) = serde_json::to_value(&policy.execution_policy) {
extra_fields.insert(ROUTING_EXECUTION_POLICY_REPORT_FIELD.to_string(), value);
}
}
@@ -179,7 +179,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -224,7 +224,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -257,7 +257,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -302,7 +302,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -109,7 +109,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptS
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -182,7 +182,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -232,7 +232,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttempt
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -124,7 +124,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttem
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
@@ -97,7 +97,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSo
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -729,117 +729,6 @@ fn update_normalization_codex_capabilities_digest(
update_normalization_string_vec_digest(digest, &capabilities.supported_service_tiers);
}
#[cfg(test)]
mod continuation_fingerprint_tests {
use http::HeaderValue;
use serde_json::json;
use super::ResponsesWebSocketBodyNormalization;
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
#[test]
fn normalization_fingerprint_is_stable_for_json_object_key_order() {
let first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}, "x": 1}));
let second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"x": 1, "reasoning": {"effort": "high"}}));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_changes_with_effective_contract() {
let base = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let changed_policy = base.clone().with_reasoning_replay_policy_for_tests(
OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
assert_ne!(
base.continuation_fingerprint(),
changed_policy.continuation_fingerprint()
);
let changed_patch = base
.clone()
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "low"}}));
assert_ne!(
base.continuation_fingerprint(),
changed_patch.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_ignores_unrelated_volatile_request_headers() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "x-contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
first
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-1"));
first
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-1"));
let mut second = first.clone();
second
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-2"));
second
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-2"));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"headers that no body-rule condition reads must not invalidate a persisted continuation"
);
}
#[test]
fn normalization_fingerprint_tracks_headers_used_by_body_rule_conditions() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "X-Contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules.clone());
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
let mut second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
second
.request_headers
.insert("x-contract", HeaderValue::from_static("disabled"));
assert_ne!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"a header that controls an effective body-rule condition remains part of the contract"
);
}
}
/// Builds one upstream decision for a Responses WebSocket turn. The session
/// reuses this decision for same-model turns and invokes the planner again when
/// a later `response.create` changes the public model.
@@ -1058,3 +947,114 @@ async fn release_responses_websocket_planning_lease(
}
}
}
#[cfg(test)]
mod continuation_fingerprint_tests {
use http::HeaderValue;
use serde_json::json;
use super::ResponsesWebSocketBodyNormalization;
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
#[test]
fn normalization_fingerprint_is_stable_for_json_object_key_order() {
let first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}, "x": 1}));
let second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"x": 1, "reasoning": {"effort": "high"}}));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_changes_with_effective_contract() {
let base = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let changed_policy = base.clone().with_reasoning_replay_policy_for_tests(
OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
assert_ne!(
base.continuation_fingerprint(),
changed_policy.continuation_fingerprint()
);
let changed_patch = base
.clone()
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "low"}}));
assert_ne!(
base.continuation_fingerprint(),
changed_patch.continuation_fingerprint()
);
}
#[test]
fn normalization_fingerprint_ignores_unrelated_volatile_request_headers() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "x-contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
first
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-1"));
first
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-1"));
let mut second = first.clone();
second
.request_headers
.insert("x-request-id", HeaderValue::from_static("request-2"));
second
.request_headers
.insert("cf-ray", HeaderValue::from_static("edge-2"));
assert_eq!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"headers that no body-rule condition reads must not invalidate a persisted continuation"
);
}
#[test]
fn normalization_fingerprint_tracks_headers_used_by_body_rule_conditions() {
let body_rules = json!([{
"action": "set",
"path": "store",
"value": false,
"condition": {
"source": "request_headers",
"path": "X-Contract",
"op": "eq",
"value": "enabled"
}
}]);
let mut first = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules.clone());
first
.request_headers
.insert("x-contract", HeaderValue::from_static("enabled"));
let mut second = ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_body_rules_for_tests(body_rules);
second
.request_headers
.insert("x-contract", HeaderValue::from_static("disabled"));
assert_ne!(
first.continuation_fingerprint(),
second.continuation_fingerprint(),
"a header that controls an effective body-rule condition remains part of the contract"
);
}
}
@@ -166,7 +166,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
@@ -216,7 +216,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStream
self.input
.routing_policy
.as_ref()
.map(|policy| policy.execution_policy)
.map(|policy| policy.execution_policy.clone())
}
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
+14 -14
View File
@@ -23,7 +23,6 @@ const MAX_BARK_TEMPLATE_BYTES: usize = 256 * 1024;
const MAX_BARK_TITLE_BYTES: usize = 512;
const MAX_BARK_BODY_BYTES: usize = 2 * 1024 * 1024;
const MAX_BARK_RENDERED_BODY_BYTES: usize = 2 * 1024 * 1024;
const MAX_BARK_RESOLVED_ADDRESSES: usize = 32;
#[derive(Clone)]
pub(crate) struct BarkPushConfig {
@@ -208,19 +207,20 @@ async fn build_bark_push_client_and_url(
let port = push_url
.port_or_known_default()
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少端口".to_string()))?;
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
tokio::time::timeout(
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
tokio::net::lookup_host((host.as_str(), port)),
)
.await
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析超时".to_string()))?
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析失败".to_string()))?
.take(MAX_BARK_RESOLVED_ADDRESSES)
.collect::<Vec<_>>()
};
let addresses = aether_http::lookup_host_with_limits(
host.as_str(),
port,
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
)
.await
.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "Bark 服务器 DNS 解析超时",
std::io::ErrorKind::InvalidData => "Bark 服务器 DNS 解析返回过多地址",
_ => "Bark 服务器 DNS 解析失败",
};
GatewayError::Internal(message.to_string())
})?;
let allow_benchmarking_ip = push_url.scheme() == "https"
&& push_url.port_or_known_default() == Some(443)
&& host.eq_ignore_ascii_case("api.day.app");
+194 -43
View File
@@ -136,14 +136,16 @@ pub(crate) async fn send_smtp_email(
email: ComposedEmail,
) -> Result<(), GatewayError> {
validate_smtp_delivery_inputs(&config, &email)?;
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email))
let stream = connect_tcp_stream(&config).await?;
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email, stream))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
}
pub(crate) async fn probe_smtp_connection(config: SmtpDeliveryConfig) -> Result<(), GatewayError> {
validate_smtp_config(&config)?;
tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config))
let stream = connect_tcp_stream(&config).await?;
tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config, stream))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
}
@@ -328,43 +330,58 @@ fn resolve_server_name(host: &str) -> Result<rustls::pki_types::ServerName<'stat
.map_err(|err| GatewayError::Internal(err.to_string()))
}
fn connect_tcp_stream(config: &SmtpDeliveryConfig) -> Result<std::net::TcpStream, GatewayError> {
use std::net::ToSocketAddrs;
let addresses = (config.host.as_str(), config.port)
.to_socket_addrs()
.map_err(|err| GatewayError::Internal(err.to_string()))?
.take(16)
.collect::<Vec<_>>();
if addresses.is_empty() {
return Err(GatewayError::Internal(
"smtp host did not resolve to an address".to_string(),
));
}
let deadline = std::time::Instant::now()
.checked_add(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))
.unwrap_or_else(std::time::Instant::now);
let mut last_error = None;
let mut stream = None;
for address in addresses {
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
if remaining.is_zero() {
break;
async fn connect_tcp_stream(
config: &SmtpDeliveryConfig,
) -> Result<std::net::TcpStream, GatewayError> {
connect_tcp_stream_with_dns(
aether_http::lookup_host_with_limits(
&config.host,
config.port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
),
std::time::Duration::from_secs(SMTP_TIMEOUT_SECS),
)
.await
}
async fn connect_tcp_stream_with_dns(
lookup: impl std::future::Future<Output = std::io::Result<Vec<std::net::SocketAddr>>>,
timeout: std::time::Duration,
) -> Result<std::net::TcpStream, GatewayError> {
let stream = tokio::time::timeout(timeout, async {
let addresses = lookup.await.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "smtp DNS resolution timed out",
std::io::ErrorKind::InvalidData => {
"smtp DNS resolution returned too many addresses"
}
_ => "smtp DNS resolution failed",
};
GatewayError::Internal(message.to_string())
})?;
if addresses.is_empty() {
return Err(GatewayError::Internal(
"smtp host did not resolve to an address".to_string(),
));
}
match std::net::TcpStream::connect_timeout(&address, remaining) {
Ok(candidate) => {
stream = Some(candidate);
break;
}
Err(err) => last_error = Some(err),
}
}
let stream = stream.ok_or_else(|| {
GatewayError::Internal(
last_error
.map(|err| err.to_string())
.unwrap_or_else(|| "smtp connection timed out".to_string()),
)
})?;
let attempts = addresses
.into_iter()
.map(|address| Box::pin(tokio::net::TcpStream::connect(address)));
futures_util::future::select_ok(attempts)
.await
.map(|(stream, _)| stream)
.map_err(|error| {
GatewayError::Internal(format!("smtp connection failed ({})", error.kind()))
})
})
.await
.map_err(|_| GatewayError::Internal("smtp DNS or TCP connection timed out".to_string()))??;
let stream = stream
.into_std()
.map_err(|err| GatewayError::Internal(err.to_string()))?;
stream
.set_nonblocking(false)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
stream
.set_read_timeout(Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS)))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
@@ -680,16 +697,15 @@ fn smtp_probe_connection<S: std::io::Read + std::io::Write>(
fn send_smtp_email_blocking(
config: SmtpDeliveryConfig,
email: ComposedEmail,
stream: std::net::TcpStream,
) -> Result<(), GatewayError> {
if config.use_ssl {
let stream = connect_tcp_stream(&config)?;
let tls_stream = wrap_tls_stream(stream, &config.host)?;
let mut reader = std::io::BufReader::new(tls_stream);
let _ = smtp_expect(&mut reader, &[220])?;
return smtp_send_message(&mut reader, &config, &email);
}
let stream = connect_tcp_stream(&config)?;
let mut reader = std::io::BufReader::new(stream);
let _ = smtp_expect(&mut reader, &[220])?;
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
@@ -705,16 +721,17 @@ fn send_smtp_email_blocking(
smtp_deliver_message(&mut reader, &config, &email)
}
fn probe_smtp_connection_blocking(config: SmtpDeliveryConfig) -> Result<(), GatewayError> {
fn probe_smtp_connection_blocking(
config: SmtpDeliveryConfig,
stream: std::net::TcpStream,
) -> Result<(), GatewayError> {
if config.use_ssl {
let stream = connect_tcp_stream(&config)?;
let tls_stream = wrap_tls_stream(stream, &config.host)?;
let mut reader = std::io::BufReader::new(tls_stream);
let _ = smtp_expect(&mut reader, &[220])?;
return smtp_probe_connection(&mut reader, &config);
}
let stream = connect_tcp_stream(&config)?;
let mut reader = std::io::BufReader::new(stream);
let _ = smtp_expect(&mut reader, &[220])?;
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
@@ -778,6 +795,140 @@ mod tests {
assert!(validate_smtp_delivery_inputs(&config(), &email()).is_ok());
}
#[tokio::test]
async fn smtp_connection_deadline_includes_a_stalled_dns_lookup() {
let error = connect_tcp_stream_with_dns(
std::future::pending(),
std::time::Duration::from_millis(5),
)
.await
.expect_err("DNS must not outlive the connection deadline");
assert!(format!("{error:?}").contains("smtp DNS or TCP connection timed out"));
}
#[tokio::test]
async fn smtp_dns_errors_and_empty_answers_fail_without_connecting() {
for (addresses, expected) in [
(Ok(Vec::new()), "smtp host did not resolve to an address"),
(
Err(std::io::Error::other("sensitive-dns-detail")),
"smtp DNS resolution failed",
),
(
Err(std::io::Error::from(std::io::ErrorKind::InvalidData)),
"smtp DNS resolution returned too many addresses",
),
(
Err(std::io::Error::from(std::io::ErrorKind::TimedOut)),
"smtp DNS resolution timed out",
),
] {
let error = connect_tcp_stream_with_dns(
std::future::ready(addresses),
std::time::Duration::from_secs(1),
)
.await
.expect_err("invalid DNS answers must fail before TCP connect");
assert!(format!("{error:?}").contains(expected));
assert!(!format!("{error:?}").contains("sensitive-dns-detail"));
}
}
#[tokio::test]
async fn smtp_connection_tries_answers_beyond_the_old_sixteen_address_limit() {
let unavailable = tokio::net::TcpSocket::new_v4().unwrap();
unavailable.bind("127.0.0.1:0".parse().unwrap()).unwrap();
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let available = listener.local_addr().unwrap();
let mut addresses = vec![unavailable.local_addr().unwrap(); 16];
addresses.push(available);
let stream = connect_tcp_stream_with_dns(
std::future::ready(Ok(addresses)),
std::time::Duration::from_secs(5),
)
.await
.expect("later DNS answers should remain available for fallback");
assert_eq!(stream.peer_addr().unwrap(), available);
assert_eq!(
stream.read_timeout().unwrap(),
Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))
);
}
#[tokio::test]
async fn smtp_probe_and_delivery_use_the_preconnected_stream() {
use tokio::io::{AsyncBufReadExt, AsyncWriteExt};
for deliver in [false, true] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let port = listener.local_addr().unwrap().port();
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.unwrap();
let mut reader = tokio::io::BufReader::new(stream);
reader
.get_mut()
.write_all(b"220 mock SMTP ready\r\n")
.await
.unwrap();
let mut delivered = false;
loop {
let mut line = String::new();
assert!(reader.read_line(&mut line).await.unwrap() > 0);
let response = if line.starts_with("EHLO ")
|| line.starts_with("MAIL FROM:")
|| line.starts_with("RCPT TO:")
{
&b"250 OK\r\n"[..]
} else if line == "DATA\r\n" {
reader
.get_mut()
.write_all(b"354 End with dot\r\n")
.await
.unwrap();
loop {
line.clear();
assert!(reader.read_line(&mut line).await.unwrap() > 0);
if line == ".\r\n" {
break;
}
}
delivered = true;
&b"250 Accepted\r\n"[..]
} else {
assert_eq!(line, "QUIT\r\n");
reader
.get_mut()
.write_all(b"221 Goodbye\r\n")
.await
.unwrap();
break;
};
reader.get_mut().write_all(response).await.unwrap();
}
assert_eq!(delivered, deliver);
});
let config = SmtpDeliveryConfig {
host: "127.0.0.1".to_string(),
port,
user: None,
password: None,
use_tls: false,
use_ssl: false,
..config()
};
tokio::time::timeout(std::time::Duration::from_secs(5), async {
if deliver {
send_smtp_email(config, email()).await.unwrap();
} else {
probe_smtp_connection(config).await.unwrap();
}
server.await.unwrap();
})
.await
.expect("local SMTP probe and delivery should complete");
}
}
#[test]
fn rejects_authentication_over_plaintext_smtp() {
let mut insecure = config();
@@ -68,7 +68,6 @@ const CHATGPT_WEB_IMAGE_PUBLIC_CONNECT_TIMEOUT_MS: u64 = 10_000;
const CHATGPT_WEB_IMAGE_PUBLIC_READ_TIMEOUT_MS: u64 = 30_000;
const CHATGPT_WEB_IMAGE_PUBLIC_TOTAL_TIMEOUT_MS: u64 = 300_000;
const CHATGPT_WEB_OPAQUE_ID_MAX_BYTES: usize = 256;
const CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES: usize = 32;
const CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES: usize = 64 * 1024;
const CHATGPT_WEB_IMAGE_UPLOAD_RESPONSE_LIMIT_BYTES: usize = 64 * 1024;
const CHATGPT_WEB_IMAGE_MAX_PROMPT_BYTES: usize = 32 * 1024;
@@ -1334,24 +1333,18 @@ async fn resolve_public_web_image_addrs(
"ChatGPT-Web image URL is missing a port".to_string(),
)
})?;
let resolved = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
tokio::time::timeout(lookup_timeout, tokio::net::lookup_host((host, port)))
.await
.map_err(|_| {
ExecutionRuntimeTransportError::UpstreamRequest(
"ChatGPT-Web image URL DNS resolution timed out".to_string(),
)
})?
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format!(
"ChatGPT-Web image URL DNS resolution failed: {err}"
))
})?
.take(CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES)
.collect::<Vec<_>>()
};
let resolved = aether_http::lookup_host_with_limits(host, port, lookup_timeout)
.await
.map_err(|error| {
let message = match error.kind() {
std::io::ErrorKind::TimedOut => "ChatGPT-Web image URL DNS resolution timed out",
std::io::ErrorKind::InvalidData => {
"ChatGPT-Web image URL DNS resolution returned too many addresses"
}
_ => "ChatGPT-Web image URL DNS resolution failed",
};
ExecutionRuntimeTransportError::UpstreamRequest(message.to_string())
})?;
if resolved.is_empty() {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"ChatGPT-Web image URL DNS resolution returned no addresses".to_string(),
@@ -14,7 +14,7 @@ fn sync_plan_kind_disables_local_candidate_failover(plan_kind: &str) -> bool {
)
}
fn openai_image_success_disables_local_success_failover(
pub(super) fn openai_image_success_disables_local_success_failover(
plan: &ExecutionPlan,
status_code: u16,
) -> bool {
@@ -1036,6 +1036,7 @@ mod tests {
policy,
LocalFailoverPolicy {
max_retries: Some(1),
routing_rules: Default::default(),
max_transfer_count: 0,
max_transfer_timeout_seconds: 0,
stop_status_codes: [503].into_iter().collect(),
@@ -1826,12 +1826,12 @@ async fn fetch_grok_attachment_url(
// a fragment from the previous URL, while an absolute Location can
// introduce either explicitly.
validate_grok_attachment_url(&url)?;
let public_addr = public_socket_addr_for_url(&url).await?;
let public_addrs = public_socket_addrs_for_url(&url).await?;
let response = reqwest::Client::builder()
.no_proxy()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.resolve_to_addrs(url.host_str().unwrap_or_default(), &[public_addr])
.resolve_to_addrs(url.host_str().unwrap_or_default(), &public_addrs)
.build()
.map_err(ExecutionRuntimeTransportError::ClientBuild)?
.get(url.clone())
@@ -1897,10 +1897,10 @@ fn validate_grok_attachment_url(url: &reqwest::Url) -> Result<(), ExecutionRunti
Ok(())
}
async fn public_socket_addr_for_url(
async fn public_socket_addrs_for_url(
url: &reqwest::Url,
) -> Result<std::net::SocketAddr, ExecutionRuntimeTransportError> {
let host = url.host().ok_or_else(|| {
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
let host = url.host_str().ok_or_else(|| {
ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL is missing a host".to_string(),
)
@@ -1910,64 +1910,34 @@ async fn public_socket_addr_for_url(
"Grok attachment URL is missing a port".to_string(),
)
})?;
let host = match host {
url::Host::Ipv4(ip) => {
let ip = IpAddr::V4(ip);
if !grok_attachment_ip_is_public(ip) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
return Ok(std::net::SocketAddr::new(ip, port));
}
url::Host::Ipv6(ip) => {
let ip = IpAddr::V6(ip);
if !grok_attachment_ip_is_public(ip) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
return Ok(std::net::SocketAddr::new(ip, port));
}
url::Host::Domain(host) => host,
};
if let Ok(ip) = host.parse::<IpAddr>() {
if !grok_attachment_ip_is_public(ip) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
return Ok(std::net::SocketAddr::new(ip, port));
}
let mut public_addr = None;
let mut resolved_any = false;
for addr in
let addresses =
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format!(
"Grok attachment URL DNS resolution failed: {err}"
))
})?
{
resolved_any = true;
if !grok_attachment_ip_is_public(addr.ip()) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
public_addr.get_or_insert(addr);
}
if !resolved_any {
})?;
validate_grok_attachment_addresses(addresses)
}
fn validate_grok_attachment_addresses(
addresses: Vec<std::net::SocketAddr>,
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
if addresses.is_empty() {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL DNS resolution returned no addresses".to_string(),
));
}
public_addr.ok_or_else(|| {
ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL has no public address".to_string(),
)
})
if addresses
.iter()
.any(|address| !grok_attachment_ip_is_public(address.ip()))
{
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
"Grok attachment URL resolves to a non-public address".to_string(),
));
}
Ok(addresses)
}
fn grok_attachment_ip_is_public(ip: IpAddr) -> bool {
@@ -3898,7 +3868,7 @@ mod tests {
grok_should_use_imagine_websocket, grok_success_frame_stream, grok_upload_url,
grok_upstream_model_name, grok_usage_estimate, grok_user_id_from_cookie_header,
materialize_grok_image_assets, maximum_base64_len_for_decoded_limit, openai_chat_body,
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addr_for_url,
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addrs_for_url,
set_grok_image_edit_config, validate_grok_attachment_url, GrokAttachmentInput,
GrokCollected, GrokImagineImage, GrokStreamAdapter,
};
@@ -4520,7 +4490,7 @@ mod tests {
] {
let url = reqwest::Url::parse(raw_url).expect("URL should parse");
assert!(
public_socket_addr_for_url(&url).await.is_err(),
public_socket_addrs_for_url(&url).await.is_err(),
"private IPv6 literal should be rejected: {raw_url}"
);
}
@@ -4528,13 +4498,31 @@ mod tests {
let url = reqwest::Url::parse("https://[2606:4700:4700::1111]/attachment")
.expect("URL should parse");
assert_eq!(
public_socket_addr_for_url(&url)
public_socket_addrs_for_url(&url)
.await
.expect("public IPv6 literal should pass"),
"[2606:4700:4700::1111]:443".parse().unwrap()
vec!["[2606:4700:4700::1111]:443".parse().unwrap()]
);
}
#[test]
fn grok_attachment_dns_keeps_all_safe_addresses_for_connection_fallback() {
let addresses = vec![
"[2606:4700:4700::1111]:443".parse().unwrap(),
"8.8.8.8:443".parse().unwrap(),
];
assert_eq!(
super::validate_grok_attachment_addresses(addresses.clone()).unwrap(),
addresses
);
assert!(super::validate_grok_attachment_addresses(Vec::new()).is_err());
for blocked in ["198.18.0.1:443", "127.0.0.1:443", "[fd00::1]:443"] {
let mut mixed = addresses.clone();
mixed.push(blocked.parse().unwrap());
assert!(super::validate_grok_attachment_addresses(mixed).is_err());
}
}
#[test]
fn grok_attachment_url_rejects_credentials_and_fragments_on_every_hop() {
for raw_url in [
@@ -11,6 +11,10 @@ const GEMINI_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750);
pub(super) enum StreamCommitPolicy {
ResponseHeaders,
FirstClassifiedBody,
FirstSseSemanticEvent {
max_bytes: usize,
max_wait: Duration,
},
FirstAnthropicSemanticEvent {
max_bytes: usize,
max_wait: Duration,
@@ -36,16 +40,21 @@ impl StreamCommitPolicy {
return Self::FirstClassifiedBody;
}
if force_prefetch {
return Self::FirstClassifiedBody;
}
let content_type = content_type
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or_default()
.to_ascii_lowercase();
if content_type.contains("text/event-stream") {
if provider_api_format.eq_ignore_ascii_case("openai:image")
|| client_api_format.eq_ignore_ascii_case("openai:image")
{
return if force_prefetch {
Self::FirstClassifiedBody
} else {
Self::ResponseHeaders
};
}
if provider_api_format.eq_ignore_ascii_case("claude:messages")
&& provider_api_format.eq_ignore_ascii_case(client_api_format)
&& !has_private_stream_normalizer
@@ -62,7 +71,14 @@ impl StreamCommitPolicy {
max_wait: GEMINI_PRECOMMIT_MAX_WAIT,
};
}
return Self::ResponseHeaders;
return Self::FirstSseSemanticEvent {
max_bytes: MAX_STREAM_PREFETCH_BYTES,
max_wait: Duration::from_secs(30),
};
}
if force_prefetch {
return Self::FirstClassifiedBody;
}
if has_private_stream_normalizer || has_local_stream_rewriter {
@@ -91,14 +107,17 @@ impl StreamCommitPolicy {
pub(super) const fn requires_bounded_frame_wait(self) -> bool {
matches!(
self,
Self::FirstAnthropicSemanticEvent { .. } | Self::FirstGeminiSemanticEvent { .. }
Self::FirstAnthropicSemanticEvent { .. }
| Self::FirstGeminiSemanticEvent { .. }
| Self::FirstSseSemanticEvent { .. }
)
}
pub(super) const fn max_precommit_wait(self) -> Option<Duration> {
match self {
Self::FirstAnthropicSemanticEvent { max_wait, .. }
| Self::FirstGeminiSemanticEvent { max_wait, .. } => Some(max_wait),
| Self::FirstGeminiSemanticEvent { max_wait, .. }
| Self::FirstSseSemanticEvent { max_wait, .. } => Some(max_wait),
Self::ResponseHeaders | Self::FirstClassifiedBody => None,
}
}
@@ -110,6 +129,16 @@ impl StreamCommitPolicy {
pub(super) const fn is_gemini(self) -> bool {
matches!(self, Self::FirstGeminiSemanticEvent { .. })
}
pub(super) fn with_precommit_wait(mut self, wait: Duration) -> Self {
match &mut self {
Self::FirstAnthropicSemanticEvent { max_wait, .. }
| Self::FirstGeminiSemanticEvent { max_wait, .. }
| Self::FirstSseSemanticEvent { max_wait, .. } => *max_wait = wait,
_ => {}
}
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -133,6 +162,7 @@ pub(super) struct StreamCommitGate {
observed_bytes: usize,
anthropic: AnthropicSsePrecommitInspector,
gemini: GeminiSsePrecommitInspector,
generic: GenericSsePrecommitInspector,
}
impl StreamCommitGate {
@@ -148,6 +178,7 @@ impl StreamCommitGate {
observed_bytes: 0,
anthropic: AnthropicSsePrecommitInspector::default(),
gemini: GeminiSsePrecommitInspector::default(),
generic: GenericSsePrecommitInspector::default(),
}
}
@@ -171,6 +202,9 @@ impl StreamCommitGate {
StreamCommitPolicy::FirstGeminiSemanticEvent { max_bytes, .. } => {
(max_bytes, self.gemini.observe(chunk, max_bytes))
}
StreamCommitPolicy::FirstSseSemanticEvent { max_bytes, .. } => {
(max_bytes, self.generic.observe(chunk, max_bytes))
}
StreamCommitPolicy::ResponseHeaders | StreamCommitPolicy::FirstClassifiedBody => {
return StreamPrecommitObservation::Pending;
}
@@ -217,6 +251,152 @@ enum SemanticSseObservation {
Error { status_code: u16, body_json: Value },
}
#[derive(Debug, Default)]
struct GenericSsePrecommitInspector {
buffered: Vec<u8>,
}
impl GenericSsePrecommitInspector {
fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation {
let remaining = max_bytes.saturating_sub(self.buffered.len());
self.buffered
.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
while let Some((record_end, separator_len)) = find_sse_record_boundary(&self.buffered) {
let record = self.buffered[..record_end].to_vec();
self.buffered.drain(..record_end + separator_len);
match classify_generic_sse_record(&record) {
SemanticSseObservation::Pending => {}
observation => return observation,
}
}
if chunk.len() > remaining {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
}
}
}
fn classify_generic_sse_record(record: &[u8]) -> SemanticSseObservation {
let Ok(record) = std::str::from_utf8(record) else {
return SemanticSseObservation::SemanticEvent;
};
let normalized = record.replace("\r\n", "\n").replace('\r', "\n");
let event_type = normalized
.lines()
.find_map(|line| line.strip_prefix("event:").map(str::trim));
let data = normalized
.lines()
.filter_map(|line| line.strip_prefix("data:").map(str::trim_start))
.collect::<Vec<_>>()
.join("\n");
if data.trim().is_empty() || matches!(event_type, Some("ping" | "heartbeat" | "keepalive")) {
return SemanticSseObservation::Pending;
}
if data.trim() == "[DONE]" {
return SemanticSseObservation::SemanticEvent;
}
let Ok(body_json) = serde_json::from_str::<Value>(data.trim()) else {
return SemanticSseObservation::SemanticEvent;
};
let payload_type = body_json.get("type").and_then(Value::as_str).or(event_type);
if payload_type.is_some_and(is_anthropic_semantic_event_type) {
return classify_anthropic_sse_record(record.as_bytes());
}
let error = body_json
.get("error")
.filter(|value| !value.is_null())
.or_else(|| {
body_json
.pointer("/response/error")
.filter(|value| !value.is_null())
});
if error.is_some()
|| matches!(payload_type, Some("error" | "response.failed"))
|| body_json.get("status").and_then(Value::as_str) == Some("failed")
{
let failure = error
.map(|error| serde_json::json!({ "error": error }))
.unwrap_or_else(|| body_json.clone());
return SemanticSseObservation::Error {
status_code: crate::execution_runtime::submission::resolve_local_sync_error_status_code(
200, &failure,
),
body_json: failure,
};
}
if matches!(
payload_type,
Some("ping" | "response.created" | "response.in_progress" | "response.queued")
) {
return SemanticSseObservation::Pending;
}
if payload_type == Some("response.output_item.added")
&& matches!(
body_json.pointer("/item/type").and_then(Value::as_str),
Some("message" | "reasoning")
)
&& body_json
.pointer("/item/content")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty)
&& body_json
.pointer("/item/summary")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty)
{
return SemanticSseObservation::Pending;
}
if matches!(
payload_type,
Some("response.content_part.added" | "response.reasoning_summary_part.added")
) && matches!(
body_json.pointer("/part/type").and_then(Value::as_str),
Some("output_text" | "summary_text" | "refusal")
) && !body_json
.pointer("/part/text")
.is_some_and(value_has_semantic_content)
&& !body_json
.pointer("/part/refusal")
.is_some_and(value_has_semantic_content)
{
return SemanticSseObservation::Pending;
}
if let Some(choices) = body_json.get("choices").and_then(Value::as_array) {
let semantic = choices.iter().any(|choice| {
choice
.get("finish_reason")
.is_some_and(|value| !value.is_null())
|| choice.get("text").is_some_and(value_has_semantic_content)
|| choice
.get("delta")
.or_else(|| choice.get("message"))
.and_then(Value::as_object)
.is_some_and(|delta| {
delta.iter().any(|(name, value)| {
name != "role" && value_has_semantic_content(value)
})
})
});
return if semantic {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
};
}
SemanticSseObservation::SemanticEvent
}
fn value_has_semantic_content(value: &Value) -> bool {
match value {
Value::Null => false,
Value::String(text) => !text.is_empty(),
Value::Array(values) => !values.is_empty(),
Value::Object(values) => !values.is_empty(),
_ => true,
}
}
#[derive(Debug, Default)]
struct AnthropicSsePrecommitInspector {
buffered: Vec<u8>,
@@ -355,7 +535,30 @@ fn classify_anthropic_sse_record(record: &[u8]) -> SemanticSseObservation {
(None, Some(payload_type)) => Some(payload_type),
_ => None,
};
if semantic_type.is_some_and(is_anthropic_semantic_event_type) {
let setup_only = match semantic_type {
Some("message_start") => body_json
.pointer("/message/content")
.and_then(Value::as_array)
.is_none_or(Vec::is_empty),
Some("content_block_start") => {
let block_type = body_json
.pointer("/content_block/type")
.and_then(Value::as_str);
matches!(block_type, Some("text" | "thinking"))
&& !body_json
.pointer("/content_block/text")
.is_some_and(value_has_semantic_content)
&& !body_json
.pointer("/content_block/thinking")
.is_some_and(value_has_semantic_content)
}
Some("content_block_stop") => true,
Some("message_delta") => body_json
.pointer("/delta/stop_reason")
.is_none_or(Value::is_null),
_ => false,
};
if !setup_only && semantic_type.is_some_and(is_anthropic_semantic_event_type) {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
@@ -507,6 +710,89 @@ pub(super) fn anthropic_error_status_code(body_json: &Value) -> u16 {
#[cfg(test)]
mod tests {
#[test]
fn image_streams_only_prefetch_when_explicitly_requested() {
for force_prefetch in [false, true] {
let policy = super::StreamCommitPolicy::for_response(
true,
Some("text/event-stream"),
"openai:image",
"openai:image",
false,
false,
force_prefetch,
);
assert_eq!(policy.commits_on_response_headers(), !force_prefetch);
assert!(!policy.requires_bounded_frame_wait());
}
}
#[test]
fn generic_sse_waits_through_setup_and_classifies_fragmented_errors() {
let setup = b"event: response.created\ndata: {\"type\":\"response.created\"}\n\n";
let failure = b"event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"type\":\"server_error\",\"message\":\"capacity exhausted\"}}}\n\n";
for split in 1..failure.len() {
let policy = super::StreamCommitPolicy::FirstSseSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
};
let mut gate = super::StreamCommitGate::new(policy);
assert_eq!(
gate.observe_provider_bytes(setup),
super::StreamPrecommitObservation::Pending
);
for control in [
b"event: ping\ndata: keepalive\n\n".as_slice(),
b"data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"reasoning\",\"summary\":[]}}\n\n".as_slice(),
b"data: {\"type\":\"response.reasoning_summary_part.added\",\"part\":{\"type\":\"summary_text\",\"text\":\"\"}}\n\n".as_slice(),
] {
assert_eq!(gate.observe_provider_bytes(control), super::StreamPrecommitObservation::Pending);
}
assert_eq!(
gate.observe_provider_bytes(&failure[..split]),
super::StreamPrecommitObservation::Pending
);
assert!(matches!(
gate.observe_provider_bytes(&failure[split..]),
super::StreamPrecommitObservation::UpstreamError { .. }
));
}
}
#[test]
fn generic_sse_commits_on_content_or_tool_call_but_not_role() {
for output in [
"data: {\"choices\":[{\"delta\":{\"content\":\"hello\"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"id\":\"call-1\"}]}}]}\n\n",
] {
let mut gate =
super::StreamCommitGate::new(super::StreamCommitPolicy::FirstSseSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
});
assert_eq!(gate.observe_provider_bytes(b"data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":\"\"}}]}\n\n"), super::StreamPrecommitObservation::Pending);
assert_eq!(
gate.observe_provider_bytes(output.as_bytes()),
super::StreamPrecommitObservation::Commit
);
assert_eq!(
gate.observe_provider_bytes(b"data: {\"error\":{\"message\":\"late error\"}}\n\n"),
super::StreamPrecommitObservation::Commit
);
}
}
#[test]
fn native_anthropic_setup_does_not_hide_an_early_error() {
let mut gate =
super::StreamCommitGate::new(super::StreamCommitPolicy::FirstAnthropicSemanticEvent {
max_bytes: 4096,
max_wait: std::time::Duration::from_secs(1),
});
assert_eq!(gate.observe_provider_bytes(b"event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"content\":[]}}\n\n"), super::StreamPrecommitObservation::Pending);
assert_eq!(gate.observe_provider_bytes(b"event: content_block_start\ndata: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n"), super::StreamPrecommitObservation::Pending);
assert!(matches!(gate.observe_provider_bytes(b"event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n"), super::StreamPrecommitObservation::UpstreamError { status_code: 529, .. }));
}
use std::time::Duration;
use super::{
@@ -553,7 +839,7 @@ mod tests {
false,
false,
)
.commits_on_response_headers());
.requires_bounded_frame_wait());
assert!(StreamCommitPolicy::for_response(
true,
Some("text/event-stream"),
@@ -563,7 +849,7 @@ mod tests {
true,
false,
)
.commits_on_response_headers());
.requires_bounded_frame_wait());
}
#[test]
@@ -745,8 +1031,8 @@ mod tests {
let mut gate = StreamCommitGate::new(native_anthropic_policy());
let observation = gate.observe_provider_bytes(
concat!(
"event: message_start\n",
"data: {\"type\":\"message_start\",\"message\":{}}\n\n",
"event: content_block_delta\n",
"data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n",
"event: error\n",
"data: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\"}}\n\n",
)
File diff suppressed because it is too large Load Diff
@@ -529,7 +529,13 @@ fn classify_local_sync_error_kind(
{
return LocalCoreSyncErrorKind::Overloaded;
}
if (500..600).contains(&status_code) {
if (500..600).contains(&status_code)
|| raw_type.is_some_and(|value| {
["server_error", "internal_error", "api_error"]
.iter()
.any(|kind| value.trim().eq_ignore_ascii_case(kind))
})
{
return LocalCoreSyncErrorKind::ServerError;
}
LocalCoreSyncErrorKind::InvalidRequest
@@ -676,6 +682,13 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize(
#[cfg(test)]
mod tests {
#[test]
fn success_http_status_does_not_misclassify_explicit_server_errors_as_bad_requests() {
for error_type in ["server_error", "internal_error", "api_error"] {
let body = serde_json::json!({ "error": { "type": error_type, "message": "failed" } });
assert_eq!(super::resolve_local_sync_error_status_code(200, &body), 500);
}
}
use axum::body::to_bytes;
use serde_json::json;
@@ -438,7 +438,7 @@ static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock<DirectHyperH2cSenderCacheMetric
LazyLock::new(DirectHyperH2cSenderCacheMetrics::default);
#[derive(Debug, Clone, Copy, Default)]
struct ExecutionSafeDnsResolver;
pub(crate) struct ExecutionSafeDnsResolver;
#[derive(Debug, Clone, Copy, Default)]
struct ExecutionSafeHyperDnsResolver;
@@ -446,10 +446,7 @@ struct ExecutionSafeHyperDnsResolver;
fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
let host = host.trim_end_matches('.');
host.eq_ignore_ascii_case("localhost")
|| host
.parse::<IpAddr>()
.map(|ip| ip.is_loopback())
.unwrap_or(false)
|| aether_http::parse_ip_literal_host(host).is_some_and(|ip| ip.is_loopback())
}
fn validate_resolved_execution_addresses(
@@ -491,12 +488,9 @@ async fn resolve_execution_target_addresses_with_policy(
port: u16,
provider_execution: bool,
) -> Result<Vec<SocketAddr>, std::io::Error> {
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
let addresses =
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await?
};
.await?;
validate_resolved_execution_addresses(host, addresses, provider_execution)
}
@@ -5149,7 +5143,7 @@ fn execution_log_url_host(url: &str) -> String {
.unwrap_or_else(|| "-".to_string())
}
fn validate_execution_upstream_url(
pub(crate) fn validate_execution_upstream_url(
raw_url: &str,
) -> Result<url::Url, ExecutionRuntimeTransportError> {
let url = url::Url::parse(raw_url).map_err(|_| {
@@ -5316,7 +5310,7 @@ pub(crate) fn build_execution_response_body(
mod tests {
use std::collections::BTreeMap;
use std::io::{Read, Write};
use std::sync::{Arc, Mutex, MutexGuard, OnceLock};
use std::sync::{Arc, Mutex};
use aether_contracts::tunnel::{
TUNNEL_RELAY_AUTH_NONCE_HEADER, TUNNEL_RELAY_AUTH_PAYLOAD_HEADER,
@@ -5440,6 +5434,8 @@ mod tests {
"93.184.216.34:443".parse().unwrap(),
];
for host in [
"chatgpt.com",
"api.openai.com",
"oauth2.googleapis.com",
"www.googleapis.com",
"custom.example.test",
@@ -5452,6 +5448,46 @@ mod tests {
}
}
#[tokio::test]
async fn execution_dns_handles_url_ipv6_without_weakening_relay_filtering() {
for provider_execution in [false, true] {
let addresses = super::resolve_execution_target_addresses_with_policy(
"[::1]",
8443,
provider_execution,
)
.await
.expect("literal IPv6 loopback should resolve without DNS");
assert_eq!(addresses, vec!["[::1]:8443".parse().unwrap()]);
}
let error = super::resolve_execution_target_addresses_with_policy("[fd00::1]", 443, false)
.await
.expect_err("private IPv6 must remain blocked for relay traffic");
assert_eq!(error.kind(), std::io::ErrorKind::PermissionDenied);
}
#[tokio::test]
async fn execution_dns_resolvers_preserve_provider_fake_ip_answers() {
for host in ["198.18.78.41", "198.19.1.2"] {
let expected = vec![format!("{host}:0").parse::<std::net::SocketAddr>().unwrap()];
let reqwest_addresses = reqwest::dns::Resolve::resolve(
&super::ExecutionSafeDnsResolver,
host.parse().unwrap(),
)
.await
.expect("HTTP provider DNS must accept Fake-IP answers")
.collect::<Vec<_>>();
let wreq_addresses =
wreq::dns::Resolve::resolve(&super::ExecutionSafeDnsResolver, host.into())
.await
.expect("WebSocket provider DNS must accept Fake-IP answers")
.collect::<Vec<_>>();
assert_eq!(reqwest_addresses, expected);
assert_eq!(wreq_addresses, expected);
}
}
#[test]
fn execution_dns_answers_keep_relay_address_filtering() {
let public = "93.184.216.34:443".parse().unwrap();
@@ -6228,16 +6264,14 @@ mod tests {
TestEnvVarGuard { key, previous }
}
fn direct_reqwest_env_lock() -> MutexGuard<'static, ()> {
static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
LOCK.get_or_init(|| Mutex::new(()))
.lock()
.expect("direct reqwest env lock")
fn direct_reqwest_env_lock() -> &'static tokio::sync::Mutex<()> {
static LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
&LOCK
}
#[test]
fn direct_reqwest_client_cache_key_includes_transport_profile() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let timeouts = ExecutionTimeouts {
connect_ms: Some(5_000),
..ExecutionTimeouts::default()
@@ -6341,7 +6375,7 @@ mod tests {
#[test]
fn direct_reqwest_client_cache_evicts_least_recently_used_entry_at_capacity() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _capacity = set_test_env_var(super::DIRECT_REQWEST_CACHE_MAX_ENTRIES_ENV, "2");
let cache_key = |suffix| {
super::direct_reqwest_client_cache_key(
@@ -6439,7 +6473,7 @@ mod tests {
#[test]
fn direct_reqwest_client_cache_key_splits_origin_only_when_enabled() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-origin".into(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(),
@@ -6626,14 +6660,14 @@ mod tests {
#[test]
fn direct_h2c_client_shards_respect_explicit_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_H2C_CLIENT_SHARDS_ENV, "7");
assert_eq!(super::direct_h2c_client_shard_count(), 7);
}
#[test]
fn direct_h2c_adaptive_window_respects_explicit_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
{
let _adaptive = set_test_env_var(super::DIRECT_H2C_ADAPTIVE_WINDOW_ENV, "0");
assert!(!super::direct_h2c_adaptive_window_enabled());
@@ -6701,7 +6735,7 @@ mod tests {
#[test]
fn direct_h2c_prewarm_urls_parse_env_list() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _urls = set_test_env_var(
super::DIRECT_H2C_PREWARM_URLS_ENV,
" http://127.0.0.1:18184/v1/chat/completions,;http://127.0.0.1:18185/v1/chat/completions\nhttp://127.0.0.1:18186/v1/chat/completions ",
@@ -6719,7 +6753,7 @@ mod tests {
#[test]
fn direct_h2c_prewarm_cache_keys_dedup_by_origin() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let urls = vec![
"http://127.0.0.1:18184/v1/chat/completions".to_string(),
"http://127.0.0.1:18184/v1/responses".to_string(),
@@ -6745,7 +6779,7 @@ mod tests {
#[test]
fn direct_h2c_client_cache_splits_by_origin_and_shards() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_H2C_CLIENT_SHARDS_ENV, "3");
super::DIRECT_H2C_CLIENT_CACHE
.lock()
@@ -6770,7 +6804,7 @@ mod tests {
#[test]
fn direct_reqwest_initial_client_shards_are_bounded_by_target() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
assert_eq!(super::direct_reqwest_initial_client_shard_count(1), 1);
assert_eq!(super::direct_reqwest_initial_client_shard_count(2), 2);
assert_eq!(
@@ -6781,7 +6815,7 @@ mod tests {
#[test]
fn direct_reqwest_initial_client_shards_cap_large_sync_env() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "128");
assert_eq!(
super::direct_reqwest_initial_client_shard_count(128),
@@ -6791,7 +6825,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_client_shards_default_to_initial() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
assert_eq!(super::direct_reqwest_prewarm_client_shard_count(1), 1);
assert_eq!(
super::direct_reqwest_prewarm_client_shard_count(96),
@@ -6801,7 +6835,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_client_shards_do_not_exceed_request_path_cap() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "4");
let _prewarm = set_test_env_var(super::DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV, "128");
@@ -6810,7 +6844,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_populates_cache_for_plan() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "4");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-prewarm".into(),
@@ -6872,7 +6906,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_plan_keeps_large_sync_env_off_request_path() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "128");
let _sync = set_test_env_var(super::DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV, "4");
let _prewarm = set_test_env_var(super::DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV, "128");
@@ -6932,7 +6966,7 @@ mod tests {
#[test]
fn direct_reqwest_prewarm_skips_h2c_fast_path() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _fast_path = set_test_env_var(super::DIRECT_H2C_FAST_PATH_ENV, "1");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-fast-path-prewarm-skip".into(),
@@ -6985,7 +7019,7 @@ mod tests {
#[test]
fn direct_reqwest_cache_metrics_expose_ready_state() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().blocking_lock();
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "1");
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-ready-metrics".into(),
@@ -8569,7 +8603,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_supports_tunnel_relay() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET);
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -8736,7 +8770,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_rejects_short_tunnel_relay_secret_before_send() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", &"x".repeat(31));
let execution_runtime = DirectSyncExecutionRuntime::new();
let error = execution_runtime
@@ -8777,7 +8811,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_requires_tunnel_relay_secret_before_send() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = unset_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET");
let execution_runtime = DirectSyncExecutionRuntime::new();
let error = execution_runtime
@@ -9104,7 +9138,7 @@ mod tests {
#[tokio::test]
async fn direct_sync_execution_runtime_forwards_http1_only_control_to_tunnel_relay() {
let _env_lock = direct_reqwest_env_lock();
let _env_lock = direct_reqwest_env_lock().lock().await;
let _relay_secret = set_test_env_var("AETHER_TUNNEL_RELAY_AUTH_SECRET", RELAY_TEST_SECRET);
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -9320,7 +9354,7 @@ mod tests {
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn direct_sync_execution_runtime_uses_h2c_prior_knowledge_on_wire() {
let _guard = direct_reqwest_env_lock();
let _guard = direct_reqwest_env_lock().lock().await;
let _shards = set_test_env_var(super::DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV, "1");
let listener = crate::test_support::bind_loopback_listener()
.await
@@ -655,6 +655,73 @@ struct ProviderTransferState {
struct ProviderTransferStateTracker {
by_provider: BTreeMap<String, ProviderTransferState>,
exhausted_provider_ids: BTreeSet<String>,
global: GlobalTransferState,
}
#[derive(Debug, Default)]
struct GlobalTransferState {
first_attempt_started_at: Option<Instant>,
last_candidate: Option<(String, String, String)>,
transfer_count: u64,
limits: Option<ProviderTransferLimits>,
exhausted: bool,
}
impl GlobalTransferState {
fn load_policy(&mut self, report_context: Option<&serde_json::Value>) {
if self.limits.is_none() {
if let Some(policy) =
crate::orchestration::routing_execution_policy_from_report_context(report_context)
{
self.limits = Some(ProviderTransferLimits {
max_transfer_count: policy.max_transfer_count,
max_transfer_timeout_seconds: policy.max_transfer_timeout_seconds,
});
}
}
}
fn changes_candidate(&self, plan: &aether_contracts::ExecutionPlan) -> bool {
self.last_candidate
.as_ref()
.is_some_and(|(provider, endpoint, key)| {
provider != &plan.provider_id
|| endpoint != &plan.endpoint_id
|| key != &plan.key_id
})
}
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.first_attempt_started_at.get_or_insert(now);
if self.changes_candidate(plan) {
self.transfer_count = self.transfer_count.saturating_add(1);
}
self.last_candidate = Some((
plan.provider_id.clone(),
plan.endpoint_id.clone(),
plan.key_id.clone(),
));
}
fn check_before_attempt(
&mut self,
plan: &aether_contracts::ExecutionPlan,
now: Instant,
) -> Option<(bool, bool)> {
let limits = self.limits?;
let started_at = self.first_attempt_started_at?;
let count_reached = self.changes_candidate(plan)
&& limits.max_transfer_count > 0
&& self.transfer_count >= limits.max_transfer_count;
let timeout_reached = limits.max_transfer_timeout_seconds > 0
&& now.saturating_duration_since(started_at)
>= Duration::from_secs(limits.max_transfer_timeout_seconds);
if !count_reached && !timeout_reached {
return None;
}
self.exhausted = true;
Some((count_reached, timeout_reached))
}
}
#[derive(Clone, Debug, Default)]
@@ -717,6 +784,7 @@ struct ProviderTransferLimitReached {
impl ProviderTransferStateTracker {
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
self.global.record_attempt_started(plan, now);
match self.by_provider.entry(plan.provider_id.clone()) {
std::collections::btree_map::Entry::Vacant(entry) => {
entry.insert(ProviderTransferState {
@@ -903,11 +971,42 @@ async fn should_skip_provider_transfer_attempt<Attempt>(
where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let reached = tracker
.state
.lock()
.await
.check_before_attempt(attempt.execution_plan(), Instant::now());
let owned_report_context = attempt
.report_context_ref()
.is_none()
.then(|| attempt.report_context())
.flatten();
let report_context = attempt
.report_context_ref()
.or(owned_report_context.as_ref());
let mut tracker = tracker.state.lock().await;
tracker.global.load_policy(report_context);
if tracker.global.exhausted {
return true;
}
let now = Instant::now();
if let Some((count_reached, timeout_reached)) = tracker
.global
.check_before_attempt(attempt.execution_plan(), now)
{
warn!(
event_name = "routing_transfer_limit_reached",
log_type = "event",
trace_id,
plan_kind,
transfer_count = tracker.global.transfer_count,
elapsed_ms = tracker
.global
.first_attempt_started_at
.map(|started| now.saturating_duration_since(started).as_millis() as u64)
.unwrap_or(0),
count_reached,
timeout_reached,
"gateway exhausted the routing strategy transfer budget"
);
return true;
}
let reached = tracker.check_before_attempt(attempt.execution_plan(), now);
let Some(reached) = reached else {
return false;
};
@@ -2465,6 +2564,130 @@ mod tests {
assert_eq!(port.unused.lock().unwrap().as_slice(), ["a-key3-retry0"]);
}
#[tokio::test]
async fn routing_transfer_budget_counts_switches_across_providers_not_same_key_retries() {
for (limit, succeeds) in [(1, false), (2, true)] {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] =
json!({ "max_transfer_count": limit });
}
let outcome = run_ai_attempt_loop(&port, attempts).await.unwrap();
assert_eq!(
matches!(outcome, AiAttemptLoopOutcome::Responded(_)),
succeeds
);
{
let executed = port.executed.lock().unwrap();
assert_eq!(
&executed[..3],
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(executed.len(), if succeeds { 4 } else { 3 });
}
assert_eq!(port.tracker.state.lock().await.global.transfer_count, limit);
}
}
#[tokio::test]
async fn dynamic_loop_honors_global_transfer_budget_across_providers() {
let state = AppState::new().unwrap();
let port = TransferTestPort::new(&state);
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let mut source = TransferTestAttemptSource {
attempts: attempts.into(),
skipped_providers: Vec::new(),
};
let outcome = run_dynamic_attempt_loop(
&port,
&mut source,
"global-budget",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
outcome,
LocalExecutionRequestOutcome::Exhausted(_)
));
assert_eq!(
port.executed.lock().unwrap().as_slice(),
["a-key1-retry0", "a-key2-retry0", "a-key2-retry1"]
);
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
}
#[test]
fn routing_time_budget_is_cumulative_and_zero_is_unlimited() {
let mut global = super::GlobalTransferState::default();
global.load_policy(Some(
&json!({ "routing_execution_policy": { "max_transfer_timeout_seconds": 60 } }),
));
let now = tokio::time::Instant::now();
let plan = test_plan(None);
global.record_attempt_started(&plan, now);
global.record_attempt_started(&plan, now + Duration::from_secs(40));
assert_eq!(global.transfer_count, 0);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(59)),
None
);
assert_eq!(
global.check_before_attempt(&plan, now + Duration::from_secs(60)),
Some((false, true))
);
let mut unlimited = super::GlobalTransferState::default();
unlimited.load_policy(Some(&json!({ "routing_execution_policy": {} })));
unlimited.record_attempt_started(&plan, now);
assert_eq!(
unlimited.check_before_attempt(&plan, now + Duration::from_secs(86_400)),
None
);
}
#[tokio::test]
async fn cloned_tracker_preserves_global_budget_across_candidate_loops() {
let state = AppState::new().unwrap();
let tracker = ProviderTransferTracker::default();
let mut attempts = transfer_test_attempts();
for attempt in &mut attempts {
attempt.report_context["routing_execution_policy"] = json!({ "max_transfer_count": 1 });
}
let remaining = attempts.split_off(3);
let first_port = TransferTestPort::with_tracker(&state, tracker.clone());
let first_outcome = run_ai_attempt_loop(&first_port, attempts).await.unwrap();
assert!(matches!(first_outcome, AiAttemptLoopOutcome::Exhausted(_)));
assert_eq!(tracker.state.lock().await.global.transfer_count, 1);
let second_port = TransferTestPort::with_tracker(&state, tracker.clone());
let mut source = TransferTestAttemptSource {
attempts: remaining.into(),
skipped_providers: Vec::new(),
};
let second_outcome = run_dynamic_attempt_loop(
&second_port,
&mut source,
"global-budget-across-loops",
"test",
Duration::from_secs(1),
)
.await
.unwrap();
assert!(matches!(
second_outcome,
LocalExecutionRequestOutcome::NoPath
));
assert!(second_port.executed.lock().unwrap().is_empty());
assert_eq!(source.skipped_providers, ["provider-a", "provider-b"]);
assert!(tracker.state.lock().await.global.exhausted);
}
#[tokio::test]
async fn cloned_tracker_preserves_transfer_budget_across_candidate_loops() {
let state = AppState::new().expect("state should build");
@@ -822,15 +822,20 @@ where
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = standard_text_sync_heartbeat_final_bytes(
let completion = standard_text_sync_heartbeat_final_bytes(
client_api_format.as_str(),
redaction_slot.as_ref(),
execute(state, parts, trace_id, decision, plan_kind, started_at).await,
)
.await;
tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execute(state, parts, trace_id, decision, plan_kind, started_at) => result,
},
);
let bytes = completion.await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -1097,23 +1102,26 @@ fn build_openai_image_sync_heartbeat_shell_response(
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
let request_diagnostics = current_request_diagnostics();
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
tokio::spawn(async move {
scope_request_diagnostics_with(request_diagnostics, async move {
let bytes = openai_image_sync_heartbeat_final_bytes(
execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
)
.await,
)
.await;
let execution = execute_openai_image_sync_heartbeat_attempts(
state,
request_path,
trace_id,
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
);
let outcome = tokio::select! {
biased;
_ = tx.closed(), if cancel_on_disconnect => return,
result = execution => result,
};
let bytes = openai_image_sync_heartbeat_final_bytes(outcome).await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
})
.await;
@@ -2331,6 +2339,45 @@ mod tests {
.expect("background completion should release admission");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_cancels_when_routing_policy_enables_it() {
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (mut release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let response = crate::request_lifecycle::run_request(async move {
crate::request_lifecycle::configure_client_disconnect(
aether_routing_core::RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
},
);
let (parts, _) = http::Request::builder()
.method("POST")
.uri("/v1/responses")
.body(())
.unwrap()
.into_parts();
build_standard_text_sync_heartbeat_shell_response(
AppState::new().unwrap(),
parts,
"trace-heartbeat-disconnect".to_string(),
test_standard_text_heartbeat_decision(),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND.to_string(),
move |_, _, _, _, _, _| async move {
started_tx.send(()).unwrap();
release_rx.await.unwrap();
Ok(LocalExecutionRequestOutcome::NoPath)
},
)
})
.await
.unwrap();
started_rx.await.unwrap();
drop(response);
tokio::time::timeout(Duration::from_secs(1), release_tx.closed())
.await
.expect("heartbeat must drop upstream execution immediately");
}
#[tokio::test]
async fn standard_text_sync_heartbeat_propagates_request_diagnostics_to_terminal_usage() {
let (state, usage_repository) = heartbeat_usage_test_state(json!({
@@ -339,57 +339,6 @@ async fn build_admin_oauth_test_payload(
}))
}
#[cfg(test)]
mod tests {
use super::{
is_fixed_linuxdo_oauth_origin, resolve_public_admin_oauth_endpoint,
validate_public_admin_oauth_resolved_addrs,
};
use std::net::SocketAddr;
#[tokio::test]
async fn oauth_test_endpoint_rejects_loopback_https_targets_before_connecting() {
let url = reqwest::Url::parse("https://127.0.0.1/oauth/token").expect("URL");
assert!(resolve_public_admin_oauth_endpoint(&url).await.is_err());
}
#[test]
fn linuxdo_builtin_origin_allows_only_benchmarking_addresses() {
let fixed = reqwest::Url::parse("https://connect.linux.do/oauth2/token")
.expect("LinuxDo URL should parse");
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
assert!(is_fixed_linuxdo_oauth_origin(&fixed));
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], true).is_ok());
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], false).is_err());
assert!(validate_public_admin_oauth_resolved_addrs(
&fixed,
&[fake, SocketAddr::from(([127, 0, 0, 1], 443))],
true,
)
.is_err());
}
#[test]
fn custom_or_non_default_oauth_origins_reject_benchmarking_addresses() {
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
for raw_url in [
"https://oauth.example.test/token",
"https://connect.linux.do:8443/oauth2/token",
"https://connect.linuxdo.org/oauth2/token",
"https://connect.linux.do.evil.test/oauth2/token",
"https://connect.linux.do/oauth2/token?tenant=unexpected",
] {
let url = reqwest::Url::parse(raw_url).expect("test URL should parse");
assert!(
!is_fixed_linuxdo_oauth_origin(&url),
"must not trust {raw_url}"
);
assert!(validate_public_admin_oauth_resolved_addrs(&url, &[fake], true).is_err());
}
}
}
pub(crate) async fn maybe_build_local_admin_oauth_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -689,3 +638,54 @@ pub(crate) async fn maybe_build_local_admin_oauth_response(
Ok(None)
}
#[cfg(test)]
mod tests {
use super::{
is_fixed_linuxdo_oauth_origin, resolve_public_admin_oauth_endpoint,
validate_public_admin_oauth_resolved_addrs,
};
use std::net::SocketAddr;
#[tokio::test]
async fn oauth_test_endpoint_rejects_loopback_https_targets_before_connecting() {
let url = reqwest::Url::parse("https://127.0.0.1/oauth/token").expect("URL");
assert!(resolve_public_admin_oauth_endpoint(&url).await.is_err());
}
#[test]
fn linuxdo_builtin_origin_allows_only_benchmarking_addresses() {
let fixed = reqwest::Url::parse("https://connect.linux.do/oauth2/token")
.expect("LinuxDo URL should parse");
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
assert!(is_fixed_linuxdo_oauth_origin(&fixed));
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], true).is_ok());
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], false).is_err());
assert!(validate_public_admin_oauth_resolved_addrs(
&fixed,
&[fake, SocketAddr::from(([127, 0, 0, 1], 443))],
true,
)
.is_err());
}
#[test]
fn custom_or_non_default_oauth_origins_reject_benchmarking_addresses() {
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
for raw_url in [
"https://oauth.example.test/token",
"https://connect.linux.do:8443/oauth2/token",
"https://connect.linuxdo.org/oauth2/token",
"https://connect.linux.do.evil.test/oauth2/token",
"https://connect.linux.do/oauth2/token?tenant=unexpected",
] {
let url = reqwest::Url::parse(raw_url).expect("test URL should parse");
assert!(
!is_fixed_linuxdo_oauth_origin(&url),
"must not trust {raw_url}"
);
assert!(validate_public_admin_oauth_resolved_addrs(&url, &[fake], true).is_err());
}
}
}
@@ -278,46 +278,6 @@ async fn build_batch_delete_global_models_response(
))
}
#[cfg(test)]
mod batch_boundary_tests {
use super::{normalize_admin_global_model_batch_ids, MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS};
#[test]
fn global_model_batch_ids_are_bounded_and_deduplicated() {
assert_eq!(
normalize_admin_global_model_batch_ids(
vec![
"model-2".to_string(),
"model-1".to_string(),
" model-2 ".to_string(),
" ".to_string(),
],
"ids",
)
.expect("valid ids"),
vec![
"model-2".to_string(),
"model-1".to_string(),
" ".to_string(),
]
);
assert!(normalize_admin_global_model_batch_ids(
(0..=MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("model-{index}"))
.collect(),
"ids",
)
.is_err());
assert!(normalize_admin_global_model_batch_ids(
(0..MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("provider-{index}"))
.collect(),
"provider_ids",
)
.is_ok());
}
}
async fn build_assign_to_providers_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -357,3 +317,43 @@ async fn build_assign_to_providers_response(
&global_model_id,
))
}
#[cfg(test)]
mod batch_boundary_tests {
use super::{normalize_admin_global_model_batch_ids, MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS};
#[test]
fn global_model_batch_ids_are_bounded_and_deduplicated() {
assert_eq!(
normalize_admin_global_model_batch_ids(
vec![
"model-2".to_string(),
"model-1".to_string(),
" model-2 ".to_string(),
" ".to_string(),
],
"ids",
)
.expect("valid ids"),
vec![
"model-2".to_string(),
"model-1".to_string(),
" ".to_string(),
]
);
assert!(normalize_admin_global_model_batch_ids(
(0..=MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("model-{index}"))
.collect(),
"ids",
)
.is_err());
assert!(normalize_admin_global_model_batch_ids(
(0..MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("provider-{index}"))
.collect(),
"provider_ids",
)
.is_ok());
}
}
@@ -58,6 +58,7 @@ async fn admin_monitoring_trace_request_returns_local_payload() {
.expect("route should be handled locally");
assert_eq!(response.status(), http::StatusCode::OK);
assert!(response.headers().contains_key("x-aether-build-version"));
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
@@ -93,6 +94,15 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
Some(33),
Some(200),
),
sample_candidate(
"cand-other-attempt",
"trace-1",
1,
RequestCandidateStatus::Failed,
Some(100),
Some(20),
Some(502),
),
]));
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
@@ -110,6 +120,8 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
100,
);
usage.id = "usage-row-1".to_string();
usage.request_body_state = Some(UsageBodyCaptureState::Reference);
usage.response_body_state = Some(UsageBodyCaptureState::Reference);
usage.candidate_id = Some("cand-used".to_string());
usage.request_headers = Some(json!({
"x-trace-id": "trace-1"
@@ -140,6 +152,17 @@ async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["request_id"], json!("trace-1"));
assert_eq!(payload["diagnostic_request"]["usage_id"], "usage-row-1");
assert_eq!(
payload["candidates"][0]["extra_data"]["diagnostic_context"]["usage_id"],
"usage-row-1"
);
assert_eq!(
payload["candidates"][0]["extra_data"]["diagnostic_context"]["body_states"]
["response_body"],
"reference"
);
assert!(payload["candidates"][1]["extra_data"]["diagnostic_context"].is_null());
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
assert_eq!(
payload["candidates"][0]["extra_data"]["first_byte_time_ms"],
@@ -67,13 +67,19 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
let key_accounts =
build_admin_monitoring_key_account_display_map(admin_state, &resolved.trace).await?;
Ok(
build_admin_monitoring_trace_request_payload_response_with_key_accounts(
&resolved.trace,
resolved.usage.as_ref(),
&key_accounts,
),
)
let mut response = build_admin_monitoring_trace_request_payload_response_with_key_accounts(
&resolved.trace,
resolved.usage.as_ref(),
&key_accounts,
);
if let Ok(version) = axum::http::HeaderValue::from_str(
option_env!("AETHER_BUILD_VERSION").unwrap_or(env!("CARGO_PKG_VERSION")),
) {
response
.headers_mut()
.insert("x-aether-build-version", version);
}
Ok(response)
}
async fn resolve_admin_monitoring_trace(
@@ -66,43 +66,6 @@ fn admin_provider_oauth_kiro_refresh_error(
}
}
#[cfg(test)]
mod refresh_error_tests {
use super::admin_provider_oauth_kiro_refresh_error;
use crate::handlers::admin::request::AdminKiroAuthConfig;
use aether_oauth::core::OAuthError;
#[test]
fn kiro_refresh_error_does_not_reflect_upstream_body() {
let auth_config = AdminKiroAuthConfig {
auth_method: None,
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let detail = admin_provider_oauth_kiro_refresh_error(
&auth_config,
OAuthError::HttpStatus {
status_code: 502,
body_excerpt: "authorization=Bearer upstream-secret".to_string(),
},
);
assert_eq!(detail, "social refresh 失败: HTTP 502");
assert!(!detail.contains("upstream-secret"));
}
}
pub(super) async fn refresh_admin_provider_oauth_kiro_auth_config(
state: &AdminAppState<'_>,
auth_config: &AdminKiroAuthConfig,
@@ -240,3 +203,40 @@ pub(super) async fn fetch_admin_provider_oauth_kiro_email(
aether_admin::provider::quota::parse_kiro_usage_response(&payload, current_unix_secs())?;
json_non_empty_string(metadata.get("email"))
}
#[cfg(test)]
mod refresh_error_tests {
use super::admin_provider_oauth_kiro_refresh_error;
use crate::handlers::admin::request::AdminKiroAuthConfig;
use aether_oauth::core::OAuthError;
#[test]
fn kiro_refresh_error_does_not_reflect_upstream_body() {
let auth_config = AdminKiroAuthConfig {
auth_method: None,
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let detail = admin_provider_oauth_kiro_refresh_error(
&auth_config,
OAuthError::HttpStatus {
status_code: 502,
body_excerpt: "authorization=Bearer upstream-secret".to_string(),
},
);
assert_eq!(detail, "social refresh 失败: HTTP 502");
assert!(!detail.contains("upstream-secret"));
}
}
@@ -1035,7 +1035,7 @@ pub(crate) async fn proxy_request(
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
request: Request,
) -> Result<Response<Body>, GatewayError> {
crate::request_diagnostics::scope_request_diagnostics(Box::pin(proxy_request_inner(
crate::request_lifecycle::run_request(Box::pin(proxy_request_inner(
state,
remote_addr,
request,
@@ -3228,7 +3228,7 @@ mod tests {
async fn request_body_buffer_caps_decompressed_body_at_shared_budget() {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder
.write_all(&vec![b'a'; 128])
.write_all(&[b'a'; 128])
.expect("test gzip body should encode");
let encoded = encoder.finish().expect("test gzip body should finish");
assert!(
@@ -68,7 +68,12 @@ pub(super) async fn relay_bound_connection(
state: &AppState,
context: &WebSocketRequestContext,
) {
let mut client_connected = true;
loop {
if !client_connected && !bound.turn_state.response_in_flight() {
close_bound_upstream(bound).await;
break;
}
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
tokio::select! {
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
@@ -100,8 +105,12 @@ pub(super) async fn relay_bound_connection(
).await;
break;
}
client_message = client_socket.next() => {
client_message = client_socket.next(), if client_connected => {
let Some(client_message) = client_message else {
if retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
finalize_active_turn(
bound,
state,
@@ -111,6 +120,10 @@ pub(super) async fn relay_bound_connection(
break;
};
let Ok(client_message) = client_message else {
if retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
warn!(
event_name = "responses_websocket_client_receive_failed",
log_type = "ops",
@@ -127,6 +140,12 @@ pub(super) async fn relay_bound_connection(
close_bound_upstream(bound).await;
break;
};
if matches!(client_message, AxumWsMessage::Close(_))
&& retain_disconnected_turn(bound)
{
client_connected = false;
continue;
}
match Box::pin(forward_client_message(
client_message,
bound,
@@ -559,6 +578,7 @@ pub(super) async fn relay_bound_connection(
let mut relay_send_error = None;
let mut relay_serialization_failed = false;
match relay_directive {
_ if !client_connected => {}
Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => {
let client_frame = match parsed_upstream_frame.as_ref().map(|frame| {
bound
@@ -673,6 +693,10 @@ pub(super) async fn relay_bound_connection(
break;
}
if let Some(error) = relay_send_error {
if terminal_outcome.is_none() && retain_disconnected_turn(bound) {
client_connected = false;
continue;
}
warn!(
event_name = "responses_websocket_client_send_failed",
log_type = "ops",
@@ -737,6 +761,20 @@ pub(super) async fn relay_bound_connection(
}
}
fn retain_disconnected_turn(bound: &mut BoundResponsesConnection) -> bool {
if bound
.turn_state
.attempt()
.is_none_or(|attempt| attempt.cancel_on_client_disconnect())
{
return false;
}
bound
.turn_state
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
true
}
struct PendingContinuationRegistration {
user_id: String,
api_key_id: String,
@@ -845,6 +845,13 @@ fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> Gatew
}
impl ResponsesProviderAttempt {
pub(super) fn cancel_on_client_disconnect(&self) -> bool {
crate::orchestration::routing_execution_policy_from_report_context(
self.lifecycle.report_context(),
)
.is_some_and(|policy| policy.cancel_on_client_disconnect)
}
/// Releases all per-turn capacity before terminal persistence starts.
/// Provider-pool runtime tokens normally use an awaited removal. The
/// bounded wait prevents a broken runtime backend from stalling the relay;
@@ -24,7 +24,7 @@ use wreq::ws::message::{CloseFrame as WreqCloseFrame, Message as WreqWsMessage};
use crate::ai_serving::AiExecutionDecision;
use crate::execution_runtime::transport::{
build_browser_wreq_client, build_request_headers, normalize_execution_proxy_url,
ExecutionTransportControls,
validate_execution_upstream_url, ExecutionSafeDnsResolver, ExecutionTransportControls,
};
use crate::frontdoor_loop_guard::gateway_frontdoor_self_loop_guard_error;
use crate::handlers::proxy::websocket::session::{
@@ -66,7 +66,7 @@ pub(crate) async fn connect_upstream_websocket(
)?;
let headers =
websocket_handshake_headers(&decision.provider_request_headers, errors.headers_invalid)?;
let client = build_websocket_client(decision, &upstream_url, errors).await?;
let client = build_websocket_client(decision, errors)?;
let response = client
.websocket(upstream_url.as_str())
.headers(headers)
@@ -149,31 +149,14 @@ pub(crate) fn websocket_upstream_url(
invalid_code: &'static str,
) -> Result<Url, &'static str> {
let mut url = Url::parse(raw).map_err(|_| invalid_code)?;
if url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
|| url.fragment().is_some()
{
return Err(invalid_code);
}
let websocket_scheme = match url.scheme() {
"https" | "wss" => "wss",
"http" | "ws" => "ws",
let (http_scheme, websocket_scheme) = match url.scheme() {
"https" | "wss" => ("https", "wss"),
"http" | "ws" => ("http", "ws"),
_ => return Err(invalid_code),
};
url.set_scheme(http_scheme).map_err(|_| invalid_code)?;
let mut url = validate_execution_upstream_url(url.as_str()).map_err(|_| invalid_code)?;
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
if url.scheme() == "ws" {
let literal_ip = match url.host() {
Some(url::Host::Ipv4(address)) => Some(std::net::IpAddr::V4(address)),
Some(url::Host::Ipv6(address)) => Some(std::net::IpAddr::V6(address)),
_ => None,
};
if literal_ip.is_some_and(|address| {
aether_http::is_private_or_reserved_ip(address) && !address.is_loopback()
}) {
return Err(invalid_code);
}
}
Ok(url)
}
@@ -229,9 +212,8 @@ pub(crate) fn websocket_handshake_headers(
Ok(headers)
}
async fn build_websocket_client(
fn build_websocket_client(
decision: &AiExecutionDecision,
upstream_url: &Url,
errors: UpstreamWebSocketErrorCodes,
) -> Result<wreq::Client, &'static str> {
let timeouts = websocket_timeouts(decision);
@@ -255,41 +237,7 @@ async fn build_websocket_client(
let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?;
builder = builder.proxy(proxy);
} else {
// Pin every direct WebSocket connection to the DNS answers validated
// here. This also covers the explicitly permitted loopback `ws://`
// form; otherwise the client would perform a second lookup and a
// rebinding could escape the loopback-only policy.
let host = upstream_url.host_str().ok_or(errors.upstream_url_invalid)?;
let port = upstream_url
.port_or_known_default()
.ok_or(errors.upstream_url_invalid)?;
let addresses = if let Ok(ip) = host.parse::<std::net::IpAddr>() {
vec![std::net::SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host,
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| errors.upstream_url_invalid)?
};
let allows_loopback = host.trim_end_matches('.').eq_ignore_ascii_case("localhost")
|| host
.parse::<std::net::IpAddr>()
.map(|ip| ip.is_loopback())
.unwrap_or(false);
let unsafe_answer = if allows_loopback {
addresses.iter().any(|address| !address.ip().is_loopback())
} else {
addresses
.iter()
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()))
};
if addresses.is_empty() || unsafe_answer {
return Err(errors.upstream_url_invalid);
}
builder = builder.resolve_to_addrs(host.to_string(), addresses.iter().copied());
builder = builder.dns_resolver(ExecutionSafeDnsResolver);
}
builder.build().map_err(|_| errors.client_build_failed)
}
@@ -684,15 +632,17 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16
#[cfg(test)]
mod tests {
use super::{
bounded_send, guarded_websocket_upstream_url, resolve_websocket_proxy_url,
responses_websocket_error_event, responses_websocket_error_event_with_stream_id,
websocket_handshake_headers, websocket_relay_frame_queue, websocket_response_headers,
websocket_upstream_url, UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl,
WebSocketRelayQueueError, WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY,
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
bounded_send, build_websocket_client, guarded_websocket_upstream_url,
resolve_websocket_proxy_url, responses_websocket_error_event,
responses_websocket_error_event_with_stream_id, websocket_handshake_headers,
websocket_relay_frame_queue, websocket_response_headers, websocket_upstream_url,
UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl, WebSocketRelayQueueError,
WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY, RELAY_WRITE_TIMEOUT,
TEARDOWN_WRITE_TIMEOUT,
};
use crate::ai_serving::AiExecutionDecision;
use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url;
use aether_contracts::ProxySnapshot;
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
use axum::http::HeaderMap;
use std::collections::BTreeMap;
use std::time::Duration;
@@ -876,7 +826,9 @@ mod tests {
"ws://example.test:8080/v1/responses",
"http://example.test:8080/v1/responses",
"http://8.8.8.8:8080/v1/responses",
"wss://8.8.8.8/v1/responses",
"ws://[2606:4700:4700::1111]:8080/v1/responses",
"wss://[2606:4700:4700::1111]/v1/responses",
"ws://localhost:8080/v1/responses",
"http://127.42.0.1:8080/v1/responses",
"ws://[::1]:8080/v1/responses",
@@ -888,6 +840,14 @@ mod tests {
}
for rejected in [
"http://10.0.0.1/v1/responses",
"wss://10.0.0.1/v1/responses",
"wss://127.0.0.1/v1/responses",
"wss://[::1]/v1/responses",
"wss://[fd00::1]/v1/responses",
"wss://[::ffff:127.0.0.1]/v1/responses",
"wss://169.254.169.254/v1/responses",
"wss://198.18.78.41/v1/responses",
"wss://198.19.1.2/v1/responses",
"ws://0.0.0.0:8080/v1/responses",
"ws://[::ffff:127.0.0.1]:8080/v1/responses",
"wss://example.test/v1/responses#secret",
@@ -903,6 +863,60 @@ mod tests {
}
}
#[tokio::test]
async fn websocket_client_build_defers_provider_dns_for_all_transport_profiles() {
let errors = UpstreamWebSocketErrorCodes {
upstream_url_missing: "missing",
upstream_url_invalid: "upstream_invalid",
frontdoor_self_loop: "frontdoor_self_loop",
headers_invalid: "headers_invalid",
client_build_failed: "client_build_failed",
proxy_invalid: "proxy_invalid",
tunnel_proxy_unsupported: "tunnel_unsupported",
handshake_failed: "handshake_failed",
upgrade_rejected: "upgrade_rejected",
upgrade_failed: "upgrade_failed",
};
for profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for proxy in [
None,
Some(ProxySnapshot {
enabled: Some(false),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("http://proxy.invalid:8080".to_string()),
..Default::default()
}),
Some(ProxySnapshot {
enabled: Some(true),
url: Some("socks5h://proxy.invalid:1080".to_string()),
..Default::default()
}),
] {
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "wss://upstream.invalid/v1/responses"
}))
.expect("minimal provider decision should deserialize");
decision.transport_profile = profile.clone();
decision.proxy = proxy;
build_websocket_client(&decision, errors)
.expect("building a client must not resolve the provider or proxy hostname");
}
}
}
#[test]
fn active_websocket_proxy_without_a_target_fails_closed() {
let errors = UpstreamWebSocketErrorCodes {
@@ -937,6 +951,94 @@ mod tests {
);
}
#[tokio::test]
async fn websocket_handshake_keeps_provider_dns_remote_for_http_and_socks_proxies() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let errors = UpstreamWebSocketErrorCodes {
upstream_url_missing: "missing",
upstream_url_invalid: "upstream_invalid",
frontdoor_self_loop: "frontdoor_self_loop",
headers_invalid: "headers_invalid",
client_build_failed: "client_build_failed",
proxy_invalid: "proxy_invalid",
tunnel_proxy_unsupported: "tunnel_unsupported",
handshake_failed: "handshake_failed",
upgrade_rejected: "upgrade_rejected",
upgrade_failed: "upgrade_failed",
};
for profile in [
None,
Some(ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
..Default::default()
}),
] {
for scheme in ["http", "socks5", "socks5h"] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let proxy_addr = listener.local_addr().unwrap();
let (release, released) = tokio::sync::oneshot::channel::<()>();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
if scheme != "http" {
let mut greeting = [0; 2];
stream.read_exact(&mut greeting).await.unwrap();
assert_eq!(greeting[0], 5);
let mut methods = vec![0; greeting[1] as usize];
stream.read_exact(&mut methods).await.unwrap();
assert!(methods.contains(&0));
stream.write_all(&[5, 0]).await.unwrap();
let mut request = [0; 4];
stream.read_exact(&mut request).await.unwrap();
assert_eq!(
request,
[5, 1, 0, 3],
"proxy must receive a domain, not an IP"
);
let host_len = stream.read_u8().await.unwrap();
let mut host = vec![0; host_len as usize];
stream.read_exact(&mut host).await.unwrap();
assert_eq!(host, b"provider-dns.invalid");
assert_eq!(stream.read_u16().await.unwrap(), 80);
stream
.write_all(&[5, 0, 0, 1, 127, 0, 0, 1, 0, 80])
.await
.unwrap();
}
let socket = tokio_tungstenite::accept_async(stream).await.unwrap();
let _ = released.await;
drop(socket);
});
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
"action": "proxy",
"upstream_url": "ws://provider-dns.invalid/v1/responses",
"proxy": {"enabled": true, "url": format!("{scheme}://{proxy_addr}")}
}))
.unwrap();
decision.transport_profile = profile.clone();
let connection = tokio::time::timeout(
Duration::from_secs(5),
super::connect_upstream_websocket(
&decision,
crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS,
errors,
),
)
.await
.expect("proxied handshake must not wait for local provider DNS")
.unwrap_or_else(|error| panic!("{scheme} handshake failed: {error}"));
release.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(5), server)
.await
.unwrap()
.unwrap();
drop(connection);
}
}
}
#[test]
fn rejects_responses_websocket_frontdoor_self_loop_before_connecting() {
let base_url = configured_gateway_frontdoor_base_url();
@@ -1,5 +1,5 @@
use std::collections::BTreeMap;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use axum::{
@@ -10,6 +10,10 @@ use axum::{
};
use serde_json::json;
use crate::execution_runtime::transport::{
validate_execution_upstream_url, ExecutionSafeDnsResolver,
};
use super::test_connection_shared::select_test_connection_provider;
use super::{
provider_catalog_key_supports_format, query_param_value, AppState, GatewayPublicRequestContext,
@@ -18,98 +22,16 @@ use super::{
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
const MAX_TEST_CONNECTION_RESPONSE_BYTES: usize = 256 * 1024;
#[cfg(test)]
fn build_test_connection_client() -> Result<reqwest::Client, reqwest::Error> {
reqwest::Client::builder()
.no_proxy()
.dns_resolver(Arc::new(ExecutionSafeDnsResolver))
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(10))
.http2_adaptive_window(true)
.build()
}
#[derive(Debug)]
struct ResolvedTestConnectionTarget {
url: reqwest::Url,
host: String,
addresses: Vec<SocketAddr>,
}
/// Resolve the provider endpoint once and pin reqwest to that answer. The
/// test-connection route is reachable through the public front door, so it
/// must not perform an unbounded DNS lookup on every connect (which would
/// permit DNS rebinding into private/reserved networks).
async fn resolve_test_connection_target(
raw_url: &str,
allow_private_targets: bool,
) -> Result<ResolvedTestConnectionTarget, &'static str> {
let url = reqwest::Url::parse(raw_url).map_err(|_| "provider endpoint URL is invalid")?;
let literal_loopback = aether_http::url_has_literal_loopback_host(&url);
if !matches!(url.scheme(), "http" | "https")
|| url.host_str().is_none()
|| !url.username().is_empty()
|| url.password().is_some()
|| url.fragment().is_some()
{
return Err("provider endpoint must be an HTTP(S) URL without credentials or fragment");
}
let host = url
.host_str()
.ok_or("provider endpoint is missing a host")?
.to_string();
let literal_ip = host.parse::<IpAddr>().ok();
let port = url
.port_or_known_default()
.ok_or("provider endpoint is missing a port")?;
let addresses = if let Some(ip) = literal_ip {
vec![SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host.as_str(),
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| "provider endpoint DNS resolution failed")?
};
if addresses.is_empty() {
return Err("provider endpoint DNS resolution returned no addresses");
}
let has_private_answer = addresses
.iter()
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()));
// `allow_private_targets` is only enabled for in-process test fixtures.
// Keep that escape hatch narrowly scoped to literal loopback URLs whose
// every DNS answer is loopback; otherwise a test-only build (or an
// accidentally reused helper) could turn this public route into a
// private-network HTTP client.
let test_loopback_target = allow_private_targets
&& literal_loopback
&& addresses.iter().all(|address| address.ip().is_loopback());
if has_private_answer && !test_loopback_target {
return Err("provider endpoint resolves to a private or reserved address");
}
Ok(ResolvedTestConnectionTarget {
url,
host,
addresses,
})
}
fn build_pinned_test_connection_client(
target: &ResolvedTestConnectionTarget,
) -> Result<reqwest::Client, reqwest::Error> {
let mut builder = reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(Duration::from_secs(10))
.http2_adaptive_window(true);
if target.host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(&target.host, &target.addresses);
}
builder.build()
}
pub(super) async fn maybe_build_local_test_connection_route_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -384,18 +306,14 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
// Resolve and pin the endpoint before constructing the request. This
// keeps the public health-check route subject to the same DNS/SSRF
// boundary as the main execution transport. Unit-test fixtures may use
// loopback listeners; production requests never opt into private targets.
let target = match resolve_test_connection_target(&upstream_url, cfg!(test)).await {
Ok(target) => target,
let upstream_url = match validate_execution_upstream_url(&upstream_url) {
Ok(url) => url,
Err(reason) => {
tracing::warn!(
event_name = "provider_test_connection_target_rejected",
provider_id = %provider.id,
endpoint_id = %endpoint.id,
reason,
reason = %reason,
"provider connection test target was rejected"
);
return Some(
@@ -407,7 +325,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let test_client = match build_pinned_test_connection_client(&target) {
let test_client = match build_test_connection_client() {
Ok(client) => client,
Err(_) => {
return Some(
@@ -419,7 +337,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
);
}
};
let mut upstream_request = test_client.post(target.url);
let mut upstream_request = test_client.post(upstream_url);
for (name, value) in &provider_request_headers {
upstream_request = upstream_request.header(name, value);
}
@@ -495,7 +413,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
#[cfg(test)]
mod tests {
use super::{build_test_connection_client, resolve_test_connection_target};
use super::{build_test_connection_client, validate_execution_upstream_url};
use axum::{
body::Body,
http::{header, Request, StatusCode},
@@ -567,74 +485,59 @@ mod tests {
redirected_server.abort();
}
#[tokio::test]
async fn test_connection_target_rejects_private_addresses_in_production_mode() {
#[test]
fn test_connection_target_rejects_private_literals_like_provider_requests() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat/completions",
"http://10.0.0.1/v1/chat/completions",
"http://169.254.169.254/v1/chat/completions",
"https://10.0.0.1/v1/chat/completions",
"https://127.0.0.1/v1/chat/completions",
"https://[::1]/v1/chat/completions",
"https://localhost/v1/chat/completions",
"https://198.18.78.41/v1/chat/completions",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).is_err(),
"private provider target should be rejected: {raw_url}"
);
}
}
#[tokio::test]
async fn test_connection_target_accepts_public_http_and_https_addresses() {
for allow_private_targets in [false, true] {
for (raw_url, expected_port) in [
("http://8.8.8.8/v1/chat", 80),
("http://8.8.8.8:8080/v1/chat", 8080),
("https://8.8.8.8/v1/chat", 443),
] {
let target = resolve_test_connection_target(raw_url, allow_private_targets)
.await
.expect("public HTTP(S) provider target should resolve");
assert_eq!(target.url.as_str(), raw_url);
assert_eq!(target.host, "8.8.8.8");
assert_eq!(target.addresses.len(), 1);
assert_eq!(target.addresses[0].ip().to_string(), "8.8.8.8");
assert_eq!(target.addresses[0].port(), expected_port);
}
#[test]
fn test_connection_target_accepts_public_http_and_https_addresses() {
for (raw_url, expected_port) in [
("http://8.8.8.8/v1/chat", 80),
("http://8.8.8.8:8080/v1/chat", 8080),
("https://8.8.8.8/v1/chat", 443),
("https://[2606:4700:4700::1111]/v1/chat", 443),
] {
let url = validate_execution_upstream_url(raw_url)
.expect("public HTTP(S) provider target should be valid");
assert_eq!(url.as_str(), raw_url);
assert_eq!(url.port_or_known_default(), Some(expected_port));
}
}
#[tokio::test]
async fn test_connection_target_allows_loopback_only_for_test_fixtures() {
let target = resolve_test_connection_target("http://127.0.0.1:8080/v1/chat", true)
.await
.expect("test fixture target should resolve");
assert_eq!(target.host, "127.0.0.1");
assert_eq!(target.addresses.len(), 1);
assert!(
resolve_test_connection_target("http://10.0.0.1/v1/chat", true)
.await
.is_err(),
"test mode must not make private non-loopback HTTP endpoints acceptable"
);
assert!(
resolve_test_connection_target("https://10.0.0.1/v1/chat", true)
.await
.is_err(),
"test mode must not make private non-loopback endpoints acceptable"
);
assert!(
resolve_test_connection_target("http://localhost:8080/v1/chat", true)
.await
.is_ok(),
"literal localhost should remain available for local fixtures"
);
async fn test_connection_target_defers_dns_and_accepts_provider_loopback_urls() {
for raw_url in [
"http://127.0.0.1:8080/v1/chat",
"http://[::1]:8080/v1/chat",
"http://localhost:8080/v1/chat",
"https://provider-dns.invalid/v1/chat",
] {
let url = validate_execution_upstream_url(raw_url)
.expect("target validation must not depend on the current DNS answer");
let request = build_test_connection_client()
.expect("client should build without DNS")
.post(url)
.build()
.expect("provider request should build without DNS");
assert_eq!(request.url().as_str(), raw_url);
}
}
#[tokio::test]
async fn test_connection_target_rejects_url_credentials_and_fragments() {
#[test]
fn test_connection_target_rejects_url_credentials_and_fragments() {
for raw_url in [
"https://user:[email protected]/v1/chat",
"https://example.com/v1/chat#fragment",
@@ -643,9 +546,7 @@ mod tests {
"ftp://example.com/v1/chat",
] {
assert!(
resolve_test_connection_target(raw_url, false)
.await
.is_err(),
validate_execution_upstream_url(raw_url).is_err(),
"unsafe provider target should be rejected: {raw_url}"
);
}
@@ -168,52 +168,6 @@ fn wallet_public_refund_payload(mut payload: serde_json::Value) -> serde_json::V
payload
}
#[cfg(test)]
mod tests {
use super::wallet_refund_payload_from_record;
use aether_data::repository::wallet::StoredAdminWalletRefund;
use serde_json::json;
#[test]
fn public_refund_projection_excludes_payout_proof_and_upstream_payload() {
let record = StoredAdminWalletRefund {
id: "refund-1".to_string(),
refund_no: "rf_1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
payment_order_id: Some("order-1".to_string()),
source_type: "payment_order".to_string(),
source_id: Some("order-1".to_string()),
refund_mode: "original_channel".to_string(),
amount_usd: 10.0,
status: "processing".to_string(),
reason: Some("requested".to_string()),
failure_reason: None,
gateway_refund_id: Some("gateway-refund-1".to_string()),
payout_method: None,
payout_reference: None,
payout_proof: Some(json!({
"gateway_refund": {
"id": "gateway-refund-1",
"payload": {"payer": "sensitive", "credential": "secret"}
}
})),
requested_by: Some("user-1".to_string()),
approved_by: Some("admin-1".to_string()),
processed_by: Some("admin-1".to_string()),
created_at_unix_ms: 1,
updated_at_unix_secs: 1,
processed_at_unix_secs: Some(1),
completed_at_unix_secs: None,
};
let payload = wallet_refund_payload_from_record(&record);
assert!(payload.get("payout_proof").is_none());
assert_eq!(payload["status"], "processing");
assert_eq!(payload["gateway_refund_id"], "gateway-refund-1");
}
}
pub(super) async fn handle_wallet_refunds_list(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -659,3 +613,49 @@ pub(super) async fn handle_wallet_create_refund(
}
}
}
#[cfg(test)]
mod tests {
use super::wallet_refund_payload_from_record;
use aether_data::repository::wallet::StoredAdminWalletRefund;
use serde_json::json;
#[test]
fn public_refund_projection_excludes_payout_proof_and_upstream_payload() {
let record = StoredAdminWalletRefund {
id: "refund-1".to_string(),
refund_no: "rf_1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
payment_order_id: Some("order-1".to_string()),
source_type: "payment_order".to_string(),
source_id: Some("order-1".to_string()),
refund_mode: "original_channel".to_string(),
amount_usd: 10.0,
status: "processing".to_string(),
reason: Some("requested".to_string()),
failure_reason: None,
gateway_refund_id: Some("gateway-refund-1".to_string()),
payout_method: None,
payout_reference: None,
payout_proof: Some(json!({
"gateway_refund": {
"id": "gateway-refund-1",
"payload": {"payer": "sensitive", "credential": "secret"}
}
})),
requested_by: Some("user-1".to_string()),
approved_by: Some("admin-1".to_string()),
processed_by: Some("admin-1".to_string()),
created_at_unix_ms: 1,
updated_at_unix_secs: 1,
processed_at_unix_secs: Some(1),
completed_at_unix_secs: None,
};
let payload = wallet_refund_payload_from_record(&record);
assert!(payload.get("payout_proof").is_none());
assert_eq!(payload["status"], "processing");
assert_eq!(payload["gateway_refund_id"], "gateway-refund-1");
}
}
@@ -292,7 +292,7 @@ mod tests {
assert!(controls.is_err());
let template = "{{value}}".repeat(100_000);
let variables = BTreeMap::from([(String::from("value"), String::from("x".repeat(64)))]);
let variables = BTreeMap::from([(String::from("value"), "x".repeat(64))]);
let error = render_admin_email_template_html(&template, &variables)
.expect_err("rendered output must remain bounded");
assert!(format!("{error:?}").contains("exceeds"));
@@ -199,10 +199,7 @@ pub(crate) fn normalize_ldap_transport_server_url(raw: &str, use_starttls: bool)
// Gateway unit/integration fixtures use an in-process mock endpoint. Keep
// this exception behind the gateway test configuration; production code
// always uses the strict parser without custom schemes.
return aether_admin::system::normalize_ldap_transport_server_url_for_tests(
raw,
use_starttls,
);
aether_admin::system::normalize_ldap_transport_server_url_for_tests(raw, use_starttls)
}
#[cfg(not(test))]
{
+1
View File
@@ -71,6 +71,7 @@ mod rate_limit;
mod request_candidate_queue;
mod request_candidate_runtime;
mod request_diagnostics;
mod request_lifecycle;
mod roles;
mod router;
mod routing;
+1 -1
View File
@@ -44,7 +44,7 @@ pub(crate) fn local_auth_jwt_secret() -> Result<String, String> {
Err(std::env::VarError::NotPresent) => {
#[cfg(test)]
{
return Ok(TEST_JWT_SECRET.to_string());
Ok(TEST_JWT_SECRET.to_string())
}
#[cfg(not(test))]
@@ -3083,7 +3083,7 @@ mod tests {
.await
.expect("stale LKG read must not wait for retention lock");
assert_eq!(stale.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
assert_eq!(stale.stale_targets(), &[target.clone()]);
assert_eq!(stale.stale_targets(), std::slice::from_ref(&target));
assert_eq!(runtime.execution_count(), 1);
assert!(runtime
@@ -3111,7 +3111,7 @@ mod tests {
let load = load_one(&runtime, &client_version).await;
assert_eq!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
assert_eq!(load.stale_targets(), &[target.clone()]);
assert_eq!(load.stale_targets(), std::slice::from_ref(&target));
assert!(runtime
.state
.kv_get(&catalog_lkg_key(&target, client_version.as_str()))
@@ -3142,7 +3142,7 @@ mod tests {
let load = load_one(&runtime, &client_version).await;
assert_eq!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID), Some(&seeded));
assert_eq!(load.stale_targets(), &[target.clone()]);
assert_eq!(load.stale_targets(), std::slice::from_ref(&target));
assert_eq!(runtime.execution_count(), 1);
}
@@ -300,6 +300,34 @@ pub(crate) fn classify_local_failover(
policy: &LocalFailoverPolicy,
input: LocalFailoverInput<'_>,
) -> LocalFailoverClassification {
if input.status_code >= 400
&& policy.routing_rules.error_stop_patterns.iter().any(|rule| {
failover_pattern_matches(
&rule.pattern,
&rule.status_codes,
input.response_text,
input.status_code,
)
})
{
return LocalFailoverClassification::StopErrorPattern;
}
if input.status_code == 200
&& policy
.routing_rules
.success_failover_patterns
.iter()
.any(|rule| {
failover_pattern_matches(
&rule.pattern,
&rule.status_codes,
input.response_text,
input.status_code,
)
})
{
return LocalFailoverClassification::RetrySuccessPattern;
}
if policy.stop_status_codes.contains(&input.status_code) {
return LocalFailoverClassification::StopStatusCode;
}
@@ -487,13 +515,27 @@ fn local_failover_regex_rule_matches(
response_text: Option<&str>,
status_code: u16,
) -> bool {
if !rule.status_codes.is_empty() && !rule.status_codes.contains(&status_code) {
failover_pattern_matches(
&rule.pattern,
&rule.status_codes,
response_text,
status_code,
)
}
fn failover_pattern_matches(
pattern: &str,
status_codes: &std::collections::BTreeSet<u16>,
response_text: Option<&str>,
status_code: u16,
) -> bool {
if !status_codes.is_empty() && !status_codes.contains(&status_code) {
return false;
}
let pattern = rule.pattern.trim();
let pattern = pattern.trim();
if pattern.is_empty() {
return !rule.status_codes.is_empty();
return !status_codes.is_empty();
}
let Some(response_text) = response_text else {
@@ -507,6 +549,71 @@ fn local_failover_regex_rule_matches(
#[cfg(test)]
mod tests {
#[test]
fn routing_rules_precede_provider_rules_and_keep_provider_fallback() {
let policy = super::LocalFailoverPolicy {
routing_rules: aether_routing_core::RoutingFailoverRules {
success_failover_patterns: vec![aether_routing_core::RoutingFailoverRule {
pattern: "(?i)capacity.*exhausted".to_string(),
..Default::default()
}],
error_stop_patterns: vec![aether_routing_core::RoutingFailoverRule {
pattern: "invalid.*parameter".to_string(),
status_codes: [400].into_iter().collect(),
}],
},
stop_status_codes: [200, 403].into_iter().collect(),
continue_status_codes: [400].into_iter().collect(),
..Default::default()
};
for (status, body, expected) in [
(
200,
"CAPACITY exhausted",
super::LocalFailoverClassification::RetrySuccessPattern,
),
(
400,
"invalid request parameter",
super::LocalFailoverClassification::StopErrorPattern,
),
(
400,
"capacity exhausted",
super::LocalFailoverClassification::RetryStatusCode,
),
(
403,
"permission denied",
super::LocalFailoverClassification::StopStatusCode,
),
(
429,
"rate limited",
super::LocalFailoverClassification::RetryUpstreamFailure,
),
] {
assert_eq!(
super::classify_local_failover(
&policy,
super::LocalFailoverInput::new(status, Some(body))
),
expected
);
}
}
#[test]
fn provider_transport_stop_rule_is_respected() {
let policy = super::LocalFailoverPolicy {
stop_on_transport_errors: true,
..Default::default()
};
assert_eq!(
super::classify_local_transport_error(&policy),
super::LocalTransportFailoverClassification::StopTransportError
);
}
use std::collections::BTreeSet;
use super::{
@@ -4,7 +4,7 @@ use aether_contracts::ExecutionPlan;
use serde_json::{json, Value};
use tracing::debug;
use aether_routing_core::RoutingExecutionPolicy;
use aether_routing_core::{RoutingExecutionPolicy, RoutingFailoverRules};
use crate::provider_transport::GatewayProviderTransportSnapshot;
use crate::AppState;
@@ -14,6 +14,7 @@ pub(crate) const ROUTING_EXECUTION_POLICY_REPORT_FIELD: &str = "routing_executio
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct LocalFailoverPolicy {
pub(crate) routing_rules: RoutingFailoverRules,
pub(crate) max_retries: Option<u64>,
pub(crate) max_transfer_count: u64,
pub(crate) max_transfer_timeout_seconds: u64,
@@ -29,6 +30,7 @@ pub(crate) struct LocalFailoverPolicy {
impl Default for LocalFailoverPolicy {
fn default() -> Self {
Self {
routing_rules: RoutingFailoverRules::default(),
max_retries: None,
max_transfer_count: 0,
max_transfer_timeout_seconds: 0,
@@ -61,8 +63,10 @@ pub(crate) async fn resolve_local_failover_policy(
Ok(Some(transport)) => local_failover_policy_from_transport(&transport),
Ok(None) | Err(_) => LocalFailoverPolicy::default(),
};
let cyber_continue_failover = routing_execution_policy_from_report_context(report_context)
.is_some_and(|policy| policy.cyber_continue_failover);
let routing_policy =
routing_execution_policy_from_report_context(report_context).unwrap_or_default();
let cyber_continue_failover = routing_policy.cyber_continue_failover;
policy.routing_rules = routing_policy.failover_rules;
policy.stop_cyber_policy_errors = !cyber_continue_failover;
debug!(
event_name = "local_failover_policy_loaded",
@@ -80,6 +84,8 @@ pub(crate) async fn resolve_local_failover_policy(
stop_on_transport_errors = policy.stop_on_transport_errors,
success_failover_pattern_count = policy.success_failover_patterns.len(),
error_stop_pattern_count = policy.error_stop_patterns.len(),
global_success_pattern_count = policy.routing_rules.success_failover_patterns.len(),
global_stop_pattern_count = policy.routing_rules.error_stop_patterns.len(),
cyber_continue_failover,
"gateway loaded local failover policy from transport snapshot"
);
@@ -122,6 +128,7 @@ pub(crate) fn local_failover_policy_from_transport(
});
LocalFailoverPolicy {
routing_rules: RoutingFailoverRules::default(),
max_retries,
max_transfer_count: provider_config
.and_then(|value| value.get("max_transfer_count"))
@@ -184,6 +191,10 @@ pub(crate) fn local_failover_policy_from_report_context(
.as_object()?;
Some(LocalFailoverPolicy {
routing_rules: object
.get("routing_rules")
.and_then(|value| serde_json::from_value(value.clone()).ok())
.unwrap_or_default(),
max_retries: object.get("max_retries").and_then(parse_u64_value),
max_transfer_count: object
.get("max_transfer_count")
@@ -267,6 +278,7 @@ fn parse_status_code_list(value: &Value) -> BTreeSet<u16> {
fn local_failover_policy_to_value(policy: &LocalFailoverPolicy) -> Value {
json!({
"routing_rules": policy.routing_rules,
"max_retries": policy.max_retries,
"max_transfer_count": policy.max_transfer_count,
"max_transfer_timeout_seconds": policy.max_transfer_timeout_seconds,
@@ -525,6 +537,7 @@ mod tests {
assert_eq!(
local_failover_policy_from_report_context(Some(&report_context)),
Some(LocalFailoverPolicy {
routing_rules: Default::default(),
max_retries: Some(2),
max_transfer_count: 10,
max_transfer_timeout_seconds: 60,
@@ -3764,18 +3764,19 @@ mod tests {
.await;
assert_eq!(normal_batch.len(), 1);
let retry_states = metrics
.retry_states
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
assert_eq!(
retry_states
.get(&(0, RequestCandidateQueueLane::Normal))
.map(|state| state.attempt),
Some(1)
);
assert!(!retry_states.contains_key(&(0, RequestCandidateQueueLane::Active)));
drop(retry_states);
{
let retry_states = metrics
.retry_states
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
assert_eq!(
retry_states
.get(&(0, RequestCandidateQueueLane::Normal))
.map(|state| state.attempt),
Some(1)
);
assert!(!retry_states.contains_key(&(0, RequestCandidateQueueLane::Active)));
}
assert!(request_candidate_retry_is_ready(
&metrics,
0,
@@ -0,0 +1,336 @@
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll};
use aether_routing_core::RoutingExecutionPolicy;
use axum::body::{Body, Bytes, HttpBody};
use http::Response;
use http_body::{Frame, SizeHint};
use http_body_util::BodyExt;
use crate::request_diagnostics::{scope_request_diagnostics_with, RequestDiagnostics};
use crate::GatewayError;
tokio::task_local! {
static CANCEL_ON_CLIENT_DISCONNECT: Arc<AtomicBool>;
}
pub(crate) fn configure_client_disconnect(policy: RoutingExecutionPolicy) {
let _ = CANCEL_ON_CLIENT_DISCONNECT.try_with(|cancel| {
cancel.store(policy.cancel_on_client_disconnect, Ordering::Release);
});
}
pub(crate) fn cancel_on_client_disconnect() -> bool {
CANCEL_ON_CLIENT_DISCONNECT
.try_with(|cancel| cancel.load(Ordering::Acquire))
.unwrap_or(false)
}
pub(crate) async fn run_request<F>(future: F) -> Result<Response<Body>, GatewayError>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
let cancel = Arc::new(AtomicBool::new(true));
let diagnostics = Arc::new(RequestDiagnostics::default());
let cancel_for_response = Arc::clone(&cancel);
let future = CANCEL_ON_CLIENT_DISCONNECT.scope(
Arc::clone(&cancel),
scope_request_diagnostics_with(Some(Arc::clone(&diagnostics)), async move {
let response = future.await?;
if cancel_for_response.load(Ordering::Acquire) {
return Ok(response);
}
Ok(response.map(|body| {
Body::new(CompleteOnDisconnectBody {
body: Some(body),
diagnostics,
})
}))
}),
);
CompleteOnDisconnectRequest {
future: Some(Box::pin(future)),
cancel,
}
.await
}
struct CompleteOnDisconnectRequest<F>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
future: Option<Pin<Box<F>>>,
cancel: Arc<AtomicBool>,
}
impl<F> Future for CompleteOnDisconnectRequest<F>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
type Output = Result<Response<Body>, GatewayError>;
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
let result = self
.future
.as_mut()
.expect("request future")
.as_mut()
.poll(context);
if result.is_ready() {
self.future.take();
}
result
}
}
impl<F> Drop for CompleteOnDisconnectRequest<F>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
fn drop(&mut self) {
if self.cancel.load(Ordering::Acquire) {
return;
}
if let (Some(future), Ok(runtime)) =
(self.future.take(), tokio::runtime::Handle::try_current())
{
runtime.spawn(async move {
if let Ok(response) = future.await {
drain_body(response.into_body()).await;
}
});
}
}
}
struct CompleteOnDisconnectBody {
body: Option<Body>,
diagnostics: Arc<RequestDiagnostics>,
}
impl HttpBody for CompleteOnDisconnectBody {
type Data = Bytes;
type Error = axum::Error;
fn poll_frame(
mut self: Pin<&mut Self>,
context: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
let Some(body) = self.body.as_mut() else {
return Poll::Ready(None);
};
let result = Pin::new(body).poll_frame(context);
if matches!(result, Poll::Ready(None | Some(Err(_)))) {
self.body.take();
}
result
}
fn is_end_stream(&self) -> bool {
self.body.as_ref().is_none_or(HttpBody::is_end_stream)
}
fn size_hint(&self) -> SizeHint {
self.body
.as_ref()
.map(HttpBody::size_hint)
.unwrap_or_else(|| SizeHint::with_exact(0))
}
}
impl Drop for CompleteOnDisconnectBody {
fn drop(&mut self) {
let Some(body) = self.body.take().filter(|body| !body.is_end_stream()) else {
return;
};
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
runtime.spawn(scope_request_diagnostics_with(
Some(Arc::clone(&self.diagnostics)),
drain_body(body),
));
}
}
}
async fn drain_body(mut body: Body) {
while let Some(frame) = body.frame().await {
if frame.is_err() {
break;
}
}
}
#[cfg(test)]
mod tests {
use std::io;
use std::time::Duration;
use futures_util::stream;
use http::HeaderMap;
use http_body_util::StreamBody;
use tokio::sync::{mpsc, oneshot};
use super::*;
#[tokio::test]
async fn disconnected_request_finishes_and_keeps_admission_and_diagnostics() {
let gate = aether_runtime::ConcurrencyGate::new("disconnect_request", 1);
let permit = gate.try_acquire().unwrap();
let (started_tx, started_rx) = oneshot::channel();
let (release_tx, release_rx) = oneshot::channel();
let (finished_tx, finished_rx) = oneshot::channel();
let request = tokio::spawn(run_request(async move {
let _permit = permit;
configure_client_disconnect(RoutingExecutionPolicy::default());
started_tx.send(()).unwrap();
release_rx.await.unwrap();
assert!(crate::request_diagnostics::current_request_diagnostics().is_some());
finished_tx.send(()).unwrap();
Ok(Response::new(Body::empty()))
}));
started_rx.await.unwrap();
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert_eq!(gate.snapshot().in_flight, 1);
release_tx.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(1), finished_rx)
.await
.unwrap()
.unwrap();
assert_eq!(gate.snapshot().in_flight, 0);
}
#[tokio::test]
async fn enabled_cancellation_and_unresolved_requests_drop_immediately() {
for resolve_policy in [false, true] {
let (started_tx, started_rx) = oneshot::channel();
let (release_tx, release_rx) = oneshot::channel::<()>();
let request = tokio::spawn(run_request(async move {
if resolve_policy {
configure_client_disconnect(RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
});
}
started_tx.send(()).unwrap();
release_rx.await.unwrap();
Ok(Response::new(Body::empty()))
}));
started_rx.await.unwrap();
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert!(release_tx.send(()).is_err());
}
}
#[tokio::test]
async fn disconnected_body_drains_without_buffering_and_holds_admission() {
for consume_first_chunk in [false, true] {
let gate = aether_runtime::ConcurrencyGate::new("disconnect_body", 1);
let permit = gate.try_acquire().unwrap();
let (sender, receiver) = mpsc::channel(1);
let (finished_tx, finished_rx) = oneshot::channel();
let response = run_request(async move {
configure_client_disconnect(RoutingExecutionPolicy::default());
let body = Body::from_stream(stream::unfold(
(receiver, finished_tx, permit),
|(mut receiver, finished_tx, permit)| async move {
match receiver.recv().await {
Some(bytes) => {
Some((Ok::<_, io::Error>(bytes), (receiver, finished_tx, permit)))
}
None => {
assert!(crate::request_diagnostics::current_request_diagnostics()
.is_some());
finished_tx.send(()).unwrap();
None
}
}
},
));
Ok(Response::new(body))
})
.await
.unwrap();
let mut body = response.into_body();
if consume_first_chunk {
sender.send(Bytes::from_static(b"first")).await.unwrap();
assert_eq!(
body.frame().await.unwrap().unwrap().into_data().unwrap(),
"first"
);
}
drop(body);
assert_eq!(gate.snapshot().in_flight, 1);
tokio::time::timeout(Duration::from_secs(1), async {
for _ in 0..100 {
sender.send(Bytes::from_static(b"remaining")).await.unwrap();
}
drop(sender);
finished_rx.await.unwrap();
})
.await
.unwrap();
assert_eq!(gate.snapshot().in_flight, 0);
}
}
#[tokio::test]
async fn enabled_cancellation_drops_stream_receiver() {
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
let response = run_request(async move {
configure_client_disconnect(RoutingExecutionPolicy {
cancel_on_client_disconnect: true,
..Default::default()
});
Ok(Response::new(Body::from_stream(stream::unfold(
receiver,
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
))))
})
.await
.unwrap();
drop(response);
assert!(sender.is_closed());
}
#[tokio::test]
async fn connected_response_preserves_headers_size_hint_and_trailers() {
let response = run_request(async {
configure_client_disconnect(RoutingExecutionPolicy::default());
Ok(Response::builder()
.status(201)
.header("x-test", "unchanged")
.body(Body::from("hello"))
.unwrap())
})
.await
.unwrap();
assert_eq!(response.status(), 201);
assert_eq!(response.headers()["x-test"], "unchanged");
assert_eq!(response.body().size_hint().exact(), Some(5));
assert_eq!(
response.into_body().collect().await.unwrap().to_bytes(),
"hello"
);
let mut trailers = HeaderMap::new();
trailers.insert("x-finished", "yes".parse().unwrap());
let response = run_request(async move {
configure_client_disconnect(RoutingExecutionPolicy::default());
let frames = stream::iter([
Ok::<_, io::Error>(Frame::data(Bytes::from_static(b"hello"))),
Ok(Frame::trailers(trailers)),
]);
Ok(Response::new(Body::new(StreamBody::new(frames))))
})
.await
.unwrap();
let collected = response.into_body().collect().await.unwrap();
assert_eq!(collected.trailers().unwrap()["x-finished"], "yes");
assert_eq!(collected.to_bytes(), "hello");
}
}
+31 -32
View File
@@ -57,7 +57,7 @@ pub(crate) fn resolve_gateway_routing_policy(
let config = serde_json::from_value::<RoutingGroupConfig>(input.group_config_json.clone())
.map_err(|_| invalid_routing_group_config())?;
resolve_routing_policy(
let policy = resolve_routing_policy(
&config,
RoutingPolicyInput {
group_id: input.group_id,
@@ -73,7 +73,9 @@ pub(crate) fn resolve_gateway_routing_policy(
phase: input.phase,
},
)
.map_err(routing_policy_error)
.map_err(routing_policy_error)?;
crate::request_lifecycle::configure_client_disconnect(policy.execution_policy.clone());
Ok(policy)
}
pub(crate) fn resolve_gateway_static_default_routing_policy(
@@ -82,6 +84,7 @@ pub(crate) fn resolve_gateway_static_default_routing_policy(
let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else {
return Ok(None);
};
crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy.clone());
Ok(Some(ResolvedRoutingPolicy {
group_id: input.group_id.map(str::to_string),
@@ -142,28 +145,11 @@ fn static_default_policy_fields(
.ok_or_else(invalid_routing_group_config)?,
None => DEFAULT_STICKY_KEY_ATTEMPTS,
};
let enable_cf_heartbeat = routing_bool_field(
default_policy.get("enable_cf_heartbeat"),
"enable_cf_heartbeat",
)?;
// Older strategies stored separate image/text heartbeat flags. Treat
// either legacy flag as enabling the unified CF heartbeat setting while
// allowing newly saved strategies to use only the canonical key.
let legacy_image_heartbeat = routing_bool_field(
default_policy.get("enable_openai_image_sync_heartbeat"),
"enable_openai_image_sync_heartbeat",
)?;
let legacy_text_heartbeat = routing_bool_field(
default_policy.get("enable_standard_text_sync_heartbeat"),
"enable_standard_text_sync_heartbeat",
)?;
let execution_policy = aether_routing_core::RoutingExecutionPolicy {
enable_cf_heartbeat: enable_cf_heartbeat || legacy_image_heartbeat || legacy_text_heartbeat,
cyber_continue_failover: routing_bool_field(
default_policy.get("cyber_continue_failover"),
"cyber_continue_failover",
)?,
};
let execution_policy: aether_routing_core::RoutingExecutionPolicy =
serde_json::from_value(Value::Object(default_policy.clone()))
.map_err(|_| invalid_routing_group_config())?;
aether_routing_core::validate_routing_failover_rules(&execution_policy.failover_rules)
.map_err(|_| invalid_routing_group_config())?;
Ok(Some(RoutingDefaultPolicy {
priority_mode,
@@ -174,13 +160,6 @@ fn static_default_policy_fields(
}))
}
fn routing_bool_field(value: Option<&Value>, _field: &str) -> Result<bool, GatewayError> {
match value {
Some(value) => value.as_bool().ok_or_else(invalid_routing_group_config),
None => Ok(false),
}
}
fn routing_array_field_is_missing_or_empty(
object: &serde_json::Map<String, Value>,
key: &str,
@@ -238,7 +217,14 @@ mod tests {
"default_policy": {
"priority_mode": "global_key",
"scheduling_mode": "load_balance",
"keep_priority_on_conversion": true
"keep_priority_on_conversion": true,
"cancel_on_client_disconnect": true,
"max_transfer_count": 3,
"max_transfer_timeout_seconds": 90,
"failover_rules": {
"success_failover_patterns": [{"pattern": "(?i)capacity.*exhausted"}],
"error_stop_patterns": [{"status_codes": [400]}]
}
},
"allowed_models": ["legacy-model"],
"model_policies": [],
@@ -274,6 +260,19 @@ mod tests {
.expect("full policy should resolve");
assert_eq!(static_policy, full_policy);
assert_eq!(static_policy.execution_policy.max_transfer_count, 3);
assert_eq!(
static_policy.execution_policy.max_transfer_timeout_seconds,
90
);
assert_eq!(
static_policy
.execution_policy
.failover_rules
.error_stop_patterns
.len(),
1
);
assert_eq!(
static_policy.priority_mode,
RoutingSetPriorityMode::GlobalKey
+14 -2
View File
@@ -274,7 +274,13 @@ mod tests {
"priority_mode": "provider",
"scheduling_mode": "cache_affinity",
"keep_priority_on_conversion": false,
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS,
"max_transfer_count": 0,
"max_transfer_timeout_seconds": 0,
"failover_rules": {
"success_failover_patterns": [],
"error_stop_patterns": []
}
})
);
@@ -323,7 +329,13 @@ mod tests {
"priority_mode": "provider",
"scheduling_mode": "cache_affinity",
"keep_priority_on_conversion": false,
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS
"sticky_key_attempts": DEFAULT_STICKY_KEY_ATTEMPTS,
"max_transfer_count": 0,
"max_transfer_timeout_seconds": 0,
"failover_rules": {
"success_failover_patterns": [],
"error_stop_patterns": []
}
})
);
}
@@ -196,11 +196,12 @@ impl AppState {
{
let session = session.into();
#[cfg(test)]
if self.auth_session_store.is_some() && self.auth_user_store.is_some() {
if let (Some(session_store), Some(user_store)) = (
self.auth_session_store.as_ref(),
self.auth_user_store.as_ref(),
) {
let existing = {
self.auth_user_store
.as_ref()
.expect("checked auth user store")
user_store
.lock()
.expect("auth user store should lock")
.get(&session.user_id)
@@ -217,12 +218,7 @@ impl AppState {
let Some(existing) = existing else {
return Ok(None);
};
let mut users = self
.auth_user_store
.as_ref()
.expect("checked auth user store")
.lock()
.expect("auth user store should lock");
let mut users = user_store.lock().expect("auth user store should lock");
let user = users.entry(session.user_id.clone()).or_insert(existing);
if user.password_hash.as_deref() != Some(expected_password_hash)
|| !user.auth_source.eq_ignore_ascii_case("local")
@@ -238,10 +234,7 @@ impl AppState {
.or(session.last_seen_at)
.unwrap_or_else(chrono::Utc::now);
user.last_login_at = Some(now);
let mut sessions = self
.auth_session_store
.as_ref()
.expect("checked auth session store")
let mut sessions = session_store
.lock()
.expect("auth session store should lock");
for existing in sessions.values_mut() {
@@ -802,7 +802,7 @@ impl AppState {
}
return Ok(Some(LdapAuthProvisioningResult {
user,
owned_wallet_id: initialized.created.then(|| initialized.wallet.id),
owned_wallet_id: initialized.created.then_some(initialized.wallet.id),
}));
}
@@ -1,7 +1,7 @@
use super::{
any, build_router_with_state, build_state_with_execution_runtime_override, json, start_server,
Arc, Body, Bytes, HeaderValue, Infallible, Json, Mutex, Request, Response, Router, StatusCode,
TRACE_ID_HEADER,
LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER,
};
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
@@ -55,6 +55,24 @@ fn hash_api_key(value: &str) -> String {
format!("{:x}", hasher.finalize())
}
async fn build_cancelling_gateway(state: crate::AppState) -> Router {
state
.data
.update_routing_group(
"system-default",
aether_data_contracts::repository::routing_profiles::UpdateRoutingGroupRecord {
config_json: Some(json!({"default_policy": {"cancel_on_client_disconnect": true}})),
version: Some(2),
updated_at: 2,
..Default::default()
},
)
.await
.expect("routing policy should update")
.expect("default strategy should exist");
build_router_with_state(state)
}
fn sample_local_openai_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
StoredAuthApiKeySnapshot::new(
user_id.to_string(),
@@ -384,7 +402,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
vec![sample_local_openai_key()],
));
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let gateway = build_router_with_state(
let gateway = build_cancelling_gateway(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
@@ -395,7 +413,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
DEVELOPMENT_ENCRYPTION_KEY,
),
),
);
).await;
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
@@ -467,7 +485,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
vec![sample_local_openai_key()],
));
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let gateway = build_router_with_state(
let gateway = build_cancelling_gateway(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
@@ -478,7 +496,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
DEVELOPMENT_ENCRYPTION_KEY,
),
),
);
).await;
let (gateway_url, gateway_handle) = start_server(gateway).await;
let request = reqwest::Client::new()
@@ -608,7 +626,7 @@ async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error_
auth_repository,
candidate_selection_repository,
provider_catalog_repository,
request_candidate_repository,
Arc::clone(&request_candidate_repository),
DEVELOPMENT_ENCRYPTION_KEY,
),
),
@@ -631,17 +649,41 @@ async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error_
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
response
.headers()
.get(http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok()),
Some("text/event-stream")
Some("application/json")
);
let body_text = response.text().await.expect("response body should read");
assert!(body_text.contains("\"rate_limit_error\""));
assert!(body_text.contains("\"slow down\""));
assert_eq!(
response
.headers()
.get(LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER)
.and_then(|value| value.to_str().ok()),
Some("execution_runtime_candidates_exhausted")
);
let body_json: serde_json::Value = response.json().await.expect("response body should parse");
assert_eq!(body_json["error"]["type"], "http_error");
let stored_candidates = request_candidate_repository
.list_by_request_id("trace-openai-chat-stream-prefetch-error-123")
.await
.expect("request candidate trace should read");
let failed_candidate = stored_candidates
.iter()
.find(|candidate| candidate.status == RequestCandidateStatus::Failed)
.expect("prefetched error should mark the attempted candidate as failed");
assert!(stored_candidates
.iter()
.all(|candidate| candidate.status != RequestCandidateStatus::Success));
assert_eq!(failed_candidate.status_code, Some(429));
assert_eq!(
failed_candidate.error_type.as_deref(),
Some("rate_limit_error")
);
assert_eq!(failed_candidate.error_message.as_deref(), Some("slow down"));
assert!(failed_candidate.finished_at_unix_ms.is_some());
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
@@ -349,7 +349,7 @@ async fn gateway_executes_gemini_chat_stream_via_local_decision_gate_with_local_
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":33,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -419,7 +419,7 @@ async fn gateway_executes_gemini_chat_stream_via_local_decision_gate_with_local_
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_execution_runtime_request = seen_execution_runtime
@@ -326,7 +326,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_with_local_s
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -396,7 +396,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_with_local_s
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_execution_runtime_request = seen_execution_runtime
@@ -846,7 +846,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"response\\\":{\\\"candidates\\\":[]},\\\"remainingCredits\\\":42,\\\"consumedCredits\\\":1,\\\"traceId\\\":\\\"trace-upstream-1\\\"}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"response\\\":{\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]},\\\"remainingCredits\\\":42,\\\"consumedCredits\\\":1,\\\"traceId\\\":\\\"trace-upstream-1\\\"}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -934,7 +934,7 @@ async fn gateway_executes_gemini_cli_stream_via_local_decision_gate_after_oauth_
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_refresh_request = seen_refresh
@@ -1354,7 +1354,7 @@ async fn gateway_executes_vertex_ai_gemini_cli_stream_via_local_decision_gate_wi
});
let frames = concat!(
"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[]}\\n\\n\"}}\n",
"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"candidates\\\":[{\\\"content\\\":{\\\"parts\\\":[{\\\"text\\\":\\\"ok\\\"}]},\\\"finishReason\\\":\\\"STOP\\\"}]}\\n\\n\"}}\n",
"{\"type\":\"telemetry\",\"payload\":{\"kind\":\"telemetry\",\"telemetry\":{\"elapsed_ms\":34,\"upstream_bytes\":26}}}\n",
"{\"type\":\"eof\",\"payload\":{\"kind\":\"eof\"}}\n"
);
@@ -1422,7 +1422,7 @@ async fn gateway_executes_vertex_ai_gemini_cli_stream_via_local_decision_gate_wi
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
"data: {\"candidates\":[]}\n\n"
"data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"ok\"}]},\"finishReason\":\"STOP\"}]}\n\n"
);
let seen_execution_runtime_request = seen_execution_runtime
@@ -96,12 +96,11 @@ fn aether_data_backend_pool_modules_do_not_own_maintenance_sql() {
#[test]
fn wallet_maintenance_sql_is_partitioned_by_driver() {
let facade = read_workspace_file("crates/aether-data/runtime/src/backend/wallet.rs");
for module in ["mod postgres;"] {
assert!(
facade.contains(module),
"wallet facade should declare {module}"
);
}
let module = "mod postgres;";
assert!(
facade.contains(module),
"wallet facade should declare {module}"
);
for forbidden in [
"sqlx::",
"PostgresBackend",
@@ -115,24 +114,22 @@ fn wallet_maintenance_sql_is_partitioned_by_driver() {
);
}
for (driver, backend) in [("postgres", "PostgresBackend")] {
let path = format!("crates/aether-data/runtime/src/backend/wallet/{driver}.rs");
let source = read_workspace_file(&path);
assert!(source.contains(&format!("impl {backend}")));
assert!(source.contains("aggregate_wallet_daily_usage"));
assert!(source.contains("sqlx::query"));
}
let (driver, backend) = ("postgres", "PostgresBackend");
let path = format!("crates/aether-data/runtime/src/backend/wallet/{driver}.rs");
let source = read_workspace_file(&path);
assert!(source.contains(&format!("impl {backend}")));
assert!(source.contains("aggregate_wallet_daily_usage"));
assert!(source.contains("sqlx::query"));
}
#[test]
fn table_maintenance_is_partitioned_for_each_driver() {
let facade = read_workspace_file("crates/aether-data/runtime/src/backend/maintenance.rs");
for module in ["mod postgres;"] {
assert!(
facade.contains(module),
"maintenance facade should declare {module}"
);
}
let module = "mod postgres;";
assert!(
facade.contains(module),
"maintenance facade should declare {module}"
);
for forbidden in [
"impl PostgresBackend",
"VACUUM ANALYZE",
@@ -156,12 +153,11 @@ fn table_maintenance_is_partitioned_for_each_driver() {
#[test]
fn system_driver_database_operations_are_partitioned() {
let facade = read_workspace_file("crates/aether-data/runtime/src/backend/system.rs");
for module in ["mod postgres;"] {
assert!(
facade.contains(module),
"system facade should declare {module}"
);
}
let module = "mod postgres;";
assert!(
facade.contains(module),
"system facade should declare {module}"
);
for forbidden in [
"impl PostgresBackend",
"fn map_postgres_stats_daily_aggregate(",
@@ -1521,12 +1517,11 @@ fn lifecycle_migrations_are_partitioned_by_driver() {
types.contains(required),
"migrate/types.rs should own {required}"
);
for forbidden in ["PgPool"] {
assert!(
!types.contains(forbidden),
"migrate/types.rs should remain driver-independent from {forbidden}"
);
}
let forbidden = "PgPool";
assert!(
!types.contains(forbidden),
"migrate/types.rs should remain driver-independent from {forbidden}"
);
let postgres =
read_workspace_file("crates/aether-data/runtime/src/lifecycle/migrate/postgres.rs");
@@ -1905,12 +1900,11 @@ fn gateway_system_config_types_are_owned_by_aether_data() {
}
let data_backends =
read_workspace_file("crates/aether-data/runtime/src/backend/maintenance.rs");
for pattern in ["postgres.list_system_config_entries().await"] {
assert!(
data_backends.contains(pattern),
"aether-data backends should own driver-specific system config dispatch {pattern}"
);
}
let pattern = "postgres.list_system_config_entries().await";
assert!(
data_backends.contains(pattern),
"aether-data backends should own driver-specific system config dispatch {pattern}"
);
for pattern in [
"|(key, value, description, updated_at_unix_secs)|",
"Ok((0, 0, 0, 0))",
@@ -3705,9 +3705,9 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p
scores[0].hard_state.schedulable(),
"OAuth completion should replace AuthInvalid with a schedulable score"
);
let decrypted_api_key = decrypt_persisted_provider_api_key(&persisted);
let decrypted_api_key = decrypt_persisted_provider_api_key(persisted);
assert_eq!(decrypted_api_key, "new-codex-access-token");
let decrypted_auth_config = decrypt_persisted_provider_auth_config(&persisted);
let decrypted_auth_config = decrypt_persisted_provider_auth_config(persisted);
let auth_config: serde_json::Value =
serde_json::from_str(&decrypted_auth_config).expect("auth config json should parse");
assert_eq!(auth_config["provider_type"], "codex");
@@ -3834,12 +3834,12 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
Some(&Value::Null)
);
assert_eq!(
decrypt_test_provider_catalog_credential(&key, ProviderCatalogCredentialField::ApiKey,),
decrypt_test_provider_catalog_credential(key, ProviderCatalogCredentialField::ApiKey,),
"oauth-access-token-new"
);
let auth_config =
decrypt_test_provider_catalog_credential(&key, ProviderCatalogCredentialField::AuthConfig);
decrypt_test_provider_catalog_credential(key, ProviderCatalogCredentialField::AuthConfig);
let auth_config: Value =
serde_json::from_str(&auth_config).expect("oauth auth config json should parse");
assert_eq!(auth_config["provider_type"], "codex");
@@ -2659,9 +2659,9 @@ fn set_test_env_var(key: &'static str, value: &str) -> TestEnvVarGuard {
}
#[cfg(test)]
fn payment_callback_env_lock() -> &'static std::sync::Mutex<()> {
static LOCK: std::sync::OnceLock<std::sync::Mutex<()>> = std::sync::OnceLock::new();
LOCK.get_or_init(|| std::sync::Mutex::new(()))
fn payment_callback_env_lock() -> &'static tokio::sync::Mutex<()> {
static LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
&LOCK
}
const TEST_PAYMENT_CALLBACK_SECRET: &str = "test-callback-secret-0123456789abcdef";
@@ -11875,9 +11875,7 @@ async fn gateway_does_not_report_logout_success_when_session_revoke_is_rejected(
#[tokio::test]
async fn gateway_handles_payment_callback_route_locally_without_proxying_upstream() {
let _env_lock = payment_callback_env_lock()
.lock()
.expect("payment callback test env lock should not be poisoned");
let _env_lock = payment_callback_env_lock().lock().await;
let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", TEST_PAYMENT_CALLBACK_SECRET);
let now = Utc::now();
let user = StoredUserAuthRecord::new(
@@ -12020,9 +12018,7 @@ async fn gateway_handles_payment_callback_route_locally_without_proxying_upstrea
#[tokio::test]
async fn gateway_rejects_payment_callback_with_mismatched_payment_method_locally() {
let _env_lock = payment_callback_env_lock()
.lock()
.expect("payment callback test env lock should not be poisoned");
let _env_lock = payment_callback_env_lock().lock().await;
let _secret_guard = set_test_env_var("PAYMENT_CALLBACK_SECRET", TEST_PAYMENT_CALLBACK_SECRET);
let now = Utc::now();
let user = StoredUserAuthRecord::new(
+1 -1
View File
@@ -1076,7 +1076,7 @@ async fn gateway_forwards_public_request_to_remote_tunnel_owner_before_fallback_
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider_with_request_timeout(
"provider-owner",
Some(0.1),
Some(2.0),
)],
vec![sample_endpoint("endpoint-owner", "provider-owner")],
vec![sample_bound_key(
+1 -1
View File
@@ -1100,7 +1100,7 @@ async fn sync_transport_error_policy_stops_or_retries_candidates_end_to_end_impl
let mut second_candidate = sample_local_openai_candidate_row();
second_candidate.key_id = "key-openai-usage-local-2".to_string();
second_candidate.key_name = "secondary".to_string();
second_candidate.key_internal_priority = second_candidate.key_internal_priority - 1;
second_candidate.key_internal_priority -= 1;
let candidate_selection_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_local_openai_candidate_row(),
+31 -6
View File
@@ -170,12 +170,13 @@ Linux/macOS 可运行 `sudo aether-tunnel upgrade [version]`。自更新只接
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--upstream-connect-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_CONNECT_TIMEOUT_SECS` | `30` | 上游建连超时(秒) |
| `--upstream-connect-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_CONNECT_TIMEOUT` | `30` | 上游建连超时(秒) |
| `--upstream-pool-max-idle-per-host` | `AETHER_TUNNEL_UPSTREAM_POOL_MAX_IDLE_PER_HOST` | `64` | 每 Host 最大空闲连接数 |
| `--upstream-pool-idle-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_POOL_IDLE_TIMEOUT_SECS` | `300` | 连接池空闲超时(秒) |
| `--upstream-tcp-keepalive-secs` | `AETHER_TUNNEL_UPSTREAM_TCP_KEEPALIVE_SECS` | `60` | TCP keepalive(秒,0 关闭) |
| `--upstream-pool-idle-timeout-secs` | `AETHER_TUNNEL_UPSTREAM_POOL_IDLE_TIMEOUT` | `300` | 连接池空闲超时(秒) |
| `--upstream-tcp-keepalive-secs` | `AETHER_TUNNEL_UPSTREAM_TCP_KEEPALIVE` | `60` | TCP keepalive(秒,0 关闭) |
| `--upstream-tcp-nodelay` | `AETHER_TUNNEL_UPSTREAM_TCP_NODELAY` | `true` | 启用 TCP_NODELAY |
| `--upstream-proxy-url` | `AETHER_TUNNEL_UPSTREAM_PROXY_URL` | 空 | 仅 provider 上游请求使用的出口代理 |
| `--upstream-proxy-remote-dns` | `AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS` | `false` | 显式信任 HTTP/SOCKS5h 代理解析供应商域名并执行目标 IP 访问控制;需重启 |
启用 `follow_redirects` 后,同源 307/308 会在请求体不超过 5 MiB 时重放。首个上游请求始终流式传输;超过重放预算时不会拒绝或截断原请求,而是将 307/308 响应原样返回给调用方。
@@ -185,14 +186,38 @@ Linux/macOS 可运行 `sudo aether-tunnel upgrade [version]`。自更新只接
upstream_proxy_url = "socks5h://microwarp:1080"
```
默认仍由隧道本机解析供应商域名、执行端口/IP ACL,再把已校验的 IP 交给代理;仅配置
`socks5h://` 不会跳过本地 DNS。这保留现有的防 DNS 重绑定及内网访问边界。
如果隧道本机 DNS 不可用、被污染或返回不可路由的 Fake-IP,可显式委托**受信任且配置了
目的地址访问控制的代理**解析域名。在 TOML 顶层(第一个 `[[servers]]` 之前)配置:
```toml
upstream_proxy_url = "socks5h://microwarp:1080"
upstream_proxy_remote_dns = true
```
也可启用环境变量 `AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS=true`、CLI 参数
`--upstream-proxy-remote-dns` 或 setup 中的 `Proxy Remote DNS` 开关,保存后重启。
该模式仅支持 `http://` 和 `socks5h://`,不支持本地解析语义的 `socks5://`;未配置代理时
启动会报错。域名原样交给 HTTP CONNECT/SOCKS5h,HTTP Host 和 TLS SNI/证书校验仍使用
原域名,不会在失败时偷偷回退到本地 DNS。
**安全边界:**普通 HTTP CONNECT/SOCKS5 不能让隧道校验代理最终解析出的目标 IP,因此
启用该模式代表把域名目标的 IP ACL 委托给代理,而不只是换一个 DNS 服务器。隧道仍检查
端口、URL 凭据/fragment、`localhost` 和 IP 字面地址;默认继续拒绝私网/保留 IP 字面地址。
这不需要打开 `allow_private_targets`。代理本身的域名仍需本地解析;如果本地 DNS 完全
不可用,使用代理 IP 地址或修复本地解析。代理 DNS、TCP、CONNECT/SOCKS 和 TLS 握手共同
受 `upstream_connect_timeout_secs` 限制。
如果需要让 Aether 管理 API 和 WebSocket tunnel 也走代理,使用 `aether_outbound_proxy_url`。
#### Aether API 客户端
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--aether-request-timeout-secs` | `AETHER_TUNNEL_AETHER_REQUEST_TIMEOUT_SECS` | `10` | 请求总超时(秒) |
| `--aether-connect-timeout-secs` | `AETHER_TUNNEL_AETHER_CONNECT_TIMEOUT_SECS` | `10` | 建连超时(秒) |
| `--aether-request-timeout-secs` | `AETHER_TUNNEL_AETHER_REQUEST_TIMEOUT` | `10` | 请求总超时(秒) |
| `--aether-connect-timeout-secs` | `AETHER_TUNNEL_AETHER_CONNECT_TIMEOUT` | `10` | 建连超时(秒) |
| `--aether-outbound-proxy-url` | `AETHER_TUNNEL_AETHER_OUTBOUND_PROXY_URL` | 空 | Aether 注册、心跳和 WebSocket tunnel 回连使用的出口代理(默认不走代理) |
| `--aether-retry-max-attempts` | `AETHER_TUNNEL_AETHER_RETRY_MAX_ATTEMPTS` | `3` | 最大重试次数 |
@@ -201,7 +226,7 @@ upstream_proxy_url = "socks5h://microwarp:1080"
| 参数 | 环境变量 | 默认值 | 说明 |
|------|----------|--------|------|
| `--allow-private-targets` | `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS` | `false` | 默认拦截 private/reserved 目标地址;仅在明确需要访问内网服务时设为 `true`,且仅影响重启后的进程 |
| `--dns-cache-ttl-secs` | `AETHER_TUNNEL_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL(秒) |
| `--dns-cache-ttl-secs` | `AETHER_TUNNEL_DNS_CACHE_TTL` | `60` | DNS 缓存 TTL(秒) |
| `--dns-cache-capacity` | `AETHER_TUNNEL_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) |
#### 日志
+6
View File
@@ -126,11 +126,16 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
if let Ok(proxy) = crate::egress_proxy::UpstreamProxyConfig::parse(proxy_url) {
info!(
upstream_proxy_url = %proxy.redacted_url(),
upstream_proxy_remote_dns = config.upstream_proxy_remote_dns,
"provider upstream egress proxy configured"
);
}
}
if config.upstream_proxy_remote_dns {
warn!("provider hostname DNS resolution and destination IP access controls are delegated to the trusted upstream proxy");
}
// Resolve public IP (best-effort for region info)
let public_ip = match &config.public_ip {
Some(ip) => ip.clone(),
@@ -1420,6 +1425,7 @@ mod tests {
upstream_tcp_keepalive_secs: 60,
upstream_tcp_nodelay: true,
upstream_proxy_url: None,
upstream_proxy_remote_dns: false,
legacy_redirect_replay_budget_bytes_ignored: None,
emit_proxy_timing_header: true,
log_level: "info".to_string(),
+75
View File
@@ -480,6 +480,14 @@ pub struct Config {
#[arg(long, env = "AETHER_TUNNEL_UPSTREAM_PROXY_URL")]
pub upstream_proxy_url: Option<String>,
#[arg(
long,
env = "AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS",
default_value_t = false,
help = "Trust an HTTP or SOCKS5h upstream proxy to resolve hostnames and enforce destination IP access controls"
)]
pub upstream_proxy_remote_dns: bool,
/// Accepted only so older launch commands and environments keep working.
/// Redirect request bodies are always replayed without a cumulative size limit.
#[arg(
@@ -820,6 +828,16 @@ impl Config {
crate::egress_proxy::UpstreamProxyConfig::parse(proxy_url)
.map_err(|err| anyhow::anyhow!("upstream_proxy_url invalid: {err}"))?;
}
if self.upstream_proxy_remote_dns {
let proxy_url = normalized_proxy_url(&self.upstream_proxy_url).ok_or_else(|| {
anyhow::anyhow!("upstream_proxy_remote_dns requires upstream_proxy_url")
})?;
let proxy = crate::egress_proxy::UpstreamProxyConfig::parse(proxy_url)
.map_err(anyhow::Error::msg)?;
if !proxy.supports_remote_target_dns() {
anyhow::bail!("upstream_proxy_remote_dns requires an http:// or socks5h:// proxy");
}
}
if matches!(self.max_in_flight_streams, Some(0)) {
anyhow::bail!("max_in_flight_streams must be > 0");
}
@@ -1087,6 +1105,8 @@ pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_proxy_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub upstream_proxy_remote_dns: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub emit_proxy_timing_header: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub log_level: Option<String>,
@@ -1306,6 +1326,10 @@ impl ConfigFile {
self.upstream_tcp_nodelay
);
set!("AETHER_TUNNEL_UPSTREAM_PROXY_URL", self.upstream_proxy_url);
set!(
"AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS",
self.upstream_proxy_remote_dns
);
set!(
"AETHER_TUNNEL_EMIT_PROXY_TIMING_HEADER",
self.emit_proxy_timing_header
@@ -1670,6 +1694,57 @@ mod tests {
use super::*;
use crate::hardware::HardwareInfo;
#[test]
fn proxy_remote_dns_requires_explicit_trust_and_a_remote_dns_proxy() {
let mut config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
]);
assert!(!config.upstream_proxy_remote_dns);
let argument = Config::command()
.get_arguments()
.find(|argument| argument.get_id() == "upstream_proxy_remote_dns")
.unwrap()
.clone();
assert_eq!(
argument.get_env().unwrap(),
"AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS"
);
config.upstream_proxy_remote_dns = true;
for proxy in [None, Some(" "), Some("socks5://127.0.0.1:1080")] {
config.upstream_proxy_url = proxy.map(str::to_string);
assert!(config
.validate()
.unwrap_err()
.to_string()
.contains("upstream_proxy_remote_dns"));
}
for proxy in ["http://127.0.0.1:8080", "socks5h://127.0.0.1:1080"] {
config.upstream_proxy_url = Some(proxy.to_string());
config
.validate()
.expect("explicit remote DNS configuration should validate");
}
}
#[test]
fn config_file_round_trips_proxy_remote_dns() {
let config = parse_config_file_content(
"upstream_proxy_url = \"socks5h://127.0.0.1:1080\"\nupstream_proxy_remote_dns = true",
)
.unwrap();
assert_eq!(config.upstream_proxy_remote_dns, Some(true));
let round_trip: ConfigFile = toml::from_str(&toml::to_string(&config).unwrap()).unwrap();
assert_eq!(round_trip.upstream_proxy_remote_dns, Some(true));
assert_eq!(ConfigFile::default().upstream_proxy_remote_dns, None);
}
fn config_save_test_dir(label: &str) -> std::path::PathBuf {
let path = std::env::temp_dir().join(format!(
"aether-tunnel-config-{label}-{}",
+21 -1
View File
@@ -142,6 +142,13 @@ impl UpstreamProxyConfig {
self.scheme == UpstreamProxyScheme::Socks5h
}
pub(crate) fn supports_remote_target_dns(&self) -> bool {
matches!(
self.scheme,
UpstreamProxyScheme::Http | UpstreamProxyScheme::Socks5h
)
}
pub(crate) fn basic_auth_header(&self) -> Option<String> {
let username = self.username()?;
let mut credentials = String::with_capacity(
@@ -445,7 +452,7 @@ pub(crate) async fn socks5_target_address(
remote_dns: bool,
) -> io::Result<Vec<u8>> {
let mut request = vec![0x05, 0x01, 0x00];
if let Ok(ip) = target_host.parse::<IpAddr>() {
if let Some(ip) = aether_http::parse_ip_literal_host(target_host) {
push_socks5_ip_address(&mut request, ip);
} else if remote_dns {
let host = target_host.as_bytes();
@@ -524,6 +531,19 @@ fn non_empty_url_part(value: &str) -> Option<String> {
mod tests {
use super::*;
#[tokio::test]
async fn socks_proxy_encodes_bracketed_ipv6_as_an_ip_for_both_dns_modes() {
for remote_dns in [false, true] {
let expected = socks5_target_address("::1", 443, remote_dns).await.unwrap();
let actual = socks5_target_address("[::1]", 443, remote_dns)
.await
.unwrap();
assert_eq!(actual, expected);
assert_eq!(&actual[..4], &[5, 1, 0, 4]);
assert_eq!(&actual[20..], &443u16.to_be_bytes());
}
}
#[test]
fn parses_http_proxy_with_default_port() {
let proxy = UpstreamProxyConfig::parse("http://proxy.example").expect("proxy should parse");
+55 -1
View File
@@ -214,6 +214,14 @@ impl App {
required: false,
help: "Heartbeat interval in seconds; default is 5",
},
Field {
label: "Proxy Remote DNS",
key: "upstream_proxy_remote_dns",
value: "false".into(),
kind: FieldKind::Bool,
required: false,
help: "Trust HTTP/SOCKS5h egress proxy to resolve provider hosts and enforce destination IP ACLs; restart required",
},
],
selected: 0,
mode: Mode::Normal,
@@ -288,6 +296,9 @@ impl App {
"allow_private_targets" => cfg.allow_private_targets.map(|v| v.to_string()),
"heartbeat_interval" => cfg.heartbeat_interval.map(|v| v.to_string()),
"upstream_proxy_url" => cfg.upstream_proxy_url.clone(),
"upstream_proxy_remote_dns" => {
cfg.upstream_proxy_remote_dns.map(|value| value.to_string())
}
_ => None,
};
if let Some(v) = val {
@@ -396,11 +407,23 @@ impl App {
let get_tab = |tab: &ServerTab, key: &str| -> Option<String> { Self::get_tab(tab, key) };
let save_logs_to_file = self.toggle_enabled("save_logs_to_file");
let upstream_proxy_url = self.parse_optional_upstream_proxy_url()?;
let upstream_proxy_remote_dns = self.toggle_enabled("upstream_proxy_remote_dns");
if upstream_proxy_remote_dns {
let proxy_url = upstream_proxy_url
.as_deref()
.ok_or_else(|| anyhow::anyhow!("Proxy Remote DNS requires an egress proxy"))?;
let proxy = UpstreamProxyConfig::parse(proxy_url).map_err(anyhow::Error::msg)?;
if !proxy.supports_remote_target_dns() {
anyhow::bail!("Proxy Remote DNS requires an http:// or socks5h:// proxy");
}
}
let mut cfg = ConfigFile {
log_level: get_global("log_level"),
allow_private_targets: Some(self.toggle_enabled("allow_private_targets")),
heartbeat_interval: self.parse_optional_heartbeat_interval()?,
upstream_proxy_url: self.parse_optional_upstream_proxy_url()?,
upstream_proxy_url,
upstream_proxy_remote_dns: Some(upstream_proxy_remote_dns),
log_destination: Some(if save_logs_to_file {
TunnelLogDestinationArg::Both
} else {
@@ -1086,6 +1109,37 @@ mod tests {
app
}
#[test]
fn proxy_remote_dns_toggle_round_trips_and_requires_a_trusted_proxy() {
let mut app = sample_app();
assert_eq!(
app.to_config().unwrap().upstream_proxy_remote_dns,
Some(false)
);
set_global_field(&mut app, "upstream_proxy_remote_dns", "true");
assert!(app
.to_config()
.unwrap_err()
.to_string()
.contains("requires an egress proxy"));
set_global_field(&mut app, "upstream_proxy_url", "socks5://127.0.0.1:1080");
assert!(app
.to_config()
.unwrap_err()
.to_string()
.contains("socks5h://"));
for proxy in ["http://127.0.0.1:8080", "socks5h://127.0.0.1:1080"] {
set_global_field(&mut app, "upstream_proxy_url", proxy);
let config = app.to_config().unwrap();
let mut restored = sample_app();
restored.apply_config(&config);
let round_trip = restored.to_config().unwrap();
assert_eq!(round_trip.upstream_proxy_remote_dns, Some(true));
assert_eq!(round_trip.upstream_proxy_url.as_deref(), Some(proxy));
}
}
fn unique_temp_config_path(name: &str) -> PathBuf {
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
+19 -6
View File
@@ -226,21 +226,34 @@ pub async fn validate_target(
allow_private: bool,
dns_cache: &DnsCache,
) -> Result<Vec<SocketAddr>, FilterError> {
// Port whitelist check
if let Some(address) = validate_target_literal(host, port, allowed_ports, allow_private)? {
return Ok(vec![address]);
}
resolve_public_addrs(host, port, allow_private, dns_cache).await
}
pub(crate) fn validate_target_literal(
host: &str,
port: u16,
allowed_ports: &HashSet<u16>,
allow_private: bool,
) -> Result<Option<SocketAddr>, FilterError> {
if !allowed_ports.contains(&port) {
return Err(FilterError::PortNotAllowed(port));
}
// Try parsing as IP directly (no DNS needed)
if let Ok(ip) = host.parse::<IpAddr>() {
if let Some(ip) = aether_http::parse_ip_literal_host(host) {
if !allow_private && is_private_ip(&ip) {
return Err(FilterError::PrivateIp(ip));
}
return Ok(vec![SocketAddr::new(ip, port)]);
return Ok(Some(SocketAddr::new(ip, port)));
}
// Resolve and return the exact addresses authorized for this request.
resolve_public_addrs(host, port, allow_private, dns_cache).await
if !allow_private && host.trim_end_matches('.').eq_ignore_ascii_case("localhost") {
return Err(FilterError::NoPublicAddrs(host.to_string()));
}
Ok(None)
}
#[cfg(test)]
+1
View File
@@ -737,6 +737,7 @@ mod tests {
upstream_tcp_keepalive_secs: 60,
upstream_tcp_nodelay: true,
upstream_proxy_url: None,
upstream_proxy_remote_dns: false,
legacy_redirect_replay_budget_bytes_ignored: None,
emit_proxy_timing_header: true,
log_level: "info".to_string(),
+123 -13
View File
@@ -1276,6 +1276,40 @@ fn resolve_redirect<B>(
}
}
async fn resolve_upstream_target(
current_url: &url::Url,
allowed_ports: &std::collections::HashSet<u16>,
allow_private_targets: bool,
proxy_remote_dns: bool,
dns_cache: &target_filter::DnsCache,
) -> Result<upstream_client::ValidatedUpstreamTarget, String> {
validate_tunnel_upstream_url(current_url, allow_private_targets).map_err(str::to_string)?;
let host = current_url
.host_str()
.ok_or_else(|| "missing host in URL".to_string())?;
let port = current_url
.port_or_known_default()
.ok_or_else(|| "missing port in URL".to_string())?;
let addresses = if proxy_remote_dns {
match target_filter::validate_target_literal(
host,
port,
allowed_ports,
allow_private_targets,
)
.map_err(|_| "upstream target blocked".to_string())?
{
Some(address) => vec![address],
None => return upstream_client::ValidatedUpstreamTarget::proxy_resolved(current_url),
}
} else {
target_filter::validate_target(host, port, allowed_ports, allow_private_targets, dns_cache)
.await
.map_err(|_| "upstream target blocked".to_string())?
};
upstream_client::ValidatedUpstreamTarget::new(current_url, addresses)
}
#[allow(clippy::too_many_arguments)]
async fn execute_upstream_request(
state: &AppState,
@@ -1288,24 +1322,19 @@ async fn execute_upstream_request(
timeout: Duration,
http1_only: bool,
) -> Result<UpstreamResponseContext, String> {
let host = current_url
.host_str()
.ok_or_else(|| "missing host in URL".to_string())?;
let port = current_url.port_or_known_default().unwrap_or(443);
let dns_start = Instant::now();
let validated_addrs = {
let validated_target = {
let allowed_ports = Arc::clone(&server.dynamic.load().allowed_ports);
match target_filter::validate_target(
host,
port,
match resolve_upstream_target(
current_url,
&allowed_ports,
state.config.allow_private_targets,
state.config.upstream_proxy_remote_dns,
&state.dns_cache,
)
.await
{
Ok(addrs) => addrs,
Ok(target) => target,
Err(_error) => {
server.metrics.dns_failures.fetch_add(1, Ordering::Release);
// Keep the detailed filter error out of the tunnel response;
@@ -1317,9 +1346,6 @@ async fn execute_upstream_request(
};
let dns_ms = dns_start.elapsed().as_millis() as u64;
let validated_target =
upstream_client::ValidatedUpstreamTarget::new(current_url, validated_addrs)?;
let client_key = upstream_client::upstream_client_pool_key(
meta.provider_id.as_deref(),
meta.endpoint_id.as_deref(),
@@ -2326,6 +2352,89 @@ fn build_prefixed_request_body(
#[cfg(test)]
mod tests {
#[tokio::test]
async fn remote_dns_target_resolution_skips_local_dns_and_keeps_literal_acl() {
let cache = target_filter::DnsCache::new(Duration::from_secs(60), 16);
let ports = [80, 443].into_iter().collect();
let url = url::Url::parse("https://remote-dns-test.invalid/path").unwrap();
let target = resolve_upstream_target(&url, &ports, false, true, &cache)
.await
.expect("trusted proxy should receive an unresolved hostname");
assert!(target.uses_proxy_dns());
assert!(cache.get("remote-dns-test.invalid", 443).await.is_none());
for address in ["https://8.8.8.8/", "https://[2606:4700:4700::1111]/"] {
let target = resolve_upstream_target(
&url::Url::parse(address).unwrap(),
&ports,
false,
true,
&cache,
)
.await
.unwrap();
assert!(!target.uses_proxy_dns(), "IP literals must remain pinned");
}
for address in [
"https://127.0.0.1/",
"https://10.0.0.1/",
"https://198.18.0.1/",
"https://[::1]/",
"https://[::ffff:127.0.0.1]/",
"https://localhost/",
"https://LOCALHOST./",
"https://remote-dns-test.invalid:25/",
"https://user:[email protected]/",
"https://remote-dns-test.invalid/#fragment",
"ftp://remote-dns-test.invalid/",
] {
assert!(
resolve_upstream_target(
&url::Url::parse(address).unwrap(),
&ports,
false,
true,
&cache,
)
.await
.is_err(),
"target should remain blocked: {address}"
);
}
}
#[tokio::test]
async fn strict_dns_targets_stay_pinned_and_separate_from_remote_dns_targets() {
let cache = target_filter::DnsCache::new(Duration::from_secs(60), 16);
let ports = [443].into_iter().collect();
let url = url::Url::parse("https://remote-dns-test.invalid/").unwrap();
cache
.insert(
"remote-dns-test.invalid",
443,
Arc::new(vec!["8.8.8.8:443".parse().unwrap()]),
)
.await;
let strict = resolve_upstream_target(&url, &ports, false, false, &cache)
.await
.unwrap();
let remote = resolve_upstream_target(&url, &ports, false, true, &cache)
.await
.unwrap();
assert!(!strict.uses_proxy_dns());
assert!(remote.uses_proxy_dns());
assert_ne!(strict, remote);
let private_url = url::Url::parse("http://[::1]/").unwrap();
let private_ports = [80].into_iter().collect();
let explicitly_allowed =
resolve_upstream_target(&private_url, &private_ports, true, true, &cache)
.await
.unwrap();
assert!(!explicitly_allowed.uses_proxy_dns());
}
#[tokio::test(start_paused = true)]
async fn window_updates_wait_for_capacity_instead_of_disappearing() {
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(1);
@@ -3820,6 +3929,7 @@ mod tests {
upstream_tcp_keepalive_secs: 60,
upstream_tcp_nodelay: true,
upstream_proxy_url: None,
upstream_proxy_remote_dns: false,
legacy_redirect_replay_budget_bytes_ignored: None,
emit_proxy_timing_header: true,
log_level: "info".to_string(),
+350 -42
View File
@@ -34,7 +34,8 @@ use tower_service::Service;
use crate::config::Config;
use crate::egress_proxy::{
connect_validated_target_via_proxy, ProxyConnectOptions, UpstreamProxyConfig,
connect_target_via_proxy, connect_validated_target_via_proxy, ProxyConnectOptions,
UpstreamProxyConfig,
};
use crate::target_filter::DnsCache;
@@ -66,11 +67,42 @@ pub struct ValidatedUpstreamTarget {
scheme: String,
host: String,
port: u16,
addrs: Vec<SocketAddr>,
resolution: UpstreamTargetResolution,
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
enum UpstreamTargetResolution {
Pinned(Vec<SocketAddr>),
ProxyDns,
}
impl ValidatedUpstreamTarget {
pub fn new(target_url: &url::Url, mut addrs: Vec<SocketAddr>) -> Result<Self, String> {
if addrs.is_empty() {
return Err("validated upstream target has no addresses".to_string());
}
addrs.sort_unstable();
addrs.dedup();
Self::with_resolution(target_url, UpstreamTargetResolution::Pinned(addrs))
}
pub(crate) fn proxy_resolved(target_url: &url::Url) -> Result<Self, String> {
if !matches!(target_url.host(), Some(url::Host::Domain(_))) {
return Err("IP literal targets must use pinned addresses".to_string());
}
Self::with_resolution(target_url, UpstreamTargetResolution::ProxyDns)
}
fn with_resolution(
target_url: &url::Url,
resolution: UpstreamTargetResolution,
) -> Result<Self, String> {
if !target_url.username().is_empty()
|| target_url.password().is_some()
|| target_url.fragment().is_some()
{
return Err("upstream target must not contain credentials or a fragment".to_string());
}
let scheme = target_url.scheme().to_ascii_lowercase();
if !matches!(scheme.as_str(), "http" | "https") {
return Err(format!("unsupported upstream scheme {scheme}"));
@@ -86,19 +118,16 @@ impl ValidatedUpstreamTarget {
let port = target_url
.port_or_known_default()
.ok_or_else(|| "missing port in upstream URL".to_string())?;
if addrs.is_empty() {
return Err("validated upstream target has no addresses".to_string());
if let UpstreamTargetResolution::Pinned(addrs) = &resolution {
if addrs.iter().any(|addr| addr.port() != port) {
return Err("validated upstream target address has the wrong port".to_string());
}
}
if addrs.iter().any(|addr| addr.port() != port) {
return Err("validated upstream target address has the wrong port".to_string());
}
addrs.sort_unstable();
addrs.dedup();
Ok(Self {
scheme,
host,
port,
addrs,
resolution,
})
}
@@ -119,8 +148,8 @@ impl ValidatedUpstreamTarget {
Ok(())
}
fn addrs(&self) -> &[SocketAddr] {
&self.addrs
pub(crate) fn uses_proxy_dns(&self) -> bool {
matches!(self.resolution, UpstreamTargetResolution::ProxyDns)
}
}
@@ -333,9 +362,14 @@ impl Service<Name> for PinnedResolver {
"DNS request does not match the validated upstream host",
));
}
Ok(ValidatedAddrs {
inner: target.addrs.into_iter(),
})
match target.resolution {
UpstreamTargetResolution::Pinned(addrs) => Ok(ValidatedAddrs {
inner: addrs.into_iter(),
}),
UpstreamTargetResolution::ProxyDns => Err(io::Error::other(
"proxy-resolved target must not fall back to local DNS",
)),
}
})
}
}
@@ -376,16 +410,25 @@ impl Service<Uri> for InstrumentedConnector {
};
let connect_start = std::time::Instant::now();
return Box::pin(async move {
connect_via_proxy(
dst,
scheme,
tls_config,
proxy,
validated_target,
options,
connect_start,
tokio::time::timeout(
options.connect_timeout,
connect_via_proxy(
dst,
scheme,
tls_config,
proxy,
validated_target,
options,
connect_start,
),
)
.await
.map_err(|_| {
Box::new(io::Error::new(
io::ErrorKind::TimedOut,
"upstream proxy connection timed out",
)) as BoxError
})?
});
}
let connecting = self.http.call(dst.clone());
@@ -441,20 +484,40 @@ async fn connect_via_proxy(
connect_start: std::time::Instant,
) -> Result<TimedConn, BoxError> {
let scheme = scheme.ok_or_else(|| io::Error::other("missing scheme"))?;
let mut last_error = None;
let mut connected = None;
for target_addr in validated_target.addrs().iter().copied() {
match connect_validated_target_via_proxy(&proxy, target_addr, options).await {
Ok(tcp) => {
connected = Some(tcp);
break;
let tcp = match &validated_target.resolution {
UpstreamTargetResolution::ProxyDns => {
if !proxy.supports_remote_target_dns() {
return Err(
io::Error::other("upstream proxy does not support remote target DNS").into(),
);
}
Err(error) => last_error = Some(error),
connect_target_via_proxy(
&proxy,
&validated_target.host,
validated_target.port,
options,
)
.await?
}
}
let tcp = connected.ok_or_else(|| {
last_error.unwrap_or_else(|| io::Error::other("validated upstream target has no addresses"))
})?;
UpstreamTargetResolution::Pinned(addrs) => {
let mut last_error = None;
let mut connected = None;
for target_addr in addrs.iter().copied() {
match connect_validated_target_via_proxy(&proxy, target_addr, options).await {
Ok(tcp) => {
connected = Some(tcp);
break;
}
Err(error) => last_error = Some(error),
}
}
connected.ok_or_else(|| {
last_error.unwrap_or_else(|| {
io::Error::other("validated upstream target has no addresses")
})
})?
}
};
let connect_ms = connect_start.elapsed().as_millis() as u64;
@@ -512,6 +575,24 @@ fn build_upstream_client_with_protocol(
http1_only: bool,
h2c_prior_knowledge: bool,
) -> Result<UpstreamClient, String> {
let proxy = config
.upstream_proxy_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(UpstreamProxyConfig::parse)
.transpose()?;
if validated_target.uses_proxy_dns()
&& (!config.upstream_proxy_remote_dns
|| !proxy
.as_ref()
.is_some_and(UpstreamProxyConfig::supports_remote_target_dns))
{
return Err(
"proxy-resolved upstream requires explicit remote DNS and an HTTP or SOCKS5h proxy"
.to_string(),
);
}
let mut http = HttpConnector::new_with_resolver(PinnedResolver::new(validated_target.clone()));
http.enforce_http(false);
http.set_connect_timeout(Some(Duration::from_secs(
@@ -530,13 +611,7 @@ fn build_upstream_client_with_protocol(
http,
tls_config: build_tls_config(http1_only),
validated_target,
proxy: config
.upstream_proxy_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(UpstreamProxyConfig::parse)
.transpose()?,
proxy,
connect_timeout: Duration::from_secs(config.upstream_connect_timeout_secs),
tcp_nodelay: config.upstream_tcp_nodelay,
tcp_keepalive: (config.upstream_tcp_keepalive_secs > 0)
@@ -1035,6 +1110,239 @@ mod tests {
);
}
#[tokio::test]
async fn trusted_http_proxy_resolves_hostname_without_local_dns() {
let (proxy_url, connect_rx, request_rx) = spawn_http_proxy().await;
let client = remote_dns_client(&proxy_url, "http://remote-dns-test.invalid/");
let request = hyper::Request::builder()
.uri("http://remote-dns-test.invalid/remote-dns")
.body(full_request_body(Bytes::new()))
.unwrap();
let response = tokio::time::timeout(Duration::from_secs(5), client.request(request))
.await
.unwrap()
.expect("proxy should resolve the target without local DNS");
assert_eq!(response.status(), hyper::StatusCode::OK);
assert_eq!(
response.into_body().collect().await.unwrap().to_bytes(),
"ok"
);
assert!(connect_rx
.await
.unwrap()
.starts_with("CONNECT remote-dns-test.invalid:80 HTTP/1.1\r\n"));
let request = request_rx.await.unwrap().to_ascii_lowercase();
assert!(request.starts_with("get /remote-dns http/1.1\r\n"));
assert!(request.contains("\r\nhost: remote-dns-test.invalid\r\n"));
}
#[tokio::test]
async fn trusted_socks5h_proxy_receives_hostname_not_a_locally_resolved_ip() {
let (proxy_url, target_rx, request_rx) = spawn_remote_dns_socks_proxy().await;
let client = remote_dns_client(&proxy_url, "http://remote-dns-test.invalid/");
let request = hyper::Request::builder()
.uri("http://remote-dns-test.invalid/remote-dns")
.body(full_request_body(Bytes::new()))
.unwrap();
let response = tokio::time::timeout(Duration::from_secs(5), client.request(request))
.await
.unwrap()
.expect("SOCKS proxy should receive the unresolved target");
assert_eq!(response.status(), hyper::StatusCode::OK);
assert_eq!(
response.into_body().collect().await.unwrap().to_bytes(),
"ok"
);
assert_eq!(
target_rx.await.unwrap(),
("remote-dns-test.invalid".to_string(), 80)
);
assert!(request_rx
.await
.unwrap()
.to_ascii_lowercase()
.contains("\r\nhost: remote-dns-test.invalid\r\n"));
}
#[tokio::test]
async fn remote_dns_https_preserves_hostname_for_connect_and_sni() {
let (proxy_url, connect_rx) = spawn_connect_only_http_proxy().await;
let client = remote_dns_client(&proxy_url, "https://remote-dns-test.invalid/");
let uri: Uri = "https://remote-dns-test.invalid/secure".parse().unwrap();
let request = hyper::Request::builder()
.uri(uri.clone())
.body(full_request_body(Bytes::new()))
.unwrap();
let _ = tokio::time::timeout(Duration::from_secs(5), client.request(request))
.await
.unwrap();
assert!(connect_rx
.await
.unwrap()
.starts_with("CONNECT remote-dns-test.invalid:443 HTTP/1.1\r\n"));
match resolve_server_name(&uri).unwrap() {
ServerName::DnsName(name) => assert_eq!(name.as_ref(), "remote-dns-test.invalid"),
other => panic!("expected hostname for TLS verification, got {other:?}"),
}
}
#[tokio::test]
async fn remote_dns_targets_cannot_fall_back_to_local_dns_or_change_origin() {
let url = url::Url::parse("https://remote-dns-test.invalid/").unwrap();
let target = ValidatedUpstreamTarget::proxy_resolved(&url).unwrap();
let mut resolver = PinnedResolver::new(target.clone());
let error = resolver
.call("remote-dns-test.invalid".parse().unwrap())
.await
.err()
.unwrap();
assert!(error
.to_string()
.contains("must not fall back to local DNS"));
for uri in [
"http://remote-dns-test.invalid/",
"https://another-target.invalid/",
"https://remote-dns-test.invalid:8443/",
] {
assert!(target.ensure_matches_uri(&uri.parse().unwrap()).is_err());
}
for url in ["http://127.0.0.1/", "https://[::1]/"] {
assert!(
ValidatedUpstreamTarget::proxy_resolved(&url::Url::parse(url).unwrap()).is_err()
);
}
}
#[test]
fn remote_dns_clients_require_opt_in_and_do_not_share_pinned_pool_entries() {
let mut config = remote_dns_config("http://127.0.0.1:8080");
let url = url::Url::parse("https://remote-dns-test.invalid/").unwrap();
let remote = ValidatedUpstreamTarget::proxy_resolved(&url).unwrap();
config.upstream_proxy_remote_dns = false;
assert!(build_upstream_client_with_protocol(&config, remote.clone(), true, false).is_err());
config.upstream_proxy_remote_dns = true;
for proxy in [None, Some("socks5://127.0.0.1:1080")] {
config.upstream_proxy_url = proxy.map(str::to_string);
assert!(
build_upstream_client_with_protocol(&config, remote.clone(), true, false).is_err()
);
}
config.upstream_proxy_url = Some("http://127.0.0.1:8080".to_string());
let pinned =
ValidatedUpstreamTarget::new(&url, vec!["8.8.8.8:443".parse().unwrap()]).unwrap();
let remote_key = upstream_client_pool_key(None, None, None, None, false, remote);
let pinned_key = upstream_client_pool_key(None, None, None, None, false, pinned);
assert_ne!(remote_key, pinned_key);
let pool = UpstreamClientPool::new(
Arc::new(config),
Arc::new(DnsCache::new(Duration::from_secs(60), 16)),
);
pool.get_or_build(remote_key).unwrap();
pool.get_or_build(pinned_key).unwrap();
assert_eq!(pool.clients.lock().unwrap().len(), 2);
}
#[tokio::test]
async fn remote_dns_proxy_connect_timeout_covers_connect_and_tls_handshakes() {
for tls in [false, true] {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_url = format!("http://{}", listener.local_addr().unwrap());
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
read_http_headers(&mut stream).await;
if tls {
stream
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await
.unwrap();
}
std::future::pending::<()>().await;
drop(stream);
});
let target_url = if tls {
"https://remote-dns-test.invalid/"
} else {
"http://remote-dns-test.invalid/"
};
let client = remote_dns_client(&proxy_url, target_url);
let request = hyper::Request::builder()
.uri(target_url)
.body(full_request_body(Bytes::new()))
.unwrap();
let result =
tokio::time::timeout(Duration::from_secs(5), client.request(request)).await;
server.abort();
let error = result
.expect("configured connect timeout must include proxy and TLS handshakes")
.unwrap_err();
assert!(error.is_connect());
}
}
fn remote_dns_config(proxy_url: &str) -> Config {
let _ = rustls::crypto::ring::default_provider().install_default();
Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
"--upstream-proxy-url",
proxy_url,
"--upstream-proxy-remote-dns",
"--upstream-connect-timeout-secs",
"1",
])
}
fn remote_dns_client(proxy_url: &str, target_url: &str) -> UpstreamClient {
let config = remote_dns_config(proxy_url);
let target =
ValidatedUpstreamTarget::proxy_resolved(&url::Url::parse(target_url).unwrap()).unwrap();
build_upstream_client_with_protocol(&config, target, true, false).unwrap()
}
async fn spawn_remote_dns_socks_proxy() -> (
String,
tokio::sync::oneshot::Receiver<(String, u16)>,
tokio::sync::oneshot::Receiver<String>,
) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_url = format!("socks5h://{}", listener.local_addr().unwrap());
let (target_tx, target_rx) = tokio::sync::oneshot::channel();
let (request_tx, request_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut greeting = [0u8; 3];
stream.read_exact(&mut greeting).await.unwrap();
assert_eq!(greeting, [0x05, 0x01, 0x00]);
stream.write_all(&[0x05, 0x00]).await.unwrap();
let mut header = [0u8; 5];
stream.read_exact(&mut header).await.unwrap();
assert_eq!(&header[..4], &[0x05, 0x01, 0x00, 0x03]);
let mut hostname = vec![0; header[4] as usize];
stream.read_exact(&mut hostname).await.unwrap();
let port = stream.read_u16().await.unwrap();
target_tx
.send((String::from_utf8(hostname).unwrap(), port))
.unwrap();
stream
.write_all(&[0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0])
.await
.unwrap();
request_tx
.send(read_http_headers(&mut stream).await)
.unwrap();
stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
.await
.unwrap();
});
(proxy_url, target_rx, request_rx)
}
fn proxied_client(
proxy_url: &str,
target_url: &str,
@@ -312,6 +312,10 @@ pub fn build_admin_monitoring_trace_request_payload_response_with_key_accounts(
"total_candidates": trace.total_candidates,
"final_status": trace.final_status,
"total_latency_ms": trace.total_latency_ms,
"diagnostic_request": usage.filter(|usage| admin_monitoring_usage_matches_trace(usage, &trace.request_id)).map(|usage| json!({
"usage_id": usage.id,
"body_state": usage.request_body_state.map(|state| state.as_str()),
})),
"candidates": candidates,
}))
.into_response()
@@ -337,6 +341,7 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts(
item.sanitize_for_admin();
let candidate = &item.candidate;
let sanitized_extra_data = build_admin_monitoring_trace_candidate_extra_data(
&candidate.id,
candidate.extra_data.as_ref(),
candidate.status_code,
usage,
@@ -518,6 +523,7 @@ fn build_admin_monitoring_trace_candidate_ranking(existing: Option<&Value>) -> V
}
fn build_admin_monitoring_trace_candidate_extra_data(
candidate_id: &str,
existing: Option<&Value>,
candidate_status_code: Option<u16>,
usage: Option<&StoredRequestUsageAudit>,
@@ -527,6 +533,19 @@ fn build_admin_monitoring_trace_candidate_extra_data(
if let Some(usage) = usage {
let extra_object = extra_data.get_or_insert_with(serde_json::Map::new);
if usage.routing_candidate_id() == Some(candidate_id) {
extra_object.insert("diagnostic_context".to_string(), json!({
"usage_id": usage.id,
"model": usage.model,
"target_model": usage.target_model,
"body_states": {
"request_body": usage.request_body_state.map(|state| state.as_str()),
"provider_request_body": usage.provider_request_body_state.map(|state| state.as_str()),
"response_body": usage.response_body_state.map(|state| state.as_str()),
"client_response_body": usage.client_response_body_state.map(|state| state.as_str()),
}
}));
}
if let Some(first_byte_time_ms) = usage.first_byte_time_ms {
extra_object
.entry("first_byte_time_ms".to_string())
@@ -577,8 +596,21 @@ fn build_admin_monitoring_trace_candidate_extra_data(
}
}
sanitize_request_candidate_extra_data_for_persistence(extra_data.map(Value::Object))
.unwrap_or(Value::Null)
let diagnostic_context = extra_data
.as_mut()
.and_then(|extra| extra.remove("diagnostic_context"));
let mut sanitized =
sanitize_request_candidate_extra_data_for_persistence(extra_data.map(Value::Object))
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
if let Some(context) = diagnostic_context {
sanitized.insert("diagnostic_context".to_string(), context);
}
if sanitized.is_empty() {
Value::Null
} else {
Value::Object(sanitized)
}
}
fn admin_monitoring_trace_response_data(
@@ -238,3 +238,124 @@ impl fmt::Display for FormatError {
}
impl Error for FormatError {}
impl FormatError {
pub fn diagnostic(&self) -> Value {
let (code, operation, field, reason) = match self {
Self::UnsupportedFormat(_) => ("unsupported_format", "select_format", None, None),
Self::RequestParseFailed { .. } => {
("request_parse_failed", "parse_request", None, None)
}
Self::RequestEmitFailed { .. } => ("request_emit_failed", "emit_request", None, None),
Self::ResponseParseFailed { .. } => {
("response_parse_failed", "parse_response", None, None)
}
Self::ResponseEmitFailed { .. } => {
("response_emit_failed", "emit_response", None, None)
}
Self::UnsupportedField { field, reason, .. } => {
("unsupported_field", "validate", Some(field), Some(reason))
}
Self::UnauditedField { field, reason, .. } => {
("unaudited_field", "validate", Some(field), Some(reason))
}
Self::InvalidEnumValue { field, .. } => {
("invalid_enum_value", "validate", Some(field), None)
}
Self::LossyConversionBlocked { field, reason, .. } => (
"lossy_conversion_blocked",
"convert",
Some(field),
Some(reason),
),
Self::InvalidTargetField { field, reason, .. } => (
"invalid_target_field",
"validate_target",
Some(field),
Some(reason),
),
};
let path = field.map(|field| {
let path = if field.starts_with('$') {
field.clone()
} else {
format!("$.{field}")
};
path.replace("[]", "[*]")
});
let mut diagnostic = json!({
"code": code,
"operation": operation,
"path": path.as_deref().unwrap_or("$"),
"path_source": if path.is_some() { "structured" } else { "unavailable" },
"reason": reason,
"expected": reason,
"actual": null,
"missing_context": if path.is_some() { vec![] } else { vec!["field_path", "underlying_cause"] }
});
if let Self::InvalidEnumValue { value, .. } = self {
diagnostic["actual"] = json!(value);
}
match self {
Self::UnsupportedFormat(format)
| Self::RequestParseFailed { format }
| Self::RequestEmitFailed { format }
| Self::ResponseParseFailed { format }
| Self::ResponseEmitFailed { format }
| Self::UnsupportedField { format, .. }
| Self::InvalidEnumValue { format, .. }
| Self::InvalidTargetField { format, .. } => diagnostic["format"] = json!(format),
_ => {}
}
if let Self::UnauditedField {
source_format,
target_format,
..
}
| Self::LossyConversionBlocked {
source_format,
target_format,
..
} = self
{
diagnostic["source_format"] = json!(source_format);
diagnostic["target_format"] = json!(target_format);
}
diagnostic
}
}
#[cfg(test)]
mod diagnostic_tests {
use super::FormatError;
use serde_json::json;
#[test]
fn enum_diagnostic_retains_full_path_and_actual_value() {
let diagnostic = FormatError::InvalidEnumValue {
format: "openai:chat".to_string(),
field: "choices[].finish_reason".to_string(),
value: "future_reason".to_string(),
}
.diagnostic();
assert_eq!(diagnostic["code"], "invalid_enum_value");
assert_eq!(diagnostic["path"], "$.choices[*].finish_reason");
assert_eq!(diagnostic["actual"], "future_reason");
assert_eq!(diagnostic["format"], "openai:chat");
}
#[test]
fn generic_parse_failure_reports_missing_cause_without_a_fake_path() {
let diagnostic = FormatError::ResponseParseFailed {
format: "claude:messages".to_string(),
}
.diagnostic();
assert_eq!(diagnostic["code"], "response_parse_failed");
assert_eq!(diagnostic["operation"], "parse_response");
assert_eq!(diagnostic["path_source"], "unavailable");
assert_eq!(
diagnostic["missing_context"],
json!(["field_path", "underlying_cause"])
);
}
}
@@ -23,6 +23,9 @@ use crate::formats::shared::stream_core::common::{
};
use crate::formats::shared::AiSurfaceFinalizeError;
const PROVIDER_STREAM_FINISH_ERROR_MESSAGE: &str =
"Upstream stream ended with finish reason: error";
#[derive(Default)]
pub struct StreamingStandardFormatMatrix {
provider: Option<ProviderStreamParser>,
@@ -129,7 +132,7 @@ impl StreamingStandardFormatMatrix {
{
if !canonical_stream_finish_reason_is_supported(finish_reason) {
self.terminated = true;
out.extend(client.emit_unsupported_finish_reason(finish_reason)?);
out.extend(client.emit_finish_reason_error(finish_reason)?);
break;
}
}
@@ -388,7 +391,13 @@ impl StreamingStandardTerminalObserver {
if let Some(parser_error) = finish_reason
.as_deref()
.filter(|reason| !canonical_stream_finish_reason_is_supported(reason))
.map(|reason| format!("unsupported provider stream finish reason: {reason}"))
.map(|reason| {
if reason.trim() == "error" {
PROVIDER_STREAM_FINISH_ERROR_MESSAGE.to_string()
} else {
format!("unsupported provider stream finish reason: {reason}")
}
})
{
summary.parser_error.get_or_insert(parser_error);
}
@@ -660,17 +669,28 @@ impl ClientStreamEmitter {
self.emit_error(error_body)
}
fn emit_unsupported_finish_reason(
fn emit_finish_reason_error(
&mut self,
finish_reason: &str,
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
let (message, code) = if finish_reason.trim() == "error" {
(
PROVIDER_STREAM_FINISH_ERROR_MESSAGE.to_string(),
"stream_terminal_error",
)
} else {
(
format!(
"Unsupported provider stream finish reason cannot be converted losslessly: field $.finish_reason = {}",
serde_json::json!(finish_reason)
),
"unsupported_finish_reason",
)
};
let Some(error_body) = build_core_error_body_for_client_format(
self.api_format(),
&format!(
"Unsupported provider stream finish reason cannot be converted losslessly: field $.finish_reason = {}",
serde_json::json!(finish_reason)
),
Some("unsupported_finish_reason"),
&message,
Some(code),
LocalCoreSyncErrorKind::ServerError,
) else {
return Ok(Vec::new());
@@ -2121,6 +2141,163 @@ mod tests {
assert!(sse.contains("\"stop_reason\":\"tool_use\""), "{sse}");
}
#[test]
fn terminal_observer_treats_claude_error_finish_reason_as_upstream_failure() {
for upstream_message in [None, Some("Provider temporarily overloaded")] {
let context = report_context("claude:messages", "claude:messages");
let mut observer = StreamingStandardTerminalObserver::default();
observer
.push_line(
&context,
data_line(json!({
"type": "message_start",
"message": {
"id": "msg_error_finish",
"model": "claude-sonnet-4-5",
"usage": {
"input_tokens": 22,
"cache_read_input_tokens": 7,
"cache_creation_input_tokens": 3,
"cache_creation": { "ephemeral_5m_input_tokens": 3 }
}
}
})),
)
.expect("message start should be observed");
if let Some(message) = upstream_message {
observer
.push_line(
&context,
data_line(json!({
"type": "error",
"error": { "type": "overloaded_error", "message": message }
})),
)
.expect("upstream error should be observed");
}
observer
.push_line(
&context,
data_line(json!({
"type": "message_delta",
"delta": { "stop_reason": "error" },
"usage": { "output_tokens": 5 }
})),
)
.expect("error finish reason should be observed");
let summary = observer
.finish(&context)
.expect("terminal observation should finish")
.expect("failed stream should have a summary");
assert!(summary.observed_finish);
assert_eq!(summary.finish_reason.as_deref(), Some("error"));
assert_eq!(
summary.parser_error.as_deref(),
Some(upstream_message.unwrap_or("Upstream stream ended with finish reason: error"))
);
let usage = summary
.standardized_usage
.expect("usage should be retained");
assert_eq!(usage.input_tokens, 22);
assert_eq!(usage.output_tokens, 5);
assert_eq!(usage.cache_read_tokens, 7);
assert_eq!(usage.cache_creation_tokens, 3);
assert_eq!(usage.cache_creation_ephemeral_5m_tokens, 3);
}
}
#[test]
fn transforms_claude_error_finish_reason_to_terminal_errors() {
let cases = [
(
"openai:chat",
"data: {\"error\":",
"\"code\":\"stream_terminal_error\"",
),
(
"openai:responses",
"event: response.failed\n",
"\"code\":\"stream_terminal_error\"",
),
(
"claude:messages",
"event: error\n",
"\"code\":\"stream_terminal_error\"",
),
(
"gemini:generate_content",
"data: {\"error\":",
"\"status\":\"INTERNAL\"",
),
];
for (client_api_format, prefix, marker) in cases {
let context = report_context("claude:messages", client_api_format);
let mut matrix = StreamingStandardFormatMatrix::default();
let mut output = matrix
.transform_line(
&context,
data_line(json!({
"type": "content_block_delta",
"index": 0,
"delta": { "type": "text_delta", "text": "Partial answer" }
})),
)
.expect("partial response should be emitted");
output.extend(
matrix
.transform_line(
&context,
data_line(json!({
"type": "message_delta",
"delta": { "stop_reason": "error" },
"usage": { "output_tokens": 5 }
})),
)
.expect("error finish reason should emit a terminal error"),
);
let sse = String::from_utf8(output).expect("sse should be utf8");
assert!(sse.contains("Partial answer"), "{client_api_format}: {sse}");
assert!(sse.contains(prefix), "{client_api_format}: {sse}");
assert!(sse.contains(marker), "{client_api_format}: {sse}");
assert!(
sse.contains("Upstream stream ended with finish reason: error"),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("unsupported_finish_reason"),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("response.completed"),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("\"stop_reason\":\"end_turn\""),
"{client_api_format}: {sse}"
);
assert!(
!sse.contains("\"finish_reason\":\"stop\""),
"{client_api_format}: {sse}"
);
assert!(matrix
.transform_line(
&context,
data_line(json!({
"type": "message_delta",
"delta": { "stop_reason": "end_turn" }
}))
)
.expect("events after the error should be ignored")
.is_empty());
assert!(matrix
.finish(&context)
.expect("failed matrix should stay terminated")
.is_empty());
}
}
#[test]
fn transforms_unknown_stream_finish_reasons_to_visible_client_errors() {
let cases = [
@@ -34,6 +34,7 @@ pub struct CandidateFailureDiagnostic {
client_api_format: Option<String>,
provider_api_format: Option<String>,
safe_to_show: bool,
details: Option<Value>,
}
impl CandidateFailureDiagnostic {
@@ -50,6 +51,7 @@ impl CandidateFailureDiagnostic {
client_api_format: None,
provider_api_format: None,
safe_to_show: true,
details: None,
}
}
@@ -58,6 +60,11 @@ impl CandidateFailureDiagnostic {
self
}
pub fn details(mut self, details: Value) -> Self {
self.details = Some(details);
self
}
pub fn formats(
mut self,
client_api_format: impl Into<String>,
@@ -219,6 +226,10 @@ impl CandidateFailureDiagnostic {
"client_api_format": self.client_api_format,
"provider_api_format": self.provider_api_format,
"safe_to_show": self.safe_to_show,
"details": self.details,
"stage": "request",
"source_format": self.client_api_format,
"target_format": self.provider_api_format,
})
}
}
@@ -224,6 +224,7 @@ fn diagnostic_from_format_error(
format_error_path(error),
format_error_message(error, client_api_format, provider_api_format),
)
.details(error.diagnostic())
}
fn format_error_path(error: &FormatError) -> String {
@@ -982,6 +983,23 @@ mod tests {
"request_conversion"
);
assert_eq!(diagnostic["failure_diagnostic"]["path"], "$.n");
assert_eq!(diagnostic["failure_diagnostic"]["stage"], "request");
assert_eq!(
diagnostic["failure_diagnostic"]["details"]["code"],
"lossy_conversion_blocked"
);
assert_eq!(
diagnostic["failure_diagnostic"]["details"]["path_source"],
"structured"
);
assert_eq!(
diagnostic["failure_diagnostic"]["source_format"],
"openai:chat"
);
assert_eq!(
diagnostic["failure_diagnostic"]["target_format"],
"openai:responses"
);
assert_eq!(diagnostic["request_conversion_error"]["path"], "$.n");
assert!(diagnostic["failure_diagnostic"]["message"]
.as_str()
+124 -62
View File
@@ -1,8 +1,8 @@
use aether_data_contracts::repository::billing::StoredBillingModelContext;
use aether_data_contracts::repository::usage::{
extract_provider_cache_ttl_minutes_from_metadata, resolve_provider_cache_ttl_minutes,
resolve_provider_service_tier_from_request_capture, USAGE_AVAILABLE_METADATA_KEY,
USAGE_PRICING_AVAILABLE_METADATA_KEY,
resolve_provider_service_tier_from_request_capture, CANCELLED_REQUEST_FEE_METADATA_KEY,
USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY,
};
use aether_data_contracts::DataLayerError;
use aether_usage_runtime::{UsageEvent, UsageEventType};
@@ -40,6 +40,18 @@ pub async fn enrich_usage_event_with_billing(
data: &dyn BillingModelContextLookup,
event: &mut UsageEvent,
) -> Result<(), DataLayerError> {
if matches!(event.event_type, UsageEventType::Cancelled) {
event.data.total_cost_usd = Some(0.0);
event.data.actual_total_cost_usd = Some(0.0);
if let Some(metadata) = event
.data
.request_metadata
.as_mut()
.and_then(Value::as_object_mut)
{
metadata.remove(CANCELLED_REQUEST_FEE_METADATA_KEY);
}
}
// Session transports such as Codex Live expose lifecycle telemetry but no
// authoritative token/cost object. Do not run request-based pricing with
// zero default tokens: that would turn "unknown" into a fabricated charge.
@@ -65,7 +77,10 @@ pub async fn enrich_usage_event_with_billing(
clear_usage_costs(event);
return Ok(());
}
if !matches!(event.event_type, UsageEventType::Completed) {
if !matches!(
event.event_type,
UsageEventType::Completed | UsageEventType::Cancelled
) {
event.data.total_cost_usd = Some(0.0);
event.data.actual_total_cost_usd = Some(0.0);
return Ok(());
@@ -189,7 +204,10 @@ fn calculate_billing_computation(
} else {
usage_event_image_count(&event.data).unwrap_or(0)
};
let request_count = if failed {
let cancelled = matches!(event.event_type, UsageEventType::Cancelled);
let request_count = if cancelled {
1
} else if failed {
0
} else if is_image_usage && image_count > 0 {
image_count
@@ -197,7 +215,7 @@ fn calculate_billing_computation(
1
};
let processing_tiers = usage_event_processing_tiers(&event.data);
let input = BillingUsageInput {
let mut input = BillingUsageInput {
task_type: if is_image_usage {
"image".to_string()
} else {
@@ -237,6 +255,16 @@ fn calculate_billing_computation(
.or(pricing.provider_api_key_cache_ttl_minutes),
};
if cancelled {
input.input_tokens = 0;
input.output_tokens = 0;
input.cache_creation_tokens = 0;
input.cache_creation_ephemeral_5m_tokens = 0;
input.cache_creation_ephemeral_1h_tokens = 0;
input.cache_read_tokens = 0;
input.image_count = 0;
}
BillingService::new()
.calculate(pricing, &input)
.map_err(|err| {
@@ -356,9 +384,32 @@ fn apply_billing_computation(
pricing: &BillingModelPricingSnapshot,
computation: BillingComputation,
) -> Result<(), DataLayerError> {
let cancelled = matches!(event.event_type, UsageEventType::Cancelled);
if cancelled
&& !computation
.pricing_resolution
.price_per_request
.is_some_and(|price| price > 0.0)
{
return Ok(());
}
event.data.total_cost_usd = Some(computation.cost_result.cost);
event.data.actual_total_cost_usd = Some(computation.actual_total_cost);
merge_billing_snapshot_metadata(&mut event.data.request_metadata, pricing, &computation)
merge_billing_snapshot_metadata(&mut event.data.request_metadata, pricing, &computation)?;
if cancelled {
if let Some(metadata) = event
.data
.request_metadata
.as_mut()
.and_then(Value::as_object_mut)
{
metadata.insert(
CANCELLED_REQUEST_FEE_METADATA_KEY.to_string(),
Value::Bool(true),
);
}
}
Ok(())
}
fn map_pricing_context(context: StoredBillingModelContext) -> BillingModelPricingSnapshot {
@@ -1272,8 +1323,11 @@ mod tests {
}
#[tokio::test]
async fn cancelled_usage_event_remains_unbilled() {
let lookup = TestLookup {
async fn cancelled_usage_bills_only_configured_request_fee() {
for (request_type, request_price) in
[("chat", None), ("chat", Some(0.02)), ("image", Some(0.02))]
{
let lookup = TestLookup {
name_context: Some(
StoredBillingModelContext::new(
"provider-1".to_string(),
@@ -1284,7 +1338,7 @@ mod tests {
"global-model-1".to_string(),
"gpt-5".to_string(),
None,
Some(0.02),
request_price,
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":3.0,"output_price_per_1m":15.0,"cache_creation_price_per_1m":3.75,"cache_read_price_per_1m":0.30}]})),
Some("model-1".to_string()),
Some("gpt-5-upstream".to_string()),
@@ -1296,61 +1350,69 @@ mod tests {
),
model_id_context: None,
};
let mut event = UsageEvent::new(
UsageEventType::Cancelled,
"req-billing-cancelled-1",
UsageEventData {
provider_name: "OpenAI".to_string(),
model: "gpt-5".to_string(),
provider_id: Some("provider-1".to_string()),
provider_api_key_id: Some("key-1".to_string()),
request_type: Some("chat".to_string()),
api_format: Some("openai:responses".to_string()),
endpoint_api_format: Some("openai:responses".to_string()),
input_tokens: Some(1_000),
output_tokens: Some(500),
cache_read_input_tokens: Some(100),
status_code: Some(499),
..UsageEventData::default()
},
);
let mut event = UsageEvent::new(
UsageEventType::Cancelled,
"req-billing-cancelled-1",
UsageEventData {
provider_name: "OpenAI".to_string(),
model: "gpt-5".to_string(),
provider_id: Some("provider-1".to_string()),
provider_api_key_id: Some("key-1".to_string()),
request_type: Some(request_type.to_string()),
api_format: Some("openai:responses".to_string()),
endpoint_api_format: Some("openai:responses".to_string()),
input_tokens: Some(1_000),
output_tokens: Some(500),
cache_read_input_tokens: Some(100),
status_code: Some(499),
request_metadata: Some(
json!({"cancelled_request_fee": true, "image_count": 3}),
),
..UsageEventData::default()
},
);
enrich_usage_event_with_billing(&lookup, &mut event)
.await
.expect("billing should succeed");
enrich_usage_event_with_billing(&lookup, &mut event)
.await
.expect("billing should succeed");
assert_eq!(event.data.total_cost_usd, Some(0.0));
assert_eq!(event.data.actual_total_cost_usd, Some(0.0));
assert_eq!(
event
.data
.request_metadata
.as_ref()
.and_then(|value| value.get("billing_snapshot"))
.and_then(|value| value.get("status"))
.and_then(Value::as_str),
None
);
assert_eq!(
event
.data
.request_metadata
.as_ref()
.and_then(|value| value.get("billing_dimensions"))
.and_then(|value| value.get("input_tokens"))
.and_then(Value::as_i64),
None
);
assert_eq!(
event
.data
.request_metadata
.as_ref()
.and_then(|value| value.get("billing_dimensions"))
.and_then(|value| value.get("cache_read_tokens"))
.and_then(Value::as_i64),
None
);
let expected_cost = request_price.unwrap_or(0.0);
assert_eq!(event.data.total_cost_usd, Some(expected_cost));
assert_eq!(event.data.actual_total_cost_usd, Some(expected_cost * 0.5));
assert_eq!(event.data.input_tokens, Some(1_000));
assert_eq!(event.data.output_tokens, Some(500));
let metadata = event.data.request_metadata.as_ref().unwrap();
assert_eq!(
aether_data_contracts::repository::usage::cancelled_request_fee_is_billable(Some(
metadata
)),
request_price.is_some()
);
if request_price.is_some() {
assert_eq!(
metadata.pointer("/billing_snapshot/cost_breakdown/request_cost"),
Some(&json!(expected_cost))
);
assert_eq!(
metadata.pointer("/billing_dimensions/input_tokens"),
Some(&json!(0))
);
assert_eq!(
metadata.pointer("/billing_dimensions/output_tokens"),
Some(&json!(0))
);
assert_eq!(
metadata.pointer("/billing_dimensions/cache_read_tokens"),
Some(&json!(0))
);
assert_eq!(
metadata.pointer("/billing_dimensions/request_count"),
Some(&json!(1))
);
} else {
assert!(metadata.get("billing_snapshot").is_none());
}
}
}
#[tokio::test]
@@ -4677,36 +4677,11 @@ VALUES (
.await
.map_postgres_err()?;
let order_row = sqlx::query(
r#"
SELECT
id,
order_no,
wallet_id,
user_id,
CAST(amount_usd AS DOUBLE PRECISION) AS amount_usd,
CAST(pay_amount AS DOUBLE PRECISION) AS pay_amount,
pay_currency,
CAST(exchange_rate AS DOUBLE PRECISION) AS exchange_rate,
CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd,
CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd,
payment_method,
gateway_order_id,
gateway_response,
status,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM paid_at) AS BIGINT) AS paid_at_unix_secs,
CAST(EXTRACT(EPOCH FROM credited_at) AS BIGINT) AS credited_at_unix_secs,
CAST(EXTRACT(EPOCH FROM expires_at) AS BIGINT) AS expires_at_unix_secs
FROM payment_orders
WHERE id = $1
LIMIT 1
"#,
)
.bind(&order_id)
.fetch_one(&mut **tx)
.await
.map_postgres_err()?;
let order_row = sqlx::query(FIND_ADMIN_PAYMENT_ORDER_SQL)
.bind(&order_id)
.fetch_one(&mut **tx)
.await
.map_postgres_err()?;
Ok(Some((wallet, map_admin_payment_order_row(&order_row)?)))
})
})
@@ -5894,6 +5869,8 @@ RETURNING
CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd,
CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd,
payment_method,
payment_provider,
order_kind,
gateway_order_id,
gateway_response,
status,
@@ -5939,6 +5916,8 @@ SELECT
CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd,
CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd,
payment_method,
payment_provider,
order_kind,
gateway_order_id,
gateway_response,
status,
@@ -5993,6 +5972,8 @@ RETURNING
CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd,
CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd,
payment_method,
payment_provider,
order_kind,
gateway_order_id,
gateway_response,
status,
@@ -6259,6 +6240,8 @@ RETURNING
CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd,
CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd,
payment_method,
payment_provider,
order_kind,
gateway_order_id,
gateway_response,
status,
@@ -6458,6 +6441,8 @@ RETURNING
CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd,
CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd,
payment_method,
payment_provider,
order_kind,
gateway_order_id,
gateway_response,
status,
@@ -7298,6 +7283,8 @@ RETURNING
CAST(refunded_amount_usd AS DOUBLE PRECISION) AS refunded_amount_usd,
CAST(refundable_amount_usd AS DOUBLE PRECISION) AS refundable_amount_usd,
payment_method,
payment_provider,
order_kind,
gateway_order_id,
gateway_response,
status,
@@ -8490,9 +8477,543 @@ VALUES ($1, $2, 'gift', 'gift_initial', $3, 0, $3, 0, 0, 0, $3, 'system_task', $
#[cfg(test)]
mod tests {
use aether_data_contracts::repository::wallet::{
CreateManualWalletRechargeInput, CreditAdminPaymentOrderInput, RedeemWalletCodeInput,
RedeemWalletCodeOutcome, WalletLookupKey, WalletMutationOutcome, WalletReadRepository,
WalletWriteRepository,
};
use sqlx::Row;
use super::SqlxWalletRepository;
use crate::{PostgresPoolConfig, PostgresPoolFactory};
#[test]
fn payment_order_sql_projections_cover_mapper_columns() {
let source = include_str!("wallet.rs")
.split("#[cfg(test)]")
.next()
.expect("wallet implementation should exist");
let mapper = source
.split("fn map_admin_payment_order_row(")
.nth(1)
.expect("payment order mapper should exist")
.split("\nfn ")
.next()
.expect("payment order mapper body should exist");
let required_columns = mapper
.split("row_get(row, \"")
.skip(1)
.map(|read| read.split('"').next().expect("column name should exist"))
.collect::<Vec<_>>();
assert!(required_columns.contains(&"payment_provider"));
assert!(required_columns.contains(&"order_kind"));
let mut projections_checked = 0;
for fragment in source.split("r#\"").skip(1) {
let Some((sql, _)) = fragment.split_once("\"#") else {
continue;
};
if !sql.contains("payment_orders")
|| !sql.contains("AS created_at_unix_ms")
|| !sql.contains("pay_currency")
{
continue;
}
let projection = match sql.rsplit_once("RETURNING") {
Some((_, projection)) => projection.to_string(),
None => sql
.lines()
.take_while(|line| !line.trim_start().starts_with("FROM payment_orders"))
.collect::<Vec<_>>()
.join("\n"),
};
let tokens = projection
.split(|character: char| !character.is_ascii_alphanumeric() && character != '_')
.collect::<Vec<_>>();
for column in &required_columns {
assert!(
tokens.contains(column),
"payment order projection omits {column}: {sql}"
);
}
projections_checked += 1;
}
assert!(
projections_checked > 0,
"payment order projections should exist"
);
}
async fn isolated_wallet_test_pool() -> sqlx::PgPool {
let database_url = std::env::var("AETHER_TEST_DATABASE_URL")
.expect("AETHER_TEST_DATABASE_URL must point at the test database");
let options = database_url
.parse::<sqlx::postgres::PgConnectOptions>()
.expect("test database URL should parse")
.options([("search_path", "pg_temp")]);
let pool = sqlx::postgres::PgPoolOptions::new()
.max_connections(1)
.idle_timeout(None)
.max_lifetime(None)
.connect_with(options)
.await
.expect("test database should connect");
for table in [
"wallets",
"payment_orders",
"wallet_transactions",
"user_plan_entitlements",
"redeem_code_batches",
"redeem_codes",
] {
sqlx::query(&format!(
"CREATE TEMP TABLE {table} (LIKE public.{table} INCLUDING ALL)"
))
.execute(&pool)
.await
.expect("isolated wallet table should be created");
}
pool
}
async fn seed_wallet(pool: &sqlx::PgPool) -> (String, String) {
let wallet_id = uuid::Uuid::new_v4().to_string();
let user_id = uuid::Uuid::new_v4().to_string();
sqlx::query(
"INSERT INTO wallets (id, user_id, balance, gift_balance, total_recharged, created_at, updated_at) VALUES ($1, $2, 10, 3, 20, NOW(), NOW())",
)
.bind(&wallet_id)
.bind(&user_id)
.execute(pool)
.await
.expect("test wallet should be created");
(wallet_id, user_id)
}
async fn seed_pending_order(
pool: &sqlx::PgPool,
wallet_id: &str,
user_id: &str,
order_kind: &str,
) -> String {
let order_id = uuid::Uuid::new_v4().to_string();
let plan_snapshot = (order_kind == "plan_purchase").then(|| {
serde_json::json!({
"id": "test-plan",
"duration_days": 30,
"purchase_limit_scope": "unlimited",
"entitlements": [{
"type": "wallet_credit",
"amount_usd": 4.0,
"balance_bucket": "gift",
}],
})
});
sqlx::query(
"INSERT INTO payment_orders (id, order_no, wallet_id, user_id, amount_usd, pay_amount, pay_currency, payment_method, payment_provider, order_kind, product_id, product_snapshot, status, created_at, expires_at) VALUES ($1, $2, $3, $4, 5, 5, 'USD', 'stripe', 'stripe', $5, $6, $7, 'pending', NOW(), NOW() + INTERVAL '1 hour')",
)
.bind(&order_id)
.bind(format!("order-{order_id}"))
.bind(wallet_id)
.bind(user_id)
.bind(order_kind)
.bind(plan_snapshot.as_ref().map(|_| "test-plan"))
.bind(plan_snapshot)
.execute(pool)
.await
.expect("pending payment order should be created");
order_id
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL bootstrap schema"]
async fn live_manual_recharge_commits_wallet_order_and_transaction() {
let pool = isolated_wallet_test_pool().await;
let (wallet_id, user_id) = seed_wallet(&pool).await;
let repository = SqlxWalletRepository::new(pool.clone());
let input = CreateManualWalletRechargeInput {
wallet_id: wallet_id.clone(),
amount_usd: 5.0,
payment_method: "admin_manual".to_string(),
operator_id: Some(uuid::Uuid::new_v4().to_string()),
description: Some("manual recharge regression".to_string()),
order_no: format!("manual-{}", uuid::Uuid::new_v4()),
};
let (wallet, order) = repository
.create_manual_wallet_recharge(input.clone())
.await
.expect("manual recharge should commit")
.expect("test wallet should exist");
assert_eq!(wallet.id, wallet_id);
assert_eq!(wallet.user_id.as_deref(), Some(user_id.as_str()));
assert_eq!(wallet.balance, 15.0);
assert_eq!(wallet.gift_balance, 3.0);
assert_eq!(wallet.total_recharged, 25.0);
assert_eq!(order.order_no, input.order_no);
assert_eq!(order.wallet_id, wallet_id);
assert_eq!(order.user_id.as_deref(), Some(user_id.as_str()));
assert_eq!(order.amount_usd, 5.0);
assert_eq!(order.refunded_amount_usd, 0.0);
assert_eq!(order.refundable_amount_usd, 5.0);
assert_eq!(order.payment_method, "admin_manual");
assert_eq!(order.payment_provider, None);
assert_eq!(order.order_kind, "wallet_recharge");
assert_eq!(order.status, "credited");
assert!(order.paid_at_unix_secs.is_some());
assert!(order.credited_at_unix_secs.is_some());
assert_eq!(
order.gateway_response,
Some(serde_json::json!({
"source": "manual",
"operator_id": input.operator_id,
"description": input.description,
}))
);
assert!(repository
.create_manual_wallet_recharge(input.clone())
.await
.is_err());
for amount_usd in [0.0, -1.0, f64::NAN, f64::INFINITY] {
assert!(repository
.create_manual_wallet_recharge(CreateManualWalletRechargeInput {
amount_usd,
..input.clone()
})
.await
.is_err());
}
assert!(repository
.create_manual_wallet_recharge(CreateManualWalletRechargeInput {
wallet_id: uuid::Uuid::new_v4().to_string(),
..input.clone()
})
.await
.expect("missing wallet should not fail")
.is_none());
let persisted_wallet = repository
.find(WalletLookupKey::WalletId(&wallet_id))
.await
.expect("wallet should be readable after commit")
.expect("wallet should persist");
assert_eq!(persisted_wallet, wallet);
let persisted_order = repository
.find_admin_payment_order(&order.id)
.await
.expect("payment order should be readable after commit")
.expect("payment order should persist");
assert_eq!(persisted_order, order);
let transaction = sqlx::query(
"SELECT category, reason_code, CAST(amount AS DOUBLE PRECISION) AS amount, CAST(balance_before AS DOUBLE PRECISION) AS balance_before, CAST(balance_after AS DOUBLE PRECISION) AS balance_after, CAST(recharge_balance_before AS DOUBLE PRECISION) AS recharge_balance_before, CAST(recharge_balance_after AS DOUBLE PRECISION) AS recharge_balance_after, CAST(gift_balance_before AS DOUBLE PRECISION) AS gift_balance_before, CAST(gift_balance_after AS DOUBLE PRECISION) AS gift_balance_after, link_type, link_id, operator_id, description FROM wallet_transactions WHERE wallet_id = $1",
)
.bind(&wallet_id)
.fetch_one(&pool)
.await
.expect("recharge transaction should persist");
assert_eq!(transaction.get::<String, _>("category"), "recharge");
assert_eq!(
transaction.get::<String, _>("reason_code"),
"topup_admin_manual"
);
assert_eq!(transaction.get::<f64, _>("amount"), 5.0);
assert_eq!(transaction.get::<f64, _>("balance_before"), 13.0);
assert_eq!(transaction.get::<f64, _>("balance_after"), 18.0);
assert_eq!(transaction.get::<f64, _>("recharge_balance_before"), 10.0);
assert_eq!(transaction.get::<f64, _>("recharge_balance_after"), 15.0);
assert_eq!(transaction.get::<f64, _>("gift_balance_before"), 3.0);
assert_eq!(transaction.get::<f64, _>("gift_balance_after"), 3.0);
assert_eq!(transaction.get::<String, _>("link_type"), "payment_order");
assert_eq!(transaction.get::<String, _>("link_id"), order.id);
assert_eq!(
transaction.get::<Option<String>, _>("operator_id"),
input.operator_id
);
assert_eq!(
transaction.get::<Option<String>, _>("description"),
input.description
);
for table in ["wallets", "payment_orders", "wallet_transactions"] {
let count: i64 = sqlx::query_scalar(&format!("SELECT COUNT(*) FROM {table}"))
.fetch_one(&pool)
.await
.expect("wallet record count should be readable");
assert_eq!(count, 1, "rejected recharges must not add {table} rows");
}
pool.close().await;
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL bootstrap schema"]
async fn live_admin_order_state_changes_preserve_metadata() {
let pool = isolated_wallet_test_pool().await;
let (wallet_id, user_id) = seed_wallet(&pool).await;
let repository = SqlxWalletRepository::new(pool.clone());
for target_status in ["expired", "failed"] {
let order_id = seed_pending_order(&pool, &wallet_id, &user_id, "wallet_recharge").await;
let order = if target_status == "expired" {
let outcome = repository
.expire_admin_payment_order(&order_id)
.await
.expect("order expiry should commit");
let WalletMutationOutcome::Applied((order, changed)) = outcome else {
panic!("pending order should expire");
};
assert!(changed);
assert!(matches!(
repository.expire_admin_payment_order(&order_id).await,
Ok(WalletMutationOutcome::Applied((_, false)))
));
order
} else {
let outcome = repository
.fail_admin_payment_order(&order_id)
.await
.expect("order failure should commit");
let WalletMutationOutcome::Applied(order) = outcome else {
panic!("pending order should be marked failed");
};
order
};
assert_eq!(order.status, target_status);
assert_eq!(order.payment_provider.as_deref(), Some("stripe"));
assert_eq!(order.order_kind, "wallet_recharge");
assert_eq!(
repository
.find_admin_payment_order(&order_id)
.await
.unwrap(),
Some(order)
);
}
let wallet = repository
.find(WalletLookupKey::WalletId(&wallet_id))
.await
.unwrap()
.unwrap();
assert_eq!(
(wallet.balance, wallet.gift_balance, wallet.total_recharged),
(10.0, 3.0, 20.0)
);
let transaction_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions")
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(transaction_count, 0);
pool.close().await;
}
async fn assert_admin_payment_order_credit(order_kind: &str) {
let pool = isolated_wallet_test_pool().await;
let (wallet_id, user_id) = seed_wallet(&pool).await;
let order_id = seed_pending_order(&pool, &wallet_id, &user_id, order_kind).await;
let repository = SqlxWalletRepository::new(pool.clone());
let input = CreditAdminPaymentOrderInput {
order_id: order_id.clone(),
gateway_order_id: None,
pay_amount: None,
pay_currency: None,
exchange_rate: None,
gateway_response_patch: None,
operator_id: Some(uuid::Uuid::new_v4().to_string()),
};
let outcome = repository
.credit_admin_payment_order(input.clone())
.await
.expect("admin credit should commit");
let WalletMutationOutcome::Applied((order, changed)) = outcome else {
panic!("pending payment order should be credited");
};
assert!(changed);
assert_eq!(order.status, "credited");
assert_eq!(order.payment_provider.as_deref(), Some("stripe"));
assert_eq!(order.order_kind, order_kind);
assert!(order.paid_at_unix_secs.is_some());
assert!(order.credited_at_unix_secs.is_some());
assert_eq!(
repository
.find_admin_payment_order(&order_id)
.await
.unwrap(),
Some(order.clone())
);
assert!(matches!(
repository.credit_admin_payment_order(input).await,
Ok(WalletMutationOutcome::Applied((_, false)))
));
assert!(matches!(
repository.expire_admin_payment_order(&order_id).await,
Ok(WalletMutationOutcome::Invalid(_))
));
assert!(matches!(
repository.fail_admin_payment_order(&order_id).await,
Ok(WalletMutationOutcome::Invalid(_))
));
let wallet = repository
.find(WalletLookupKey::WalletId(&wallet_id))
.await
.unwrap()
.unwrap();
let expected_balances = if order_kind == "plan_purchase" {
assert_eq!(order.refundable_amount_usd, 0.0);
let fulfillment: String =
sqlx::query_scalar("SELECT fulfillment_status FROM payment_orders WHERE id = $1")
.bind(&order_id)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(fulfillment, "fulfilled");
(10.0, 7.0, 20.0)
} else {
assert_eq!(order.refundable_amount_usd, 5.0);
(15.0, 3.0, 25.0)
};
assert_eq!(
(wallet.balance, wallet.gift_balance, wallet.total_recharged),
expected_balances
);
let transaction_count: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = $1")
.bind(&wallet_id)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(transaction_count, 1);
let entitlement_count: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM user_plan_entitlements WHERE payment_order_id = $1",
)
.bind(&order_id)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(entitlement_count, i64::from(order_kind == "plan_purchase"));
pool.close().await;
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL bootstrap schema"]
async fn live_admin_wallet_order_credit_commits_once() {
assert_admin_payment_order_credit("wallet_recharge").await;
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL bootstrap schema"]
async fn live_admin_plan_order_credit_commits_once() {
assert_admin_payment_order_credit("plan_purchase").await;
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL bootstrap schema"]
async fn live_redeem_code_commits_order_and_wallet_once_for_each_bucket() {
let pool = isolated_wallet_test_pool().await;
let repository = SqlxWalletRepository::new(pool.clone());
for balance_bucket in ["recharge", "gift"] {
let (wallet_id, user_id) = seed_wallet(&pool).await;
let batch_id = uuid::Uuid::new_v4().to_string();
let code_id = uuid::Uuid::new_v4().to_string();
let code = super::generate_redeem_code_normalized();
sqlx::query(
"INSERT INTO redeem_code_batches (id, name, amount_usd, balance_bucket, total_count, created_at, updated_at) VALUES ($1, 'regression batch', 5, $2, 1, NOW(), NOW())",
)
.bind(&batch_id)
.bind(balance_bucket)
.execute(&pool)
.await
.unwrap();
sqlx::query(
"INSERT INTO redeem_codes (id, batch_id, code_hash, code_prefix, code_suffix, created_at, updated_at) VALUES ($1, $2, $3, $4, $5, NOW(), NOW())",
)
.bind(&code_id)
.bind(&batch_id)
.bind(super::hash_redeem_code(&code))
.bind(super::redeem_code_prefix(&code))
.bind(super::redeem_code_suffix(&code))
.execute(&pool)
.await
.unwrap();
let input = RedeemWalletCodeInput {
code: super::format_redeem_code(&code),
user_id,
order_no: format!("redeem-{}", uuid::Uuid::new_v4()),
};
let outcome = repository
.redeem_wallet_code(input.clone())
.await
.expect("redeem code recharge should commit");
let RedeemWalletCodeOutcome::Redeemed {
wallet,
order,
amount_usd,
..
} = outcome
else {
panic!("active code should be redeemed");
};
assert_eq!(amount_usd, 5.0);
assert_eq!(order.status, "credited");
assert_eq!(order.order_kind, "wallet_recharge");
assert_eq!(order.payment_provider, None);
let expected_balances = if balance_bucket == "recharge" {
assert_eq!(order.payment_method, "card_code");
assert_eq!(order.refundable_amount_usd, 5.0);
(15.0, 3.0, 25.0)
} else {
assert_eq!(order.payment_method, "gift_code");
assert_eq!(order.refundable_amount_usd, 0.0);
(10.0, 8.0, 25.0)
};
assert_eq!(
(wallet.balance, wallet.gift_balance, wallet.total_recharged),
expected_balances
);
assert!(matches!(
repository.redeem_wallet_code(input).await,
Ok(RedeemWalletCodeOutcome::CodeRedeemed)
));
assert_eq!(
repository
.find(WalletLookupKey::WalletId(&wallet_id))
.await
.unwrap(),
Some(wallet)
);
assert_eq!(
repository
.find_admin_payment_order(&order.id)
.await
.unwrap(),
Some(order.clone())
);
let redeemed = sqlx::query(
"SELECT status, redeemed_wallet_id, redeemed_payment_order_id FROM redeem_codes WHERE id = $1",
)
.bind(&code_id)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(redeemed.get::<String, _>("status"), "redeemed");
assert_eq!(redeemed.get::<String, _>("redeemed_wallet_id"), wallet_id);
assert_eq!(
redeemed.get::<String, _>("redeemed_payment_order_id"),
order.id
);
for table in ["payment_orders", "wallet_transactions"] {
let count: i64 = sqlx::query_scalar(&format!(
"SELECT COUNT(*) FROM {table} WHERE wallet_id = $1"
))
.bind(&wallet_id)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(count, 1, "redeeming twice must not duplicate {table}");
}
}
pool.close().await;
}
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {
let factory = PostgresPoolFactory::new(PostgresPoolConfig {
@@ -857,7 +857,18 @@ pub fn sanitize_request_candidate_extra_data_for_persistence(
("error_flow", &["message"][..]),
(
"failure_diagnostic",
&["path", "field_path", "message", "type", "reason"][..],
&[
"path",
"field_path",
"message",
"type",
"reason",
"details",
"stage",
"source_format",
"target_format",
"safe_to_show",
][..],
),
(
"request_conversion_error",
@@ -2445,7 +2456,7 @@ mod tests {
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"},
"failure_diagnostic": {"path": "$.input", "message": "private conversion failure", "safe_to_show": false, "stage": "request", "details": {"code": "invalid_enum_value", "actual": "private-value"}},
"request_body": {"input": "private prompt"}
});
let admin = super::sanitize_request_candidate_extra_data_for_persistence(Some(raw))
@@ -2456,6 +2467,12 @@ mod tests {
assert!(body.len() <= 65_536);
assert!(body.ends_with("...[truncated]"));
assert!(admin.get("request_body").is_none());
assert_eq!(admin["failure_diagnostic"]["safe_to_show"], false);
assert_eq!(admin["failure_diagnostic"]["stage"], "request");
assert_eq!(
admin["failure_diagnostic"]["details"]["code"],
"invalid_enum_value"
);
assert_eq!(
super::sanitize_request_candidate_extra_data_for_persistence(Some(admin.clone())),
Some(admin.clone()),
@@ -20,6 +20,14 @@ use super::{
const UPSTREAM_IS_STREAM_KEY: &str = "upstream_is_stream";
const PLAN_USAGE_RESERVATION_TOKEN_KEY: &str = "plan_usage_reservation_token";
const BODY_SIZE_BASIS: &str = "serialized gateway request bodies after normalization";
pub const CANCELLED_REQUEST_FEE_METADATA_KEY: &str = "cancelled_request_fee";
pub fn cancelled_request_fee_is_billable(metadata: Option<&Value>) -> bool {
metadata
.and_then(|metadata| metadata.get(CANCELLED_REQUEST_FEE_METADATA_KEY))
.and_then(Value::as_bool)
.unwrap_or(false)
}
/// Projects request metadata onto the persistence contract. Unknown fields and malformed values
/// are discarded instead of being recursively copied into an audit row.
@@ -48,6 +56,7 @@ pub fn sanitize_usage_request_metadata_object(source: &Map<String, Value>) -> Op
PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY,
"transport_error",
"is_free_tier",
CANCELLED_REQUEST_FEE_METADATA_KEY,
USAGE_AVAILABLE_METADATA_KEY,
USAGE_PRICING_AVAILABLE_METADATA_KEY,
] {
+60 -3
View File
@@ -1,5 +1,5 @@
use std::io;
use std::net::SocketAddr;
use std::net::{IpAddr, Ipv6Addr, SocketAddr};
use std::time::Duration;
/// Maximum number of addresses accepted from one hostname lookup.
@@ -13,6 +13,16 @@ pub const MAX_DNS_RESOLVED_ADDRESSES: usize = 32;
/// Upper bound used by callers that do not have a tighter request deadline.
pub const DEFAULT_DNS_LOOKUP_TIMEOUT: Duration = Duration::from_secs(10);
pub fn parse_ip_literal_host(host: &str) -> Option<IpAddr> {
host.parse().ok().or_else(|| {
host.strip_prefix('[')?
.strip_suffix(']')?
.parse::<Ipv6Addr>()
.ok()
.map(IpAddr::V6)
})
}
/// Resolve a host while bounding both resolver wait time and answer count.
///
/// The iterator is consumed one item past the allowed count so an answer set
@@ -30,6 +40,9 @@ pub async fn lookup_host_with_limits(
"DNS lookup timeout must be non-zero",
));
}
if let Some(ip) = parse_ip_literal_host(host) {
return Ok(vec![SocketAddr::new(ip, port)]);
}
let mut resolved = tokio::time::timeout(timeout, tokio::net::lookup_host((host, port)))
.await
@@ -59,10 +72,32 @@ mod tests {
use std::net::SocketAddr;
use super::{
collect_resolved_addresses_with_limit, lookup_host_with_limits, DEFAULT_DNS_LOOKUP_TIMEOUT,
MAX_DNS_RESOLVED_ADDRESSES,
collect_resolved_addresses_with_limit, lookup_host_with_limits, parse_ip_literal_host,
DEFAULT_DNS_LOOKUP_TIMEOUT, MAX_DNS_RESOLVED_ADDRESSES,
};
#[test]
fn ip_literal_parser_does_not_turn_bracketed_names_into_hosts() {
for host in [
"localhost",
"[localhost]",
"[127.0.0.1]",
"[::1",
"::1]",
"[[::1]]",
] {
assert_eq!(parse_ip_literal_host(host), None, "{host}");
}
for (host, expected) in [
("198.18.78.41", "198.18.78.41"),
("::1", "::1"),
("[::1]", "::1"),
("[::ffff:127.0.0.1]", "::ffff:127.0.0.1"),
] {
assert_eq!(parse_ip_literal_host(host), Some(expected.parse().unwrap()));
}
}
#[tokio::test]
async fn rejects_zero_dns_timeout_before_resolving() {
let error = lookup_host_with_limits("localhost", 80, std::time::Duration::ZERO)
@@ -71,6 +106,16 @@ mod tests {
assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
}
#[tokio::test]
async fn resolves_bracketed_ipv6_without_a_hostname_lookup() {
for host in ["[::1]", "[2606:4700:4700::1111]", "[fd00::1]"] {
let addresses = lookup_host_with_limits(host, 8443, DEFAULT_DNS_LOOKUP_TIMEOUT)
.await
.expect("URL-form IPv6 literals must not be sent to DNS");
assert_eq!(addresses, vec![format!("{host}:8443").parse().unwrap()]);
}
}
#[tokio::test]
async fn resolves_within_shared_address_bound() {
let addresses = lookup_host_with_limits("localhost", 80, DEFAULT_DNS_LOOKUP_TIMEOUT)
@@ -80,6 +125,18 @@ mod tests {
assert!(addresses.len() <= MAX_DNS_RESOLVED_ADDRESSES);
}
#[test]
fn preserves_every_answer_at_the_shared_bound() {
let expected = (0..MAX_DNS_RESOLVED_ADDRESSES)
.map(|index| SocketAddr::from(([198, 18, 0, index as u8], 443)))
.collect::<Vec<_>>();
let mut resolved = expected.clone().into_iter();
assert_eq!(
collect_resolved_addresses_with_limit(&mut resolved).unwrap(),
expected
);
}
#[test]
fn rejects_an_answer_set_larger_than_the_shared_bound() {
let mut resolved = (0..=MAX_DNS_RESOLVED_ADDRESSES)
+4 -1
View File
@@ -7,7 +7,10 @@ mod retry;
pub use client::{apply_http_client_config, build_http_client, build_http_client_with_headers};
pub use config::{HttpClientConfig, HttpRetryConfig};
pub use dns::{lookup_host_with_limits, DEFAULT_DNS_LOOKUP_TIMEOUT, MAX_DNS_RESOLVED_ADDRESSES};
pub use dns::{
lookup_host_with_limits, parse_ip_literal_host, DEFAULT_DNS_LOOKUP_TIMEOUT,
MAX_DNS_RESOLVED_ADDRESSES,
};
pub use header_security::{
connection_declared_header_names, is_https_or_loopback_http_url, is_ipv4_benchmarking_fake_ip,
is_private_or_reserved_ip, url_has_literal_loopback_host,
+119
View File
@@ -0,0 +1,119 @@
use std::collections::BTreeSet;
use regex::Regex;
use serde::{Deserialize, Serialize};
pub const MAX_ROUTING_FAILOVER_RULES: usize = 64;
pub const MAX_ROUTING_FAILOVER_PATTERN_BYTES: usize = 4096;
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default)]
pub struct RoutingFailoverRule {
pub pattern: String,
pub status_codes: BTreeSet<u16>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default)]
pub struct RoutingFailoverRules {
pub success_failover_patterns: Vec<RoutingFailoverRule>,
pub error_stop_patterns: Vec<RoutingFailoverRule>,
}
pub fn validate_routing_failover_rules(rules: &RoutingFailoverRules) -> Result<(), String> {
for (name, entries, success) in [
(
"success_failover_patterns",
&rules.success_failover_patterns,
true,
),
("error_stop_patterns", &rules.error_stop_patterns, false),
] {
if entries.len() > MAX_ROUTING_FAILOVER_RULES {
return Err(format!("{name} exceeds {MAX_ROUTING_FAILOVER_RULES} rules"));
}
for (index, rule) in entries.iter().enumerate() {
let pattern = rule.pattern.trim();
if pattern.is_empty() && (success || rule.status_codes.is_empty()) {
return Err(format!(
"{name}[{index}] requires a pattern or error status codes"
));
}
if pattern.len() > MAX_ROUTING_FAILOVER_PATTERN_BYTES {
return Err(format!(
"{name}[{index}] pattern exceeds {MAX_ROUTING_FAILOVER_PATTERN_BYTES} bytes"
));
}
if !pattern.is_empty() {
Regex::new(pattern)
.map_err(|error| format!("{name}[{index}] invalid regex: {error}"))?;
}
if rule.status_codes.iter().any(|status| {
if success {
*status != 200
} else {
!(400..=599).contains(status)
}
}) {
return Err(format!("{name}[{index}] contains invalid status codes"));
}
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validates_regex_and_status_only_stop_rules() {
let rules = RoutingFailoverRules {
success_failover_patterns: vec![RoutingFailoverRule {
pattern: "(?i)capacity.*exhausted".to_string(),
..Default::default()
}],
error_stop_patterns: vec![RoutingFailoverRule {
status_codes: [400, 413].into_iter().collect(),
..Default::default()
}],
};
assert!(validate_routing_failover_rules(&rules).is_ok());
}
#[test]
fn rejects_invalid_or_unbounded_rule_configuration() {
for rule in [
RoutingFailoverRule::default(),
RoutingFailoverRule {
pattern: "[".to_string(),
..Default::default()
},
RoutingFailoverRule {
pattern: "error".to_string(),
status_codes: [429].into_iter().collect(),
},
RoutingFailoverRule {
pattern: "a".repeat(MAX_ROUTING_FAILOVER_PATTERN_BYTES + 1),
..Default::default()
},
] {
let rules = RoutingFailoverRules {
success_failover_patterns: vec![rule],
..Default::default()
};
assert!(validate_routing_failover_rules(&rules).is_err());
}
let rules = RoutingFailoverRules {
error_stop_patterns: vec![
RoutingFailoverRule {
status_codes: [400].into_iter().collect(),
..Default::default()
};
MAX_ROUTING_FAILOVER_RULES + 1
],
..Default::default()
};
assert!(validate_routing_failover_rules(&rules).is_err());
}
}
+5
View File
@@ -1,5 +1,6 @@
mod actions;
mod conditions;
mod failover;
mod model;
mod mutations;
mod policy;
@@ -12,6 +13,10 @@ pub use actions::{
RoutingSchedulingMode, RoutingSetPriorityMode,
};
pub use conditions::{RoutingCondition, RoutingConditionContext, RoutingConditionOp};
pub use failover::{
validate_routing_failover_rules, RoutingFailoverRule, RoutingFailoverRules,
MAX_ROUTING_FAILOVER_PATTERN_BYTES, MAX_ROUTING_FAILOVER_RULES,
};
pub use model::{
RoutingDefaultPolicy, RoutingExecutionPolicy, RoutingGroupBinding, RoutingGroupBindingSubject,
RoutingGroupConfig, RoutingGroupRecord, RoutingGroupVersionRecord, RoutingModelPolicy,
+75 -1
View File
@@ -7,6 +7,7 @@ use crate::actions::{
RoutingAction, RoutingRulePhase, RoutingSchedulingMode, RoutingSetPriorityMode,
};
use crate::conditions::RoutingCondition;
use crate::RoutingFailoverRules;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RoutingSchedulingPreset {
@@ -33,12 +34,20 @@ pub const DEFAULT_STICKY_KEY_ATTEMPTS: u32 = 2;
/// transport configuration. A resolved policy is snapshotted for the request
/// and can therefore be consumed by execution without rereading mutable
/// system settings.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Default)]
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Default)]
pub struct RoutingExecutionPolicy {
#[serde(default, skip_serializing_if = "is_false")]
pub enable_cf_heartbeat: bool,
#[serde(default, skip_serializing_if = "is_false")]
pub cyber_continue_failover: bool,
#[serde(default, skip_serializing_if = "is_false")]
pub cancel_on_client_disconnect: bool,
#[serde(default)]
pub max_transfer_count: u64,
#[serde(default)]
pub max_transfer_timeout_seconds: u64,
#[serde(default)]
pub failover_rules: RoutingFailoverRules,
}
impl<'de> Deserialize<'de> for RoutingExecutionPolicy {
@@ -56,6 +65,14 @@ impl<'de> Deserialize<'de> for RoutingExecutionPolicy {
enable_standard_text_sync_heartbeat: bool,
#[serde(default)]
cyber_continue_failover: bool,
#[serde(default)]
cancel_on_client_disconnect: bool,
#[serde(default)]
max_transfer_count: u64,
#[serde(default)]
max_transfer_timeout_seconds: u64,
#[serde(default)]
failover_rules: RoutingFailoverRules,
}
let value = LegacyCompatibleExecutionPolicy::deserialize(deserializer)?;
@@ -64,6 +81,10 @@ impl<'de> Deserialize<'de> for RoutingExecutionPolicy {
|| value.enable_openai_image_sync_heartbeat
|| value.enable_standard_text_sync_heartbeat,
cyber_continue_failover: value.cyber_continue_failover,
cancel_on_client_disconnect: value.cancel_on_client_disconnect,
max_transfer_count: value.max_transfer_count,
max_transfer_timeout_seconds: value.max_transfer_timeout_seconds,
failover_rules: value.failover_rules,
})
}
}
@@ -107,6 +128,59 @@ fn is_false(value: &bool) -> bool {
!*value
}
#[cfg(test)]
mod execution_policy_tests {
use super::*;
#[test]
fn routing_failover_configuration_round_trips_and_validates() {
let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({
"default_policy": {
"max_transfer_count": 3,
"max_transfer_timeout_seconds": 90,
"failover_rules": {
"success_failover_patterns": [{ "pattern": "(?i)capacity" }],
"error_stop_patterns": [{ "status_codes": [400, 413] }]
}
}
}))
.unwrap();
crate::validate_routing_group_config(&config).unwrap();
assert_eq!(config.default_policy.execution_policy.max_transfer_count, 3);
let value = serde_json::to_value(&config).unwrap();
assert_eq!(value["default_policy"]["max_transfer_timeout_seconds"], 90);
assert_eq!(
serde_json::from_value::<RoutingGroupConfig>(value).unwrap(),
config
);
}
#[test]
fn cancellation_defaults_off_and_round_trips_with_legacy_heartbeat() {
let default: RoutingDefaultPolicy = serde_json::from_str("{}").unwrap();
assert!(!default.execution_policy.cancel_on_client_disconnect);
let policy: RoutingDefaultPolicy = serde_json::from_value(serde_json::json!({
"cancel_on_client_disconnect": true,
"enable_standard_text_sync_heartbeat": true
}))
.unwrap();
assert!(policy.execution_policy.cancel_on_client_disconnect);
assert!(policy.execution_policy.enable_cf_heartbeat);
let encoded = serde_json::to_value(&policy).unwrap();
assert_eq!(encoded["cancel_on_client_disconnect"], true);
assert_eq!(
serde_json::from_value::<RoutingDefaultPolicy>(encoded).unwrap(),
policy
);
assert!(
serde_json::from_value::<RoutingDefaultPolicy>(serde_json::json!({
"cancel_on_client_disconnect": "true"
}))
.is_err()
);
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct RoutingModelPolicy {
pub model: String,
+1 -1
View File
@@ -89,7 +89,7 @@ pub fn resolve_routing_policy(
scheduling_mode: config.default_policy.scheduling_mode,
keep_priority_on_conversion: config.default_policy.keep_priority_on_conversion,
sticky_key_attempts: config.default_policy.sticky_key_attempts,
execution_policy: config.default_policy.execution_policy,
execution_policy: config.default_policy.execution_policy.clone(),
ranking_overlay: RankingOverlay::default(),
mutation_plan: MutationPlan::default(),
pool_policy_overrides: BTreeMap::new(),
@@ -28,6 +28,8 @@ const ROUTING_POOL_PRESETS: &[&str] = &[
#[derive(Debug, Error, Clone, PartialEq, Eq)]
pub enum RoutingValidationError {
#[error("routing failover rules are invalid: {0}")]
InvalidFailoverRules(String),
#[error("routing rule id is empty")]
EmptyRuleId,
#[error("duplicate routing rule id: {0}")]
@@ -69,6 +71,8 @@ pub enum RoutingValidationError {
pub fn validate_routing_group_config(
config: &RoutingGroupConfig,
) -> Result<(), RoutingValidationError> {
crate::validate_routing_failover_rules(&config.default_policy.execution_policy.failover_rules)
.map_err(RoutingValidationError::InvalidFailoverRules)?;
let mut rule_ids = BTreeSet::new();
for model_policy in &config.model_policies {
if model_policy.model.trim().is_empty() {
@@ -1310,16 +1310,18 @@ mod tests {
#[test]
fn profile_values_depend_only_on_seed_sequence_and_domain() {
let mut config = Config::default();
config.seed = 0xfeed_beef;
config.first_byte_delay = Duration::from_millis(100);
config.first_byte_jitter = Duration::from_millis(50);
config.chunk_delay = Duration::from_millis(30);
config.chunk_delay_jitter = Duration::from_millis(10);
config.payload_bytes = 128;
config.payload_bytes_jitter = 64;
config.fault_truncate_stream_bps = BASIS_POINTS;
config.chunks = 9;
let config = Config {
seed: 0xfeed_beef,
first_byte_delay: Duration::from_millis(100),
first_byte_jitter: Duration::from_millis(50),
chunk_delay: Duration::from_millis(30),
chunk_delay_jitter: Duration::from_millis(10),
payload_bytes: 128,
payload_bytes_jitter: 64,
fault_truncate_stream_bps: BASIS_POINTS,
chunks: 9,
..Default::default()
};
let first = request_profile(&config, 1234);
let unrelated = request_profile(&config, 9999);
@@ -1345,11 +1347,13 @@ mod tests {
#[test]
fn fault_buckets_are_disjoint_and_ordered() {
let mut config = Config::default();
config.fault_429_bps = 100;
config.fault_500_bps = 200;
config.fault_timeout_bps = 300;
config.fault_truncate_stream_bps = 400;
let config = Config {
fault_429_bps: 100,
fault_500_bps: 200,
fault_timeout_bps: 300,
fault_truncate_stream_bps: 400,
..Default::default()
};
assert_eq!(select_fault(&config, 0), Fault::Status429);
assert_eq!(select_fault(&config, 99), Fault::Status429);
@@ -1435,12 +1439,14 @@ mod tests {
let bind = listener
.local_addr()
.expect("test listener should have a local address");
let mut config = Config::default();
config.binds = vec![bind];
config.chunks = 0;
config.chunk_delay = Duration::ZERO;
config.fault_truncate_stream_bps = BASIS_POINTS;
config.seed = 0x1234_5678;
let config = Config {
binds: vec![bind],
chunks: 0,
chunk_delay: Duration::ZERO,
fault_truncate_stream_bps: BASIS_POINTS,
seed: 0x1234_5678,
..Default::default()
};
let metrics = Arc::new(Metrics::for_binds(&config.binds));
let router = build_router(config, Arc::clone(&metrics), bind);
let server = tokio::spawn(async move {
@@ -25,6 +25,9 @@ use aether_data_contracts::repository::global_models::{
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::routing_profiles::{
RoutingGroupLookupKey, UpdateRoutingGroupRecord,
};
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageAuditListQuery};
use aether_gateway::{build_router_with_state, AppState, GatewayDataConfig, UsageRuntimeConfig};
use aether_testkit::{ManagedPostgresServer, SpawnedServer};
@@ -491,8 +494,7 @@ async fn disabling_the_downstream_key_is_enforced_on_the_next_turn_of_the_same_s
Ok(())
}
/// A client that walks away before the provider produced anything must settle
/// as a void row: nothing was produced, so nothing is billed.
/// Immediate cancellation voids token billing before the provider completes.
///
/// This is the path with no protocol event to announce it: the relay loop owns
/// the turn, and losing the client is an exit the upstream never reports.
@@ -506,7 +508,14 @@ async fn disabling_the_downstream_key_is_enforced_on_the_next_turn_of_the_same_s
/// `a_closed_client_socket_before_any_terminal_still_voids_the_bill`.
#[tokio::test]
async fn client_disconnect_before_any_provider_output_settles_a_void_row() -> Result<(), BoxError> {
let harness = Harness::start(UpstreamBehavior::StallAfterCreated).await?;
let harness = Harness::start_configured(
UpstreamBehavior::StallAfterCreated,
ProviderFixture::SingleOpenAiKey,
PiiRedaction::Disabled,
true,
None,
)
.await?;
let mut client = harness.connect().await?;
client
@@ -542,6 +551,73 @@ async fn client_disconnect_before_any_provider_output_settles_a_void_row() -> Re
Ok(())
}
#[tokio::test]
async fn client_disconnect_defaults_to_completing_and_billing_the_turn() -> Result<(), BoxError> {
let harness = Harness::start(UpstreamBehavior::CompleteAfterRelease).await?;
let mut client = harness.connect().await?;
client
.send(response_create(json!({"input": "finish without client"})))
.await?;
receive_event(&mut client, "response.created").await?;
drop(client);
tokio::time::sleep(Duration::from_millis(50)).await;
harness.upstream.release_completion.notify_one();
let audits = harness
.usage_audits_where(1, "completed disconnected turn", |audit| {
audit.status == "completed" && audit.billing_status == "settled"
})
.await?;
assert_eq!(audits.len(), 1);
assert_eq!(audits[0].input_tokens, INPUT_TOKENS);
assert_eq!(audits[0].output_tokens, OUTPUT_TOKENS);
assert_eq!(audits[0].status_code, Some(200));
Ok(())
}
#[tokio::test]
async fn client_disconnect_still_settles_the_per_request_fee_when_aborted() -> Result<(), BoxError>
{
let harness = Harness::start_configured(
UpstreamBehavior::StallAfterCreated,
ProviderFixture::SingleOpenAiKey,
PiiRedaction::Disabled,
true,
Some(0.02),
)
.await?;
let mut client = harness.connect().await?;
client
.send(response_create(json!({"input": "cancel with request fee"})))
.await?;
receive_event(&mut client, "response.created").await?;
drop(client);
let audits = harness
.usage_audits_where(1, "cancelled request fee settlement", |audit| {
audit.status == "cancelled" && audit.billing_status == "settled"
})
.await?;
assert_eq!(audits.len(), 1);
assert_eq!(audits[0].status_code, Some(499));
assert_eq!(audits[0].total_tokens, 0);
assert_eq!(audits[0].total_cost_usd, 0.02);
assert_eq!(audits[0].actual_total_cost_usd, 0.02);
let backends = DataBackends::from_config(DataLayerConfig::from_database(
harness.database.config.clone(),
))?;
let detail = backends
.read()
.usage()
.ok_or("usage reader unavailable")?
.find_by_request_id(&audits[0].request_id)
.await?
.ok_or("usage detail unavailable")?;
assert_eq!(
detail.request_metadata.as_ref().unwrap()["cancelled_request_fee"],
true
);
Ok(())
}
/// An upstream that dies mid-turn must surface an error and still settle.
#[tokio::test]
async fn upstream_drop_mid_turn_reports_an_error_and_settles_the_usage_row() -> Result<(), BoxError>
@@ -850,6 +926,16 @@ impl Harness {
behavior: UpstreamBehavior,
fixture: ProviderFixture,
redaction: PiiRedaction,
) -> Result<Self, BoxError> {
Self::start_configured(behavior, fixture, redaction, false, None).await
}
async fn start_configured(
behavior: UpstreamBehavior,
fixture: ProviderFixture,
redaction: PiiRedaction,
cancel_on_client_disconnect: bool,
request_price: Option<f64>,
) -> Result<Self, BoxError> {
let upstream = Arc::new(MockUpstreamState::new(behavior));
let upstream_server =
@@ -864,6 +950,16 @@ impl Harness {
)
.await?;
if let Some(price) = request_price {
let pool = sqlx::PgPool::connect(&database.config.url).await?;
sqlx::query("UPDATE models SET price_per_request = $1 WHERE id = $2")
.bind(price)
.bind(PROVIDER_MODEL_ID)
.execute(&pool)
.await?;
pool.close().await;
}
let data_config = GatewayDataConfig::from_database_config(database.config.clone())
.with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY);
let state = AppState::new()?
@@ -877,6 +973,32 @@ impl Harness {
..UsageRuntimeConfig::default()
})?;
state.ensure_system_default_routing_group().await?;
if cancel_on_client_disconnect {
let backends =
DataBackends::from_config(DataLayerConfig::from_database(database.config.clone()))?;
let mut group = backends
.read()
.routing_groups()
.ok_or("routing reader unavailable")?
.find_routing_group(RoutingGroupLookupKey::SystemDefault)
.await?
.ok_or("default routing group unavailable")?;
group.config_json["default_policy"]["cancel_on_client_disconnect"] = json!(true);
backends
.write()
.routing_groups()
.ok_or("routing writer unavailable")?
.update_routing_group(
&group.id,
UpdateRoutingGroupRecord {
config_json: Some(group.config_json),
version: Some(group.version + 1),
updated_at: group.updated_at + 1,
..Default::default()
},
)
.await?;
}
let gateway_server = SpawnedServer::start(build_router_with_state(state)).await?;
let websocket_url = format!(
"{}/v1/responses",
@@ -1093,6 +1215,7 @@ where
enum UpstreamBehavior {
/// Announce, stream one delta, and complete — the ordinary turn.
CompleteEveryTurn,
CompleteAfterRelease,
/// Announce the response and then go quiet, leaving the turn in flight.
StallAfterCreated,
/// Announce the response and then hang up mid-turn.
@@ -1118,6 +1241,7 @@ struct MockUpstreamState {
events: Mutex<Vec<Value>>,
authorization_headers: Mutex<Vec<Option<String>>>,
handshakes: Mutex<Vec<ObservedUpstreamHandshake>>,
release_completion: tokio::sync::Notify,
}
#[derive(Debug, Clone)]
@@ -1134,6 +1258,7 @@ impl MockUpstreamState {
events: Mutex::new(Vec::new()),
authorization_headers: Mutex::new(Vec::new()),
handshakes: Mutex::new(Vec::new()),
release_completion: tokio::sync::Notify::new(),
}
}
@@ -1215,6 +1340,15 @@ async fn run_mock_upstream(
break;
}
}
UpstreamBehavior::CompleteAfterRelease => {
if send_mock_created(&mut socket, &response_id).await.is_err() {
break;
}
state.release_completion.notified().await;
if send_mock_turn(&mut socket, &response_id).await.is_err() {
break;
}
}
UpstreamBehavior::StallAfterCreated => {
if send_mock_created(&mut socket, &response_id).await.is_err() {
break;
@@ -1265,10 +1399,13 @@ async fn run_mock_upstream(
}
}
}
AxumWsMessage::Ping(payload) => {
if socket.send(AxumWsMessage::Pong(payload)).await.is_err() {
break;
}
AxumWsMessage::Ping(payload)
if socket
.send(AxumWsMessage::Pong(payload.clone()))
.await
.is_err() =>
{
break;
}
AxumWsMessage::Close(_) => break,
_ => {}
+32
View File
@@ -221,6 +221,13 @@ fn lifecycle_status_and_billing(
}
UsageEventType::Completed => ("completed", "pending"),
UsageEventType::Failed => ("failed", "void"),
UsageEventType::Cancelled
if aether_data_contracts::repository::usage::cancelled_request_fee_is_billable(
request_metadata,
) =>
{
("cancelled", "pending")
}
UsageEventType::Cancelled => ("cancelled", "void"),
}
}
@@ -491,6 +498,31 @@ mod tests {
assert_eq!(record.first_byte_time_ms, Some(50));
}
#[test]
fn cancelled_request_fee_keeps_cancelled_status_and_pending_billing() {
let event = UsageEvent::new(
UsageEventType::Cancelled,
"req-cancelled-fee",
UsageEventData {
provider_name: "OpenAI".to_string(),
model: "gpt-5".to_string(),
total_cost_usd: Some(0.02),
actual_total_cost_usd: Some(0.01),
request_metadata: Some(serde_json::json!({"cancelled_request_fee": true})),
..Default::default()
},
);
let record = build_upsert_usage_record_from_event(&event).unwrap();
assert_eq!(record.status, "cancelled");
assert_eq!(record.billing_status, "pending");
assert_eq!(record.total_cost_usd, Some(0.02));
assert_eq!(record.actual_total_cost_usd, Some(0.01));
assert_eq!(
record.request_metadata.unwrap()["cancelled_request_fee"],
true
);
}
#[test]
fn completed_unmetered_session_audit_is_void_without_fabricated_usage() {
let record = build_upsert_usage_record_from_event(&UsageEvent {
+71 -3
View File
@@ -5,8 +5,10 @@ use aether_data_contracts::repository::settlement::{
ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement,
UsagePolicyCostReservationState, UsageSettlementInput,
};
use aether_data_contracts::repository::usage::StoredRequestUsageAudit;
use aether_data_contracts::repository::usage::PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY;
use aether_data_contracts::repository::usage::{
cancelled_request_fee_is_billable, StoredRequestUsageAudit,
};
use aether_data_contracts::{DataLayerError, DataLayerError::InvalidInput};
use async_trait::async_trait;
@@ -39,6 +41,11 @@ pub async fn reconcile_usage_policy_cost_for_event(
}
let terminal_state = match event.event_type {
UsageEventType::Completed => UsagePolicyCostReservationState::Finalized,
UsageEventType::Cancelled
if cancelled_request_fee_is_billable(event.data.request_metadata.as_ref()) =>
{
UsagePolicyCostReservationState::Finalized
}
UsageEventType::Failed | UsageEventType::Cancelled => {
UsagePolicyCostReservationState::Released
}
@@ -107,7 +114,10 @@ pub async fn settle_usage_if_needed(
usage.user_id.as_deref().and_then(non_empty_trimmed),
usage_policy_reservation_token(usage),
) {
let (terminal_state, actual_cost_units) = if usage.status == "completed" {
let (terminal_state, actual_cost_units) = if usage.status == "completed"
|| (usage.status == "cancelled"
&& cancelled_request_fee_is_billable(usage.request_metadata.as_ref()))
{
(
UsagePolicyCostReservationState::Finalized,
nonnegative_usd_to_usage_policy_cost_units(
@@ -136,7 +146,10 @@ pub async fn settle_usage_if_needed(
}
}
if usage.status == "cancelled" || usage.billing_status != "pending" {
if usage.billing_status != "pending"
|| (usage.status == "cancelled"
&& !cancelled_request_fee_is_billable(usage.request_metadata.as_ref()))
{
return Ok(());
}
let input = UsageSettlementInput {
@@ -429,6 +442,61 @@ mod tests {
);
}
#[tokio::test]
async fn cancelled_request_fee_settles_wallet_and_finalizes_cost_reservation() {
let writer = TestSettlementWriter {
has_writer: true,
..Default::default()
};
let mut usage = sample_usage();
usage.status = "cancelled".to_string();
usage.status_code = Some(499);
usage.request_metadata.as_mut().unwrap()["cancelled_request_fee"] = json!(true);
settle_usage_if_needed(&writer, &usage).await.unwrap();
let inputs = writer.inputs.lock().unwrap();
assert_eq!(inputs.len(), 1);
assert_eq!(inputs[0].status, "cancelled");
assert_eq!(inputs[0].actual_total_cost_usd, usage.actual_total_cost_usd);
let reconciliations = writer.reconciliations.lock().unwrap();
assert_eq!(reconciliations.len(), 1);
assert_eq!(
reconciliations[0].terminal_state,
UsagePolicyCostReservationState::Finalized
);
assert_eq!(reconciliations[0].actual_cost_units, 75_000_000);
}
#[tokio::test]
async fn cancelled_request_fee_event_finalizes_cost_reservation() {
let writer = TestSettlementWriter {
has_writer: true,
..Default::default()
};
let event = UsageEvent::new(
UsageEventType::Cancelled,
"req-cancelled-fee",
UsageEventData {
user_id: Some("user-1".to_string()),
actual_total_cost_usd: Some(0.01),
request_metadata: Some(json!({
"cancelled_request_fee": true,
"plan_usage_reservation_token": "server-token"
})),
..Default::default()
},
);
reconcile_usage_policy_cost_for_event(&writer, &event)
.await
.unwrap();
let reconciliations = writer.reconciliations.lock().unwrap();
assert_eq!(reconciliations.len(), 1);
assert_eq!(
reconciliations[0].terminal_state,
UsagePolicyCostReservationState::Finalized
);
assert_eq!(reconciliations[0].actual_cost_units, 1_000_000);
}
#[tokio::test]
async fn releases_failed_usage_before_void_settlement() {
let writer = TestSettlementWriter {
@@ -0,0 +1,39 @@
# 调度策略:取消请求立即打断
管理入口:**调度策略配置 → 系统配置 → 取消请求立即打断**。
配置保存在当前策略的 `config_json.default_policy.cancel_on_client_disconnect`,默认 `false`。
旧策略缺失此字段也按关闭处理;不需要数据库结构迁移,不读取同名全局系统设置。
策略选定后,请求沿用该次解析的配置,不因管理员随后修改策略而改变断连处理。
| 配置 | 客户端取消/断连时 | 计费 |
| --- | --- | --- |
| 关闭(默认) | 已选定策略的请求继续执行,后台读取响应直到正常完成或原有超时/上游错误 | 按实际完成结果正常结算 |
| 开启 | 中止仍在执行的请求、停止读取上游响应 | Token、缓存 Token 和图片产出费用不收取;配置了 `price_per_request` 的请求保留一次请求费用及对应倍率 |
已经取得上游终态的请求不会因最后一跳投递失败而撤销已完成的结算。
取消按次收费的记录仍显示 `cancelled`/499,但计费状态可为 `settled`;不要仅凭请求状态判断是否收费。
## 执行边界
- HTTP 同步、SSE、同格式直通和 CF 心跳响应共用请求生命周期保护。
- Responses WebSocket 断连后只完成当前进行中的 turn,不无限维持空闲连接;开启开关则直接终止当前 turn。
- 鉴权、请求体接收等尚未选定策略的阶段仍可直接取消。
- Live/Realtime 长连接会话关闭、异步任务显式取消不是有限 HTTP/Responses 请求的断连续跑,不转成后台常驻会话。
- 原有上游总超时、首字节超时及故障处理仍有效;这里不创建持久化后台作业,进程退出不能保证继续执行。
## 代码调整
- `request_lifecycle.rs` 将 HTTP 请求 Future 和响应 Body 的所有权与客户端连接解耦。正常连接保留原来的逐帧路径、响应头、长度提示和 trailers;仅断连时启动后台接管,逐帧丢弃待发送内容,不聚合完整响应。
- 请求准入凭证随原 Future/Body 保留至完成,防止断连提前释放并发额度;诊断上下文同时保留。
- CF 心跳后台执行显式监听响应接收端关闭,避免打开开关后仍继续执行。
- Responses WebSocket 将“客户端已断开”和“上游已完成”分别处理,复用既有终态观察、超时和结算逻辑。
- 计费统一在 billing enrichment 中计算取消请求的单次费用。`cancelled_request_fee` 由服务端计算后标记,贯穿审计持久化、钱包结算及套餐成本预留结算,避免有费用却仍写为 `void` 或释放成本预留。
## 回归覆盖
- 新旧配置默认关闭、布尔校验、策略保存及模型级调度编辑保留开关。
- 响应头前断连、首帧前/首帧后断连、立即取消、并发额度、诊断信息、响应头/trailers 透传。
- 真实流式执行链路的完成/取消 usage 与候选状态,以及心跳取消。
- 取消时按次/按 Token/混合定价、图片请求单次费、倍率、审计状态、钱包及成本预留结算。
- Responses WebSocket 使用真实网关和临时 PostgreSQL 验证断连继续完成、开启后不收费,以及按次费用实际结算。
@@ -151,6 +151,32 @@ ws.send(json.dumps({
- Direct provider proxy settings are honored through the selected transport
profile. Tunnel-mode proxy nodes are not supported for this bridge yet.
### DNS and proxy behavior
The gateway's plain and browser-profile WebSocket clients share the HTTP/SSE
provider DNS resolver. Provider hostname answers are not filtered by address
range, including Fake-IP answers such as `198.18.0.0/15`. DNS lookup timeout,
answer-count limits, and rejection of empty answers still apply. Resolution
happens during connection establishment, not while building the client; DNS
failures therefore surface as upstream handshake failures rather than invalid
upstream URLs.
This policy applies only to configured provider hostnames. URL validation still
rejects credentials, fragments, and literal private/reserved IP targets (except
loopback `ws://`), and the gateway frontdoor self-loop guard remains active.
Tunnel owner-relay DNS address filtering is unchanged.
Configure an explicit HTTP(S) or SOCKS proxy on the provider when needed; these
clients do not automatically use system proxy environment variables. A proxy
connection does not trigger a separate gateway-side lookup of the provider
hostname. Existing SOCKS remote-DNS normalization remains in effect.
The [DNS egress audit](dns-egress-audit-2026-09-08.md) documents the separate
tunnel and untrusted-download policies, regression coverage, and remaining
deployment limitations; not every outbound path uses the provider DNS policy.
### Usage and logging
Usage and audit finalization now runs for every accepted `response.create`.
Existing usage body-capture and header-redaction policies apply to the resulting
records. Newly created WebSocket usage records expose `is_websocket=true`, and
@@ -0,0 +1,42 @@
# 格式转换失败诊断导出
## 使用方法
在请求详情的失败或跳过节点中,点击「失败诊断」面板的复制按钮。
复制时才会读取已采集的正文;页面预览本身不会批量加载正文。
请分享整个 JSON,而不是只分享 `summary`。复制成功标志仅在剪贴板写入成功后出现。
## Schema v2
- `diagnostic`:错误码、请求/响应/流式阶段、源/目标格式、转换器标识、完整 JSON 路径及期望约束/实际值。
- `path_source`:`structured` 是后端结构化路径,`message_inference` 是历史文案推断,`protocol_inference` 是根据上游协议推断的原始字段路径,`unavailable` 表示没有可靠字段路径。通用 `$.finish_reason` 会按协议定位到具体原始字段,同时保留 `reported_path`。
- `stage_source`:区分后端阶段信息与历史记录推断。请求转换方向为客户端到上游;响应/流式转换方向相反。
- `versions`:前端版本、导出时网关版本、失败时运行版本。历史记录没有运行版本时保留 `null`,不能把导出版本当作失败版本。
- `request` / `node`:请求、候选、重试、模型及时间等定位信息。
- `reproduction.sources`:脱敏正文片段、字段样本和流式失败事件窗口。数组通配路径的样本带具体下标。
- `reproduction.missing_context`:未采集、无权限、正文过大、读取失败、缺少失败事件或路径等缺口。
后端只为明确匹配 `candidate_id` 的候选提供上游正文记录。历史记录仅有候选索引时,不把最后一次重试的正文猜成当前失败的正文。
原始客户端请求可在同一请求内共享,但不会把其他候选的上游请求/响应当作失败现场。
`body_ref` 不是下载地址;前端不访问其中的 URL,而是使用现有、受权限保护的正文接口。
## 完整性与安全边界
- `not_loaded`:尚未补取上下文。
- `sanitized_context`:必要来源已取得,但仍然经过脱敏、大小限制或事件窗口裁剪。
- `insufficient_context`:还缺少明确列出的信息,不能据此假设能够完整复现。
- `replay_ready: false`:导出的是供排查的证据包,不是可以无条件自动执行的请求。修改前应根据样本建立最小回归测试。
每份正文的下载/解码处理上限为 1 MiB,读取超时为 5 秒;导出 JSON 上限为 64 Ki 字符。
字符串、数组、对象深度与节点数量也有限制。流式窗口保留匹配失败的事件、帧序号及邻近事件;匹配不到时明确标记,而不是认定流尾就是故障点。
默认移除常见认证头、密钥、令牌、Cookie、密码、签名 URL 参数、正文文本和二进制数据。
脱敏是规则化处理,分享前仍需检查自定义字段和错误消息是否含业务敏感信息。
不会为了诊断绕过正文采集策略、授权或存储限制;也不会把 `error` 或未知结束原因映射成正常成功。
## 建议处理流程
1. 检查 `diagnostic.stage`、`path_source` 和 `missing_context`,区分转换器缺陷、合法的无损转换拒绝和上游失败。
2. 对照源/目标格式以及字段样本,建立最小失败输入;流式问题同时保留必要的前序事件。
3. 先补失败回归测试,再修复转换规则。
4. 验证原有正常映射、失败闭合以及凭据脱敏没有回退。
@@ -0,0 +1,171 @@
# DNS 与出站连接审计(2026-09-08)
## 范围与结论边界
本轮基于当前工作区检查 `apps/`、`crates/` 中的 DNS 查询、IP 地址校验、
客户端构建、代理配置以及 TCP/WebSocket 连接入口,沿调用关系区分供应商请求、
身份认证、任意 URL 下载和隧道转发。不是仅搜索 `WebSocket` 或 `chatgpt.com`。
第一轮 WS 修复不足以说明所有路径已经一致。本轮又发现连接测试的旧地址过滤、
URL 形式 IPv6 误入 DNS、两处 DNS 答案静默截断,以及附件下载只保留首个地址。
这些问题,以及后续复核发现的 SMTP 无界解析和隧道缺少显式远程 DNS 模式,
已在工作区修复。本文不表示生产服务器已经部署,也不保证真实上游的
DNS、TLS、TUN 路由或出口代理一定可用。
## 已修复问题
### 1. 普通供应商出口的策略重复
- 普通 HTTP/SSE、浏览器指纹 HTTP、H2C 已使用供应商解析策略;普通 WS 仍有独立过滤,
已在第一轮改为使用 `ExecutionSafeDnsResolver`。
- `/v1/test-connection` 的本地快捷路径仍逐项拒绝私网/保留 DNS 答案,导致同一个
配置好的供应商正式请求能成功、连接测试却失败。本轮删除该重复策略,复用正式
HTTP 客户端的解析器。
- HTTP、WS 和连接测试统一使用 `validate_execution_upstream_url` 校验供应商 URL。
域名的 DNS 答案不按地址段过滤,不等于允许在 URL 中直接填写任意私网 IP。
- URL 中的凭据、fragment、私网/保留 IP 字面地址仍被拒绝;HTTP/WS 字面 loopback
保持正式执行路径已有的兼容策略。禁用重定向、供应商显式代理和 WS 自循环检查保留。
相关文件:
- `apps/aether-gateway/src/execution_runtime/transport.rs`
- `apps/aether-gateway/src/handlers/proxy/websocket/transport.rs`
- `apps/aether-gateway/src/handlers/public/support/test_connection/route.rs`
### 2. IPv6 字面地址被当作域名
`Url::host_str()` 可提供 `[::1]` 形式的主机名,而 `IpAddr::from_str` 和
`lookup_host((host, port))` 的原有调用没有正确消化这个形式。修复前新增测试实际失败,
错误为 `failed to lookup address information: nodename nor servname provided, or not known`。
- 公共解析器现在先识别 IPv4、裸 IPv6 和合法的方括号 IPv6,直接生成 socket 地址。
- 不接受 `[localhost]`、`[127.0.0.1]` 等伪造的方括号主机名。
- relay 的 loopback 判断也使用同一解析函数,防止解析成功后又误判 `[::1]`。
- 隧道 SOCKS 地址编码复用该函数;IPv6 字面地址不再在远程 DNS 模式下被当作域名发送。
- 私网过滤仍由每个调用方的安全策略决定,公共解析器本身不扩大地址权限。
相关文件:`crates/aether-http/src/dns.rs`、
`apps/aether-tunnel/src/egress_proxy.rs`、网关执行传输模块。
### 3. DNS 答案静默截断
Bark 推送和 ChatGPT-Web 图片解析仍直接调用系统 DNS,然后仅取前 32 个答案。
本轮改用共享的 `lookup_host_with_limits`:保留原超时预算,超过 32 个答案直接报错,
不再静默忽略剩余答案。两条路径的私网校验、地址固定及官方来源 Fake-IP 例外不变。
owner gateway 转发的独立实现已取第 33 个答案并拒绝超限,因此不是同类遗漏。
### 4. Grok 附件下载缺少多地址回退
原实现校验全部 DNS 答案后只固定第一个公网地址;首个地址不可连接时,客户端无法
尝试 DNS 返回的其它公网地址。本轮改为把全部经过校验的地址交给客户端,保留双栈和
多地址回退能力。空答案、Fake-IP、私网及公网/私网混合答案仍整体拒绝。
### 5. SMTP DNS 不受连接超时控制
SMTP 发送和连接探测现在都先用共享异步解析器建立 TCP 连接,再把已连接的 blocking
socket 交给原来的 SMTP/TLS 协议实现,不在阻塞任务中重新解析或连接。
- DNS 最长 10 秒;DNS 与 TCP 尝试共用 30 秒总预算。
- 保留全部不超过 32 个答案,不再静默截断为前 16 个;超限直接报错。
- TCP 地址竞争取首个成功连接并释放其它尝试,避免首个黑洞地址吃完整个预算,
使后续可达地址根本没有机会建连。SMTP/TLS 仅在最终选中的连接上运行。
- 解析失败、空答案、解析超时及总连接超时有明确错误,不把底层 DNS 细节暴露给调用方。
- TLS 继续验证原始 SMTP 主机名;读写超时仍为 30 秒,内网邮件服务器策略不变。
相关文件:`apps/aether-gateway/src/email_delivery.rs`。
### 6. 隧道供应商出口可显式委托代理解析
新增默认关闭的 `upstream_proxy_remote_dns`(CLI `--upstream-proxy-remote-dns`,环境变量
`AETHER_TUNNEL_UPSTREAM_PROXY_REMOTE_DNS`,setup 的 `Proxy Remote DNS` 开关)。
- 默认路径仍本地解析、执行 ACL、固定 IP,`socks5h://` 本身不改变既有安全策略。
- 显式启用后,HTTP CONNECT 或 SOCKS5h 接收原始域名,不再预先查询隧道本机 DNS;
Host、SNI 和证书校验仍保留原域名。
- 必须配置 HTTP/SOCKS5h 代理;无代理或 `socks5://` 会在启动和客户端构建时拒绝。
- 仍执行端口白名单、URL 校验以及 IP 字面地址/`localhost` 限制;字面 IP 保持固定。
- 远程 DNS 和固定 IP 使用不同连接池键;远程模式解析器明确拒绝本地 DNS 回退。
- 代理 DNS、TCP、CONNECT/SOCKS 和 TLS 握手共同受上游连接超时限制。
- 启用时打印安全提示:**域名目标的最终 IP ACL 由受信任代理负责**。普通 CONNECT/
SOCKS5 协议无法让隧道校验代理最终连接的 IP,不能宣称远程解析仍保留本地逐 IP 检查。
配置示例及部署边界见 `apps/aether-tunnel/README.md` 的“上游 HTTP 请求”章节。
该文档中 DNS 缓存和连接超时等环境变量误写的 `_SECS` 后缀也已纠正,避免按文档
设置后实际未被程序读取;CLI/TOML 参数名称不变。
## 必须保留的策略差异
| 路径 | DNS / 代理策略 | 本轮处理 |
| --- | --- | --- |
| 普通供应商 HTTP/SSE、浏览器指纹、H2C、WS、连接测试 | 域名答案不按地址段过滤;显式代理优先;供应商客户端不自动使用系统代理环境变量 | 统一遗留分支 |
| 供应商操作类 OAuth、模型获取 | 经执行计划进入供应商运行时;不能与用户登录的身份 OAuth 混为一谈 | 核对调用关系 |
| 身份 OAuth / 管理端 OAuth 探测 | 独立敏感出口;校验目标、固定地址;部分内置官方来源允许窄范围 Fake-IP | 保留,不全局放开 |
| Grok 用户附件、公共视频 URL | 不可信 URL;保留公网限制和固定地址,不能套用供应商域名策略 | Grok 保留全部安全地址 |
| ChatGPT-Web 图片下载与上传 | 普通 URL 严格过滤;可信存储来源有专门 Fake-IP 例外 | 公共有界解析器 |
| 支付出口 | 独立公网校验;固定 Stripe 来源有专门 Fake-IP 例外 | 保留 |
| 系统更新、外部模型目录、Server Chan、Bark | 各自的可信来源例外;自定义目的地不能自动获得同样权限 | Bark 公共有界解析器 |
| gateway owner / internal relay | 独立私网策略、可信 relay 配置和地址固定 | 保留;修复 IPv6 判断 |
| 隧道承载的供应商 HTTP 流量 | 默认本地端口/IP ACL 和固定 IP;显式远程模式委托受信任代理解析及执行域名 IP ACL | 新增默认关闭的远程 DNS 模式 |
| 隧道到 gateway 的控制连接 | 与隧道供应商出口分开;可配置专门出口代理和 IP family | IPv6 / SOCKS 编码复用公共函数 |
| 独立 Responses WS probe | 独立直连诊断程序,不使用供应商代理配置 | 不应当作生产代理路径的等价验证 |
## 仍需注意的实际限制
1. **Fake-IP 只是地址,不提供路由。** 取消供应商 DNS 地址过滤后,进程所在网络仍必须
能通过对应的 TUN/透明代理处理 Fake-IP;否则会变成 TCP 超时,而不是过滤报错。
2. **隧道代理不等于网关直连代理。** 默认仍先本地解析并通过 ACL;只有显式启用
`upstream_proxy_remote_dns` 才委托代理解析供应商域名。代理端点自身的域名仍需
本地 DNS;本地解析完全不可用时,代理 URL 应使用可达 IP。
3. 隧道默认 `AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS=false` 仍会拒绝 Fake-IP。
该开关是扩大内网访问权限,不是建议普遍启用的 DNS 修复;应优先让隧道主机得到
可路由的真实 DNS 答案,或配置受信任代理的显式远程 DNS 模式。
4. 更新客户端支持自己的代理环境变量和 `NO_PROXY`;不能把这一点推广到供应商请求。
网关供应商 SOCKS 配置会归一化为远程 DNS 语义,但其它明确区分 SOCKS5/SOCKS5h
的独立工具仍遵循各自配置。
5. SMTP 等辅助服务不是供应商解析器的调用方。SMTP 已修复 DNS 超时和答案截断,
但不自动继承供应商出口代理。系统解析器由 Tokio 阻塞池承载,异步超时会停止等待,
不等于操作系统正在执行的 DNS 调用能够被强制终止。
6. 本轮没有使用生产凭据、发送真实模型请求、修改系统 DNS、关闭 TUN 或重启服务器。
生产验证仍需在实际容器/进程的网络命名空间内进行。
## 回归验证
- 公共 DNS:合法/非法 IPv6 主机形式、端口、零超时、答案上限、超限拒绝。
- 供应商 DNS:Fake-IP 保留,HTTP 与 wreq 解析结果一致;relay 私网过滤不变。
- WS:普通与浏览器指纹客户端,经本机 HTTP、SOCKS5、SOCKS5h mock 代理,使用
`provider-dns.invalid` 完成真实 WS upgrade;SOCKS mock 断言收到域名而非本地解析 IP。
这是本地明文 WS 的代理路径测试,不代替生产 WSS 的 TLS/SNI 检查。
- 连接测试:供应商 URL 校验一致,无预解析构建请求,重定向不转发凭据。
- 附件与辅助出口:多地址保留、私网/混合答案拒绝及官方 Fake-IP 例外回归。
- 隧道:IPv6 SOCKS 编码、地址 ACL、缓存策略隔离、默认固定 IP 代理连接;新增远程
DNS 模式的 HTTP CONNECT/SOCKS5h 域名握手、Host/SNI 主机名、禁止本地解析回退、
连接池隔离、代理/TLS 握手超时,以及 CLI/TOML/TUI 配置验证和持久化。
- SMTP:DNS 超时、空答案/错误信息、前 16 个地址不可用时使用第 17 个地址,以及
本机 mock SMTP 探测和邮件投递;不发送真实邮件。多地址测试先复现串行建连超时,
改成有界地址竞争后通过。
最终重新执行结果:**809 项测试通过,0 失败**。
- 网关:594 项,覆盖完整 WS 模块、执行传输、Grok、ChatGPT-Web 图片、连接测试,
DNS/Fake-IP/解析地址校验,以及 SMTP 发送、探测和相关配置回归。
- 隧道:197 项,全量单元/本机集成测试。
- `aether-http`:18 项,包含先失败、后修复通过的方括号 IPv6 回归。
- 单独执行的 13 项 SMTP 回归及前轮测试均为上述集合的子集,不重复计入总数。
- Rust 格式检查与 `git diff --check` 通过。
涉及监听器的测试在获准的沙箱外绑定本机回环端口;最初沙箱内的端口权限失败不作为
功能失败,也未通过跳过测试来规避。所有 Cargo 测试使用 `--offline`,代理上游是
本机 mock,不使用生产凭据。
复现命令:
```bash
cargo test -p aether-gateway --lib --offline -- --quiet \
handlers::proxy::websocket:: execution_runtime::transport::tests \
execution_runtime::grok::tests execution_runtime::chatgpt_web_image::tests \
test_connection bark_push::tests server_chan_push::tests \
dns fake_ip benchmarking resolved_addrs email_delivery smtp
cargo test -p aether-tunnel --offline -- --quiet
cargo test -p aether-http --offline
```
+57
View File
@@ -0,0 +1,57 @@
# 调度策略级故障转移
调度策略的 `default_policy` 支持跨提供商的转移预算与错误规则。它们跟随当前请求的 `routing_execution_policy` 快照进入执行器,不依赖运行中修改全局系统设置。已有策略缺省为不限次数、不限累计时间、无全局错误规则。
```json
{
"default_policy": {
"sticky_key_attempts": 2,
"max_transfer_count": 3,
"max_transfer_timeout_seconds": 90,
"failover_rules": {
"success_failover_patterns": [
{ "pattern": "(?i)capacity.*exhausted" }
],
"error_stop_patterns": [
{ "status_codes": [400, 413], "pattern": "invalid.*parameter" },
{ "status_codes": [422] }
]
}
}
}
```
## 预算语义
- `sticky_key_attempts` 是首个粘性候选上的总尝试次数,`2` 表示首次请求加一次同 Key 重试。该行为保持不变。
- 全局 `max_transfer_count` 统计切换候选的次数,首次尝试不计数。同一提供商、端点、Key 上的重试不计数;改变该组合计一次。`3` 最多允许首次候选之后再切换三次。
- 全局 `max_transfer_timeout_seconds` 从首次候选开始执行时计时,覆盖后续重试与切换间的累计耗时。它在准备下一次尝试时检查,不会强制打断已经执行中的调用或已提交给客户端的流;单次连接、首字节、读取和非流式完整调用超时仍独立生效。
- 两个全局预算的 `0` 都表示不限制。被筛除、禁用或被提供商级预算跳过而未执行的候选不计数。
- 提供商自身的转移次数和时间预算继续生效。提供商预算耗尽只跳过该提供商,仍可尝试其他提供商;全局预算耗尽则不再执行任何提供商的新尝试。
- 预算仅约束当前请求,不能重置或替代客户端取消策略、权限校验及本地执行异常的终止行为。
## 错误规则
先匹配调度策略的全局显式规则;未匹配时继续使用提供商的规则及协议默认行为。本地执行函数真正返回 `Err` 时仍然终止,不做兜底重放。
- **成功转移规则**:仅当 HTTP 200 响应匹配配置的正则时继续转移,不是对所有 200 进行重试。非流式请求匹配响应体;流式请求只匹配尚未交付业务输出的有界预读取内容。
- **结构化错误优先**:标准流式请求统一在预读取阶段先解析完整 SSE 事件或 JSON 中的错误,再判断成功正则;不在半截错误载荷上提前触发成功转移,避免旧的 JSON 提前探测绕过错误终止规则。普通文本响应仍支持跨分片匹配。
- **图片成功保护**:`openai:image` 的成功响应保留不重放行为,不因全局或提供商的成功正则再次生成图片;正常错误响应仍按错误规则处理。
- **错误终止规则**:适用于 400–599 错误。状态码与正则都填写时要求同时满足;只填状态码表示该状态一律终止;只填正则表示在所有错误状态上匹配。流内错误使用解析后的错误状态,而不是外层 200。
- **网络错误**:没有上游 HTTP 状态的连接、TLS、DNS、提交前超时等错误统一继续转移;提供商级别若单独配置了停止规则,仍按提供商规则处理。
- 正则使用 Rust `regex` 语法,支持 `(?i)` 等内联标志。服务端拒绝无效正则、无意义的空规则以及错误状态范围。每组最多 64 条,每条表达式最多 4096 字节。
## 流式 200 的恢复窗口
上游 HTTP 200 响应头不再默认关闭标准文本 SSE 的恢复窗口。执行器先缓冲协议开场事件,例如 Responses 的 `response.created`、Chat 的 role-only 增量、Anthropic 的空 `message_start` / 文本块起始事件。`openai:image` 图片专用流保留原来的响应头提交行为,成功正则不打开重放窗口。
首个业务内容之前的结构化错误、过早 EOF、首字节超时以及 200 正则命中会进入统一故障转移判断。真实文本、思考、工具调用或正常结束事件确定后,缓冲内容按原顺序交付;之后发生的错误保持终止,不重新执行原请求。
预读取受单次首字节时间和既有字节上限约束。达到字节上限时保守提交,避免无限缓冲;这意味着不能承诺识别响应任意位置的错误或正则。原始 HTTP 状态与最终执行结果是不同观测值,不能因流内失败而伪改已经发送的 HTTP 状态码。
## 排查
- `routing_transfer_limit_reached`:当前调度策略的累计次数或时间预算耗尽。
- `provider_transfer_limit_reached`:提供商自身的预算耗尽。
- `local_stream_candidate_retry_scheduled`:输出前的流内错误或 200 规则触发了继续调度。
- `local_stream_transport_retry_scheduled`:输出前传输错误触发了继续调度。
@@ -0,0 +1,259 @@
# Aether 系统、测试与 CI 瘦身审计
- 审计日期:2026-09-08(Asia/Shanghai)
- 源码基线:`8b766930b`
范围:仓库结构、依赖图、本地构建产物、GitHub Actions 实际日志、测试与发布边界。
本次只新增审计报告,没有修改业务代码、测试或工作流,没有清理文件,也没有访问生产服务器。没有采集生产 RSS、CPU、数据库体积或请求延迟,因此下文不能作为生产内存泄漏或吞吐退化的结论。
## 一、结论
优先减掉的是**重复构建、过重的测试夹具、没有实际用例的构建任务和失效的模块边界**,不是先删功能或减少回归断言。
- 普通 Rust CI 最近 30 条记录中,20 次成功运行的总耗时中位数约 **16 分 33 秒**,范围 **11 分 42 秒~19 分 03 秒**。样本覆盖 2026-09-05~2026-09-08,包含 push 和 PR;没有将失败、取消或未完成运行计入。
- 成功样本中 Gateway 是关键路径:一次耗时 **11 分 39 秒**,另一次 **16 分 46 秒**;既有巨型测试目标编译,也有数分钟测试执行,不能只归咎于缓存或测试数量。
- “Workspace Rest” 虽然排除了 Gateway 测试目标,仍经由 Tunnel 的开发依赖编译完整 Gateway。分 job 没有实现真正的依赖隔离。
- Nightly 文档测试花了 **252 秒**,41 个库实际执行 **0 个 doctest**;发布默认编译 4 个二进制,只上传其中 1 个。
- 本地 `target/` 约 **168GB**,其中 `debug/incremental/` 约 **116GB**。这是构建缓存,不是生产镜像或业务源码体积。
所有改进收益都需要前后对照验证。本文不会把并行 job 的节省简单相加为流水线总时长,也不会承诺尚未测量的提速比例。
## 二、实测基线
### 2.1 CI 运行记录
数据来自 `fawney19/Aether` 的 GitHub Actions API、各 job 的步骤时间戳及原始日志。
| 运行 ID | 类型、源码 | 观察结果 |
| --- | --- | --- |
| `34174603131` | Rust CI,`7113d04f`,成功 | 总耗时 11:55;17 个 job 累计执行 39:19 |
| `34153166516` | Rust CI,`099b810a`,成功 | 总耗时 17:02;17 个 job 累计执行 46:10 |
| `34132961583` | 正式发布,`v0.7.17`,成功 | 总耗时 33:25;最慢为 macOS Intel 构建 |
| `34163371099` | Nightly,`099b810a`,失败 | 总耗时 47:42;检查和编译成功,GHCR 发布 job 在启动阶段失败 |
普通 CI 的总耗时按 `updated_at - run_started_at` 统计,包含调度和收尾;各 job 耗时按自身开始、完成时间统计,累计值不是计费分钟。详细成功样本与本地基线的 Rust CI、Nightly、Release、Cargo profile 和 Vitest 配置没有差异,业务改动会影响测试数量与耗时。
`34174603131` 中各主要测试步骤:
| 任务 | 编译/链接 | 测试执行 | 步骤或 job 耗时 |
| --- | --- | --- | --- |
| Gateway lib | 4:57 | 5,139 个测试,213.807 秒 | Test lib 步骤 514 秒 |
| Gateway bins | 2:27 | 78 个测试,0.270 秒 | Test bins 步骤 151 秒 |
| Workspace Rest | 5:36 | 3,377 个测试,26.357 秒,另有 16 个跳过 | Test 步骤 366 秒 |
| Integration Scenarios | 5:16 | 15 个 bin 内单测及 11 个 E2E 用例,约 8 秒 | Test 步骤 325 秒 |
| Data | 未进一步拆分 | 保留数据库相关保障 | Test 步骤 55 秒,job 88 秒 |
另一轮 `34153166516` 的 Gateway lib 编译 6:22、执行 390.039 秒,bins 编译 3:05、执行 0.329 秒。运行环境和缓存差异明显,不能只用最快一轮估算收益。
### 2.2 源码与磁盘
源码只统计 Git 跟踪文件,行数为物理行,包含注释和测试,不等于生产代码行数。
| 项目 | 规模 |
| --- | --- |
| Cargo workspace | 42 个 package |
| Rust 源码 | 1,874 个文件,1,114,527 行 |
| Gateway package 的 Rust 文件 | 1,176 个文件,679,538 行,约占全部 Rust 行数 61% |
| 以 tests/test/testkit 等命名识别的 Rust 测试及支持内容 | 214,086 行;未加上大量内联 `#[cfg(test)]` 模块 |
| Gateway 架构守卫测试 | 13 个文件,14,918 行,208 个 `#[test]` |
| 前端测试 | 208 个文件,31,945 行;Nightly 实跑 1,486 个用例 |
| 本地 `target/debug/incremental/` | 约 116GB |
| 本地 `target/debug/deps/` | 约 49GB |
| 本地 `frontend/node_modules/` / `frontend/dist/` | 约 311MB / 8MB |
| 本地历史 `htmlcov/` / `logs/` | 约 75MB / 121MB,均被 Git 忽略 |
## 三、优先处理:不降低保障的浪费
### P1-1:解除 Tunnel 测试对完整 Gateway 的反向依赖
**证据:** `apps/aether-tunnel/Cargo.toml:48` 将带 `testkit` 的 Gateway 列为 dev-dependency。`.github/workflows/rust-ci.yml:208` 和 `.github/workflows/rust-ci.yml:394` 虽然排除 Gateway 目标,但 Rest 的真实日志仍出现编译 `aether-gateway`。Gateway、Rest、Integration 因此在独立 runner 中重复付出重型构建成本。
**建议:** 将 Tunnel 中需要完整 Gateway 的端到端场景迁到独立集成测试目标;Tunnel 的协议、状态机、配置等单测只依赖轻量契约和测试支持。迁移之后比较测试清单,确保没有丢失端到端场景。
**验收:** 普通 Rest 单测及 Clippy 的依赖闭包不再包含 Gateway;Tunnel 跨端集成场景仍在专门任务执行。此项主要降低累计 runner 工作量,是否缩短总时长取决于 Gateway 关键路径是否也得到优化。
### P1-2:发布只编译真正发布的二进制
**证据:** Cargo metadata 显示 Gateway 有 4 个 bin:服务主程序、backup-restore、两个 WebSocket probe。`.github/workflows/release.yml:227`、`.github/workflows/release.yml:229` 和 `.github/workflows/nightly.yml:296` 未选择具体 bin,但 `.github/workflows/release.yml:236` 只上传主程序。
2026-09-07 的正式发布中,macOS Intel job 耗时 31:24,编译/链接日志耗时 28:52;Nightly 对应 job 耗时 33:45。发布慢不能全算在测试头上。
**建议:** 正式发布与 Nightly 的 `cargo build` / `cross build` 添加 `--bin aether-gateway`;probe 和其他运维程序保留独立检查、测试或按需打包入口。保留当前 release 优化策略,先测减少目标的收益,再讨论 LTO 调整。
**边界:** 没有逐 bin 链接计时,不能声称这会让整个发布缩短四分之三。
### P1-3:取消空 doctest 构建,合并重复的驱动检查
- `.github/workflows/nightly.yml:100` 的 workspace doctest 步骤耗时 252 秒,41 个库合计 0 个用例。建议对确实无 doctest、且不计划承载文档示例的库显式管理 `doctest`,或通过独立清单检查是否存在可执行文档示例后再调度。未来新增示例必须能重新纳入检查,不能永久盲目跳过全部文档测试。
- `crates/aether-data/runtime/Cargo.toml:10` 中 `default = ["postgres"]`,`all-drivers = ["postgres"]`,源码没有单独依赖 `all-drivers` 的条件分支。`.github/workflows/rust-ci.yml:334` 的两个 feature job 实际没有覆盖两个不同数据库驱动。
- `.github/workflows/rust-ci.yml:394` 的 Rest 已执行 Postgres adapter 的测试,`.github/workflows/rust-ci.yml:403` 又运行一次。若保留独立 adapter job,应从 Rest 排除它;否则直接以 Rest 承担这份覆盖。
**收益口径:** 在样本中,去掉空 doctest 可省约 4:12 的该 job 时间;feature 与 adapter 重复工作为几十秒量级。它们多数不在普通 CI 关键路径上,不应当作主 CI 总时长的等额收益。
### P1-4:区分集成测试与压测工具的构建
**证据:** `crates/aether-testing/integration/Cargo.toml:1` 所属 package 自动发现 14 个场景 bin;其中 11 个没有测试函数。`.github/workflows/rust-ci.yml:467` 对整个 package 执行 `--bins --tests`,耗时 5:25,实际测试约 8 秒。
这里并没有执行那些压测程序的 `main()` 来验证容量、恢复时间或性能;空测试 harness 的成功不能视为压测成功。
**建议:** 将 E2E 测试与 benchmark/probe 工具分开;把三个工具内的 15 个单测迁入适合的库或独立目标。工具源码继续接受 Clippy/check,真实压测由手动或计划任务运行。必要时使用 `required-features` 门控工具 bin。
**注意:** 仅给 bin 添加 `test = false` 不足以阻止所有额外构建;Cargo 构建 integration test 时还可能自动构建同 package 的普通 bin。应从目标和 package 边界解决,而不是仅换命令拼写。
### P1-5:重做按变更类型路由,同时补覆盖缺口
**证据:** `.github/workflows/rust-ci.yml:9` 将 README、安装脚本、Compose 与 Rust 源码共同触发整条 Rust 流水线;job 内没有进一步区分。当前五份工作流没有 frontend PR 工作流;前端完整检查在 Nightly 执行。
另一个现存缺口:`apps/aether-gateway/tests/admin_unsigned_identity_headers.rs:16` 的普通 integration test 不在 Gateway 的 `--lib`、`--bins` 两条 nextest 命令内;独立 Integration Scenarios job 选择的是另一个 package。Nightly 的 `check --all-targets` 只检查编译,不会替代执行此安全用例。
**建议:**
1. 安装脚本、Compose、发布工作流改动优先运行对应安全 fixtures;纯 README 文档修改不必编译完整 Gateway。
2. Rust 改动执行相关测试;Cargo、工具链、公共契约和测试基础设施变更应保守扩大到完整检查。
3. 前端改动运行自己的类型检查和测试,不必等待 Nightly。
4. 显式加入 Gateway integration test,包括上述身份头安全用例。
5. 保留稳定的最终 `check` 门禁,并验证预期跳过的 job。不要让路径过滤造成 required check 永久 pending,或把失败当成允许跳过。
工具链触发项还应补查 `rust-toolchain.toml`、`.cargo/**`;这些文件目前不在该工作流的路径列表中。
## 四、关键路径:Gateway 测试要减“重量”而非减断言
### P1-6:改造昂贵夹具与不必要的全应用初始化
成功样本中最慢的用例包括:
| 用例 | 时间 | 代码位置 |
| --- | --- | --- |
| 跳过大量 blocked account 的扫描预算 | 22.668 秒 | `apps/aether-gateway/src/dispatch/pool_scheduler.rs:3940` |
| v1 备份兼容及历史密钥尝试 | 14.324 秒 | `apps/aether-gateway/src/backup/executor.rs:1223` |
| 跳过大量 exhausted account | 8.481 秒 | `apps/aether-gateway/src/dispatch/pool_scheduler.rs:3859` |
| 大池 LRU 和动态跳过 | 8.147 秒 | `apps/aether-gateway/src/dispatch/pool_scheduler.rs:4424` |
池调度测试会构造 1,700 个账号,并在 `apps/aether-gateway/src/dispatch/pool_scheduler.rs:4868` 为每个账号调用真实凭据封装。可以将扫描预算、分页、跳过逻辑用轻量 repository/credential fixture 验证,另保留少量真实加解密联调用例和大规模边界场景;不要简单把大池规模缩小到失去原来的回归条件。
`crates/aether-crypto/src/python_fernet.rs:302` 的历史密钥派生包含进程内缓存及 100,000 次 PBKDF2。真实 nextest 用例是分进程执行的,跨用例不能指望共享这份缓存。是否构成主要开销仍需针对性计时;可为不测试历史派生逻辑的夹具选择合法的固定测试密钥或预制密文,历史兼容和真实密码学用例必须保留生产强度。
**不要做:** 为通过 CI 下调生产加密迭代数、删掉备份兼容测试、跨测试共享可变 AppState、将关键安全测试统一 ignore。
### P1-7:把巨大单一测试目标拆成真正独立的边界
`apps/aether-gateway/src/lib.rs:234` 将广泛的内部测试树纳入同一个 lib test binary。单纯把大文件拆成几个 `mod` 文件,不会使它们成为独立 Cargo 编译单元。
建议优先迁出无需访问私有业务状态的架构守卫测试,再逐步将调度策略、协议转换、计费纯函数测试迁到所属 crate;HTTP 行为与跨模块场景放到有明确支持 API 的 integration target。避免为了迁移测试而把所有内部类型公开。
架构守卫现有 208 个测试、约 1.49 万行,很多是源码字符串和依赖规则检查。例如 `apps/aether-gateway/src/tests/architecture/workspace_tiers.rs:3` 按 manifest 字符串判断依赖边界。它们应继续存在,但可进入不依赖 Gateway 的小工具/测试目标;依赖规则优先检查 Cargo metadata,行为正确性继续由行为测试负责。
若采用 nextest 分片,应先复用一次构建产物,再分发执行;直接给 N 个 runner 各自重新编译巨型 Gateway 会放大成本。先衡量编译与执行占比,再确定是否分片。
**特别注意:** 仅将 `--lib` 与 `--bins` 合为一条命令,并不消除普通库与 `cfg(test)` 测试库的两种构建。不能把样本中 2:27 的 bins 编译时间直接记作可全部省掉。
### P1-8:缓存按实际构建方式组织
所有 Rust job 使用同一个 `shared-key`。实测 Gateway、Rest、Integration 恢复了同一份约 380MB 的缓存;日志显示 `cache-workspace-crates: false`。但 Gateway 的 mold `RUSTFLAGS` 只存在于测试 step,见 `.github/workflows/rust-ci.yml:265`,其余任务又使用不同的 Clippy/check/test 和 feature 组合。
Gateway 样本的 Rust sccache 命中率达到 95.48%,仍花了数分钟构建,并存在 76 次标记为 `crate-type` 的不可缓存调用。这说明“再装一个缓存”不是根治;同时也不能据此断言缓存无效。
建议把影响缓存选择的环境提前到 job 级,按 lint/check 与 test、target、toolchain、必要 feature 区分缓存用途;相同构建尽量统一。先测恢复、保存、不可缓存调用与编译时间,不要给每个细碎目标无限新建 cache key,也不要直接缓存完整的巨型 `target/`。
现有 nextest、sccache、Gateway mold 和 CI 的关闭 debug 信息配置已经到位,不列为“尚未实施”的建议。
## 五、源码和依赖的长期瘦身
### P2-1:继续完成已有模块边界,而不是继续增加空壳 crate
Gateway 仍承载约 67.95 万行 Rust。部分现有边界很薄:Gateway execution crate 82 行、control crate 100 行、provider core 91 行、usage core 110 行,而核心实现仍留在应用中。
值得分批治理的集中点:
| 文件 | 物理行数 | 建议拆分依据 |
| --- | --- | --- |
| `apps/aether-gateway/src/execution_runtime/stream/execution.rs:1` | 15,088 | 流状态机、传输适配、计费收尾、对应测试 |
| `crates/aether-data/adapters/postgres/src/usage/mod.rs:1` | 13,833 | 写入、查询、统计聚合、审计存储 |
| `crates/aether-usage/runtime/src/runtime.rs:1` | 12,809 | 状态推进、结算策略、持久化适配、测试 |
| `apps/aether-gateway/src/handlers/admin/request/system/import.rs:1` | 9,490 | 导入校验、版本兼容、执行与回滚 |
这些数字包含内联测试,不是生产实现行数。目标应是缩小依赖闭包、变更影响面和测试目标,而非追求拆出更多文件。
`apps/aether-gateway/src/state/app.rs:376` 的 AppState 有 91 个字段,其中 28 个字段名包含 cache;`crates/aether-data/contracts/src/repository/usage/types.rs:1984` 的 UpsertUsageRecord 有 67 个字段。这反映了测试构造和模块依赖面较宽,但不能据此直接判定运行时占用过大。可以引入按职责的窄上下文和统一 fixture builder,避免每个测试复制完整对象。
### P2-2:从基础契约中剥离重型格式实现
`crates/aether-data/contracts/Cargo.toml:10` 依赖整个 `aether-ai-formats`;后者约 8.23 万行 Rust。contracts 中实际使用包括格式权限、别名与少量 usage 元数据策略,见 `crates/aether-data/contracts/src/repository/auth.rs:295`。
建议把稳定的格式标识、权限和小型元数据契约下沉到现有基础契约层,完整 request/response/stream 转换留在 formats。避免为了一个格式权限判断,让数据库契约持续依赖整个转换实现。这比任意合并 crate 更有价值。
### P2-3:依赖体积优化必须保留兼容能力
`cargo tree --offline --locked` 确认 Gateway 同时包含:
- 主请求链路的 reqwest 0.12 与 `object_store` 引入的 reqwest 0.13。
- rustls 的 ring 和 aws-lc 路径,以及 wreq 的 boring2 路径。
这些是进一步分析构建和二进制体积的候选,不是已证实可直接删除的依赖。`aws-lc-rs` 还有直接密码学用途,wreq 承担专门传输能力;移除前必须核查调用和握手、指纹、代理兼容测试。
先获取 release 的 Cargo timings 和二进制符号/section 体积,再评估版本统一或可选能力 profile;不要仅根据依赖名字或锁文件重复条目盲目替换。
## 六、前端测试与构建
### P1/P2:测试环境按需要加载
Nightly 前端测试实测 106.01 秒,208 个文件、1,486 个用例全部通过。Vitest 报告 environment 146.40 秒、import 51.29 秒、tests 72.78 秒;这些分项含并行累计时间,不能相加当作墙钟时间。
`frontend/vitest.config.ts:10` 为全部测试使用 jsdom;静态扫描有 126/208 个测试文件未直接出现常见 DOM 操作标记,但这并不证明它们的传递依赖不需要 DOM。`frontend/src/tests/vitest.setup.ts:61` 还在每个用例前加载 i18n 并重设语言。
建议通过独立 Vitest project 或显式环境标记,将已确认的纯函数/解析器/数据转换测试放到 node 环境,DOM 组件保留 jsdom;按测试类别拆分 setup,不要全局取消隔离。先迁一组、验证测试数与结果,再扩展。
### P1:同一 web 项目重复构建
`.github/workflows/release.yml:84` 和 `.github/workflows/deploy-pages.yml:58` 已单独构建 VSCodex web;随后 frontend 的 `prebuild` 又调用 `frontend/scripts/sync-vscodex.mjs:64` 无条件重建一次。
建议分开“安装/构建嵌入 web”与“复制已构建产物”,在同一个任务内只构建一次,消费明确来源的 artifact。不要仅通过 `dist` 存在就认定源码和产物一致。Nightly 前端这里没有相同的预先重复 build,不应错误地宣称所有流水线都重复。
### P2:低风险依赖清理候选
当前前端源码未发现 `three` 的模块导入,但 package 声明了 `three` 和 `@types/three`;本地二者合计约 36MB。确认无动态/外部消费者后可移除并更新锁文件,验证 type-check、测试和构建。
这主要减少安装和维护成本;不能保证生产 bundle 同样减少 36MB。现有 chart、pinyin、Stripe 均发现使用,不能一并判作无用依赖。
MarkdownViewer 使用完整 `highlight.js`,而 CodeHighlight 已按语言导入,可统一按需高亮策略;现有前端产物约 8MB,优先级低于 Rust 重复构建和测试初始化。
## 七、本地构建产物治理
当前最值得清理的是 116GB 的 `target/debug/incremental/`,其次是 49GB 的 `target/debug/deps/`。全量 `cargo clean` 虽能回收空间,也会迫使下次重建全部依赖。
建议先确认没有进行中的 Cargo/rustc 任务,对长期未使用的增量会话、历史目标产物做定期清理;稳定本地 profile、工具链和 `RUSTFLAGS`,避免频繁产生不同构建组合。保留仍在使用的依赖缓存,不要每次构建前清空 target。
`htmlcov/` 属于历史 Python 覆盖率产物,可在确认不再需要后清理,但仅几十 MB,不是主要收益。不要误删仍有效的 Python 安装/Compose 安全测试,也不要把日志、备份或数据目录当构建缓存删除。
生产 Dockerfile 已使用预构建二进制、前端产物和 distroless runtime,并通过 `.dockerignore` 排除 target、node_modules 等开发内容;不建议把“换更小基础镜像”列为当前第一优先级。生产镜像实际体积还需单独测量。
## 八、建议实施顺序与验收
| 批次 | 改造 | 验收标准 |
| --- | --- | --- |
| A:小改动去空转 | 发布指定主 bin;管理空 doctest;消除重复 adapter 与等价 feature 任务;嵌入 web 只构建一次 | 保留现有有效用例与工件;对照任务耗时和测试清单 |
| B:加快反馈并补漏 | 变更路由、稳定 gate、前端 PR 检查、Gateway integration 安全用例 | 纯文档/脚本不编译全栈;公共变更仍完整检查;安全用例实际执行 |
| C:解除编译耦合 | Tunnel 跨端测试迁移;工具与 E2E 分离;架构守卫轻量化 | Rest 依赖闭包无 Gateway;无用 bin 不参与 E2E 构建;守卫持续有效 |
| D:减少测试初始化 | 重型夹具拆层;纯前端测试切 node;窄上下文和统一 builder | 慢用例时间降低,测试数与断言目的不减少,无新增污染或 flaky |
| E:结构与依赖治理 | 真实职责迁入现有 crate;基础格式契约下沉;审慎统一依赖 | 普通修改触发的重编译面缩小,性能和兼容性基线不回退 |
每批先使用相同 SHA、runner 类型、toolchain 和 feature 集合做对照,至少区分冷缓存/热缓存与 job 执行/流水线总耗时。持续保存编译 timings、nextest 测试清单和结果、慢测试列表、缓存统计、frontend environment 时间、发布工件体积。
首要结果指标:普通 PR 更快得到正确反馈、累计重复构建下降、测试保障不退化。不要把删除测试数量、增大并发数或缩短单个非关键任务当作最终目标。
## 九、复查入口
本次没有重新运行全量构建或全量测试,使用了真实 CI 日志、离线 Cargo metadata/tree 和只读源码统计。临时 API JSON、job 日志与依赖树保存在本机 `/tmp/aether-slim-audit/`,未纳入 Git;临时目录可能被系统清理,运行 ID 可用于再次取证。
可重复的只读命令:
```sh
cargo metadata --offline --locked --no-deps --format-version 1
cargo tree --offline --locked -p aether-gateway -e normal -d
gh api 'repos/fawney19/Aether/actions/runs/34174603131/jobs?per_page=100'
gh api 'repos/fawney19/Aether/actions/jobs/101901417322/logs'
gh api 'repos/fawney19/Aether/actions/jobs/101869558414/logs'
du -h -d 1 target/debug
```
Cargo 关于默认构建目标及 integration test 自动构建 bin 的语义,另对照本机 Rust 1.95.0 随附的 Cargo `cargo-build`、`cargo-test`、`cargo-targets` 官方文档。并行测试时间、缓存命中率和源码体积均按各自定义解释,未将其混用为生产性能结论。
@@ -4,6 +4,7 @@ const { getMock } = vi.hoisted(() => ({ getMock: vi.fn() }))
vi.mock('@/api/client', () => ({ default: { get: getMock } }))
import { dashboardApi } from '@/api/dashboard'
import { requestTraceApi } from '@/api/requestTrace'
import { cache } from '@/utils/cache'
beforeEach(() => {
@@ -13,6 +14,20 @@ beforeEach(() => {
})
describe('dashboard body loading', () => {
it('reports download progress so diagnostic exports can abort oversized bodies', async () => {
const onProgress = vi.fn()
getMock.mockResolvedValue({ data: new ArrayBuffer(0), headers: { 'x-aether-body-encoding': 'json', 'x-aether-usage-id': 'usage-1', 'x-aether-body-field': 'response_body' } })
await dashboardApi.getRequestBody('usage-1', 'response_body', undefined, onProgress)
getMock.mock.calls[0][1].onDownloadProgress({ loaded: 2 * 1024 * 1024 })
expect(onProgress).toHaveBeenCalledWith(2 * 1024 * 1024)
})
it('records the gateway version at export without confusing it with failure-time metadata', async () => {
getMock.mockResolvedValue({ data: { request_id: 'request-1', candidates: [] }, headers: { 'x-aether-build-version': 'test-build' } })
expect(await requestTraceApi.getRequestTrace('request-1')).toMatchObject({ gateway_version: 'test-build' })
getMock.mockResolvedValue({ data: { request_id: 'request-1', candidates: [] } })
expect(await requestTraceApi.getRequestTrace('request-1')).toMatchObject({ gateway_version: null })
})
it('requests opaque body bytes, outside the JSON detail cache', async () => {
const bytes = new ArrayBuffer(20)
const controller = new AbortController()
+2 -1
View File
@@ -481,11 +481,12 @@ export const dashboardApi = {
return options.signal ? fetchDetail() : cachedRequest(cacheKey, fetchDetail, cacheTtlMs)
},
async getRequestBody(requestId: string, field: RequestBodyField, signal?: AbortSignal) {
async getRequestBody(requestId: string, field: RequestBodyField, signal?: AbortSignal, onProgress?: (loaded: number) => void) {
const response = await apiClient.get<ArrayBuffer>(`/api/admin/usage/${requestId}`, {
params: { include_bodies: true, body_field: field, body_format: 'raw' },
responseType: 'arraybuffer',
signal,
...(onProgress ? { onDownloadProgress: (event: { loaded: number }) => onProgress(event.loaded) } : {}),
})
const encoding = response.headers['x-aether-body-encoding']
if ((encoding !== 'gzip' && encoding !== 'json') || response.headers['x-aether-usage-id'] !== requestId || response.headers['x-aether-body-field'] !== field) {
+3 -1
View File
@@ -121,6 +121,8 @@ export interface CandidateRecord {
}
export interface RequestTrace {
gateway_version?: string | null
diagnostic_request?: { usage_id: string, body_state?: string | null } | null
request_id: string
request_path?: string
request_query_string?: string
@@ -154,7 +156,7 @@ export const requestTraceApi = {
const response = await apiClient.get<RequestTrace>(`/api/admin/monitoring/trace/${requestId}`, {
params: { attempted_only: attemptedOnly },
})
return response.data
return { ...response.data, gateway_version: response.headers?.['x-aether-build-version'] ?? response.data.gateway_version ?? null }
},
/**
@@ -1,10 +1,32 @@
import { describe, expect, it } from 'vitest'
import { createSSRApp, h } from 'vue'
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import { createApp, createSSRApp, h, nextTick, type App } from 'vue'
import { renderToString } from '@vue/server-renderer'
import type { ProviderWithEndpointsSummary } from '@/api/endpoints'
import type { EndpointAPIKey } from '@/api/endpoints/keys'
import ModelMappingTab from '../provider-tabs/ModelMappingTab.vue'
const keyMocks = vi.hoisted(() => ({ getProviderKeys: vi.fn() }))
vi.mock('@/api/endpoints/keys', () => keyMocks)
const testMocks = vi.hoisted(() => ({
testModel: vi.fn(),
getRequestTrace: vi.fn(),
showError: vi.fn(),
showSuccess: vi.fn(),
}))
vi.mock('@/api/endpoints/providers', async importOriginal => ({
...await importOriginal<typeof import('@/api/endpoints/providers')>(),
testModel: testMocks.testModel,
}))
vi.mock('@/api/requestTrace', () => ({
requestTraceApi: { getRequestTrace: testMocks.getRequestTrace },
}))
vi.mock('@/composables/useToast', () => ({
useToast: () => ({ error: testMocks.showError, success: testMocks.showSuccess }),
}))
const provider: ProviderWithEndpointsSummary = {
id: 'provider-demo',
name: 'Demo Provider',
@@ -29,7 +51,212 @@ const provider: ProviderWithEndpointsSummary = {
updated_at: '2026-01-01T00:00:00Z',
}
type MappingTabProps = InstanceType<typeof ModelMappingTab>['$props']
type MappingTestState = {
runMappingTest: (key: string, model: string) => void
handleSelectTestEndpoint: (id: string) => void
selectedTestKeyIds: string[]
testKeyOptions: Array<{ value: string; label: string }>
loadingTestKeys: boolean
handleTestDialogClose: () => void
handleStartMappingTest: () => Promise<void>
}
const endpoints = [
{ id: 'chat', api_format: 'openai:chat', base_url: 'https://example.com', is_active: true, active_keys: 1 },
{ id: 'claude', api_format: 'claude:messages', base_url: 'https://example.com', is_active: true, active_keys: 1 },
] as MappingTabProps['endpoints']
function createTestKey(overrides: Partial<EndpointAPIKey>): EndpointAPIKey {
return {
id: 'test-key',
provider_id: provider.id,
name: 'Test Key',
api_formats: [],
api_key_masked: '',
auth_type: 'api_key',
internal_priority: 0,
cache_ttl_minutes: 0,
max_probe_interval_minutes: 1,
health_score: 1,
consecutive_failures: 0,
request_count: 0,
success_count: 0,
error_count: 0,
success_rate: 0,
avg_response_time_ms: 0,
is_active: true,
created_at: provider.created_at,
updated_at: provider.updated_at,
...overrides,
}
}
const testKeys = [
createTestKey({ id: 'chat-key', name: 'Chat Key', api_key_masked: 'sk-****chat', api_formats: ['openai:chat'], internal_priority: 0 }),
createTestKey({ id: 'claude-key', name: 'Claude Key', api_formats: ['claude:messages'], internal_priority: 1 }),
createTestKey({ id: 'disabled-key', is_active: false, internal_priority: 2 }),
]
const mounted: Array<{ app: App; root: HTMLElement }> = []
function mountMappingTab(overrides: Partial<MappingTabProps> = {}) {
const root = document.createElement('div')
document.body.appendChild(root)
const app = createApp(ModelMappingTab, { provider, endpoints, models: [], ...overrides })
const instance = app.mount(root)
const state = (instance.$ as unknown as { setupState: MappingTestState }).setupState
mounted.push({ app, root })
return state
}
function buttonWithText(text: string): HTMLButtonElement {
const button = [...document.querySelectorAll('button')]
.find(element => element.textContent?.trim() === text)
if (!button) throw new Error(`Missing button: ${text}`)
return button
}
async function openMappingTest(state: MappingTestState) {
state.runMappingTest('mapping', 'test-model')
await Promise.resolve()
await nextTick()
}
beforeEach(() => {
vi.resetAllMocks()
keyMocks.getProviderKeys.mockResolvedValue(testKeys)
testMocks.testModel.mockResolvedValue({ success: true, model: 'test-model' })
testMocks.getRequestTrace.mockResolvedValue(null)
})
afterEach(() => {
for (const { app, root } of mounted.splice(0)) {
app.unmount()
root.remove()
}
})
describe('ModelMappingTab response contracts', () => {
it('loads test keys and removes incompatible selections when switching endpoints', async () => {
const state = mountMappingTab()
await openMappingTest(state)
expect(keyMocks.getProviderKeys).toHaveBeenCalledWith(provider.id)
expect(state.testKeyOptions).toEqual([
{ value: 'chat-key', label: 'Chat Key · sk-****chat · api_key' },
])
expect(document.body.textContent).toContain('测试 Key')
buttonWithText('默认调度(不指定 Key)').click()
await nextTick()
const option = [...document.querySelectorAll<HTMLInputElement>('input[type="checkbox"]')]
.find(element => element.parentElement?.textContent?.includes('Chat Key'))
if (!option) throw new Error('Missing Chat Key option')
option.click()
await nextTick()
expect(state.selectedTestKeyIds).toEqual(['chat-key'])
state.handleSelectTestEndpoint('claude')
expect(state.selectedTestKeyIds).toEqual([])
expect(state.testKeyOptions.map(option => option.value)).toEqual(['claude-key'])
state.selectedTestKeyIds = ['claude-key']
state.handleTestDialogClose()
expect(state.selectedTestKeyIds).toEqual([])
})
it.each([
{ selectedKeyIds: ['chat-key'] },
{ selectedKeyIds: ['chat-key', 'shared-key'] },
])('passes the selected keys to the test request: $selectedKeyIds', async ({ selectedKeyIds }) => {
keyMocks.getProviderKeys.mockResolvedValue([
...testKeys,
createTestKey({ id: 'shared-key', name: 'Shared Key', internal_priority: 3 }),
])
const state = mountMappingTab()
await openMappingTest(state)
state.selectedTestKeyIds = [...selectedKeyIds, selectedKeyIds[0], 'disabled-key', 'claude-key']
await state.handleStartMappingTest()
expect(testMocks.testModel).toHaveBeenCalledExactlyOnceWith(expect.objectContaining({
provider_id: provider.id,
mode: 'direct',
model_name: 'test-model',
endpoint_id: 'chat',
api_format: 'openai:chat',
api_key_ids: selectedKeyIds,
}), expect.objectContaining({ signal: expect.any(AbortSignal) }))
})
it('keeps default scheduling when no key is selected', async () => {
const state = mountMappingTab()
await openMappingTest(state)
await state.handleStartMappingTest()
expect(testMocks.testModel).toHaveBeenCalledOnce()
expect(testMocks.testModel.mock.calls[0][0]).not.toHaveProperty('api_key_ids')
})
it('waits for keys to load before allowing a test', async () => {
let resolveKeys!: (keys: EndpointAPIKey[]) => void
keyMocks.getProviderKeys.mockReturnValue(new Promise<EndpointAPIKey[]>(resolve => {
resolveKeys = resolve
}))
const state = mountMappingTab()
await openMappingTest(state)
expect(buttonWithText('正在加载 Key').disabled).toBe(true)
expect(buttonWithText('开始测试').disabled).toBe(true)
await state.handleStartMappingTest()
expect(testMocks.testModel).not.toHaveBeenCalled()
resolveKeys(testKeys)
await Promise.resolve()
await nextTick()
expect(buttonWithText('开始测试').disabled).toBe(false)
})
it('keeps the selector visible when no compatible keys are available', async () => {
keyMocks.getProviderKeys.mockResolvedValue([testKeys[1], testKeys[2]])
const state = mountMappingTab()
await openMappingTest(state)
expect(document.body.textContent).toContain('测试 Key')
buttonWithText('默认调度(不指定 Key)').click()
await nextTick()
expect(document.body.textContent).toContain('暂无可选 Key')
})
it('keeps provided keys usable after a loading failure', async () => {
keyMocks.getProviderKeys.mockRejectedValue(new Error('Key service unavailable'))
const state = mountMappingTab({ providerKeys: testKeys })
await openMappingTest(state)
expect(testMocks.showError).toHaveBeenCalledOnce()
expect(state.loadingTestKeys).toBe(false)
expect(state.testKeyOptions.map(option => option.value)).toEqual(['chat-key'])
expect(buttonWithText('开始测试').disabled).toBe(false)
})
it('ignores key responses from a closed dialog', async () => {
let resolveKeys!: (keys: EndpointAPIKey[]) => void
keyMocks.getProviderKeys.mockReturnValueOnce(new Promise<EndpointAPIKey[]>(resolve => {
resolveKeys = resolve
}))
const state = mountMappingTab()
await openMappingTest(state)
state.handleTestDialogClose()
expect(state.loadingTestKeys).toBe(false)
keyMocks.getProviderKeys.mockResolvedValue([])
await openMappingTest(state)
resolveKeys(testKeys)
await Promise.resolve()
await nextTick()
expect(state.testKeyOptions).toEqual([])
})
it('keeps the module visible when a legacy or malformed preview reaches the component', async () => {
const props: InstanceType<typeof ModelMappingTab>['$props'] = {
provider,
@@ -345,18 +345,22 @@
:request-body-draft="testRequestBodyDraft"
:request-body-reset-value="testRequestBodyResetValue"
:request-body-error="testRequestBodyError"
:start-disabled="!selectedTestEndpoint || !!testRequestHeadersError || !!testRequestBodyError"
:key-options="testKeyOptions"
:selected-key-ids="selectedTestKeyIds"
:key-options-loading="loadingTestKeys"
:start-disabled="loadingTestKeys || !selectedTestEndpoint || !!testRequestHeadersError || !!testRequestBodyError"
@close="handleTestDialogClose"
@back="handleTestDialogBack"
@select-endpoint="handleSelectTestEndpoint"
@start="handleStartMappingTest"
@update:request-headers-draft="testRequestHeadersDraft = $event"
@update:request-body-draft="testRequestBodyDraft = $event"
@update:selected-key-ids="selectedTestKeyIds = $event"
/>
</template>
<script setup lang="ts">
import { ref, computed } from 'vue'
import { ref, computed, watch } from 'vue'
import { useSmartPagination } from '@/composables/useSmartPagination'
import { useModelTest } from '@/composables/useModelTest'
import { Tag, Plus, Edit, Trash2, ChevronRight, Loader2, Play } from 'lucide-vue-next'
@@ -374,7 +378,7 @@ import {
type ProviderMappingPreviewResponse,
} from '@/api/endpoints'
import { formatApiFormat } from '@/api/endpoints/types/api-format'
import { type EndpointAPIKey } from '@/api/endpoints/keys'
import { getProviderKeys, type EndpointAPIKey } from '@/api/endpoints/keys'
import { updateModel } from '@/api/endpoints/models'
import { useI18n } from '@/i18n'
import { parseApiError } from '@/utils/errorParser'
@@ -391,6 +395,7 @@ import {
isModelTestableApiFormat,
isModelTestableEndpoint,
modelTestMappingScopeMatchesEndpoint,
modelTestKeySupportsEndpoint,
parseModelTestRequestHeadersDraft,
parseModelTestRequestBodyDraft,
selectPreferredModelTestEndpoint,
@@ -452,6 +457,56 @@ const testingModelName = ref<string | null>(null)
const testingSourceModel = ref<Model | null>(null)
const preselectedModelId = ref<string | null>(null)
const selectedTestEndpoint = ref<ProviderEndpoint | null>(null)
const selectedTestKeyIds = ref<string[]>([])
const testKeys = ref<EndpointAPIKey[] | null>(null)
const loadingTestKeys = ref(false)
let testKeysLoadVersion = 0
const testKeyOptions = computed(() => {
const endpoint = selectedTestEndpoint.value
if (!endpoint) return []
return [...new Map((testKeys.value ?? props.providerKeys ?? []).map(key => [key.id, key])).values()]
.filter(key => modelTestKeySupportsEndpoint(key, endpoint, props.provider.provider_type))
.sort((left, right) => left.internal_priority - right.internal_priority)
.map(key => ({
value: key.id,
label: [
key.name?.trim() || key.api_key_masked?.trim() || key.id,
key.name?.trim() ? key.api_key_masked?.trim() : '',
key.auth_type?.trim(),
].filter(Boolean).join(' · '),
}))
})
function pruneSelectedTestKeyIds() {
const allowed = new Set(testKeyOptions.value.map(option => option.value))
selectedTestKeyIds.value = [...new Set(selectedTestKeyIds.value.filter(id => allowed.has(id)))]
}
async function loadTestKeys() {
const version = ++testKeysLoadVersion
const providerId = props.provider.id
loadingTestKeys.value = true
try {
const keys = await getProviderKeys(providerId)
if (version === testKeysLoadVersion && providerId === props.provider.id) {
testKeys.value = keys
}
} catch (err: unknown) {
if (version === testKeysLoadVersion && providerId === props.provider.id) {
showError(parseApiError(err, '加载测试 Key 失败'), '错误')
}
} finally {
if (version === testKeysLoadVersion) loadingTestKeys.value = false
}
}
watch(testKeyOptions, pruneSelectedTestKeyIds)
watch(() => props.provider.id, () => {
testKeysLoadVersion += 1
testKeys.value = null
selectedTestKeyIds.value = []
loadingTestKeys.value = false
})
const testRequestHeadersDraft = ref('')
const testRequestHeadersResetValue = ref('')
const testRequestBodyDraft = ref('')
@@ -789,11 +844,15 @@ async function onDialogSaved() {
function handleTestDialogClose() {
modelTest.resetState()
testKeysLoadVersion += 1
loadingTestKeys.value = false
testKeys.value = null
pendingMappingKey.value = null
testingModelName.value = null
testingSourceModel.value = null
testingMapping.value = null
selectedTestEndpoint.value = null
selectedTestKeyIds.value = []
mappingTestEndpoints.value = null
testRequestHeadersDraft.value = ''
testRequestHeadersResetValue.value = ''
@@ -811,6 +870,7 @@ function handleSelectTestEndpoint(endpointId: string) {
const endpoint = selectableTestEndpoints.value.find(item => item.id === endpointId)
if (!endpoint) return
selectedTestEndpoint.value = endpoint
pruneSelectedTestKeyIds()
syncMappingTestRequestBody()
}
@@ -839,6 +899,8 @@ function runMappingTest(
return
}
pendingMappingKey.value = testingKey
selectedTestKeyIds.value = []
void loadTestKeys()
modelTest.testResult.value = null
modelTest.dialogOpen.value = true
testingMapping.value = null
@@ -884,7 +946,7 @@ function syncMappingTestRequestBody() {
}
async function handleStartMappingTest() {
if (modelTest.testing.value || !testingModelName.value) return
if (modelTest.testing.value || loadingTestKeys.value || !testingModelName.value) return
const endpoint = selectedTestEndpoint.value || selectableTestEndpoints.value[0]
if (!endpoint) {
showError('请选择要测试的端点')
@@ -905,6 +967,7 @@ async function handleStartMappingTest() {
const currentMappingKey = pendingMappingKey.value || testingModelName.value
testingMapping.value = pendingMappingKey.value ? currentMappingKey : null
pruneSelectedTestKeyIds()
await modelTest.startTest({
mode: 'direct',
modelName: testingModelName.value,
@@ -912,6 +975,7 @@ async function handleStartMappingTest() {
apiFormat: endpoint.api_format,
endpointId: endpoint.id,
endpointBaseUrl: endpoint.base_url,
apiKeyIds: selectedTestKeyIds.value,
requestHeaders,
requestBody,
})
@@ -875,7 +875,7 @@ const modelMappingAvailable = computed(
() => props.modelMappingAvailable === true && modelMappingOptions.value.length > 0,
)
const showKeySelector = computed(() => (
keyOptionsLoading.value || keyOptions.value.length > 0 || selectedKeyIds.value.length > 0
props.keyOptions !== undefined || keyOptionsLoading.value || selectedKeyIds.value.length > 0
))
const keySelectorPlaceholder = computed(() => (
keyOptionsLoading.value && keyOptions.value.length === 0 ? '正在加载 Key' : '默认调度(不指定 Key)'
@@ -0,0 +1,227 @@
import { afterEach, describe, expect, it } from 'vitest'
import { createApp, h, nextTick, ref, type App } from 'vue'
import RoutingFailoverPolicyEditor from '../components/RoutingFailoverPolicyEditor.vue'
import { normalizeRoutingFailoverPolicy, type RoutingFailoverPolicy } from '../utils/routingFailover'
const mounted: Array<{ app: App, root: HTMLElement }> = []
function mountEditor() {
const policy = ref(normalizeRoutingFailoverPolicy())
const editor = ref<{ commitJsonDrafts: () => boolean } | null>(null)
const pending = ref(false)
const generation = ref(0)
const disabled = ref(false)
const root = document.createElement('div')
document.body.appendChild(root)
const app = createApp({
setup: () => () => h(RoutingFailoverPolicyEditor, {
ref: editor,
key: generation.value,
disabled: disabled.value,
modelValue: policy.value,
'onUpdate:modelValue': (value: RoutingFailoverPolicy) => { policy.value = value },
onPendingChange: (value: boolean) => { pending.value = value },
}),
})
app.mount(root)
mounted.push({ app, root })
return { root, policy, editor, pending, generation, disabled }
}
async function input(element: HTMLInputElement | HTMLTextAreaElement, value: string) {
element.value = value
element.dispatchEvent(new Event('input', { bubbles: true }))
await nextTick()
}
function control<T extends HTMLElement>(root: HTMLElement, label: string): T {
const element = root.querySelector<T>(`[aria-label="${label}"]`)
if (!element) throw new Error(`Missing control: ${label}`)
return element
}
afterEach(() => {
for (const { app, root } of mounted.splice(0)) {
app.unmount()
root.remove()
}
})
describe('RoutingFailoverPolicyEditor', () => {
it.each(['{}', '{"success_failover_pattern":[]}', '{"failover_rules":{"error_stop_patterns":[]}}'])('rejects missing JSON sections instead of silently clearing rules: %s', async draft => {
const { root, policy, editor } = mountEditor()
policy.value.failover_rules.success_failover_patterns = [{ pattern: 'capacity', status_codes: [] }]
await nextTick()
control<HTMLButtonElement>(root, '切到成功转移规则 JSON').click()
await nextTick()
await input(control<HTMLTextAreaElement>(root, '成功转移规则 JSON'), draft)
expect(editor.value?.commitJsonDrafts()).toBe(false)
expect(policy.value.failover_rules.success_failover_patterns).toEqual([{ pattern: 'capacity', status_codes: [] }])
})
it('accepts a named JSON section nested in a complete policy', async () => {
const { root, policy, editor } = mountEditor()
control<HTMLButtonElement>(root, '切到成功转移规则 JSON').click()
await nextTick()
await input(control<HTMLTextAreaElement>(root, '成功转移规则 JSON'), '{"failover_rules":{"success_failover_patterns":[{"pattern":"capacity"}]}}')
expect(editor.value?.commitJsonDrafts()).toBe(true)
expect(policy.value.failover_rules.success_failover_patterns).toEqual([{ pattern: 'capacity', status_codes: [] }])
})
it('keeps status drafts attached to their rules after a row is deleted', async () => {
const { root, policy, editor } = mountEditor()
for (const index of [1, 2]) {
control<HTMLButtonElement>(root, '添加错误终止规则').click()
await nextTick()
await input(control<HTMLInputElement>(root, `终止规则 ${index} 状态码`), index === 1 ? '400,' : '429, 503')
}
control<HTMLButtonElement>(root, '删除错误终止规则 1').click()
await nextTick()
expect(control<HTMLInputElement>(root, '终止规则 1 状态码').value).toBe('429, 503')
expect(editor.value?.commitJsonDrafts()).toBe(true)
expect(policy.value.failover_rules.error_stop_patterns).toEqual([{ pattern: '', status_codes: [429, 503] }])
})
it('disables JSON mode switches and formatting while a save is in flight', async () => {
const { root, policy, editor, disabled } = mountEditor()
control<HTMLButtonElement>(root, '切到成功转移规则 JSON').click()
await nextTick()
await input(control<HTMLTextAreaElement>(root, '成功转移规则 JSON'), '[{"pattern":"capacity"}]')
disabled.value = true
await nextTick()
expect(control<HTMLButtonElement>(root, '切回成功转移规则表单').disabled).toBe(true)
for (const button of root.querySelectorAll<HTMLButtonElement>('button')) expect(button.disabled).toBe(true)
expect(editor.value?.commitJsonDrafts()).toBe(false)
expect(policy.value.failover_rules.success_failover_patterns).toEqual([])
})
it('commits both JSON sections atomically when saving without returning to the form', async () => {
const { root, policy, editor, pending } = mountEditor()
control<HTMLButtonElement>(root, '切到成功转移规则 JSON').click()
control<HTMLButtonElement>(root, '切到错误终止规则 JSON').click()
await nextTick()
const [successJson, errorJson] = root.querySelectorAll<HTMLTextAreaElement>('textarea')
await input(successJson, '[{"pattern":"(?i)capacity"}]')
await input(errorJson, '[{"status_codes":[400,413]}]')
expect(editor.value?.commitJsonDrafts()).toBe(true)
await nextTick()
expect(policy.value.failover_rules.success_failover_patterns).toEqual([{ pattern: '(?i)capacity', status_codes: [] }])
expect(policy.value.failover_rules.error_stop_patterns).toEqual([{ pattern: '', status_codes: [400, 413] }])
expect(pending.value).toBe(false)
})
it('marks JSON-only edits pending and never partially applies invalid drafts', async () => {
const { root, policy, editor, pending } = mountEditor()
control<HTMLButtonElement>(root, '切到成功转移规则 JSON').click()
control<HTMLButtonElement>(root, '切到错误终止规则 JSON').click()
await nextTick()
const [successJson, errorJson] = root.querySelectorAll<HTMLTextAreaElement>('textarea')
await input(successJson, '[{"pattern":"capacity"}]')
expect(pending.value).toBe(true)
await input(errorJson, '{')
expect(editor.value?.commitJsonDrafts()).toBe(false)
expect(policy.value.failover_rules.success_failover_patterns).toEqual([])
await input(errorJson, '[{"status_codes":[429]}]')
expect(editor.value?.commitJsonDrafts()).toBe(true)
await nextTick()
expect(root.querySelector('[role="alert"]')).toBeNull()
})
it('preserves separators during status-code typing and rejects invalid local input at save time', async () => {
const { root, policy, editor } = mountEditor()
control<HTMLButtonElement>(root, '添加错误终止规则').click()
await nextTick()
const statuses = control<HTMLInputElement>(root, '终止规则 1 状态码')
for (const value of ['4', '40', '400', '400,', '400, ', '400, 4', '400, 41', '400, 413']) {
await input(statuses, value)
expect(statuses.value).toBe(value)
}
expect(editor.value?.commitJsonDrafts()).toBe(true)
await nextTick()
expect(policy.value.failover_rules.error_stop_patterns[0].status_codes).toEqual([400, 413])
await input(statuses, 'oops')
expect(editor.value?.commitJsonDrafts()).toBe(false)
})
it('rejects non-finite limits instead of normalizing them to unlimited', async () => {
const { policy, editor } = mountEditor()
policy.value.max_transfer_count = Number.NaN
await nextTick()
expect(editor.value?.commitJsonDrafts()).toBe(false)
})
it('edits independent global budgets and documents sticky retry exclusion', async () => {
const { root, policy } = mountEditor()
expect(root.textContent).toContain('首次尝试和粘性同 Key 重试不计入')
expect(root.textContent).toContain('不会中断已开始的调用')
const count = control<HTMLInputElement>(root, '全局最大转移次数')
count.value = '4'
count.dispatchEvent(new Event('input', { bubbles: true }))
await nextTick()
expect(policy.value.max_transfer_count).toBe(4)
expect(policy.value.max_transfer_timeout_seconds).toBe(0)
})
it('adds regex and status-only rules and reports invalid drafts', async () => {
const { root, policy } = mountEditor()
control<HTMLButtonElement>(root, '添加成功转移规则').click()
await nextTick()
expect(root.querySelector('[role="alert"]')?.textContent).toContain('正则表达式')
const regex = control<HTMLInputElement>(root, '成功转移规则 1 正则')
regex.value = '(?i)capacity.*exhausted'
regex.dispatchEvent(new Event('input', { bubbles: true }))
await nextTick()
expect(policy.value.failover_rules.success_failover_patterns[0].pattern).toBe('(?i)capacity.*exhausted')
control<HTMLButtonElement>(root, '添加错误终止规则').click()
await nextTick()
const statuses = control<HTMLInputElement>(root, '终止规则 1 状态码')
statuses.value = '400, 413'
statuses.dispatchEvent(new Event('input', { bubbles: true }))
await nextTick()
expect(policy.value.failover_rules.error_stop_patterns[0].status_codes).toEqual([400, 413])
expect(root.querySelector('[role="alert"]')).toBeNull()
control<HTMLButtonElement>(root, '删除成功转移规则 1').click()
await nextTick()
expect(policy.value.failover_rules.success_failover_patterns).toHaveLength(0)
})
it('edits and applies both rule groups through JSON mode', async () => {
const { root, policy } = mountEditor()
control<HTMLButtonElement>(root, '切到成功转移规则 JSON').click()
await nextTick()
const successJson = root.querySelector<HTMLTextAreaElement>('textarea')
if (!successJson) throw new Error('Missing success JSON editor')
successJson.value = '[{"pattern":"capacity"}]'
successJson.dispatchEvent(new Event('input', { bubbles: true }))
await nextTick()
control<HTMLButtonElement>(root, '切回成功转移规则表单').click()
await nextTick()
expect(policy.value.failover_rules.success_failover_patterns).toEqual([{ pattern: 'capacity', status_codes: [] }])
control<HTMLButtonElement>(root, '切到错误终止规则 JSON').click()
await nextTick()
const errorJson = root.querySelector<HTMLTextAreaElement>('textarea')
if (!errorJson) throw new Error('Missing error JSON editor')
errorJson.value = '[{"status_codes":[429,500],"pattern":"rate"}]'
errorJson.dispatchEvent(new Event('input', { bubbles: true }))
await nextTick()
control<HTMLButtonElement>(root, '切回错误终止规则表单').click()
await nextTick()
expect(policy.value.failover_rules.error_stop_patterns).toEqual([{ pattern: 'rate', status_codes: [429, 500] }])
})
it('keeps invalid JSON visible until it is corrected', async () => {
const { root, policy } = mountEditor()
control<HTMLButtonElement>(root, '切到错误终止规则 JSON').click()
await nextTick()
const editor = root.querySelector<HTMLTextAreaElement>('textarea')
if (!editor) throw new Error('Missing JSON editor')
editor.value = '{'
editor.dispatchEvent(new Event('input', { bubbles: true }))
await nextTick()
control<HTMLButtonElement>(root, '切回错误终止规则表单').click()
await nextTick()
expect(root.querySelector('[role="alert"]')?.textContent).toContain('JSON')
expect(policy.value.failover_rules.error_stop_patterns).toHaveLength(0)
})
})
@@ -0,0 +1,49 @@
import { describe, expect, it } from 'vitest'
import { normalizeRoutingFailoverPolicy, validateRoutingFailoverPolicy } from '../utils/routingFailover'
import { createEmptyRoutingGroupConfig, getModelScheduling, normalizeRoutingGroupConfig, upsertModelSchedulingRule } from '../utils/routingPolicy'
describe('routing failover policy', () => {
it('keeps legacy strategies unlimited with empty global rules', () => {
const policy = normalizeRoutingGroupConfig({}).default_policy
expect(policy.max_transfer_count).toBe(0)
expect(policy.max_transfer_timeout_seconds).toBe(0)
expect(policy.failover_rules).toEqual({ success_failover_patterns: [], error_stop_patterns: [] })
expect(validateRoutingFailoverPolicy(policy)).toBeNull()
})
it('preserves global limits and rules across model edits without sharing mutable arrays', () => {
const config = createEmptyRoutingGroupConfig()
Object.assign(config.default_policy, {
max_transfer_count: 3,
max_transfer_timeout_seconds: 90,
failover_rules: {
success_failover_patterns: [{ pattern: '(?i)capacity', status_codes: [] }],
error_stop_patterns: [{ pattern: '', status_codes: [400, 413] }],
},
})
const updated = upsertModelSchedulingRule(config, 'model-a', { priority_mode: 'provider', scheduling_mode: 'fixed_order' })
const policy = getModelScheduling(updated, 'model-a')
expect(policy.max_transfer_count).toBe(3)
expect(policy.max_transfer_timeout_seconds).toBe(90)
expect(policy.failover_rules).toEqual(config.default_policy.failover_rules)
expect(validateRoutingFailoverPolicy(policy)).toBeNull()
policy.failover_rules.error_stop_patterns[0].status_codes.push(422)
expect(config.default_policy.failover_rules.error_stop_patterns[0].status_codes).toEqual([400, 413])
})
it('rejects invalid budgets and ambiguous empty rules before saving', () => {
const policy = normalizeRoutingFailoverPolicy()
policy.max_transfer_count = -1
expect(validateRoutingFailoverPolicy(policy)).toContain('非负整数')
policy.max_transfer_count = 0
policy.max_transfer_timeout_seconds = 0.5
expect(validateRoutingFailoverPolicy(policy)).toContain('非负整数')
policy.max_transfer_timeout_seconds = 0
policy.failover_rules.error_stop_patterns.push({ pattern: '', status_codes: [] })
expect(validateRoutingFailoverPolicy(policy)).toContain('状态码或正则')
policy.failover_rules.error_stop_patterns[0].status_codes = [200]
expect(validateRoutingFailoverPolicy(policy)).toContain('400–599')
policy.failover_rules.error_stop_patterns[0].status_codes = [400]
expect(validateRoutingFailoverPolicy(policy)).toBeNull()
})
})
@@ -24,6 +24,20 @@ describe('routingPolicy', () => {
expect(config.default_policy.priority_mode).toBe('provider')
expect(config.default_policy.scheduling_mode).toBe('cache_affinity')
expect(config.default_policy.cancel_on_client_disconnect).toBe(false)
})
it('preserves cancellation policy across model scheduling edits', () => {
const config = createEmptyRoutingGroupConfig()
config.default_policy.cancel_on_client_disconnect = true
const updated = upsertModelSchedulingRule(config, 'gpt-5', {
priority_mode: 'global_key',
scheduling_mode: 'fixed_order',
})
expect(normalizeRoutingGroupConfig(updated).default_policy.cancel_on_client_disconnect).toBe(true)
expect(getModelScheduling(updated, 'gpt-5').cancel_on_client_disconnect).toBe(true)
expect(getModelScheduling(updated, 'other-model').cancel_on_client_disconnect).toBe(true)
expect(createEmptyRoutingGroupConfig().default_policy.cancel_on_client_disconnect).toBe(false)
})
it('drops the legacy group model allowlist while normalizing config', () => {
@@ -77,7 +91,7 @@ describe('routingPolicy', () => {
expect(createEmptyRoutingGroupConfig().default_policy.sticky_key_attempts).toBe(2)
expect(normalizeRoutingGroupConfig({}).default_policy.sticky_key_attempts).toBe(2)
expect(normalizeRoutingGroupConfig({
default_policy: { priority_mode: 'provider', scheduling_mode: 'cache_affinity', keep_priority_on_conversion: false, sticky_key_attempts: 3, enable_cf_heartbeat: false, cyber_continue_failover: false },
default_policy: { ...createEmptyRoutingGroupConfig().default_policy, priority_mode: 'provider', scheduling_mode: 'cache_affinity', keep_priority_on_conversion: false, sticky_key_attempts: 3, enable_cf_heartbeat: false, cyber_continue_failover: false, cancel_on_client_disconnect: false },
}).default_policy.sticky_key_attempts).toBe(3)
expect(normalizeStickyKeyAttempts('5')).toBe(5)
expect(normalizeStickyKeyAttempts(-1)).toBe(2)

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