Merge pull request #289 from AAEE86/rust

fix(usage): 统一使用记录与仪表盘的缓存命中率计算口径
This commit is contained in:
fawney19
2026-04-12 16:12:22 +08:00
committed by GitHub
33 changed files with 3333 additions and 456 deletions

View File

@@ -26,7 +26,7 @@ fn applies_codex_defaults_when_body_rules_do_not_handle_fields() {
}
#[test]
fn defers_to_user_body_rules_for_handled_fields() {
fn strips_store_for_compact_even_when_body_rules_handle_it() {
let body_rules = json!([
{"action":"set","path":"store","value":true},
{"action":"set","path":"instructions","value":"Keep custom"},
@@ -51,7 +51,7 @@ fn defers_to_user_body_rules_for_handled_fields() {
);
assert!(body.get("max_output_tokens").is_none());
assert_eq!(body["store"], true);
assert!(body.get("store").is_none());
assert_eq!(body["instructions"], "Keep custom");
assert_eq!(body["metadata"]["mode"], "custom");
assert_eq!(body["top_p"], 0.5);

View File

@@ -267,4 +267,32 @@ mod tests {
assert_eq!(converted["store"], false);
assert_eq!(converted["instructions"], "You are GPT-5.");
}
#[test]
fn strips_store_for_openai_compact_requests() {
let request = json!({
"model": "gpt-5",
"messages": [{
"role": "user",
"content": "Hello from OpenAI Chat"
}],
"store": true
});
let converted = build_standard_request_body(
&request,
"openai:chat",
"gpt-5",
"openai",
"openai:compact",
"/v1/chat/completions",
false,
None,
None,
)
.expect("openai chat should convert to openai compact");
assert_eq!(converted["model"], "gpt-5");
assert!(converted.get("store").is_none());
}
}

View File

@@ -1,6 +1,5 @@
use serde_json::Value;
use super::super::codex::apply_codex_openai_cli_special_body_edits;
use crate::ai_pipeline::conversion::{request_conversion_kind, RequestConversionKind};
use crate::ai_pipeline::transport::apply_local_body_rules;
use crate::ai_pipeline::transport::url::{
@@ -8,6 +7,7 @@ use crate::ai_pipeline::transport::url::{
build_openai_cli_url, build_passthrough_path_url,
};
use crate::ai_pipeline::{
apply_codex_openai_cli_special_body_edits, apply_openai_compact_special_body_edits,
build_cross_format_openai_chat_request_body as pipeline_build_cross_format_openai_chat_request_body,
build_local_openai_chat_request_body as pipeline_build_local_openai_chat_request_body,
GatewayProviderTransportSnapshot,
@@ -74,6 +74,7 @@ pub(crate) fn build_cross_format_openai_chat_request_body(
body_rules,
user_api_key_id,
);
apply_openai_compact_special_body_edits(&mut provider_request_body, provider_api_format);
Some(provider_request_body)
}

View File

@@ -3,7 +3,6 @@ use std::collections::BTreeMap;
use serde_json::Value;
use url::form_urlencoded;
use super::super::codex::apply_codex_openai_cli_special_body_edits;
use crate::ai_pipeline::conversion::{request_conversion_kind, RequestConversionKind};
use crate::ai_pipeline::transport::antigravity::{
build_antigravity_v1internal_url, AntigravityRequestUrlAction,
@@ -14,6 +13,7 @@ use crate::ai_pipeline::transport::url::{
build_openai_cli_url, build_passthrough_path_url,
};
use crate::ai_pipeline::{
apply_codex_openai_cli_special_body_edits, apply_openai_compact_special_body_edits,
build_cross_format_openai_cli_request_body as pipeline_build_cross_format_openai_cli_request_body,
build_local_openai_cli_request_body as pipeline_build_local_openai_cli_request_body,
GatewayProviderTransportSnapshot,
@@ -40,6 +40,7 @@ pub(crate) fn build_local_openai_cli_request_body(
body_rules,
user_api_key_id,
);
apply_openai_compact_special_body_edits(&mut provider_request_body, provider_api_format);
Some(provider_request_body)
}
@@ -70,6 +71,7 @@ pub(crate) fn build_cross_format_openai_cli_request_body(
body_rules,
user_api_key_id,
);
apply_openai_compact_special_body_edits(&mut provider_request_body, provider_api_format);
Some(provider_request_body)
}

View File

@@ -81,6 +81,28 @@ fn local_openai_cli_wrapper_preserves_body_order_after_edits() {
);
}
#[test]
fn local_openai_compact_wrapper_strips_store_for_same_format_requests() {
let body_json = json!({
"model": "gpt-5.4",
"input": [],
"store": true
});
let provider_request_body = build_local_openai_cli_request_body(
&body_json,
"gpt-5.4",
false,
"openai",
"openai:compact",
None,
None,
)
.expect("local openai compact body should build");
assert!(provider_request_body.get("store").is_none());
}
#[test]
fn strips_metadata_for_codex_openai_cli_requests() {
let body_json = json!({

View File

@@ -3,15 +3,15 @@ pub(crate) use aether_ai_pipeline::api::{
aggregate_openai_chat_stream_sync_response, aggregate_openai_cli_stream_sync_response,
aggregate_standard_chat_stream_sync_response, aggregate_standard_cli_stream_sync_response,
apply_codex_openai_cli_special_body_edits, apply_codex_openai_cli_special_headers,
augment_sync_report_context, build_core_error_body_for_client_format,
build_cross_format_openai_chat_request_body, build_cross_format_openai_cli_request_body,
build_generated_tool_call_id, build_kiro_final_message_sse_events,
build_kiro_initial_sse_events, build_kiro_stream_error_sse_events,
build_local_openai_chat_request_body, build_local_openai_cli_request_body,
build_local_success_background_report, build_local_success_conversion_background_report,
build_openai_cli_response, build_standard_request_body,
build_standard_request_body_from_canonical, build_standard_upstream_url,
calculate_kiro_context_input_tokens, canonicalize_tool_arguments,
apply_openai_compact_special_body_edits, augment_sync_report_context,
build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body,
build_cross_format_openai_cli_request_body, build_generated_tool_call_id,
build_kiro_final_message_sse_events, build_kiro_initial_sse_events,
build_kiro_stream_error_sse_events, build_local_openai_chat_request_body,
build_local_openai_cli_request_body, build_local_success_background_report,
build_local_success_conversion_background_report, build_openai_cli_response,
build_standard_request_body, build_standard_request_body_from_canonical,
build_standard_upstream_url, calculate_kiro_context_input_tokens, canonicalize_tool_arguments,
convert_claude_chat_response_to_openai_chat, convert_claude_cli_response_to_openai_cli,
convert_gemini_chat_response_to_openai_chat, convert_gemini_cli_response_to_openai_cli,
convert_openai_chat_request_to_claude_request, convert_openai_chat_request_to_gemini_request,

View File

@@ -4,8 +4,9 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::query_param_value;
use crate::GatewayError;
use aether_admin::observability::usage::{
admin_usage_bad_request_response, admin_usage_data_unavailable_response,
admin_usage_matches_optional_id, admin_usage_parse_recent_hours,
admin_usage_bad_request_response, admin_usage_cache_creation_tokens,
admin_usage_data_unavailable_response, admin_usage_matches_optional_id,
admin_usage_parse_recent_hours, admin_usage_total_input_context,
ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
};
use axum::{
@@ -50,10 +51,9 @@ pub(super) async fn build_admin_usage_cache_affinity_hit_analysis_response(
.iter()
.map(|item| item.cache_read_input_tokens)
.sum();
let total_cache_creation_tokens: u64 = filtered
.iter()
.map(|item| item.cache_creation_input_tokens)
.sum();
let total_cache_creation_tokens: u64 =
filtered.iter().map(admin_usage_cache_creation_tokens).sum();
let total_input_context: u64 = filtered.iter().map(admin_usage_total_input_context).sum();
let total_cache_read_cost: f64 = filtered.iter().map(|item| item.cache_read_cost_usd).sum();
let total_cache_creation_cost: f64 = filtered
.iter()
@@ -63,12 +63,11 @@ pub(super) async fn build_admin_usage_cache_affinity_hit_analysis_response(
.iter()
.filter(|item| item.cache_read_input_tokens > 0)
.count();
let total_context_tokens = total_input_tokens.saturating_add(total_cache_read_tokens);
let token_cache_hit_rate = if total_context_tokens == 0 {
let token_cache_hit_rate = if total_input_context == 0 {
0.0
} else {
round_to(
total_cache_read_tokens as f64 / total_context_tokens as f64 * 100.0,
total_cache_read_tokens as f64 / total_input_context as f64 * 100.0,
2,
)
};

View File

@@ -16,12 +16,12 @@ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_model_fetch::{
aggregate_models_for_cache, build_models_fetch_execution_plan,
endpoint_supports_rust_models_fetch, extract_error_message, parse_models_response,
aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
preset_models_for_provider, selected_models_fetch_endpoints,
};
use axum::{body::Body, http::Response, response::IntoResponse, Json};
use serde_json::{json, Value};
use std::collections::{BTreeMap, BTreeSet};
use std::collections::BTreeSet;
pub(crate) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE: &str =
"Rust local provider-query model test is not configured";
@@ -31,22 +31,15 @@ const ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL: &str =
"No active endpoints found for this provider";
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL: &str =
"No models returned from any endpoint";
const PROVIDER_QUERY_FETCH_FORMAT_PRIORITY: &[&[&str]] = &[
&[
"openai:chat",
"openai:responses",
"openai:cli",
"openai:compact",
],
&["claude:chat", "claude:cli"],
&["gemini:chat", "gemini:cli"],
];
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL: &str = "No models returned from any key";
const ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX: &str = "upstream_models_provider:";
#[derive(Debug)]
struct ProviderQueryKeyFetchResult {
models: Vec<Value>,
error: Option<String>,
from_cache: bool,
has_success: bool,
}
fn provider_query_provider_payload(provider: &StoredProviderCatalogProvider) -> Value {
@@ -66,64 +59,6 @@ fn provider_query_key_display_name(key: &StoredProviderCatalogKey) -> String {
}
}
fn provider_query_normalize_api_format(value: &str) -> String {
value.trim().to_ascii_lowercase()
}
fn provider_query_selected_fetch_endpoints(
endpoints: &[StoredProviderCatalogEndpoint],
key: &StoredProviderCatalogKey,
) -> Vec<StoredProviderCatalogEndpoint> {
let allowed_api_formats = key
.api_formats
.as_ref()
.and_then(Value::as_array)
.map(|items| {
items
.iter()
.filter_map(Value::as_str)
.map(provider_query_normalize_api_format)
.filter(|value| !value.is_empty())
.collect::<BTreeSet<_>>()
})
.filter(|items| !items.is_empty());
let mut by_format = BTreeMap::<String, StoredProviderCatalogEndpoint>::new();
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
let api_format = provider_query_normalize_api_format(&endpoint.api_format);
if api_format.is_empty() || !endpoint_supports_rust_models_fetch(&api_format) {
continue;
}
if allowed_api_formats
.as_ref()
.is_some_and(|formats| !formats.contains(&api_format))
{
continue;
}
by_format.insert(api_format, endpoint.clone());
}
// 与 Python 版本保持一致:同族优先使用 chat 端点,其次才回退到其他抓取格式。
let covered_formats = PROVIDER_QUERY_FETCH_FORMAT_PRIORITY
.iter()
.flat_map(|items| items.iter().copied())
.collect::<BTreeSet<_>>();
let mut selected = PROVIDER_QUERY_FETCH_FORMAT_PRIORITY
.iter()
.filter_map(|candidates| {
candidates
.iter()
.find_map(|api_format| by_format.remove(*api_format))
})
.collect::<Vec<_>>();
selected.extend(
by_format
.into_iter()
.filter(|(api_format, _)| !covered_formats.contains(api_format.as_str()))
.map(|(_, endpoint)| endpoint),
);
selected
}
async fn provider_query_read_cached_models(
state: &AdminAppState<'_>,
provider_id: &str,
@@ -147,38 +82,97 @@ async fn provider_query_read_cached_models(
Some(aggregate_models_for_cache(&parsed))
}
async fn provider_query_fetch_models_from_transport(
async fn provider_query_read_provider_cached_models(
state: &AdminAppState<'_>,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
) -> Result<Vec<Value>, String> {
let plan = build_models_fetch_execution_plan(state.app(), transport).await?;
let result = execution_runtime::execute_execution_runtime_sync_plan(state.app(), None, &plan)
provider_id: &str,
) -> Option<Vec<Value>> {
let runner = state.app().redis_kv_runner()?;
let cache_key = runner.keyspace().key(&format!(
"{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}"
));
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| format!("{err:?}"))?;
.ok()?;
let raw = redis::cmd("GET")
.arg(&cache_key)
.query_async::<Option<String>>(&mut connection)
.await
.ok()??;
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
Some(aggregate_models_for_cache(&parsed))
}
if result.status_code != 200 {
let message = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(extract_error_message)
.or_else(|| {
result.error.as_ref().and_then(|error| {
let message = error.message.trim();
(!message.is_empty()).then_some(message.to_string())
async fn provider_query_write_provider_cached_models(
state: &AdminAppState<'_>,
provider_id: &str,
models: &[Value],
) {
let Some(runner) = state.app().redis_kv_runner() else {
return;
};
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
return;
};
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
let _ = runner
.setex(
&cache_key,
&serialized,
Some(aether_model_fetch::model_fetch_interval_minutes().saturating_mul(60)),
)
.await;
}
fn provider_query_antigravity_tier_weight(raw_auth_config: Option<&str>) -> i32 {
raw_auth_config
.and_then(|value| serde_json::from_str::<Value>(value).ok())
.and_then(|value| value.get("tier").cloned())
.and_then(|value| value.as_str().map(ToOwned::to_owned))
.map(|tier| match tier.trim().to_ascii_lowercase().as_str() {
"ultra" => 3,
"pro" => 2,
"free" => 1,
_ => 0,
})
.unwrap_or(0)
}
async fn provider_query_sort_antigravity_keys(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
keys: Vec<StoredProviderCatalogKey>,
) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> {
let mut ranked = Vec::new();
for key in keys {
let availability = if key.oauth_invalid_at_unix_secs.is_some() {
0
} else {
1
};
let tier_weight = if let Some(endpoint) = selected_models_fetch_endpoints(endpoints, &key)
.into_iter()
.next()
{
state
.app()
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
.map(|transport| {
provider_query_antigravity_tier_weight(
transport.key.decrypted_auth_config.as_deref(),
)
})
})
.unwrap_or_else(|| format!("upstream returned status {}", result.status_code));
return Err(message);
.unwrap_or(0)
} else {
0
};
ranked.push(((availability, tier_weight), key));
}
let body_json = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.ok_or_else(|| "models fetch response body is missing JSON payload".to_string())?;
let parsed = parse_models_response(&transport.endpoint.api_format, body_json)?;
Ok(parsed.cached_models)
ranked.sort_by(|left, right| right.0.cmp(&left.0));
Ok(ranked.into_iter().map(|(_, key)| key).collect())
}
async fn provider_query_fetch_models_for_key(
@@ -196,20 +190,30 @@ async fn provider_query_fetch_models_for_key(
models: cached_models,
error: None,
from_cache: true,
has_success: true,
});
}
}
let selected_endpoints = provider_query_selected_fetch_endpoints(endpoints, key);
let selected_endpoints = selected_models_fetch_endpoints(endpoints, key);
if selected_endpoints.is_empty() {
if let Some(models) = preset_models_for_provider(&provider.provider_type) {
return Ok(ProviderQueryKeyFetchResult {
models: aggregate_models_for_cache(&models),
error: None,
from_cache: false,
has_success: true,
});
}
return Ok(ProviderQueryKeyFetchResult {
models: Vec::new(),
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
from_cache: false,
has_success: false,
});
}
let mut all_models = Vec::new();
let mut transports = Vec::new();
let mut all_errors = Vec::new();
for endpoint in selected_endpoints {
let Some(transport) = state
@@ -223,14 +227,34 @@ async fn provider_query_fetch_models_for_key(
));
continue;
};
match provider_query_fetch_models_from_transport(state, &transport).await {
Ok(models) => all_models.extend(models),
Err(err) => all_errors.push(err),
}
transports.push(transport);
}
let unique_models = aggregate_models_for_cache(&all_models);
if !unique_models.is_empty() {
if transports.is_empty() {
return Ok(ProviderQueryKeyFetchResult {
models: Vec::new(),
error: Some(all_errors.join("; ")),
from_cache: false,
has_success: false,
});
}
let outcome = match fetch_models_from_transports(state.app(), &transports).await {
Ok(outcome) => outcome,
Err(err) => {
all_errors.push(err);
return Ok(ProviderQueryKeyFetchResult {
models: Vec::new(),
error: Some(all_errors.join("; ")),
from_cache: false,
has_success: false,
});
}
};
all_errors.extend(outcome.errors);
let unique_models = aggregate_models_for_cache(&outcome.cached_models);
if outcome.has_success && !unique_models.is_empty() {
<AppState as ModelFetchRuntimeState>::write_upstream_models_cache(
state.app(),
&provider.id,
@@ -253,6 +277,7 @@ async fn provider_query_fetch_models_for_key(
models: unique_models,
error,
from_cache: false,
has_success: outcome.has_success,
})
}
@@ -317,7 +342,10 @@ pub(crate) async fn build_admin_provider_query_models_response(
.into_response());
}
let active_keys = keys.iter().filter(|key| key.is_active).collect::<Vec<_>>();
let active_keys = keys
.into_iter()
.filter(|key| key.is_active)
.collect::<Vec<_>>();
if active_keys.is_empty() {
return Ok(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
@@ -325,11 +353,45 @@ pub(crate) async fn build_admin_provider_query_models_response(
}
let active_key_count = active_keys.len();
if provider
.provider_type
.trim()
.eq_ignore_ascii_case("antigravity")
&& !force_refresh
{
if let Some(models) = provider_query_read_provider_cached_models(state, &provider.id).await
{
return Ok(Json(json!({
"success": !models.is_empty(),
"data": {
"models": models,
"error": serde_json::Value::Null,
"from_cache": true,
"keys_total": active_key_count,
"keys_cached": active_key_count,
"keys_fetched": 0,
},
"provider": provider_query_provider_payload(&provider),
}))
.into_response());
}
}
let ordered_keys = if provider
.provider_type
.trim()
.eq_ignore_ascii_case("antigravity")
{
provider_query_sort_antigravity_keys(state, &provider, &endpoints, active_keys).await?
} else {
active_keys
};
let mut all_models = Vec::new();
let mut all_errors = Vec::new();
let mut cache_hit_count = 0usize;
let mut fetch_count = 0usize;
for key in active_keys {
for key in &ordered_keys {
let result =
provider_query_fetch_models_for_key(state, &provider, &endpoints, key, force_refresh)
.await?;
@@ -346,9 +408,25 @@ pub(crate) async fn build_admin_provider_query_models_response(
} else {
fetch_count += 1;
}
if provider
.provider_type
.trim()
.eq_ignore_ascii_case("antigravity")
&& result.has_success
{
break;
}
}
let models = aggregate_models_for_cache(&all_models);
if provider
.provider_type
.trim()
.eq_ignore_ascii_case("antigravity")
&& !models.is_empty()
{
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
@@ -356,7 +434,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
Some(all_errors.join("; "))
};
if !success && error.is_none() {
error = Some("No models returned from any key".to_string());
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
}
Ok(Json(json!({

View File

@@ -2,6 +2,7 @@ use super::{
build_auth_error_response, query_param_value, resolve_authenticated_local_user, AppState,
GatewayError, GatewayPublicRequestContext,
};
use aether_billing::normalize_total_input_context_for_cache_hit_rate;
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageAuditListQuery};
use axum::{
body::Body,
@@ -28,6 +29,7 @@ struct DashboardUsageTotals {
total_tokens: u64,
cache_creation_tokens: u64,
cache_read_tokens: u64,
cache_hit_total_input_context: u64,
cache_creation_cost_usd: f64,
cache_read_cost_usd: f64,
total_cost_usd: f64,
@@ -69,12 +71,14 @@ pub(super) fn decision_route_kind(request_context: &GatewayPublicRequestContext)
impl DashboardUsageTotals {
fn record(&mut self, item: &StoredRequestUsageAudit) {
let cache_creation_tokens = dashboard_cache_creation_tokens(item);
self.requests += 1;
self.input_tokens += item.input_tokens;
self.output_tokens += item.output_tokens;
self.total_tokens += item.total_tokens;
self.cache_creation_tokens += item.cache_creation_input_tokens;
self.cache_creation_tokens += cache_creation_tokens;
self.cache_read_tokens += item.cache_read_input_tokens;
self.cache_hit_total_input_context += dashboard_total_input_context(item);
self.cache_creation_cost_usd += item.cache_creation_cost_usd;
self.cache_read_cost_usd += item.cache_read_cost_usd;
self.total_cost_usd += item.total_cost_usd;
@@ -102,15 +106,50 @@ impl DashboardUsageTotals {
}
fn cache_hit_rate(&self) -> f64 {
let total_cache_tokens = self.cache_creation_tokens + self.cache_read_tokens;
if total_cache_tokens == 0 {
if self.cache_hit_total_input_context == 0 {
0.0
} else {
dashboard_round_f64(self.cache_read_tokens as f64 / total_cache_tokens as f64, 4)
dashboard_round_f64(
self.cache_read_tokens as f64 / self.cache_hit_total_input_context as f64 * 100.0,
2,
)
}
}
}
fn dashboard_usage_should_count_in_summary(item: &StoredRequestUsageAudit) -> bool {
!matches!(item.status.as_str(), "pending" | "streaming")
&& !matches!(item.provider_name.as_str(), "unknown" | "pending")
}
fn dashboard_cache_creation_tokens(item: &StoredRequestUsageAudit) -> u64 {
let classified = item
.cache_creation_ephemeral_5m_input_tokens
.saturating_add(item.cache_creation_ephemeral_1h_input_tokens);
if item.cache_creation_input_tokens == 0 && classified > 0 {
classified
} else {
item.cache_creation_input_tokens
}
}
fn dashboard_total_input_context(item: &StoredRequestUsageAudit) -> u64 {
let api_format = item
.endpoint_api_format
.as_deref()
.or(item.api_format.as_deref());
let input_tokens = i64::try_from(item.input_tokens).unwrap_or(i64::MAX);
let cache_creation_tokens =
i64::try_from(dashboard_cache_creation_tokens(item)).unwrap_or(i64::MAX);
let cache_read_tokens = i64::try_from(item.cache_read_input_tokens).unwrap_or(i64::MAX);
normalize_total_input_context_for_cache_hit_rate(
api_format,
input_tokens,
cache_creation_tokens,
cache_read_tokens,
) as u64
}
fn dashboard_round_f64(value: f64, decimals: u32) -> f64 {
let factor = 10_f64.powi(i32::try_from(decimals).unwrap_or_default());
(value * factor).round() / factor
@@ -128,6 +167,39 @@ fn dashboard_format_integer(value: u64) -> String {
formatted.chars().rev().collect()
}
fn dashboard_trimmed_decimal(value: f64, decimals: usize) -> String {
let mut formatted = format!("{value:.decimals$}");
while formatted.contains('.') && formatted.ends_with('0') {
formatted.pop();
}
if formatted.ends_with('.') {
formatted.pop();
}
formatted
}
fn dashboard_format_token_compact(value: u64) -> String {
if value < 1_000 {
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);
}
let decimals = if thousands >= 10.0 { 1 } else { 2 };
return format!("{}K", dashboard_trimmed_decimal(thousands, decimals));
}
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))
}
fn dashboard_format_usd(value: f64) -> String {
format!("${:.2}", dashboard_round_f64(value, 2))
}
@@ -144,6 +216,16 @@ fn dashboard_format_token_subvalue(totals: &DashboardUsageTotals) -> String {
)
}
fn dashboard_format_today_token_subvalue(totals: &DashboardUsageTotals) -> String {
format!(
"输入 {} / 输出 {} · 写缓存 {} / 读缓存 {}",
dashboard_format_token_compact(totals.input_tokens),
dashboard_format_token_compact(totals.output_tokens),
dashboard_format_token_compact(totals.cache_creation_tokens),
dashboard_format_token_compact(totals.cache_read_tokens)
)
}
fn dashboard_parse_tz_offset_minutes(query: Option<&str>) -> Result<i32, String> {
query_param_value(query, "tz_offset_minutes")
.map(|value| {
@@ -392,7 +474,10 @@ async fn dashboard_list_usage_for_range(
})
.await
{
Ok(value) => Ok(value),
Ok(mut value) => {
value.retain(dashboard_usage_should_count_in_summary);
Ok(value)
}
Err(err) => Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
@@ -578,8 +663,8 @@ pub(super) async fn handle_dashboard_stats_get(
},
{
"name": "今日 Token",
"value": dashboard_format_integer(today_totals.total_tokens),
"subValue": dashboard_format_token_subvalue(&today_totals),
"value": dashboard_format_token_compact(today_totals.total_tokens),
"subValue": dashboard_format_today_token_subvalue(&today_totals),
"icon": "Zap",
},
{

View File

@@ -1,6 +1,8 @@
use std::collections::{BTreeMap, BTreeSet};
use aether_billing::normalize_input_tokens_for_billing;
use aether_billing::{
normalize_input_tokens_for_billing, normalize_total_input_context_for_cache_hit_rate,
};
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageAuditListQuery};
use axum::{
body::Body,
@@ -101,9 +103,20 @@ fn users_me_usage_cache_creation_tokens(item: &StoredRequestUsageAudit) -> u64 {
}
fn users_me_usage_total_input_context(item: &StoredRequestUsageAudit) -> u64 {
item.input_tokens
.saturating_add(users_me_usage_cache_creation_tokens(item))
.saturating_add(item.cache_read_input_tokens)
let api_format = item
.endpoint_api_format
.as_deref()
.or(item.api_format.as_deref());
let input_tokens = i64::try_from(item.input_tokens).unwrap_or(i64::MAX);
let cache_creation_tokens =
i64::try_from(users_me_usage_cache_creation_tokens(item)).unwrap_or(i64::MAX);
let cache_read_tokens = i64::try_from(item.cache_read_input_tokens).unwrap_or(i64::MAX);
normalize_total_input_context_for_cache_hit_rate(
api_format,
input_tokens,
cache_creation_tokens,
cache_read_tokens,
) as u64
}
fn users_me_usage_effective_input_tokens(item: &StoredRequestUsageAudit) -> u64 {

View File

@@ -2,21 +2,18 @@ use std::collections::HashMap;
use std::time::Duration;
use std::time::{SystemTime, UNIX_EPOCH};
use aether_contracts::ExecutionResult;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_model_fetch::{
apply_model_filters, build_models_fetch_execution_plan, extract_error_message,
json_string_list, model_fetch_interval_minutes, model_fetch_startup_delay_seconds,
model_fetch_startup_enabled, parse_models_response, select_models_fetch_endpoint,
apply_model_filters, fetch_models_from_transports, json_string_list, merge_upstream_metadata,
model_fetch_interval_minutes, model_fetch_startup_delay_seconds, model_fetch_startup_enabled,
preset_models_for_provider, selected_models_fetch_endpoints,
sync_provider_model_whitelist_associations, ModelFetchAssociationStore, ModelFetchRunSummary,
ModelsFetchSuccess,
};
use serde_json::{json, Value};
use tracing::{debug, info, warn};
use crate::provider_transport::GatewayProviderTransportSnapshot;
use crate::{AppState, GatewayError};
pub(crate) mod state;
@@ -26,8 +23,8 @@ use self::state::ModelFetchRuntimeState;
#[derive(Debug, Clone)]
struct SelectedFetchTarget {
provider: StoredProviderCatalogProvider,
endpoint: StoredProviderCatalogEndpoint,
key: StoredProviderCatalogKey,
endpoints: Vec<StoredProviderCatalogEndpoint>,
}
pub(crate) fn spawn_model_fetch_worker(state: AppState) -> Option<tokio::task::JoinHandle<()>> {
@@ -131,38 +128,12 @@ where
if !key.is_active || !key.auto_fetch_models {
continue;
}
if let Some(endpoint) = select_models_fetch_endpoint(&endpoints, &key) {
targets.push(SelectedFetchTarget {
provider: provider.clone(),
endpoint,
key,
});
} else {
targets.push(SelectedFetchTarget {
provider: provider.clone(),
endpoint: StoredProviderCatalogEndpoint::new(
"__unsupported__".to_string(),
provider.id.clone(),
"__unsupported__".to_string(),
None,
None,
false,
)
.expect("unsupported sentinel endpoint should build")
.with_transport_fields(
"https://unsupported.invalid".to_string(),
None,
None,
None,
None,
None,
None,
None,
)
.expect("unsupported sentinel endpoint transport should build"),
key,
});
}
let selected_endpoints = selected_models_fetch_endpoints(&endpoints, &key);
targets.push(SelectedFetchTarget {
provider: provider.clone(),
key,
endpoints: selected_endpoints,
});
}
}
@@ -215,7 +186,34 @@ async fn fetch_and_persist_key_models(
target: &SelectedFetchTarget,
) -> Result<KeyFetchDisposition, GatewayError> {
let now_unix_secs = now_unix_secs();
if target.endpoint.api_format == "__unsupported__" {
if target.endpoints.is_empty() {
if let Some(models) = preset_models_for_provider(&target.provider.provider_type) {
let fetched_model_ids = models
.iter()
.filter_map(|model| model.get("id"))
.filter_map(Value::as_str)
.map(ToOwned::to_owned)
.collect::<Vec<_>>();
let filtered_models = apply_model_filters(
&fetched_model_ids,
json_string_list(target.key.locked_models.as_ref()),
json_string_list(target.key.model_include_patterns.as_ref()),
json_string_list(target.key.model_exclude_patterns.as_ref()),
);
persist_key_fetch_success(state, &target.key, now_unix_secs, &filtered_models, None)
.await?;
state
.write_upstream_models_cache(&target.provider.id, &target.key.id, &models)
.await;
sync_provider_model_whitelist_associations(
state,
&target.provider.id,
&filtered_models,
)
.await
.map_err(GatewayError::Internal)?;
return Ok(KeyFetchDisposition::Succeeded);
}
persist_key_fetch_failure(
state,
&target.key,
@@ -226,10 +224,25 @@ async fn fetch_and_persist_key_models(
return Ok(KeyFetchDisposition::Skipped);
}
let Some(transport) = state
.read_provider_transport_snapshot(&target.provider.id, &target.endpoint.id, &target.key.id)
.await?
else {
let mut transports = Vec::new();
for endpoint in &target.endpoints {
match state
.read_provider_transport_snapshot(&target.provider.id, &endpoint.id, &target.key.id)
.await?
{
Some(transport) => transports.push(transport),
None => {
warn!(
provider_id = %target.provider.id,
endpoint_id = %endpoint.id,
key_id = %target.key.id,
"gateway model fetch transport snapshot unavailable"
);
}
}
}
if transports.is_empty() {
persist_key_fetch_failure(
state,
&target.key,
@@ -238,9 +251,9 @@ async fn fetch_and_persist_key_models(
)
.await?;
return Ok(KeyFetchDisposition::Skipped);
};
}
let result = match execute_models_fetch_request(state, &transport).await {
let result = match fetch_models_from_transports(state, &transports).await {
Ok(result) => result,
Err(err) => {
persist_key_fetch_failure(state, &target.key, now_unix_secs, err.clone()).await?;
@@ -254,6 +267,22 @@ async fn fetch_and_persist_key_models(
}
};
if !result.has_success {
let error = if result.errors.is_empty() {
"Upstream models fetch failed".to_string()
} else {
result.errors.join("; ")
};
persist_key_fetch_failure(state, &target.key, now_unix_secs, error.clone()).await?;
warn!(
provider_id = %target.provider.id,
key_id = %target.key.id,
message = %error,
"gateway model fetch failed"
);
return Ok(KeyFetchDisposition::Failed);
}
let filtered_models = apply_model_filters(
&result.fetched_model_ids,
json_string_list(target.key.locked_models.as_ref()),
@@ -261,7 +290,14 @@ async fn fetch_and_persist_key_models(
json_string_list(target.key.model_exclude_patterns.as_ref()),
);
persist_key_fetch_success(state, &target.key, now_unix_secs, &filtered_models).await?;
persist_key_fetch_success(
state,
&target.key,
now_unix_secs,
&filtered_models,
result.upstream_metadata.as_ref(),
)
.await?;
state
.write_upstream_models_cache(&target.provider.id, &target.key.id, &result.cached_models)
.await;
@@ -271,41 +307,6 @@ async fn fetch_and_persist_key_models(
Ok(KeyFetchDisposition::Succeeded)
}
async fn execute_models_fetch_request(
state: &(impl ModelFetchRuntimeState + ?Sized),
transport: &GatewayProviderTransportSnapshot,
) -> Result<ModelsFetchSuccess, String> {
let plan = build_models_fetch_execution_plan(state, transport).await?;
let result = state
.execute_execution_runtime_sync_plan(&plan)
.await
.map_err(|err| format!("{err:?}"))?;
if result.status_code != 200 {
let message = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(extract_error_message)
.or_else(|| {
result.error.as_ref().and_then(|error| {
let message = error.message.trim();
(!message.is_empty()).then_some(message.to_string())
})
})
.unwrap_or_else(|| format!("upstream returned status {}", result.status_code));
return Err(message);
}
let body_json = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.ok_or_else(|| "models fetch response body is missing JSON payload".to_string())?;
parse_models_response(&transport.endpoint.api_format, body_json)
}
async fn persist_key_fetch_failure(
state: &(impl ModelFetchRuntimeState + ?Sized),
key: &StoredProviderCatalogKey,
@@ -325,6 +326,7 @@ async fn persist_key_fetch_success(
key: &StoredProviderCatalogKey,
now_unix_secs: u64,
allowed_models: &[String],
upstream_metadata: Option<&Value>,
) -> Result<(), GatewayError> {
let mut updated = key.clone();
updated.allowed_models = if allowed_models.is_empty() {
@@ -332,6 +334,12 @@ async fn persist_key_fetch_success(
} else {
Some(json!(allowed_models))
};
if let Some(upstream_metadata) = upstream_metadata {
updated.upstream_metadata = Some(merge_upstream_metadata(
updated.upstream_metadata.as_ref(),
upstream_metadata,
));
}
updated.last_models_fetch_at_unix_secs = Some(now_unix_secs);
updated.last_models_fetch_error = None;
updated.updated_at_unix_secs = Some(now_unix_secs);
@@ -345,3 +353,559 @@ fn now_unix_secs() -> u64 {
.unwrap_or_default()
.as_secs()
}
#[cfg(test)]
mod tests {
use super::{perform_model_fetch_once_with_state, state::ModelFetchRuntimeState};
use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot};
use aether_data_contracts::repository::global_models::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, StoredAdminGlobalModelPage,
StoredAdminProviderModel, UpsertAdminProviderModelRecord,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_model_fetch::{
build_models_fetch_execution_plan, ModelFetchAssociationStore, ModelFetchTransportRuntime,
};
use async_trait::async_trait;
use serde_json::{json, Value};
use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, Mutex};
use crate::provider_transport::LocalResolvedOAuthRequestAuth;
use crate::GatewayError;
use aether_provider_transport::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
#[derive(Clone, Default)]
struct TestState {
providers: Arc<Vec<StoredProviderCatalogProvider>>,
endpoints: Arc<Vec<StoredProviderCatalogEndpoint>>,
keys: Arc<Mutex<Vec<StoredProviderCatalogKey>>>,
transports: Arc<HashMap<(String, String, String), GatewayProviderTransportSnapshot>>,
execution_results: Arc<Mutex<VecDeque<ExecutionResult>>>,
cached_models: Arc<Mutex<HashMap<(String, String), Vec<Value>>>>,
}
impl TestState {
fn new(
providers: Vec<StoredProviderCatalogProvider>,
endpoints: Vec<StoredProviderCatalogEndpoint>,
keys: Vec<StoredProviderCatalogKey>,
transports: HashMap<(String, String, String), GatewayProviderTransportSnapshot>,
execution_results: Vec<ExecutionResult>,
) -> Self {
Self {
providers: Arc::new(providers),
endpoints: Arc::new(endpoints),
keys: Arc::new(Mutex::new(keys)),
transports: Arc::new(transports),
execution_results: Arc::new(Mutex::new(VecDeque::from(execution_results))),
cached_models: Arc::new(Mutex::new(HashMap::new())),
}
}
fn key(&self, key_id: &str) -> StoredProviderCatalogKey {
self.keys
.lock()
.expect("keys mutex")
.iter()
.find(|key| key.id == key_id)
.cloned()
.expect("key should exist")
}
}
#[async_trait]
impl ModelFetchTransportRuntime for TestState {
async fn resolve_local_oauth_request_auth(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Result<Option<LocalResolvedOAuthRequestAuth>, String> {
if transport.key.auth_type.trim().eq_ignore_ascii_case("oauth") {
return Ok(Some(LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: "Bearer oauth-token".to_string(),
}));
}
Ok(None)
}
async fn resolve_model_fetch_proxy(
&self,
_transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
None
}
async fn execute_model_fetch_execution_plan(
&self,
_plan: &ExecutionPlan,
) -> Result<ExecutionResult, String> {
self.execution_results
.lock()
.expect("execution result mutex")
.pop_front()
.ok_or_else(|| "missing execution result".to_string())
}
}
#[async_trait]
impl ModelFetchAssociationStore for TestState {
type Error = String;
fn has_global_model_reader(&self) -> bool {
false
}
fn has_global_model_writer(&self) -> bool {
false
}
fn model_fetch_internal_error(&self, message: String) -> Self::Error {
message
}
async fn list_admin_provider_models(
&self,
_query: &AdminProviderModelListQuery,
) -> Result<Vec<StoredAdminProviderModel>, Self::Error> {
Ok(Vec::new())
}
async fn list_admin_global_models(
&self,
_query: &AdminGlobalModelListQuery,
) -> Result<StoredAdminGlobalModelPage, Self::Error> {
Ok(StoredAdminGlobalModelPage {
items: Vec::new(),
total: 0,
})
}
async fn create_admin_provider_model(
&self,
_record: &UpsertAdminProviderModelRecord,
) -> Result<Option<StoredAdminProviderModel>, Self::Error> {
Ok(None)
}
async fn list_provider_catalog_keys_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, Self::Error> {
Ok(self
.keys
.lock()
.expect("keys mutex")
.iter()
.filter(|key| {
provider_ids
.iter()
.any(|provider_id| provider_id == &key.provider_id)
})
.cloned()
.collect())
}
async fn delete_admin_provider_model(
&self,
_provider_id: &str,
_model_id: &str,
) -> Result<bool, Self::Error> {
Ok(false)
}
}
#[async_trait]
impl ModelFetchRuntimeState for TestState {
fn has_provider_catalog_data_reader(&self) -> bool {
true
}
fn has_provider_catalog_data_writer(&self) -> bool {
true
}
async fn list_provider_catalog_providers(
&self,
_active_only: bool,
) -> Result<Vec<StoredProviderCatalogProvider>, GatewayError> {
Ok(self.providers.as_ref().clone())
}
async fn list_provider_catalog_endpoints_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, GatewayError> {
Ok(self
.endpoints
.iter()
.filter(|endpoint| {
provider_ids
.iter()
.any(|provider_id| provider_id == &endpoint.provider_id)
})
.cloned()
.collect())
}
async fn read_provider_transport_snapshot(
&self,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
) -> Result<Option<GatewayProviderTransportSnapshot>, GatewayError> {
Ok(self
.transports
.get(&(
provider_id.to_string(),
endpoint_id.to_string(),
key_id.to_string(),
))
.cloned())
}
async fn execute_execution_runtime_sync_plan(
&self,
_plan: &ExecutionPlan,
) -> Result<ExecutionResult, GatewayError> {
Err(GatewayError::Internal(
"execute_execution_runtime_sync_plan should not be called".to_string(),
))
}
async fn update_provider_catalog_key(
&self,
key: &StoredProviderCatalogKey,
) -> Result<(), GatewayError> {
let mut keys = self.keys.lock().expect("keys mutex");
let Some(slot) = keys.iter_mut().find(|item| item.id == key.id) else {
return Err(GatewayError::Internal("key not found".to_string()));
};
*slot = key.clone();
Ok(())
}
async fn write_upstream_models_cache(
&self,
provider_id: &str,
key_id: &str,
cached_models: &[Value],
) {
self.cached_models.lock().expect("cache mutex").insert(
(provider_id.to_string(), key_id.to_string()),
cached_models.to_vec(),
);
}
}
fn sample_provider(provider_id: &str, provider_type: &str) -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
provider_id.to_string(),
provider_id.to_string(),
None,
provider_type.to_string(),
)
.expect("provider should build")
.with_transport_fields(true, false, false, None, None, None, None, None, None)
}
fn sample_endpoint(
endpoint_id: &str,
provider_id: &str,
api_format: &str,
) -> StoredProviderCatalogEndpoint {
StoredProviderCatalogEndpoint::new(
endpoint_id.to_string(),
provider_id.to_string(),
api_format.to_string(),
None,
None,
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://cloudcode-pa.googleapis.com".to_string(),
None,
None,
None,
None,
None,
None,
None,
)
.expect("endpoint transport should build")
}
fn sample_key(
key_id: &str,
provider_id: &str,
auth_type: &str,
api_formats: &[&str],
) -> StoredProviderCatalogKey {
let mut key = StoredProviderCatalogKey::new(
key_id.to_string(),
provider_id.to_string(),
"primary".to_string(),
auth_type.to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(json!(api_formats)),
"encrypted".to_string(),
None,
None,
None,
None,
None,
None,
None,
)
.expect("key transport should build");
key.auto_fetch_models = true;
key
}
fn sample_transport(
provider_type: &str,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
api_format: &str,
auth_type: &str,
decrypted_auth_config: Option<&str>,
) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: provider_id.to_string(),
name: provider_id.to_string(),
provider_type: provider_type.to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: endpoint_id.to_string(),
provider_id: provider_id.to_string(),
api_format: api_format.to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://cloudcode-pa.googleapis.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: key_id.to_string(),
provider_id: provider_id.to_string(),
name: "primary".to_string(),
auth_type: auth_type.to_string(),
is_active: true,
api_formats: Some(vec![api_format.to_string()]),
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: decrypted_auth_config.map(ToOwned::to_owned),
},
}
}
fn execution_result(body: Value) -> ExecutionResult {
ExecutionResult {
request_id: "req-1".to_string(),
candidate_id: None,
status_code: 200,
headers: Default::default(),
body: Some(aether_contracts::ResponseBody {
json_body: Some(body),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
}
}
#[tokio::test]
async fn gateway_runtime_state_supports_shared_models_fetch_plan_builder() {
let state = TestState::default();
let transport = sample_transport(
"openai",
"provider-openai",
"endpoint-openai-chat",
"key-openai-chat",
"openai:chat",
"api_key",
None,
);
let plan = build_models_fetch_execution_plan(&state, &transport)
.await
.expect("shared models fetch plan should build");
assert_eq!(plan.method, "GET");
assert_eq!(plan.provider_id, "provider-openai");
assert_eq!(plan.endpoint_id, "endpoint-openai-chat");
assert_eq!(plan.key_id, "key-openai-chat");
assert_eq!(plan.model_name.as_deref(), Some("models"));
}
#[tokio::test]
async fn model_fetch_uses_preset_models_without_endpoint() {
let provider = sample_provider("provider-codex", "codex");
let key = sample_key("key-codex", "provider-codex", "api_key", &["openai:cli"]);
let state = TestState::new(vec![provider], vec![], vec![key], HashMap::new(), vec![]);
let summary = perform_model_fetch_once_with_state(&state)
.await
.expect("fetch should succeed");
assert_eq!(summary.attempted, 1);
assert_eq!(summary.succeeded, 1);
let updated = state.key("key-codex");
let allowed_models = updated
.allowed_models
.and_then(|value| value.as_array().cloned())
.expect("allowed_models should be set");
assert!(allowed_models.iter().any(|model| model == "gpt-5.4"));
assert!(state
.cached_models
.lock()
.expect("cache mutex")
.contains_key(&("provider-codex".to_string(), "key-codex".to_string())));
}
#[tokio::test]
async fn model_fetch_merges_antigravity_metadata_and_preserves_reset_time() {
let provider = sample_provider("provider-antigravity", "antigravity");
let endpoint = sample_endpoint(
"endpoint-antigravity",
"provider-antigravity",
"gemini:chat",
);
let mut key = sample_key(
"key-antigravity",
"provider-antigravity",
"oauth",
&["gemini:chat"],
);
key.upstream_metadata = Some(json!({
"antigravity": {
"quota_by_model": {
"gemini-2.5-pro": {
"reset_time": "2026-04-12T00:00:00Z"
}
}
}
}));
let transport = sample_transport(
"antigravity",
"provider-antigravity",
"endpoint-antigravity",
"key-antigravity",
"gemini:chat",
"oauth",
Some(r#"{"project_id":"project-1","client_version":"1.2.3","session_id":"sess-1"}"#),
);
let state = TestState::new(
vec![provider],
vec![endpoint],
vec![key],
HashMap::from([(
(
"provider-antigravity".to_string(),
"endpoint-antigravity".to_string(),
"key-antigravity".to_string(),
),
transport,
)]),
vec![execution_result(json!({
"models": {
"gemini-2.5-pro": {
"displayName": "Gemini 2.5 Pro",
"quotaInfo": {
"remainingFraction": 0.25
}
}
}
}))],
);
let summary = perform_model_fetch_once_with_state(&state)
.await
.expect("fetch should succeed");
assert_eq!(summary.succeeded, 1);
let updated = state.key("key-antigravity");
assert_eq!(updated.allowed_models, Some(json!(["gemini-2.5-pro"])));
assert_eq!(
updated
.upstream_metadata
.as_ref()
.and_then(|value| value.get("antigravity"))
.and_then(|value| value.get("quota_by_model"))
.and_then(|value| value.get("gemini-2.5-pro"))
.and_then(|value| value.get("reset_time")),
Some(&json!("2026-04-12T00:00:00Z"))
);
}
#[tokio::test]
async fn model_fetch_failure_keeps_existing_allowed_models() {
let provider = sample_provider("provider-openai", "openai");
let endpoint = sample_endpoint(
"endpoint-openai-responses",
"provider-openai",
"openai:responses",
);
let mut key = sample_key(
"key-openai-responses",
"provider-openai",
"api_key",
&["openai:responses"],
);
key.allowed_models = Some(json!(["gpt-old"]));
let state = TestState::new(
vec![provider],
vec![endpoint],
vec![key],
HashMap::new(),
vec![],
);
let summary = perform_model_fetch_once_with_state(&state)
.await
.expect("fetch should finish");
assert_eq!(summary.succeeded, 0);
assert_eq!(summary.skipped, 1);
let updated = state.key("key-openai-responses");
assert_eq!(updated.allowed_models, Some(json!(["gpt-old"])));
assert_eq!(
updated.last_models_fetch_error.as_deref(),
Some("No supported endpoint for Rust models fetch")
);
}
}

View File

@@ -92,6 +92,19 @@ impl ModelFetchTransportRuntime for AppState {
) -> Option<ProxySnapshot> {
resolve_transport_proxy_snapshot_with_tunnel_affinity(self, transport).await
}
async fn execute_model_fetch_execution_plan(
&self,
plan: &ExecutionPlan,
) -> Result<ExecutionResult, String> {
execution_runtime::execute_execution_runtime_sync_plan(self, None, plan)
.await
.map_err(|err| match err {
GatewayError::UpstreamUnavailable { message, .. }
| GatewayError::ControlUnavailable { message, .. }
| GatewayError::Internal(message) => message,
})
}
}
#[async_trait]

View File

@@ -278,16 +278,16 @@ async fn gateway_handles_admin_provider_query_models_with_openai_responses_endpo
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["success"], json!(true));
assert_eq!(payload["data"]["error"], serde_json::Value::Null);
assert_eq!(payload["data"]["from_cache"], json!(false));
assert_eq!(payload["success"], json!(false));
assert_eq!(
payload["data"]["models"][0]["api_formats"],
json!(["openai:responses"])
payload["data"]["error"],
json!("No active endpoints found for this provider")
);
assert_eq!(payload["data"]["from_cache"], json!(false));
assert_eq!(payload["data"]["models"], json!([]));
assert_eq!(
*execution_runtime_hits.lock().expect("mutex should lock"),
1
0
);
gateway_handle.abort();
@@ -557,7 +557,7 @@ async fn gateway_handles_admin_provider_query_models_aggregating_active_keys() {
.iter()
.map(|model| model["id"].as_str().expect("id should exist"))
.collect::<Vec<_>>();
assert_eq!(model_ids, vec!["gpt-5", "gpt-4.1"]);
assert_eq!(model_ids, vec!["gpt-4.1", "gpt-5"]);
assert_eq!(
*execution_runtime_hits.lock().expect("mutex should lock"),
2
@@ -567,6 +567,81 @@ async fn gateway_handles_admin_provider_query_models_aggregating_active_keys() {
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_handles_admin_provider_query_models_for_fixed_provider_without_endpoint() {
let execution_runtime_hits = Arc::new(Mutex::new(0usize));
let execution_runtime_hits_clone = Arc::clone(&execution_runtime_hits);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |_request: Request| {
let execution_runtime_hits_inner = Arc::clone(&execution_runtime_hits_clone);
async move {
*execution_runtime_hits_inner
.lock()
.expect("mutex should lock") += 1;
Json(json!({
"request_id": "unexpected",
"status_code": 500
}))
}
}),
);
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let mut provider = sample_provider("provider-codex", "Codex", 10);
provider.provider_type = "codex".to_string();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![],
vec![sample_key(
"key-codex-oauth",
"provider-codex",
"openai:cli",
"sk-test-codex",
)],
));
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(GatewayDataState::with_provider_transport_reader_for_tests(
provider_catalog_repository,
DEVELOPMENT_ENCRYPTION_KEY.to_string(),
)),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/api/admin/provider-query/models"))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"provider_id": "provider-codex",
"api_key_id": "key-codex-oauth"
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["success"], json!(true));
assert_eq!(payload["data"]["error"], serde_json::Value::Null);
assert_eq!(payload["data"]["from_cache"], json!(false));
let models = payload["data"]["models"]
.as_array()
.expect("models should be an array");
assert!(models.iter().any(|model| model["id"] == "gpt-5.4"));
assert_eq!(
*execution_runtime_hits.lock().expect("mutex should lock"),
0
);
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_handles_admin_provider_query_test_model_locally_with_trusted_admin_principal() {
assert_admin_provider_query_route(

View File

@@ -393,11 +393,11 @@ async fn gateway_handles_admin_usage_aggregation_stats_locally_with_trusted_admi
assert_eq!(items[0]["request_count"], 2);
assert_eq!(items[0]["output_tokens"], 40);
assert_eq!(items[0]["effective_input_tokens"], 150);
assert_eq!(items[0]["total_input_context"], 200);
assert_eq!(items[0]["total_input_context"], 160);
assert_eq!(items[0]["cache_creation_tokens"], 30);
assert_eq!(items[0]["cache_creation_ephemeral_5m_tokens"], 12);
assert_eq!(items[0]["cache_creation_ephemeral_1h_tokens"], 18);
assert_eq!(items[0]["cache_hit_rate"], 5.0);
assert_eq!(items[0]["cache_hit_rate"], 6.25);
assert_eq!(items[1]["model"], "claude-3-7");
assert_eq!(items[1]["output_tokens"], 20);
@@ -1498,7 +1498,7 @@ async fn gateway_handles_admin_usage_cache_affinity_hit_analysis_locally_with_tr
assert_eq!(payload["total_input_tokens"], 140);
assert_eq!(payload["total_cache_read_tokens"], 50);
assert_eq!(payload["total_cache_creation_tokens"], 15);
assert_eq!(payload["token_cache_hit_rate"], 26.32);
assert_eq!(payload["token_cache_hit_rate"], 35.71);
assert_eq!(payload["total_cache_read_cost_usd"], 0.02);
assert_eq!(payload["total_cache_creation_cost_usd"], 0.015);
assert_eq!(payload["estimated_savings_usd"], 0.18);

View File

@@ -4763,7 +4763,7 @@ async fn gateway_handles_users_me_usage_locally_without_proxying_upstream() {
payload["summary_by_model"][0]["effective_input_tokens"],
105
);
assert_eq!(payload["summary_by_model"][0]["total_input_context"], 145);
assert_eq!(payload["summary_by_model"][0]["total_input_context"], 120);
assert_eq!(payload["billing"]["id"], "wallet-auth-1");
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);

View File

@@ -172,7 +172,7 @@ async fn gateway_handles_dashboard_stats_locally_without_proxying_upstream() {
#[tokio::test]
async fn gateway_handles_admin_dashboard_stats_locally_without_proxying_upstream() {
let now = Utc::now();
let now = stable_dashboard_now();
let admin = StoredUserAuthRecord::new(
"admin-auth-1".to_string(),
Some("admin@example.com".to_string()),
@@ -213,25 +213,54 @@ async fn gateway_handles_admin_dashboard_stats_locally_without_proxying_upstream
"refresh-dashboard-stats-admin",
now,
);
let mut openai_usage = sample_user_usage_audit(
"usage-dashboard-admin-1",
"req-dashboard-admin-1",
"user-auth-1",
"gpt-5",
"openai",
"completed",
now - chrono::Duration::minutes(10),
);
openai_usage.input_tokens = 12_000;
openai_usage.output_tokens = 3_000;
openai_usage.total_tokens = 15_000;
openai_usage.cache_creation_input_tokens = 1_200;
openai_usage.cache_creation_ephemeral_5m_input_tokens = 600;
openai_usage.cache_creation_ephemeral_1h_input_tokens = 600;
openai_usage.cache_read_input_tokens = 800;
let mut claude_usage = sample_user_usage_audit(
"usage-dashboard-admin-2",
"req-dashboard-admin-2",
"user-auth-2",
"claude-3-7",
"claude",
"completed",
now - chrono::Duration::minutes(5),
);
claude_usage.input_tokens = 900;
claude_usage.output_tokens = 100;
claude_usage.total_tokens = 1_000;
claude_usage.cache_creation_input_tokens = 50;
claude_usage.cache_creation_ephemeral_5m_input_tokens = 20;
claude_usage.cache_creation_ephemeral_1h_input_tokens = 30;
claude_usage.cache_read_input_tokens = 200;
let streaming_usage = sample_user_usage_audit(
"usage-dashboard-admin-3",
"req-dashboard-admin-3",
"user-auth-3",
"gpt-4.1",
"openai",
"streaming",
now - chrono::Duration::minutes(1),
);
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![
sample_user_usage_audit(
"usage-dashboard-admin-1",
"req-dashboard-admin-1",
"user-auth-1",
"gpt-5",
"openai",
"completed",
now - chrono::Duration::minutes(10),
),
sample_user_usage_audit(
"usage-dashboard-admin-2",
"req-dashboard-admin-2",
"user-auth-2",
"claude-3-7",
"claude",
"completed",
now - chrono::Duration::minutes(5),
),
openai_usage,
claude_usage,
streaming_usage,
]));
let user_repository = Arc::new(
InMemoryUserReadRepository::seed_auth_users(vec![admin.clone()]).with_export_users(vec![
@@ -367,6 +396,15 @@ async fn gateway_handles_admin_dashboard_stats_locally_without_proxying_upstream
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["today"]["requests"], 2);
assert_eq!(payload["today"]["tokens"], 16_000);
assert_eq!(payload["today"]["cost"], json!(2.5));
assert_eq!(payload["stats"][0]["value"], json!("2"));
assert_eq!(payload["stats"][1]["value"], json!("16K"));
assert_eq!(
payload["stats"][1]["subValue"],
json!("输入 12.9K / 输出 3.1K · 写缓存 1.25K / 读缓存 1K")
);
assert_eq!(payload["users"]["total"], 2);
assert_eq!(payload["users"]["active"], 1);
assert_eq!(payload["api_keys"]["total"], 3);