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

3
Cargo.lock generated
View File

@@ -206,8 +206,11 @@ dependencies = [
"aether-provider-transport", "aether-provider-transport",
"aether-scheduler-core", "aether-scheduler-core",
"async-trait", "async-trait",
"base64 0.22.1",
"regex", "regex",
"rsa",
"serde_json", "serde_json",
"sha2",
"tokio", "tokio",
"uuid", "uuid",
] ]

View File

@@ -26,7 +26,7 @@ fn applies_codex_defaults_when_body_rules_do_not_handle_fields() {
} }
#[test] #[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!([ let body_rules = json!([
{"action":"set","path":"store","value":true}, {"action":"set","path":"store","value":true},
{"action":"set","path":"instructions","value":"Keep custom"}, {"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!(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["instructions"], "Keep custom");
assert_eq!(body["metadata"]["mode"], "custom"); assert_eq!(body["metadata"]["mode"], "custom");
assert_eq!(body["top_p"], 0.5); assert_eq!(body["top_p"], 0.5);

View File

@@ -267,4 +267,32 @@ mod tests {
assert_eq!(converted["store"], false); assert_eq!(converted["store"], false);
assert_eq!(converted["instructions"], "You are GPT-5."); 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 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::conversion::{request_conversion_kind, RequestConversionKind};
use crate::ai_pipeline::transport::apply_local_body_rules; use crate::ai_pipeline::transport::apply_local_body_rules;
use crate::ai_pipeline::transport::url::{ use crate::ai_pipeline::transport::url::{
@@ -8,6 +7,7 @@ use crate::ai_pipeline::transport::url::{
build_openai_cli_url, build_passthrough_path_url, build_openai_cli_url, build_passthrough_path_url,
}; };
use crate::ai_pipeline::{ 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_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, build_local_openai_chat_request_body as pipeline_build_local_openai_chat_request_body,
GatewayProviderTransportSnapshot, GatewayProviderTransportSnapshot,
@@ -74,6 +74,7 @@ pub(crate) fn build_cross_format_openai_chat_request_body(
body_rules, body_rules,
user_api_key_id, user_api_key_id,
); );
apply_openai_compact_special_body_edits(&mut provider_request_body, provider_api_format);
Some(provider_request_body) Some(provider_request_body)
} }

View File

@@ -3,7 +3,6 @@ use std::collections::BTreeMap;
use serde_json::Value; use serde_json::Value;
use url::form_urlencoded; 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::conversion::{request_conversion_kind, RequestConversionKind};
use crate::ai_pipeline::transport::antigravity::{ use crate::ai_pipeline::transport::antigravity::{
build_antigravity_v1internal_url, AntigravityRequestUrlAction, build_antigravity_v1internal_url, AntigravityRequestUrlAction,
@@ -14,6 +13,7 @@ use crate::ai_pipeline::transport::url::{
build_openai_cli_url, build_passthrough_path_url, build_openai_cli_url, build_passthrough_path_url,
}; };
use crate::ai_pipeline::{ 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_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, build_local_openai_cli_request_body as pipeline_build_local_openai_cli_request_body,
GatewayProviderTransportSnapshot, GatewayProviderTransportSnapshot,
@@ -40,6 +40,7 @@ pub(crate) fn build_local_openai_cli_request_body(
body_rules, body_rules,
user_api_key_id, user_api_key_id,
); );
apply_openai_compact_special_body_edits(&mut provider_request_body, provider_api_format);
Some(provider_request_body) Some(provider_request_body)
} }
@@ -70,6 +71,7 @@ pub(crate) fn build_cross_format_openai_cli_request_body(
body_rules, body_rules,
user_api_key_id, user_api_key_id,
); );
apply_openai_compact_special_body_edits(&mut provider_request_body, provider_api_format);
Some(provider_request_body) 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] #[test]
fn strips_metadata_for_codex_openai_cli_requests() { fn strips_metadata_for_codex_openai_cli_requests() {
let body_json = json!({ 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_openai_chat_stream_sync_response, aggregate_openai_cli_stream_sync_response,
aggregate_standard_chat_stream_sync_response, aggregate_standard_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, apply_codex_openai_cli_special_body_edits, apply_codex_openai_cli_special_headers,
augment_sync_report_context, build_core_error_body_for_client_format, apply_openai_compact_special_body_edits, augment_sync_report_context,
build_cross_format_openai_chat_request_body, build_cross_format_openai_cli_request_body, build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body,
build_generated_tool_call_id, build_kiro_final_message_sse_events, build_cross_format_openai_cli_request_body, build_generated_tool_call_id,
build_kiro_initial_sse_events, build_kiro_stream_error_sse_events, build_kiro_final_message_sse_events, build_kiro_initial_sse_events,
build_local_openai_chat_request_body, build_local_openai_cli_request_body, build_kiro_stream_error_sse_events, build_local_openai_chat_request_body,
build_local_success_background_report, build_local_success_conversion_background_report, build_local_openai_cli_request_body, build_local_success_background_report,
build_openai_cli_response, build_standard_request_body, build_local_success_conversion_background_report, build_openai_cli_response,
build_standard_request_body_from_canonical, build_standard_upstream_url, build_standard_request_body, build_standard_request_body_from_canonical,
calculate_kiro_context_input_tokens, canonicalize_tool_arguments, 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_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_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, 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::handlers::admin::shared::query_param_value;
use crate::GatewayError; use crate::GatewayError;
use aether_admin::observability::usage::{ use aether_admin::observability::usage::{
admin_usage_bad_request_response, admin_usage_data_unavailable_response, admin_usage_bad_request_response, admin_usage_cache_creation_tokens,
admin_usage_matches_optional_id, admin_usage_parse_recent_hours, 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, ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
}; };
use axum::{ use axum::{
@@ -50,10 +51,9 @@ pub(super) async fn build_admin_usage_cache_affinity_hit_analysis_response(
.iter() .iter()
.map(|item| item.cache_read_input_tokens) .map(|item| item.cache_read_input_tokens)
.sum(); .sum();
let total_cache_creation_tokens: u64 = filtered let total_cache_creation_tokens: u64 =
.iter() filtered.iter().map(admin_usage_cache_creation_tokens).sum();
.map(|item| item.cache_creation_input_tokens) let total_input_context: u64 = filtered.iter().map(admin_usage_total_input_context).sum();
.sum();
let total_cache_read_cost: f64 = filtered.iter().map(|item| item.cache_read_cost_usd).sum(); let total_cache_read_cost: f64 = filtered.iter().map(|item| item.cache_read_cost_usd).sum();
let total_cache_creation_cost: f64 = filtered let total_cache_creation_cost: f64 = filtered
.iter() .iter()
@@ -63,12 +63,11 @@ pub(super) async fn build_admin_usage_cache_affinity_hit_analysis_response(
.iter() .iter()
.filter(|item| item.cache_read_input_tokens > 0) .filter(|item| item.cache_read_input_tokens > 0)
.count(); .count();
let total_context_tokens = total_input_tokens.saturating_add(total_cache_read_tokens); let token_cache_hit_rate = if total_input_context == 0 {
let token_cache_hit_rate = if total_context_tokens == 0 {
0.0 0.0
} else { } else {
round_to( 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, 2,
) )
}; };

View File

@@ -16,12 +16,12 @@ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
}; };
use aether_model_fetch::{ use aether_model_fetch::{
aggregate_models_for_cache, build_models_fetch_execution_plan, aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
endpoint_supports_rust_models_fetch, extract_error_message, parse_models_response, preset_models_for_provider, selected_models_fetch_endpoints,
}; };
use axum::{body::Body, http::Response, response::IntoResponse, Json}; use axum::{body::Body, http::Response, response::IntoResponse, Json};
use serde_json::{json, Value}; 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 = pub(crate) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE: &str =
"Rust local provider-query model test is not configured"; "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"; "No active endpoints found for this provider";
const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL: &str = const ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL: &str =
"No models returned from any endpoint"; "No models returned from any endpoint";
const PROVIDER_QUERY_FETCH_FORMAT_PRIORITY: &[&[&str]] = &[ 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:";
"openai:chat",
"openai:responses",
"openai:cli",
"openai:compact",
],
&["claude:chat", "claude:cli"],
&["gemini:chat", "gemini:cli"],
];
#[derive(Debug)] #[derive(Debug)]
struct ProviderQueryKeyFetchResult { struct ProviderQueryKeyFetchResult {
models: Vec<Value>, models: Vec<Value>,
error: Option<String>, error: Option<String>,
from_cache: bool, from_cache: bool,
has_success: bool,
} }
fn provider_query_provider_payload(provider: &StoredProviderCatalogProvider) -> Value { 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( async fn provider_query_read_cached_models(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider_id: &str, provider_id: &str,
@@ -147,38 +82,97 @@ async fn provider_query_read_cached_models(
Some(aggregate_models_for_cache(&parsed)) Some(aggregate_models_for_cache(&parsed))
} }
async fn provider_query_fetch_models_from_transport( async fn provider_query_read_provider_cached_models(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot, provider_id: &str,
) -> Result<Vec<Value>, String> { ) -> Option<Vec<Value>> {
let plan = build_models_fetch_execution_plan(state.app(), transport).await?; let runner = state.app().redis_kv_runner()?;
let result = execution_runtime::execute_execution_runtime_sync_plan(state.app(), None, &plan) let cache_key = runner.keyspace().key(&format!(
"{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}"
));
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await .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 { async fn provider_query_write_provider_cached_models(
let message = result state: &AdminAppState<'_>,
.body provider_id: &str,
.as_ref() models: &[Value],
.and_then(|body| body.json_body.as_ref()) ) {
.and_then(extract_error_message) let Some(runner) = state.app().redis_kv_runner() else {
.or_else(|| { return;
result.error.as_ref().and_then(|error| { };
let message = error.message.trim(); let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
(!message.is_empty()).then_some(message.to_string()) 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(0)
.unwrap_or_else(|| format!("upstream returned status {}", result.status_code)); } else {
return Err(message); 0
};
ranked.push(((availability, tier_weight), key));
} }
ranked.sort_by(|left, right| right.0.cmp(&left.0));
let body_json = result Ok(ranked.into_iter().map(|(_, key)| key).collect())
.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)
} }
async fn provider_query_fetch_models_for_key( async fn provider_query_fetch_models_for_key(
@@ -196,20 +190,30 @@ async fn provider_query_fetch_models_for_key(
models: cached_models, models: cached_models,
error: None, error: None,
from_cache: true, 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 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 { return Ok(ProviderQueryKeyFetchResult {
models: Vec::new(), models: Vec::new(),
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()), error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
from_cache: false, from_cache: false,
has_success: false,
}); });
} }
let mut all_models = Vec::new(); let mut transports = Vec::new();
let mut all_errors = Vec::new(); let mut all_errors = Vec::new();
for endpoint in selected_endpoints { for endpoint in selected_endpoints {
let Some(transport) = state let Some(transport) = state
@@ -223,14 +227,34 @@ async fn provider_query_fetch_models_for_key(
)); ));
continue; continue;
}; };
match provider_query_fetch_models_from_transport(state, &transport).await { transports.push(transport);
Ok(models) => all_models.extend(models),
Err(err) => all_errors.push(err),
}
} }
let unique_models = aggregate_models_for_cache(&all_models); if transports.is_empty() {
if !unique_models.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( <AppState as ModelFetchRuntimeState>::write_upstream_models_cache(
state.app(), state.app(),
&provider.id, &provider.id,
@@ -253,6 +277,7 @@ async fn provider_query_fetch_models_for_key(
models: unique_models, models: unique_models,
error, error,
from_cache: false, from_cache: false,
has_success: outcome.has_success,
}) })
} }
@@ -317,7 +342,10 @@ pub(crate) async fn build_admin_provider_query_models_response(
.into_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() { if active_keys.is_empty() {
return Ok(build_admin_provider_query_bad_request_response( return Ok(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL, 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(); 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_models = Vec::new();
let mut all_errors = Vec::new(); let mut all_errors = Vec::new();
let mut cache_hit_count = 0usize; let mut cache_hit_count = 0usize;
let mut fetch_count = 0usize; let mut fetch_count = 0usize;
for key in active_keys { for key in &ordered_keys {
let result = let result =
provider_query_fetch_models_for_key(state, &provider, &endpoints, key, force_refresh) provider_query_fetch_models_for_key(state, &provider, &endpoints, key, force_refresh)
.await?; .await?;
@@ -346,9 +408,25 @@ pub(crate) async fn build_admin_provider_query_models_response(
} else { } else {
fetch_count += 1; 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); 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 success = !models.is_empty();
let mut error = if all_errors.is_empty() { let mut error = if all_errors.is_empty() {
None None
@@ -356,7 +434,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
Some(all_errors.join("; ")) Some(all_errors.join("; "))
}; };
if !success && error.is_none() { 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!({ Ok(Json(json!({

View File

@@ -2,6 +2,7 @@ use super::{
build_auth_error_response, query_param_value, resolve_authenticated_local_user, AppState, build_auth_error_response, query_param_value, resolve_authenticated_local_user, AppState,
GatewayError, GatewayPublicRequestContext, GatewayError, GatewayPublicRequestContext,
}; };
use aether_billing::normalize_total_input_context_for_cache_hit_rate;
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageAuditListQuery}; use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageAuditListQuery};
use axum::{ use axum::{
body::Body, body::Body,
@@ -28,6 +29,7 @@ struct DashboardUsageTotals {
total_tokens: u64, total_tokens: u64,
cache_creation_tokens: u64, cache_creation_tokens: u64,
cache_read_tokens: u64, cache_read_tokens: u64,
cache_hit_total_input_context: u64,
cache_creation_cost_usd: f64, cache_creation_cost_usd: f64,
cache_read_cost_usd: f64, cache_read_cost_usd: f64,
total_cost_usd: f64, total_cost_usd: f64,
@@ -69,12 +71,14 @@ pub(super) fn decision_route_kind(request_context: &GatewayPublicRequestContext)
impl DashboardUsageTotals { impl DashboardUsageTotals {
fn record(&mut self, item: &StoredRequestUsageAudit) { fn record(&mut self, item: &StoredRequestUsageAudit) {
let cache_creation_tokens = dashboard_cache_creation_tokens(item);
self.requests += 1; self.requests += 1;
self.input_tokens += item.input_tokens; self.input_tokens += item.input_tokens;
self.output_tokens += item.output_tokens; self.output_tokens += item.output_tokens;
self.total_tokens += item.total_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_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_creation_cost_usd += item.cache_creation_cost_usd;
self.cache_read_cost_usd += item.cache_read_cost_usd; self.cache_read_cost_usd += item.cache_read_cost_usd;
self.total_cost_usd += item.total_cost_usd; self.total_cost_usd += item.total_cost_usd;
@@ -102,15 +106,50 @@ impl DashboardUsageTotals {
} }
fn cache_hit_rate(&self) -> f64 { fn cache_hit_rate(&self) -> f64 {
let total_cache_tokens = self.cache_creation_tokens + self.cache_read_tokens; if self.cache_hit_total_input_context == 0 {
if total_cache_tokens == 0 {
0.0 0.0
} else { } 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 { fn dashboard_round_f64(value: f64, decimals: u32) -> f64 {
let factor = 10_f64.powi(i32::try_from(decimals).unwrap_or_default()); let factor = 10_f64.powi(i32::try_from(decimals).unwrap_or_default());
(value * factor).round() / factor (value * factor).round() / factor
@@ -128,6 +167,39 @@ fn dashboard_format_integer(value: u64) -> String {
formatted.chars().rev().collect() 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 { fn dashboard_format_usd(value: f64) -> String {
format!("${:.2}", dashboard_round_f64(value, 2)) 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> { fn dashboard_parse_tz_offset_minutes(query: Option<&str>) -> Result<i32, String> {
query_param_value(query, "tz_offset_minutes") query_param_value(query, "tz_offset_minutes")
.map(|value| { .map(|value| {
@@ -392,7 +474,10 @@ async fn dashboard_list_usage_for_range(
}) })
.await .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( Err(err) => Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR, http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"), format!("{error_context}: {err:?}"),
@@ -578,8 +663,8 @@ pub(super) async fn handle_dashboard_stats_get(
}, },
{ {
"name": "今日 Token", "name": "今日 Token",
"value": dashboard_format_integer(today_totals.total_tokens), "value": dashboard_format_token_compact(today_totals.total_tokens),
"subValue": dashboard_format_token_subvalue(&today_totals), "subValue": dashboard_format_today_token_subvalue(&today_totals),
"icon": "Zap", "icon": "Zap",
}, },
{ {

View File

@@ -1,6 +1,8 @@
use std::collections::{BTreeMap, BTreeSet}; 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 aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageAuditListQuery};
use axum::{ use axum::{
body::Body, 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 { fn users_me_usage_total_input_context(item: &StoredRequestUsageAudit) -> u64 {
item.input_tokens let api_format = item
.saturating_add(users_me_usage_cache_creation_tokens(item)) .endpoint_api_format
.saturating_add(item.cache_read_input_tokens) .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 { 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::Duration;
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
use aether_contracts::ExecutionResult;
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
}; };
use aether_model_fetch::{ use aether_model_fetch::{
apply_model_filters, build_models_fetch_execution_plan, extract_error_message, apply_model_filters, fetch_models_from_transports, json_string_list, merge_upstream_metadata,
json_string_list, model_fetch_interval_minutes, model_fetch_startup_delay_seconds, model_fetch_interval_minutes, model_fetch_startup_delay_seconds, model_fetch_startup_enabled,
model_fetch_startup_enabled, parse_models_response, select_models_fetch_endpoint, preset_models_for_provider, selected_models_fetch_endpoints,
sync_provider_model_whitelist_associations, ModelFetchAssociationStore, ModelFetchRunSummary, sync_provider_model_whitelist_associations, ModelFetchAssociationStore, ModelFetchRunSummary,
ModelsFetchSuccess,
}; };
use serde_json::{json, Value}; use serde_json::{json, Value};
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use crate::provider_transport::GatewayProviderTransportSnapshot;
use crate::{AppState, GatewayError}; use crate::{AppState, GatewayError};
pub(crate) mod state; pub(crate) mod state;
@@ -26,8 +23,8 @@ use self::state::ModelFetchRuntimeState;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
struct SelectedFetchTarget { struct SelectedFetchTarget {
provider: StoredProviderCatalogProvider, provider: StoredProviderCatalogProvider,
endpoint: StoredProviderCatalogEndpoint,
key: StoredProviderCatalogKey, key: StoredProviderCatalogKey,
endpoints: Vec<StoredProviderCatalogEndpoint>,
} }
pub(crate) fn spawn_model_fetch_worker(state: AppState) -> Option<tokio::task::JoinHandle<()>> { 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 { if !key.is_active || !key.auto_fetch_models {
continue; continue;
} }
if let Some(endpoint) = select_models_fetch_endpoint(&endpoints, &key) { let selected_endpoints = selected_models_fetch_endpoints(&endpoints, &key);
targets.push(SelectedFetchTarget { targets.push(SelectedFetchTarget {
provider: provider.clone(), provider: provider.clone(),
endpoint, key,
key, endpoints: selected_endpoints,
}); });
} 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,
});
}
} }
} }
@@ -215,7 +186,34 @@ async fn fetch_and_persist_key_models(
target: &SelectedFetchTarget, target: &SelectedFetchTarget,
) -> Result<KeyFetchDisposition, GatewayError> { ) -> Result<KeyFetchDisposition, GatewayError> {
let now_unix_secs = now_unix_secs(); 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( persist_key_fetch_failure(
state, state,
&target.key, &target.key,
@@ -226,10 +224,25 @@ async fn fetch_and_persist_key_models(
return Ok(KeyFetchDisposition::Skipped); return Ok(KeyFetchDisposition::Skipped);
} }
let Some(transport) = state let mut transports = Vec::new();
.read_provider_transport_snapshot(&target.provider.id, &target.endpoint.id, &target.key.id) for endpoint in &target.endpoints {
.await? match state
else { .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( persist_key_fetch_failure(
state, state,
&target.key, &target.key,
@@ -238,9 +251,9 @@ async fn fetch_and_persist_key_models(
) )
.await?; .await?;
return Ok(KeyFetchDisposition::Skipped); 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, Ok(result) => result,
Err(err) => { Err(err) => {
persist_key_fetch_failure(state, &target.key, now_unix_secs, err.clone()).await?; 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( let filtered_models = apply_model_filters(
&result.fetched_model_ids, &result.fetched_model_ids,
json_string_list(target.key.locked_models.as_ref()), 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()), 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 state
.write_upstream_models_cache(&target.provider.id, &target.key.id, &result.cached_models) .write_upstream_models_cache(&target.provider.id, &target.key.id, &result.cached_models)
.await; .await;
@@ -271,41 +307,6 @@ async fn fetch_and_persist_key_models(
Ok(KeyFetchDisposition::Succeeded) 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( async fn persist_key_fetch_failure(
state: &(impl ModelFetchRuntimeState + ?Sized), state: &(impl ModelFetchRuntimeState + ?Sized),
key: &StoredProviderCatalogKey, key: &StoredProviderCatalogKey,
@@ -325,6 +326,7 @@ async fn persist_key_fetch_success(
key: &StoredProviderCatalogKey, key: &StoredProviderCatalogKey,
now_unix_secs: u64, now_unix_secs: u64,
allowed_models: &[String], allowed_models: &[String],
upstream_metadata: Option<&Value>,
) -> Result<(), GatewayError> { ) -> Result<(), GatewayError> {
let mut updated = key.clone(); let mut updated = key.clone();
updated.allowed_models = if allowed_models.is_empty() { updated.allowed_models = if allowed_models.is_empty() {
@@ -332,6 +334,12 @@ async fn persist_key_fetch_success(
} else { } else {
Some(json!(allowed_models)) 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_at_unix_secs = Some(now_unix_secs);
updated.last_models_fetch_error = None; updated.last_models_fetch_error = None;
updated.updated_at_unix_secs = Some(now_unix_secs); updated.updated_at_unix_secs = Some(now_unix_secs);
@@ -345,3 +353,559 @@ fn now_unix_secs() -> u64 {
.unwrap_or_default() .unwrap_or_default()
.as_secs() .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> { ) -> Option<ProxySnapshot> {
resolve_transport_proxy_snapshot_with_tunnel_affinity(self, transport).await 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] #[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); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["success"], json!(true)); assert_eq!(payload["success"], json!(false));
assert_eq!(payload["data"]["error"], serde_json::Value::Null);
assert_eq!(payload["data"]["from_cache"], json!(false));
assert_eq!( assert_eq!(
payload["data"]["models"][0]["api_formats"], payload["data"]["error"],
json!(["openai:responses"]) json!("No active endpoints found for this provider")
); );
assert_eq!(payload["data"]["from_cache"], json!(false));
assert_eq!(payload["data"]["models"], json!([]));
assert_eq!( assert_eq!(
*execution_runtime_hits.lock().expect("mutex should lock"), *execution_runtime_hits.lock().expect("mutex should lock"),
1 0
); );
gateway_handle.abort(); gateway_handle.abort();
@@ -557,7 +557,7 @@ async fn gateway_handles_admin_provider_query_models_aggregating_active_keys() {
.iter() .iter()
.map(|model| model["id"].as_str().expect("id should exist")) .map(|model| model["id"].as_str().expect("id should exist"))
.collect::<Vec<_>>(); .collect::<Vec<_>>();
assert_eq!(model_ids, vec!["gpt-5", "gpt-4.1"]); assert_eq!(model_ids, vec!["gpt-4.1", "gpt-5"]);
assert_eq!( assert_eq!(
*execution_runtime_hits.lock().expect("mutex should lock"), *execution_runtime_hits.lock().expect("mutex should lock"),
2 2
@@ -567,6 +567,81 @@ async fn gateway_handles_admin_provider_query_models_aggregating_active_keys() {
execution_runtime_handle.abort(); 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] #[tokio::test]
async fn gateway_handles_admin_provider_query_test_model_locally_with_trusted_admin_principal() { async fn gateway_handles_admin_provider_query_test_model_locally_with_trusted_admin_principal() {
assert_admin_provider_query_route( 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]["request_count"], 2);
assert_eq!(items[0]["output_tokens"], 40); assert_eq!(items[0]["output_tokens"], 40);
assert_eq!(items[0]["effective_input_tokens"], 150); 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_tokens"], 30);
assert_eq!(items[0]["cache_creation_ephemeral_5m_tokens"], 12); 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_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]["model"], "claude-3-7");
assert_eq!(items[1]["output_tokens"], 20); 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_input_tokens"], 140);
assert_eq!(payload["total_cache_read_tokens"], 50); assert_eq!(payload["total_cache_read_tokens"], 50);
assert_eq!(payload["total_cache_creation_tokens"], 15); 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_read_cost_usd"], 0.02);
assert_eq!(payload["total_cache_creation_cost_usd"], 0.015); assert_eq!(payload["total_cache_creation_cost_usd"], 0.015);
assert_eq!(payload["estimated_savings_usd"], 0.18); 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"], payload["summary_by_model"][0]["effective_input_tokens"],
105 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!(payload["billing"]["id"], "wallet-auth-1");
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); 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] #[tokio::test]
async fn gateway_handles_admin_dashboard_stats_locally_without_proxying_upstream() { async fn gateway_handles_admin_dashboard_stats_locally_without_proxying_upstream() {
let now = Utc::now(); let now = stable_dashboard_now();
let admin = StoredUserAuthRecord::new( let admin = StoredUserAuthRecord::new(
"admin-auth-1".to_string(), "admin-auth-1".to_string(),
Some("admin@example.com".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", "refresh-dashboard-stats-admin",
now, 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![ let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![
sample_user_usage_audit( openai_usage,
"usage-dashboard-admin-1", claude_usage,
"req-dashboard-admin-1", streaming_usage,
"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),
),
])); ]));
let user_repository = Arc::new( let user_repository = Arc::new(
InMemoryUserReadRepository::seed_auth_users(vec![admin.clone()]).with_export_users(vec![ 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); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["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"]["total"], 2);
assert_eq!(payload["users"]["active"], 1); assert_eq!(payload["users"]["active"], 1);
assert_eq!(payload["api_keys"]["total"], 3); assert_eq!(payload["api_keys"]["total"], 3);

View File

@@ -1,5 +1,7 @@
use crate::observability::stats::{aggregate_usage_stats, parse_bounded_u32, round_to}; use crate::observability::stats::{aggregate_usage_stats, parse_bounded_u32, round_to};
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::repository::users::StoredUserSummary; use aether_data::repository::users::StoredUserSummary;
use aether_data_contracts::repository::{ use aether_data_contracts::repository::{
provider_catalog::{StoredProviderCatalogEndpoint, StoredProviderCatalogProvider}, provider_catalog::{StoredProviderCatalogEndpoint, StoredProviderCatalogProvider},
@@ -318,9 +320,20 @@ pub fn admin_usage_cache_creation_tokens(item: &StoredRequestUsageAudit) -> u64
} }
pub fn admin_usage_total_input_context(item: &StoredRequestUsageAudit) -> u64 { pub fn admin_usage_total_input_context(item: &StoredRequestUsageAudit) -> u64 {
item.input_tokens let api_format = item
.saturating_add(admin_usage_cache_creation_tokens(item)) .endpoint_api_format
.saturating_add(item.cache_read_input_tokens) .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(admin_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
} }
pub fn admin_usage_effective_input_tokens(item: &StoredRequestUsageAudit) -> u64 { pub fn admin_usage_effective_input_tokens(item: &StoredRequestUsageAudit) -> u64 {
@@ -358,13 +371,15 @@ pub fn admin_usage_aggregation_by_model_json(
limit: usize, limit: usize,
) -> Value { ) -> Value {
#[allow(clippy::type_complexity)] #[allow(clippy::type_complexity)]
let mut grouped: BTreeMap<String, (u64, u64, u64, u64, u64, u64, u64, u64, u64, f64, f64)> = let mut grouped: BTreeMap<
BTreeMap::new(); String,
(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, f64, f64),
> = BTreeMap::new();
for item in usage { for item in usage {
let key = item.model.clone(); let key = item.model.clone();
let entry = grouped let entry = grouped
.entry(key) .entry(key)
.or_insert((0, 0, 0, 0, 0, 0, 0, 0, 0, 0.0, 0.0)); .or_insert((0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0.0, 0.0));
entry.0 = entry.0.saturating_add(1); entry.0 = entry.0.saturating_add(1);
entry.1 = entry.1.saturating_add(item.total_tokens); entry.1 = entry.1.saturating_add(item.total_tokens);
entry.2 = entry.2.saturating_add(item.input_tokens); entry.2 = entry.2.saturating_add(item.input_tokens);
@@ -374,16 +389,19 @@ pub fn admin_usage_aggregation_by_model_json(
.saturating_add(admin_usage_effective_input_tokens(item)); .saturating_add(admin_usage_effective_input_tokens(item));
entry.5 = entry entry.5 = entry
.5 .5
.saturating_add(admin_usage_cache_creation_tokens(item)); .saturating_add(admin_usage_total_input_context(item));
entry.6 = entry entry.6 = entry
.6 .6
.saturating_add(item.cache_creation_ephemeral_5m_input_tokens); .saturating_add(admin_usage_cache_creation_tokens(item));
entry.7 = entry entry.7 = entry
.7 .7
.saturating_add(item.cache_creation_ephemeral_5m_input_tokens);
entry.8 = entry
.8
.saturating_add(item.cache_creation_ephemeral_1h_input_tokens); .saturating_add(item.cache_creation_ephemeral_1h_input_tokens);
entry.8 = entry.8.saturating_add(item.cache_read_input_tokens); entry.9 = entry.9.saturating_add(item.cache_read_input_tokens);
entry.9 += item.total_cost_usd; entry.10 += item.total_cost_usd;
entry.10 += item.actual_total_cost_usd; entry.11 += item.actual_total_cost_usd;
} }
let mut items: Vec<Value> = grouped let mut items: Vec<Value> = grouped
@@ -394,9 +412,10 @@ pub fn admin_usage_aggregation_by_model_json(
( (
request_count, request_count,
total_tokens, total_tokens,
input_tokens, _input_tokens,
output_tokens, output_tokens,
effective_input_tokens, effective_input_tokens,
total_input_context,
cache_creation_tokens, cache_creation_tokens,
cache_creation_ephemeral_5m_tokens, cache_creation_ephemeral_5m_tokens,
cache_creation_ephemeral_1h_tokens, cache_creation_ephemeral_1h_tokens,
@@ -410,9 +429,7 @@ pub fn admin_usage_aggregation_by_model_json(
"request_count": request_count, "request_count": request_count,
"total_tokens": total_tokens, "total_tokens": total_tokens,
"effective_input_tokens": effective_input_tokens, "effective_input_tokens": effective_input_tokens,
"total_input_context": input_tokens "total_input_context": total_input_context,
.saturating_add(cache_creation_tokens)
.saturating_add(cache_read_tokens),
"output_tokens": output_tokens, "output_tokens": output_tokens,
"total_cost": round_to(total_cost, 6), "total_cost": round_to(total_cost, 6),
"actual_cost": round_to(actual_cost, 6), "actual_cost": round_to(actual_cost, 6),
@@ -421,9 +438,7 @@ pub fn admin_usage_aggregation_by_model_json(
"cache_creation_ephemeral_1h_tokens": cache_creation_ephemeral_1h_tokens, "cache_creation_ephemeral_1h_tokens": cache_creation_ephemeral_1h_tokens,
"cache_read_tokens": cache_read_tokens, "cache_read_tokens": cache_read_tokens,
"cache_hit_rate": admin_usage_token_cache_hit_rate( "cache_hit_rate": admin_usage_token_cache_hit_rate(
input_tokens total_input_context,
.saturating_add(cache_creation_tokens)
.saturating_add(cache_read_tokens),
cache_read_tokens, cache_read_tokens,
), ),
}) })
@@ -464,6 +479,7 @@ pub fn admin_usage_aggregation_by_provider_json(
u64, u64,
u64, u64,
u64, u64,
u64,
f64, f64,
f64, f64,
u64, u64,
@@ -488,6 +504,7 @@ pub fn admin_usage_aggregation_by_provider_json(
0, 0,
0, 0,
0, 0,
0,
0.0, 0.0,
0.0, 0.0,
0, 0,
@@ -505,21 +522,24 @@ pub fn admin_usage_aggregation_by_provider_json(
.saturating_add(admin_usage_effective_input_tokens(item)); .saturating_add(admin_usage_effective_input_tokens(item));
entry.6 = entry entry.6 = entry
.6 .6
.saturating_add(admin_usage_cache_creation_tokens(item)); .saturating_add(admin_usage_total_input_context(item));
entry.7 = entry entry.7 = entry
.7 .7
.saturating_add(item.cache_creation_ephemeral_5m_input_tokens); .saturating_add(admin_usage_cache_creation_tokens(item));
entry.8 = entry entry.8 = entry
.8 .8
.saturating_add(item.cache_creation_ephemeral_5m_input_tokens);
entry.9 = entry
.9
.saturating_add(item.cache_creation_ephemeral_1h_input_tokens); .saturating_add(item.cache_creation_ephemeral_1h_input_tokens);
entry.9 = entry.9.saturating_add(item.cache_read_input_tokens); entry.10 = entry.10.saturating_add(item.cache_read_input_tokens);
entry.10 += item.total_cost_usd; entry.11 += item.total_cost_usd;
entry.11 += item.actual_total_cost_usd; entry.12 += item.actual_total_cost_usd;
entry.12 = entry
.12
.saturating_add(item.response_time_ms.unwrap_or_default());
entry.13 = entry entry.13 = entry
.13 .13
.saturating_add(item.response_time_ms.unwrap_or_default());
entry.14 = entry
.14
.saturating_add(if admin_usage_is_success(item) { 1 } else { 0 }); .saturating_add(if admin_usage_is_success(item) { 1 } else { 0 });
} }
@@ -532,9 +552,10 @@ pub fn admin_usage_aggregation_by_provider_json(
provider_name, provider_name,
request_count, request_count,
total_tokens, total_tokens,
input_tokens, _input_tokens,
output_tokens, output_tokens,
effective_input_tokens, effective_input_tokens,
total_input_context,
cache_creation_tokens, cache_creation_tokens,
cache_creation_ephemeral_5m_tokens, cache_creation_ephemeral_5m_tokens,
cache_creation_ephemeral_1h_tokens, cache_creation_ephemeral_1h_tokens,
@@ -562,9 +583,7 @@ pub fn admin_usage_aggregation_by_provider_json(
"request_count": request_count, "request_count": request_count,
"total_tokens": total_tokens, "total_tokens": total_tokens,
"effective_input_tokens": effective_input_tokens, "effective_input_tokens": effective_input_tokens,
"total_input_context": input_tokens "total_input_context": total_input_context,
.saturating_add(cache_creation_tokens)
.saturating_add(cache_read_tokens),
"output_tokens": output_tokens, "output_tokens": output_tokens,
"total_cost": round_to(total_cost, 6), "total_cost": round_to(total_cost, 6),
"actual_cost": round_to(actual_cost, 6), "actual_cost": round_to(actual_cost, 6),
@@ -576,9 +595,7 @@ pub fn admin_usage_aggregation_by_provider_json(
"cache_creation_ephemeral_1h_tokens": cache_creation_ephemeral_1h_tokens, "cache_creation_ephemeral_1h_tokens": cache_creation_ephemeral_1h_tokens,
"cache_read_tokens": cache_read_tokens, "cache_read_tokens": cache_read_tokens,
"cache_hit_rate": admin_usage_token_cache_hit_rate( "cache_hit_rate": admin_usage_token_cache_hit_rate(
input_tokens total_input_context,
.saturating_add(cache_creation_tokens)
.saturating_add(cache_read_tokens),
cache_read_tokens, cache_read_tokens,
), ),
}) })
@@ -608,7 +625,21 @@ pub fn admin_usage_aggregation_by_api_format_json(
#[allow(clippy::type_complexity)] #[allow(clippy::type_complexity)]
let mut grouped: BTreeMap< let mut grouped: BTreeMap<
String, String,
(u64, u64, u64, u64, u64, u64, u64, u64, u64, f64, f64, u64), (
u64,
u64,
u64,
u64,
u64,
u64,
u64,
u64,
u64,
u64,
f64,
f64,
u64,
),
> = BTreeMap::new(); > = BTreeMap::new();
for item in usage { for item in usage {
let key = item let key = item
@@ -617,7 +648,7 @@ pub fn admin_usage_aggregation_by_api_format_json(
.unwrap_or_else(|| "unknown".to_string()); .unwrap_or_else(|| "unknown".to_string());
let entry = grouped let entry = grouped
.entry(key) .entry(key)
.or_insert((0, 0, 0, 0, 0, 0, 0, 0, 0, 0.0, 0.0, 0)); .or_insert((0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0.0, 0.0, 0));
entry.0 = entry.0.saturating_add(1); entry.0 = entry.0.saturating_add(1);
entry.1 = entry.1.saturating_add(item.total_tokens); entry.1 = entry.1.saturating_add(item.total_tokens);
entry.2 = entry.2.saturating_add(item.input_tokens); entry.2 = entry.2.saturating_add(item.input_tokens);
@@ -627,18 +658,21 @@ pub fn admin_usage_aggregation_by_api_format_json(
.saturating_add(admin_usage_effective_input_tokens(item)); .saturating_add(admin_usage_effective_input_tokens(item));
entry.5 = entry entry.5 = entry
.5 .5
.saturating_add(admin_usage_cache_creation_tokens(item)); .saturating_add(admin_usage_total_input_context(item));
entry.6 = entry entry.6 = entry
.6 .6
.saturating_add(item.cache_creation_ephemeral_5m_input_tokens); .saturating_add(admin_usage_cache_creation_tokens(item));
entry.7 = entry entry.7 = entry
.7 .7
.saturating_add(item.cache_creation_ephemeral_5m_input_tokens);
entry.8 = entry
.8
.saturating_add(item.cache_creation_ephemeral_1h_input_tokens); .saturating_add(item.cache_creation_ephemeral_1h_input_tokens);
entry.8 = entry.8.saturating_add(item.cache_read_input_tokens); entry.9 = entry.9.saturating_add(item.cache_read_input_tokens);
entry.9 += item.total_cost_usd; entry.10 += item.total_cost_usd;
entry.10 += item.actual_total_cost_usd; entry.11 += item.actual_total_cost_usd;
entry.11 = entry entry.12 = entry
.11 .12
.saturating_add(item.response_time_ms.unwrap_or_default()); .saturating_add(item.response_time_ms.unwrap_or_default());
} }
@@ -650,9 +684,10 @@ pub fn admin_usage_aggregation_by_api_format_json(
( (
request_count, request_count,
total_tokens, total_tokens,
input_tokens, _input_tokens,
output_tokens, output_tokens,
effective_input_tokens, effective_input_tokens,
total_input_context,
cache_creation_tokens, cache_creation_tokens,
cache_creation_ephemeral_5m_tokens, cache_creation_ephemeral_5m_tokens,
cache_creation_ephemeral_1h_tokens, cache_creation_ephemeral_1h_tokens,
@@ -672,9 +707,7 @@ pub fn admin_usage_aggregation_by_api_format_json(
"request_count": request_count, "request_count": request_count,
"total_tokens": total_tokens, "total_tokens": total_tokens,
"effective_input_tokens": effective_input_tokens, "effective_input_tokens": effective_input_tokens,
"total_input_context": input_tokens "total_input_context": total_input_context,
.saturating_add(cache_creation_tokens)
.saturating_add(cache_read_tokens),
"output_tokens": output_tokens, "output_tokens": output_tokens,
"total_cost": round_to(total_cost, 6), "total_cost": round_to(total_cost, 6),
"actual_cost": round_to(actual_cost, 6), "actual_cost": round_to(actual_cost, 6),
@@ -684,9 +717,7 @@ pub fn admin_usage_aggregation_by_api_format_json(
"cache_creation_ephemeral_1h_tokens": cache_creation_ephemeral_1h_tokens, "cache_creation_ephemeral_1h_tokens": cache_creation_ephemeral_1h_tokens,
"cache_read_tokens": cache_read_tokens, "cache_read_tokens": cache_read_tokens,
"cache_hit_rate": admin_usage_token_cache_hit_rate( "cache_hit_rate": admin_usage_token_cache_hit_rate(
input_tokens total_input_context,
.saturating_add(cache_creation_tokens)
.saturating_add(cache_read_tokens),
cache_read_tokens, cache_read_tokens,
), ),
}) })

View File

@@ -139,9 +139,9 @@ pub use crate::planner::specialized::{
}; };
pub use crate::planner::standard::{ pub use crate::planner::standard::{
apply_codex_openai_cli_special_body_edits, apply_codex_openai_cli_special_headers, apply_codex_openai_cli_special_body_edits, apply_codex_openai_cli_special_headers,
build_cross_format_openai_chat_request_body, build_cross_format_openai_cli_request_body, apply_openai_compact_special_body_edits, build_cross_format_openai_chat_request_body,
build_local_openai_chat_request_body, build_local_openai_cli_request_body, build_cross_format_openai_cli_request_body, build_local_openai_chat_request_body,
build_standard_request_body, build_standard_upstream_url, build_local_openai_cli_request_body, build_standard_request_body, build_standard_upstream_url,
claude::{ claude::{
resolve_stream_spec as resolve_claude_stream_spec, resolve_stream_spec as resolve_claude_stream_spec,
resolve_sync_spec as resolve_claude_sync_spec, resolve_sync_spec as resolve_claude_sync_spec,

View File

@@ -20,6 +20,12 @@ fn is_codex_openai_cli_request(provider_type: &str, provider_api_format: &str) -
) )
} }
fn is_openai_compact_request(provider_api_format: &str) -> bool {
provider_api_format
.trim()
.eq_ignore_ascii_case("openai:compact")
}
fn build_stable_codex_prompt_cache_key(user_api_key_id: &str) -> Option<String> { fn build_stable_codex_prompt_cache_key(user_api_key_id: &str) -> Option<String> {
let normalized = user_api_key_id.trim(); let normalized = user_api_key_id.trim();
if normalized.is_empty() { if normalized.is_empty() {
@@ -135,6 +141,22 @@ fn maybe_inject_codex_prompt_cache_key(
); );
} }
pub fn apply_openai_compact_special_body_edits(
provider_request_body: &mut Value,
provider_api_format: &str,
) {
if !is_openai_compact_request(provider_api_format) {
return;
}
let Some(body_object) = provider_request_body.as_object_mut() else {
return;
};
// `/v1/responses/compact` does not accept `store`.
body_object.remove("store");
}
pub fn apply_codex_openai_cli_special_body_edits( pub fn apply_codex_openai_cli_special_body_edits(
provider_request_body: &mut Value, provider_request_body: &mut Value,
provider_type: &str, provider_type: &str,
@@ -162,7 +184,9 @@ pub fn apply_codex_openai_cli_special_body_edits(
if !body_rules_handle_path(body_rules, "metadata") { if !body_rules_handle_path(body_rules, "metadata") {
body_object.remove("metadata"); body_object.remove("metadata");
} }
if !body_rules_handle_path(body_rules, "store") { if is_openai_compact_request(provider_api_format) {
body_object.remove("store");
} else if !body_rules_handle_path(body_rules, "store") {
body_object.insert("store".to_string(), json!(false)); body_object.insert("store".to_string(), json!(false));
} }
if !body_rules_handle_path(body_rules, "instructions") if !body_rules_handle_path(body_rules, "instructions")

View File

@@ -14,7 +14,7 @@ use crate::conversion::request::{
}; };
use super::{ use super::{
codex::apply_codex_openai_cli_special_body_edits, apply_openai_compact_special_body_edits, codex::apply_codex_openai_cli_special_body_edits,
normalize::build_local_openai_chat_request_body, normalize::build_local_openai_chat_request_body,
}; };
@@ -52,6 +52,7 @@ pub fn build_standard_request_body(
body_rules, body_rules,
user_api_key_id, user_api_key_id,
); );
apply_openai_compact_special_body_edits(&mut provider_request_body, provider_api_format);
Some(provider_request_body) Some(provider_request_body)
} }

View File

@@ -8,6 +8,7 @@ pub mod openai_cli;
pub use codex::{ pub use codex::{
apply_codex_openai_cli_special_body_edits, apply_codex_openai_cli_special_headers, apply_codex_openai_cli_special_body_edits, apply_codex_openai_cli_special_headers,
apply_openai_compact_special_body_edits,
}; };
pub use family::{LocalStandardSourceFamily, LocalStandardSourceMode, LocalStandardSpec}; pub use family::{LocalStandardSourceFamily, LocalStandardSourceMode, LocalStandardSpec};
pub use matrix::{ pub use matrix::{

View File

@@ -27,4 +27,6 @@ pub use schema::{
BillingSnapshot, BillingSnapshotStatus, CostResult, BILLING_SNAPSHOT_SCHEMA_VERSION, BillingSnapshot, BillingSnapshotStatus, CostResult, BILLING_SNAPSHOT_SCHEMA_VERSION,
}; };
pub use service::BillingService; pub use service::BillingService;
pub use token_normalization::normalize_input_tokens_for_billing; pub use token_normalization::{
normalize_input_tokens_for_billing, normalize_total_input_context_for_cache_hit_rate,
};

View File

@@ -43,9 +43,42 @@ pub fn normalize_input_tokens_for_billing(
} }
} }
pub fn normalize_total_input_context_for_cache_hit_rate(
api_format: Option<&str>,
input_tokens: i64,
cache_creation_tokens: i64,
cache_read_tokens: i64,
) -> i64 {
let normalized_input_tokens = input_tokens.max(0);
let normalized_cache_creation_tokens = cache_creation_tokens.max(0);
let normalized_cache_read_tokens = cache_read_tokens.max(0);
let fresh_input_tokens = match parse_api_family(api_format) {
ApiFamily::Claude => {
normalized_input_tokens.saturating_add(normalized_cache_creation_tokens)
}
ApiFamily::OpenAi | ApiFamily::Gemini => normalize_input_tokens_for_billing(
api_format,
normalized_input_tokens,
normalized_cache_read_tokens,
),
ApiFamily::Unknown => {
if normalized_cache_creation_tokens > 0 {
normalized_input_tokens.saturating_add(normalized_cache_creation_tokens)
} else {
normalized_input_tokens
}
}
};
fresh_input_tokens.saturating_add(normalized_cache_read_tokens)
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::normalize_input_tokens_for_billing; use super::{
normalize_input_tokens_for_billing, normalize_total_input_context_for_cache_hit_rate,
};
#[test] #[test]
fn subtracts_cache_tokens_for_openai_and_gemini() { fn subtracts_cache_tokens_for_openai_and_gemini() {
@@ -66,4 +99,36 @@ mod tests {
100 100
); );
} }
#[test]
fn normalizes_cache_hit_context_for_openai_and_gemini() {
assert_eq!(
normalize_total_input_context_for_cache_hit_rate(Some("openai:chat"), 120, 10, 15),
120
);
assert_eq!(
normalize_total_input_context_for_cache_hit_rate(Some("gemini:chat"), 120, 10, 15),
120
);
}
#[test]
fn includes_cache_creation_for_claude_cache_hit_context() {
assert_eq!(
normalize_total_input_context_for_cache_hit_rate(Some("claude:chat"), 60, 15, 5),
80
);
}
#[test]
fn falls_back_to_creation_aware_context_for_unknown_formats() {
assert_eq!(
normalize_total_input_context_for_cache_hit_rate(None, 20, 10, 5),
35
);
assert_eq!(
normalize_total_input_context_for_cache_hit_rate(None, 20, 0, 5),
25
);
}
} }

View File

@@ -11,8 +11,11 @@ aether-data-contracts.workspace = true
aether-provider-transport.workspace = true aether-provider-transport.workspace = true
aether-scheduler-core.workspace = true aether-scheduler-core.workspace = true
async-trait.workspace = true async-trait.workspace = true
base64.workspace = true
regex.workspace = true regex.workspace = true
rsa = "0.9.10"
serde_json.workspace = true serde_json.workspace = true
sha2 = { workspace = true, features = ["oid"] }
uuid.workspace = true uuid.workspace = true
[dev-dependencies] [dev-dependencies]

View File

@@ -1,6 +1,7 @@
mod association_sync; mod association_sync;
mod config; mod config;
mod logic; mod logic;
mod strategy;
mod transport; mod transport;
pub use association_sync::{ pub use association_sync::{
@@ -12,6 +13,13 @@ pub use config::{
pub use logic::{ pub use logic::{
aggregate_models_for_cache, apply_model_filters, build_models_fetch_url, aggregate_models_for_cache, apply_model_filters, build_models_fetch_url,
endpoint_supports_rust_models_fetch, extract_error_message, json_string_list, endpoint_supports_rust_models_fetch, extract_error_message, json_string_list,
parse_models_response, select_models_fetch_endpoint, ModelFetchRunSummary, ModelsFetchSuccess, merge_upstream_metadata, parse_models_response, parse_models_response_page,
preset_models_for_provider, provider_type_uses_preset_models, select_models_fetch_endpoint,
selected_models_fetch_endpoints, ModelFetchRunSummary, ModelsFetchPage, ModelsFetchSuccess,
};
pub use strategy::{fetch_models_from_transports, ModelsFetchOutcome};
pub use transport::{
build_antigravity_fetch_available_models_plan, build_gemini_cli_load_code_assist_plan,
build_models_fetch_execution_plan, build_standard_models_fetch_execution_plan,
build_vertex_models_fetch_execution_plan, ModelFetchTransportRuntime,
}; };
pub use transport::{build_models_fetch_execution_plan, ModelFetchTransportRuntime};

View File

@@ -3,9 +3,14 @@ use std::collections::{BTreeMap, BTreeSet};
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
}; };
use aether_provider_transport::provider_types::provider_type_supports_model_fetch;
use regex::Regex; use regex::Regex;
use serde_json::Value; use serde_json::{json, Value};
const MODEL_FETCH_FORMAT_PRIORITY: &[&[&str]] = &[
&["openai:chat", "openai:cli", "openai:compact"],
&["claude:chat", "claude:cli"],
&["gemini:chat", "gemini:cli"],
];
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ModelFetchRunSummary { pub struct ModelFetchRunSummary {
@@ -21,6 +26,14 @@ pub struct ModelsFetchSuccess {
pub cached_models: Vec<Value>, pub cached_models: Vec<Value>,
} }
#[derive(Debug, Clone, PartialEq)]
pub struct ModelsFetchPage {
pub fetched_model_ids: Vec<String>,
pub cached_models: Vec<Value>,
pub has_more: bool,
pub next_after_id: Option<String>,
}
pub fn extract_error_message(value: &Value) -> Option<String> { pub fn extract_error_message(value: &Value) -> Option<String> {
value value
.get("error") .get("error")
@@ -41,21 +54,20 @@ pub fn extract_error_message(value: &Value) -> Option<String> {
} }
pub fn build_models_fetch_url( pub fn build_models_fetch_url(
provider_type: &str, _provider_type: &str,
endpoint_api_format: &str, endpoint_api_format: &str,
base_url: &str, base_url: &str,
) -> Option<(String, String)> { ) -> Option<(String, String)> {
let api_format = normalize_api_format(endpoint_api_format); let api_format = normalize_api_format(endpoint_api_format);
if !provider_type_supports_model_fetch(provider_type) { if !endpoint_supports_rust_models_fetch(&api_format) {
return None; return None;
} }
let url = if api_format.starts_with("openai:") || api_format.starts_with("claude:") { let url = if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
build_v1_models_url(base_url) build_v1_models_url(base_url)
} else if api_format.starts_with("gemini:") { } else if api_format.starts_with("gemini:") {
build_gemini_models_url(base_url) build_gemini_models_url(base_url)
} else { } else {
return None; None
}?; }?;
Some((url, api_format)) Some((url, api_format))
} }
@@ -64,13 +76,38 @@ pub fn parse_models_response(
endpoint_api_format: &str, endpoint_api_format: &str,
body: &Value, body: &Value,
) -> Result<ModelsFetchSuccess, String> { ) -> Result<ModelsFetchSuccess, String> {
let parsed = parse_models_response_page(endpoint_api_format, body)?;
Ok(ModelsFetchSuccess {
fetched_model_ids: parsed.fetched_model_ids,
cached_models: parsed.cached_models,
})
}
pub fn parse_models_response_page(
endpoint_api_format: &str,
body: &Value,
) -> Result<ModelsFetchPage, String> {
let api_format = normalize_api_format(endpoint_api_format); let api_format = normalize_api_format(endpoint_api_format);
let mut cached_models = Vec::new(); let mut cached_models = Vec::new();
let mut fetched_model_ids = Vec::new(); let mut fetched_model_ids = Vec::new();
let mut seen = BTreeSet::new(); let mut seen = BTreeSet::new();
let mut has_more = false;
let mut next_after_id = None;
if api_format.starts_with("openai:") || api_format.starts_with("claude:") { if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
let items = if let Some(items) = body.get("data").and_then(Value::as_array) { let items = if let Some(items) = body.get("data").and_then(Value::as_array) {
has_more = body
.get("has_more")
.and_then(Value::as_bool)
.unwrap_or(false);
if api_format.starts_with("claude:") && has_more {
next_after_id = body
.get("last_id")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
}
items items
} else if let Some(items) = body.as_array() { } else if let Some(items) = body.as_array() {
items items
@@ -117,29 +154,54 @@ pub fn parse_models_response(
return Err("models response parser does not support this provider format".to_string()); return Err("models response parser does not support this provider format".to_string());
} }
Ok(ModelsFetchSuccess { Ok(ModelsFetchPage {
fetched_model_ids, fetched_model_ids,
cached_models, cached_models,
has_more,
next_after_id,
}) })
} }
pub fn selected_models_fetch_endpoints(
endpoints: &[StoredProviderCatalogEndpoint],
key: &StoredProviderCatalogKey,
) -> Vec<StoredProviderCatalogEndpoint> {
let key_formats = json_string_list(key.api_formats.as_ref())
.into_iter()
.map(|value| normalize_api_format(&value))
.collect::<BTreeSet<_>>();
let mut by_format = BTreeMap::<String, StoredProviderCatalogEndpoint>::new();
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
let api_format = normalize_api_format(&endpoint.api_format);
if api_format.is_empty() || !endpoint_supports_rust_models_fetch(&api_format) {
continue;
}
if !key_formats.is_empty() && !key_formats.contains(&api_format) {
continue;
}
by_format
.entry(api_format)
.or_insert_with(|| endpoint.clone());
}
MODEL_FETCH_FORMAT_PRIORITY
.iter()
.filter_map(|candidates| {
candidates
.iter()
.find_map(|api_format| by_format.remove(*api_format))
})
.collect()
}
pub fn select_models_fetch_endpoint( pub fn select_models_fetch_endpoint(
endpoints: &[StoredProviderCatalogEndpoint], endpoints: &[StoredProviderCatalogEndpoint],
key: &StoredProviderCatalogKey, key: &StoredProviderCatalogKey,
) -> Option<StoredProviderCatalogEndpoint> { ) -> Option<StoredProviderCatalogEndpoint> {
let key_formats = json_string_list(key.api_formats.as_ref()) selected_models_fetch_endpoints(endpoints, key)
.into_iter() .into_iter()
.map(|value| normalize_api_format(&value)) .next()
.collect::<BTreeSet<_>>();
endpoints
.iter()
.filter(|endpoint| endpoint.is_active)
.find(|endpoint| {
let api_format = normalize_api_format(&endpoint.api_format);
(key_formats.is_empty() || key_formats.contains(&api_format))
&& endpoint_supports_rust_models_fetch(&endpoint.api_format)
})
.cloned()
} }
pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool { pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
@@ -148,7 +210,6 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
api_format.as_str(), api_format.as_str(),
"openai:chat" "openai:chat"
| "openai:cli" | "openai:cli"
| "openai:responses"
| "openai:compact" | "openai:compact"
| "claude:chat" | "claude:chat"
| "claude:cli" | "claude:cli"
@@ -157,6 +218,100 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
) )
} }
pub fn provider_type_uses_preset_models(provider_type: &str) -> bool {
matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"codex" | "kiro" | "claude_code" | "gemini_cli"
)
}
#[rustfmt::skip]
pub fn preset_models_for_provider(provider_type: &str) -> Option<Vec<Value>> {
let models = match provider_type.trim().to_ascii_lowercase().as_str() {
"gemini_cli" => vec![
preset_model("gemini-2.5-pro", "google", "Gemini 2.5 Pro", "gemini:cli"),
preset_model("gemini-2.5-flash", "google", "Gemini 2.5 Flash", "gemini:cli"),
preset_model("gemini-3-pro-preview", "google", "Gemini 3 Pro Preview", "gemini:cli"),
preset_model("gemini-3-flash-preview", "google", "Gemini 3 Flash Preview", "gemini:cli"),
preset_model("gemini-3.1-pro-preview", "google", "Gemini 3.1 Pro Preview", "gemini:cli"),
],
"kiro" => vec![
preset_model("claude-sonnet-4.5", "anthropic", "Claude Sonnet 4.5", "claude:cli"),
preset_model("claude-sonnet-4.6", "anthropic", "Claude Sonnet 4.6", "claude:cli"),
preset_model("claude-opus-4.5", "anthropic", "Claude Opus 4.5", "claude:cli"),
preset_model("claude-opus-4.6", "anthropic", "Claude Opus 4.6", "claude:cli"),
preset_model("claude-haiku-4.5", "anthropic", "Claude Haiku 4.5", "claude:cli"),
],
"claude_code" => vec![
preset_model("claude-opus-4-5-20251101", "anthropic", "Claude Opus 4.5", "claude:cli"),
preset_model("claude-opus-4-6", "anthropic", "Claude Opus 4.6", "claude:cli"),
preset_model("claude-sonnet-4-6", "anthropic", "Claude Sonnet 4.6", "claude:cli"),
preset_model("claude-sonnet-4-5-20250929", "anthropic", "Claude Sonnet 4.5", "claude:cli"),
preset_model("claude-haiku-4-5-20251001", "anthropic", "Claude Haiku 4.5", "claude:cli"),
],
"codex" => vec![
preset_model("gpt-5", "openai", "GPT-5", "openai:cli"),
preset_model("gpt-5-codex", "openai", "GPT-5 Codex", "openai:cli"),
preset_model("gpt-5-codex-mini", "openai", "GPT-5 Codex Mini", "openai:cli"),
preset_model("gpt-5.1", "openai", "GPT-5.1", "openai:cli"),
preset_model("gpt-5.1-codex", "openai", "GPT-5.1 Codex", "openai:cli"),
preset_model("gpt-5.1-codex-mini", "openai", "GPT-5.1 Codex Mini", "openai:cli"),
preset_model("gpt-5.1-codex-max", "openai", "GPT-5.1 Codex Max", "openai:cli"),
preset_model("gpt-5.2", "openai", "GPT-5.2", "openai:cli"),
preset_model("gpt-5.2-codex", "openai", "GPT-5.2 Codex", "openai:cli"),
preset_model("gpt-5.3-codex", "openai", "GPT-5.3 Codex", "openai:cli"),
preset_model("gpt-5.4", "openai", "GPT-5.4", "openai:cli"),
],
_ => return None,
};
Some(models)
}
pub fn merge_upstream_metadata(current: Option<&Value>, incoming: &Value) -> Value {
let mut merged = current
.and_then(Value::as_object)
.cloned()
.unwrap_or_default();
let Some(incoming_object) = incoming.as_object() else {
return Value::Object(merged);
};
for (namespace, value) in incoming_object {
let mut next_value = value.clone();
if let (Some(next_namespace), Some(old_namespace)) = (
next_value.as_object_mut(),
merged.get(namespace).and_then(Value::as_object),
) {
if let (Some(new_quota), Some(old_quota)) = (
next_namespace
.get_mut("quota_by_model")
.and_then(Value::as_object_mut),
old_namespace
.get("quota_by_model")
.and_then(Value::as_object),
) {
for (model_id, new_info) in new_quota.iter_mut() {
let Some(new_info_object) = new_info.as_object_mut() else {
continue;
};
let Some(old_info_object) = old_quota.get(model_id).and_then(Value::as_object)
else {
continue;
};
if !new_info_object.contains_key("reset_time") {
if let Some(reset_time) = old_info_object.get("reset_time") {
new_info_object.insert("reset_time".to_string(), reset_time.clone());
}
}
}
}
}
merged.insert(namespace.clone(), next_value);
}
Value::Object(merged)
}
pub fn apply_model_filters( pub fn apply_model_filters(
fetched_model_ids: &[String], fetched_model_ids: &[String],
locked_models: Vec<String>, locked_models: Vec<String>,
@@ -211,7 +366,6 @@ pub fn json_string_list(value: Option<&Value>) -> Vec<String> {
pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> { pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
let mut aggregated = BTreeMap::<String, serde_json::Map<String, Value>>::new(); let mut aggregated = BTreeMap::<String, serde_json::Map<String, Value>>::new();
let mut order = Vec::<String>::new();
for model in models { for model in models {
let Some(object) = model.as_object() else { let Some(object) = model.as_object() else {
@@ -227,7 +381,6 @@ pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
}; };
let entry = aggregated.entry(model_id.to_string()).or_insert_with(|| { let entry = aggregated.entry(model_id.to_string()).or_insert_with(|| {
order.push(model_id.to_string());
let mut cloned = object.clone(); let mut cloned = object.clone();
cloned.remove("api_format"); cloned.remove("api_format");
cloned cloned
@@ -286,11 +439,7 @@ pub fn aggregate_models_for_cache(models: &[Value]) -> Vec<Value> {
} }
} }
order aggregated.into_values().map(Value::Object).collect()
.into_iter()
.filter_map(|model_id| aggregated.remove(&model_id))
.map(Value::Object)
.collect()
} }
fn build_v1_models_url(base_url: &str) -> Option<String> { fn build_v1_models_url(base_url: &str) -> Option<String> {
@@ -347,10 +496,37 @@ fn normalize_cached_model(item: &Value, model_id: &str, api_format: &str) -> Val
"api_formats".to_string(), "api_formats".to_string(),
Value::Array(vec![Value::String(api_format.to_string())]), Value::Array(vec![Value::String(api_format.to_string())]),
); );
if api_format.starts_with("gemini:") {
object
.entry("owned_by".to_string())
.or_insert_with(|| Value::String("google".to_string()));
if !object.contains_key("display_name") {
let display_name = item
.get("displayName")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(model_id);
object.insert(
"display_name".to_string(),
Value::String(display_name.to_string()),
);
}
}
object.remove("api_format"); object.remove("api_format");
Value::Object(object) Value::Object(object)
} }
fn preset_model(model_id: &str, owned_by: &str, display_name: &str, api_format: &str) -> Value {
json!({
"id": model_id,
"object": "model",
"owned_by": owned_by,
"display_name": display_name,
"api_formats": [api_format],
})
}
fn wildcard_matches(pattern: &str, model_id: &str) -> bool { fn wildcard_matches(pattern: &str, model_id: &str) -> bool {
let mut regex = String::from("^"); let mut regex = String::from("^");
for ch in pattern.chars() { for ch in pattern.chars() {
@@ -379,7 +555,8 @@ mod tests {
use super::{ use super::{
aggregate_models_for_cache, apply_model_filters, build_gemini_models_url, aggregate_models_for_cache, apply_model_filters, build_gemini_models_url,
build_models_fetch_url, parse_models_response, select_models_fetch_endpoint, build_models_fetch_url, merge_upstream_metadata, parse_models_response,
parse_models_response_page, preset_models_for_provider, selected_models_fetch_endpoints,
}; };
fn sample_endpoint( fn sample_endpoint(
@@ -410,7 +587,11 @@ mod tests {
.expect("endpoint transport should build") .expect("endpoint transport should build")
} }
fn sample_key(provider_id: &str, key_id: &str) -> StoredProviderCatalogKey { fn sample_key(
provider_id: &str,
key_id: &str,
api_formats: &[&str],
) -> StoredProviderCatalogKey {
StoredProviderCatalogKey::new( StoredProviderCatalogKey::new(
key_id.to_string(), key_id.to_string(),
provider_id.to_string(), provider_id.to_string(),
@@ -421,7 +602,7 @@ mod tests {
) )
.expect("key should build") .expect("key should build")
.with_transport_fields( .with_transport_fields(
Some(json!(["openai:chat"])), Some(json!(api_formats)),
"encrypted".to_string(), "encrypted".to_string(),
None, None,
None, None,
@@ -453,12 +634,15 @@ mod tests {
} }
#[test] #[test]
fn aggregate_models_for_cache_merges_api_formats_by_model_id() { fn aggregate_models_for_cache_merges_api_formats_and_sorts_by_model_id() {
let aggregated = aggregate_models_for_cache(&[ let aggregated = aggregate_models_for_cache(&[
json!({"id":"gpt-5","api_formats":["openai:chat"]}), json!({"id":"zeta","api_formats":["openai:chat"]}),
json!({"id":"gpt-5","api_formats":["openai:cli"]}), json!({"id":"alpha","api_formats":["openai:cli"]}),
json!({"id":"alpha","api_formats":["openai:chat"]}),
]); ]);
assert_eq!(aggregated.len(), 1); assert_eq!(aggregated.len(), 2);
assert_eq!(aggregated[0]["id"], "alpha");
assert_eq!(aggregated[1]["id"], "zeta");
assert_eq!( assert_eq!(
aggregated[0]["api_formats"], aggregated[0]["api_formats"],
json!(["openai:chat", "openai:cli"]) json!(["openai:chat", "openai:cli"])
@@ -488,9 +672,9 @@ mod tests {
} }
#[test] #[test]
fn build_models_fetch_url_rejects_provider_types_without_fetch_support() { fn build_models_fetch_url_excludes_openai_responses() {
assert_eq!( assert_eq!(
build_models_fetch_url("vertex_ai", "gemini:chat", "https://example.com"), build_models_fetch_url("openai", "openai:responses", "https://example.com"),
None None
); );
} }
@@ -510,9 +694,30 @@ mod tests {
} }
#[test] #[test]
fn select_models_fetch_endpoint_respects_key_api_formats() { fn parse_models_response_page_reads_claude_pagination_state() {
let key = sample_key("provider-1", "key-1"); let parsed = parse_models_response_page(
"claude:chat",
&json!({
"data": [{"id": "claude-sonnet-4"}],
"has_more": true,
"last_id": "cursor-2"
}),
)
.expect("response should parse");
assert!(parsed.has_more);
assert_eq!(parsed.next_after_id.as_deref(), Some("cursor-2"));
}
#[test]
fn selected_models_fetch_endpoints_prefers_chat_and_excludes_responses() {
let key = sample_key("provider-1", "key-1", &["openai:chat", "openai:responses"]);
let endpoints = vec![ let endpoints = vec![
sample_endpoint(
"provider-1",
"endpoint-responses",
"openai:responses",
"https://example.com",
),
sample_endpoint( sample_endpoint(
"provider-1", "provider-1",
"endpoint-cli", "endpoint-cli",
@@ -526,8 +731,50 @@ mod tests {
"https://example.com", "https://example.com",
), ),
]; ];
let selected = let selected = selected_models_fetch_endpoints(&endpoints, &key);
select_models_fetch_endpoint(&endpoints, &key).expect("endpoint should be selected"); assert_eq!(selected.len(), 1);
assert_eq!(selected.id, "endpoint-chat"); assert_eq!(selected[0].id, "endpoint-chat");
}
#[test]
fn merge_upstream_metadata_keeps_existing_reset_time_for_returned_models() {
let merged = merge_upstream_metadata(
Some(&json!({
"antigravity": {
"quota_by_model": {
"gemini-2.5-pro": {
"remaining_fraction": 0.3,
"reset_time": "2026-04-12T00:00:00Z"
},
"stale-model": {
"remaining_fraction": 0.1,
"reset_time": "old"
}
}
}
})),
&json!({
"antigravity": {
"quota_by_model": {
"gemini-2.5-pro": {
"remaining_fraction": 0.6
}
}
}
}),
);
assert_eq!(
merged["antigravity"]["quota_by_model"]["gemini-2.5-pro"]["reset_time"],
"2026-04-12T00:00:00Z"
);
assert!(merged["antigravity"]["quota_by_model"]
.get("stale-model")
.is_none());
}
#[test]
fn preset_models_cover_codex_catalog() {
let models = preset_models_for_provider("codex").expect("preset models should exist");
assert!(models.iter().any(|model| model["id"] == "gpt-5.4"));
} }
} }

File diff suppressed because it is too large Load Diff

View File

@@ -1,20 +1,47 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use aether_contracts::{ExecutionPlan, ProxySnapshot, RequestBody}; use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot, RequestBody};
use aether_provider_transport::auth::{ use aether_provider_transport::antigravity::{
resolve_local_gemini_auth, resolve_local_openai_chat_auth, resolve_local_standard_auth, build_antigravity_static_identity_headers, resolve_local_antigravity_request_auth,
AntigravityRequestAuthSupport, ANTIGRAVITY_REQUEST_USER_AGENT,
};
use aether_provider_transport::auth::{
ensure_upstream_auth_header, resolve_local_gemini_auth, resolve_local_openai_chat_auth,
resolve_local_standard_auth,
}; };
use aether_provider_transport::url::build_passthrough_path_url;
use aether_provider_transport::vertex::resolve_local_vertex_api_key_query_auth; use aether_provider_transport::vertex::resolve_local_vertex_api_key_query_auth;
use aether_provider_transport::{ use aether_provider_transport::{
apply_local_header_rules, ensure_upstream_auth_header, resolve_transport_execution_timeouts, apply_local_header_rules, resolve_transport_execution_timeouts, resolve_transport_tls_profile,
resolve_transport_tls_profile, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
}; };
use async_trait::async_trait; use async_trait::async_trait;
use serde_json::json; use serde_json::json;
use crate::build_models_fetch_url; use crate::build_models_fetch_url;
const OPENAI_CLI_USER_AGENT: &str = "openai-codex/1.0";
const CLAUDE_CLI_USER_AGENT: &str = "claude-code/1.0.1";
const GEMINI_CLI_USER_AGENT: &str = "GeminiCLI/0.1.5 (Windows; AMD64)";
const CLAUDE_VERSION_HEADER: &str = "2023-06-01";
const ANTIGRAVITY_FETCH_PROVIDER_API_FORMAT: &str = "antigravity:fetch_available_models";
const GEMINI_CLI_LOAD_CODE_ASSIST_PROVIDER_API_FORMAT: &str = "gemini_cli:load_code_assist";
const BROWSER_FINGERPRINT_HEADERS: &[(&str, &str)] = &[
(
"user-agent",
"Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/140.0.7339.249 Electron/38.7.0 Safari/537.36",
),
("accept", "application/json"),
("accept-encoding", "gzip, deflate, br"),
("accept-language", "zh-CN"),
("sec-ch-ua", "\"Not=A?Brand\";v=\"24\", \"Chromium\";v=\"140\""),
("sec-ch-ua-mobile", "?0"),
("sec-ch-ua-platform", "\"macOS\""),
("sec-fetch-site", "cross-site"),
("sec-fetch-mode", "cors"),
("sec-fetch-dest", "empty"),
];
#[async_trait] #[async_trait]
pub trait ModelFetchTransportRuntime: Send + Sync { pub trait ModelFetchTransportRuntime: Send + Sync {
async fn resolve_local_oauth_request_auth( async fn resolve_local_oauth_request_auth(
@@ -26,113 +53,438 @@ pub trait ModelFetchTransportRuntime: Send + Sync {
&self, &self,
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot>; ) -> Option<ProxySnapshot>;
async fn execute_model_fetch_execution_plan(
&self,
plan: &ExecutionPlan,
) -> Result<ExecutionResult, String>;
} }
pub async fn build_models_fetch_execution_plan( pub async fn build_models_fetch_execution_plan(
runtime: &(impl ModelFetchTransportRuntime + ?Sized), runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,
) -> Result<ExecutionPlan, String> { ) -> Result<ExecutionPlan, String> {
let (upstream_url, provider_api_format) = build_models_fetch_url( build_standard_models_fetch_execution_plan(runtime, transport, None).await
&transport.provider.provider_type, }
&transport.endpoint.api_format,
&transport.endpoint.base_url,
)
.ok_or_else(|| "Rust models fetch does not support this provider format yet".to_string())?;
let (auth_header_name, auth_header_value) = resolve_models_fetch_auth(runtime, transport)
.await?
.ok_or_else(|| {
"Rust models fetch auth resolution is not supported for this key".to_string()
})?;
let mut headers = BTreeMap::from([(auth_header_name.clone(), auth_header_value.clone())]); struct ModelFetchExecutionPlanRequest {
if !apply_local_header_rules( method: String,
&mut headers, url: String,
transport.endpoint.header_rules.as_ref(), headers: BTreeMap<String, String>,
&[auth_header_name.as_str()], content_type: Option<String>,
&json!({}), body: RequestBody,
None, client_api_format: String,
) { provider_api_format: String,
return Err("Endpoint header_rules application failed".to_string()); model_name: Option<String>,
}
pub async fn build_standard_models_fetch_execution_plan(
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot,
after_id: Option<&str>,
) -> Result<ExecutionPlan, String> {
let api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
let provider_api_format = api_format.clone();
let mut headers = standard_models_fetch_headers(&api_format);
let mut protected_headers = Vec::<String>::new();
if api_format.starts_with("openai:") || api_format.starts_with("claude:") {
let (auth_header_name, auth_header_value) =
resolve_standard_header_auth(runtime, transport)
.await?
.ok_or_else(|| {
"Rust models fetch auth resolution is not supported for this key".to_string()
})?;
protected_headers.push(auth_header_name.clone());
headers.insert(auth_header_name.clone(), auth_header_value.clone());
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
ensure_upstream_auth_header(&mut headers, &auth_header_name, &auth_header_value);
} else {
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
} }
ensure_upstream_auth_header(&mut headers, &auth_header_name, &auth_header_value);
let upstream_url = build_standard_models_fetch_url(transport, after_id)?;
build_execution_plan(
runtime,
transport,
ModelFetchExecutionPlanRequest {
method: "GET".to_string(),
url: upstream_url,
headers,
content_type: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
client_api_format: provider_api_format.clone(),
provider_api_format,
model_name: Some("models".to_string()),
},
)
.await
}
pub async fn build_antigravity_fetch_available_models_plan(
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot,
base_url: &str,
project_id: &str,
) -> Result<ExecutionPlan, String> {
let authorization = resolve_oauth_header_auth(runtime, transport)
.await?
.ok_or_else(|| "Antigravity fetch requires OAuth authorization header".to_string())?;
let identity_auth = match resolve_local_antigravity_request_auth(transport) {
AntigravityRequestAuthSupport::Supported(auth) => auth,
AntigravityRequestAuthSupport::Unsupported(reason) => {
return Err(format!(
"Antigravity fetch auth resolution is not supported: {reason:?}"
))
}
};
let mut headers = build_antigravity_static_identity_headers(&identity_auth);
headers.insert(authorization.0.clone(), authorization.1.clone());
headers.insert("content-type".to_string(), "application/json".to_string());
headers.insert("accept".to_string(), "application/json".to_string());
headers
.entry("user-agent".to_string())
.or_insert_with(|| ANTIGRAVITY_REQUEST_USER_AGENT.to_string());
let protected_headers = vec![authorization.0];
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
let url = format!(
"{}{}",
base_url.trim_end_matches('/'),
"/v1internal:fetchAvailableModels"
);
build_execution_plan(
runtime,
transport,
ModelFetchExecutionPlanRequest {
method: "POST".to_string(),
url,
headers,
content_type: Some("application/json".to_string()),
body: RequestBody::from_json(json!({ "project": project_id })),
client_api_format: "gemini:chat".to_string(),
provider_api_format: ANTIGRAVITY_FETCH_PROVIDER_API_FORMAT.to_string(),
model_name: Some("fetchAvailableModels".to_string()),
},
)
.await
}
pub async fn build_gemini_cli_load_code_assist_plan(
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot,
) -> Result<ExecutionPlan, String> {
let authorization = resolve_bearer_or_oauth_header_auth(runtime, transport)
.await?
.ok_or_else(|| "GeminiCLI loadCodeAssist requires bearer or OAuth auth".to_string())?;
let mut headers = BTreeMap::from([
("user-agent".to_string(), GEMINI_CLI_USER_AGENT.to_string()),
("accept-encoding".to_string(), "identity".to_string()),
("content-type".to_string(), "application/json".to_string()),
]);
headers.insert(authorization.0.clone(), authorization.1.clone());
let protected_headers = vec![authorization.0];
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
build_execution_plan(
runtime,
transport,
ModelFetchExecutionPlanRequest {
method: "POST".to_string(),
url: "https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist".to_string(),
headers,
content_type: Some("application/json".to_string()),
body: RequestBody::from_json(json!({
"metadata": {
"ideType": "ANTIGRAVITY",
"platform": "PLATFORM_UNSPECIFIED",
"pluginType": "GEMINI",
}
})),
client_api_format: "gemini:cli".to_string(),
provider_api_format: GEMINI_CLI_LOAD_CODE_ASSIST_PROVIDER_API_FORMAT.to_string(),
model_name: Some("loadCodeAssist".to_string()),
},
)
.await
}
pub async fn build_vertex_models_fetch_execution_plan(
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot,
url: &str,
api_format: &str,
auth_header: Option<(String, String)>,
) -> Result<ExecutionPlan, String> {
let mut headers = standard_models_fetch_headers(api_format);
let mut protected_headers = Vec::<String>::new();
if let Some((name, value)) = auth_header {
protected_headers.push(name.clone());
headers.insert(name.clone(), value.clone());
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
ensure_upstream_auth_header(&mut headers, &name, &value);
} else {
headers = apply_fetch_header_rules(transport, headers, &protected_headers)?;
}
build_execution_plan(
runtime,
transport,
ModelFetchExecutionPlanRequest {
method: "GET".to_string(),
url: url.trim().to_string(),
headers,
content_type: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
client_api_format: api_format.to_string(),
provider_api_format: api_format.to_string(),
model_name: Some("models".to_string()),
},
)
.await
}
async fn build_execution_plan(
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot,
request: ModelFetchExecutionPlanRequest,
) -> Result<ExecutionPlan, String> {
let ModelFetchExecutionPlanRequest {
method,
url,
headers,
content_type,
body,
client_api_format,
provider_api_format,
model_name,
} = request;
Ok(ExecutionPlan { Ok(ExecutionPlan {
request_id: format!("req-model-fetch-{}", transport.key.id), request_id: format!(
"req-model-fetch-{}-{}",
transport.key.id,
provider_api_format.replace(':', "-")
),
candidate_id: None, candidate_id: None,
provider_name: Some(transport.provider.name.clone()), provider_name: Some(transport.provider.name.clone()),
provider_id: transport.provider.id.clone(), provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(), endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(), key_id: transport.key.id.clone(),
method: "GET".to_string(), method,
url: upstream_url, url,
headers, headers,
content_type: None, content_type,
content_encoding: None, content_encoding: None,
body: RequestBody { body,
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: false, stream: false,
client_api_format: provider_api_format.clone(), client_api_format,
provider_api_format, provider_api_format,
model_name: None, model_name,
proxy: runtime.resolve_model_fetch_proxy(transport).await, proxy: runtime.resolve_model_fetch_proxy(transport).await,
tls_profile: resolve_transport_tls_profile(transport), tls_profile: resolve_transport_tls_profile(transport),
timeouts: resolve_transport_execution_timeouts(transport), timeouts: resolve_transport_execution_timeouts(transport),
}) })
} }
async fn resolve_models_fetch_auth( async fn resolve_standard_header_auth(
runtime: &(impl ModelFetchTransportRuntime + ?Sized), runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,
) -> Result<Option<(String, String)>, String> { ) -> Result<Option<(String, String)>, String> {
if transport.key.auth_type.trim().eq_ignore_ascii_case("oauth") if transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|| transport.key.auth_type.trim().eq_ignore_ascii_case("kiro") || transport.key.auth_type.trim().eq_ignore_ascii_case("kiro")
{ {
return match runtime.resolve_local_oauth_request_auth(transport).await { return resolve_oauth_header_auth(runtime, transport).await;
Ok(Some(LocalResolvedOAuthRequestAuth::Header { name, value })) => {
Ok(Some((name, value)))
}
Ok(Some(LocalResolvedOAuthRequestAuth::Kiro(_))) => Ok(None),
Ok(None) => Ok(None),
Err(err) => Err(err),
};
} }
if let Some(auth) = resolve_local_openai_chat_auth(transport) { let api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
return Ok(Some(auth)); if api_format.starts_with("openai:") {
return Ok(resolve_local_openai_chat_auth(transport));
} }
if let Some(auth) = resolve_local_standard_auth(transport) { if api_format.starts_with("claude:") {
return Ok(Some(auth)); return Ok(resolve_local_standard_auth(transport));
}
if let Some(auth) = resolve_local_gemini_auth(transport) {
return Ok(Some(auth));
}
if let Some(query_auth) = resolve_local_vertex_api_key_query_auth(transport) {
let url = build_passthrough_path_url(
&transport.endpoint.base_url,
"/v1/publishers/google/models",
Some(&format!("{}={}", query_auth.name, query_auth.value)),
&[],
);
if url.is_some() {
return Ok(None);
}
} }
Ok(None) Ok(None)
} }
async fn resolve_oauth_header_auth(
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot,
) -> Result<Option<(String, String)>, String> {
match runtime.resolve_local_oauth_request_auth(transport).await {
Ok(Some(LocalResolvedOAuthRequestAuth::Header { name, value })) => Ok(Some((name, value))),
Ok(Some(LocalResolvedOAuthRequestAuth::Kiro(_))) => Ok(None),
Ok(None) => Ok(None),
Err(err) => Err(err),
}
}
async fn resolve_bearer_or_oauth_header_auth(
runtime: &(impl ModelFetchTransportRuntime + ?Sized),
transport: &GatewayProviderTransportSnapshot,
) -> Result<Option<(String, String)>, String> {
if let Some(auth) = resolve_oauth_header_auth(runtime, transport).await? {
return Ok(Some(auth));
}
if let Some((name, value)) = resolve_local_openai_chat_auth(transport) {
return Ok(Some((name, value)));
}
if transport
.key
.auth_type
.trim()
.eq_ignore_ascii_case("bearer")
{
let secret = transport.key.decrypted_api_key.trim();
if !secret.is_empty() {
return Ok(Some((
"authorization".to_string(),
format!("Bearer {secret}"),
)));
}
}
Ok(None)
}
fn apply_fetch_header_rules(
transport: &GatewayProviderTransportSnapshot,
mut headers: BTreeMap<String, String>,
protected_headers: &[String],
) -> Result<BTreeMap<String, String>, String> {
let protected = protected_headers
.iter()
.map(String::as_str)
.collect::<Vec<_>>();
if !apply_local_header_rules(
&mut headers,
transport.endpoint.header_rules.as_ref(),
&protected,
&json!({}),
None,
) {
return Err("Endpoint header_rules application failed".to_string());
}
Ok(headers)
}
fn standard_models_fetch_headers(api_format: &str) -> BTreeMap<String, String> {
let api_format = api_format.trim().to_ascii_lowercase();
match api_format.as_str() {
"openai:cli" | "openai:compact" => {
BTreeMap::from([("user-agent".to_string(), OPENAI_CLI_USER_AGENT.to_string())])
}
"claude:chat" => BTreeMap::from([(
"anthropic-version".to_string(),
CLAUDE_VERSION_HEADER.to_string(),
)]),
"claude:cli" => BTreeMap::from([
("user-agent".to_string(), CLAUDE_CLI_USER_AGENT.to_string()),
(
"anthropic-version".to_string(),
CLAUDE_VERSION_HEADER.to_string(),
),
]),
"gemini:chat" => BROWSER_FINGERPRINT_HEADERS
.iter()
.map(|(key, value)| (key.to_string(), value.to_string()))
.collect(),
"gemini:cli" => {
let mut headers = BROWSER_FINGERPRINT_HEADERS
.iter()
.map(|(key, value)| (key.to_string(), value.to_string()))
.collect::<BTreeMap<_, _>>();
headers.insert("user-agent".to_string(), GEMINI_CLI_USER_AGENT.to_string());
headers
}
_ => BTreeMap::new(),
}
}
fn build_standard_models_fetch_url(
transport: &GatewayProviderTransportSnapshot,
after_id: Option<&str>,
) -> Result<String, String> {
let api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
if api_format.starts_with("gemini:") {
let secret = resolve_local_vertex_api_key_query_auth(transport)
.map(|auth| auth.value)
.or_else(|| {
resolve_local_gemini_auth(transport).and_then(|(name, value)| {
name.eq_ignore_ascii_case("x-goog-api-key").then_some(value)
})
})
.or_else(|| {
let secret = transport.key.decrypted_api_key.trim();
(!secret.is_empty()).then_some(secret.to_string())
})
.ok_or_else(|| "Gemini models fetch requires an API key".to_string())?;
let (url, _) = build_models_fetch_url(
&transport.provider.provider_type,
&transport.endpoint.api_format,
&transport.endpoint.base_url,
)
.ok_or_else(|| "Rust models fetch does not support this provider format yet".to_string())?;
return Ok(append_query_param(url, "key", &secret));
}
let (mut url, _) = build_models_fetch_url(
&transport.provider.provider_type,
&transport.endpoint.api_format,
&transport.endpoint.base_url,
)
.ok_or_else(|| "Rust models fetch does not support this provider format yet".to_string())?;
if api_format.starts_with("claude:") {
url = append_query_param(url, "limit", "100");
if let Some(after_id) = after_id.map(str::trim).filter(|value| !value.is_empty()) {
url = append_query_param(url, "after_id", after_id);
}
}
Ok(url)
}
fn append_query_param(mut url: String, key: &str, value: &str) -> String {
if key.trim().is_empty() || value.trim().is_empty() {
return url;
}
let separator = if url.contains('?') { '&' } else { '?' };
url.push(separator);
url.push_str(key.trim());
url.push('=');
url.push_str(value.trim());
url
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use aether_contracts::ProxySnapshot; use aether_contracts::{ExecutionPlan, ExecutionResult, ProxySnapshot};
use aether_provider_transport::snapshot::{ use aether_provider_transport::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot, GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
}; };
use async_trait::async_trait; use async_trait::async_trait;
use serde_json::json;
use super::{build_models_fetch_execution_plan, ModelFetchTransportRuntime}; use super::{
build_antigravity_fetch_available_models_plan, build_gemini_cli_load_code_assist_plan,
build_models_fetch_execution_plan, build_standard_models_fetch_execution_plan,
build_vertex_models_fetch_execution_plan, ModelFetchTransportRuntime,
};
struct TestRuntime { struct TestRuntime {
oauth_auth: Option<aether_provider_transport::LocalResolvedOAuthRequestAuth>, oauth_auth: Option<aether_provider_transport::LocalResolvedOAuthRequestAuth>,
@@ -155,14 +507,25 @@ mod tests {
) -> Option<ProxySnapshot> { ) -> Option<ProxySnapshot> {
self.proxy.clone() self.proxy.clone()
} }
async fn execute_model_fetch_execution_plan(
&self,
_plan: &ExecutionPlan,
) -> Result<ExecutionResult, String> {
unreachable!("tests only validate plan construction")
}
} }
fn sample_transport(api_format: &str, auth_type: &str) -> GatewayProviderTransportSnapshot { fn sample_transport(
provider_type: &str,
api_format: &str,
auth_type: &str,
) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot { GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider { provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(), id: "provider-1".to_string(),
name: "Provider One".to_string(), name: "Provider One".to_string(),
provider_type: "openai".to_string(), provider_type: provider_type.to_string(),
website: None, website: None,
is_active: true, is_active: true,
keep_priority_on_conversion: false, keep_priority_on_conversion: false,
@@ -196,7 +559,7 @@ mod tests {
name: "key".to_string(), name: "key".to_string(),
auth_type: auth_type.to_string(), auth_type: auth_type.to_string(),
is_active: true, is_active: true,
api_formats: None, api_formats: Some(vec![api_format.to_string()]),
allowed_models: None, allowed_models: None,
capabilities: None, capabilities: None,
rate_multipliers: None, rate_multipliers: None,
@@ -205,35 +568,86 @@ mod tests {
proxy: None, proxy: None,
fingerprint: None, fingerprint: None,
decrypted_api_key: "secret".to_string(), decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None, decrypted_auth_config: Some(
r#"{"project_id":"project-1","client_version":"1.2.3","session_id":"sess-1"}"#
.to_string(),
),
}, },
} }
} }
#[tokio::test] #[tokio::test]
async fn builds_openai_models_fetch_plan_from_transport_snapshot() { async fn builds_openai_cli_models_fetch_plan_with_cli_user_agent() {
let runtime = TestRuntime { let runtime = TestRuntime {
oauth_auth: None, oauth_auth: None,
proxy: None, proxy: None,
}; };
let plan = build_models_fetch_execution_plan( let mut transport = sample_transport("openai", "openai:cli", "api_key");
&runtime, transport.key.decrypted_auth_config = None;
&sample_transport("openai:chat", "api_key"), let plan = build_models_fetch_execution_plan(&runtime, &transport)
) .await
.await .expect("plan");
.expect("plan");
assert_eq!(plan.method, "GET");
assert_eq!(plan.url, "https://example.com/v1/models"); assert_eq!(plan.url, "https://example.com/v1/models");
assert_eq!(
plan.headers.get("user-agent").map(String::as_str),
Some("openai-codex/1.0")
);
assert_eq!( assert_eq!(
plan.headers.get("authorization").map(String::as_str), plan.headers.get("authorization").map(String::as_str),
Some("Bearer secret") Some("Bearer secret")
); );
assert_eq!(plan.provider_api_format, "openai:chat");
} }
#[tokio::test] #[tokio::test]
async fn builds_oauth_models_fetch_plan_from_runtime_auth() { async fn builds_claude_models_fetch_plan_with_pagination() {
let runtime = TestRuntime {
oauth_auth: None,
proxy: None,
};
let mut transport = sample_transport("custom", "claude:chat", "api_key");
transport.key.decrypted_auth_config = None;
let plan =
build_standard_models_fetch_execution_plan(&runtime, &transport, Some("cursor-1"))
.await
.expect("plan");
assert_eq!(
plan.url,
"https://example.com/v1/models?limit=100&after_id=cursor-1"
);
assert_eq!(
plan.headers.get("anthropic-version").map(String::as_str),
Some("2023-06-01")
);
assert_eq!(
plan.headers.get("x-api-key").map(String::as_str),
Some("secret")
);
}
#[tokio::test]
async fn builds_gemini_models_fetch_plan_with_browser_headers_and_query_auth() {
let runtime = TestRuntime {
oauth_auth: None,
proxy: None,
};
let mut transport = sample_transport("custom", "gemini:chat", "api_key");
transport.key.decrypted_auth_config = None;
let plan = build_models_fetch_execution_plan(&runtime, &transport)
.await
.expect("plan");
assert_eq!(plan.url, "https://example.com/v1beta/models?key=secret");
assert!(plan.headers.contains_key("sec-ch-ua"));
assert_eq!(
plan.headers.get("accept-language").map(String::as_str),
Some("zh-CN")
);
}
#[tokio::test]
async fn builds_antigravity_fetch_available_models_plan() {
let runtime = TestRuntime { let runtime = TestRuntime {
oauth_auth: Some( oauth_auth: Some(
aether_provider_transport::LocalResolvedOAuthRequestAuth::Header { aether_provider_transport::LocalResolvedOAuthRequestAuth::Header {
@@ -241,27 +655,85 @@ mod tests {
value: "Bearer oauth-token".to_string(), value: "Bearer oauth-token".to_string(),
}, },
), ),
proxy: Some(ProxySnapshot { proxy: None,
enabled: Some(true),
mode: Some("fixed".to_string()),
node_id: None,
label: None,
url: Some("http://proxy.internal".to_string()),
extra: None,
}),
}; };
let plan = let transport = sample_transport("antigravity", "gemini:chat", "oauth");
build_models_fetch_execution_plan(&runtime, &sample_transport("openai:chat", "oauth")) let plan = build_antigravity_fetch_available_models_plan(
.await &runtime,
.expect("plan"); &transport,
"https://daily-cloudcode-pa.sandbox.googleapis.com",
"project-1",
)
.await
.expect("plan");
assert_eq!(plan.method, "POST");
assert_eq!( assert_eq!(
plan.headers.get("authorization").map(String::as_str), plan.url,
Some("Bearer oauth-token") "https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:fetchAvailableModels"
); );
assert_eq!( assert_eq!(
plan.proxy.as_ref().and_then(|proxy| proxy.url.as_deref()), plan.provider_api_format,
Some("http://proxy.internal") "antigravity:fetch_available_models"
);
assert_eq!(
plan.body
.json_body
.as_ref()
.and_then(|value| value.get("project")),
Some(&json!("project-1"))
); );
} }
#[tokio::test]
async fn builds_gemini_cli_load_code_assist_plan() {
let runtime = TestRuntime {
oauth_auth: Some(
aether_provider_transport::LocalResolvedOAuthRequestAuth::Header {
name: "authorization".to_string(),
value: "Bearer oauth-token".to_string(),
},
),
proxy: None,
};
let transport = sample_transport("gemini_cli", "gemini:cli", "oauth");
let plan = build_gemini_cli_load_code_assist_plan(&runtime, &transport)
.await
.expect("plan");
assert_eq!(plan.method, "POST");
assert_eq!(
plan.url,
"https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist"
);
assert_eq!(
plan.headers.get("user-agent").map(String::as_str),
Some("GeminiCLI/0.1.5 (Windows; AMD64)")
);
}
#[tokio::test]
async fn builds_vertex_models_fetch_plan_with_auth_override() {
let runtime = TestRuntime {
oauth_auth: None,
proxy: None,
};
let mut transport = sample_transport("vertex_ai", "claude:chat", "api_key");
transport.key.decrypted_auth_config = None;
let plan = build_vertex_models_fetch_execution_plan(
&runtime,
&transport,
"https://aiplatform.googleapis.com/v1/publishers/google/models?key=secret",
"gemini:chat",
None,
)
.await
.expect("plan");
assert_eq!(
plan.url,
"https://aiplatform.googleapis.com/v1/publishers/google/models?key=secret"
);
assert!(plan.headers.contains_key("sec-ch-ua"));
}
} }

View File

@@ -197,7 +197,7 @@ export const MOCK_DASHBOARD_STATS: DashboardStatsResponse = {
cache_read_tokens: 200000, cache_read_tokens: 200000,
cache_creation_cost: 0.25, cache_creation_cost: 0.25,
cache_read_cost: 0.10, cache_read_cost: 0.10,
cache_hit_rate: 0.35, cache_hit_rate: 35.0,
total_cache_tokens: 250000 total_cache_tokens: 250000
}, },
users: { users: {

View File

@@ -0,0 +1,19 @@
import { describe, expect, it } from 'vitest'
import { parseDateLike } from '../date'
describe('parseDateLike', () => {
it('parses date-only strings as local calendar dates', () => {
const date = parseDateLike('2026-04-12')
expect(date.getFullYear()).toBe(2026)
expect(date.getMonth()).toBe(3)
expect(date.getDate()).toBe(12)
})
it('keeps timestamp strings delegated to native Date parsing', () => {
const date = parseDateLike('2026-04-12T15:30:00Z')
expect(Number.isNaN(date.getTime())).toBe(false)
expect(date.toISOString()).toBe('2026-04-12T15:30:00.000Z')
})
})

View File

@@ -0,0 +1,15 @@
const DATE_ONLY_PATTERN = /^(\d{4})-(\d{2})-(\d{2})$/
/**
* 将 `YYYY-MM-DD` 解析为本地时区日期,避免浏览器按 UTC 解析后串天。
* 其他带时间/时区的信息仍交给原生 Date 处理。
*/
export function parseDateLike(dateString: string): Date {
const matched = DATE_ONLY_PATTERN.exec(dateString)
if (!matched) {
return new Date(dateString)
}
const [, year, month, day] = matched
return new Date(Number(year), Number(month) - 1, Number(day))
}

View File

@@ -819,6 +819,7 @@ import {
DollarSign, DollarSign,
Key, Key,
Hash, Hash,
Zap,
Bell, Bell,
AlertCircle, AlertCircle,
AlertTriangle, AlertTriangle,
@@ -830,6 +831,7 @@ import {
Shuffle Shuffle
} from 'lucide-vue-next' } from 'lucide-vue-next'
import { formatTokens, formatCurrency } from '@/utils/format' import { formatTokens, formatCurrency } from '@/utils/format'
import { parseDateLike } from '@/utils/date'
import { marked } from 'marked' import { marked } from 'marked'
import { sanitizeMarkdown } from '@/utils/sanitize' import { sanitizeMarkdown } from '@/utils/sanitize'
import type { ChartData, ChartOptions, ChartDataset, TooltipItem } from 'chart.js' import type { ChartData, ChartOptions, ChartDataset, TooltipItem } from 'chart.js'
@@ -1000,7 +1002,7 @@ const selectedAnnouncement = ref<Announcement | null>(null)
const detailDialogOpen = ref(false) const detailDialogOpen = ref(false)
const iconMap: Record<string, unknown> = { const iconMap: Record<string, unknown> = {
Users, Activity, TrendingUp, DollarSign, Key, Hash, Database Users, Activity, TrendingUp, DollarSign, Key, Hash, Zap, Database
} }
// 空状态占位卡片 // 空状态占位卡片
@@ -1315,7 +1317,10 @@ onBeforeUnmount(() => {
async function loadDashboardData() { async function loadDashboardData() {
loading.value = true loading.value = true
try { try {
const statsData = await dashboardApi.getStats() const statsData = await dashboardApi.getStats({
timezone: dailyTimeRange.value.timezone,
tz_offset_minutes: dailyTimeRange.value.tz_offset_minutes
})
stats.value = statsData.stats.map(stat => ({ stats.value = statsData.stats.map(stat => ({
...stat, ...stat,
icon: iconMap[stat.icon] || Activity icon: iconMap[stat.icon] || Activity
@@ -1382,7 +1387,7 @@ function scheduleDailyStatsLoad() {
watch(dailyTimeRange, scheduleDailyStatsLoad, { deep: true }) watch(dailyTimeRange, scheduleDailyStatsLoad, { deep: true })
function formatDate(dateString: string): string { function formatDate(dateString: string): string {
const date = new Date(dateString) const date = parseDateLike(dateString)
const today = new Date() const today = new Date()
const yesterday = new Date(today) const yesterday = new Date(today)
yesterday.setDate(yesterday.getDate() - 1) yesterday.setDate(yesterday.getDate() - 1)
@@ -1392,7 +1397,7 @@ function formatDate(dateString: string): string {
} }
function formatDateForChart(dateString: string): string { function formatDateForChart(dateString: string): string {
const date = new Date(dateString) const date = parseDateLike(dateString)
const today = new Date() const today = new Date()
const yesterday = new Date(today) const yesterday = new Date(today)
yesterday.setDate(yesterday.getDate() - 1) yesterday.setDate(yesterday.getDate() - 1)