mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 02:47:45 +08:00
Merge upstream main into feat/500-api-key-ip-whitelist
This commit is contained in:
@@ -38,7 +38,7 @@ pub(super) async fn build_admin_create_api_key_install_session_response(
|
||||
Err(_) => {
|
||||
return Ok(build_admin_api_keys_bad_request_response(
|
||||
"请求数据验证失败",
|
||||
))
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -28,6 +28,8 @@ pub(crate) struct AdminOAuthProviderUpsertRequest {
|
||||
#[serde(default)]
|
||||
pub(super) extra_config: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) icon_url: Option<String>,
|
||||
#[serde(default)]
|
||||
pub(super) is_enabled: bool,
|
||||
#[serde(default)]
|
||||
pub(super) force: bool,
|
||||
@@ -70,6 +72,7 @@ pub(super) fn build_admin_oauth_provider_payload(
|
||||
"frontend_callback_url": provider.frontend_callback_url,
|
||||
"attribute_mapping": provider.attribute_mapping,
|
||||
"extra_config": provider.extra_config,
|
||||
"icon_url": provider.icon_url,
|
||||
"is_enabled": provider.is_enabled,
|
||||
})
|
||||
}
|
||||
@@ -364,6 +367,10 @@ pub(super) fn build_admin_oauth_upsert_record(
|
||||
frontend_callback_url: frontend_callback_url.to_string(),
|
||||
attribute_mapping: payload.attribute_mapping,
|
||||
extra_config: payload.extra_config,
|
||||
icon_url: payload.icon_url.and_then(|value| {
|
||||
let value = value.trim().to_string();
|
||||
(!value.is_empty()).then_some(value)
|
||||
}),
|
||||
is_enabled: payload.is_enabled,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -16,7 +16,7 @@ use axum::{
|
||||
use serde_json::json;
|
||||
use std::time::Duration;
|
||||
|
||||
const ADMIN_OAUTH_TEST_TIMEOUT_SECS: u64 = 5;
|
||||
const ADMIN_OAUTH_TEST_TIMEOUT_SECS: u64 = 10;
|
||||
const LINUXDO_AUTHORIZATION_URL: &str = "https://connect.linux.do/oauth2/authorize";
|
||||
const LINUXDO_TOKEN_URL: &str = "https://connect.linux.do/oauth2/token";
|
||||
|
||||
@@ -54,10 +54,7 @@ async fn admin_oauth_endpoint_reachable(client: &reqwest::Client, url: &str) ->
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => {
|
||||
let status = response.status();
|
||||
status != reqwest::StatusCode::NOT_FOUND && status.as_u16() < 500
|
||||
}
|
||||
Ok(response) => response.status().as_u16() < 500,
|
||||
Err(_) => false,
|
||||
}
|
||||
}
|
||||
@@ -120,10 +117,16 @@ async fn build_admin_oauth_test_payload(
|
||||
}));
|
||||
};
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
let proxy_snapshot = state.app().resolve_system_proxy_snapshot().await;
|
||||
let mut client_builder = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(ADMIN_OAUTH_TEST_TIMEOUT_SECS))
|
||||
.redirect(reqwest::redirect::Policy::limited(3))
|
||||
.build();
|
||||
.redirect(reqwest::redirect::Policy::limited(3));
|
||||
if let Some(proxy_url) = proxy_snapshot.as_ref().and_then(|p| p.url.as_deref()) {
|
||||
if let Ok(proxy) = reqwest::Proxy::all(proxy_url) {
|
||||
client_builder = client_builder.proxy(proxy);
|
||||
}
|
||||
}
|
||||
let client = client_builder.build();
|
||||
let Ok(client) = client else {
|
||||
return Ok(json!({
|
||||
"authorization_url_reachable": false,
|
||||
|
||||
@@ -16,6 +16,7 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use tracing::warn;
|
||||
|
||||
pub(super) async fn maybe_build_local_admin_payment_orders_response(
|
||||
state: &AdminAppState<'_>,
|
||||
@@ -211,6 +212,19 @@ async fn build_admin_payment_credit_order_response(
|
||||
.await?
|
||||
{
|
||||
crate::AdminWalletMutationOutcome::Applied((order, credited)) => {
|
||||
if credited {
|
||||
if let Err(err) = state
|
||||
.app()
|
||||
.apply_referral_rewards_for_payment_order_id(&order.id)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
error = ?err,
|
||||
order_id = %order.id,
|
||||
"failed to apply referral rewards for admin-credited payment order"
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(attach_admin_audit_response(
|
||||
Json(json!({
|
||||
"order": build_admin_payment_order_payload(&order),
|
||||
|
||||
@@ -14,6 +14,7 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use tracing::warn;
|
||||
|
||||
pub(in super::super) async fn build_admin_wallet_complete_refund_response(
|
||||
state: &AdminAppState<'_>,
|
||||
@@ -86,6 +87,20 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
|
||||
.await?
|
||||
{
|
||||
crate::AdminWalletMutationOutcome::Applied(refund) => {
|
||||
if let Some(order_id) = refund.payment_order_id.as_deref() {
|
||||
if let Err(err) = state
|
||||
.app()
|
||||
.reverse_referral_rewards_for_order(order_id, refund.amount_usd)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
error = ?err,
|
||||
order_id = %order_id,
|
||||
refund_id = %refund.id,
|
||||
"failed to reverse referral rewards for completed refund"
|
||||
);
|
||||
}
|
||||
}
|
||||
let response = Json(json!({
|
||||
"refund": build_admin_wallet_refund_payload(&wallet, &owner, &refund),
|
||||
}))
|
||||
|
||||
@@ -6,6 +6,7 @@ pub(super) mod features;
|
||||
mod model;
|
||||
pub(super) mod observability;
|
||||
pub(super) mod provider;
|
||||
mod referrals;
|
||||
mod routing;
|
||||
mod system;
|
||||
mod users;
|
||||
|
||||
@@ -107,6 +107,13 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
let Some(provider) = providers.get(&model.provider_id) else {
|
||||
continue;
|
||||
};
|
||||
let provider_model_mapping_names =
|
||||
provider_model_mapping_names_for_routing(model.provider_model_mappings.as_ref());
|
||||
let key_match_model_names = key_match_model_names_for_routing(
|
||||
&global_model.name,
|
||||
&model.provider_model_name,
|
||||
&provider_model_mapping_names,
|
||||
);
|
||||
let mut endpoint_payloads = Vec::new();
|
||||
let mut active_endpoints = 0usize;
|
||||
for endpoint in endpoints_by_provider
|
||||
@@ -132,7 +139,7 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
.filter(|key| {
|
||||
key_allowed_models_match_global_model_for_routing(
|
||||
key.allowed_models.as_ref(),
|
||||
&global_model.name,
|
||||
&key_match_model_names,
|
||||
&global_model_mappings,
|
||||
)
|
||||
})
|
||||
@@ -313,7 +320,7 @@ pub(crate) async fn build_admin_global_model_routing_payload(
|
||||
|
||||
fn key_allowed_models_match_global_model_for_routing(
|
||||
raw_allowed_models: Option<&serde_json::Value>,
|
||||
global_model_name: &str,
|
||||
model_names: &[String],
|
||||
global_model_mappings: &[String],
|
||||
) -> bool {
|
||||
// 兼容 Python 预览逻辑:None/[] 视为“不限制”,在链路预览中保留该 Key。
|
||||
@@ -322,14 +329,16 @@ fn key_allowed_models_match_global_model_for_routing(
|
||||
return true;
|
||||
}
|
||||
|
||||
if allowed_models
|
||||
.iter()
|
||||
.any(|value| value == global_model_name)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
|
||||
for allowed_model in &allowed_models {
|
||||
for allowed_model in allowed_models.iter().map(String::as_str).map(str::trim) {
|
||||
if allowed_model.is_empty() {
|
||||
continue;
|
||||
}
|
||||
if model_names
|
||||
.iter()
|
||||
.any(|model_name| model_name.eq_ignore_ascii_case(allowed_model))
|
||||
{
|
||||
return true;
|
||||
}
|
||||
for pattern in global_model_mappings {
|
||||
if matches_model_mapping(pattern, allowed_model) {
|
||||
return true;
|
||||
@@ -340,6 +349,54 @@ fn key_allowed_models_match_global_model_for_routing(
|
||||
false
|
||||
}
|
||||
|
||||
fn provider_model_mapping_names_for_routing(
|
||||
raw_mappings: Option<&serde_json::Value>,
|
||||
) -> Vec<String> {
|
||||
raw_mappings
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
item.as_str()
|
||||
.or_else(|| item.get("name").and_then(serde_json::Value::as_str))
|
||||
})
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn key_match_model_names_for_routing(
|
||||
global_model_name: &str,
|
||||
provider_model_name: &str,
|
||||
provider_model_mapping_names: &[String],
|
||||
) -> Vec<String> {
|
||||
let mut names = Vec::new();
|
||||
push_unique_model_name(&mut names, global_model_name);
|
||||
push_unique_model_name(&mut names, provider_model_name);
|
||||
for mapping_name in provider_model_mapping_names {
|
||||
push_unique_model_name(&mut names, mapping_name);
|
||||
}
|
||||
names
|
||||
}
|
||||
|
||||
fn push_unique_model_name(names: &mut Vec<String>, value: &str) {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
return;
|
||||
}
|
||||
if names
|
||||
.iter()
|
||||
.any(|existing| existing.eq_ignore_ascii_case(value))
|
||||
{
|
||||
return;
|
||||
}
|
||||
names.push(value.to_string());
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_assign_global_model_to_providers_payload(
|
||||
state: &AdminAppState<'_>,
|
||||
global_model_id: &str,
|
||||
|
||||
@@ -6,6 +6,7 @@ use super::route_filters::{
|
||||
};
|
||||
use crate::constants::INTERNAL_GATEWAY_PATH_PREFIXES;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::build_admin_usage_counter_health_payload;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::observability::monitoring::{
|
||||
admin_monitoring_bad_request_response, admin_monitoring_user_behavior_user_id_from_path,
|
||||
@@ -189,6 +190,13 @@ pub(super) async fn build_admin_monitoring_system_status_response(
|
||||
)
|
||||
.unwrap_or(usize::MAX);
|
||||
let tunnel = state.tunnel.stats();
|
||||
let usage_counter_snapshot = state
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let usage_counter =
|
||||
build_admin_usage_counter_health_payload(&usage_counter_snapshot, now_unix_secs);
|
||||
|
||||
Ok(build_admin_monitoring_system_status_payload_response(
|
||||
now,
|
||||
@@ -206,5 +214,6 @@ pub(super) async fn build_admin_monitoring_system_status_response(
|
||||
tunnel.active_streams,
|
||||
INTERNAL_GATEWAY_PATH_PREFIXES,
|
||||
recent_errors,
|
||||
usage_counter,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -249,9 +249,10 @@ async fn admin_monitoring_resilience_status_returns_local_payload() {
|
||||
let recommendations = payload["recommendations"]
|
||||
.as_array()
|
||||
.expect("recommendations should be array");
|
||||
assert!(recommendations.iter().any(|item| item
|
||||
.as_str()
|
||||
.is_some_and(|value| value.contains("prod-key"))));
|
||||
assert!(recommendations.iter().any(|item| {
|
||||
item.as_str()
|
||||
.is_some_and(|value| value.contains("prod-key"))
|
||||
}));
|
||||
assert!(payload["timestamp"].as_str().is_some());
|
||||
}
|
||||
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use super::range::{build_comparison_range, parse_bounded_u32};
|
||||
use super::resolve_admin_usage_time_range;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
|
||||
use crate::handlers::admin::shared::{
|
||||
build_admin_usage_counter_health_payload, query_param_optional_bool, query_param_value,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::observability::stats::{
|
||||
admin_stats_bad_request_response, admin_stats_comparison_empty_response,
|
||||
@@ -21,6 +23,22 @@ use aether_data_contracts::repository::usage::{
|
||||
};
|
||||
use axum::{body::Body, http, response::Response};
|
||||
|
||||
async fn build_usage_counter_health_payload(
|
||||
state: &AdminAppState<'_>,
|
||||
) -> Result<serde_json::Value, GatewayError> {
|
||||
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
|
||||
let snapshot = state
|
||||
.as_ref()
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
Ok(build_admin_usage_counter_health_payload(
|
||||
&snapshot,
|
||||
now_unix_secs,
|
||||
))
|
||||
}
|
||||
|
||||
fn usage_summary_to_admin_stats_aggregate(
|
||||
summary: &aether_data_contracts::repository::usage::StoredUsageAuditSummary,
|
||||
) -> AdminStatsAggregate {
|
||||
@@ -211,13 +229,18 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
|
||||
Ok(value) => u64::from(value.unwrap_or(10_000)),
|
||||
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
|
||||
};
|
||||
let usage_counter = build_usage_counter_health_payload(state).await?;
|
||||
if !state.has_usage_data_reader() {
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response()));
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response(
|
||||
usage_counter,
|
||||
)));
|
||||
}
|
||||
|
||||
let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds()
|
||||
else {
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response()));
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response(
|
||||
usage_counter,
|
||||
)));
|
||||
};
|
||||
let performance = state
|
||||
.summarize_usage_provider_performance(&UsageProviderPerformanceQuery {
|
||||
@@ -240,6 +263,7 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
|
||||
.await?;
|
||||
return Ok(Some(build_admin_stats_provider_performance_response(
|
||||
&performance,
|
||||
usage_counter,
|
||||
)));
|
||||
}
|
||||
|
||||
|
||||
@@ -5,12 +5,13 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::query_param_value;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::observability::usage::{
|
||||
admin_usage_bad_request_response, admin_usage_data_unavailable_response,
|
||||
admin_usage_has_fallback, admin_usage_is_failed, admin_usage_matches_search,
|
||||
admin_usage_matches_username, admin_usage_parse_ids, admin_usage_parse_limit,
|
||||
admin_usage_parse_offset, admin_usage_provider_key_name, admin_usage_record_json,
|
||||
build_admin_usage_active_requests_response, build_admin_usage_records_response,
|
||||
build_admin_usage_summary_stats_response_from_summary, ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
|
||||
admin_usage_bad_request_response, admin_usage_client_family,
|
||||
admin_usage_data_unavailable_response, admin_usage_has_fallback, admin_usage_is_failed,
|
||||
admin_usage_matches_search, admin_usage_matches_username, admin_usage_parse_ids,
|
||||
admin_usage_parse_limit, admin_usage_parse_offset, admin_usage_provider_key_name,
|
||||
admin_usage_record_json, build_admin_usage_active_requests_response,
|
||||
build_admin_usage_records_response, build_admin_usage_summary_stats_response_from_summary,
|
||||
ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
|
||||
};
|
||||
use aether_data::repository::users::StoredUserSummary;
|
||||
use aether_data_contracts::repository::{
|
||||
@@ -263,6 +264,19 @@ fn admin_usage_matches_attempt_status(
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_usage_matches_client_family(
|
||||
item: &StoredRequestUsageAudit,
|
||||
client_family: Option<&str>,
|
||||
) -> bool {
|
||||
let Some(client_family) = client_family
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return true;
|
||||
};
|
||||
admin_usage_client_family(item).is_some_and(|value| value.eq_ignore_ascii_case(client_family))
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn build_admin_usage_records_response_with_attempt_flags(
|
||||
items: &[StoredRequestUsageAudit],
|
||||
@@ -600,6 +614,7 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
|
||||
admin_usage_attempt_status_filter(query_param_value(query, "status").as_deref());
|
||||
let search = query_param_value(query, "search");
|
||||
let username_filter = query_param_value(query, "username");
|
||||
let client_family_filter = query_param_value(query, "client_family");
|
||||
let limit = match admin_usage_parse_limit(query) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(Some(admin_usage_bad_request_response(detail))),
|
||||
@@ -634,7 +649,12 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
|
||||
let active_username_filter = username_filter
|
||||
.as_deref()
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
let (usage, total) = if let Some(attempt_status) = attempt_status_filter {
|
||||
let active_client_family_filter = client_family_filter
|
||||
.as_deref()
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
let (usage, total) = if attempt_status_filter.is_some()
|
||||
|| active_client_family_filter.is_some()
|
||||
{
|
||||
let mut usage = state.list_usage_audits(&base_query).await?;
|
||||
let user_ids: Vec<String> = usage
|
||||
.iter()
|
||||
@@ -664,12 +684,14 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
|
||||
active_username_filter,
|
||||
&users_by_id,
|
||||
state.has_auth_user_data_reader(),
|
||||
) && admin_usage_matches_attempt_status(
|
||||
item,
|
||||
attempt_status,
|
||||
&attempt_flags_by_usage_id,
|
||||
request_candidate_reader_available,
|
||||
)
|
||||
) && attempt_status_filter.is_none_or(|attempt_status| {
|
||||
admin_usage_matches_attempt_status(
|
||||
item,
|
||||
attempt_status,
|
||||
&attempt_flags_by_usage_id,
|
||||
request_candidate_reader_available,
|
||||
)
|
||||
}) && admin_usage_matches_client_family(item, active_client_family_filter)
|
||||
});
|
||||
sort_usage_newest_first(&mut usage);
|
||||
let total = usage.len();
|
||||
|
||||
@@ -23,6 +23,10 @@ pub(super) fn endpoint_key_counts_by_format(
|
||||
admin_provider_endpoints_pure::endpoint_key_counts_by_format(provider_type, endpoints, keys)
|
||||
}
|
||||
|
||||
pub(super) fn normalize_endpoint_api_format(api_format: &str) -> String {
|
||||
admin_provider_endpoints_pure::normalize_endpoint_api_format(api_format)
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_endpoint_response(
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
provider_name: &str,
|
||||
|
||||
@@ -4,7 +4,10 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use super::payloads::{build_admin_provider_endpoint_response, endpoint_key_counts_by_format};
|
||||
use super::payloads::{
|
||||
build_admin_provider_endpoint_response, endpoint_key_counts_by_format,
|
||||
normalize_endpoint_api_format,
|
||||
};
|
||||
|
||||
pub(crate) async fn build_admin_provider_endpoints_payload(
|
||||
state: &AdminAppState<'_>,
|
||||
@@ -52,8 +55,7 @@ pub(crate) async fn build_admin_provider_endpoints_payload(
|
||||
.skip(skip)
|
||||
.take(limit)
|
||||
.map(|endpoint| {
|
||||
let endpoint_api_format =
|
||||
crate::ai_serving::normalize_api_format_alias(&endpoint.api_format);
|
||||
let endpoint_api_format = normalize_endpoint_api_format(&endpoint.api_format);
|
||||
build_admin_provider_endpoint_response(
|
||||
&endpoint,
|
||||
&provider.name,
|
||||
@@ -105,7 +107,7 @@ pub(crate) async fn build_admin_endpoint_payload(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let endpoint_api_format = crate::ai_serving::normalize_api_format_alias(&endpoint.api_format);
|
||||
let endpoint_api_format = normalize_endpoint_api_format(&endpoint.api_format);
|
||||
|
||||
Some(build_admin_provider_endpoint_response(
|
||||
&endpoint,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use super::extractors::admin_endpoint_id;
|
||||
use super::payloads::{
|
||||
build_admin_provider_endpoint_response, endpoint_key_counts_by_format,
|
||||
AdminProviderEndpointUpdatePatch,
|
||||
normalize_endpoint_api_format, AdminProviderEndpointUpdatePatch,
|
||||
};
|
||||
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
@@ -152,7 +152,7 @@ pub(super) async fn maybe_handle(
|
||||
std::slice::from_ref(&updated),
|
||||
&keys,
|
||||
);
|
||||
let updated_api_format = crate::ai_serving::normalize_api_format_alias(&updated.api_format);
|
||||
let updated_api_format = normalize_endpoint_api_format(&updated.api_format);
|
||||
|
||||
Ok(Some(
|
||||
Json(build_admin_provider_endpoint_response(
|
||||
|
||||
@@ -32,7 +32,7 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(_provider) = state
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
@@ -129,7 +129,9 @@ pub(super) async fn maybe_handle(
|
||||
Json(serde_json::Value::Array(
|
||||
created
|
||||
.iter()
|
||||
.map(|model| build_admin_provider_model_response(model, now_unix_secs))
|
||||
.map(|model| {
|
||||
build_admin_provider_model_response(&provider, model, now_unix_secs)
|
||||
})
|
||||
.collect(),
|
||||
))
|
||||
.into_response(),
|
||||
|
||||
@@ -31,7 +31,7 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(_provider) = state
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
@@ -90,8 +90,12 @@ pub(super) async fn maybe_handle(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
Json(build_admin_provider_model_response(&created, now_unix_secs))
|
||||
.into_response()
|
||||
Json(build_admin_provider_model_response(
|
||||
&provider,
|
||||
&created,
|
||||
now_unix_secs,
|
||||
))
|
||||
.into_response()
|
||||
}
|
||||
None => (
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
use crate::handlers::admin::provider::shared::model_test_capabilities::{
|
||||
admin_provider_model_supports_image_generation, admin_provider_model_test_capabilities_payload,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::models as admin_provider_models_pure;
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminProviderModelListQuery, StoredAdminProviderModel,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) fn admin_provider_model_effective_input_price(
|
||||
@@ -26,10 +30,32 @@ pub(super) fn admin_provider_model_effective_capability(
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_model_response(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
model: &StoredAdminProviderModel,
|
||||
now_unix_secs: u64,
|
||||
) -> serde_json::Value {
|
||||
admin_provider_models_pure::build_admin_provider_model_response(model, now_unix_secs)
|
||||
let mut payload =
|
||||
admin_provider_models_pure::build_admin_provider_model_response(model, now_unix_secs);
|
||||
let fallback_supports_image_generation = payload
|
||||
.get("effective_supports_image_generation")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let supports_image_generation = admin_provider_model_supports_image_generation(
|
||||
&provider.provider_type,
|
||||
&model.provider_model_name,
|
||||
fallback_supports_image_generation,
|
||||
);
|
||||
if let Some(object) = payload.as_object_mut() {
|
||||
object.insert(
|
||||
"model_test_capabilities".to_string(),
|
||||
admin_provider_model_test_capabilities_payload(
|
||||
&provider.provider_type,
|
||||
&model.provider_model_name,
|
||||
supports_image_generation,
|
||||
),
|
||||
);
|
||||
}
|
||||
payload
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_provider_models_payload(
|
||||
@@ -48,9 +74,10 @@ pub(super) async fn build_admin_provider_models_payload(
|
||||
.ok()?
|
||||
.into_iter()
|
||||
.next()?;
|
||||
let provider_id = provider.id.clone();
|
||||
let mut models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider.id,
|
||||
provider_id,
|
||||
is_active,
|
||||
offset: skip,
|
||||
limit,
|
||||
@@ -70,7 +97,7 @@ pub(super) async fn build_admin_provider_models_payload(
|
||||
Some(serde_json::Value::Array(
|
||||
models
|
||||
.iter()
|
||||
.map(|model| build_admin_provider_model_response(model, now_unix_secs))
|
||||
.map(|model| build_admin_provider_model_response(&provider, model, now_unix_secs))
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
@@ -80,9 +107,15 @@ pub(super) async fn build_admin_provider_model_payload(
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Option<serde_json::Value> {
|
||||
if !state.has_global_model_data_reader() {
|
||||
if !state.has_provider_catalog_data_reader() || !state.has_global_model_data_reader() {
|
||||
return None;
|
||||
}
|
||||
let provider = state
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await
|
||||
.ok()?
|
||||
.into_iter()
|
||||
.next()?;
|
||||
let model = state
|
||||
.get_admin_provider_model(provider_id, model_id)
|
||||
.await
|
||||
@@ -92,7 +125,11 @@ pub(super) async fn build_admin_provider_model_payload(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
Some(build_admin_provider_model_response(&model, now_unix_secs))
|
||||
Some(build_admin_provider_model_response(
|
||||
&provider,
|
||||
&model,
|
||||
now_unix_secs,
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) async fn admin_provider_model_name_exists(
|
||||
|
||||
@@ -33,6 +33,20 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
|
||||
)
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(existing) = state
|
||||
.get_admin_provider_model(&provider_id, &model_id)
|
||||
.await?
|
||||
@@ -110,8 +124,12 @@ pub(super) async fn maybe_handle(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
Json(build_admin_provider_model_response(&updated, now_unix_secs))
|
||||
.into_response()
|
||||
Json(build_admin_provider_model_response(
|
||||
&provider,
|
||||
&updated,
|
||||
now_unix_secs,
|
||||
))
|
||||
.into_response()
|
||||
}
|
||||
None => (
|
||||
http::StatusCode::NOT_FOUND,
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use super::super::helpers::admin_provider_oauth_key_name_from_auth_config;
|
||||
use super::super::token_import::{
|
||||
build_provider_access_token_import_auth_config, provider_type_supports_access_token_import,
|
||||
};
|
||||
@@ -24,13 +25,11 @@ use crate::handlers::admin::provider::oauth::runtime::{
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, exchange_admin_provider_oauth_refresh_token,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::oauth::parse_admin_provider_oauth_kiro_batch_import_entries;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
struct AdminProviderOAuthResolvedBatchImport {
|
||||
access_token: String,
|
||||
@@ -45,7 +44,7 @@ pub(super) fn estimate_admin_provider_oauth_batch_import_total(
|
||||
if provider_type.eq_ignore_ascii_case("kiro") {
|
||||
parse_admin_provider_oauth_kiro_batch_import_entries(raw_credentials).len()
|
||||
} else {
|
||||
parse_admin_provider_oauth_batch_import_entries(raw_credentials).len()
|
||||
parse_admin_provider_oauth_batch_import_entries(provider_type, raw_credentials).len()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -67,7 +66,8 @@ pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(raw_credentials);
|
||||
let entries =
|
||||
parse_admin_provider_oauth_batch_import_entries(provider_type, raw_credentials);
|
||||
execute_admin_provider_oauth_batch_import(
|
||||
state,
|
||||
provider_id,
|
||||
@@ -82,7 +82,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
|
||||
async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
state: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
template: Option<AdminProviderOAuthTemplate>,
|
||||
provider_type: &str,
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
@@ -99,6 +99,29 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
if let Some(refresh_token) = refresh_token {
|
||||
let Some(template) = template else {
|
||||
if provider_type_supports_access_token_import(provider_type) {
|
||||
if let Some(access_token) = access_token {
|
||||
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
|
||||
provider_type,
|
||||
access_token,
|
||||
Some(refresh_token),
|
||||
entry.expires_at,
|
||||
Some("Provider 不支持 Refresh Token 交换,已回退为 Session Token 导入"),
|
||||
);
|
||||
return Ok(AdminProviderOAuthResolvedBatchImport {
|
||||
access_token: access_token.to_string(),
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
}
|
||||
return Err(
|
||||
"该 Provider 不支持 Refresh Token 导入,请提供 sso_token 或 access_token"
|
||||
.to_string(),
|
||||
);
|
||||
};
|
||||
|
||||
let token_payload = match exchange_admin_provider_oauth_refresh_token(
|
||||
state,
|
||||
template,
|
||||
@@ -152,7 +175,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
|
||||
if let Some(access_token) = access_token {
|
||||
if !provider_type_supports_access_token_import(provider_type) {
|
||||
return Err("Access Token 导入仅支持 Codex / ChatGPT Web Provider".to_string());
|
||||
return Err("Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider".to_string());
|
||||
}
|
||||
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
|
||||
provider_type,
|
||||
@@ -204,25 +227,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
});
|
||||
};
|
||||
|
||||
let Some(template) = admin_provider_oauth_template(provider_type) else {
|
||||
return Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success: 0,
|
||||
failed: entries.len(),
|
||||
results: entries
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, _)| {
|
||||
json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
|
||||
"replaced": false,
|
||||
})
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
};
|
||||
let template = admin_provider_oauth_template(provider_type);
|
||||
|
||||
let endpoint_resolution =
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, provider_type).await?;
|
||||
@@ -340,24 +345,11 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let key_name = auth_config
|
||||
.get("email")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|email| format!("{provider_type}_{email}"))
|
||||
.unwrap_or_else(|| {
|
||||
format!(
|
||||
"{}_{}_{}",
|
||||
provider_type,
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0),
|
||||
index
|
||||
)
|
||||
});
|
||||
let key_name = admin_provider_oauth_key_name_from_auth_config(
|
||||
provider_type,
|
||||
&auth_config,
|
||||
Some(index),
|
||||
);
|
||||
match create_provider_oauth_catalog_key(
|
||||
state,
|
||||
provider_id,
|
||||
|
||||
+1
-5
@@ -8,7 +8,7 @@ use super::parse::{
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_provider_id;
|
||||
@@ -60,10 +60,6 @@ pub(in super::super) async fn handle_admin_provider_oauth_batch_import(
|
||||
"该 Provider 不是固定类型,无法使用 provider-oauth",
|
||||
));
|
||||
}
|
||||
if provider_type != "kiro" && admin_provider_oauth_template(&provider_type).is_none() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use super::super::token_import::{import_tokens_from_raw_token, normalize_single_import_tokens};
|
||||
use super::super::token_import::{import_tokens_from_raw_token, normalize_provider_import_tokens};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{current_unix_secs, json_u64_value};
|
||||
use axum::{
|
||||
@@ -25,8 +25,15 @@ pub(super) struct AdminProviderOAuthBatchImportEntry {
|
||||
pub account_id: Option<String>,
|
||||
pub account_user_id: Option<String>,
|
||||
pub plan_type: Option<String>,
|
||||
pub pool_tier: Option<String>,
|
||||
pub user_id: Option<String>,
|
||||
pub email: Option<String>,
|
||||
pub account_name: Option<String>,
|
||||
pub sso_rw_token: Option<String>,
|
||||
pub cf_cookies: Option<String>,
|
||||
pub cf_clearance: Option<String>,
|
||||
pub user_agent: Option<String>,
|
||||
pub browser_profile: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -67,16 +74,72 @@ fn coerce_admin_provider_oauth_import_str(value: Option<&serde_json::Value>) ->
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn grok_cookie_value(raw: &str, name: &str) -> Option<String> {
|
||||
raw.trim()
|
||||
.strip_prefix("Cookie:")
|
||||
.unwrap_or_else(|| raw.trim())
|
||||
.split(';')
|
||||
.filter_map(|segment| segment.trim().split_once('='))
|
||||
.find_map(|(cookie_name, cookie_value)| {
|
||||
cookie_name
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(name)
|
||||
.then(|| cookie_value.trim())
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
fn grok_cookie_profile(raw: &str) -> Option<String> {
|
||||
let raw = raw
|
||||
.trim()
|
||||
.strip_prefix("Cookie:")
|
||||
.unwrap_or_else(|| raw.trim());
|
||||
let parts = raw
|
||||
.split(';')
|
||||
.filter_map(|segment| {
|
||||
let (cookie_name, cookie_value) = segment.trim().split_once('=')?;
|
||||
let cookie_name = cookie_name.trim();
|
||||
let cookie_value = cookie_value.trim();
|
||||
if cookie_name.is_empty()
|
||||
|| cookie_value.is_empty()
|
||||
|| cookie_name.eq_ignore_ascii_case("sso")
|
||||
|| cookie_name.eq_ignore_ascii_case("sso-rw")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(format!("{cookie_name}={cookie_value}"))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
(!parts.is_empty()).then(|| parts.join("; "))
|
||||
}
|
||||
|
||||
fn grok_cookie_session_token(provider_type: &str, raw: &str) -> Option<String> {
|
||||
provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok")
|
||||
.then(|| grok_cookie_value(raw, "sso"))
|
||||
.flatten()
|
||||
}
|
||||
|
||||
fn extract_admin_provider_oauth_batch_import_entry(
|
||||
provider_type: &str,
|
||||
item: &serde_json::Value,
|
||||
) -> Option<AdminProviderOAuthBatchImportEntry> {
|
||||
match item {
|
||||
serde_json::Value::String(value) => {
|
||||
let refresh_token = value.trim();
|
||||
if refresh_token.is_empty() {
|
||||
let raw_token = value.trim();
|
||||
if raw_token.is_empty() {
|
||||
None
|
||||
} else {
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(refresh_token);
|
||||
let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token);
|
||||
let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token);
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(token_input);
|
||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||
provider_type,
|
||||
refresh_token.as_deref(),
|
||||
access_token.as_deref(),
|
||||
);
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
@@ -84,8 +147,15 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
pool_tier: None,
|
||||
user_id: grok_cookie_value(raw_token, "x-userid"),
|
||||
email: None,
|
||||
account_name: None,
|
||||
sso_rw_token: grok_cookie_value(raw_token, "sso-rw"),
|
||||
cf_cookies: grok_cookie_profile(raw_token),
|
||||
cf_clearance: grok_cookie_value(raw_token, "cf_clearance"),
|
||||
user_agent: None,
|
||||
browser_profile: None,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -100,8 +170,34 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
.get("access_token")
|
||||
.or_else(|| object.get("accessToken")),
|
||||
);
|
||||
let (refresh_token, access_token) =
|
||||
normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref());
|
||||
let grok_token_alias = if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
object.get("token")
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let grok_cookie = if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
object.get("cookie").or_else(|| object.get("cookieHeader")),
|
||||
)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let session_token = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("sso_token")
|
||||
.or_else(|| object.get("ssoToken"))
|
||||
.or(grok_token_alias),
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "sso"))
|
||||
});
|
||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||
provider_type,
|
||||
refresh_token.as_deref(),
|
||||
access_token.as_deref().or(session_token.as_deref()),
|
||||
);
|
||||
if refresh_token.is_none() && access_token.is_none() {
|
||||
return None;
|
||||
}
|
||||
@@ -129,14 +225,65 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
.or_else(|| object.get("chatgptPlanType")),
|
||||
)
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
let pool_tier = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("pool_tier")
|
||||
.or_else(|| object.get("poolTier"))
|
||||
.or_else(|| object.get("tier")),
|
||||
)
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
let user_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("user_id")
|
||||
.or_else(|| object.get("userId"))
|
||||
.or_else(|| object.get("chatgpt_user_id"))
|
||||
.or_else(|| object.get("chatgptUserId")),
|
||||
);
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "x-userid"))
|
||||
});
|
||||
let email = coerce_admin_provider_oauth_import_str(object.get("email"));
|
||||
let account_name = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_name")
|
||||
.or_else(|| object.get("accountName")),
|
||||
);
|
||||
let sso_rw_token = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("sso_rw_token")
|
||||
.or_else(|| object.get("ssoRwToken")),
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "sso-rw"))
|
||||
});
|
||||
let cf_clearance = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("cf_clearance")
|
||||
.or_else(|| object.get("cfClearance")),
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "cf_clearance"))
|
||||
});
|
||||
let cf_cookies = coerce_admin_provider_oauth_import_str(
|
||||
object.get("cf_cookies").or_else(|| object.get("cfCookies")),
|
||||
)
|
||||
.or_else(|| grok_cookie.as_deref().and_then(grok_cookie_profile));
|
||||
let user_agent = coerce_admin_provider_oauth_import_str(
|
||||
object.get("user_agent").or_else(|| object.get("userAgent")),
|
||||
);
|
||||
let browser_profile = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("browser_profile")
|
||||
.or_else(|| object.get("browserProfile"))
|
||||
.or_else(|| object.get("browser"))
|
||||
.or_else(|| object.get("impersonate")),
|
||||
);
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
@@ -144,8 +291,15 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
account_id,
|
||||
account_user_id,
|
||||
plan_type,
|
||||
pool_tier,
|
||||
user_id,
|
||||
email,
|
||||
account_name,
|
||||
sso_rw_token,
|
||||
cf_cookies,
|
||||
cf_clearance,
|
||||
user_agent,
|
||||
browser_profile,
|
||||
})
|
||||
}
|
||||
_ => None,
|
||||
@@ -153,6 +307,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
provider_type: &str,
|
||||
raw_credentials: &str,
|
||||
) -> Vec<AdminProviderOAuthBatchImportEntry> {
|
||||
let raw = raw_credentials.trim();
|
||||
@@ -165,7 +320,9 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
{
|
||||
return items
|
||||
.iter()
|
||||
.filter_map(extract_admin_provider_oauth_batch_import_entry)
|
||||
.filter_map(|item| {
|
||||
extract_admin_provider_oauth_batch_import_entry(provider_type, item)
|
||||
})
|
||||
.collect();
|
||||
}
|
||||
}
|
||||
@@ -174,7 +331,7 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
if let Ok(value @ serde_json::Value::Object(_)) =
|
||||
serde_json::from_str::<serde_json::Value>(raw)
|
||||
{
|
||||
return extract_admin_provider_oauth_batch_import_entry(&value)
|
||||
return extract_admin_provider_oauth_batch_import_entry(provider_type, &value)
|
||||
.into_iter()
|
||||
.collect();
|
||||
}
|
||||
@@ -183,18 +340,19 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
raw.lines()
|
||||
.map(str::trim)
|
||||
.filter(|line| !line.is_empty() && !line.starts_with('#'))
|
||||
.map(|token| {
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(token);
|
||||
AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
expires_at: None,
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
email: None,
|
||||
.filter_map(|line| {
|
||||
if line.starts_with('{') {
|
||||
return serde_json::from_str::<serde_json::Value>(line)
|
||||
.ok()
|
||||
.and_then(|value| {
|
||||
extract_admin_provider_oauth_batch_import_entry(provider_type, &value)
|
||||
});
|
||||
}
|
||||
|
||||
extract_admin_provider_oauth_batch_import_entry(
|
||||
provider_type,
|
||||
&serde_json::Value::String(line.to_string()),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -204,10 +362,8 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
auth_config: &mut serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
if !matches!(
|
||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||
"codex" | "chatgpt_web"
|
||||
) {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") {
|
||||
return;
|
||||
}
|
||||
if let Some(account_id) = entry.account_id.as_ref() {
|
||||
@@ -225,6 +381,11 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
.entry("plan_type".to_string())
|
||||
.or_insert_with(|| json!(plan_type));
|
||||
}
|
||||
if let Some(pool_tier) = entry.pool_tier.as_ref() {
|
||||
auth_config
|
||||
.entry("pool_tier".to_string())
|
||||
.or_insert_with(|| json!(pool_tier));
|
||||
}
|
||||
if let Some(user_id) = entry.user_id.as_ref() {
|
||||
auth_config
|
||||
.entry("user_id".to_string())
|
||||
@@ -235,6 +396,36 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
.entry("email".to_string())
|
||||
.or_insert_with(|| json!(email));
|
||||
}
|
||||
if let Some(account_name) = entry.account_name.as_ref() {
|
||||
auth_config
|
||||
.entry("account_name".to_string())
|
||||
.or_insert_with(|| json!(account_name));
|
||||
}
|
||||
if let Some(sso_rw_token) = entry.sso_rw_token.as_ref() {
|
||||
auth_config
|
||||
.entry("sso_rw_token".to_string())
|
||||
.or_insert_with(|| json!(sso_rw_token));
|
||||
}
|
||||
if let Some(cf_cookies) = entry.cf_cookies.as_ref() {
|
||||
auth_config
|
||||
.entry("cf_cookies".to_string())
|
||||
.or_insert_with(|| json!(cf_cookies));
|
||||
}
|
||||
if let Some(cf_clearance) = entry.cf_clearance.as_ref() {
|
||||
auth_config
|
||||
.entry("cf_clearance".to_string())
|
||||
.or_insert_with(|| json!(cf_clearance));
|
||||
}
|
||||
if let Some(user_agent) = entry.user_agent.as_ref() {
|
||||
auth_config
|
||||
.entry("user_agent".to_string())
|
||||
.or_insert_with(|| json!(user_agent));
|
||||
}
|
||||
if let Some(browser_profile) = entry.browser_profile.as_ref() {
|
||||
auth_config
|
||||
.entry("browser_profile".to_string())
|
||||
.or_insert_with(|| json!(browser_profile));
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn extract_admin_provider_oauth_batch_error_detail(
|
||||
@@ -337,6 +528,7 @@ mod tests {
|
||||
#[test]
|
||||
fn parses_access_token_only_entry() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"codex",
|
||||
r#"[{"accessToken":"at_1","expiresAt":2100000000,"accountId":"acc-1","email":"[email protected]"}]"#,
|
||||
);
|
||||
|
||||
@@ -356,10 +548,89 @@ mod tests {
|
||||
"exp": 2_000_000_000u64,
|
||||
}));
|
||||
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(&token);
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries("codex", &token);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some(token.as_str()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_jsonl_session_entries() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
r#"{"sso_token":"sso-1","cf_clearance":"cf-1","pool_tier":"heavy","email":"[email protected]","browser_profile":"chrome136"}"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
|
||||
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
|
||||
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
|
||||
assert_eq!(entries[0].browser_profile.as_deref(), Some("chrome136"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_token_alias_with_account_traits() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
r#"[{"token":"sso-1","planType":"super","tier":"heavy","accountName":"Grok Heavy"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].plan_type.as_deref(), Some("super"));
|
||||
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
|
||||
assert_eq!(entries[0].account_name.as_deref(), Some("Grok Heavy"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_plain_line_as_session_token() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries("grok", "opaque-sso-token");
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("opaque-sso-token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_cookie_line_as_session_metadata() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
"i18nextLng=zh; cf_clearance=cf-1; sso-rw=rw-1; sso=sso-1; x-userid=user-1",
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].sso_rw_token.as_deref(), Some("rw-1"));
|
||||
assert_eq!(
|
||||
entries[0].cf_cookies.as_deref(),
|
||||
Some("i18nextLng=zh; cf_clearance=cf-1; x-userid=user-1")
|
||||
);
|
||||
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
|
||||
assert_eq!(entries[0].user_id.as_deref(), Some("user-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_cookie_object_as_session_metadata() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
r#"[{"cookie":"cf_clearance=cf-1; sso-rw=rw-1; sso=sso-1; x-userid=user-1","tier":"heavy"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].sso_rw_token.as_deref(), Some("rw-1"));
|
||||
assert_eq!(
|
||||
entries[0].cf_cookies.as_deref(),
|
||||
Some("cf_clearance=cf-1; x-userid=user-1")
|
||||
);
|
||||
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
|
||||
assert_eq!(entries[0].user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,7 +10,7 @@ use super::progress::{
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_provider_id;
|
||||
@@ -124,10 +124,6 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
"该 Provider 不是固定类型,无法使用 provider-oauth",
|
||||
));
|
||||
}
|
||||
if provider_type != "kiro" && admin_provider_oauth_template(&provider_type).is_none() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
|
||||
@@ -3,6 +3,8 @@ use axum::{
|
||||
body::Body,
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) fn attach_admin_provider_oauth_audit_response(
|
||||
response: Response<Body>,
|
||||
@@ -19,3 +21,79 @@ pub(super) fn attach_admin_provider_oauth_audit_response(
|
||||
};
|
||||
attach_admin_audit_response(response, event_name, action, target_type, &target_id)
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_oauth_key_name_from_auth_config(
|
||||
provider_type: &str,
|
||||
auth_config: &Map<String, Value>,
|
||||
batch_index: Option<usize>,
|
||||
) -> String {
|
||||
let provider_type = provider_type.trim();
|
||||
if let Some(email) = trimmed_auth_config_string(auth_config, "email") {
|
||||
return format!("{provider_type}_{email}");
|
||||
}
|
||||
if provider_type.eq_ignore_ascii_case("grok") {
|
||||
if let Some(user_id) = trimmed_auth_config_string(auth_config, "user_id") {
|
||||
return format!("grok_{user_id}");
|
||||
}
|
||||
}
|
||||
|
||||
let timestamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
match batch_index {
|
||||
Some(index) => format!("{provider_type}_{timestamp}_{index}"),
|
||||
None => format!("账号_{timestamp}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn trimmed_auth_config_string(auth_config: &Map<String, Value>, key: &str) -> Option<String> {
|
||||
auth_config
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::{json, Map};
|
||||
|
||||
#[test]
|
||||
fn grok_default_key_name_uses_full_user_id() {
|
||||
let mut auth_config = Map::new();
|
||||
auth_config.insert(
|
||||
"user_id".to_string(),
|
||||
json!("1619039a-0191-4e0a-a490-8f4ad21262c9"),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
|
||||
"grok_1619039a-0191-4e0a-a490-8f4ad21262c9"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_key_name_prefers_email_over_grok_user_id() {
|
||||
let mut auth_config = Map::new();
|
||||
auth_config.insert("email".to_string(), json!("[email protected]"));
|
||||
auth_config.insert("user_id".to_string(), json!("user-1"));
|
||||
|
||||
assert_eq!(
|
||||
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
|
||||
"[email protected]"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn batch_default_key_name_keeps_existing_timestamp_shape() {
|
||||
let auth_config = Map::new();
|
||||
let name = admin_provider_oauth_key_name_from_auth_config("codex", &auth_config, Some(3));
|
||||
|
||||
assert!(name.starts_with("codex_"));
|
||||
assert!(name.ends_with("_3"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,8 +14,9 @@ use super::super::state::{
|
||||
exchange_admin_provider_oauth_refresh_token, is_fixed_provider_type_for_provider_oauth,
|
||||
json_u64_value,
|
||||
};
|
||||
use super::helpers::admin_provider_oauth_key_name_from_auth_config;
|
||||
use super::token_import::{
|
||||
build_provider_access_token_import_auth_config, normalize_single_import_tokens,
|
||||
build_provider_access_token_import_auth_config, normalize_provider_import_tokens,
|
||||
provider_type_supports_access_token_import,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_import_provider_id;
|
||||
@@ -31,7 +32,6 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
struct AdminProviderOAuthSingleImportTokens {
|
||||
access_token: String,
|
||||
@@ -72,7 +72,8 @@ fn apply_single_import_hints(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
auth_config: &mut serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
if !provider_type_supports_access_token_import(provider_type) {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -110,19 +111,40 @@ fn apply_single_import_hints(
|
||||
&["user_id", "userId", "chatgpt_user_id", "chatgptUserId"][..],
|
||||
),
|
||||
("account_name", &["account_name", "accountName"][..]),
|
||||
("sso_rw_token", &["sso_rw_token", "ssoRwToken"][..]),
|
||||
(
|
||||
"cf_cookies",
|
||||
&["cf_cookies", "cfCookies", "cookie", "cookieHeader"][..],
|
||||
),
|
||||
("cf_clearance", &["cf_clearance", "cfClearance"][..]),
|
||||
("user_agent", &["user_agent", "userAgent"][..]),
|
||||
(
|
||||
"browser_profile",
|
||||
&[
|
||||
"browser_profile",
|
||||
"browserProfile",
|
||||
"browser",
|
||||
"impersonate",
|
||||
][..],
|
||||
),
|
||||
("pool_tier", &["pool_tier", "poolTier", "tier"][..]),
|
||||
] {
|
||||
let Some(value) = import_payload_string_any(payload, keys) else {
|
||||
continue;
|
||||
};
|
||||
auth_config
|
||||
.entry(target.to_string())
|
||||
.or_insert_with(|| json!(value));
|
||||
auth_config.entry(target.to_string()).or_insert_with(|| {
|
||||
if target == "plan_type" || target == "pool_tier" {
|
||||
json!(value.to_ascii_lowercase())
|
||||
} else {
|
||||
json!(value)
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async fn resolve_admin_provider_oauth_single_import_tokens(
|
||||
state: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
template: Option<AdminProviderOAuthTemplate>,
|
||||
provider_type: &str,
|
||||
refresh_token: Option<&str>,
|
||||
access_token: Option<&str>,
|
||||
@@ -133,6 +155,32 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let Some(template) = template else {
|
||||
if provider_type_supports_access_token_import(provider_type) {
|
||||
if let Some(access_token) = access_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
|
||||
provider_type,
|
||||
access_token,
|
||||
Some(refresh_token),
|
||||
imported_expires_at,
|
||||
Some("Provider 不支持 Refresh Token 交换,已回退为 Session Token 导入"),
|
||||
);
|
||||
return Ok(AdminProviderOAuthSingleImportTokens {
|
||||
access_token: access_token.to_string(),
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
}
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"该 Provider 不支持 Refresh Token 导入,请提供 sso_token 或 access_token",
|
||||
));
|
||||
};
|
||||
|
||||
let token_payload = match state
|
||||
.exchange_admin_provider_oauth_refresh_token(
|
||||
template,
|
||||
@@ -200,7 +248,7 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
|
||||
if !provider_type_supports_access_token_import(provider_type) {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Access Token 导入仅支持 Codex / ChatGPT Web Provider",
|
||||
"Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider",
|
||||
));
|
||||
}
|
||||
|
||||
@@ -248,18 +296,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
}
|
||||
};
|
||||
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
||||
let access_token_input = import_payload_string(&raw_payload, "access_token", "accessToken");
|
||||
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
|
||||
let (refresh_token_input, access_token_input) = normalize_single_import_tokens(
|
||||
refresh_token_input.as_deref(),
|
||||
access_token_input.as_deref(),
|
||||
let access_token_input = import_payload_string_any(
|
||||
&raw_payload,
|
||||
&["access_token", "accessToken", "sso_token", "ssoToken"],
|
||||
);
|
||||
if refresh_token_input.is_none() && access_token_input.is_none() {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token 或 Access Token 不能为空",
|
||||
));
|
||||
}
|
||||
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
|
||||
let name = raw_payload
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
@@ -285,6 +326,17 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
));
|
||||
};
|
||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||
let (refresh_token_input, access_token_input) = normalize_provider_import_tokens(
|
||||
&provider_type,
|
||||
refresh_token_input.as_deref(),
|
||||
access_token_input.as_deref(),
|
||||
);
|
||||
if refresh_token_input.is_none() && access_token_input.is_none() {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token、Access Token 或 sso_token 不能为空",
|
||||
));
|
||||
}
|
||||
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
@@ -297,9 +349,10 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
"Kiro 不支持单条 Refresh Token 导入,请使用批量导入或设备授权。",
|
||||
));
|
||||
}
|
||||
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
||||
let template = admin_provider_oauth_template(&provider_type);
|
||||
if template.is_none() && !provider_type_supports_access_token_import(&provider_type) {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
};
|
||||
}
|
||||
let endpoint_resolution =
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
|
||||
let endpoints = endpoint_resolution.endpoints;
|
||||
@@ -380,25 +433,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let name = name
|
||||
.or_else(|| {
|
||||
auth_config
|
||||
.get("email")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
.unwrap_or_else(|| {
|
||||
format!(
|
||||
"账号_{}",
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0)
|
||||
)
|
||||
});
|
||||
let name = name.unwrap_or_else(|| {
|
||||
admin_provider_oauth_key_name_from_auth_config(&provider_type, &auth_config, None)
|
||||
});
|
||||
match state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&provider_id,
|
||||
|
||||
@@ -81,6 +81,28 @@ pub(super) fn normalize_single_import_tokens(
|
||||
(refresh_token, access_token)
|
||||
}
|
||||
|
||||
pub(super) fn normalize_provider_import_tokens(
|
||||
provider_type: &str,
|
||||
refresh_token: Option<&str>,
|
||||
access_token: Option<&str>,
|
||||
) -> (Option<String>, Option<String>) {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
let refresh_token = refresh_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let access_token = access_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
|
||||
if provider_type == "grok" {
|
||||
return (None, access_token.or(refresh_token));
|
||||
}
|
||||
|
||||
normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref())
|
||||
}
|
||||
|
||||
pub(super) fn import_tokens_from_raw_token(token: &str) -> (Option<String>, Option<String>) {
|
||||
if looks_like_access_token(token) {
|
||||
(None, Some(token.trim().to_string()))
|
||||
@@ -98,7 +120,7 @@ pub(super) fn decode_access_token_expires_at(access_token: &str) -> Option<u64>
|
||||
pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool {
|
||||
matches!(
|
||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||
"codex" | "chatgpt_web"
|
||||
"codex" | "chatgpt_web" | "grok"
|
||||
)
|
||||
}
|
||||
|
||||
@@ -123,6 +145,11 @@ pub(super) fn build_provider_access_token_import_auth_config(
|
||||
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
|
||||
}
|
||||
|
||||
if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
auth_config.insert("sso_token".to_string(), json!(access_token));
|
||||
auth_config.insert("auth_method".to_string(), json!("sso_token"));
|
||||
}
|
||||
|
||||
auth_config.insert(
|
||||
"access_token_import_temporary".to_string(),
|
||||
json!(refresh_token.is_none()),
|
||||
@@ -149,7 +176,7 @@ pub(super) fn build_provider_access_token_import_auth_config(
|
||||
mod tests {
|
||||
use super::{
|
||||
build_provider_access_token_import_auth_config, decode_access_token_expires_at,
|
||||
looks_like_access_token, normalize_single_import_tokens,
|
||||
looks_like_access_token, normalize_provider_import_tokens, normalize_single_import_tokens,
|
||||
};
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||
use serde_json::json;
|
||||
@@ -250,4 +277,34 @@ mod tests {
|
||||
Some(&json!(true))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_grok_import_treats_opaque_session_as_access_token() {
|
||||
let (refresh_token, access_token) =
|
||||
normalize_provider_import_tokens("grok", Some("sso_session_token"), None);
|
||||
assert!(refresh_token.is_none());
|
||||
assert_eq!(access_token.as_deref(), Some("sso_session_token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_grok_auth_config_from_session_token() {
|
||||
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
|
||||
"grok",
|
||||
"sso_session_token",
|
||||
None,
|
||||
Some(2_200_000_000),
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(expires_at, Some(2_200_000_000));
|
||||
assert_eq!(
|
||||
auth_config.get("sso_token"),
|
||||
Some(&json!("sso_session_token"))
|
||||
);
|
||||
assert_eq!(auth_config.get("auth_method"), Some(&json!("sso_token")));
|
||||
assert_eq!(
|
||||
auth_config.get("expires_at"),
|
||||
Some(&json!(2_200_000_000u64))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,8 +8,10 @@ use crate::GatewayError;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_provider_transport::provider_types::provider_type_is_fixed;
|
||||
use serde_json::json;
|
||||
use aether_provider_transport::{
|
||||
grok_browser_transport_fingerprint_from_auth_config, provider_types::provider_type_is_fixed,
|
||||
};
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
@@ -92,6 +94,16 @@ pub(crate) fn build_provider_oauth_auth_config_from_token_payload(
|
||||
(auth_config, access_token, refresh_token, expires_at)
|
||||
}
|
||||
|
||||
fn grok_oauth_catalog_key_fingerprint(
|
||||
provider_type: &str,
|
||||
auth_config: &Map<String, Value>,
|
||||
) -> Option<Value> {
|
||||
if !provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
return None;
|
||||
}
|
||||
grok_browser_transport_fingerprint_from_auth_config(auth_config)
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_oauth_catalog_key(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
@@ -136,7 +148,7 @@ pub(crate) async fn create_provider_oauth_catalog_key(
|
||||
None,
|
||||
expires_at_unix_secs,
|
||||
proxy,
|
||||
None,
|
||||
grok_oauth_catalog_key_fingerprint(provider_type, auth_config),
|
||||
)
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
record.internal_priority = 50;
|
||||
@@ -193,6 +205,9 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
|
||||
updated.expires_at_unix_secs = expires_at_unix_secs;
|
||||
updated.oauth_invalid_at_unix_secs = None;
|
||||
updated.oauth_invalid_reason = None;
|
||||
if updated.fingerprint.is_none() {
|
||||
updated.fingerprint = grok_oauth_catalog_key_fingerprint(provider_type, auth_config);
|
||||
}
|
||||
updated.health_by_format = Some(json!({}));
|
||||
updated.circuit_breaker_by_format = Some(json!({}));
|
||||
updated.error_count = Some(0);
|
||||
@@ -223,7 +238,9 @@ fn provider_oauth_catalog_key_api_formats(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::provider_oauth_token_payload_expires_at_unix_secs;
|
||||
use super::{
|
||||
grok_oauth_catalog_key_fingerprint, provider_oauth_token_payload_expires_at_unix_secs,
|
||||
};
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||
use serde_json::json;
|
||||
|
||||
@@ -273,4 +290,60 @@ mod tests {
|
||||
Some(2_000_000_000)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_oauth_catalog_key_fingerprint_uses_browser_wreq_profile() {
|
||||
let auth_config = json!({
|
||||
"sso_token": "abc",
|
||||
"browser_profile": "chrome-137",
|
||||
});
|
||||
let auth_config = auth_config.as_object().expect("object");
|
||||
|
||||
let fingerprint = grok_oauth_catalog_key_fingerprint("grok", auth_config)
|
||||
.expect("fingerprint should resolve");
|
||||
|
||||
assert_eq!(
|
||||
fingerprint["transport_profile"]["profile_id"],
|
||||
json!("chrome137")
|
||||
);
|
||||
assert_eq!(
|
||||
fingerprint["transport_profile"]["backend"],
|
||||
json!("browser_wreq")
|
||||
);
|
||||
assert_eq!(
|
||||
fingerprint["transport_profile"]["extra"]["browser_profile"],
|
||||
json!("chrome137")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_oauth_catalog_key_fingerprint_infers_profile_from_user_agent() {
|
||||
let auth_config = json!({
|
||||
"sso_token": "abc",
|
||||
"user_agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/137.0.0.0 Safari/537.36",
|
||||
});
|
||||
let auth_config = auth_config.as_object().expect("object");
|
||||
|
||||
let fingerprint = grok_oauth_catalog_key_fingerprint("grok", auth_config)
|
||||
.expect("fingerprint should resolve");
|
||||
|
||||
assert_eq!(
|
||||
fingerprint["transport_profile"]["profile_id"],
|
||||
json!("chrome137")
|
||||
);
|
||||
assert_eq!(
|
||||
fingerprint["transport_profile"]["extra"]["browser_profile"],
|
||||
json!("chrome137")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_oauth_catalog_key_fingerprint_ignores_non_grok_providers() {
|
||||
let auth_config = json!({
|
||||
"browser_profile": "chrome136",
|
||||
});
|
||||
let auth_config = auth_config.as_object().expect("object");
|
||||
|
||||
assert!(grok_oauth_catalog_key_fingerprint("openai", auth_config).is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ use std::pin::Pin;
|
||||
use super::antigravity::refresh_antigravity_provider_quota_locally;
|
||||
use super::chatgpt_web::refresh_chatgpt_web_provider_quota_locally;
|
||||
use super::codex::refresh_codex_provider_quota_locally;
|
||||
use super::grok::refresh_grok_provider_quota_locally;
|
||||
use super::kiro::refresh_kiro_provider_quota_locally;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
@@ -33,6 +34,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] =
|
||||
refresh_chatgpt_web_provider_quota_locally_boxed,
|
||||
),
|
||||
("codex", refresh_codex_provider_quota_locally_boxed),
|
||||
("grok", refresh_grok_provider_quota_locally_boxed),
|
||||
("kiro", refresh_kiro_provider_quota_locally_boxed),
|
||||
];
|
||||
|
||||
@@ -117,3 +119,19 @@ fn refresh_kiro_provider_quota_locally_boxed<'a>(
|
||||
proxy_override,
|
||||
))
|
||||
}
|
||||
|
||||
fn refresh_grok_provider_quota_locally_boxed<'a>(
|
||||
state: &'a AdminAppState<'a>,
|
||||
provider: &'a StoredProviderCatalogProvider,
|
||||
endpoint: &'a StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> ProviderQuotaRefreshFuture<'a> {
|
||||
Box::pin(refresh_grok_provider_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
keys,
|
||||
proxy_override,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,826 @@
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, extract_execution_error_message,
|
||||
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionResult, ProxySnapshot, RequestBody, ResolvedTransportProfile,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_provider_pool::{
|
||||
grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier,
|
||||
};
|
||||
use aether_provider_transport::grok_browser_profile_metadata_from_resolved_transport_profile;
|
||||
use base64::Engine as _;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
const GROK_DEFAULT_BASE_URL: &str = "https://grok.com";
|
||||
const GROK_RATE_LIMITS_PATH: &str = "/rest/rate-limits";
|
||||
const GROK_STATSIG_ID: &str = "ZTpUeXBlRXJyb3I6IENhbm5vdCByZWFkIHByb3BlcnRpZXMgb2YgdW5kZWZpbmVkIChyZWFkaW5nICdjaGlsZE5vZGVzJyk=";
|
||||
|
||||
fn grok_base_url(endpoint: &StoredProviderCatalogEndpoint) -> String {
|
||||
let base_url = endpoint.base_url.trim().trim_end_matches('/');
|
||||
if base_url.is_empty() {
|
||||
GROK_DEFAULT_BASE_URL.to_string()
|
||||
} else {
|
||||
base_url.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
fn grok_auth_config(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
) -> Option<serde_json::Value> {
|
||||
transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| serde_json::from_str::<serde_json::Value>(value).ok())
|
||||
}
|
||||
|
||||
fn grok_auth_string(auth_config: Option<&serde_json::Value>, fields: &[&str]) -> Option<String> {
|
||||
let object = auth_config.and_then(serde_json::Value::as_object)?;
|
||||
fields.iter().find_map(|field| {
|
||||
object
|
||||
.get(*field)
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
fn build_grok_quota_headers(
|
||||
auth_config: Option<&serde_json::Value>,
|
||||
transport_profile: Option<&ResolvedTransportProfile>,
|
||||
base_url: &str,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
let cookie = build_grok_quota_cookie(auth_config).unwrap_or_default();
|
||||
let browser_profile =
|
||||
grok_browser_profile_metadata_from_resolved_transport_profile(transport_profile?)?;
|
||||
Some(BTreeMap::from([
|
||||
("accept".to_string(), "*/*".to_string()),
|
||||
(
|
||||
"accept-language".to_string(),
|
||||
"zh-CN,zh;q=0.9,en;q=0.8".to_string(),
|
||||
),
|
||||
(
|
||||
"baggage".to_string(),
|
||||
"sentry-environment=production,sentry-release=d6add6fb0460641fd482d767a335ef72b9b6abb8,sentry-public_key=b311e0f2690c81f25e2c4cf6d4f7ce1c".to_string(),
|
||||
),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
("origin".to_string(), base_url.to_string()),
|
||||
("priority".to_string(), "u=1, i".to_string()),
|
||||
("referer".to_string(), format!("{base_url}/")),
|
||||
("sec-ch-ua".to_string(), browser_profile.sec_ch_ua),
|
||||
("sec-ch-ua-mobile".to_string(), "?0".to_string()),
|
||||
("sec-ch-ua-model".to_string(), String::new()),
|
||||
(
|
||||
"sec-ch-ua-platform".to_string(),
|
||||
browser_profile.sec_ch_ua_platform,
|
||||
),
|
||||
("sec-fetch-dest".to_string(), "empty".to_string()),
|
||||
("sec-fetch-mode".to_string(), "cors".to_string()),
|
||||
("sec-fetch-site".to_string(), "same-origin".to_string()),
|
||||
("user-agent".to_string(), browser_profile.user_agent),
|
||||
("cookie".to_string(), cookie),
|
||||
("x-statsig-id".to_string(), GROK_STATSIG_ID.to_string()),
|
||||
("x-xai-request-id".to_string(), Uuid::new_v4().to_string()),
|
||||
]))
|
||||
}
|
||||
|
||||
fn build_grok_quota_cookie(auth_config: Option<&serde_json::Value>) -> Option<String> {
|
||||
let token = grok_auth_string(auth_config, &["sso_token", "access_token", "token"])?;
|
||||
let token = strip_cookie_prefix(token.trim(), "sso=");
|
||||
if token.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let sso_rw = grok_auth_string(auth_config, &["sso_rw_token", "ssoRwToken"])
|
||||
.map(|value| strip_cookie_prefix(value.trim(), "sso-rw="))
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_else(|| token.clone());
|
||||
|
||||
let mut parts = vec![format!("sso={token}"), format!("sso-rw={sso_rw}")];
|
||||
if let Some(extra_cookies) =
|
||||
grok_auth_string(auth_config, &["cf_cookies", "cfCookies", "cookie"])
|
||||
.and_then(|value| normalize_grok_extra_cookies(value.as_str()))
|
||||
{
|
||||
parts.push(extra_cookies);
|
||||
}
|
||||
let cf_clearance = grok_auth_string(auth_config, &["cf_clearance", "cfClearance"])
|
||||
.map(|value| strip_cookie_prefix(value.trim(), "cf_clearance="))
|
||||
.filter(|value| !value.is_empty());
|
||||
if let Some(cf_clearance) = cf_clearance {
|
||||
if !parts.iter().any(|part| part.contains("cf_clearance=")) {
|
||||
parts.push(format!("cf_clearance={cf_clearance}"));
|
||||
}
|
||||
}
|
||||
Some(parts.join("; "))
|
||||
}
|
||||
|
||||
fn strip_cookie_prefix(value: &str, prefix: &str) -> String {
|
||||
value
|
||||
.strip_prefix(prefix)
|
||||
.map(str::trim)
|
||||
.unwrap_or(value)
|
||||
.to_string()
|
||||
}
|
||||
|
||||
fn normalize_grok_extra_cookies(value: &str) -> Option<String> {
|
||||
let parts = value
|
||||
.trim()
|
||||
.trim_matches(';')
|
||||
.split(';')
|
||||
.filter_map(|segment| {
|
||||
let (name, value) = segment.trim().split_once('=')?;
|
||||
let name = name.trim();
|
||||
let value = value.trim();
|
||||
if name.is_empty()
|
||||
|| value.is_empty()
|
||||
|| name.eq_ignore_ascii_case("sso")
|
||||
|| name.eq_ignore_ascii_case("sso-rw")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(format!("{name}={value}"))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
(!parts.is_empty()).then(|| parts.join("; "))
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
struct GrokRateLimitSnapshot {
|
||||
remaining: f64,
|
||||
total: f64,
|
||||
window_seconds: u64,
|
||||
wait_time_seconds: Option<u64>,
|
||||
}
|
||||
|
||||
impl GrokRateLimitSnapshot {
|
||||
fn reset_after_seconds(self) -> u64 {
|
||||
self.wait_time_seconds.unwrap_or(self.window_seconds)
|
||||
}
|
||||
|
||||
fn reset_at_source(self) -> &'static str {
|
||||
if self.wait_time_seconds.is_some() {
|
||||
"grok_rate_limits_wait_time"
|
||||
} else {
|
||||
"grok_rate_limits_window"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_grok_rate_limits(body: &serde_json::Value) -> Option<GrokRateLimitSnapshot> {
|
||||
let remaining = body
|
||||
.get("remainingQueries")
|
||||
.and_then(serde_json::Value::as_f64)?;
|
||||
let total = body
|
||||
.get("totalQueries")
|
||||
.and_then(serde_json::Value::as_f64)
|
||||
.unwrap_or(remaining.max(0.0));
|
||||
let window_seconds = body
|
||||
.get("windowSizeSeconds")
|
||||
.and_then(serde_json::Value::as_u64)
|
||||
.unwrap_or(72_000);
|
||||
let wait_time_seconds = body
|
||||
.get("waitTimeSeconds")
|
||||
.and_then(serde_json::Value::as_u64);
|
||||
Some(GrokRateLimitSnapshot {
|
||||
remaining,
|
||||
total,
|
||||
window_seconds,
|
||||
wait_time_seconds,
|
||||
})
|
||||
}
|
||||
|
||||
fn grok_pool_tier_hint_for_refresh(
|
||||
key: &StoredProviderCatalogKey,
|
||||
auth_config: Option<&serde_json::Value>,
|
||||
) -> Option<&'static str> {
|
||||
key.status_snapshot
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|snapshot| snapshot.get("quota"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(grok_pool_tier_from_quota_bucket)
|
||||
.or_else(|| {
|
||||
key.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|metadata| metadata.get("grok"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(grok_pool_tier_from_quota_bucket)
|
||||
})
|
||||
.or_else(|| {
|
||||
auth_config
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(grok_pool_tier_from_quota_bucket)
|
||||
})
|
||||
}
|
||||
|
||||
async fn execute_grok_quota_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
body: serde_json::Value,
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||
let proxy = match proxy_override {
|
||||
Some(proxy) => Some(proxy.clone()),
|
||||
None => {
|
||||
state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let transport_profile = state.resolve_transport_profile(transport);
|
||||
let base_url = grok_base_url(endpoint);
|
||||
let headers = build_grok_quota_headers(
|
||||
grok_auth_config(transport).as_ref(),
|
||||
transport_profile.as_ref(),
|
||||
&base_url,
|
||||
)
|
||||
.ok_or_else(|| {
|
||||
GatewayError::Internal("unsupported Grok browser transport profile".to_string())
|
||||
})?;
|
||||
let plan = ExecutionPlan {
|
||||
request_id: format!("grok-quota:{}", transport.key.id),
|
||||
candidate_id: None,
|
||||
provider_name: Some("grok".to_string()),
|
||||
provider_id: transport.provider.id.clone(),
|
||||
endpoint_id: transport.endpoint.id.clone(),
|
||||
key_id: transport.key.id.clone(),
|
||||
method: "POST".to_string(),
|
||||
url: format!(
|
||||
"{}/{}",
|
||||
base_url,
|
||||
GROK_RATE_LIMITS_PATH.trim_start_matches('/')
|
||||
),
|
||||
headers,
|
||||
content_type: Some("application/json".to_string()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(body),
|
||||
stream: false,
|
||||
client_api_format: "openai:responses".to_string(),
|
||||
provider_api_format: "grok:rate_limits".to_string(),
|
||||
model_name: Some("grok-quota".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts,
|
||||
};
|
||||
|
||||
execute_provider_quota_plan(state, transport, plan, "grok").await
|
||||
}
|
||||
|
||||
fn grok_quota_error_detail(result: &ExecutionResult) -> Option<String> {
|
||||
extract_execution_error_message(result).or_else(|| {
|
||||
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
|
||||
let decoded = base64::engine::general_purpose::STANDARD
|
||||
.decode(body)
|
||||
.ok()?;
|
||||
let text = String::from_utf8_lossy(&decoded).trim().to_string();
|
||||
(!text.is_empty()).then_some(text)
|
||||
})
|
||||
}
|
||||
|
||||
fn grok_is_cloudflare_challenge(message: &str) -> bool {
|
||||
let lowered = message.to_ascii_lowercase();
|
||||
lowered.contains("cloudflare")
|
||||
|| lowered.contains("just a moment")
|
||||
|| lowered.contains("__cf_chl")
|
||||
|| lowered.contains("cf-ray")
|
||||
}
|
||||
|
||||
fn grok_quota_invalid_reason(status_code: u16, upstream_message: Option<&str>) -> String {
|
||||
let message = upstream_message.unwrap_or_default().trim();
|
||||
if status_code == 403 && grok_is_cloudflare_challenge(message) {
|
||||
return format!(
|
||||
"{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败,请重新从同一浏览器复制最新 Cookie 和 User-Agent,或配置可通过 Cloudflare 的代理运行时"
|
||||
);
|
||||
}
|
||||
let detail = if message.is_empty() {
|
||||
match status_code {
|
||||
401 => "Grok Token 无效或已过期",
|
||||
403 => "Grok 账户访问受限",
|
||||
_ => "Grok 请求失败",
|
||||
}
|
||||
} else {
|
||||
message
|
||||
};
|
||||
match status_code {
|
||||
401 => format!("{OAUTH_EXPIRED_PREFIX}{detail}"),
|
||||
403 => format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}"),
|
||||
_ => detail.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn grok_quota_result_message(reason: &str) -> String {
|
||||
for prefix in [
|
||||
OAUTH_REFRESH_FAILED_PREFIX,
|
||||
OAUTH_EXPIRED_PREFIX,
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX,
|
||||
] {
|
||||
if let Some(message) = reason.strip_prefix(prefix) {
|
||||
return message.trim().to_string();
|
||||
}
|
||||
}
|
||||
reason.trim().to_string()
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_grok_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> Result<Option<serde_json::Value>, GatewayError> {
|
||||
let mut results = Vec::new();
|
||||
let mut success_count = 0usize;
|
||||
let mut failed_count = 0usize;
|
||||
|
||||
for key in keys {
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
{
|
||||
Some(transport) => transport,
|
||||
None => {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "Provider transport snapshot unavailable",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if grok_auth_config(&transport).is_none() {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "缺少 Grok 账号会话信息,请先导入 Token",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
let auth_config = grok_auth_config(&transport);
|
||||
let mut quota_by_model = serde_json::Map::new();
|
||||
let mut refreshed = false;
|
||||
let mut invalid_reason = None::<String>;
|
||||
let mut invalid_at = key.oauth_invalid_at_unix_secs;
|
||||
let mut last_status_code = None::<u16>;
|
||||
let mut last_error_message = None::<String>;
|
||||
let mut metadata_update = serde_json::Map::new();
|
||||
let base_url = grok_base_url(endpoint);
|
||||
|
||||
let supported_windows = grok_supported_quota_windows_for_tier(
|
||||
grok_pool_tier_hint_for_refresh(&key, auth_config.as_ref()),
|
||||
);
|
||||
for (quota_key, mode_name) in supported_windows.iter().copied() {
|
||||
let result = match execute_grok_quota_plan(
|
||||
state,
|
||||
&transport,
|
||||
endpoint,
|
||||
json!({ "modelName": mode_name }),
|
||||
proxy_override.as_ref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderQuotaExecutionOutcome::Response(result) => result,
|
||||
ProviderQuotaExecutionOutcome::Failure(detail) => {
|
||||
last_error_message = Some(format!("rate-limits 请求执行失败: {detail}"));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
last_status_code = Some(result.status_code);
|
||||
|
||||
if result.status_code == 200 {
|
||||
if let Some(body_json) = result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.as_ref())
|
||||
{
|
||||
if let Some(rate_limit) = parse_grok_rate_limits(body_json) {
|
||||
refreshed = true;
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let reset_after_seconds = rate_limit.reset_after_seconds();
|
||||
let reset_at = now_unix_secs.saturating_add(reset_after_seconds);
|
||||
quota_by_model.insert(
|
||||
(*quota_key).to_string(),
|
||||
json!({
|
||||
"display_name": *mode_name,
|
||||
"remaining_fraction": if rate_limit.total > 0.0 { Some((rate_limit.remaining / rate_limit.total).clamp(0.0, 1.0)) } else { None::<f64> },
|
||||
"used_percent": if rate_limit.total > 0.0 { Some(((rate_limit.total - rate_limit.remaining).max(0.0) / rate_limit.total * 100.0).clamp(0.0, 100.0)) } else { None::<f64> },
|
||||
"remaining": rate_limit.remaining,
|
||||
"total": rate_limit.total,
|
||||
"window_seconds": rate_limit.window_seconds,
|
||||
"wait_time_seconds": rate_limit.wait_time_seconds,
|
||||
"reset_after_seconds": reset_after_seconds,
|
||||
"reset_at": reset_at,
|
||||
"next_reset_at": reset_at,
|
||||
"reset_at_source": rate_limit.reset_at_source(),
|
||||
"is_exhausted": rate_limit.remaining <= 0.0,
|
||||
}),
|
||||
);
|
||||
} else {
|
||||
last_error_message = Some(
|
||||
"Grok rate-limits 未返回 remainingQueries/totalQueries".to_string(),
|
||||
);
|
||||
}
|
||||
} else {
|
||||
last_error_message = Some("Grok rate-limits 未返回 JSON 数据".to_string());
|
||||
}
|
||||
} else if matches!(result.status_code, 401 | 403) {
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
invalid_at = Some(now_unix_secs);
|
||||
let error_detail = grok_quota_error_detail(&result);
|
||||
invalid_reason = Some(grok_quota_invalid_reason(
|
||||
result.status_code,
|
||||
error_detail.as_deref(),
|
||||
));
|
||||
last_error_message = invalid_reason.as_deref().map(grok_quota_result_message);
|
||||
} else {
|
||||
let error_detail =
|
||||
grok_quota_error_detail(&result).unwrap_or_else(|| "Grok 请求失败".to_string());
|
||||
last_error_message = Some(format!(
|
||||
"Grok rate-limits 请求失败({}): {error_detail}",
|
||||
result.status_code
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
if refreshed {
|
||||
if let Some(pool_tier) = grok_pool_tier_from_quota_bucket("a_by_model)
|
||||
.or_else(|| grok_pool_tier_hint_for_refresh(&key, auth_config.as_ref()))
|
||||
{
|
||||
let pool_tier_value = json!(pool_tier);
|
||||
metadata_update.insert("pool_tier".to_string(), pool_tier_value.clone());
|
||||
metadata_update
|
||||
.entry("plan_type".to_string())
|
||||
.or_insert(pool_tier_value);
|
||||
}
|
||||
metadata_update.insert(
|
||||
"updated_at".to_string(),
|
||||
json!(SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0)),
|
||||
);
|
||||
metadata_update.insert("base_url".to_string(), json!(base_url));
|
||||
metadata_update.insert("quota_by_model".to_string(), json!(quota_by_model));
|
||||
}
|
||||
|
||||
let metadata_update_value = if metadata_update.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(serde_json::Value::Object({
|
||||
let mut map = serde_json::Map::new();
|
||||
map.insert(
|
||||
"grok".to_string(),
|
||||
serde_json::Value::Object(metadata_update.clone()),
|
||||
);
|
||||
map
|
||||
}))
|
||||
};
|
||||
|
||||
if !persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update_value.as_ref(),
|
||||
invalid_at,
|
||||
invalid_reason,
|
||||
None,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "Key 状态写入失败",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
if refreshed {
|
||||
success_count += 1;
|
||||
} else {
|
||||
failed_count += 1;
|
||||
}
|
||||
|
||||
let mut payload = serde_json::Map::new();
|
||||
payload.insert("key_id".to_string(), json!(key.id));
|
||||
payload.insert("key_name".to_string(), json!(key.name));
|
||||
payload.insert(
|
||||
"status".to_string(),
|
||||
json!(if refreshed { "success" } else { "error" }),
|
||||
);
|
||||
if let Some(metadata) = metadata_update.get("quota_by_model").cloned() {
|
||||
payload.insert("metadata".to_string(), metadata);
|
||||
}
|
||||
if let Some(quota_snapshot) = build_quota_snapshot_payload(
|
||||
"grok",
|
||||
key.status_snapshot.as_ref(),
|
||||
metadata_update_value.as_ref(),
|
||||
) {
|
||||
payload.insert("quota_snapshot".to_string(), quota_snapshot);
|
||||
}
|
||||
if !refreshed {
|
||||
payload.insert(
|
||||
"message".to_string(),
|
||||
json!(last_error_message.unwrap_or_else(|| {
|
||||
"Grok rate-limits 未返回可用配额数据".to_string()
|
||||
})),
|
||||
);
|
||||
if let Some(status_code) = last_status_code {
|
||||
payload.insert("status_code".to_string(), json!(status_code));
|
||||
}
|
||||
}
|
||||
results.push(serde_json::Value::Object(payload));
|
||||
}
|
||||
|
||||
Ok(Some(json!({
|
||||
"success": success_count,
|
||||
"failed": failed_count,
|
||||
"total": success_count + failed_count,
|
||||
"results": results,
|
||||
"message": format!("已处理 {} 个 Key", success_count + failed_count),
|
||||
"auto_removed": 0,
|
||||
})))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_grok_quota_cookie, build_grok_quota_headers, grok_pool_tier_hint_for_refresh,
|
||||
grok_quota_error_detail, grok_quota_invalid_reason, grok_quota_result_message,
|
||||
parse_grok_rate_limits,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::OAUTH_REFRESH_FAILED_PREFIX;
|
||||
use aether_contracts::{ExecutionResult, ResponseBody};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use base64::Engine as _;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
fn sample_key(
|
||||
status_snapshot: Option<serde_json::Value>,
|
||||
upstream_metadata: Option<serde_json::Value>,
|
||||
) -> StoredProviderCatalogKey {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"key-1".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
key.status_snapshot = status_snapshot;
|
||||
key.upstream_metadata = upstream_metadata;
|
||||
key
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_cookie_preserves_grok_session_and_clearance() {
|
||||
let auth_config = json!({
|
||||
"sso_token": "sso=abc",
|
||||
"sso_rw_token": "sso-rw=rw",
|
||||
"cf_clearance": "cf"
|
||||
});
|
||||
|
||||
let cookie = build_grok_quota_cookie(Some(&auth_config)).expect("cookie should build");
|
||||
|
||||
assert_eq!(cookie, "sso=abc; sso-rw=rw; cf_clearance=cf");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_cookie_removes_duplicate_session_cookies_from_cf_profile() {
|
||||
let auth_config = json!({
|
||||
"sso_token": "abc",
|
||||
"sso_rw_token": "rw",
|
||||
"cf_cookies": "i18nextLng=zh; sso=ignored; sso-rw=ignored-rw; cf_clearance=cf"
|
||||
});
|
||||
|
||||
let cookie = build_grok_quota_cookie(Some(&auth_config)).expect("cookie should build");
|
||||
|
||||
assert_eq!(cookie, "sso=abc; sso-rw=rw; i18nextLng=zh; cf_clearance=cf");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_headers_use_resolved_transport_profile_user_agent() {
|
||||
let auth_config = json!({
|
||||
"sso_token": "abc",
|
||||
"user_agent": "Mozilla/5.0 custom"
|
||||
});
|
||||
let transport_profile = aether_provider_transport::grok_browser_resolved_transport_profile(
|
||||
Some("chrome137"),
|
||||
"test",
|
||||
)
|
||||
.expect("profile should resolve");
|
||||
|
||||
let headers = build_grok_quota_headers(
|
||||
Some(&auth_config),
|
||||
Some(&transport_profile),
|
||||
"https://grok.com",
|
||||
)
|
||||
.expect("headers should build");
|
||||
|
||||
assert!(headers
|
||||
.get("user-agent")
|
||||
.is_some_and(|value| value.contains("Chrome/137.0.0.0")));
|
||||
assert_eq!(
|
||||
headers.get("sec-ch-ua"),
|
||||
Some(
|
||||
&r#""Google Chrome";v="137", "Chromium";v="137", "Not(A:Brand";v="24""#.to_string()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_headers_default_to_chrome136_clearance_profile() {
|
||||
let auth_config = json!({
|
||||
"sso_token": "abc"
|
||||
});
|
||||
|
||||
let transport_profile =
|
||||
aether_provider_transport::grok_browser_resolved_transport_profile(None, "test")
|
||||
.expect("profile should resolve");
|
||||
let headers = build_grok_quota_headers(
|
||||
Some(&auth_config),
|
||||
Some(&transport_profile),
|
||||
"https://grok.com",
|
||||
)
|
||||
.expect("headers should build");
|
||||
|
||||
assert!(headers
|
||||
.get("user-agent")
|
||||
.is_some_and(|value| value.contains("Chrome/136.0.0.0")));
|
||||
assert_eq!(
|
||||
headers.get("sec-ch-ua"),
|
||||
Some(
|
||||
&r#""Google Chrome";v="136", "Chromium";v="136", "Not(A:Brand";v="24""#.to_string()
|
||||
)
|
||||
);
|
||||
assert_eq!(
|
||||
headers.get("sec-ch-ua-platform"),
|
||||
Some(&r#""macOS""#.to_string())
|
||||
);
|
||||
assert!(headers.contains_key("x-statsig-id"));
|
||||
assert!(headers.contains_key("x-xai-request-id"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_headers_do_not_mark_rate_limits_as_grok_app_chat_runtime() {
|
||||
let auth_config = json!({
|
||||
"sso_token": "abc"
|
||||
});
|
||||
|
||||
let transport_profile =
|
||||
aether_provider_transport::grok_browser_resolved_transport_profile(None, "test")
|
||||
.expect("profile should resolve");
|
||||
let headers = build_grok_quota_headers(
|
||||
Some(&auth_config),
|
||||
Some(&transport_profile),
|
||||
"https://grok.com",
|
||||
)
|
||||
.expect("headers should build");
|
||||
|
||||
assert!(!headers.contains_key(aether_provider_transport::GROK_INTERNAL_HEADER));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_wait_time_seconds_as_authoritative_reset_delay() {
|
||||
let body = json!({
|
||||
"windowSizeSeconds": 86_400,
|
||||
"remainingQueries": 0,
|
||||
"waitTimeSeconds": 12_648,
|
||||
"totalQueries": 30,
|
||||
"lowEffortRateLimits": null,
|
||||
"highEffortRateLimits": null
|
||||
});
|
||||
|
||||
let rate_limits = parse_grok_rate_limits(&body).expect("rate limits should parse");
|
||||
|
||||
assert_eq!(rate_limits.remaining, 0.0);
|
||||
assert_eq!(rate_limits.total, 30.0);
|
||||
assert_eq!(rate_limits.window_seconds, 86_400);
|
||||
assert_eq!(rate_limits.wait_time_seconds, Some(12_648));
|
||||
assert_eq!(rate_limits.reset_after_seconds(), 12_648);
|
||||
assert_eq!(rate_limits.reset_at_source(), "grok_rate_limits_wait_time");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_rate_limits_falls_back_to_window_when_wait_time_is_absent() {
|
||||
let body = json!({
|
||||
"windowSizeSeconds": 86_400,
|
||||
"remainingQueries": 12,
|
||||
"totalQueries": 30
|
||||
});
|
||||
|
||||
let rate_limits = parse_grok_rate_limits(&body).expect("rate limits should parse");
|
||||
|
||||
assert_eq!(rate_limits.remaining, 12.0);
|
||||
assert_eq!(rate_limits.total, 30.0);
|
||||
assert_eq!(rate_limits.window_seconds, 86_400);
|
||||
assert_eq!(rate_limits.wait_time_seconds, None);
|
||||
assert_eq!(rate_limits.reset_after_seconds(), 86_400);
|
||||
assert_eq!(rate_limits.reset_at_source(), "grok_rate_limits_window");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_grok_pool_tier_from_live_quota_totals() {
|
||||
let key = sample_key(
|
||||
Some(json!({
|
||||
"quota": {
|
||||
"pool_tier": "heavy"
|
||||
}
|
||||
})),
|
||||
None,
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
grok_pool_tier_hint_for_refresh(&key, Some(&json!({}))),
|
||||
Some("heavy")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn infers_basic_grok_pool_tier_from_fast_quota_when_auto_is_absent() {
|
||||
let key = sample_key(
|
||||
None,
|
||||
Some(json!({
|
||||
"grok": {
|
||||
"plan_type": "basic"
|
||||
}
|
||||
})),
|
||||
);
|
||||
|
||||
assert_eq!(grok_pool_tier_hint_for_refresh(&key, None), Some("basic"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cloudflare_challenge_403_is_not_account_block() {
|
||||
let body = "<!DOCTYPE html><html><head><title>Just a moment...</title></head><body>Cloudflare</body></html>";
|
||||
let result = ExecutionResult {
|
||||
request_id: "grok-quota:test".to_string(),
|
||||
candidate_id: None,
|
||||
status_code: 403,
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
};
|
||||
|
||||
let detail = grok_quota_error_detail(&result).expect("html body should be decoded");
|
||||
let reason = grok_quota_invalid_reason(result.status_code, Some(&detail));
|
||||
|
||||
assert!(reason.starts_with("[REFRESH_FAILED] "));
|
||||
assert!(!reason.starts_with("[ACCOUNT_BLOCK] "));
|
||||
assert!(reason.contains("Cloudflare"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_result_message_removes_status_prefix() {
|
||||
let reason = format!("{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败");
|
||||
|
||||
assert_eq!(
|
||||
grok_quota_result_message(&reason),
|
||||
"Grok Cloudflare 验证失败"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -2,5 +2,6 @@ pub(crate) mod antigravity;
|
||||
pub(crate) mod chatgpt_web;
|
||||
pub(crate) mod codex;
|
||||
pub(crate) mod dispatch;
|
||||
pub(crate) mod grok;
|
||||
pub(crate) mod kiro;
|
||||
pub(crate) mod shared;
|
||||
|
||||
@@ -63,6 +63,13 @@ fn select_provider_oauth_runtime_endpoint(
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("openai:image")
|
||||
}),
|
||||
"grok" => matching_endpoint(endpoints, include_inactive, |endpoint| {
|
||||
endpoint
|
||||
.api_format
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("openai:chat")
|
||||
})
|
||||
.or_else(|| matching_endpoint(endpoints, include_inactive, |_| true)),
|
||||
"antigravity" => matching_endpoint(endpoints, include_inactive, |endpoint| {
|
||||
endpoint
|
||||
.api_format
|
||||
|
||||
+1
-1
@@ -40,7 +40,7 @@ pub(super) async fn admin_provider_ops_sub2api_balance_payload(
|
||||
"query_balance",
|
||||
message,
|
||||
None,
|
||||
)
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -57,7 +57,7 @@ pub(super) async fn handle_admin_provider_ops_action(
|
||||
Err(_) => {
|
||||
return Ok(Some(bad_request_detail_response(
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
)))
|
||||
)));
|
||||
}
|
||||
};
|
||||
let payload =
|
||||
@@ -67,7 +67,7 @@ pub(super) async fn handle_admin_provider_ops_action(
|
||||
Err(_) => {
|
||||
return Ok(Some(bad_request_detail_response(
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
)))
|
||||
)));
|
||||
}
|
||||
};
|
||||
payload.config
|
||||
|
||||
@@ -13,6 +13,7 @@ use crate::provider_pool_demand::{
|
||||
provider_pool_burst_pending, read_provider_pool_demand_snapshot,
|
||||
};
|
||||
use aether_runtime_state::{DataLayerError, RuntimeState};
|
||||
use futures_util::future::join_all;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use tracing::warn;
|
||||
@@ -32,15 +33,16 @@ pub(crate) async fn read_admin_provider_pool_cooldown_counts(
|
||||
runtime: &RuntimeState,
|
||||
provider_ids: &[String],
|
||||
) -> BTreeMap<String, usize> {
|
||||
let mut counts = BTreeMap::new();
|
||||
for provider_id in provider_ids {
|
||||
join_all(provider_ids.iter().map(|provider_id| async move {
|
||||
let count = runtime
|
||||
.set_len(&pool_cooldown_index_key(provider_id))
|
||||
.await
|
||||
.unwrap_or(0);
|
||||
counts.insert(provider_id.clone(), count);
|
||||
}
|
||||
counts
|
||||
(provider_id.clone(), count)
|
||||
}))
|
||||
.await
|
||||
.into_iter()
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
@@ -170,11 +172,15 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
}
|
||||
|
||||
let now = current_unix_secs();
|
||||
for (key_id, cost_key) in key_ids.iter().zip(cost_keys) {
|
||||
let window_start = now.saturating_sub(pool_config.cost_window_seconds) as f64;
|
||||
let total = runtime
|
||||
.score_range_by_min(&cost_key, window_start)
|
||||
.await
|
||||
let cost_window_start = now.saturating_sub(pool_config.cost_window_seconds) as f64;
|
||||
let cost_results = join_all(
|
||||
cost_keys
|
||||
.iter()
|
||||
.map(|cost_key| runtime.score_range_by_min(cost_key, cost_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(cost_results) {
|
||||
let total = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
.map(|member| parse_pool_cost_member(member))
|
||||
@@ -184,11 +190,15 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
}
|
||||
}
|
||||
|
||||
for (key_id, latency_key) in key_ids.iter().zip(latency_keys) {
|
||||
let window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
|
||||
let samples = runtime
|
||||
.score_range_by_min(&latency_key, window_start)
|
||||
.await
|
||||
let latency_window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
|
||||
let latency_results = join_all(
|
||||
latency_keys
|
||||
.iter()
|
||||
.map(|latency_key| runtime.score_range_by_min(latency_key, latency_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(latency_results) {
|
||||
let samples = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
.map(|member| parse_pool_latency_member(member))
|
||||
|
||||
@@ -4,7 +4,8 @@ use super::keys::{
|
||||
};
|
||||
use crate::handlers::admin::provider::pool::config::admin_provider_pool_cache_affinity_enabled;
|
||||
use crate::handlers::admin::provider::shared::support::{
|
||||
AdminProviderPoolConfig, AdminProviderPoolUnschedulableRule,
|
||||
admin_provider_pool_quota_probe_active_members_key, AdminProviderPoolConfig,
|
||||
AdminProviderPoolUnschedulableRule,
|
||||
};
|
||||
use aether_runtime_state::RuntimeState;
|
||||
use regex::Regex;
|
||||
@@ -311,6 +312,27 @@ async fn set_pool_cooldown(
|
||||
std::time::Duration::from_secs(ttl_seconds.saturating_add(60)),
|
||||
)
|
||||
.await;
|
||||
spawn_remove_pool_active_probe_member(runtime, provider_id, key_id);
|
||||
}
|
||||
|
||||
fn spawn_remove_pool_active_probe_member(runtime: &RuntimeState, provider_id: &str, key_id: &str) {
|
||||
let runtime = runtime.clone();
|
||||
let provider_id = provider_id.to_string();
|
||||
let key_id = key_id.to_string();
|
||||
tokio::spawn(async move {
|
||||
if let Err(err) = runtime
|
||||
.set_remove(
|
||||
&admin_provider_pool_quota_probe_active_members_key(&provider_id),
|
||||
&key_id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
"gateway admin provider pool: failed to remove active probe member for provider {provider_id} key {key_id}: {:?}",
|
||||
err
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
async fn invalidate_pool_oauth_cache(runtime: &RuntimeState, key_id: &str) {
|
||||
@@ -423,10 +445,12 @@ pub(crate) async fn record_admin_provider_pool_error(
|
||||
|
||||
if status_code == 401 {
|
||||
invalidate_pool_oauth_cache(runtime, key_id).await;
|
||||
spawn_remove_pool_active_probe_member(runtime, provider_id, key_id);
|
||||
return;
|
||||
}
|
||||
|
||||
if status_code == 402 {
|
||||
spawn_remove_pool_active_probe_member(runtime, provider_id, key_id);
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -435,6 +459,7 @@ pub(crate) async fn record_admin_provider_pool_error(
|
||||
.iter()
|
||||
.any(|pattern| error_message.contains(pattern))
|
||||
{
|
||||
spawn_remove_pool_active_probe_member(runtime, provider_id, key_id);
|
||||
return;
|
||||
}
|
||||
set_pool_cooldown(
|
||||
@@ -569,8 +594,8 @@ mod tests {
|
||||
};
|
||||
use crate::handlers::admin::provider::pool::runtime::reads::read_admin_provider_pool_runtime_state;
|
||||
use crate::handlers::admin::provider::shared::support::{
|
||||
AdminProviderPoolConfig, AdminProviderPoolSchedulingPreset,
|
||||
AdminProviderPoolUnschedulableRule,
|
||||
admin_provider_pool_quota_probe_active_members_key, AdminProviderPoolConfig,
|
||||
AdminProviderPoolSchedulingPreset, AdminProviderPoolUnschedulableRule,
|
||||
};
|
||||
use crate::AppState;
|
||||
use aether_runtime_state::{RedisClientConfig, RuntimeState, RuntimeStateConfig};
|
||||
@@ -642,6 +667,24 @@ mod tests {
|
||||
.with_runtime_state(std::sync::Arc::new(runtime_state))
|
||||
}
|
||||
|
||||
async fn wait_for_active_probe_members_empty(runtime: &RuntimeState, set_key: &str) {
|
||||
for _ in 0..20 {
|
||||
let members = runtime
|
||||
.set_members(set_key)
|
||||
.await
|
||||
.expect("active members should read");
|
||||
if members.is_empty() {
|
||||
return;
|
||||
}
|
||||
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
|
||||
}
|
||||
let members = runtime
|
||||
.set_members(set_key)
|
||||
.await
|
||||
.expect("active members should read");
|
||||
assert!(members.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_google_quota_cooldown_from_reset_timestamp() {
|
||||
let now_unix_secs = chrono::DateTime::parse_from_rfc3339("2026-04-17T10:00:00Z")
|
||||
@@ -910,6 +953,50 @@ mod tests {
|
||||
.is_some_and(|ttl| *ttl <= 120 && *ttl >= 100));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn error_feedback_removes_active_probe_member_when_key_becomes_unschedulable() {
|
||||
let Some(redis) = start_managed_redis_or_skip().await else {
|
||||
return;
|
||||
};
|
||||
let app = build_runner_app(redis.redis_url(), "pool_runtime_evict_active_probe").await;
|
||||
let runtime = app.runtime_state.as_ref();
|
||||
let pool_config = sample_pool_config();
|
||||
let set_key = admin_provider_pool_quota_probe_active_members_key("provider-1");
|
||||
runtime
|
||||
.set_add(&set_key, "key-2")
|
||||
.await
|
||||
.expect("active member should insert");
|
||||
|
||||
record_admin_provider_pool_error(
|
||||
runtime,
|
||||
"provider-1",
|
||||
"key-2",
|
||||
&pool_config,
|
||||
429,
|
||||
Some(r#"{"error":{"message":"rate limited"}}"#),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
wait_for_active_probe_members_empty(runtime, &set_key).await;
|
||||
|
||||
runtime
|
||||
.set_add(&set_key, "key-402")
|
||||
.await
|
||||
.expect("active member should insert");
|
||||
record_admin_provider_pool_error(
|
||||
runtime,
|
||||
"provider-1",
|
||||
"key-402",
|
||||
&pool_config,
|
||||
402,
|
||||
Some(r#"{"error":{"message":"quota exhausted"}}"#),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
wait_for_active_probe_members_empty(runtime, &set_key).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn error_feedback_uses_google_quota_cooldown_when_retry_after_missing() {
|
||||
let Some(redis) = start_managed_redis_or_skip().await else {
|
||||
|
||||
@@ -3,7 +3,9 @@ use crate::handlers::admin::provider::shared::support::{
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::{provider_key_status_snapshot_payload, unix_secs_to_rfc3339};
|
||||
use crate::provider_key_auth::{provider_key_auth_semantics, provider_key_effective_api_formats};
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_semantics, provider_key_can_refresh_oauth, provider_key_effective_api_formats,
|
||||
};
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
use aether_data_contracts::repository::pool_scores::StoredPoolMemberScore;
|
||||
@@ -597,6 +599,185 @@ fn admin_pool_build_antigravity_account_quota_from_snapshot(
|
||||
))
|
||||
}
|
||||
|
||||
fn admin_pool_grok_quota_window_label(
|
||||
window: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> String {
|
||||
let raw_code = window
|
||||
.get("code")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default()
|
||||
.trim_start_matches("model:")
|
||||
.to_ascii_lowercase();
|
||||
let raw_label = window
|
||||
.get("label")
|
||||
.or_else(|| window.get("model"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or(raw_code.as_str())
|
||||
.to_ascii_lowercase();
|
||||
match raw_label.as_str() {
|
||||
"quota_auto" | "auto" => "Auto".to_string(),
|
||||
"quota_fast" | "fast" => "Fast".to_string(),
|
||||
"quota_expert" | "expert" => "Expert".to_string(),
|
||||
"quota_heavy" | "heavy" => "Heavy".to_string(),
|
||||
"quota_grok_4_3" | "grok-420-computer-use-sa" => "Grok 4.3".to_string(),
|
||||
_ => match raw_code.as_str() {
|
||||
"quota_auto" | "auto" => "Auto".to_string(),
|
||||
"quota_fast" | "fast" => "Fast".to_string(),
|
||||
"quota_expert" | "expert" => "Expert".to_string(),
|
||||
"quota_heavy" | "heavy" => "Heavy".to_string(),
|
||||
"quota_grok_4_3" | "grok-420-computer-use-sa" => "Grok 4.3".to_string(),
|
||||
_ => window
|
||||
.get("label")
|
||||
.or_else(|| window.get("model"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("模式")
|
||||
.to_string(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_pool_quota_window_remaining_percent(
|
||||
window: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Option<f64> {
|
||||
admin_pool_json_to_f64(window.get("remaining_ratio"))
|
||||
.map(|value| (value * 100.0).clamp(0.0, 100.0))
|
||||
.or_else(|| {
|
||||
admin_pool_json_to_f64(window.get("used_ratio"))
|
||||
.map(|value| ((1.0 - value) * 100.0).clamp(0.0, 100.0))
|
||||
})
|
||||
.or_else(|| {
|
||||
admin_pool_json_to_f64(window.get("remaining_value"))
|
||||
.zip(admin_pool_json_to_f64(window.get("limit_value")))
|
||||
.and_then(|(remaining, limit)| {
|
||||
(limit > 0.0).then_some((remaining / limit * 100.0).clamp(0.0, 100.0))
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
admin_pool_json_to_f64(window.get("used_value"))
|
||||
.zip(admin_pool_json_to_f64(window.get("limit_value")))
|
||||
.and_then(|(used, limit)| {
|
||||
(limit > 0.0).then_some(((1.0 - used / limit) * 100.0).clamp(0.0, 100.0))
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn admin_pool_quota_window_value_text(
|
||||
window: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Option<String> {
|
||||
let limit_value =
|
||||
admin_pool_json_to_f64(window.get("limit_value")).filter(|value| *value > 0.0)?;
|
||||
if let Some(remaining_value) = admin_pool_json_to_f64(window.get("remaining_value")) {
|
||||
return Some(format!(
|
||||
"{}/{}",
|
||||
admin_pool_format_quota_value(remaining_value),
|
||||
admin_pool_format_quota_value(limit_value),
|
||||
));
|
||||
}
|
||||
admin_pool_json_to_f64(window.get("used_value")).map(|used_value| {
|
||||
format!(
|
||||
"{}/{}",
|
||||
admin_pool_format_quota_value((limit_value - used_value).max(0.0)),
|
||||
admin_pool_format_quota_value(limit_value),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn admin_pool_build_grok_account_quota_from_snapshot(
|
||||
quota_snapshot: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Option<String> {
|
||||
let code = quota_snapshot
|
||||
.get("code")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default();
|
||||
if code.eq_ignore_ascii_case("banned") {
|
||||
return quota_snapshot
|
||||
.get("label")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| Some("账号已封禁".to_string()));
|
||||
}
|
||||
if code.eq_ignore_ascii_case("forbidden") {
|
||||
return quota_snapshot
|
||||
.get("label")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| Some("访问受限".to_string()));
|
||||
}
|
||||
|
||||
let model_parts = admin_pool_quota_windows(quota_snapshot)
|
||||
.into_iter()
|
||||
.filter(|window| {
|
||||
window
|
||||
.get("scope")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some_and(|scope| scope.eq_ignore_ascii_case("model"))
|
||||
})
|
||||
.filter_map(|window| {
|
||||
let remaining_percent = admin_pool_quota_window_remaining_percent(window)?;
|
||||
let mut part = format!(
|
||||
"{}剩余 {}",
|
||||
admin_pool_grok_quota_window_label(window),
|
||||
admin_pool_format_percent(remaining_percent),
|
||||
);
|
||||
if let Some(value_text) = admin_pool_quota_window_value_text(window) {
|
||||
part.push_str(&format!(" ({value_text})"));
|
||||
}
|
||||
Some(part)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
if !model_parts.is_empty() {
|
||||
return Some(model_parts.join(" | "));
|
||||
}
|
||||
|
||||
let window = admin_pool_quota_window(quota_snapshot, "usage")
|
||||
.or_else(|| admin_pool_quota_windows(quota_snapshot).into_iter().next())?;
|
||||
let remaining_value = admin_pool_json_to_f64(window.get("remaining_value"));
|
||||
let limit_value = admin_pool_json_to_f64(window.get("limit_value"));
|
||||
if let (Some(remaining_value), Some(limit_value)) = (remaining_value, limit_value) {
|
||||
if limit_value > 0.0 && remaining_value <= 0.0 {
|
||||
return Some(format!(
|
||||
"剩余 {}/{}",
|
||||
admin_pool_format_quota_value(remaining_value),
|
||||
admin_pool_format_quota_value(limit_value),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(remaining_percent) = admin_pool_quota_window_remaining_percent(window) {
|
||||
if let Some(value_text) = admin_pool_quota_window_value_text(window) {
|
||||
return Some(format!(
|
||||
"剩余 {} ({value_text})",
|
||||
admin_pool_format_percent(remaining_percent),
|
||||
));
|
||||
}
|
||||
return Some(format!(
|
||||
"剩余 {}",
|
||||
admin_pool_format_percent(remaining_percent),
|
||||
));
|
||||
}
|
||||
|
||||
match (remaining_value, limit_value) {
|
||||
(Some(remaining_value), Some(limit_value)) if limit_value > 0.0 => Some(format!(
|
||||
"剩余 {}/{}",
|
||||
admin_pool_format_quota_value(remaining_value),
|
||||
admin_pool_format_quota_value(limit_value),
|
||||
)),
|
||||
_ => quota_snapshot
|
||||
.get("label")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_pool_build_gemini_cli_account_quota_from_snapshot(
|
||||
quota_snapshot: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Option<String> {
|
||||
@@ -699,6 +880,13 @@ fn admin_pool_build_account_quota(
|
||||
return Some(account_quota);
|
||||
}
|
||||
}
|
||||
"grok" => {
|
||||
if let Some(account_quota) =
|
||||
admin_pool_build_grok_account_quota_from_snapshot(quota_snapshot)
|
||||
{
|
||||
return Some(account_quota);
|
||||
}
|
||||
}
|
||||
"gemini_cli" => {
|
||||
if let Some(account_quota) =
|
||||
admin_pool_build_gemini_cli_account_quota_from_snapshot(quota_snapshot)
|
||||
@@ -962,7 +1150,10 @@ pub(super) fn build_admin_pool_key_payload(
|
||||
);
|
||||
payload.insert(
|
||||
"can_refresh_oauth".to_string(),
|
||||
json!(auth_semantics.can_refresh_oauth()),
|
||||
json!(provider_key_can_refresh_oauth(
|
||||
auth_semantics,
|
||||
auth_config.as_ref()
|
||||
)),
|
||||
);
|
||||
payload.insert(
|
||||
"can_export_oauth".to_string(),
|
||||
@@ -1225,4 +1416,39 @@ mod tests {
|
||||
assert_eq!(usage["total_tokens"], json!(375));
|
||||
assert_eq!(usage["total_cost_usd"], json!("0.60000000"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_model_quota_is_rendered_for_pool_rows() {
|
||||
let quota_snapshot = json!({
|
||||
"provider_type": "grok",
|
||||
"code": "ok",
|
||||
"exhausted": false,
|
||||
"plan_type": "heavy",
|
||||
"pool_tier": "heavy",
|
||||
"windows": [
|
||||
{
|
||||
"code": "model:quota_auto",
|
||||
"label": "auto",
|
||||
"scope": "model",
|
||||
"remaining_ratio": 0.4,
|
||||
"used_value": 90,
|
||||
"limit_value": 150
|
||||
},
|
||||
{
|
||||
"code": "model:quota_heavy",
|
||||
"label": "heavy",
|
||||
"scope": "model",
|
||||
"remaining_ratio": 0.0,
|
||||
"used_value": 20,
|
||||
"limit_value": 20
|
||||
}
|
||||
]
|
||||
});
|
||||
let quota_snapshot = quota_snapshot.as_object().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
admin_pool_build_account_quota("grok", Some(quota_snapshot)),
|
||||
Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+61
-41
@@ -19,6 +19,7 @@ use axum::{
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use futures_util::future::join_all;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub(super) async fn build_admin_pool_overview_response(
|
||||
@@ -67,47 +68,66 @@ pub(super) async fn build_admin_pool_overview_response(
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
|
||||
let probe_config = PoolQuotaProbeWorkerConfig::from_env();
|
||||
let mut runtime_metrics_by_provider = BTreeMap::new();
|
||||
for (provider, pool_config) in &pool_enabled_providers {
|
||||
let active_keys = key_stats_by_provider
|
||||
.get(&provider.id)
|
||||
.map(|item| item.active_keys as usize)
|
||||
.unwrap_or(0);
|
||||
let hot_count = if pool_config.probing_enabled {
|
||||
state
|
||||
.runtime_state()
|
||||
.set_len(&admin_provider_pool_quota_probe_active_members_key(
|
||||
&provider.id,
|
||||
))
|
||||
.await
|
||||
.unwrap_or(0)
|
||||
} else {
|
||||
0
|
||||
};
|
||||
let demand_snapshot = read_provider_pool_demand_snapshot(
|
||||
state.runtime_state(),
|
||||
&provider.id,
|
||||
active_keys,
|
||||
probe_config.max_keys_per_provider,
|
||||
)
|
||||
.await;
|
||||
let burst_pending = pool_config.probing_enabled
|
||||
&& provider_pool_burst_pending(state.runtime_state(), &provider.id).await;
|
||||
runtime_metrics_by_provider.insert(
|
||||
provider.id.clone(),
|
||||
json!({
|
||||
"provider_hot_count": hot_count,
|
||||
"provider_desired_hot": if pool_config.probing_enabled {
|
||||
demand_snapshot.desired_hot
|
||||
} else {
|
||||
0
|
||||
},
|
||||
"provider_in_flight": demand_snapshot.in_flight,
|
||||
"provider_ema_in_flight": demand_snapshot.ema_in_flight,
|
||||
"provider_burst_pending": burst_pending,
|
||||
}),
|
||||
);
|
||||
}
|
||||
let runtime_metrics_by_provider = join_all(pool_enabled_providers.iter().map(
|
||||
|(provider, pool_config)| {
|
||||
let provider_id = provider.id.clone();
|
||||
let probing_enabled = pool_config.probing_enabled;
|
||||
let active_keys = key_stats_by_provider
|
||||
.get(&provider.id)
|
||||
.map(|item| item.active_keys as usize)
|
||||
.unwrap_or(0);
|
||||
let max_keys_per_provider = probe_config.max_keys_per_provider;
|
||||
|
||||
async move {
|
||||
let hot_count_future = async {
|
||||
if probing_enabled {
|
||||
state
|
||||
.runtime_state()
|
||||
.set_len(&admin_provider_pool_quota_probe_active_members_key(
|
||||
&provider_id,
|
||||
))
|
||||
.await
|
||||
.unwrap_or(0)
|
||||
} else {
|
||||
0
|
||||
}
|
||||
};
|
||||
let demand_snapshot_future = read_provider_pool_demand_snapshot(
|
||||
state.runtime_state(),
|
||||
&provider_id,
|
||||
active_keys,
|
||||
max_keys_per_provider,
|
||||
);
|
||||
let burst_pending_future = async {
|
||||
probing_enabled
|
||||
&& provider_pool_burst_pending(state.runtime_state(), &provider_id).await
|
||||
};
|
||||
let (hot_count, demand_snapshot, burst_pending) = tokio::join!(
|
||||
hot_count_future,
|
||||
demand_snapshot_future,
|
||||
burst_pending_future
|
||||
);
|
||||
|
||||
(
|
||||
provider_id,
|
||||
json!({
|
||||
"provider_hot_count": hot_count,
|
||||
"provider_desired_hot": if probing_enabled {
|
||||
demand_snapshot.desired_hot
|
||||
} else {
|
||||
0
|
||||
},
|
||||
"provider_in_flight": demand_snapshot.in_flight,
|
||||
"provider_ema_in_flight": demand_snapshot.ema_in_flight,
|
||||
"provider_burst_pending": burst_pending,
|
||||
}),
|
||||
)
|
||||
}
|
||||
},
|
||||
))
|
||||
.await
|
||||
.into_iter()
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
|
||||
let providers = pool_enabled_providers
|
||||
.into_iter()
|
||||
|
||||
+3
-2
@@ -3,7 +3,7 @@ use super::{
|
||||
AdminPoolResolveSelectionRequest, ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::provider_key_auth::provider_key_auth_semantics;
|
||||
use crate::provider_key_auth::{provider_key_auth_semantics, provider_key_can_refresh_oauth};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::pool as admin_provider_pool_pure;
|
||||
use axum::{
|
||||
@@ -94,6 +94,7 @@ pub(super) async fn build_admin_pool_resolve_selection_response(
|
||||
.iter()
|
||||
.map(|key| {
|
||||
let auth_semantics = provider_key_auth_semantics(key, &provider_type);
|
||||
let auth_config = state.parse_catalog_auth_config_json(key);
|
||||
json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
@@ -102,7 +103,7 @@ pub(super) async fn build_admin_pool_resolve_selection_response(
|
||||
"credential_kind": auth_semantics.credential_kind().as_str(),
|
||||
"runtime_auth_kind": auth_semantics.runtime_auth_kind().as_str(),
|
||||
"oauth_managed": auth_semantics.oauth_managed(),
|
||||
"can_refresh_oauth": auth_semantics.can_refresh_oauth(),
|
||||
"can_refresh_oauth": provider_key_can_refresh_oauth(auth_semantics, auth_config.as_ref()),
|
||||
"can_export_oauth": auth_semantics.can_export_oauth(),
|
||||
"can_edit_oauth": auth_semantics.can_edit_oauth(),
|
||||
})
|
||||
|
||||
@@ -17,6 +17,9 @@ use crate::ai_serving::{
|
||||
};
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::execution_runtime;
|
||||
use crate::handlers::admin::provider::shared::model_test_capabilities::{
|
||||
admin_provider_model_supports_image_generation, admin_provider_model_test_capabilities_payload,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::handlers::shared::provider_pool::{
|
||||
admin_provider_pool_config_from_config_value, read_admin_provider_pool_runtime_state,
|
||||
@@ -100,6 +103,181 @@ struct ProviderQueryKeyFetchResult {
|
||||
has_success: bool,
|
||||
}
|
||||
|
||||
fn provider_query_model_id(model: &Value) -> Option<&str> {
|
||||
model
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn provider_query_grok_required_tier_rank(model_id: &str) -> Option<u8> {
|
||||
match model_id.trim() {
|
||||
"grok-4.20-0309-non-reasoning" | "grok-4.20-fast" | "grok-imagine-image-lite" => Some(0),
|
||||
"grok-4.20-0309"
|
||||
| "grok-4.20-0309-reasoning"
|
||||
| "grok-4.20-0309-non-reasoning-super"
|
||||
| "grok-4.20-0309-super"
|
||||
| "grok-4.20-0309-reasoning-super"
|
||||
| "grok-4.20-auto"
|
||||
| "grok-4.20-expert"
|
||||
| "grok-4.3-beta"
|
||||
| "grok-imagine-image"
|
||||
| "grok-imagine-image-pro"
|
||||
| "grok-imagine-image-edit" => Some(1),
|
||||
"grok-4.20-0309-non-reasoning-heavy"
|
||||
| "grok-4.20-0309-heavy"
|
||||
| "grok-4.20-0309-reasoning-heavy"
|
||||
| "grok-4.20-multi-agent-0309"
|
||||
| "grok-4.20-heavy" => Some(2),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_normalize_grok_pool_tier(value: Option<&str>) -> Option<&'static str> {
|
||||
match value?.trim().to_ascii_lowercase().as_str() {
|
||||
"basic" => Some("basic"),
|
||||
"super" => Some("super"),
|
||||
"heavy" => Some("heavy"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_grok_pool_tier_rank(value: Option<&str>) -> u8 {
|
||||
match provider_query_normalize_grok_pool_tier(value).unwrap_or("basic") {
|
||||
"heavy" => 2,
|
||||
"super" => 1,
|
||||
_ => 0,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_grok_quota_string(quota: &Map<String, Value>, fields: &[&str]) -> Option<String> {
|
||||
fields.iter().find_map(|field| {
|
||||
quota
|
||||
.get(*field)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_grok_window_limit(quota: &Map<String, Value>, model_name: &str) -> Option<f64> {
|
||||
quota
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)?
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.find(|window| {
|
||||
window
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| value.trim() == model_name)
|
||||
})
|
||||
.and_then(|window| window.get("limit_value"))
|
||||
.and_then(Value::as_f64)
|
||||
.filter(|value| value.is_finite() && *value > 0.0)
|
||||
}
|
||||
|
||||
fn provider_query_grok_pool_tier_from_quota(quota: &Map<String, Value>) -> Option<&'static str> {
|
||||
if let Some(tier) =
|
||||
provider_query_grok_quota_string(quota, &["pool_tier", "tier", "plan_type", "plan"])
|
||||
.and_then(|value| provider_query_normalize_grok_pool_tier(Some(&value)))
|
||||
{
|
||||
return Some(tier);
|
||||
}
|
||||
|
||||
if let Some(auto_total) = provider_query_grok_window_limit(quota, "quota_auto") {
|
||||
if (auto_total - 150.0).abs() < f64::EPSILON {
|
||||
return Some("heavy");
|
||||
}
|
||||
if (auto_total - 50.0).abs() < f64::EPSILON {
|
||||
return Some("super");
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(fast_total) = provider_query_grok_window_limit(quota, "quota_fast") {
|
||||
if (fast_total - 400.0).abs() < f64::EPSILON {
|
||||
return Some("heavy");
|
||||
}
|
||||
if (fast_total - 140.0).abs() < f64::EPSILON {
|
||||
return Some("super");
|
||||
}
|
||||
if (fast_total - 30.0).abs() < f64::EPSILON {
|
||||
return Some("basic");
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn provider_query_grok_key_pool_tier(key: &StoredProviderCatalogKey) -> Option<&'static str> {
|
||||
key.status_snapshot
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|snapshot| snapshot.get("quota"))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(provider_query_grok_pool_tier_from_quota)
|
||||
}
|
||||
|
||||
fn provider_query_filter_models_for_key(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
key: &StoredProviderCatalogKey,
|
||||
models: Vec<Value>,
|
||||
) -> Vec<Value> {
|
||||
if !provider.provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
return models;
|
||||
}
|
||||
|
||||
let allowed_rank = provider_query_grok_pool_tier_rank(provider_query_grok_key_pool_tier(key));
|
||||
models
|
||||
.into_iter()
|
||||
.filter(|model| {
|
||||
provider_query_model_id(model)
|
||||
.and_then(provider_query_grok_required_tier_rank)
|
||||
.is_some_and(|required_rank| required_rank <= allowed_rank)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn provider_query_attach_model_test_capabilities(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
models: Vec<Value>,
|
||||
) -> Vec<Value> {
|
||||
models
|
||||
.into_iter()
|
||||
.map(|mut model| {
|
||||
let Some(object) = model.as_object_mut() else {
|
||||
return model;
|
||||
};
|
||||
let model_id = object
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
let supports_image_generation = admin_provider_model_supports_image_generation(
|
||||
&provider.provider_type,
|
||||
&model_id,
|
||||
object
|
||||
.get("supports_image_generation")
|
||||
.or_else(|| object.get("effective_supports_image_generation"))
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false),
|
||||
);
|
||||
object.insert(
|
||||
"model_test_capabilities".to_string(),
|
||||
admin_provider_model_test_capabilities_payload(
|
||||
&provider.provider_type,
|
||||
&model_id,
|
||||
supports_image_generation,
|
||||
),
|
||||
);
|
||||
model
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn provider_query_codex_preset_fallback(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Option<ProviderQueryKeyFetchResult> {
|
||||
@@ -245,8 +423,9 @@ async fn provider_query_fetch_models_for_key(
|
||||
if let Some(cached_models) =
|
||||
provider_query_read_cached_models(state, &provider.id, &key.id).await
|
||||
{
|
||||
let models = provider_query_filter_models_for_key(provider, key, cached_models);
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: cached_models,
|
||||
models,
|
||||
error: None,
|
||||
from_cache: true,
|
||||
has_success: true,
|
||||
@@ -257,8 +436,13 @@ async fn provider_query_fetch_models_for_key(
|
||||
let selected_endpoints = selected_models_fetch_endpoints(endpoints, key);
|
||||
if selected_endpoints.is_empty() {
|
||||
if let Some(models) = preset_models_for_provider(&provider.provider_type) {
|
||||
let models = provider_query_filter_models_for_key(
|
||||
provider,
|
||||
key,
|
||||
aggregate_models_for_cache(&models),
|
||||
);
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: aggregate_models_for_cache(&models),
|
||||
models,
|
||||
error: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
@@ -342,7 +526,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
}
|
||||
|
||||
Ok(ProviderQueryKeyFetchResult {
|
||||
models: unique_models,
|
||||
models: provider_query_filter_models_for_key(provider, key, unique_models),
|
||||
error,
|
||||
from_cache: false,
|
||||
has_success: outcome.has_success,
|
||||
@@ -397,11 +581,12 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
force_refresh,
|
||||
)
|
||||
.await?;
|
||||
let success = !result.models.is_empty();
|
||||
let models = provider_query_attach_model_test_capabilities(&provider, result.models);
|
||||
let success = !models.is_empty();
|
||||
return Ok(Json(json!({
|
||||
"success": success,
|
||||
"data": {
|
||||
"models": result.models,
|
||||
"models": models,
|
||||
"error": result.error,
|
||||
"from_cache": result.from_cache,
|
||||
},
|
||||
@@ -429,6 +614,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
{
|
||||
if let Some(models) = provider_query_read_provider_cached_models(state, &provider.id).await
|
||||
{
|
||||
let models = provider_query_attach_model_test_capabilities(&provider, models);
|
||||
return Ok(Json(json!({
|
||||
"success": !models.is_empty(),
|
||||
"data": {
|
||||
@@ -504,6 +690,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
if !success && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
|
||||
}
|
||||
let models = provider_query_attach_model_test_capabilities(&provider, models);
|
||||
|
||||
Ok(Json(json!({
|
||||
"success": success,
|
||||
@@ -519,3 +706,152 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
|
||||
fn grok_provider() -> StoredProviderCatalogProvider {
|
||||
let mut provider = StoredProviderCatalogProvider::new(
|
||||
"provider-1".to_string(),
|
||||
"Grok".to_string(),
|
||||
None,
|
||||
"grok".to_string(),
|
||||
)
|
||||
.expect("provider should build");
|
||||
provider.provider_type = "grok".to_string();
|
||||
provider
|
||||
}
|
||||
|
||||
fn grok_key_with_quota(quota: Value) -> StoredProviderCatalogKey {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"key-1".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
key.status_snapshot = Some(json!({ "quota": quota }));
|
||||
key
|
||||
}
|
||||
|
||||
fn model(id: &str) -> Value {
|
||||
json!({ "id": id })
|
||||
}
|
||||
|
||||
fn filtered_ids(key: &StoredProviderCatalogKey) -> Vec<String> {
|
||||
provider_query_filter_models_for_key(
|
||||
&grok_provider(),
|
||||
key,
|
||||
vec![
|
||||
model("grok-4.20-0309-non-reasoning"),
|
||||
model("grok-4.20-auto"),
|
||||
model("grok-4.20-heavy"),
|
||||
model("grok-imagine-image-lite"),
|
||||
model("grok-imagine-image"),
|
||||
model("grok-imagine-image-edit"),
|
||||
],
|
||||
)
|
||||
.into_iter()
|
||||
.filter_map(|item| item.get("id").and_then(Value::as_str).map(str::to_string))
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_basic_tier_hides_super_and_heavy_models() {
|
||||
let key = grok_key_with_quota(json!({ "pool_tier": "basic" }));
|
||||
|
||||
assert_eq!(
|
||||
filtered_ids(&key),
|
||||
["grok-4.20-0309-non-reasoning", "grok-imagine-image-lite"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_super_tier_hides_heavy_models() {
|
||||
let key = grok_key_with_quota(json!({ "plan_type": "super" }));
|
||||
|
||||
assert_eq!(
|
||||
filtered_ids(&key),
|
||||
[
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
"grok-4.20-auto",
|
||||
"grok-imagine-image-lite",
|
||||
"grok-imagine-image",
|
||||
"grok-imagine-image-edit"
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_heavy_tier_keeps_full_non_video_catalog() {
|
||||
let key = grok_key_with_quota(json!({ "pool_tier": "heavy" }));
|
||||
|
||||
assert_eq!(
|
||||
filtered_ids(&key),
|
||||
[
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
"grok-4.20-auto",
|
||||
"grok-4.20-heavy",
|
||||
"grok-imagine-image-lite",
|
||||
"grok-imagine-image",
|
||||
"grok-imagine-image-edit"
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_tier_falls_back_to_live_quota_windows() {
|
||||
let key = grok_key_with_quota(json!({
|
||||
"windows": [
|
||||
{ "model": "quota_fast", "limit_value": 140.0 }
|
||||
]
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
filtered_ids(&key),
|
||||
[
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
"grok-4.20-auto",
|
||||
"grok-imagine-image-lite",
|
||||
"grok-imagine-image",
|
||||
"grok-imagine-image-edit"
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_attaches_model_test_capabilities_to_models() {
|
||||
let models = provider_query_attach_model_test_capabilities(
|
||||
&grok_provider(),
|
||||
vec![
|
||||
model("grok-4.20-fast"),
|
||||
model("grok-imagine-image"),
|
||||
model("grok-imagine-image-edit"),
|
||||
],
|
||||
);
|
||||
|
||||
assert!(models[0]["model_test_capabilities"]["openai:image"].is_null());
|
||||
assert_eq!(
|
||||
models[1]["model_test_capabilities"]["openai:image"]["max_generation_count"],
|
||||
json!(4)
|
||||
);
|
||||
assert_eq!(
|
||||
models[1]["model_test_capabilities"]["openai:image"]["supports_generation"],
|
||||
json!(true)
|
||||
);
|
||||
assert_eq!(
|
||||
models[2]["model_test_capabilities"]["openai:image"]["supports_generation"],
|
||||
json!(false)
|
||||
);
|
||||
assert_eq!(
|
||||
models[2]["model_test_capabilities"]["openai:image"]["supports_edit"],
|
||||
json!(true)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@ use crate::ai_serving::{
|
||||
};
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::execution_runtime;
|
||||
use crate::handlers::admin::provider::write::provider::reconcile_admin_fixed_provider_template_endpoints;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::handlers::shared::provider_pool::{
|
||||
admin_provider_pool_config_from_config_value, read_admin_provider_pool_runtime_state,
|
||||
@@ -81,6 +82,7 @@ use tracing::{debug, warn};
|
||||
use uuid::Uuid;
|
||||
|
||||
mod adapter;
|
||||
mod capabilities;
|
||||
mod model_mapping;
|
||||
mod summary;
|
||||
|
||||
@@ -88,13 +90,17 @@ use self::adapter::{
|
||||
provider_query_antigravity_test_unsupported_reason,
|
||||
provider_query_antigravity_unsupported_reason,
|
||||
provider_query_default_antigravity_endpoint_test_body,
|
||||
provider_query_model_test_endpoint_priority, provider_query_normalize_api_format_alias,
|
||||
provider_query_standard_test_client_api_format,
|
||||
provider_query_grok_test_unsupported_reason, provider_query_model_test_endpoint_priority,
|
||||
provider_query_normalize_api_format_alias, provider_query_standard_test_client_api_format,
|
||||
provider_query_standard_test_unsupported_reason,
|
||||
provider_query_test_adapter_for_provider_api_format,
|
||||
provider_query_transport_supports_model_test_execution,
|
||||
provider_query_unsupported_test_api_format_message, ProviderQueryTestAdapter,
|
||||
};
|
||||
use self::capabilities::{
|
||||
provider_query_openai_image_normalize_failure_message,
|
||||
provider_query_openai_image_normalize_options,
|
||||
};
|
||||
use self::model_mapping::{
|
||||
provider_query_resolve_explicit_mapped_effective_model,
|
||||
provider_query_resolve_global_effective_model,
|
||||
@@ -530,12 +536,20 @@ fn provider_query_build_test_request_body_for_route(
|
||||
provider_query_build_test_request_body_with_model_policy(payload, model, override_custom_model)
|
||||
}
|
||||
|
||||
fn provider_query_build_test_request_body_with_model_policy(
|
||||
fn provider_query_build_test_request_body_for_api_format(
|
||||
payload: &Value,
|
||||
model: &str,
|
||||
override_custom_model: bool,
|
||||
route_path: &str,
|
||||
client_api_format: &str,
|
||||
) -> Value {
|
||||
let client_api_format = provider_query_normalize_api_format_alias(client_api_format);
|
||||
let override_custom_model = route_path.ends_with("/test-model-failover")
|
||||
|| provider_query_extract_mapped_model_name(payload).is_some();
|
||||
if let Some(mut body) = provider_query_extract_request_body(payload) {
|
||||
let has_conversation = provider_query_request_body_has_conversation_for_api_format(
|
||||
&body,
|
||||
client_api_format.as_str(),
|
||||
);
|
||||
if let Some(object) = body.as_object_mut() {
|
||||
if override_custom_model {
|
||||
object.insert("model".to_string(), Value::String(model.to_string()));
|
||||
@@ -544,6 +558,123 @@ fn provider_query_build_test_request_body_with_model_policy(
|
||||
.entry("model".to_string())
|
||||
.or_insert_with(|| Value::String(model.to_string()));
|
||||
}
|
||||
if !has_conversation {
|
||||
provider_query_insert_default_test_conversation(
|
||||
object,
|
||||
client_api_format.as_str(),
|
||||
payload,
|
||||
);
|
||||
}
|
||||
}
|
||||
return body;
|
||||
}
|
||||
|
||||
let message = provider_query_extract_message(payload)
|
||||
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string());
|
||||
match client_api_format.as_str() {
|
||||
"openai:responses" | "openai:responses:compact" => json!({
|
||||
"model": model,
|
||||
"input": message,
|
||||
"max_output_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
}),
|
||||
"claude:messages" => json!({
|
||||
"model": model,
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": message
|
||||
}],
|
||||
"max_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
}),
|
||||
_ => json!({
|
||||
"model": model,
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": message
|
||||
}],
|
||||
"max_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_build_grok_test_request_body_for_api_format(
|
||||
payload: &Value,
|
||||
model: &str,
|
||||
route_path: &str,
|
||||
client_api_format: &str,
|
||||
) -> Value {
|
||||
provider_query_build_test_request_body_for_api_format(
|
||||
payload,
|
||||
model,
|
||||
route_path,
|
||||
client_api_format,
|
||||
)
|
||||
}
|
||||
|
||||
fn provider_query_insert_default_test_conversation(
|
||||
object: &mut Map<String, Value>,
|
||||
client_api_format: &str,
|
||||
payload: &Value,
|
||||
) {
|
||||
let message = provider_query_extract_message(payload)
|
||||
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string());
|
||||
match client_api_format {
|
||||
"openai:responses" | "openai:responses:compact" => {
|
||||
object.insert("input".to_string(), Value::String(message));
|
||||
}
|
||||
"claude:messages" => {
|
||||
object.insert(
|
||||
"messages".to_string(),
|
||||
json!([{ "role": "user", "content": message }]),
|
||||
);
|
||||
}
|
||||
_ => {
|
||||
object.insert(
|
||||
"messages".to_string(),
|
||||
json!([{ "role": "user", "content": message }]),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_grok_test_client_api_format(provider_api_format: &str) -> &'static str {
|
||||
match provider_query_normalize_api_format_alias(provider_api_format).as_str() {
|
||||
"openai:responses" | "openai:responses:compact" => "openai:responses",
|
||||
"claude:messages" => "claude:messages",
|
||||
_ => "openai:chat",
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_build_test_request_body_with_model_policy(
|
||||
payload: &Value,
|
||||
model: &str,
|
||||
override_custom_model: bool,
|
||||
) -> Value {
|
||||
if let Some(mut body) = provider_query_extract_request_body(payload) {
|
||||
let has_conversation = provider_query_request_body_has_conversation(&body);
|
||||
if let Some(object) = body.as_object_mut() {
|
||||
if override_custom_model {
|
||||
object.insert("model".to_string(), Value::String(model.to_string()));
|
||||
} else {
|
||||
object
|
||||
.entry("model".to_string())
|
||||
.or_insert_with(|| Value::String(model.to_string()));
|
||||
}
|
||||
if !has_conversation {
|
||||
object.insert(
|
||||
"messages".to_string(),
|
||||
json!([{
|
||||
"role": "user",
|
||||
"content": provider_query_extract_message(payload)
|
||||
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string())
|
||||
}]),
|
||||
);
|
||||
}
|
||||
}
|
||||
return body;
|
||||
}
|
||||
@@ -555,12 +686,78 @@ fn provider_query_build_test_request_body_with_model_policy(
|
||||
"content": provider_query_extract_message(payload)
|
||||
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string())
|
||||
}],
|
||||
"max_tokens": 30,
|
||||
"temperature": 0.7,
|
||||
"stream": true,
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_request_body_has_conversation(body: &Value) -> bool {
|
||||
body.get("messages")
|
||||
.and_then(Value::as_array)
|
||||
.map(|messages| {
|
||||
messages
|
||||
.iter()
|
||||
.any(|message| value_has_non_empty_text(message.get("content")))
|
||||
})
|
||||
.unwrap_or(false)
|
||||
|| value_has_non_empty_text(body.get("input"))
|
||||
|| value_has_non_empty_text(body.get("prompt"))
|
||||
|| value_has_non_empty_text(body.get("query"))
|
||||
|| value_has_non_empty_text(body.get("system"))
|
||||
}
|
||||
|
||||
fn provider_query_request_body_has_conversation_for_api_format(
|
||||
body: &Value,
|
||||
client_api_format: &str,
|
||||
) -> bool {
|
||||
match provider_query_normalize_api_format_alias(client_api_format).as_str() {
|
||||
"openai:responses" | "openai:responses:compact" => {
|
||||
value_has_non_empty_text(body.get("input"))
|
||||
|| value_has_non_empty_text(body.get("prompt"))
|
||||
}
|
||||
"claude:messages" => {
|
||||
body.get("messages")
|
||||
.and_then(Value::as_array)
|
||||
.map(|messages| {
|
||||
messages
|
||||
.iter()
|
||||
.any(|message| value_has_non_empty_text(message.get("content")))
|
||||
})
|
||||
.unwrap_or(false)
|
||||
|| value_has_non_empty_text(body.get("system"))
|
||||
}
|
||||
_ => provider_query_request_body_has_conversation(body),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_request_body_is_openai_responses_shape(body: &Value) -> bool {
|
||||
let Some(object) = body.as_object() else {
|
||||
return false;
|
||||
};
|
||||
[
|
||||
"input",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"instructions",
|
||||
"previous_response_id",
|
||||
]
|
||||
.iter()
|
||||
.any(|key| object.contains_key(*key))
|
||||
}
|
||||
|
||||
fn value_has_non_empty_text(value: Option<&Value>) -> bool {
|
||||
match value {
|
||||
Some(Value::String(value)) => !value.trim().is_empty(),
|
||||
Some(Value::Array(values)) => values
|
||||
.iter()
|
||||
.any(|value| value_has_non_empty_text(Some(value))),
|
||||
Some(Value::Object(values)) => values
|
||||
.values()
|
||||
.any(|value| value_has_non_empty_text(Some(value))),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_request_body_model<'a>(request_body: &'a Value, fallback: &'a str) -> &'a str {
|
||||
request_body
|
||||
.get("model")
|
||||
@@ -1022,6 +1219,8 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
payload: &Value,
|
||||
requested_model_override: Option<&str>,
|
||||
) -> Result<Vec<ProviderQueryTestCandidate>, Response<Body>> {
|
||||
provider_query_reconcile_fixed_provider_endpoints_for_test_model(state, provider).await?;
|
||||
|
||||
let provider_ids = vec![provider.id.clone()];
|
||||
let endpoints = state
|
||||
.app()
|
||||
@@ -1227,6 +1426,33 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
Ok(candidates)
|
||||
}
|
||||
|
||||
async fn provider_query_reconcile_fixed_provider_endpoints_for_test_model(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Result<(), Response<Body>> {
|
||||
if state
|
||||
.fixed_provider_template(&provider.provider_type)
|
||||
.is_none()
|
||||
|| !state.has_provider_catalog_data_writer()
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
reconcile_admin_fixed_provider_template_endpoints(state, provider)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
warn!(
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider.provider_type,
|
||||
error = ?err,
|
||||
"admin provider-query test-model: failed to reconcile fixed provider endpoints"
|
||||
);
|
||||
build_admin_provider_query_bad_request_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_decode_execution_body(
|
||||
result: &aether_contracts::ExecutionResult,
|
||||
) -> Option<Vec<u8>> {
|
||||
@@ -1256,7 +1482,7 @@ fn provider_query_standard_execution_response_body(
|
||||
provider_api_format: &str,
|
||||
result: &aether_contracts::ExecutionResult,
|
||||
) -> Option<Value> {
|
||||
result
|
||||
let body = result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.clone())
|
||||
@@ -1264,7 +1490,15 @@ fn provider_query_standard_execution_response_body(
|
||||
provider_query_decode_execution_body(result).and_then(|body| {
|
||||
provider_query_aggregate_standard_stream_sync_response(provider_api_format, &body)
|
||||
})
|
||||
})
|
||||
})?;
|
||||
if result.status_code < 400
|
||||
&& provider_query_normalize_api_format_alias(provider_api_format)
|
||||
== "gemini:generate_content"
|
||||
&& !crate::ai_serving::gemini_generate_content_response_has_visible_output(&body)
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(body)
|
||||
}
|
||||
|
||||
fn provider_query_extract_error_message(
|
||||
@@ -1562,6 +1796,32 @@ fn provider_query_chatgpt_web_image_internal_url(base_url: &str) -> String {
|
||||
format!("{base_url}/__aether/chatgpt-web-image")
|
||||
}
|
||||
|
||||
fn provider_query_openai_image_test_upstream_url(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
request_query: Option<&str>,
|
||||
) -> String {
|
||||
if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("chatgpt_web")
|
||||
{
|
||||
provider_query_chatgpt_web_image_internal_url(&transport.endpoint.base_url)
|
||||
} else if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok")
|
||||
{
|
||||
crate::provider_transport::build_grok_upstream_url(
|
||||
transport,
|
||||
crate::provider_transport::GROK_CHAT_PATH,
|
||||
)
|
||||
} else {
|
||||
crate::provider_transport::build_openai_image_upstream_url(transport, request_query)
|
||||
}
|
||||
}
|
||||
|
||||
async fn provider_query_finalize_openai_image_result(
|
||||
route_path: &str,
|
||||
trace_id: &str,
|
||||
@@ -1661,12 +1921,16 @@ async fn provider_query_execute_openai_image_test_candidate(
|
||||
*synthetic_request.headers_mut() = incoming_request_headers;
|
||||
let (parts, _) = synthetic_request.into_parts();
|
||||
|
||||
let Some(normalized_request) =
|
||||
crate::ai_serving::normalize_openai_image_request(&parts, &request_body, None)
|
||||
else {
|
||||
let provider_type = transport.provider.provider_type.as_str();
|
||||
let Some(normalized_request) = crate::ai_serving::normalize_openai_image_request_with_options(
|
||||
&parts,
|
||||
&request_body,
|
||||
None,
|
||||
provider_query_openai_image_normalize_options(provider_type),
|
||||
) else {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
request_body.clone(),
|
||||
"Provider request body could not be normalized for openai:image",
|
||||
provider_query_openai_image_normalize_failure_message(provider_type, &request_body),
|
||||
));
|
||||
};
|
||||
|
||||
@@ -1675,6 +1939,11 @@ async fn provider_query_execute_openai_image_test_candidate(
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("chatgpt_web");
|
||||
let is_grok = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok");
|
||||
let mut provider_request_body = if is_chatgpt_web {
|
||||
match crate::ai_serving::build_chatgpt_web_image_request_body(&parts, &request_body, None) {
|
||||
Ok(body) => body,
|
||||
@@ -1702,17 +1971,33 @@ async fn provider_query_execute_openai_image_test_candidate(
|
||||
"Provider auth is unavailable for openai:image",
|
||||
));
|
||||
};
|
||||
let transport_profile = state.resolve_transport_profile(&transport);
|
||||
|
||||
let Some(mut request_headers) = crate::provider_transport::build_openai_image_headers(
|
||||
crate::provider_transport::ProviderOpenAiImageHeadersInput {
|
||||
headers: &parts.headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &request_body,
|
||||
},
|
||||
) else {
|
||||
let Some(mut request_headers) = (if is_grok {
|
||||
crate::provider_transport::build_grok_browser_headers(
|
||||
crate::provider_transport::GrokHeaderInput {
|
||||
transport: &transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(&parts.headers),
|
||||
content_type: "application/json",
|
||||
accept: "*/*",
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &request_body,
|
||||
},
|
||||
)
|
||||
} else {
|
||||
crate::provider_transport::build_openai_image_headers(
|
||||
crate::provider_transport::ProviderOpenAiImageHeadersInput {
|
||||
headers: &parts.headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &request_body,
|
||||
},
|
||||
)
|
||||
}) else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "failed",
|
||||
skip_reason: None,
|
||||
@@ -1728,6 +2013,7 @@ async fn provider_query_execute_openai_image_test_candidate(
|
||||
};
|
||||
if is_chatgpt_web {
|
||||
request_headers.insert("x-aether-chatgpt-web-image".to_string(), "1".to_string());
|
||||
} else if is_grok {
|
||||
} else {
|
||||
crate::ai_serving::apply_codex_openai_responses_special_headers(
|
||||
&mut request_headers,
|
||||
@@ -1761,16 +2047,12 @@ async fn provider_query_execute_openai_image_test_candidate(
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or(request_model.as_str())
|
||||
.to_string();
|
||||
let image_request = if is_chatgpt_web {
|
||||
let image_request = if is_chatgpt_web || is_grok {
|
||||
provider_request_body.clone()
|
||||
} else {
|
||||
normalized_request.summary_json.clone()
|
||||
};
|
||||
let request_url = if is_chatgpt_web {
|
||||
provider_query_chatgpt_web_image_internal_url(&transport.endpoint.base_url)
|
||||
} else {
|
||||
crate::provider_transport::build_openai_image_upstream_url(&transport, parts.uri.query())
|
||||
};
|
||||
let request_url = provider_query_openai_image_test_upstream_url(&transport, parts.uri.query());
|
||||
let upstream_is_stream = provider_request_body
|
||||
.get("stream")
|
||||
.and_then(Value::as_bool)
|
||||
@@ -1796,11 +2078,27 @@ async fn provider_query_execute_openai_image_test_candidate(
|
||||
proxy: state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
|
||||
.await,
|
||||
transport_profile: state.resolve_transport_profile(&transport),
|
||||
transport_profile: transport_profile.clone(),
|
||||
timeouts: state.resolve_transport_execution_timeouts(&transport),
|
||||
};
|
||||
|
||||
let result = if is_chatgpt_web {
|
||||
let result = if is_grok {
|
||||
let report_context = json!({
|
||||
"client_api_format": "openai:image",
|
||||
"provider_api_format": "openai:image",
|
||||
"provider_type": "grok",
|
||||
"model": request_model,
|
||||
"mapped_model": mapped_model,
|
||||
"image_request": image_request.clone(),
|
||||
});
|
||||
state
|
||||
.execute_execution_runtime_sync_plan_with_report_context(
|
||||
Some(trace_id),
|
||||
&plan,
|
||||
Some(&report_context),
|
||||
)
|
||||
.await?
|
||||
} else if is_chatgpt_web {
|
||||
let report_context = json!({
|
||||
"client_api_format": "openai:image",
|
||||
"provider_api_format": "openai:image",
|
||||
@@ -2115,6 +2413,177 @@ async fn provider_query_execute_antigravity_test_candidate(
|
||||
})
|
||||
}
|
||||
|
||||
async fn provider_query_execute_grok_test_candidate(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
candidate: &ProviderQueryTestCandidate,
|
||||
payload: &Value,
|
||||
route_path: &str,
|
||||
trace_id: &str,
|
||||
) -> Result<ProviderQueryExecutionOutcome, GatewayError> {
|
||||
let Some(transport) = state
|
||||
.read_provider_transport_snapshot(&provider.id, &candidate.endpoint.id, &candidate.key.id)
|
||||
.await?
|
||||
else {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
Value::Null,
|
||||
"Provider transport snapshot is unavailable",
|
||||
));
|
||||
};
|
||||
|
||||
let provider_api_format =
|
||||
provider_query_normalize_api_format_alias(&candidate.endpoint.api_format);
|
||||
let client_api_format = provider_query_grok_test_client_api_format(&provider_api_format);
|
||||
let request_body = provider_query_build_grok_test_request_body_for_api_format(
|
||||
payload,
|
||||
&candidate.effective_model,
|
||||
route_path,
|
||||
client_api_format,
|
||||
);
|
||||
if let Some(reason) =
|
||||
provider_query_grok_test_unsupported_reason(&transport, &provider_api_format)
|
||||
{
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
request_body,
|
||||
format!(
|
||||
"{} ({reason})",
|
||||
provider_query_unsupported_test_api_format_message(&candidate.endpoint.api_format)
|
||||
),
|
||||
));
|
||||
}
|
||||
|
||||
let incoming_request_headers = provider_query_extract_request_headers(payload);
|
||||
let mut synthetic_request = http::Request::builder()
|
||||
.uri(route_path)
|
||||
.body(())
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
*synthetic_request.headers_mut() = incoming_request_headers;
|
||||
let (parts, _) = synthetic_request.into_parts();
|
||||
|
||||
let request_model =
|
||||
provider_query_request_body_model(&request_body, &candidate.effective_model);
|
||||
let request_url = crate::provider_transport::build_grok_upstream_url(
|
||||
&transport,
|
||||
crate::provider_transport::GROK_CHAT_PATH,
|
||||
);
|
||||
let provider_request_body = crate::provider_transport::build_grok_app_chat_body(
|
||||
client_api_format,
|
||||
Some(request_model),
|
||||
&request_body,
|
||||
);
|
||||
let report_context = json!({
|
||||
"provider_type": provider.provider_type,
|
||||
"provider_api_format": provider_api_format,
|
||||
"client_api_format": client_api_format,
|
||||
"model": request_model,
|
||||
"mapped_model": candidate.effective_model,
|
||||
"request_path": route_path,
|
||||
"request_body": request_body,
|
||||
});
|
||||
let transport_profile = state.resolve_transport_profile(&transport);
|
||||
let Some(request_headers) = crate::provider_transport::build_grok_browser_headers(
|
||||
crate::provider_transport::GrokHeaderInput {
|
||||
transport: &transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(&parts.headers),
|
||||
content_type: "application/json",
|
||||
accept: "text/event-stream",
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: &request_body,
|
||||
},
|
||||
) else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "failed",
|
||||
skip_reason: None,
|
||||
error_message: Some("provider request headers build failed".to_string()),
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url,
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body: provider_request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
};
|
||||
|
||||
let plan = ExecutionPlan {
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: Some(format!("provider-query-{}", candidate.key.id)),
|
||||
provider_name: Some(provider.name.clone()),
|
||||
provider_id: provider.id.clone(),
|
||||
endpoint_id: candidate.endpoint.id.clone(),
|
||||
key_id: candidate.key.id.clone(),
|
||||
method: "POST".to_string(),
|
||||
url: request_url.clone(),
|
||||
headers: request_headers.clone(),
|
||||
content_type: Some("application/json".to_string()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(request_body.clone()),
|
||||
stream: true,
|
||||
client_api_format: client_api_format.to_string(),
|
||||
provider_api_format: provider_api_format.clone(),
|
||||
model_name: Some(request_model.to_string()),
|
||||
proxy: state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
|
||||
.await,
|
||||
transport_profile,
|
||||
timeouts: state.resolve_transport_execution_timeouts(&transport),
|
||||
};
|
||||
|
||||
let result = match state
|
||||
.execute_execution_runtime_sync_plan_with_report_context(
|
||||
Some(trace_id),
|
||||
&plan,
|
||||
Some(&report_context),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(err) => {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "failed",
|
||||
skip_reason: None,
|
||||
error_message: Some(format!("model test execution failed: {err:?}")),
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url,
|
||||
request_headers,
|
||||
request_body: provider_request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
}
|
||||
};
|
||||
let response_body = result.body.as_ref().and_then(|body| body.json_body.clone());
|
||||
let did_fail = result.status_code >= 400 || response_body.is_none();
|
||||
let error_message = if did_fail {
|
||||
provider_query_extract_error_message(&result).or_else(|| {
|
||||
response_body.is_none().then(|| {
|
||||
format!(
|
||||
"Provider returned HTTP {} without a model-test response body",
|
||||
result.status_code
|
||||
)
|
||||
})
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
Ok(ProviderQueryExecutionOutcome {
|
||||
status: if did_fail { "failed" } else { "success" },
|
||||
skip_reason: None,
|
||||
error_message,
|
||||
status_code: Some(result.status_code),
|
||||
latency_ms: result.telemetry.as_ref().and_then(|value| value.elapsed_ms),
|
||||
request_url,
|
||||
request_headers,
|
||||
request_body: provider_request_body,
|
||||
response_headers: result.headers,
|
||||
response_body,
|
||||
})
|
||||
}
|
||||
|
||||
async fn provider_query_execute_standard_test_candidate(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -2132,22 +2601,25 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
"Provider transport snapshot is unavailable",
|
||||
));
|
||||
};
|
||||
let original_request_body = provider_query_build_test_request_body_for_route(
|
||||
let provider_api_format = candidate.endpoint.api_format.as_str();
|
||||
let normalized_provider_api_format =
|
||||
crate::ai_serving::normalize_api_format_alias(provider_api_format);
|
||||
let client_api_format =
|
||||
provider_query_standard_test_client_api_format(normalized_provider_api_format.as_str());
|
||||
let original_request_body = provider_query_build_test_request_body_for_api_format(
|
||||
payload,
|
||||
&candidate.effective_model,
|
||||
route_path,
|
||||
client_api_format,
|
||||
);
|
||||
if !provider_query_transport_supports_model_test_execution(
|
||||
state,
|
||||
&transport,
|
||||
candidate.endpoint.api_format.as_str(),
|
||||
provider_api_format,
|
||||
) {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
original_request_body,
|
||||
provider_query_standard_test_unsupported_reason(
|
||||
&transport,
|
||||
candidate.endpoint.api_format.as_str(),
|
||||
),
|
||||
provider_query_standard_test_unsupported_reason(&transport, provider_api_format),
|
||||
));
|
||||
}
|
||||
|
||||
@@ -2159,11 +2631,6 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
let request_model =
|
||||
provider_query_request_body_model(&request_body, &candidate.effective_model);
|
||||
|
||||
let provider_api_format = candidate.endpoint.api_format.as_str();
|
||||
let normalized_provider_api_format =
|
||||
crate::ai_serving::normalize_api_format_alias(provider_api_format);
|
||||
let client_api_format =
|
||||
provider_query_standard_test_client_api_format(normalized_provider_api_format.as_str());
|
||||
let upstream_is_stream = provider_query_resolve_standard_test_upstream_is_stream(
|
||||
transport.endpoint.config.as_ref(),
|
||||
transport.provider.provider_type.as_str(),
|
||||
@@ -2229,12 +2696,20 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
}
|
||||
"openai:responses" | "openai:responses:compact" => {
|
||||
let Some(mut provider_request_body) =
|
||||
crate::ai_serving::build_cross_format_openai_chat_request_body(
|
||||
&request_body,
|
||||
request_model,
|
||||
normalized_provider_api_format.as_str(),
|
||||
upstream_is_stream,
|
||||
)
|
||||
(if provider_query_request_body_is_openai_responses_shape(&request_body) {
|
||||
crate::ai_serving::build_local_openai_responses_request_body(
|
||||
&request_body,
|
||||
request_model,
|
||||
upstream_is_stream,
|
||||
)
|
||||
} else {
|
||||
crate::ai_serving::build_cross_format_openai_chat_request_body(
|
||||
&request_body,
|
||||
request_model,
|
||||
normalized_provider_api_format.as_str(),
|
||||
upstream_is_stream,
|
||||
)
|
||||
})
|
||||
else {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
request_body.clone(),
|
||||
@@ -2267,7 +2742,7 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
}
|
||||
"openai:embedding" | "gemini:embedding" | "jina:embedding" | "doubao:embedding"
|
||||
| "openai:rerank" | "jina:rerank" => {
|
||||
let Some(provider_request_body) =
|
||||
let Some(mut provider_request_body) =
|
||||
crate::ai_serving::build_standard_request_body_with_model_directives_and_request_headers(
|
||||
&request_body,
|
||||
client_api_format,
|
||||
@@ -2287,6 +2762,18 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
format!("Provider request body could not be built for {provider_api_format}"),
|
||||
));
|
||||
};
|
||||
if let Err(err) = crate::provider_transport::apply_transport_request_body_semantics(
|
||||
&mut provider_request_body,
|
||||
&transport,
|
||||
normalized_provider_api_format.as_str(),
|
||||
) {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
provider_request_body,
|
||||
format!(
|
||||
"Provider request body is not compatible with transport semantics: {err}"
|
||||
),
|
||||
));
|
||||
}
|
||||
provider_request_body
|
||||
}
|
||||
_ => {
|
||||
@@ -2367,7 +2854,7 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
*synthetic_request.headers_mut() = incoming_request_headers;
|
||||
let (parts, _) = synthetic_request.into_parts();
|
||||
|
||||
let request_url = crate::provider_transport::build_transport_request_url(
|
||||
let request_url = crate::provider_transport::build_transport_request_url_for_request_body(
|
||||
&transport,
|
||||
crate::provider_transport::TransportRequestUrlParams {
|
||||
provider_api_format,
|
||||
@@ -2376,6 +2863,7 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
request_query: parts.uri.query(),
|
||||
kiro_api_region: None,
|
||||
},
|
||||
Some(&provider_request_body),
|
||||
);
|
||||
let Some(request_url) = request_url else {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
@@ -2659,6 +3147,12 @@ async fn build_admin_provider_query_kiro_failover_response(
|
||||
)
|
||||
.await
|
||||
}
|
||||
Some(ProviderQueryTestAdapter::Grok) => {
|
||||
provider_query_execute_grok_test_candidate(
|
||||
state, &provider, candidate, payload, route_path, &trace_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
Some(ProviderQueryTestAdapter::Standard) => {
|
||||
provider_query_execute_standard_test_candidate(
|
||||
state, &provider, candidate, payload, route_path, &trace_id,
|
||||
@@ -2718,7 +3212,10 @@ async fn build_admin_provider_query_kiro_failover_response(
|
||||
));
|
||||
if is_success {
|
||||
success_body = response_body;
|
||||
success_stream = matches!(adapter, Some(ProviderQueryTestAdapter::Kiro));
|
||||
success_stream = matches!(
|
||||
adapter,
|
||||
Some(ProviderQueryTestAdapter::Kiro | ProviderQueryTestAdapter::Grok)
|
||||
);
|
||||
winning_candidate_index = Some(candidate_index);
|
||||
break;
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ use serde_json::{json, Value};
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) enum ProviderQueryTestAdapter {
|
||||
Standard,
|
||||
Grok,
|
||||
Kiro,
|
||||
OpenAiImage,
|
||||
Antigravity,
|
||||
@@ -62,10 +63,10 @@ pub(super) fn provider_query_standard_test_unsupported_reason(
|
||||
api_format,
|
||||
)
|
||||
}
|
||||
"gemini:generate_content"
|
||||
if crate::provider_transport::is_vertex_api_key_transport_context(transport) =>
|
||||
"gemini:generate_content" | "gemini:embedding"
|
||||
if crate::provider_transport::is_vertex_transport_context(transport) =>
|
||||
{
|
||||
aether_provider_transport::vertex::local_vertex_api_key_gemini_transport_unsupported_reason_with_network(
|
||||
aether_provider_transport::vertex::local_vertex_gemini_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
)
|
||||
}
|
||||
@@ -134,6 +135,65 @@ pub(super) fn provider_query_antigravity_test_unsupported_reason(
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_grok_test_unsupported_reason(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
if !transport.provider.is_active {
|
||||
return Some("provider_inactive");
|
||||
}
|
||||
if !transport.endpoint.is_active {
|
||||
return Some("endpoint_inactive");
|
||||
}
|
||||
if !transport.key.is_active {
|
||||
return Some("key_inactive");
|
||||
}
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok")
|
||||
{
|
||||
return Some("transport_provider_type_unsupported");
|
||||
}
|
||||
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
|
||||
if !matches!(
|
||||
normalized_api_format.as_str(),
|
||||
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
|
||||
) {
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if provider_query_normalize_api_format_alias(&transport.endpoint.api_format)
|
||||
!= normalized_api_format
|
||||
{
|
||||
return Some("transport_api_format_mismatch");
|
||||
}
|
||||
if crate::provider_transport::resolve_grok_session_auth(transport).is_none() {
|
||||
return Some("transport_oauth_resolution_unsupported");
|
||||
}
|
||||
if !crate::provider_transport::header_rules_are_locally_supported(
|
||||
transport.endpoint.header_rules.as_ref(),
|
||||
) {
|
||||
return Some("transport_header_rules_unsupported");
|
||||
}
|
||||
if !crate::provider_transport::body_rules_are_locally_supported(
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
) {
|
||||
return Some("transport_body_rules_unsupported");
|
||||
}
|
||||
if !crate::provider_transport::transport_proxy_is_locally_supported(transport) {
|
||||
return Some("transport_proxy_unsupported");
|
||||
}
|
||||
if crate::provider_transport::transport_profile_is_configured(transport)
|
||||
&& crate::provider_transport::resolve_transport_profile(transport).is_none()
|
||||
{
|
||||
return Some("transport_profile_unsupported");
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_normalize_api_format_alias(value: &str) -> String {
|
||||
crate::ai_serving::normalize_api_format_alias(value)
|
||||
}
|
||||
@@ -147,6 +207,15 @@ pub(super) fn provider_query_test_adapter_for_provider_api_format(
|
||||
}
|
||||
|
||||
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
|
||||
if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
return match normalized_api_format.as_str() {
|
||||
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages" => {
|
||||
Some(ProviderQueryTestAdapter::Grok)
|
||||
}
|
||||
"openai:image" => Some(ProviderQueryTestAdapter::OpenAiImage),
|
||||
_ => None,
|
||||
};
|
||||
}
|
||||
if normalized_api_format == "openai:image" {
|
||||
return Some(ProviderQueryTestAdapter::OpenAiImage);
|
||||
}
|
||||
@@ -182,6 +251,16 @@ pub(super) fn provider_query_model_test_endpoint_priority(
|
||||
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
|
||||
match provider_query_test_adapter_for_provider_api_format(provider_type, api_format)? {
|
||||
ProviderQueryTestAdapter::Kiro => Some(0),
|
||||
ProviderQueryTestAdapter::Grok => {
|
||||
if matches!(
|
||||
normalized_api_format.as_str(),
|
||||
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
|
||||
) {
|
||||
Some(0)
|
||||
} else {
|
||||
Some(2)
|
||||
}
|
||||
}
|
||||
ProviderQueryTestAdapter::Antigravity => Some(1),
|
||||
ProviderQueryTestAdapter::OpenAiImage => Some(2),
|
||||
ProviderQueryTestAdapter::Standard => {
|
||||
@@ -234,6 +313,9 @@ pub(super) fn provider_query_transport_supports_model_test_execution(
|
||||
)
|
||||
.is_none()
|
||||
}
|
||||
Some(ProviderQueryTestAdapter::Grok) => {
|
||||
provider_query_grok_test_unsupported_reason(transport, api_format).is_none()
|
||||
}
|
||||
Some(ProviderQueryTestAdapter::Standard) => match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"openai:chat" => {
|
||||
crate::provider_transport::policy::supports_local_openai_chat_transport(transport)
|
||||
|
||||
+50
@@ -0,0 +1,50 @@
|
||||
use crate::handlers::admin::provider::shared::model_test_capabilities::{
|
||||
admin_provider_openai_image_normalize_options, admin_provider_openai_image_test_capability,
|
||||
AdminProviderOpenAiImageTestCapability,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) struct ProviderQueryOpenAiImageTestCapability(AdminProviderOpenAiImageTestCapability);
|
||||
|
||||
pub(super) fn provider_query_openai_image_test_capability(
|
||||
provider_type: &str,
|
||||
) -> ProviderQueryOpenAiImageTestCapability {
|
||||
ProviderQueryOpenAiImageTestCapability(admin_provider_openai_image_test_capability(
|
||||
provider_type,
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_openai_image_normalize_options(
|
||||
provider_type: &str,
|
||||
) -> crate::ai_serving::OpenAiImageNormalizeOptions {
|
||||
admin_provider_openai_image_normalize_options(provider_type)
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_openai_image_requested_count(request_body: &Value) -> Option<u64> {
|
||||
request_body.get("n").and_then(|value| {
|
||||
value.as_u64().or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| value.parse::<u64>().ok())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_openai_image_normalize_failure_message(
|
||||
provider_type: &str,
|
||||
request_body: &Value,
|
||||
) -> String {
|
||||
let capability = provider_query_openai_image_test_capability(provider_type);
|
||||
if provider_query_openai_image_requested_count(request_body)
|
||||
.is_some_and(|value| !capability.0.supports_generation_count(value))
|
||||
{
|
||||
return format!(
|
||||
"Provider request body could not be normalized for openai:image: selected provider supports n=1..{} for generation",
|
||||
capability.0.max_generation_count
|
||||
);
|
||||
}
|
||||
"Provider request body could not be normalized for openai:image".to_string()
|
||||
}
|
||||
+209
-2
@@ -1,3 +1,5 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::super::provider_query_key_display_name;
|
||||
use super::{ProviderQueryExecutionOutcome, ProviderQueryTestCandidate};
|
||||
use serde_json::{json, Value};
|
||||
@@ -7,11 +9,29 @@ pub(super) fn provider_query_test_attempt_payload(
|
||||
candidate: &ProviderQueryTestCandidate,
|
||||
execution: &ProviderQueryExecutionOutcome,
|
||||
) -> Value {
|
||||
let endpoint_route = provider_query_endpoint_route_payload(candidate, execution);
|
||||
let endpoint_product = endpoint_route
|
||||
.get("product")
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
let endpoint_variant = endpoint_route
|
||||
.get("variant")
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
let endpoint_action = endpoint_route.get("action").cloned().unwrap_or(Value::Null);
|
||||
let endpoint_batch_strategy = endpoint_route
|
||||
.get("batch_strategy")
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null);
|
||||
json!({
|
||||
"candidate_index": candidate_index,
|
||||
"retry_index": 0,
|
||||
"endpoint_api_format": candidate.endpoint.api_format,
|
||||
"endpoint_base_url": candidate.endpoint.base_url,
|
||||
"endpoint_product": endpoint_product,
|
||||
"endpoint_variant": endpoint_variant,
|
||||
"endpoint_action": endpoint_action,
|
||||
"endpoint_batch_strategy": endpoint_batch_strategy,
|
||||
"key_name": provider_query_key_display_name(&candidate.key),
|
||||
"key_id": candidate.key.id,
|
||||
"auth_type": candidate.key.auth_type,
|
||||
@@ -22,13 +42,166 @@ pub(super) fn provider_query_test_attempt_payload(
|
||||
"status_code": execution.status_code,
|
||||
"latency_ms": execution.latency_ms,
|
||||
"request_url": execution.request_url,
|
||||
"request_headers": execution.request_headers,
|
||||
"request_headers": provider_query_redact_diagnostic_headers(&execution.request_headers),
|
||||
"request_body": execution.request_body,
|
||||
"response_headers": execution.response_headers,
|
||||
"response_headers": provider_query_redact_diagnostic_headers(&execution.response_headers),
|
||||
"response_body": execution.response_body,
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_endpoint_route_payload(
|
||||
candidate: &ProviderQueryTestCandidate,
|
||||
execution: &ProviderQueryExecutionOutcome,
|
||||
) -> Value {
|
||||
let api_format = crate::ai_serving::normalize_api_format_alias(&candidate.endpoint.api_format);
|
||||
let request_url = execution.request_url.to_ascii_lowercase();
|
||||
let base_url = candidate.endpoint.base_url.to_ascii_lowercase();
|
||||
let is_vertex = request_url.contains("aiplatform.googleapis.com")
|
||||
|| base_url.contains("aiplatform.googleapis.com");
|
||||
let is_gemini_api = request_url.contains("generativelanguage.googleapis.com")
|
||||
|| base_url.contains("generativelanguage.googleapis.com");
|
||||
let is_openai_compat =
|
||||
request_url.contains("/endpoints/openapi") || request_url.contains("/openai/");
|
||||
let is_batch = execution
|
||||
.request_body
|
||||
.get("requests")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|items| !items.is_empty());
|
||||
let vertex_instance_count = execution
|
||||
.request_body
|
||||
.get("instances")
|
||||
.and_then(Value::as_array)
|
||||
.map(Vec::len)
|
||||
.unwrap_or(0);
|
||||
|
||||
let (product, variant, action, batch_strategy) = match api_format.as_str() {
|
||||
"gemini:embedding" if is_vertex => (
|
||||
"Vertex AI",
|
||||
"vertex_native",
|
||||
"predict",
|
||||
if vertex_instance_count > 1 {
|
||||
"predict_instances"
|
||||
} else {
|
||||
"single_instance"
|
||||
},
|
||||
),
|
||||
"gemini:embedding" if is_gemini_api => (
|
||||
"Gemini API",
|
||||
"gemini_native",
|
||||
if is_batch {
|
||||
"batchEmbedContents"
|
||||
} else {
|
||||
"embedContent"
|
||||
},
|
||||
if is_batch {
|
||||
"native_batch"
|
||||
} else {
|
||||
"single_native"
|
||||
},
|
||||
),
|
||||
"gemini:embedding" => (
|
||||
"Gemini native",
|
||||
"gemini_native",
|
||||
if is_batch {
|
||||
"batchEmbedContents"
|
||||
} else {
|
||||
"embedContent"
|
||||
},
|
||||
if is_batch {
|
||||
"native_batch"
|
||||
} else {
|
||||
"single_native"
|
||||
},
|
||||
),
|
||||
"gemini:generate_content" if is_vertex => {
|
||||
("Vertex AI", "vertex_native", "generateContent", "")
|
||||
}
|
||||
"gemini:generate_content" if is_gemini_api => {
|
||||
("Gemini API", "gemini_native", "generateContent", "")
|
||||
}
|
||||
"gemini:generate_content" => ("Gemini native", "gemini_native", "generateContent", ""),
|
||||
"openai:embedding" if is_vertex && is_openai_compat => (
|
||||
"Vertex AI OpenAI-compatible",
|
||||
"openai_compatible",
|
||||
"embeddings",
|
||||
"openai_batch",
|
||||
),
|
||||
"openai:embedding" if is_gemini_api && is_openai_compat => (
|
||||
"Gemini API OpenAI-compatible",
|
||||
"openai_compatible",
|
||||
"embeddings",
|
||||
"openai_batch",
|
||||
),
|
||||
"openai:embedding" => (
|
||||
"OpenAI-compatible",
|
||||
"openai_compatible",
|
||||
"embeddings",
|
||||
"openai_batch",
|
||||
),
|
||||
"openai:chat" if is_vertex && is_openai_compat => (
|
||||
"Vertex AI OpenAI-compatible",
|
||||
"openai_compatible",
|
||||
"chat/completions",
|
||||
"",
|
||||
),
|
||||
"openai:chat" if is_gemini_api && is_openai_compat => (
|
||||
"Gemini API OpenAI-compatible",
|
||||
"openai_compatible",
|
||||
"chat/completions",
|
||||
"",
|
||||
),
|
||||
"openai:chat" => (
|
||||
"OpenAI-compatible",
|
||||
"openai_compatible",
|
||||
"chat/completions",
|
||||
"",
|
||||
),
|
||||
_ => (
|
||||
"Provider endpoint",
|
||||
"provider_native",
|
||||
"provider_request",
|
||||
"",
|
||||
),
|
||||
};
|
||||
|
||||
json!({
|
||||
"product": product,
|
||||
"variant": variant,
|
||||
"action": action,
|
||||
"batch_strategy": batch_strategy,
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_redact_diagnostic_headers(
|
||||
headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
headers
|
||||
.iter()
|
||||
.map(|(name, value)| {
|
||||
if provider_query_header_is_sensitive(name) {
|
||||
(name.clone(), "<redacted>".to_string())
|
||||
} else {
|
||||
(name.clone(), value.clone())
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn provider_query_header_is_sensitive(name: &str) -> bool {
|
||||
matches!(
|
||||
name.trim().to_ascii_lowercase().as_str(),
|
||||
"authorization"
|
||||
| "proxy-authorization"
|
||||
| "cookie"
|
||||
| "set-cookie"
|
||||
| "x-api-key"
|
||||
| "api-key"
|
||||
| "x-goog-api-key"
|
||||
| "anthropic-api-key"
|
||||
| "openai-api-key"
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_candidate_summary_payload(
|
||||
total_candidates: usize,
|
||||
total_attempts: usize,
|
||||
@@ -133,3 +306,37 @@ pub(super) fn provider_query_candidate_summary_payload(
|
||||
.unwrap_or(Value::Null),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn provider_query_diagnostic_headers_redact_credentials() {
|
||||
let headers = BTreeMap::from([
|
||||
("cookie".to_string(), "sso=secret".to_string()),
|
||||
("authorization".to_string(), "Bearer secret".to_string()),
|
||||
("x-goog-api-key".to_string(), "secret".to_string()),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
]);
|
||||
|
||||
let redacted = provider_query_redact_diagnostic_headers(&headers);
|
||||
|
||||
assert_eq!(
|
||||
redacted.get("cookie").map(String::as_str),
|
||||
Some("<redacted>")
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.get("authorization").map(String::as_str),
|
||||
Some("<redacted>")
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.get("x-goog-api-key").map(String::as_str),
|
||||
Some("<redacted>")
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.get("content-type").map(String::as_str),
|
||||
Some("application/json")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,68 @@
|
||||
use super::*;
|
||||
use crate::handlers::admin::request::AdminGatewayProviderTransportSnapshot;
|
||||
use serde_json::json;
|
||||
|
||||
fn sample_openai_image_transport(provider_type: &str) -> AdminGatewayProviderTransportSnapshot {
|
||||
AdminGatewayProviderTransportSnapshot {
|
||||
provider: crate::provider_transport::snapshot::GatewayProviderTransportProvider {
|
||||
id: "provider-1".to_string(),
|
||||
name: "Provider".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: crate::provider_transport::snapshot::GatewayProviderTransportEndpoint {
|
||||
id: "endpoint-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
api_format: "openai:image".to_string(),
|
||||
api_family: None,
|
||||
endpoint_kind: None,
|
||||
is_active: true,
|
||||
base_url: "https://grok.com/".to_string(),
|
||||
header_rules: None,
|
||||
body_rules: None,
|
||||
max_retries: None,
|
||||
custom_path: None,
|
||||
config: None,
|
||||
format_acceptance_config: None,
|
||||
proxy: None,
|
||||
},
|
||||
key: crate::provider_transport::snapshot::GatewayProviderTransportKey {
|
||||
id: "key-1".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
name: "key".to_string(),
|
||||
auth_type: "oauth".to_string(),
|
||||
is_active: true,
|
||||
api_formats: None,
|
||||
auth_type_by_format: None,
|
||||
allow_auth_channel_mismatch_formats: None,
|
||||
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: String::new(),
|
||||
decrypted_auth_config: Some(
|
||||
json!({
|
||||
"sso_token": "abc",
|
||||
"sso_rw_token": "rw"
|
||||
})
|
||||
.to_string(),
|
||||
),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_preserves_custom_model() {
|
||||
let payload = json!({
|
||||
@@ -28,6 +90,52 @@ fn provider_query_test_request_body_defaults_missing_model() {
|
||||
assert_eq!(body["model"], json!("fallback-model"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_fills_empty_conversation() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "custom-upstream-model",
|
||||
"messages": []
|
||||
}
|
||||
});
|
||||
|
||||
let body = provider_query_build_test_request_body(&payload, "fallback-model");
|
||||
|
||||
assert_eq!(body["model"], json!("custom-upstream-model"));
|
||||
assert_eq!(
|
||||
body["messages"],
|
||||
json!([{ "role": "user", "content": DEFAULT_PROVIDER_QUERY_TEST_MESSAGE }])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_keeps_non_empty_conversation() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "custom-upstream-model",
|
||||
"messages": [{ "role": "user", "content": "custom prompt" }]
|
||||
}
|
||||
});
|
||||
|
||||
let body = provider_query_build_test_request_body(&payload, "fallback-model");
|
||||
|
||||
assert_eq!(
|
||||
body["messages"],
|
||||
json!([{ "role": "user", "content": "custom prompt" }])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_default_test_request_body_does_not_set_max_tokens() {
|
||||
let body = provider_query_build_test_request_body(&json!({}), "fallback-model");
|
||||
|
||||
assert_eq!(body["model"], json!("fallback-model"));
|
||||
assert!(
|
||||
body.get("max_tokens").is_none(),
|
||||
"admin model test must not silently force a low max_tokens value"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_failover_request_body_overrides_custom_model() {
|
||||
let payload = json!({
|
||||
@@ -168,6 +276,88 @@ fn provider_query_standard_test_aggregates_responses_stream_body() {
|
||||
assert_eq!(body["output"][0]["content"][0]["text"], json!("Hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_aggregates_responses_image_generation_call() {
|
||||
let stream_body = concat!(
|
||||
"event: response.created\n",
|
||||
"data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_img_123\",\"object\":\"response\",\"model\":\"gpt-5.4-mini\",\"status\":\"in_progress\",\"output\":[]}}\n\n",
|
||||
"event: response.output_item.done\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"ig_123\",\"type\":\"image_generation_call\",\"status\":\"completed\",\"output_format\":\"png\",\"result\":\"aGVsbG8=\"}}\n\n",
|
||||
"event: response.completed\n",
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_img_123\",\"object\":\"response\",\"model\":\"gpt-5.4-mini\",\"status\":\"completed\",\"output\":[]}}\n\n",
|
||||
);
|
||||
let result = aether_contracts::ExecutionResult {
|
||||
request_id: "provider-test".to_string(),
|
||||
candidate_id: Some("candidate-0".to_string()),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(
|
||||
base64::engine::general_purpose::STANDARD.encode(stream_body.as_bytes()),
|
||||
),
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
};
|
||||
|
||||
let body = provider_query_standard_execution_response_body("openai:responses", &result)
|
||||
.expect("responses image stream body should aggregate");
|
||||
|
||||
assert_eq!(body["output"][0]["type"], json!("image_generation_call"));
|
||||
assert_eq!(body["output"][0]["result"], json!("aGVsbG8="));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_responses_test_request_body_defaults_to_responses_input() {
|
||||
let payload = json!({"message": "hello from responses"});
|
||||
|
||||
let body = provider_query_build_test_request_body_for_api_format(
|
||||
&payload,
|
||||
"gpt-5.4-mini",
|
||||
"/api/admin/provider-query/test-model",
|
||||
"openai:responses",
|
||||
);
|
||||
|
||||
assert_eq!(body["model"], json!("gpt-5.4-mini"));
|
||||
assert_eq!(body["input"], json!("hello from responses"));
|
||||
assert!(body.get("messages").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_rejects_gemini_success_without_visible_output() {
|
||||
let result = aether_contracts::ExecutionResult {
|
||||
request_id: "provider-test".to_string(),
|
||||
candidate_id: Some("candidate-0".to_string()),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(aether_contracts::ResponseBody {
|
||||
json_body: Some(json!({
|
||||
"candidates": [{
|
||||
"content": {"role": "model"},
|
||||
"finishReason": "MAX_TOKENS"
|
||||
}],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 8,
|
||||
"candidatesTokenCount": 1,
|
||||
"thoughtsTokenCount": 25,
|
||||
"totalTokenCount": 34
|
||||
},
|
||||
"modelVersion": "gemini-3-flash-preview",
|
||||
"responseId": "resp-empty"
|
||||
})),
|
||||
body_bytes_b64: None,
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
};
|
||||
|
||||
assert!(
|
||||
provider_query_standard_execution_response_body("gemini:generate_content", &result)
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_adapter_routes_fixed_provider_endpoint_types() {
|
||||
assert_eq!(
|
||||
@@ -204,6 +394,22 @@ fn provider_query_test_adapter_routes_fixed_provider_endpoint_types() {
|
||||
),
|
||||
Some(ProviderQueryTestAdapter::Antigravity)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("grok", "openai:chat"),
|
||||
Some(ProviderQueryTestAdapter::Grok)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("grok", "openai:responses"),
|
||||
Some(ProviderQueryTestAdapter::Grok)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("grok", "claude:messages"),
|
||||
Some(ProviderQueryTestAdapter::Grok)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("grok", "openai:image"),
|
||||
Some(ProviderQueryTestAdapter::OpenAiImage)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_test_adapter_for_provider_api_format("custom", "openai:embedding"),
|
||||
Some(ProviderQueryTestAdapter::Standard)
|
||||
@@ -244,12 +450,154 @@ fn provider_query_endpoint_priority_prefers_text_before_cli_and_image() {
|
||||
provider_query_model_test_endpoint_priority("chatgpt_web", "openai:image"),
|
||||
Some(2)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_model_test_endpoint_priority("grok", "openai:chat"),
|
||||
Some(0)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_model_test_endpoint_priority("grok", "openai:responses"),
|
||||
Some(0)
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_model_test_endpoint_priority("antigravity", "gemini:generate_content"),
|
||||
Some(1)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_model_test_body_maps_non_reasoning_model_to_fast_mode() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "grok-4.20-0309-non-reasoning",
|
||||
"messages": [
|
||||
{"role": "system", "content": "be concise"},
|
||||
{"role": "user", "content": "hello"}
|
||||
]
|
||||
}
|
||||
});
|
||||
let request_body = provider_query_build_test_request_body_for_route(
|
||||
&payload,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
"/api/admin/provider-query/test-model",
|
||||
);
|
||||
|
||||
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
|
||||
"openai:chat",
|
||||
Some(provider_query_request_body_model(
|
||||
&request_body,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
)),
|
||||
&request_body,
|
||||
);
|
||||
|
||||
assert_eq!(upstream_body["modeId"], json!("fast"));
|
||||
assert_eq!(
|
||||
upstream_body["message"],
|
||||
json!("[system]: be concise\n\n[user]: hello")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_model_test_uses_responses_client_body_for_responses_endpoint() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "grok-4.20-0309-non-reasoning",
|
||||
"input": "hello from responses body"
|
||||
}
|
||||
});
|
||||
let request_body = provider_query_build_grok_test_request_body_for_api_format(
|
||||
&payload,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
"/api/admin/provider-query/test-model",
|
||||
"openai:responses",
|
||||
);
|
||||
|
||||
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
|
||||
provider_query_grok_test_client_api_format("openai:responses"),
|
||||
Some(provider_query_request_body_model(
|
||||
&request_body,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
)),
|
||||
&request_body,
|
||||
);
|
||||
|
||||
assert_eq!(upstream_body["modeId"], json!("fast"));
|
||||
assert_eq!(upstream_body["message"], json!("hello from responses body"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_model_test_uses_responses_input_when_existing_body_has_messages() {
|
||||
let payload = json!({
|
||||
"request_body": {
|
||||
"model": "grok-4.20-0309-non-reasoning",
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": "hello from stale chat body"
|
||||
}]
|
||||
}
|
||||
});
|
||||
let request_body = provider_query_build_grok_test_request_body_for_api_format(
|
||||
&payload,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
"/api/admin/provider-query/test-model",
|
||||
"openai:responses",
|
||||
);
|
||||
|
||||
assert_eq!(request_body["model"], json!("grok-4.20-0309-non-reasoning"));
|
||||
assert_eq!(
|
||||
request_body["input"],
|
||||
json!("Hello! This is a test message.")
|
||||
);
|
||||
assert!(request_body.get("messages").is_some());
|
||||
|
||||
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
|
||||
provider_query_grok_test_client_api_format("openai:responses"),
|
||||
Some(provider_query_request_body_model(
|
||||
&request_body,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
)),
|
||||
&request_body,
|
||||
);
|
||||
|
||||
assert_eq!(upstream_body["modeId"], json!("fast"));
|
||||
assert_eq!(
|
||||
upstream_body["message"],
|
||||
json!("Hello! This is a test message.")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_model_test_defaults_claude_messages_body_for_claude_endpoint() {
|
||||
let payload = json!({});
|
||||
let request_body = provider_query_build_grok_test_request_body_for_api_format(
|
||||
&payload,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
"/api/admin/provider-query/test-model",
|
||||
"claude:messages",
|
||||
);
|
||||
|
||||
assert_eq!(request_body["model"], json!("grok-4.20-0309-non-reasoning"));
|
||||
assert_eq!(
|
||||
request_body["messages"],
|
||||
json!([{ "role": "user", "content": DEFAULT_PROVIDER_QUERY_TEST_MESSAGE }])
|
||||
);
|
||||
|
||||
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
|
||||
provider_query_grok_test_client_api_format("claude:messages"),
|
||||
Some(provider_query_request_body_model(
|
||||
&request_body,
|
||||
"grok-4.20-0309-non-reasoning",
|
||||
)),
|
||||
&request_body,
|
||||
);
|
||||
|
||||
assert_eq!(upstream_body["modeId"], json!("fast"));
|
||||
assert_eq!(
|
||||
upstream_body["message"],
|
||||
json!("[user]: Hello! This is a test message.")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_candidate_summary_marks_unused_after_first_success() {
|
||||
let attempts = vec![json!({
|
||||
@@ -360,3 +708,76 @@ fn provider_query_failover_image_test_request_body_overrides_model() {
|
||||
|
||||
assert_eq!(body["model"], json!("new-image-model"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_image_test_allows_multi_generation_count() {
|
||||
let request = http::Request::builder()
|
||||
.uri("/v1/images/generations")
|
||||
.body(())
|
||||
.expect("request should build");
|
||||
let (parts, _) = request.into_parts();
|
||||
let body = json!({
|
||||
"model": "grok-imagine-image",
|
||||
"prompt": "draw",
|
||||
"n": 2
|
||||
});
|
||||
|
||||
let normalized = crate::ai_serving::normalize_openai_image_request_with_options(
|
||||
&parts,
|
||||
&body,
|
||||
None,
|
||||
provider_query_openai_image_normalize_options("grok"),
|
||||
)
|
||||
.expect("grok image model tests should allow multi-image generation");
|
||||
let provider_body = crate::ai_serving::build_openai_image_provider_request_body(&normalized);
|
||||
|
||||
assert_eq!(provider_body["n"], json!(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_grok_image_test_uses_grok_app_chat_upstream_url() {
|
||||
let transport = sample_openai_image_transport("grok");
|
||||
|
||||
assert_eq!(
|
||||
provider_query_openai_image_test_upstream_url(&transport, Some("trace=1")),
|
||||
"https://grok.com/rest/app-chat/conversations/new"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_chatgpt_web_image_test_uses_internal_upstream_url() {
|
||||
let transport = sample_openai_image_transport("chatgpt_web");
|
||||
|
||||
assert_eq!(
|
||||
provider_query_openai_image_test_upstream_url(&transport, Some("trace=1")),
|
||||
"https://grok.com/__aether/chatgpt-web-image"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_non_grok_image_test_keeps_single_generation_boundary() {
|
||||
let request = http::Request::builder()
|
||||
.uri("/v1/images/generations")
|
||||
.body(())
|
||||
.expect("request should build");
|
||||
let (parts, _) = request.into_parts();
|
||||
let body = json!({
|
||||
"model": "gpt-image-2",
|
||||
"prompt": "draw",
|
||||
"n": 2
|
||||
});
|
||||
|
||||
assert!(
|
||||
crate::ai_serving::normalize_openai_image_request_with_options(
|
||||
&parts,
|
||||
&body,
|
||||
None,
|
||||
provider_query_openai_image_normalize_options("chatgpt_web"),
|
||||
)
|
||||
.is_none()
|
||||
);
|
||||
assert_eq!(
|
||||
provider_query_openai_image_normalize_failure_message("chatgpt_web", &body),
|
||||
"Provider request body could not be normalized for openai:image: selected provider supports n=1..1 for generation"
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
pub(crate) mod model_test_capabilities;
|
||||
pub(crate) mod paths;
|
||||
pub(crate) mod payloads;
|
||||
pub(crate) mod support;
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
use crate::image_capabilities::{
|
||||
openai_image_normalize_options_for_provider, openai_image_provider_max_generation_count,
|
||||
};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
const GROK_IMAGE_MODEL_IDS: &[&str] = &[
|
||||
"grok-imagine-image-lite",
|
||||
"grok-imagine-image",
|
||||
"grok-imagine-image-pro",
|
||||
"grok-imagine-image-edit",
|
||||
];
|
||||
const GROK_IMAGE_EDIT_MODEL_ID: &str = "grok-imagine-image-edit";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) struct AdminProviderOpenAiImageTestCapability {
|
||||
pub(crate) max_generation_count: u64,
|
||||
}
|
||||
|
||||
impl AdminProviderOpenAiImageTestCapability {
|
||||
pub(crate) fn supports_generation_count(self, count: u64) -> bool {
|
||||
count >= 1 && count <= self.max_generation_count
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_openai_image_test_capability(
|
||||
provider_type: &str,
|
||||
) -> AdminProviderOpenAiImageTestCapability {
|
||||
AdminProviderOpenAiImageTestCapability {
|
||||
max_generation_count: openai_image_provider_max_generation_count(provider_type),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_openai_image_normalize_options(
|
||||
provider_type: &str,
|
||||
) -> crate::ai_serving::OpenAiImageNormalizeOptions {
|
||||
openai_image_normalize_options_for_provider(provider_type)
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_model_test_capabilities_payload(
|
||||
provider_type: &str,
|
||||
model_id: &str,
|
||||
supports_image_generation: bool,
|
||||
) -> Value {
|
||||
let provider_type = provider_type.trim();
|
||||
let model_id = model_id.trim();
|
||||
let is_grok_image_edit =
|
||||
provider_type.eq_ignore_ascii_case("grok") && model_id == GROK_IMAGE_EDIT_MODEL_ID;
|
||||
let openai_image = if supports_image_generation {
|
||||
Some(json!({
|
||||
"max_generation_count": admin_provider_openai_image_test_capability(provider_type).max_generation_count,
|
||||
"supports_generation": !is_grok_image_edit,
|
||||
"supports_edit": is_grok_image_edit,
|
||||
}))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
json!({
|
||||
"openai:image": openai_image,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_model_supports_image_generation(
|
||||
provider_type: &str,
|
||||
model_id: &str,
|
||||
fallback_supports_image_generation: bool,
|
||||
) -> bool {
|
||||
if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
let model_id = model_id.trim();
|
||||
return GROK_IMAGE_MODEL_IDS
|
||||
.iter()
|
||||
.any(|candidate| model_id.eq_ignore_ascii_case(candidate));
|
||||
}
|
||||
fallback_supports_image_generation
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn grok_image_generation_models_expose_multi_image_capability() {
|
||||
let payload =
|
||||
admin_provider_model_test_capabilities_payload("grok", "grok-imagine-image", true);
|
||||
|
||||
assert_eq!(payload["openai:image"]["max_generation_count"], 4);
|
||||
assert_eq!(payload["openai:image"]["supports_generation"], true);
|
||||
assert_eq!(payload["openai:image"]["supports_edit"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_image_edit_model_is_edit_only_for_generation_tests() {
|
||||
let payload =
|
||||
admin_provider_model_test_capabilities_payload("grok", "grok-imagine-image-edit", true);
|
||||
|
||||
assert_eq!(payload["openai:image"]["max_generation_count"], 4);
|
||||
assert_eq!(payload["openai:image"]["supports_generation"], false);
|
||||
assert_eq!(payload["openai:image"]["supports_edit"], true);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_image_models_report_null_image_test_capability() {
|
||||
let payload = admin_provider_model_test_capabilities_payload("openai", "gpt-5.5", false);
|
||||
|
||||
assert!(payload["openai:image"].is_null());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_image_support_uses_catalog_model_ids_not_global_fallback() {
|
||||
assert!(admin_provider_model_supports_image_generation(
|
||||
"grok",
|
||||
"grok-imagine-image-pro",
|
||||
false,
|
||||
));
|
||||
assert!(!admin_provider_model_supports_image_generation(
|
||||
"grok",
|
||||
"grok-4.20-fast",
|
||||
true,
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -132,6 +132,12 @@ pub(crate) fn build_admin_provider_summary_value(
|
||||
.and_then(|cfg| cfg.get("architecture_id"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned);
|
||||
let kiro_simulated_cache_enabled = config
|
||||
.and_then(|cfg| cfg.get("kiro"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|cfg| cfg.get("simulated_cache_enabled"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let billing_type = quota_snapshot
|
||||
.map(|quota| quota.billing_type.clone())
|
||||
.or_else(|| provider.billing_type.clone());
|
||||
@@ -190,6 +196,7 @@ pub(crate) fn build_admin_provider_summary_value(
|
||||
"endpoint_health_details": endpoint_health_details,
|
||||
"ops_configured": ops_configured,
|
||||
"ops_architecture_id": ops_architecture_id,
|
||||
"kiro_simulated_cache_enabled": kiro_simulated_cache_enabled,
|
||||
"created_at": endpoint_timestamp_or_now(provider.created_at_unix_ms, now_unix_secs),
|
||||
"updated_at": endpoint_timestamp_or_now(provider.updated_at_unix_secs, now_unix_secs),
|
||||
})
|
||||
|
||||
@@ -2,7 +2,7 @@ use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyCreateRe
|
||||
use crate::handlers::admin::provider::write::normalize::{
|
||||
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
|
||||
normalize_api_format_list, normalize_auth_type, normalize_auth_type_by_format,
|
||||
validate_vertex_api_formats,
|
||||
normalize_max_probe_interval_minutes, validate_vertex_api_formats,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::{
|
||||
@@ -181,7 +181,8 @@ pub(crate) async fn build_admin_create_provider_key_record(
|
||||
key.rpm_limit = payload.rpm_limit;
|
||||
key.concurrent_limit = normalize_optional_api_key_concurrent_limit(payload.concurrent_limit)?;
|
||||
key.cache_ttl_minutes = payload.cache_ttl_minutes.unwrap_or(5);
|
||||
key.max_probe_interval_minutes = payload.max_probe_interval_minutes.unwrap_or(32);
|
||||
key.max_probe_interval_minutes =
|
||||
normalize_max_probe_interval_minutes(payload.max_probe_interval_minutes.unwrap_or(32))?;
|
||||
key.request_count = Some(0);
|
||||
key.success_count = Some(0);
|
||||
key.error_count = Some(0);
|
||||
|
||||
@@ -41,8 +41,9 @@ async fn build_admin_provider_key_items_payload(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let items = key_page
|
||||
.items
|
||||
let keys = key_page.items;
|
||||
|
||||
let items = keys
|
||||
.into_iter()
|
||||
.map(|key| {
|
||||
let api_formats =
|
||||
|
||||
@@ -2,7 +2,7 @@ use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePa
|
||||
use crate::handlers::admin::provider::write::normalize::{
|
||||
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
|
||||
normalize_api_format_list, normalize_auth_type, normalize_auth_type_by_format,
|
||||
validate_vertex_api_formats,
|
||||
normalize_max_probe_interval_minutes, validate_vertex_api_formats,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::{
|
||||
@@ -293,7 +293,8 @@ pub(crate) async fn build_admin_update_provider_key_record(
|
||||
updated.cache_ttl_minutes = cache_ttl_minutes;
|
||||
}
|
||||
if let Some(max_probe_interval_minutes) = payload.max_probe_interval_minutes {
|
||||
updated.max_probe_interval_minutes = max_probe_interval_minutes;
|
||||
updated.max_probe_interval_minutes =
|
||||
normalize_max_probe_interval_minutes(max_probe_interval_minutes)?;
|
||||
}
|
||||
if let Some(is_active) = payload.is_active {
|
||||
updated.is_active = is_active;
|
||||
|
||||
@@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, Strin
|
||||
let normalized = value.trim().to_ascii_lowercase();
|
||||
match normalized.as_str() {
|
||||
"custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli"
|
||||
| "antigravity" | "vertex_ai" => Ok(normalized),
|
||||
| "antigravity" | "vertex_ai" | "grok" => Ok(normalized),
|
||||
_ => Err(
|
||||
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai"
|
||||
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok"
|
||||
.to_string(),
|
||||
),
|
||||
}
|
||||
@@ -115,6 +115,14 @@ pub(crate) fn normalize_auth_type(value: Option<&str>) -> Result<String, String>
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_max_probe_interval_minutes(value: i32) -> Result<i32, String> {
|
||||
if (0..=32).contains(&value) {
|
||||
Ok(value)
|
||||
} else {
|
||||
Err("max_probe_interval_minutes 必须在 0 到 32 之间".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_pool_advanced_config(
|
||||
value: Option<serde_json::Value>,
|
||||
) -> Result<Option<serde_json::Value>, String> {
|
||||
@@ -161,8 +169,12 @@ pub(crate) fn validate_vertex_api_formats(
|
||||
}
|
||||
|
||||
let allowed = match auth_type {
|
||||
"api_key" => &["gemini:generate_content"][..],
|
||||
"service_account" | "vertex_ai" => &["claude:messages", "gemini:generate_content"][..],
|
||||
"api_key" => &["gemini:generate_content", "gemini:embedding"][..],
|
||||
"service_account" | "vertex_ai" => &[
|
||||
"claude:messages",
|
||||
"gemini:generate_content",
|
||||
"gemini:embedding",
|
||||
][..],
|
||||
_ => return Ok(()),
|
||||
};
|
||||
let invalid = api_formats
|
||||
@@ -260,6 +272,14 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_provider_type_supports_grok() {
|
||||
assert_eq!(
|
||||
normalize_provider_type_input(" Grok ").expect("type should normalize"),
|
||||
"grok"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_api_format_list_dedupes_canonical_formats() {
|
||||
assert_eq!(
|
||||
@@ -367,4 +387,27 @@ mod tests {
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_vertex_api_formats_allows_gemini_embedding() {
|
||||
assert!(validate_vertex_api_formats(
|
||||
"vertex_ai",
|
||||
"api_key",
|
||||
&[
|
||||
"gemini:generate_content".to_string(),
|
||||
"gemini:embedding".to_string()
|
||||
],
|
||||
)
|
||||
.is_ok());
|
||||
assert!(validate_vertex_api_formats(
|
||||
"vertex_ai",
|
||||
"service_account",
|
||||
&[
|
||||
"claude:messages".to_string(),
|
||||
"gemini:generate_content".to_string(),
|
||||
"gemini:embedding".to_string()
|
||||
],
|
||||
)
|
||||
.is_ok());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -205,7 +205,7 @@ pub(crate) async fn build_admin_update_provider_record(
|
||||
updated.stream_first_byte_timeout_secs = match payload.stream_first_byte_timeout {
|
||||
Some(value) if (1.0..=300.0).contains(&value) => Some(value),
|
||||
Some(_) => {
|
||||
return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string())
|
||||
return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string());
|
||||
}
|
||||
None => None,
|
||||
};
|
||||
|
||||
@@ -0,0 +1,278 @@
|
||||
use crate::data::state::{ReferralRelationshipListQuery, ReferralRewardListQuery};
|
||||
use crate::handlers::admin::request::{AdminRouteRequest, AdminRouteResult};
|
||||
use crate::handlers::admin::shared::{attach_admin_audit_response, query_param_value};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
|
||||
#[derive(Debug, Default, Deserialize)]
|
||||
struct ReferralAdminMutationRequest {
|
||||
note: Option<String>,
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_referrals_response(
|
||||
request: AdminRouteRequest<'_>,
|
||||
) -> AdminRouteResult {
|
||||
let request_context = request.request_context();
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if decision.route_family.as_deref() != Some("referrals_manage") {
|
||||
return Ok(None);
|
||||
}
|
||||
let response = match decision.route_kind.as_deref() {
|
||||
Some("list_referrals") => {
|
||||
build_admin_referrals_list_response(&request.state(), &request_context).await?
|
||||
}
|
||||
Some("list_referral_rewards") => {
|
||||
build_admin_referral_rewards_list_response(&request.state(), &request_context).await?
|
||||
}
|
||||
Some("retry_referral_reward") => {
|
||||
build_admin_referral_reward_retry_response(
|
||||
&request.state(),
|
||||
&request_context,
|
||||
request.request_body(),
|
||||
)
|
||||
.await?
|
||||
}
|
||||
Some("void_referral_reward") => {
|
||||
build_admin_referral_reward_void_response(
|
||||
&request.state(),
|
||||
&request_context,
|
||||
request.request_body(),
|
||||
)
|
||||
.await?
|
||||
}
|
||||
_ => build_admin_referrals_unavailable_response(),
|
||||
};
|
||||
Ok(Some(response))
|
||||
}
|
||||
|
||||
fn admin_referrals_bad_request(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn build_admin_referrals_unavailable_response() -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
Json(json!({ "detail": "Admin referral data unavailable" })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn parse_limit(query: Option<&str>) -> Result<usize, String> {
|
||||
match query_param_value(query, "limit") {
|
||||
Some(value) => value
|
||||
.parse::<usize>()
|
||||
.map(|value| value.clamp(1, 200))
|
||||
.map_err(|_| "limit 必须是正整数".to_string()),
|
||||
None => Ok(50),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_offset(query: Option<&str>) -> Result<usize, String> {
|
||||
match query_param_value(query, "offset") {
|
||||
Some(value) => value
|
||||
.parse::<usize>()
|
||||
.map_err(|_| "offset 必须是非负整数".to_string()),
|
||||
None => Ok(0),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_optional_bool(query: Option<&str>, key: &str) -> Result<Option<bool>, String> {
|
||||
let Some(value) = query_param_value(query, key) else {
|
||||
return Ok(None);
|
||||
};
|
||||
match value.trim().to_ascii_lowercase().as_str() {
|
||||
"true" | "1" | "yes" => Ok(Some(true)),
|
||||
"false" | "0" | "no" => Ok(Some(false)),
|
||||
_ => Err(format!("{key} 必须是布尔值")),
|
||||
}
|
||||
}
|
||||
|
||||
fn operator_id(
|
||||
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
|
||||
) -> Option<String> {
|
||||
request_context
|
||||
.decision()
|
||||
.and_then(|decision| decision.admin_principal.as_ref())
|
||||
.map(|principal| principal.user_id.clone())
|
||||
}
|
||||
|
||||
fn reward_id_from_path(path: &str, suffix: &str) -> Option<String> {
|
||||
let trimmed = path.trim_end_matches('/');
|
||||
let rest = trimmed.strip_prefix("/api/admin/referral-rewards/")?;
|
||||
let id = rest.strip_suffix(suffix)?.trim_end_matches('/');
|
||||
(!id.is_empty()).then_some(id.to_string())
|
||||
}
|
||||
|
||||
fn parse_mutation_note(body: Option<&axum::body::Bytes>) -> Result<Option<String>, String> {
|
||||
let Some(body) = body.filter(|body| !body.is_empty()) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let payload = serde_json::from_slice::<ReferralAdminMutationRequest>(body)
|
||||
.map_err(|_| "请求数据验证失败".to_string())?;
|
||||
Ok(payload
|
||||
.note
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty()))
|
||||
}
|
||||
|
||||
async fn build_admin_referrals_list_response(
|
||||
state: &crate::handlers::admin::request::AdminAppState<'_>,
|
||||
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let query = request_context.query_string();
|
||||
let limit = match parse_limit(query) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(admin_referrals_bad_request(detail)),
|
||||
};
|
||||
let offset = match parse_offset(query) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(admin_referrals_bad_request(detail)),
|
||||
};
|
||||
let first_paid = match parse_optional_bool(query, "first_paid") {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(admin_referrals_bad_request(detail)),
|
||||
};
|
||||
let Some((items, total, stats)) = state
|
||||
.app()
|
||||
.list_admin_referral_relationships(ReferralRelationshipListQuery {
|
||||
inviter: query_param_value(query, "inviter"),
|
||||
invitee: query_param_value(query, "invitee"),
|
||||
invite_code: query_param_value(query, "invite_code"),
|
||||
first_paid,
|
||||
limit,
|
||||
offset,
|
||||
})
|
||||
.await?
|
||||
else {
|
||||
return Ok(build_admin_referrals_unavailable_response());
|
||||
};
|
||||
Ok(Json(json!({
|
||||
"items": items,
|
||||
"total": total,
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
"stats": stats,
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
async fn build_admin_referral_rewards_list_response(
|
||||
state: &crate::handlers::admin::request::AdminAppState<'_>,
|
||||
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let query = request_context.query_string();
|
||||
let limit = match parse_limit(query) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(admin_referrals_bad_request(detail)),
|
||||
};
|
||||
let offset = match parse_offset(query) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(admin_referrals_bad_request(detail)),
|
||||
};
|
||||
let Some((items, total, stats)) = state
|
||||
.app()
|
||||
.list_admin_referral_rewards(ReferralRewardListQuery {
|
||||
order_id: query_param_value(query, "order_id"),
|
||||
reward_type: query_param_value(query, "reward_type"),
|
||||
status: query_param_value(query, "status"),
|
||||
limit,
|
||||
offset,
|
||||
})
|
||||
.await?
|
||||
else {
|
||||
return Ok(build_admin_referrals_unavailable_response());
|
||||
};
|
||||
Ok(Json(json!({
|
||||
"items": items,
|
||||
"total": total,
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
"stats": stats,
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
async fn build_admin_referral_reward_retry_response(
|
||||
state: &crate::handlers::admin::request::AdminAppState<'_>,
|
||||
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
|
||||
request_body: Option<&axum::body::Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(reward_id) = reward_id_from_path(request_context.path(), "/retry") else {
|
||||
return Ok(admin_referrals_bad_request("返利记录不存在"));
|
||||
};
|
||||
let note = match parse_mutation_note(request_body) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(admin_referrals_bad_request(detail)),
|
||||
};
|
||||
match state
|
||||
.app()
|
||||
.retry_referral_reward(
|
||||
&reward_id,
|
||||
operator_id(request_context).as_deref(),
|
||||
note.as_deref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(reward) => Ok(attach_admin_audit_response(
|
||||
Json(json!({ "reward": reward })).into_response(),
|
||||
"admin_referral_reward_retry",
|
||||
"retry_referral_reward",
|
||||
"referral_reward",
|
||||
&reward_id,
|
||||
)),
|
||||
None => Ok((
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": "Referral reward not found" })),
|
||||
)
|
||||
.into_response()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn build_admin_referral_reward_void_response(
|
||||
state: &crate::handlers::admin::request::AdminAppState<'_>,
|
||||
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
|
||||
request_body: Option<&axum::body::Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(reward_id) = reward_id_from_path(request_context.path(), "/void") else {
|
||||
return Ok(admin_referrals_bad_request("返利记录不存在"));
|
||||
};
|
||||
let note = match parse_mutation_note(request_body) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(admin_referrals_bad_request(detail)),
|
||||
};
|
||||
match state
|
||||
.app()
|
||||
.void_referral_reward(
|
||||
&reward_id,
|
||||
operator_id(request_context).as_deref(),
|
||||
note.as_deref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(reward) => Ok(attach_admin_audit_response(
|
||||
Json(json!({ "reward": reward })).into_response(),
|
||||
"admin_referral_reward_void",
|
||||
"void_referral_reward",
|
||||
"referral_reward",
|
||||
&reward_id,
|
||||
)),
|
||||
None => Ok((
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": "Referral reward not found" })),
|
||||
)
|
||||
.into_response()),
|
||||
}
|
||||
}
|
||||
@@ -250,4 +250,19 @@ impl<'a> AdminAppState<'a> {
|
||||
crate::execution_runtime::execute_execution_runtime_sync_plan(self.app, trace_id, plan)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_execution_runtime_sync_plan_with_report_context(
|
||||
&self,
|
||||
trace_id: Option<&str>,
|
||||
plan: &aether_contracts::ExecutionPlan,
|
||||
report_context: Option<&serde_json::Value>,
|
||||
) -> Result<aether_contracts::ExecutionResult, GatewayError> {
|
||||
crate::execution_runtime::execute_execution_runtime_sync_plan_with_report_context(
|
||||
self.app,
|
||||
trace_id,
|
||||
plan,
|
||||
report_context,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,8 +18,7 @@ use std::collections::BTreeMap;
|
||||
use std::io::Read;
|
||||
use url::Url;
|
||||
|
||||
const KIRO_IDC_AMZ_USER_AGENT: &str =
|
||||
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
|
||||
const KIRO_IDC_AMZ_USER_AGENT: &str = "aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
|
||||
const ADMIN_PROVIDER_OAUTH_TIMEOUT_MS: u64 = 30_000;
|
||||
const ADMIN_PROVIDER_OAUTH_PROXY_TIMEOUT_MS: u64 = 60_000;
|
||||
|
||||
|
||||
@@ -84,7 +84,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Json(json!({ "detail": "请求数据验证失败" })),
|
||||
)
|
||||
.into_response(),
|
||||
))
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -1197,7 +1197,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Err(_) => {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"Provider '{provider_name}' 配置格式无效"
|
||||
))))
|
||||
))));
|
||||
}
|
||||
};
|
||||
let mut updated = invalid!(
|
||||
@@ -1230,7 +1230,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Err(_) => {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"Provider '{provider_name}' 配置格式无效"
|
||||
))))
|
||||
))));
|
||||
}
|
||||
};
|
||||
let (mut record, shift_existing_priorities_from) =
|
||||
@@ -1301,7 +1301,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Err(_) => {
|
||||
return Ok(Err(invalid_request(
|
||||
"Provider Endpoint 配置格式无效",
|
||||
)))
|
||||
)));
|
||||
}
|
||||
};
|
||||
let (fields, payload) = patch.into_parts();
|
||||
@@ -1504,7 +1504,7 @@ impl<'a> AdminAppState<'a> {
|
||||
) {
|
||||
Ok(patch) => patch,
|
||||
Err(_) => {
|
||||
return Ok(Err(invalid_request("Provider Key 配置格式无效")))
|
||||
return Ok(Err(invalid_request("Provider Key 配置格式无效")));
|
||||
}
|
||||
};
|
||||
let mut updated = invalid!(
|
||||
@@ -1885,6 +1885,7 @@ impl<'a> AdminAppState<'a> {
|
||||
oauth_provider.extra_config,
|
||||
"extra_config",
|
||||
)),
|
||||
icon_url: None,
|
||||
is_enabled: oauth_provider.is_enabled,
|
||||
};
|
||||
invalid!(record.validate().map_err(|err| err.to_string()));
|
||||
@@ -1985,7 +1986,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Err(_) => {
|
||||
return Ok(Err(invalid_request(
|
||||
"merge_mode 仅支持 skip / overwrite / error",
|
||||
)))
|
||||
)));
|
||||
}
|
||||
};
|
||||
let empty = Vec::new();
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::{
|
||||
announcements, auth, billing, endpoint, features, model, observability, provider, request,
|
||||
routing, system, users,
|
||||
announcements, auth, billing, endpoint, features, model, observability, provider, referrals,
|
||||
request, routing, system, users,
|
||||
};
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_response(
|
||||
@@ -44,6 +44,10 @@ pub(crate) async fn maybe_build_local_admin_response(
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) = referrals::maybe_build_local_admin_referrals_response(request).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) = features::maybe_build_local_admin_features_response(request).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
mod paths;
|
||||
mod payloads;
|
||||
mod proxy_errors;
|
||||
mod usage_counter;
|
||||
|
||||
pub(crate) use self::paths::*;
|
||||
pub(crate) use self::payloads::*;
|
||||
pub(crate) use self::proxy_errors::build_proxy_error_response;
|
||||
pub(crate) use self::usage_counter::build_admin_usage_counter_health_payload;
|
||||
pub(crate) use crate::handlers::shared::{
|
||||
attach_admin_audit_response, build_admin_provider_key_response,
|
||||
decrypt_catalog_secret_with_fallbacks, default_provider_key_status_snapshot,
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
use aether_data_contracts::repository::usage::UsageCounterHealthSnapshot;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub(crate) fn build_admin_usage_counter_health_payload(
|
||||
snapshot: &UsageCounterHealthSnapshot,
|
||||
now_unix_secs: u64,
|
||||
) -> Value {
|
||||
let oldest_pending_age_secs = snapshot
|
||||
.oldest_pending_created_at_unix_secs
|
||||
.map(|created_at| now_unix_secs.saturating_sub(created_at));
|
||||
let status = match (snapshot.pending_rows, oldest_pending_age_secs) {
|
||||
(0, _) => "idle",
|
||||
(_, Some(age)) if age >= 60 => "backlogged",
|
||||
_ => "catching_up",
|
||||
};
|
||||
|
||||
json!({
|
||||
"status": status,
|
||||
"outbox_pending_rows": snapshot.pending_rows,
|
||||
"outbox_processed_rows": snapshot.processed_rows,
|
||||
"oldest_pending_created_at_unix_secs": snapshot.oldest_pending_created_at_unix_secs,
|
||||
"oldest_pending_age_secs": oldest_pending_age_secs,
|
||||
"latest_processed_at_unix_secs": snapshot.latest_processed_at_unix_secs,
|
||||
"pending_by_kind": snapshot.pending_by_kind,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_data_contracts::repository::usage::UsageCounterHealthSnapshot;
|
||||
use serde_json::json;
|
||||
|
||||
use super::build_admin_usage_counter_health_payload;
|
||||
|
||||
#[test]
|
||||
fn usage_counter_health_payload_reports_idle_without_pending_rows() {
|
||||
let snapshot = UsageCounterHealthSnapshot {
|
||||
processed_rows: 42,
|
||||
latest_processed_at_unix_secs: Some(1_000),
|
||||
..UsageCounterHealthSnapshot::default()
|
||||
};
|
||||
|
||||
let payload = build_admin_usage_counter_health_payload(&snapshot, 1_100);
|
||||
|
||||
assert_eq!(payload["status"], json!("idle"));
|
||||
assert_eq!(payload["outbox_pending_rows"], json!(0));
|
||||
assert_eq!(payload["outbox_processed_rows"], json!(42));
|
||||
assert_eq!(payload["oldest_pending_age_secs"], json!(null));
|
||||
assert_eq!(payload["latest_processed_at_unix_secs"], json!(1_000));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_counter_health_payload_reports_catching_up_for_fresh_backlog() {
|
||||
let mut pending_by_kind = BTreeMap::new();
|
||||
pending_by_kind.insert("api_key".to_string(), 3);
|
||||
let snapshot = UsageCounterHealthSnapshot {
|
||||
pending_rows: 3,
|
||||
oldest_pending_created_at_unix_secs: Some(1_050),
|
||||
pending_by_kind,
|
||||
..UsageCounterHealthSnapshot::default()
|
||||
};
|
||||
|
||||
let payload = build_admin_usage_counter_health_payload(&snapshot, 1_100);
|
||||
|
||||
assert_eq!(payload["status"], json!("catching_up"));
|
||||
assert_eq!(payload["oldest_pending_age_secs"], json!(50));
|
||||
assert_eq!(payload["pending_by_kind"]["api_key"], json!(3));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_counter_health_payload_reports_backlogged_for_old_backlog() {
|
||||
let snapshot = UsageCounterHealthSnapshot {
|
||||
pending_rows: 1,
|
||||
oldest_pending_created_at_unix_secs: Some(1_000),
|
||||
..UsageCounterHealthSnapshot::default()
|
||||
};
|
||||
|
||||
let payload = build_admin_usage_counter_health_payload(&snapshot, 1_060);
|
||||
|
||||
assert_eq!(payload["status"], json!("backlogged"));
|
||||
assert_eq!(payload["oldest_pending_age_secs"], json!(60));
|
||||
}
|
||||
}
|
||||
@@ -380,7 +380,7 @@ pub(crate) async fn maybe_build_local_admin_proxy_nodes_response(
|
||||
};
|
||||
if !existing.tunnel_mode {
|
||||
return Ok(Some(bad_request_response(
|
||||
"non-tunnel mode is no longer supported, please upgrade aether-proxy to use tunnel mode",
|
||||
"non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode",
|
||||
)));
|
||||
}
|
||||
let Some(node) = state.apply_proxy_node_heartbeat(&mutation).await? else {
|
||||
@@ -1049,7 +1049,7 @@ async fn test_proxy_node_connectivity(
|
||||
None,
|
||||
None,
|
||||
Some(
|
||||
"non-tunnel mode is no longer supported, please upgrade aether-proxy to use tunnel mode"
|
||||
"non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode"
|
||||
.to_string(),
|
||||
),
|
||||
);
|
||||
@@ -1559,7 +1559,8 @@ fn admin_proxy_node_test_node_id_from_path(path: &str) -> Option<String> {
|
||||
fn normalize_proxy_upgrade_version(value: &str) -> String {
|
||||
value
|
||||
.trim()
|
||||
.strip_prefix("proxy-v")
|
||||
.strip_prefix("tunnel-v")
|
||||
.or_else(|| value.trim().strip_prefix("proxy-v"))
|
||||
.unwrap_or(value.trim())
|
||||
.to_ascii_lowercase()
|
||||
}
|
||||
@@ -2155,7 +2156,7 @@ async fn create_proxy_install_management_token(
|
||||
user,
|
||||
token_hash: hash_proxy_install_management_token(&raw_token),
|
||||
token_prefix: proxy_install_management_token_prefix(&raw_token),
|
||||
name: format!("aether-proxy {node_name} {short_id}"),
|
||||
name: format!("aether-tunnel {node_name} {short_id}"),
|
||||
description: Some("Created by proxy node one-click installer".to_string()),
|
||||
allowed_ips: None,
|
||||
permissions: Some(json!(["admin:proxy_nodes:write"])),
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::shared::build_admin_usage_counter_health_payload;
|
||||
use crate::handlers::shared::{system_config_bool, system_config_string};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::system::{
|
||||
@@ -114,6 +115,15 @@ pub(crate) async fn build_admin_system_stats_payload(
|
||||
.filter(|provider| provider.is_active)
|
||||
.count() as u64;
|
||||
let stats = state.read_admin_system_stats().await?;
|
||||
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
|
||||
let usage_counter_snapshot = state
|
||||
.as_ref()
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let usage_counter =
|
||||
build_admin_usage_counter_health_payload(&usage_counter_snapshot, now_unix_secs);
|
||||
|
||||
Ok(build_admin_system_stats_payload_pure(
|
||||
stats.total_users,
|
||||
@@ -122,6 +132,7 @@ pub(crate) async fn build_admin_system_stats_payload(
|
||||
active_providers,
|
||||
stats.total_api_keys,
|
||||
stats.total_requests,
|
||||
usage_counter,
|
||||
))
|
||||
}
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
#[derive(Debug, serde::Deserialize)]
|
||||
struct AdminUserGroupPayload {
|
||||
@@ -209,14 +210,18 @@ pub(in super::super) async fn build_admin_replace_user_group_members_response(
|
||||
if state.find_user_group_by_id(&group_id).await?.is_none() {
|
||||
return Ok(not_found("用户分组不存在"));
|
||||
}
|
||||
if read_default_user_group_id(state).await?.as_deref() == Some(group_id.as_str()) {
|
||||
return Ok(bad_request_owned("默认用户组成员由系统维护".to_string()));
|
||||
}
|
||||
let payload = match parse_members_payload(request_body) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(bad_request_owned(detail)),
|
||||
};
|
||||
let user_ids = normalize_ids(payload.user_ids);
|
||||
if read_default_user_group_id(state).await?.as_deref() == Some(group_id.as_str()) {
|
||||
if let Some(response) =
|
||||
validate_default_group_member_replacement(state, &group_id, &user_ids).await?
|
||||
{
|
||||
return Ok(response);
|
||||
}
|
||||
}
|
||||
let known_users = state.resolve_auth_user_summaries_by_ids(&user_ids).await?;
|
||||
if known_users.len() != user_ids.len() {
|
||||
return Ok(bad_request_owned("成员包含不存在的用户".to_string()));
|
||||
@@ -245,6 +250,52 @@ pub(in super::super) async fn build_admin_replace_user_group_members_response(
|
||||
))
|
||||
}
|
||||
|
||||
async fn validate_default_group_member_replacement(
|
||||
state: &AdminAppState<'_>,
|
||||
group_id: &str,
|
||||
next_user_ids: &[String],
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let next_user_ids = next_user_ids.iter().cloned().collect::<BTreeSet<String>>();
|
||||
let removed_user_ids = state
|
||||
.list_user_group_members(group_id)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|member| !next_user_ids.contains(&member.user_id))
|
||||
.map(|member| member.user_id)
|
||||
.collect::<Vec<_>>();
|
||||
if removed_user_ids.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let summaries = state
|
||||
.resolve_auth_user_summaries_by_ids(&removed_user_ids)
|
||||
.await?;
|
||||
let users_with_other_groups = state
|
||||
.list_user_group_memberships_by_user_ids(&removed_user_ids)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|membership| membership.group_id != group_id)
|
||||
.map(|membership| membership.user_id)
|
||||
.collect::<BTreeSet<_>>();
|
||||
|
||||
for user_id in removed_user_ids {
|
||||
let Some(summary) = summaries.get(&user_id) else {
|
||||
continue;
|
||||
};
|
||||
if crate::roles::can_access_admin_console(&summary.role) {
|
||||
continue;
|
||||
}
|
||||
if !users_with_other_groups.contains(&user_id) {
|
||||
return Ok(Some(bad_request_owned(format!(
|
||||
"用户 {} 移出默认组后将不属于任何用户组",
|
||||
summary.username
|
||||
))));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
pub(in super::super) async fn build_admin_set_default_user_group_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_body: Option<&axum::body::Bytes>,
|
||||
|
||||
@@ -17,6 +17,7 @@ use tracing::{info, trace, warn};
|
||||
|
||||
pub(super) fn request_wants_stream(
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
headers: &http::HeaderMap,
|
||||
body: &axum::body::Bytes,
|
||||
) -> bool {
|
||||
if request_context
|
||||
@@ -34,7 +35,11 @@ pub(super) fn request_wants_stream(
|
||||
{
|
||||
return false;
|
||||
}
|
||||
serde_json::from_slice::<serde_json::Value>(body)
|
||||
let body = match crate::headers::decoded_request_body_bytes(headers, body.as_ref()) {
|
||||
Ok(body) => body,
|
||||
Err(_) => return false,
|
||||
};
|
||||
serde_json::from_slice::<serde_json::Value>(body.as_ref())
|
||||
.ok()
|
||||
.and_then(|value| value.get("stream").and_then(|stream| stream.as_bool()))
|
||||
.unwrap_or(false)
|
||||
@@ -258,11 +263,11 @@ pub(super) fn finalize_gateway_response_with_context(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::finalize_gateway_response;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use super::{finalize_gateway_response, request_wants_stream};
|
||||
use crate::control::{GatewayControlDecision, GatewayPublicRequestContext};
|
||||
use crate::AppState;
|
||||
use axum::body::Body;
|
||||
use axum::http::{Method, Response, StatusCode};
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::http::{HeaderMap, HeaderValue, Method, Response, StatusCode};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Instant;
|
||||
use tracing_subscriber::filter::LevelFilter;
|
||||
@@ -306,6 +311,32 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_wants_stream_reads_zstd_encoded_json_body() {
|
||||
let request_context = GatewayPublicRequestContext {
|
||||
trace_id: "trace-zstd-stream".to_string(),
|
||||
request_method: Method::POST,
|
||||
request_path: "/v1/responses".to_string(),
|
||||
request_query_string: None,
|
||||
request_content_type: Some("application/json".to_string()),
|
||||
host_header: None,
|
||||
control_decision: None,
|
||||
};
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::CONTENT_ENCODING,
|
||||
HeaderValue::from_static("zstd"),
|
||||
);
|
||||
let encoded = zstd::stream::encode_all(br#"{"stream":true}"#.as_slice(), 0)
|
||||
.expect("zstd body should encode");
|
||||
|
||||
assert!(request_wants_stream(
|
||||
&request_context,
|
||||
&headers,
|
||||
&Bytes::from(encoded),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn finalize_gateway_response_logs_sanitized_path_and_query() {
|
||||
let state = AppState::new().expect("gateway state should build");
|
||||
|
||||
@@ -21,12 +21,13 @@ use crate::constants::{
|
||||
EXECUTION_PATH_EXECUTION_RUNTIME_STREAM, EXECUTION_PATH_EXECUTION_RUNTIME_SYNC,
|
||||
EXECUTION_PATH_LOCAL_AI_PUBLIC, EXECUTION_PATH_LOCAL_API_KEY_CONCURRENCY_LIMITED,
|
||||
EXECUTION_PATH_LOCAL_AUTH_DENIED, EXECUTION_PATH_LOCAL_EXECUTION_LOOP_DETECTED,
|
||||
EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, EXECUTION_PATH_LOCAL_OVERLOADED,
|
||||
EXECUTION_PATH_LOCAL_PROXY_PASSTHROUGH_REMOVED, EXECUTION_PATH_LOCAL_RATE_LIMITED,
|
||||
EXECUTION_PATH_LOCAL_ROUTE_NOT_FOUND, EXECUTION_PATH_PUBLIC_PROXY_PASSTHROUGH,
|
||||
EXECUTION_RUNTIME_LOOP_GUARD_HEADER, FORWARDED_FOR_HEADER, FORWARDED_HOST_HEADER,
|
||||
FORWARDED_PROTO_HEADER, GATEWAY_HEADER, LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER,
|
||||
TRACE_ID_HEADER, TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, TRUSTED_AUTH_API_KEY_ID_HEADER,
|
||||
EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS, EXECUTION_PATH_LOCAL_INVALID_REQUEST,
|
||||
EXECUTION_PATH_LOCAL_OVERLOADED, EXECUTION_PATH_LOCAL_PROXY_PASSTHROUGH_REMOVED,
|
||||
EXECUTION_PATH_LOCAL_RATE_LIMITED, EXECUTION_PATH_LOCAL_ROUTE_NOT_FOUND,
|
||||
EXECUTION_PATH_PUBLIC_PROXY_PASSTHROUGH, EXECUTION_RUNTIME_LOOP_GUARD_HEADER,
|
||||
FORWARDED_FOR_HEADER, FORWARDED_HOST_HEADER, FORWARDED_PROTO_HEADER, GATEWAY_HEADER,
|
||||
LOCAL_EXECUTION_RUNTIME_MISS_REASON_HEADER, TRACE_ID_HEADER,
|
||||
TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, TRUSTED_AUTH_API_KEY_ID_HEADER,
|
||||
TRUSTED_AUTH_BALANCE_HEADER, TRUSTED_AUTH_USER_ID_HEADER, TUNNEL_AFFINITY_FORWARDED_BY_HEADER,
|
||||
TUNNEL_AFFINITY_OWNER_INSTANCE_HEADER,
|
||||
};
|
||||
@@ -51,7 +52,7 @@ use crate::handlers::shared::{
|
||||
};
|
||||
use crate::headers::{
|
||||
extract_or_generate_trace_id, request_origin_from_headers_and_remote_addr,
|
||||
should_skip_request_header,
|
||||
should_skip_request_header, RequestBodyNormalizationError,
|
||||
};
|
||||
use crate::router::RequestAdmissionError;
|
||||
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode};
|
||||
@@ -92,6 +93,69 @@ const EXECUTION_PATH_TUNNEL_AFFINITY_FORWARD: &str = "tunnel_affinity_forward";
|
||||
const MANAGEMENT_TOKEN_PREFIX: &str = "ae-";
|
||||
const LEGACY_MANAGEMENT_TOKEN_PREFIX: &str = "ae_";
|
||||
|
||||
fn build_request_body_normalization_error_response(
|
||||
trace_id: &str,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
error: &RequestBodyNormalizationError,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
warn!(
|
||||
event_name = "frontdoor_request_body_normalization_failed",
|
||||
log_type = "ops",
|
||||
trace_id,
|
||||
method = %request_context.request_method,
|
||||
path = %request_context.request_path_and_query(),
|
||||
error = %error,
|
||||
"gateway rejected request with invalid encoded body"
|
||||
);
|
||||
build_local_http_error_response(
|
||||
trace_id,
|
||||
request_context.control_decision.as_ref(),
|
||||
error.http_status(),
|
||||
error.client_message().as_str(),
|
||||
)
|
||||
}
|
||||
|
||||
async fn buffer_and_normalize_request_body(
|
||||
request_body: &mut Option<Body>,
|
||||
headers: &mut http::HeaderMap,
|
||||
body_owner_expectation: &'static str,
|
||||
) -> Result<Result<Bytes, RequestBodyNormalizationError>, GatewayError> {
|
||||
if let Err(err) = crate::headers::check_request_content_length(headers) {
|
||||
return Ok(Err(err));
|
||||
}
|
||||
let body = to_bytes(
|
||||
request_body.take().expect(body_owner_expectation),
|
||||
usize::MAX,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
Ok(crate::headers::normalize_request_body_headers_and_bytes(
|
||||
headers, body,
|
||||
))
|
||||
}
|
||||
|
||||
fn finalize_request_body_normalization_rejection(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
remote_addr: &std::net::SocketAddr,
|
||||
started_at: &std::time::Instant,
|
||||
trace_id: &str,
|
||||
request_permit: Option<aether_runtime::AdmissionPermit>,
|
||||
error: &RequestBodyNormalizationError,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let response =
|
||||
build_request_body_normalization_error_response(trace_id, request_context, error)?;
|
||||
Ok(finalize_gateway_response_with_context(
|
||||
state,
|
||||
response,
|
||||
remote_addr,
|
||||
request_context,
|
||||
EXECUTION_PATH_LOCAL_INVALID_REQUEST,
|
||||
started_at,
|
||||
request_permit,
|
||||
))
|
||||
}
|
||||
|
||||
fn local_execution_outcome_label(outcome: &LocalExecutionRequestOutcome) -> &'static str {
|
||||
match outcome {
|
||||
LocalExecutionRequestOutcome::Responded(_) => "responded",
|
||||
@@ -342,8 +406,11 @@ async fn maybe_forward_public_request_to_tunnel_owner(
|
||||
) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json =
|
||||
buffered_body.and_then(|body| serde_json::from_slice::<serde_json::Value>(body).ok());
|
||||
let body_json = buffered_body.and_then(|body| {
|
||||
let body =
|
||||
crate::headers::decoded_request_body_bytes(&parts.headers, body.as_ref()).ok()?;
|
||||
serde_json::from_slice::<serde_json::Value>(body.as_ref()).ok()
|
||||
});
|
||||
let client_session_affinity =
|
||||
crate::client_session_affinity::client_session_affinity_from_parts(
|
||||
parts,
|
||||
@@ -473,6 +540,7 @@ async fn maybe_forward_public_request_to_tunnel_owner(
|
||||
|
||||
let mut response = build_sync_aware_affinity_forward_response(
|
||||
request_context,
|
||||
&parts.headers,
|
||||
buffered_body,
|
||||
decision,
|
||||
upstream_response,
|
||||
@@ -732,6 +800,7 @@ fn build_stream_sse_proxy_response(
|
||||
|
||||
async fn build_sync_aware_affinity_forward_response(
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
request_headers: &http::HeaderMap,
|
||||
buffered_body: Option<&Bytes>,
|
||||
decision: &GatewayControlDecision,
|
||||
upstream_response: reqwest::Response,
|
||||
@@ -739,7 +808,7 @@ async fn build_sync_aware_affinity_forward_response(
|
||||
let Some(buffered_body) = buffered_body else {
|
||||
return build_client_response(upstream_response, &request_context.trace_id, Some(decision));
|
||||
};
|
||||
let stream_request = request_wants_stream(request_context, buffered_body);
|
||||
let stream_request = request_wants_stream(request_context, request_headers, buffered_body);
|
||||
let upstream_is_sse = upstream_response_is_sse(upstream_response.headers());
|
||||
if (!stream_request && !upstream_is_sse) || (stream_request && upstream_is_sse) {
|
||||
return build_client_response(upstream_response, &request_context.trace_id, Some(decision));
|
||||
@@ -980,16 +1049,26 @@ pub(crate) async fn proxy_request(
|
||||
}
|
||||
let mut request_body = Some(body);
|
||||
let local_proxy_body = if local_proxy_route_requires_buffered_body(&request_context) {
|
||||
Some(
|
||||
to_bytes(
|
||||
request_body
|
||||
.take()
|
||||
.expect("local proxy body buffering should own request body"),
|
||||
usize::MAX,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?,
|
||||
let body = buffer_and_normalize_request_body(
|
||||
&mut request_body,
|
||||
&mut parts.headers,
|
||||
"local proxy body buffering should own request body",
|
||||
)
|
||||
.await?;
|
||||
match body {
|
||||
Ok(body) => Some(body),
|
||||
Err(err) => {
|
||||
return finalize_request_body_normalization_rejection(
|
||||
&state,
|
||||
&request_context,
|
||||
&remote_addr,
|
||||
&started_at,
|
||||
&trace_id,
|
||||
request_permit.take(),
|
||||
&err,
|
||||
);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
@@ -1150,16 +1229,26 @@ pub(crate) async fn proxy_request(
|
||||
&& request_enables_control_execute(&parts.headers);
|
||||
|
||||
let buffered_body = if should_buffer_body {
|
||||
Some(
|
||||
to_bytes(
|
||||
request_body
|
||||
.take()
|
||||
.expect("buffered auth/execution runtime path should own request body"),
|
||||
usize::MAX,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?,
|
||||
let body = buffer_and_normalize_request_body(
|
||||
&mut request_body,
|
||||
&mut parts.headers,
|
||||
"buffered auth/execution runtime path should own request body",
|
||||
)
|
||||
.await?;
|
||||
match body {
|
||||
Ok(body) => Some(body),
|
||||
Err(err) => {
|
||||
return finalize_request_body_normalization_rejection(
|
||||
&state,
|
||||
&request_context,
|
||||
&remote_addr,
|
||||
&started_at,
|
||||
&trace_id,
|
||||
request_permit.take(),
|
||||
&err,
|
||||
);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
@@ -1301,7 +1390,7 @@ pub(crate) async fn proxy_request(
|
||||
let buffered_body = buffered_body
|
||||
.as_ref()
|
||||
.expect("execution runtime/control auth gate should have buffered request body");
|
||||
let stream_request = request_wants_stream(&request_context, buffered_body);
|
||||
let stream_request = request_wants_stream(&request_context, &parts.headers, buffered_body);
|
||||
let mut local_execution_exhaustion = None;
|
||||
if stream_request {
|
||||
let stream_outcome = maybe_execute_stream_request(
|
||||
@@ -1778,9 +1867,9 @@ fn local_execution_runtime_miss_all_candidates_skipped_detail(
|
||||
(_, Some(summary), Some(model)) => format!(
|
||||
"支持模型 {model} 的候选提供商全部不可用:{summary}(原因代码: all_candidates_skipped)"
|
||||
),
|
||||
(_, Some(summary), None) => format!(
|
||||
"候选提供商全部不可用:{summary}(原因代码: all_candidates_skipped)"
|
||||
),
|
||||
(_, Some(summary), None) => {
|
||||
format!("候选提供商全部不可用:{summary}(原因代码: all_candidates_skipped)")
|
||||
}
|
||||
(count, None, Some(model)) if count > 0 => format!(
|
||||
"找到 {count} 个支持模型 {model} 的候选提供商,但都不满足本次{request_mode}请求要求(原因代码: all_candidates_skipped)"
|
||||
),
|
||||
@@ -1838,6 +1927,7 @@ fn local_execution_runtime_miss_skip_reason_label(reason: &str) -> &str {
|
||||
"key_inactive" => "API Key 未启用",
|
||||
"key_model_disabled" => "API Key 未允许该模型",
|
||||
"mapped_model_missing" => "模型映射缺失",
|
||||
"pool_active_probe_sealed" => "池内账号未进入主动探测热池",
|
||||
"pool_cooldown" => "池内账号处于冷却中",
|
||||
"pool_cost_limit_reached" => "池内账号成本额度已用尽",
|
||||
"pool_group_exhausted" => "池化提供商没有可调度账号",
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
use crate::async_task::CancelVideoTaskError;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::image_capabilities::{
|
||||
openai_image_gateway_max_generation_count, openai_image_gateway_max_generation_count_for_model,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
use aether_data_contracts::repository::video_tasks::{
|
||||
StoredVideoTask, VideoTaskQueryFilter, VideoTaskStatus,
|
||||
@@ -18,9 +21,6 @@ const AI_PUBLIC_METHOD_NOT_ALLOWED_DETAIL: &str = "Method not allowed";
|
||||
const AI_PUBLIC_UNAUTHORIZED_DETAIL: &str = "Unauthorized";
|
||||
const OPENAI_IMAGE_PROMPT_DETAIL: &str = "图片生成/编辑请求缺少 prompt";
|
||||
const OPENAI_IMAGE_EDIT_INPUT_DETAIL: &str = "图片编辑请求至少需要 1 张输入图片";
|
||||
const OPENAI_IMAGE_VARIATION_INPUT_DETAIL: &str = "图片变体请求需要 image 文件";
|
||||
const OPENAI_IMAGE_N_DETAIL: &str = "当前 Codex 图片反代仅支持 n=1";
|
||||
const OPENAI_IMAGE_STREAM_VARIATION_DETAIL: &str = "图片变体接口当前仅支持同步响应";
|
||||
const OPENAI_IMAGE_PARTIAL_IMAGES_DETAIL: &str =
|
||||
"partial_images 仅支持 0-3,且必须配合 stream=true";
|
||||
const OPENAI_IMAGE_STYLE_DETAIL: &str = "当前 Codex 图片反代暂不支持 style 参数";
|
||||
@@ -57,7 +57,6 @@ const OPENAI_RERANK_STREAM_UNSUPPORTED_DETAIL: &str = "Rerank requests do not su
|
||||
enum OpenAiImageOperation {
|
||||
Generate,
|
||||
Edit,
|
||||
Variation,
|
||||
}
|
||||
|
||||
impl OpenAiImageOperation {
|
||||
@@ -65,7 +64,6 @@ impl OpenAiImageOperation {
|
||||
match path {
|
||||
"/v1/images/generations" => Some(Self::Generate),
|
||||
"/v1/images/edits" => Some(Self::Edit),
|
||||
"/v1/images/variations" => Some(Self::Variation),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -202,7 +200,7 @@ fn maybe_build_local_openai_request_validation_response(
|
||||
if decision.route_kind.as_deref() != Some("image")
|
||||
|| !matches!(
|
||||
request_context.request_path.as_str(),
|
||||
"/v1/images/generations" | "/v1/images/edits" | "/v1/images/variations"
|
||||
"/v1/images/generations" | "/v1/images/edits"
|
||||
)
|
||||
{
|
||||
return None;
|
||||
@@ -240,19 +238,13 @@ fn maybe_build_local_openai_request_validation_response(
|
||||
OPENAI_IMAGE_EDIT_INPUT_DETAIL,
|
||||
));
|
||||
}
|
||||
OpenAiImageOperation::Variation if validation.image_count == 0 => {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
OPENAI_IMAGE_VARIATION_INPUT_DETAIL,
|
||||
));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
if validation.n.is_some_and(|value| value != 1) {
|
||||
if let Some(detail) = validate_openai_image_n(&validation) {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
OPENAI_IMAGE_N_DETAIL,
|
||||
detail,
|
||||
));
|
||||
}
|
||||
|
||||
@@ -272,15 +264,6 @@ fn maybe_build_local_openai_request_validation_response(
|
||||
));
|
||||
}
|
||||
|
||||
if validation.stream {
|
||||
if operation == OpenAiImageOperation::Variation {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
OPENAI_IMAGE_STREAM_VARIATION_DETAIL,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
if validation
|
||||
.response_format
|
||||
.as_deref()
|
||||
@@ -360,6 +343,23 @@ fn maybe_build_local_openai_request_validation_response(
|
||||
None
|
||||
}
|
||||
|
||||
fn openai_image_n_detail(max_generation_count: u64) -> String {
|
||||
if max_generation_count >= openai_image_gateway_max_generation_count() {
|
||||
format!("当前图片反代仅支持 n=1..{max_generation_count}")
|
||||
} else {
|
||||
format!("当前图片模型仅支持 n=1..{max_generation_count}")
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_openai_image_n(validation: &OpenAiImageValidationInput) -> Option<String> {
|
||||
let max_generation_count =
|
||||
openai_image_gateway_max_generation_count_for_model(validation.model.as_deref());
|
||||
validation
|
||||
.n
|
||||
.is_some_and(|value| value == 0 || value > max_generation_count)
|
||||
.then(|| openai_image_n_detail(max_generation_count))
|
||||
}
|
||||
|
||||
fn validate_openai_embedding_request(
|
||||
content_type: Option<&str>,
|
||||
request_body: &Bytes,
|
||||
@@ -535,7 +535,6 @@ fn parse_openai_image_validation_input(
|
||||
OpenAiImageOperation::Generate | OpenAiImageOperation::Edit => {
|
||||
OPENAI_IMAGE_PROMPT_DETAIL
|
||||
}
|
||||
OpenAiImageOperation::Variation => OPENAI_IMAGE_VARIATION_INPUT_DETAIL,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -1260,7 +1259,8 @@ fn estimate_text_tokens(text: &str) -> u64 {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
estimate_claude_count_tokens, parse_openai_image_validation_input, OpenAiImageOperation,
|
||||
estimate_claude_count_tokens, parse_openai_image_validation_input, validate_openai_image_n,
|
||||
OpenAiImageOperation,
|
||||
};
|
||||
use axum::body::Bytes;
|
||||
use serde_json::json;
|
||||
@@ -1344,4 +1344,31 @@ mod tests {
|
||||
assert_eq!(validation.prompt.as_deref(), Some("edit this image"));
|
||||
assert_eq!(validation.image_count, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn image_validation_restricts_multi_image_count_to_grok_models() {
|
||||
let openai_body = Bytes::from_static(br#"{"model":"gpt-image-2","prompt":"draw","n":2}"#);
|
||||
let openai_validation = parse_openai_image_validation_input(
|
||||
OpenAiImageOperation::Generate,
|
||||
Some("application/json"),
|
||||
&openai_body,
|
||||
)
|
||||
.expect("valid image payload should parse");
|
||||
|
||||
assert_eq!(
|
||||
validate_openai_image_n(&openai_validation).as_deref(),
|
||||
Some("当前图片模型仅支持 n=1..1")
|
||||
);
|
||||
|
||||
let grok_body =
|
||||
Bytes::from_static(br#"{"model":"grok-imagine-image-lite","prompt":"draw","n":4}"#);
|
||||
let grok_validation = parse_openai_image_validation_input(
|
||||
OpenAiImageOperation::Generate,
|
||||
Some("application/json"),
|
||||
&grok_body,
|
||||
)
|
||||
.expect("valid grok image payload should parse");
|
||||
|
||||
assert!(validate_openai_image_n(&grok_validation).is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -32,6 +32,7 @@ struct AdminAnnouncementCreateRequest {
|
||||
kind: String,
|
||||
priority: Option<i32>,
|
||||
is_pinned: Option<bool>,
|
||||
requires_ack: Option<bool>,
|
||||
start_time: Option<String>,
|
||||
end_time: Option<String>,
|
||||
}
|
||||
@@ -45,6 +46,7 @@ struct AdminAnnouncementUpdateRequest {
|
||||
priority: Option<i32>,
|
||||
is_active: Option<bool>,
|
||||
is_pinned: Option<bool>,
|
||||
requires_ack: Option<bool>,
|
||||
start_time: Option<String>,
|
||||
end_time: Option<String>,
|
||||
}
|
||||
@@ -168,6 +170,7 @@ fn build_create_record(
|
||||
kind: payload.kind,
|
||||
priority: payload.priority.unwrap_or(0),
|
||||
is_pinned: payload.is_pinned.unwrap_or(false),
|
||||
requires_ack: payload.requires_ack.unwrap_or(false),
|
||||
author_id: operator_id,
|
||||
start_time_unix_secs: parse_optional_rfc3339_unix_secs(
|
||||
payload.start_time.as_deref(),
|
||||
@@ -194,6 +197,7 @@ fn build_update_record(
|
||||
priority: payload.priority,
|
||||
is_active: payload.is_active,
|
||||
is_pinned: payload.is_pinned,
|
||||
requires_ack: payload.requires_ack,
|
||||
start_time_unix_secs: parse_optional_rfc3339_unix_secs(
|
||||
payload.start_time.as_deref(),
|
||||
"start_time",
|
||||
|
||||
@@ -78,6 +78,7 @@ pub(super) fn build_public_announcement_payload(
|
||||
"priority": announcement.priority,
|
||||
"is_active": announcement.is_active,
|
||||
"is_pinned": announcement.is_pinned,
|
||||
"requires_ack": announcement.requires_ack,
|
||||
"author": {
|
||||
"id": announcement.author_id,
|
||||
"username": announcement.author_username,
|
||||
|
||||
@@ -13,7 +13,7 @@ use super::super::{build_unhandled_public_support_response, resolve_authenticate
|
||||
use super::announcements_shared::{
|
||||
announcements_bad_request_response, announcements_internal_detail,
|
||||
announcements_internal_error_response, announcements_not_found_response,
|
||||
read_status_announcement_id_from_path,
|
||||
build_public_announcement_payload, read_status_announcement_id_from_path,
|
||||
};
|
||||
|
||||
#[derive(Debug, serde::Deserialize)]
|
||||
@@ -75,6 +75,37 @@ pub(crate) async fn maybe_build_local_announcement_user_response(
|
||||
};
|
||||
Some(Json(json!({ "unread_count": unread_count })).into_response())
|
||||
}
|
||||
Some("required_unread")
|
||||
if request_context.request_method == http::Method::GET
|
||||
&& matches!(
|
||||
request_context.request_path.as_str(),
|
||||
"/api/announcements/users/me/required-unread"
|
||||
| "/api/announcements/users/me/required-unread/"
|
||||
) =>
|
||||
{
|
||||
let items = match state
|
||||
.list_required_unread_active_announcements(&auth.user.id, now_unix_secs, 20)
|
||||
.await
|
||||
{
|
||||
Ok(value) => value,
|
||||
Err(err) => {
|
||||
return Some(announcements_internal_error_response(
|
||||
announcements_internal_detail(err),
|
||||
))
|
||||
}
|
||||
};
|
||||
let payload_items = items
|
||||
.iter()
|
||||
.map(build_public_announcement_payload)
|
||||
.collect::<Vec<_>>();
|
||||
Some(
|
||||
Json(json!({
|
||||
"items": payload_items,
|
||||
"total": payload_items.len(),
|
||||
}))
|
||||
.into_response(),
|
||||
)
|
||||
}
|
||||
Some("read_all")
|
||||
if request_context.request_method == http::Method::POST
|
||||
&& matches!(
|
||||
|
||||
@@ -26,6 +26,18 @@ pub(crate) async fn build_auth_registration_settings_payload(
|
||||
let turnstile_site_key_config = state
|
||||
.read_system_config_json_value("turnstile_site_key")
|
||||
.await?;
|
||||
let privacy_enabled_config = state
|
||||
.read_system_config_json_value("registration_privacy_policy_enabled")
|
||||
.await?;
|
||||
let privacy_format_config = state
|
||||
.read_system_config_json_value("registration_privacy_policy_format")
|
||||
.await?;
|
||||
let privacy_content_config = state
|
||||
.read_system_config_json_value("registration_privacy_policy_content")
|
||||
.await?;
|
||||
let privacy_version_config = state
|
||||
.read_system_config_json_value("registration_privacy_policy_version")
|
||||
.await?;
|
||||
|
||||
let email_configured = smtp_host
|
||||
.as_ref()
|
||||
@@ -48,6 +60,15 @@ pub(crate) async fn build_auth_registration_settings_payload(
|
||||
};
|
||||
let turnstile_enabled = system_config_bool(turnstile_enabled_config.as_ref(), false);
|
||||
let turnstile_site_key = system_config_string(turnstile_site_key_config.as_ref());
|
||||
let privacy_policy_enabled = system_config_bool(privacy_enabled_config.as_ref(), false);
|
||||
let privacy_policy_format = match system_config_string(privacy_format_config.as_ref()) {
|
||||
Some(value) if matches!(value.as_str(), "markdown" | "html") => value,
|
||||
_ => "markdown".to_string(),
|
||||
};
|
||||
let privacy_policy_content =
|
||||
system_config_string(privacy_content_config.as_ref()).unwrap_or_default();
|
||||
let privacy_policy_version =
|
||||
system_config_string(privacy_version_config.as_ref()).unwrap_or_else(|| "1".to_string());
|
||||
|
||||
Ok(json!({
|
||||
"enable_registration": enable_registration,
|
||||
@@ -57,6 +78,12 @@ pub(crate) async fn build_auth_registration_settings_payload(
|
||||
"turnstile_enabled": turnstile_enabled,
|
||||
"turnstile_site_key": turnstile_site_key,
|
||||
"turnstile_required_actions": ["send_verification_code", "register"],
|
||||
"privacy_policy": {
|
||||
"enabled": privacy_policy_enabled,
|
||||
"format": privacy_policy_format,
|
||||
"content": privacy_policy_content,
|
||||
"version": privacy_policy_version,
|
||||
},
|
||||
}))
|
||||
}
|
||||
|
||||
|
||||
@@ -18,6 +18,9 @@ struct AuthRegisterRequest {
|
||||
username: String,
|
||||
password: String,
|
||||
turnstile_token: Option<String>,
|
||||
invite_code: Option<String>,
|
||||
privacy_policy_accepted: Option<bool>,
|
||||
privacy_policy_version: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -131,6 +134,26 @@ pub(crate) fn validate_auth_register_password(password: &str, policy: &str) -> R
|
||||
Ok(())
|
||||
}
|
||||
|
||||
struct RegistrationPrivacyPolicySettings {
|
||||
enabled: bool,
|
||||
version: String,
|
||||
}
|
||||
|
||||
async fn read_registration_privacy_policy_settings(
|
||||
state: &AppState,
|
||||
) -> Result<RegistrationPrivacyPolicySettings, GatewayError> {
|
||||
let enabled = state
|
||||
.read_system_config_json_value("registration_privacy_policy_enabled")
|
||||
.await?;
|
||||
let version = state
|
||||
.read_system_config_json_value("registration_privacy_policy_version")
|
||||
.await?;
|
||||
Ok(RegistrationPrivacyPolicySettings {
|
||||
enabled: system_config_bool(enabled.as_ref(), false),
|
||||
version: system_config_string(version.as_ref()).unwrap_or_else(|| "1".to_string()),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn auth_password_policy_level(state: &AppState) -> Result<String, GatewayError> {
|
||||
let config = state
|
||||
.read_system_config_json_value("password_policy_level")
|
||||
@@ -388,6 +411,31 @@ pub(super) async fn handle_auth_register(
|
||||
if !enable_registration {
|
||||
return build_auth_error_response(http::StatusCode::FORBIDDEN, "系统暂不开放注册", false);
|
||||
}
|
||||
let privacy_policy = match read_registration_privacy_policy_settings(state).await {
|
||||
Ok(value) => value,
|
||||
Err(err) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("auth settings lookup failed: {err:?}"),
|
||||
false,
|
||||
);
|
||||
}
|
||||
};
|
||||
if privacy_policy.enabled {
|
||||
let accepted = payload.privacy_policy_accepted.unwrap_or(false);
|
||||
let accepted_version = payload
|
||||
.privacy_policy_version
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.unwrap_or_default();
|
||||
if !accepted || accepted_version != privacy_policy.version {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"请先阅读并同意当前版本的隐私政策",
|
||||
false,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if let Err(response) = verify_auth_turnstile(
|
||||
state,
|
||||
@@ -555,6 +603,63 @@ pub(super) async fn handle_auth_register(
|
||||
false,
|
||||
);
|
||||
}
|
||||
if privacy_policy.enabled {
|
||||
match state
|
||||
.record_user_privacy_policy_acceptance(&user.id, &privacy_policy.version)
|
||||
.await
|
||||
{
|
||||
Ok(true) => {}
|
||||
Ok(false) => {
|
||||
let _ = state.delete_local_auth_user(&user.id).await;
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
AUTH_REGISTRATION_STORAGE_UNAVAILABLE_DETAIL,
|
||||
false,
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
let _ = state.delete_local_auth_user(&user.id).await;
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("auth privacy policy acceptance failed: {err:?}"),
|
||||
false,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
let invite_code = payload
|
||||
.invite_code
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
if invite_code.is_some() {
|
||||
let source = json!({
|
||||
"channel": "registration",
|
||||
"ip": cf_connecting_ip,
|
||||
"user_agent": headers
|
||||
.get(http::header::USER_AGENT)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
});
|
||||
if let Err(err) = state
|
||||
.bind_referral_invite_after_registration(
|
||||
&user.id,
|
||||
user.email_verified,
|
||||
invite_code,
|
||||
Some(source),
|
||||
)
|
||||
.await
|
||||
{
|
||||
let _ = state.delete_local_auth_user(&user.id).await;
|
||||
let (status, detail) = match err {
|
||||
GatewayError::Client { status, message } => (status, message),
|
||||
other => (
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("auth referral binding failed: {other:?}"),
|
||||
),
|
||||
};
|
||||
return build_auth_error_response(status, detail, false);
|
||||
}
|
||||
}
|
||||
|
||||
if require_verification {
|
||||
if let Some(email) = email.as_deref() {
|
||||
|
||||
@@ -14,11 +14,11 @@ use super::{
|
||||
|
||||
const INSTALL_SESSION_TTL_SECS: u64 = 15 * 60;
|
||||
const INSTALL_SESSION_KEY_PREFIX: &str = "install:session:";
|
||||
const PROXY_INSTALL_SESSION_KEY_PREFIX: &str = "proxy-install:session:";
|
||||
const PROXY_INSTALL_UNIX_SCRIPT_URL: &str =
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-proxy/install.sh";
|
||||
const PROXY_INSTALL_POWERSHELL_SCRIPT_URL: &str =
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-proxy/install.ps1";
|
||||
const TUNNEL_INSTALL_SESSION_KEY_PREFIX: &str = "tunnel-install:session:";
|
||||
const TUNNEL_INSTALL_UNIX_SCRIPT_URL: &str =
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh";
|
||||
const TUNNEL_INSTALL_POWERSHELL_SCRIPT_URL: &str =
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1";
|
||||
|
||||
#[derive(Debug, Clone, Copy, Deserialize, Serialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
@@ -55,7 +55,7 @@ struct StoredInstallSession {
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
struct StoredProxyInstallSession {
|
||||
struct StoredTunnelInstallSession {
|
||||
aether_url: String,
|
||||
management_token: String,
|
||||
node_name: String,
|
||||
@@ -91,9 +91,10 @@ fn install_code_from_path(request_path: &str) -> Option<(String, bool)> {
|
||||
(!code.is_empty()).then(|| (code.to_string(), is_powershell))
|
||||
}
|
||||
|
||||
fn proxy_install_code_from_path(request_path: &str) -> Option<(String, bool)> {
|
||||
fn tunnel_install_code_from_path(request_path: &str) -> Option<(String, bool)> {
|
||||
let raw = request_path
|
||||
.strip_prefix("/install-proxy/")?
|
||||
.strip_prefix("/install-tunnel/")
|
||||
.or_else(|| request_path.strip_prefix("/install-proxy/"))?
|
||||
.trim()
|
||||
.trim_matches('/');
|
||||
if raw.is_empty() || raw.contains('/') {
|
||||
@@ -108,8 +109,8 @@ fn install_session_runtime_key(code: &str) -> String {
|
||||
format!("{INSTALL_SESSION_KEY_PREFIX}{code}")
|
||||
}
|
||||
|
||||
fn proxy_install_session_runtime_key(code: &str) -> String {
|
||||
format!("{PROXY_INSTALL_SESSION_KEY_PREFIX}{code}")
|
||||
fn tunnel_install_session_runtime_key(code: &str) -> String {
|
||||
format!("{TUNNEL_INSTALL_SESSION_KEY_PREFIX}{code}")
|
||||
}
|
||||
|
||||
fn generate_install_code() -> String {
|
||||
@@ -164,42 +165,42 @@ fn powershell_single_quote(value: &str) -> String {
|
||||
format!("'{}'", value.replace('\'', "''"))
|
||||
}
|
||||
|
||||
fn build_proxy_unix_script(session: &StoredProxyInstallSession) -> String {
|
||||
fn build_tunnel_unix_script(session: &StoredTunnelInstallSession) -> String {
|
||||
format!(
|
||||
r###"#!/bin/sh
|
||||
set -eu
|
||||
export AETHER_PROXY_AETHER_URL={aether_url}
|
||||
export AETHER_PROXY_MANAGEMENT_TOKEN={management_token}
|
||||
export AETHER_PROXY_NODE_NAME={node_name}
|
||||
export AETHER_TUNNEL_AETHER_URL={aether_url}
|
||||
export AETHER_TUNNEL_MANAGEMENT_TOKEN={management_token}
|
||||
export AETHER_TUNNEL_NODE_NAME={node_name}
|
||||
|
||||
if command -v curl >/dev/null 2>&1; then
|
||||
curl -fsSL {script_url} | sh
|
||||
elif command -v wget >/dev/null 2>&1; then
|
||||
wget -qO- {script_url} | sh
|
||||
else
|
||||
printf '%s\n' "[Aether Proxy] 需要 curl 或 wget 下载安装脚本" >&2
|
||||
printf '%s\n' "[Aether Tunnel] 需要 curl 或 wget 下载安装脚本" >&2
|
||||
exit 1
|
||||
fi
|
||||
"###,
|
||||
aether_url = shell_single_quote(&session.aether_url),
|
||||
management_token = shell_single_quote(&session.management_token),
|
||||
node_name = shell_single_quote(&session.node_name),
|
||||
script_url = shell_single_quote(PROXY_INSTALL_UNIX_SCRIPT_URL),
|
||||
script_url = shell_single_quote(TUNNEL_INSTALL_UNIX_SCRIPT_URL),
|
||||
)
|
||||
}
|
||||
|
||||
fn build_proxy_powershell_script(session: &StoredProxyInstallSession) -> String {
|
||||
fn build_tunnel_powershell_script(session: &StoredTunnelInstallSession) -> String {
|
||||
format!(
|
||||
r###"$ErrorActionPreference = 'Stop'
|
||||
$env:AETHER_PROXY_AETHER_URL = {aether_url}
|
||||
$env:AETHER_PROXY_MANAGEMENT_TOKEN = {management_token}
|
||||
$env:AETHER_PROXY_NODE_NAME = {node_name}
|
||||
$env:AETHER_TUNNEL_AETHER_URL = {aether_url}
|
||||
$env:AETHER_TUNNEL_MANAGEMENT_TOKEN = {management_token}
|
||||
$env:AETHER_TUNNEL_NODE_NAME = {node_name}
|
||||
irm {script_url} | iex
|
||||
"###,
|
||||
aether_url = powershell_single_quote(&session.aether_url),
|
||||
management_token = powershell_single_quote(&session.management_token),
|
||||
node_name = powershell_single_quote(&session.node_name),
|
||||
script_url = powershell_single_quote(PROXY_INSTALL_POWERSHELL_SCRIPT_URL),
|
||||
script_url = powershell_single_quote(TUNNEL_INSTALL_POWERSHELL_SCRIPT_URL),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -659,7 +660,7 @@ pub(crate) async fn build_proxy_node_install_session_response(
|
||||
) -> Response<Body> {
|
||||
let code = generate_install_code();
|
||||
let expires_at_unix_secs = unix_secs_now().saturating_add(INSTALL_SESSION_TTL_SECS);
|
||||
let session = StoredProxyInstallSession {
|
||||
let session = StoredTunnelInstallSession {
|
||||
aether_url: base_url_from_request(headers, request_context),
|
||||
management_token,
|
||||
node_name,
|
||||
@@ -670,14 +671,14 @@ pub(crate) async fn build_proxy_node_install_session_response(
|
||||
Err(err) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("proxy install session serialize failed: {err:?}"),
|
||||
format!("tunnel install session serialize failed: {err:?}"),
|
||||
false,
|
||||
)
|
||||
}
|
||||
};
|
||||
if let Err(err) = state
|
||||
.runtime_kv_setex(
|
||||
&proxy_install_session_runtime_key(&code),
|
||||
&tunnel_install_session_runtime_key(&code),
|
||||
&serialized,
|
||||
INSTALL_SESSION_TTL_SECS,
|
||||
)
|
||||
@@ -685,7 +686,7 @@ pub(crate) async fn build_proxy_node_install_session_response(
|
||||
{
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("proxy install session create failed: {err:?}"),
|
||||
format!("tunnel install session create failed: {err:?}"),
|
||||
false,
|
||||
);
|
||||
}
|
||||
@@ -697,8 +698,8 @@ pub(crate) async fn build_proxy_node_install_session_response(
|
||||
"expires_in_seconds": INSTALL_SESSION_TTL_SECS,
|
||||
"node_name": session.node_name,
|
||||
"aether_url": session.aether_url,
|
||||
"unix_command": format!("curl -fsSL {base_url}/install-proxy/{code} | sh"),
|
||||
"powershell_command": format!("irm {base_url}/install-proxy/{code}.ps1 | iex"),
|
||||
"unix_command": format!("curl -fsSL {base_url}/install-tunnel/{code} | sh"),
|
||||
"powershell_command": format!("irm {base_url}/install-tunnel/{code}.ps1 | iex"),
|
||||
}))
|
||||
.into_response()
|
||||
}
|
||||
@@ -711,8 +712,10 @@ pub(super) async fn maybe_build_local_install_response(
|
||||
if decision.route_family.as_deref() != Some("install") {
|
||||
return None;
|
||||
}
|
||||
if request_context.request_path.starts_with("/install-proxy/") {
|
||||
return Some(maybe_build_local_proxy_install_response(state, request_context).await);
|
||||
if request_context.request_path.starts_with("/install-tunnel/")
|
||||
|| request_context.request_path.starts_with("/install-proxy/")
|
||||
{
|
||||
return Some(maybe_build_local_tunnel_install_response(state, request_context).await);
|
||||
}
|
||||
let Some((code, wants_powershell)) = install_code_from_path(&request_context.request_path)
|
||||
else {
|
||||
@@ -789,45 +792,45 @@ pub(super) async fn maybe_build_local_install_response(
|
||||
Some(response)
|
||||
}
|
||||
|
||||
async fn maybe_build_local_proxy_install_response(
|
||||
async fn maybe_build_local_tunnel_install_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
) -> Response<Body> {
|
||||
let Some((code, wants_powershell)) =
|
||||
proxy_install_code_from_path(&request_context.request_path)
|
||||
tunnel_install_code_from_path(&request_context.request_path)
|
||||
else {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"proxy install code 不存在或已失效",
|
||||
"tunnel install code 不存在或已失效",
|
||||
false,
|
||||
);
|
||||
};
|
||||
let raw = match state
|
||||
.runtime_kv_getdel(&proxy_install_session_runtime_key(&code))
|
||||
.runtime_kv_getdel(&tunnel_install_session_runtime_key(&code))
|
||||
.await
|
||||
{
|
||||
Ok(Some(value)) => value,
|
||||
Ok(None) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"proxy install code 不存在、已过期或已使用",
|
||||
"tunnel install code 不存在、已过期或已使用",
|
||||
false,
|
||||
)
|
||||
}
|
||||
Err(err) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("proxy install session lookup failed: {err:?}"),
|
||||
format!("tunnel install session lookup failed: {err:?}"),
|
||||
false,
|
||||
)
|
||||
}
|
||||
};
|
||||
let session = match serde_json::from_str::<StoredProxyInstallSession>(&raw) {
|
||||
let session = match serde_json::from_str::<StoredTunnelInstallSession>(&raw) {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"proxy install code 数据无效",
|
||||
"tunnel install code 数据无效",
|
||||
false,
|
||||
)
|
||||
}
|
||||
@@ -835,14 +838,14 @@ async fn maybe_build_local_proxy_install_response(
|
||||
if session.expires_at_unix_secs <= unix_secs_now() {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"proxy install code 已过期",
|
||||
"tunnel install code 已过期",
|
||||
false,
|
||||
);
|
||||
}
|
||||
let body = if wants_powershell {
|
||||
build_proxy_powershell_script(&session)
|
||||
build_tunnel_powershell_script(&session)
|
||||
} else {
|
||||
build_proxy_unix_script(&session)
|
||||
build_tunnel_unix_script(&session)
|
||||
};
|
||||
let content_type = if wants_powershell {
|
||||
"text/plain; charset=utf-8"
|
||||
@@ -885,8 +888,8 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn test_proxy_session() -> StoredProxyInstallSession {
|
||||
StoredProxyInstallSession {
|
||||
fn test_tunnel_session() -> StoredTunnelInstallSession {
|
||||
StoredTunnelInstallSession {
|
||||
aether_url: "https://aether.example".to_string(),
|
||||
management_token: "ae-test-token".to_string(),
|
||||
node_name: "jp-proxy-01".to_string(),
|
||||
@@ -895,41 +898,45 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_install_path_accepts_shell_and_powershell_codes() {
|
||||
fn tunnel_install_path_accepts_shell_and_powershell_codes() {
|
||||
assert_eq!(
|
||||
proxy_install_code_from_path("/install-proxy/abc123"),
|
||||
tunnel_install_code_from_path("/install-tunnel/abc123"),
|
||||
Some(("abc123".to_string(), false))
|
||||
);
|
||||
assert_eq!(
|
||||
proxy_install_code_from_path("/install-proxy/abc123.ps1"),
|
||||
tunnel_install_code_from_path("/install-tunnel/abc123.ps1"),
|
||||
Some(("abc123".to_string(), true))
|
||||
);
|
||||
assert_eq!(proxy_install_code_from_path("/install-proxy/a/b"), None);
|
||||
assert_eq!(
|
||||
tunnel_install_code_from_path("/install-proxy/abc123"),
|
||||
Some(("abc123".to_string(), false))
|
||||
);
|
||||
assert_eq!(tunnel_install_code_from_path("/install-tunnel/a/b"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_unix_script_exports_session_values_and_reuses_proxy_installer() {
|
||||
let script = build_proxy_unix_script(&test_proxy_session());
|
||||
fn tunnel_unix_script_exports_session_values_and_reuses_tunnel_installer() {
|
||||
let script = build_tunnel_unix_script(&test_tunnel_session());
|
||||
|
||||
assert!(script.contains("export AETHER_PROXY_AETHER_URL='https://aether.example'"));
|
||||
assert!(script.contains("export AETHER_PROXY_MANAGEMENT_TOKEN='ae-test-token'"));
|
||||
assert!(script.contains("export AETHER_PROXY_NODE_NAME='jp-proxy-01'"));
|
||||
assert!(script.contains("export AETHER_TUNNEL_AETHER_URL='https://aether.example'"));
|
||||
assert!(script.contains("export AETHER_TUNNEL_MANAGEMENT_TOKEN='ae-test-token'"));
|
||||
assert!(script.contains("export AETHER_TUNNEL_NODE_NAME='jp-proxy-01'"));
|
||||
assert!(script.contains(
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-proxy/install.sh"
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.sh"
|
||||
));
|
||||
assert!(!script.contains("aether-rust-pioneer"));
|
||||
assert!(!script.contains("[[servers]]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn proxy_powershell_script_exports_session_values_and_reuses_proxy_installer() {
|
||||
let script = build_proxy_powershell_script(&test_proxy_session());
|
||||
fn tunnel_powershell_script_exports_session_values_and_reuses_tunnel_installer() {
|
||||
let script = build_tunnel_powershell_script(&test_tunnel_session());
|
||||
|
||||
assert!(script.contains("$env:AETHER_PROXY_AETHER_URL = 'https://aether.example'"));
|
||||
assert!(script.contains("$env:AETHER_PROXY_MANAGEMENT_TOKEN = 'ae-test-token'"));
|
||||
assert!(script.contains("$env:AETHER_PROXY_NODE_NAME = 'jp-proxy-01'"));
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_AETHER_URL = 'https://aether.example'"));
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_MANAGEMENT_TOKEN = 'ae-test-token'"));
|
||||
assert!(script.contains("$env:AETHER_TUNNEL_NODE_NAME = 'jp-proxy-01'"));
|
||||
assert!(script.contains(
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-proxy/install.ps1"
|
||||
"https://raw.githubusercontent.com/fawney19/Aether/main/apps/aether-tunnel/install.ps1"
|
||||
));
|
||||
assert!(!script.contains("aether-rust-pioneer"));
|
||||
assert!(!script.contains("[[servers]]"));
|
||||
|
||||
@@ -3,6 +3,7 @@ use std::collections::BTreeMap;
|
||||
use axum::{body::Body, http, response::Response};
|
||||
use md5::{Digest, Md5};
|
||||
use serde_json::json;
|
||||
use tracing::warn;
|
||||
|
||||
use super::{payment_shared::payment_callback_payload_hash, AppState, GatewayPublicRequestContext};
|
||||
|
||||
@@ -373,9 +374,20 @@ pub(super) async fn handle_epay_notify(
|
||||
|
||||
match outcome {
|
||||
Ok(Some(aether_data::repository::wallet::ProcessPaymentCallbackOutcome::Applied {
|
||||
order,
|
||||
order_id,
|
||||
..
|
||||
}))
|
||||
| Ok(Some(
|
||||
})) => {
|
||||
if let Err(err) = state.apply_referral_rewards_for_paid_order(&order).await {
|
||||
warn!(
|
||||
error = ?err,
|
||||
order_id = %order_id,
|
||||
"failed to apply referral rewards for epay callback"
|
||||
);
|
||||
}
|
||||
epay_plain(http::StatusCode::OK, "success")
|
||||
}
|
||||
Ok(Some(
|
||||
aether_data::repository::wallet::ProcessPaymentCallbackOutcome::AlreadyCredited {
|
||||
..
|
||||
},
|
||||
|
||||
@@ -10,6 +10,7 @@ use super::{
|
||||
build_auth_error_response, build_payment_callback_storage_unavailable_response, AppState,
|
||||
GatewayPublicRequestContext,
|
||||
};
|
||||
use tracing::warn;
|
||||
|
||||
pub(super) async fn handle_payment_callback_with_wallet_repository(
|
||||
state: &AppState,
|
||||
@@ -109,21 +110,30 @@ pub(super) async fn handle_payment_callback_with_wallet_repository(
|
||||
order_no,
|
||||
wallet_id,
|
||||
order,
|
||||
} => build_auth_json_response(
|
||||
http::StatusCode::OK,
|
||||
json!({
|
||||
"ok": true,
|
||||
"duplicate": duplicate,
|
||||
"credited": true,
|
||||
"order_id": order_id,
|
||||
"order_no": order_no,
|
||||
"status": order.status,
|
||||
"wallet_id": wallet_id,
|
||||
"payment_method": payment_method,
|
||||
"request_path": request_context.request_path,
|
||||
}),
|
||||
None,
|
||||
),
|
||||
} => {
|
||||
if let Err(err) = state.apply_referral_rewards_for_paid_order(&order).await {
|
||||
warn!(
|
||||
error = ?err,
|
||||
order_id = %order_id,
|
||||
"failed to apply referral rewards for credited payment order"
|
||||
);
|
||||
}
|
||||
build_auth_json_response(
|
||||
http::StatusCode::OK,
|
||||
json!({
|
||||
"ok": true,
|
||||
"duplicate": duplicate,
|
||||
"credited": true,
|
||||
"order_id": order_id,
|
||||
"order_no": order_no,
|
||||
"status": order.status,
|
||||
"wallet_id": wallet_id,
|
||||
"payment_method": payment_method,
|
||||
"request_path": request_context.request_path,
|
||||
}),
|
||||
None,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -169,7 +169,11 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
}
|
||||
|
||||
let mut provider_request_body = match format_value.as_str() {
|
||||
"openai:chat" | "claude:messages" => json!({
|
||||
"openai:chat" => json!({
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "Health check"}],
|
||||
}),
|
||||
"claude:messages" => json!({
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "Health check"}],
|
||||
"max_tokens": 5,
|
||||
@@ -179,9 +183,6 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
"role": "user",
|
||||
"parts": [{"text": "Health check"}],
|
||||
}],
|
||||
"generationConfig": {
|
||||
"maxOutputTokens": 5,
|
||||
},
|
||||
}),
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
@@ -31,6 +31,9 @@ use user_me_catalog::*;
|
||||
#[path = "user_me_preferences.rs"]
|
||||
mod user_me_preferences;
|
||||
use user_me_preferences::*;
|
||||
#[path = "user_me_referral.rs"]
|
||||
mod user_me_referral;
|
||||
use user_me_referral::*;
|
||||
#[path = "user_me_profile.rs"]
|
||||
mod user_me_profile;
|
||||
use user_me_profile::*;
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
use super::{
|
||||
build_auth_error_response, resolve_authenticated_local_user, AppState,
|
||||
GatewayPublicRequestContext,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn handle_users_me_referral_get(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
headers: &http::HeaderMap,
|
||||
) -> Response<Body> {
|
||||
let auth = match resolve_authenticated_local_user(state, request_context, headers).await {
|
||||
Ok(value) => value,
|
||||
Err(response) => return response,
|
||||
};
|
||||
if !state.has_referral_data_backend() {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"邀请返利数据暂不可用",
|
||||
false,
|
||||
);
|
||||
}
|
||||
let dashboard = match state.referral_dashboard(&auth.user.id).await {
|
||||
Ok(Some(value)) => value,
|
||||
Ok(None) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"邀请返利数据暂不可用",
|
||||
false,
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("referral dashboard failed: {err:?}"),
|
||||
false,
|
||||
);
|
||||
}
|
||||
};
|
||||
let base = headers
|
||||
.get("origin")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_default();
|
||||
let invitation_link = if base.is_empty() {
|
||||
format!("/register?invite={}", dashboard.invite_code)
|
||||
} else {
|
||||
format!("{base}/register?invite={}", dashboard.invite_code)
|
||||
};
|
||||
Json(json!({
|
||||
"invite_code": dashboard.invite_code,
|
||||
"invitation_link": invitation_link,
|
||||
"summary": {
|
||||
"total_invites": dashboard.total_invites,
|
||||
"effective_invites": dashboard.effective_invites,
|
||||
"paid_reward_usd": dashboard.paid_reward_usd,
|
||||
"pending_reward_usd": dashboard.pending_reward_usd,
|
||||
"reversed_reward_usd": dashboard.reversed_reward_usd,
|
||||
}
|
||||
}))
|
||||
.into_response()
|
||||
}
|
||||
@@ -15,11 +15,12 @@ use super::{
|
||||
handle_users_me_management_tokens_list, handle_users_me_model_capabilities_get,
|
||||
handle_users_me_model_capabilities_put, handle_users_me_password_patch,
|
||||
handle_users_me_preferences_get, handle_users_me_preferences_put,
|
||||
handle_users_me_providers_get, handle_users_me_sessions_get, handle_users_me_update_session,
|
||||
handle_users_me_usage_active_get, handle_users_me_usage_get, handle_users_me_usage_heatmap_get,
|
||||
handle_users_me_usage_interval_timeline_get, users_me_api_key_capabilities_path_matches,
|
||||
users_me_api_key_detail_path_matches, users_me_api_key_install_sessions_path_matches,
|
||||
users_me_api_key_providers_path_matches, users_me_management_token_detail_path_matches,
|
||||
handle_users_me_providers_get, handle_users_me_referral_get, handle_users_me_sessions_get,
|
||||
handle_users_me_update_session, handle_users_me_usage_active_get, handle_users_me_usage_get,
|
||||
handle_users_me_usage_heatmap_get, handle_users_me_usage_interval_timeline_get,
|
||||
users_me_api_key_capabilities_path_matches, users_me_api_key_detail_path_matches,
|
||||
users_me_api_key_install_sessions_path_matches, users_me_api_key_providers_path_matches,
|
||||
users_me_management_token_detail_path_matches,
|
||||
users_me_management_token_regenerate_path_matches,
|
||||
users_me_management_token_toggle_path_matches, users_me_management_tokens_root,
|
||||
users_me_session_detail_path_matches, AppState, GatewayPublicRequestContext,
|
||||
@@ -211,6 +212,9 @@ pub(crate) async fn maybe_build_local_users_me_response(
|
||||
Some("preferences") if request_context.request_path == "/api/users/me/preferences" => {
|
||||
Some(handle_users_me_preferences_get(state, request_context, headers).await)
|
||||
}
|
||||
Some("referral") if request_context.request_path == "/api/users/me/referral" => {
|
||||
Some(handle_users_me_referral_get(state, request_context, headers).await)
|
||||
}
|
||||
Some("available_models")
|
||||
if request_context.request_path == "/api/users/me/available-models" =>
|
||||
{
|
||||
|
||||
@@ -321,6 +321,75 @@ fn users_me_usage_upstream_is_stream(item: &StoredRequestUsageAudit) -> bool {
|
||||
.unwrap_or(item.is_stream)
|
||||
}
|
||||
|
||||
fn users_me_usage_metadata_string<'a>(
|
||||
item: &'a StoredRequestUsageAudit,
|
||||
key: &str,
|
||||
) -> Option<&'a str> {
|
||||
item.request_metadata
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|metadata| metadata.get(key))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn infer_client_family_from_user_agent(user_agent: &str) -> Option<&'static str> {
|
||||
let normalized = user_agent.trim().to_ascii_lowercase();
|
||||
if normalized.is_empty() {
|
||||
return None;
|
||||
}
|
||||
if normalized.starts_with("codex_vscode") {
|
||||
return Some("codex_vscode");
|
||||
}
|
||||
if normalized.starts_with("codex") {
|
||||
return Some("codex");
|
||||
}
|
||||
if normalized.contains("claude-code") || normalized.contains("claude_code") {
|
||||
return Some("claude_code");
|
||||
}
|
||||
if normalized.contains("opencode") {
|
||||
return Some("opencode");
|
||||
}
|
||||
if normalized.contains("geminicli") || normalized.contains("gemini-cli") {
|
||||
return Some("gemini_cli");
|
||||
}
|
||||
if normalized.starts_with("openai/js") {
|
||||
return Some("openai_js_sdk");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn users_me_usage_client_family(item: &StoredRequestUsageAudit) -> Option<&str> {
|
||||
item.client_family
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.or_else(|| {
|
||||
item.request_metadata
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|metadata| {
|
||||
metadata
|
||||
.get("client_session_affinity")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|affinity| affinity.get("client_family"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.or_else(|| {
|
||||
metadata
|
||||
.get("client_family")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
})
|
||||
})
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
.or_else(|| {
|
||||
users_me_usage_metadata_string(item, "user_agent")
|
||||
.and_then(infer_client_family_from_user_agent)
|
||||
})
|
||||
}
|
||||
|
||||
fn build_users_me_usage_record_payload(
|
||||
item: &StoredRequestUsageAudit,
|
||||
include_actual_cost: bool,
|
||||
@@ -353,6 +422,11 @@ fn build_users_me_usage_record_payload(
|
||||
"upstream_is_stream": upstream_is_stream,
|
||||
"client_requested_stream": client_is_stream,
|
||||
"client_is_stream": client_is_stream,
|
||||
"client_family": users_me_usage_client_family(item),
|
||||
"client_ip": users_me_usage_metadata_string(item, "client_ip"),
|
||||
"user_agent": users_me_usage_metadata_string(item, "user_agent"),
|
||||
"request_path": users_me_usage_metadata_string(item, "request_path"),
|
||||
"request_path_and_query": users_me_usage_metadata_string(item, "request_path_and_query"),
|
||||
"status": item.status,
|
||||
"has_fallback": item.has_fallback(),
|
||||
"created_at": unix_secs_to_rfc3339(item.created_at_unix_ms),
|
||||
@@ -411,6 +485,9 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
|
||||
"client_requested_stream": client_is_stream,
|
||||
"client_is_stream": client_is_stream,
|
||||
"has_format_conversion": item.has_format_conversion,
|
||||
"client_family": users_me_usage_client_family(item),
|
||||
"client_ip": users_me_usage_metadata_string(item, "client_ip"),
|
||||
"user_agent": users_me_usage_metadata_string(item, "user_agent"),
|
||||
"target_model": item.target_model,
|
||||
"has_fallback": item.has_fallback(),
|
||||
});
|
||||
@@ -1368,6 +1445,41 @@ mod tests {
|
||||
assert_eq!(active_payload["client_is_stream"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_usage_payload_infers_client_family_from_user_agent() {
|
||||
let item = StoredRequestUsageAudit {
|
||||
request_metadata: Some(json!({
|
||||
"client_ip": "192.168.0.28",
|
||||
"user_agent": "codex_vscode/0.131.0-alpha.9 (Windows 10.0.26200; x86_64)"
|
||||
})),
|
||||
..sample_usage("completed")
|
||||
};
|
||||
|
||||
let record_payload =
|
||||
build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false);
|
||||
let active_payload = build_users_me_usage_active_payload(&item);
|
||||
|
||||
assert_eq!(record_payload["client_family"], "codex_vscode");
|
||||
assert_eq!(record_payload["client_ip"], "192.168.0.28");
|
||||
assert_eq!(active_payload["client_family"], "codex_vscode");
|
||||
assert_eq!(active_payload["client_ip"], "192.168.0.28");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_usage_payload_labels_openai_js_user_agent_as_sdk() {
|
||||
let item = StoredRequestUsageAudit {
|
||||
request_metadata: Some(json!({
|
||||
"user_agent": "OpenAI/JS 6.34.0"
|
||||
})),
|
||||
..sample_usage("completed")
|
||||
};
|
||||
|
||||
let record_payload =
|
||||
build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false);
|
||||
|
||||
assert_eq!(record_payload["client_family"], "openai_js_sdk");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_usage_stream_inference_falls_back_to_request_body_stream_flag() {
|
||||
let item = StoredRequestUsageAudit {
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use crate::handlers::shared::{json_string_list, unix_secs_to_rfc3339};
|
||||
use crate::provider_key_auth::{
|
||||
provider_key_auth_semantics, provider_key_configured_api_formats,
|
||||
provider_key_inherits_provider_api_formats,
|
||||
provider_key_auth_semantics, provider_key_can_refresh_oauth,
|
||||
provider_key_configured_api_formats, provider_key_inherits_provider_api_formats,
|
||||
};
|
||||
use crate::AppState;
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
@@ -10,6 +10,9 @@ use aether_admin::provider::status as admin_provider_status_pure;
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use aether_crypto::{decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use aether_provider_pool::{
|
||||
grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier,
|
||||
};
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::borrow::Cow;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -461,6 +464,24 @@ fn model_quota_window_snapshot(
|
||||
item: &Map<String, Value>,
|
||||
observed_at_unix_secs: Option<u64>,
|
||||
) -> Option<Value> {
|
||||
let remaining_value = item
|
||||
.get("remaining")
|
||||
.or_else(|| item.get("remaining_value"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let limit_value = item
|
||||
.get("total")
|
||||
.or_else(|| item.get("limit_value"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||
.filter(|value| *value > 0.0);
|
||||
let used_value = item
|
||||
.get("used")
|
||||
.or_else(|| item.get("used_value"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||
.or_else(|| {
|
||||
remaining_value
|
||||
.zip(limit_value)
|
||||
.map(|(remaining, limit)| (limit - remaining).max(0.0))
|
||||
});
|
||||
let used_ratio = item
|
||||
.get("used_percent")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||
@@ -489,6 +510,8 @@ fn model_quota_window_snapshot(
|
||||
&& reset_at.is_none()
|
||||
&& reset_seconds.is_none()
|
||||
&& is_exhausted.is_none()
|
||||
&& remaining_value.is_none()
|
||||
&& limit_value.is_none()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
@@ -507,12 +530,29 @@ fn model_quota_window_snapshot(
|
||||
window.insert("model".to_string(), json!(model_name));
|
||||
window.insert("used_ratio".to_string(), json!(used_ratio));
|
||||
window.insert("remaining_ratio".to_string(), json!(remaining_ratio));
|
||||
window.insert("used_value".to_string(), json!(used_value));
|
||||
window.insert("remaining_value".to_string(), json!(remaining_value));
|
||||
window.insert("limit_value".to_string(), json!(limit_value));
|
||||
window.insert("reset_at".to_string(), json!(reset_at));
|
||||
window.insert("reset_seconds".to_string(), json!(reset_seconds));
|
||||
window.insert("is_exhausted".to_string(), json!(is_exhausted));
|
||||
Some(Value::Object(window))
|
||||
}
|
||||
|
||||
fn provider_quota_metadata_string(
|
||||
metadata: &Map<String, Value>,
|
||||
fields: &[&str],
|
||||
) -> Option<String> {
|
||||
fields.iter().find_map(|field| {
|
||||
metadata
|
||||
.get(*field)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
fn quota_windows_usage_ratio(windows: &[Value]) -> Option<f64> {
|
||||
windows
|
||||
.iter()
|
||||
@@ -1126,6 +1166,72 @@ fn build_antigravity_quota_status_snapshot(
|
||||
}))
|
||||
}
|
||||
|
||||
fn build_grok_quota_status_snapshot(
|
||||
upstream_metadata: Option<&Value>,
|
||||
source: &str,
|
||||
) -> Option<Value> {
|
||||
let metadata = provider_quota_metadata_bucket(upstream_metadata, "grok")?;
|
||||
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
|
||||
let inferred_pool_tier = grok_pool_tier_from_quota_bucket(metadata);
|
||||
let pool_tier = provider_quota_metadata_string(metadata, &["pool_tier", "tier"])
|
||||
.or_else(|| inferred_pool_tier.map(ToOwned::to_owned));
|
||||
let plan_type = provider_quota_metadata_string(metadata, &["plan_type", "plan"])
|
||||
.or_else(|| pool_tier.clone());
|
||||
let supported_windows = grok_supported_quota_windows_for_tier(pool_tier.as_deref());
|
||||
let windows = provider_quota_model_bucket(metadata)
|
||||
.map(|models| {
|
||||
models
|
||||
.iter()
|
||||
.filter_map(|(model_name, item)| {
|
||||
if !supported_windows
|
||||
.iter()
|
||||
.any(|(quota_key, _)| *quota_key == model_name.as_str())
|
||||
{
|
||||
return None;
|
||||
}
|
||||
model_quota_window_snapshot(
|
||||
model_name,
|
||||
item.as_object()?,
|
||||
observed_at_unix_secs,
|
||||
)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
if windows.is_empty() && observed_at_unix_secs.is_none() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let usage_ratio = quota_windows_usage_ratio(&windows);
|
||||
let reset_seconds = quota_windows_min_reset_seconds(&windows);
|
||||
let reset_at = quota_windows_min_reset_at(&windows);
|
||||
let exhausted = quota_windows_all_exhausted(&windows);
|
||||
|
||||
Some(json!({
|
||||
"version": 2,
|
||||
"provider_type": "grok",
|
||||
"code": if exhausted { "exhausted" } else { "ok" },
|
||||
"label": if exhausted { Some("额度耗尽") } else { None::<&str> },
|
||||
"reason": if exhausted {
|
||||
Some("所有 Grok 模式额度已耗尽")
|
||||
} else {
|
||||
None::<&str>
|
||||
},
|
||||
"freshness": "fresh",
|
||||
"source": source,
|
||||
"observed_at": observed_at_unix_secs,
|
||||
"exhausted": exhausted,
|
||||
"usage_ratio": usage_ratio,
|
||||
"updated_at": observed_at_unix_secs,
|
||||
"reset_at": reset_at,
|
||||
"reset_seconds": reset_seconds,
|
||||
"plan_type": plan_type,
|
||||
"pool_tier": pool_tier,
|
||||
"windows": windows,
|
||||
}))
|
||||
}
|
||||
|
||||
fn build_gemini_cli_quota_status_snapshot(
|
||||
upstream_metadata: Option<&Value>,
|
||||
source: &str,
|
||||
@@ -1228,6 +1334,7 @@ pub(crate) fn sync_provider_key_quota_status_snapshot(
|
||||
"kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source),
|
||||
"chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source),
|
||||
"antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source),
|
||||
"grok" => build_grok_quota_status_snapshot(upstream_metadata, source),
|
||||
"gemini_cli" => build_gemini_cli_quota_status_snapshot(upstream_metadata, source),
|
||||
_ => None,
|
||||
}?;
|
||||
@@ -1590,7 +1697,10 @@ pub(crate) fn build_admin_provider_key_response(
|
||||
);
|
||||
payload.insert(
|
||||
"can_refresh_oauth".to_string(),
|
||||
json!(auth_semantics.can_refresh_oauth()),
|
||||
json!(provider_key_can_refresh_oauth(
|
||||
auth_semantics,
|
||||
auth_config.as_ref()
|
||||
)),
|
||||
);
|
||||
payload.insert(
|
||||
"can_export_oauth".to_string(),
|
||||
@@ -2099,6 +2209,69 @@ mod tests {
|
||||
assert_eq!(window.get("remaining_ratio"), Some(&json!(0.96)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_backfills_grok_model_quota() {
|
||||
let mut key = sample_catalog_key();
|
||||
key.upstream_metadata = Some(json!({
|
||||
"grok": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"pool_tier": "heavy",
|
||||
"plan_type": "heavy",
|
||||
"quota_by_model": {
|
||||
"quota_auto": {
|
||||
"display_name": "auto",
|
||||
"remaining_fraction": 0.4,
|
||||
"used_percent": 60.0,
|
||||
"remaining": 60.0,
|
||||
"total": 150.0,
|
||||
"reset_at": 1_778_157_172u64,
|
||||
"is_exhausted": false
|
||||
},
|
||||
"quota_heavy": {
|
||||
"display_name": "heavy",
|
||||
"remaining_fraction": 0.0,
|
||||
"used_percent": 100.0,
|
||||
"reset_at": 1_778_157_172u64,
|
||||
"is_exhausted": true
|
||||
}
|
||||
}
|
||||
}
|
||||
}));
|
||||
|
||||
let payload = provider_key_status_snapshot_payload(&key, "grok");
|
||||
let quota = payload
|
||||
.get("quota")
|
||||
.and_then(Value::as_object)
|
||||
.expect("quota snapshot should be object");
|
||||
let windows = quota
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)
|
||||
.expect("grok quota windows should exist");
|
||||
|
||||
assert_eq!(quota.get("provider_type"), Some(&json!("grok")));
|
||||
assert_eq!(quota.get("code"), Some(&json!("ok")));
|
||||
assert_eq!(quota.get("plan_type"), Some(&json!("heavy")));
|
||||
assert_eq!(quota.get("pool_tier"), Some(&json!("heavy")));
|
||||
assert_eq!(quota.get("exhausted"), Some(&json!(false)));
|
||||
assert_eq!(quota.get("usage_ratio"), Some(&json!(1.0)));
|
||||
assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64)));
|
||||
assert_eq!(windows.len(), 2);
|
||||
assert!(windows.iter().any(|window| {
|
||||
window
|
||||
.get("code")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|code| code == "model:quota_auto")
|
||||
}));
|
||||
let auto = windows
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.find(|window| window.get("code") == Some(&json!("model:quota_auto")))
|
||||
.expect("auto quota window should exist");
|
||||
assert_eq!(auto.get("remaining_value"), Some(&json!(60.0)));
|
||||
assert_eq!(auto.get("limit_value"), Some(&json!(150.0)));
|
||||
assert_eq!(auto.get("used_value"), Some(&json!(90.0)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_preserves_existing_materialized_quota_snapshot() {
|
||||
let mut key = sample_catalog_key();
|
||||
|
||||
Reference in New Issue
Block a user