mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
Merge origin/main into fix/gemini-cli-v1internal
This commit is contained in:
@@ -1,6 +1,9 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
use aether_scheduler_core::count_recent_rpm_requests_for_provider_key_since;
|
||||
use aether_scheduler_core::{
|
||||
count_recent_rpm_requests_for_provider_key_since,
|
||||
provider_key_circuit_payload_is_active_open_at,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
@@ -18,6 +21,10 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.await
|
||||
.ok()
|
||||
.and_then(|mut keys| keys.drain(..).next())?;
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or_default();
|
||||
let provider = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
|
||||
.await
|
||||
@@ -81,10 +88,10 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.and_then(|value| value.get("last_failure_at"))
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
payload["circuit_breaker_open"] = json!(circuit_data
|
||||
.and_then(|value| value.get("open"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false));
|
||||
payload["circuit_breaker_open"] =
|
||||
json!(circuit_data.is_some_and(
|
||||
|value| provider_key_circuit_payload_is_active_open_at(value, now_unix_secs)
|
||||
));
|
||||
payload["circuit_breaker_open_at"] = circuit_data
|
||||
.and_then(|value| value.get("open_at"))
|
||||
.cloned()
|
||||
@@ -107,15 +114,16 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.unwrap_or(0));
|
||||
} else {
|
||||
let mut formats_payload = serde_json::Map::new();
|
||||
let mut any_circuit_open = false;
|
||||
for format_name in
|
||||
provider_key_effective_api_formats(&key, &provider.provider_type, &endpoints)
|
||||
{
|
||||
let health_data = health_by_format.and_then(|formats| formats.get(&format_name));
|
||||
let circuit_data = circuit_by_format.and_then(|formats| formats.get(&format_name));
|
||||
let is_open = circuit_data
|
||||
.and_then(|value| value.get("open"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let active_open = circuit_data.is_some_and(|value| {
|
||||
provider_key_circuit_payload_is_active_open_at(value, now_unix_secs)
|
||||
});
|
||||
any_circuit_open |= active_open;
|
||||
formats_payload.insert(
|
||||
format_name.clone(),
|
||||
json!({
|
||||
@@ -134,8 +142,8 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null),
|
||||
"circuit_breaker": {
|
||||
"state": if is_open { "open" } else { "closed" },
|
||||
"open": is_open,
|
||||
"state": if active_open { "open" } else { "closed" },
|
||||
"open": active_open,
|
||||
"open_at": circuit_data
|
||||
.and_then(|value| value.get("open_at"))
|
||||
.cloned()
|
||||
@@ -167,13 +175,6 @@ pub(crate) async fn build_admin_key_health_payload(
|
||||
.filter_map(serde_json::Value::as_f64)
|
||||
.reduce(f64::min)
|
||||
.unwrap_or(1.0);
|
||||
let any_circuit_open = formats_payload.values().any(|value| {
|
||||
value
|
||||
.get("circuit_breaker")
|
||||
.and_then(|circuit| circuit.get("open"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
});
|
||||
|
||||
payload["key_health_score"] = json!(key_health_score);
|
||||
payload["any_circuit_open"] = json!(any_circuit_open);
|
||||
|
||||
@@ -4,7 +4,9 @@ use crate::handlers::public::{api_format_display_name, build_public_health_timel
|
||||
use crate::handlers::shared::unix_ms_to_rfc3339;
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
use aether_data_contracts::repository::candidates::PublicHealthTimelineBucket;
|
||||
use aether_scheduler_core::{is_provider_key_circuit_open, provider_key_health_score};
|
||||
use aether_scheduler_core::{
|
||||
any_provider_key_circuit_open_at, is_provider_key_circuit_open_at, provider_key_health_score,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -99,7 +101,9 @@ pub(crate) async fn build_admin_endpoint_health_status_payload(
|
||||
.entry(api_format.clone())
|
||||
.or_default()
|
||||
.insert(key.id.clone());
|
||||
if key.is_active && !is_provider_key_circuit_open(&key, &api_format) {
|
||||
if key.is_active
|
||||
&& !is_provider_key_circuit_open_at(&key, &api_format, now_unix_secs)
|
||||
{
|
||||
let key_health_score =
|
||||
provider_key_health_score(&key, &api_format).unwrap_or(1.0);
|
||||
active_keys_by_format
|
||||
@@ -229,6 +233,10 @@ pub(crate) async fn build_admin_health_summary_payload(
|
||||
return None;
|
||||
}
|
||||
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or_default();
|
||||
let providers = state
|
||||
.list_provider_catalog_providers(false)
|
||||
.await
|
||||
@@ -286,20 +294,7 @@ pub(crate) async fn build_admin_health_summary_payload(
|
||||
.count();
|
||||
let circuit_open_keys = keys
|
||||
.iter()
|
||||
.filter(|key| {
|
||||
key.circuit_breaker_by_format
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.map(|formats| {
|
||||
formats.values().any(|circuit| {
|
||||
circuit
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
})
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.filter(|key| any_provider_key_circuit_open_at(key, now_unix_secs))
|
||||
.count();
|
||||
|
||||
Some(json!({
|
||||
|
||||
@@ -8,10 +8,12 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_scheduler_core::{
|
||||
is_provider_key_circuit_open, matches_model_mapping, provider_key_health_score,
|
||||
is_provider_key_circuit_open_at, matches_model_mapping,
|
||||
provider_key_circuit_payload_is_active_open_at, provider_key_health_score,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
@@ -86,6 +88,10 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
.flatten()
|
||||
.and_then(|value| value.as_bool())
|
||||
.unwrap_or(false);
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or_default();
|
||||
|
||||
let global_model_mappings = global_model
|
||||
.config
|
||||
@@ -163,9 +169,11 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
entries
|
||||
.iter()
|
||||
.filter_map(|(api_format, value)| {
|
||||
value.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.filter(|is_open| *is_open)
|
||||
provider_key_circuit_payload_is_active_open_at(
|
||||
value,
|
||||
now_unix_secs,
|
||||
)
|
||||
.then_some(())
|
||||
.map(|_| api_format.clone())
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
@@ -188,7 +196,7 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
"effective_rpm": effective_rpm,
|
||||
"allowed_models": allowed_models,
|
||||
"health_score": provider_key_health_score(key, &endpoint.api_format),
|
||||
"circuit_breaker_open": is_provider_key_circuit_open(key, &endpoint.api_format),
|
||||
"circuit_breaker_open": is_provider_key_circuit_open_at(key, &endpoint.api_format, now_unix_secs),
|
||||
"circuit_breaker_formats": circuit_breaker_formats,
|
||||
"next_probe_at": next_probe_at,
|
||||
});
|
||||
|
||||
+8
-7
@@ -1,6 +1,6 @@
|
||||
use super::super::usage_helpers::admin_monitoring_usage_is_error;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::{provider_key_health_summary, unix_secs_to_rfc3339};
|
||||
use crate::handlers::admin::shared::{provider_key_health_summary_at, unix_secs_to_rfc3339};
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::{
|
||||
provider_catalog::StoredProviderCatalogKey, usage::UsageMonitoringErrorListQuery,
|
||||
@@ -99,7 +99,7 @@ pub(super) async fn build_admin_monitoring_resilience_snapshot(
|
||||
last_failure_at,
|
||||
circuit_breaker_open,
|
||||
circuit_by_format,
|
||||
) = provider_key_health_summary(key);
|
||||
) = provider_key_health_summary_at(key, now.timestamp().max(0) as u64);
|
||||
if health_score < 0.8 {
|
||||
degraded_keys += 1;
|
||||
}
|
||||
@@ -110,11 +110,12 @@ pub(super) async fn build_admin_monitoring_resilience_snapshot(
|
||||
let open_formats = circuit_by_format
|
||||
.iter()
|
||||
.filter_map(|(api_format, value)| {
|
||||
value
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.filter(|open| *open)
|
||||
.map(|_| api_format.clone())
|
||||
aether_scheduler_core::provider_key_circuit_payload_is_active_open_at(
|
||||
value,
|
||||
now.timestamp().max(0) as u64,
|
||||
)
|
||||
.then_some(())
|
||||
.map(|_| api_format.clone())
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
|
||||
@@ -1237,7 +1237,8 @@ async fn admin_monitoring_circuit_history_returns_local_payload() {
|
||||
"openai:chat": {
|
||||
"open": true,
|
||||
"open_at": "2026-03-30T12:00:00+00:00",
|
||||
"next_probe_at": "2026-03-30T12:05:00+00:00",
|
||||
"next_probe_at": "2099-03-30T12:05:00+00:00",
|
||||
"recovery_seconds": 300,
|
||||
"reason": "错误率过高"
|
||||
}
|
||||
})),
|
||||
|
||||
@@ -81,6 +81,133 @@ async fn admin_monitoring_trace_request_returns_local_payload() {
|
||||
assert_eq!(payload["candidates"][0]["status_code"], json!(502));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_resolves_usage_id_to_header_trace_id() {
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
sample_candidate(
|
||||
"cand-used",
|
||||
"trace-1",
|
||||
0,
|
||||
RequestCandidateStatus::Success,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(200),
|
||||
),
|
||||
]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
vec![sample_endpoint()],
|
||||
vec![sample_key()],
|
||||
));
|
||||
let mut usage = sample_usage(
|
||||
"usage-request-1",
|
||||
"provider-1",
|
||||
"OpenAI",
|
||||
40,
|
||||
0.02,
|
||||
"completed",
|
||||
Some(200),
|
||||
100,
|
||||
);
|
||||
usage.id = "usage-row-1".to_string();
|
||||
usage.candidate_id = Some("cand-used".to_string());
|
||||
usage.request_headers = Some(json!({
|
||||
"x-trace-id": "trace-1"
|
||||
}));
|
||||
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,
|
||||
)
|
||||
.with_provider_catalog_reader(provider_catalog);
|
||||
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/usage-row-1?attempted_only=true",
|
||||
);
|
||||
|
||||
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");
|
||||
assert_eq!(payload["request_id"], json!("trace-1"));
|
||||
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
|
||||
assert_eq!(
|
||||
payload["candidates"][0]["extra_data"]["first_byte_time_ms"],
|
||||
json!(30)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_resolves_usage_request_id_to_metadata_trace_id() {
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
sample_candidate(
|
||||
"cand-used",
|
||||
"trace-2",
|
||||
0,
|
||||
RequestCandidateStatus::Success,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(200),
|
||||
),
|
||||
]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
vec![sample_endpoint()],
|
||||
vec![sample_key()],
|
||||
));
|
||||
let mut usage = sample_usage(
|
||||
"usage-request-2",
|
||||
"provider-1",
|
||||
"OpenAI",
|
||||
40,
|
||||
0.02,
|
||||
"completed",
|
||||
Some(200),
|
||||
100,
|
||||
);
|
||||
usage.candidate_id = Some("cand-used".to_string());
|
||||
usage.request_metadata = Some(json!({
|
||||
"trace_id": "trace-2"
|
||||
}));
|
||||
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,
|
||||
)
|
||||
.with_provider_catalog_reader(provider_catalog);
|
||||
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/usage-request-2",
|
||||
);
|
||||
|
||||
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");
|
||||
assert_eq!(payload["request_id"], json!("trace-2"));
|
||||
assert_eq!(payload["candidates"][0]["id"], json!("cand-used"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_returns_oauth_account_label_from_auth_config() {
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![
|
||||
|
||||
@@ -12,6 +12,7 @@ use aether_admin::observability::monitoring::{
|
||||
use aether_data_contracts::repository::{
|
||||
candidates::{DecisionTrace, RequestCandidateStatus},
|
||||
provider_catalog::StoredProviderCatalogKey,
|
||||
usage::StoredRequestUsageAudit,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
@@ -21,12 +22,16 @@ use serde_json::{Map, Value};
|
||||
use std::collections::BTreeMap;
|
||||
use tracing::debug;
|
||||
|
||||
struct ResolvedAdminMonitoringTrace {
|
||||
trace: DecisionTrace,
|
||||
usage: Option<StoredRequestUsageAudit>,
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_monitoring_trace_request_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let admin_state = state;
|
||||
let state = state.as_ref();
|
||||
let Some(request_id) =
|
||||
admin_monitoring_trace_request_id_from_path(&request_context.request_path)
|
||||
else {
|
||||
@@ -39,11 +44,8 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
|
||||
Err(detail) => return Ok(admin_monitoring_bad_request_response(detail)),
|
||||
};
|
||||
|
||||
let Some(trace) = state
|
||||
.data
|
||||
.read_decision_trace(&request_id, attempted_only)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
let Some(resolved) =
|
||||
resolve_admin_monitoring_trace(admin_state, &request_id, attempted_only).await?
|
||||
else {
|
||||
debug!(
|
||||
event_name = "admin_monitoring_request_trace_not_found",
|
||||
@@ -58,22 +60,113 @@ pub(super) async fn build_admin_monitoring_trace_request_response(
|
||||
attempted_only,
|
||||
));
|
||||
};
|
||||
let usage = state
|
||||
.data
|
||||
.read_request_usage_audit(&request_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let key_accounts = build_admin_monitoring_key_account_display_map(admin_state, &trace).await?;
|
||||
let key_accounts =
|
||||
build_admin_monitoring_key_account_display_map(admin_state, &resolved.trace).await?;
|
||||
|
||||
Ok(
|
||||
build_admin_monitoring_trace_request_payload_response_with_key_accounts(
|
||||
&trace,
|
||||
usage.as_ref(),
|
||||
&resolved.trace,
|
||||
resolved.usage.as_ref(),
|
||||
&key_accounts,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
async fn resolve_admin_monitoring_trace(
|
||||
state: &AdminAppState<'_>,
|
||||
request_id: &str,
|
||||
attempted_only: bool,
|
||||
) -> Result<Option<ResolvedAdminMonitoringTrace>, GatewayError> {
|
||||
let app = state.as_ref();
|
||||
if let Some(trace) = app
|
||||
.data
|
||||
.read_decision_trace(request_id, attempted_only)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
{
|
||||
let usage = app
|
||||
.data
|
||||
.read_request_usage_audit(request_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
return Ok(Some(ResolvedAdminMonitoringTrace { trace, usage }));
|
||||
}
|
||||
|
||||
let mut usage_candidates = Vec::new();
|
||||
if let Some(usage) = app
|
||||
.data
|
||||
.read_request_usage_audit(request_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
{
|
||||
usage_candidates.push(usage);
|
||||
}
|
||||
if let Some(usage) = state.find_request_usage_by_id(request_id).await? {
|
||||
if !usage_candidates.iter().any(|item| item.id == usage.id) {
|
||||
usage_candidates.push(usage);
|
||||
}
|
||||
}
|
||||
|
||||
for usage in usage_candidates {
|
||||
for trace_request_id in admin_monitoring_usage_trace_request_ids(&usage) {
|
||||
if trace_request_id == request_id {
|
||||
continue;
|
||||
}
|
||||
if let Some(trace) = app
|
||||
.data
|
||||
.read_decision_trace(&trace_request_id, attempted_only)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
{
|
||||
return Ok(Some(ResolvedAdminMonitoringTrace {
|
||||
trace,
|
||||
usage: Some(usage),
|
||||
}));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn admin_monitoring_usage_trace_request_ids(usage: &StoredRequestUsageAudit) -> Vec<String> {
|
||||
let mut ids = Vec::new();
|
||||
push_non_empty_unique(&mut ids, usage.request_id.as_str());
|
||||
if let Some(trace_id) = usage.trace_id() {
|
||||
push_non_empty_unique(&mut ids, trace_id);
|
||||
}
|
||||
if let Some(trace_id) = usage_trace_id_from_headers(usage.request_headers.as_ref()) {
|
||||
push_non_empty_unique(&mut ids, trace_id.as_str());
|
||||
}
|
||||
if let Some(trace_id) = usage_trace_id_from_headers(usage.provider_request_headers.as_ref()) {
|
||||
push_non_empty_unique(&mut ids, trace_id.as_str());
|
||||
}
|
||||
ids
|
||||
}
|
||||
|
||||
fn usage_trace_id_from_headers(headers: Option<&Value>) -> Option<String> {
|
||||
let object = headers?.as_object()?;
|
||||
object.iter().find_map(|(key, value)| {
|
||||
key.eq_ignore_ascii_case(crate::constants::TRACE_ID_HEADER)
|
||||
.then(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
.flatten()
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
fn push_non_empty_unique(values: &mut Vec<String>, value: &str) {
|
||||
let value = value.trim();
|
||||
if value.is_empty() || values.iter().any(|existing| existing == value) {
|
||||
return;
|
||||
}
|
||||
values.push(value.to_string());
|
||||
}
|
||||
|
||||
async fn build_admin_monitoring_key_account_display_map(
|
||||
state: &AdminAppState<'_>,
|
||||
trace: &DecisionTrace,
|
||||
|
||||
@@ -91,6 +91,26 @@ fn coerce_admin_provider_oauth_import_project_id(
|
||||
}
|
||||
}
|
||||
|
||||
fn json_import_expiry_value(value: Option<&serde_json::Value>) -> Option<u64> {
|
||||
let value = value?;
|
||||
json_u64_value(Some(value)).or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok())
|
||||
.and_then(|value| u64::try_from(value.timestamp()).ok())
|
||||
})
|
||||
}
|
||||
|
||||
fn json_import_expiry_from_keys(
|
||||
object: &serde_json::Map<String, serde_json::Value>,
|
||||
keys: &[&str],
|
||||
) -> Option<u64> {
|
||||
keys.iter()
|
||||
.find_map(|key| json_import_expiry_value(object.get(*key)))
|
||||
}
|
||||
|
||||
fn grok_cookie_value(raw: &str, name: &str) -> Option<String> {
|
||||
raw.trim()
|
||||
.strip_prefix("Cookie:")
|
||||
@@ -204,6 +224,8 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
object
|
||||
.get("sso_token")
|
||||
.or_else(|| object.get("ssoToken"))
|
||||
.or_else(|| object.get("session_token"))
|
||||
.or_else(|| object.get("sessionToken"))
|
||||
.or(grok_token_alias),
|
||||
)
|
||||
.or_else(|| {
|
||||
@@ -254,7 +276,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
refresh_token
|
||||
};
|
||||
let expires_at =
|
||||
json_u64_value(object.get("expires_at").or_else(|| object.get("expiresAt")));
|
||||
json_import_expiry_from_keys(object, &["expires_at", "expiresAt", "expired"]);
|
||||
let account_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_id")
|
||||
@@ -660,6 +682,21 @@ mod tests {
|
||||
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_common_chatgpt_web_json_aliases() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"chatgpt_web",
|
||||
r#"[{"session_token":"session-1","expired":"2030-01-01T00:00:00Z","chatgpt_account_id":"acc-1","chatgpt_plan_type":"plus"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("session-1"));
|
||||
assert_eq!(entries[0].expires_at, Some(1_893_456_000));
|
||||
assert_eq!(entries[0].account_id.as_deref(), Some("acc-1"));
|
||||
assert_eq!(entries[0].plan_type.as_deref(), Some("plus"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_plain_jwt_line_as_access_token() {
|
||||
let token = unsigned_jwt(json!({
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::provider_oauth_token_payload_expires_at_unix_secs;
|
||||
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::super::super::runtime::resolve_provider_oauth_runtime_endpoints;
|
||||
use super::super::super::runtime::{
|
||||
resolve_provider_oauth_runtime_endpoints,
|
||||
spawn_provider_oauth_account_state_refresh_after_update,
|
||||
};
|
||||
use super::super::super::state::{
|
||||
admin_provider_oauth_template, enrich_admin_provider_oauth_auth_config,
|
||||
is_fixed_provider_type_for_provider_oauth, json_non_empty_string,
|
||||
@@ -219,50 +221,20 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
));
|
||||
}
|
||||
|
||||
let mut account_state_recheck_attempted = false;
|
||||
let mut account_state_recheck_error = None::<String>;
|
||||
if provider_type == "codex" {
|
||||
if let Some(endpoint) = runtime_endpoint {
|
||||
let refreshed_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
.unwrap_or_else(|| key.clone());
|
||||
if let Some(result) = refresh_codex_provider_quota_locally(
|
||||
state,
|
||||
&provider,
|
||||
&endpoint,
|
||||
vec![refreshed_key],
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
account_state_recheck_attempted = true;
|
||||
let success = result
|
||||
.get("success")
|
||||
.and_then(serde_json::Value::as_u64)
|
||||
.unwrap_or(0);
|
||||
if success == 0 {
|
||||
account_state_recheck_error = result
|
||||
.get("results")
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.and_then(|results| results.first())
|
||||
.and_then(|value| value.get("message"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
provider.clone(),
|
||||
key_id.clone(),
|
||||
request_proxy.clone(),
|
||||
);
|
||||
|
||||
Ok(Json(json!({
|
||||
"provider_type": provider_type,
|
||||
"expires_at": expires_at,
|
||||
"has_refresh_token": refresh_token.is_some(),
|
||||
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
|
||||
"account_state_recheck_attempted": account_state_recheck_attempted,
|
||||
"account_state_recheck_error": account_state_recheck_error,
|
||||
"account_state_recheck_attempted": false,
|
||||
"account_state_recheck_error": serde_json::Value::Null,
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
@@ -828,13 +828,7 @@ async fn handle_admin_provider_oauth_windsurf_browser_device_poll(
|
||||
};
|
||||
callback_token.to_string()
|
||||
} else {
|
||||
let token = token.unwrap_or_default();
|
||||
if windsurf_raw_api_key(token).is_none() {
|
||||
return Ok(windsurf_browser_poll_error_response(
|
||||
"浏览器授权请提交包含 state 的回调 URL;纯 token 请使用导入授权",
|
||||
));
|
||||
}
|
||||
token.to_string()
|
||||
token.unwrap_or_default().to_string()
|
||||
};
|
||||
|
||||
let mut raw_credentials = serde_json::Map::new();
|
||||
|
||||
@@ -106,12 +106,21 @@ fn import_payload_string_any(
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn import_payload_u64(
|
||||
fn import_payload_u64_any(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
snake_case: &str,
|
||||
camel_case: &str,
|
||||
keys: &[&str],
|
||||
) -> Option<u64> {
|
||||
json_u64_value(payload.get(snake_case).or_else(|| payload.get(camel_case)))
|
||||
keys.iter().find_map(|key| {
|
||||
let value = payload.get(*key)?;
|
||||
json_u64_value(Some(value)).or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok())
|
||||
.and_then(|value| u64::try_from(value.timestamp()).ok())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn apply_single_import_hints(
|
||||
@@ -409,9 +418,17 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
||||
let access_token_input = import_payload_string_any(
|
||||
&raw_payload,
|
||||
&["access_token", "accessToken", "sso_token", "ssoToken"],
|
||||
&[
|
||||
"access_token",
|
||||
"accessToken",
|
||||
"sso_token",
|
||||
"ssoToken",
|
||||
"session_token",
|
||||
"sessionToken",
|
||||
],
|
||||
);
|
||||
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
|
||||
let imported_expires_at =
|
||||
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
|
||||
let name = raw_payload
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
@@ -622,8 +639,41 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::sanitize_windsurf_import_error;
|
||||
use super::{
|
||||
import_payload_string_any, import_payload_u64_any, sanitize_windsurf_import_error,
|
||||
};
|
||||
use aether_oauth::core::OAuthError;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn single_import_accepts_session_token_alias() {
|
||||
let payload = json!({
|
||||
"session_token": "session-1",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(
|
||||
import_payload_string_any(&payload, &["access_token", "session_token"]).as_deref(),
|
||||
Some("session-1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_import_accepts_iso_expired_alias() {
|
||||
let payload = json!({
|
||||
"expired": "2030-01-01T00:00:00Z",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(
|
||||
import_payload_u64_any(&payload, &["expires_at", "expiresAt", "expired"]),
|
||||
Some(1_893_456_000)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_import_error_redacts_http_body() {
|
||||
|
||||
+50
-40
@@ -81,48 +81,42 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
.await?;
|
||||
if provider_auto_remove_banned_keys(provider.config.as_ref()) {
|
||||
let now_unix_secs = helpers::unix_now_secs();
|
||||
let latest_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next();
|
||||
if latest_key.as_ref().is_some_and(|latest_key| {
|
||||
should_auto_remove_oauth_invalid_key(
|
||||
latest_key,
|
||||
None,
|
||||
false,
|
||||
now_unix_secs,
|
||||
)
|
||||
}) {
|
||||
state
|
||||
.clear_admin_provider_pool_cooldown(&provider.id, &key_id)
|
||||
.await;
|
||||
state
|
||||
.reset_admin_provider_pool_cost(&provider.id, &key_id)
|
||||
.await;
|
||||
if state.delete_provider_catalog_key(&key_id).await? {
|
||||
let deleted_key_ids = [key_id.clone()];
|
||||
state
|
||||
.cleanup_deleted_provider_catalog_refs(
|
||||
&provider.id,
|
||||
&[],
|
||||
&deleted_key_ids,
|
||||
let auto_removed = state
|
||||
.cleanup_provider_catalog_key_if_current(
|
||||
&provider,
|
||||
&key_id,
|
||||
|latest_key| {
|
||||
should_auto_remove_oauth_invalid_key(
|
||||
latest_key,
|
||||
Some(&failure_reason),
|
||||
false,
|
||||
now_unix_secs,
|
||||
)
|
||||
.await?;
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if auto_removed {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
}
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_failed_retained",
|
||||
"gateway manual provider oauth refresh failure retained key"
|
||||
);
|
||||
}
|
||||
}
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
@@ -171,9 +165,25 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
};
|
||||
|
||||
if !helpers::key_is_account_blocked(&key, OAUTH_ACCOUNT_BLOCK_PREFIX) {
|
||||
let _ = state
|
||||
let previous_oauth_refresh_issue =
|
||||
key.oauth_invalid_reason.as_deref().is_some_and(|reason| {
|
||||
reason.lines().map(str::trim).any(|line| {
|
||||
line.starts_with("[OAUTH_EXPIRED]") || line.starts_with("[REFRESH_FAILED]")
|
||||
})
|
||||
});
|
||||
let cleared = state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
|
||||
.await?;
|
||||
if cleared && previous_oauth_refresh_issue {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_fixed",
|
||||
"gateway manual provider oauth refresh cleared oauth invalid marker"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let refreshed_key = state
|
||||
|
||||
@@ -201,10 +201,10 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let mut updated = existing_key.clone();
|
||||
updated.is_active = true;
|
||||
updated.encrypted_api_key = Some(encrypted_api_key);
|
||||
updated.encrypted_auth_config = Some(encrypted_auth_config);
|
||||
updated.api_formats = provider_oauth_catalog_key_api_formats(provider_type, api_formats);
|
||||
updated.is_active = true;
|
||||
updated.expires_at_unix_secs = expires_at_unix_secs;
|
||||
updated.oauth_invalid_at_unix_secs = None;
|
||||
updated.oauth_invalid_reason = None;
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload, coerce_json_f64,
|
||||
coerce_json_string, default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
extract_execution_error_message, oauth_refresh_auto_removed_result,
|
||||
persist_provider_quota_refresh_state, quota_key_auto_removed,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
coerce_json_string, execute_provider_quota_plan, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -33,11 +33,10 @@ async fn execute_antigravity_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec = build_antigravity_pool_quota_request(
|
||||
&transport.key.id,
|
||||
&transport.endpoint.base_url,
|
||||
|
||||
@@ -1,16 +1,19 @@
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, extract_execution_error_message,
|
||||
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::quota::parse_chatgpt_web_conversation_init_response;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_contracts::{
|
||||
ExecutionResult, ProxySnapshot, ResolvedTransportProfile, TRANSPORT_BACKEND_BROWSER_WREQ,
|
||||
TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
@@ -18,10 +21,12 @@ use aether_provider_pool::{
|
||||
build_chatgpt_web_pool_quota_request, enrich_chatgpt_web_quota_metadata,
|
||||
normalize_chatgpt_web_image_quota_limit,
|
||||
};
|
||||
use base64::Engine as _;
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
|
||||
const CHATGPT_WEB_BROWSER_PROFILE: &str = "chrome143";
|
||||
|
||||
fn chatgpt_web_auth_config(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
@@ -67,26 +72,105 @@ async fn execute_chatgpt_web_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec =
|
||||
build_chatgpt_web_pool_quota_request(&transport.key.id, &endpoint.base_url, authorization);
|
||||
let resolved_transport_profile = state.resolve_transport_profile(transport);
|
||||
let plan = super::shared::build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
proxy,
|
||||
state.resolve_transport_profile(transport),
|
||||
chatgpt_web_quota_transport_profile(resolved_transport_profile.as_ref()),
|
||||
timeouts,
|
||||
);
|
||||
|
||||
execute_provider_quota_plan(state, transport, plan, "chatgpt_web").await
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_transport_profile(
|
||||
transport_profile: Option<&ResolvedTransportProfile>,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
match transport_profile {
|
||||
Some(profile)
|
||||
if profile
|
||||
.backend
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ) =>
|
||||
{
|
||||
Some(profile.clone())
|
||||
}
|
||||
_ => Some(default_chatgpt_web_quota_transport_profile()),
|
||||
}
|
||||
}
|
||||
|
||||
fn default_chatgpt_web_quota_transport_profile() -> ResolvedTransportProfile {
|
||||
ResolvedTransportProfile {
|
||||
profile_id: CHATGPT_WEB_BROWSER_PROFILE.to_string(),
|
||||
backend: TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: Some(json!({
|
||||
"browser_profile": CHATGPT_WEB_BROWSER_PROFILE,
|
||||
"source": "chatgpt_web_quota_default",
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_error_detail(result: &ExecutionResult) -> Option<String> {
|
||||
extract_execution_error_message(result).or_else(|| {
|
||||
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
|
||||
let decoded = base64::engine::general_purpose::STANDARD
|
||||
.decode(body)
|
||||
.ok()?;
|
||||
let text = String::from_utf8_lossy(&decoded).trim().to_string();
|
||||
(!text.is_empty()).then_some(text)
|
||||
})
|
||||
}
|
||||
|
||||
fn chatgpt_web_is_structured_account_block(message: &str) -> bool {
|
||||
let lowered = message.to_ascii_lowercase();
|
||||
[
|
||||
"account has been disabled",
|
||||
"account disabled",
|
||||
"account has been deactivated",
|
||||
"account_deactivated",
|
||||
"account deactivated",
|
||||
"organization has been disabled",
|
||||
"organization_disabled",
|
||||
"deactivated_workspace",
|
||||
"account suspended",
|
||||
"account banned",
|
||||
"account_block",
|
||||
"account blocked",
|
||||
"访问被禁止",
|
||||
"账户访问被禁止",
|
||||
"账户已封禁",
|
||||
"封禁",
|
||||
"封号",
|
||||
"被封",
|
||||
]
|
||||
.iter()
|
||||
.any(|keyword| lowered.contains(keyword))
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_403_refresh_failed_reason(message: Option<&str>) -> String {
|
||||
let detail = message
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.filter(|value| !value.contains('<'))
|
||||
.unwrap_or("ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制");
|
||||
format!("{OAUTH_REFRESH_FAILED_PREFIX}{detail}")
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<&str>) -> String {
|
||||
let message = upstream_message.unwrap_or_default().trim();
|
||||
if status_code == 403 && !chatgpt_web_is_structured_account_block(message) {
|
||||
return chatgpt_web_quota_403_refresh_failed_reason(upstream_message);
|
||||
}
|
||||
let detail = if message.is_empty() {
|
||||
match status_code {
|
||||
401 => "ChatGPT Web Token 无效或已过期",
|
||||
@@ -103,6 +187,19 @@ fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<&
|
||||
}
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_result_message(reason: &str) -> String {
|
||||
for prefix in [
|
||||
OAUTH_REFRESH_FAILED_PREFIX,
|
||||
OAUTH_EXPIRED_PREFIX,
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX,
|
||||
] {
|
||||
if let Some(message) = reason.strip_prefix(prefix) {
|
||||
return message.trim().to_string();
|
||||
}
|
||||
}
|
||||
reason.trim().to_string()
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -216,8 +313,20 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
message = Some("响应中未包含 ChatGPT Web 生图限额信息".to_string());
|
||||
}
|
||||
} else {
|
||||
let err_msg = extract_execution_error_message(&result);
|
||||
message = Some(match err_msg.as_deref() {
|
||||
let err_msg = chatgpt_web_quota_error_detail(&result);
|
||||
let invalid_reason = if matches!(result.status_code, 401 | 403) {
|
||||
Some(chatgpt_web_quota_invalid_reason(
|
||||
result.status_code,
|
||||
err_msg.as_deref(),
|
||||
))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let display_detail = invalid_reason
|
||||
.as_deref()
|
||||
.map(chatgpt_web_quota_result_message)
|
||||
.or_else(|| err_msg.clone());
|
||||
message = Some(match display_detail.as_deref() {
|
||||
Some(detail) if !detail.is_empty() => {
|
||||
format!(
|
||||
"conversation/init 返回状态码 {}: {}",
|
||||
@@ -229,12 +338,14 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
|
||||
if matches!(result.status_code, 401 | 403) {
|
||||
oauth_invalid_at_unix_secs = Some(now_unix_secs);
|
||||
oauth_invalid_reason = Some(chatgpt_web_quota_invalid_reason(
|
||||
result.status_code,
|
||||
err_msg.as_deref(),
|
||||
));
|
||||
oauth_invalid_reason = invalid_reason;
|
||||
status = if result.status_code == 401 {
|
||||
"auth_invalid".to_string()
|
||||
} else if oauth_invalid_reason
|
||||
.as_deref()
|
||||
.is_some_and(|reason| reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX))
|
||||
{
|
||||
"refresh_failed".to_string()
|
||||
} else {
|
||||
"forbidden".to_string()
|
||||
};
|
||||
@@ -303,3 +414,81 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
"auto_removed": auto_removed_count,
|
||||
})))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aether_contracts::{ResponseBody, TRANSPORT_BACKEND_REQWEST_RUSTLS};
|
||||
use base64::Engine as _;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn quota_refresh_defaults_to_browser_wreq_transport() {
|
||||
let profile = chatgpt_web_quota_transport_profile(None).expect("transport profile");
|
||||
|
||||
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
|
||||
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
|
||||
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_AUTO);
|
||||
assert_eq!(profile.pool_scope, TRANSPORT_POOL_SCOPE_KEY);
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("browser_profile"))
|
||||
.and_then(serde_json::Value::as_str),
|
||||
Some(CHATGPT_WEB_BROWSER_PROFILE)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_refresh_overrides_non_browser_transport() {
|
||||
let reqwest_profile = ResolvedTransportProfile {
|
||||
profile_id: "chrome_136".to_string(),
|
||||
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: None,
|
||||
};
|
||||
|
||||
let profile =
|
||||
chatgpt_web_quota_transport_profile(Some(&reqwest_profile)).expect("transport profile");
|
||||
|
||||
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
|
||||
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn browser_challenge_403_is_not_account_block() {
|
||||
let body = "<!DOCTYPE html><html><head><title>Just a moment...</title></head><body>Cloudflare</body></html>";
|
||||
let result = ExecutionResult {
|
||||
request_id: "chatgpt-web-quota:test".to_string(),
|
||||
candidate_id: None,
|
||||
status_code: 403,
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
};
|
||||
|
||||
let detail = chatgpt_web_quota_error_detail(&result).expect("html body should decode");
|
||||
let reason = chatgpt_web_quota_invalid_reason(result.status_code, Some(&detail));
|
||||
|
||||
assert!(reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX));
|
||||
assert!(!reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX));
|
||||
assert_eq!(
|
||||
chatgpt_web_quota_result_message(&reason),
|
||||
"ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_account_block_403_remains_account_block() {
|
||||
let reason = chatgpt_web_quota_invalid_reason(403, Some("account has been deactivated"));
|
||||
|
||||
assert!(reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -44,6 +44,15 @@ fn merge_codex_quota_metadata(
|
||||
serde_json::Value::Object(merged)
|
||||
}
|
||||
|
||||
fn codex_oauth_refresh_issue_reason(reason: Option<&str>) -> bool {
|
||||
reason.is_some_and(|reason| {
|
||||
reason
|
||||
.lines()
|
||||
.map(str::trim)
|
||||
.any(|line| line.starts_with("[OAUTH_EXPIRED]") || line.starts_with("[REFRESH_FAILED]"))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -56,8 +65,13 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
let mut success_count = 0usize;
|
||||
let mut failed_count = 0usize;
|
||||
let mut auto_removed_count = 0usize;
|
||||
let mut refresh_fixed_count = 0usize;
|
||||
let mut refresh_failed_retained_count = 0usize;
|
||||
let mut auto_removed_hard_banned_count = 0usize;
|
||||
|
||||
for key in keys {
|
||||
let had_oauth_refresh_issue =
|
||||
codex_oauth_refresh_issue_reason(key.oauth_invalid_reason.as_deref());
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
@@ -276,13 +290,9 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
}
|
||||
|
||||
let auto_removed = auto_remove_abnormal_keys
|
||||
let auto_remove_candidate = auto_remove_abnormal_keys
|
||||
&& should_auto_remove_structured_reason(oauth_invalid_reason.as_deref());
|
||||
if auto_removed {
|
||||
if state.delete_provider_catalog_key(&key.id).await? {
|
||||
auto_removed_count += 1;
|
||||
}
|
||||
} else if !persist_provider_quota_refresh_state(
|
||||
let persisted = persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
@@ -290,8 +300,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
oauth_invalid_reason.clone(),
|
||||
None,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
.await?;
|
||||
if !persisted {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
@@ -301,6 +311,29 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
let auto_removed = if auto_remove_candidate {
|
||||
state
|
||||
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
|
||||
should_auto_remove_structured_reason(latest_key.oauth_invalid_reason.as_deref())
|
||||
})
|
||||
.await?
|
||||
} else {
|
||||
false
|
||||
};
|
||||
if auto_removed {
|
||||
auto_removed_count += 1;
|
||||
auto_removed_hard_banned_count += 1;
|
||||
}
|
||||
let refresh_fixed =
|
||||
status == "success" && had_oauth_refresh_issue && oauth_invalid_reason.is_none();
|
||||
if refresh_fixed {
|
||||
refresh_fixed_count += 1;
|
||||
}
|
||||
let refresh_failed_retained =
|
||||
status != "success" && oauth_invalid_reason.is_some() && !auto_removed;
|
||||
if refresh_failed_retained {
|
||||
refresh_failed_retained_count += 1;
|
||||
}
|
||||
|
||||
if status == "success" {
|
||||
success_count += 1;
|
||||
@@ -336,6 +369,13 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
if auto_removed {
|
||||
payload.insert("auto_removed".to_string(), json!(true));
|
||||
payload.insert("auto_removed_hard_banned".to_string(), json!(true));
|
||||
}
|
||||
if refresh_fixed {
|
||||
payload.insert("refresh_fixed".to_string(), json!(true));
|
||||
}
|
||||
if refresh_failed_retained {
|
||||
payload.insert("refresh_failed_retained".to_string(), json!(true));
|
||||
}
|
||||
results.push(serde_json::Value::Object(payload));
|
||||
}
|
||||
@@ -346,5 +386,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
"total": results.len(),
|
||||
"results": results,
|
||||
"auto_removed": auto_removed_count,
|
||||
"refresh_fixed": refresh_fixed_count,
|
||||
"refresh_failed_retained": refresh_failed_retained_count,
|
||||
"auto_removed_hard_banned": auto_removed_hard_banned_count,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::super::shared::{
|
||||
build_provider_quota_execution_plan, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, ProviderQuotaExecutionOutcome,
|
||||
build_provider_quota_execution_plan, execute_provider_quota_plan,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -38,11 +38,10 @@ pub(super) async fn execute_codex_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, extract_execution_error_message,
|
||||
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
|
||||
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
@@ -245,11 +244,10 @@ async fn execute_grok_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let transport_profile = state.resolve_transport_profile(transport);
|
||||
let base_url = grok_base_url(endpoint);
|
||||
let headers = build_grok_quota_headers(
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::super::shared::{
|
||||
build_provider_quota_execution_plan, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, ProviderQuotaExecutionOutcome,
|
||||
build_provider_quota_execution_plan, execute_provider_quota_plan,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{
|
||||
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminKiroRequestAuth,
|
||||
@@ -23,11 +23,10 @@ pub(super) async fn execute_kiro_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec = build_kiro_pool_quota_request(
|
||||
&transport.key.id,
|
||||
&KiroPoolQuotaAuthInput {
|
||||
|
||||
@@ -46,6 +46,23 @@ pub(super) fn default_provider_quota_execution_timeouts(
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider_quota_execution_timeouts(
|
||||
configured: Option<ExecutionTimeouts>,
|
||||
proxy: Option<&ProxySnapshot>,
|
||||
) -> ExecutionTimeouts {
|
||||
let defaults = default_provider_quota_execution_timeouts(proxy);
|
||||
let Some(mut timeouts) = configured else {
|
||||
return defaults;
|
||||
};
|
||||
timeouts.connect_ms = timeouts.connect_ms.or(defaults.connect_ms);
|
||||
timeouts.read_ms = timeouts.read_ms.or(defaults.read_ms);
|
||||
timeouts.write_ms = timeouts.write_ms.or(defaults.write_ms);
|
||||
timeouts.pool_ms = timeouts.pool_ms.or(defaults.pool_ms);
|
||||
timeouts.total_ms = timeouts.total_ms.or(defaults.total_ms);
|
||||
timeouts.first_byte_ms = timeouts.first_byte_ms.or(defaults.first_byte_ms);
|
||||
timeouts
|
||||
}
|
||||
|
||||
pub(crate) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
|
||||
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
|
||||
}
|
||||
@@ -317,12 +334,7 @@ pub(super) async fn execute_provider_quota_plan(
|
||||
match state.execute_execution_runtime_sync_plan(None, &plan).await {
|
||||
Ok(result) => Ok(ProviderQuotaExecutionOutcome::Response(result)),
|
||||
Err(err) => {
|
||||
let error = match err {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
};
|
||||
let error = err.into_message();
|
||||
let proxy_node_id = plan
|
||||
.proxy
|
||||
.as_ref()
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload,
|
||||
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload, execute_provider_quota_plan,
|
||||
extract_execution_error_message, persist_provider_quota_refresh_state,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
quota_refresh_success_invalid_state, resolve_provider_quota_execution_timeouts,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -33,11 +33,10 @@ async fn execute_windsurf_probe_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
|
||||
@@ -301,12 +301,7 @@ fn admin_provider_ops_decode_response_bytes(
|
||||
}
|
||||
|
||||
fn admin_provider_ops_gateway_error_message(error: GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
error.into_message()
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_verify_execution_error_message(error: &str) -> String {
|
||||
|
||||
@@ -16,7 +16,12 @@ use aether_runtime_state::{DataLayerError, RuntimeState};
|
||||
use futures_util::future::join_all;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use tracing::warn;
|
||||
use tracing::{info, warn};
|
||||
|
||||
const DEFAULT_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT: usize = 512;
|
||||
const MAX_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT: usize = 10_000;
|
||||
const POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT_ENV: &str =
|
||||
"AETHER_GATEWAY_ADMIN_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT";
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
SystemTime::now()
|
||||
@@ -29,6 +34,20 @@ fn should_load_active_probe_members(pool_config: &AdminProviderPoolConfig) -> bo
|
||||
pool_config.probing_enabled
|
||||
}
|
||||
|
||||
fn pool_runtime_window_metric_key_limit() -> usize {
|
||||
std::env::var(POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT_ENV)
|
||||
.ok()
|
||||
.and_then(|value| value.trim().parse::<usize>().ok())
|
||||
.filter(|value| *value > 0)
|
||||
.unwrap_or(DEFAULT_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT)
|
||||
.clamp(1, MAX_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT)
|
||||
}
|
||||
|
||||
fn bounded_runtime_window_metric_key_ids(key_ids: &[String], limit: usize) -> &[String] {
|
||||
let end = key_ids.len().min(limit.max(1));
|
||||
&key_ids[..end]
|
||||
}
|
||||
|
||||
pub(crate) async fn read_admin_provider_pool_cooldown_counts(
|
||||
runtime: &RuntimeState,
|
||||
provider_ids: &[String],
|
||||
@@ -54,8 +73,21 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
) -> AdminProviderPoolRuntimeState {
|
||||
let mut state = AdminProviderPoolRuntimeState::default();
|
||||
let cooldown_keys = pool_cooldown_keys(provider_id, key_ids);
|
||||
let cost_keys = pool_cost_keys(provider_id, key_ids);
|
||||
let latency_keys = pool_latency_keys(provider_id, key_ids);
|
||||
let metric_key_limit = pool_runtime_window_metric_key_limit();
|
||||
let metric_key_ids = bounded_runtime_window_metric_key_ids(key_ids, metric_key_limit);
|
||||
if metric_key_ids.len() < key_ids.len() {
|
||||
info!(
|
||||
event_name = "admin_pool_runtime_window_metrics_truncated",
|
||||
log_type = "event",
|
||||
provider_id,
|
||||
total_key_count = key_ids.len(),
|
||||
scanned_key_count = metric_key_ids.len(),
|
||||
metric_key_limit,
|
||||
"gateway limited admin pool runtime cost/latency window reads"
|
||||
);
|
||||
}
|
||||
let cost_keys = pool_cost_keys(provider_id, metric_key_ids);
|
||||
let latency_keys = pool_latency_keys(provider_id, metric_key_ids);
|
||||
let sticky_sessions_enabled = pool_config.sticky_session_ttl_seconds > 0
|
||||
&& admin_provider_pool_cache_affinity_enabled(pool_config);
|
||||
|
||||
@@ -179,7 +211,7 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
.map(|cost_key| runtime.score_range_by_min(cost_key, cost_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(cost_results) {
|
||||
for (key_id, members) in metric_key_ids.iter().zip(cost_results) {
|
||||
let total = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
@@ -197,7 +229,7 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
.map(|latency_key| runtime.score_range_by_min(latency_key, latency_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(latency_results) {
|
||||
for (key_id, members) in metric_key_ids.iter().zip(latency_results) {
|
||||
let samples = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
@@ -265,3 +297,30 @@ pub(crate) async fn read_admin_provider_pool_key_cooldown_reason(
|
||||
.kv_get(&pool_cooldown_key(provider_id, key_id))
|
||||
.await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::bounded_runtime_window_metric_key_ids;
|
||||
|
||||
#[test]
|
||||
fn runtime_window_metric_key_ids_are_bounded() {
|
||||
let key_ids = vec![
|
||||
"key-1".to_string(),
|
||||
"key-2".to_string(),
|
||||
"key-3".to_string(),
|
||||
];
|
||||
|
||||
let bounded = bounded_runtime_window_metric_key_ids(&key_ids, 2);
|
||||
|
||||
assert_eq!(bounded, &key_ids[..2]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_window_metric_key_ids_keep_at_least_one_key() {
|
||||
let key_ids = vec!["key-1".to_string(), "key-2".to_string()];
|
||||
|
||||
let bounded = bounded_runtime_window_metric_key_ids(&key_ids, 0);
|
||||
|
||||
assert_eq!(bounded, &key_ids[..1]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_data_contracts::repository::usage::StoredProviderApiKeyWindowUsageSummary;
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -933,19 +934,14 @@ fn admin_pool_health_score(key: &StoredProviderCatalogKey) -> f64 {
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_pool_circuit_breaker_open(key: &StoredProviderCatalogKey) -> bool {
|
||||
fn admin_pool_circuit_breaker_open(key: &StoredProviderCatalogKey, now_unix_secs: u64) -> bool {
|
||||
key.circuit_breaker_by_format
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.map(|formats| {
|
||||
formats
|
||||
.values()
|
||||
.filter_map(serde_json::Value::as_object)
|
||||
.any(|item| {
|
||||
item.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.any(|item| provider_key_circuit_payload_is_active_open_at(item, now_unix_secs))
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
@@ -1044,7 +1040,7 @@ pub(super) fn build_admin_pool_key_payload(
|
||||
.as_ref()
|
||||
.and_then(|_| runtime.cooldown_ttl_by_key.get(&key.id).copied());
|
||||
let health_score = admin_pool_health_score(key);
|
||||
let circuit_breaker_open = admin_pool_circuit_breaker_open(key);
|
||||
let circuit_breaker_open = admin_pool_circuit_breaker_open(key, now_unix_secs);
|
||||
let auth_semantics = provider_key_auth_semantics(key, provider_type);
|
||||
let account_quota_exhausted = pool_config
|
||||
.as_ref()
|
||||
|
||||
@@ -99,6 +99,7 @@ static PROVIDER_QUERY_POOL_LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
struct ProviderQueryKeyFetchResult {
|
||||
models: Vec<Value>,
|
||||
error: Option<String>,
|
||||
warning: Option<String>,
|
||||
from_cache: bool,
|
||||
has_success: bool,
|
||||
}
|
||||
@@ -288,6 +289,7 @@ fn provider_query_codex_preset_fallback(
|
||||
Some(ProviderQueryKeyFetchResult {
|
||||
models: aggregate_models_for_cache(&models),
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
})
|
||||
@@ -427,6 +429,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models,
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: true,
|
||||
has_success: true,
|
||||
});
|
||||
@@ -444,6 +447,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models,
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
});
|
||||
@@ -451,6 +455,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -477,6 +482,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -492,6 +498,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -528,18 +535,25 @@ async fn provider_query_fetch_models_for_key(
|
||||
}
|
||||
}
|
||||
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
} else {
|
||||
let has_models = !unique_models.is_empty();
|
||||
let mut error = if !has_models && !all_errors.is_empty() {
|
||||
Some(all_errors.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if unique_models.is_empty() && error.is_none() {
|
||||
let warning = if has_models && !all_errors.is_empty() {
|
||||
Some(all_errors.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if !has_models && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL.to_string());
|
||||
}
|
||||
|
||||
Ok(ProviderQueryKeyFetchResult {
|
||||
models: provider_query_filter_models_for_key(provider, key, unique_models),
|
||||
error,
|
||||
warning,
|
||||
from_cache: false,
|
||||
has_success: outcome.has_success,
|
||||
})
|
||||
@@ -600,6 +614,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": result.error,
|
||||
"warning": result.warning,
|
||||
"from_cache": result.from_cache,
|
||||
},
|
||||
"provider": provider_query_provider_payload(&provider),
|
||||
@@ -632,6 +647,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": serde_json::Value::Null,
|
||||
"warning": serde_json::Value::Null,
|
||||
"from_cache": true,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": active_key_count,
|
||||
@@ -655,6 +671,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
|
||||
let mut all_models = Vec::new();
|
||||
let mut all_errors = Vec::new();
|
||||
let mut all_warnings = Vec::new();
|
||||
let mut cache_hit_count = 0usize;
|
||||
let mut fetch_count = 0usize;
|
||||
for key in &ordered_keys {
|
||||
@@ -669,6 +686,13 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
error
|
||||
));
|
||||
}
|
||||
if let Some(warning) = result.warning {
|
||||
all_warnings.push(format!(
|
||||
"Key {}: {}",
|
||||
provider_query_key_display_name(key),
|
||||
warning
|
||||
));
|
||||
}
|
||||
if result.from_cache {
|
||||
cache_hit_count += 1;
|
||||
} else {
|
||||
@@ -694,10 +718,17 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
provider_query_write_provider_cached_models(state, &provider.id, &models).await;
|
||||
}
|
||||
let success = !models.is_empty();
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
let mut all_issues = all_errors;
|
||||
all_issues.extend(all_warnings);
|
||||
let mut error = if !success && !all_issues.is_empty() {
|
||||
Some(all_issues.join("; "))
|
||||
} else {
|
||||
Some(all_errors.join("; "))
|
||||
None
|
||||
};
|
||||
let warning = if success && !all_issues.is_empty() {
|
||||
Some(all_issues.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if !success && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
|
||||
@@ -709,6 +740,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": error,
|
||||
"warning": warning,
|
||||
"from_cache": fetch_count == 0 && cache_hit_count > 0,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": cache_hit_count,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use super::super::payload::{
|
||||
provider_query_extract_api_key_id, provider_query_extract_force_refresh,
|
||||
provider_query_extract_api_key_ids, provider_query_extract_force_refresh,
|
||||
provider_query_extract_model, provider_query_extract_provider_id,
|
||||
provider_query_extract_request_id,
|
||||
};
|
||||
@@ -68,6 +68,7 @@ use aether_model_fetch::{
|
||||
aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
|
||||
preset_models_for_provider, selected_models_fetch_endpoints,
|
||||
};
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
http::{self, HeaderMap, HeaderName, HeaderValue},
|
||||
@@ -122,6 +123,7 @@ const ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL: &str =
|
||||
"No active endpoint or API key found";
|
||||
const ADMIN_PROVIDER_QUERY_INVALID_MAPPED_MODEL_DETAIL: &str =
|
||||
"mapped_model_name is not valid for the selected model and endpoint";
|
||||
const PROVIDER_QUERY_KEY_MODEL_NOT_ALLOWED_SKIP_REASON: &str = "key_model_not_allowed";
|
||||
const ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX: &str = "upstream_models_provider:";
|
||||
const DEFAULT_PROVIDER_QUERY_TEST_MESSAGE: &str = "Hello! This is a test message.";
|
||||
static PROVIDER_QUERY_POOL_LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
@@ -859,7 +861,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
selected_key_id: Option<&str>,
|
||||
selected_key_ids: Option<&BTreeSet<String>>,
|
||||
) -> Option<StoredProviderCatalogEndpoint> {
|
||||
for priority in 0..=2 {
|
||||
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
|
||||
@@ -872,7 +874,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
}
|
||||
for key in keys {
|
||||
if !key.is_active
|
||||
|| selected_key_id.is_some_and(|value| value != key.id.as_str())
|
||||
|| !provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
|
||||
|| !provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
@@ -904,7 +906,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
endpoint.is_active
|
||||
&& keys.iter().any(|key| {
|
||||
key.is_active
|
||||
&& selected_key_id.is_none_or(|value| value == key.id.as_str())
|
||||
&& provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
|
||||
&& provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
@@ -916,10 +918,56 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn provider_query_selected_key_ids_allow_key(
|
||||
selected_key_ids: Option<&BTreeSet<String>>,
|
||||
key_id: &str,
|
||||
) -> bool {
|
||||
selected_key_ids.is_none_or(|ids| ids.contains(key_id))
|
||||
}
|
||||
|
||||
fn provider_query_selected_key_ids_all_exist(
|
||||
selected_key_ids: &BTreeSet<String>,
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
) -> bool {
|
||||
selected_key_ids
|
||||
.iter()
|
||||
.all(|id| keys.iter().any(|key| key.id == *id))
|
||||
}
|
||||
|
||||
fn provider_query_model_name_matches(left: &str, right: &str) -> bool {
|
||||
let left = left.trim();
|
||||
let right = right.trim();
|
||||
!left.is_empty() && !right.is_empty() && left.eq_ignore_ascii_case(right)
|
||||
}
|
||||
|
||||
fn provider_query_key_allows_effective_test_model(
|
||||
key: &StoredProviderCatalogKey,
|
||||
requested_model: &str,
|
||||
effective_model: &str,
|
||||
) -> bool {
|
||||
let allowed_models = json_string_list(key.allowed_models.as_ref());
|
||||
if key.allowed_models.is_none() || allowed_models.is_empty() {
|
||||
return true;
|
||||
}
|
||||
|
||||
let requested_base_model = crate::ai_serving::model_directive_base_model(requested_model);
|
||||
allowed_models
|
||||
.iter()
|
||||
.map(String::as_str)
|
||||
.any(|allowed_model| {
|
||||
provider_query_model_name_matches(allowed_model, requested_model)
|
||||
|| provider_query_model_name_matches(allowed_model, effective_model)
|
||||
|| requested_base_model.as_deref().is_some_and(|base_model| {
|
||||
provider_query_model_name_matches(allowed_model, base_model)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_test_key_sort_key(
|
||||
provider_type: &str,
|
||||
key: &StoredProviderCatalogKey,
|
||||
endpoint_api_format: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> (u8, u8, i32, u64, i32) {
|
||||
let quota_exhausted =
|
||||
admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type);
|
||||
@@ -928,10 +976,7 @@ fn provider_query_test_key_sort_key(
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get(endpoint_api_format))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("open"))
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
.is_some_and(|value| provider_key_circuit_payload_is_active_open_at(value, now_unix_secs));
|
||||
let health_score = key
|
||||
.health_by_format
|
||||
.as_ref()
|
||||
@@ -1263,7 +1308,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
)
|
||||
})?;
|
||||
let selected_key_id = provider_query_extract_api_key_id(payload);
|
||||
let selected_key_ids = provider_query_extract_api_key_ids(payload);
|
||||
let requested_endpoint_id = provider_query_extract_endpoint_id(payload);
|
||||
let requested_api_format = provider_query_extract_api_format(payload);
|
||||
let endpoint = if requested_endpoint_id.is_none()
|
||||
@@ -1275,7 +1320,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider,
|
||||
&endpoints,
|
||||
&all_keys,
|
||||
selected_key_id.as_deref(),
|
||||
selected_key_ids.as_ref(),
|
||||
)
|
||||
.await
|
||||
.ok_or_else(|| {
|
||||
@@ -1308,22 +1353,11 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(api_key_id) = selected_key_id.as_deref() {
|
||||
let Some(key) = all_keys.iter().find(|key| key.id == api_key_id) else {
|
||||
if let Some(selected_key_ids) = selected_key_ids.as_ref() {
|
||||
if !provider_query_selected_key_ids_all_exist(selected_key_ids, &all_keys) {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
|
||||
));
|
||||
};
|
||||
if !key.is_active
|
||||
|| !provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
&endpoint.api_format,
|
||||
)
|
||||
{
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1378,20 +1412,43 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
.unwrap_or(requested_model.clone())
|
||||
};
|
||||
|
||||
let mut keys = all_keys
|
||||
let now_unix_secs = current_unix_ms() / 1000;
|
||||
let mut keys = Vec::new();
|
||||
let mut model_skipped_candidates = Vec::new();
|
||||
|
||||
for key in all_keys
|
||||
.into_iter()
|
||||
.filter(|key| key.is_active)
|
||||
.filter(|key| {
|
||||
selected_key_id
|
||||
.as_deref()
|
||||
.is_none_or(|value| value == key.id.as_str())
|
||||
})
|
||||
.filter(|key| provider_query_selected_key_ids_allow_key(selected_key_ids.as_ref(), &key.id))
|
||||
.filter(|key| {
|
||||
provider_query_key_supports_endpoint(key, &provider.provider_type, &endpoint.api_format)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
{
|
||||
if provider_query_key_allows_effective_test_model(&key, &requested_model, &effective_model)
|
||||
{
|
||||
keys.push(key);
|
||||
} else {
|
||||
model_skipped_candidates.push(ProviderQueryTestCandidate {
|
||||
endpoint: endpoint.clone(),
|
||||
key,
|
||||
effective_model: effective_model.clone(),
|
||||
scheduler_skip_reason: Some(
|
||||
PROVIDER_QUERY_KEY_MODEL_NOT_ALLOWED_SKIP_REASON.to_string(),
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let candidates = if test_mode.eq_ignore_ascii_case("pool") {
|
||||
model_skipped_candidates.sort_by_key(|candidate| {
|
||||
provider_query_test_key_sort_key(
|
||||
provider.provider_type.as_str(),
|
||||
&candidate.key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
|
||||
let scheduled_candidates = if test_mode.eq_ignore_ascii_case("pool") {
|
||||
if let Some(pool_config) =
|
||||
admin_provider_pool_config_from_config_value(provider.config.as_ref())
|
||||
{
|
||||
@@ -1411,6 +1468,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider.provider_type.as_str(),
|
||||
key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
keys.into_iter()
|
||||
@@ -1428,6 +1486,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider.provider_type.as_str(),
|
||||
key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
keys.into_iter()
|
||||
@@ -1439,6 +1498,8 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
let mut candidates = model_skipped_candidates;
|
||||
candidates.extend(scheduled_candidates);
|
||||
|
||||
if candidates.is_empty() {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
|
||||
@@ -64,6 +64,68 @@ fn sample_openai_image_transport(provider_type: &str) -> AdminGatewayProviderTra
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_catalog_key_with_allowed_models(
|
||||
allowed_models: Option<serde_json::Value>,
|
||||
) -> aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey {
|
||||
let mut key =
|
||||
aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"key".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("sample provider key should build");
|
||||
key.allowed_models = allowed_models;
|
||||
key
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_allows_keys_without_model_restrictions() {
|
||||
let unrestricted = sample_catalog_key_with_allowed_models(None);
|
||||
let empty = sample_catalog_key_with_allowed_models(Some(json!([])));
|
||||
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&unrestricted,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&empty,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_filters_key_disallowed_for_requested_model() {
|
||||
let key = sample_catalog_key_with_allowed_models(Some(json!(["model-a"])));
|
||||
|
||||
assert!(!provider_query_key_allows_effective_test_model(
|
||||
&key,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_allows_key_for_requested_or_mapped_model() {
|
||||
let requested_allowed = sample_catalog_key_with_allowed_models(Some(json!(["model-b"])));
|
||||
let mapped_allowed = sample_catalog_key_with_allowed_models(Some(json!(["MODEL-B-UPSTREAM"])));
|
||||
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&requested_allowed,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&mapped_allowed,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_preserves_custom_model() {
|
||||
let payload = json!({
|
||||
@@ -232,6 +294,27 @@ fn provider_query_request_body_model_uses_non_empty_string_only() {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_extracts_multiple_selected_key_ids() {
|
||||
let payload = json!({
|
||||
"api_key_ids": [" key-b ", "", "key-a", "key-b"],
|
||||
"api_key_id": "key-c"
|
||||
});
|
||||
|
||||
let ids = provider_query_extract_api_key_ids(&payload)
|
||||
.expect("non-empty key selection should be extracted")
|
||||
.into_iter()
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(ids, vec!["key-a", "key-b", "key-c"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_empty_selected_key_ids_keep_default_selection() {
|
||||
assert!(provider_query_extract_api_key_ids(&json!({})).is_none());
|
||||
assert!(provider_query_extract_api_key_ids(&json!({ "api_key_ids": [] })).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() {
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use axum::body::Bytes;
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
pub(crate) fn parse_admin_provider_query_body(
|
||||
request_body: Option<&Bytes>,
|
||||
@@ -36,6 +37,47 @@ pub(crate) fn provider_query_extract_api_key_id(payload: &serde_json::Value) ->
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn provider_query_insert_api_key_id(ids: &mut BTreeSet<String>, value: &str) {
|
||||
let value = value.trim();
|
||||
if !value.is_empty() {
|
||||
ids.insert(value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn provider_query_extract_api_key_ids(
|
||||
payload: &serde_json::Value,
|
||||
) -> Option<BTreeSet<String>> {
|
||||
let mut ids = BTreeSet::new();
|
||||
|
||||
if let Some(value) = payload
|
||||
.get("api_key_ids")
|
||||
.or_else(|| payload.get("provider_key_ids"))
|
||||
.or_else(|| payload.get("key_ids"))
|
||||
{
|
||||
match value {
|
||||
serde_json::Value::Array(items) => {
|
||||
for item in items {
|
||||
if let Some(value) = item.as_str() {
|
||||
provider_query_insert_api_key_id(&mut ids, value);
|
||||
}
|
||||
}
|
||||
}
|
||||
serde_json::Value::String(value) => {
|
||||
for item in value.split(',') {
|
||||
provider_query_insert_api_key_id(&mut ids, item);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(api_key_id) = provider_query_extract_api_key_id(payload) {
|
||||
ids.insert(api_key_id);
|
||||
}
|
||||
|
||||
(!ids.is_empty()).then_some(ids)
|
||||
}
|
||||
|
||||
pub(crate) fn provider_query_extract_force_refresh(payload: &serde_json::Value) -> bool {
|
||||
payload
|
||||
.get("force_refresh")
|
||||
|
||||
@@ -662,10 +662,5 @@ fn admin_provider_oauth_decode_response_bytes(
|
||||
}
|
||||
|
||||
fn admin_provider_oauth_gateway_error_message(error: GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
error.into_message()
|
||||
}
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
use super::*;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -193,8 +196,6 @@ impl<'a> AdminAppState<'a> {
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
|
||||
let Some(provider) = self
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id.to_string()))
|
||||
.await?
|
||||
@@ -208,6 +209,30 @@ impl<'a> AdminAppState<'a> {
|
||||
.into_response());
|
||||
};
|
||||
|
||||
let affected = self
|
||||
.cleanup_known_banned_provider_catalog_keys(&provider)
|
||||
.await?;
|
||||
if affected == 0 {
|
||||
return Ok(Json(
|
||||
aether_admin::provider::pool::build_admin_pool_cleanup_empty_payload(
|
||||
"未发现可清理的异常账号",
|
||||
),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
|
||||
Ok(
|
||||
Json(aether_admin::provider::pool::build_admin_pool_cleanup_result_payload(affected))
|
||||
.into_response(),
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_known_banned_provider_catalog_keys(
|
||||
&self,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Result<usize, GatewayError> {
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
|
||||
let banned_keys = self
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await?
|
||||
@@ -215,12 +240,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.filter(admin_provider_pool_pure::admin_pool_key_is_known_banned)
|
||||
.collect::<Vec<_>>();
|
||||
if banned_keys.is_empty() {
|
||||
return Ok(Json(
|
||||
admin_provider_pool_pure::build_admin_pool_cleanup_empty_payload(
|
||||
"未发现可清理的异常账号",
|
||||
),
|
||||
)
|
||||
.into_response());
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
let deleted_key_ids = banned_keys
|
||||
@@ -243,10 +263,42 @@ impl<'a> AdminAppState<'a> {
|
||||
self.cleanup_deleted_provider_catalog_refs(&provider.id, &[], &deleted_key_ids)
|
||||
.await?;
|
||||
|
||||
Ok(
|
||||
Json(admin_provider_pool_pure::build_admin_pool_cleanup_result_payload(affected))
|
||||
.into_response(),
|
||||
)
|
||||
Ok(affected)
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_provider_catalog_key_if_current<F>(
|
||||
&self,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
key_id: &str,
|
||||
should_delete: F,
|
||||
) -> Result<bool, GatewayError>
|
||||
where
|
||||
F: FnOnce(&StoredProviderCatalogKey) -> bool,
|
||||
{
|
||||
let key_ids = [key_id.to_string()];
|
||||
let Some(key) = self
|
||||
.read_provider_catalog_keys_by_ids(&key_ids)
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if key.provider_id != provider.id || !should_delete(&key) {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
self.clear_admin_provider_pool_cooldown(&provider.id, &key.id)
|
||||
.await;
|
||||
self.reset_admin_provider_pool_cost(&provider.id, &key.id)
|
||||
.await;
|
||||
let deleted = self.delete_provider_catalog_key(&key.id).await?;
|
||||
if deleted {
|
||||
let deleted_key_ids = [key.id.clone()];
|
||||
self.cleanup_deleted_provider_catalog_refs(&provider.id, &[], &deleted_key_ids)
|
||||
.await?;
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_pool_batch_action_response(
|
||||
|
||||
@@ -40,6 +40,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.map(|model| AdminSystemConfigGlobalModel {
|
||||
name: model.name.clone(),
|
||||
display_name: model.display_name.clone(),
|
||||
usage_count: Some(model.usage_count),
|
||||
default_price_per_request: model.default_price_per_request,
|
||||
default_tiered_pricing: model.default_tiered_pricing.clone(),
|
||||
supported_capabilities: model.supported_capabilities.as_ref().and_then(|value| {
|
||||
@@ -169,6 +170,13 @@ impl<'a> AdminAppState<'a> {
|
||||
) -> Result<serde_json::Value, GatewayError> {
|
||||
let users = self.list_non_admin_export_users().await?;
|
||||
let user_ids = users.iter().map(|user| user.id.clone()).collect::<Vec<_>>();
|
||||
let user_usage_totals = self
|
||||
.app
|
||||
.summarize_usage_totals_by_user_ids(&user_ids)
|
||||
.await?
|
||||
.into_iter()
|
||||
.map(|totals| (totals.user_id.clone(), totals))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let user_wallets = self.list_wallet_snapshots_by_user_ids(&user_ids).await?;
|
||||
let user_api_keys = self
|
||||
.list_auth_api_key_export_records_by_user_ids(&user_ids)
|
||||
@@ -185,6 +193,7 @@ impl<'a> AdminAppState<'a> {
|
||||
let standalone_wallets = self
|
||||
.list_wallet_snapshots_by_api_key_ids(&standalone_api_key_ids)
|
||||
.await?;
|
||||
let usage_aggregates = self.export_admin_system_usage_aggregates().await?;
|
||||
|
||||
let wallets_by_user_id = user_wallets
|
||||
.into_iter()
|
||||
@@ -260,8 +269,10 @@ impl<'a> AdminAppState<'a> {
|
||||
self.build_admin_system_users_export_api_key_payload(key, None, true)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let usage_totals = user_usage_totals.get(&user.id);
|
||||
|
||||
json!({
|
||||
"id": user.id.clone(),
|
||||
"email": user.email.clone(),
|
||||
"email_verified": user.email_verified,
|
||||
"username": user.username.clone(),
|
||||
@@ -284,6 +295,12 @@ impl<'a> AdminAppState<'a> {
|
||||
.unwrap_or(false),
|
||||
"wallet": wallet_payload,
|
||||
"is_active": user.is_active,
|
||||
"request_count": usage_totals
|
||||
.map(|totals| totals.request_count)
|
||||
.unwrap_or(0),
|
||||
"total_tokens": usage_totals
|
||||
.map(|totals| totals.total_tokens)
|
||||
.unwrap_or(0),
|
||||
"api_keys": api_keys_payload,
|
||||
})
|
||||
})
|
||||
@@ -306,6 +323,7 @@ impl<'a> AdminAppState<'a> {
|
||||
"user_groups": user_groups_data,
|
||||
"users": users_data,
|
||||
"standalone_keys": standalone_keys_data,
|
||||
"usage_aggregates": usage_aggregates,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -330,6 +348,7 @@ impl<'a> AdminAppState<'a> {
|
||||
include_is_standalone: bool,
|
||||
) -> serde_json::Value {
|
||||
let mut payload = serde_json::Map::from_iter([
|
||||
("api_key_id".to_string(), json!(key.api_key_id.clone())),
|
||||
("key_hash".to_string(), json!(key.key_hash.clone())),
|
||||
("name".to_string(), json!(key.name.clone())),
|
||||
(
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -70,6 +70,26 @@ impl<'a> AdminAppState<'a> {
|
||||
self.app.purge_admin_system_data(target).await
|
||||
}
|
||||
|
||||
pub(crate) async fn export_admin_system_usage_aggregates(
|
||||
&self,
|
||||
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateSnapshot, GatewayError>
|
||||
{
|
||||
self.app.export_admin_system_usage_aggregates().await
|
||||
}
|
||||
|
||||
pub(crate) async fn import_admin_system_usage_aggregates(
|
||||
&self,
|
||||
snapshot: &aether_data::repository::system::AdminSystemUsageAggregateSnapshot,
|
||||
user_id_map: &std::collections::BTreeMap<String, String>,
|
||||
api_key_id_map: &std::collections::BTreeMap<String, String>,
|
||||
mode: aether_data::repository::system::AdminSystemUsageAggregateImportMode,
|
||||
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateImportSummary, GatewayError>
|
||||
{
|
||||
self.app
|
||||
.import_admin_system_usage_aggregates(snapshot, user_id_map, api_key_id_map, mode)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn run_admin_system_cleanup_once(
|
||||
&self,
|
||||
) -> Result<crate::maintenance::AdminSystemCleanupSummary, GatewayError> {
|
||||
|
||||
@@ -706,6 +706,19 @@ impl<'a> AdminAppState<'a> {
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn set_api_key_usage_totals(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
total_requests: u64,
|
||||
total_tokens: u64,
|
||||
total_cost_usd: f64,
|
||||
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
|
||||
{
|
||||
self.app
|
||||
.set_api_key_usage_totals(api_key_id, total_requests, total_tokens, total_cost_usd)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_user_api_key(
|
||||
&self,
|
||||
user_id: &str,
|
||||
|
||||
@@ -13,7 +13,8 @@ pub(crate) use crate::handlers::shared::{
|
||||
effective_catalog_encryption_key, encrypt_catalog_secret_with_fallbacks, json_string_list,
|
||||
masked_catalog_api_key, normalize_json_array, normalize_json_object, normalize_string_list,
|
||||
parse_catalog_auth_config_json, provider_catalog_key_supports_format,
|
||||
provider_key_health_summary, provider_key_status_snapshot_payload, query_param_bool,
|
||||
query_param_optional_bool, query_param_value, take_secret_prefix, take_secret_suffix,
|
||||
unix_secs_to_rfc3339, OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
|
||||
provider_key_health_summary, provider_key_health_summary_at,
|
||||
provider_key_status_snapshot_payload, query_param_bool, query_param_optional_bool,
|
||||
query_param_value, take_secret_prefix, take_secret_suffix, unix_secs_to_rfc3339,
|
||||
OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
|
||||
};
|
||||
|
||||
@@ -13,10 +13,16 @@ use crate::handlers::admin::system::shared::paths::{
|
||||
};
|
||||
use crate::handlers::admin::system::shared::settings::{
|
||||
apply_admin_system_settings_update, build_admin_api_formats_payload,
|
||||
build_admin_system_check_update_payload_from_release, build_admin_system_settings_payload,
|
||||
build_admin_system_stats_payload, current_aether_version, fetch_latest_admin_system_release,
|
||||
build_admin_system_check_update_payload_from_release, build_admin_system_releases_list_payload,
|
||||
build_admin_system_settings_payload, build_admin_system_stats_payload, current_aether_version,
|
||||
fetch_admin_system_releases, fetch_latest_admin_system_release, resolve_update_target,
|
||||
};
|
||||
use crate::handlers::admin::system::shared::smtp::build_admin_smtp_test_payload;
|
||||
use crate::handlers::admin::system::shared::update::{
|
||||
build_admin_system_update_capability_payload, current_self_update_blocker,
|
||||
prepare_admin_system_update_task, read_update_history, read_update_task_status,
|
||||
self_update_supported, start_admin_system_rollback_task, start_admin_system_update_task,
|
||||
};
|
||||
use crate::important_notification::build_important_notification_test_payload;
|
||||
use crate::maintenance::{ManualUsageCleanupMode, ManualUsageCleanupOptions};
|
||||
use crate::GatewayError;
|
||||
@@ -58,7 +64,8 @@ pub(super) async fn maybe_build_local_admin_core_system_response(
|
||||
&& request_method == http::Method::GET
|
||||
&& request_path == "/api/admin/system/check-update"
|
||||
{
|
||||
let (latest_release, error) = fetch_latest_admin_system_release().await;
|
||||
let force = query_flag(request_context.query_string(), "force");
|
||||
let (latest_release, error) = fetch_latest_admin_system_release(force).await;
|
||||
return Ok(Some(
|
||||
Json(build_admin_system_check_update_payload_from_release(
|
||||
latest_release,
|
||||
@@ -68,6 +75,133 @@ pub(super) async fn maybe_build_local_admin_core_system_response(
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("releases")
|
||||
&& request_method == http::Method::GET
|
||||
&& request_path == "/api/admin/system/releases"
|
||||
{
|
||||
let force = query_flag(request_context.query_string(), "force");
|
||||
let (releases, error) = fetch_admin_system_releases(force).await;
|
||||
return Ok(Some(
|
||||
Json(build_admin_system_releases_list_payload(releases, error)).into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("update_capability")
|
||||
&& request_method == http::Method::GET
|
||||
&& request_path == "/api/admin/system/update-capability"
|
||||
{
|
||||
return Ok(Some(
|
||||
Json(build_admin_system_update_capability_payload()).into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("prepare_update")
|
||||
&& request_method == http::Method::POST
|
||||
&& request_path == "/api/admin/system/prepare-update"
|
||||
{
|
||||
if !self_update_supported() {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::PRECONDITION_REQUIRED,
|
||||
Json(json!({ "detail": current_self_update_blocker() })),
|
||||
)
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
let target_version = request_body
|
||||
.filter(|b| !b.is_empty())
|
||||
.and_then(|body| serde_json::from_slice::<serde_json::Value>(body).ok())
|
||||
.and_then(|v| v.get("version").and_then(|v| v.as_str().map(String::from)));
|
||||
|
||||
let (version, tarball_url, sha256sums_url) =
|
||||
match resolve_update_target(target_version).await {
|
||||
Ok(result) => result,
|
||||
Err((status, payload)) => {
|
||||
return Ok(Some((status, Json(payload)).into_response()));
|
||||
}
|
||||
};
|
||||
|
||||
return Ok(Some(
|
||||
match prepare_admin_system_update_task(version, tarball_url, sha256sums_url).await? {
|
||||
Ok(payload) => attach_admin_audit_response(
|
||||
Json(payload).into_response(),
|
||||
"admin_system_update_prepared",
|
||||
"prepare_system_update",
|
||||
"system_update",
|
||||
"global",
|
||||
),
|
||||
Err((status, payload)) => (status, Json(payload)).into_response(),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("apply_update")
|
||||
&& request_method == http::Method::POST
|
||||
&& request_path == "/api/admin/system/apply-update"
|
||||
{
|
||||
let version = request_body
|
||||
.filter(|b| !b.is_empty())
|
||||
.and_then(|body| serde_json::from_slice::<serde_json::Value>(body).ok())
|
||||
.and_then(|v| v.get("version").and_then(|v| v.as_str().map(String::from)));
|
||||
|
||||
return Ok(Some(
|
||||
match start_admin_system_update_task(version).await? {
|
||||
Ok(payload) => attach_admin_audit_response(
|
||||
Json(payload).into_response(),
|
||||
"admin_system_update_started",
|
||||
"apply_system_update",
|
||||
"system_update",
|
||||
"global",
|
||||
),
|
||||
Err((status, payload)) => (status, Json(payload)).into_response(),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("rollback")
|
||||
&& request_method == http::Method::POST
|
||||
&& request_path == "/api/admin/system/rollback"
|
||||
{
|
||||
return Ok(Some(match start_admin_system_rollback_task().await? {
|
||||
Ok(payload) => attach_admin_audit_response(
|
||||
Json(payload).into_response(),
|
||||
"admin_system_rollback_started",
|
||||
"rollback_system_update",
|
||||
"system_rollback",
|
||||
"global",
|
||||
),
|
||||
Err((status, payload)) => (status, Json(payload)).into_response(),
|
||||
}));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("update_status")
|
||||
&& request_method == http::Method::GET
|
||||
&& request_path == "/api/admin/system/update-status"
|
||||
{
|
||||
let status = read_update_task_status();
|
||||
return Ok(Some(
|
||||
Json(json!({
|
||||
"phase": status.phase,
|
||||
"error": status.error,
|
||||
"output": status.output,
|
||||
"progress_label": status.progress_label,
|
||||
"downloaded_bytes": status.downloaded_bytes,
|
||||
"total_bytes": status.total_bytes,
|
||||
"progress_percent": status.progress_percent,
|
||||
}))
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("update_history")
|
||||
&& request_method == http::Method::GET
|
||||
&& request_path == "/api/admin/system/update-history"
|
||||
{
|
||||
let entries = read_update_history();
|
||||
return Ok(Some(Json(json!({ "entries": entries })).into_response()));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("aws_regions")
|
||||
&& request_method == http::Method::GET
|
||||
&& request_path == "/api/admin/system/aws-regions"
|
||||
@@ -233,6 +367,39 @@ pub(super) async fn maybe_build_local_admin_core_system_response(
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("s3_backup_run")
|
||||
&& request_method == http::Method::POST
|
||||
&& request_path == "/api/admin/system/backups/s3/run"
|
||||
{
|
||||
return Ok(Some(
|
||||
match crate::backup::task::start_s3_backup_task(
|
||||
state.cloned_app(),
|
||||
"manual",
|
||||
decision
|
||||
.admin_principal
|
||||
.as_ref()
|
||||
.map(|principal| principal.user_id.as_str()),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(task) => attach_admin_audit_response(
|
||||
Json(json!({
|
||||
"message": "S3 备份任务已提交",
|
||||
"task": task,
|
||||
}))
|
||||
.into_response(),
|
||||
"admin_system_s3_backup_task_started",
|
||||
"run_s3_backup",
|
||||
"s3_backup",
|
||||
"global",
|
||||
),
|
||||
Err(error) => {
|
||||
(error.status(), Json(json!({ "detail": error.detail() }))).into_response()
|
||||
}
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("smtp_test")
|
||||
&& request_method == http::Method::POST
|
||||
&& request_path == "/api/admin/system/smtp/test"
|
||||
@@ -942,6 +1109,15 @@ fn query_param(query_string: Option<&str>, name: &str) -> Option<String> {
|
||||
.find_map(|(key, value)| (key == name && !value.is_empty()).then(|| value.into_owned()))
|
||||
}
|
||||
|
||||
fn query_flag(query_string: Option<&str>, name: &str) -> bool {
|
||||
query_param(query_string, name).is_some_and(|value| {
|
||||
matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
"1" | "true" | "yes" | "on"
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_older_than_days_query(query_string: Option<&str>) -> Result<Option<u32>, Response<Body>> {
|
||||
let Some(value) = query_param(query_string, "older_than_days") else {
|
||||
return Ok(None);
|
||||
|
||||
@@ -63,6 +63,10 @@ struct ProxyNodeRegisterRequest {
|
||||
proxy_version: Option<String>,
|
||||
#[serde(default)]
|
||||
tunnel_mode: Option<bool>,
|
||||
#[serde(default)]
|
||||
tunnel_security: Option<String>,
|
||||
#[serde(default)]
|
||||
tunnel_encryption_key: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -349,9 +353,21 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
|
||||
Ok(mutation) => mutation,
|
||||
Err(response) => return Ok(Some(response)),
|
||||
};
|
||||
let tunnel_encryption_key = mutation
|
||||
.proxy_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key"))
|
||||
.and_then(|value| value.as_str())
|
||||
.map(str::to_string);
|
||||
let Some(node) = state.register_proxy_node(&mutation).await? else {
|
||||
return Ok(Some(build_admin_proxy_nodes_data_unavailable_response()));
|
||||
};
|
||||
if let Some(key) = tunnel_encryption_key {
|
||||
state
|
||||
.app()
|
||||
.tunnel
|
||||
.register_secure_tunnel_key(node.id.clone(), key);
|
||||
}
|
||||
return Ok(Some(
|
||||
Json(json!({
|
||||
"node_id": node.id,
|
||||
@@ -1358,6 +1374,9 @@ fn build_tunnel_probe_relay_envelope(
|
||||
method: "GET".to_string(),
|
||||
url: probe_url.trim().to_string(),
|
||||
headers: std::collections::HashMap::new(),
|
||||
stream: false,
|
||||
request_timeout_ms: None,
|
||||
stream_first_byte_timeout_ms: None,
|
||||
timeout: timeout_secs,
|
||||
follow_redirects: Some(false),
|
||||
http1_only: false,
|
||||
@@ -1420,12 +1439,43 @@ fn validate_register_request(
|
||||
}
|
||||
validate_optional_object(input.hardware_info.as_ref(), "hardware_info")?;
|
||||
validate_optional_object(input.proxy_metadata.as_ref(), "proxy_metadata")?;
|
||||
let tunnel_security =
|
||||
normalize_optional_string(input.tunnel_security.as_deref(), "tunnel_security", 64)?;
|
||||
let tunnel_encryption_key = normalize_optional_string(
|
||||
input.tunnel_encryption_key.as_deref(),
|
||||
"tunnel_encryption_key",
|
||||
128,
|
||||
)?;
|
||||
|
||||
let registered_by = request_context
|
||||
.decision()
|
||||
.and_then(|decision| decision.admin_principal.as_ref())
|
||||
.map(|principal| principal.user_id.clone());
|
||||
|
||||
let mut proxy_metadata = input.proxy_metadata;
|
||||
if tunnel_security.as_deref()
|
||||
== Some(aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED)
|
||||
{
|
||||
let key = tunnel_encryption_key.as_deref().ok_or_else(|| {
|
||||
bad_request_response(
|
||||
"tunnel_encryption_key is required when tunnel_security=non_tls_required",
|
||||
)
|
||||
})?;
|
||||
aether_contracts::tunnel_security::decode_psk(key)
|
||||
.map_err(|err| bad_request_response(err.to_string()))?;
|
||||
let mut metadata = proxy_metadata
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default();
|
||||
metadata.insert(
|
||||
"tunnel_security".to_string(),
|
||||
json!({
|
||||
"mode": aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED,
|
||||
"encryption_key": key,
|
||||
}),
|
||||
);
|
||||
proxy_metadata = Some(Value::Object(metadata));
|
||||
}
|
||||
|
||||
Ok(
|
||||
aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation {
|
||||
name,
|
||||
@@ -1438,7 +1488,7 @@ fn validate_register_request(
|
||||
avg_latency_ms: input.avg_latency_ms,
|
||||
hardware_info: input.hardware_info,
|
||||
estimated_max_concurrency: input.estimated_max_concurrency,
|
||||
proxy_metadata: input.proxy_metadata,
|
||||
proxy_metadata,
|
||||
proxy_version: normalize_optional_string(
|
||||
input.proxy_version.as_deref(),
|
||||
"proxy_version",
|
||||
|
||||
@@ -4,3 +4,5 @@ pub(crate) mod modules;
|
||||
pub(crate) mod paths;
|
||||
pub(crate) mod settings;
|
||||
pub(crate) mod smtp;
|
||||
pub(crate) mod update;
|
||||
pub(crate) mod update_client;
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::backup::config::S3BackupConfig;
|
||||
use crate::bark_push::bark_push_configured;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::shared::{module_available_from_env, system_config_bool};
|
||||
@@ -112,7 +113,7 @@ pub(crate) const ADMIN_MODULE_DEFINITIONS: &[AdminModuleDefinition] = &[
|
||||
AdminModuleDefinition {
|
||||
name: "model_directives",
|
||||
display_name: "模型后缀参数",
|
||||
description: "允许通过模型名后缀覆盖推理参数",
|
||||
description: "允许通过模型名后缀覆盖推理参数或服务层级",
|
||||
category: "integration",
|
||||
env_key: "MODEL_DIRECTIVES_AVAILABLE",
|
||||
default_available: true,
|
||||
@@ -121,6 +122,18 @@ pub(crate) const ADMIN_MODULE_DEFINITIONS: &[AdminModuleDefinition] = &[
|
||||
admin_menu_group: None,
|
||||
admin_menu_order: 59,
|
||||
},
|
||||
AdminModuleDefinition {
|
||||
name: "s3_backup",
|
||||
display_name: "S3 备份",
|
||||
description: "将配置、用户或完整数据定期备份到 S3-compatible 对象存储",
|
||||
category: "integration",
|
||||
env_key: "S3_BACKUP_AVAILABLE",
|
||||
default_available: true,
|
||||
admin_route: Some("/admin/modules/s3-backup"),
|
||||
admin_menu_icon: Some("CloudUpload"),
|
||||
admin_menu_group: None,
|
||||
admin_menu_order: 60,
|
||||
},
|
||||
AdminModuleDefinition {
|
||||
name: "gemini_files",
|
||||
display_name: "文件缓存",
|
||||
@@ -183,6 +196,7 @@ pub(crate) struct AdminModuleRuntimeState {
|
||||
important_notification_configured: bool,
|
||||
server_chan_push_configured: bool,
|
||||
bark_push_configured: bool,
|
||||
s3_backup_configured: bool,
|
||||
}
|
||||
|
||||
pub(crate) fn admin_module_by_name(name: &str) -> Option<&'static AdminModuleDefinition> {
|
||||
@@ -209,6 +223,8 @@ pub(crate) fn admin_module_enabled_config_key(module: &AdminModuleDefinition) ->
|
||||
ENABLE_MODEL_DIRECTIVES_CONFIG_KEY.to_string()
|
||||
} else if module.name == "important_notification" {
|
||||
IMPORTANT_NOTIFICATION_ENABLED_KEY.to_string()
|
||||
} else if module.name == "s3_backup" {
|
||||
crate::backup::S3_BACKUP_ENABLED_KEY.to_string()
|
||||
} else {
|
||||
format!("module.{}.enabled", module.name)
|
||||
}
|
||||
@@ -272,6 +288,7 @@ pub(crate) async fn build_admin_module_runtime_state(
|
||||
let notification_configured = important_notification_configured(state.app()).await?;
|
||||
let server_chan_configured = server_chan_push_configured(state.app()).await?;
|
||||
let bark_configured = bark_push_configured(state.app()).await?;
|
||||
let backup_configured = s3_backup_configured(state.app()).await;
|
||||
|
||||
Ok(AdminModuleRuntimeState {
|
||||
oauth_providers,
|
||||
@@ -280,21 +297,36 @@ pub(crate) async fn build_admin_module_runtime_state(
|
||||
important_notification_configured: notification_configured,
|
||||
server_chan_push_configured: server_chan_configured,
|
||||
bark_push_configured: bark_configured,
|
||||
s3_backup_configured: backup_configured,
|
||||
})
|
||||
}
|
||||
|
||||
async fn s3_backup_configured(app: &crate::AppState) -> bool {
|
||||
let Ok(mut values) = crate::backup::task::load_s3_backup_config_values(app).await else {
|
||||
return false;
|
||||
};
|
||||
values.insert(
|
||||
crate::backup::S3_BACKUP_ENABLED_KEY.to_string(),
|
||||
json!(true),
|
||||
);
|
||||
S3BackupConfig::from_json_map(&values).is_ok()
|
||||
}
|
||||
|
||||
pub(crate) fn build_admin_module_validation_result(
|
||||
module: &AdminModuleDefinition,
|
||||
runtime: &AdminModuleRuntimeState,
|
||||
) -> (bool, Option<String>) {
|
||||
admin_system_kernel::build_admin_module_validation_result(
|
||||
module.name,
|
||||
&runtime.oauth_providers,
|
||||
runtime.ldap_config.as_ref(),
|
||||
runtime.gemini_files_has_capable_key,
|
||||
runtime.important_notification_configured,
|
||||
runtime.server_chan_push_configured,
|
||||
runtime.bark_push_configured,
|
||||
admin_system_kernel::AdminModuleValidationInput {
|
||||
module_name: module.name,
|
||||
oauth_providers: &runtime.oauth_providers,
|
||||
ldap_config: runtime.ldap_config.as_ref(),
|
||||
gemini_files_has_capable_key: runtime.gemini_files_has_capable_key,
|
||||
important_notification_configured: runtime.important_notification_configured,
|
||||
server_chan_push_configured: runtime.server_chan_push_configured,
|
||||
bark_push_configured: runtime.bark_push_configured,
|
||||
s3_backup_configured: runtime.s3_backup_configured,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,11 +1,18 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::build_admin_usage_counter_health_payload;
|
||||
use crate::handlers::admin::system::shared::update::{
|
||||
current_self_update_blocker, self_update_supported,
|
||||
};
|
||||
use crate::handlers::admin::system::shared::update_client::{
|
||||
build_direct_update_http_client, build_update_http_client, has_explicit_update_proxy_env,
|
||||
update_github_token_from_env,
|
||||
};
|
||||
use crate::handlers::shared::{system_config_bool, system_config_string};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::system::{
|
||||
build_admin_api_formats_payload as build_admin_api_formats_payload_pure,
|
||||
build_admin_system_check_update_payload as build_admin_system_check_update_payload_pure,
|
||||
build_admin_system_check_update_payload_with_release,
|
||||
build_admin_system_check_update_payload_with_release, build_admin_system_releases_payload,
|
||||
build_admin_system_settings_payload as build_admin_system_settings_payload_pure,
|
||||
build_admin_system_settings_updated_payload,
|
||||
build_admin_system_stats_payload as build_admin_system_stats_payload_pure,
|
||||
@@ -15,12 +22,18 @@ use axum::body::Bytes;
|
||||
use axum::http;
|
||||
#[cfg(not(test))]
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
#[cfg(not(test))]
|
||||
use serde_json::{json, Value};
|
||||
use std::time::Duration;
|
||||
|
||||
const AETHER_RELEASES_API_URL: &str =
|
||||
"https://api.github.com/repos/fawney19/Aether/releases?per_page=20";
|
||||
const AETHER_RELEASE_TAG_URL_BASE: &str = "https://github.com/fawney19/Aether/releases/tag";
|
||||
const SOURCE_BUILD_UPDATE_BLOCKER: &str = "当前为源码构建,请使用 git pull 后重新编译。";
|
||||
const SOURCE_BUILD_RELEASE_BLOCKER: &str = "当前为源码构建,请手动切换到对应标签后重新编译。";
|
||||
|
||||
/// Minimum interval between actual GitHub API requests. Within this window
|
||||
/// the cached result is reused.
|
||||
const RELEASE_CACHE_TTL: Duration = Duration::from_secs(1200);
|
||||
|
||||
pub(crate) fn current_aether_version() -> String {
|
||||
option_env!("AETHER_BUILD_VERSION")
|
||||
@@ -37,56 +50,424 @@ pub(crate) fn build_admin_system_check_update_payload_from_release(
|
||||
latest_release: Option<AdminSystemUpdateRelease>,
|
||||
error: Option<String>,
|
||||
) -> serde_json::Value {
|
||||
build_admin_system_check_update_payload_with_release(
|
||||
let mut payload = build_admin_system_check_update_payload_with_release(
|
||||
current_aether_version(),
|
||||
latest_release,
|
||||
error,
|
||||
)
|
||||
);
|
||||
apply_self_update_check_update_override(&mut payload, self_update_supported());
|
||||
payload
|
||||
}
|
||||
|
||||
pub(crate) fn build_admin_system_releases_list_payload(
|
||||
releases: Vec<AdminSystemUpdateRelease>,
|
||||
error: Option<String>,
|
||||
) -> serde_json::Value {
|
||||
let mut payload =
|
||||
build_admin_system_releases_payload(current_aether_version(), releases, error);
|
||||
apply_self_update_releases_override(&mut payload, self_update_supported());
|
||||
payload
|
||||
}
|
||||
|
||||
fn current_build_is_release() -> bool {
|
||||
option_env!("AETHER_BUILD_TYPE").unwrap_or("source") == "release"
|
||||
}
|
||||
|
||||
fn apply_self_update_check_update_override(payload: &mut Value, supported: bool) {
|
||||
apply_self_update_check_update_override_with_blocker(
|
||||
payload,
|
||||
supported,
|
||||
current_self_update_blocker(),
|
||||
);
|
||||
}
|
||||
|
||||
fn apply_self_update_check_update_override_with_blocker(
|
||||
payload: &mut Value,
|
||||
supported: bool,
|
||||
blocker: &str,
|
||||
) {
|
||||
if supported {
|
||||
return;
|
||||
}
|
||||
if payload.get("has_update").and_then(Value::as_bool) != Some(true) {
|
||||
return;
|
||||
}
|
||||
|
||||
payload["updatable"] = json!(false);
|
||||
payload["update_blocker"] = json!(blocker);
|
||||
}
|
||||
|
||||
fn apply_self_update_releases_override(payload: &mut Value, supported: bool) {
|
||||
apply_self_update_releases_override_with_blocker(
|
||||
payload,
|
||||
supported,
|
||||
current_self_update_release_blocker(),
|
||||
);
|
||||
}
|
||||
|
||||
fn apply_self_update_releases_override_with_blocker(
|
||||
payload: &mut Value,
|
||||
supported: bool,
|
||||
blocker: &str,
|
||||
) {
|
||||
if supported {
|
||||
return;
|
||||
}
|
||||
let Some(releases) = payload.get_mut("releases").and_then(Value::as_array_mut) else {
|
||||
return;
|
||||
};
|
||||
|
||||
for release in releases {
|
||||
if release.get("is_current").and_then(Value::as_bool) == Some(true) {
|
||||
continue;
|
||||
}
|
||||
release["updatable"] = json!(false);
|
||||
if release.get("update_blocker").is_none() || release["update_blocker"].is_null() {
|
||||
release["update_blocker"] = json!(blocker);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn current_self_update_release_blocker() -> &'static str {
|
||||
if !current_build_is_release() {
|
||||
return SOURCE_BUILD_RELEASE_BLOCKER;
|
||||
}
|
||||
|
||||
if self_update_supported() {
|
||||
""
|
||||
} else {
|
||||
current_self_update_blocker()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
struct CachedReleases {
|
||||
all: Vec<AdminSystemUpdateRelease>,
|
||||
latest: Option<AdminSystemUpdateRelease>,
|
||||
error: Option<String>,
|
||||
fetched_at: std::time::Instant,
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
fn releases_cache() -> &'static std::sync::Mutex<Option<CachedReleases>> {
|
||||
static CACHE: std::sync::OnceLock<std::sync::Mutex<Option<CachedReleases>>> =
|
||||
std::sync::OnceLock::new();
|
||||
CACHE.get_or_init(|| std::sync::Mutex::new(None))
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
async fn ensure_releases_cached(force: bool) {
|
||||
{
|
||||
if let Ok(guard) = releases_cache().lock() {
|
||||
if let Some(cached) = guard.as_ref() {
|
||||
if should_reuse_releases_cache(
|
||||
force,
|
||||
cached.error.is_some(),
|
||||
cached.fetched_at.elapsed(),
|
||||
) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let (all, latest, error) = match fetch_admin_system_releases_inner().await {
|
||||
Ok(releases) => {
|
||||
let latest = releases.first().cloned();
|
||||
(releases, latest, None)
|
||||
}
|
||||
Err(err) => (Vec::new(), None, Some(err)),
|
||||
};
|
||||
|
||||
if let Ok(mut guard) = releases_cache().lock() {
|
||||
*guard = Some(CachedReleases {
|
||||
all,
|
||||
latest,
|
||||
error,
|
||||
fetched_at: std::time::Instant::now(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
fn should_reuse_releases_cache(force: bool, has_error: bool, age: Duration) -> bool {
|
||||
!force && !has_error && age < RELEASE_CACHE_TTL
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
pub(crate) async fn fetch_latest_admin_system_release(
|
||||
force: bool,
|
||||
) -> (Option<AdminSystemUpdateRelease>, Option<String>) {
|
||||
match fetch_latest_admin_system_release_inner().await {
|
||||
Ok(release) => (release, None),
|
||||
Err(err) => (None, Some(err)),
|
||||
ensure_releases_cached(force).await;
|
||||
if let Ok(guard) = releases_cache().lock() {
|
||||
if let Some(cached) = guard.as_ref() {
|
||||
return (cached.latest.clone(), cached.error.clone());
|
||||
}
|
||||
}
|
||||
(None, None)
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
pub(crate) async fn fetch_admin_system_releases(
|
||||
force: bool,
|
||||
) -> (Vec<AdminSystemUpdateRelease>, Option<String>) {
|
||||
ensure_releases_cached(force).await;
|
||||
if let Ok(guard) = releases_cache().lock() {
|
||||
if let Some(cached) = guard.as_ref() {
|
||||
return (cached.all.clone(), cached.error.clone());
|
||||
}
|
||||
}
|
||||
(Vec::new(), None)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) async fn fetch_admin_system_releases(
|
||||
_force: bool,
|
||||
) -> (Vec<AdminSystemUpdateRelease>, Option<String>) {
|
||||
(
|
||||
Vec::new(),
|
||||
Some("测试环境未请求 GitHub Releases".to_string()),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) async fn fetch_latest_admin_system_release(
|
||||
_force: bool,
|
||||
) -> (Option<AdminSystemUpdateRelease>, Option<String>) {
|
||||
(None, Some("测试环境未请求 GitHub Releases".to_string()))
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
async fn fetch_latest_admin_system_release_inner(
|
||||
) -> Result<Option<AdminSystemUpdateRelease>, String> {
|
||||
let releases: Vec<GitHubRelease> = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(8))
|
||||
.build()
|
||||
.map_err(|err| format!("创建更新检查客户端失败: {err}"))?
|
||||
.get(AETHER_RELEASES_API_URL)
|
||||
.header(reqwest::header::USER_AGENT, "Aether-Gateway update-check")
|
||||
.header(reqwest::header::ACCEPT, "application/vnd.github+json")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| format!("请求 GitHub Releases 失败: {err}"))?
|
||||
.error_for_status()
|
||||
.map_err(|err| format!("GitHub Releases 返回错误: {err}"))?
|
||||
.json()
|
||||
.await
|
||||
.map_err(|err| format!("解析 GitHub Releases 失败: {err}"))?;
|
||||
pub(crate) async fn resolve_update_target(
|
||||
version: Option<String>,
|
||||
) -> Result<(String, String, Option<String>), (http::StatusCode, serde_json::Value)> {
|
||||
let (releases, error) = fetch_admin_system_releases(false).await;
|
||||
if releases.is_empty() {
|
||||
return Err((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
json!({ "detail": error.unwrap_or_else(|| "无法获取版本信息".to_string()) }),
|
||||
));
|
||||
}
|
||||
|
||||
let release = match version {
|
||||
Some(ref v) => releases.into_iter().find(|r| r.version == *v),
|
||||
None => releases.into_iter().next(),
|
||||
};
|
||||
|
||||
let release = release.ok_or_else(|| {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
json!({ "detail": "未找到指定版本" }),
|
||||
)
|
||||
})?;
|
||||
|
||||
let tarball_url = release.tarball_url.ok_or_else(|| {
|
||||
(
|
||||
http::StatusCode::PRECONDITION_REQUIRED,
|
||||
json!({ "detail": format!("版本 {} 没有适用于当前平台的安装包", release.version) }),
|
||||
)
|
||||
})?;
|
||||
|
||||
let sha256sums_url = release.sha256sums_url.ok_or_else(|| {
|
||||
(
|
||||
http::StatusCode::PRECONDITION_REQUIRED,
|
||||
json!({ "detail": format!("版本 {} 缺少 SHA256SUMS 校验文件,已拒绝在线更新", release.version) }),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok((release.version, tarball_url, Some(sha256sums_url)))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) async fn resolve_update_target(
|
||||
_version: Option<String>,
|
||||
) -> Result<(String, String, Option<String>), (http::StatusCode, serde_json::Value)> {
|
||||
Err((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
json!({ "detail": "测试环境不支持更新" }),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
async fn fetch_admin_system_releases_inner() -> Result<Vec<AdminSystemUpdateRelease>, String> {
|
||||
let current_channel = update_channel_for_version(¤t_aether_version());
|
||||
let timeout = Duration::from_secs(8);
|
||||
let github_token = update_github_token_from_env();
|
||||
let client = build_update_http_client(timeout, "更新检查")?;
|
||||
let releases = match fetch_github_releases_with_client(&client, github_token.as_deref()).await {
|
||||
Ok(releases) => releases,
|
||||
Err(err) if err.rate_limited && !has_explicit_update_proxy_env() => {
|
||||
let direct_client = build_direct_update_http_client(timeout, "更新检查直连重试")?;
|
||||
fetch_github_releases_with_client(&direct_client, github_token.as_deref())
|
||||
.await
|
||||
.map_err(|retry_err| retry_err.message)?
|
||||
}
|
||||
Err(err) => return Err(err.message),
|
||||
};
|
||||
|
||||
Ok(releases
|
||||
.into_iter()
|
||||
.find(|release| !release.draft && release.tag_name.starts_with('v'))
|
||||
.map(|release| AdminSystemUpdateRelease {
|
||||
version: release.tag_name,
|
||||
release_url: Some(release.html_url),
|
||||
release_notes: release.body.filter(|body| !body.trim().is_empty()),
|
||||
published_at: release.published_at,
|
||||
}))
|
||||
.filter(|release| should_include_release_for_channel(release, current_channel))
|
||||
.map(|release| {
|
||||
let (tarball_url, sha256sums_url) = select_release_tarball_urls(&release);
|
||||
AdminSystemUpdateRelease {
|
||||
version: release.tag_name.clone(),
|
||||
release_url: Some(github_release_tag_url(&release.tag_name)),
|
||||
release_notes: release.body.filter(|body| !body.trim().is_empty()),
|
||||
published_at: release.published_at,
|
||||
tarball_url,
|
||||
sha256sums_url,
|
||||
}
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct GitHubReleaseFetchError {
|
||||
message: String,
|
||||
rate_limited: bool,
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
async fn fetch_github_releases_with_client(
|
||||
client: &reqwest::Client,
|
||||
github_token: Option<&str>,
|
||||
) -> Result<Vec<GitHubRelease>, GitHubReleaseFetchError> {
|
||||
let mut request = client
|
||||
.get(AETHER_RELEASES_API_URL)
|
||||
.header(reqwest::header::USER_AGENT, "Aether-Gateway update-check")
|
||||
.header(reqwest::header::ACCEPT, "application/vnd.github+json");
|
||||
if let Some(token) = github_token {
|
||||
request = request.bearer_auth(token);
|
||||
}
|
||||
|
||||
let response = request
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| GitHubReleaseFetchError {
|
||||
message: format!("请求 GitHub Releases 失败: {err}"),
|
||||
rate_limited: false,
|
||||
})?;
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
return Err(github_release_response_error(status, &body));
|
||||
}
|
||||
|
||||
response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|err| GitHubReleaseFetchError {
|
||||
message: format!("解析 GitHub Releases 失败: {err}"),
|
||||
rate_limited: false,
|
||||
})
|
||||
}
|
||||
|
||||
fn github_release_response_error(
|
||||
status: reqwest::StatusCode,
|
||||
response_body: &str,
|
||||
) -> GitHubReleaseFetchError {
|
||||
if is_github_rate_limit_error(status, response_body) {
|
||||
return GitHubReleaseFetchError {
|
||||
message: "GitHub Releases API 已触发限流;当前共享代理出口的匿名额度已用尽,更新检查将自动尝试直连。若仍失败,请配置 AETHER_UPDATE_GITHUB_TOKEN / GITHUB_TOKEN / GH_TOKEN,或为 GitHub 更新检查单独设置可用代理。".to_string(),
|
||||
rate_limited: true,
|
||||
};
|
||||
}
|
||||
|
||||
let detail = parse_github_error_message(response_body);
|
||||
let message = if detail.is_empty() {
|
||||
format!("GitHub Releases 返回错误: HTTP {status}")
|
||||
} else {
|
||||
format!("GitHub Releases 返回错误: HTTP {status}; {detail}")
|
||||
};
|
||||
GitHubReleaseFetchError {
|
||||
message,
|
||||
rate_limited: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_github_rate_limit_error(status: reqwest::StatusCode, response_body: &str) -> bool {
|
||||
status == reqwest::StatusCode::FORBIDDEN
|
||||
&& response_body
|
||||
.to_ascii_lowercase()
|
||||
.contains("rate limit exceeded")
|
||||
}
|
||||
|
||||
fn parse_github_error_message(response_body: &str) -> String {
|
||||
serde_json::from_str::<serde_json::Value>(response_body)
|
||||
.ok()
|
||||
.and_then(|value| {
|
||||
value
|
||||
.get("message")
|
||||
.and_then(|message| message.as_str())
|
||||
.map(str::trim)
|
||||
.filter(|message| !message.is_empty())
|
||||
.map(ToString::to_string)
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
fn should_include_release_for_channel(
|
||||
release: &GitHubRelease,
|
||||
current_channel: UpdateChannel,
|
||||
) -> bool {
|
||||
if release.draft || !release.tag_name.starts_with('v') {
|
||||
return false;
|
||||
}
|
||||
if !release.prerelease {
|
||||
return true;
|
||||
}
|
||||
current_channel.allows_prerelease(&release.tag_name)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum UpdateChannel {
|
||||
Stable,
|
||||
Rc,
|
||||
Beta,
|
||||
OtherPrerelease,
|
||||
}
|
||||
|
||||
impl UpdateChannel {
|
||||
fn allows_prerelease(self, release_version: &str) -> bool {
|
||||
match self {
|
||||
Self::Stable => false,
|
||||
Self::Rc => update_channel_for_version(release_version) == Self::Rc,
|
||||
Self::Beta => update_channel_for_version(release_version) == Self::Beta,
|
||||
Self::OtherPrerelease => {
|
||||
matches!(
|
||||
update_channel_for_version(release_version),
|
||||
Self::Rc | Self::Beta | Self::OtherPrerelease
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn update_channel_for_version(version: &str) -> UpdateChannel {
|
||||
let normalized = version
|
||||
.trim()
|
||||
.strip_prefix('v')
|
||||
.or_else(|| version.trim().strip_prefix('V'))
|
||||
.unwrap_or(version.trim());
|
||||
let Some((_, prerelease)) = normalized.split_once('-') else {
|
||||
return UpdateChannel::Stable;
|
||||
};
|
||||
let prerelease = prerelease.to_ascii_lowercase();
|
||||
if prerelease.starts_with("rc") {
|
||||
UpdateChannel::Rc
|
||||
} else if prerelease.starts_with("beta") {
|
||||
UpdateChannel::Beta
|
||||
} else {
|
||||
UpdateChannel::OtherPrerelease
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct GitHubReleaseAsset {
|
||||
name: String,
|
||||
browser_download_url: String,
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
@@ -100,6 +481,43 @@ struct GitHubRelease {
|
||||
published_at: Option<String>,
|
||||
#[serde(default)]
|
||||
draft: bool,
|
||||
#[serde(default)]
|
||||
prerelease: bool,
|
||||
#[serde(default)]
|
||||
assets: Vec<GitHubReleaseAsset>,
|
||||
}
|
||||
|
||||
fn github_release_tag_url(tag_name: &str) -> String {
|
||||
format!("{AETHER_RELEASE_TAG_URL_BASE}/{tag_name}")
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
fn select_release_tarball_urls(release: &GitHubRelease) -> (Option<String>, Option<String>) {
|
||||
let platform = if cfg!(target_os = "macos") {
|
||||
"macos"
|
||||
} else {
|
||||
"linux"
|
||||
};
|
||||
let arch = if cfg!(target_arch = "aarch64") {
|
||||
"arm64"
|
||||
} else {
|
||||
"amd64"
|
||||
};
|
||||
let expected_name = format!("aether-{}-{}-{}.tar.gz", release.tag_name, platform, arch);
|
||||
|
||||
let tarball_url = release
|
||||
.assets
|
||||
.iter()
|
||||
.find(|a| a.name == expected_name)
|
||||
.map(|a| a.browser_download_url.clone());
|
||||
|
||||
let sha256sums_url = release
|
||||
.assets
|
||||
.iter()
|
||||
.find(|a| a.name == "SHA256SUMS")
|
||||
.map(|a| a.browser_download_url.clone());
|
||||
|
||||
(tarball_url, sha256sums_url)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_system_stats_payload(
|
||||
@@ -249,3 +667,145 @@ pub(crate) async fn apply_admin_system_settings_update(
|
||||
pub(crate) fn build_admin_api_formats_payload() -> serde_json::Value {
|
||||
build_admin_api_formats_payload_pure()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn update_channel_detects_stable_rc_beta_and_other_prerelease() {
|
||||
assert_eq!(update_channel_for_version("v1.2.3"), UpdateChannel::Stable);
|
||||
assert_eq!(update_channel_for_version("1.2.3-rc1"), UpdateChannel::Rc);
|
||||
assert_eq!(
|
||||
update_channel_for_version("1.2.3-beta.2"),
|
||||
UpdateChannel::Beta
|
||||
);
|
||||
assert_eq!(
|
||||
update_channel_for_version("1.2.3-alpha.1"),
|
||||
UpdateChannel::OtherPrerelease
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stable_channel_does_not_allow_prereleases() {
|
||||
assert!(!UpdateChannel::Stable.allows_prerelease("v1.2.3-rc1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prerelease_channels_only_follow_matching_channel() {
|
||||
assert!(UpdateChannel::Rc.allows_prerelease("v1.2.3-rc2"));
|
||||
assert!(!UpdateChannel::Rc.allows_prerelease("v1.2.3-beta.1"));
|
||||
assert!(UpdateChannel::Beta.allows_prerelease("v1.2.3-beta.2"));
|
||||
assert!(!UpdateChannel::Beta.allows_prerelease("v1.2.3-rc1"));
|
||||
assert!(UpdateChannel::OtherPrerelease.allows_prerelease("v1.2.3-alpha.1"));
|
||||
assert!(UpdateChannel::OtherPrerelease.allows_prerelease("v1.2.3-rc1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn github_release_tag_url_points_to_explicit_tag_page() {
|
||||
assert_eq!(
|
||||
github_release_tag_url("v0.7.3"),
|
||||
"https://github.com/fawney19/Aether/releases/tag/v0.7.3"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn successful_release_cache_is_reused_within_ttl() {
|
||||
assert!(should_reuse_releases_cache(
|
||||
false,
|
||||
false,
|
||||
Duration::from_secs(60)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn failed_release_cache_is_not_reused() {
|
||||
assert!(!should_reuse_releases_cache(
|
||||
false,
|
||||
true,
|
||||
Duration::from_secs(60)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn force_refresh_bypasses_release_cache() {
|
||||
assert!(!should_reuse_releases_cache(
|
||||
true,
|
||||
false,
|
||||
Duration::from_secs(60)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn github_rate_limit_response_is_marked_retryable() {
|
||||
let err = github_release_response_error(
|
||||
reqwest::StatusCode::FORBIDDEN,
|
||||
r#"{"message":"API rate limit exceeded for 1.2.3.4."}"#,
|
||||
);
|
||||
assert!(err.rate_limited);
|
||||
assert!(err.message.contains("GitHub Releases API 已触发限流"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_rate_limit_github_response_keeps_http_status_message() {
|
||||
let err = github_release_response_error(
|
||||
reqwest::StatusCode::FORBIDDEN,
|
||||
r#"{"message":"Resource not accessible"}"#,
|
||||
);
|
||||
assert!(!err.rate_limited);
|
||||
assert!(err.message.contains("HTTP 403 Forbidden"));
|
||||
assert!(err.message.contains("Resource not accessible"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn self_update_check_update_override_marks_latest_release_non_updatable() {
|
||||
let mut payload = json!({
|
||||
"has_update": true,
|
||||
"updatable": true,
|
||||
"update_blocker": serde_json::Value::Null
|
||||
});
|
||||
|
||||
apply_self_update_check_update_override_with_blocker(
|
||||
&mut payload,
|
||||
false,
|
||||
SOURCE_BUILD_UPDATE_BLOCKER,
|
||||
);
|
||||
|
||||
assert_eq!(payload["updatable"], false);
|
||||
assert_eq!(payload["update_blocker"], SOURCE_BUILD_UPDATE_BLOCKER);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn self_update_releases_override_marks_non_current_entries_non_updatable() {
|
||||
let mut payload = json!({
|
||||
"releases": [
|
||||
{
|
||||
"version": "v0.7.3",
|
||||
"is_current": false,
|
||||
"updatable": true,
|
||||
"update_blocker": serde_json::Value::Null
|
||||
},
|
||||
{
|
||||
"version": "v0.7.2",
|
||||
"is_current": true,
|
||||
"updatable": true,
|
||||
"update_blocker": "当前版本"
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
apply_self_update_releases_override_with_blocker(
|
||||
&mut payload,
|
||||
false,
|
||||
SOURCE_BUILD_RELEASE_BLOCKER,
|
||||
);
|
||||
|
||||
assert_eq!(payload["releases"][0]["updatable"], false);
|
||||
assert_eq!(
|
||||
payload["releases"][0]["update_blocker"],
|
||||
SOURCE_BUILD_RELEASE_BLOCKER
|
||||
);
|
||||
assert_eq!(payload["releases"][1]["updatable"], true);
|
||||
assert_eq!(payload["releases"][1]["update_blocker"], "当前版本");
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,68 @@
|
||||
use std::time::Duration;
|
||||
|
||||
const EXPLICIT_UPDATE_PROXY_ENV_KEYS: &[&str] = &["AETHER_UPDATE_PROXY_URL", "UPDATE_PROXY_URL"];
|
||||
|
||||
const UPDATE_PROXY_ENV_KEYS: &[&str] = &[
|
||||
"AETHER_UPDATE_PROXY_URL",
|
||||
"UPDATE_PROXY_URL",
|
||||
"HTTPS_PROXY",
|
||||
"https_proxy",
|
||||
"ALL_PROXY",
|
||||
"all_proxy",
|
||||
"HTTP_PROXY",
|
||||
"http_proxy",
|
||||
];
|
||||
|
||||
const UPDATE_GITHUB_TOKEN_ENV_KEYS: &[&str] =
|
||||
&["AETHER_UPDATE_GITHUB_TOKEN", "GITHUB_TOKEN", "GH_TOKEN"];
|
||||
|
||||
pub(crate) fn build_update_http_client(
|
||||
timeout: Duration,
|
||||
label: &str,
|
||||
) -> Result<reqwest::Client, String> {
|
||||
let mut builder = base_update_http_client_builder(timeout);
|
||||
if let Some(proxy_url) = update_proxy_url_from_env() {
|
||||
let proxy = reqwest::Proxy::all(proxy_url)
|
||||
.map_err(|_| format!("创建{label}代理失败,请检查更新代理环境变量"))?
|
||||
.no_proxy(reqwest::NoProxy::from_env());
|
||||
builder = builder.proxy(proxy);
|
||||
}
|
||||
builder
|
||||
.build()
|
||||
.map_err(|err| format!("创建{label}客户端失败: {err}"))
|
||||
}
|
||||
|
||||
pub(crate) fn build_direct_update_http_client(
|
||||
timeout: Duration,
|
||||
label: &str,
|
||||
) -> Result<reqwest::Client, String> {
|
||||
base_update_http_client_builder(timeout)
|
||||
.no_proxy()
|
||||
.build()
|
||||
.map_err(|err| format!("创建{label}客户端失败: {err}"))
|
||||
}
|
||||
|
||||
pub(crate) fn has_explicit_update_proxy_env() -> bool {
|
||||
read_nonempty_env_value(EXPLICIT_UPDATE_PROXY_ENV_KEYS).is_some()
|
||||
}
|
||||
|
||||
fn base_update_http_client_builder(timeout: Duration) -> reqwest::ClientBuilder {
|
||||
reqwest::Client::builder().timeout(timeout)
|
||||
}
|
||||
|
||||
fn update_proxy_url_from_env() -> Option<String> {
|
||||
read_nonempty_env_value(UPDATE_PROXY_ENV_KEYS)
|
||||
}
|
||||
|
||||
pub(crate) fn update_github_token_from_env() -> Option<String> {
|
||||
read_nonempty_env_value(UPDATE_GITHUB_TOKEN_ENV_KEYS)
|
||||
}
|
||||
|
||||
fn read_nonempty_env_value(keys: &[&str]) -> Option<String> {
|
||||
keys.iter().find_map(|key| {
|
||||
std::env::var(key)
|
||||
.ok()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
}
|
||||
@@ -417,6 +417,7 @@ async fn resolve_admin_user_selection(
|
||||
group_id: filters
|
||||
.as_ref()
|
||||
.and_then(|filters| filters.group_id.clone()),
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.map_err(|_| "用户数据不可用".to_string())?
|
||||
|
||||
@@ -37,6 +37,13 @@ pub(in super::super) async fn build_admin_list_users_response(
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
let sort_by = query_param_value(request_context.query_string(), "sort_by")
|
||||
.and_then(|value| aether_data::repository::users::UserExportSortBy::parse(&value))
|
||||
.unwrap_or_default();
|
||||
let sort_order = query_param_value(request_context.query_string(), "sort_order")
|
||||
.and_then(|value| aether_data::repository::users::UserExportSortOrder::parse(&value))
|
||||
.unwrap_or_default();
|
||||
|
||||
let query = aether_data::repository::users::UserExportListQuery {
|
||||
skip,
|
||||
limit,
|
||||
@@ -44,6 +51,8 @@ pub(in super::super) async fn build_admin_list_users_response(
|
||||
is_active,
|
||||
search,
|
||||
group_id,
|
||||
sort_by,
|
||||
sort_order,
|
||||
};
|
||||
let (paged_rows_result, total_result) = tokio::join!(
|
||||
state.list_export_users_page(&query),
|
||||
|
||||
@@ -152,12 +152,7 @@ pub(in super::super) async fn build_admin_update_user_response(
|
||||
};
|
||||
let effective_role = role.as_deref().unwrap_or(existing_user.role.as_str());
|
||||
let group_ids = if field_presence.contains("group_ids") {
|
||||
let requested_group_ids = normalize_admin_user_group_ids(payload.group_ids);
|
||||
Some(
|
||||
state
|
||||
.include_default_user_group_ids_for_role(&requested_group_ids, effective_role)
|
||||
.await?,
|
||||
)
|
||||
Some(normalize_admin_user_group_ids(payload.group_ids))
|
||||
} else if role.is_some() {
|
||||
let requested_group_ids = state
|
||||
.list_user_groups_for_user(&user_id)
|
||||
|
||||
@@ -354,12 +354,7 @@ pub(crate) async fn maybe_build_internal_finalize_video_response(
|
||||
}
|
||||
|
||||
pub(crate) fn gateway_error_message(error: GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
error.into_message()
|
||||
}
|
||||
|
||||
pub(crate) fn build_internal_tunnel_heartbeat_ack(
|
||||
|
||||
@@ -21,13 +21,13 @@ use crate::constants::{
|
||||
EXECUTION_PATH_EXECUTION_RUNTIME_STREAM, EXECUTION_PATH_EXECUTION_RUNTIME_SYNC,
|
||||
EXECUTION_PATH_LOCAL_AI_PUBLIC, EXECUTION_PATH_LOCAL_API_KEY_CONCURRENCY_LIMITED,
|
||||
EXECUTION_PATH_LOCAL_AUTH_DENIED, EXECUTION_PATH_LOCAL_EXECUTION_LOOP_DETECTED,
|
||||
EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, EXECUTION_PATH_LOCAL_INVALID_REQUEST,
|
||||
EXECUTION_PATH_LOCAL_OVERLOADED, EXECUTION_PATH_LOCAL_PROXY_PASSTHROUGH_REMOVED,
|
||||
EXECUTION_PATH_LOCAL_RATE_LIMITED, EXECUTION_PATH_LOCAL_ROUTE_NOT_FOUND,
|
||||
EXECUTION_PATH_PUBLIC_PROXY_PASSTHROUGH, EXECUTION_RUNTIME_LOOP_GUARD_HEADER,
|
||||
FORWARDED_FOR_HEADER, FORWARDED_HOST_HEADER, FORWARDED_PROTO_HEADER, GATEWAY_HEADER,
|
||||
LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER,
|
||||
TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, TRUSTED_AUTH_API_KEY_ID_HEADER,
|
||||
EXECUTION_PATH_LOCAL_EXECUTION_PLANNING_TIMEOUT, EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS,
|
||||
EXECUTION_PATH_LOCAL_INVALID_REQUEST, EXECUTION_PATH_LOCAL_OVERLOADED,
|
||||
EXECUTION_PATH_LOCAL_PROXY_PASSTHROUGH_REMOVED, EXECUTION_PATH_LOCAL_RATE_LIMITED,
|
||||
EXECUTION_PATH_LOCAL_ROUTE_NOT_FOUND, EXECUTION_PATH_PUBLIC_PROXY_PASSTHROUGH,
|
||||
EXECUTION_RUNTIME_LOOP_GUARD_HEADER, FORWARDED_FOR_HEADER, FORWARDED_HOST_HEADER,
|
||||
FORWARDED_PROTO_HEADER, GATEWAY_HEADER, LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER,
|
||||
TRACE_ID_HEADER, TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, TRUSTED_AUTH_API_KEY_ID_HEADER,
|
||||
TRUSTED_AUTH_BALANCE_HEADER, TRUSTED_AUTH_USER_ID_HEADER, TUNNEL_AFFINITY_FORWARDED_BY_HEADER,
|
||||
TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER,
|
||||
};
|
||||
@@ -55,6 +55,10 @@ use crate::headers::{
|
||||
should_skip_request_header, RequestBodyNormalizationError,
|
||||
};
|
||||
use crate::router::RequestAdmissionError;
|
||||
use crate::scheduler::candidate::{
|
||||
is_auth_api_key_concurrency_limit_skip_reason, AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON,
|
||||
LEGACY_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON,
|
||||
};
|
||||
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode};
|
||||
use crate::{
|
||||
AppState, FrontdoorUserRpmOutcome, GatewayError, GatewayFallbackMetricKind,
|
||||
@@ -65,7 +69,11 @@ use axum::extract::{ConnectInfo, Request, State};
|
||||
use axum::http::{self, header::HeaderName, header::HeaderValue, Response};
|
||||
use futures_util::StreamExt;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::{collections::BTreeMap, time::Instant};
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
error::Error as StdError,
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
const OPENAI_CHAT_LOCAL_EXECUTION_RUNTIME_MISS_DETAIL: &str =
|
||||
@@ -88,24 +96,193 @@ const LOCAL_PROXY_PASSTHROUGH_REMOVED_DETAIL: &str =
|
||||
const LOCAL_EXECUTION_LOOP_DETECTED_DETAIL: &str =
|
||||
"Gateway detected an execution runtime request loop back into the local frontdoor";
|
||||
const AUTH_API_KEY_CONCURRENCY_LIMIT_REACHED_DETAIL: &str =
|
||||
"当前 API Key 并发请求数已达上限,请稍后重试";
|
||||
"当前调用方 API Key 并发请求数已达上限,请稍后重试";
|
||||
const REQUEST_BODY_READ_TIMEOUT_DETAIL: &str =
|
||||
"Request body read timed out before the gateway could route the request";
|
||||
const REQUEST_BODY_READ_FAILED_DETAIL: &str = "Failed to read request body";
|
||||
const LOCAL_EXECUTION_PLANNING_TIMEOUT_DETAIL: &str =
|
||||
"当前 AI 请求在本地执行规划阶段超时,请稍后重试";
|
||||
const EXECUTION_PATH_TUNNEL_AFFINITY_FORWARD: &str = "tunnel_affinity_forward";
|
||||
const MANAGEMENT_TOKEN_PREFIX: &str = "ae-";
|
||||
const LEGACY_MANAGEMENT_TOKEN_PREFIX: &str = "ae_";
|
||||
|
||||
fn build_request_body_normalization_error_response(
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct RequestBodyBufferPolicy {
|
||||
max_bytes: u64,
|
||||
read_timeout: Duration,
|
||||
}
|
||||
|
||||
impl RequestBodyBufferPolicy {
|
||||
fn from_state(state: &AppState) -> Self {
|
||||
Self {
|
||||
max_bytes: crate::headers::max_request_body_bytes(),
|
||||
read_timeout: state.frontdoor_runtime_guards.request_body_read_timeout,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn for_tests(max_bytes: u64, read_timeout: Duration) -> Self {
|
||||
Self {
|
||||
max_bytes,
|
||||
read_timeout,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
enum RequestBodyBufferError {
|
||||
Normalization(RequestBodyNormalizationError),
|
||||
TooLarge { limit_bytes: u64 },
|
||||
Timeout { timeout_ms: u64 },
|
||||
ReadFailed { message: String },
|
||||
}
|
||||
|
||||
impl RequestBodyBufferError {
|
||||
fn http_status(&self) -> http::StatusCode {
|
||||
match self {
|
||||
Self::Normalization(error) => error.http_status(),
|
||||
Self::TooLarge { .. } => http::StatusCode::PAYLOAD_TOO_LARGE,
|
||||
Self::Timeout { .. } => http::StatusCode::REQUEST_TIMEOUT,
|
||||
Self::ReadFailed { .. } => http::StatusCode::BAD_REQUEST,
|
||||
}
|
||||
}
|
||||
|
||||
fn client_message(&self) -> String {
|
||||
match self {
|
||||
Self::Normalization(error) => error.client_message(),
|
||||
Self::TooLarge { limit_bytes } => format!("Request body exceeds {limit_bytes} bytes"),
|
||||
Self::Timeout { .. } => REQUEST_BODY_READ_TIMEOUT_DETAIL.to_string(),
|
||||
Self::ReadFailed { .. } => REQUEST_BODY_READ_FAILED_DETAIL.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn reason(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Normalization(error) => match error {
|
||||
RequestBodyNormalizationError::UnsupportedContentEncoding(_) => {
|
||||
"unsupported_content_encoding"
|
||||
}
|
||||
RequestBodyNormalizationError::DecodeFailed { .. } => "decode_failed",
|
||||
RequestBodyNormalizationError::DecompressedBodyTooLarge { .. } => {
|
||||
"decompressed_body_too_large"
|
||||
}
|
||||
RequestBodyNormalizationError::RequestBodyTooLarge { .. } => {
|
||||
"request_body_too_large"
|
||||
}
|
||||
},
|
||||
Self::TooLarge { .. } => "request_body_too_large",
|
||||
Self::Timeout { .. } => "request_body_read_timeout",
|
||||
Self::ReadFailed { .. } => "request_body_read_failed",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn buffer_and_normalize_request_body(
|
||||
request_body: &mut Option<Body>,
|
||||
headers: &mut http::HeaderMap,
|
||||
body_owner_expectation: &'static str,
|
||||
trace_id: &str,
|
||||
method: &http::Method,
|
||||
path_and_query: &str,
|
||||
phase: &'static str,
|
||||
policy: RequestBodyBufferPolicy,
|
||||
) -> Result<Bytes, RequestBodyBufferError> {
|
||||
if let Err(err) =
|
||||
crate::headers::check_request_content_length_with_limit(headers, policy.max_bytes)
|
||||
{
|
||||
return Err(RequestBodyBufferError::Normalization(err));
|
||||
}
|
||||
|
||||
let read_started_at = Instant::now();
|
||||
let timeout_ms = policy.read_timeout.as_millis() as u64;
|
||||
info!(
|
||||
event_name = "frontdoor_request_body_buffer_started",
|
||||
log_type = "event",
|
||||
trace_id,
|
||||
method = %method,
|
||||
path = %path_and_query,
|
||||
phase,
|
||||
max_body_bytes = policy.max_bytes,
|
||||
timeout_ms,
|
||||
"gateway started buffering request body"
|
||||
);
|
||||
|
||||
let body_limit = usize::try_from(policy.max_bytes).unwrap_or(usize::MAX);
|
||||
let body = match tokio::time::timeout(
|
||||
policy.read_timeout,
|
||||
to_bytes(
|
||||
request_body.take().expect(body_owner_expectation),
|
||||
body_limit,
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(body)) => body,
|
||||
Ok(Err(err)) if request_body_collection_exceeded_limit(&err) => {
|
||||
return Err(RequestBodyBufferError::TooLarge {
|
||||
limit_bytes: policy.max_bytes,
|
||||
});
|
||||
}
|
||||
Ok(Err(err)) => {
|
||||
return Err(RequestBodyBufferError::ReadFailed {
|
||||
message: err.to_string(),
|
||||
});
|
||||
}
|
||||
Err(_) => {
|
||||
return Err(RequestBodyBufferError::Timeout { timeout_ms });
|
||||
}
|
||||
};
|
||||
|
||||
let normalized = crate::headers::normalize_request_body_headers_and_bytes_with_limit(
|
||||
headers,
|
||||
body,
|
||||
policy.max_bytes,
|
||||
)
|
||||
.map_err(RequestBodyBufferError::Normalization)?;
|
||||
info!(
|
||||
event_name = "frontdoor_request_body_buffer_completed",
|
||||
log_type = "event",
|
||||
trace_id,
|
||||
method = %method,
|
||||
path = %path_and_query,
|
||||
phase,
|
||||
body_bytes = normalized.len(),
|
||||
elapsed_ms = read_started_at.elapsed().as_millis() as u64,
|
||||
"gateway completed request body buffering"
|
||||
);
|
||||
Ok(normalized)
|
||||
}
|
||||
|
||||
fn request_body_collection_exceeded_limit(error: &(dyn StdError + 'static)) -> bool {
|
||||
let mut current = Some(error);
|
||||
while let Some(error) = current {
|
||||
if error.to_string().contains("length limit exceeded") {
|
||||
return true;
|
||||
}
|
||||
current = error.source();
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
fn build_request_body_buffer_error_response(
|
||||
trace_id: &str,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
error: &RequestBodyNormalizationError,
|
||||
error: &RequestBodyBufferError,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
warn!(
|
||||
event_name = "frontdoor_request_body_normalization_failed",
|
||||
event_name = "frontdoor_request_body_buffer_failed",
|
||||
log_type = "ops",
|
||||
trace_id,
|
||||
method = %request_context.request_method,
|
||||
path = %request_context.request_path_and_query(),
|
||||
error = %error,
|
||||
"gateway rejected request with invalid encoded body"
|
||||
status_code = error.http_status().as_u16(),
|
||||
reason = error.reason(),
|
||||
detail = %error.client_message(),
|
||||
read_error = match error {
|
||||
RequestBodyBufferError::ReadFailed { message } => message.as_str(),
|
||||
_ => "",
|
||||
},
|
||||
"gateway rejected request body before local execution planning"
|
||||
);
|
||||
build_local_http_error_response(
|
||||
trace_id,
|
||||
@@ -115,36 +292,16 @@ fn build_request_body_normalization_error_response(
|
||||
)
|
||||
}
|
||||
|
||||
async fn buffer_and_normalize_request_body(
|
||||
request_body: &mut Option<Body>,
|
||||
headers: &mut http::HeaderMap,
|
||||
body_owner_expectation: &'static str,
|
||||
) -> Result<Result<Bytes, RequestBodyNormalizationError>, GatewayError> {
|
||||
if let Err(err) = crate::headers::check_request_content_length(headers) {
|
||||
return Ok(Err(err));
|
||||
}
|
||||
let body = to_bytes(
|
||||
request_body.take().expect(body_owner_expectation),
|
||||
usize::MAX,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
Ok(crate::headers::normalize_request_body_headers_and_bytes(
|
||||
headers, body,
|
||||
))
|
||||
}
|
||||
|
||||
fn finalize_request_body_normalization_rejection(
|
||||
fn finalize_request_body_buffer_rejection(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
remote_addr: &std::net::SocketAddr,
|
||||
started_at: &std::time::Instant,
|
||||
trace_id: &str,
|
||||
request_permit: Option<aether_runtime::AdmissionPermit>,
|
||||
error: &RequestBodyNormalizationError,
|
||||
error: &RequestBodyBufferError,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let response =
|
||||
build_request_body_normalization_error_response(trace_id, request_context, error)?;
|
||||
let response = build_request_body_buffer_error_response(trace_id, request_context, error)?;
|
||||
Ok(finalize_gateway_response_with_context(
|
||||
state,
|
||||
response,
|
||||
@@ -156,6 +313,59 @@ fn finalize_request_body_normalization_rejection(
|
||||
))
|
||||
}
|
||||
|
||||
fn local_execution_planning_timeout_parts(error: &GatewayError) -> Option<(&'static str, u64)> {
|
||||
match error {
|
||||
GatewayError::LocalExecutionPlanningTimeout {
|
||||
phase, timeout_ms, ..
|
||||
} => Some((*phase, *timeout_ms)),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn finalize_local_execution_planning_timeout(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
remote_addr: &std::net::SocketAddr,
|
||||
started_at: &std::time::Instant,
|
||||
trace_id: &str,
|
||||
request_permit: Option<aether_runtime::AdmissionPermit>,
|
||||
control_decision: Option<&GatewayControlDecision>,
|
||||
phase: &'static str,
|
||||
timeout_ms: u64,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
warn!(
|
||||
event_name = "frontdoor_local_execution_planning_timeout",
|
||||
log_type = "ops",
|
||||
trace_id,
|
||||
method = %request_context.request_method,
|
||||
path = %request_context.request_path_and_query(),
|
||||
route_family = control_decision
|
||||
.and_then(|decision| decision.route_family.as_deref())
|
||||
.unwrap_or("-"),
|
||||
route_kind = control_decision
|
||||
.and_then(|decision| decision.route_kind.as_deref())
|
||||
.unwrap_or("-"),
|
||||
phase,
|
||||
timeout_ms,
|
||||
"gateway failed local execution before a candidate could be selected"
|
||||
);
|
||||
let response = build_local_http_error_response(
|
||||
trace_id,
|
||||
control_decision,
|
||||
http::StatusCode::GATEWAY_TIMEOUT,
|
||||
LOCAL_EXECUTION_PLANNING_TIMEOUT_DETAIL,
|
||||
)?;
|
||||
Ok(finalize_gateway_response_with_context(
|
||||
state,
|
||||
response,
|
||||
remote_addr,
|
||||
request_context,
|
||||
EXECUTION_PATH_LOCAL_EXECUTION_PLANNING_TIMEOUT,
|
||||
started_at,
|
||||
request_permit,
|
||||
))
|
||||
}
|
||||
|
||||
fn local_execution_outcome_label(outcome: &LocalExecutionRequestOutcome) -> &'static str {
|
||||
match outcome {
|
||||
LocalExecutionRequestOutcome::Responded(_) => "responded",
|
||||
@@ -991,16 +1201,22 @@ pub(crate) async fn proxy_request(
|
||||
}
|
||||
let mut request_body = Some(body);
|
||||
let local_proxy_body = if local_proxy_route_requires_buffered_body(&request_context) {
|
||||
let body_buffer_policy = RequestBodyBufferPolicy::from_state(&state);
|
||||
let body = buffer_and_normalize_request_body(
|
||||
&mut request_body,
|
||||
&mut parts.headers,
|
||||
"local proxy body buffering should own request body",
|
||||
&trace_id,
|
||||
&parts.method,
|
||||
&request_context.request_path_and_query(),
|
||||
"local_proxy",
|
||||
body_buffer_policy,
|
||||
)
|
||||
.await?;
|
||||
.await;
|
||||
match body {
|
||||
Ok(body) => Some(body),
|
||||
Err(err) => {
|
||||
return finalize_request_body_normalization_rejection(
|
||||
return finalize_request_body_buffer_rejection(
|
||||
&state,
|
||||
&request_context,
|
||||
&remote_addr,
|
||||
@@ -1171,16 +1387,22 @@ pub(crate) async fn proxy_request(
|
||||
&& request_enables_control_execute(&parts.headers);
|
||||
|
||||
let buffered_body = if should_buffer_body {
|
||||
let body_buffer_policy = RequestBodyBufferPolicy::from_state(&state);
|
||||
let body = buffer_and_normalize_request_body(
|
||||
&mut request_body,
|
||||
&mut parts.headers,
|
||||
"buffered auth/execution runtime path should own request body",
|
||||
&trace_id,
|
||||
&parts.method,
|
||||
&request_context.request_path_and_query(),
|
||||
"auth_execution",
|
||||
body_buffer_policy,
|
||||
)
|
||||
.await?;
|
||||
.await;
|
||||
match body {
|
||||
Ok(body) => Some(body),
|
||||
Err(err) => {
|
||||
return finalize_request_body_normalization_rejection(
|
||||
return finalize_request_body_buffer_rejection(
|
||||
&state,
|
||||
&request_context,
|
||||
&remote_addr,
|
||||
@@ -1335,14 +1557,34 @@ pub(crate) async fn proxy_request(
|
||||
let stream_request = request_wants_stream(&request_context, &parts.headers, buffered_body);
|
||||
let mut local_execution_exhaustion = None;
|
||||
if stream_request {
|
||||
let stream_outcome = maybe_execute_stream_request(
|
||||
let stream_outcome = match maybe_execute_stream_request(
|
||||
&state,
|
||||
&parts,
|
||||
buffered_body,
|
||||
&trace_id,
|
||||
control_decision,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
{
|
||||
Ok(outcome) => outcome,
|
||||
Err(err) => {
|
||||
if let Some((phase, timeout_ms)) = local_execution_planning_timeout_parts(&err)
|
||||
{
|
||||
return finalize_local_execution_planning_timeout(
|
||||
&state,
|
||||
&request_context,
|
||||
&remote_addr,
|
||||
&started_at,
|
||||
&trace_id,
|
||||
request_permit.take(),
|
||||
control_decision,
|
||||
phase,
|
||||
timeout_ms,
|
||||
);
|
||||
}
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
debug!(
|
||||
event_name = "proxy_stream_local_execute_outcome",
|
||||
log_type = "debug",
|
||||
@@ -1380,9 +1622,34 @@ pub(crate) async fn proxy_request(
|
||||
LocalExecutionRequestOutcome::NoPath => {}
|
||||
}
|
||||
}
|
||||
match maybe_execute_sync_request(&state, &parts, buffered_body, &trace_id, control_decision)
|
||||
.await?
|
||||
let sync_outcome = match maybe_execute_sync_request(
|
||||
&state,
|
||||
&parts,
|
||||
buffered_body,
|
||||
&trace_id,
|
||||
control_decision,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(outcome) => outcome,
|
||||
Err(err) => {
|
||||
if let Some((phase, timeout_ms)) = local_execution_planning_timeout_parts(&err) {
|
||||
return finalize_local_execution_planning_timeout(
|
||||
&state,
|
||||
&request_context,
|
||||
&remote_addr,
|
||||
&started_at,
|
||||
&trace_id,
|
||||
request_permit.take(),
|
||||
control_decision,
|
||||
phase,
|
||||
timeout_ms,
|
||||
);
|
||||
}
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
match sync_outcome {
|
||||
LocalExecutionRequestOutcome::Responded(execution_runtime_response) => {
|
||||
let execution_runtime_response = restore_redacted_sync_execution_response(
|
||||
execution_runtime_response,
|
||||
@@ -1406,15 +1673,35 @@ pub(crate) async fn proxy_request(
|
||||
LocalExecutionRequestOutcome::NoPath => {}
|
||||
}
|
||||
if parts.method != http::Method::POST {
|
||||
match maybe_execute_stream_request(
|
||||
let stream_outcome = match maybe_execute_stream_request(
|
||||
&state,
|
||||
&parts,
|
||||
buffered_body,
|
||||
&trace_id,
|
||||
control_decision,
|
||||
)
|
||||
.await?
|
||||
.await
|
||||
{
|
||||
Ok(outcome) => outcome,
|
||||
Err(err) => {
|
||||
if let Some((phase, timeout_ms)) = local_execution_planning_timeout_parts(&err)
|
||||
{
|
||||
return finalize_local_execution_planning_timeout(
|
||||
&state,
|
||||
&request_context,
|
||||
&remote_addr,
|
||||
&started_at,
|
||||
&trace_id,
|
||||
request_permit.take(),
|
||||
control_decision,
|
||||
phase,
|
||||
timeout_ms,
|
||||
);
|
||||
}
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
match stream_outcome {
|
||||
LocalExecutionRequestOutcome::Responded(execution_runtime_response) => {
|
||||
let execution_runtime_response = restore_redacted_stream_execution_response(
|
||||
execution_runtime_response,
|
||||
@@ -1507,7 +1794,9 @@ pub(crate) async fn proxy_request(
|
||||
let auth_api_key_concurrency_limited = diagnostic_is_auth_api_key_concurrency_limited(
|
||||
local_execution_runtime_miss_diagnostic.as_ref(),
|
||||
) || local_execution_runtime_miss_context
|
||||
.all_candidates_skipped_for_reason("api_key_concurrency_limit_reached");
|
||||
.all_candidates_skipped_for_reason(AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON)
|
||||
|| local_execution_runtime_miss_context
|
||||
.all_candidates_skipped_for_reason(LEGACY_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON);
|
||||
let local_execution_runtime_miss_detail = local_execution_runtime_miss_detail(
|
||||
control_decision,
|
||||
local_execution_runtime_miss_diagnostic.as_ref(),
|
||||
@@ -1634,7 +1923,7 @@ pub(crate) async fn proxy_request(
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| {
|
||||
auth_api_key_concurrency_limited
|
||||
.then_some("api_key_concurrency_limit_reached".to_string())
|
||||
.then_some(AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON.to_string())
|
||||
});
|
||||
if let Some(reason) = local_execution_runtime_miss_reason {
|
||||
response.headers_mut().insert(
|
||||
@@ -1859,7 +2148,9 @@ fn local_execution_runtime_miss_skip_reasons_summary(
|
||||
|
||||
fn local_execution_runtime_miss_skip_reason_label(reason: &str) -> &str {
|
||||
match reason {
|
||||
"api_key_concurrency_limit_reached" => "API Key 并发已达上限",
|
||||
"auth_api_key_concurrency_limit_reached" | "api_key_concurrency_limit_reached" => {
|
||||
"调用方 API Key 并发已达上限"
|
||||
}
|
||||
"auth_channel_mismatch" => "认证通道不匹配",
|
||||
"auth_snapshot_missing" => "API Key 本地执行配置缺失",
|
||||
"endpoint_api_format_changed" => "端点 API 格式已变更",
|
||||
@@ -1874,7 +2165,9 @@ fn local_execution_runtime_miss_skip_reason_label(reason: &str) -> &str {
|
||||
"pool_cost_limit_reached" => "池内账号成本额度已用尽",
|
||||
"pool_group_exhausted" => "池化提供商没有可调度账号",
|
||||
"pool_key_lease_busy" => "池内账号正被其他请求占用",
|
||||
"provider_concurrency_limit_reached" => "上游提供商并发已达上限",
|
||||
"provider_inactive" => "提供商未启用",
|
||||
"provider_key_concurrency_limit_reached" => "上游账号并发已达上限",
|
||||
"provider_request_body_missing" => "无法构建上游请求体",
|
||||
"provider_request_body_build_failed" => "上游请求体转换失败",
|
||||
"transport_api_format_mismatch" => "传输层 API 格式不匹配",
|
||||
@@ -1934,15 +2227,12 @@ fn diagnostic_is_auth_api_key_concurrency_limited(
|
||||
let Some(diagnostic) = diagnostic else {
|
||||
return false;
|
||||
};
|
||||
diagnostic.reason == "api_key_concurrency_limit_reached"
|
||||
is_auth_api_key_concurrency_limit_skip_reason(diagnostic.reason.as_str())
|
||||
|| (diagnostic.reason == "all_candidates_skipped"
|
||||
&& diagnostic.skip_reasons.len() == 1
|
||||
&& diagnostic
|
||||
.skip_reasons
|
||||
.get("api_key_concurrency_limit_reached")
|
||||
.copied()
|
||||
.unwrap_or(0)
|
||||
> 0)
|
||||
&& diagnostic.skip_reasons.iter().any(|(reason, count)| {
|
||||
is_auth_api_key_concurrency_limit_skip_reason(reason.as_str()) && *count > 0
|
||||
}))
|
||||
}
|
||||
|
||||
fn local_execution_runtime_miss_route_detail(
|
||||
@@ -1977,14 +2267,17 @@ fn local_execution_runtime_miss_route_detail(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::time::Duration;
|
||||
|
||||
use super::{
|
||||
api_key_remote_ip_allowed, diagnostic_is_auth_api_key_concurrency_limited,
|
||||
local_execution_runtime_miss_detail, restore_redacted_stream_execution_response,
|
||||
restore_redacted_sync_execution_response, GatewayControlDecision,
|
||||
LocalExecutionRuntimeMissDiagnostic,
|
||||
api_key_remote_ip_allowed, buffer_and_normalize_request_body,
|
||||
diagnostic_is_auth_api_key_concurrency_limited, local_execution_runtime_miss_detail,
|
||||
restore_redacted_stream_execution_response, restore_redacted_sync_execution_response,
|
||||
GatewayControlDecision, LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError,
|
||||
RequestBodyBufferPolicy,
|
||||
};
|
||||
use axum::body::{to_bytes, Body};
|
||||
use axum::http::{header, Response};
|
||||
use axum::body::{to_bytes, Body, Bytes};
|
||||
use axum::http::{header, HeaderMap, Method, Response};
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
@@ -2112,6 +2405,58 @@ mod tests {
|
||||
assert!(!message.contains(&sentinel));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_body_buffer_rejects_chunked_body_when_limit_is_exceeded() {
|
||||
let mut body = Some(Body::from(Bytes::from_static(b"abcdef")));
|
||||
let mut headers = HeaderMap::new();
|
||||
|
||||
let err = buffer_and_normalize_request_body(
|
||||
&mut body,
|
||||
&mut headers,
|
||||
"test owns body",
|
||||
"trace-body-large",
|
||||
&Method::POST,
|
||||
"/v1/responses",
|
||||
"test",
|
||||
RequestBodyBufferPolicy::for_tests(5, Duration::from_secs(1)),
|
||||
)
|
||||
.await
|
||||
.expect_err("body exceeding the ingress limit should fail");
|
||||
|
||||
assert!(matches!(
|
||||
err,
|
||||
RequestBodyBufferError::TooLarge { limit_bytes: 5 }
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_body_buffer_times_out_instead_of_waiting_forever() {
|
||||
let stream = async_stream::stream! {
|
||||
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(b"{"));
|
||||
std::future::pending::<()>().await;
|
||||
};
|
||||
let mut body = Some(Body::from_stream(stream));
|
||||
let mut headers = HeaderMap::new();
|
||||
|
||||
let err = buffer_and_normalize_request_body(
|
||||
&mut body,
|
||||
&mut headers,
|
||||
"test owns body",
|
||||
"trace-body-timeout",
|
||||
&Method::POST,
|
||||
"/v1/responses",
|
||||
"test",
|
||||
RequestBodyBufferPolicy::for_tests(1024, Duration::from_millis(5)),
|
||||
)
|
||||
.await
|
||||
.expect_err("body buffering should time out");
|
||||
|
||||
assert!(matches!(
|
||||
err,
|
||||
RequestBodyBufferError::Timeout { timeout_ms: 5 }
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_miss_detail_returns_model_specific_stream_message_when_candidates_are_unavailable() {
|
||||
let decision = GatewayControlDecision::synthetic(
|
||||
@@ -2176,7 +2521,7 @@ mod tests {
|
||||
let diagnostic = LocalExecutionRuntimeMissDiagnostic {
|
||||
reason: "all_candidates_skipped".to_string(),
|
||||
skip_reasons: std::collections::BTreeMap::from([(
|
||||
"api_key_concurrency_limit_reached".to_string(),
|
||||
"auth_api_key_concurrency_limit_reached".to_string(),
|
||||
1,
|
||||
)]),
|
||||
requested_model: Some("gpt-5.4".to_string()),
|
||||
@@ -2188,7 +2533,7 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
detail.as_deref(),
|
||||
Some("当前 API Key 并发请求数已达上限,请稍后重试")
|
||||
Some("当前调用方 API Key 并发请求数已达上限,请稍后重试")
|
||||
);
|
||||
assert!(diagnostic_is_auth_api_key_concurrency_limited(Some(
|
||||
&diagnostic
|
||||
@@ -2219,7 +2564,7 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
detail.as_deref(),
|
||||
Some("当前 API Key 并发请求数已达上限,请稍后重试")
|
||||
Some("当前调用方 API Key 并发请求数已达上限,请稍后重试")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,6 +52,12 @@ const OPENAI_RERANK_TOP_N_DETAIL: &str = "Rerank request top_n must be a positiv
|
||||
const OPENAI_RERANK_CHAT_PAYLOAD_DETAIL: &str =
|
||||
"Rerank request must use query/documents, not chat messages";
|
||||
const OPENAI_RERANK_STREAM_UNSUPPORTED_DETAIL: &str = "Rerank requests do not support streaming";
|
||||
const ANTIGRAVITY_USER_SETTINGS_MISSING_BODY_DETAIL: &str =
|
||||
"Antigravity setUserSettings request body is required";
|
||||
const ANTIGRAVITY_USER_SETTINGS_INVALID_JSON_DETAIL: &str =
|
||||
"Antigravity setUserSettings request JSON body is invalid";
|
||||
const ANTIGRAVITY_USER_SETTINGS_INVALID_DETAIL: &str =
|
||||
"Antigravity setUserSettings request must include object userSettings";
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
enum OpenAiImageOperation {
|
||||
@@ -103,7 +109,9 @@ pub(crate) fn ai_public_local_requires_buffered_body(
|
||||
&& request_context.request_path == "/v1/embeddings")
|
||||
|| (decision.route_family.as_deref() == Some("openai")
|
||||
&& decision.route_kind.as_deref() == Some("rerank")
|
||||
&& request_context.request_path == "/v1/rerank"))
|
||||
&& request_context.request_path == "/v1/rerank")
|
||||
|| (decision.route_family.as_deref() == Some("antigravity")
|
||||
&& decision.route_kind.as_deref() != Some("stream_generate_content")))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -133,6 +141,12 @@ pub(crate) async fn maybe_build_local_ai_public_response(
|
||||
return Some(response);
|
||||
}
|
||||
|
||||
if let Some(response) =
|
||||
maybe_build_local_antigravity_v1internal_response(request_context, request_body)
|
||||
{
|
||||
return Some(response);
|
||||
}
|
||||
|
||||
maybe_build_local_gemini_video_operations_response(state, request_context, decision).await
|
||||
}
|
||||
|
||||
@@ -861,6 +875,238 @@ fn maybe_build_local_claude_count_tokens_response(
|
||||
Some(Json(json!({ "input_tokens": input_tokens })).into_response())
|
||||
}
|
||||
|
||||
fn maybe_build_local_antigravity_v1internal_response(
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Option<Response<Body>> {
|
||||
let decision = request_context.control_decision.as_ref()?;
|
||||
if decision.route_family.as_deref() != Some("antigravity")
|
||||
|| request_context.request_method != http::Method::POST
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
match decision.route_kind.as_deref()? {
|
||||
"load_code_assist" => {
|
||||
Some(Json(build_antigravity_load_code_assist_payload()).into_response())
|
||||
}
|
||||
"fetch_available_models" => {
|
||||
Some(Json(build_antigravity_fetch_available_models_payload()).into_response())
|
||||
}
|
||||
"fetch_user_info" => {
|
||||
Some(Json(build_antigravity_fetch_user_info_payload()).into_response())
|
||||
}
|
||||
"fetch_admin_controls" => Some(Json(json!({})).into_response()),
|
||||
"list_experiments" => Some(
|
||||
Json(json!({
|
||||
"experimentIds": [],
|
||||
"flags": []
|
||||
}))
|
||||
.into_response(),
|
||||
),
|
||||
"record_code_assist_metrics" => Some(Json(json!({})).into_response()),
|
||||
"set_user_settings" => Some(build_antigravity_set_user_settings_response(request_body)),
|
||||
"stream_generate_content" => None,
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_antigravity_set_user_settings_response(request_body: Option<&Bytes>) -> Response<Body> {
|
||||
let Some(request_body) = request_body else {
|
||||
return build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
ANTIGRAVITY_USER_SETTINGS_MISSING_BODY_DETAIL,
|
||||
);
|
||||
};
|
||||
let payload = match serde_json::from_slice::<Value>(request_body) {
|
||||
Ok(payload) => payload,
|
||||
Err(_) => {
|
||||
return build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
ANTIGRAVITY_USER_SETTINGS_INVALID_JSON_DETAIL,
|
||||
);
|
||||
}
|
||||
};
|
||||
let Some(user_settings) = payload
|
||||
.get("userSettings")
|
||||
.filter(|value| value.is_object())
|
||||
.cloned()
|
||||
else {
|
||||
return build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
ANTIGRAVITY_USER_SETTINGS_INVALID_DETAIL,
|
||||
);
|
||||
};
|
||||
|
||||
Json(json!({ "userSettings": user_settings })).into_response()
|
||||
}
|
||||
|
||||
fn build_antigravity_load_code_assist_payload() -> Value {
|
||||
json!({
|
||||
"allowedTiers": [
|
||||
antigravity_free_tier_payload(true),
|
||||
antigravity_standard_tier_payload()
|
||||
],
|
||||
"cloudaicompanionProject": "aether-antigravity-local",
|
||||
"currentTier": antigravity_free_tier_payload(false),
|
||||
"gcpManaged": false,
|
||||
"paidTier": antigravity_paid_tier_payload(),
|
||||
"upgradeSubscriptionUri": "https://codeassist.google.com/upgrade"
|
||||
})
|
||||
}
|
||||
|
||||
fn antigravity_free_tier_payload(include_default_marker: bool) -> Value {
|
||||
if include_default_marker {
|
||||
json!({
|
||||
"id": "free-tier",
|
||||
"name": "Antigravity",
|
||||
"description": "Gemini-powered code suggestions and chat in multiple IDEs",
|
||||
"privacyNotice": {
|
||||
"showNotice": false
|
||||
},
|
||||
"isDefault": true
|
||||
})
|
||||
} else {
|
||||
json!({
|
||||
"id": "free-tier",
|
||||
"name": "Antigravity",
|
||||
"description": "Gemini-powered code suggestions and chat in multiple IDEs",
|
||||
"privacyNotice": {
|
||||
"showNotice": false
|
||||
},
|
||||
"upgradeSubscriptionUri": "https://codeassist.google.com/upgrade",
|
||||
"upgradeSubscriptionText": "Upgrade for higher Antigravity request limits",
|
||||
"upgradeSubscriptionType": "GDP_HELIUM"
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn antigravity_standard_tier_payload() -> Value {
|
||||
json!({
|
||||
"id": "standard-tier",
|
||||
"name": "Antigravity",
|
||||
"description": "Unlimited coding assistant with the most powerful Gemini models",
|
||||
"userDefinedCloudaicompanionProject": true,
|
||||
"privacyNotice": {},
|
||||
"usesGcpTos": true
|
||||
})
|
||||
}
|
||||
|
||||
fn antigravity_paid_tier_payload() -> Value {
|
||||
json!({
|
||||
"id": "g1-pro-tier",
|
||||
"name": "Google AI Pro",
|
||||
"description": "Google AI Pro",
|
||||
"upgradeSubscriptionUri": "https://antigravity.google/g1-upgrade",
|
||||
"upgradeSubscriptionText": "Upgrade for the highest Antigravity request limits"
|
||||
})
|
||||
}
|
||||
|
||||
fn build_antigravity_fetch_user_info_payload() -> Value {
|
||||
json!({
|
||||
"regionCode": "US",
|
||||
"userSettings": build_antigravity_default_user_settings_payload()
|
||||
})
|
||||
}
|
||||
|
||||
fn build_antigravity_default_user_settings_payload() -> Value {
|
||||
json!({
|
||||
"preferredModelId": "gemini-3.1-flash-lite"
|
||||
})
|
||||
}
|
||||
|
||||
fn build_antigravity_fetch_available_models_payload() -> Value {
|
||||
json!({
|
||||
"models": {
|
||||
"gemini-3.5-flash-low": antigravity_model_payload("gemini-3.5-flash-low", "Gemini 3.5 Flash Low"),
|
||||
"gemini-3-flash-agent": antigravity_model_payload("gemini-3-flash-agent", "Gemini 3 Flash Agent"),
|
||||
"gemini-3.1-flash-lite": antigravity_model_payload("gemini-3.1-flash-lite", "Gemini 3.1 Flash Lite"),
|
||||
"gemini-3.1-pro-low": antigravity_model_payload("gemini-3.1-pro-low", "Gemini 3.1 Pro Low"),
|
||||
"gemini-3-flash": antigravity_model_payload("gemini-3-flash", "Gemini 3 Flash"),
|
||||
"gemini-2.5-flash": antigravity_model_payload("gemini-2.5-flash", "Gemini 2.5 Flash"),
|
||||
"gemini-2.5-flash-lite": antigravity_model_payload("gemini-2.5-flash-lite", "Gemini 2.5 Flash Lite"),
|
||||
"gemini-2.5-flash-thinking": antigravity_model_payload("gemini-2.5-flash-thinking", "Gemini 2.5 Flash Thinking"),
|
||||
"gemini-2.5-pro": antigravity_model_payload("gemini-2.5-pro", "Gemini 2.5 Pro"),
|
||||
"gemini-3.1-flash-image": antigravity_model_payload("gemini-3.1-flash-image", "Gemini 3.1 Flash Image"),
|
||||
"tab_flash_lite_preview": antigravity_model_payload("tab_flash_lite_preview", "Tab Flash Lite Preview"),
|
||||
"tab_jump_flash_lite_preview": antigravity_model_payload("tab_jump_flash_lite_preview", "Tab Jump Flash Lite Preview"),
|
||||
"models/proactive-observer": antigravity_model_payload("models/proactive-observer", "Proactive Observer")
|
||||
},
|
||||
"agentModelSorts": [
|
||||
{
|
||||
"displayName": "Recommended",
|
||||
"groups": [
|
||||
{
|
||||
"modelIds": [
|
||||
"gemini-3.1-flash-lite",
|
||||
"gemini-3-flash-agent",
|
||||
"gemini-3.1-pro-low",
|
||||
"gemini-3.5-flash-low"
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"audioTranscriptionModelIds": ["models/proactive-observer"],
|
||||
"commandModelIds": ["gemini-3-flash"],
|
||||
"commitMessageModelIds": ["gemini-3.1-flash-lite"],
|
||||
"defaultAgentModelId": "gemini-3.1-flash-lite",
|
||||
"deprecatedModelIds": {},
|
||||
"experimentIds": [],
|
||||
"imageGenerationModelIds": ["gemini-3.1-flash-image"],
|
||||
"mqueryModelIds": ["gemini-3.1-flash-lite"],
|
||||
"tabModelIds": ["tab_flash_lite_preview", "tab_jump_flash_lite_preview"],
|
||||
"tieredModelIds": {
|
||||
"flash": ["gemini-3-flash-agent"],
|
||||
"flashLite": ["gemini-3.1-flash-lite"],
|
||||
"pro": ["gemini-3.1-pro-low"]
|
||||
},
|
||||
"webSearchModelIds": ["gemini-3.1-flash-lite"]
|
||||
})
|
||||
}
|
||||
|
||||
fn antigravity_model_payload(id: &str, display_name: &str) -> Value {
|
||||
let model = match id {
|
||||
"gemini-2.5-flash" => "MODEL_GOOGLE_GEMINI_2_5_FLASH",
|
||||
"gemini-2.5-flash-lite" => "MODEL_GOOGLE_GEMINI_2_5_FLASH_LITE",
|
||||
"gemini-2.5-flash-thinking" => "MODEL_GOOGLE_GEMINI_2_5_FLASH_THINKING",
|
||||
"gemini-2.5-pro" => "MODEL_GOOGLE_GEMINI_2_5_PRO",
|
||||
"gemini-3-flash" => "MODEL_PLACEHOLDER_M18",
|
||||
"gemini-3-flash-agent" => "MODEL_PLACEHOLDER_M132",
|
||||
"gemini-3.1-flash-image" => "MODEL_PLACEHOLDER_M21",
|
||||
"gemini-3.1-flash-lite" => "MODEL_PLACEHOLDER_M50",
|
||||
"gemini-3.1-pro-low" => "MODEL_PLACEHOLDER_M36",
|
||||
"gemini-3.5-flash-low" => "MODEL_PLACEHOLDER_M20",
|
||||
"models/proactive-observer" => "MODEL_PLACEHOLDER_M70",
|
||||
"tab_flash_lite_preview" => "MODEL_PLACEHOLDER_M19",
|
||||
"tab_jump_flash_lite_preview" => "MODEL_PLACEHOLDER_M28",
|
||||
_ => "MODEL_PLACEHOLDER_M20",
|
||||
};
|
||||
json!({
|
||||
"apiProvider": "API_PROVIDER_GOOGLE_GEMINI",
|
||||
"displayName": display_name,
|
||||
"maxOutputTokens": 65536,
|
||||
"maxTokens": 1048576,
|
||||
"minThinkingBudget": 32,
|
||||
"model": model,
|
||||
"modelProvider": "MODEL_PROVIDER_GOOGLE",
|
||||
"recommended": id == "gemini-3.1-flash-lite",
|
||||
"supportedMimeTypes": {
|
||||
"application/json": true,
|
||||
"application/pdf": true,
|
||||
"image/jpeg": true,
|
||||
"image/png": true,
|
||||
"text/markdown": true,
|
||||
"text/plain": true
|
||||
},
|
||||
"supportsImages": true,
|
||||
"supportsThinking": true,
|
||||
"supportsVideo": true,
|
||||
"thinkingBudget": 4000,
|
||||
"tokenizerType": "LLAMA_WITH_SPECIAL"
|
||||
})
|
||||
}
|
||||
|
||||
async fn maybe_build_local_gemini_video_operations_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
|
||||
@@ -136,12 +136,7 @@ pub(super) fn announcements_internal_error_response(detail: impl Into<String>) -
|
||||
}
|
||||
|
||||
pub(super) fn announcements_internal_detail(err: GatewayError) -> String {
|
||||
match err {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
err.into_message()
|
||||
}
|
||||
|
||||
pub(super) fn parse_optional_rfc3339_unix_secs(
|
||||
|
||||
@@ -160,21 +160,26 @@ fn dashboard_format_token_compact(value: u64) -> String {
|
||||
return dashboard_format_integer(value);
|
||||
}
|
||||
|
||||
if value < 1_000_000 {
|
||||
let thousands = value as f64 / 1_000.0;
|
||||
if thousands >= 100.0 {
|
||||
return format!("{}K", thousands.round() as u64);
|
||||
const UNITS: &[(u64, &str)] = &[
|
||||
(1_000_000_000_000, "T"),
|
||||
(1_000_000_000, "B"),
|
||||
(1_000_000, "M"),
|
||||
(1_000, "K"),
|
||||
];
|
||||
|
||||
for (divisor, suffix) in UNITS {
|
||||
if value < *divisor {
|
||||
continue;
|
||||
}
|
||||
let decimals = if thousands >= 10.0 { 1 } else { 2 };
|
||||
return format!("{}K", dashboard_trimmed_decimal(thousands, decimals));
|
||||
let scaled = value as f64 / *divisor as f64;
|
||||
if scaled >= 100.0 {
|
||||
return format!("{}{}", scaled.round() as u64, suffix);
|
||||
}
|
||||
let decimals = if scaled >= 10.0 { 1 } else { 2 };
|
||||
return format!("{}{}", dashboard_trimmed_decimal(scaled, decimals), suffix);
|
||||
}
|
||||
|
||||
let millions = value as f64 / 1_000_000.0;
|
||||
if millions >= 100.0 {
|
||||
return format!("{}M", millions.round() as u64);
|
||||
}
|
||||
let decimals = if millions >= 10.0 { 1 } else { 2 };
|
||||
format!("{}M", dashboard_trimmed_decimal(millions, decimals))
|
||||
dashboard_format_integer(value)
|
||||
}
|
||||
|
||||
fn dashboard_format_usd(value: f64) -> String {
|
||||
@@ -1500,3 +1505,17 @@ fn dashboard_format_time_hhmm(unix_secs: u64) -> Option<String> {
|
||||
let datetime = chrono::DateTime::<chrono::Utc>::from_timestamp(timestamp, 0)?;
|
||||
Some(datetime.format("%H:%M").to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::dashboard_format_token_compact;
|
||||
|
||||
#[test]
|
||||
fn dashboard_format_token_compact_promotes_above_millions() {
|
||||
assert_eq!(dashboard_format_token_compact(999), "999");
|
||||
assert_eq!(dashboard_format_token_compact(1_250), "1.25K");
|
||||
assert_eq!(dashboard_format_token_compact(12_500_000), "12.5M");
|
||||
assert_eq!(dashboard_format_token_compact(1_250_000_000), "1.25B");
|
||||
assert_eq!(dashboard_format_token_compact(12_500_000_000_000), "12.5T");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,7 +16,7 @@ const INSTALL_SESSION_TTL_SECS: u64 = 15 * 60;
|
||||
const INSTALL_SESSION_KEY_PREFIX: &str = "install:session:";
|
||||
const TUNNEL_INSTALL_SESSION_KEY_PREFIX: &str = "tunnel-install:session:";
|
||||
const TUNNEL_INSTALL_UNIX_SCRIPT_URL: &str =
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh";
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/refs/heads/main/apps/aether-tunnel/install.sh";
|
||||
const TUNNEL_INSTALL_POWERSHELL_SCRIPT_URL: &str =
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1";
|
||||
|
||||
@@ -59,6 +59,8 @@ struct StoredTunnelInstallSession {
|
||||
aether_url: String,
|
||||
management_token: String,
|
||||
node_name: String,
|
||||
tunnel_security: String,
|
||||
tunnel_encryption_key: String,
|
||||
expires_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
@@ -93,8 +95,7 @@ fn install_code_from_path(request_path: &str) -> Option<(String, bool)> {
|
||||
|
||||
fn tunnel_install_code_from_path(request_path: &str) -> Option<(String, bool)> {
|
||||
let raw = request_path
|
||||
.strip_prefix("/install-tunnel/")
|
||||
.or_else(|| request_path.strip_prefix("/install-proxy/"))?
|
||||
.strip_prefix("/install-tunnel/")?
|
||||
.trim()
|
||||
.trim_matches('/');
|
||||
if raw.is_empty() || raw.contains('/') {
|
||||
@@ -122,6 +123,17 @@ fn generate_install_code() -> String {
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn generate_tunnel_encryption_key() -> String {
|
||||
use base64::Engine;
|
||||
|
||||
let first = uuid::Uuid::new_v4();
|
||||
let second = uuid::Uuid::new_v4();
|
||||
let mut key = [0_u8; 32];
|
||||
key[..16].copy_from_slice(first.as_bytes());
|
||||
key[16..].copy_from_slice(second.as_bytes());
|
||||
base64::engine::general_purpose::STANDARD.encode(key)
|
||||
}
|
||||
|
||||
fn unix_secs_now() -> u64 {
|
||||
chrono::Utc::now().timestamp().max(0) as u64
|
||||
}
|
||||
@@ -172,6 +184,8 @@ set -eu
|
||||
export AETHER_TUNNEL_AETHER_URL={aether_url}
|
||||
export AETHER_TUNNEL_MANAGEMENT_TOKEN={management_token}
|
||||
export AETHER_TUNNEL_NODE_NAME={node_name}
|
||||
export AETHER_TUNNEL_SECURITY={tunnel_security}
|
||||
export AETHER_TUNNEL_ENCRYPTION_KEY={tunnel_encryption_key}
|
||||
|
||||
if command -v curl >/dev/null 2>&1; then
|
||||
curl -fsSL {script_url} | sh
|
||||
@@ -185,6 +199,8 @@ fi
|
||||
aether_url = shell_single_quote(&session.aether_url),
|
||||
management_token = shell_single_quote(&session.management_token),
|
||||
node_name = shell_single_quote(&session.node_name),
|
||||
tunnel_security = shell_single_quote(&session.tunnel_security),
|
||||
tunnel_encryption_key = shell_single_quote(&session.tunnel_encryption_key),
|
||||
script_url = shell_single_quote(TUNNEL_INSTALL_UNIX_SCRIPT_URL),
|
||||
)
|
||||
}
|
||||
@@ -195,11 +211,15 @@ fn build_tunnel_powershell_script(session: &StoredTunnelInstallSession) -> Strin
|
||||
$env:AETHER_TUNNEL_AETHER_URL = {aether_url}
|
||||
$env:AETHER_TUNNEL_MANAGEMENT_TOKEN = {management_token}
|
||||
$env:AETHER_TUNNEL_NODE_NAME = {node_name}
|
||||
$env:AETHER_TUNNEL_SECURITY = {tunnel_security}
|
||||
$env:AETHER_TUNNEL_ENCRYPTION_KEY = {tunnel_encryption_key}
|
||||
irm {script_url} | iex
|
||||
"###,
|
||||
aether_url = powershell_single_quote(&session.aether_url),
|
||||
management_token = powershell_single_quote(&session.management_token),
|
||||
node_name = powershell_single_quote(&session.node_name),
|
||||
tunnel_security = powershell_single_quote(&session.tunnel_security),
|
||||
tunnel_encryption_key = powershell_single_quote(&session.tunnel_encryption_key),
|
||||
script_url = powershell_single_quote(TUNNEL_INSTALL_POWERSHELL_SCRIPT_URL),
|
||||
)
|
||||
}
|
||||
@@ -664,6 +684,8 @@ pub(crate) async fn build_proxy_node_install_session_response(
|
||||
aether_url: base_url_from_request(headers, request_context),
|
||||
management_token,
|
||||
node_name,
|
||||
tunnel_security: "non_tls_required".to_string(),
|
||||
tunnel_encryption_key: generate_tunnel_encryption_key(),
|
||||
expires_at_unix_secs,
|
||||
};
|
||||
let serialized = match serde_json::to_string(&session) {
|
||||
@@ -712,9 +734,7 @@ pub(super) async fn maybe_build_local_install_response(
|
||||
if decision.route_family.as_deref() != Some("install") {
|
||||
return None;
|
||||
}
|
||||
if request_context.request_path.starts_with("/install-tunnel/")
|
||||
|| request_context.request_path.starts_with("/install-proxy/")
|
||||
{
|
||||
if request_context.request_path.starts_with("/install-tunnel/") {
|
||||
return Some(maybe_build_local_tunnel_install_response(state, request_context).await);
|
||||
}
|
||||
let Some((code, wants_powershell)) = install_code_from_path(&request_context.request_path)
|
||||
@@ -893,6 +913,8 @@ mod tests {
|
||||
aether_url: "https://aether.example".to_string(),
|
||||
management_token: "ae-test-token".to_string(),
|
||||
node_name: "jp-proxy-01".to_string(),
|
||||
tunnel_security: "non_tls_required".to_string(),
|
||||
tunnel_encryption_key: "base64-32-bytes".to_string(),
|
||||
expires_at_unix_secs: u64::MAX,
|
||||
}
|
||||
}
|
||||
@@ -907,10 +929,6 @@ mod tests {
|
||||
tunnel_install_code_from_path("/install-tunnel/abc123.ps1"),
|
||||
Some(("abc123".to_string(), true))
|
||||
);
|
||||
assert_eq!(
|
||||
tunnel_install_code_from_path("/install-proxy/abc123"),
|
||||
Some(("abc123".to_string(), false))
|
||||
);
|
||||
assert_eq!(tunnel_install_code_from_path("/install-tunnel/a/b"), None);
|
||||
}
|
||||
|
||||
@@ -921,8 +939,10 @@ mod tests {
|
||||
assert!(script.contains("export AETHER_TUNNEL_AETHER_URL='https://aether.example'"));
|
||||
assert!(script.contains("export AETHER_TUNNEL_MANAGEMENT_TOKEN='ae-test-token'"));
|
||||
assert!(script.contains("export AETHER_TUNNEL_NODE_NAME='jp-proxy-01'"));
|
||||
assert!(script.contains("export AETHER_TUNNEL_SECURITY='non_tls_required'"));
|
||||
assert!(script.contains("export AETHER_TUNNEL_ENCRYPTION_KEY='base64-32-bytes'"));
|
||||
assert!(script.contains(
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh"
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/refs/heads/main/apps/aether-tunnel/install.sh"
|
||||
));
|
||||
assert!(!script.contains("aether-rust-pioneer"));
|
||||
assert!(!script.contains("[[servers]]"));
|
||||
@@ -935,6 +955,8 @@ mod tests {
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_AETHER_URL = 'https://aether.example'"));
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_MANAGEMENT_TOKEN = 'ae-test-token'"));
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_NODE_NAME = 'jp-proxy-01'"));
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_SECURITY = 'non_tls_required'"));
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_ENCRYPTION_KEY = 'base64-32-bytes'"));
|
||||
assert!(script.contains(
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1"
|
||||
));
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use std::fmt::Debug;
|
||||
use std::future::Future;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
|
||||
use axum::{body::Body, response::Response};
|
||||
use tokio::time::timeout;
|
||||
use tracing::warn;
|
||||
|
||||
use super::models_responses::{
|
||||
build_claude_model_detail_response, build_claude_models_list_response,
|
||||
@@ -15,6 +19,59 @@ use super::models_shared::{
|
||||
};
|
||||
use super::{query_param_value, AppState, GatewayPublicRequestContext};
|
||||
|
||||
#[cfg(not(test))]
|
||||
const MODELS_ROUTE_READ_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
#[cfg(test)]
|
||||
const MODELS_ROUTE_READ_TIMEOUT: Duration = Duration::from_millis(50);
|
||||
|
||||
async fn await_models_route_read<T, E, Fut>(operation: &'static str, future: Fut) -> Option<T>
|
||||
where
|
||||
E: Debug,
|
||||
Fut: Future<Output = Result<T, E>>,
|
||||
{
|
||||
match timeout(MODELS_ROUTE_READ_TIMEOUT, future).await {
|
||||
Ok(Ok(value)) => Some(value),
|
||||
Ok(Err(error)) => {
|
||||
warn!(
|
||||
event_name = "models_route_read_error",
|
||||
log_type = "ops",
|
||||
operation,
|
||||
error = ?error,
|
||||
"gateway local models route read failed"
|
||||
);
|
||||
None
|
||||
}
|
||||
Err(_) => {
|
||||
warn!(
|
||||
event_name = "models_route_read_timeout",
|
||||
log_type = "ops",
|
||||
operation,
|
||||
timeout_ms = MODELS_ROUTE_READ_TIMEOUT.as_millis() as u64,
|
||||
"gateway local models route read timed out"
|
||||
);
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn build_models_read_fallback_response(
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
api_format: &str,
|
||||
) -> Response<Body> {
|
||||
let route_kind = request_context
|
||||
.control_decision
|
||||
.as_ref()
|
||||
.and_then(|decision| decision.route_kind.as_deref());
|
||||
match route_kind {
|
||||
Some("detail") => {
|
||||
let model_id = models_detail_id(&request_context.request_path)
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
build_models_not_found_response(&model_id, api_format)
|
||||
}
|
||||
_ => build_empty_models_list_response(api_format),
|
||||
}
|
||||
}
|
||||
|
||||
fn sort_and_dedup_model_rows(
|
||||
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
@@ -47,10 +104,11 @@ async fn list_model_rows_for_client_format(
|
||||
) -> Option<Vec<StoredMinimalCandidateSelectionRow>> {
|
||||
let mut collected = Vec::new();
|
||||
for query_format in models_query_api_formats(api_format) {
|
||||
let rows = state
|
||||
.list_minimal_candidate_selection_rows_for_api_format(query_format)
|
||||
.await
|
||||
.ok()?;
|
||||
let rows = await_models_route_read(
|
||||
"candidate_selection_by_api_format",
|
||||
state.list_minimal_candidate_selection_rows_for_api_format(query_format),
|
||||
)
|
||||
.await?;
|
||||
let mut filtered = filter_rows_for_models(rows, auth_snapshot, query_format);
|
||||
collected.append(&mut filtered);
|
||||
}
|
||||
@@ -65,13 +123,14 @@ async fn list_model_rows_for_client_format_and_global_model(
|
||||
) -> Option<Vec<StoredMinimalCandidateSelectionRow>> {
|
||||
let mut collected = Vec::new();
|
||||
for query_format in models_query_api_formats(api_format) {
|
||||
let rows = state
|
||||
.list_minimal_candidate_selection_rows_for_api_format_and_global_model(
|
||||
let rows = await_models_route_read(
|
||||
"candidate_selection_by_global_model",
|
||||
state.list_minimal_candidate_selection_rows_for_api_format_and_global_model(
|
||||
query_format,
|
||||
global_model_name,
|
||||
)
|
||||
.await
|
||||
.ok()?;
|
||||
),
|
||||
)
|
||||
.await?;
|
||||
let mut filtered = filter_rows_for_models(rows, auth_snapshot, query_format);
|
||||
collected.append(&mut filtered);
|
||||
}
|
||||
@@ -96,21 +155,38 @@ pub(super) async fn maybe_build_local_models_route_response(
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs();
|
||||
let auth_snapshot = state
|
||||
.data
|
||||
.read_auth_api_key_snapshot(
|
||||
let auth_snapshot = match await_models_route_read(
|
||||
"auth_api_key_snapshot",
|
||||
state.data.read_auth_api_key_snapshot(
|
||||
&auth_context.user_id,
|
||||
&auth_context.api_key_id,
|
||||
now_unix_secs,
|
||||
)
|
||||
.await
|
||||
.ok()
|
||||
.flatten();
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Some(snapshot) => snapshot,
|
||||
None => {
|
||||
return Some(build_models_read_fallback_response(
|
||||
request_context,
|
||||
api_format,
|
||||
))
|
||||
}
|
||||
};
|
||||
let auth_snapshot = auth_snapshot.as_ref();
|
||||
|
||||
match decision.route_kind.as_deref() {
|
||||
Some("list") => {
|
||||
let rows = list_model_rows_for_client_format(state, api_format, auth_snapshot).await?;
|
||||
let rows =
|
||||
match list_model_rows_for_client_format(state, api_format, auth_snapshot).await {
|
||||
Some(rows) => rows,
|
||||
None => {
|
||||
return Some(build_models_read_fallback_response(
|
||||
request_context,
|
||||
api_format,
|
||||
))
|
||||
}
|
||||
};
|
||||
if rows.is_empty() {
|
||||
return Some(build_empty_models_list_response(api_format));
|
||||
}
|
||||
@@ -156,13 +232,22 @@ pub(super) async fn maybe_build_local_models_route_response(
|
||||
}
|
||||
Some("detail") => {
|
||||
let model_id = models_detail_id(&request_context.request_path)?;
|
||||
let rows = list_model_rows_for_client_format_and_global_model(
|
||||
let rows = match list_model_rows_for_client_format_and_global_model(
|
||||
state,
|
||||
api_format,
|
||||
&model_id,
|
||||
auth_snapshot,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
{
|
||||
Some(rows) => rows,
|
||||
None => {
|
||||
return Some(build_models_read_fallback_response(
|
||||
request_context,
|
||||
api_format,
|
||||
))
|
||||
}
|
||||
};
|
||||
let Some(row) = rows.first() else {
|
||||
return Some(build_models_not_found_response(&model_id, api_format));
|
||||
};
|
||||
|
||||
@@ -14,6 +14,8 @@ use super::{
|
||||
provider_catalog_key_supports_format, query_param_value, AppState, GatewayPublicRequestContext,
|
||||
};
|
||||
|
||||
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
|
||||
|
||||
pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
@@ -287,12 +289,10 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
for (name, value) in &provider_request_headers {
|
||||
upstream_request = upstream_request.header(name, value);
|
||||
}
|
||||
if let Some(total_ms) =
|
||||
crate::provider_transport::resolve_transport_execution_timeouts(&transport)
|
||||
.and_then(|timeouts| timeouts.total_ms.or(timeouts.first_byte_ms))
|
||||
{
|
||||
upstream_request = upstream_request.timeout(Duration::from_millis(total_ms));
|
||||
}
|
||||
let total_ms = crate::provider_transport::resolve_transport_execution_timeouts(&transport)
|
||||
.and_then(|timeouts| timeouts.total_ms)
|
||||
.unwrap_or(DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS);
|
||||
upstream_request = upstream_request.timeout(Duration::from_millis(total_ms));
|
||||
|
||||
let response = match upstream_request.json(&provider_request_body).send().await {
|
||||
Ok(response) => response,
|
||||
|
||||
@@ -355,10 +355,58 @@ fn infer_client_family_from_user_agent(user_agent: &str) -> Option<&'static str>
|
||||
if normalized.contains("geminicli") || normalized.contains("gemini-cli") {
|
||||
return Some("gemini_cli");
|
||||
}
|
||||
if normalized.contains("qwencode") {
|
||||
return Some("qwen_code");
|
||||
}
|
||||
if normalized.contains("roo-code") || normalized.contains("roocode") {
|
||||
return Some("roo_code");
|
||||
}
|
||||
if normalized.contains("kilo-code") || normalized.contains("kilocode") {
|
||||
return Some("kilocode");
|
||||
}
|
||||
if normalized.contains("cherrystudio") || normalized.contains("cherry-studio") {
|
||||
return Some("cherrystudio");
|
||||
}
|
||||
if normalized.contains("openui-agent-manager") || normalized.contains("openui") {
|
||||
return Some("openui");
|
||||
}
|
||||
if normalized.contains("cursor") {
|
||||
return Some("cursor");
|
||||
}
|
||||
if normalized.contains("windsurf") {
|
||||
return Some("windsurf");
|
||||
}
|
||||
if normalized.contains("continue") {
|
||||
return Some("continue");
|
||||
}
|
||||
if normalized.contains("cline") {
|
||||
return Some("cline");
|
||||
}
|
||||
if normalized.contains("aider") {
|
||||
return Some("aider");
|
||||
}
|
||||
if normalized.contains("langchain") {
|
||||
return Some("langchain");
|
||||
}
|
||||
if normalized.contains("llamaindex") || normalized.contains("llama-index") {
|
||||
return Some("llamaindex");
|
||||
}
|
||||
if normalized.starts_with("openai/js") {
|
||||
return Some("openai_js_sdk");
|
||||
}
|
||||
None
|
||||
if normalized.starts_with("openai/python") {
|
||||
return Some("openai_python_sdk");
|
||||
}
|
||||
if normalized.starts_with("anthropic/js") || normalized.contains("anthropic-sdk-typescript") {
|
||||
return Some("anthropic_js_sdk");
|
||||
}
|
||||
if normalized.starts_with("anthropic/python") || normalized.contains("anthropic-sdk-python") {
|
||||
return Some("anthropic_python_sdk");
|
||||
}
|
||||
if normalized.contains("/js ") || normalized.contains("/python ") {
|
||||
return Some("sdk");
|
||||
}
|
||||
Some("unknown")
|
||||
}
|
||||
|
||||
fn users_me_usage_client_family(item: &StoredRequestUsageAudit) -> Option<&str> {
|
||||
@@ -451,6 +499,12 @@ fn build_users_me_usage_record_payload(
|
||||
if item.target_model.is_some() {
|
||||
payload["target_model"] = json!(item.target_model.clone());
|
||||
}
|
||||
if let Some(reasoning_effort) = item.provider_reasoning_effort() {
|
||||
payload["reasoning_effort"] = json!(reasoning_effort);
|
||||
}
|
||||
if let Some(service_tier) = item.provider_service_tier() {
|
||||
payload["service_tier"] = json!(service_tier);
|
||||
}
|
||||
if include_actual_cost {
|
||||
payload["actual_cost"] = json!(round_to(item.actual_total_cost_usd, 6));
|
||||
payload["rate_multiplier"] = json!(rate_multiplier);
|
||||
@@ -510,6 +564,12 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
|
||||
.expect("object")
|
||||
.remove("target_model");
|
||||
}
|
||||
if let Some(reasoning_effort) = item.provider_reasoning_effort() {
|
||||
payload["reasoning_effort"] = json!(reasoning_effort);
|
||||
}
|
||||
if let Some(service_tier) = item.provider_service_tier() {
|
||||
payload["service_tier"] = json!(service_tier);
|
||||
}
|
||||
payload
|
||||
}
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ use super::enabled_key_capability_short_names;
|
||||
use crate::handlers::shared::unix_secs_to_rfc3339;
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
use crate::AppState;
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use serde_json::json;
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -173,10 +174,10 @@ pub(crate) async fn build_admin_keys_grouped_by_format_payload(
|
||||
.get("health_score")
|
||||
.and_then(serde_json::Value::as_f64)
|
||||
.unwrap_or(1.0),
|
||||
"circuit_breaker_open": format_circuit
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false),
|
||||
"circuit_breaker_open": provider_key_circuit_payload_is_active_open_at(
|
||||
&format_circuit,
|
||||
now_unix_secs,
|
||||
),
|
||||
"last_used_at": key.last_used_at_unix_secs.and_then(unix_secs_to_rfc3339),
|
||||
"created_at": unix_secs_to_rfc3339(key.created_at_unix_ms.unwrap_or(now_unix_secs)),
|
||||
"updated_at": unix_secs_to_rfc3339(key.updated_at_unix_secs.unwrap_or(now_unix_secs)),
|
||||
|
||||
@@ -13,6 +13,7 @@ use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKe
|
||||
use aether_provider_pool::{
|
||||
grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier,
|
||||
};
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::borrow::Cow;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -438,27 +439,53 @@ fn chatgpt_web_image_quota_limit(
|
||||
metadata: &Map<String, Value>,
|
||||
remaining: Option<f64>,
|
||||
) -> Option<f64> {
|
||||
let explicit_limit = metadata
|
||||
.get("image_quota_total")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||
.filter(|value| *value > 0.0);
|
||||
let plan_type = metadata
|
||||
.get("plan_type")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
if plan_type.as_deref() == Some("free") {
|
||||
return Some(25.0);
|
||||
}
|
||||
|
||||
let explicit_limit = metadata
|
||||
.get("image_quota_total")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||
.filter(|value| *value > 0.0);
|
||||
let limit_source = metadata
|
||||
.get("image_quota_limit_source")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
if let Some(limit) = explicit_limit {
|
||||
return Some(limit);
|
||||
if !chatgpt_web_image_quota_limit_is_legacy_free_default(
|
||||
limit,
|
||||
limit_source,
|
||||
plan_type.as_deref(),
|
||||
remaining,
|
||||
) {
|
||||
return Some(limit);
|
||||
}
|
||||
}
|
||||
|
||||
remaining.filter(|value| *value > 0.0)
|
||||
}
|
||||
|
||||
fn chatgpt_web_image_quota_limit_is_legacy_free_default(
|
||||
limit: f64,
|
||||
limit_source: Option<&str>,
|
||||
plan_type: Option<&str>,
|
||||
remaining: Option<f64>,
|
||||
) -> bool {
|
||||
let plan_type_is_free = plan_type
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case("free"));
|
||||
if !plan_type_is_free || limit_source.is_some() {
|
||||
return false;
|
||||
}
|
||||
if (limit - 25.0).abs() > f64::EPSILON {
|
||||
return false;
|
||||
}
|
||||
remaining.is_none_or(|value| value < limit)
|
||||
}
|
||||
|
||||
fn model_quota_window_snapshot(
|
||||
model_name: &str,
|
||||
item: &Map<String, Value>,
|
||||
@@ -1961,6 +1988,39 @@ pub(crate) fn provider_key_health_summary(
|
||||
Option<String>,
|
||||
bool,
|
||||
serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
provider_key_health_summary_with_circuit_predicate(key, |value| {
|
||||
value
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn provider_key_health_summary_at(
|
||||
key: &StoredProviderCatalogKey,
|
||||
now_unix_secs: u64,
|
||||
) -> (
|
||||
f64,
|
||||
i64,
|
||||
Option<String>,
|
||||
bool,
|
||||
serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
provider_key_health_summary_with_circuit_predicate(key, |value| {
|
||||
provider_key_circuit_payload_is_active_open_at(value, now_unix_secs)
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_key_health_summary_with_circuit_predicate(
|
||||
key: &StoredProviderCatalogKey,
|
||||
circuit_is_open: impl Fn(&serde_json::Value) -> bool,
|
||||
) -> (
|
||||
f64,
|
||||
i64,
|
||||
Option<String>,
|
||||
bool,
|
||||
serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
let health_by_format = key
|
||||
.health_by_format
|
||||
@@ -2003,12 +2063,7 @@ pub(crate) fn provider_key_health_summary(
|
||||
}
|
||||
}
|
||||
|
||||
let any_circuit_open = circuit_by_format.values().any(|value| {
|
||||
value
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
});
|
||||
let any_circuit_open = circuit_by_format.values().any(circuit_is_open);
|
||||
|
||||
(
|
||||
if health_by_format.is_empty() {
|
||||
@@ -2157,15 +2212,10 @@ pub(crate) fn build_admin_provider_key_response(
|
||||
last_failure_at,
|
||||
circuit_breaker_open,
|
||||
circuit_by_format,
|
||||
) = provider_key_health_summary(key);
|
||||
) = provider_key_health_summary_at(key, now_unix_secs);
|
||||
let circuit_sample = circuit_by_format
|
||||
.values()
|
||||
.find(|value| {
|
||||
value
|
||||
.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.find(|value| provider_key_circuit_payload_is_active_open_at(value, now_unix_secs))
|
||||
.or_else(|| circuit_by_format.values().next());
|
||||
let is_adaptive = key.rpm_limit.is_none();
|
||||
let effective_limit = if is_adaptive {
|
||||
@@ -2718,12 +2768,39 @@ mod tests {
|
||||
assert_eq!(quota.get("code"), Some(&json!("ok")));
|
||||
assert_eq!(quota.get("plan_type"), Some(&json!("free")));
|
||||
assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64)));
|
||||
assert_eq!(quota.get("usage_ratio"), Some(&json!(0.04)));
|
||||
assert_eq!(quota.get("usage_ratio"), Some(&json!(0.0)));
|
||||
assert_eq!(window.get("code"), Some(&json!("image_gen")));
|
||||
assert_eq!(window.get("remaining_value"), Some(&json!(24.0)));
|
||||
assert_eq!(window.get("limit_value"), Some(&json!(25.0)));
|
||||
assert_eq!(window.get("used_value"), Some(&json!(1.0)));
|
||||
assert_eq!(window.get("remaining_ratio"), Some(&json!(0.96)));
|
||||
assert_eq!(window.get("limit_value"), Some(&json!(24.0)));
|
||||
assert_eq!(window.get("used_value"), Some(&json!(0.0)));
|
||||
assert_eq!(window.get("remaining_ratio"), Some(&json!(1.0)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_ignores_chatgpt_web_legacy_free_25_limit() {
|
||||
let mut key = sample_catalog_key();
|
||||
key.upstream_metadata = Some(json!({
|
||||
"chatgpt_web": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"plan_type": "free",
|
||||
"image_quota_remaining": 19.0,
|
||||
"image_quota_total": 25.0
|
||||
}
|
||||
}));
|
||||
|
||||
let payload = provider_key_status_snapshot_payload(&key, "chatgpt_web");
|
||||
let window = payload
|
||||
.get("quota")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|quota| quota.get("windows"))
|
||||
.and_then(Value::as_array)
|
||||
.and_then(|windows| windows.first())
|
||||
.and_then(Value::as_object)
|
||||
.expect("image quota window should exist");
|
||||
|
||||
assert_eq!(window.get("remaining_value"), Some(&json!(19.0)));
|
||||
assert_eq!(window.get("limit_value"), Some(&json!(19.0)));
|
||||
assert_eq!(window.get("used_value"), Some(&json!(0.0)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -26,8 +26,9 @@ pub(crate) use self::catalog::{
|
||||
default_provider_key_status_snapshot, effective_catalog_encryption_key,
|
||||
encrypt_catalog_secret_with_fallbacks, masked_catalog_api_key, parse_catalog_auth_config_json,
|
||||
provider_catalog_key_supports_format, provider_key_health_summary,
|
||||
provider_key_status_snapshot_payload, sync_provider_key_oauth_status_snapshot,
|
||||
sync_provider_key_quota_status_snapshot, take_secret_prefix, take_secret_suffix,
|
||||
provider_key_health_summary_at, provider_key_status_snapshot_payload,
|
||||
sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot,
|
||||
take_secret_prefix, take_secret_suffix,
|
||||
};
|
||||
pub(crate) use self::email_templates::{
|
||||
admin_email_template_definition, admin_email_template_html_key,
|
||||
|
||||
Reference in New Issue
Block a user