mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
Merge origin/main into fix/gemini-cli-v1internal
This commit is contained in:
@@ -91,6 +91,26 @@ fn coerce_admin_provider_oauth_import_project_id(
|
||||
}
|
||||
}
|
||||
|
||||
fn json_import_expiry_value(value: Option<&serde_json::Value>) -> Option<u64> {
|
||||
let value = value?;
|
||||
json_u64_value(Some(value)).or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok())
|
||||
.and_then(|value| u64::try_from(value.timestamp()).ok())
|
||||
})
|
||||
}
|
||||
|
||||
fn json_import_expiry_from_keys(
|
||||
object: &serde_json::Map<String, serde_json::Value>,
|
||||
keys: &[&str],
|
||||
) -> Option<u64> {
|
||||
keys.iter()
|
||||
.find_map(|key| json_import_expiry_value(object.get(*key)))
|
||||
}
|
||||
|
||||
fn grok_cookie_value(raw: &str, name: &str) -> Option<String> {
|
||||
raw.trim()
|
||||
.strip_prefix("Cookie:")
|
||||
@@ -204,6 +224,8 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
object
|
||||
.get("sso_token")
|
||||
.or_else(|| object.get("ssoToken"))
|
||||
.or_else(|| object.get("session_token"))
|
||||
.or_else(|| object.get("sessionToken"))
|
||||
.or(grok_token_alias),
|
||||
)
|
||||
.or_else(|| {
|
||||
@@ -254,7 +276,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
refresh_token
|
||||
};
|
||||
let expires_at =
|
||||
json_u64_value(object.get("expires_at").or_else(|| object.get("expiresAt")));
|
||||
json_import_expiry_from_keys(object, &["expires_at", "expiresAt", "expired"]);
|
||||
let account_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_id")
|
||||
@@ -660,6 +682,21 @@ mod tests {
|
||||
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_common_chatgpt_web_json_aliases() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"chatgpt_web",
|
||||
r#"[{"session_token":"session-1","expired":"2030-01-01T00:00:00Z","chatgpt_account_id":"acc-1","chatgpt_plan_type":"plus"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("session-1"));
|
||||
assert_eq!(entries[0].expires_at, Some(1_893_456_000));
|
||||
assert_eq!(entries[0].account_id.as_deref(), Some("acc-1"));
|
||||
assert_eq!(entries[0].plan_type.as_deref(), Some("plus"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_plain_jwt_line_as_access_token() {
|
||||
let token = unsigned_jwt(json!({
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::provider_oauth_token_payload_expires_at_unix_secs;
|
||||
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::super::super::runtime::resolve_provider_oauth_runtime_endpoints;
|
||||
use super::super::super::runtime::{
|
||||
resolve_provider_oauth_runtime_endpoints,
|
||||
spawn_provider_oauth_account_state_refresh_after_update,
|
||||
};
|
||||
use super::super::super::state::{
|
||||
admin_provider_oauth_template, enrich_admin_provider_oauth_auth_config,
|
||||
is_fixed_provider_type_for_provider_oauth, json_non_empty_string,
|
||||
@@ -219,50 +221,20 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
));
|
||||
}
|
||||
|
||||
let mut account_state_recheck_attempted = false;
|
||||
let mut account_state_recheck_error = None::<String>;
|
||||
if provider_type == "codex" {
|
||||
if let Some(endpoint) = runtime_endpoint {
|
||||
let refreshed_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
.unwrap_or_else(|| key.clone());
|
||||
if let Some(result) = refresh_codex_provider_quota_locally(
|
||||
state,
|
||||
&provider,
|
||||
&endpoint,
|
||||
vec![refreshed_key],
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
account_state_recheck_attempted = true;
|
||||
let success = result
|
||||
.get("success")
|
||||
.and_then(serde_json::Value::as_u64)
|
||||
.unwrap_or(0);
|
||||
if success == 0 {
|
||||
account_state_recheck_error = result
|
||||
.get("results")
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.and_then(|results| results.first())
|
||||
.and_then(|value| value.get("message"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
provider.clone(),
|
||||
key_id.clone(),
|
||||
request_proxy.clone(),
|
||||
);
|
||||
|
||||
Ok(Json(json!({
|
||||
"provider_type": provider_type,
|
||||
"expires_at": expires_at,
|
||||
"has_refresh_token": refresh_token.is_some(),
|
||||
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
|
||||
"account_state_recheck_attempted": account_state_recheck_attempted,
|
||||
"account_state_recheck_error": account_state_recheck_error,
|
||||
"account_state_recheck_attempted": false,
|
||||
"account_state_recheck_error": serde_json::Value::Null,
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
@@ -828,13 +828,7 @@ async fn handle_admin_provider_oauth_windsurf_browser_device_poll(
|
||||
};
|
||||
callback_token.to_string()
|
||||
} else {
|
||||
let token = token.unwrap_or_default();
|
||||
if windsurf_raw_api_key(token).is_none() {
|
||||
return Ok(windsurf_browser_poll_error_response(
|
||||
"浏览器授权请提交包含 state 的回调 URL;纯 token 请使用导入授权",
|
||||
));
|
||||
}
|
||||
token.to_string()
|
||||
token.unwrap_or_default().to_string()
|
||||
};
|
||||
|
||||
let mut raw_credentials = serde_json::Map::new();
|
||||
|
||||
@@ -106,12 +106,21 @@ fn import_payload_string_any(
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn import_payload_u64(
|
||||
fn import_payload_u64_any(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
snake_case: &str,
|
||||
camel_case: &str,
|
||||
keys: &[&str],
|
||||
) -> Option<u64> {
|
||||
json_u64_value(payload.get(snake_case).or_else(|| payload.get(camel_case)))
|
||||
keys.iter().find_map(|key| {
|
||||
let value = payload.get(*key)?;
|
||||
json_u64_value(Some(value)).or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok())
|
||||
.and_then(|value| u64::try_from(value.timestamp()).ok())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn apply_single_import_hints(
|
||||
@@ -409,9 +418,17 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
||||
let access_token_input = import_payload_string_any(
|
||||
&raw_payload,
|
||||
&["access_token", "accessToken", "sso_token", "ssoToken"],
|
||||
&[
|
||||
"access_token",
|
||||
"accessToken",
|
||||
"sso_token",
|
||||
"ssoToken",
|
||||
"session_token",
|
||||
"sessionToken",
|
||||
],
|
||||
);
|
||||
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
|
||||
let imported_expires_at =
|
||||
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
|
||||
let name = raw_payload
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
@@ -622,8 +639,41 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::sanitize_windsurf_import_error;
|
||||
use super::{
|
||||
import_payload_string_any, import_payload_u64_any, sanitize_windsurf_import_error,
|
||||
};
|
||||
use aether_oauth::core::OAuthError;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn single_import_accepts_session_token_alias() {
|
||||
let payload = json!({
|
||||
"session_token": "session-1",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(
|
||||
import_payload_string_any(&payload, &["access_token", "session_token"]).as_deref(),
|
||||
Some("session-1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_import_accepts_iso_expired_alias() {
|
||||
let payload = json!({
|
||||
"expired": "2030-01-01T00:00:00Z",
|
||||
})
|
||||
.as_object()
|
||||
.cloned()
|
||||
.expect("payload should be an object");
|
||||
|
||||
assert_eq!(
|
||||
import_payload_u64_any(&payload, &["expires_at", "expiresAt", "expired"]),
|
||||
Some(1_893_456_000)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_import_error_redacts_http_body() {
|
||||
|
||||
+50
-40
@@ -81,48 +81,42 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
.await?;
|
||||
if provider_auto_remove_banned_keys(provider.config.as_ref()) {
|
||||
let now_unix_secs = helpers::unix_now_secs();
|
||||
let latest_key = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next();
|
||||
if latest_key.as_ref().is_some_and(|latest_key| {
|
||||
should_auto_remove_oauth_invalid_key(
|
||||
latest_key,
|
||||
None,
|
||||
false,
|
||||
now_unix_secs,
|
||||
)
|
||||
}) {
|
||||
state
|
||||
.clear_admin_provider_pool_cooldown(&provider.id, &key_id)
|
||||
.await;
|
||||
state
|
||||
.reset_admin_provider_pool_cost(&provider.id, &key_id)
|
||||
.await;
|
||||
if state.delete_provider_catalog_key(&key_id).await? {
|
||||
let deleted_key_ids = [key_id.clone()];
|
||||
state
|
||||
.cleanup_deleted_provider_catalog_refs(
|
||||
&provider.id,
|
||||
&[],
|
||||
&deleted_key_ids,
|
||||
let auto_removed = state
|
||||
.cleanup_provider_catalog_key_if_current(
|
||||
&provider,
|
||||
&key_id,
|
||||
|latest_key| {
|
||||
should_auto_remove_oauth_invalid_key(
|
||||
latest_key,
|
||||
Some(&failure_reason),
|
||||
false,
|
||||
now_unix_secs,
|
||||
)
|
||||
.await?;
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if auto_removed {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "auto_removed_oauth_refresh_failed",
|
||||
"gateway manual provider oauth refresh auto-removed unusable key"
|
||||
);
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_auto_removed_response(&error_reason),
|
||||
));
|
||||
}
|
||||
}
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_failed_retained",
|
||||
"gateway manual provider oauth refresh failure retained key"
|
||||
);
|
||||
}
|
||||
}
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
@@ -171,9 +165,25 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
};
|
||||
|
||||
if !helpers::key_is_account_blocked(&key, OAUTH_ACCOUNT_BLOCK_PREFIX) {
|
||||
let _ = state
|
||||
let previous_oauth_refresh_issue =
|
||||
key.oauth_invalid_reason.as_deref().is_some_and(|reason| {
|
||||
reason.lines().map(str::trim).any(|line| {
|
||||
line.starts_with("[OAUTH_EXPIRED]") || line.starts_with("[REFRESH_FAILED]")
|
||||
})
|
||||
});
|
||||
let cleared = state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
|
||||
.await?;
|
||||
if cleared && previous_oauth_refresh_issue {
|
||||
tracing::info!(
|
||||
trace_id = %trace_id,
|
||||
key_id = %key_id,
|
||||
provider_id = %provider.id,
|
||||
provider_type = %provider_type,
|
||||
event_name = "refresh_fixed",
|
||||
"gateway manual provider oauth refresh cleared oauth invalid marker"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let refreshed_key = state
|
||||
|
||||
@@ -201,10 +201,10 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let mut updated = existing_key.clone();
|
||||
updated.is_active = true;
|
||||
updated.encrypted_api_key = Some(encrypted_api_key);
|
||||
updated.encrypted_auth_config = Some(encrypted_auth_config);
|
||||
updated.api_formats = provider_oauth_catalog_key_api_formats(provider_type, api_formats);
|
||||
updated.is_active = true;
|
||||
updated.expires_at_unix_secs = expires_at_unix_secs;
|
||||
updated.oauth_invalid_at_unix_secs = None;
|
||||
updated.oauth_invalid_reason = None;
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload, coerce_json_f64,
|
||||
coerce_json_string, default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
extract_execution_error_message, oauth_refresh_auto_removed_result,
|
||||
persist_provider_quota_refresh_state, quota_key_auto_removed,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
coerce_json_string, execute_provider_quota_plan, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -33,11 +33,10 @@ async fn execute_antigravity_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec = build_antigravity_pool_quota_request(
|
||||
&transport.key.id,
|
||||
&transport.endpoint.base_url,
|
||||
|
||||
@@ -1,16 +1,19 @@
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, extract_execution_error_message,
|
||||
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
|
||||
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
quota_key_auto_removed, quota_refresh_success_invalid_state,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::quota::parse_chatgpt_web_conversation_init_response;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_contracts::{
|
||||
ExecutionResult, ProxySnapshot, ResolvedTransportProfile, TRANSPORT_BACKEND_BROWSER_WREQ,
|
||||
TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
@@ -18,10 +21,12 @@ use aether_provider_pool::{
|
||||
build_chatgpt_web_pool_quota_request, enrich_chatgpt_web_quota_metadata,
|
||||
normalize_chatgpt_web_image_quota_limit,
|
||||
};
|
||||
use base64::Engine as _;
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
|
||||
const CHATGPT_WEB_BROWSER_PROFILE: &str = "chrome143";
|
||||
|
||||
fn chatgpt_web_auth_config(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
@@ -67,26 +72,105 @@ async fn execute_chatgpt_web_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec =
|
||||
build_chatgpt_web_pool_quota_request(&transport.key.id, &endpoint.base_url, authorization);
|
||||
let resolved_transport_profile = state.resolve_transport_profile(transport);
|
||||
let plan = super::shared::build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
proxy,
|
||||
state.resolve_transport_profile(transport),
|
||||
chatgpt_web_quota_transport_profile(resolved_transport_profile.as_ref()),
|
||||
timeouts,
|
||||
);
|
||||
|
||||
execute_provider_quota_plan(state, transport, plan, "chatgpt_web").await
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_transport_profile(
|
||||
transport_profile: Option<&ResolvedTransportProfile>,
|
||||
) -> Option<ResolvedTransportProfile> {
|
||||
match transport_profile {
|
||||
Some(profile)
|
||||
if profile
|
||||
.backend
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ) =>
|
||||
{
|
||||
Some(profile.clone())
|
||||
}
|
||||
_ => Some(default_chatgpt_web_quota_transport_profile()),
|
||||
}
|
||||
}
|
||||
|
||||
fn default_chatgpt_web_quota_transport_profile() -> ResolvedTransportProfile {
|
||||
ResolvedTransportProfile {
|
||||
profile_id: CHATGPT_WEB_BROWSER_PROFILE.to_string(),
|
||||
backend: TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: Some(json!({
|
||||
"browser_profile": CHATGPT_WEB_BROWSER_PROFILE,
|
||||
"source": "chatgpt_web_quota_default",
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_error_detail(result: &ExecutionResult) -> Option<String> {
|
||||
extract_execution_error_message(result).or_else(|| {
|
||||
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
|
||||
let decoded = base64::engine::general_purpose::STANDARD
|
||||
.decode(body)
|
||||
.ok()?;
|
||||
let text = String::from_utf8_lossy(&decoded).trim().to_string();
|
||||
(!text.is_empty()).then_some(text)
|
||||
})
|
||||
}
|
||||
|
||||
fn chatgpt_web_is_structured_account_block(message: &str) -> bool {
|
||||
let lowered = message.to_ascii_lowercase();
|
||||
[
|
||||
"account has been disabled",
|
||||
"account disabled",
|
||||
"account has been deactivated",
|
||||
"account_deactivated",
|
||||
"account deactivated",
|
||||
"organization has been disabled",
|
||||
"organization_disabled",
|
||||
"deactivated_workspace",
|
||||
"account suspended",
|
||||
"account banned",
|
||||
"account_block",
|
||||
"account blocked",
|
||||
"访问被禁止",
|
||||
"账户访问被禁止",
|
||||
"账户已封禁",
|
||||
"封禁",
|
||||
"封号",
|
||||
"被封",
|
||||
]
|
||||
.iter()
|
||||
.any(|keyword| lowered.contains(keyword))
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_403_refresh_failed_reason(message: Option<&str>) -> String {
|
||||
let detail = message
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.filter(|value| !value.contains('<'))
|
||||
.unwrap_or("ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制");
|
||||
format!("{OAUTH_REFRESH_FAILED_PREFIX}{detail}")
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<&str>) -> String {
|
||||
let message = upstream_message.unwrap_or_default().trim();
|
||||
if status_code == 403 && !chatgpt_web_is_structured_account_block(message) {
|
||||
return chatgpt_web_quota_403_refresh_failed_reason(upstream_message);
|
||||
}
|
||||
let detail = if message.is_empty() {
|
||||
match status_code {
|
||||
401 => "ChatGPT Web Token 无效或已过期",
|
||||
@@ -103,6 +187,19 @@ fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<&
|
||||
}
|
||||
}
|
||||
|
||||
fn chatgpt_web_quota_result_message(reason: &str) -> String {
|
||||
for prefix in [
|
||||
OAUTH_REFRESH_FAILED_PREFIX,
|
||||
OAUTH_EXPIRED_PREFIX,
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX,
|
||||
] {
|
||||
if let Some(message) = reason.strip_prefix(prefix) {
|
||||
return message.trim().to_string();
|
||||
}
|
||||
}
|
||||
reason.trim().to_string()
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -216,8 +313,20 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
message = Some("响应中未包含 ChatGPT Web 生图限额信息".to_string());
|
||||
}
|
||||
} else {
|
||||
let err_msg = extract_execution_error_message(&result);
|
||||
message = Some(match err_msg.as_deref() {
|
||||
let err_msg = chatgpt_web_quota_error_detail(&result);
|
||||
let invalid_reason = if matches!(result.status_code, 401 | 403) {
|
||||
Some(chatgpt_web_quota_invalid_reason(
|
||||
result.status_code,
|
||||
err_msg.as_deref(),
|
||||
))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let display_detail = invalid_reason
|
||||
.as_deref()
|
||||
.map(chatgpt_web_quota_result_message)
|
||||
.or_else(|| err_msg.clone());
|
||||
message = Some(match display_detail.as_deref() {
|
||||
Some(detail) if !detail.is_empty() => {
|
||||
format!(
|
||||
"conversation/init 返回状态码 {}: {}",
|
||||
@@ -229,12 +338,14 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
|
||||
if matches!(result.status_code, 401 | 403) {
|
||||
oauth_invalid_at_unix_secs = Some(now_unix_secs);
|
||||
oauth_invalid_reason = Some(chatgpt_web_quota_invalid_reason(
|
||||
result.status_code,
|
||||
err_msg.as_deref(),
|
||||
));
|
||||
oauth_invalid_reason = invalid_reason;
|
||||
status = if result.status_code == 401 {
|
||||
"auth_invalid".to_string()
|
||||
} else if oauth_invalid_reason
|
||||
.as_deref()
|
||||
.is_some_and(|reason| reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX))
|
||||
{
|
||||
"refresh_failed".to_string()
|
||||
} else {
|
||||
"forbidden".to_string()
|
||||
};
|
||||
@@ -303,3 +414,81 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
|
||||
"auto_removed": auto_removed_count,
|
||||
})))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aether_contracts::{ResponseBody, TRANSPORT_BACKEND_REQWEST_RUSTLS};
|
||||
use base64::Engine as _;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn quota_refresh_defaults_to_browser_wreq_transport() {
|
||||
let profile = chatgpt_web_quota_transport_profile(None).expect("transport profile");
|
||||
|
||||
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
|
||||
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
|
||||
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_AUTO);
|
||||
assert_eq!(profile.pool_scope, TRANSPORT_POOL_SCOPE_KEY);
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("browser_profile"))
|
||||
.and_then(serde_json::Value::as_str),
|
||||
Some(CHATGPT_WEB_BROWSER_PROFILE)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_refresh_overrides_non_browser_transport() {
|
||||
let reqwest_profile = ResolvedTransportProfile {
|
||||
profile_id: "chrome_136".to_string(),
|
||||
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: None,
|
||||
};
|
||||
|
||||
let profile =
|
||||
chatgpt_web_quota_transport_profile(Some(&reqwest_profile)).expect("transport profile");
|
||||
|
||||
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
|
||||
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn browser_challenge_403_is_not_account_block() {
|
||||
let body = "<!DOCTYPE html><html><head><title>Just a moment...</title></head><body>Cloudflare</body></html>";
|
||||
let result = ExecutionResult {
|
||||
request_id: "chatgpt-web-quota:test".to_string(),
|
||||
candidate_id: None,
|
||||
status_code: 403,
|
||||
headers: BTreeMap::new(),
|
||||
body: Some(ResponseBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
};
|
||||
|
||||
let detail = chatgpt_web_quota_error_detail(&result).expect("html body should decode");
|
||||
let reason = chatgpt_web_quota_invalid_reason(result.status_code, Some(&detail));
|
||||
|
||||
assert!(reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX));
|
||||
assert!(!reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX));
|
||||
assert_eq!(
|
||||
chatgpt_web_quota_result_message(&reason),
|
||||
"ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_account_block_403_remains_account_block() {
|
||||
let reason = chatgpt_web_quota_invalid_reason(403, Some("account has been deactivated"));
|
||||
|
||||
assert!(reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -44,6 +44,15 @@ fn merge_codex_quota_metadata(
|
||||
serde_json::Value::Object(merged)
|
||||
}
|
||||
|
||||
fn codex_oauth_refresh_issue_reason(reason: Option<&str>) -> bool {
|
||||
reason.is_some_and(|reason| {
|
||||
reason
|
||||
.lines()
|
||||
.map(str::trim)
|
||||
.any(|line| line.starts_with("[OAUTH_EXPIRED]") || line.starts_with("[REFRESH_FAILED]"))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -56,8 +65,13 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
let mut success_count = 0usize;
|
||||
let mut failed_count = 0usize;
|
||||
let mut auto_removed_count = 0usize;
|
||||
let mut refresh_fixed_count = 0usize;
|
||||
let mut refresh_failed_retained_count = 0usize;
|
||||
let mut auto_removed_hard_banned_count = 0usize;
|
||||
|
||||
for key in keys {
|
||||
let had_oauth_refresh_issue =
|
||||
codex_oauth_refresh_issue_reason(key.oauth_invalid_reason.as_deref());
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
@@ -276,13 +290,9 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
}
|
||||
|
||||
let auto_removed = auto_remove_abnormal_keys
|
||||
let auto_remove_candidate = auto_remove_abnormal_keys
|
||||
&& should_auto_remove_structured_reason(oauth_invalid_reason.as_deref());
|
||||
if auto_removed {
|
||||
if state.delete_provider_catalog_key(&key.id).await? {
|
||||
auto_removed_count += 1;
|
||||
}
|
||||
} else if !persist_provider_quota_refresh_state(
|
||||
let persisted = persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
@@ -290,8 +300,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
oauth_invalid_reason.clone(),
|
||||
None,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
.await?;
|
||||
if !persisted {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
@@ -301,6 +311,29 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
let auto_removed = if auto_remove_candidate {
|
||||
state
|
||||
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
|
||||
should_auto_remove_structured_reason(latest_key.oauth_invalid_reason.as_deref())
|
||||
})
|
||||
.await?
|
||||
} else {
|
||||
false
|
||||
};
|
||||
if auto_removed {
|
||||
auto_removed_count += 1;
|
||||
auto_removed_hard_banned_count += 1;
|
||||
}
|
||||
let refresh_fixed =
|
||||
status == "success" && had_oauth_refresh_issue && oauth_invalid_reason.is_none();
|
||||
if refresh_fixed {
|
||||
refresh_fixed_count += 1;
|
||||
}
|
||||
let refresh_failed_retained =
|
||||
status != "success" && oauth_invalid_reason.is_some() && !auto_removed;
|
||||
if refresh_failed_retained {
|
||||
refresh_failed_retained_count += 1;
|
||||
}
|
||||
|
||||
if status == "success" {
|
||||
success_count += 1;
|
||||
@@ -336,6 +369,13 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
}
|
||||
if auto_removed {
|
||||
payload.insert("auto_removed".to_string(), json!(true));
|
||||
payload.insert("auto_removed_hard_banned".to_string(), json!(true));
|
||||
}
|
||||
if refresh_fixed {
|
||||
payload.insert("refresh_fixed".to_string(), json!(true));
|
||||
}
|
||||
if refresh_failed_retained {
|
||||
payload.insert("refresh_failed_retained".to_string(), json!(true));
|
||||
}
|
||||
results.push(serde_json::Value::Object(payload));
|
||||
}
|
||||
@@ -346,5 +386,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
"total": results.len(),
|
||||
"results": results,
|
||||
"auto_removed": auto_removed_count,
|
||||
"refresh_fixed": refresh_fixed_count,
|
||||
"refresh_failed_retained": refresh_failed_retained_count,
|
||||
"auto_removed_hard_banned": auto_removed_hard_banned_count,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::super::shared::{
|
||||
build_provider_quota_execution_plan, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, ProviderQuotaExecutionOutcome,
|
||||
build_provider_quota_execution_plan, execute_provider_quota_plan,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -38,11 +38,10 @@ pub(super) async fn execute_codex_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
use super::shared::{
|
||||
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, extract_execution_error_message,
|
||||
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
|
||||
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
@@ -245,11 +244,10 @@ async fn execute_grok_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let transport_profile = state.resolve_transport_profile(transport);
|
||||
let base_url = grok_base_url(endpoint);
|
||||
let headers = build_grok_quota_headers(
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use super::super::shared::{
|
||||
build_provider_quota_execution_plan, default_provider_quota_execution_timeouts,
|
||||
execute_provider_quota_plan, ProviderQuotaExecutionOutcome,
|
||||
build_provider_quota_execution_plan, execute_provider_quota_plan,
|
||||
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{
|
||||
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminKiroRequestAuth,
|
||||
@@ -23,11 +23,10 @@ pub(super) async fn execute_kiro_quota_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let spec = build_kiro_pool_quota_request(
|
||||
&transport.key.id,
|
||||
&KiroPoolQuotaAuthInput {
|
||||
|
||||
@@ -46,6 +46,23 @@ pub(super) fn default_provider_quota_execution_timeouts(
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider_quota_execution_timeouts(
|
||||
configured: Option<ExecutionTimeouts>,
|
||||
proxy: Option<&ProxySnapshot>,
|
||||
) -> ExecutionTimeouts {
|
||||
let defaults = default_provider_quota_execution_timeouts(proxy);
|
||||
let Some(mut timeouts) = configured else {
|
||||
return defaults;
|
||||
};
|
||||
timeouts.connect_ms = timeouts.connect_ms.or(defaults.connect_ms);
|
||||
timeouts.read_ms = timeouts.read_ms.or(defaults.read_ms);
|
||||
timeouts.write_ms = timeouts.write_ms.or(defaults.write_ms);
|
||||
timeouts.pool_ms = timeouts.pool_ms.or(defaults.pool_ms);
|
||||
timeouts.total_ms = timeouts.total_ms.or(defaults.total_ms);
|
||||
timeouts.first_byte_ms = timeouts.first_byte_ms.or(defaults.first_byte_ms);
|
||||
timeouts
|
||||
}
|
||||
|
||||
pub(crate) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
|
||||
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
|
||||
}
|
||||
@@ -317,12 +334,7 @@ pub(super) async fn execute_provider_quota_plan(
|
||||
match state.execute_execution_runtime_sync_plan(None, &plan).await {
|
||||
Ok(result) => Ok(ProviderQuotaExecutionOutcome::Response(result)),
|
||||
Err(err) => {
|
||||
let error = match err {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
};
|
||||
let error = err.into_message();
|
||||
let proxy_node_id = plan
|
||||
.proxy
|
||||
.as_ref()
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload,
|
||||
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload, execute_provider_quota_plan,
|
||||
extract_execution_error_message, persist_provider_quota_refresh_state,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
quota_refresh_success_invalid_state, resolve_provider_quota_execution_timeouts,
|
||||
ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
@@ -33,11 +33,10 @@ async fn execute_windsurf_probe_plan(
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let timeouts = Some(resolve_provider_quota_execution_timeouts(
|
||||
state.resolve_transport_execution_timeouts(transport),
|
||||
proxy.as_ref(),
|
||||
));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
|
||||
@@ -301,12 +301,7 @@ fn admin_provider_ops_decode_response_bytes(
|
||||
}
|
||||
|
||||
fn admin_provider_ops_gateway_error_message(error: GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::UpstreamUnavailable { message, .. }
|
||||
| GatewayError::ControlUnavailable { message, .. }
|
||||
| GatewayError::Client { message, .. }
|
||||
| GatewayError::Internal(message) => message,
|
||||
}
|
||||
error.into_message()
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_verify_execution_error_message(error: &str) -> String {
|
||||
|
||||
@@ -16,7 +16,12 @@ use aether_runtime_state::{DataLayerError, RuntimeState};
|
||||
use futures_util::future::join_all;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use tracing::warn;
|
||||
use tracing::{info, warn};
|
||||
|
||||
const DEFAULT_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT: usize = 512;
|
||||
const MAX_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT: usize = 10_000;
|
||||
const POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT_ENV: &str =
|
||||
"AETHER_GATEWAY_ADMIN_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT";
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
SystemTime::now()
|
||||
@@ -29,6 +34,20 @@ fn should_load_active_probe_members(pool_config: &AdminProviderPoolConfig) -> bo
|
||||
pool_config.probing_enabled
|
||||
}
|
||||
|
||||
fn pool_runtime_window_metric_key_limit() -> usize {
|
||||
std::env::var(POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT_ENV)
|
||||
.ok()
|
||||
.and_then(|value| value.trim().parse::<usize>().ok())
|
||||
.filter(|value| *value > 0)
|
||||
.unwrap_or(DEFAULT_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT)
|
||||
.clamp(1, MAX_POOL_RUNTIME_WINDOW_METRIC_KEY_LIMIT)
|
||||
}
|
||||
|
||||
fn bounded_runtime_window_metric_key_ids(key_ids: &[String], limit: usize) -> &[String] {
|
||||
let end = key_ids.len().min(limit.max(1));
|
||||
&key_ids[..end]
|
||||
}
|
||||
|
||||
pub(crate) async fn read_admin_provider_pool_cooldown_counts(
|
||||
runtime: &RuntimeState,
|
||||
provider_ids: &[String],
|
||||
@@ -54,8 +73,21 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
) -> AdminProviderPoolRuntimeState {
|
||||
let mut state = AdminProviderPoolRuntimeState::default();
|
||||
let cooldown_keys = pool_cooldown_keys(provider_id, key_ids);
|
||||
let cost_keys = pool_cost_keys(provider_id, key_ids);
|
||||
let latency_keys = pool_latency_keys(provider_id, key_ids);
|
||||
let metric_key_limit = pool_runtime_window_metric_key_limit();
|
||||
let metric_key_ids = bounded_runtime_window_metric_key_ids(key_ids, metric_key_limit);
|
||||
if metric_key_ids.len() < key_ids.len() {
|
||||
info!(
|
||||
event_name = "admin_pool_runtime_window_metrics_truncated",
|
||||
log_type = "event",
|
||||
provider_id,
|
||||
total_key_count = key_ids.len(),
|
||||
scanned_key_count = metric_key_ids.len(),
|
||||
metric_key_limit,
|
||||
"gateway limited admin pool runtime cost/latency window reads"
|
||||
);
|
||||
}
|
||||
let cost_keys = pool_cost_keys(provider_id, metric_key_ids);
|
||||
let latency_keys = pool_latency_keys(provider_id, metric_key_ids);
|
||||
let sticky_sessions_enabled = pool_config.sticky_session_ttl_seconds > 0
|
||||
&& admin_provider_pool_cache_affinity_enabled(pool_config);
|
||||
|
||||
@@ -179,7 +211,7 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
.map(|cost_key| runtime.score_range_by_min(cost_key, cost_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(cost_results) {
|
||||
for (key_id, members) in metric_key_ids.iter().zip(cost_results) {
|
||||
let total = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
@@ -197,7 +229,7 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
|
||||
.map(|latency_key| runtime.score_range_by_min(latency_key, latency_window_start)),
|
||||
)
|
||||
.await;
|
||||
for (key_id, members) in key_ids.iter().zip(latency_results) {
|
||||
for (key_id, members) in metric_key_ids.iter().zip(latency_results) {
|
||||
let samples = members
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
@@ -265,3 +297,30 @@ pub(crate) async fn read_admin_provider_pool_key_cooldown_reason(
|
||||
.kv_get(&pool_cooldown_key(provider_id, key_id))
|
||||
.await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::bounded_runtime_window_metric_key_ids;
|
||||
|
||||
#[test]
|
||||
fn runtime_window_metric_key_ids_are_bounded() {
|
||||
let key_ids = vec![
|
||||
"key-1".to_string(),
|
||||
"key-2".to_string(),
|
||||
"key-3".to_string(),
|
||||
];
|
||||
|
||||
let bounded = bounded_runtime_window_metric_key_ids(&key_ids, 2);
|
||||
|
||||
assert_eq!(bounded, &key_ids[..2]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_window_metric_key_ids_keep_at_least_one_key() {
|
||||
let key_ids = vec!["key-1".to_string(), "key-2".to_string()];
|
||||
|
||||
let bounded = bounded_runtime_window_metric_key_ids(&key_ids, 0);
|
||||
|
||||
assert_eq!(bounded, &key_ids[..1]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_data_contracts::repository::usage::StoredProviderApiKeyWindowUsageSummary;
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
@@ -933,19 +934,14 @@ fn admin_pool_health_score(key: &StoredProviderCatalogKey) -> f64 {
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_pool_circuit_breaker_open(key: &StoredProviderCatalogKey) -> bool {
|
||||
fn admin_pool_circuit_breaker_open(key: &StoredProviderCatalogKey, now_unix_secs: u64) -> bool {
|
||||
key.circuit_breaker_by_format
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.map(|formats| {
|
||||
formats
|
||||
.values()
|
||||
.filter_map(serde_json::Value::as_object)
|
||||
.any(|item| {
|
||||
item.get("open")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.any(|item| provider_key_circuit_payload_is_active_open_at(item, now_unix_secs))
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
@@ -1044,7 +1040,7 @@ pub(super) fn build_admin_pool_key_payload(
|
||||
.as_ref()
|
||||
.and_then(|_| runtime.cooldown_ttl_by_key.get(&key.id).copied());
|
||||
let health_score = admin_pool_health_score(key);
|
||||
let circuit_breaker_open = admin_pool_circuit_breaker_open(key);
|
||||
let circuit_breaker_open = admin_pool_circuit_breaker_open(key, now_unix_secs);
|
||||
let auth_semantics = provider_key_auth_semantics(key, provider_type);
|
||||
let account_quota_exhausted = pool_config
|
||||
.as_ref()
|
||||
|
||||
@@ -99,6 +99,7 @@ static PROVIDER_QUERY_POOL_LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
struct ProviderQueryKeyFetchResult {
|
||||
models: Vec<Value>,
|
||||
error: Option<String>,
|
||||
warning: Option<String>,
|
||||
from_cache: bool,
|
||||
has_success: bool,
|
||||
}
|
||||
@@ -288,6 +289,7 @@ fn provider_query_codex_preset_fallback(
|
||||
Some(ProviderQueryKeyFetchResult {
|
||||
models: aggregate_models_for_cache(&models),
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
})
|
||||
@@ -427,6 +429,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models,
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: true,
|
||||
has_success: true,
|
||||
});
|
||||
@@ -444,6 +447,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models,
|
||||
error: None,
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: true,
|
||||
});
|
||||
@@ -451,6 +455,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(ADMIN_PROVIDER_QUERY_NO_ACTIVE_ENDPOINT_DETAIL.to_string()),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -477,6 +482,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -492,6 +498,7 @@ async fn provider_query_fetch_models_for_key(
|
||||
return Ok(ProviderQueryKeyFetchResult {
|
||||
models: Vec::new(),
|
||||
error: Some(all_errors.join("; ")),
|
||||
warning: None,
|
||||
from_cache: false,
|
||||
has_success: false,
|
||||
});
|
||||
@@ -528,18 +535,25 @@ async fn provider_query_fetch_models_for_key(
|
||||
}
|
||||
}
|
||||
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
} else {
|
||||
let has_models = !unique_models.is_empty();
|
||||
let mut error = if !has_models && !all_errors.is_empty() {
|
||||
Some(all_errors.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if unique_models.is_empty() && error.is_none() {
|
||||
let warning = if has_models && !all_errors.is_empty() {
|
||||
Some(all_errors.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if !has_models && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_ENDPOINT_DETAIL.to_string());
|
||||
}
|
||||
|
||||
Ok(ProviderQueryKeyFetchResult {
|
||||
models: provider_query_filter_models_for_key(provider, key, unique_models),
|
||||
error,
|
||||
warning,
|
||||
from_cache: false,
|
||||
has_success: outcome.has_success,
|
||||
})
|
||||
@@ -600,6 +614,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": result.error,
|
||||
"warning": result.warning,
|
||||
"from_cache": result.from_cache,
|
||||
},
|
||||
"provider": provider_query_provider_payload(&provider),
|
||||
@@ -632,6 +647,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": serde_json::Value::Null,
|
||||
"warning": serde_json::Value::Null,
|
||||
"from_cache": true,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": active_key_count,
|
||||
@@ -655,6 +671,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
|
||||
let mut all_models = Vec::new();
|
||||
let mut all_errors = Vec::new();
|
||||
let mut all_warnings = Vec::new();
|
||||
let mut cache_hit_count = 0usize;
|
||||
let mut fetch_count = 0usize;
|
||||
for key in &ordered_keys {
|
||||
@@ -669,6 +686,13 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
error
|
||||
));
|
||||
}
|
||||
if let Some(warning) = result.warning {
|
||||
all_warnings.push(format!(
|
||||
"Key {}: {}",
|
||||
provider_query_key_display_name(key),
|
||||
warning
|
||||
));
|
||||
}
|
||||
if result.from_cache {
|
||||
cache_hit_count += 1;
|
||||
} else {
|
||||
@@ -694,10 +718,17 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
provider_query_write_provider_cached_models(state, &provider.id, &models).await;
|
||||
}
|
||||
let success = !models.is_empty();
|
||||
let mut error = if all_errors.is_empty() {
|
||||
None
|
||||
let mut all_issues = all_errors;
|
||||
all_issues.extend(all_warnings);
|
||||
let mut error = if !success && !all_issues.is_empty() {
|
||||
Some(all_issues.join("; "))
|
||||
} else {
|
||||
Some(all_errors.join("; "))
|
||||
None
|
||||
};
|
||||
let warning = if success && !all_issues.is_empty() {
|
||||
Some(all_issues.join("; "))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if !success && error.is_none() {
|
||||
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
|
||||
@@ -709,6 +740,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"data": {
|
||||
"models": models,
|
||||
"error": error,
|
||||
"warning": warning,
|
||||
"from_cache": fetch_count == 0 && cache_hit_count > 0,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": cache_hit_count,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use super::super::payload::{
|
||||
provider_query_extract_api_key_id, provider_query_extract_force_refresh,
|
||||
provider_query_extract_api_key_ids, provider_query_extract_force_refresh,
|
||||
provider_query_extract_model, provider_query_extract_provider_id,
|
||||
provider_query_extract_request_id,
|
||||
};
|
||||
@@ -68,6 +68,7 @@ use aether_model_fetch::{
|
||||
aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
|
||||
preset_models_for_provider, selected_models_fetch_endpoints,
|
||||
};
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
http::{self, HeaderMap, HeaderName, HeaderValue},
|
||||
@@ -122,6 +123,7 @@ const ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL: &str =
|
||||
"No active endpoint or API key found";
|
||||
const ADMIN_PROVIDER_QUERY_INVALID_MAPPED_MODEL_DETAIL: &str =
|
||||
"mapped_model_name is not valid for the selected model and endpoint";
|
||||
const PROVIDER_QUERY_KEY_MODEL_NOT_ALLOWED_SKIP_REASON: &str = "key_model_not_allowed";
|
||||
const ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX: &str = "upstream_models_provider:";
|
||||
const DEFAULT_PROVIDER_QUERY_TEST_MESSAGE: &str = "Hello! This is a test message.";
|
||||
static PROVIDER_QUERY_POOL_LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
@@ -859,7 +861,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
selected_key_id: Option<&str>,
|
||||
selected_key_ids: Option<&BTreeSet<String>>,
|
||||
) -> Option<StoredProviderCatalogEndpoint> {
|
||||
for priority in 0..=2 {
|
||||
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
|
||||
@@ -872,7 +874,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
}
|
||||
for key in keys {
|
||||
if !key.is_active
|
||||
|| selected_key_id.is_some_and(|value| value != key.id.as_str())
|
||||
|| !provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
|
||||
|| !provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
@@ -904,7 +906,7 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
endpoint.is_active
|
||||
&& keys.iter().any(|key| {
|
||||
key.is_active
|
||||
&& selected_key_id.is_none_or(|value| value == key.id.as_str())
|
||||
&& provider_query_selected_key_ids_allow_key(selected_key_ids, &key.id)
|
||||
&& provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
@@ -916,10 +918,56 @@ async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn provider_query_selected_key_ids_allow_key(
|
||||
selected_key_ids: Option<&BTreeSet<String>>,
|
||||
key_id: &str,
|
||||
) -> bool {
|
||||
selected_key_ids.is_none_or(|ids| ids.contains(key_id))
|
||||
}
|
||||
|
||||
fn provider_query_selected_key_ids_all_exist(
|
||||
selected_key_ids: &BTreeSet<String>,
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
) -> bool {
|
||||
selected_key_ids
|
||||
.iter()
|
||||
.all(|id| keys.iter().any(|key| key.id == *id))
|
||||
}
|
||||
|
||||
fn provider_query_model_name_matches(left: &str, right: &str) -> bool {
|
||||
let left = left.trim();
|
||||
let right = right.trim();
|
||||
!left.is_empty() && !right.is_empty() && left.eq_ignore_ascii_case(right)
|
||||
}
|
||||
|
||||
fn provider_query_key_allows_effective_test_model(
|
||||
key: &StoredProviderCatalogKey,
|
||||
requested_model: &str,
|
||||
effective_model: &str,
|
||||
) -> bool {
|
||||
let allowed_models = json_string_list(key.allowed_models.as_ref());
|
||||
if key.allowed_models.is_none() || allowed_models.is_empty() {
|
||||
return true;
|
||||
}
|
||||
|
||||
let requested_base_model = crate::ai_serving::model_directive_base_model(requested_model);
|
||||
allowed_models
|
||||
.iter()
|
||||
.map(String::as_str)
|
||||
.any(|allowed_model| {
|
||||
provider_query_model_name_matches(allowed_model, requested_model)
|
||||
|| provider_query_model_name_matches(allowed_model, effective_model)
|
||||
|| requested_base_model.as_deref().is_some_and(|base_model| {
|
||||
provider_query_model_name_matches(allowed_model, base_model)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_test_key_sort_key(
|
||||
provider_type: &str,
|
||||
key: &StoredProviderCatalogKey,
|
||||
endpoint_api_format: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> (u8, u8, i32, u64, i32) {
|
||||
let quota_exhausted =
|
||||
admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type);
|
||||
@@ -928,10 +976,7 @@ fn provider_query_test_key_sort_key(
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get(endpoint_api_format))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|value| value.get("open"))
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
.is_some_and(|value| provider_key_circuit_payload_is_active_open_at(value, now_unix_secs));
|
||||
let health_score = key
|
||||
.health_by_format
|
||||
.as_ref()
|
||||
@@ -1263,7 +1308,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
)
|
||||
})?;
|
||||
let selected_key_id = provider_query_extract_api_key_id(payload);
|
||||
let selected_key_ids = provider_query_extract_api_key_ids(payload);
|
||||
let requested_endpoint_id = provider_query_extract_endpoint_id(payload);
|
||||
let requested_api_format = provider_query_extract_api_format(payload);
|
||||
let endpoint = if requested_endpoint_id.is_none()
|
||||
@@ -1275,7 +1320,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider,
|
||||
&endpoints,
|
||||
&all_keys,
|
||||
selected_key_id.as_deref(),
|
||||
selected_key_ids.as_ref(),
|
||||
)
|
||||
.await
|
||||
.ok_or_else(|| {
|
||||
@@ -1308,22 +1353,11 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(api_key_id) = selected_key_id.as_deref() {
|
||||
let Some(key) = all_keys.iter().find(|key| key.id == api_key_id) else {
|
||||
if let Some(selected_key_ids) = selected_key_ids.as_ref() {
|
||||
if !provider_query_selected_key_ids_all_exist(selected_key_ids, &all_keys) {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
|
||||
));
|
||||
};
|
||||
if !key.is_active
|
||||
|| !provider_query_key_supports_endpoint(
|
||||
key,
|
||||
&provider.provider_type,
|
||||
&endpoint.api_format,
|
||||
)
|
||||
{
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_TEST_CANDIDATE_DETAIL,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1378,20 +1412,43 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
.unwrap_or(requested_model.clone())
|
||||
};
|
||||
|
||||
let mut keys = all_keys
|
||||
let now_unix_secs = current_unix_ms() / 1000;
|
||||
let mut keys = Vec::new();
|
||||
let mut model_skipped_candidates = Vec::new();
|
||||
|
||||
for key in all_keys
|
||||
.into_iter()
|
||||
.filter(|key| key.is_active)
|
||||
.filter(|key| {
|
||||
selected_key_id
|
||||
.as_deref()
|
||||
.is_none_or(|value| value == key.id.as_str())
|
||||
})
|
||||
.filter(|key| provider_query_selected_key_ids_allow_key(selected_key_ids.as_ref(), &key.id))
|
||||
.filter(|key| {
|
||||
provider_query_key_supports_endpoint(key, &provider.provider_type, &endpoint.api_format)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
{
|
||||
if provider_query_key_allows_effective_test_model(&key, &requested_model, &effective_model)
|
||||
{
|
||||
keys.push(key);
|
||||
} else {
|
||||
model_skipped_candidates.push(ProviderQueryTestCandidate {
|
||||
endpoint: endpoint.clone(),
|
||||
key,
|
||||
effective_model: effective_model.clone(),
|
||||
scheduler_skip_reason: Some(
|
||||
PROVIDER_QUERY_KEY_MODEL_NOT_ALLOWED_SKIP_REASON.to_string(),
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let candidates = if test_mode.eq_ignore_ascii_case("pool") {
|
||||
model_skipped_candidates.sort_by_key(|candidate| {
|
||||
provider_query_test_key_sort_key(
|
||||
provider.provider_type.as_str(),
|
||||
&candidate.key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
|
||||
let scheduled_candidates = if test_mode.eq_ignore_ascii_case("pool") {
|
||||
if let Some(pool_config) =
|
||||
admin_provider_pool_config_from_config_value(provider.config.as_ref())
|
||||
{
|
||||
@@ -1411,6 +1468,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider.provider_type.as_str(),
|
||||
key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
keys.into_iter()
|
||||
@@ -1428,6 +1486,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
provider.provider_type.as_str(),
|
||||
key,
|
||||
&endpoint.api_format,
|
||||
now_unix_secs,
|
||||
)
|
||||
});
|
||||
keys.into_iter()
|
||||
@@ -1439,6 +1498,8 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
let mut candidates = model_skipped_candidates;
|
||||
candidates.extend(scheduled_candidates);
|
||||
|
||||
if candidates.is_empty() {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
|
||||
@@ -64,6 +64,68 @@ fn sample_openai_image_transport(provider_type: &str) -> AdminGatewayProviderTra
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_catalog_key_with_allowed_models(
|
||||
allowed_models: Option<serde_json::Value>,
|
||||
) -> aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey {
|
||||
let mut key =
|
||||
aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"key".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("sample provider key should build");
|
||||
key.allowed_models = allowed_models;
|
||||
key
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_allows_keys_without_model_restrictions() {
|
||||
let unrestricted = sample_catalog_key_with_allowed_models(None);
|
||||
let empty = sample_catalog_key_with_allowed_models(Some(json!([])));
|
||||
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&unrestricted,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&empty,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_filters_key_disallowed_for_requested_model() {
|
||||
let key = sample_catalog_key_with_allowed_models(Some(json!(["model-a"])));
|
||||
|
||||
assert!(!provider_query_key_allows_effective_test_model(
|
||||
&key,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_allows_key_for_requested_or_mapped_model() {
|
||||
let requested_allowed = sample_catalog_key_with_allowed_models(Some(json!(["model-b"])));
|
||||
let mapped_allowed = sample_catalog_key_with_allowed_models(Some(json!(["MODEL-B-UPSTREAM"])));
|
||||
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&requested_allowed,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
assert!(provider_query_key_allows_effective_test_model(
|
||||
&mapped_allowed,
|
||||
"model-b",
|
||||
"model-b-upstream",
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_test_request_body_preserves_custom_model() {
|
||||
let payload = json!({
|
||||
@@ -232,6 +294,27 @@ fn provider_query_request_body_model_uses_non_empty_string_only() {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_extracts_multiple_selected_key_ids() {
|
||||
let payload = json!({
|
||||
"api_key_ids": [" key-b ", "", "key-a", "key-b"],
|
||||
"api_key_id": "key-c"
|
||||
});
|
||||
|
||||
let ids = provider_query_extract_api_key_ids(&payload)
|
||||
.expect("non-empty key selection should be extracted")
|
||||
.into_iter()
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(ids, vec!["key-a", "key-b", "key-c"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_model_test_empty_selected_key_ids_keep_default_selection() {
|
||||
assert!(provider_query_extract_api_key_ids(&json!({})).is_none());
|
||||
assert!(provider_query_extract_api_key_ids(&json!({ "api_key_ids": [] })).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() {
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use axum::body::Bytes;
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
pub(crate) fn parse_admin_provider_query_body(
|
||||
request_body: Option<&Bytes>,
|
||||
@@ -36,6 +37,47 @@ pub(crate) fn provider_query_extract_api_key_id(payload: &serde_json::Value) ->
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn provider_query_insert_api_key_id(ids: &mut BTreeSet<String>, value: &str) {
|
||||
let value = value.trim();
|
||||
if !value.is_empty() {
|
||||
ids.insert(value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn provider_query_extract_api_key_ids(
|
||||
payload: &serde_json::Value,
|
||||
) -> Option<BTreeSet<String>> {
|
||||
let mut ids = BTreeSet::new();
|
||||
|
||||
if let Some(value) = payload
|
||||
.get("api_key_ids")
|
||||
.or_else(|| payload.get("provider_key_ids"))
|
||||
.or_else(|| payload.get("key_ids"))
|
||||
{
|
||||
match value {
|
||||
serde_json::Value::Array(items) => {
|
||||
for item in items {
|
||||
if let Some(value) = item.as_str() {
|
||||
provider_query_insert_api_key_id(&mut ids, value);
|
||||
}
|
||||
}
|
||||
}
|
||||
serde_json::Value::String(value) => {
|
||||
for item in value.split(',') {
|
||||
provider_query_insert_api_key_id(&mut ids, item);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(api_key_id) = provider_query_extract_api_key_id(payload) {
|
||||
ids.insert(api_key_id);
|
||||
}
|
||||
|
||||
(!ids.is_empty()).then_some(ids)
|
||||
}
|
||||
|
||||
pub(crate) fn provider_query_extract_force_refresh(payload: &serde_json::Value) -> bool {
|
||||
payload
|
||||
.get("force_refresh")
|
||||
|
||||
Reference in New Issue
Block a user