mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 17:07:46 +08:00
Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7b8048c6ae | ||
|
|
ec95f2ca1f | ||
|
|
aa7dbe67d3 | ||
|
|
a26680f460 | ||
|
|
522b979052 | ||
|
|
808946312a | ||
|
|
741107bf71 | ||
|
|
6962731220 | ||
|
|
062e111c03 | ||
|
|
470c59e197 | ||
|
|
2f929e74c7 | ||
|
|
fc0417ceb9 | ||
|
|
44174a31e0 | ||
|
|
b599fb7354 | ||
|
|
14f96c9fa0 |
+7
-2
@@ -111,8 +111,13 @@ ADMIN_USERNAME=admin123456
|
||||
# AETHER_BARK_ALLOW_HTTP=false
|
||||
# AETHER_BARK_ALLOW_PRIVATE_TARGETS=false
|
||||
|
||||
# 可选 Provider OAuth 客户端。使用 Gemini CLI / Antigravity 浏览器授权时必须配置
|
||||
# 对应的 client secret;client ID 未配置时使用内置的公开 native-app client ID。
|
||||
# 普通 Provider 反代(包括 Provider OAuth)不按 DNS 地址过滤上游,兼容任意
|
||||
# Fake-IP 域名及内网 DNS。仅信任管理员配置的上游;没有严格 DNS 过滤开关。
|
||||
# URL 协议、字面 IP、TLS 证书,以及隧道中继和登录 OAuth 的校验仍保留。
|
||||
|
||||
# 可选 Provider OAuth 客户端。Gemini CLI 授权及刷新必须配置 client secret。
|
||||
# Antigravity 默认使用内置 native-app 客户端凭据;自定义 client ID 时必须同时配置
|
||||
# 对应的 client secret。未配置 client ID 时使用内置的公开 native-app client ID。
|
||||
# AETHER_GEMINI_CLI_OAUTH_CLIENT_ID=
|
||||
# AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET=
|
||||
# AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID=
|
||||
|
||||
@@ -13,6 +13,10 @@
|
||||
.plans
|
||||
.playwright-mcp/
|
||||
|
||||
docs/architecture
|
||||
!docs/architecture/architecture-dark.svg
|
||||
!docs/architecture/architecture-light.svg
|
||||
|
||||
### Python ###
|
||||
*.db
|
||||
*.db-*
|
||||
|
||||
Generated
+1
-1
@@ -661,7 +661,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "aether-tunnel"
|
||||
version = "0.3.16"
|
||||
version = "0.3.17"
|
||||
dependencies = [
|
||||
"aether-contracts",
|
||||
"aether-gateway",
|
||||
|
||||
@@ -17,13 +17,13 @@ fn sanitize_request_candidate_rows(
|
||||
mut candidates: Vec<StoredRequestCandidate>,
|
||||
) -> Vec<StoredRequestCandidate> {
|
||||
for candidate in &mut candidates {
|
||||
candidate.sanitize_sensitive_diagnostics();
|
||||
candidate.sanitize_for_persistence();
|
||||
}
|
||||
candidates
|
||||
}
|
||||
|
||||
fn sanitize_request_candidate_row(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
|
||||
candidate.sanitize_sensitive_diagnostics();
|
||||
candidate.sanitize_for_persistence();
|
||||
candidate
|
||||
}
|
||||
|
||||
@@ -1048,7 +1048,10 @@ mod request_candidate_security_tests {
|
||||
fn assert_candidate_is_sanitized(candidate: &StoredRequestCandidate) {
|
||||
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
||||
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
||||
assert!(candidate.error_message.is_none());
|
||||
assert_eq!(
|
||||
candidate.error_message.as_deref(),
|
||||
Some("Bearer candidate-secret")
|
||||
);
|
||||
assert_eq!(
|
||||
candidate.extra_data,
|
||||
Some(json!({"gateway_execution_runtime": true}))
|
||||
@@ -1057,13 +1060,15 @@ mod request_candidate_security_tests {
|
||||
candidate.required_capabilities,
|
||||
Some(json!({"vision": true}))
|
||||
);
|
||||
assert!(!serde_json::to_string(candidate)
|
||||
let mut public_candidate = candidate.clone();
|
||||
public_candidate.sanitize_sensitive_diagnostics();
|
||||
assert!(!serde_json::to_string(&public_candidate)
|
||||
.expect("candidate should serialize")
|
||||
.contains("candidate-secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gateway_candidate_boundary_sanitizes_repository_rows_and_write_results() {
|
||||
fn gateway_candidate_boundary_preserves_admin_errors_and_removes_request_payloads() {
|
||||
let candidate = sanitize_request_candidate_row(untrusted_candidate());
|
||||
assert_candidate_is_sanitized(&candidate);
|
||||
|
||||
|
||||
@@ -14682,9 +14682,10 @@ mod tests {
|
||||
candidate_extra["upstream_response"]["status_code"],
|
||||
json!(302)
|
||||
);
|
||||
assert!(candidate_extra["upstream_response"]
|
||||
.get("headers")
|
||||
.is_none());
|
||||
assert_eq!(
|
||||
candidate_extra["upstream_response"]["headers"]["location"],
|
||||
"/"
|
||||
);
|
||||
assert!(candidate_extra["upstream_response"].get("body").is_none());
|
||||
assert!(candidate_extra.get("client_response").is_none());
|
||||
|
||||
|
||||
@@ -22,8 +22,8 @@ use aether_contracts::{
|
||||
};
|
||||
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
|
||||
use aether_http::{
|
||||
apply_http_client_config, is_https_or_loopback_http_url, is_ipv4_benchmarking_fake_ip,
|
||||
is_private_or_reserved_ip, HttpClientConfig,
|
||||
apply_http_client_config, is_https_or_loopback_http_url, is_private_or_reserved_ip,
|
||||
HttpClientConfig,
|
||||
};
|
||||
use aether_runtime::{MetricKind, MetricSample};
|
||||
use axum::body::Bytes;
|
||||
@@ -63,8 +63,6 @@ use crate::upstream_admission::UpstreamTargetAdmissionPermit;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const HUB_RELAY_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope";
|
||||
pub(crate) const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY: &str =
|
||||
aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY;
|
||||
const HUB_RELAY_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
||||
const MAX_SAFE_REDIRECTS: usize = 10;
|
||||
const MAX_UPSTREAM_ERROR_DETAIL_BYTES: usize = 2_048;
|
||||
@@ -442,209 +440,12 @@ struct DirectHyperH2cSenderCacheMetrics {
|
||||
static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock<DirectHyperH2cSenderCacheMetrics> =
|
||||
LazyLock::new(DirectHyperH2cSenderCacheMetrics::default);
|
||||
|
||||
/// DNS resolver used for direct provider connections.
|
||||
///
|
||||
/// Provider endpoint URLs are frequently user/configuration supplied. The
|
||||
/// platform resolver may return a different answer on every lookup, so merely
|
||||
/// checking a URL's host (or resolving it once before constructing a client)
|
||||
/// is not sufficient to prevent DNS rebinding. This resolver validates every
|
||||
/// answer at the point reqwest/wreq asks for it. Explicit loopback targets are
|
||||
/// retained for the supported local-provider workflow, but a hostname that is
|
||||
/// not itself `localhost` can never resolve to a loopback/private address.
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct ExecutionSafeDnsResolver;
|
||||
|
||||
/// Resolver adapter for the legacy Hyper client retained for compatibility
|
||||
/// with the non-fast-path H2C cache. Keep this path subject to the same
|
||||
/// private-address and rebinding checks as reqwest/wreq clients.
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct ExecutionSafeHyperDnsResolver;
|
||||
|
||||
// Local DNS interception tools may use RFC 2544's 198.18.0.0/15 range for
|
||||
// synthetic answers. This exception is deliberately an allowlist rather
|
||||
// than a property of the address range itself: a custom provider hostname
|
||||
// must not be able to turn a local synthetic mapping into an SSRF primitive.
|
||||
// Keep this list limited to origins that Aether constructs as built-in
|
||||
// provider/model-fetch targets. In particular, do not use a
|
||||
// suffix match for ordinary hosts (for example, `evil.chatgpt.com`).
|
||||
const TRUSTED_EXECUTION_BENCHMARKING_DNS_EXACT_HOSTS: &[&str] = &[
|
||||
"aiplatform.googleapis.com",
|
||||
"antigravity.googleapis.com",
|
||||
"api.openai.com",
|
||||
"api.anthropic.com",
|
||||
"api.deepseek.com",
|
||||
"chatgpt.com",
|
||||
"cloudcode-pa.googleapis.com",
|
||||
"daily-cloudcode-pa.googleapis.com",
|
||||
"daily-cloudcode-pa.sandbox.googleapis.com",
|
||||
"dashscope.aliyuncs.com",
|
||||
"generativelanguage.googleapis.com",
|
||||
"grok.com",
|
||||
"open.bigmodel.cn",
|
||||
"q.us-iso-east-1.c2s.ic.gov",
|
||||
"q.us-isob-east-1.sc2s.sgov.gov",
|
||||
"q.us-isof-east-1.csp.hci.ic.gov",
|
||||
"q.us-isof-south-1.csp.hci.ic.gov",
|
||||
"server.codeium.com",
|
||||
];
|
||||
|
||||
const TRUSTED_EXECUTION_VERTEX_DNS_REGIONS: &[&str] = &[
|
||||
"africa-south1",
|
||||
"asia-east1",
|
||||
"asia-east2",
|
||||
"asia-northeast1",
|
||||
"asia-northeast2",
|
||||
"asia-northeast3",
|
||||
"asia-south1",
|
||||
"asia-south2",
|
||||
"asia-southeast1",
|
||||
"asia-southeast2",
|
||||
"australia-southeast1",
|
||||
"australia-southeast2",
|
||||
"europe-central2",
|
||||
"europe-north1",
|
||||
"europe-southwest1",
|
||||
"europe-west1",
|
||||
"europe-west2",
|
||||
"europe-west3",
|
||||
"europe-west4",
|
||||
"europe-west6",
|
||||
"europe-west8",
|
||||
"europe-west9",
|
||||
"europe-west10",
|
||||
"europe-west12",
|
||||
"me-central1",
|
||||
"me-central2",
|
||||
"me-west1",
|
||||
"northamerica-northeast1",
|
||||
"northamerica-northeast2",
|
||||
"southamerica-east1",
|
||||
"southamerica-west1",
|
||||
"us-central1",
|
||||
"us-east1",
|
||||
"us-east4",
|
||||
"us-east5",
|
||||
"us-south1",
|
||||
"us-west1",
|
||||
"us-west2",
|
||||
"us-west3",
|
||||
"us-west4",
|
||||
];
|
||||
|
||||
const TRUSTED_EXECUTION_AWS_DNS_REGIONS: &[&str] = &[
|
||||
"af-south-1",
|
||||
"ap-east-1",
|
||||
"ap-northeast-1",
|
||||
"ap-northeast-2",
|
||||
"ap-northeast-3",
|
||||
"ap-south-1",
|
||||
"ap-south-2",
|
||||
"ap-southeast-1",
|
||||
"ap-southeast-2",
|
||||
"ap-southeast-3",
|
||||
"ap-southeast-4",
|
||||
"ca-central-1",
|
||||
"ca-west-1",
|
||||
"eu-central-1",
|
||||
"eu-central-2",
|
||||
"eu-north-1",
|
||||
"eu-south-1",
|
||||
"eu-south-2",
|
||||
"eu-west-1",
|
||||
"eu-west-2",
|
||||
"eu-west-3",
|
||||
"il-central-1",
|
||||
"me-central-1",
|
||||
"me-south-1",
|
||||
"mx-central-1",
|
||||
"sa-east-1",
|
||||
"us-east-1",
|
||||
"us-east-2",
|
||||
"us-gov-east-1",
|
||||
"us-gov-west-1",
|
||||
"us-west-1",
|
||||
"us-west-2",
|
||||
];
|
||||
|
||||
static EXECUTION_EXTRA_TRUSTED_DNS_HOSTS: LazyLock<StdRwLock<BTreeSet<String>>> =
|
||||
LazyLock::new(|| StdRwLock::new(BTreeSet::new()));
|
||||
|
||||
pub(crate) fn refresh_execution_extra_trusted_dns_hosts(value: Option<&Value>) {
|
||||
let hosts = value
|
||||
.cloned()
|
||||
.and_then(|value| {
|
||||
aether_admin::system::normalize_execution_extra_trusted_dns_hosts_config_value(value)
|
||||
.ok()
|
||||
})
|
||||
.and_then(|value| {
|
||||
value.as_array().map(|hosts| {
|
||||
hosts
|
||||
.iter()
|
||||
.filter_map(Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<BTreeSet<_>>()
|
||||
})
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
if let Ok(mut current) = EXECUTION_EXTRA_TRUSTED_DNS_HOSTS.write() {
|
||||
*current = hosts;
|
||||
}
|
||||
}
|
||||
|
||||
/// Return whether `host` is one of the fixed provider origins for which a
|
||||
/// local RFC-2544 synthetic answer can be accepted. The resolver receives only
|
||||
/// a hostname (not the URL scheme/path), so all policy that can be expressed
|
||||
/// here is intentionally host based. URL validation still requires HTTPS for
|
||||
/// non-loopback upstreams before this resolver is used.
|
||||
fn execution_host_allows_benchmarking_dns_answer(host: &str) -> bool {
|
||||
let extra_hosts = EXECUTION_EXTRA_TRUSTED_DNS_HOSTS
|
||||
.read()
|
||||
.map(|hosts| hosts.clone())
|
||||
.unwrap_or_default();
|
||||
execution_host_allows_benchmarking_dns_answer_with_extra_hosts(host, &extra_hosts)
|
||||
}
|
||||
|
||||
fn execution_host_allows_benchmarking_dns_answer_with_extra_hosts(
|
||||
host: &str,
|
||||
extra_hosts: &BTreeSet<String>,
|
||||
) -> bool {
|
||||
let host = host.trim().trim_end_matches('.').to_ascii_lowercase();
|
||||
if extra_hosts.contains(&host)
|
||||
|| TRUSTED_EXECUTION_BENCHMARKING_DNS_EXACT_HOSTS
|
||||
.iter()
|
||||
.any(|trusted| *trusted == host)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
|
||||
// Vertex service-account requests use `<region>-aiplatform.googleapis.com`.
|
||||
// Keep this compatibility exception limited to known provider regions.
|
||||
if let Some(region) = host.strip_suffix("-aiplatform.googleapis.com") {
|
||||
return TRUSTED_EXECUTION_VERTEX_DNS_REGIONS.contains(®ion);
|
||||
}
|
||||
|
||||
// Kiro uses a small, fixed set of regional service origins. Match each
|
||||
// supported AWS partition explicitly; never use a broad suffix check that
|
||||
// could accept an attacker-controlled subdomain.
|
||||
matches_regional_service_host(&host, "q", ".amazonaws.com")
|
||||
|| matches_regional_service_host(&host, "q-fips", ".amazonaws.com")
|
||||
|| matches_regional_service_host(&host, "codewhisperer", ".amazonaws.com")
|
||||
|| matches_regional_service_host(&host, "oidc", ".amazonaws.com")
|
||||
|| matches_regional_service_host(&host, "prod", ".auth.desktop.kiro.dev")
|
||||
}
|
||||
|
||||
fn matches_regional_service_host(host: &str, service: &str, suffix: &str) -> bool {
|
||||
let Some(region) = host
|
||||
.strip_prefix(service)
|
||||
.and_then(|value| value.strip_prefix('.'))
|
||||
.and_then(|value| value.strip_suffix(suffix))
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
TRUSTED_EXECUTION_AWS_DNS_REGIONS.contains(®ion)
|
||||
}
|
||||
|
||||
fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
|
||||
let host = host.trim_end_matches('.');
|
||||
host.eq_ignore_ascii_case("localhost")
|
||||
@@ -654,17 +455,10 @@ fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn validate_execution_dns_answers(
|
||||
fn validate_resolved_execution_addresses(
|
||||
host: &str,
|
||||
addresses: Vec<SocketAddr>,
|
||||
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
||||
validate_execution_dns_answers_with_policy(host, addresses, true)
|
||||
}
|
||||
|
||||
fn validate_execution_dns_answers_with_policy(
|
||||
host: &str,
|
||||
addresses: Vec<SocketAddr>,
|
||||
allow_trusted_benchmarking_dns_answer: bool,
|
||||
provider_execution: bool,
|
||||
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
||||
if addresses.is_empty() {
|
||||
return Err(std::io::Error::new(
|
||||
@@ -672,25 +466,22 @@ fn validate_execution_dns_answers_with_policy(
|
||||
"upstream DNS resolution returned no addresses",
|
||||
));
|
||||
}
|
||||
|
||||
if provider_execution {
|
||||
return Ok(addresses);
|
||||
}
|
||||
let allows_loopback = dns_host_explicitly_allows_loopback(host);
|
||||
let allows_benchmarking_dns_answer = allow_trusted_benchmarking_dns_answer
|
||||
&& execution_host_allows_benchmarking_dns_answer(host);
|
||||
let unsafe_answer = addresses.iter().any(|address| {
|
||||
if addresses.iter().any(|address| {
|
||||
if allows_loopback {
|
||||
!address.ip().is_loopback()
|
||||
} else {
|
||||
is_private_or_reserved_ip(address.ip())
|
||||
&& !(allows_benchmarking_dns_answer && is_ipv4_benchmarking_fake_ip(address.ip()))
|
||||
}
|
||||
});
|
||||
if unsafe_answer {
|
||||
}) {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::PermissionDenied,
|
||||
"upstream DNS resolution returned a private or reserved address",
|
||||
"tunnel relay DNS resolution returned a private or reserved address",
|
||||
));
|
||||
}
|
||||
|
||||
Ok(addresses)
|
||||
}
|
||||
|
||||
@@ -701,7 +492,7 @@ async fn resolve_execution_dns_addresses(host: &str) -> Result<Vec<SocketAddr>,
|
||||
async fn resolve_execution_target_addresses_with_policy(
|
||||
host: &str,
|
||||
port: u16,
|
||||
allow_trusted_benchmarking_dns_answer: bool,
|
||||
provider_execution: bool,
|
||||
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
||||
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
vec![SocketAddr::new(ip, port)]
|
||||
@@ -709,11 +500,7 @@ async fn resolve_execution_target_addresses_with_policy(
|
||||
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
|
||||
.await?
|
||||
};
|
||||
validate_execution_dns_answers_with_policy(
|
||||
host,
|
||||
addresses,
|
||||
allow_trusted_benchmarking_dns_answer,
|
||||
)
|
||||
validate_resolved_execution_addresses(host, addresses, provider_execution)
|
||||
}
|
||||
|
||||
impl reqwest::dns::Resolve for ExecutionSafeDnsResolver {
|
||||
@@ -3439,10 +3226,6 @@ async fn resolve_relay_target_addresses(
|
||||
let port = url.port_or_known_default().ok_or_else(|| {
|
||||
ExecutionRuntimeTransportError::RelayError("tunnel relay URL has no port".to_string())
|
||||
})?;
|
||||
// Relay destinations remain strict even when their hostname happens to be
|
||||
// an official provider origin. The RFC-2544 compatibility exception is
|
||||
// only for direct provider execution; allowing it here would weaken the
|
||||
// relay SSRF guard.
|
||||
let addresses = resolve_execution_target_addresses_with_policy(host, port, false)
|
||||
.await
|
||||
.map_err(|error| match error.kind() {
|
||||
@@ -5650,115 +5433,75 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn execution_dns_answers_reject_private_addresses_and_allow_explicit_loopback() {
|
||||
let public = "93.184.216.34:443".parse().unwrap();
|
||||
let private = "10.0.0.8:443".parse().unwrap();
|
||||
let loopback_v4 = "127.0.0.1:8080".parse().unwrap();
|
||||
let loopback_v6 = "[::1]:8080".parse().unwrap();
|
||||
|
||||
assert!(super::validate_execution_dns_answers("api.example.test", vec![public]).is_ok());
|
||||
assert!(super::validate_execution_dns_answers("api.example.test", vec![private]).is_err());
|
||||
assert!(
|
||||
super::validate_execution_dns_answers("localhost", vec![loopback_v4, loopback_v6])
|
||||
.is_ok()
|
||||
);
|
||||
assert!(super::validate_execution_dns_answers("localhost", vec![private]).is_err());
|
||||
assert!(super::validate_execution_dns_answers("api.example.test", Vec::new()).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn execution_dns_answers_allow_benchmarking_range_only_for_fixed_provider_hosts() {
|
||||
let fake = "198.18.75.234:443".parse().unwrap();
|
||||
fn execution_dns_answers_allow_all_provider_hosts_without_address_filtering() {
|
||||
let addresses = vec![
|
||||
"198.18.78.41:443".parse().unwrap(),
|
||||
"10.0.0.8:443".parse().unwrap(),
|
||||
"127.0.0.1:443".parse().unwrap(),
|
||||
"169.254.169.254:443".parse().unwrap(),
|
||||
"[fd00::1]:443".parse().unwrap(),
|
||||
"93.184.216.34:443".parse().unwrap(),
|
||||
];
|
||||
for host in [
|
||||
"api.openai.com",
|
||||
"CHATGPT.COM.",
|
||||
"us-central1-aiplatform.googleapis.com",
|
||||
"me-central2-aiplatform.googleapis.com",
|
||||
"q.us-east-1.amazonaws.com",
|
||||
"q-fips.us-gov-west-1.amazonaws.com",
|
||||
"codewhisperer.us-west-2.amazonaws.com",
|
||||
"oidc.us-east-1.amazonaws.com",
|
||||
"prod.us-east-1.auth.desktop.kiro.dev",
|
||||
"q.us-iso-east-1.c2s.ic.gov",
|
||||
"q.us-isob-east-1.sc2s.sgov.gov",
|
||||
"q.us-isof-east-1.csp.hci.ic.gov",
|
||||
"oauth2.googleapis.com",
|
||||
"www.googleapis.com",
|
||||
"custom.example.test",
|
||||
] {
|
||||
assert!(
|
||||
super::validate_execution_dns_answers(host, vec![fake]).is_ok(),
|
||||
"fixed provider host should accept a benchmarking DNS answer: {host}"
|
||||
);
|
||||
}
|
||||
|
||||
for host in [
|
||||
"api.example.test",
|
||||
"evil.chatgpt.com",
|
||||
"api.openai.com.evil.test",
|
||||
"q.us-east-1.evil.amazonaws.com",
|
||||
"q.us-east-1.amazonaws.com.attacker.test",
|
||||
"q.localhost.amazonaws.com",
|
||||
"evil-1-aiplatform.googleapis.com",
|
||||
"q.evil-1.amazonaws.com",
|
||||
"q-fips.evil-1.amazonaws.com",
|
||||
"codewhisperer.evil-1.amazonaws.com",
|
||||
"prod.evil-1.auth.desktop.kiro.dev",
|
||||
"oidc.evil-1.amazonaws.com",
|
||||
"q.us-central1.amazonaws.com",
|
||||
"us-east-1-aiplatform.googleapis.com",
|
||||
"q.us-east-1.c2s.ic.gov",
|
||||
"q.us-iso-east-1.sc2s.sgov.gov",
|
||||
"q-fips.us-gov-west-1.evil.amazonaws.com",
|
||||
"codewhisperer.us-west-2.evil.amazonaws.com",
|
||||
"oidc.us-east-1.evil.amazonaws.com",
|
||||
"prod.us-east-1.auth.desktop.kiro.dev.attacker.test",
|
||||
"prod.us-east-1.evil.auth.desktop.kiro.dev",
|
||||
"q.us-iso-east-1.evil.c2s.ic.gov",
|
||||
"q.us-iso-east-1.c2s.ic.gov.attacker.test",
|
||||
"q.us-iso-east-1.c2s.ic.gov.evil",
|
||||
"198.18.75.234",
|
||||
] {
|
||||
assert!(
|
||||
super::validate_execution_dns_answers(host, vec![fake]).is_err(),
|
||||
"untrusted or lookalike host must reject a benchmarking DNS answer: {host}"
|
||||
assert_eq!(
|
||||
super::validate_resolved_execution_addresses(host, addresses.clone(), true)
|
||||
.expect("provider DNS answers should pass through"),
|
||||
addresses
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn execution_dns_answers_allow_benchmarking_range_for_configured_exact_hosts() {
|
||||
let fake = "198.18.75.234:443".parse().unwrap();
|
||||
super::refresh_execution_extra_trusted_dns_hosts(Some(&json!(["custom.example.com",])));
|
||||
|
||||
assert!(super::validate_execution_dns_answers("custom.example.com", vec![fake]).is_ok());
|
||||
assert!(
|
||||
super::validate_execution_dns_answers("api.custom.example.com", vec![fake]).is_err()
|
||||
);
|
||||
|
||||
super::refresh_execution_extra_trusted_dns_hosts(None);
|
||||
assert!(super::validate_execution_dns_answers("custom.example.com", vec![fake]).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn execution_dns_answers_reject_mixed_private_results_and_strict_relay_policy() {
|
||||
let fake = "198.18.75.234:443".parse().unwrap();
|
||||
fn execution_dns_answers_keep_relay_address_filtering() {
|
||||
let public = "93.184.216.34:443".parse().unwrap();
|
||||
let private = "10.0.0.8:443".parse().unwrap();
|
||||
|
||||
// A trusted host may have a synthetic answer alongside a genuine public
|
||||
// answer, but any real private answer still fails closed.
|
||||
for host in ["oauth2.googleapis.com", "custom.example.test"] {
|
||||
assert!(
|
||||
super::validate_resolved_execution_addresses(host, vec![public], false).is_ok()
|
||||
);
|
||||
for blocked in [
|
||||
"198.18.78.41:443",
|
||||
"10.0.0.8:443",
|
||||
"127.0.0.1:443",
|
||||
"169.254.169.254:443",
|
||||
"[fd00::1]:443",
|
||||
] {
|
||||
let blocked = blocked.parse().unwrap();
|
||||
assert!(
|
||||
super::validate_resolved_execution_addresses(host, vec![blocked], false)
|
||||
.is_err()
|
||||
);
|
||||
assert!(super::validate_resolved_execution_addresses(
|
||||
host,
|
||||
vec![public, blocked],
|
||||
false
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
let loopback = vec![
|
||||
"127.0.0.1:443".parse().unwrap(),
|
||||
"[::1]:443".parse().unwrap(),
|
||||
];
|
||||
assert!(super::validate_resolved_execution_addresses("localhost", loopback, false).is_ok());
|
||||
assert!(
|
||||
super::validate_execution_dns_answers("api.openai.com", vec![fake, public]).is_ok()
|
||||
super::validate_resolved_execution_addresses("localhost", vec![public], false).is_err()
|
||||
);
|
||||
assert!(
|
||||
super::validate_execution_dns_answers("api.openai.com", vec![fake, private]).is_err()
|
||||
);
|
||||
|
||||
// Tunnel relay resolution opts out of the compatibility exception.
|
||||
assert!(super::validate_execution_dns_answers_with_policy(
|
||||
"api.openai.com",
|
||||
vec![fake],
|
||||
false,
|
||||
)
|
||||
.is_err());
|
||||
for provider_execution in [false, true] {
|
||||
assert_eq!(
|
||||
super::validate_resolved_execution_addresses(
|
||||
"custom.example.test",
|
||||
Vec::new(),
|
||||
provider_execution
|
||||
)
|
||||
.expect_err("empty DNS answers must fail")
|
||||
.kind(),
|
||||
std::io::ErrorKind::NotFound
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -659,7 +659,7 @@ async fn admin_monitoring_trace_request_exposes_request_path_from_usage_audit()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_redacts_failed_candidate_response_payloads() {
|
||||
async fn admin_monitoring_trace_request_exposes_failed_candidate_response_payloads() {
|
||||
let mut candidate = sample_candidate(
|
||||
"cand-used",
|
||||
"request-1",
|
||||
@@ -746,14 +746,96 @@ async fn admin_monitoring_trace_request_redacts_failed_candidate_response_payloa
|
||||
json!("upstream_response")
|
||||
);
|
||||
assert!(extra["upstream_response"].get("headers").is_none());
|
||||
assert!(extra["upstream_response"].get("body").is_none());
|
||||
assert_eq!(
|
||||
extra["upstream_response"]["body"]["error"]["message"],
|
||||
"redirect blocked"
|
||||
);
|
||||
assert!(extra["upstream_response"].get("body_ref").is_none());
|
||||
assert!(extra.get("client_response").is_none());
|
||||
assert!(extra.get("provider_response").is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_response_body() {
|
||||
async fn admin_monitoring_trace_request_does_not_replace_attempt_status_with_usage_status() {
|
||||
for (candidate_status, upstream_status, expected_status) in [
|
||||
(Some(400), Some(400), Some(400)),
|
||||
(Some(400), None, Some(400)),
|
||||
(Some(502), Some(200), Some(200)),
|
||||
(None, None, None),
|
||||
] {
|
||||
let mut candidate = sample_candidate(
|
||||
"cand-used",
|
||||
"request-failover-status",
|
||||
0,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(33),
|
||||
candidate_status,
|
||||
);
|
||||
if let Some(status_code) = upstream_status {
|
||||
candidate.extra_data = Some(json!({
|
||||
"upstream_response": {
|
||||
"status_code": status_code,
|
||||
"body": {"error": {"message": "sensitive upstream error"}}
|
||||
}
|
||||
}));
|
||||
}
|
||||
let request_candidates =
|
||||
Arc::new(InMemoryRequestCandidateRepository::seed(vec![candidate]));
|
||||
let mut usage = sample_usage(
|
||||
"request-failover-status",
|
||||
"provider-1",
|
||||
"OpenAI",
|
||||
0,
|
||||
0.0,
|
||||
"failed",
|
||||
Some(503),
|
||||
100,
|
||||
);
|
||||
usage.candidate_id = Some("cand-used".to_string());
|
||||
usage.response_body_state = Some(UsageBodyCaptureState::Reference);
|
||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
|
||||
request_candidates,
|
||||
usage_repository,
|
||||
);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let context = request_context(
|
||||
http::Method::GET,
|
||||
"/api/admin/monitoring/trace/request-failover-status",
|
||||
);
|
||||
|
||||
let response = local_monitoring_response(&state, &context)
|
||||
.await
|
||||
.expect("handler should not error")
|
||||
.expect("route should be handled locally");
|
||||
assert_eq!(response.status(), http::StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
let payload: serde_json::Value =
|
||||
serde_json::from_slice(&body).expect("json body should parse");
|
||||
let candidate = &payload["candidates"][0];
|
||||
assert_eq!(candidate["status_code"], json!(candidate_status));
|
||||
assert_eq!(
|
||||
candidate["extra_data"]["upstream_response"]["status_code"],
|
||||
json!(expected_status),
|
||||
);
|
||||
if upstream_status.is_some() {
|
||||
assert_eq!(
|
||||
candidate["extra_data"]["upstream_response"]["body"]["error"]["message"],
|
||||
"sensitive upstream error"
|
||||
);
|
||||
}
|
||||
assert!(candidate["error_message"].is_null());
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_keeps_candidate_errors_without_hydrating_usage_bodies() {
|
||||
let mut candidate = sample_candidate(
|
||||
"cand-used",
|
||||
"request-ref-body",
|
||||
@@ -771,8 +853,7 @@ async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_respon
|
||||
"x-request-id": "stale-request-like-body"
|
||||
},
|
||||
"body": {
|
||||
"model": "gpt-5.6-sol",
|
||||
"input": [{"role": "user", "content": "request prompt"}]
|
||||
"error": {"message": "candidate-specific upstream failure"}
|
||||
}
|
||||
}
|
||||
}));
|
||||
@@ -837,8 +918,14 @@ async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_respon
|
||||
assert_eq!(upstream_response["status_code"], json!(400));
|
||||
assert_eq!(upstream_response["source"], json!("upstream_response"));
|
||||
assert_eq!(upstream_response["body_state"], json!("reference"));
|
||||
assert!(upstream_response.get("headers").is_none());
|
||||
assert!(upstream_response.get("body").is_none());
|
||||
assert_eq!(
|
||||
upstream_response["headers"]["content-type"],
|
||||
"text/event-stream"
|
||||
);
|
||||
assert_eq!(
|
||||
upstream_response["body"]["error"]["message"],
|
||||
"candidate-specific upstream failure"
|
||||
);
|
||||
assert!(upstream_response.get("body_ref").is_none());
|
||||
}
|
||||
|
||||
|
||||
@@ -1940,12 +1940,16 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn user_usage_active_override_uses_terminal_candidate_latency() {
|
||||
let candidate = sample_candidate(
|
||||
let mut candidate = sample_candidate(
|
||||
RequestCandidateStatus::Success,
|
||||
Some(200),
|
||||
Some(9_210),
|
||||
None,
|
||||
);
|
||||
candidate.error_message = Some("private upstream diagnostic".to_string());
|
||||
candidate.extra_data = Some(json!({
|
||||
"upstream_response": {"body": {"error": {"message": "private upstream diagnostic"}}}
|
||||
}));
|
||||
|
||||
let payload =
|
||||
users_me_usage_terminal_candidate_state_override(&[candidate]).expect("override");
|
||||
@@ -1953,6 +1957,7 @@ mod tests {
|
||||
assert_eq!(payload["status"], "completed");
|
||||
assert_eq!(payload["response_time_ms"], 9_210);
|
||||
assert_eq!(payload["status_code"], 200);
|
||||
assert!(!payload.to_string().contains("private upstream diagnostic"));
|
||||
assert_eq!(
|
||||
payload["response_time_updated_at"],
|
||||
"1970-01-01T00:00:10.210+00:00"
|
||||
|
||||
@@ -120,8 +120,16 @@ pub(crate) async fn decrypt_or_migrate_smtp_password(
|
||||
}
|
||||
let plaintext = decrypt_system_config_secret(state, "smtp_password", stored.trim())
|
||||
.or_else(|| {
|
||||
(!stored.trim().is_empty() && !looks_like_python_fernet_ciphertext(stored.trim()))
|
||||
.then(|| stored.trim().to_string())
|
||||
if stored_secret_uses_known_envelope_family(stored.trim()) {
|
||||
return None;
|
||||
}
|
||||
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), stored.trim()).or_else(
|
||||
|| {
|
||||
(!stored.trim().is_empty()
|
||||
&& !looks_like_python_fernet_ciphertext(stored.trim()))
|
||||
.then(|| stored.trim().to_string())
|
||||
},
|
||||
)
|
||||
})
|
||||
.ok_or_else(|| system_config_secret_error("stored SMTP password cannot be decrypted"))?;
|
||||
if plaintext.contains('\0') {
|
||||
@@ -700,11 +708,13 @@ mod tests {
|
||||
use super::{
|
||||
bark_device_key_binding, decrypt_bark_device_key_v2, decrypt_ldap_bind_password_v2,
|
||||
decrypt_ldap_bind_password_v3, decrypt_or_migrate_bark_device_key,
|
||||
decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_system_config_secret,
|
||||
decrypt_or_migrate_ldap_bind_password, decrypt_or_migrate_smtp_password,
|
||||
decrypt_or_migrate_system_config_secret,
|
||||
decrypt_or_migrate_system_config_secret_with_before_compare, decrypt_system_config_secret,
|
||||
encrypt_bark_device_key, encrypt_ldap_bind_password, encrypt_system_config_secret,
|
||||
ldap_module_config_is_valid, normalize_ldap_transport_server_url,
|
||||
LDAP_BIND_PASSWORD_V2_PREFIX, LDAP_BIND_PASSWORD_V3_PREFIX, SYSTEM_CONFIG_SECRET_V2_PREFIX,
|
||||
encrypt_bark_device_key, encrypt_ldap_bind_password, encrypt_smtp_password,
|
||||
encrypt_system_config_secret, ldap_module_config_is_valid,
|
||||
normalize_ldap_transport_server_url, smtp_password_binding, LDAP_BIND_PASSWORD_V2_PREFIX,
|
||||
LDAP_BIND_PASSWORD_V3_PREFIX, SMTP_PASSWORD_V3_PREFIX, SYSTEM_CONFIG_SECRET_V2_PREFIX,
|
||||
};
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::AppState;
|
||||
@@ -740,6 +750,165 @@ mod tests {
|
||||
state
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn smtp_password_migrates_legacy_formats_to_bound_v3() {
|
||||
let binding = smtp_password_binding(
|
||||
"smtp.example.com",
|
||||
587,
|
||||
Some("[email protected]"),
|
||||
true,
|
||||
false,
|
||||
)
|
||||
.expect("SMTP binding should build");
|
||||
let fixture_state = state_with_stored_secret(TEST_SECRET);
|
||||
let legacy_values = [
|
||||
TEST_SECRET.to_string(),
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, TEST_SECRET)
|
||||
.expect("legacy SMTP password should encrypt"),
|
||||
encrypt_system_config_secret(&fixture_state, TEST_KEY, TEST_SECRET)
|
||||
.expect("v2 SMTP password should encrypt"),
|
||||
];
|
||||
|
||||
for legacy in legacy_values {
|
||||
let state = state_with_stored_secret(&legacy);
|
||||
let plaintext = decrypt_or_migrate_smtp_password(&state, &binding, legacy.clone())
|
||||
.await
|
||||
.expect("legacy SMTP password should migrate");
|
||||
assert_eq!(plaintext, TEST_SECRET);
|
||||
let migrated = state
|
||||
.read_system_config_json_value_strong(TEST_KEY)
|
||||
.await
|
||||
.expect("SMTP password should read")
|
||||
.and_then(|value| value.as_str().map(ToOwned::to_owned))
|
||||
.expect("SMTP password should be a string");
|
||||
assert!(migrated.starts_with(SMTP_PASSWORD_V3_PREFIX));
|
||||
assert_ne!(migrated, legacy);
|
||||
assert_eq!(
|
||||
decrypt_or_migrate_smtp_password(&state, &binding, migrated.clone())
|
||||
.await
|
||||
.expect("migrated SMTP password should decrypt"),
|
||||
TEST_SECRET
|
||||
);
|
||||
assert_eq!(
|
||||
state
|
||||
.read_system_config_json_value_strong(TEST_KEY)
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(json!(migrated))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn smtp_password_rejects_invalid_ciphertext_without_rewriting() {
|
||||
let binding = smtp_password_binding(
|
||||
"smtp.example.com",
|
||||
587,
|
||||
Some("[email protected]"),
|
||||
true,
|
||||
false,
|
||||
)
|
||||
.expect("SMTP binding should build");
|
||||
let fixture_state = state_with_stored_secret(TEST_SECRET);
|
||||
let mut tampered = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, TEST_SECRET)
|
||||
.expect("legacy SMTP password should encrypt");
|
||||
tampered.replace_range(tampered.len() - 2.., "AA");
|
||||
let invalid_values = [
|
||||
tampered,
|
||||
encrypt_python_fernet_plaintext("unavailable-historical-key", TEST_SECRET)
|
||||
.expect("wrong-key SMTP password should encrypt"),
|
||||
encrypt_system_config_secret(&fixture_state, "other_secret", TEST_SECRET)
|
||||
.expect("wrong-purpose secret should encrypt"),
|
||||
"aether-system-config-secret-v2:invalid".to_string(),
|
||||
"aether-smtp-password-v3:invalid".to_string(),
|
||||
"aether-runtime-secret-v1:invalid".to_string(),
|
||||
"aether-unknown-secret-v4:invalid".to_string(),
|
||||
];
|
||||
|
||||
for stored in invalid_values {
|
||||
let state = state_with_stored_secret(&stored);
|
||||
let error = decrypt_or_migrate_smtp_password(&state, &binding, stored.clone())
|
||||
.await
|
||||
.expect_err("invalid ciphertext must not become an SMTP password");
|
||||
assert_eq!(
|
||||
error.into_message(),
|
||||
"stored SMTP password cannot be decrypted"
|
||||
);
|
||||
assert_eq!(
|
||||
state
|
||||
.read_system_config_json_value_strong(TEST_KEY)
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(json!(stored))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn smtp_password_v3_rejects_changed_transport_binding() {
|
||||
let binding = smtp_password_binding(
|
||||
"smtp.example.com",
|
||||
587,
|
||||
Some("[email protected]"),
|
||||
true,
|
||||
false,
|
||||
)
|
||||
.expect("SMTP binding should build");
|
||||
let stored = encrypt_smtp_password(
|
||||
&state_with_stored_secret(TEST_SECRET),
|
||||
&binding,
|
||||
TEST_SECRET,
|
||||
)
|
||||
.expect("SMTP password should encrypt");
|
||||
let state = state_with_stored_secret(&stored);
|
||||
for changed_binding in [
|
||||
smtp_password_binding(
|
||||
"other.example.com",
|
||||
587,
|
||||
Some("[email protected]"),
|
||||
true,
|
||||
false,
|
||||
),
|
||||
smtp_password_binding(
|
||||
"smtp.example.com",
|
||||
465,
|
||||
Some("[email protected]"),
|
||||
true,
|
||||
false,
|
||||
),
|
||||
smtp_password_binding(
|
||||
"smtp.example.com",
|
||||
587,
|
||||
Some("[email protected]"),
|
||||
true,
|
||||
false,
|
||||
),
|
||||
smtp_password_binding(
|
||||
"smtp.example.com",
|
||||
587,
|
||||
Some("[email protected]"),
|
||||
false,
|
||||
false,
|
||||
),
|
||||
smtp_password_binding("smtp.example.com", 587, Some("[email protected]"), true, true),
|
||||
] {
|
||||
assert!(decrypt_or_migrate_smtp_password(
|
||||
&state,
|
||||
&changed_binding.expect("changed binding should build"),
|
||||
stored.clone(),
|
||||
)
|
||||
.await
|
||||
.is_err());
|
||||
}
|
||||
assert_eq!(
|
||||
state
|
||||
.read_system_config_json_value_strong(TEST_KEY)
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(json!(stored))
|
||||
);
|
||||
}
|
||||
|
||||
fn ldap_config(bind_password: &str) -> StoredLdapModuleConfig {
|
||||
StoredLdapModuleConfig {
|
||||
server_url: "ldaps://ldap.example.com".to_string(),
|
||||
|
||||
@@ -238,9 +238,11 @@ async fn read_notification_channel_readiness(
|
||||
state: &AppState,
|
||||
config: &ImportantNotificationConfig,
|
||||
) -> Result<NotificationChannelReadiness, GatewayError> {
|
||||
let smtp_config = read_smtp_delivery_config(state).await?;
|
||||
let email = config.email_enabled
|
||||
&& !config.email_recipients.is_empty()
|
||||
&& matches!(read_smtp_delivery_config(state).await, Ok(Some(_)));
|
||||
Ok(NotificationChannelReadiness {
|
||||
email: config.email_enabled && !config.email_recipients.is_empty() && smtp_config.is_some(),
|
||||
email,
|
||||
server_chan: config.server_chan.enabled && config.server_chan.send_key.is_some(),
|
||||
bark: config.bark.enabled && config.bark.device_key.is_some(),
|
||||
})
|
||||
@@ -840,13 +842,50 @@ fn escape_html(value: &str) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
apply_notification_item_template, parse_channel_filter, parse_notification_items,
|
||||
parse_recipient_list, ImportantNotification, ImportantNotificationChannelFilter,
|
||||
MAX_NOTIFICATION_ITEMS, MAX_NOTIFICATION_RECIPIENTS, MAX_NOTIFICATION_RECIPIENT_BYTES,
|
||||
apply_notification_item_template, important_notification_configured, parse_channel_filter,
|
||||
parse_notification_items, parse_recipient_list, ImportantNotification,
|
||||
ImportantNotificationChannelFilter, IMPORTANT_NOTIFICATION_EMAIL_ENABLED_KEY,
|
||||
IMPORTANT_NOTIFICATION_EMAIL_RECIPIENTS_KEY, MAX_NOTIFICATION_ITEMS,
|
||||
MAX_NOTIFICATION_RECIPIENTS, MAX_NOTIFICATION_RECIPIENT_BYTES,
|
||||
MAX_NOTIFICATION_TEMPLATE_BYTES,
|
||||
};
|
||||
use crate::{data::GatewayDataState, AppState};
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use serde_json::json;
|
||||
|
||||
#[tokio::test]
|
||||
async fn unused_email_channel_does_not_load_or_migrate_smtp_password() {
|
||||
for (email_enabled, recipients) in [(false, "[email protected]"), (true, "")] {
|
||||
let data = GatewayDataState::disabled()
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||
.with_system_config_values_for_tests(vec![
|
||||
(
|
||||
IMPORTANT_NOTIFICATION_EMAIL_ENABLED_KEY.to_string(),
|
||||
json!(email_enabled),
|
||||
),
|
||||
(
|
||||
IMPORTANT_NOTIFICATION_EMAIL_RECIPIENTS_KEY.to_string(),
|
||||
json!(recipients),
|
||||
),
|
||||
("smtp_host".to_string(), json!("smtp.example.com")),
|
||||
("smtp_user".to_string(), json!("[email protected]")),
|
||||
("smtp_password".to_string(), json!("unused-smtp-password")),
|
||||
("smtp_from_email".to_string(), json!("[email protected]")),
|
||||
]);
|
||||
let state = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data);
|
||||
assert!(!important_notification_configured(&state).await.unwrap());
|
||||
assert_eq!(
|
||||
state
|
||||
.read_system_config_json_value_strong("smtp_password")
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(json!("unused-smtp-password"))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_recipient_list_accepts_arrays_and_delimiters() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -2408,10 +2408,6 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
|
||||
);
|
||||
}
|
||||
}
|
||||
match state.prewarm_execution_extra_trusted_dns_hosts().await {
|
||||
Ok(_) => info!("prewarmed execution Fake-IP DNS allowlist"),
|
||||
Err(err) => warn!(error = %err, "failed to prewarm execution Fake-IP DNS allowlist"),
|
||||
}
|
||||
match prewarm_direct_h2c_sender_cache_from_env_for_startup().await {
|
||||
Ok(Some(report)) => {
|
||||
if report.failed_targets > 0 {
|
||||
|
||||
@@ -533,12 +533,10 @@ impl AppState {
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
|
||||
let keys = self
|
||||
.data
|
||||
self.data
|
||||
.list_provider_catalog_key_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
self.open_provider_catalog_keys(keys).await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn list_provider_catalog_key_maintenance_summaries_by_provider_ids(
|
||||
|
||||
@@ -272,6 +272,50 @@ mod tests {
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn app_state_reads_redacted_key_summaries_without_opening_credentials() {
|
||||
for (api_key, auth_config) in [
|
||||
(Some("summary"), None),
|
||||
(None, Some("{}")),
|
||||
(Some("summary"), Some("{}")),
|
||||
] {
|
||||
let health = serde_json::json!({"openai:chat": {"health_score": 0.75}});
|
||||
let key = sample_key(
|
||||
"key-1",
|
||||
"provider-1",
|
||||
api_key.map(ToOwned::to_owned),
|
||||
auth_config.map(ToOwned::to_owned),
|
||||
)
|
||||
.with_health_fields(Some(health.clone()), None);
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-1")],
|
||||
Vec::new(),
|
||||
vec![key],
|
||||
));
|
||||
let state = AppState::new()
|
||||
.expect("test state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_reader_for_tests(repository)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
let provider_ids = ["provider-1".to_string()];
|
||||
|
||||
let summaries = state
|
||||
.list_provider_catalog_key_summaries_by_provider_ids(&provider_ids)
|
||||
.await
|
||||
.expect("redacted summaries should not require credential authentication");
|
||||
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0].health_by_format.as_ref(), Some(&health));
|
||||
assert_eq!(summaries[0].encrypted_api_key.as_deref(), api_key);
|
||||
assert_eq!(summaries[0].encrypted_auth_config.as_deref(), auth_config);
|
||||
assert!(state
|
||||
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
|
||||
.await
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn app_state_migrates_both_legacy_fields_with_one_exact_cas() {
|
||||
let legacy_api =
|
||||
|
||||
@@ -154,15 +154,6 @@ impl AppState {
|
||||
.map_err(|err| format!("{err:?}"))
|
||||
}
|
||||
|
||||
pub async fn prewarm_execution_extra_trusted_dns_hosts(&self) -> Result<(), String> {
|
||||
self.read_system_config_json_value(
|
||||
aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY,
|
||||
)
|
||||
.await
|
||||
.map(|_| ())
|
||||
.map_err(|err| format!("{err:?}"))
|
||||
}
|
||||
|
||||
fn usage_worker_queue_for(
|
||||
runtime_state: &Arc<RuntimeState>,
|
||||
) -> Option<Arc<dyn RuntimeQueueStore>> {
|
||||
@@ -778,18 +769,6 @@ impl AppState {
|
||||
.expect("admin monitoring error stats reset cache should lock")
|
||||
}
|
||||
|
||||
fn refresh_execution_extra_trusted_dns_hosts(
|
||||
&self,
|
||||
key: &str,
|
||||
value: Option<&serde_json::Value>,
|
||||
) {
|
||||
if key.eq_ignore_ascii_case(
|
||||
aether_admin::system::EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY,
|
||||
) {
|
||||
crate::execution_runtime::transport::refresh_execution_extra_trusted_dns_hosts(value);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn mark_admin_monitoring_error_stats_reset(&self, now_unix_secs: u64) {
|
||||
let mut reset_at = self
|
||||
.admin_monitoring_error_stats_reset_at
|
||||
@@ -809,7 +788,6 @@ impl AppState {
|
||||
SYSTEM_CONFIG_CACHE_MAX_STALENESS,
|
||||
)
|
||||
.await?;
|
||||
self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref());
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
@@ -822,7 +800,6 @@ impl AppState {
|
||||
.find_system_config_value_strong(key)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref());
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
@@ -961,7 +938,6 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
self.system_config_cache
|
||||
.insert(key.to_string(), None, SYSTEM_CONFIG_CACHE_MAX_STALENESS);
|
||||
self.refresh_execution_extra_trusted_dns_hosts(key, None);
|
||||
if deleted && system_config_key_affects_scheduler(key) {
|
||||
self.invalidate_scheduler_affinity_cache();
|
||||
}
|
||||
@@ -1043,7 +1019,6 @@ impl AppState {
|
||||
}
|
||||
|
||||
fn remember_system_config_write(&self, key: &str, value: Option<serde_json::Value>) {
|
||||
self.refresh_execution_extra_trusted_dns_hosts(key, value.as_ref());
|
||||
self.system_config_cache
|
||||
.insert(key.to_string(), value, SYSTEM_CONFIG_CACHE_MAX_STALENESS);
|
||||
if system_config_key_affects_scheduler(key) {
|
||||
@@ -1091,7 +1066,6 @@ impl AppState {
|
||||
| aether_data::repository::system::AdminSystemPurgeTarget::Stats
|
||||
) {
|
||||
self.system_config_cache.clear();
|
||||
crate::execution_runtime::transport::refresh_execution_extra_trusted_dns_hosts(None);
|
||||
self.invalidate_provider_routing_caches();
|
||||
}
|
||||
Ok(summary)
|
||||
|
||||
@@ -2460,15 +2460,21 @@ async fn gateway_retries_next_local_openai_chat_stream_candidate_after_retryable
|
||||
failed_candidate.error_type.as_deref(),
|
||||
Some("retryable_upstream_status")
|
||||
);
|
||||
assert!(failed_candidate.error_message.is_none());
|
||||
assert!(failed_candidate.error_message.is_some());
|
||||
let failed_upstream_response = failed_candidate
|
||||
.extra_data
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upstream_response"))
|
||||
.expect("failed stream candidate should keep its upstream response");
|
||||
assert_eq!(failed_upstream_response["status_code"], json!(429));
|
||||
assert!(failed_upstream_response.get("headers").is_none());
|
||||
assert!(failed_upstream_response.get("body").is_none());
|
||||
assert_eq!(
|
||||
failed_upstream_response["headers"]["content-type"],
|
||||
"application/json"
|
||||
);
|
||||
assert_eq!(
|
||||
failed_upstream_response["body"]["error"]["message"],
|
||||
"rate limited"
|
||||
);
|
||||
assert_eq!(success_candidate.status, RequestCandidateStatus::Success);
|
||||
assert_eq!(success_candidate.status_code, Some(200));
|
||||
assert!(success_candidate.started_at_unix_ms.is_some());
|
||||
|
||||
@@ -1166,15 +1166,21 @@ async fn gateway_retries_next_local_openai_chat_sync_candidate_after_auth_failur
|
||||
assert_eq!(failed_candidate.retry_index, retry_index as u32);
|
||||
assert_eq!(failed_candidate.status, RequestCandidateStatus::Failed);
|
||||
assert_eq!(failed_candidate.status_code, Some(401));
|
||||
assert!(failed_candidate.error_message.is_none());
|
||||
assert!(failed_candidate.error_message.is_some());
|
||||
let failed_upstream_response = failed_candidate
|
||||
.extra_data
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("upstream_response"))
|
||||
.expect("failed candidate should keep its upstream response");
|
||||
assert_eq!(failed_upstream_response["status_code"], json!(401));
|
||||
assert!(failed_upstream_response.get("headers").is_none());
|
||||
assert!(failed_upstream_response.get("body").is_none());
|
||||
assert_eq!(
|
||||
failed_upstream_response["headers"]["content-type"],
|
||||
"application/json"
|
||||
);
|
||||
assert_eq!(
|
||||
failed_upstream_response["body"]["error"]["message"],
|
||||
"invalid auth token"
|
||||
);
|
||||
}
|
||||
assert_eq!(stored_candidates[2].candidate_index, 1);
|
||||
assert_eq!(stored_candidates[2].status, RequestCandidateStatus::Success);
|
||||
|
||||
@@ -534,6 +534,24 @@ async fn gateway_exposes_request_audit_bundle_via_internal_audit_endpoint() {
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() {
|
||||
let mut failed_candidate = sample_request_candidate(
|
||||
"cand-2",
|
||||
"req-trace-1",
|
||||
1,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(37),
|
||||
Some(502),
|
||||
);
|
||||
failed_candidate.error_message = Some("private upstream diagnostic".to_string());
|
||||
failed_candidate.extra_data = Some(json!({
|
||||
"upstream_response": {
|
||||
"status_code": 502,
|
||||
"headers": {"x-request-id": "private-upstream-id"},
|
||||
"body": {"error": {"message": "private upstream diagnostic"}}
|
||||
},
|
||||
"error_flow": {"status_code": 502, "message": "private upstream diagnostic"}
|
||||
}));
|
||||
let repository = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
sample_request_candidate(
|
||||
"cand-1",
|
||||
@@ -544,15 +562,7 @@ async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() {
|
||||
None,
|
||||
None,
|
||||
),
|
||||
sample_request_candidate(
|
||||
"cand-2",
|
||||
"req-trace-1",
|
||||
1,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(37),
|
||||
Some(502),
|
||||
),
|
||||
failed_candidate,
|
||||
]));
|
||||
|
||||
let gateway_state = AppState::new()
|
||||
@@ -582,6 +592,14 @@ async fn gateway_exposes_request_candidate_trace_via_internal_audit_endpoint() {
|
||||
);
|
||||
assert_eq!(payload["candidates"][0]["id"], "cand-2");
|
||||
assert_eq!(payload["candidates"][0]["status"], "failed");
|
||||
assert!(payload["candidates"][0]["error_message"].is_null());
|
||||
assert_eq!(
|
||||
payload["candidates"][0]["extra_data"]["upstream_response"]["status_code"],
|
||||
502
|
||||
);
|
||||
let serialized = payload.to_string();
|
||||
assert!(!serialized.contains("private upstream diagnostic"));
|
||||
assert!(!serialized.contains("private-upstream-id"));
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
use aether_data::repository::auth_modules::InMemoryAuthModuleReadRepository;
|
||||
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
|
||||
use aether_data::repository::management_tokens::{
|
||||
@@ -32,6 +32,155 @@ use crate::data::GatewayDataState;
|
||||
const ADMIN_ENDPOINT_HEALTH_DATA_UNAVAILABLE_DETAIL: &str =
|
||||
"Admin endpoint health data unavailable";
|
||||
|
||||
async fn assert_admin_modules_status_with_smtp_password(
|
||||
stored_password: &str,
|
||||
notification_ready: bool,
|
||||
server_chan_enabled: bool,
|
||||
) -> AppState {
|
||||
let data = GatewayDataState::with_auth_module_reader_for_tests(Arc::new(
|
||||
InMemoryAuthModuleReadRepository::seed(Vec::new(), None),
|
||||
))
|
||||
.with_provider_catalog_reader(Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
)))
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||
.with_system_config_values_for_tests(vec![
|
||||
("module.management_tokens.enabled".to_string(), json!(true)),
|
||||
(
|
||||
"module.important_notification.enabled".to_string(),
|
||||
json!(true),
|
||||
),
|
||||
(
|
||||
"module.important_notification.email_enabled".to_string(),
|
||||
json!(true),
|
||||
),
|
||||
(
|
||||
"module.important_notification.email_recipients".to_string(),
|
||||
json!("[email protected]"),
|
||||
),
|
||||
(
|
||||
"module.server_chan_push.enabled".to_string(),
|
||||
json!(server_chan_enabled),
|
||||
),
|
||||
(
|
||||
"module.server_chan_push.send_key".to_string(),
|
||||
json!(if server_chan_enabled {
|
||||
"SCT-test-send-key"
|
||||
} else {
|
||||
""
|
||||
}),
|
||||
),
|
||||
("smtp_host".to_string(), json!("smtp.example.com")),
|
||||
("smtp_port".to_string(), json!(587)),
|
||||
("smtp_user".to_string(), json!("[email protected]")),
|
||||
("smtp_password".to_string(), json!(stored_password)),
|
||||
("smtp_use_tls".to_string(), json!(true)),
|
||||
("smtp_from_email".to_string(), json!("[email protected]")),
|
||||
]);
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data);
|
||||
let (gateway_url, gateway_handle) = start_server(build_router_with_state(state.clone())).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
for path in [
|
||||
"/api/admin/modules/status",
|
||||
"/api/admin/modules/status/important_notification",
|
||||
] {
|
||||
let response = client
|
||||
.get(format!("{gateway_url}{path}"))
|
||||
.header(GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.send()
|
||||
.await
|
||||
.expect("module status request should succeed");
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("module status should parse");
|
||||
assert!(!payload.to_string().contains(stored_password));
|
||||
let notification = if path == "/api/admin/modules/status" {
|
||||
assert_eq!(
|
||||
payload
|
||||
.as_object()
|
||||
.expect("module list should be an object")
|
||||
.len(),
|
||||
14
|
||||
);
|
||||
assert_eq!(payload["management_tokens"]["active"], json!(true));
|
||||
&payload["important_notification"]
|
||||
} else {
|
||||
&payload
|
||||
};
|
||||
assert_eq!(notification["enabled"], json!(true));
|
||||
assert_eq!(notification["config_validated"], json!(notification_ready));
|
||||
assert_eq!(notification["active"], json!(notification_ready));
|
||||
assert_eq!(notification["config_error"].is_null(), notification_ready);
|
||||
}
|
||||
gateway_handle.abort();
|
||||
|
||||
assert_eq!(
|
||||
crate::important_notification::important_notification_dispatch_ready_for_item(
|
||||
&state,
|
||||
crate::important_notification::PROVIDER_QUOTA_ALERT_ITEM_KEY,
|
||||
)
|
||||
.await
|
||||
.expect("SMTP errors should not abort notification readiness"),
|
||||
notification_ready
|
||||
);
|
||||
let summary = crate::maintenance::perform_provider_quota_alert_once(&state)
|
||||
.await
|
||||
.expect("SMTP errors should not abort the quota alert worker");
|
||||
assert_eq!(summary.failed, 0);
|
||||
assert_eq!(summary.alerted, 0);
|
||||
state
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_modules_status_with_legacy_smtp_password() {
|
||||
let ciphertext =
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-smtp-password")
|
||||
.expect("legacy SMTP password should encrypt");
|
||||
let state = assert_admin_modules_status_with_smtp_password(&ciphertext, true, false).await;
|
||||
let stored = state
|
||||
.read_system_config_json_value_strong("smtp_password")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert!(stored
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.starts_with("aether-smtp-password-v3:"));
|
||||
let smtp = crate::email_delivery::read_smtp_delivery_config(&state)
|
||||
.await
|
||||
.expect("migrated SMTP config should load")
|
||||
.expect("SMTP should be configured");
|
||||
assert_eq!(smtp.password.as_deref(), Some("legacy-smtp-password"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_modules_status_with_invalid_smtp_password() {
|
||||
let ciphertext =
|
||||
encrypt_python_fernet_plaintext("unavailable-historical-key", "legacy-smtp-password")
|
||||
.expect("unknown-key SMTP password should encrypt");
|
||||
let state = assert_admin_modules_status_with_smtp_password(&ciphertext, false, false).await;
|
||||
assert_eq!(
|
||||
state
|
||||
.read_system_config_json_value_strong("smtp_password")
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(json!(ciphertext))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_modules_status_with_invalid_smtp_and_working_push() {
|
||||
assert_admin_modules_status_with_smtp_password("aether-smtp-password-v3:invalid", true, true)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_returns_service_unavailable_for_admin_health_api_formats_when_readers_unavailable()
|
||||
{
|
||||
|
||||
@@ -473,6 +473,23 @@ async fn gateway_handles_admin_monitoring_trace_request_locally_with_trusted_adm
|
||||
}),
|
||||
);
|
||||
|
||||
let mut failed_candidate = sample_candidate(
|
||||
"cand-used",
|
||||
"request-1",
|
||||
1,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(502),
|
||||
);
|
||||
failed_candidate.error_message = Some("private upstream diagnostic".to_string());
|
||||
failed_candidate.extra_data = Some(json!({
|
||||
"upstream_response": {
|
||||
"status_code": 502,
|
||||
"headers": {"x-request-id": "upstream-diagnostic-id"},
|
||||
"body": {"error": {"message": "private upstream diagnostic"}}
|
||||
}
|
||||
}));
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
sample_candidate(
|
||||
"cand-unused",
|
||||
@@ -483,15 +500,7 @@ async fn gateway_handles_admin_monitoring_trace_request_locally_with_trusted_adm
|
||||
None,
|
||||
None,
|
||||
),
|
||||
sample_candidate(
|
||||
"cand-used",
|
||||
"request-1",
|
||||
1,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(502),
|
||||
),
|
||||
failed_candidate,
|
||||
]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
@@ -538,6 +547,37 @@ async fn gateway_handles_admin_monitoring_trace_request_locally_with_trusted_adm
|
||||
assert_eq!(payload["candidates"][0]["key_auth_type"], json!("api_key"));
|
||||
assert_eq!(payload["candidates"][0]["latency_ms"], json!(33));
|
||||
assert_eq!(payload["candidates"][0]["status_code"], json!(502));
|
||||
assert_eq!(
|
||||
payload["candidates"][0]["error_message"],
|
||||
"private upstream diagnostic"
|
||||
);
|
||||
assert_eq!(
|
||||
payload["candidates"][0]["extra_data"]["upstream_response"]["body"]["error"]["message"],
|
||||
"private upstream diagnostic"
|
||||
);
|
||||
for role in [None, Some("user")] {
|
||||
let mut request = reqwest::Client::new().get(format!(
|
||||
"{gateway_url}/api/admin/monitoring/trace/request-1"
|
||||
));
|
||||
if let Some(role) = role {
|
||||
request = request
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "regular-user")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, role)
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "regular-session");
|
||||
}
|
||||
let response = request
|
||||
.send()
|
||||
.await
|
||||
.expect("unauthorized probe should complete");
|
||||
assert!(matches!(
|
||||
response.status(),
|
||||
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN
|
||||
));
|
||||
let body = response.text().await.expect("denial body should read");
|
||||
assert!(!body.contains("private upstream diagnostic"));
|
||||
assert!(!body.contains("upstream-diagnostic-id"));
|
||||
}
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
|
||||
@@ -70,6 +70,66 @@ async fn provider_health_summary(
|
||||
payload["items"][0].clone()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_provider_summary_health_preserves_redacted_key_summaries() {
|
||||
let endpoint = sample_endpoint(
|
||||
"endpoint-chat",
|
||||
"provider-openai",
|
||||
"openai:chat",
|
||||
"https://api.openai.example",
|
||||
);
|
||||
let keys = [
|
||||
("key-api", "api_key", None, 0.25),
|
||||
("key-oauth", "oauth", Some("{}"), 0.75),
|
||||
]
|
||||
.into_iter()
|
||||
.map(|(key_id, auth_type, auth_config, score)| {
|
||||
let mut key = sample_key(key_id, "provider-openai", "openai:chat", "test")
|
||||
.with_health_fields(Some(json!({"openai:chat": {"health_score": score}})), None);
|
||||
key.auth_type = auth_type.to_string();
|
||||
key.encrypted_api_key = Some("summary".to_string());
|
||||
key.encrypted_auth_config = auth_config.map(ToOwned::to_owned);
|
||||
key
|
||||
})
|
||||
.collect();
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-openai", "openai", 10)],
|
||||
vec![endpoint],
|
||||
keys,
|
||||
));
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||
repository,
|
||||
));
|
||||
|
||||
for uri in [
|
||||
"/api/admin/providers/summary",
|
||||
"/api/admin/providers/provider-openai/summary",
|
||||
] {
|
||||
let response = local_admin_providers_response(&state, http::Method::GET, uri, None).await;
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = axum::body::to_bytes(response.into_body(), 1024 * 1024)
|
||||
.await
|
||||
.expect("summary body should read");
|
||||
let payload: serde_json::Value =
|
||||
serde_json::from_slice(&body).expect("summary should parse");
|
||||
let summary = if uri == "/api/admin/providers/summary" {
|
||||
&payload["items"][0]
|
||||
} else {
|
||||
&payload
|
||||
};
|
||||
|
||||
assert_eq!(summary["total_keys"], 2);
|
||||
assert_eq!(summary["active_keys"], 2);
|
||||
assert_eq!(summary["endpoint_health_details"][0]["total_keys"], 2);
|
||||
assert_eq!(summary["endpoint_health_details"][0]["active_keys"], 2);
|
||||
assert_eq!(summary["endpoint_health_details"][0]["health_score"], 0.5);
|
||||
assert_eq!(summary["avg_health_score"], 0.5);
|
||||
assert_eq!(summary["unhealthy_endpoints"], 0);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_provider_summary_health_ignores_disabled_keys() {
|
||||
let endpoint = sample_endpoint(
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::Arc;
|
||||
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use parking_lot::Mutex;
|
||||
use tokio::sync::Notify;
|
||||
|
||||
const CHUNK_BYTES: usize = 32 * 1024;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum LocalBodyEvent {
|
||||
Chunk(Bytes),
|
||||
End,
|
||||
Error(String),
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct BufferState {
|
||||
chunks: VecDeque<BytesMut>,
|
||||
bytes: usize,
|
||||
terminal: Option<Result<(), String>>,
|
||||
receiver_taken: bool,
|
||||
receiver_closed: bool,
|
||||
}
|
||||
|
||||
pub(super) struct ResponseBuffer {
|
||||
state: Mutex<BufferState>,
|
||||
notify: Notify,
|
||||
capacity: usize,
|
||||
}
|
||||
|
||||
impl ResponseBuffer {
|
||||
pub(super) fn new(capacity: usize) -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
state: Mutex::new(BufferState::default()),
|
||||
notify: Notify::new(),
|
||||
capacity,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn take_receiver(self: &Arc<Self>) -> Option<BodyReceiver> {
|
||||
let mut state = self.state.lock();
|
||||
if state.receiver_taken {
|
||||
return None;
|
||||
}
|
||||
state.receiver_taken = true;
|
||||
Some(BodyReceiver {
|
||||
buffer: Arc::clone(self),
|
||||
finished: false,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn push(&self, mut payload: Bytes) -> bool {
|
||||
let mut state = self.state.lock();
|
||||
if state.terminal.is_some()
|
||||
|| state.receiver_closed
|
||||
|| payload.len() > self.capacity.saturating_sub(state.bytes)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
state.bytes += payload.len();
|
||||
while !payload.is_empty() {
|
||||
if let Some(tail) = state
|
||||
.chunks
|
||||
.back_mut()
|
||||
.filter(|chunk| chunk.len() < CHUNK_BYTES)
|
||||
{
|
||||
let count = payload.len().min(CHUNK_BYTES - tail.len());
|
||||
tail.extend_from_slice(&payload.split_to(count));
|
||||
} else {
|
||||
let count = payload.len().min(CHUNK_BYTES);
|
||||
let chunk = payload.split_to(count);
|
||||
state.chunks.push_back(
|
||||
chunk
|
||||
.try_into_mut()
|
||||
.unwrap_or_else(|chunk| BytesMut::from(chunk.as_ref())),
|
||||
);
|
||||
}
|
||||
}
|
||||
drop(state);
|
||||
self.notify.notify_waiters();
|
||||
true
|
||||
}
|
||||
|
||||
pub(super) fn finish(&self, result: Result<(), String>) {
|
||||
let mut state = self.state.lock();
|
||||
if state.terminal.is_none() {
|
||||
state.terminal = Some(result);
|
||||
}
|
||||
drop(state);
|
||||
self.notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct BodyReceiver {
|
||||
buffer: Arc<ResponseBuffer>,
|
||||
finished: bool,
|
||||
}
|
||||
|
||||
impl BodyReceiver {
|
||||
pub(super) async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||
if self.finished {
|
||||
return None;
|
||||
}
|
||||
loop {
|
||||
let notified = self.buffer.notify.notified();
|
||||
tokio::pin!(notified);
|
||||
notified.as_mut().enable();
|
||||
{
|
||||
let mut state = self.buffer.state.lock();
|
||||
if let Some(chunk) = state.chunks.pop_front() {
|
||||
state.bytes -= chunk.len();
|
||||
return Some(LocalBodyEvent::Chunk(chunk.freeze()));
|
||||
}
|
||||
if let Some(terminal) = state.terminal.take() {
|
||||
self.finished = true;
|
||||
state.receiver_closed = true;
|
||||
return Some(match terminal {
|
||||
Ok(()) => LocalBodyEvent::End,
|
||||
Err(error) => LocalBodyEvent::Error(error),
|
||||
});
|
||||
}
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for BodyReceiver {
|
||||
fn drop(&mut self) {
|
||||
let mut state = self.buffer.state.lock();
|
||||
state.receiver_closed = true;
|
||||
state.chunks.clear();
|
||||
state.bytes = 0;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn error_survives_a_full_buffer() {
|
||||
let buffer = ResponseBuffer::new(CHUNK_BYTES);
|
||||
let mut receiver = buffer.take_receiver().unwrap();
|
||||
assert!(buffer.push(Bytes::from(vec![b'x'; CHUNK_BYTES])));
|
||||
buffer.finish(Err("proxy disconnected".into()));
|
||||
assert!(matches!(
|
||||
receiver.recv().await,
|
||||
Some(LocalBodyEvent::Chunk(_))
|
||||
));
|
||||
assert!(
|
||||
matches!(receiver.recv().await, Some(LocalBodyEvent::Error(error)) if error == "proxy disconnected")
|
||||
);
|
||||
assert!(receiver.recv().await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn small_frames_are_coalesced_within_the_byte_budget() {
|
||||
let buffer = ResponseBuffer::new(4096);
|
||||
let mut receiver = buffer.take_receiver().unwrap();
|
||||
for _ in 0..4096 {
|
||||
assert!(buffer.push(Bytes::from_static(b"x")));
|
||||
}
|
||||
assert!(!buffer.push(Bytes::from_static(b"x")));
|
||||
assert_eq!(buffer.state.lock().chunks.len(), 1);
|
||||
buffer.finish(Ok(()));
|
||||
assert!(
|
||||
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 4096)
|
||||
);
|
||||
assert!(matches!(receiver.recv().await, Some(LocalBodyEvent::End)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn terminal_wakes_an_empty_receiver_and_is_not_overwritten() {
|
||||
let buffer = ResponseBuffer::new(1024);
|
||||
let mut receiver = buffer.take_receiver().unwrap();
|
||||
let task = tokio::spawn(async move { receiver.recv().await });
|
||||
tokio::task::yield_now().await;
|
||||
buffer.finish(Err("cancelled".into()));
|
||||
buffer.finish(Ok(()));
|
||||
assert!(
|
||||
matches!(task.await.unwrap(), Some(LocalBodyEvent::Error(error)) if error == "cancelled")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,250 @@
|
||||
use super::*;
|
||||
|
||||
async fn fixture(
|
||||
window: u32,
|
||||
capacity: usize,
|
||||
) -> (
|
||||
Arc<HubRouter>,
|
||||
Arc<ProxyConn>,
|
||||
Arc<LocalStream>,
|
||||
aether_runtime::BoundedQueueReceiver<Message>,
|
||||
) {
|
||||
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||
let (sender, receiver) = bounded_queue(capacity);
|
||||
let (close_tx, _) = watch::channel(false);
|
||||
let connection = Arc::new(
|
||||
ProxyConn::new(
|
||||
99,
|
||||
"flow-test".into(),
|
||||
"flow-test".into(),
|
||||
sender,
|
||||
close_tx,
|
||||
16,
|
||||
3,
|
||||
)
|
||||
.with_settings(protocol::SettingsPayload {
|
||||
initial_stream_window_bytes: window,
|
||||
min_window_update_bytes: (window / 4).max(1),
|
||||
drain_deadline_ms: 1000,
|
||||
}),
|
||||
);
|
||||
hub.register_proxy(Arc::clone(&connection));
|
||||
let stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||
(hub, connection, stream, receiver)
|
||||
}
|
||||
|
||||
fn meta() -> protocol::RequestMeta {
|
||||
protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
method: "GET".into(),
|
||||
url: "https://example.com".into(),
|
||||
headers: HashMap::new(),
|
||||
stream: true,
|
||||
request_timeout_ms: None,
|
||||
stream_first_byte_timeout_ms: None,
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
}
|
||||
}
|
||||
|
||||
async fn headers(hub: &Arc<HubRouter>, stream: &LocalStream) {
|
||||
let payload = serde_json::to_vec(&protocol::ResponseMeta {
|
||||
status: 200,
|
||||
headers: vec![],
|
||||
})
|
||||
.unwrap();
|
||||
let mut frame = protocol::encode_frame(
|
||||
stream.proxy_stream_id,
|
||||
protocol::RESPONSE_HEADERS,
|
||||
0,
|
||||
&payload,
|
||||
);
|
||||
hub.handle_proxy_frame(stream.proxy_conn_id, &mut frame)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn window_credit_is_retried_after_queue_pressure_and_cancelled_receive() {
|
||||
let (hub, _, stream, mut outbound) = fixture(128, 1).await;
|
||||
headers(&hub, &stream).await;
|
||||
assert!(stream.push_body_chunk(Bytes::from(vec![b'x'; 64])));
|
||||
let mut receiver = stream.take_body_receiver().unwrap();
|
||||
assert!(
|
||||
tokio::time::timeout(Duration::from_millis(10), receiver.recv())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
assert_eq!(*stream.response_consumed_since_update.lock(), 64);
|
||||
outbound.recv().await.unwrap();
|
||||
let event = tokio::time::timeout(Duration::from_secs(1), receiver.recv())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(event, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 64));
|
||||
assert_eq!(*stream.response_consumed_since_update.lock(), 0);
|
||||
let Message::Binary(data) = outbound.recv().await.unwrap() else {
|
||||
panic!("expected binary update")
|
||||
};
|
||||
let frame = aether_contracts::tunnel::Frame::decode(data).unwrap();
|
||||
let update: protocol::WindowUpdatePayload = serde_json::from_slice(&frame.payload).unwrap();
|
||||
assert_eq!(
|
||||
frame.msg_type,
|
||||
aether_contracts::tunnel::MsgType::WindowUpdate
|
||||
);
|
||||
assert_eq!(update.delta_bytes, 64);
|
||||
hub.cancel_local_stream(stream.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn response_credit_is_not_returned_until_consumed() {
|
||||
let (hub, _, stream, mut outbound) = fixture(128, 4).await;
|
||||
outbound.recv().await.unwrap();
|
||||
let mut body = protocol::encode_frame(
|
||||
stream.proxy_stream_id,
|
||||
protocol::RESPONSE_BODY,
|
||||
0,
|
||||
&[b'x'; 128],
|
||||
);
|
||||
hub.handle_proxy_frame(99, &mut body).await;
|
||||
assert!(outbound.try_recv().is_err());
|
||||
let mut receiver = stream.take_body_receiver().unwrap();
|
||||
assert!(
|
||||
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 128)
|
||||
);
|
||||
assert!(outbound.try_recv().is_ok());
|
||||
hub.cancel_local_stream(stream.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelled_stream_open_releases_slot_without_resetting_connection() {
|
||||
let (hub, connection, first_stream, mut outbound) = fixture(128, 1).await;
|
||||
let opening_hub = Arc::clone(&hub);
|
||||
let opening =
|
||||
tokio::spawn(async move { opening_hub.open_local_stream("flow-test", &meta()).await });
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while hub.local_streams.len() != 2 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
opening.abort();
|
||||
assert!(matches!(opening.await, Err(error) if error.is_cancelled()));
|
||||
assert_eq!(connection.stream_count.load(Ordering::Relaxed), 1);
|
||||
assert_eq!(hub.local_streams.len(), 1);
|
||||
assert_eq!(hub.proxy_to_local.len(), 1);
|
||||
assert!(connection.is_available());
|
||||
outbound.recv().await.unwrap();
|
||||
assert!(outbound.try_recv().is_err());
|
||||
let next_stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||
outbound.recv().await.unwrap();
|
||||
hub.cancel_local_stream(first_stream.id, "test complete");
|
||||
outbound.recv().await.unwrap();
|
||||
hub.cancel_local_stream(next_stream.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn full_response_buffer_preserves_disconnect_error() {
|
||||
let (hub, connection, stream, mut outbound) = fixture(4 * 1024 * 1024, 512).await;
|
||||
outbound.recv().await.unwrap();
|
||||
headers(&hub, &stream).await;
|
||||
let mut receiver = stream.take_body_receiver().unwrap();
|
||||
for _ in 0..128 {
|
||||
let mut frame = protocol::encode_frame(
|
||||
stream.proxy_stream_id,
|
||||
protocol::RESPONSE_BODY,
|
||||
0,
|
||||
&vec![b'x'; 32 * 1024],
|
||||
);
|
||||
hub.handle_proxy_frame(99, &mut frame).await;
|
||||
}
|
||||
hub.unregister_proxy(connection.id, &connection.node_id);
|
||||
let mut bytes = 0;
|
||||
loop {
|
||||
match receiver.recv().await {
|
||||
Some(LocalBodyEvent::Chunk(chunk)) => bytes += chunk.len(),
|
||||
Some(LocalBodyEvent::Error(error)) => {
|
||||
assert!(error.contains("disconnected"));
|
||||
break;
|
||||
}
|
||||
event => panic!("disconnect must not become normal EOF: {event:?}"),
|
||||
}
|
||||
}
|
||||
assert_eq!(bytes, 4 * 1024 * 1024);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slow_stream_does_not_block_another_stream_on_the_same_connection() {
|
||||
let (hub, _, slow, mut outbound) = fixture(128, 512).await;
|
||||
outbound.recv().await.unwrap();
|
||||
let fast = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||
assert!(slow.push_body_chunk(Bytes::from(vec![b'x'; 128])));
|
||||
let mut overflowing = protocol::encode_frame(
|
||||
slow.proxy_stream_id,
|
||||
protocol::RESPONSE_BODY,
|
||||
0,
|
||||
b"overflow",
|
||||
);
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
hub.handle_proxy_frame(99, &mut overflowing).await;
|
||||
headers(&hub, &fast).await;
|
||||
assert_eq!(
|
||||
fast.wait_headers(Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap()
|
||||
.status,
|
||||
200
|
||||
);
|
||||
})
|
||||
.await
|
||||
.expect("slow stream must not block connection reader");
|
||||
assert!(!hub.local_streams.contains_key(&slow.id));
|
||||
assert!(hub.local_streams.contains_key(&fast.id));
|
||||
hub.cancel_local_stream(fast.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelling_a_stream_wakes_request_window_waiters() {
|
||||
let (_, _, stream, _) = fixture(128, 512).await;
|
||||
*stream.request_window.available.lock() = 0;
|
||||
let waiter = tokio::spawn({
|
||||
let stream = Arc::clone(&stream);
|
||||
async move {
|
||||
stream
|
||||
.acquire_request_window(1, Duration::from_secs(30))
|
||||
.await
|
||||
}
|
||||
});
|
||||
tokio::task::yield_now().await;
|
||||
stream.fail("cancelled");
|
||||
assert!(tokio::time::timeout(Duration::from_secs(1), waiter)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn concurrent_headers_and_credit_updates_do_not_lose_notifications() {
|
||||
for index in 0..256 {
|
||||
let stream = Arc::new(LocalStream::new(index, "test".into(), 1, 1, 1));
|
||||
let window = Arc::new(StreamFlowWindow::new(0));
|
||||
let waiter = tokio::spawn({
|
||||
let stream = Arc::clone(&stream);
|
||||
let window = Arc::clone(&window);
|
||||
async move {
|
||||
stream.wait_headers(Duration::from_secs(1)).await.unwrap();
|
||||
window.acquire(1, Duration::from_secs(1)).await.unwrap();
|
||||
}
|
||||
});
|
||||
stream.set_response_headers(protocol::ResponseMeta {
|
||||
status: 200,
|
||||
headers: vec![],
|
||||
});
|
||||
window.add(1);
|
||||
waiter.await.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -12,10 +12,11 @@ use axum::extract::ws::Message;
|
||||
use bytes::Bytes;
|
||||
use dashmap::DashMap;
|
||||
use parking_lot::{Mutex, RwLock};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::sync::{watch, Notify};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
pub use super::body::LocalBodyEvent;
|
||||
use super::body::{BodyReceiver, ResponseBuffer};
|
||||
use super::control_plane::ControlPlaneClient;
|
||||
use super::protocol;
|
||||
|
||||
@@ -29,6 +30,10 @@ const DEFAULT_DRAIN_DEADLINE_MS: u64 = 30_000;
|
||||
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
|
||||
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "flow_control_tests.rs"]
|
||||
mod flow_control_tests;
|
||||
|
||||
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
||||
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
|
||||
.ok()
|
||||
@@ -53,11 +58,16 @@ static NODE_STATUS_QUEUE_CAPACITY: LazyLock<usize> = LazyLock::new(|| {
|
||||
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
|
||||
});
|
||||
|
||||
static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
||||
STREAM_INITIAL_WINDOW_BYTES
|
||||
.saturating_div(4)
|
||||
.clamp(1, 1024 * 1024)
|
||||
});
|
||||
pub(super) fn local_settings() -> protocol::SettingsPayload {
|
||||
protocol::SettingsPayload {
|
||||
initial_stream_window_bytes: (*STREAM_INITIAL_WINDOW_BYTES)
|
||||
.min(aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u32),
|
||||
min_window_update_bytes: STREAM_INITIAL_WINDOW_BYTES
|
||||
.saturating_div(4)
|
||||
.clamp(1, 1024 * 1024),
|
||||
drain_deadline_ms: *DRAIN_DEADLINE_MS,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SendStatus {
|
||||
@@ -91,6 +101,7 @@ impl ConnHealthState {
|
||||
struct StreamFlowWindow {
|
||||
available: Mutex<u64>,
|
||||
notify: Notify,
|
||||
closed: AtomicBool,
|
||||
}
|
||||
|
||||
impl StreamFlowWindow {
|
||||
@@ -98,6 +109,7 @@ impl StreamFlowWindow {
|
||||
Self {
|
||||
available: Mutex::new(u64::from(initial)),
|
||||
notify: Notify::new(),
|
||||
closed: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -109,6 +121,12 @@ impl StreamFlowWindow {
|
||||
let requested = bytes as u64;
|
||||
let started_at = Instant::now();
|
||||
loop {
|
||||
let notified = self.notify.notified();
|
||||
tokio::pin!(notified);
|
||||
notified.as_mut().enable();
|
||||
if self.closed.load(Ordering::Acquire) {
|
||||
return Err(());
|
||||
}
|
||||
{
|
||||
let mut available = self.available.lock();
|
||||
if *available >= requested {
|
||||
@@ -120,10 +138,7 @@ impl StreamFlowWindow {
|
||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||
return Err(());
|
||||
};
|
||||
if tokio::time::timeout(remaining, self.notify.notified())
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
if tokio::time::timeout(remaining, notified).await.is_err() {
|
||||
return Err(());
|
||||
}
|
||||
}
|
||||
@@ -138,6 +153,11 @@ impl StreamFlowWindow {
|
||||
drop(available);
|
||||
self.notify.notify_waiters();
|
||||
}
|
||||
|
||||
fn close(&self) {
|
||||
self.closed.store(true, Ordering::Release);
|
||||
self.notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
@@ -209,6 +229,10 @@ impl BoundedOutbound {
|
||||
pub fn snapshot(&self) -> QueueSnapshot {
|
||||
self.tx.snapshot()
|
||||
}
|
||||
|
||||
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||
self.close_tx.subscribe()
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ProxyConn {
|
||||
@@ -231,6 +255,7 @@ pub struct ProxyConn {
|
||||
flow_window_blocked_ms: AtomicU64,
|
||||
write_latency_last_us: AtomicU64,
|
||||
write_latency_ewma_us: AtomicU64,
|
||||
settings: Mutex<protocol::SettingsPayload>,
|
||||
}
|
||||
|
||||
impl ProxyConn {
|
||||
@@ -244,6 +269,7 @@ impl ProxyConn {
|
||||
protocol_version: u8,
|
||||
) -> Self {
|
||||
Self {
|
||||
settings: Mutex::new(local_settings()),
|
||||
id,
|
||||
node_id,
|
||||
node_name,
|
||||
@@ -271,6 +297,11 @@ impl ProxyConn {
|
||||
self
|
||||
}
|
||||
|
||||
pub(super) fn with_settings(mut self, settings: protocol::SettingsPayload) -> Self {
|
||||
*self.settings.get_mut() = settings;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
||||
self.node_generation = tunnel_generation;
|
||||
self
|
||||
@@ -565,13 +596,6 @@ pub struct LocalResponseHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum LocalBodyEvent {
|
||||
Chunk(Bytes),
|
||||
End,
|
||||
Error(String),
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct LocalWaitState {
|
||||
response: Option<LocalResponseHead>,
|
||||
@@ -585,10 +609,11 @@ pub struct LocalStream {
|
||||
proxy_stream_id: u32,
|
||||
request_window: StreamFlowWindow,
|
||||
response_consumed_since_update: Mutex<u64>,
|
||||
min_window_update_bytes: u32,
|
||||
response_connection: Mutex<Option<std::sync::Weak<ProxyConn>>>,
|
||||
wait_state: Mutex<LocalWaitState>,
|
||||
headers_notify: Notify,
|
||||
body_tx: mpsc::Sender<LocalBodyEvent>,
|
||||
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
|
||||
body: Arc<ResponseBuffer>,
|
||||
terminal: AtomicBool,
|
||||
}
|
||||
|
||||
@@ -600,7 +625,6 @@ impl LocalStream {
|
||||
proxy_stream_id: u32,
|
||||
initial_window_bytes: u32,
|
||||
) -> Self {
|
||||
let (body_tx, body_rx) = mpsc::channel(128);
|
||||
Self {
|
||||
id,
|
||||
tunnel_generation,
|
||||
@@ -608,10 +632,11 @@ impl LocalStream {
|
||||
proxy_stream_id,
|
||||
request_window: StreamFlowWindow::new(initial_window_bytes),
|
||||
response_consumed_since_update: Mutex::new(0),
|
||||
min_window_update_bytes: (initial_window_bytes / 4).clamp(1, 1024 * 1024),
|
||||
response_connection: Mutex::new(None),
|
||||
wait_state: Mutex::new(LocalWaitState::default()),
|
||||
headers_notify: Notify::new(),
|
||||
body_tx,
|
||||
body_rx: Mutex::new(Some(body_rx)),
|
||||
body: ResponseBuffer::new(initial_window_bytes as usize),
|
||||
terminal: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
@@ -632,26 +657,51 @@ impl LocalStream {
|
||||
self.request_window.add(delta);
|
||||
}
|
||||
|
||||
fn response_window_update_delta(&self, bytes: usize) -> Option<u32> {
|
||||
if bytes == 0 {
|
||||
return None;
|
||||
async fn flush_response_credit(&self) -> Result<(), String> {
|
||||
if self.terminal.load(Ordering::Acquire) {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let mut consumed = self.response_consumed_since_update.lock();
|
||||
*consumed = consumed.saturating_add(bytes as u64);
|
||||
let threshold = u64::from(*STREAM_MIN_WINDOW_UPDATE_BYTES);
|
||||
if *consumed < threshold {
|
||||
return None;
|
||||
let connection = self
|
||||
.response_connection
|
||||
.lock()
|
||||
.as_ref()
|
||||
.and_then(std::sync::Weak::upgrade);
|
||||
let Some(connection) = connection else {
|
||||
return Ok(());
|
||||
};
|
||||
if connection.protocol_version() < 3 {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let delta = (*consumed).min(u64::from(u32::MAX)) as u32;
|
||||
*consumed = consumed.saturating_sub(u64::from(delta));
|
||||
Some(delta)
|
||||
let delta = {
|
||||
let consumed = self.response_consumed_since_update.lock();
|
||||
if *consumed < u64::from(self.min_window_update_bytes) {
|
||||
return Ok(());
|
||||
}
|
||||
(*consumed).min(u64::from(u32::MAX)) as u32
|
||||
};
|
||||
let frame = protocol::encode_window_update(self.proxy_stream_id, delta);
|
||||
if connection
|
||||
.send_wait(Message::Binary(frame.into()), OUTBOUND_BACKPRESSURE_TIMEOUT)
|
||||
.await
|
||||
== SendStatus::Queued
|
||||
{
|
||||
let mut consumed = self.response_consumed_since_update.lock();
|
||||
*consumed = consumed.saturating_sub(u64::from(delta));
|
||||
return Ok(());
|
||||
}
|
||||
if self.terminal.load(Ordering::Acquire) {
|
||||
return Ok(());
|
||||
}
|
||||
connection.request_close();
|
||||
Err("proxy flow-control update failed".to_string())
|
||||
}
|
||||
|
||||
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
|
||||
tokio::time::timeout(timeout, async {
|
||||
loop {
|
||||
let notified = self.headers_notify.notified();
|
||||
tokio::pin!(notified);
|
||||
notified.as_mut().enable();
|
||||
let outcome = {
|
||||
let state = self.wait_state.lock();
|
||||
if let Some(response) = &state.response {
|
||||
@@ -662,15 +712,20 @@ impl LocalStream {
|
||||
if let Some(error) = outcome {
|
||||
return Err(error);
|
||||
}
|
||||
self.headers_notify.notified().await;
|
||||
notified.await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.map_err(|_| "timed out waiting for response headers".to_string())?
|
||||
}
|
||||
|
||||
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
|
||||
self.body_rx.lock().take()
|
||||
pub fn take_body_receiver(self: &Arc<Self>) -> Option<LocalBodyReceiver> {
|
||||
self.body.take_receiver().map(|receiver| LocalBodyReceiver {
|
||||
receiver,
|
||||
stream: Arc::clone(self),
|
||||
failed: false,
|
||||
pending: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
|
||||
@@ -690,21 +745,11 @@ impl LocalStream {
|
||||
}
|
||||
}
|
||||
|
||||
async fn push_body_chunk(&self, payload: Bytes) -> bool {
|
||||
fn push_body_chunk(&self, payload: Bytes) -> bool {
|
||||
if self.terminal.load(Ordering::Acquire) {
|
||||
return false;
|
||||
}
|
||||
// Use a timeout to prevent a slow consumer from blocking the shared
|
||||
// proxy-connection reader (head-of-line blocking across streams).
|
||||
match tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
self.body_tx.send(LocalBodyEvent::Chunk(payload)),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(())) => true,
|
||||
_ => false,
|
||||
}
|
||||
self.body.push(payload)
|
||||
}
|
||||
|
||||
fn finish(&self) {
|
||||
@@ -722,7 +767,8 @@ impl LocalStream {
|
||||
if notify {
|
||||
self.headers_notify.notify_waiters();
|
||||
}
|
||||
let _ = self.body_tx.try_send(LocalBodyEvent::End);
|
||||
self.request_window.close();
|
||||
self.body.finish(Ok(()));
|
||||
}
|
||||
|
||||
fn fail(&self, error: impl Into<String>) {
|
||||
@@ -742,7 +788,38 @@ impl LocalStream {
|
||||
if notify {
|
||||
self.headers_notify.notify_waiters();
|
||||
}
|
||||
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
|
||||
self.request_window.close();
|
||||
self.body.finish(Err(error));
|
||||
}
|
||||
}
|
||||
|
||||
pub struct LocalBodyReceiver {
|
||||
receiver: BodyReceiver,
|
||||
stream: Arc<LocalStream>,
|
||||
failed: bool,
|
||||
pending: Option<LocalBodyEvent>,
|
||||
}
|
||||
|
||||
impl LocalBodyReceiver {
|
||||
pub async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||
if self.failed {
|
||||
return None;
|
||||
}
|
||||
if self.pending.is_none() {
|
||||
let event = self.receiver.recv().await?;
|
||||
if let LocalBodyEvent::Chunk(chunk) = &event {
|
||||
let mut consumed = self.stream.response_consumed_since_update.lock();
|
||||
*consumed = consumed.saturating_add(chunk.len() as u64);
|
||||
}
|
||||
self.pending = Some(event);
|
||||
}
|
||||
if matches!(self.pending, Some(LocalBodyEvent::Chunk(_))) {
|
||||
if let Err(error) = self.stream.flush_response_credit().await {
|
||||
self.failed = true;
|
||||
return Some(LocalBodyEvent::Error(error));
|
||||
}
|
||||
}
|
||||
self.pending.take()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -765,6 +842,21 @@ pub struct HubRouter {
|
||||
drain_reasons: Mutex<HashMap<String, u64>>,
|
||||
}
|
||||
|
||||
struct PendingStreamGuard<'router> {
|
||||
hub: &'router HubRouter,
|
||||
connection: &'router ProxyConn,
|
||||
stream_id: u64,
|
||||
committed: bool,
|
||||
}
|
||||
|
||||
impl Drop for PendingStreamGuard<'_> {
|
||||
fn drop(&mut self) {
|
||||
if !self.committed && self.hub.cleanup_local_stream(self.stream_id) {
|
||||
self.connection.release_stream();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct NodeStatusEvent {
|
||||
node_id: String,
|
||||
authenticated_key: Option<String>,
|
||||
@@ -1166,17 +1258,27 @@ impl HubRouter {
|
||||
|
||||
// Frames encoded successfully -- now register the stream.
|
||||
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
||||
let local_stream = Arc::new(LocalStream::new(
|
||||
let settings = proxy_conn.settings.lock().clone();
|
||||
let mut local_stream = LocalStream::new(
|
||||
local_stream_id,
|
||||
proxy_conn.node_generation.clone(),
|
||||
proxy_conn.id,
|
||||
proxy_stream_id,
|
||||
*STREAM_INITIAL_WINDOW_BYTES,
|
||||
));
|
||||
settings.initial_stream_window_bytes,
|
||||
);
|
||||
local_stream.min_window_update_bytes = settings.min_window_update_bytes;
|
||||
*local_stream.response_connection.get_mut() = Some(Arc::downgrade(&proxy_conn));
|
||||
let local_stream = Arc::new(local_stream);
|
||||
self.local_streams
|
||||
.insert(local_stream_id, local_stream.clone());
|
||||
self.proxy_to_local
|
||||
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
||||
let mut pending_stream = PendingStreamGuard {
|
||||
hub: self,
|
||||
connection: &proxy_conn,
|
||||
stream_id: local_stream_id,
|
||||
committed: false,
|
||||
};
|
||||
|
||||
let send_status = proxy_conn
|
||||
.send_wait(
|
||||
@@ -1195,10 +1297,11 @@ impl HubRouter {
|
||||
"open_local_stream dispatched"
|
||||
);
|
||||
match send_status {
|
||||
SendStatus::Queued => Ok(local_stream),
|
||||
SendStatus::Queued => {
|
||||
pending_stream.committed = true;
|
||||
Ok(local_stream)
|
||||
}
|
||||
SendStatus::Closed | SendStatus::Congested => {
|
||||
self.cleanup_local_stream(local_stream_id);
|
||||
proxy_conn.release_stream();
|
||||
Err("proxy connection congested".to_string())
|
||||
}
|
||||
}
|
||||
@@ -1243,7 +1346,9 @@ impl HubRouter {
|
||||
.map(|entry| entry.value().clone())
|
||||
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
||||
|
||||
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
|
||||
let chunk_size = MAX_REQUEST_BODY_FRAME_SIZE
|
||||
.min(proxy_conn.settings.lock().initial_stream_window_bytes as usize);
|
||||
let total_chunks = payload.len().div_ceil(chunk_size);
|
||||
let result = if total_chunks == 0 {
|
||||
if end_stream {
|
||||
self.send_request_body_frame(&proxy_conn, &stream, &[], true)
|
||||
@@ -1252,7 +1357,7 @@ impl HubRouter {
|
||||
Ok(())
|
||||
}
|
||||
} else {
|
||||
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
|
||||
for (index, chunk) in payload.chunks(chunk_size).enumerate() {
|
||||
let is_last_chunk = index + 1 == total_chunks;
|
||||
if let Err(error) = self
|
||||
.send_request_body_frame(
|
||||
@@ -1352,17 +1457,20 @@ impl HubRouter {
|
||||
} else {
|
||||
protocol::encode_stream_error(stream.proxy_stream_id, reason)
|
||||
};
|
||||
let _ = pc.send(Message::Binary(frame.into()));
|
||||
if pc.send(Message::Binary(frame.into())) != SendStatus::Queued {
|
||||
pc.request_close();
|
||||
}
|
||||
}
|
||||
stream.fail(reason.to_string());
|
||||
}
|
||||
|
||||
fn cleanup_local_stream(&self, local_stream_id: u64) {
|
||||
fn cleanup_local_stream(&self, local_stream_id: u64) -> bool {
|
||||
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
||||
return;
|
||||
return false;
|
||||
};
|
||||
self.proxy_to_local
|
||||
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
||||
true
|
||||
}
|
||||
|
||||
pub async fn handle_proxy_frame(self: &Arc<Self>, proxy_conn_id: u64, data: &mut [u8]) {
|
||||
@@ -1433,9 +1541,7 @@ impl HubRouter {
|
||||
.get(&proxy_conn_id)
|
||||
.map(|entry| entry.value().clone());
|
||||
if let Some(pc) = pc {
|
||||
let _ = pc
|
||||
.send_wait(Message::Binary(pong.into()), Duration::from_millis(250))
|
||||
.await;
|
||||
let _ = pc.send(Message::Binary(pong.into()));
|
||||
}
|
||||
}
|
||||
protocol::PONG => {}
|
||||
@@ -1509,11 +1615,36 @@ impl HubRouter {
|
||||
);
|
||||
}
|
||||
protocol::SETTINGS => {
|
||||
debug!(
|
||||
msg_type = header.msg_type,
|
||||
proxy_conn_id = proxy_conn_id,
|
||||
"received tunnel protocol v3 SETTINGS from proxy"
|
||||
);
|
||||
let settings = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::SettingsPayload>(&payload).ok()
|
||||
})
|
||||
.filter(|settings| settings.is_valid());
|
||||
if let Some(connection) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
if header.stream_id != 0 || header.flags != 0 {
|
||||
connection.request_close();
|
||||
return;
|
||||
}
|
||||
let Some(settings) = settings else {
|
||||
connection.request_close();
|
||||
return;
|
||||
};
|
||||
let local = local_settings();
|
||||
let settings = settings
|
||||
.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||
let mut current = connection.settings.lock();
|
||||
if connection.stream_count.load(Ordering::Acquire) > 0 && *current != settings {
|
||||
drop(current);
|
||||
connection.request_close();
|
||||
return;
|
||||
}
|
||||
*current = settings;
|
||||
}
|
||||
}
|
||||
protocol::WINDOW_UPDATE => {
|
||||
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
|
||||
@@ -1739,19 +1870,8 @@ impl HubRouter {
|
||||
None => return,
|
||||
};
|
||||
|
||||
let payload_len = payload.len();
|
||||
if !stream.push_body_chunk(Bytes::from(payload)).await {
|
||||
if !stream.push_body_chunk(Bytes::from(payload)) {
|
||||
self.cancel_local_stream(local_id, "local relay response congested");
|
||||
return;
|
||||
}
|
||||
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
if pc.protocol_version() >= 3 {
|
||||
if let Some(delta) = stream.response_window_update_delta(payload_len) {
|
||||
let frame = protocol::encode_window_update(header.stream_id, delta);
|
||||
let _ = pc.send(Message::Binary(frame.into()));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -9,14 +9,13 @@ use axum::body::{Body, Bytes};
|
||||
use axum::extract::{ConnectInfo, Path, Request, State};
|
||||
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||
use axum::response::IntoResponse;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::api::response::apply_streaming_response_headers;
|
||||
use crate::headers::should_skip_response_header;
|
||||
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
|
||||
|
||||
use super::hub::{LocalBodyEvent, LocalStream};
|
||||
use super::hub::{LocalBodyEvent, LocalBodyReceiver, LocalStream};
|
||||
use super::protocol;
|
||||
use super::{AppState, RelayRequestAuthenticated};
|
||||
|
||||
@@ -40,7 +39,7 @@ impl Drop for StreamGuard {
|
||||
pub(crate) struct DirectRelayResponse {
|
||||
status: u16,
|
||||
headers: Vec<(String, String)>,
|
||||
body_rx: mpsc::Receiver<LocalBodyEvent>,
|
||||
body_rx: LocalBodyReceiver,
|
||||
request_guard: StreamGuard,
|
||||
_request_permit: Option<AdmissionPermit>,
|
||||
}
|
||||
@@ -55,10 +54,13 @@ impl DirectRelayResponse {
|
||||
}
|
||||
|
||||
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, String> {
|
||||
if self.request_guard.finished {
|
||||
return Ok(None);
|
||||
}
|
||||
let event = self.body_rx.recv().await;
|
||||
match event {
|
||||
Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)),
|
||||
Some(LocalBodyEvent::End) | None => {
|
||||
Some(LocalBodyEvent::End) => {
|
||||
self.request_guard.finished = true;
|
||||
Ok(None)
|
||||
}
|
||||
@@ -66,6 +68,7 @@ impl DirectRelayResponse {
|
||||
self.request_guard.finished = true;
|
||||
Err(error)
|
||||
}
|
||||
None => Err("tunnel response ended without a terminal frame".to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -84,6 +87,11 @@ pub(crate) async fn open_direct_relay_stream(
|
||||
.open_authorized_local_stream(node_id, &meta)
|
||||
.await
|
||||
.map_err(|error| format!("connect: {error}"))?;
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
if let Err(error) = state
|
||||
.hub
|
||||
.push_local_request_body(stream.id, body, true)
|
||||
@@ -126,11 +134,7 @@ pub(crate) async fn open_direct_relay_stream(
|
||||
status: response_head.status,
|
||||
headers: response_head.headers,
|
||||
body_rx,
|
||||
request_guard: StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
},
|
||||
request_guard,
|
||||
_request_permit: request_permit,
|
||||
})
|
||||
}
|
||||
@@ -259,6 +263,11 @@ pub async fn relay_request(
|
||||
);
|
||||
}
|
||||
};
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
let body_stream = match spool.body_stream().await {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
@@ -306,12 +315,6 @@ pub async fn relay_request(
|
||||
);
|
||||
}
|
||||
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
|
||||
let wait_timeout = relay_header_timeout(&meta);
|
||||
let response_head = match stream.wait_headers(wait_timeout).await {
|
||||
Ok(response) => response,
|
||||
@@ -373,6 +376,9 @@ pub async fn relay_request(
|
||||
}
|
||||
}
|
||||
}
|
||||
if !guard.finished {
|
||||
yield Err(io::Error::other("tunnel response ended without a terminal frame"));
|
||||
}
|
||||
guard.finished = true;
|
||||
};
|
||||
|
||||
@@ -563,6 +569,100 @@ mod tests {
|
||||
request
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelled_relays_reset_streams_during_upload_and_header_wait() {
|
||||
for direct in [true, false] {
|
||||
for during_upload in [true, false] {
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||
sample_connected_proxy_node("node-123"),
|
||||
]));
|
||||
let data = Arc::new(
|
||||
GatewayDataState::with_proxy_node_repository_for_tests(repository)
|
||||
.with_system_config_values_for_tests(
|
||||
Vec::<(String, serde_json::Value)>::new(),
|
||||
)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
let state = test_app_state().with_data(data);
|
||||
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
let connection = Arc::new(
|
||||
ProxyConn::new(
|
||||
500,
|
||||
"node-123".into(),
|
||||
"Node 123".into(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
3,
|
||||
)
|
||||
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
|
||||
.with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string())
|
||||
.with_settings(protocol::SettingsPayload {
|
||||
initial_stream_window_bytes: 128,
|
||||
min_window_update_bytes: 32,
|
||||
drain_deadline_ms: 1000,
|
||||
}),
|
||||
);
|
||||
state.hub.register_proxy(Arc::clone(&connection));
|
||||
let meta = protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
method: "POST".into(),
|
||||
url: "https://example.com/".into(),
|
||||
headers: HashMap::new(),
|
||||
stream: true,
|
||||
request_timeout_ms: None,
|
||||
stream_first_byte_timeout_ms: None,
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
};
|
||||
let body = Bytes::from(vec![b'x'; if during_upload { 256 } else { 0 }]);
|
||||
let relay = tokio::spawn(async move {
|
||||
if direct {
|
||||
let _response =
|
||||
super::open_direct_relay_stream(&state, "node-123", meta, body)
|
||||
.await
|
||||
.unwrap();
|
||||
} else {
|
||||
let request =
|
||||
authenticated_request(encode_relay_envelope(&meta, &body)).await;
|
||||
let _response = relay_request(
|
||||
Path("node-123".into()),
|
||||
State(state),
|
||||
ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4242))),
|
||||
request,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
});
|
||||
recv_tunnel_test_frame(&mut proxy_rx, "request headers").await;
|
||||
recv_tunnel_test_frame(&mut proxy_rx, "request body").await;
|
||||
relay.abort();
|
||||
assert!(relay.await.unwrap_err().is_cancelled());
|
||||
let Message::Binary(frame) = recv_tunnel_test_frame(&mut proxy_rx, "reset").await
|
||||
else {
|
||||
panic!("expected binary reset frame")
|
||||
};
|
||||
let frame = aether_contracts::tunnel::Frame::decode(frame).unwrap();
|
||||
assert_eq!(
|
||||
frame.msg_type,
|
||||
aether_contracts::tunnel::MsgType::ResetStream
|
||||
);
|
||||
assert_eq!(
|
||||
connection
|
||||
.stream_count
|
||||
.load(std::sync::atomic::Ordering::Relaxed),
|
||||
0
|
||||
);
|
||||
assert!(connection.is_available());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
|
||||
let meta = protocol::RequestMeta {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
mod body;
|
||||
mod control_plane;
|
||||
mod hub;
|
||||
mod local_relay;
|
||||
|
||||
@@ -9,6 +9,7 @@ use aether_runtime::bounded_queue;
|
||||
use axum::extract::ws::{Message, WebSocket};
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use tokio::sync::watch;
|
||||
use tokio::task::JoinSet;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus};
|
||||
@@ -84,6 +85,35 @@ pub async fn handle_proxy_connection(
|
||||
|
||||
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
|
||||
let (close_tx, mut close_rx) = watch::channel(false);
|
||||
let settings = if protocol_version >= 3 {
|
||||
let Some(settings) = read_proxy_settings(
|
||||
&mut ws_tx,
|
||||
&mut ws_rx,
|
||||
security.as_deref(),
|
||||
protocol_version,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
warn!(conn_id, "proxy SETTINGS negotiation failed");
|
||||
return;
|
||||
};
|
||||
let local = super::hub::local_settings();
|
||||
let negotiated =
|
||||
settings.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||
let message = Message::Binary(protocol::encode_settings(&negotiated).into());
|
||||
let Ok(message) = encrypt_message(message, security.as_deref()) else {
|
||||
return;
|
||||
};
|
||||
if !matches!(
|
||||
tokio::time::timeout(PROXY_HELLO_TIMEOUT, ws_tx.send(message)).await,
|
||||
Ok(Ok(()))
|
||||
) {
|
||||
return;
|
||||
}
|
||||
negotiated
|
||||
} else {
|
||||
super::hub::local_settings()
|
||||
};
|
||||
let conn = ProxyConn::new(
|
||||
conn_id,
|
||||
node_id.clone(),
|
||||
@@ -93,7 +123,8 @@ pub async fn handle_proxy_connection(
|
||||
max_streams,
|
||||
protocol_version,
|
||||
)
|
||||
.with_tunnel_generation(node_generation);
|
||||
.with_tunnel_generation(node_generation)
|
||||
.with_settings(settings);
|
||||
let conn = match (security_key.clone(), management_token_credential) {
|
||||
(Some(key), None) => Arc::new(conn.with_authenticated_key(key)),
|
||||
(None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)),
|
||||
@@ -416,19 +447,22 @@ async fn run_proxy_reader(
|
||||
let idle_enabled = !idle_timeout.is_zero();
|
||||
let mut oversized_count = 0u32;
|
||||
let mut frames_received: u64 = 0;
|
||||
let mut close_rx = conn.outbound.subscribe_close();
|
||||
let mut heartbeats = JoinSet::new();
|
||||
loop {
|
||||
let msg = if idle_enabled {
|
||||
tokio::select! {
|
||||
msg = ws_rx.next() => msg,
|
||||
_ = tokio::time::sleep(idle_timeout) => {
|
||||
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
||||
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
|
||||
conn.request_close();
|
||||
break;
|
||||
}
|
||||
if conn.outbound.is_closing() {
|
||||
break;
|
||||
}
|
||||
while heartbeats.try_join_next().is_some() {}
|
||||
let msg = tokio::select! {
|
||||
biased;
|
||||
_ = close_rx.changed() => break,
|
||||
msg = ws_rx.next() => msg,
|
||||
_ = tokio::time::sleep(idle_timeout), if idle_enabled => {
|
||||
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
||||
conn.request_close();
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
ws_rx.next().await
|
||||
};
|
||||
|
||||
match msg {
|
||||
@@ -463,7 +497,27 @@ async fn run_proxy_reader(
|
||||
continue;
|
||||
}
|
||||
|
||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||
let is_heartbeat = protocol::FrameHeader::parse(&data)
|
||||
.is_some_and(|header| header.msg_type == protocol::HEARTBEAT_DATA);
|
||||
if is_heartbeat {
|
||||
if heartbeats.is_empty() {
|
||||
let heartbeat_hub = Arc::clone(&hub);
|
||||
let conn_id = conn.id;
|
||||
heartbeats.spawn(async move {
|
||||
if tokio::time::timeout(
|
||||
Duration::from_secs(10),
|
||||
heartbeat_hub.handle_proxy_frame(conn_id, &mut data),
|
||||
)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
warn!(conn_id, "proxy heartbeat processing timed out");
|
||||
}
|
||||
});
|
||||
}
|
||||
} else {
|
||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||
}
|
||||
}
|
||||
Some(Ok(Message::Close(_))) | None => {
|
||||
info!(
|
||||
@@ -489,6 +543,56 @@ async fn run_proxy_reader(
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
heartbeats.shutdown().await;
|
||||
}
|
||||
|
||||
async fn read_proxy_settings(
|
||||
ws_tx: &mut futures_util::stream::SplitSink<WebSocket, Message>,
|
||||
ws_rx: &mut futures_util::stream::SplitStream<WebSocket>,
|
||||
security: Option<&SecureFrameCodec>,
|
||||
protocol_version: u8,
|
||||
) -> Option<protocol::SettingsPayload> {
|
||||
tokio::time::timeout(PROXY_HELLO_TIMEOUT, async {
|
||||
let mut hello_received = security.is_some();
|
||||
for _ in 0..MAX_PREAUTH_PINGS {
|
||||
match ws_rx.next().await? {
|
||||
Ok(Message::Binary(data)) => {
|
||||
if data.len() > 256 * 1024 {
|
||||
return None;
|
||||
}
|
||||
let data = decrypt_message(data, security).ok()?;
|
||||
let frame = Frame::decode(data.into()).ok()?;
|
||||
if frame.stream_id != 0 || frame.flags != 0 {
|
||||
return None;
|
||||
}
|
||||
match frame.msg_type {
|
||||
MsgType::Hello if !hello_received => {
|
||||
let hello =
|
||||
serde_json::from_slice::<HelloPayload>(&frame.payload).ok()?;
|
||||
if hello.protocol_version != protocol_version {
|
||||
return None;
|
||||
}
|
||||
hello_received = true;
|
||||
}
|
||||
MsgType::Settings if hello_received => {
|
||||
let settings =
|
||||
serde_json::from_slice::<protocol::SettingsPayload>(&frame.payload)
|
||||
.ok()?;
|
||||
return settings.is_valid().then_some(settings);
|
||||
}
|
||||
_ => return None,
|
||||
}
|
||||
}
|
||||
Ok(Message::Ping(payload)) => ws_tx.send(Message::Pong(payload)).await.ok()?,
|
||||
Ok(Message::Pong(_)) => {}
|
||||
_ => return None,
|
||||
}
|
||||
}
|
||||
None
|
||||
})
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
}
|
||||
|
||||
fn encrypt_message(
|
||||
@@ -523,6 +627,158 @@ fn decrypt_message(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[cfg(feature = "testkit")]
|
||||
#[tokio::test]
|
||||
async fn slow_heartbeat_does_not_block_response_frames() {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use tokio_tungstenite::tungstenite::{client::IntoClientRequest, Message as ClientMessage};
|
||||
|
||||
let called = Arc::new(AtomicUsize::new(0));
|
||||
let callback_called = Arc::clone(&called);
|
||||
let control_plane = super::super::control_plane::ControlPlaneClient::local(
|
||||
move |_, _| {
|
||||
callback_called.fetch_add(1, Ordering::SeqCst);
|
||||
Box::pin(std::future::pending())
|
||||
},
|
||||
|_, _, _, _| Box::pin(async { Ok(()) }),
|
||||
);
|
||||
let data = crate::data::GatewayDataState::with_tunnel_management_auth_for_testkit(
|
||||
"heartbeat-test",
|
||||
"heartbeat-generation",
|
||||
"ae-tunnel-harness-management-token",
|
||||
aether_crypto::DEVELOPMENT_ENCRYPTION_KEY,
|
||||
)
|
||||
.unwrap();
|
||||
let state = super::super::AppState::new(
|
||||
control_plane,
|
||||
ConnConfig {
|
||||
ping_interval: Duration::from_secs(60),
|
||||
idle_timeout: Duration::ZERO,
|
||||
outbound_queue_capacity: 128,
|
||||
},
|
||||
16,
|
||||
)
|
||||
.with_data(Arc::new(data));
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let router = super::super::build_router_with_state(state.clone());
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(
|
||||
listener,
|
||||
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
});
|
||||
let mut request = format!("ws://{address}/api/internal/proxy-tunnel")
|
||||
.into_client_request()
|
||||
.unwrap();
|
||||
let headers = request.headers_mut();
|
||||
headers.insert("x-node-id", "heartbeat-test".parse().unwrap());
|
||||
headers.insert(
|
||||
aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER,
|
||||
"heartbeat-generation".parse().unwrap(),
|
||||
);
|
||||
headers.insert(
|
||||
aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER,
|
||||
"3".parse().unwrap(),
|
||||
);
|
||||
headers.insert(
|
||||
"authorization",
|
||||
"Bearer ae-tunnel-harness-management-token".parse().unwrap(),
|
||||
);
|
||||
let (mut websocket, _) = tokio_tungstenite::connect_async(request).await.unwrap();
|
||||
let hello = HelloPayload {
|
||||
protocol_version: 3,
|
||||
capabilities: vec![],
|
||||
session_id: None,
|
||||
replica_id: None,
|
||||
};
|
||||
websocket
|
||||
.send(ClientMessage::Binary(protocol::encode_hello(&hello).into()))
|
||||
.await
|
||||
.unwrap();
|
||||
websocket
|
||||
.send(ClientMessage::Binary(
|
||||
protocol::encode_settings(&super::super::hub::local_settings()).into(),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let ClientMessage::Binary(settings) = websocket.next().await.unwrap().unwrap() else {
|
||||
panic!("expected SETTINGS")
|
||||
};
|
||||
assert_eq!(Frame::decode(settings).unwrap().msg_type, MsgType::Settings);
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while !state.hub.has_local_proxy("heartbeat-test") {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let meta: protocol::RequestMeta = serde_json::from_value(serde_json::json!({
|
||||
"method": "GET", "url": "https://example.com", "headers": {}, "stream": true, "timeout": 10
|
||||
})).unwrap();
|
||||
let stream = state
|
||||
.hub
|
||||
.open_local_stream("heartbeat-test", &meta)
|
||||
.await
|
||||
.unwrap();
|
||||
let ClientMessage::Binary(request) = websocket.next().await.unwrap().unwrap() else {
|
||||
panic!("expected request headers")
|
||||
};
|
||||
let stream_id = Frame::decode(request).unwrap().stream_id;
|
||||
let heartbeat = Frame::control(
|
||||
MsgType::HeartbeatData,
|
||||
serde_json::to_vec(&serde_json::json!({"node_id": "heartbeat-test"})).unwrap(),
|
||||
);
|
||||
websocket
|
||||
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while called.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
for _ in 0..8 {
|
||||
websocket
|
||||
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let response = Frame::new(
|
||||
stream_id,
|
||||
MsgType::ResponseHeaders,
|
||||
0,
|
||||
serde_json::to_vec(&serde_json::json!({"status": 200, "headers": []})).unwrap(),
|
||||
);
|
||||
websocket
|
||||
.send(ClientMessage::Binary(response.encode()))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
stream
|
||||
.wait_headers(Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap()
|
||||
.status,
|
||||
200
|
||||
);
|
||||
assert_eq!(called.load(Ordering::SeqCst), 1);
|
||||
state.hub.request_close_all_proxies();
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while state.hub.has_local_proxy("heartbeat-test") {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
server.abort();
|
||||
let _ = server.await;
|
||||
}
|
||||
|
||||
use super::*;
|
||||
|
||||
const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
||||
|
||||
@@ -1308,7 +1308,13 @@ mod tests {
|
||||
stored[0].error_type.as_deref(),
|
||||
Some("stream_missing_terminal_event")
|
||||
);
|
||||
assert!(stored[0].error_message.is_none());
|
||||
assert_eq!(
|
||||
stored[0].error_message.as_deref(),
|
||||
Some(super::STREAM_MISSING_TERMINAL_EVENT_MESSAGE)
|
||||
);
|
||||
let mut public_candidate = stored[0].clone();
|
||||
public_candidate.sanitize_sensitive_diagnostics();
|
||||
assert!(public_candidate.error_message.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "aether-tunnel"
|
||||
version = "0.3.16"
|
||||
version = "0.3.17"
|
||||
edition = "2021"
|
||||
description = "Tunnel agent for Aether"
|
||||
|
||||
@@ -47,3 +47,4 @@ uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
aether-gateway = { workspace = true, features = ["testkit"] }
|
||||
tokio = { version = "1", features = ["test-util"] }
|
||||
|
||||
@@ -4,6 +4,15 @@ Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道
|
||||
|
||||
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
|
||||
|
||||
## 流式传输与升级注意事项
|
||||
|
||||
- 协议 v3 连接在 `HELLO` / `SETTINGS` 协商后才接收业务请求。实际双向流窗口取 gateway 与 agent 配置的较小值,信用更新阈值不超过该窗口的四分之一;单帧也不会超过协商窗口。
|
||||
- 响应缓冲按字节限额并合并小帧,结束和错误状态独立保存。慢消费者不会阻塞同一隧道其他流的读取;超出窗口或缓冲预算的流会被明确终止,不会静默截断。
|
||||
- 信用更新在消费数据后可靠入队;启用重定向重放时,进入有界重放缓存也视为请求体消费。持续无法投递关键控制帧时会关闭连接并向在途请求报告错误。
|
||||
- 客户端取消会终止对应上游请求,断连会回收 session 的 writer、heartbeat 和请求任务。正常 drain 在配置期限内继续处理已有流,期限到达后终止残留任务。
|
||||
- 建议先升级 gateway,再升级 agent。既有 v3 agent 已发送 `HELLO` / `SETTINGS`,可连接新 gateway;自定义 v3 节点必须完成这两步握手。协议 v1/v2 保留旧握手。与旧 gateway 混用时应保持默认窗口配置,不能依赖旧 gateway 应用新的窗口协商。
|
||||
- 自动重连恢复后续请求,不会自动续传已经输出的 SSE,也不会无条件重放已经发送的请求。
|
||||
|
||||
## 安装
|
||||
|
||||
`aether-tunnel` 会根据宿主机自动选择服务管理器:
|
||||
|
||||
@@ -764,6 +764,13 @@ impl Config {
|
||||
if self.tunnel_stream_initial_window_bytes == 0 {
|
||||
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
|
||||
}
|
||||
if u64::from(self.tunnel_stream_initial_window_bytes)
|
||||
> aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64
|
||||
{
|
||||
anyhow::bail!(
|
||||
"tunnel_stream_initial_window_bytes exceeds the maximum tunnel payload size"
|
||||
);
|
||||
}
|
||||
if self.tunnel_drain_deadline_ms == 0 {
|
||||
anyhow::bail!("tunnel_drain_deadline_ms must be > 0");
|
||||
}
|
||||
|
||||
@@ -203,19 +203,34 @@ pub async fn connect_and_run(
|
||||
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
|
||||
|
||||
// Spawn writer task (with WebSocket ping keepalive)
|
||||
let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
||||
let (frame_tx, writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
||||
ws_sink,
|
||||
ping_interval,
|
||||
Some(Arc::clone(&server.tunnel_metrics)),
|
||||
security.clone(),
|
||||
);
|
||||
let mut writer_handle = super::task::SessionTask::new(writer_handle);
|
||||
send_protocol_v3_hello(&frame_tx, &security_session, state).await;
|
||||
let drain_signal = spawn_drain_signal(
|
||||
let (session_drain_tx, session_drain_rx) = watch::channel(*drain.borrow());
|
||||
let forward_drain_tx = session_drain_tx.clone();
|
||||
let mut external_drain = drain;
|
||||
let forward_drain = super::task::SessionTask::new(tokio::spawn(async move {
|
||||
loop {
|
||||
if *external_drain.borrow() {
|
||||
let _ = forward_drain_tx.send(true);
|
||||
break;
|
||||
}
|
||||
if external_drain.changed().await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}));
|
||||
let drain_signal = super::task::SessionTask::new(spawn_drain_signal(
|
||||
conn_idx,
|
||||
frame_tx.clone(),
|
||||
drain.clone(),
|
||||
session_drain_rx.clone(),
|
||||
state.config.tunnel_drain_deadline_ms,
|
||||
);
|
||||
));
|
||||
|
||||
// Spawn heartbeat task (only for primary connection to avoid
|
||||
// resetting shared atomic metrics via swap(0))
|
||||
@@ -237,16 +252,19 @@ pub async fn connect_and_run(
|
||||
// ensures we detect this and trigger a reconnect promptly.
|
||||
let state_clone = Arc::clone(state);
|
||||
let server_clone = Arc::clone(server);
|
||||
let outcome = tokio::select! {
|
||||
result = dispatcher::run_with_security(
|
||||
let outcome = {
|
||||
let dispatch = dispatcher::run_with_security(
|
||||
state_clone,
|
||||
server_clone,
|
||||
ws_read,
|
||||
frame_tx.clone(),
|
||||
hb_handle,
|
||||
drain.clone(),
|
||||
session_drain_rx,
|
||||
security.clone(),
|
||||
) => {
|
||||
);
|
||||
tokio::pin!(dispatch);
|
||||
tokio::select! {
|
||||
result = &mut dispatch => {
|
||||
match result {
|
||||
Ok(()) => Ok(TunnelOutcome::Disconnected),
|
||||
Err(e) => {
|
||||
@@ -258,6 +276,8 @@ pub async fn connect_and_run(
|
||||
}
|
||||
}
|
||||
writer_result = &mut writer_handle => {
|
||||
frame_tx.close();
|
||||
let _ = tokio::time::timeout(Duration::from_secs(1), &mut dispatch).await;
|
||||
match writer_result {
|
||||
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
|
||||
Err(e) => {
|
||||
@@ -278,24 +298,37 @@ pub async fn connect_and_run(
|
||||
}
|
||||
_ = shutdown.changed() => {
|
||||
debug!("shutdown during tunnel dispatch");
|
||||
let _ = session_drain_tx.send(true);
|
||||
let deadline = Duration::from_millis(state.config.tunnel_drain_deadline_ms).saturating_add(Duration::from_secs(1));
|
||||
let _ = tokio::time::timeout(deadline, &mut dispatch).await;
|
||||
Ok(TunnelOutcome::Shutdown)
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Drop our sender; the writer will exit once all stream handler clones
|
||||
// are also dropped (i.e. after they finish their in-flight work).
|
||||
drop(frame_tx);
|
||||
forward_drain.abort();
|
||||
let _ = forward_drain.await;
|
||||
if !drain_signal.is_finished() {
|
||||
drain_signal.abort();
|
||||
let _ = drain_signal.await;
|
||||
}
|
||||
|
||||
// Wait for the writer task to finish with a generous timeout — the
|
||||
// dispatcher already waits up to 30s for stream handlers, so 35s here
|
||||
// covers that plus a small margin.
|
||||
// Skip if the writer already exited (the select branch that fired).
|
||||
if !writer_handle.is_finished() {
|
||||
let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await;
|
||||
let flush_timeout = if *session_drain_tx.borrow() {
|
||||
Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||
} else {
|
||||
Duration::from_secs(1)
|
||||
};
|
||||
if tokio::time::timeout(flush_timeout, &mut writer_handle)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
writer_handle.abort();
|
||||
let _ = writer_handle.await;
|
||||
}
|
||||
}
|
||||
|
||||
let connected_for = connected_at.elapsed();
|
||||
|
||||
@@ -9,7 +9,7 @@ use std::time::Duration;
|
||||
use bytes::Bytes;
|
||||
use futures_util::StreamExt;
|
||||
use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore};
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio::task::{AbortHandle, JoinSet};
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
use tracing::{debug, error, info, warn};
|
||||
|
||||
@@ -41,13 +41,25 @@ impl AsRef<[u8]> for BudgetedFramePayload {
|
||||
enum StreamDispatchStatus {
|
||||
Delivered,
|
||||
Closed,
|
||||
TimedOut,
|
||||
Congested,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct StreamDispatchTarget {
|
||||
body_tx: mpsc::Sender<Frame>,
|
||||
response_window: Arc<StreamSendWindow>,
|
||||
handler: Option<AbortHandle>,
|
||||
}
|
||||
|
||||
struct StreamCompletion {
|
||||
stream_id: u32,
|
||||
finished_tx: mpsc::UnboundedSender<u32>,
|
||||
}
|
||||
|
||||
impl Drop for StreamCompletion {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.finished_tx.send(self.stream_id);
|
||||
}
|
||||
}
|
||||
|
||||
/// A request stream is identified by a non-zero id and may only be opened
|
||||
@@ -109,7 +121,7 @@ where
|
||||
// reopen the same id and bypass the stream admission limit.
|
||||
let mut active_handler_ids: HashSet<u32> = HashSet::new();
|
||||
// Track spawned stream handlers so we can wait for them on shutdown
|
||||
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
|
||||
let mut handler_handles = JoinSet::new();
|
||||
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
|
||||
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
|
||||
let mut frames_since_cleanup: u32 = 0;
|
||||
@@ -121,30 +133,48 @@ where
|
||||
// Track last time we received any data to detect stale connections
|
||||
let mut last_data_at = tokio::time::Instant::now();
|
||||
let mut draining = *drain.borrow();
|
||||
let mut drain_open = true;
|
||||
let mut drain_deadline = draining.then(|| {
|
||||
tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||
});
|
||||
let mut initial_window_bytes = state.config.tunnel_stream_initial_window_bytes;
|
||||
let mut close_rx = frame_tx.subscribe_close();
|
||||
|
||||
let read_err = loop {
|
||||
if *close_rx.borrow() {
|
||||
break None;
|
||||
}
|
||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||
info!("tunnel drained after in-flight streams completed");
|
||||
break None;
|
||||
}
|
||||
|
||||
let msg_result = tokio::select! {
|
||||
_ = close_rx.changed() => break None,
|
||||
msg = ws_stream.next() => {
|
||||
match msg {
|
||||
Some(r) => r,
|
||||
None => break None,
|
||||
}
|
||||
}
|
||||
changed = drain.changed() => {
|
||||
changed = drain.changed(), if drain_open => {
|
||||
if changed.is_err() {
|
||||
drain_open = false;
|
||||
continue;
|
||||
}
|
||||
if *drain.borrow() {
|
||||
info!("tunnel drain requested, waiting for in-flight streams");
|
||||
draining = true;
|
||||
drain_deadline.get_or_insert_with(|| tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms));
|
||||
}
|
||||
continue;
|
||||
}
|
||||
_ = async {
|
||||
match drain_deadline {
|
||||
Some(deadline) => tokio::time::sleep_until(deadline).await,
|
||||
None => std::future::pending().await,
|
||||
}
|
||||
} => break None,
|
||||
finished = handler_finished_rx.recv() => {
|
||||
if let Some(stream_id) = finished {
|
||||
active_handler_ids.remove(&stream_id);
|
||||
@@ -238,20 +268,7 @@ where
|
||||
continue;
|
||||
}
|
||||
if draining {
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from("tunnel draining"),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped during drain"
|
||||
);
|
||||
}
|
||||
try_send_stream_error(&frame_tx, frame.stream_id, "tunnel draining");
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -263,6 +280,11 @@ where
|
||||
Ok(p) => p,
|
||||
Err(e) => {
|
||||
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"invalid request metadata",
|
||||
);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
@@ -270,21 +292,11 @@ where
|
||||
Ok(m) => m,
|
||||
Err(e) => {
|
||||
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
|
||||
// Use try_send to avoid blocking the read loop
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from(format!("invalid request metadata: {e}")),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped"
|
||||
);
|
||||
}
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"invalid request metadata",
|
||||
);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
@@ -294,33 +306,27 @@ where
|
||||
stream_id = frame.stream_id,
|
||||
"max concurrent streams reached"
|
||||
);
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from("max concurrent streams reached"),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped"
|
||||
);
|
||||
}
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"max concurrent streams reached",
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Create body channel and spawn handler
|
||||
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
|
||||
let response_window = Arc::new(StreamSendWindow::new(
|
||||
state.config.tunnel_stream_initial_window_bytes,
|
||||
));
|
||||
let body_capacity = (initial_window_bytes as usize)
|
||||
.div_ceil(32 * 1024)
|
||||
.saturating_add(1)
|
||||
.max(64);
|
||||
let (body_tx, body_rx) = mpsc::channel::<Frame>(body_capacity);
|
||||
let response_window = Arc::new(StreamSendWindow::new(initial_window_bytes));
|
||||
streams.insert(
|
||||
frame.stream_id,
|
||||
StreamDispatchTarget {
|
||||
body_tx,
|
||||
response_window: Arc::clone(&response_window),
|
||||
handler: None,
|
||||
},
|
||||
);
|
||||
active_handler_ids.insert(frame.stream_id);
|
||||
@@ -329,9 +335,13 @@ where
|
||||
let state_clone = Arc::clone(&state);
|
||||
let server_clone = Arc::clone(&server);
|
||||
let tx_clone = frame_tx.clone();
|
||||
let finished_tx = handler_finished_tx.clone();
|
||||
let sid = frame.stream_id;
|
||||
let handle = tokio::spawn(async move {
|
||||
let completion = StreamCompletion {
|
||||
stream_id: sid,
|
||||
finished_tx: handler_finished_tx.clone(),
|
||||
};
|
||||
let handle = handler_handles.spawn(async move {
|
||||
let _completion = completion;
|
||||
stream_handler::handle_stream(
|
||||
state_clone,
|
||||
server_clone,
|
||||
@@ -342,9 +352,8 @@ where
|
||||
response_window,
|
||||
)
|
||||
.await;
|
||||
let _ = finished_tx.send(sid);
|
||||
});
|
||||
handler_handles.push(handle);
|
||||
streams.get_mut(&sid).expect("new stream exists").handler = Some(handle);
|
||||
|
||||
if request_headers_end_stream {
|
||||
if let Some(target) = streams.get(&sid) {
|
||||
@@ -365,19 +374,21 @@ where
|
||||
let is_end = frame.is_end_stream();
|
||||
let sid = frame.stream_id;
|
||||
let dispatch = dispatch_stream_frame(&target.body_tx, frame).await;
|
||||
if dispatch != StreamDispatchStatus::Delivered {
|
||||
streams.remove(&sid);
|
||||
if dispatch == StreamDispatchStatus::TimedOut {
|
||||
server.tunnel_metrics.record_error(
|
||||
"stream_dispatch_timeout",
|
||||
&format!("request body dispatch timed out for stream {}", sid),
|
||||
);
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
sid,
|
||||
"tunnel request body dispatch stalled",
|
||||
);
|
||||
if dispatch == StreamDispatchStatus::Congested {
|
||||
if let Some(target) = streams.remove(&sid) {
|
||||
if let Some(handler) = target.handler {
|
||||
handler.abort();
|
||||
}
|
||||
}
|
||||
server.tunnel_metrics.record_error(
|
||||
"stream_dispatch_timeout",
|
||||
&format!("request body dispatch congested for stream {}", sid),
|
||||
);
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
sid,
|
||||
"tunnel request body dispatch stalled",
|
||||
);
|
||||
if is_end && draining && streams.is_empty() && active_handler_ids.is_empty()
|
||||
{
|
||||
info!("tunnel drained after request body completion");
|
||||
@@ -387,10 +398,29 @@ where
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::StreamEnd | MsgType::StreamError | MsgType::ResetStream => {
|
||||
MsgType::StreamEnd => {
|
||||
if let Some(target) = streams.get(&frame.stream_id) {
|
||||
if dispatch_stream_frame(&target.body_tx, frame.clone()).await
|
||||
== StreamDispatchStatus::Congested
|
||||
{
|
||||
if let Some(handler) = &target.handler {
|
||||
handler.abort();
|
||||
}
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"tunnel request body dispatch stalled",
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::StreamError | MsgType::ResetStream => {
|
||||
// Client-side cancellation or end
|
||||
if let Some(target) = streams.remove(&frame.stream_id) {
|
||||
let _ = dispatch_stream_frame(&target.body_tx, frame).await;
|
||||
if let Some(handler) = target.handler {
|
||||
handler.abort();
|
||||
}
|
||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||
info!("tunnel drained after stream termination");
|
||||
break None;
|
||||
@@ -409,12 +439,22 @@ where
|
||||
}
|
||||
|
||||
MsgType::HeartbeatAck => {
|
||||
heartbeat.on_ack(frame.payload).await;
|
||||
heartbeat.on_ack(frame.payload);
|
||||
}
|
||||
|
||||
MsgType::GoAway => {
|
||||
info!("received GOAWAY");
|
||||
break None;
|
||||
draining = true;
|
||||
let deadline_ms =
|
||||
serde_json::from_slice::<aether_contracts::tunnel::GoAwayPayload>(
|
||||
&frame.payload,
|
||||
)
|
||||
.map(|payload| payload.drain_deadline_ms)
|
||||
.unwrap_or(state.config.tunnel_drain_deadline_ms)
|
||||
.min(state.config.tunnel_drain_deadline_ms);
|
||||
drain_deadline.get_or_insert_with(|| {
|
||||
tokio::time::Instant::now() + Duration::from_millis(deadline_ms)
|
||||
});
|
||||
}
|
||||
|
||||
MsgType::WindowUpdate => {
|
||||
@@ -433,7 +473,31 @@ where
|
||||
);
|
||||
}
|
||||
|
||||
MsgType::Hello | MsgType::Settings | MsgType::LoadReport => {
|
||||
MsgType::Settings => {
|
||||
if frame.stream_id != 0 || frame.flags != 0 {
|
||||
break None;
|
||||
}
|
||||
let settings = serde_json::from_slice::<aether_contracts::tunnel::SettingsPayload>(
|
||||
&frame.payload,
|
||||
)
|
||||
.ok()
|
||||
.filter(|settings| settings.is_valid());
|
||||
let Some(settings) = settings else {
|
||||
warn!("invalid tunnel SETTINGS");
|
||||
break None;
|
||||
};
|
||||
if !streams.is_empty()
|
||||
&& settings.initial_stream_window_bytes != initial_window_bytes
|
||||
{
|
||||
warn!("tunnel SETTINGS changed with active streams");
|
||||
break None;
|
||||
}
|
||||
initial_window_bytes = settings
|
||||
.initial_stream_window_bytes
|
||||
.min(state.config.tunnel_stream_initial_window_bytes);
|
||||
}
|
||||
|
||||
MsgType::Hello | MsgType::LoadReport => {
|
||||
debug!(
|
||||
msg_type = ?frame.msg_type,
|
||||
stream_id = frame.stream_id,
|
||||
@@ -455,7 +519,7 @@ where
|
||||
// Trigger every 64 frames OR when the count exceeds max_streams.
|
||||
frames_since_cleanup += 1;
|
||||
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
|
||||
handler_handles.retain(|h| !h.is_finished());
|
||||
while handler_handles.try_join_next().is_some() {}
|
||||
frames_since_cleanup = 0;
|
||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||
info!("tunnel drained after cleanup");
|
||||
@@ -467,9 +531,7 @@ where
|
||||
// Drop body senders so stream handlers waiting on body_rx will unblock
|
||||
streams.clear();
|
||||
|
||||
// Wait for active stream handlers to finish so their frame_tx clones
|
||||
// are dropped before the writer closes the sink.
|
||||
drain_handlers(handler_handles).await;
|
||||
handler_handles.shutdown().await;
|
||||
|
||||
match read_err {
|
||||
Some(e) => Err(e.into()),
|
||||
@@ -478,30 +540,13 @@ where
|
||||
}
|
||||
|
||||
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
|
||||
let stream_id = frame.stream_id;
|
||||
let dispatched = tokio::time::timeout(stream_frame_dispatch_timeout(), async {
|
||||
let frame = attach_request_body_queue_budget(frame).await?;
|
||||
tx.send(frame).await.ok()?;
|
||||
Some(())
|
||||
})
|
||||
.await;
|
||||
match dispatched {
|
||||
Ok(Some(())) => StreamDispatchStatus::Delivered,
|
||||
Ok(None) => {
|
||||
warn!(
|
||||
stream_id,
|
||||
"stream handler channel or request body budget closed while dispatching tunnel frame"
|
||||
);
|
||||
StreamDispatchStatus::Closed
|
||||
}
|
||||
Err(_) => {
|
||||
warn!(
|
||||
stream_id,
|
||||
timeout_ms = stream_frame_dispatch_timeout().as_millis(),
|
||||
"stream handler channel blocked while dispatching tunnel frame"
|
||||
);
|
||||
StreamDispatchStatus::TimedOut
|
||||
}
|
||||
let Some(frame) = attach_request_body_queue_budget(frame).await else {
|
||||
return StreamDispatchStatus::Congested;
|
||||
};
|
||||
match tx.try_send(frame) {
|
||||
Ok(()) => StreamDispatchStatus::Delivered,
|
||||
Err(mpsc::error::TrySendError::Closed(_)) => StreamDispatchStatus::Closed,
|
||||
Err(mpsc::error::TrySendError::Full(_)) => StreamDispatchStatus::Congested,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -523,7 +568,7 @@ async fn attach_request_body_queue_budget_with(
|
||||
return Some(frame);
|
||||
}
|
||||
let permits = request_body_queue_permits(&frame, budget_bytes)?;
|
||||
let permit = budget.acquire_many_owned(permits).await.ok()?;
|
||||
let permit = budget.try_acquire_many_owned(permits).ok()?;
|
||||
frame.payload = Bytes::from_owner(BudgetedFramePayload {
|
||||
bytes: frame.payload,
|
||||
_permit: permit,
|
||||
@@ -549,20 +594,6 @@ fn request_body_queue_permits(frame: &Frame, budget_bytes: usize) -> Option<u32>
|
||||
u32::try_from(retained_bytes).ok()
|
||||
}
|
||||
|
||||
/// Bound how long a single stream handler is allowed to block the shared
|
||||
/// WebSocket read loop while receiving request-body frames.
|
||||
fn stream_frame_dispatch_timeout() -> Duration {
|
||||
#[cfg(test)]
|
||||
{
|
||||
Duration::from_millis(25)
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
{
|
||||
Duration::from_millis(500)
|
||||
}
|
||||
}
|
||||
|
||||
fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) {
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
@@ -573,6 +604,7 @@ fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'stat
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
frame_tx.close();
|
||||
warn!(
|
||||
stream_id,
|
||||
"writer channel full, StreamError dropped while aborting stalled stream"
|
||||
@@ -587,21 +619,6 @@ fn prune_closed_stream_senders(streams: &mut HashMap<u32, StreamDispatchTarget>)
|
||||
before.saturating_sub(streams.len())
|
||||
}
|
||||
|
||||
/// Wait for all active stream handlers to finish (with a timeout).
|
||||
async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
|
||||
if handles.is_empty() {
|
||||
return;
|
||||
}
|
||||
let count = handles.len();
|
||||
debug!(count, "waiting for active stream handlers to finish");
|
||||
let _ = tokio::time::timeout(Duration::from_secs(30), async {
|
||||
for h in handles {
|
||||
let _ = h.await;
|
||||
}
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -633,7 +650,7 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
stalled_send.await.expect("dispatch task should join"),
|
||||
StreamDispatchStatus::TimedOut
|
||||
StreamDispatchStatus::Congested
|
||||
);
|
||||
|
||||
let retained = rx
|
||||
@@ -737,6 +754,7 @@ mod tests {
|
||||
StreamDispatchTarget {
|
||||
body_tx: closed_tx,
|
||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||
handler: None,
|
||||
},
|
||||
),
|
||||
(
|
||||
@@ -744,6 +762,7 @@ mod tests {
|
||||
StreamDispatchTarget {
|
||||
body_tx: open_tx,
|
||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||
handler: None,
|
||||
},
|
||||
),
|
||||
]);
|
||||
@@ -763,6 +782,7 @@ mod tests {
|
||||
StreamDispatchTarget {
|
||||
body_tx: tx,
|
||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||
handler: None,
|
||||
},
|
||||
)]);
|
||||
let mut active_handler_ids = HashSet::from([7]);
|
||||
|
||||
@@ -31,14 +31,22 @@ enum AckDecision {
|
||||
}
|
||||
|
||||
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
||||
#[derive(Clone)]
|
||||
pub struct HeartbeatHandle {
|
||||
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
|
||||
task: Option<tokio::task::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl HeartbeatHandle {
|
||||
pub async fn on_ack(&self, payload: Bytes) {
|
||||
let _ = self.ack_tx.send(payload).await;
|
||||
pub fn on_ack(&self, payload: Bytes) {
|
||||
let _ = self.ack_tx.try_send(payload);
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for HeartbeatHandle {
|
||||
fn drop(&mut self) {
|
||||
if let Some(task) = self.task.take() {
|
||||
task.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -48,7 +56,7 @@ impl HeartbeatHandle {
|
||||
pub fn spawn_noop() -> HeartbeatHandle {
|
||||
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
|
||||
// receiver is immediately dropped; on_ack() calls will silently fail
|
||||
HeartbeatHandle { ack_tx }
|
||||
HeartbeatHandle { ack_tx, task: None }
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
@@ -74,7 +82,7 @@ pub fn spawn(
|
||||
) -> HeartbeatHandle {
|
||||
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
|
||||
|
||||
tokio::spawn(async move {
|
||||
let task = tokio::spawn(async move {
|
||||
// Read initial interval from dynamic config (may be updated by remote config).
|
||||
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
||||
let mut current_interval = initial_interval;
|
||||
@@ -151,7 +159,8 @@ pub fn spawn(
|
||||
current_interval = new_interval;
|
||||
}
|
||||
}
|
||||
Some(ack_payload) = ack_rx.recv() => {
|
||||
ack_payload = ack_rx.recv() => {
|
||||
let Some(ack_payload) = ack_payload else { break; };
|
||||
match handle_ack(&server, &ack_payload) {
|
||||
AckDecision::Accept {
|
||||
heartbeat_id: ack_id,
|
||||
@@ -179,7 +188,10 @@ pub fn spawn(
|
||||
}
|
||||
});
|
||||
|
||||
HeartbeatHandle { ack_tx }
|
||||
HeartbeatHandle {
|
||||
ack_tx,
|
||||
task: Some(task),
|
||||
}
|
||||
}
|
||||
|
||||
async fn build_heartbeat_payload(
|
||||
|
||||
@@ -3,6 +3,7 @@ pub mod dispatcher;
|
||||
pub mod heartbeat;
|
||||
pub mod protocol;
|
||||
pub mod stream_handler;
|
||||
mod task;
|
||||
pub mod writer;
|
||||
|
||||
use std::sync::Arc;
|
||||
@@ -332,9 +333,9 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
||||
tokio::time::sleep(Duration::from_millis(200)).await;
|
||||
gateway_handle.abort();
|
||||
let _ = (&mut gateway_handle).await;
|
||||
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
||||
|
||||
let (_restarted_gateway_state, restarted_gateway_handle) =
|
||||
start_gateway_on_port_retry(gateway_port)
|
||||
@@ -349,6 +350,7 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(server.tunnel_metrics.snapshot().connect_successes >= 2);
|
||||
let _ = shutdown_tx.send(true);
|
||||
tokio::time::timeout(Duration::from_secs(5), tunnel_task)
|
||||
.await
|
||||
@@ -380,7 +382,17 @@ mod tests {
|
||||
gateway_base_url: &str,
|
||||
node_id: &str,
|
||||
) -> Option<(StatusCode, String)> {
|
||||
let payload = relay_probe_envelope();
|
||||
let response = relay_response(gateway_base_url, node_id, relay_probe_envelope()).await?;
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
Some((status, body))
|
||||
}
|
||||
|
||||
async fn relay_response(
|
||||
gateway_base_url: &str,
|
||||
node_id: &str,
|
||||
payload: Vec<u8>,
|
||||
) -> Option<reqwest::Response> {
|
||||
let timestamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.expect("test clock should be after epoch")
|
||||
@@ -398,7 +410,7 @@ mod tests {
|
||||
&nonce,
|
||||
&digest,
|
||||
);
|
||||
let response = reqwest::Client::new()
|
||||
reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
|
||||
))
|
||||
@@ -421,10 +433,7 @@ mod tests {
|
||||
.body(payload)
|
||||
.send()
|
||||
.await
|
||||
.ok()?;
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
Some((status, body))
|
||||
.ok()
|
||||
}
|
||||
|
||||
fn relay_probe_envelope() -> Vec<u8> {
|
||||
@@ -456,30 +465,150 @@ mod tests {
|
||||
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
|
||||
// The embedded gateway now fails closed when relay authentication is
|
||||
// not configured. Keep this integration fixture explicitly authenticated.
|
||||
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
||||
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
||||
std::env::set_var(
|
||||
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
||||
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
||||
);
|
||||
std::env::set_var(
|
||||
"AETHER_GATEWAY_INSTANCE_ID",
|
||||
"tunnel-reconnect-test-gateway",
|
||||
);
|
||||
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
||||
aether_gateway::configure_test_tunnel_security(
|
||||
&mut state,
|
||||
"node-recovery",
|
||||
"test-generation-1",
|
||||
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
||||
);
|
||||
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
||||
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
||||
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
|
||||
let state = {
|
||||
let _guard = ENV_LOCK.lock().unwrap();
|
||||
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
||||
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
||||
std::env::set_var(
|
||||
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
||||
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
||||
);
|
||||
std::env::set_var(
|
||||
"AETHER_GATEWAY_INSTANCE_ID",
|
||||
"tunnel-reconnect-test-gateway",
|
||||
);
|
||||
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
||||
aether_gateway::configure_test_tunnel_security(
|
||||
&mut state,
|
||||
"node-recovery",
|
||||
"test-generation-1",
|
||||
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
||||
);
|
||||
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
||||
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
||||
state
|
||||
};
|
||||
let router = build_router_with_state(state.clone());
|
||||
let handle = spawn_router_on_port(port, router).await?;
|
||||
Ok((state, handle))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn negotiated_small_window_streams_large_responses_and_cancels_idle_upstream() {
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::routing::get;
|
||||
use futures_util::StreamExt;
|
||||
|
||||
ensure_rustls_provider();
|
||||
let upstream_port = reserve_local_port().unwrap();
|
||||
let upstream = Router::new()
|
||||
.route(
|
||||
"/large",
|
||||
get(|| async { Body::from(vec![b'x'; 2 * 1024 * 1024]) }),
|
||||
)
|
||||
.route(
|
||||
"/idle",
|
||||
get(|| async {
|
||||
let first = futures_util::stream::once(async {
|
||||
Ok::<_, std::io::Error>(Bytes::from_static(b"data: started\n\n"))
|
||||
});
|
||||
(
|
||||
[("content-type", "text/event-stream")],
|
||||
Body::from_stream(first.chain(futures_util::stream::pending())),
|
||||
)
|
||||
}),
|
||||
);
|
||||
let upstream_task = super::task::SessionTask::new(
|
||||
spawn_router_on_port(upstream_port, upstream).await.unwrap(),
|
||||
);
|
||||
let gateway_port = reserve_local_port().unwrap();
|
||||
let gateway_url = format!("http://127.0.0.1:{gateway_port}");
|
||||
let (_, gateway_task) = start_gateway_on_port(gateway_port).await.unwrap();
|
||||
let gateway_task = super::task::SessionTask::new(gateway_task);
|
||||
let mut config = sample_config(&gateway_url);
|
||||
config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired;
|
||||
config.tunnel_encryption_key = Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".into());
|
||||
config.tunnel_stream_initial_window_bytes = 512 * 1024;
|
||||
config.tunnel_drain_deadline_ms = 100;
|
||||
config.allow_private_targets = true;
|
||||
config.allowed_ports.push(upstream_port);
|
||||
let state = sample_state(config);
|
||||
let server = sample_server(&state, "node-recovery");
|
||||
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||
let (_drain_tx, drain_rx) = watch::channel(false);
|
||||
let tunnel_task = super::task::SessionTask::new(tokio::spawn({
|
||||
let state = Arc::clone(&state);
|
||||
let server = Arc::clone(&server);
|
||||
async move {
|
||||
run(&state, &server, 0, shutdown_rx, drain_rx).await;
|
||||
}
|
||||
}));
|
||||
wait_until_relay_status(&gateway_url, "node-recovery", StatusCode::GATEWAY_TIMEOUT).await;
|
||||
|
||||
let envelope = |path: &str| {
|
||||
let mut meta: protocol::RequestMeta =
|
||||
serde_json::from_slice(&relay_probe_envelope()[4..]).unwrap();
|
||||
meta.url = format!("http://127.0.0.1:{upstream_port}/{path}");
|
||||
meta.stream = true;
|
||||
meta.timeout = 10;
|
||||
meta.stream_first_byte_timeout_ms = Some(10_000);
|
||||
let encoded = serde_json::to_vec(&meta).unwrap();
|
||||
let mut result = (encoded.len() as u32).to_be_bytes().to_vec();
|
||||
result.extend(encoded);
|
||||
result
|
||||
};
|
||||
let response = relay_response(&gateway_url, "node-recovery", envelope("large"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = tokio::time::timeout(Duration::from_secs(10), response.bytes())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(body.len(), 2 * 1024 * 1024);
|
||||
assert!(body.iter().all(|byte| *byte == b'x'));
|
||||
|
||||
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
response.chunk().await.unwrap().unwrap(),
|
||||
"data: started\n\n"
|
||||
);
|
||||
drop(response);
|
||||
tokio::time::timeout(Duration::from_secs(3), async {
|
||||
while server
|
||||
.active_connections
|
||||
.load(std::sync::atomic::Ordering::Acquire)
|
||||
!= 0
|
||||
{
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("cancelled SSE must release the upstream handler");
|
||||
|
||||
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(response.chunk().await.unwrap().is_some());
|
||||
shutdown_tx.send(true).unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(3), tunnel_task)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
server
|
||||
.active_connections
|
||||
.load(std::sync::atomic::Ordering::Acquire),
|
||||
0
|
||||
);
|
||||
drop(response);
|
||||
drop(gateway_task);
|
||||
drop(upstream_task);
|
||||
}
|
||||
|
||||
fn restore_test_env(key: &str, value: Option<std::ffi::OsString>) {
|
||||
if let Some(value) = value {
|
||||
std::env::set_var(key, value);
|
||||
|
||||
@@ -52,6 +52,7 @@ static REDIRECT_REPLAY_BUFFERED_BYTES: AtomicUsize = AtomicUsize::new(0);
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) struct StreamSendWindow {
|
||||
initial_window_bytes: u32,
|
||||
available: Mutex<u64>,
|
||||
notify: Notify,
|
||||
}
|
||||
@@ -59,6 +60,7 @@ pub(crate) struct StreamSendWindow {
|
||||
impl StreamSendWindow {
|
||||
pub(crate) fn new(initial_window_bytes: u32) -> Self {
|
||||
Self {
|
||||
initial_window_bytes: initial_window_bytes.max(1),
|
||||
available: Mutex::new(u64::from(initial_window_bytes.max(1))),
|
||||
notify: Notify::new(),
|
||||
}
|
||||
@@ -82,6 +84,9 @@ impl StreamSendWindow {
|
||||
let requested = bytes as u64;
|
||||
let started_at = Instant::now();
|
||||
loop {
|
||||
let notified = self.notify.notified();
|
||||
tokio::pin!(notified);
|
||||
notified.as_mut().enable();
|
||||
{
|
||||
let mut available = self.available.lock().expect("stream window lock poisoned");
|
||||
if *available >= requested {
|
||||
@@ -93,10 +98,7 @@ impl StreamSendWindow {
|
||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||
return Err(());
|
||||
};
|
||||
if tokio::time::timeout(remaining, self.notify.notified())
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
if tokio::time::timeout(remaining, notified).await.is_err() {
|
||||
return Err(());
|
||||
}
|
||||
}
|
||||
@@ -173,31 +175,33 @@ fn safe_stream_error_message(message: &str) -> &'static str {
|
||||
"upstream request failed"
|
||||
}
|
||||
|
||||
fn try_send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) {
|
||||
async fn send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) -> bool {
|
||||
if bytes == 0 {
|
||||
return;
|
||||
return true;
|
||||
}
|
||||
let delta = bytes.min(u32::MAX as usize) as u32;
|
||||
if frame_tx
|
||||
.try_send(TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::WindowUpdate,
|
||||
0,
|
||||
Bytes::from(
|
||||
serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload {
|
||||
delta_bytes: delta,
|
||||
})
|
||||
.expect("window update payload should serialize"),
|
||||
),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id,
|
||||
delta_bytes = delta,
|
||||
"writer channel full, WINDOW_UPDATE dropped"
|
||||
);
|
||||
if matches!(
|
||||
tokio::time::timeout(
|
||||
FLOW_CONTROL_WAIT_TIMEOUT,
|
||||
frame_tx.send(TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::WindowUpdate,
|
||||
0,
|
||||
Bytes::from(
|
||||
serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload {
|
||||
delta_bytes: delta,
|
||||
})
|
||||
.expect("window update payload should serialize"),
|
||||
),
|
||||
))
|
||||
)
|
||||
.await,
|
||||
Ok(Ok(()))
|
||||
) {
|
||||
return true;
|
||||
}
|
||||
frame_tx.close();
|
||||
false
|
||||
}
|
||||
|
||||
/// Match reqwest's default redirect budget so direct execution and tunnel relay
|
||||
@@ -242,6 +246,23 @@ enum ReplayableRequestBody {
|
||||
struct PreparedRequestBody {
|
||||
first_request_body: Option<upstream_client::UpstreamRequestBody>,
|
||||
replay_body: ReplayableRequestBody,
|
||||
spool_task: Option<tokio::task::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl Drop for PreparedRequestBody {
|
||||
fn drop(&mut self) {
|
||||
if let Some(task) = self.spool_task.take() {
|
||||
task.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct ActiveStreamGuard(Arc<ServerContext>);
|
||||
|
||||
impl Drop for ActiveStreamGuard {
|
||||
fn drop(&mut self) {
|
||||
self.0.active_connections.fetch_sub(1, Ordering::Release);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
@@ -331,7 +352,10 @@ impl hyper::body::Body for ReplayRequestBody {
|
||||
|
||||
#[derive(Debug)]
|
||||
enum SpoolBodyEvent {
|
||||
Data(Bytes),
|
||||
Data {
|
||||
payload: Bytes,
|
||||
credit_returned: bool,
|
||||
},
|
||||
Error(String),
|
||||
End,
|
||||
}
|
||||
@@ -563,8 +587,9 @@ impl RequestBodyReplayState {
|
||||
}
|
||||
}
|
||||
|
||||
fn push_chunk(&self, payload: Bytes) {
|
||||
fn push_chunk(&self, payload: Bytes) -> bool {
|
||||
let mut disable_replay = false;
|
||||
let mut retained = false;
|
||||
let mut state = self.state.lock().expect("request body replay state lock");
|
||||
if let RequestBodyReplayStatus::Collecting {
|
||||
chunks,
|
||||
@@ -577,7 +602,7 @@ impl RequestBodyReplayState {
|
||||
drop(state);
|
||||
self.release_reserved_bytes();
|
||||
self.ready.notify_waiters();
|
||||
return;
|
||||
return false;
|
||||
};
|
||||
let accounted_bytes = payload.len().checked_add(std::mem::size_of::<Bytes>());
|
||||
if next_len > self.budget_bytes
|
||||
@@ -590,6 +615,7 @@ impl RequestBodyReplayState {
|
||||
} else {
|
||||
*buffered_len = next_len;
|
||||
chunks.push(payload);
|
||||
retained = true;
|
||||
}
|
||||
}
|
||||
drop(state);
|
||||
@@ -597,6 +623,7 @@ impl RequestBodyReplayState {
|
||||
self.release_reserved_bytes();
|
||||
self.ready.notify_waiters();
|
||||
}
|
||||
retained
|
||||
}
|
||||
|
||||
fn try_reserve_bytes(&self, bytes: usize) -> bool {
|
||||
@@ -951,9 +978,6 @@ pub(super) fn decode_request_body_frame(frame: TunnelFrame) -> Result<Bytes, std
|
||||
Ok(frame.payload)
|
||||
}
|
||||
|
||||
// Drain tunnel body frames on a detached task so the shared dispatcher is no
|
||||
// longer coupled to upstream body polling. Redirect replay retains a bounded
|
||||
// copy; crossing either replay budget only disables replay for this request.
|
||||
fn prepare_request_body(
|
||||
stream_id: u32,
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
@@ -973,19 +997,20 @@ fn prepare_request_body(
|
||||
None => ReplayableRequestBody::NonReplayable,
|
||||
};
|
||||
|
||||
tokio::spawn(spool_request_body(
|
||||
let spool_task = tokio::spawn(spool_request_body(
|
||||
stream_id,
|
||||
body_rx,
|
||||
spool_tx,
|
||||
replay_state,
|
||||
body_size,
|
||||
deadline,
|
||||
frame_tx,
|
||||
frame_tx.clone(),
|
||||
));
|
||||
|
||||
PreparedRequestBody {
|
||||
first_request_body: Some(build_spooled_request_body(spool_rx)),
|
||||
first_request_body: Some(build_spooled_request_body(spool_rx, stream_id, frame_tx)),
|
||||
replay_body,
|
||||
spool_task: Some(spool_task),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1001,6 +1026,7 @@ fn prepare_bodyless_request_body(
|
||||
} else {
|
||||
ReplayableRequestBody::NonReplayable
|
||||
},
|
||||
spool_task: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1056,11 +1082,16 @@ async fn spool_request_body(
|
||||
};
|
||||
|
||||
let Some(frame) = frame else {
|
||||
let message = "tunnel request body closed before stream end".to_string();
|
||||
if let Some(state) = &replay_state {
|
||||
state.finish();
|
||||
state.fail(message.clone());
|
||||
}
|
||||
let _ =
|
||||
send_spool_event(&mut spool_tx, SpoolBodyEvent::End, replay_state.as_ref()).await;
|
||||
let _ = send_spool_event(
|
||||
&mut spool_tx,
|
||||
SpoolBodyEvent::Error(message),
|
||||
replay_state.as_ref(),
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
};
|
||||
|
||||
@@ -1086,13 +1117,23 @@ async fn spool_request_body(
|
||||
|
||||
if !payload.is_empty() {
|
||||
body_size.fetch_add(payload.len(), Ordering::Relaxed);
|
||||
try_send_window_update(&frame_tx, stream_id, payload.len());
|
||||
if let Some(state) = &replay_state {
|
||||
state.push_chunk(payload.clone());
|
||||
let credit_returned = replay_state
|
||||
.as_ref()
|
||||
.is_some_and(|state| state.push_chunk(payload.clone()));
|
||||
if credit_returned
|
||||
&& !send_window_update(&frame_tx, stream_id, payload.len()).await
|
||||
{
|
||||
if let Some(state) = &replay_state {
|
||||
state.fail("tunnel flow-control update failed".to_string());
|
||||
}
|
||||
return;
|
||||
}
|
||||
if send_spool_event(
|
||||
&mut spool_tx,
|
||||
SpoolBodyEvent::Data(payload),
|
||||
SpoolBodyEvent::Data {
|
||||
payload,
|
||||
credit_returned,
|
||||
},
|
||||
replay_state.as_ref(),
|
||||
)
|
||||
.await
|
||||
@@ -1479,6 +1520,7 @@ where
|
||||
}
|
||||
|
||||
let mut stream = response.into_body().into_data_stream();
|
||||
let chunk_size = MAX_CHUNK_SIZE.min(response_window.initial_window_bytes as usize);
|
||||
loop {
|
||||
let chunk_result = if let Some(deadline) = response_body_deadline {
|
||||
let Some(remaining) = remaining_timeout(deadline) else {
|
||||
@@ -1531,7 +1573,7 @@ where
|
||||
|
||||
match chunk_result {
|
||||
Ok(chunk) => {
|
||||
if chunk.len() <= MAX_CHUNK_SIZE {
|
||||
if chunk.len() <= chunk_size {
|
||||
let (payload, extra_flags) = raw_payload(chunk);
|
||||
if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len())
|
||||
.await
|
||||
@@ -1561,7 +1603,7 @@ where
|
||||
} else {
|
||||
let mut offset = 0;
|
||||
while offset < chunk.len() {
|
||||
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
|
||||
let end = (offset + chunk_size).min(chunk.len());
|
||||
let slice = chunk.slice(offset..end);
|
||||
let (payload, extra_flags) = raw_payload(slice);
|
||||
if !acquire_response_credit(
|
||||
@@ -1735,6 +1777,7 @@ pub async fn handle_stream(
|
||||
};
|
||||
|
||||
server.active_connections.fetch_add(1, Ordering::Release);
|
||||
let _active_stream = ActiveStreamGuard(Arc::clone(&server));
|
||||
|
||||
let stream_io = StreamIo {
|
||||
body_rx,
|
||||
@@ -1745,7 +1788,6 @@ pub async fn handle_stream(
|
||||
|
||||
let connect_elapsed = handle_stream_inner(&state, &server, stream_id, meta, stream_io).await;
|
||||
|
||||
server.active_connections.fetch_sub(1, Ordering::Release);
|
||||
if let Some(d) = connect_elapsed {
|
||||
server.metrics.record_request(d);
|
||||
}
|
||||
@@ -1772,6 +1814,18 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
||||
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
|
||||
"writer channel stalled for body frame, abandoning stream"
|
||||
);
|
||||
let reset = TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::ResetStream,
|
||||
0,
|
||||
Bytes::from_static(b"{\"reason\":\"tunnel writer stalled\"}"),
|
||||
);
|
||||
if !matches!(
|
||||
tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(reset)).await,
|
||||
Ok(Ok(()))
|
||||
) {
|
||||
tx.close();
|
||||
}
|
||||
false
|
||||
}
|
||||
Ok(Err(QueueSendError::Full(_))) => {
|
||||
@@ -1781,7 +1835,10 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
||||
} else {
|
||||
match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
||||
Ok(Ok(())) => true,
|
||||
Ok(Err(_)) => false,
|
||||
Ok(Err(_)) => {
|
||||
tx.close();
|
||||
false
|
||||
}
|
||||
Err(_) => {
|
||||
warn!(
|
||||
stream_id,
|
||||
@@ -1789,6 +1846,7 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
||||
flags = flags,
|
||||
"control frame send timeout (writer congested), abandoning stream"
|
||||
);
|
||||
tx.close();
|
||||
false
|
||||
}
|
||||
}
|
||||
@@ -2126,7 +2184,6 @@ async fn handle_stream_inner(
|
||||
}
|
||||
|
||||
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
||||
// Error frames use best-effort delivery — don't block if writer is congested
|
||||
let safe_message = safe_stream_error_message(msg);
|
||||
let _ = send_frame(
|
||||
tx,
|
||||
@@ -2163,22 +2220,42 @@ fn build_streaming_request_body(
|
||||
|
||||
fn build_spooled_request_body(
|
||||
spool_rx: mpsc::Receiver<SpoolBodyEvent>,
|
||||
stream_id: u32,
|
||||
frame_tx: FrameSender,
|
||||
) -> upstream_client::UpstreamRequestBody {
|
||||
let body_stream = stream::unfold((spool_rx, false), |(mut spool_rx, finished)| async move {
|
||||
if finished {
|
||||
return None;
|
||||
}
|
||||
let body_stream = stream::unfold(
|
||||
(spool_rx, frame_tx, false),
|
||||
move |(mut spool_rx, frame_tx, finished)| async move {
|
||||
if finished {
|
||||
return None;
|
||||
}
|
||||
|
||||
match spool_rx.recv().await {
|
||||
Some(SpoolBodyEvent::Data(payload)) => {
|
||||
Some((Ok(BodyFrame::data(payload)), (spool_rx, false)))
|
||||
match spool_rx.recv().await {
|
||||
Some(SpoolBodyEvent::Data {
|
||||
payload,
|
||||
credit_returned,
|
||||
}) => {
|
||||
if !credit_returned
|
||||
&& !send_window_update(&frame_tx, stream_id, payload.len()).await
|
||||
{
|
||||
return Some((
|
||||
Err(io::Error::other("tunnel flow-control update failed")),
|
||||
(spool_rx, frame_tx, true),
|
||||
));
|
||||
}
|
||||
Some((Ok(BodyFrame::data(payload)), (spool_rx, frame_tx, false)))
|
||||
}
|
||||
Some(SpoolBodyEvent::Error(message)) => {
|
||||
Some((Err(io::Error::other(message)), (spool_rx, frame_tx, true)))
|
||||
}
|
||||
Some(SpoolBodyEvent::End) => None,
|
||||
None => Some((
|
||||
Err(io::Error::other("tunnel request body ended unexpectedly")),
|
||||
(spool_rx, frame_tx, true),
|
||||
)),
|
||||
}
|
||||
Some(SpoolBodyEvent::Error(message)) => {
|
||||
Some((Err(io::Error::other(message)), (spool_rx, true)))
|
||||
}
|
||||
Some(SpoolBodyEvent::End) | None => None,
|
||||
}
|
||||
});
|
||||
},
|
||||
);
|
||||
|
||||
upstream_client::stream_request_body(body_stream)
|
||||
}
|
||||
@@ -2249,6 +2326,105 @@ fn build_prefixed_request_body(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn window_updates_wait_for_capacity_instead_of_disappearing() {
|
||||
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(1);
|
||||
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1);
|
||||
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||
sender
|
||||
.try_send(TunnelFrame::control(MsgType::Ping, Bytes::new()))
|
||||
.unwrap();
|
||||
let task = tokio::spawn(async move { send_window_update(&sender, 7, 1024).await });
|
||||
tokio::time::sleep(Duration::from_secs(1)).await;
|
||||
assert!(!task.is_finished());
|
||||
high_rx.recv().await.unwrap();
|
||||
assert!(task.await.unwrap());
|
||||
let update = high_rx.recv().await.unwrap();
|
||||
assert_eq!(update.msg_type, MsgType::WindowUpdate);
|
||||
let payload: aether_contracts::tunnel::WindowUpdatePayload =
|
||||
serde_json::from_slice(&update.payload).unwrap();
|
||||
assert_eq!(payload.delta_bytes, 1024);
|
||||
}
|
||||
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn stalled_body_delivery_emits_a_reset() {
|
||||
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4);
|
||||
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1);
|
||||
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||
sender
|
||||
.try_send(TunnelFrame::new(
|
||||
7,
|
||||
MsgType::ResponseBody,
|
||||
0,
|
||||
Bytes::from_static(b"first"),
|
||||
))
|
||||
.unwrap();
|
||||
assert!(
|
||||
!send_frame(
|
||||
&sender,
|
||||
TunnelFrame::new(7, MsgType::ResponseBody, 0, Bytes::from_static(b"second"))
|
||||
)
|
||||
.await
|
||||
);
|
||||
let reset = high_rx.recv().await.unwrap();
|
||||
assert_eq!(reset.msg_type, MsgType::ResetStream);
|
||||
assert_eq!(reset.stream_id, 7);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_credit_follows_consumption_without_redirect_replay() {
|
||||
let (body_tx, body_rx) = mpsc::channel(4);
|
||||
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4);
|
||||
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4);
|
||||
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||
let mut prepared = prepare_request_body(
|
||||
7,
|
||||
body_rx,
|
||||
Arc::new(AtomicUsize::new(0)),
|
||||
Instant::now() + Duration::from_secs(10),
|
||||
false,
|
||||
sender,
|
||||
);
|
||||
body_tx
|
||||
.send(TunnelFrame::new(
|
||||
7,
|
||||
MsgType::RequestBody,
|
||||
flags::END_STREAM,
|
||||
Bytes::from_static(b"body"),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::task::yield_now().await;
|
||||
assert!(high_rx.try_recv().is_err());
|
||||
let mut body = prepared.take_first_request_body();
|
||||
assert!(body.frame().await.unwrap().is_ok());
|
||||
assert_eq!(
|
||||
high_rx.recv().await.unwrap().msg_type,
|
||||
MsgType::WindowUpdate
|
||||
);
|
||||
assert!(body.frame().await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_prepared_body_cancels_its_spooler() {
|
||||
let (body_tx, body_rx) = mpsc::channel(4);
|
||||
let (high_tx, _high_rx) = aether_runtime::bounded_queue(4);
|
||||
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4);
|
||||
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||
let prepared = prepare_request_body(
|
||||
7,
|
||||
body_rx,
|
||||
Arc::new(AtomicUsize::new(0)),
|
||||
Instant::now() + Duration::from_secs(3600),
|
||||
false,
|
||||
sender,
|
||||
);
|
||||
drop(prepared);
|
||||
tokio::time::timeout(Duration::from_secs(1), body_tx.closed())
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::net::SocketAddr;
|
||||
use std::pin::Pin;
|
||||
@@ -2379,7 +2555,7 @@ mod tests {
|
||||
let (tx, rx) = mpsc::channel(4);
|
||||
let (frame_tx, sent, writer_handle) = spawn_test_writer();
|
||||
let body_size = Arc::new(AtomicUsize::new(0));
|
||||
let prepared = prepare_request_body(
|
||||
let mut prepared = prepare_request_body(
|
||||
1,
|
||||
rx,
|
||||
Arc::clone(&body_size),
|
||||
@@ -2389,6 +2565,7 @@ mod tests {
|
||||
);
|
||||
let mut body = prepared
|
||||
.first_request_body
|
||||
.take()
|
||||
.expect("first request body should be present");
|
||||
|
||||
tx.send(TunnelFrame::new(
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use tokio::task::{JoinError, JoinHandle};
|
||||
|
||||
pub(super) struct SessionTask<T>(JoinHandle<T>);
|
||||
|
||||
impl<T> SessionTask<T> {
|
||||
pub(super) fn new(handle: JoinHandle<T>) -> Self {
|
||||
Self(handle)
|
||||
}
|
||||
pub(super) fn abort(&self) {
|
||||
self.0.abort();
|
||||
}
|
||||
pub(super) fn is_finished(&self) -> bool {
|
||||
self.0.is_finished()
|
||||
}
|
||||
}
|
||||
|
||||
impl<T> Future for SessionTask<T> {
|
||||
type Output = Result<T, JoinError>;
|
||||
|
||||
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
|
||||
Pin::new(&mut self.0).poll(context)
|
||||
}
|
||||
}
|
||||
|
||||
impl<T> Drop for SessionTask<T> {
|
||||
fn drop(&mut self) {
|
||||
self.0.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_a_session_task_aborts_its_child() {
|
||||
let child = tokio::spawn(std::future::pending::<()>());
|
||||
let abort = child.abort_handle();
|
||||
drop(SessionTask::new(child));
|
||||
tokio::time::timeout(std::time::Duration::from_secs(1), async {
|
||||
while !abort.is_finished() {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -13,6 +13,7 @@ use aether_contracts::tunnel::{MsgType, HEADER_SIZE};
|
||||
use aether_runtime::QueueSnapshot;
|
||||
use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError};
|
||||
use futures_util::SinkExt;
|
||||
use tokio::sync::watch;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
use tracing::{debug, error, trace};
|
||||
@@ -24,6 +25,8 @@ use aether_contracts::tunnel_security::SecureFrameCodec;
|
||||
|
||||
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
|
||||
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
|
||||
const WRITE_TIMEOUT: Duration = Duration::from_secs(15);
|
||||
const CLOSE_TIMEOUT: Duration = Duration::from_secs(1);
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum FramePriority {
|
||||
@@ -43,9 +46,18 @@ pub struct FrameQueueSnapshots {
|
||||
pub struct FrameSender {
|
||||
high_tx: BoundedQueueSender<Frame>,
|
||||
normal_tx: BoundedQueueSender<Frame>,
|
||||
close_tx: watch::Sender<bool>,
|
||||
}
|
||||
|
||||
impl FrameSender {
|
||||
pub fn close(&self) {
|
||||
let _ = self.close_tx.send(true);
|
||||
}
|
||||
|
||||
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||
self.close_tx.subscribe()
|
||||
}
|
||||
|
||||
pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
|
||||
match classify_frame_priority(&frame) {
|
||||
FramePriority::High => self.high_tx.send(frame).await,
|
||||
@@ -73,7 +85,12 @@ impl FrameSender {
|
||||
high_tx: BoundedQueueSender<Frame>,
|
||||
normal_tx: BoundedQueueSender<Frame>,
|
||||
) -> Self {
|
||||
Self { high_tx, normal_tx }
|
||||
let (close_tx, _) = watch::channel(false);
|
||||
Self {
|
||||
high_tx,
|
||||
normal_tx,
|
||||
close_tx,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -113,15 +130,24 @@ where
|
||||
{
|
||||
let (high_tx, mut high_rx) = bounded_queue::<Frame>(HIGH_PRIORITY_QUEUE_CAPACITY);
|
||||
let (normal_tx, mut normal_rx) = bounded_queue::<Frame>(NORMAL_PRIORITY_QUEUE_CAPACITY);
|
||||
let tx = FrameSender { high_tx, normal_tx };
|
||||
let (close_tx, mut close_rx) = watch::channel(false);
|
||||
let tx = FrameSender {
|
||||
high_tx,
|
||||
normal_tx,
|
||||
close_tx,
|
||||
};
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
let mut ping_ticker = tokio::time::interval(ping_interval);
|
||||
let mut high_open = true;
|
||||
let mut normal_open = true;
|
||||
let mut close_open = true;
|
||||
ping_ticker.tick().await; // skip first immediate tick
|
||||
|
||||
loop {
|
||||
if *close_rx.borrow() {
|
||||
break;
|
||||
}
|
||||
if let Ok(frame) = high_rx.try_recv() {
|
||||
if !write_frame(
|
||||
&mut sink,
|
||||
@@ -141,6 +167,10 @@ where
|
||||
|
||||
tokio::select! {
|
||||
biased;
|
||||
changed = close_rx.changed(), if close_open => {
|
||||
if changed.is_err() { close_open = false; }
|
||||
if *close_rx.borrow() { break; }
|
||||
},
|
||||
frame = high_rx.recv(), if high_open => {
|
||||
match frame {
|
||||
Some(frame) => {
|
||||
@@ -152,7 +182,7 @@ where
|
||||
}
|
||||
}
|
||||
_ = ping_ticker.tick(), if high_open || normal_open => {
|
||||
if let Err(e) = sink.send(Message::Ping(vec![])).await {
|
||||
if let Err(e) = send_message(&mut sink, Message::Ping(vec![])).await {
|
||||
error!(error = %e, "failed to send WebSocket ping");
|
||||
if let Some(metrics) = tunnel_metrics.as_deref() {
|
||||
metrics.record_error("ws_ping_error", &e.to_string());
|
||||
@@ -174,7 +204,7 @@ where
|
||||
}
|
||||
}
|
||||
debug!("writer task exiting");
|
||||
let _ = sink.close().await;
|
||||
let _ = tokio::time::timeout(CLOSE_TIMEOUT, sink.close()).await;
|
||||
});
|
||||
|
||||
(tx, handle)
|
||||
@@ -228,7 +258,7 @@ where
|
||||
None => frame.encode(),
|
||||
};
|
||||
let wire_len = data.len().max(HEADER_SIZE);
|
||||
if let Err(e) = sink.send(Message::Binary(data.into())).await {
|
||||
if let Err(e) = send_message(sink, Message::Binary(data.into())).await {
|
||||
error!(
|
||||
stream_id = stream_id,
|
||||
msg_type = ?msg_type,
|
||||
@@ -248,8 +278,92 @@ where
|
||||
true
|
||||
}
|
||||
|
||||
async fn send_message<S>(
|
||||
sink: &mut S,
|
||||
message: Message,
|
||||
) -> Result<(), tokio_tungstenite::tungstenite::Error>
|
||||
where
|
||||
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin,
|
||||
{
|
||||
tokio::time::timeout(WRITE_TIMEOUT, sink.send(message))
|
||||
.await
|
||||
.map_err(|_| {
|
||||
tokio_tungstenite::tungstenite::Error::Io(std::io::Error::new(
|
||||
std::io::ErrorKind::TimedOut,
|
||||
"tunnel WebSocket write timed out",
|
||||
))
|
||||
})?
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[tokio::test]
|
||||
async fn dropping_last_sender_flushes_queued_body_and_end_frames() {
|
||||
let sink = VecSink::default();
|
||||
let sent = Arc::clone(&sink.sent);
|
||||
let (sender, task) = spawn_writer(sink, Duration::from_secs(60));
|
||||
sender
|
||||
.send(Frame::new(
|
||||
7,
|
||||
MsgType::ResponseBody,
|
||||
0,
|
||||
bytes::Bytes::from_static(b"late"),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
sender
|
||||
.send(Frame::new(7, MsgType::StreamEnd, 0, bytes::Bytes::new()))
|
||||
.await
|
||||
.unwrap();
|
||||
drop(sender);
|
||||
task.await.unwrap();
|
||||
let frames = sent.lock().unwrap();
|
||||
assert_eq!(frames.len(), 2);
|
||||
let Message::Binary(body) = &frames[0] else {
|
||||
panic!("expected body")
|
||||
};
|
||||
assert_eq!(
|
||||
Frame::decode(body.clone().into()).unwrap().payload,
|
||||
b"late".as_slice()
|
||||
);
|
||||
}
|
||||
|
||||
struct StalledSink;
|
||||
|
||||
impl futures_util::Sink<Message> for StalledSink {
|
||||
type Error = Error;
|
||||
fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||
Poll::Pending
|
||||
}
|
||||
fn start_send(self: Pin<&mut Self>, _: Message) -> Result<(), Error> {
|
||||
Ok(())
|
||||
}
|
||||
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||
Poll::Pending
|
||||
}
|
||||
fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||
Poll::Pending
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn stalled_socket_write_and_close_are_bounded() {
|
||||
let (sender, task) = spawn_writer(StalledSink, Duration::from_secs(60));
|
||||
sender
|
||||
.send(Frame::new(
|
||||
1,
|
||||
MsgType::ResponseBody,
|
||||
0,
|
||||
bytes::Bytes::from_static(b"data"),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(20), task)
|
||||
.await
|
||||
.expect("writer should time out")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use aether_data_contracts::repository::{
|
||||
candidates::{
|
||||
sanitize_request_candidate_api_formats, sanitize_request_candidate_error_type,
|
||||
sanitize_request_candidate_extra_data, sanitize_request_candidate_required_capabilities,
|
||||
sanitize_request_candidate_skip_reason, DecisionTrace, DecisionTraceCandidate,
|
||||
RequestCandidateStatus,
|
||||
sanitize_request_candidate_extra_data_for_persistence,
|
||||
sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason,
|
||||
DecisionTrace, DecisionTraceCandidate, RequestCandidateStatus,
|
||||
},
|
||||
provider_catalog::StoredProviderCatalogKey,
|
||||
usage::StoredRequestUsageAudit,
|
||||
@@ -334,10 +334,13 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts(
|
||||
key_accounts: &BTreeMap<String, AdminMonitoringKeyAccountDisplay>,
|
||||
) -> Value {
|
||||
let mut item = item.clone();
|
||||
item.sanitize_sensitive_diagnostics();
|
||||
item.sanitize_for_admin();
|
||||
let candidate = &item.candidate;
|
||||
let sanitized_extra_data =
|
||||
build_admin_monitoring_trace_candidate_extra_data(candidate.extra_data.as_ref(), usage);
|
||||
let sanitized_extra_data = build_admin_monitoring_trace_candidate_extra_data(
|
||||
candidate.extra_data.as_ref(),
|
||||
candidate.status_code,
|
||||
usage,
|
||||
);
|
||||
let sanitized_extra_data_ref =
|
||||
(!sanitized_extra_data.is_null()).then_some(&sanitized_extra_data);
|
||||
let sanitized_key_api_formats =
|
||||
@@ -382,7 +385,7 @@ pub fn build_admin_monitoring_trace_request_candidate_payload_with_key_accounts(
|
||||
"is_cached": candidate.is_cached,
|
||||
"status_code": candidate.status_code,
|
||||
"error_type": sanitize_request_candidate_error_type(candidate.error_type.clone()),
|
||||
"error_message": serde_json::Value::Null,
|
||||
"error_message": candidate.error_message,
|
||||
"latency_ms": candidate.latency_ms,
|
||||
"concurrent_requests": candidate.concurrent_requests,
|
||||
"ranking": build_admin_monitoring_trace_candidate_ranking(sanitized_extra_data_ref),
|
||||
@@ -516,9 +519,10 @@ fn build_admin_monitoring_trace_candidate_ranking(existing: Option<&Value>) -> V
|
||||
|
||||
fn build_admin_monitoring_trace_candidate_extra_data(
|
||||
existing: Option<&Value>,
|
||||
candidate_status_code: Option<u16>,
|
||||
usage: Option<&StoredRequestUsageAudit>,
|
||||
) -> Value {
|
||||
let mut extra_data = sanitize_request_candidate_extra_data(existing.cloned())
|
||||
let mut extra_data = sanitize_request_candidate_extra_data_for_persistence(existing.cloned())
|
||||
.and_then(|value| value.as_object().cloned());
|
||||
|
||||
if let Some(usage) = usage {
|
||||
@@ -546,7 +550,7 @@ fn build_admin_monitoring_trace_candidate_extra_data(
|
||||
if admin_monitoring_usage_is_error_node(usage) {
|
||||
if let Some(response) = admin_monitoring_trace_response_data(
|
||||
"upstream_response",
|
||||
usage.status_code,
|
||||
candidate_status_code,
|
||||
usage.response_body_state,
|
||||
) {
|
||||
merge_admin_monitoring_trace_response(extra_object, "upstream_response", response);
|
||||
@@ -573,7 +577,8 @@ fn build_admin_monitoring_trace_candidate_extra_data(
|
||||
}
|
||||
}
|
||||
|
||||
sanitize_request_candidate_extra_data(extra_data.map(Value::Object)).unwrap_or(Value::Null)
|
||||
sanitize_request_candidate_extra_data_for_persistence(extra_data.map(Value::Object))
|
||||
.unwrap_or(Value::Null)
|
||||
}
|
||||
|
||||
fn admin_monitoring_trace_response_data(
|
||||
@@ -607,6 +612,9 @@ fn merge_admin_monitoring_trace_response(
|
||||
};
|
||||
|
||||
for (field, value) in response_object {
|
||||
if field == "status_code" && existing_object.get(field).is_some_and(Value::is_number) {
|
||||
continue;
|
||||
}
|
||||
if admin_monitoring_trace_response_value_empty(value)
|
||||
&& existing_object
|
||||
.get(field)
|
||||
|
||||
@@ -67,17 +67,6 @@ pub fn admin_email_template_html_is_valid(value: &str) -> bool {
|
||||
pub const ADMIN_SYSTEM_CONFIG_EXPORT_VERSION: &str = "2.3";
|
||||
pub const ADMIN_SYSTEM_CONFIG_SUPPORTED_VERSIONS: &[&str] =
|
||||
&["2.0", "2.1", "2.2", ADMIN_SYSTEM_CONFIG_EXPORT_VERSION];
|
||||
pub const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY: &str = "execution_extra_trusted_dns_hosts";
|
||||
pub const EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_MAX_ENTRIES: usize = 128;
|
||||
pub const EXECUTION_EXTRA_TRUSTED_DNS_HOST_MAX_BYTES: usize = 253;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ExecutionExtraTrustedDnsHostsConfigError {
|
||||
InvalidValue,
|
||||
TooManyEntries,
|
||||
InvalidHost,
|
||||
}
|
||||
|
||||
pub const ADMIN_SYSTEM_USERS_EXPORT_VERSION: &str = "1.6";
|
||||
pub const ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS: &[&str] =
|
||||
&["1.3", "1.4", "1.5", ADMIN_SYSTEM_USERS_EXPORT_VERSION];
|
||||
@@ -2289,7 +2278,6 @@ pub fn admin_system_config_default_value(key: &str) -> Option<serde_json::Value>
|
||||
"email_suffix_mode" => Some(json!("none")),
|
||||
"email_suffix_list" => Some(json!([])),
|
||||
"enable_format_conversion" => Some(json!(false)),
|
||||
EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY => Some(json!([])),
|
||||
"enable_model_directives" => Some(json!(false)),
|
||||
// Failover after a provider-side Cyber policy refusal is an explicit
|
||||
// opt-in. Keep the system-config fallback aligned with the routing
|
||||
@@ -2329,58 +2317,6 @@ pub fn admin_system_config_default_value(key: &str) -> Option<serde_json::Value>
|
||||
}
|
||||
}
|
||||
|
||||
pub fn normalize_execution_extra_trusted_dns_hosts_config_value(
|
||||
value: serde_json::Value,
|
||||
) -> Result<serde_json::Value, ExecutionExtraTrustedDnsHostsConfigError> {
|
||||
let values = match value {
|
||||
Value::Null => Vec::new(),
|
||||
Value::Array(values) => values,
|
||||
_ => return Err(ExecutionExtraTrustedDnsHostsConfigError::InvalidValue),
|
||||
};
|
||||
if values.len() > EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_MAX_ENTRIES {
|
||||
return Err(ExecutionExtraTrustedDnsHostsConfigError::TooManyEntries);
|
||||
}
|
||||
|
||||
let mut hosts = BTreeSet::new();
|
||||
for value in values {
|
||||
let host = value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.ok_or(ExecutionExtraTrustedDnsHostsConfigError::InvalidHost)?;
|
||||
let host = host.trim_end_matches('.').to_ascii_lowercase();
|
||||
if !execution_extra_trusted_dns_host_is_valid(&host) {
|
||||
return Err(ExecutionExtraTrustedDnsHostsConfigError::InvalidHost);
|
||||
}
|
||||
hosts.insert(host);
|
||||
}
|
||||
|
||||
Ok(Value::Array(hosts.into_iter().map(Value::String).collect()))
|
||||
}
|
||||
|
||||
fn execution_extra_trusted_dns_host_is_valid(host: &str) -> bool {
|
||||
if host.is_empty()
|
||||
|| host.len() > EXECUTION_EXTRA_TRUSTED_DNS_HOST_MAX_BYTES
|
||||
|| !host.is_ascii()
|
||||
|| host.parse::<std::net::IpAddr>().is_ok()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
let labels = host.split('.').collect::<Vec<_>>();
|
||||
if labels.len() < 2 {
|
||||
return false;
|
||||
}
|
||||
labels.iter().all(|label| {
|
||||
!label.is_empty()
|
||||
&& label.len() <= 63
|
||||
&& !label.starts_with('-')
|
||||
&& !label.ends_with('-')
|
||||
&& label
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-')
|
||||
})
|
||||
}
|
||||
|
||||
pub fn build_admin_system_configs_payload(
|
||||
entries: &[StoredSystemConfigEntry],
|
||||
) -> serde_json::Value {
|
||||
@@ -2872,15 +2808,6 @@ pub fn parse_admin_system_config_update(
|
||||
)
|
||||
})?;
|
||||
}
|
||||
EXECUTION_EXTRA_TRUSTED_DNS_HOSTS_CONFIG_KEY => {
|
||||
value =
|
||||
normalize_execution_extra_trusted_dns_hosts_config_value(value).map_err(|_| {
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
json!({ "detail": "额外可信 Fake-IP 域名配置格式无效" }),
|
||||
)
|
||||
})?;
|
||||
}
|
||||
"module.important_notification.default_channel" => {
|
||||
value = normalize_notification_channel_value(value).map_err(|_| {
|
||||
(
|
||||
@@ -4621,36 +4548,6 @@ mod tests {
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extra_trusted_dns_hosts_update_normalizes_exact_hostnames() {
|
||||
let update = parse_admin_system_config_update(
|
||||
"execution_extra_trusted_dns_hosts",
|
||||
br#"{"value":[" API.Example.COM. ","api.example.com"]}"#,
|
||||
)
|
||||
.expect("valid extra trusted DNS hosts should parse");
|
||||
assert_eq!(update.value, json!(["api.example.com"]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extra_trusted_dns_hosts_update_rejects_non_exact_hostnames() {
|
||||
for value in [
|
||||
r#"["*.example.com"]"#,
|
||||
r#"["example.com:443"]"#,
|
||||
r#"["https://example.com/path"]"#,
|
||||
r#"["10.0.0.1"]"#,
|
||||
r#"["example..com"]"#,
|
||||
r#"["localhost"]"#,
|
||||
] {
|
||||
let body = format!(r#"{{"value":{value}}}"#);
|
||||
let error = parse_admin_system_config_update(
|
||||
"execution_extra_trusted_dns_hosts",
|
||||
body.as_bytes(),
|
||||
)
|
||||
.expect_err("non-exact hostname should be rejected");
|
||||
assert_eq!(error.0, http::StatusCode::BAD_REQUEST);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_notification_email_config_key_normalizes_to_important_notification() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -588,6 +588,68 @@ pub struct SettingsPayload {
|
||||
pub drain_deadline_ms: u64,
|
||||
}
|
||||
|
||||
impl SettingsPayload {
|
||||
pub fn is_valid(&self) -> bool {
|
||||
self.initial_stream_window_bytes > 0
|
||||
&& u64::from(self.initial_stream_window_bytes)
|
||||
<= MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64
|
||||
&& self.min_window_update_bytes > 0
|
||||
&& self.min_window_update_bytes <= self.initial_stream_window_bytes
|
||||
&& self.drain_deadline_ms > 0
|
||||
}
|
||||
|
||||
pub fn negotiate(&self, initial_window_bytes: u32, drain_deadline_ms: u64) -> Self {
|
||||
let window = self
|
||||
.initial_stream_window_bytes
|
||||
.min(initial_window_bytes)
|
||||
.max(1);
|
||||
Self {
|
||||
initial_stream_window_bytes: window,
|
||||
min_window_update_bytes: self.min_window_update_bytes.min((window / 4).max(1)),
|
||||
drain_deadline_ms: self.drain_deadline_ms.min(drain_deadline_ms),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod settings_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn negotiation_bounds_window_updates_by_the_smaller_window() {
|
||||
let settings = SettingsPayload {
|
||||
initial_stream_window_bytes: 512 * 1024,
|
||||
min_window_update_bytes: 128 * 1024,
|
||||
drain_deadline_ms: 30_000,
|
||||
};
|
||||
let negotiated = settings.negotiate(4 * 1024 * 1024, 1000);
|
||||
assert_eq!(negotiated.initial_stream_window_bytes, 512 * 1024);
|
||||
assert_eq!(negotiated.min_window_update_bytes, 128 * 1024);
|
||||
assert_eq!(negotiated.drain_deadline_ms, 1000);
|
||||
assert!(negotiated.is_valid());
|
||||
let tiny = settings.negotiate(1, 1);
|
||||
assert_eq!(tiny.initial_stream_window_bytes, 1);
|
||||
assert_eq!(tiny.min_window_update_bytes, 1);
|
||||
assert!(tiny.is_valid());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_window_settings_are_rejected() {
|
||||
let mut settings = SettingsPayload {
|
||||
initial_stream_window_bytes: 1024,
|
||||
min_window_update_bytes: 256,
|
||||
drain_deadline_ms: 1,
|
||||
};
|
||||
settings.min_window_update_bytes = 1025;
|
||||
assert!(!settings.is_valid());
|
||||
settings.min_window_update_bytes = 0;
|
||||
assert!(!settings.is_valid());
|
||||
settings.initial_stream_window_bytes = u32::MAX;
|
||||
settings.min_window_update_bytes = 1;
|
||||
assert!(!settings.is_valid());
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct WindowUpdatePayload {
|
||||
pub delta_bytes: u32,
|
||||
|
||||
@@ -172,7 +172,7 @@ DO UPDATE SET
|
||||
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
||||
ELSE EXCLUDED.error_type
|
||||
END,
|
||||
error_message = NULL,
|
||||
error_message = __AETHER_CANDIDATE_ERROR_MESSAGE__,
|
||||
latency_ms = CASE
|
||||
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||
AND EXCLUDED.status <> request_candidates.status
|
||||
@@ -184,7 +184,7 @@ DO UPDATE SET
|
||||
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
||||
END,
|
||||
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
||||
extra_data = EXCLUDED.extra_data,
|
||||
extra_data = __AETHER_CANDIDATE_EXTRA_DATA__,
|
||||
required_capabilities = EXCLUDED.required_capabilities,
|
||||
created_at = CASE
|
||||
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
||||
@@ -275,7 +275,7 @@ DO UPDATE SET
|
||||
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
||||
ELSE EXCLUDED.error_type
|
||||
END,
|
||||
error_message = NULL,
|
||||
error_message = __AETHER_CANDIDATE_ERROR_MESSAGE__,
|
||||
latency_ms = CASE
|
||||
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||
AND EXCLUDED.status <> request_candidates.status
|
||||
@@ -287,7 +287,7 @@ DO UPDATE SET
|
||||
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
||||
END,
|
||||
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
||||
extra_data = EXCLUDED.extra_data,
|
||||
extra_data = __AETHER_CANDIDATE_EXTRA_DATA__,
|
||||
required_capabilities = EXCLUDED.required_capabilities,
|
||||
created_at = CASE
|
||||
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
||||
@@ -353,7 +353,7 @@ DO UPDATE SET
|
||||
THEN __AETHER_SANITIZED_LEGACY_ERROR_TYPE__
|
||||
ELSE EXCLUDED.error_type
|
||||
END,
|
||||
error_message = NULL,
|
||||
error_message = __AETHER_CANDIDATE_ERROR_MESSAGE__,
|
||||
latency_ms = CASE
|
||||
WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||
AND EXCLUDED.status <> request_candidates.status
|
||||
@@ -365,7 +365,7 @@ DO UPDATE SET
|
||||
ELSE COALESCE(EXCLUDED.latency_ms, request_candidates.latency_ms)
|
||||
END,
|
||||
concurrent_requests = COALESCE(EXCLUDED.concurrent_requests, request_candidates.concurrent_requests),
|
||||
extra_data = EXCLUDED.extra_data,
|
||||
extra_data = __AETHER_CANDIDATE_EXTRA_DATA__,
|
||||
required_capabilities = EXCLUDED.required_capabilities,
|
||||
created_at = CASE
|
||||
WHEN request_candidates.created_at <= TO_TIMESTAMP(1)
|
||||
@@ -446,6 +446,23 @@ fn postgres_candidate_upsert_sql(template: &str) -> String {
|
||||
)
|
||||
.as_str(),
|
||||
)
|
||||
.replace(
|
||||
"__AETHER_CANDIDATE_ERROR_MESSAGE__",
|
||||
"CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') \
|
||||
AND EXCLUDED.status <> request_candidates.status \
|
||||
THEN request_candidates.error_message \
|
||||
WHEN request_candidates.status = 'pending' AND EXCLUDED.status IN ('available', 'unused') \
|
||||
THEN request_candidates.error_message \
|
||||
WHEN request_candidates.status = 'streaming' AND EXCLUDED.status IN ('available', 'unused', 'pending') \
|
||||
THEN request_candidates.error_message \
|
||||
ELSE COALESCE(EXCLUDED.error_message, request_candidates.error_message) END",
|
||||
)
|
||||
.replace(
|
||||
"__AETHER_CANDIDATE_EXTRA_DATA__",
|
||||
"CASE WHEN request_candidates.status IN ('success', 'failed', 'cancelled', 'skipped') \
|
||||
AND (EXCLUDED.status <> request_candidates.status OR EXCLUDED.extra_data IS NULL) \
|
||||
THEN request_candidates.extra_data ELSE EXCLUDED.extra_data END",
|
||||
)
|
||||
}
|
||||
|
||||
fn postgres_sanitized_legacy_diagnostic_sql(
|
||||
@@ -1381,33 +1398,35 @@ mod tests {
|
||||
finished_at_unix_ms: Some(2),
|
||||
};
|
||||
|
||||
assert_eq!(sanitize_request_candidate_for_postgres(&mut candidate), 0);
|
||||
assert_eq!(sanitize_request_candidate_for_postgres(&mut candidate), 1);
|
||||
assert!(candidate.username.is_none());
|
||||
assert!(candidate.api_key_name.is_none());
|
||||
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
||||
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
||||
assert!(candidate.error_message.is_none());
|
||||
assert_eq!(candidate.error_message.as_deref(), Some("bad�message"));
|
||||
assert!(candidate.extra_data.is_none());
|
||||
assert!(candidate.required_capabilities.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_postgres_candidate_conflict_path_discards_legacy_diagnostics() {
|
||||
fn every_postgres_candidate_conflict_path_preserves_errors_without_unrelated_legacy_data() {
|
||||
for sql in [
|
||||
UPSERT_SQL.as_str(),
|
||||
UPSERT_CONFLICT_SQL.as_str(),
|
||||
UPSERT_CONFLICT_INHERIT_IS_CACHED_SQL.as_str(),
|
||||
] {
|
||||
assert!(sql.contains("error_message = NULL"));
|
||||
assert!(sql.contains("extra_data = EXCLUDED.extra_data"));
|
||||
assert!(
|
||||
sql.contains("COALESCE(EXCLUDED.error_message, request_candidates.error_message)")
|
||||
);
|
||||
assert!(sql.contains("THEN request_candidates.extra_data ELSE EXCLUDED.extra_data END"));
|
||||
assert!(sql.contains("required_capabilities = EXCLUDED.required_capabilities"));
|
||||
assert!(!sql.contains("request_candidates.error_message"));
|
||||
assert!(!sql.contains("request_candidates.extra_data"));
|
||||
assert!(!sql.contains("COALESCE(request_candidates.extra_data"));
|
||||
assert!(!sql.contains("request_candidates.required_capabilities"));
|
||||
assert!(sql.contains("ELSE 'unclassified_skip' END"));
|
||||
assert!(sql.contains("ELSE 'unclassified_error' END"));
|
||||
assert!(sql.contains("THEN 'first_byte_timeout'"));
|
||||
assert!(!sql.contains("__AETHER_SANITIZED_LEGACY_"));
|
||||
assert!(!sql.contains("__AETHER_CANDIDATE_"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1541,10 +1560,11 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW())
|
||||
.fetch_one(repository.pool())
|
||||
.await
|
||||
.expect("raw candidate diagnostics should load");
|
||||
assert!(
|
||||
assert_eq!(
|
||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "error_message")
|
||||
.expect("error_message should decode")
|
||||
.is_none()
|
||||
.as_deref(),
|
||||
Some("bad�message")
|
||||
);
|
||||
assert_eq!(
|
||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "skip_reason")
|
||||
@@ -1576,7 +1596,7 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW())
|
||||
.expect("sanitized candidate should be readable");
|
||||
assert_eq!(rows.len(), 1);
|
||||
assert_eq!(rows[0].status, RequestCandidateStatus::Success);
|
||||
assert!(rows[0].error_message.is_none());
|
||||
assert_eq!(rows[0].error_message.as_deref(), Some("bad�message"));
|
||||
assert!(rows[0].extra_data.is_none());
|
||||
assert!(rows[0].required_capabilities.is_none());
|
||||
}
|
||||
@@ -1589,12 +1609,92 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW())
|
||||
1
|
||||
);
|
||||
|
||||
let mut cleanup_request_ids = vec![single_request_id, batch_request_id, healthy_request_id];
|
||||
for write_path in 0..3 {
|
||||
let request_id = uuid::Uuid::new_v4().to_string();
|
||||
cleanup_request_ids.push(request_id.clone());
|
||||
let mut failed = candidate(&request_id, uuid::Uuid::new_v4().to_string());
|
||||
failed.status = RequestCandidateStatus::Failed;
|
||||
failed.status_code = Some(400);
|
||||
failed.error_message = Some("original upstream failure".to_string());
|
||||
failed.is_cached = (write_path != 2).then_some(false);
|
||||
failed.extra_data = Some(json!({
|
||||
"upstream_response": {
|
||||
"status_code": 400,
|
||||
"headers": {"x-request-id": "original-upstream-id"},
|
||||
"body": {"error": {"message": "original upstream failure", "param": "model"}}
|
||||
},
|
||||
"error_flow": {"status_code": 400, "message": "original upstream failure"}
|
||||
}));
|
||||
let mut pending = failed.clone();
|
||||
pending.status = RequestCandidateStatus::Pending;
|
||||
pending.status_code = None;
|
||||
pending.error_message = None;
|
||||
pending.extra_data = None;
|
||||
repository
|
||||
.upsert(pending.clone())
|
||||
.await
|
||||
.expect("pending seed should persist");
|
||||
if write_path == 0 {
|
||||
repository
|
||||
.upsert(failed)
|
||||
.await
|
||||
.expect("single failure should persist");
|
||||
} else {
|
||||
repository
|
||||
.upsert_many(vec![failed])
|
||||
.await
|
||||
.expect("batch failure should persist");
|
||||
}
|
||||
pending.status_code = Some(200);
|
||||
pending.error_message = Some("late unrelated error".to_string());
|
||||
pending.extra_data = Some(json!({
|
||||
"upstream_response": {"status_code": 200, "body": "late unrelated response"}
|
||||
}));
|
||||
if write_path == 0 {
|
||||
repository
|
||||
.upsert(pending)
|
||||
.await
|
||||
.expect("late update should persist");
|
||||
} else {
|
||||
repository
|
||||
.upsert_many(vec![pending])
|
||||
.await
|
||||
.expect("late batch should persist");
|
||||
}
|
||||
let stored = repository
|
||||
.list_by_request_id(&request_id)
|
||||
.await
|
||||
.expect("failure should read");
|
||||
assert_eq!(stored[0].status, RequestCandidateStatus::Failed);
|
||||
assert_eq!(stored[0].status_code, Some(400));
|
||||
assert_eq!(
|
||||
stored[0].error_message.as_deref(),
|
||||
Some("original upstream failure")
|
||||
);
|
||||
let extra = stored[0]
|
||||
.extra_data
|
||||
.as_ref()
|
||||
.expect("failure details should remain");
|
||||
assert_eq!(extra["upstream_response"]["status_code"], 400);
|
||||
assert_eq!(
|
||||
extra["upstream_response"]["headers"]["x-request-id"],
|
||||
"original-upstream-id"
|
||||
);
|
||||
assert_eq!(
|
||||
extra["upstream_response"]["body"]["error"]["param"],
|
||||
"model"
|
||||
);
|
||||
assert_eq!(extra["error_flow"]["message"], "original upstream failure");
|
||||
let mut public = stored[0].clone();
|
||||
public.sanitize_sensitive_diagnostics();
|
||||
assert!(!serde_json::to_string(&public)
|
||||
.expect("public record should serialize")
|
||||
.contains("original upstream failure"));
|
||||
}
|
||||
|
||||
sqlx::query("DELETE FROM request_candidates WHERE request_id = ANY($1)")
|
||||
.bind(vec![
|
||||
single_request_id,
|
||||
batch_request_id,
|
||||
healthy_request_id,
|
||||
])
|
||||
.bind(cleanup_request_ids)
|
||||
.execute(repository.pool())
|
||||
.await
|
||||
.expect("candidate NUL test rows should clean up");
|
||||
|
||||
@@ -4,6 +4,7 @@ pub use types::{
|
||||
build_decision_trace, derive_request_candidate_final_status,
|
||||
request_candidate_lifecycle_would_regress, sanitize_request_candidate_api_formats,
|
||||
sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data,
|
||||
sanitize_request_candidate_extra_data_for_persistence,
|
||||
sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason,
|
||||
DecisionTrace, DecisionTraceCandidate, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
RequestCandidateFinalStatus, RequestCandidateReadRepository, RequestCandidateRepository,
|
||||
|
||||
@@ -224,6 +224,21 @@ pub struct StoredRequestCandidate {
|
||||
}
|
||||
|
||||
impl StoredRequestCandidate {
|
||||
pub fn sanitize_for_persistence(&mut self) {
|
||||
self.username = None;
|
||||
self.api_key_name = None;
|
||||
self.skip_reason = sanitize_request_candidate_skip_reason(self.skip_reason.take());
|
||||
self.error_type = sanitize_request_candidate_error_type(self.error_type.take());
|
||||
self.error_message = self
|
||||
.error_message
|
||||
.take()
|
||||
.map(limit_candidate_diagnostic_text);
|
||||
self.extra_data =
|
||||
sanitize_request_candidate_extra_data_for_persistence(self.extra_data.take());
|
||||
self.required_capabilities =
|
||||
sanitize_request_candidate_required_capabilities(self.required_capabilities.take());
|
||||
}
|
||||
|
||||
pub fn sanitize_sensitive_diagnostics(&mut self) {
|
||||
self.username = None;
|
||||
self.api_key_name = None;
|
||||
@@ -349,7 +364,7 @@ impl StoredRequestCandidate {
|
||||
started_at_unix_ms,
|
||||
finished_at_unix_ms,
|
||||
};
|
||||
candidate.sanitize_sensitive_diagnostics();
|
||||
candidate.sanitize_for_persistence();
|
||||
Ok(candidate)
|
||||
}
|
||||
}
|
||||
@@ -386,7 +401,7 @@ impl RequestCandidateTrace {
|
||||
attempted_only: bool,
|
||||
) -> Option<Self> {
|
||||
for candidate in &mut all_candidates {
|
||||
candidate.sanitize_sensitive_diagnostics();
|
||||
candidate.sanitize_for_persistence();
|
||||
}
|
||||
if all_candidates.is_empty() {
|
||||
return None;
|
||||
@@ -510,8 +525,17 @@ pub struct DecisionTraceCandidate {
|
||||
}
|
||||
|
||||
impl DecisionTraceCandidate {
|
||||
pub fn sanitize_for_admin(&mut self) {
|
||||
self.candidate.sanitize_for_persistence();
|
||||
self.sanitize_catalog_metadata();
|
||||
}
|
||||
|
||||
pub fn sanitize_sensitive_diagnostics(&mut self) {
|
||||
self.candidate.sanitize_sensitive_diagnostics();
|
||||
self.sanitize_catalog_metadata();
|
||||
}
|
||||
|
||||
fn sanitize_catalog_metadata(&mut self) {
|
||||
self.provider_website = self
|
||||
.provider_website
|
||||
.take()
|
||||
@@ -574,7 +598,9 @@ pub fn build_decision_trace(
|
||||
})
|
||||
.collect(),
|
||||
};
|
||||
trace.sanitize_sensitive_diagnostics();
|
||||
for item in &mut trace.candidates {
|
||||
item.sanitize_for_admin();
|
||||
}
|
||||
trace
|
||||
}
|
||||
|
||||
@@ -727,8 +753,12 @@ impl UpsertRequestCandidateRecord {
|
||||
self.api_key_name = None;
|
||||
self.skip_reason = sanitize_request_candidate_skip_reason(self.skip_reason.take());
|
||||
self.error_type = sanitize_request_candidate_error_type(self.error_type.take());
|
||||
self.error_message = None;
|
||||
self.extra_data = sanitize_request_candidate_extra_data(self.extra_data.take());
|
||||
self.error_message = self
|
||||
.error_message
|
||||
.take()
|
||||
.map(limit_candidate_diagnostic_text);
|
||||
self.extra_data =
|
||||
sanitize_request_candidate_extra_data_for_persistence(self.extra_data.take());
|
||||
self.required_capabilities =
|
||||
sanitize_request_candidate_required_capabilities(self.required_capabilities.take());
|
||||
}
|
||||
@@ -788,6 +818,78 @@ pub fn sanitize_request_candidate_error_type(value: Option<String>) -> Option<St
|
||||
Some(safe)
|
||||
}
|
||||
|
||||
const MAX_CANDIDATE_DIAGNOSTIC_BYTES: usize = 65_536;
|
||||
|
||||
fn limit_candidate_diagnostic_text(mut text: String) -> String {
|
||||
const SUFFIX: &str = "...[truncated]";
|
||||
if text.len() > MAX_CANDIDATE_DIAGNOSTIC_BYTES {
|
||||
let mut boundary = MAX_CANDIDATE_DIAGNOSTIC_BYTES - SUFFIX.len();
|
||||
while !text.is_char_boundary(boundary) {
|
||||
boundary -= 1;
|
||||
}
|
||||
text.truncate(boundary);
|
||||
text.push_str(SUFFIX);
|
||||
}
|
||||
text
|
||||
}
|
||||
|
||||
fn limit_candidate_diagnostic_value(value: &serde_json::Value) -> serde_json::Value {
|
||||
if let Some(text) = value.as_str() {
|
||||
return serde_json::Value::String(limit_candidate_diagnostic_text(text.to_string()));
|
||||
}
|
||||
let serialized = value.to_string();
|
||||
if serialized.len() > MAX_CANDIDATE_DIAGNOSTIC_BYTES {
|
||||
serde_json::Value::String(limit_candidate_diagnostic_text(serialized))
|
||||
} else {
|
||||
value.clone()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn sanitize_request_candidate_extra_data_for_persistence(
|
||||
extra_data: Option<serde_json::Value>,
|
||||
) -> Option<serde_json::Value> {
|
||||
let object = extra_data.as_ref()?.as_object()?;
|
||||
let mut sanitized = sanitize_request_candidate_extra_data(extra_data.clone())
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default();
|
||||
for (key, fields) in [
|
||||
("upstream_response", &["headers", "body"][..]),
|
||||
("error_flow", &["message"][..]),
|
||||
(
|
||||
"failure_diagnostic",
|
||||
&["path", "field_path", "message", "type", "reason"][..],
|
||||
),
|
||||
(
|
||||
"request_conversion_error",
|
||||
&["path", "field_path", "message", "type", "reason"][..],
|
||||
),
|
||||
(
|
||||
"request_body_build_error",
|
||||
&["path", "field_path", "message", "type", "reason"][..],
|
||||
),
|
||||
] {
|
||||
let Some(diagnostic) = object.get(key).and_then(serde_json::Value::as_object) else {
|
||||
continue;
|
||||
};
|
||||
let mut summary = sanitized
|
||||
.remove(key)
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default();
|
||||
for field in fields {
|
||||
if let Some(value) = diagnostic.get(*field).filter(|value| !value.is_null()) {
|
||||
summary.insert(
|
||||
(*field).to_string(),
|
||||
limit_candidate_diagnostic_value(value),
|
||||
);
|
||||
}
|
||||
}
|
||||
if !summary.is_empty() {
|
||||
sanitized.insert(key.to_string(), serde_json::Value::Object(summary));
|
||||
}
|
||||
}
|
||||
(!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized))
|
||||
}
|
||||
|
||||
pub fn sanitize_request_candidate_extra_data(
|
||||
extra_data: Option<serde_json::Value>,
|
||||
) -> Option<serde_json::Value> {
|
||||
@@ -1949,7 +2051,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn candidate_persistence_removes_credentials_and_raw_payloads() {
|
||||
fn candidate_persistence_keeps_admin_errors_but_removes_request_credentials() {
|
||||
let mut record = UpsertRequestCandidateRecord {
|
||||
id: "candidate-1".to_string(),
|
||||
request_id: "request-1".to_string(),
|
||||
@@ -2079,7 +2181,7 @@ mod tests {
|
||||
);
|
||||
assert!(record.username.is_none());
|
||||
assert!(record.api_key_name.is_none());
|
||||
assert!(record.error_message.is_none());
|
||||
assert_eq!(record.error_message.as_deref(), Some("unauthorized"));
|
||||
let extra = record
|
||||
.extra_data
|
||||
.as_ref()
|
||||
@@ -2112,7 +2214,10 @@ mod tests {
|
||||
assert_eq!(extra["error_flow"]["stage"], "upstream");
|
||||
assert_eq!(extra["error_flow"]["retryable"], true);
|
||||
assert_eq!(extra["error_flow"]["status_code"], 401);
|
||||
assert!(extra["error_flow"].get("message").is_none());
|
||||
assert_eq!(
|
||||
extra["error_flow"]["message"],
|
||||
"token vertex-secret rejected"
|
||||
);
|
||||
assert_eq!(extra["gateway_execution_runtime"], true);
|
||||
assert_eq!(extra["client_api_format"], "openai:responses");
|
||||
assert_eq!(extra["provider_api_format"], "claude:messages");
|
||||
@@ -2127,8 +2232,14 @@ mod tests {
|
||||
assert_eq!(extra["upstream_response"]["source"], "upstream_response");
|
||||
assert_eq!(extra["upstream_response"]["status_code"], 401);
|
||||
assert_eq!(extra["upstream_response"]["body_state"], "inline");
|
||||
assert!(extra["upstream_response"].get("headers").is_none());
|
||||
assert!(extra["upstream_response"].get("body").is_none());
|
||||
assert_eq!(
|
||||
extra["upstream_response"]["headers"]["set-cookie"],
|
||||
"session=secret"
|
||||
);
|
||||
assert_eq!(
|
||||
extra["upstream_response"]["body"]["error"]["message"],
|
||||
"token vertex-secret rejected"
|
||||
);
|
||||
assert_eq!(
|
||||
extra["image_progress"]["last_client_visible_event"],
|
||||
"image_generation.partial_image"
|
||||
@@ -2178,7 +2289,6 @@ mod tests {
|
||||
|
||||
let serialized = serde_json::to_string(&record).expect("candidate should serialize");
|
||||
for sensitive in [
|
||||
"vertex-secret",
|
||||
"client-secret",
|
||||
"credential-label-secret",
|
||||
"header-rule-secret",
|
||||
@@ -2199,6 +2309,11 @@ mod tests {
|
||||
"candidate must not retain {sensitive}"
|
||||
);
|
||||
}
|
||||
let public_extra = super::sanitize_request_candidate_extra_data(record.extra_data);
|
||||
let serialized =
|
||||
serde_json::to_string(&public_extra).expect("public data should serialize");
|
||||
assert!(!serialized.contains("vertex-secret"));
|
||||
assert!(!serialized.contains("session=secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -2278,8 +2393,8 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn candidate_database_read_sanitizes_legacy_diagnostic_text() {
|
||||
let candidate = StoredRequestCandidate::new(
|
||||
fn candidate_database_read_preserves_errors_until_public_projection() {
|
||||
let mut candidate = StoredRequestCandidate::new(
|
||||
"candidate-1".to_string(),
|
||||
"request-1".to_string(),
|
||||
None,
|
||||
@@ -2315,9 +2430,42 @@ mod tests {
|
||||
candidate.error_type.as_deref(),
|
||||
Some(UNCLASSIFIED_CANDIDATE_ERROR_TYPE)
|
||||
);
|
||||
assert!(candidate.error_message.is_none());
|
||||
assert_eq!(
|
||||
candidate.error_message.as_deref(),
|
||||
Some("legacy secret message")
|
||||
);
|
||||
assert!(candidate.username.is_none());
|
||||
assert!(candidate.api_key_name.is_none());
|
||||
candidate.sanitize_sensitive_diagnostics();
|
||||
assert!(candidate.error_message.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn admin_diagnostics_are_bounded_and_public_projection_removes_them() {
|
||||
let raw = json!({
|
||||
"upstream_response": {"status_code": 400, "body": "错误内容".repeat(20_000)},
|
||||
"error_flow": {"status_code": 400, "message": "private upstream failure"},
|
||||
"failure_diagnostic": {"path": "$.input", "message": "private conversion failure"},
|
||||
"request_body": {"input": "private prompt"}
|
||||
});
|
||||
let admin = super::sanitize_request_candidate_extra_data_for_persistence(Some(raw))
|
||||
.expect("admin diagnostics should remain");
|
||||
let body = admin["upstream_response"]["body"]
|
||||
.as_str()
|
||||
.expect("body should be text");
|
||||
assert!(body.len() <= 65_536);
|
||||
assert!(body.ends_with("...[truncated]"));
|
||||
assert!(admin.get("request_body").is_none());
|
||||
assert_eq!(
|
||||
super::sanitize_request_candidate_extra_data_for_persistence(Some(admin.clone())),
|
||||
Some(admin.clone()),
|
||||
);
|
||||
let public = super::sanitize_request_candidate_extra_data(Some(admin))
|
||||
.expect("public status should remain");
|
||||
assert_eq!(public["upstream_response"]["status_code"], 400);
|
||||
assert!(public["upstream_response"].get("body").is_none());
|
||||
assert!(public["error_flow"].get("message").is_none());
|
||||
assert!(public.get("failure_diagnostic").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -10,19 +10,31 @@ use crate::DataLayerError;
|
||||
use async_trait::async_trait;
|
||||
|
||||
fn sanitize_stored_candidate(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
|
||||
candidate.sanitize_sensitive_diagnostics();
|
||||
candidate.sanitize_for_persistence();
|
||||
candidate
|
||||
}
|
||||
|
||||
fn merge_extra_data(
|
||||
existing: Option<serde_json::Value>,
|
||||
overlay: Option<serde_json::Value>,
|
||||
preserve_error_details: bool,
|
||||
) -> Option<serde_json::Value> {
|
||||
match (existing, overlay) {
|
||||
(
|
||||
Some(serde_json::Value::Object(mut existing_object)),
|
||||
Some(serde_json::Value::Object(overlay_object)),
|
||||
Some(serde_json::Value::Object(mut overlay_object)),
|
||||
) => {
|
||||
if preserve_error_details {
|
||||
for key in [
|
||||
"upstream_response",
|
||||
"error_flow",
|
||||
"failure_diagnostic",
|
||||
"request_conversion_error",
|
||||
"request_body_build_error",
|
||||
] {
|
||||
overlay_object.remove(key);
|
||||
}
|
||||
}
|
||||
existing_object.extend(overlay_object);
|
||||
Some(serde_json::Value::Object(existing_object))
|
||||
}
|
||||
@@ -397,7 +409,13 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
.error_type
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.error_type.clone()))
|
||||
},
|
||||
error_message: None,
|
||||
error_message: if preserve_existing_lifecycle {
|
||||
existing.as_ref().and_then(|row| row.error_message.clone())
|
||||
} else {
|
||||
candidate
|
||||
.error_message
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.error_message.clone()))
|
||||
},
|
||||
latency_ms: if preserve_existing_lifecycle {
|
||||
existing.as_ref().and_then(|row| row.latency_ms)
|
||||
} else {
|
||||
@@ -411,6 +429,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
extra_data: merge_extra_data(
|
||||
existing.as_ref().and_then(|row| row.extra_data.clone()),
|
||||
candidate.extra_data,
|
||||
preserve_existing_lifecycle,
|
||||
),
|
||||
required_capabilities: candidate.required_capabilities.or_else(|| {
|
||||
existing
|
||||
@@ -589,7 +608,10 @@ mod tests {
|
||||
let candidate = stored
|
||||
.get("cand-raw")
|
||||
.expect("seeded candidate should exist");
|
||||
assert!(candidate.error_message.is_none());
|
||||
assert_eq!(
|
||||
candidate.error_message.as_deref(),
|
||||
Some("Bearer secret-token")
|
||||
);
|
||||
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
||||
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
||||
assert_eq!(
|
||||
@@ -619,7 +641,10 @@ mod tests {
|
||||
.iter()
|
||||
.find(|candidate| candidate.id == "cand-bypassed")
|
||||
.expect("bypassed candidate should be returned");
|
||||
assert!(candidate.error_message.is_none());
|
||||
assert_eq!(
|
||||
candidate.error_message.as_deref(),
|
||||
Some("Bearer secret-token")
|
||||
);
|
||||
assert_eq!(
|
||||
candidate.extra_data,
|
||||
Some(json!({"gateway_execution_runtime": true}))
|
||||
@@ -665,7 +690,7 @@ mod tests {
|
||||
.await
|
||||
.expect("candidate merge should succeed");
|
||||
assert_eq!(merged.id, "cand-bypassed");
|
||||
assert!(merged.error_message.is_none());
|
||||
assert_eq!(merged.error_message.as_deref(), Some("Bearer secret-token"));
|
||||
assert_eq!(
|
||||
merged.extra_data,
|
||||
Some(json!({
|
||||
@@ -839,7 +864,13 @@ mod tests {
|
||||
Some("retryable upstream failure".to_string()),
|
||||
Some(45),
|
||||
Some(1),
|
||||
Some(json!({"stream_completed": true})),
|
||||
Some(json!({
|
||||
"stream_completed": true,
|
||||
"upstream_response": {
|
||||
"status_code": 503,
|
||||
"body": {"error": {"message": "original upstream failure"}}
|
||||
}
|
||||
})),
|
||||
None,
|
||||
100,
|
||||
Some(101),
|
||||
@@ -869,7 +900,10 @@ mod tests {
|
||||
error_message: None,
|
||||
latency_ms: Some(9_999),
|
||||
concurrent_requests: Some(2),
|
||||
extra_data: Some(json!({"gateway_execution_runtime": true})),
|
||||
extra_data: Some(json!({
|
||||
"gateway_execution_runtime": true,
|
||||
"upstream_response": {"status_code": 200, "body": "late unrelated response"}
|
||||
})),
|
||||
required_capabilities: None,
|
||||
created_at_unix_ms: None,
|
||||
started_at_unix_ms: Some(102),
|
||||
@@ -882,7 +916,10 @@ mod tests {
|
||||
assert_eq!(updated.status, RequestCandidateStatus::Failed);
|
||||
assert_eq!(updated.status_code, Some(503));
|
||||
assert_eq!(updated.error_type.as_deref(), Some("upstream_error"));
|
||||
assert!(updated.error_message.is_none());
|
||||
assert_eq!(
|
||||
updated.error_message.as_deref(),
|
||||
Some("retryable upstream failure")
|
||||
);
|
||||
assert_eq!(updated.latency_ms, Some(45));
|
||||
assert_eq!(updated.concurrent_requests, Some(2));
|
||||
assert_eq!(updated.finished_at_unix_ms, Some(145));
|
||||
@@ -890,7 +927,11 @@ mod tests {
|
||||
updated.extra_data,
|
||||
Some(json!({
|
||||
"gateway_execution_runtime": true,
|
||||
"stream_completed": true
|
||||
"stream_completed": true,
|
||||
"upstream_response": {
|
||||
"status_code": 503,
|
||||
"body": {"error": {"message": "original upstream failure"}}
|
||||
}
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
@@ -238,7 +238,22 @@ impl GenericProviderOAuthAdapter {
|
||||
return required_client_secret(env_name, non_empty_owned(value)).map(Some);
|
||||
}
|
||||
|
||||
required_client_secret(env_name, non_empty_environment_value(env_name)).map(Some)
|
||||
let configured = non_empty_environment_value(env_name);
|
||||
self.resolve_client_secret(configured).map(Some)
|
||||
}
|
||||
|
||||
fn resolve_client_secret(&self, configured: Option<String>) -> Result<String, OAuthError> {
|
||||
let configured = configured.or_else(|| {
|
||||
(self.template.provider_type == "antigravity"
|
||||
&& self.client_id() == self.template.client_id)
|
||||
.then(|| "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".to_string())
|
||||
});
|
||||
required_client_secret(
|
||||
self.template
|
||||
.client_secret_env
|
||||
.unwrap_or(ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV),
|
||||
configured,
|
||||
)
|
||||
}
|
||||
|
||||
async fn exchange_grant(
|
||||
@@ -957,6 +972,57 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn antigravity_default_credentials_support_authorization_and_refresh() {
|
||||
let mut template = template_for_provider_type("antigravity").expect("template");
|
||||
template.client_id_env = None;
|
||||
template.client_secret_env = Some("AETHER_TEST_UNUSED_ANTIGRAVITY_SECRET");
|
||||
let adapter = GenericProviderOAuthAdapter::new(template);
|
||||
let ctx = transport_context("antigravity");
|
||||
adapter
|
||||
.build_authorize_url(&ctx, "state", None)
|
||||
.expect("authorize");
|
||||
let seen_request = Arc::new(Mutex::new(None));
|
||||
let executor = StaticExecutor {
|
||||
seen_request: Arc::clone(&seen_request),
|
||||
response_payload: json!({"access_token": "new-token", "expires_in": 3600}),
|
||||
};
|
||||
adapter
|
||||
.refresh(&executor, &ctx, &oauth_account("antigravity"))
|
||||
.await
|
||||
.expect("refresh");
|
||||
let seen = seen_request.lock().expect("lock").clone().expect("request");
|
||||
let body = seen.body_bytes.expect("body");
|
||||
let fields = url::form_urlencoded::parse(&body)
|
||||
.into_owned()
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
assert_eq!(fields["client_id"], template.client_id);
|
||||
assert_eq!(
|
||||
fields["client_secret"],
|
||||
adapter.resolve_client_secret(None).expect("default")
|
||||
);
|
||||
assert_eq!(fields["refresh_token"], "old-refresh-token");
|
||||
assert!(!format!("{adapter:?}").contains(&fields["client_secret"]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn antigravity_custom_client_requires_its_own_secret() {
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type("antigravity")
|
||||
.expect("adapter")
|
||||
.with_oauth_credentials_for_tests("custom-client", "custom-secret");
|
||||
assert!(adapter.resolve_client_secret(None).is_err());
|
||||
assert_eq!(
|
||||
adapter
|
||||
.resolve_client_secret(Some("custom-secret".to_string()))
|
||||
.expect("configured secret"),
|
||||
"custom-secret"
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.client_secret().expect("override"),
|
||||
Some("custom-secret".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generic_adapter_debug_redacts_oauth_credentials() {
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli")
|
||||
|
||||
@@ -17,6 +17,15 @@
|
||||
|
||||
`total_keys` 和 `active_keys` 仍反映密钥配置数量,不因缺少健康数据而减少。
|
||||
|
||||
## 数据读取
|
||||
|
||||
摘要从密钥的轻量投影读取 API 格式、启用状态和 `health_by_format`。
|
||||
PostgreSQL 投影中的凭据字段使用 `summary` / `{}` 等脱敏占位值,并非真实密文,
|
||||
因此摘要读取不执行凭据解密、认证或迁移。完整密钥读取仍保留原有的凭据安全校验。
|
||||
|
||||
若将这些占位值送入凭据校验,读取会失败,并被摘要聚合当作空密钥列表,
|
||||
导致已配置密钥的端点也被错误显示为灰色;不能通过给缺失分数默认填 `100%` 来修复。
|
||||
|
||||
## 页面展示
|
||||
|
||||
桌面表格和手机卡片使用相同规则:
|
||||
|
||||
@@ -392,8 +392,7 @@ class ApiClient {
|
||||
|
||||
this.isRefreshing = true
|
||||
const requestAuthStateVersion = this.authStateVersion
|
||||
let restorePromise!: Promise<string>
|
||||
restorePromise = (async () => {
|
||||
const restorePromise = (async () => {
|
||||
const accessToken = await this.coordinatedRefresh()
|
||||
if (requestAuthStateVersion !== this.authStateVersion) {
|
||||
throw new Error('Auth state changed during session restore')
|
||||
|
||||
@@ -6,6 +6,7 @@
|
||||
|
||||
<script setup lang="ts">
|
||||
import { ref, onMounted, onUnmounted, watch, nextTick } from 'vue'
|
||||
import { useI18n } from '@/i18n'
|
||||
import {
|
||||
Chart as ChartJS,
|
||||
CategoryScale,
|
||||
@@ -43,6 +44,7 @@ interface Props {
|
||||
}
|
||||
|
||||
const chartRef = ref<HTMLCanvasElement>()
|
||||
const { locale } = useI18n()
|
||||
let chart: ChartJS<'bar'> | null = null
|
||||
|
||||
const defaultOptions: ChartOptions<'bar'> = {
|
||||
@@ -91,9 +93,7 @@ const defaultOptions: ChartOptions<'bar'> = {
|
||||
}
|
||||
}
|
||||
|
||||
function createChart() {
|
||||
if (!chartRef.value) return
|
||||
|
||||
function buildChartOptions(): ChartOptions<'bar'> {
|
||||
const stackedOptions = props.stacked ? {
|
||||
scales: {
|
||||
x: { ...defaultOptions.scales?.x, stacked: true },
|
||||
@@ -106,14 +106,21 @@ function createChart() {
|
||||
}
|
||||
}
|
||||
|
||||
return {
|
||||
...defaultOptions,
|
||||
...stackedOptions,
|
||||
locale: locale.value,
|
||||
...props.options
|
||||
}
|
||||
}
|
||||
|
||||
function createChart() {
|
||||
if (!chartRef.value) return
|
||||
|
||||
chart = new ChartJS(chartRef.value, {
|
||||
type: 'bar',
|
||||
data: props.data,
|
||||
options: {
|
||||
...defaultOptions,
|
||||
...stackedOptions,
|
||||
...props.options
|
||||
}
|
||||
options: buildChartOptions()
|
||||
})
|
||||
}
|
||||
|
||||
@@ -137,12 +144,9 @@ onUnmounted(() => {
|
||||
})
|
||||
|
||||
watch(() => props.data, updateChart, { deep: true })
|
||||
watch(() => props.options, () => {
|
||||
watch([() => props.options, () => props.stacked, locale], () => {
|
||||
if (chart) {
|
||||
chart.options = {
|
||||
...defaultOptions,
|
||||
...props.options
|
||||
}
|
||||
chart.options = buildChartOptions()
|
||||
chart.update()
|
||||
}
|
||||
}, { deep: true })
|
||||
|
||||
@@ -6,6 +6,7 @@
|
||||
|
||||
<script setup lang="ts">
|
||||
import { ref, onMounted, onUnmounted, watch, nextTick } from 'vue'
|
||||
import { useI18n } from '@/i18n'
|
||||
import {
|
||||
Chart as ChartJS,
|
||||
ArcElement,
|
||||
@@ -37,6 +38,7 @@ interface Props {
|
||||
}
|
||||
|
||||
const chartRef = ref<HTMLCanvasElement>()
|
||||
const { locale } = useI18n()
|
||||
let chart: ChartJS<'doughnut'> | null = null
|
||||
|
||||
const defaultOptions: ChartOptions<'doughnut'> = {
|
||||
@@ -79,6 +81,7 @@ function createChart() {
|
||||
data: props.data,
|
||||
options: {
|
||||
...defaultOptions,
|
||||
locale: locale.value,
|
||||
...props.options
|
||||
}
|
||||
})
|
||||
@@ -104,9 +107,9 @@ onUnmounted(() => {
|
||||
})
|
||||
|
||||
watch(() => props.data, updateChart, { deep: true })
|
||||
watch(() => props.options, () => {
|
||||
watch([() => props.options, locale], () => {
|
||||
if (chart) {
|
||||
chart.options = { ...defaultOptions, ...props.options }
|
||||
chart.options = { ...defaultOptions, locale: locale.value, ...props.options }
|
||||
chart.update()
|
||||
}
|
||||
}, { deep: true })
|
||||
|
||||
@@ -6,6 +6,7 @@
|
||||
|
||||
<script setup lang="ts">
|
||||
import { ref, onMounted, onUnmounted, watch, nextTick } from 'vue'
|
||||
import { useI18n } from '@/i18n'
|
||||
import {
|
||||
Chart as ChartJS,
|
||||
CategoryScale,
|
||||
@@ -44,11 +45,13 @@ interface Props {
|
||||
}
|
||||
|
||||
const chartRef = ref<HTMLCanvasElement>()
|
||||
const { locale } = useI18n()
|
||||
let chart: ChartJS<'line'> | null = null
|
||||
|
||||
function buildChartOptions(): ChartOptions<'line'> {
|
||||
return {
|
||||
...defaultOptions,
|
||||
locale: locale.value,
|
||||
...props.options
|
||||
}
|
||||
}
|
||||
@@ -121,7 +124,7 @@ onUnmounted(() => {
|
||||
|
||||
// 监听引用变化,避免深监听触发整图重算
|
||||
watch(() => props.data, updateChart)
|
||||
watch(() => props.options, () => {
|
||||
watch([() => props.options, locale], () => {
|
||||
if (chart) {
|
||||
chart.options = buildChartOptions()
|
||||
chart.update('none')
|
||||
|
||||
@@ -3,17 +3,17 @@
|
||||
<canvas ref="chartRef" />
|
||||
<div
|
||||
v-if="crosshairStats"
|
||||
class="absolute top-2 right-2 bg-gray-800/90 text-gray-100 px-3 py-2 rounded-lg text-sm shadow-lg border border-gray-600"
|
||||
class="absolute top-2 right-2 max-w-[calc(100%-1rem)] break-words bg-gray-800/90 text-gray-100 px-3 py-2 rounded-lg text-sm shadow-lg border border-gray-600"
|
||||
>
|
||||
<div class="font-medium text-yellow-400">
|
||||
Y = {{ crosshairStats.yValue.toFixed(1) }} 分钟
|
||||
{{ t('chart.crosshairValue', { value: crosshairStats.yValue.toFixed(1) }) }}
|
||||
</div>
|
||||
<!-- 单个 dataset 时显示简单统计 -->
|
||||
<div
|
||||
v-if="crosshairStats.datasets.length === 1"
|
||||
class="mt-1"
|
||||
>
|
||||
<span class="text-green-400">{{ crosshairStats.datasets[0].belowCount }}</span> / {{ crosshairStats.datasets[0].totalCount }} 点在横线以下
|
||||
<span class="text-green-400">{{ crosshairStats.datasets[0].belowCount }}</span> / {{ crosshairStats.datasets[0].totalCount }} {{ t('chart.pointsBelow') }}
|
||||
<span class="ml-2 text-blue-400">({{ crosshairStats.datasets[0].belowPercent.toFixed(1) }}%)</span>
|
||||
</div>
|
||||
<!-- 多个 dataset 时按模型分别显示 -->
|
||||
@@ -36,7 +36,7 @@
|
||||
</div>
|
||||
<!-- 总计 -->
|
||||
<div class="flex items-center gap-2 pt-1 border-t border-gray-600 mt-1">
|
||||
<span class="text-gray-300">总计:</span>
|
||||
<span class="text-gray-300">{{ t('chart.total') }}:</span>
|
||||
<span class="text-green-400">{{ crosshairStats.totalBelowCount }}</span>/<span class="text-gray-400">{{ crosshairStats.totalCount }}</span>
|
||||
<span class="text-blue-400">({{ crosshairStats.totalBelowPercent.toFixed(1) }}%)</span>
|
||||
</div>
|
||||
@@ -46,6 +46,7 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { getI18nLocale, useI18n } from '@/i18n'
|
||||
import { ref, onMounted, onUnmounted, watch, nextTick, computed } from 'vue'
|
||||
import {
|
||||
Chart as ChartJS,
|
||||
@@ -62,6 +63,7 @@ import {
|
||||
type Scale
|
||||
} from 'chart.js'
|
||||
import 'chartjs-adapter-date-fns'
|
||||
import { enUS, zhCN } from 'date-fns/locale'
|
||||
|
||||
const props = withDefaults(defineProps<Props>(), {
|
||||
height: 300,
|
||||
@@ -114,6 +116,7 @@ interface GapInfo {
|
||||
}
|
||||
|
||||
const chartRef = ref<HTMLCanvasElement>()
|
||||
const { locale, t } = useI18n()
|
||||
let chart: ChartJS<'scatter'> | null = null
|
||||
|
||||
const crosshairY = ref<number | null>(null)
|
||||
@@ -151,7 +154,7 @@ const crosshairStats = computed<CrosshairStats | null>(() => {
|
||||
|
||||
if (dsTotal > 0) {
|
||||
datasetStats.push({
|
||||
label: dataset.label || 'Unknown',
|
||||
label: dataset.label || t('chart.unknown'),
|
||||
color: (dataset.backgroundColor as string) || 'rgba(59, 130, 246, 0.7)',
|
||||
belowCount,
|
||||
totalCount: dsTotal,
|
||||
@@ -325,7 +328,8 @@ function formatDuration(ms: number): string {
|
||||
return `${minutes}m`
|
||||
}
|
||||
|
||||
const defaultOptions: ChartOptions<'scatter'> = {
|
||||
const defaultOptions = computed<ChartOptions<'scatter'>>(() => ({
|
||||
locale: locale.value,
|
||||
responsive: true,
|
||||
maintainAspectRatio: false,
|
||||
interaction: {
|
||||
@@ -335,6 +339,9 @@ const defaultOptions: ChartOptions<'scatter'> = {
|
||||
scales: {
|
||||
x: {
|
||||
type: 'time',
|
||||
adapters: {
|
||||
date: { locale: locale.value === 'zh-CN' ? zhCN : enUS }
|
||||
},
|
||||
time: {
|
||||
displayFormats: {
|
||||
hour: 'HH:mm'
|
||||
@@ -378,7 +385,7 @@ const defaultOptions: ChartOptions<'scatter'> = {
|
||||
},
|
||||
title: {
|
||||
display: true,
|
||||
text: '间隔 (分钟)',
|
||||
text: t('chart.intervalAxis'),
|
||||
color: 'rgb(107, 114, 128)'
|
||||
},
|
||||
afterBuildTicks(scale: Scale) {
|
||||
@@ -407,7 +414,7 @@ const defaultOptions: ChartOptions<'scatter'> = {
|
||||
const point = contexts[0].raw as { x: string; _originalX?: string }
|
||||
const timeStr = point._originalX || point.x
|
||||
const date = new Date(timeStr)
|
||||
return date.toLocaleString('zh-CN', {
|
||||
return date.toLocaleString(getI18nLocale(), {
|
||||
month: 'numeric',
|
||||
day: 'numeric',
|
||||
hour: '2-digit',
|
||||
@@ -417,7 +424,7 @@ const defaultOptions: ChartOptions<'scatter'> = {
|
||||
label: (context) => {
|
||||
const point = context.raw as { x: string; y: number; _originalY?: number }
|
||||
const realY = point._originalY ?? toRealValue(point.y)
|
||||
return `间隔: ${realY.toFixed(1)} 分钟`
|
||||
return t('chart.intervalTooltip', { value: realY.toFixed(1) })
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -451,7 +458,7 @@ const defaultOptions: ChartOptions<'scatter'> = {
|
||||
|
||||
chartInstance.draw()
|
||||
}
|
||||
}
|
||||
}))
|
||||
|
||||
// 修改 crosshairPlugin 使用显示值
|
||||
const crosshairPluginWithTransform: Plugin<'scatter'> = {
|
||||
@@ -544,7 +551,7 @@ function createChart() {
|
||||
type: 'scatter',
|
||||
data: chartData,
|
||||
options: {
|
||||
...defaultOptions,
|
||||
...defaultOptions.value,
|
||||
...props.options
|
||||
},
|
||||
plugins: [crosshairPluginWithTransform, gapMarkerPlugin]
|
||||
@@ -586,10 +593,10 @@ watch(
|
||||
],
|
||||
updateChart
|
||||
)
|
||||
watch(() => props.options, () => {
|
||||
watch([() => props.options, defaultOptions], () => {
|
||||
if (chart) {
|
||||
chart.options = {
|
||||
...defaultOptions,
|
||||
...defaultOptions.value,
|
||||
...props.options
|
||||
}
|
||||
chart.update('none')
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
import { afterEach, describe, expect, it, vi } from 'vitest'
|
||||
import { createApp, h, nextTick, type App } from 'vue'
|
||||
import type { ChartConfiguration, ChartData, ChartOptions } from 'chart.js'
|
||||
|
||||
import BarChart from '../BarChart.vue'
|
||||
import ScatterChart from '../ScatterChart.vue'
|
||||
import CostForecastChart from '@/components/stats/CostForecastChart.vue'
|
||||
import { setI18nLocale } from '@/i18n'
|
||||
|
||||
const { chartConstructor } = vi.hoisted(() => ({ chartConstructor: vi.fn() }))
|
||||
|
||||
vi.mock('chartjs-adapter-date-fns', () => ({}))
|
||||
vi.mock('chart.js', async importOriginal => {
|
||||
const original = await importOriginal<typeof import('chart.js')>()
|
||||
return {
|
||||
...original,
|
||||
Chart: class {
|
||||
static register = vi.fn()
|
||||
data: ChartData
|
||||
options: ChartOptions
|
||||
update = vi.fn()
|
||||
destroy = vi.fn()
|
||||
|
||||
constructor(canvas: HTMLCanvasElement, config: ChartConfiguration) {
|
||||
this.data = config.data
|
||||
this.options = config.options ?? {}
|
||||
chartConstructor(canvas, config, this)
|
||||
}
|
||||
},
|
||||
}
|
||||
})
|
||||
|
||||
const mountedApps: Array<{ app: App, root: HTMLElement }> = []
|
||||
|
||||
async function mountChart(app: App) {
|
||||
const root = document.createElement('div')
|
||||
document.body.appendChild(root)
|
||||
app.mount(root)
|
||||
mountedApps.push({ app, root })
|
||||
await nextTick()
|
||||
await nextTick()
|
||||
}
|
||||
|
||||
function renderedChart() {
|
||||
return chartConstructor.mock.calls[chartConstructor.mock.calls.length - 1]?.[2] as {
|
||||
data: ChartData
|
||||
options: ChartOptions<'scatter'>
|
||||
update: ReturnType<typeof vi.fn>
|
||||
}
|
||||
}
|
||||
|
||||
afterEach(() => {
|
||||
for (const { app, root } of mountedApps.splice(0)) {
|
||||
app.unmount()
|
||||
root.remove()
|
||||
}
|
||||
chartConstructor.mockClear()
|
||||
})
|
||||
|
||||
describe('Chart locale updates', () => {
|
||||
it('redraws the scatter axis when the locale changes without changing its data', async () => {
|
||||
await mountChart(createApp({
|
||||
render: () => h(ScatterChart, {
|
||||
data: { datasets: [{ label: 'model-a', data: [{ x: 1000, y: 2 }] }] },
|
||||
}),
|
||||
}))
|
||||
const chart = renderedChart()
|
||||
const initialData = chart.data
|
||||
expect(chart.options.scales?.y?.title?.text).toBe('间隔 (分钟)')
|
||||
|
||||
setI18nLocale('en-US')
|
||||
await nextTick()
|
||||
|
||||
expect(chart.options.locale).toBe('en-US')
|
||||
expect(chart.options.scales?.y?.title?.text).toBe('Interval (minutes)')
|
||||
expect(chart.data).toBe(initialData)
|
||||
expect(chart.update).toHaveBeenCalledWith('none')
|
||||
})
|
||||
|
||||
it('updates forecast legend labels while preserving cost values', async () => {
|
||||
await mountChart(createApp({
|
||||
render: () => h(CostForecastChart, {
|
||||
title: 'Forecast',
|
||||
history: [{ date: '2026-09-01', total_cost: 12.5 }],
|
||||
forecast: [{ date: '2026-09-02', total_cost: 13 }],
|
||||
}),
|
||||
}))
|
||||
const chart = renderedChart()
|
||||
expect(chart.data.datasets.map(dataset => dataset.label)).toEqual(['实际成本', '预测成本'])
|
||||
|
||||
setI18nLocale('en-US')
|
||||
await nextTick()
|
||||
|
||||
expect(chart.data.datasets.map(dataset => dataset.label)).toEqual(['Actual cost', 'Forecast cost'])
|
||||
expect(chart.data.datasets.map(dataset => dataset.data)).toEqual([[12.5, null], [null, 13]])
|
||||
})
|
||||
|
||||
it('preserves unstacked bars when the locale changes', async () => {
|
||||
await mountChart(createApp({
|
||||
render: () => h(BarChart, {
|
||||
stacked: false,
|
||||
data: { labels: ['model-a'], datasets: [{ data: [2] }] },
|
||||
}),
|
||||
}))
|
||||
|
||||
setI18nLocale('en-US')
|
||||
await nextTick()
|
||||
|
||||
const chart = renderedChart()
|
||||
expect(chart.options.locale).toBe('en-US')
|
||||
expect(chart.options.scales?.x?.stacked).toBe(false)
|
||||
expect(chart.options.scales?.y?.stacked).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -2,7 +2,7 @@
|
||||
<DropdownMenu>
|
||||
<DropdownMenuTrigger as-child>
|
||||
<button
|
||||
class="flex h-9 w-9 items-center justify-center rounded-lg text-muted-foreground transition hover:bg-muted/50 hover:text-foreground"
|
||||
class="flex h-9 w-9 shrink-0 items-center justify-center rounded-lg text-muted-foreground transition hover:bg-muted/50 hover:text-foreground"
|
||||
:aria-label="t('common.language')"
|
||||
:title="t('common.language')"
|
||||
type="button"
|
||||
@@ -18,12 +18,13 @@
|
||||
v-for="option in options"
|
||||
:key="option.value"
|
||||
class="justify-between gap-3"
|
||||
:lang="option.value"
|
||||
@select="setLocale(option.value)"
|
||||
>
|
||||
<span>{{ option.label }}</span>
|
||||
<Check
|
||||
v-if="locale === option.value"
|
||||
class="h-4 w-4 text-primary"
|
||||
class="h-4 w-4 shrink-0 text-primary"
|
||||
/>
|
||||
</DropdownMenuItem>
|
||||
</DropdownMenuContent>
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
<template>
|
||||
<div class="flex flex-wrap items-center gap-2">
|
||||
<div class="flex max-w-full flex-wrap items-center gap-2">
|
||||
<Select
|
||||
v-model="selectedPreset"
|
||||
>
|
||||
<SelectTrigger
|
||||
class="h-8 w-32 text-xs border-border/60"
|
||||
class="h-8 w-40 text-xs border-border/60"
|
||||
:class="[presetTriggerClass]"
|
||||
>
|
||||
<SelectValue :placeholder="legacyT('选择时间段')" />
|
||||
@@ -22,7 +22,7 @@
|
||||
|
||||
<div
|
||||
v-if="selectedPreset === 'custom'"
|
||||
class="flex items-center gap-2"
|
||||
class="flex max-w-full flex-wrap items-center gap-2"
|
||||
>
|
||||
<Input
|
||||
v-model="startDate"
|
||||
|
||||
@@ -5,8 +5,8 @@
|
||||
:class="headerClasses"
|
||||
>
|
||||
<slot name="header">
|
||||
<div class="flex items-center justify-between">
|
||||
<div>
|
||||
<div class="flex flex-wrap items-start justify-between gap-4">
|
||||
<div class="min-w-0 flex-1 basis-64">
|
||||
<h3
|
||||
v-if="title"
|
||||
class="text-lg font-medium leading-6 text-foreground"
|
||||
@@ -20,7 +20,10 @@
|
||||
{{ description }}
|
||||
</p>
|
||||
</div>
|
||||
<div v-if="$slots.actions">
|
||||
<div
|
||||
v-if="$slots.actions"
|
||||
class="max-w-full shrink-0 [&>div]:flex-wrap [&_button]:shrink-0"
|
||||
>
|
||||
<slot name="actions" />
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
<template>
|
||||
<div class="flex flex-col gap-3 sm:flex-row sm:items-center sm:justify-between">
|
||||
<div class="flex-1">
|
||||
<div class="flex items-center gap-3">
|
||||
<div class="min-w-0 flex-1">
|
||||
<div class="flex min-w-0 items-center gap-3">
|
||||
<slot name="icon">
|
||||
<div
|
||||
v-if="icon"
|
||||
class="flex h-10 w-10 items-center justify-center rounded-xl bg-primary/10"
|
||||
class="flex h-10 w-10 shrink-0 items-center justify-center rounded-xl bg-primary/10"
|
||||
>
|
||||
<component
|
||||
:is="icon"
|
||||
@@ -14,13 +14,13 @@
|
||||
</div>
|
||||
</slot>
|
||||
|
||||
<div>
|
||||
<h1 class="text-2xl font-semibold text-foreground sm:text-3xl">
|
||||
<div class="min-w-0">
|
||||
<h1 class="break-words text-2xl font-semibold text-foreground sm:text-3xl">
|
||||
{{ title }}
|
||||
</h1>
|
||||
<p
|
||||
v-if="description"
|
||||
class="mt-1 text-sm text-muted-foreground"
|
||||
class="mt-1 break-words text-sm text-muted-foreground"
|
||||
>
|
||||
{{ description }}
|
||||
</p>
|
||||
@@ -30,7 +30,7 @@
|
||||
|
||||
<div
|
||||
v-if="$slots.actions"
|
||||
class="flex items-center gap-2"
|
||||
class="flex min-w-0 flex-wrap items-center gap-2 sm:max-w-[50%] sm:justify-end"
|
||||
>
|
||||
<slot name="actions" />
|
||||
</div>
|
||||
|
||||
@@ -21,7 +21,7 @@
|
||||
:class="index > 0 ? 'pt-1' : ''"
|
||||
>
|
||||
<span class="text-[10px] font-medium text-muted-foreground/50 font-mono tabular-nums">{{ String(index + 1).padStart(2, '0') }}</span>
|
||||
<span class="text-[10px] font-semibold text-muted-foreground/70 uppercase tracking-[0.1em]">{{ group.title }}</span>
|
||||
<span class="min-w-0 break-words text-[10px] font-semibold leading-4 text-muted-foreground/70 uppercase tracking-normal">{{ group.title }}</span>
|
||||
</div>
|
||||
|
||||
<!-- Links -->
|
||||
@@ -37,7 +37,7 @@
|
||||
<TooltipTrigger as-child>
|
||||
<RouterLink
|
||||
:to="item.href"
|
||||
class="group relative flex items-center rounded-lg"
|
||||
class="group relative flex min-w-0 items-center gap-2 rounded-lg"
|
||||
:class="[
|
||||
collapsed
|
||||
? 'h-9 justify-center px-0 transition-colors duration-150'
|
||||
@@ -62,7 +62,7 @@
|
||||
/>
|
||||
<span
|
||||
v-if="!collapsed"
|
||||
class="truncate text-[13px] tracking-tight"
|
||||
class="min-w-0 break-words text-[13px] leading-5 tracking-normal"
|
||||
>{{ item.name }}</span>
|
||||
</div>
|
||||
|
||||
|
||||
@@ -4,26 +4,27 @@
|
||||
<Teleport to="body">
|
||||
<div
|
||||
v-if="tooltip.visible && tooltip.day"
|
||||
class="fixed z-50 rounded-lg border border-border/70 bg-background px-3 py-2 text-xs shadow-lg backdrop-blur pointer-events-none"
|
||||
ref="tooltipRef"
|
||||
class="fixed z-50 w-[200px] max-w-[calc(100vw-1rem)] break-words rounded-lg border border-border/70 bg-background px-3 py-2 text-xs shadow-lg backdrop-blur pointer-events-none"
|
||||
:style="tooltipStyle"
|
||||
>
|
||||
<p class="font-medium">
|
||||
{{ tooltip.day.date }}
|
||||
{{ formatDay(tooltip.day.date) }}
|
||||
</p>
|
||||
<p class="mt-0.5">
|
||||
{{ tooltip.day.requests }} 次请求 · {{ formatTokens(tooltip.day.total_tokens) }}
|
||||
{{ t('heatmap.requests', { count: tooltip.day.requests }) }} · {{ formatTokens(tooltip.day.total_tokens) }}
|
||||
</p>
|
||||
<p class="text-[11px] text-muted-foreground">
|
||||
成本 {{ formatCurrency(tooltip.day.total_cost) }}
|
||||
{{ t('heatmap.cost', { value: formatCurrency(tooltip.day.total_cost) }) }}
|
||||
</p>
|
||||
</div>
|
||||
</Teleport>
|
||||
|
||||
<div
|
||||
v-if="showHeader"
|
||||
class="flex items-center justify-between gap-4"
|
||||
class="flex flex-wrap items-center justify-between gap-4"
|
||||
>
|
||||
<div class="flex-shrink-0">
|
||||
<div class="min-w-0 break-words">
|
||||
<p class="text-sm font-semibold">
|
||||
{{ title }}
|
||||
</p>
|
||||
@@ -38,20 +39,20 @@
|
||||
v-if="weekColumns.length > 0"
|
||||
class="flex items-center gap-1 text-[11px] text-muted-foreground flex-shrink-0"
|
||||
>
|
||||
<span class="flex-shrink-0">少</span>
|
||||
<span class="flex-shrink-0">{{ t('heatmap.less') }}</span>
|
||||
<div
|
||||
v-for="(level, index) in legendLevels"
|
||||
:key="index"
|
||||
class="w-3 h-3 rounded-[3px] flex-shrink-0"
|
||||
:style="getLegendStyle(level)"
|
||||
/>
|
||||
<span class="flex-shrink-0">多</span>
|
||||
<span class="flex-shrink-0">{{ t('heatmap.more') }}</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div
|
||||
v-if="weekColumns.length > 0"
|
||||
class="flex w-full gap-3"
|
||||
class="flex w-full gap-3 overflow-x-auto"
|
||||
>
|
||||
<div
|
||||
class="flex flex-col text-[10px] text-muted-foreground flex-shrink-0"
|
||||
@@ -62,33 +63,12 @@
|
||||
M
|
||||
</div>
|
||||
<span
|
||||
:style="dayLabelStyle"
|
||||
class="flex items-center invisible"
|
||||
>周日</span>
|
||||
<span
|
||||
v-for="(weekday, index) in weekdayLabels"
|
||||
:key="index"
|
||||
:style="dayLabelStyle"
|
||||
class="flex items-center"
|
||||
>一</span>
|
||||
<span
|
||||
:style="dayLabelStyle"
|
||||
class="flex items-center invisible"
|
||||
>周二</span>
|
||||
<span
|
||||
:style="dayLabelStyle"
|
||||
class="flex items-center"
|
||||
>三</span>
|
||||
<span
|
||||
:style="dayLabelStyle"
|
||||
class="flex items-center invisible"
|
||||
>周四</span>
|
||||
<span
|
||||
:style="dayLabelStyle"
|
||||
class="flex items-center"
|
||||
>五</span>
|
||||
<span
|
||||
:style="dayLabelStyle"
|
||||
class="flex items-center invisible"
|
||||
>周六</span>
|
||||
:class="{ invisible: index % 2 === 0 }"
|
||||
>{{ weekday }}</span>
|
||||
</div>
|
||||
<div class="flex-1 min-w-[200px]">
|
||||
<div
|
||||
@@ -103,7 +83,7 @@
|
||||
v-for="(week, weekIndex) in weekColumns"
|
||||
:key="`month-${weekIndex}`"
|
||||
:style="monthCellStyle"
|
||||
class="text-center"
|
||||
class="whitespace-nowrap text-left"
|
||||
>
|
||||
<span v-if="monthMarkers[weekIndex]">{{ monthMarkers[weekIndex] }}</span>
|
||||
</div>
|
||||
@@ -146,15 +126,16 @@
|
||||
v-else
|
||||
class="text-xs text-muted-foreground"
|
||||
>
|
||||
暂无活跃数据
|
||||
{{ t('heatmap.empty') }}
|
||||
</p>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, onBeforeUnmount, onMounted, ref, watch } from 'vue'
|
||||
import { computed, nextTick, onBeforeUnmount, onMounted, ref, watch } from 'vue'
|
||||
import type { ActivityHeatmap, ActivityHeatmapDay } from '@/types/activity'
|
||||
import { formatCurrency, formatTokens } from '@/utils/format'
|
||||
import { useI18n } from '@/i18n'
|
||||
|
||||
const props = withDefaults(defineProps<{
|
||||
data?: ActivityHeatmap | null
|
||||
@@ -169,12 +150,23 @@ const props = withDefaults(defineProps<{
|
||||
})
|
||||
|
||||
const legendLevels = [0.08, 0.25, 0.45, 0.65, 0.85]
|
||||
const { locale, t } = useI18n()
|
||||
const weekdayLabels = computed(() => {
|
||||
const formatter = new Intl.DateTimeFormat(locale.value, { weekday: 'short', timeZone: 'UTC' })
|
||||
return Array.from({ length: 7 }, (_, day) => formatter.format(new Date(Date.UTC(2024, 0, 7 + day))))
|
||||
})
|
||||
|
||||
function formatDay(value: string): string {
|
||||
return new Intl.DateTimeFormat(locale.value, { dateStyle: 'medium', timeZone: 'UTC' })
|
||||
.format(new Date(`${value}T00:00:00Z`))
|
||||
}
|
||||
|
||||
type DayWithMeta = ActivityHeatmapDay & { dateObj: Date }
|
||||
const heatmapWrapper = ref<HTMLElement | null>(null)
|
||||
const heatmapWidth = ref(0)
|
||||
const cellSize = ref(10)
|
||||
const cellGap = ref(4)
|
||||
const tooltipRef = ref<HTMLElement | null>(null)
|
||||
const tooltip = ref<{ day: ActivityHeatmapDay | null; x: number; y: number; visible: boolean; below: boolean }>({
|
||||
day: null,
|
||||
x: 0,
|
||||
@@ -259,6 +251,7 @@ const weekColumns = computed(() => {
|
||||
const monthMarkers = computed(() => {
|
||||
const markers: Record<number, string> = {}
|
||||
const columns = weekColumns.value
|
||||
const formatter = new Intl.DateTimeFormat(locale.value, { month: 'short', timeZone: 'UTC' })
|
||||
let lastMonth: number | null = null
|
||||
|
||||
columns.forEach((week, index) => {
|
||||
@@ -270,7 +263,7 @@ const monthMarkers = computed(() => {
|
||||
if (month === lastMonth) {
|
||||
return
|
||||
}
|
||||
markers[index] = `${month + 1}月`
|
||||
markers[index] = formatter.format(firstValid.dateObj)
|
||||
lastMonth = month
|
||||
})
|
||||
|
||||
@@ -343,10 +336,14 @@ onBeforeUnmount(() => {
|
||||
}
|
||||
})
|
||||
|
||||
function handleHover(day: ActivityHeatmapDay, event: MouseEvent) {
|
||||
async function handleHover(day: ActivityHeatmapDay, event: MouseEvent) {
|
||||
const cellRect = (event.currentTarget as HTMLElement).getBoundingClientRect()
|
||||
const tooltipWidth = 200
|
||||
const tooltipHeight = 72
|
||||
tooltip.value = { day, x: cellRect.left, y: cellRect.top, visible: true, below: false }
|
||||
await nextTick()
|
||||
if (!tooltip.value.visible || tooltip.value.day?.date !== day.date) return
|
||||
|
||||
const tooltipWidth = tooltipRef.value?.offsetWidth || 200
|
||||
const tooltipHeight = tooltipRef.value?.offsetHeight || 72
|
||||
|
||||
// Calculate horizontal position (centered on cell)
|
||||
let left = cellRect.left + cellRect.width / 2
|
||||
@@ -401,11 +398,11 @@ function getCellStyle(requests: number) {
|
||||
}
|
||||
|
||||
function buildTooltip(day: ActivityHeatmapDay): string {
|
||||
const dateLabel = day.date
|
||||
const dateLabel = formatDay(day.date)
|
||||
const costLabel = formatCurrency(day.total_cost || 0)
|
||||
const parts = [`${dateLabel}`, `${day.requests} 次请求`, `${formatTokens(day.total_tokens)} tokens`, costLabel]
|
||||
const parts = [dateLabel, t('heatmap.requests', { count: day.requests }), `${formatTokens(day.total_tokens)} tokens`, costLabel]
|
||||
if (day.actual_total_cost !== undefined) {
|
||||
parts.push(`倍率: ${formatCurrency(day.actual_total_cost)}`)
|
||||
parts.push(t('heatmap.actualCost', { value: formatCurrency(day.actual_total_cost) }))
|
||||
}
|
||||
return parts.join(' · ')
|
||||
}
|
||||
|
||||
@@ -29,6 +29,7 @@
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed } from 'vue'
|
||||
import { useI18n } from '@/i18n'
|
||||
import LineChart from '@/components/charts/LineChart.vue'
|
||||
import { LoadingState } from '@/components/common'
|
||||
import { formatCurrency } from '@/utils/format'
|
||||
@@ -45,6 +46,7 @@ const props = withDefaults(defineProps<Props>(), {
|
||||
subtitle: undefined,
|
||||
loading: false
|
||||
})
|
||||
const { t } = useI18n()
|
||||
|
||||
const labels = computed(() => [
|
||||
...props.history.map(item => item.date),
|
||||
@@ -58,7 +60,7 @@ const chartData = computed(() => {
|
||||
labels: labels.value,
|
||||
datasets: [
|
||||
{
|
||||
label: '实际成本',
|
||||
label: t('chart.actualCost'),
|
||||
data: historyValues.concat(new Array(forecastValues.length).fill(null)),
|
||||
borderColor: 'rgb(59, 130, 246)',
|
||||
backgroundColor: 'rgba(59, 130, 246, 0.15)',
|
||||
@@ -66,7 +68,7 @@ const chartData = computed(() => {
|
||||
pointRadius: 2
|
||||
},
|
||||
{
|
||||
label: '预测成本',
|
||||
label: t('chart.forecastCost'),
|
||||
data: new Array(historyValues.length).fill(null).concat(forecastValues),
|
||||
borderColor: 'rgb(234, 179, 8)',
|
||||
backgroundColor: 'rgba(234, 179, 8, 0.15)',
|
||||
|
||||
@@ -59,6 +59,7 @@
|
||||
|
||||
<script setup lang="ts">
|
||||
import { Card } from '@/components/ui'
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
import { EmptyState, LoadingState } from '@/components/common'
|
||||
import { formatCurrency } from '@/utils/format'
|
||||
import type { QuotaUsageProvider } from '@/api/admin'
|
||||
@@ -76,6 +77,6 @@ withDefaults(defineProps<Props>(), {
|
||||
})
|
||||
|
||||
function formatDate(value: string) {
|
||||
return new Date(value).toLocaleDateString()
|
||||
return new Date(value).toLocaleDateString(getI18nLocale())
|
||||
}
|
||||
</script>
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
import { afterEach, describe, expect, it, vi } from 'vitest'
|
||||
import { createApp, h, nextTick, ref, type App } from 'vue'
|
||||
|
||||
import Pagination from '../pagination.vue'
|
||||
import { setI18nLocale } from '@/i18n'
|
||||
|
||||
const mountedApps: Array<{ app: App, root: HTMLElement }> = []
|
||||
|
||||
afterEach(() => {
|
||||
for (const { app, root } of mountedApps.splice(0)) {
|
||||
app.unmount()
|
||||
root.remove()
|
||||
}
|
||||
})
|
||||
|
||||
describe('Pagination', () => {
|
||||
it('updates the summary and accessible page controls when the locale changes', async () => {
|
||||
const root = document.createElement('div')
|
||||
document.body.appendChild(root)
|
||||
const updateCurrent = vi.fn()
|
||||
const app = createApp({
|
||||
render: () => h(Pagination, {
|
||||
current: 1,
|
||||
total: 1250,
|
||||
pageSize: 20,
|
||||
showPageSizeSelector: false,
|
||||
'onUpdate:current': updateCurrent,
|
||||
}),
|
||||
})
|
||||
app.mount(root)
|
||||
mountedApps.push({ app, root })
|
||||
|
||||
expect(root.querySelector('[aria-live]')?.textContent).toContain('共 1,250 条')
|
||||
expect(root.querySelector('[aria-current="page"]')?.getAttribute('aria-label')).toBe('第 1 页')
|
||||
|
||||
setI18nLocale('en-US')
|
||||
await nextTick()
|
||||
|
||||
expect(root.querySelector('[aria-live]')?.textContent).toContain('Showing 1-20 of 1,250 items')
|
||||
expect(root.querySelector('[aria-current="page"]')?.getAttribute('aria-label')).toBe('Page 1')
|
||||
expect(root.querySelector('input')?.getAttribute('aria-label')).toBe('Go to page')
|
||||
|
||||
root.querySelector<HTMLButtonElement>('[aria-label="Page 2"]')?.click()
|
||||
expect(updateCurrent).toHaveBeenCalledWith(2)
|
||||
})
|
||||
|
||||
it('shows a zero-based empty range after the final record is removed', async () => {
|
||||
const total = ref(1)
|
||||
const root = document.createElement('div')
|
||||
document.body.appendChild(root)
|
||||
const app = createApp({
|
||||
render: () => h(Pagination, {
|
||||
current: 1,
|
||||
total: total.value,
|
||||
showPageSizeSelector: false,
|
||||
}),
|
||||
})
|
||||
app.mount(root)
|
||||
mountedApps.push({ app, root })
|
||||
|
||||
total.value = 0
|
||||
setI18nLocale('en-US')
|
||||
await nextTick()
|
||||
|
||||
expect(root.querySelector('[aria-live]')?.textContent).toContain('Showing 0-0 of 0 items')
|
||||
expect(root.querySelectorAll('button')).toHaveLength(0)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,55 @@
|
||||
import { afterEach, describe, expect, it, vi } from 'vitest'
|
||||
import { createApp, h, nextTick, type App } from 'vue'
|
||||
|
||||
import Tabs from '../tabs.vue'
|
||||
import TabsList from '../tabs-list.vue'
|
||||
import TabsTrigger from '../tabs-trigger.vue'
|
||||
import { setI18nLocale, useI18n } from '@/i18n'
|
||||
|
||||
const mountedApps: Array<{ app: App, root: HTMLElement }> = []
|
||||
|
||||
afterEach(() => {
|
||||
for (const { app, root } of mountedApps.splice(0)) {
|
||||
app.unmount()
|
||||
root.remove()
|
||||
}
|
||||
vi.restoreAllMocks()
|
||||
vi.useRealTimers()
|
||||
})
|
||||
|
||||
describe('TabsList', () => {
|
||||
it('repositions the indicator after translated labels change width', async () => {
|
||||
vi.useFakeTimers()
|
||||
vi.spyOn(HTMLElement.prototype, 'getBoundingClientRect').mockImplementation(function (this: HTMLElement) {
|
||||
return { width: this.textContent === '个人设置' ? 80 : 140 } as DOMRect
|
||||
})
|
||||
vi.spyOn(HTMLElement.prototype, 'offsetLeft', 'get').mockReturnValue(4)
|
||||
|
||||
const root = document.createElement('div')
|
||||
document.body.appendChild(root)
|
||||
const app = createApp({
|
||||
setup() {
|
||||
const { t } = useI18n()
|
||||
return () => h(Tabs, { modelValue: 'settings' }, {
|
||||
default: () => h(TabsList, {}, {
|
||||
default: () => h(TabsTrigger, { value: 'settings' }, () => t('common.settings')),
|
||||
}),
|
||||
})
|
||||
},
|
||||
})
|
||||
app.mount(root)
|
||||
mountedApps.push({ app, root })
|
||||
await nextTick()
|
||||
await vi.runAllTimersAsync()
|
||||
|
||||
const indicator = root.querySelector<HTMLElement>('.tabs-indicator')
|
||||
expect(indicator?.style.width).toBe('80px')
|
||||
expect(indicator?.style.transform).toBe('translateX(4px)')
|
||||
|
||||
setI18nLocale('en-US')
|
||||
await nextTick()
|
||||
await vi.runAllTimersAsync()
|
||||
|
||||
expect(indicator?.style.width).toBe('140px')
|
||||
})
|
||||
})
|
||||
@@ -31,7 +31,7 @@ const props = withDefaults(defineProps<Props>(), {
|
||||
|
||||
const buttonClass = computed(() => {
|
||||
const baseClass =
|
||||
'inline-flex items-center justify-center rounded-xl text-sm font-semibold transition-all duration-200 ring-offset-background focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2 disabled:pointer-events-none disabled:opacity-50 active:scale-[0.98]'
|
||||
'inline-flex min-w-0 max-w-full items-center justify-center rounded-xl text-sm font-semibold leading-5 [overflow-wrap:anywhere] [&_svg]:shrink-0 transition-all duration-200 ring-offset-background focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2 disabled:pointer-events-none disabled:opacity-50 active:scale-[0.98]'
|
||||
|
||||
const variantClasses = {
|
||||
default: 'bg-primary text-white hover:bg-primary/90',
|
||||
@@ -48,7 +48,7 @@ const buttonClass = computed(() => {
|
||||
default: 'h-11 px-5',
|
||||
sm: 'h-9 rounded-lg px-3',
|
||||
lg: 'h-12 rounded-xl px-8 text-base',
|
||||
icon: 'h-11 w-11 rounded-2xl',
|
||||
icon: 'h-11 w-11 shrink-0 rounded-2xl',
|
||||
}
|
||||
|
||||
return cn(
|
||||
|
||||
@@ -57,12 +57,12 @@
|
||||
/>
|
||||
</div>
|
||||
<div class="flex-1 min-w-0">
|
||||
<h3 class="text-balance text-base font-semibold leading-tight text-foreground sm:text-lg">
|
||||
<h3 class="break-words text-balance text-base font-semibold leading-tight text-foreground sm:text-lg">
|
||||
{{ title }}
|
||||
</h3>
|
||||
<p
|
||||
v-if="description"
|
||||
class="mt-0.5 text-pretty text-xs leading-4 text-muted-foreground"
|
||||
class="mt-0.5 break-words text-pretty text-xs leading-4 text-muted-foreground"
|
||||
>
|
||||
{{ description }}
|
||||
</p>
|
||||
@@ -80,7 +80,7 @@
|
||||
<!-- Footer 区域:如果有 footer 插槽,自动添加样式 -->
|
||||
<div
|
||||
v-if="slots.footer"
|
||||
class="flex shrink-0 flex-col-reverse items-stretch gap-2 border-t border-border bg-background/95 px-4 pb-[max(0.75rem,env(safe-area-inset-bottom))] pt-3 backdrop-blur-sm [&>button]:w-full sm:flex-row-reverse sm:items-center sm:gap-3 sm:bg-muted/10 sm:px-6 sm:py-4 sm:[&>button]:w-auto"
|
||||
class="flex shrink-0 flex-col-reverse items-stretch gap-2 border-t border-border bg-background/95 px-4 pb-[max(0.75rem,env(safe-area-inset-bottom))] pt-3 backdrop-blur-sm [&>button]:min-h-min [&>button]:w-full [&>button]:whitespace-normal [&>button]:py-2 sm:flex-row-reverse sm:flex-wrap sm:items-center sm:gap-3 sm:bg-muted/10 sm:px-6 sm:py-4 sm:[&>button]:w-auto"
|
||||
>
|
||||
<slot name="footer" />
|
||||
</div>
|
||||
@@ -169,7 +169,7 @@ const maxWidthClass = computed(() => {
|
||||
})
|
||||
|
||||
const contentBodyClass = computed(() => [
|
||||
'min-h-0 overflow-y-auto overscroll-contain',
|
||||
'min-h-0 min-w-0 overflow-y-auto overscroll-contain',
|
||||
props.noPadding ? '' : 'px-4 py-3 sm:px-6',
|
||||
].filter(Boolean).join(' '))
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
<template>
|
||||
<div class="border-t border-border px-6 py-4 bg-muted/10 flex flex-row-reverse gap-3">
|
||||
<div class="flex flex-col-reverse gap-3 border-t border-border bg-muted/10 px-4 py-4 [&>button]:min-h-min [&>button]:whitespace-normal [&>button]:py-2 sm:flex-row-reverse sm:flex-wrap sm:px-6">
|
||||
<slot />
|
||||
</div>
|
||||
</template>
|
||||
</template>
|
||||
|
||||
@@ -27,7 +27,7 @@ const props = withDefaults(defineProps<Props>(), {
|
||||
|
||||
const contentClass = computed(() =>
|
||||
cn(
|
||||
'z-[200] min-w-[8rem] overflow-hidden rounded-2xl border border-border bg-card p-1 text-foreground shadow-2xl backdrop-blur-xl',
|
||||
'z-[200] min-w-[8rem] max-w-[calc(100vw-1rem)] max-h-[var(--radix-dropdown-menu-content-available-height)] overflow-y-auto overscroll-contain rounded-2xl border border-border bg-card p-1 text-foreground shadow-2xl backdrop-blur-xl',
|
||||
'data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0 data-[state=closed]:zoom-out-95 data-[state=open]:zoom-in-95',
|
||||
'data-[side=bottom]:slide-in-from-top-2 data-[side=left]:slide-in-from-right-2 data-[side=right]:slide-in-from-left-2 data-[side=top]:slide-in-from-bottom-2',
|
||||
props.class
|
||||
|
||||
@@ -16,7 +16,7 @@ defineEmits<{
|
||||
|
||||
const itemClass = computed(() =>
|
||||
cn(
|
||||
'relative flex cursor-pointer select-none items-center rounded-lg px-3 py-1.5 text-sm outline-none',
|
||||
'relative flex min-w-0 cursor-pointer select-none items-center whitespace-normal break-words rounded-lg px-3 py-1.5 text-sm leading-5 outline-none [&_svg]:shrink-0',
|
||||
'data-[highlighted]:bg-accent focus:bg-accent text-foreground',
|
||||
'transition-colors data-[disabled]:pointer-events-none data-[disabled]:opacity-50',
|
||||
props.class
|
||||
|
||||
@@ -16,7 +16,7 @@ const props = defineProps<Props>()
|
||||
|
||||
const labelClass = computed(() =>
|
||||
cn(
|
||||
'text-[11px] font-semibold uppercase tracking-[0.14em] text-muted-foreground peer-disabled:cursor-not-allowed peer-disabled:opacity-70',
|
||||
'text-[11px] font-semibold tracking-normal break-words text-muted-foreground peer-disabled:cursor-not-allowed peer-disabled:opacity-70',
|
||||
props.class
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
<template>
|
||||
<div class="flex flex-col sm:flex-row gap-3 sm:gap-4 border-t border-border/60 px-4 sm:px-6 py-3 sm:py-4 bg-muted/20">
|
||||
<div class="flex min-w-0 flex-col gap-3 border-t border-border/60 bg-muted/20 px-4 py-3 sm:flex-row sm:flex-wrap sm:items-center sm:gap-4 sm:px-6 sm:py-4">
|
||||
<!-- 左侧:记录范围和每页数量 -->
|
||||
<div class="flex items-center justify-between sm:justify-start gap-3 text-sm text-muted-foreground">
|
||||
<span class="font-medium whitespace-nowrap">
|
||||
<div class="flex min-w-0 flex-wrap items-center justify-between gap-3 text-sm text-muted-foreground sm:justify-start">
|
||||
<span
|
||||
class="min-w-0 break-words font-medium tabular-nums"
|
||||
aria-live="polite"
|
||||
>
|
||||
{{ rangeSummary }}
|
||||
</span>
|
||||
<Select
|
||||
@@ -10,7 +13,10 @@
|
||||
:model-value="String(pageSize)"
|
||||
@update:model-value="handlePageSizeChange"
|
||||
>
|
||||
<SelectTrigger class="w-[120px] h-8 sm:h-9 border-border/60 text-xs sm:text-sm">
|
||||
<SelectTrigger
|
||||
class="h-8 w-auto min-w-[120px] shrink-0 border-border/60 text-xs sm:h-9 sm:text-sm"
|
||||
:aria-label="t('pagination.pageSizeLabel')"
|
||||
>
|
||||
<span class="flex-1 text-center">
|
||||
<SelectValue />
|
||||
</span>
|
||||
@@ -31,8 +37,8 @@
|
||||
<div class="flex flex-wrap items-center justify-center gap-1.5 sm:gap-2 sm:ml-auto">
|
||||
<!-- 页码按钮(智能省略) -->
|
||||
<template
|
||||
v-for="page in pageNumbers"
|
||||
:key="page"
|
||||
v-for="(page, index) in pageNumbers"
|
||||
:key="`${page}-${index}`"
|
||||
>
|
||||
<Button
|
||||
v-if="typeof page === 'number'"
|
||||
@@ -40,9 +46,11 @@
|
||||
size="sm"
|
||||
class="h-9 min-w-[36px] px-2"
|
||||
:class="page === current ? 'shadow-sm' : ''"
|
||||
:aria-label="t('pagination.pageNumber', { page: formatNumber(page) })"
|
||||
:aria-current="page === current ? 'page' : undefined"
|
||||
@click="handlePageChange(page)"
|
||||
>
|
||||
{{ page }}
|
||||
{{ formatNumber(page) }}
|
||||
</Button>
|
||||
<span
|
||||
v-else
|
||||
@@ -55,18 +63,19 @@
|
||||
v-if="totalPages > 7"
|
||||
class="flex items-center gap-1.5 ml-2 text-sm text-muted-foreground"
|
||||
>
|
||||
<span class="hidden sm:inline">{{ jumpToLabel }}</span>
|
||||
<span class="hidden sm:inline">{{ t('pagination.goToPage') }}</span>
|
||||
<input
|
||||
v-model="jumpPageInput"
|
||||
type="text"
|
||||
inputmode="numeric"
|
||||
pattern="[0-9]*"
|
||||
:aria-label="t('pagination.jumpToPage')"
|
||||
class="w-12 h-9 px-2 text-center text-sm border border-border/60 rounded-md bg-background focus:outline-none focus:ring-2 focus:ring-primary/40 focus:border-primary/60"
|
||||
@keydown.enter="handleJumpPage"
|
||||
@blur="handleJumpPage"
|
||||
@input="filterNumericInput"
|
||||
>
|
||||
<span class="hidden sm:inline">{{ pageLabel }}</span>
|
||||
<span class="hidden sm:inline">{{ t('pagination.pageLabel') }}</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -102,28 +111,29 @@ const props = withDefaults(defineProps<Props>(), {
|
||||
const emit = defineEmits<Emits>()
|
||||
|
||||
const jumpPageInput = ref('')
|
||||
const locale = useI18n().locale
|
||||
const { locale, t } = useI18n()
|
||||
const numberFormatter = computed(() => new Intl.NumberFormat(locale.value))
|
||||
|
||||
function formatNumber(value: number): string {
|
||||
return numberFormatter.value.format(value)
|
||||
}
|
||||
|
||||
const totalPages = computed(() => Math.ceil(props.total / props.pageSize))
|
||||
|
||||
const recordRange = computed(() => {
|
||||
const start = (props.current - 1) * props.pageSize + 1
|
||||
const start = props.total === 0 ? 0 : (props.current - 1) * props.pageSize + 1
|
||||
const end = Math.min(props.current * props.pageSize, props.total)
|
||||
return { start, end }
|
||||
})
|
||||
|
||||
const rangeSummary = computed(() => {
|
||||
if (locale.value === 'en-US') {
|
||||
return `Showing ${recordRange.value.start}-${recordRange.value.end} of ${props.total} items`
|
||||
}
|
||||
return `显示 ${recordRange.value.start}-${recordRange.value.end} 条,共 ${props.total} 条`
|
||||
})
|
||||
|
||||
const jumpToLabel = computed(() => locale.value === 'en-US' ? 'Go to' : '跳至')
|
||||
const pageLabel = computed(() => locale.value === 'en-US' ? 'page' : '页')
|
||||
const rangeSummary = computed(() => t('pagination.range', {
|
||||
start: formatNumber(recordRange.value.start),
|
||||
end: formatNumber(recordRange.value.end),
|
||||
total: formatNumber(props.total),
|
||||
}))
|
||||
|
||||
function pageSizeLabel(size: number): string {
|
||||
return locale.value === 'en-US' ? `${size} / page` : `${size} 条/页`
|
||||
return t('pagination.pageSize', { size: formatNumber(size) })
|
||||
}
|
||||
|
||||
const pageNumbers = computed(() => {
|
||||
|
||||
@@ -21,7 +21,8 @@
|
||||
<Input
|
||||
ref="searchInputRef"
|
||||
v-model="searchQuery"
|
||||
:placeholder="searchPlaceholder"
|
||||
:placeholder="searchPlaceholder ?? t('common.searchPlaceholder')"
|
||||
:aria-label="searchPlaceholder ?? t('common.searchPlaceholder')"
|
||||
class="h-9 rounded-xl border-border/60 bg-background/80 pl-9 pr-3 text-sm"
|
||||
@keydown.stop
|
||||
/>
|
||||
@@ -34,7 +35,7 @@
|
||||
v-if="showEmptyState"
|
||||
class="px-3 py-2 text-sm text-muted-foreground"
|
||||
>
|
||||
未找到匹配项
|
||||
{{ t('common.noSearchResults') }}
|
||||
</div>
|
||||
</SelectViewport>
|
||||
</SelectContentPrimitive>
|
||||
@@ -66,6 +67,7 @@ import {
|
||||
type RegisteredSelectItem,
|
||||
} from './select-search-context'
|
||||
import { matchesSearchQuery, preloadPinyin } from '@/utils/search'
|
||||
import { useI18n } from '@/i18n'
|
||||
|
||||
interface Props {
|
||||
class?: string
|
||||
@@ -90,9 +92,10 @@ const props = withDefaults(defineProps<Props>(), {
|
||||
disablePortal: undefined,
|
||||
searchable: true,
|
||||
searchThreshold: 8,
|
||||
searchPlaceholder: '输入关键词搜索...',
|
||||
searchPlaceholder: undefined,
|
||||
})
|
||||
|
||||
const { t } = useI18n()
|
||||
const isInsideDialog = inject(DIALOG_CONTEXT_KEY, false)
|
||||
const shouldDisablePortal = computed(
|
||||
() => props.disablePortal ?? isInsideDialog,
|
||||
@@ -177,7 +180,7 @@ watch(showSearchInput, async (visible) => {
|
||||
|
||||
const contentClass = computed(() =>
|
||||
cn(
|
||||
'z-[200] max-h-96 min-w-[8rem] overflow-hidden rounded-2xl border border-border bg-card text-foreground shadow-2xl backdrop-blur-xl pointer-events-auto',
|
||||
'z-[200] max-h-96 min-w-[var(--radix-select-trigger-width,8rem)] max-w-[min(calc(100vw-1rem),var(--radix-select-content-available-width,100vw))] overflow-hidden rounded-2xl border border-border bg-card text-foreground shadow-2xl backdrop-blur-xl pointer-events-auto',
|
||||
'data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0 data-[state=closed]:zoom-out-95 data-[state=open]:zoom-in-95',
|
||||
'data-[side=bottom]:slide-in-from-top-2 data-[side=left]:slide-in-from-right-2 data-[side=right]:slide-in-from-left-2 data-[side=top]:slide-in-from-bottom-2',
|
||||
props.class,
|
||||
|
||||
@@ -101,7 +101,7 @@ onBeforeUnmount(() => {
|
||||
<Check class="h-4 w-4" />
|
||||
</SelectItemIndicator>
|
||||
</span>
|
||||
<SelectItemText>
|
||||
<SelectItemText class="min-w-0 whitespace-normal break-words text-left">
|
||||
<slot />
|
||||
</SelectItemText>
|
||||
</SelectItemPrimitive>
|
||||
|
||||
@@ -13,7 +13,7 @@ const props = defineProps<Props>()
|
||||
|
||||
const triggerClass = computed(() =>
|
||||
cn(
|
||||
'flex h-11 w-full items-center justify-between rounded-2xl border border-border/60 bg-card/80 px-4 py-2 text-sm shadow-sm placeholder:text-muted-foreground focus:outline-none focus:ring-2 focus:ring-primary/40 focus:border-primary/60 disabled:cursor-not-allowed disabled:opacity-50 text-foreground cursor-pointer backdrop-blur transition-all',
|
||||
'flex h-11 w-full min-w-0 items-center justify-between gap-2 rounded-2xl border border-border/60 bg-card/80 px-4 py-2 text-left text-sm shadow-sm placeholder:text-muted-foreground focus:outline-none focus:ring-2 focus:ring-primary/40 focus:border-primary/60 disabled:cursor-not-allowed disabled:opacity-50 text-foreground cursor-pointer backdrop-blur transition-all',
|
||||
props.class
|
||||
)
|
||||
)
|
||||
@@ -25,7 +25,7 @@ const triggerClass = computed(() =>
|
||||
:class="triggerClass"
|
||||
:disabled="disabled"
|
||||
>
|
||||
<span class="truncate">
|
||||
<span class="min-w-0 flex-1 truncate">
|
||||
<slot />
|
||||
</span>
|
||||
<ChevronDown class="h-4 w-4 opacity-50 pointer-events-none flex-shrink-0" />
|
||||
|
||||
@@ -3,6 +3,7 @@ import { computed, nextTick, onBeforeUnmount, onMounted, ref, useAttrs, useSlots
|
||||
import { ArrowDown, ArrowUp, ArrowUpDown, ListFilter } from 'lucide-vue-next'
|
||||
|
||||
import { cn } from '@/lib/utils'
|
||||
import { useI18n } from '@/i18n'
|
||||
import TableHead from './table-head.vue'
|
||||
|
||||
type SortDirection = 'asc' | 'desc'
|
||||
@@ -29,7 +30,7 @@ const props = withDefaults(defineProps<{
|
||||
align: 'left',
|
||||
title: undefined,
|
||||
filterActive: false,
|
||||
filterTitle: '筛选',
|
||||
filterTitle: undefined,
|
||||
filterContentClass: undefined,
|
||||
})
|
||||
|
||||
@@ -42,6 +43,7 @@ defineOptions({
|
||||
})
|
||||
|
||||
const attrs = useAttrs()
|
||||
const { t } = useI18n()
|
||||
const slots = useSlots()
|
||||
const rootRef = ref<HTMLElement | null>(null)
|
||||
const filterTriggerRef = ref<HTMLButtonElement | null>(null)
|
||||
@@ -65,12 +67,12 @@ const ariaSort = computed(() => {
|
||||
return props.direction === 'asc' ? 'ascending' : 'descending'
|
||||
})
|
||||
const wrapperClass = computed(() => cn(
|
||||
'relative flex w-full items-center gap-1.5',
|
||||
'relative flex min-w-0 w-full items-center gap-1.5',
|
||||
props.align === 'center' && 'justify-center',
|
||||
props.align === 'right' && 'justify-end',
|
||||
))
|
||||
const labelClass = computed(() => cn(
|
||||
'inline-flex min-w-0 items-center gap-1.5 text-xs font-semibold text-muted-foreground',
|
||||
'inline-flex min-w-0 items-center gap-1.5 whitespace-normal break-words text-left text-xs font-semibold leading-4 text-muted-foreground',
|
||||
props.align === 'center' && 'justify-center',
|
||||
props.align === 'right' && 'justify-end',
|
||||
))
|
||||
@@ -89,7 +91,7 @@ const filterButtonClass = computed(() => cn(
|
||||
: 'text-muted-foreground/60 hover:bg-muted/50 hover:text-foreground',
|
||||
))
|
||||
const filterPanelClass = computed(() => cn(
|
||||
'fixed z-[1000] w-64 rounded-md border bg-popover p-3 text-popover-foreground shadow-md outline-none',
|
||||
'fixed z-[1000] w-64 max-w-[calc(100vw-1rem)] rounded-md border bg-popover p-3 text-popover-foreground shadow-md outline-none',
|
||||
props.filterContentClass,
|
||||
))
|
||||
|
||||
@@ -185,7 +187,8 @@ onBeforeUnmount(() => {
|
||||
ref="filterTriggerRef"
|
||||
type="button"
|
||||
:class="filterButtonClass"
|
||||
:title="filterTitle"
|
||||
:title="filterTitle ?? t('common.filter')"
|
||||
:aria-label="filterTitle ?? t('common.filter')"
|
||||
:aria-pressed="filterActive"
|
||||
@click.stop="toggleFilter"
|
||||
>
|
||||
@@ -210,7 +213,7 @@ onBeforeUnmount(() => {
|
||||
v-if="canSort"
|
||||
type="button"
|
||||
:class="buttonClass"
|
||||
:title="title || '排序'"
|
||||
:title="title || t('common.sort')"
|
||||
@click="handleSort"
|
||||
>
|
||||
<slot />
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
role="switch"
|
||||
:aria-checked="modelValue"
|
||||
:disabled="disabled"
|
||||
class="relative inline-flex h-6 w-11 items-center rounded-full transition-colors disabled:cursor-not-allowed disabled:opacity-50"
|
||||
class="relative inline-flex h-6 w-11 shrink-0 items-center rounded-full transition-colors disabled:cursor-not-allowed disabled:opacity-50"
|
||||
:class="[
|
||||
modelValue ? 'bg-primary' : 'bg-muted'
|
||||
]"
|
||||
@@ -28,4 +28,4 @@ defineProps<{
|
||||
defineEmits<{
|
||||
'update:modelValue': [value: boolean]
|
||||
}>()
|
||||
</script>
|
||||
</script>
|
||||
|
||||
@@ -15,12 +15,14 @@
|
||||
<script setup lang="ts">
|
||||
import { computed, ref, watch, onMounted, onUnmounted, nextTick, inject, type Ref } from 'vue'
|
||||
import { cn } from '@/lib/utils'
|
||||
import { useI18n } from '@/i18n'
|
||||
|
||||
interface Props {
|
||||
class?: string
|
||||
}
|
||||
|
||||
const props = defineProps<Props>()
|
||||
const { locale } = useI18n()
|
||||
|
||||
const listRef = ref<HTMLElement | null>(null)
|
||||
const indicatorStyle = ref<Record<string, string>>({
|
||||
@@ -82,11 +84,7 @@ const updateIndicator = () => {
|
||||
// 确保按钮已渲染
|
||||
if (buttonRect.width === 0) return
|
||||
|
||||
// 计算相对位置:累加前面所有按钮的宽度
|
||||
let offsetLeft = 0
|
||||
for (let i = 0; i < newIndex; i++) {
|
||||
offsetLeft += buttons[i].getBoundingClientRect().width
|
||||
}
|
||||
const offsetLeft = activeButton.offsetLeft
|
||||
|
||||
// 判断是否需要动画:
|
||||
// 1. 首次初始化不需要动画
|
||||
@@ -123,7 +121,7 @@ const scheduleIndicatorUpdate = () => {
|
||||
|
||||
// 监听 activeTab 变化
|
||||
watch(
|
||||
() => activeTab?.value,
|
||||
() => [activeTab?.value, locale.value],
|
||||
() => {
|
||||
nextTick(() => {
|
||||
scheduleIndicatorUpdate()
|
||||
@@ -166,9 +164,11 @@ onUnmounted(() => {
|
||||
<style scoped>
|
||||
.tabs-list {
|
||||
position: relative;
|
||||
height: 2.5rem;
|
||||
min-height: 2.5rem;
|
||||
max-width: 100%;
|
||||
overflow-x: auto;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
justify-content: flex-start;
|
||||
border-radius: 0.5rem;
|
||||
background-color: hsl(var(--muted) / 0.3);
|
||||
padding: 0.25rem;
|
||||
@@ -176,6 +176,12 @@ onUnmounted(() => {
|
||||
border: 1px solid hsl(var(--border) / 0.6);
|
||||
}
|
||||
|
||||
.tabs-list.grid :deep(button[data-value]) {
|
||||
min-width: 0;
|
||||
white-space: normal;
|
||||
overflow-wrap: anywhere;
|
||||
}
|
||||
|
||||
.tabs-indicator {
|
||||
position: absolute;
|
||||
z-index: 0;
|
||||
|
||||
@@ -32,7 +32,7 @@ const handleClick = () => {
|
||||
|
||||
const triggerClass = computed(() => {
|
||||
return cn(
|
||||
'relative z-10 inline-flex items-center justify-center whitespace-nowrap rounded-md px-3 py-1.5 text-sm font-medium ring-offset-background transition-colors focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2 disabled:pointer-events-none disabled:opacity-50',
|
||||
'relative z-10 inline-flex shrink-0 items-center justify-center whitespace-nowrap rounded-md px-3 py-1.5 text-sm font-medium [&_svg]:shrink-0 ring-offset-background transition-colors focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2 disabled:pointer-events-none disabled:opacity-50',
|
||||
isActive.value
|
||||
? 'text-foreground font-semibold'
|
||||
: 'text-muted-foreground hover:text-foreground',
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
import { afterEach, describe, expect, it } from 'vitest'
|
||||
import { setI18nLocale } from '@/i18n'
|
||||
import { useToast } from '../useToast'
|
||||
import { useConfirm } from '../useConfirm'
|
||||
|
||||
afterEach(() => {
|
||||
useToast().clearAll()
|
||||
useConfirm().handleCancel()
|
||||
})
|
||||
|
||||
describe('localized feedback', () => {
|
||||
it('retranslates an existing toast in both directions without changing its identity', () => {
|
||||
const { showToast, toasts, removeToast } = useToast()
|
||||
setI18nLocale('en-US')
|
||||
const id = showToast({ title: '保存', description: '保存成功', duration: 0 })
|
||||
expect(toasts.value[0]).toMatchObject({ id, title: 'Save' })
|
||||
setI18nLocale('zh-CN')
|
||||
expect(toasts.value[0]).toMatchObject({ id, title: '保存', message: '保存成功' })
|
||||
setI18nLocale('en-US')
|
||||
expect(toasts.value[0].title).toBe('Save')
|
||||
removeToast(id)
|
||||
expect(toasts.value).toEqual([])
|
||||
})
|
||||
|
||||
it('keeps the pending confirmation while its labels change language', async () => {
|
||||
const { confirm, state, handleConfirm } = useConfirm()
|
||||
setI18nLocale('en-US')
|
||||
const result = confirm({ message: '保存', confirmText: '保存' })
|
||||
expect(state.value.confirmText).toBe('Save')
|
||||
setI18nLocale('zh-CN')
|
||||
expect(state.value).toMatchObject({ isOpen: true, message: '保存', confirmText: '保存' })
|
||||
setI18nLocale('en-US')
|
||||
expect(state.value.message).toBe('Save')
|
||||
handleConfirm()
|
||||
expect(await result).toBe(true)
|
||||
expect(state.value.isOpen).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -1,4 +1,4 @@
|
||||
import { ref } from 'vue'
|
||||
import { computed, ref } from 'vue'
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
import { translateLegacyText } from '@/i18n/messages'
|
||||
|
||||
@@ -40,10 +40,10 @@ export function useConfirm() {
|
||||
return new Promise((resolve) => {
|
||||
state.value = {
|
||||
isOpen: true,
|
||||
title: localizeConfirmText(options.title || '确认操作'),
|
||||
message: localizeConfirmText(options.message),
|
||||
confirmText: localizeConfirmText(options.confirmText || '确认'),
|
||||
cancelText: localizeConfirmText(options.cancelText || '取消'),
|
||||
title: options.title || '确认操作',
|
||||
message: options.message,
|
||||
confirmText: options.confirmText || '确认',
|
||||
cancelText: options.cancelText || '取消',
|
||||
variant: options.variant || 'question',
|
||||
resolve
|
||||
}
|
||||
@@ -107,7 +107,13 @@ export function useConfirm() {
|
||||
}
|
||||
|
||||
return {
|
||||
state,
|
||||
state: computed(() => ({
|
||||
...state.value,
|
||||
title: localizeConfirmText(state.value.title || '确认操作'),
|
||||
message: localizeConfirmText(state.value.message),
|
||||
confirmText: localizeConfirmText(state.value.confirmText || '确认'),
|
||||
cancelText: localizeConfirmText(state.value.cancelText || '取消'),
|
||||
})),
|
||||
confirm,
|
||||
confirmDanger,
|
||||
confirmWarning,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { ref } from 'vue'
|
||||
import { computed, ref } from 'vue'
|
||||
import { TOAST_CONFIG } from '@/config/constants'
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
import { translateLegacyText } from '@/i18n/messages'
|
||||
@@ -36,8 +36,8 @@ export function useToast() {
|
||||
duration: 5000,
|
||||
...toastOptions,
|
||||
variant: normalizeToastVariant(options.variant),
|
||||
title: localizeToastText(options.title),
|
||||
message: localizeToastText(options.message ?? description),
|
||||
title: options.title,
|
||||
message: options.message ?? description,
|
||||
}
|
||||
|
||||
|
||||
@@ -81,7 +81,11 @@ export function useToast() {
|
||||
}
|
||||
|
||||
return {
|
||||
toasts,
|
||||
toasts: computed(() => toasts.value.map(toast => ({
|
||||
...toast,
|
||||
title: localizeToastText(toast.title),
|
||||
message: localizeToastText(toast.message),
|
||||
}))),
|
||||
showToast,
|
||||
removeToast,
|
||||
toast: showToast,
|
||||
|
||||
@@ -533,6 +533,7 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
import { ref, watch, computed } from 'vue'
|
||||
import {
|
||||
X,
|
||||
@@ -760,7 +761,7 @@ function handleClose() {
|
||||
function formatDate(dateStr: string): string {
|
||||
if (!dateStr) return '-'
|
||||
const date = new Date(dateStr)
|
||||
return date.toLocaleDateString('zh-CN', {
|
||||
return date.toLocaleDateString(getI18nLocale(), {
|
||||
year: 'numeric',
|
||||
month: '2-digit',
|
||||
day: '2-digit'
|
||||
|
||||
@@ -0,0 +1,323 @@
|
||||
<template>
|
||||
<Card
|
||||
variant="interactive"
|
||||
class="flex min-w-0 flex-col cursor-pointer overflow-hidden"
|
||||
@mousedown="$emit('mousedown', $event)"
|
||||
@click="$emit('rowClick', $event, provider.id)"
|
||||
>
|
||||
<div class="flex items-start gap-2 p-4 pb-3">
|
||||
<slot name="drag-handle" />
|
||||
<div
|
||||
class="flex h-10 w-10 shrink-0 items-center justify-center rounded-xl text-base font-semibold"
|
||||
:class="provider.is_active ? 'bg-primary/10 text-primary' : 'bg-muted text-muted-foreground'"
|
||||
aria-hidden="true"
|
||||
>
|
||||
{{ provider.name.slice(0, 1).toUpperCase() }}
|
||||
</div>
|
||||
<div class="min-w-0 flex-1 space-y-1">
|
||||
<div class="flex items-center gap-1.5">
|
||||
<button
|
||||
type="button"
|
||||
class="min-w-0 truncate rounded text-left text-sm font-semibold text-foreground hover:text-primary focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring"
|
||||
:title="provider.name"
|
||||
@click.stop="$emit('viewDetail', provider.id)"
|
||||
>
|
||||
{{ provider.name }}
|
||||
</button>
|
||||
<a
|
||||
v-if="safeProviderWebsite"
|
||||
:href="safeProviderWebsite"
|
||||
target="_blank"
|
||||
rel="noopener noreferrer"
|
||||
class="shrink-0 text-muted-foreground transition-colors hover:text-primary"
|
||||
:title="safeProviderWebsite"
|
||||
@click.stop
|
||||
>
|
||||
<ExternalLink class="h-3.5 w-3.5" />
|
||||
</a>
|
||||
</div>
|
||||
<div
|
||||
v-if="editingDescriptionId === provider.id"
|
||||
data-desc-editor
|
||||
class="flex items-center gap-1"
|
||||
@click.stop
|
||||
>
|
||||
<input
|
||||
v-model="localDescriptionValue"
|
||||
v-auto-focus
|
||||
class="min-w-0 flex-1 rounded border border-border bg-background px-1.5 py-0.5 text-xs text-foreground focus:outline-none focus:ring-1 focus:ring-primary/50"
|
||||
:placeholder="legacyT('输入备注...')"
|
||||
:aria-label="legacyT('输入备注...')"
|
||||
@keydown="handleDescriptionKeydown"
|
||||
>
|
||||
<button
|
||||
type="button"
|
||||
class="shrink-0 rounded p-0.5 text-primary hover:bg-muted"
|
||||
:title="legacyT('保存')"
|
||||
:aria-label="legacyT('保存')"
|
||||
@click="handleSave"
|
||||
>
|
||||
<Check class="h-3.5 w-3.5" />
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
class="shrink-0 rounded p-0.5 text-muted-foreground hover:bg-muted"
|
||||
:title="legacyT('取消')"
|
||||
:aria-label="legacyT('取消')"
|
||||
@click="handleCancel"
|
||||
>
|
||||
<X class="h-3.5 w-3.5" />
|
||||
</button>
|
||||
</div>
|
||||
<button
|
||||
v-else
|
||||
type="button"
|
||||
class="group/desc flex max-w-full items-center gap-1 rounded text-xs text-muted-foreground transition-colors hover:text-foreground focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring"
|
||||
:title="provider.description || legacyT('添加备注')"
|
||||
@click="handleStartEdit"
|
||||
>
|
||||
<span class="truncate">{{ provider.description || legacyT('添加备注') }}</span>
|
||||
<Pencil class="h-3 w-3 shrink-0 opacity-0 transition-opacity group-hover/desc:opacity-60 group-focus-visible/desc:opacity-60" />
|
||||
</button>
|
||||
</div>
|
||||
<Badge
|
||||
:variant="provider.is_active ? 'success' : 'secondary'"
|
||||
class="shrink-0 text-xs"
|
||||
>
|
||||
{{ legacyT(provider.is_active ? '活跃' : '停用') }}
|
||||
</Badge>
|
||||
</div>
|
||||
|
||||
<div class="flex flex-1 flex-col gap-4 px-4 pb-4">
|
||||
<div class="space-y-2 rounded-xl border border-border/40 bg-muted/20 p-3">
|
||||
<div class="flex flex-wrap items-center justify-between gap-2">
|
||||
<span class="text-xs text-muted-foreground">{{ legacyT('余额监控') }}</span>
|
||||
<Badge
|
||||
variant="outline"
|
||||
class="border-border/50 text-[10px] font-normal"
|
||||
>
|
||||
{{ formatBillingType(provider.billing_type || 'pay_as_you_go') }}
|
||||
</Badge>
|
||||
</div>
|
||||
<ProviderBalanceCell
|
||||
:provider="provider"
|
||||
:is-balance-loading="isBalanceLoading"
|
||||
:get-provider-balance="getProviderBalance"
|
||||
:get-provider-balance-breakdown="getProviderBalanceBreakdown"
|
||||
:get-provider-balance-error="getProviderBalanceError"
|
||||
:get-provider-checkin="getProviderCheckin"
|
||||
:get-provider-cookie-expired="getProviderCookieExpired"
|
||||
:get-provider-balance-extra="getProviderBalanceExtra"
|
||||
:format-balance-display="formatBalanceDisplay"
|
||||
:format-reset-countdown="formatResetCountdown"
|
||||
:get-quota-used-color-class="getQuotaUsedColorClass"
|
||||
/>
|
||||
</div>
|
||||
|
||||
<dl class="grid grid-cols-3 divide-x divide-border/50 text-center">
|
||||
<div class="min-w-0 space-y-1 px-1">
|
||||
<dt class="text-xs text-muted-foreground">
|
||||
{{ legacyT('端点') }}
|
||||
</dt>
|
||||
<dd class="text-sm font-semibold tabular-nums">
|
||||
{{ provider.active_endpoints }}<span class="ml-0.5 text-xs font-normal text-muted-foreground">/ {{ provider.total_endpoints }}</span>
|
||||
</dd>
|
||||
</div>
|
||||
<div class="min-w-0 space-y-1 px-1">
|
||||
<dt class="text-xs text-muted-foreground">
|
||||
{{ legacyT(isKeyManagedProviderType(provider.provider_type) ? '密钥' : '账号') }}
|
||||
</dt>
|
||||
<dd class="text-sm font-semibold tabular-nums">
|
||||
{{ provider.active_keys }}<span class="ml-0.5 text-xs font-normal text-muted-foreground">/ {{ provider.total_keys }}</span>
|
||||
</dd>
|
||||
</div>
|
||||
<div class="min-w-0 space-y-1 px-1">
|
||||
<dt class="text-xs text-muted-foreground">
|
||||
{{ legacyT('模型') }}
|
||||
</dt>
|
||||
<dd class="text-sm font-semibold tabular-nums">
|
||||
{{ provider.active_models }}<span class="ml-0.5 text-xs font-normal text-muted-foreground">/ {{ provider.total_models }}</span>
|
||||
</dd>
|
||||
</div>
|
||||
</dl>
|
||||
|
||||
<div class="mt-auto space-y-2 border-t border-border/40 pt-3">
|
||||
<div class="text-xs text-muted-foreground">
|
||||
{{ legacyT('端点健康') }}
|
||||
</div>
|
||||
<div
|
||||
v-if="provider.endpoint_health_details?.length"
|
||||
class="grid grid-cols-3 gap-x-3 gap-y-2"
|
||||
>
|
||||
<div
|
||||
v-for="endpoint in sortEndpoints(provider.endpoint_health_details)"
|
||||
:key="endpoint.api_format"
|
||||
class="flex min-w-0 flex-col gap-1.5"
|
||||
:title="getEndpointTooltip(endpoint, locale)"
|
||||
>
|
||||
<div class="flex items-center justify-between gap-1 text-[10px] leading-none text-muted-foreground">
|
||||
<span class="font-medium">{{ formatApiFormatShort(endpoint.api_format) }}</span>
|
||||
<span class="tabular-nums">{{ getEndpointHealthLabel(endpoint) }}</span>
|
||||
</div>
|
||||
<div class="h-1.5 w-full overflow-hidden rounded-full bg-border dark:bg-border/80">
|
||||
<div
|
||||
class="h-full rounded-full transition-all duration-300"
|
||||
:class="getEndpointDotColor(endpoint)"
|
||||
:style="{ width: getEndpointHealthBarWidth(endpoint) }"
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<span
|
||||
v-else
|
||||
class="text-xs text-muted-foreground/60"
|
||||
>{{ legacyT('暂无端点') }}</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div
|
||||
class="flex items-center justify-between gap-2 border-t border-border/40 bg-muted/10 px-3 py-2"
|
||||
@click.stop
|
||||
>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-8 text-muted-foreground hover:text-primary"
|
||||
:title="legacyT('查看详情')"
|
||||
:aria-label="legacyT('查看详情')"
|
||||
@click="$emit('viewDetail', provider.id)"
|
||||
>
|
||||
<Eye class="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
<div class="flex shrink-0 items-center gap-0.5">
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-8 text-muted-foreground hover:text-foreground"
|
||||
:title="legacyT('编辑提供商')"
|
||||
:aria-label="legacyT('编辑提供商')"
|
||||
@click="$emit('editProvider', provider)"
|
||||
>
|
||||
<Edit class="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-8 text-muted-foreground hover:text-foreground"
|
||||
:title="legacyT('扩展操作配置')"
|
||||
:aria-label="legacyT('扩展操作配置')"
|
||||
@click="$emit('openOpsConfig', provider)"
|
||||
>
|
||||
<KeyRound class="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-8 text-muted-foreground hover:text-foreground"
|
||||
:title="legacyT(provider.is_active ? '停用提供商' : '启用提供商')"
|
||||
:aria-label="legacyT(provider.is_active ? '停用提供商' : '启用提供商')"
|
||||
@click="$emit('toggleStatus', provider)"
|
||||
>
|
||||
<Power class="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-8 text-muted-foreground hover:text-destructive"
|
||||
:title="legacyT('删除提供商')"
|
||||
:aria-label="legacyT('删除提供商')"
|
||||
@click="$emit('deleteProvider', provider)"
|
||||
>
|
||||
<Trash2 class="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
</Card>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, ref, watch } from 'vue'
|
||||
import { Check, Edit, ExternalLink, Eye, KeyRound, Pencil, Power, Trash2, X } from 'lucide-vue-next'
|
||||
import Badge from '@/components/ui/badge.vue'
|
||||
import Button from '@/components/ui/button.vue'
|
||||
import Card from '@/components/ui/card.vue'
|
||||
import ProviderBalanceCell from './ProviderBalanceCell.vue'
|
||||
import { formatApiFormatShort, type ProviderWithEndpointsSummary } from '@/api/endpoints'
|
||||
import type { BalanceExtraItem } from '@/features/providers/auth-templates'
|
||||
import {
|
||||
sortEndpoints,
|
||||
getEndpointHealthLabel,
|
||||
getEndpointHealthBarWidth,
|
||||
getEndpointDotColor,
|
||||
getEndpointTooltip,
|
||||
} from '@/features/providers/composables/useEndpointStatus'
|
||||
import { isKeyManagedProviderType } from '../utils/providerTypeUtils'
|
||||
import { formatBillingType } from '@/utils/format'
|
||||
import { safeExternalWebUrl } from '@/utils/navigationSecurity'
|
||||
import { useI18n } from '@/i18n'
|
||||
|
||||
const props = defineProps<{
|
||||
provider: ProviderWithEndpointsSummary
|
||||
editingDescriptionId: string | null
|
||||
isBalanceLoading: (providerId: string) => boolean
|
||||
getProviderBalance: (providerId: string) => { available: number | null; currency: string } | null
|
||||
getProviderBalanceBreakdown: (providerId: string) => { balance: number; points: number; currency: string } | null
|
||||
getProviderBalanceError: (providerId: string) => { status: string; message: string } | null
|
||||
getProviderCheckin: (providerId: string) => { success: boolean | null; message: string } | null
|
||||
getProviderCookieExpired: (providerId: string) => { expired: boolean; message: string } | null
|
||||
getProviderBalanceExtra: (providerId: string, architectureId?: string) => BalanceExtraItem[]
|
||||
formatBalanceDisplay: (balance: { available: number | null; currency: string } | null) => string
|
||||
formatResetCountdown: (resetsAt: number) => string
|
||||
getQuotaUsedColorClass: (provider: ProviderWithEndpointsSummary) => string
|
||||
}>()
|
||||
|
||||
const emit = defineEmits<{
|
||||
'mousedown': [event: MouseEvent]
|
||||
'rowClick': [event: MouseEvent, providerId: string]
|
||||
'viewDetail': [providerId: string]
|
||||
'editProvider': [provider: ProviderWithEndpointsSummary]
|
||||
'openOpsConfig': [provider: ProviderWithEndpointsSummary]
|
||||
'toggleStatus': [provider: ProviderWithEndpointsSummary]
|
||||
'deleteProvider': [provider: ProviderWithEndpointsSummary]
|
||||
'startEditDescription': [event: Event, provider: ProviderWithEndpointsSummary]
|
||||
'saveDescription': [event: Event, provider: ProviderWithEndpointsSummary, value: string]
|
||||
'cancelEditDescription': [event?: Event]
|
||||
}>()
|
||||
|
||||
const { legacyT, locale } = useI18n()
|
||||
const safeProviderWebsite = computed(() => safeExternalWebUrl(props.provider.website))
|
||||
const localDescriptionValue = ref('')
|
||||
const vAutoFocus = {
|
||||
mounted: (element: HTMLElement) => element.focus(),
|
||||
}
|
||||
|
||||
watch(() => props.editingDescriptionId, (providerId) => {
|
||||
if (providerId === props.provider.id) {
|
||||
localDescriptionValue.value = props.provider.description || ''
|
||||
}
|
||||
}, { immediate: true })
|
||||
|
||||
function handleStartEdit(event: Event) {
|
||||
event.stopPropagation()
|
||||
emit('startEditDescription', event, props.provider)
|
||||
}
|
||||
|
||||
function handleSave(event: Event) {
|
||||
event.stopPropagation()
|
||||
emit('saveDescription', event, props.provider, localDescriptionValue.value)
|
||||
}
|
||||
|
||||
function handleCancel(event: Event) {
|
||||
event.stopPropagation()
|
||||
emit('cancelEditDescription', event)
|
||||
}
|
||||
|
||||
function handleDescriptionKeydown(event: KeyboardEvent) {
|
||||
if (event.key === 'Enter') {
|
||||
event.preventDefault()
|
||||
handleSave(event)
|
||||
} else if (event.key === 'Escape') {
|
||||
handleCancel(event)
|
||||
}
|
||||
}
|
||||
</script>
|
||||
@@ -940,7 +940,7 @@ import {
|
||||
} from 'lucide-vue-next'
|
||||
import { parseApiError } from '@/utils/errorParser'
|
||||
import { useEscapeKey } from '@/composables/useEscapeKey'
|
||||
import { useI18n } from '@/i18n'
|
||||
import { getI18nLocale, useI18n } from '@/i18n'
|
||||
import Button from '@/components/ui/button.vue'
|
||||
import Card from '@/components/ui/card.vue'
|
||||
import { useToast } from '@/composables/useToast'
|
||||
@@ -980,6 +980,7 @@ import ProviderQuotaProgressRow from '@/features/providers/components/ProviderQu
|
||||
import ProviderQuotaSectionHeader from '@/features/providers/components/ProviderQuotaSectionHeader.vue'
|
||||
import { useProxyNodesStore } from '@/stores/proxy-nodes'
|
||||
import { resolveAntigravityQuotaGroupLabel } from '@/features/providers/utils/antigravityQuota'
|
||||
import { refreshQuotaInBackground } from '@/features/providers/utils/refreshQuotaInBackground'
|
||||
import {
|
||||
deleteEndpointKey,
|
||||
recoverKeyHealth,
|
||||
@@ -2476,7 +2477,7 @@ function isKiroBannedKey(key: EndpointAPIKey): boolean {
|
||||
function formatBanTimestamp(timestamp: number | undefined): string {
|
||||
if (!timestamp) return ''
|
||||
const date = new Date(timestamp * 1000)
|
||||
return date.toLocaleString('zh-CN', {
|
||||
return date.toLocaleString(getI18nLocale(), {
|
||||
month: 'short',
|
||||
day: 'numeric',
|
||||
hour: '2-digit',
|
||||
@@ -2862,15 +2863,21 @@ async function autoRefreshQuotaInBackground(): Promise<boolean> {
|
||||
}
|
||||
|
||||
refreshingQuota.value = true
|
||||
const isCurrent = () => props.open && props.providerId === providerId
|
||||
try {
|
||||
const result = await refreshProviderQuota(providerId)
|
||||
const result = await refreshQuotaInBackground({
|
||||
refresh: () => refreshProviderQuota(providerId),
|
||||
isCurrent,
|
||||
retryInitialEmptyQuota: providerType === 'antigravity' && !hadCachedQuota,
|
||||
})
|
||||
if (!result) return false
|
||||
const applied = applyQuotaResults(result.results)
|
||||
if (result.success <= 0 && applied === 0 && !hadCachedQuota && providerType === 'antigravity') {
|
||||
showError(legacyT('没有获取到配额信息(请检查账号是否已授权、project_id 是否存在)'), legacyT('提示'))
|
||||
showWarning(legacyT('配额暂未就绪,请稍后刷新'), legacyT('提示'))
|
||||
}
|
||||
return applied > 0
|
||||
} catch (err: unknown) {
|
||||
if (!hadCachedQuota && providerType === 'antigravity') {
|
||||
if (isCurrent() && !hadCachedQuota && providerType === 'antigravity') {
|
||||
showError(localizedApiError(err, '后台刷新配额失败'), legacyT('错误'))
|
||||
}
|
||||
return false
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
<template>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-5 touch-none select-none cursor-grab text-muted-foreground/50 hover:text-primary active:cursor-grabbing disabled:cursor-default"
|
||||
:disabled="disabled"
|
||||
:title="legacyT('拖动调整当前页展示顺序,也可使用方向键移动')"
|
||||
:aria-label="`${legacyT('调整展示顺序')}: ${providerName}`"
|
||||
data-provider-drag-handle
|
||||
@pointerdown.stop="$emit('pointerdown', $event)"
|
||||
@keydown.stop="$emit('keydown', $event)"
|
||||
@mousedown.stop
|
||||
@click.stop
|
||||
@dragstart.prevent
|
||||
>
|
||||
<GripVertical class="h-4 w-4" />
|
||||
</Button>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { GripVertical } from 'lucide-vue-next'
|
||||
import Button from '@/components/ui/button.vue'
|
||||
import { useI18n } from '@/i18n'
|
||||
|
||||
defineProps<{
|
||||
providerName: string
|
||||
disabled: boolean
|
||||
}>()
|
||||
|
||||
defineEmits<{
|
||||
pointerdown: [event: PointerEvent]
|
||||
keydown: [event: KeyboardEvent]
|
||||
}>()
|
||||
|
||||
const { legacyT } = useI18n()
|
||||
</script>
|
||||
@@ -26,7 +26,7 @@
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="grid grid-cols-2 gap-4">
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-4">
|
||||
<div class="space-y-1.5">
|
||||
<Label>{{ legacyT('提供商类型') }}</Label>
|
||||
<Select
|
||||
@@ -130,7 +130,7 @@
|
||||
</h3>
|
||||
|
||||
<!-- 超时配置 -->
|
||||
<div class="grid grid-cols-2 gap-4">
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-4">
|
||||
<div class="space-y-1.5">
|
||||
<Label>
|
||||
{{ legacyT('流式首字节超时') }}
|
||||
@@ -164,11 +164,11 @@
|
||||
</div>
|
||||
|
||||
<!-- 提供商内转移限制 -->
|
||||
<div class="grid grid-cols-2 gap-2 sm:gap-4">
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-4">
|
||||
<div class="min-w-0 space-y-1.5">
|
||||
<Label
|
||||
for="max-transfer-count"
|
||||
class="whitespace-nowrap text-xs sm:text-sm"
|
||||
class="text-xs sm:text-sm"
|
||||
>
|
||||
{{ legacyT('最大转移次数') }}
|
||||
</Label>
|
||||
@@ -185,7 +185,7 @@
|
||||
<div class="min-w-0 space-y-1.5">
|
||||
<Label
|
||||
for="max-transfer-timeout-seconds"
|
||||
class="whitespace-nowrap text-xs sm:text-sm"
|
||||
class="text-xs sm:text-sm"
|
||||
>
|
||||
{{ legacyT('最大转移超时') }}
|
||||
<span class="text-xs text-muted-foreground">{{ legacyT('(秒)') }}</span>
|
||||
|
||||
@@ -9,7 +9,10 @@
|
||||
@update:model-value="handleDialogUpdate"
|
||||
>
|
||||
<div class="space-y-3.5">
|
||||
<nav class="grid grid-cols-3 gap-1.5 rounded-xl bg-muted/40 p-1.5" aria-label="批量导入步骤">
|
||||
<nav
|
||||
class="grid grid-cols-3 gap-1.5 rounded-xl bg-muted/40 p-1.5"
|
||||
aria-label="批量导入步骤"
|
||||
>
|
||||
<button
|
||||
v-for="step in steps"
|
||||
:key="step.id"
|
||||
@@ -37,8 +40,12 @@
|
||||
<div class="flex min-w-0 items-center gap-3">
|
||||
<span class="flex h-7 w-7 shrink-0 items-center justify-center rounded-md bg-foreground text-xs font-semibold text-background">1</span>
|
||||
<div class="min-w-0">
|
||||
<h3 class="text-balance text-sm font-semibold text-foreground">粘贴名称与 Key</h3>
|
||||
<p class="text-pretty text-[11px] leading-4 text-muted-foreground">每行一条,仅接受四个短横线分隔</p>
|
||||
<h3 class="text-balance text-sm font-semibold text-foreground">
|
||||
粘贴名称与 Key
|
||||
</h3>
|
||||
<p class="text-pretty text-[11px] leading-4 text-muted-foreground">
|
||||
每行一条,仅接受四个短横线分隔
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
<Badge
|
||||
@@ -50,34 +57,40 @@
|
||||
</header>
|
||||
|
||||
<div class="min-w-0">
|
||||
<Label for="provider-key-batch-input" class="sr-only">Key 列表</Label>
|
||||
<Textarea
|
||||
id="provider-key-batch-input"
|
||||
v-model="inputText"
|
||||
class="h-[280px] min-h-[220px] max-h-[520px] !resize-y !rounded-none !border-0 !bg-transparent !px-4 !py-4 font-mono text-[13px] leading-6 !shadow-none !ring-0 focus-visible:!ring-0"
|
||||
spellcheck="false"
|
||||
placeholder="主账号----sk-xxxx 备用账号----sk-yyyy"
|
||||
/>
|
||||
<Label
|
||||
for="provider-key-batch-input"
|
||||
class="sr-only"
|
||||
>Key 列表</Label>
|
||||
<Textarea
|
||||
id="provider-key-batch-input"
|
||||
v-model="inputText"
|
||||
class="h-[280px] min-h-[220px] max-h-[520px] !resize-y !rounded-none !border-0 !bg-transparent !px-4 !py-4 font-mono text-[13px] leading-6 !shadow-none !ring-0 focus-visible:!ring-0"
|
||||
spellcheck="false"
|
||||
placeholder="主账号----sk-xxxx 备用账号----sk-yyyy"
|
||||
/>
|
||||
<div
|
||||
v-if="parsed.errors.length > 0"
|
||||
class="mx-4 mb-3 space-y-1 rounded-lg bg-destructive/5 px-3 py-2 text-[11px] text-destructive ring-1 ring-destructive/20"
|
||||
>
|
||||
<div
|
||||
v-if="parsed.errors.length > 0"
|
||||
class="mx-4 mb-3 space-y-1 rounded-lg bg-destructive/5 px-3 py-2 text-[11px] text-destructive ring-1 ring-destructive/20"
|
||||
v-for="(item, index) in parsed.errors.slice(0, 6)"
|
||||
:key="`${item.lineNumber}-${index}`"
|
||||
>
|
||||
<div
|
||||
v-for="(item, index) in parsed.errors.slice(0, 6)"
|
||||
:key="`${item.lineNumber}-${index}`"
|
||||
>
|
||||
{{ item.lineNumber ? `第 ${item.lineNumber} 行:` : '' }}{{ item.message }}
|
||||
</div>
|
||||
<div v-if="parsed.errors.length > 6" class="font-medium">
|
||||
另有 {{ parsed.errors.length - 6 }} 个问题
|
||||
</div>
|
||||
{{ item.lineNumber ? `第 ${item.lineNumber} 行:` : '' }}{{ item.message }}
|
||||
</div>
|
||||
<div class="flex flex-wrap items-center gap-2 border-t border-border/50 bg-muted/10 px-4 py-2.5 text-[11px] text-muted-foreground">
|
||||
<span class="rounded-md bg-muted px-2 py-1 font-mono text-foreground/80">名称----Key</span>
|
||||
<span>名称和 Key 都不能为空</span>
|
||||
<span class="ml-auto hidden tabular-nums sm:inline">已识别 {{ parsed.items.length }} 条</span>
|
||||
<div
|
||||
v-if="parsed.errors.length > 6"
|
||||
class="font-medium"
|
||||
>
|
||||
另有 {{ parsed.errors.length - 6 }} 个问题
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex flex-wrap items-center gap-2 border-t border-border/50 bg-muted/10 px-4 py-2.5 text-[11px] text-muted-foreground">
|
||||
<span class="rounded-md bg-muted px-2 py-1 font-mono text-foreground/80">名称----Key</span>
|
||||
<span>名称和 Key 都不能为空</span>
|
||||
<span class="ml-auto hidden tabular-nums sm:inline">已识别 {{ parsed.items.length }} 条</span>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section
|
||||
@@ -98,7 +111,10 @@
|
||||
>{{ item }}</span>
|
||||
</span>
|
||||
</span>
|
||||
<Badge :variant="selectedApiFormats.length > 0 ? 'success' : 'destructive'" class="ml-auto shrink-0 tabular-nums">
|
||||
<Badge
|
||||
:variant="selectedApiFormats.length > 0 ? 'success' : 'destructive'"
|
||||
class="ml-auto shrink-0 tabular-nums"
|
||||
>
|
||||
{{ selectedApiFormats.length }} 种格式
|
||||
</Badge>
|
||||
</header>
|
||||
@@ -123,12 +139,18 @@
|
||||
<header class="flex min-h-[72px] items-center gap-3 border-b border-border/60 bg-muted/15 px-4 py-3">
|
||||
<span class="flex h-7 w-7 shrink-0 items-center justify-center rounded-md bg-foreground text-xs font-semibold text-background">3</span>
|
||||
<div class="min-w-0 flex-1">
|
||||
<h3 class="text-balance text-sm font-semibold">逐项确认</h3>
|
||||
<p class="text-pretty text-[11px] leading-4 text-muted-foreground">展开任意 Key 可修改内容或设置单独配置</p>
|
||||
<h3 class="text-balance text-sm font-semibold">
|
||||
逐项确认
|
||||
</h3>
|
||||
<p class="text-pretty text-[11px] leading-4 text-muted-foreground">
|
||||
展开任意 Key 可修改内容或设置单独配置
|
||||
</p>
|
||||
</div>
|
||||
<div class="shrink-0 text-right text-[11px] text-muted-foreground">
|
||||
<div><span class="font-semibold tabular-nums text-foreground">{{ reviewItems.length }}</span> 个 Key</div>
|
||||
<div v-if="customizedItemCount > 0"><span class="tabular-nums">{{ customizedItemCount }}</span> 个单独配置</div>
|
||||
<div v-if="customizedItemCount > 0">
|
||||
<span class="tabular-nums">{{ customizedItemCount }}</span> 个单独配置
|
||||
</div>
|
||||
</div>
|
||||
</header>
|
||||
|
||||
@@ -156,9 +178,13 @@
|
||||
v-if="reviewErrorsByIndex.has(entry.index)"
|
||||
variant="destructive"
|
||||
class="shrink-0 text-[10px]"
|
||||
>需修正</Badge>
|
||||
>
|
||||
需修正
|
||||
</Badge>
|
||||
</div>
|
||||
<div class="mt-0.5 truncate font-mono text-[10px] text-muted-foreground">
|
||||
{{ maskSecret(entry.item.apiKey) }}
|
||||
</div>
|
||||
<div class="mt-0.5 truncate font-mono text-[10px] text-muted-foreground">{{ maskSecret(entry.item.apiKey) }}</div>
|
||||
</div>
|
||||
<div class="hidden shrink-0 items-center gap-1.5 sm:flex">
|
||||
<span class="rounded-md bg-muted px-2 py-0.5 text-[10px] text-muted-foreground">{{ effectiveAuthLabel(entry.item) }}</span>
|
||||
@@ -186,18 +212,30 @@
|
||||
<div class="grid gap-3 sm:grid-cols-2">
|
||||
<div class="space-y-1.5">
|
||||
<Label class="text-xs">名称</Label>
|
||||
<Input v-model="entry.item.name" class="h-10" placeholder="必填" />
|
||||
<Input
|
||||
v-model="entry.item.name"
|
||||
class="h-10"
|
||||
placeholder="必填"
|
||||
/>
|
||||
</div>
|
||||
<div class="space-y-1.5">
|
||||
<Label class="text-xs">Key</Label>
|
||||
<Input v-model="entry.item.apiKey" class="h-10 font-mono text-xs" placeholder="必填" />
|
||||
<Input
|
||||
v-model="entry.item.apiKey"
|
||||
class="h-10 font-mono text-xs"
|
||||
placeholder="必填"
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="flex min-h-12 items-center justify-between gap-3 rounded-lg bg-background px-3 shadow-[0_0_0_1px_rgb(0_0_0/0.06)] dark:shadow-[0_0_0_1px_rgb(255_255_255/0.08)]">
|
||||
<div>
|
||||
<div class="text-xs font-medium">单独配置此 Key</div>
|
||||
<div class="text-[11px] text-muted-foreground">开启后覆盖第二步中的统一配置</div>
|
||||
<div class="text-xs font-medium">
|
||||
单独配置此 Key
|
||||
</div>
|
||||
<div class="text-[11px] text-muted-foreground">
|
||||
开启后覆盖第二步中的统一配置
|
||||
</div>
|
||||
</div>
|
||||
<Switch
|
||||
:model-value="entry.item.customized"
|
||||
@@ -220,7 +258,12 @@
|
||||
v-if="reviewErrorsByIndex.has(entry.index)"
|
||||
class="space-y-1 rounded-lg bg-destructive/5 px-3 py-2 text-[11px] text-destructive"
|
||||
>
|
||||
<div v-for="message in reviewErrorsByIndex.get(entry.index)" :key="message">{{ message }}</div>
|
||||
<div
|
||||
v-for="message in reviewErrorsByIndex.get(entry.index)"
|
||||
:key="message"
|
||||
>
|
||||
{{ message }}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</article>
|
||||
@@ -230,12 +273,24 @@
|
||||
v-if="reviewPageCount > 1"
|
||||
class="flex min-h-12 items-center justify-between gap-3 border-t border-border/60 bg-muted/10 px-3 sm:px-4"
|
||||
>
|
||||
<Button variant="ghost" size="sm" class="h-9" :disabled="reviewPage === 1" @click="changeReviewPage(reviewPage - 1)">
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
class="h-9"
|
||||
:disabled="reviewPage === 1"
|
||||
@click="changeReviewPage(reviewPage - 1)"
|
||||
>
|
||||
<ChevronLeft class="mr-1 h-4 w-4" />
|
||||
上一页
|
||||
</Button>
|
||||
<span class="text-[11px] tabular-nums text-muted-foreground">{{ reviewPage }} / {{ reviewPageCount }}</span>
|
||||
<Button variant="ghost" size="sm" class="h-9" :disabled="reviewPage === reviewPageCount" @click="changeReviewPage(reviewPage + 1)">
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
class="h-9"
|
||||
:disabled="reviewPage === reviewPageCount"
|
||||
@click="changeReviewPage(reviewPage + 1)"
|
||||
>
|
||||
下一页
|
||||
<ChevronRight class="ml-1 h-4 w-4" />
|
||||
</Button>
|
||||
@@ -245,14 +300,35 @@
|
||||
|
||||
<template #footer>
|
||||
<div class="flex w-full flex-col-reverse gap-2 sm:flex-row sm:justify-end">
|
||||
<Button class="w-full sm:w-auto" variant="outline" :disabled="importing" @click="handleBack">
|
||||
<ArrowLeft v-if="currentStep > 1" class="mr-2 h-4 w-4" />
|
||||
<Button
|
||||
class="w-full sm:w-auto"
|
||||
variant="outline"
|
||||
:disabled="importing"
|
||||
@click="handleBack"
|
||||
>
|
||||
<ArrowLeft
|
||||
v-if="currentStep > 1"
|
||||
class="mr-2 h-4 w-4"
|
||||
/>
|
||||
{{ currentStep === 1 ? '取消' : '上一步' }}
|
||||
</Button>
|
||||
<Button class="w-full sm:w-auto" :disabled="primaryActionDisabled" @click="handlePrimaryAction">
|
||||
<Loader2 v-if="importing" class="mr-2 h-4 w-4 animate-spin" />
|
||||
<ListPlus v-else-if="currentStep === 3" class="mr-2 h-4 w-4" />
|
||||
<ArrowRight v-else class="mr-2 h-4 w-4" />
|
||||
<Button
|
||||
class="w-full sm:w-auto"
|
||||
:disabled="primaryActionDisabled"
|
||||
@click="handlePrimaryAction"
|
||||
>
|
||||
<Loader2
|
||||
v-if="importing"
|
||||
class="mr-2 h-4 w-4 animate-spin"
|
||||
/>
|
||||
<ListPlus
|
||||
v-else-if="currentStep === 3"
|
||||
class="mr-2 h-4 w-4"
|
||||
/>
|
||||
<ArrowRight
|
||||
v-else
|
||||
class="mr-2 h-4 w-4"
|
||||
/>
|
||||
{{ primaryActionLabel }}
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
@@ -28,8 +28,12 @@
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="api_key">API Key</SelectItem>
|
||||
<SelectItem value="bearer">Bearer Token</SelectItem>
|
||||
<SelectItem value="api_key">
|
||||
API Key
|
||||
</SelectItem>
|
||||
<SelectItem value="bearer">
|
||||
Bearer Token
|
||||
</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
@@ -112,8 +116,12 @@
|
||||
|
||||
<div class="flex min-h-12 items-center justify-between gap-3 rounded-lg bg-background px-3 shadow-[0_0_0_1px_rgb(0_0_0/0.06),0_1px_2px_rgb(0_0_0/0.04)] dark:shadow-[0_0_0_1px_rgb(255_255_255/0.08)]">
|
||||
<div>
|
||||
<div class="text-xs font-medium">导入后立即启用</div>
|
||||
<div class="text-[11px] text-muted-foreground">关闭后仍会创建,但不会进入调度</div>
|
||||
<div class="text-xs font-medium">
|
||||
导入后立即启用
|
||||
</div>
|
||||
<div class="text-[11px] text-muted-foreground">
|
||||
关闭后仍会创建,但不会进入调度
|
||||
</div>
|
||||
</div>
|
||||
<Switch
|
||||
:model-value="settings.is_active"
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
>
|
||||
<!-- 第一行:名称 + 状态 + 操作 -->
|
||||
<div class="flex items-start justify-between gap-3">
|
||||
<slot name="drag-handle" />
|
||||
<div class="flex-1 min-w-0 space-y-0.5">
|
||||
<div class="flex items-center gap-1.5">
|
||||
<span class="font-medium text-foreground truncate">{{ provider.name }}</span>
|
||||
|
||||
@@ -22,7 +22,7 @@
|
||||
</div>
|
||||
|
||||
<!-- 状态筛选 -->
|
||||
<div class="xl:hidden">
|
||||
<div :class="{ 'xl:hidden': !cardView }">
|
||||
<Select
|
||||
:model-value="filterStatus"
|
||||
@update:model-value="$emit('update:filterStatus', $event)"
|
||||
@@ -43,7 +43,7 @@
|
||||
</div>
|
||||
|
||||
<!-- API 格式筛选 -->
|
||||
<div class="xl:hidden">
|
||||
<div :class="{ 'xl:hidden': !cardView }">
|
||||
<Select
|
||||
:model-value="filterApiFormat"
|
||||
@update:model-value="$emit('update:filterApiFormat', $event)"
|
||||
@@ -64,7 +64,7 @@
|
||||
</div>
|
||||
|
||||
<!-- 模型筛选 -->
|
||||
<div class="xl:hidden">
|
||||
<div :class="{ 'xl:hidden': !cardView }">
|
||||
<Select
|
||||
:model-value="filterModel"
|
||||
@update:model-value="$emit('update:filterModel', $event)"
|
||||
@@ -122,13 +122,32 @@
|
||||
:loading="loading"
|
||||
@click="$emit('refresh')"
|
||||
/>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-8 w-8"
|
||||
:class="{ 'bg-primary/10 text-primary hover:bg-primary/15 hover:text-primary': cardView }"
|
||||
:title="legacyT(cardView ? '切换到列表视图' : '切换到卡片视图')"
|
||||
:aria-label="legacyT('卡片视图')"
|
||||
:aria-pressed="cardView"
|
||||
@click="$emit('toggleView')"
|
||||
>
|
||||
<List
|
||||
v-if="cardView"
|
||||
class="w-3.5 h-3.5"
|
||||
/>
|
||||
<LayoutGrid
|
||||
v-else
|
||||
class="w-3.5 h-3.5"
|
||||
/>
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { Search, Plus, FilterX, Users } from 'lucide-vue-next'
|
||||
import { Search, Plus, FilterX, Users, LayoutGrid, List } from 'lucide-vue-next'
|
||||
import Button from '@/components/ui/button.vue'
|
||||
import Input from '@/components/ui/input.vue'
|
||||
import Select from '@/components/ui/select.vue'
|
||||
@@ -150,6 +169,7 @@ defineProps<{
|
||||
modelFilters: FilterOption[]
|
||||
hasActiveFilters: boolean
|
||||
loading: boolean
|
||||
cardView: boolean
|
||||
}>()
|
||||
|
||||
defineEmits<{
|
||||
@@ -161,6 +181,7 @@ defineEmits<{
|
||||
'batchProcess': []
|
||||
'addProvider': []
|
||||
'refresh': []
|
||||
'toggleView': []
|
||||
}>()
|
||||
|
||||
const { legacyT } = useI18n()
|
||||
|
||||
@@ -4,6 +4,13 @@
|
||||
@mousedown="$emit('mousedown', $event)"
|
||||
@click="$emit('rowClick', $event, provider.id)"
|
||||
>
|
||||
<TableCell
|
||||
v-if="$slots['drag-handle']"
|
||||
class="w-9 px-2 py-3.5"
|
||||
@click.stop
|
||||
>
|
||||
<slot name="drag-handle" />
|
||||
</TableCell>
|
||||
<TableCell class="py-3.5">
|
||||
<div class="space-y-0.5">
|
||||
<div class="flex items-center gap-1.5">
|
||||
|
||||
@@ -5,6 +5,7 @@ import type { ProviderWithEndpointsSummary } from '@/api/endpoints'
|
||||
import { createI18n } from '@/i18n'
|
||||
import ProviderTableRow from '../ProviderTableRow.vue'
|
||||
import ProviderMobileCard from '../ProviderMobileCard.vue'
|
||||
import ProviderCard from '../ProviderCard.vue'
|
||||
|
||||
vi.mock('../ProviderBalanceCell.vue', () => ({
|
||||
default: { render: () => null },
|
||||
@@ -74,6 +75,7 @@ function mountProvider(component: Component, healthScore: number | null) {
|
||||
describe.each([
|
||||
['desktop provider row', ProviderTableRow],
|
||||
['mobile provider card', ProviderMobileCard],
|
||||
['grid provider card', ProviderCard],
|
||||
] as const)('%s endpoint health', (_name, component) => {
|
||||
it.each([
|
||||
{ score: null, label: '-', width: '100%', color: 'bg-muted-foreground/40' },
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
|
||||
export type HealthBadgeVariant =
|
||||
| 'default'
|
||||
| 'secondary'
|
||||
@@ -127,7 +129,7 @@ export function formatMs(value?: number | null) {
|
||||
}
|
||||
|
||||
function formatDurationNumber(value: number) {
|
||||
return new Intl.NumberFormat('zh-CN', {
|
||||
return new Intl.NumberFormat(getI18nLocale(), {
|
||||
maximumFractionDigits: Math.abs(value) < 10 ? 2 : 1
|
||||
}).format(value)
|
||||
}
|
||||
@@ -144,7 +146,7 @@ export function formatAvailability(item: HealthMonitorAvailability) {
|
||||
|
||||
export function formatTps(value?: number | null) {
|
||||
if (typeof value !== 'number' || Number.isNaN(value)) return '-'
|
||||
return `${new Intl.NumberFormat('zh-CN', {
|
||||
return `${new Intl.NumberFormat(getI18nLocale(), {
|
||||
maximumFractionDigits: value < 10 ? 2 : value < 100 ? 1 : 0
|
||||
}).format(value)} tps`
|
||||
}
|
||||
@@ -202,7 +204,7 @@ function formatTimelineRequestBreakdown(
|
||||
|
||||
function formatTimelineCount(value?: number | null) {
|
||||
if (typeof value !== 'number' || Number.isNaN(value)) return '-'
|
||||
return `${new Intl.NumberFormat('zh-CN').format(value)} 次`
|
||||
return `${new Intl.NumberFormat(getI18nLocale()).format(value)} 次`
|
||||
}
|
||||
|
||||
function formatTimelineMetricAvailability(metrics?: HealthTimelineTooltipMetrics | null) {
|
||||
@@ -213,7 +215,7 @@ function formatTimelineMetricAvailability(metrics?: HealthTimelineTooltipMetrics
|
||||
}
|
||||
|
||||
export function formatCompactNumber(value: number) {
|
||||
return new Intl.NumberFormat('zh-CN', {
|
||||
return new Intl.NumberFormat(getI18nLocale(), {
|
||||
notation: 'compact',
|
||||
maximumFractionDigits: 1
|
||||
}).format(value)
|
||||
@@ -223,7 +225,7 @@ export function formatTimestamp(timestamp?: string | null) {
|
||||
if (!timestamp) return '未知时间'
|
||||
const date = new Date(timestamp)
|
||||
if (Number.isNaN(date.getTime())) return '未知时间'
|
||||
return date.toLocaleString('zh-CN', {
|
||||
return date.toLocaleString(getI18nLocale(), {
|
||||
month: '2-digit',
|
||||
day: '2-digit',
|
||||
hour: '2-digit',
|
||||
|
||||
@@ -0,0 +1,224 @@
|
||||
import { computed, nextTick, onScopeDispose, ref, watch, type Ref } from 'vue'
|
||||
import { useEventListener, useLocalStorage, useRafFn } from '@vueuse/core'
|
||||
import { useI18n } from '@/i18n'
|
||||
|
||||
interface SortableProvider {
|
||||
id: string
|
||||
name: string
|
||||
}
|
||||
|
||||
interface ProviderPointerDrag {
|
||||
providerId: string
|
||||
pointerId: number
|
||||
startX: number
|
||||
startY: number
|
||||
handle: HTMLElement
|
||||
scrollContainer: HTMLElement | null
|
||||
}
|
||||
|
||||
export function useProviderDisplayOrder<Provider extends SortableProvider>(
|
||||
providers: () => Provider[],
|
||||
container: Ref<HTMLElement | null>,
|
||||
) {
|
||||
const { legacyT } = useI18n()
|
||||
const savedOrder = useLocalStorage<string[]>('aether-provider-display-order', [])
|
||||
const knownOrder = ref<string[]>([])
|
||||
const draggingProviderId = ref<string | null>(null)
|
||||
const dropTargetId = ref<string | null>(null)
|
||||
const pointerPosition = ref({ clientX: 0, clientY: 0 })
|
||||
const announcement = ref('')
|
||||
let pointerDrag: ProviderPointerDrag | null = null
|
||||
let suppressClickUntil = 0
|
||||
|
||||
const normalizedOrder = computed(() => Array.isArray(savedOrder.value)
|
||||
? [...new Set(savedOrder.value.filter((providerId): providerId is string => typeof providerId === 'string'))]
|
||||
: [])
|
||||
|
||||
const orderedProviders = computed(() => {
|
||||
const ranks = new Map(normalizedOrder.value.map((providerId, index) => [providerId, index]))
|
||||
return [...providers()].sort((first, second) => (
|
||||
(ranks.get(first.id) ?? ranks.size) - (ranks.get(second.id) ?? ranks.size)
|
||||
))
|
||||
})
|
||||
|
||||
const draggingProvider = computed(() => orderedProviders.value.find(provider => provider.id === draggingProviderId.value))
|
||||
const dragPreviewStyle = computed(() => ({
|
||||
left: `${Math.max(8, Math.min(pointerPosition.value.clientX + 12, window.innerWidth - 208))}px`,
|
||||
top: `${Math.max(8, Math.min(pointerPosition.value.clientY + 12, window.innerHeight - 48))}px`,
|
||||
}))
|
||||
|
||||
function moveProvider(providerId: string, targetId: string) {
|
||||
const visibleIds = orderedProviders.value.map(provider => provider.id)
|
||||
const sourceIndex = visibleIds.indexOf(providerId)
|
||||
const targetIndex = visibleIds.indexOf(targetId)
|
||||
if (sourceIndex < 0 || targetIndex < 0 || sourceIndex === targetIndex) return
|
||||
|
||||
visibleIds.splice(sourceIndex, 1)
|
||||
visibleIds.splice(targetIndex, 0, providerId)
|
||||
const visibleSet = new Set(visibleIds)
|
||||
const allIds = [...new Set([...normalizedOrder.value, ...knownOrder.value, ...visibleIds])]
|
||||
let visibleIndex = 0
|
||||
savedOrder.value = allIds.map(currentId => visibleSet.has(currentId) ? visibleIds[visibleIndex++] ?? currentId : currentId)
|
||||
announcement.value = `${legacyT('展示顺序已更新')}: ${orderedProviders.value[targetIndex]?.name} (${targetIndex + 1}/${visibleIds.length})`
|
||||
}
|
||||
|
||||
function updateDropTarget() {
|
||||
const target = document.elementFromPoint(pointerPosition.value.clientX, pointerPosition.value.clientY)
|
||||
?.closest<HTMLElement>('[data-provider-sort-id]')
|
||||
const targetId = target?.dataset.providerSortId
|
||||
dropTargetId.value = target && container.value?.contains(target)
|
||||
&& targetId !== draggingProviderId.value
|
||||
&& orderedProviders.value.some(provider => provider.id === targetId)
|
||||
? targetId ?? null
|
||||
: null
|
||||
}
|
||||
|
||||
function findScrollContainer(handle: HTMLElement): HTMLElement | null {
|
||||
let ancestor = handle.parentElement
|
||||
while (ancestor && ancestor !== document.body) {
|
||||
if (/(auto|scroll)/.test(getComputedStyle(ancestor).overflowY) && ancestor.scrollHeight > ancestor.clientHeight) {
|
||||
return ancestor
|
||||
}
|
||||
ancestor = ancestor.parentElement
|
||||
}
|
||||
return document.scrollingElement as HTMLElement | null
|
||||
}
|
||||
|
||||
const { pause, resume } = useRafFn(() => {
|
||||
const scrollContainer = pointerDrag?.scrollContainer
|
||||
if (scrollContainer) {
|
||||
const bounds = scrollContainer.getBoundingClientRect()
|
||||
const isDocument = scrollContainer === document.scrollingElement
|
||||
const top = isDocument ? 0 : Math.max(0, bounds.top)
|
||||
const bottom = isDocument ? window.innerHeight : Math.min(window.innerHeight, bounds.bottom)
|
||||
const pointerY = pointerPosition.value.clientY
|
||||
const pointerX = pointerPosition.value.clientX
|
||||
if (isDocument || (pointerX >= bounds.left && pointerX <= bounds.right)) {
|
||||
if (pointerY < top + 48) {
|
||||
scrollContainer.scrollTop -= Math.min(12, (top + 48 - pointerY) / 4)
|
||||
} else if (pointerY > bottom - 48) {
|
||||
scrollContainer.scrollTop += Math.min(12, (pointerY - bottom + 48) / 4)
|
||||
}
|
||||
}
|
||||
}
|
||||
updateDropTarget()
|
||||
}, { immediate: false })
|
||||
|
||||
function cancelDrag() {
|
||||
const previous = pointerDrag
|
||||
pointerDrag = null
|
||||
pause()
|
||||
if (draggingProviderId.value) suppressClickUntil = Date.now() + 250
|
||||
draggingProviderId.value = null
|
||||
dropTargetId.value = null
|
||||
if (previous?.handle.hasPointerCapture?.(previous.pointerId)) {
|
||||
previous.handle.releasePointerCapture(previous.pointerId)
|
||||
}
|
||||
}
|
||||
|
||||
function startDrag(providerId: string, event: PointerEvent) {
|
||||
if (event.button !== 0 || event.isPrimary === false || orderedProviders.value.length < 2) return
|
||||
if (!orderedProviders.value.some(provider => provider.id === providerId)) return
|
||||
const handle = event.currentTarget
|
||||
if (!(handle instanceof HTMLElement)) return
|
||||
|
||||
cancelDrag()
|
||||
event.preventDefault()
|
||||
handle.focus({ preventScroll: true })
|
||||
pointerDrag = {
|
||||
providerId,
|
||||
pointerId: event.pointerId,
|
||||
startX: event.clientX,
|
||||
startY: event.clientY,
|
||||
handle,
|
||||
scrollContainer: findScrollContainer(handle),
|
||||
}
|
||||
pointerPosition.value = { clientX: event.clientX, clientY: event.clientY }
|
||||
handle.setPointerCapture?.(event.pointerId)
|
||||
}
|
||||
|
||||
function handlePointerMove(event: PointerEvent) {
|
||||
if (!pointerDrag || pointerDrag.pointerId !== event.pointerId) return
|
||||
pointerPosition.value = { clientX: event.clientX, clientY: event.clientY }
|
||||
if (!draggingProviderId.value) {
|
||||
if (Math.hypot(event.clientX - pointerDrag.startX, event.clientY - pointerDrag.startY) < 5) return
|
||||
draggingProviderId.value = pointerDrag.providerId
|
||||
resume()
|
||||
}
|
||||
event.preventDefault()
|
||||
updateDropTarget()
|
||||
}
|
||||
|
||||
function handlePointerUp(event: PointerEvent) {
|
||||
if (!pointerDrag || pointerDrag.pointerId !== event.pointerId) return
|
||||
const providerId = draggingProviderId.value
|
||||
if (providerId) {
|
||||
pointerPosition.value = { clientX: event.clientX, clientY: event.clientY }
|
||||
updateDropTarget()
|
||||
}
|
||||
const targetId = dropTargetId.value
|
||||
cancelDrag()
|
||||
if (providerId && targetId) moveProvider(providerId, targetId)
|
||||
}
|
||||
|
||||
function handleSortKeydown(providerId: string, event: KeyboardEvent) {
|
||||
if (event.key === 'Escape') {
|
||||
cancelDrag()
|
||||
return
|
||||
}
|
||||
const directions: Record<string, number> = { ArrowUp: -1, ArrowLeft: -1, ArrowDown: 1, ArrowRight: 1 }
|
||||
const direction = directions[event.key]
|
||||
if (direction === undefined || pointerDrag) return
|
||||
event.preventDefault()
|
||||
const index = orderedProviders.value.findIndex(provider => provider.id === providerId)
|
||||
const target = orderedProviders.value[index + direction]
|
||||
if (index < 0 || !target) return
|
||||
const handle = event.currentTarget as HTMLElement
|
||||
moveProvider(providerId, target.id)
|
||||
void nextTick(() => handle.focus({ preventScroll: true }))
|
||||
}
|
||||
|
||||
function handleSortClick(event: MouseEvent) {
|
||||
if (!draggingProviderId.value && Date.now() >= suppressClickUntil) return
|
||||
suppressClickUntil = 0
|
||||
event.preventDefault()
|
||||
event.stopPropagation()
|
||||
}
|
||||
|
||||
function sortItemClass(providerId: string) {
|
||||
return {
|
||||
'opacity-40': draggingProviderId.value === providerId,
|
||||
'ring-2 ring-inset ring-primary/60 bg-primary/5': dropTargetId.value === providerId,
|
||||
}
|
||||
}
|
||||
|
||||
watch(() => providers().map(provider => provider.id), (providerIds) => {
|
||||
knownOrder.value = [...new Set([...knownOrder.value, ...providerIds])]
|
||||
cancelDrag()
|
||||
}, { immediate: true })
|
||||
useEventListener(window, 'pointermove', handlePointerMove, { passive: false })
|
||||
useEventListener(window, 'pointerup', handlePointerUp)
|
||||
useEventListener(window, 'pointercancel', (event) => {
|
||||
if (pointerDrag?.pointerId === event.pointerId) cancelDrag()
|
||||
})
|
||||
useEventListener(window, 'lostpointercapture', (event) => {
|
||||
if (pointerDrag?.pointerId === event.pointerId) cancelDrag()
|
||||
})
|
||||
useEventListener(window, 'blur', cancelDrag)
|
||||
useEventListener(window, 'keydown', (event) => {
|
||||
if (event.key === 'Escape') cancelDrag()
|
||||
})
|
||||
onScopeDispose(cancelDrag)
|
||||
|
||||
return {
|
||||
orderedProviders,
|
||||
draggingProvider,
|
||||
dragPreviewStyle,
|
||||
announcement,
|
||||
startDrag,
|
||||
cancelDrag,
|
||||
handleSortKeydown,
|
||||
handleSortClick,
|
||||
sortItemClass,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
import { afterEach, describe, expect, it, vi } from 'vitest'
|
||||
import type { RefreshQuotaResult } from '@/api/endpoints/keys'
|
||||
import { refreshQuotaInBackground } from '../refreshQuotaInBackground'
|
||||
|
||||
function emptyQuota(status: 'error' | 'no_metadata' | 'forbidden' = 'no_metadata'): RefreshQuotaResult {
|
||||
return {
|
||||
success: 0,
|
||||
failed: 1,
|
||||
total: 1,
|
||||
results: [{ key_id: 'key-1', key_name: 'account', status }],
|
||||
}
|
||||
}
|
||||
|
||||
const readyQuota: RefreshQuotaResult = {
|
||||
success: 1,
|
||||
failed: 0,
|
||||
total: 1,
|
||||
results: [{ key_id: 'key-1', key_name: 'account', status: 'success' }],
|
||||
}
|
||||
|
||||
afterEach(() => vi.useRealTimers())
|
||||
|
||||
describe('initial background quota refresh', () => {
|
||||
it('waits and retries an initial empty response once', async () => {
|
||||
vi.useFakeTimers()
|
||||
const refresh = vi.fn().mockResolvedValueOnce(emptyQuota()).mockResolvedValueOnce(readyQuota)
|
||||
const pending = refreshQuotaInBackground({ refresh, isCurrent: () => true, retryInitialEmptyQuota: true })
|
||||
await vi.advanceTimersByTimeAsync(999)
|
||||
expect(refresh).toHaveBeenCalledTimes(1)
|
||||
await vi.advanceTimersByTimeAsync(1)
|
||||
expect(await pending).toBe(readyQuota)
|
||||
expect(refresh).toHaveBeenCalledTimes(2)
|
||||
})
|
||||
|
||||
it('returns persistent failures after at most one retry', async () => {
|
||||
vi.useFakeTimers()
|
||||
const result = emptyQuota('error')
|
||||
const refresh = vi.fn().mockResolvedValue(result)
|
||||
const pending = refreshQuotaInBackground({ refresh, isCurrent: () => true, retryInitialEmptyQuota: true })
|
||||
await vi.runAllTimersAsync()
|
||||
expect(await pending).toBe(result)
|
||||
expect(refresh).toHaveBeenCalledTimes(2)
|
||||
})
|
||||
|
||||
it.each([readyQuota, emptyQuota('forbidden'), {
|
||||
...emptyQuota('error'),
|
||||
results: [{ ...emptyQuota('error').results[0]!, status_code: 401 }],
|
||||
}])('does not retry successful or rejected authorization results', async (result) => {
|
||||
const refresh = vi.fn().mockResolvedValue(result)
|
||||
expect(await refreshQuotaInBackground({ refresh, isCurrent: () => true, retryInitialEmptyQuota: true })).toBe(result)
|
||||
expect(refresh).toHaveBeenCalledTimes(1)
|
||||
})
|
||||
|
||||
it('does not retry when initial quota retry is disabled', async () => {
|
||||
const refresh = vi.fn().mockResolvedValue(emptyQuota())
|
||||
await refreshQuotaInBackground({ refresh, isCurrent: () => true, retryInitialEmptyQuota: false })
|
||||
expect(refresh).toHaveBeenCalledTimes(1)
|
||||
})
|
||||
|
||||
it('cancels the retry when the drawer closes or switches providers', async () => {
|
||||
vi.useFakeTimers()
|
||||
let current = true
|
||||
const refresh = vi.fn().mockResolvedValue(emptyQuota())
|
||||
const pending = refreshQuotaInBackground({ refresh, isCurrent: () => current, retryInitialEmptyQuota: true })
|
||||
await vi.advanceTimersByTimeAsync(1)
|
||||
current = false
|
||||
await vi.runAllTimersAsync()
|
||||
expect(await pending).toBeNull()
|
||||
expect(refresh).toHaveBeenCalledTimes(1)
|
||||
})
|
||||
|
||||
it('discards an in-flight response from a closed drawer', async () => {
|
||||
let current = true
|
||||
const refresh = vi.fn(async () => {
|
||||
current = false
|
||||
return readyQuota
|
||||
})
|
||||
expect(await refreshQuotaInBackground({ refresh, isCurrent: () => current, retryInitialEmptyQuota: true })).toBeNull()
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,35 @@
|
||||
import type { RefreshQuotaResult } from '@/api/endpoints/keys'
|
||||
|
||||
interface BackgroundQuotaRefreshOptions {
|
||||
refresh: () => Promise<RefreshQuotaResult>
|
||||
isCurrent: () => boolean
|
||||
retryInitialEmptyQuota: boolean
|
||||
}
|
||||
|
||||
function shouldRetryEmptyQuota(result: RefreshQuotaResult): boolean {
|
||||
return result.success === 0
|
||||
&& result.results.length > 0
|
||||
&& result.results.every(item =>
|
||||
(item.status === 'no_metadata' || item.status === 'error')
|
||||
&& item.status_code !== 401
|
||||
&& item.status_code !== 403
|
||||
&& !item.metadata
|
||||
&& !item.quota_snapshot,
|
||||
)
|
||||
}
|
||||
|
||||
export async function refreshQuotaInBackground({
|
||||
refresh,
|
||||
isCurrent,
|
||||
retryInitialEmptyQuota,
|
||||
}: BackgroundQuotaRefreshOptions): Promise<RefreshQuotaResult | null> {
|
||||
if (!isCurrent()) return null
|
||||
const result = await refresh()
|
||||
if (!isCurrent()) return null
|
||||
if (!retryInitialEmptyQuota || !shouldRetryEmptyQuota(result)) return result
|
||||
|
||||
await new Promise<void>(resolve => setTimeout(resolve, 1000))
|
||||
if (!isCurrent()) return null
|
||||
const retried = await refresh()
|
||||
return isCurrent() ? retried : null
|
||||
}
|
||||
@@ -477,6 +477,18 @@
|
||||
HTTP {{ currentAttemptRequestError.statusCode }}
|
||||
</span>
|
||||
</div>
|
||||
<div
|
||||
v-if="currentAttemptRequestError.upstreamStatusCode != null && currentAttemptRequestError.upstreamStatusCode !== currentAttemptRequestError.statusCode"
|
||||
class="error-upstream-status text-xs text-muted-foreground mb-1"
|
||||
>
|
||||
上游响应状态:HTTP {{ currentAttemptRequestError.upstreamStatusCode }}(本次尝试状态见上方)
|
||||
</div>
|
||||
<div
|
||||
v-if="currentAttemptRequestError.finalStatusCode != null"
|
||||
class="error-final-status text-xs text-muted-foreground mb-1"
|
||||
>
|
||||
请求最终状态:HTTP {{ currentAttemptRequestError.finalStatusCode }}(与本次尝试状态不同)
|
||||
</div>
|
||||
<div
|
||||
v-if="currentAttemptRequestError.message"
|
||||
class="error-msg"
|
||||
@@ -544,6 +556,7 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
import { ref, watch, computed, onBeforeUnmount } from 'vue'
|
||||
import { isAxiosError } from 'axios'
|
||||
import Card from '@/components/ui/card.vue'
|
||||
@@ -1629,21 +1642,19 @@ const shouldShowAttemptMessageWithUpstreamResponse = (
|
||||
rawMessage: string,
|
||||
upstreamResponse: Record<string, unknown> | null,
|
||||
): boolean => {
|
||||
if (!upstreamResponse) return true
|
||||
if (!upstreamResponse || !hasRenderableValue(upstreamResponse.body)) return true
|
||||
const normalized = rawMessage.trim()
|
||||
if (!normalized) return false
|
||||
if (isLocalSyncFinalizeDiagnostic(normalized)) return true
|
||||
if (isConversionDiagnosticMessage(normalized)) return true
|
||||
if (isGenericExecutionRuntimeStatusMessage(normalized)) return false
|
||||
|
||||
const hasBody = hasRenderableValue(upstreamResponse.body)
|
||||
const bodyState = (readStringField(upstreamResponse, 'body_state') ?? '').toLowerCase()
|
||||
return !hasBody && bodyState === 'disabled'
|
||||
return false
|
||||
}
|
||||
|
||||
const currentAttemptRequestError = computed<{
|
||||
message: string
|
||||
statusCode?: number
|
||||
upstreamStatusCode?: number
|
||||
finalStatusCode?: number
|
||||
upstreamResponse: Record<string, unknown> | null
|
||||
diagnostic: Record<string, unknown> | null
|
||||
} | null>(() => {
|
||||
@@ -1653,11 +1664,16 @@ const currentAttemptRequestError = computed<{
|
||||
const extra = extractObject(attempt.extra_data)
|
||||
const upstreamResponse = extractObject(extra?.upstream_response)
|
||||
const errorFlow = extractObject(extra?.error_flow)
|
||||
const statusCode = readNumberField(upstreamResponse ?? {}, 'status_code')
|
||||
const upstreamStatusCode = readNumberField(upstreamResponse ?? {}, 'status_code')
|
||||
?? readNumberField(upstreamResponse ?? {}, 'statusCode')
|
||||
const statusCode = attempt.status_code
|
||||
?? upstreamStatusCode
|
||||
?? readNumberField(errorFlow ?? {}, 'status_code')
|
||||
?? readNumberField(errorFlow ?? {}, 'statusCode')
|
||||
?? attempt.status_code
|
||||
const finalStatusCode = !['pending', 'streaming'].includes(computedFinalStatus.value)
|
||||
&& props.overrideStatusCode != null && props.overrideStatusCode !== statusCode
|
||||
? props.overrideStatusCode
|
||||
: undefined
|
||||
const flowMessage = errorFlow
|
||||
? readStringField(errorFlow, 'message')
|
||||
: ''
|
||||
@@ -1670,6 +1686,7 @@ const currentAttemptRequestError = computed<{
|
||||
const diagnosticMessage = extractVisibleDiagnosticMessage(extra)
|
||||
const rawMessage = chooseAttemptRawErrorMessage(flowMessage || '', fallbackMessage, diagnosticMessage)
|
||||
const message = formatAttemptErrorMessage(rawMessage, statusCode) || fallbackType
|
||||
|| '本次尝试失败,链路追踪未包含详细错误内容。'
|
||||
const upstreamResponseDisplay = normalizeUpstreamResponseDisplay(extra?.upstream_response)
|
||||
const visibleDiagnosticObjects = extractVisibleDiagnosticObjects(extra)
|
||||
const shouldAttachDiagnostic = Boolean(
|
||||
@@ -1694,20 +1711,16 @@ const currentAttemptRequestError = computed<{
|
||||
const response = Object.keys(upstreamResponseData).length > 0
|
||||
? upstreamResponseData
|
||||
: null
|
||||
if (
|
||||
!message
|
||||
&& statusCode == null
|
||||
&& !response
|
||||
&& !diagnostic
|
||||
) return null
|
||||
const showMessage = shouldShowAttemptMessageWithUpstreamResponse(
|
||||
rawMessage || fallbackType,
|
||||
upstreamResponseDisplay,
|
||||
)
|
||||
|
||||
return {
|
||||
message: showMessage ? (message || '未知错误') : '',
|
||||
message: showMessage ? message : '',
|
||||
statusCode,
|
||||
upstreamStatusCode,
|
||||
finalStatusCode,
|
||||
upstreamResponse: response,
|
||||
diagnostic,
|
||||
}
|
||||
@@ -2294,7 +2307,7 @@ const resolveAttemptTimeRange = (attempt: CandidateRecord | null | undefined): A
|
||||
// 格式化时间(详细)
|
||||
const formatTime = (dateStr: string) => {
|
||||
const date = new Date(dateStr)
|
||||
const timeStr = date.toLocaleTimeString('zh-CN', {
|
||||
const timeStr = date.toLocaleTimeString(getI18nLocale(), {
|
||||
hour12: false,
|
||||
hour: '2-digit',
|
||||
minute: '2-digit',
|
||||
|
||||
@@ -873,6 +873,7 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
import { ref, watch, computed, onMounted, onBeforeUnmount } from 'vue'
|
||||
import Button from '@/components/ui/button.vue'
|
||||
import { useEscapeKey } from '@/composables/useEscapeKey'
|
||||
@@ -2848,7 +2849,7 @@ onBeforeUnmount(() => {
|
||||
function formatDateTime(dateStr: string | null | undefined): string {
|
||||
if (!dateStr) return 'N/A'
|
||||
const date = new Date(dateStr)
|
||||
return date.toLocaleString('zh-CN', {
|
||||
return date.toLocaleString(getI18nLocale(), {
|
||||
year: 'numeric',
|
||||
month: '2-digit',
|
||||
day: '2-digit',
|
||||
|
||||
@@ -659,6 +659,98 @@ describe('HorizontalRequestTimeline', () => {
|
||||
expect(root.textContent).not.toContain('该错误被标记为敏感上游错误')
|
||||
})
|
||||
|
||||
it.each(['inline', 'reference', 'disabled', 'unavailable', 'none', undefined])(
|
||||
'shows a fallback for redacted errors with body state %s',
|
||||
async (bodyState) => {
|
||||
const trace = buildTrace([
|
||||
buildCandidate({
|
||||
status_code: 400,
|
||||
extra_data: {
|
||||
upstream_response: {
|
||||
status_code: 400,
|
||||
body_state: bodyState,
|
||||
},
|
||||
error_flow: {
|
||||
source: 'upstream_response',
|
||||
status_code: 400,
|
||||
decision: 'retry_next_candidate',
|
||||
},
|
||||
},
|
||||
}),
|
||||
])
|
||||
trace.final_status = 'failed'
|
||||
|
||||
const root = mountTimeline(trace, {
|
||||
overrideStatusCode: 503,
|
||||
requestStatus: 'failed',
|
||||
})
|
||||
await nextTick()
|
||||
|
||||
expect(root.querySelector('.status-tag')?.textContent?.trim()).toBe('400')
|
||||
expect(root.querySelector('.error-status-badge')?.textContent?.trim()).toBe('HTTP 400')
|
||||
expect(root.querySelector('.error-msg')?.textContent).toContain('链路追踪未包含详细错误内容')
|
||||
expect(root.querySelector('.error-final-status')?.textContent).toContain('请求最终状态:HTTP 503')
|
||||
expect(root.querySelector('.error-upstream-response-json')).toBeNull()
|
||||
},
|
||||
)
|
||||
|
||||
it('keeps the attempt status separate from the upstream transport status', async () => {
|
||||
const root = mountTimeline(buildTrace([
|
||||
buildCandidate({
|
||||
status_code: 502,
|
||||
error_type: 'stream_error',
|
||||
extra_data: {
|
||||
upstream_response: { status_code: 200, body_state: 'disabled' },
|
||||
},
|
||||
}),
|
||||
]))
|
||||
await nextTick()
|
||||
|
||||
expect(root.querySelector('.status-tag')?.textContent?.trim()).toBe('502')
|
||||
expect(root.querySelector('.error-status-badge')?.textContent?.trim()).toBe('HTTP 502')
|
||||
expect(root.querySelector('.error-upstream-status')?.textContent).toContain('上游响应状态:HTTP 200')
|
||||
expect(root.querySelector('.error-final-status')).toBeNull()
|
||||
})
|
||||
|
||||
it('does not label an active request status as final', async () => {
|
||||
const root = mountTimeline(buildTrace([
|
||||
buildCandidate({ status_code: 400 }),
|
||||
]), {
|
||||
overrideStatusCode: 200,
|
||||
requestStatus: 'streaming',
|
||||
})
|
||||
await nextTick()
|
||||
|
||||
expect(root.querySelector('.error-final-status')).toBeNull()
|
||||
})
|
||||
|
||||
it('shows failure details even when no status or error message was retained', async () => {
|
||||
const root = mountTimeline(buildTrace([buildCandidate()]))
|
||||
await nextTick()
|
||||
|
||||
expect(root.querySelector('.error-msg')?.textContent).toContain('链路追踪未包含详细错误内容')
|
||||
})
|
||||
|
||||
it('keeps a generic error visible when only upstream headers are available', async () => {
|
||||
const root = mountTimeline(buildTrace([
|
||||
buildCandidate({
|
||||
status_code: 400,
|
||||
error_message: 'execution runtime stream returned non-success status 400',
|
||||
extra_data: {
|
||||
upstream_response: {
|
||||
status_code: 400,
|
||||
headers: { 'content-type': 'application/json' },
|
||||
body_state: 'reference',
|
||||
},
|
||||
},
|
||||
}),
|
||||
]))
|
||||
await nextTick()
|
||||
|
||||
expect(root.querySelector('.error-msg')?.textContent).toContain('上游返回非成功状态 400')
|
||||
expect(root.querySelector('.error-upstream-response-json')?.textContent).toContain('application/json')
|
||||
})
|
||||
|
||||
it('keeps local sync diagnostics visible when upstream response body capture is disabled', async () => {
|
||||
const trace = buildTrace([
|
||||
buildCandidate({
|
||||
|
||||
@@ -509,6 +509,7 @@
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { getI18nLocale } from '@/i18n'
|
||||
import { computed, ref, watch } from 'vue'
|
||||
import {
|
||||
Badge,
|
||||
@@ -1027,7 +1028,7 @@ async function submitCompleteRefund() {
|
||||
|
||||
function formatDateTime(value: string | null | undefined) {
|
||||
if (!value) return '-'
|
||||
return new Date(value).toLocaleString('zh-CN', {
|
||||
return new Date(value).toLocaleString(getI18nLocale(), {
|
||||
year: 'numeric',
|
||||
month: '2-digit',
|
||||
day: '2-digit',
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import { beforeEach, describe, expect, it } from 'vitest'
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import { createApp, defineComponent, h, nextTick, ref } from 'vue'
|
||||
|
||||
import { createI18n, useI18n, useLocaleOptions } from '@/i18n'
|
||||
import { createI18n, getI18nLocale, normalizeLocale, setI18nLocale, useI18n, useLocaleOptions } from '@/i18n'
|
||||
import { formatDate, formatRelativeTime } from '@/utils/format'
|
||||
import { translateLegacyText } from '@/i18n/messages'
|
||||
import { transformLegacyTemplateI18n } from '@/i18n/legacy-template-transform'
|
||||
|
||||
@@ -118,4 +119,39 @@ describe('i18n infrastructure', () => {
|
||||
expect(translateLegacyText(' 发布于 2026-01-01 ', 'en-US')).toBe(' Published at 2026-01-01 ')
|
||||
expect(translateLegacyText('git clone https://github.com/fawney19/Aether.git', 'en-US')).toBe('git clone https://github.com/fawney19/Aether.git')
|
||||
})
|
||||
|
||||
it('normalizes saved language aliases without accepting unrelated language names', () => {
|
||||
expect(normalizeLocale('en')).toBe('en-US')
|
||||
expect(normalizeLocale('en_GB')).toBe('en-US')
|
||||
expect(normalizeLocale(' ZH-cn ')).toBe('zh-CN')
|
||||
expect(normalizeLocale('english')).toBeUndefined()
|
||||
expect(normalizeLocale('fr-FR')).toBeUndefined()
|
||||
})
|
||||
|
||||
it('continues switching language when browser storage is unavailable', () => {
|
||||
const write = vi.spyOn(localStorage, 'setItem').mockImplementation(() => {
|
||||
throw new DOMException('Storage blocked', 'SecurityError')
|
||||
})
|
||||
try {
|
||||
expect(() => setI18nLocale('en-US')).not.toThrow()
|
||||
expect(getI18nLocale()).toBe('en-US')
|
||||
expect(document.documentElement.lang).toBe('en-US')
|
||||
} finally {
|
||||
write.mockRestore()
|
||||
}
|
||||
})
|
||||
|
||||
it('updates date and relative-time formatting with the selected language', () => {
|
||||
const date = '2026-09-07T12:30:00'
|
||||
setI18nLocale('zh-CN')
|
||||
const chineseDate = formatDate(date)
|
||||
expect(formatRelativeTime(-1, 'day')).toBe('昨天')
|
||||
setI18nLocale('en-US')
|
||||
expect(formatRelativeTime(-1, 'day')).toBe('yesterday')
|
||||
expect(formatRelativeTime(-1, 'minute')).toBe('1 minute ago')
|
||||
expect(formatRelativeTime(-2, 'minute')).toBe('2 minutes ago')
|
||||
expect(formatDate(date)).not.toBe(chineseDate)
|
||||
setI18nLocale('zh-CN')
|
||||
expect(formatDate(date)).toBe(chineseDate)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import * as Vue from 'vue'
|
||||
import { compile, nextTick, ref } from 'vue'
|
||||
import { compileScript, compileTemplate, parse } from 'vue/compiler-sfc'
|
||||
|
||||
import { installLegacyDomTranslator } from '@/i18n/dom-translator'
|
||||
import { transformLegacyTemplateI18n, transformVueSource } from '@/i18n/legacy-template-transform'
|
||||
import { translateLegacyText, type Locale } from '@/i18n/messages'
|
||||
|
||||
describe('legacy template compiler', () => {
|
||||
it('translates content following nested slot templates and accepts reordered setup attributes', () => {
|
||||
const source = `<template><Panel><template #default>保存</template></Panel><footer>取消</footer></template>
|
||||
<script lang="ts" setup>const name = 'value'</script>`
|
||||
const result = transformVueSource(source)
|
||||
const { descriptor, errors } = parse(result.code)
|
||||
|
||||
expect(errors).toEqual([])
|
||||
expect(descriptor.template?.content).toContain('{{ __aetherLegacyT("取消") }}')
|
||||
expect(result.code.match(/<script/g)).toHaveLength(1)
|
||||
expect(() => compileScript(descriptor, { id: 'nested-template' })).not.toThrow()
|
||||
expect(transformVueSource(result.code).changed).toBe(false)
|
||||
})
|
||||
|
||||
it('preserves comparisons and user values while translating displayed branches', () => {
|
||||
const source = `<span>{{ count < limit ? '保存' : user.name }}</span>
|
||||
<span>{{ status === '保存' }}</span><span>{{ names['保存'] }}</span>
|
||||
<span>{{ format('保存') }}</span><span>{{ legacyT('保存') }}</span>
|
||||
<span :title="count > limit ? '关闭' : user.name">{{ user.name }}</span>`
|
||||
const result = transformLegacyTemplateI18n(source)
|
||||
|
||||
expect(result.code).toContain(`count < limit ? __aetherLegacyT("保存") : user.name`)
|
||||
expect(result.code).toContain(`{{ status === '保存' }}`)
|
||||
expect(result.code).toContain(`{{ names['保存'] }}`)
|
||||
expect(result.code).toContain(`{{ format('保存') }}`)
|
||||
expect(result.code).toContain(`{{ legacyT('保存') }}`)
|
||||
expect(result.code).toContain(`{{ user.name }}`)
|
||||
expect(compileTemplate({ source: result.code, id: 'comparisons', filename: 'comparisons.vue' }).errors).toEqual([])
|
||||
})
|
||||
|
||||
it('matches an existing plain script language when adding the setup helper', () => {
|
||||
const result = transformVueSource('<template><span>保存</span></template><script>export default { name: "Legacy" }</script>')
|
||||
const { descriptor } = parse(result.code)
|
||||
expect(descriptor.scriptSetup?.lang).toBe(descriptor.script?.lang)
|
||||
expect(() => compileScript(descriptor, { id: 'plain-script' })).not.toThrow()
|
||||
})
|
||||
|
||||
it('keeps quotes and HTML entities intact in static and bound attributes', () => {
|
||||
const result = transformLegacyTemplateI18n(`<input title="保存 "O'Reilly" & <x>" placeholder=保存 :aria-label="open ? '关闭 & 保存' : '保存'">`)
|
||||
const render = compile(result.code) as (context: Record<string, unknown>, cache: unknown[]) => Vue.VNode
|
||||
const vnode = render({ open: true, __aetherLegacyT: (value: string) => value }, [])
|
||||
|
||||
expect(vnode.props?.title).toBe(`保存 "O'Reilly" & <x>`)
|
||||
expect(vnode.props?.placeholder).toBe('保存')
|
||||
expect(vnode.props?.['aria-label']).toBe('关闭 & 保存')
|
||||
})
|
||||
|
||||
it('translates template literal text without translating interpolated values', () => {
|
||||
const result = transformLegacyTemplateI18n('<span>{{ `保存 ${user.name}` }}</span>')
|
||||
expect(result.code).toContain('`${__aetherLegacyT("保存 ")}${user.name}`')
|
||||
const render = compile(result.code) as (context: Record<string, unknown>, cache: unknown[]) => Vue.VNode
|
||||
const vnode = render({ user: { name: '取消' }, __aetherLegacyT: (value: string) => value === '保存 ' ? 'Save ' : 'unexpected' }, [])
|
||||
expect(vnode.children).toBe('Save 取消')
|
||||
})
|
||||
|
||||
it.each(['HelpHint', 'help-hint'])('translates static and displayed literal text props on %s', tag => {
|
||||
const context = {
|
||||
expanded: true,
|
||||
record: { text: '保存' },
|
||||
__aetherLegacyT: (value: string) => translateLegacyText(value, 'en-US'),
|
||||
}
|
||||
const staticResult = transformLegacyTemplateI18n(`<${tag} text="保存" />`)
|
||||
const renderStatic = compile(staticResult.code, { isCustomElement: name => name === tag }) as (context: Record<string, unknown>, cache: unknown[]) => Vue.VNode
|
||||
expect(renderStatic(context, []).props?.text).toBe('Save')
|
||||
|
||||
const dynamicResult = transformLegacyTemplateI18n(`<${tag} :text="expanded ? '关闭' : record.text" />`)
|
||||
const renderDynamic = compile(dynamicResult.code, { isCustomElement: name => name === tag }) as (context: Record<string, unknown>, cache: unknown[]) => Vue.VNode
|
||||
expect(renderDynamic(context, []).props?.text).toBe('Close')
|
||||
expect(renderDynamic({ ...context, expanded: false }, []).props?.text).toBe('保存')
|
||||
})
|
||||
|
||||
it('keeps ordinary text props and opted-out HelpHint text as application data', () => {
|
||||
const source = `<Message text="保存" /><Message :text="active ? '关闭' : record.text" />
|
||||
<div text="保存" /><HelpHint translate="no" text="保存" />
|
||||
<HelpHint data-i18n-skip :text="active ? '关闭' : record.text" />`
|
||||
expect(transformLegacyTemplateI18n(source).code).toBe(source)
|
||||
})
|
||||
|
||||
it('respects skipped subtrees including v-pre and contenteditable', () => {
|
||||
const source = `<section translate="no"><span title="关闭">保存</span></section>
|
||||
<section data-i18n-skip><span>保存</span></section>
|
||||
<section v-pre title="保存 > 标题"><span>{{ '保存' }}</span></section>
|
||||
<section contenteditable="true">保存</section><pre>保存</pre><code>保存</code>`
|
||||
const result = transformLegacyTemplateI18n(source)
|
||||
expect(result.code).toBe(source.replace('<section v-pre', '<section data-i18n-skip v-pre'))
|
||||
expect(result.needsHelper).toBe(false)
|
||||
expect(transformLegacyTemplateI18n(result.code).changed).toBe(false)
|
||||
expect(transformLegacyTemplateI18n('<span title="v-pre">保存</span>').changed).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
describe('legacy DOM translation updates', () => {
|
||||
let stop: () => void
|
||||
let locale: Vue.Ref<Locale>
|
||||
let root: HTMLDivElement
|
||||
|
||||
async function flushTranslation(): Promise<void> {
|
||||
await nextTick()
|
||||
await Promise.resolve()
|
||||
await vi.advanceTimersByTimeAsync(40)
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
vi.useFakeTimers()
|
||||
root = document.createElement('div')
|
||||
document.body.append(root)
|
||||
locale = ref<Locale>('en-US')
|
||||
stop = installLegacyDomTranslator(locale)
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
stop()
|
||||
root.remove()
|
||||
vi.useRealTimers()
|
||||
})
|
||||
|
||||
it('translates new state on a reused text node and restores the latest source', async () => {
|
||||
const text = document.createTextNode('保存')
|
||||
root.append(text)
|
||||
await flushTranslation()
|
||||
expect(text.nodeValue).toBe('Save')
|
||||
|
||||
text.nodeValue = '保存中...'
|
||||
await flushTranslation()
|
||||
expect(text.nodeValue).toBe('Saving...')
|
||||
|
||||
locale.value = 'zh-CN'
|
||||
await flushTranslation()
|
||||
expect(text.nodeValue).toBe('保存中...')
|
||||
|
||||
text.nodeValue = '取消'
|
||||
locale.value = 'en-US'
|
||||
await flushTranslation()
|
||||
expect(text.nodeValue).toBe('Cancel')
|
||||
locale.value = 'zh-CN'
|
||||
await flushTranslation()
|
||||
expect(text.nodeValue).toBe('取消')
|
||||
})
|
||||
|
||||
it('updates attributes across changes, removal, and language round trips', async () => {
|
||||
root.title = '关闭'
|
||||
await flushTranslation()
|
||||
expect(root.title).toBe('Close')
|
||||
root.title = '保存'
|
||||
await flushTranslation()
|
||||
expect(root.title).toBe('Save')
|
||||
locale.value = 'zh-CN'
|
||||
await flushTranslation()
|
||||
expect(root.title).toBe('保存')
|
||||
|
||||
root.removeAttribute('title')
|
||||
await flushTranslation()
|
||||
root.title = '取消'
|
||||
locale.value = 'en-US'
|
||||
await flushTranslation()
|
||||
expect(root.title).toBe('Cancel')
|
||||
root.title = 'User supplied English'
|
||||
locale.value = 'zh-CN'
|
||||
await flushTranslation()
|
||||
expect(root.title).toBe('User supplied English')
|
||||
})
|
||||
|
||||
it('preserves code, editable values, and explicit untranslated subtrees', async () => {
|
||||
root.innerHTML = `<code title="关闭">保存</code><pre>保存</pre><textarea>保存</textarea>
|
||||
<span contenteditable="true">保存</span><span v-pre>保存</span>
|
||||
<section translate="no"><span title="关闭">保存</span></section>
|
||||
<section data-i18n-skip><span title="关闭">保存</span></section>`
|
||||
const original = root.innerHTML
|
||||
await flushTranslation()
|
||||
expect(root.innerHTML).toBe(original)
|
||||
locale.value = 'zh-CN'
|
||||
await flushTranslation()
|
||||
expect(root.innerHTML).toBe(original)
|
||||
})
|
||||
|
||||
it('restores original text when a translated region becomes opted out', async () => {
|
||||
root.textContent = '保存'
|
||||
root.title = '关闭'
|
||||
await flushTranslation()
|
||||
expect(root.textContent).toBe('Save')
|
||||
root.setAttribute('translate', 'no')
|
||||
await flushTranslation()
|
||||
expect(root.textContent).toBe('保存')
|
||||
expect(root.title).toBe('关闭')
|
||||
})
|
||||
})
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user