mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
feat(grok): add admin oauth and quota support
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user