refactor: 大规模模块拆分与重组,新增 aether-admin crate

- 新建独立 aether-admin crate 承载 admin 相关共享契约与纯辅助函数
- 拆分 ai_pipeline 下 kiro/private_envelope/conversion/planner 等大文件为子模块目录
- 重组 admin handlers 各业务域(billing/oauth/provider/system/users 等)为目录结构,移除 shared.rs/builders.rs 等反模式
- 移除 ai_pipeline runtime adapters 旧实现(claude/openai/gemini/kiro/vertex/antigravity 等),改由 provider transport 统一承载
- 移除 control_facade/execution_facade/auth_snapshot_facade 等冗余 facade 层
- 拆分 query/billing 与 query/monitoring 模块、state/runtime/payments 与 security 模块
- 扩展架构测试覆盖 admin_billing/admin_model/admin_users 等新模块
- 删除 docs/architecture/refactor-execution-plan.md 已完成的执行计划文档
This commit is contained in:
fawney19
2026-04-09 00:10:38 +08:00
parent 4fb9882b54
commit 4fc95adfb9
663 changed files with 48471 additions and 40232 deletions
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,274 @@
use super::kiro_import::execute_admin_provider_oauth_kiro_batch_import;
use super::parse::{
apply_admin_provider_oauth_batch_import_hints, extract_admin_provider_oauth_batch_error_detail,
parse_admin_provider_oauth_batch_import_entries, AdminProviderOAuthBatchImportEntry,
AdminProviderOAuthBatchImportOutcome,
};
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
use crate::handlers::admin::provider::oauth::provisioning::build_provider_oauth_auth_config_from_token_payload;
use crate::handlers::admin::provider::oauth::provisioning::{
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
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;
use crate::GatewayError;
use aether_admin::provider::oauth::parse_admin_provider_oauth_kiro_batch_import_entries;
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) fn estimate_admin_provider_oauth_batch_import_total(
provider_type: &str,
raw_credentials: &str,
) -> usize {
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()
}
}
pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
state: &AdminAppState<'_>,
provider_id: &str,
provider_type: &str,
raw_credentials: &str,
proxy_node_id: Option<&str>,
) -> Result<AdminProviderOAuthBatchImportOutcome, GatewayError> {
if provider_type.eq_ignore_ascii_case("kiro") {
execute_admin_provider_oauth_kiro_batch_import(
state,
provider_id,
raw_credentials,
proxy_node_id,
)
.await
} else {
let entries = parse_admin_provider_oauth_batch_import_entries(raw_credentials);
execute_admin_provider_oauth_batch_import(
state,
provider_id,
provider_type,
&entries,
proxy_node_id,
)
.await
}
}
pub(super) async fn execute_admin_provider_oauth_batch_import(
state: &AdminAppState<'_>,
provider_id: &str,
provider_type: &str,
entries: &[AdminProviderOAuthBatchImportEntry],
proxy_node_id: Option<&str>,
) -> Result<AdminProviderOAuthBatchImportOutcome, GatewayError> {
let Some(provider) = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok(AdminProviderOAuthBatchImportOutcome {
total: entries.len(),
success: 0,
failed: entries.len(),
results: entries
.iter()
.enumerate()
.map(|(index, _)| {
json!({
"index": index,
"status": "error",
"error": "Provider 不存在",
"replaced": false,
})
})
.collect(),
});
};
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 endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id);
let mut results = Vec::with_capacity(entries.len());
let mut success = 0usize;
let mut failed = 0usize;
for (index, entry) in entries.iter().enumerate() {
let token_payload = match exchange_admin_provider_oauth_refresh_token(
state,
template,
entry.refresh_token.as_str(),
)
.await
{
Ok(payload) => payload,
Err(response) => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": format!(
"Token 验证失败: {}",
extract_admin_provider_oauth_batch_error_detail(response).await
),
"replaced": false,
}));
continue;
}
};
let (mut auth_config, access_token, returned_refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(provider_type, &token_payload);
let Some(access_token) = access_token else {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "Token 刷新返回缺少 access_token",
"replaced": false,
}));
continue;
};
let refresh_token = returned_refresh_token
.or_else(|| Some(entry.refresh_token.clone()))
.filter(|value| !value.trim().is_empty());
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
apply_admin_provider_oauth_batch_import_hints(provider_type, entry, &mut auth_config);
let duplicate =
match find_duplicate_provider_oauth_key(state, provider_id, &auth_config, None).await {
Ok(value) => value,
Err(detail) => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": detail,
"replaced": false,
}));
continue;
}
};
let replaced = duplicate.is_some();
let (persisted_key, key_name) = if let Some(existing_key) = duplicate {
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => (key, existing_key.name.clone()),
None => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "provider oauth write unavailable",
"replaced": true,
}));
continue;
}
}
} 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
)
});
match create_provider_oauth_catalog_key(
state,
provider_id,
key_name.as_str(),
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => (key, key_name),
None => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "provider oauth write unavailable",
"replaced": false,
}));
continue;
}
}
};
let _ =
refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
success += 1;
results.push(json!({
"index": index,
"status": "success",
"key_id": persisted_key.id,
"key_name": key_name,
"error": serde_json::Value::Null,
"replaced": replaced,
}));
}
Ok(AdminProviderOAuthBatchImportOutcome {
total: entries.len(),
success,
failed,
results,
})
}
@@ -0,0 +1,266 @@
use super::parse::{AdminProviderOAuthBatchImportEntry, AdminProviderOAuthBatchImportOutcome};
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::provisioning::{
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
use crate::handlers::admin::provider::oauth::state::decode_jwt_claims;
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
use crate::handlers::admin::request::{
AdminAppState, AdminKiroAuthConfig, AdminKiroOAuthRefreshAdapter,
};
use crate::GatewayError;
use aether_admin::provider::oauth::{
build_kiro_batch_import_key_name, coerce_admin_provider_oauth_import_str,
parse_admin_provider_oauth_kiro_batch_import_entries,
};
use serde_json::{json, Map, Value};
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
fn admin_provider_oauth_kiro_refresh_base_url_override(
state: &AdminAppState<'_>,
override_key: &str,
) -> Option<String> {
let override_url = state.provider_oauth_token_url(override_key, "");
let normalized = override_url.trim();
(!normalized.is_empty()).then(|| normalized.to_string())
}
pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
state: &AdminAppState<'_>,
provider_id: &str,
raw_credentials: &str,
proxy_node_id: Option<&str>,
) -> Result<AdminProviderOAuthBatchImportOutcome, GatewayError> {
let entries = parse_admin_provider_oauth_kiro_batch_import_entries(raw_credentials);
let Some(provider) = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok(AdminProviderOAuthBatchImportOutcome {
total: entries.len(),
success: 0,
failed: entries.len(),
results: entries
.iter()
.enumerate()
.map(|(index, _)| {
json!({
"index": index,
"status": "error",
"error": "Provider 不存在",
"replaced": false,
})
})
.collect(),
});
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
.await?;
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id);
let adapter = AdminKiroOAuthRefreshAdapter::default().with_refresh_base_urls(
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_social_refresh"),
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_idc_refresh"),
);
let mut results = Vec::with_capacity(entries.len());
let mut success = 0usize;
let mut failed = 0usize;
for (index, entry) in entries.iter().enumerate() {
let Some(mut refreshed_auth_config) = AdminKiroAuthConfig::from_json_value(entry) else {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "未找到有效的凭据数据",
"replaced": false,
}));
continue;
};
let has_refresh_token = refreshed_auth_config
.refresh_token
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty());
if !has_refresh_token {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "缺少可用的 Kiro refresh 凭据",
"replaced": false,
}));
continue;
}
refreshed_auth_config = match adapter
.refresh_auth_config(state.http_client(), &refreshed_auth_config)
.await
{
Ok(config) => config,
Err(err) => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": format!("Token 验证失败: {err:?}"),
"replaced": false,
}));
continue;
}
};
if refreshed_auth_config.auth_method.is_none() {
refreshed_auth_config.auth_method = Some(if refreshed_auth_config.is_idc_auth() {
"idc".to_string()
} else {
"social".to_string()
});
}
let mut auth_config = refreshed_auth_config
.to_json_value()
.as_object()
.cloned()
.unwrap_or_default();
auth_config.insert("provider_type".to_string(), json!("kiro"));
let email = decode_jwt_claims(
refreshed_auth_config
.access_token
.as_deref()
.unwrap_or_default(),
)
.and_then(|claims: Map<String, Value>| claims.get("email").cloned())
.and_then(|value: Value| value.as_str().map(ToOwned::to_owned))
.or_else(|| coerce_admin_provider_oauth_import_str(entry.get("email")));
if let Some(email) = email.as_ref() {
auth_config.insert("email".to_string(), json!(email));
}
let duplicate =
match find_duplicate_provider_oauth_key(state, provider_id, &auth_config, None).await {
Ok(value) => value,
Err(detail) => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": detail,
"replaced": false,
}));
continue;
}
};
let access_token = refreshed_auth_config
.access_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let Some(access_token) = access_token else {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "Token 验证失败: accessToken 为空",
"replaced": false,
}));
continue;
};
let replaced = duplicate.is_some();
let (persisted_key, key_name) = if let Some(existing_key) = duplicate {
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
refreshed_auth_config.expires_at,
)
.await?
{
Some(key) => (key, existing_key.name.clone()),
None => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "provider oauth write unavailable",
"replaced": true,
}));
continue;
}
}
} else {
let key_name = build_kiro_batch_import_key_name(
auth_config.get("email").and_then(serde_json::Value::as_str),
auth_config
.get("auth_method")
.and_then(serde_json::Value::as_str),
auth_config
.get("refresh_token")
.and_then(serde_json::Value::as_str),
);
match create_provider_oauth_catalog_key(
state,
provider_id,
&key_name,
&access_token,
&auth_config,
&provider_oauth_active_api_formats(&endpoints),
key_proxy.clone(),
refreshed_auth_config.expires_at,
)
.await?
{
Some(key) => (key, key_name),
None => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": "provider oauth write unavailable",
"replaced": false,
}));
continue;
}
}
};
let auth_method = auth_config
.get("auth_method")
.cloned()
.unwrap_or(serde_json::Value::Null);
let _ =
refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
success += 1;
results.push(json!({
"index": index,
"status": "success",
"key_id": persisted_key.id,
"key_name": key_name,
"auth_method": auth_method,
"error": serde_json::Value::Null,
"replaced": replaced,
}));
}
Ok(AdminProviderOAuthBatchImportOutcome {
total: entries.len(),
success,
failed,
results,
})
}
@@ -0,0 +1,8 @@
mod execution;
mod kiro_import;
mod orchestration;
mod parse;
mod task;
pub(super) use orchestration::handle_admin_provider_oauth_batch_import;
pub(super) use task::handle_admin_provider_oauth_start_batch_import_task;
@@ -0,0 +1,87 @@
use super::execution::{
estimate_admin_provider_oauth_batch_import_total,
execute_admin_provider_oauth_batch_import_for_provider_type,
};
use super::parse::{
build_admin_provider_oauth_batch_import_response,
parse_admin_provider_oauth_batch_import_request, AdminProviderOAuthBatchImportRequest,
};
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,
is_fixed_provider_type_for_provider_oauth,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::Bytes,
http,
response::{IntoResponse, Response},
Json,
};
pub(in super::super) async fn handle_admin_provider_oauth_batch_import(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response, GatewayError> {
let raw_state = state.cloned_app();
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) = admin_provider_oauth_batch_import_provider_id(request_context.path())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let payload = match parse_admin_provider_oauth_batch_import_request(request_body) {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 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(),
);
if total == 0 {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"未找到有效的 Token 数据",
));
}
let outcome = execute_admin_provider_oauth_batch_import_for_provider_type(
&AdminAppState::new(&raw_state),
&provider_id,
&provider_type,
payload.credentials.as_str(),
payload.proxy_node_id.as_deref(),
)
.await?;
Ok(build_admin_provider_oauth_batch_import_response(&outcome).into_response())
}
@@ -0,0 +1,289 @@
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::state::current_unix_secs;
use axum::{
body::{to_bytes, Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde::Deserialize;
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
#[derive(Debug, Clone, Deserialize)]
pub(super) struct AdminProviderOAuthBatchImportRequest {
pub credentials: String,
pub proxy_node_id: Option<String>,
}
#[derive(Debug, Clone)]
pub(super) struct AdminProviderOAuthBatchImportEntry {
pub refresh_token: String,
pub account_id: Option<String>,
pub account_user_id: Option<String>,
pub plan_type: Option<String>,
pub user_id: Option<String>,
pub email: Option<String>,
}
#[derive(Debug, Clone)]
pub(super) struct AdminProviderOAuthBatchImportOutcome {
pub total: usize,
pub success: usize,
pub failed: usize,
pub results: Vec<serde_json::Value>,
}
pub(super) fn parse_admin_provider_oauth_batch_import_request(
request_body: Option<&Bytes>,
) -> Result<AdminProviderOAuthBatchImportRequest, Response<Body>> {
let Some(request_body) = request_body else {
return Err(
crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
),
);
};
match serde_json::from_slice::<AdminProviderOAuthBatchImportRequest>(request_body) {
Ok(payload) if !payload.credentials.trim().is_empty() => Ok(payload),
_ => Err(
crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
),
),
}
}
fn coerce_admin_provider_oauth_import_str(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn extract_admin_provider_oauth_batch_import_entry(
item: &serde_json::Value,
) -> Option<AdminProviderOAuthBatchImportEntry> {
match item {
serde_json::Value::String(value) => {
let refresh_token = value.trim();
if refresh_token.is_empty() {
None
} else {
Some(AdminProviderOAuthBatchImportEntry {
refresh_token: refresh_token.to_string(),
account_id: None,
account_user_id: None,
plan_type: None,
user_id: None,
email: None,
})
}
}
serde_json::Value::Object(object) => {
let refresh_token = coerce_admin_provider_oauth_import_str(
object
.get("refresh_token")
.or_else(|| object.get("refreshToken")),
)?;
let account_id = coerce_admin_provider_oauth_import_str(
object
.get("account_id")
.or_else(|| object.get("accountId"))
.or_else(|| object.get("chatgpt_account_id"))
.or_else(|| object.get("chatgptAccountId")),
);
let account_user_id = coerce_admin_provider_oauth_import_str(
object
.get("account_user_id")
.or_else(|| object.get("accountUserId"))
.or_else(|| object.get("chatgpt_account_user_id"))
.or_else(|| object.get("chatgptAccountUserId")),
);
let plan_type = coerce_admin_provider_oauth_import_str(
object
.get("plan_type")
.or_else(|| object.get("planType"))
.or_else(|| object.get("chatgpt_plan_type"))
.or_else(|| object.get("chatgptPlanType")),
)
.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")),
);
let email = coerce_admin_provider_oauth_import_str(object.get("email"));
Some(AdminProviderOAuthBatchImportEntry {
refresh_token,
account_id,
account_user_id,
plan_type,
user_id,
email,
})
}
_ => None,
}
}
pub(super) fn parse_admin_provider_oauth_batch_import_entries(
raw_credentials: &str,
) -> Vec<AdminProviderOAuthBatchImportEntry> {
let raw = raw_credentials.trim();
if raw.is_empty() {
return Vec::new();
}
if raw.starts_with('[') {
if let Ok(serde_json::Value::Array(items)) = serde_json::from_str::<serde_json::Value>(raw)
{
return items
.iter()
.filter_map(extract_admin_provider_oauth_batch_import_entry)
.collect();
}
}
if raw.starts_with('{') {
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)
.into_iter()
.collect();
}
}
raw.lines()
.map(str::trim)
.filter(|line| !line.is_empty() && !line.starts_with('#'))
.map(|refresh_token| AdminProviderOAuthBatchImportEntry {
refresh_token: refresh_token.to_string(),
account_id: None,
account_user_id: None,
plan_type: None,
user_id: None,
email: None,
})
.collect()
}
pub(super) fn apply_admin_provider_oauth_batch_import_hints(
provider_type: &str,
entry: &AdminProviderOAuthBatchImportEntry,
auth_config: &mut serde_json::Map<String, serde_json::Value>,
) {
if !provider_type.eq_ignore_ascii_case("codex") {
return;
}
if let Some(account_id) = entry.account_id.as_ref() {
auth_config
.entry("account_id".to_string())
.or_insert_with(|| json!(account_id));
}
if let Some(account_user_id) = entry.account_user_id.as_ref() {
auth_config
.entry("account_user_id".to_string())
.or_insert_with(|| json!(account_user_id));
}
if let Some(plan_type) = entry.plan_type.as_ref() {
auth_config
.entry("plan_type".to_string())
.or_insert_with(|| json!(plan_type));
}
if let Some(user_id) = entry.user_id.as_ref() {
auth_config
.entry("user_id".to_string())
.or_insert_with(|| json!(user_id));
}
if let Some(email) = entry.email.as_ref() {
auth_config
.entry("email".to_string())
.or_insert_with(|| json!(email));
}
}
pub(super) async fn extract_admin_provider_oauth_batch_error_detail(
response: Response<Body>,
) -> String {
let status = response.status();
let raw_body = to_bytes(response.into_body(), usize::MAX).await.ok();
if let Some(raw_body) = raw_body {
if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&raw_body) {
if let Some(detail) = value.get("detail").and_then(serde_json::Value::as_str) {
let normalized = detail.trim();
if !normalized.is_empty() {
return normalized.to_string();
}
}
}
let normalized = String::from_utf8_lossy(&raw_body).trim().to_string();
if !normalized.is_empty() {
return normalized;
}
}
format!("HTTP {}", status.as_u16())
}
pub(super) fn build_admin_provider_oauth_batch_import_response(
outcome: &AdminProviderOAuthBatchImportOutcome,
) -> Json<serde_json::Value> {
Json(json!({
"total": outcome.total,
"success": outcome.success,
"failed": outcome.failed,
"results": outcome.results,
}))
}
pub(super) fn build_admin_provider_oauth_batch_task_state(
task_id: &str,
provider_id: &str,
provider_type: &str,
status: &str,
total: usize,
processed: usize,
success: usize,
failed: usize,
message: Option<&str>,
error: Option<&str>,
error_samples: Vec<serde_json::Value>,
created_at: u64,
started_at: Option<u64>,
finished_at: Option<u64>,
) -> serde_json::Value {
let updated_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(created_at);
let progress_percent = if total == 0 {
0
} else {
((processed * 100) / total).min(100) as u64
};
json!({
"task_id": task_id,
"provider_id": provider_id,
"provider_type": provider_type,
"status": status,
"total": total,
"processed": processed,
"success": success,
"failed": failed,
"progress_percent": progress_percent,
"message": message,
"error": error,
"error_samples": error_samples,
"created_at": created_at,
"started_at": started_at,
"finished_at": finished_at,
"updated_at": updated_at,
})
}
@@ -0,0 +1,249 @@
use super::execution::{
estimate_admin_provider_oauth_batch_import_total,
execute_admin_provider_oauth_batch_import_for_provider_type,
};
use super::parse::{
build_admin_provider_oauth_batch_task_state, parse_admin_provider_oauth_batch_import_request,
};
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,
is_fixed_provider_type_for_provider_oauth,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::Bytes,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
use tokio::task;
use uuid::Uuid;
const PROVIDER_OAUTH_BATCH_TASK_MAX_ERROR_SAMPLES: usize = 20;
pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_task(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_batch_import_task_provider_id(request_context.path())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let payload = match parse_admin_provider_oauth_batch_import_request(request_body) {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 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(),
);
if total == 0 {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"未找到有效的 Token 数据",
));
}
let task_id = Uuid::new_v4().to_string();
let created_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let submitted_state = build_admin_provider_oauth_batch_task_state(
&task_id,
&provider_id,
&provider_type,
"submitted",
total,
0,
0,
0,
Some("任务已提交,等待执行"),
None,
Vec::new(),
created_at,
None,
None,
);
if state
.save_provider_oauth_batch_task_payload(&task_id, &submitted_state)
.await
.is_err()
{
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth batch task redis unavailable",
));
}
let task_state = state.cloned_app();
let task_id_for_worker = task_id.clone();
let provider_id_for_worker = provider_id.clone();
let provider_type_for_worker = provider_type.clone();
let proxy_node_id = payload.proxy_node_id.clone();
let raw_credentials = payload.credentials.clone();
task::spawn(async move {
let started_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(created_at);
let processing_state = build_admin_provider_oauth_batch_task_state(
&task_id_for_worker,
&provider_id_for_worker,
&provider_type_for_worker,
"processing",
total,
0,
0,
0,
Some("任务开始执行"),
None,
Vec::new(),
created_at,
Some(started_at),
None,
);
let _ = AdminAppState::new(&task_state)
.save_provider_oauth_batch_task_payload(&task_id_for_worker, &processing_state)
.await;
match execute_admin_provider_oauth_batch_import_for_provider_type(
&AdminAppState::new(&task_state),
&provider_id_for_worker,
&provider_type_for_worker,
raw_credentials.as_str(),
proxy_node_id.as_deref(),
)
.await
{
Ok(outcome) => {
let finished_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(started_at);
let error_samples = outcome
.results
.iter()
.filter(|item| {
item.get("status").and_then(serde_json::Value::as_str) == Some("error")
})
.take(PROVIDER_OAUTH_BATCH_TASK_MAX_ERROR_SAMPLES)
.cloned()
.collect::<Vec<_>>();
let message = format!(
"导入完成:成功 {},失败 {}",
outcome.success, outcome.failed
);
let completed_state = build_admin_provider_oauth_batch_task_state(
&task_id_for_worker,
&provider_id_for_worker,
&provider_type_for_worker,
"completed",
outcome.total,
outcome.total,
outcome.success,
outcome.failed,
Some(message.as_str()),
None,
error_samples,
created_at,
Some(started_at),
Some(finished_at),
);
let _ = AdminAppState::new(&task_state)
.save_provider_oauth_batch_task_payload(&task_id_for_worker, &completed_state)
.await;
}
Err(err) => {
let finished_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(started_at);
let error_message = format!("{err:?}");
let failed_state = build_admin_provider_oauth_batch_task_state(
&task_id_for_worker,
&provider_id_for_worker,
&provider_type_for_worker,
"failed",
total,
0,
0,
0,
Some("导入任务执行失败"),
Some(error_message.as_str()),
Vec::new(),
created_at,
Some(started_at),
Some(finished_at),
);
let _ = AdminAppState::new(&task_state)
.save_provider_oauth_batch_task_payload(&task_id_for_worker, &failed_state)
.await;
tracing::warn!(
task_id = %task_id_for_worker,
provider_id = %provider_id_for_worker,
error = %error_message,
"provider oauth batch import task failed"
);
}
}
});
let submitted_response = build_admin_provider_oauth_batch_task_state(
&task_id,
&provider_id,
&provider_type,
"submitted",
total,
0,
0,
0,
Some("任务已提交,等待执行"),
None,
Vec::new(),
created_at,
None,
None,
);
Ok(Json(submitted_response).into_response())
}
@@ -1,562 +0,0 @@
use super::super::quota::codex::refresh_codex_provider_quota_locally;
use super::super::refresh::{
build_internal_control_error_response, build_provider_oauth_auth_config_from_token_payload,
create_provider_oauth_catalog_key, find_duplicate_provider_oauth_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,
};
use super::super::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
consume_provider_oauth_state, enrich_admin_provider_oauth_auth_config,
exchange_admin_provider_oauth_code, is_fixed_provider_type_for_provider_oauth,
json_non_empty_string, json_u64_value, parse_provider_oauth_callback_params,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_complete_key_id, admin_provider_oauth_complete_provider_id,
};
use crate::handlers::admin::shared::encrypt_catalog_secret_with_fallbacks;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_complete_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_complete_key_id(&request_context.request_path) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(serde_json::Value::Object(map)) => map,
_ => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let callback_url = raw_payload
.get("callback_url")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
)
});
let callback_url = match callback_url {
Ok(callback_url) => callback_url,
Err(response) => return Ok(response),
};
let params = parse_provider_oauth_callback_params(callback_url);
let Some(code) = params
.get("code")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let Some(state_nonce) = params
.get("state")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let state_data = match consume_provider_oauth_state(state, state_nonce).await {
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
if state_data.key_id != key_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next();
let Some(key) = key else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
if !key.auth_type.eq_ignore_ascii_case("oauth") {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Key 不是 oauth 认证类型",
));
}
if !state_data.provider_id.trim().is_empty() && state_data.provider_id != key.provider_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let provider_id = key.provider_id.clone();
let provider = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next();
let Some(provider) = provider else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if !state_data.provider_type.trim().is_empty()
&& !state_data
.provider_type
.eq_ignore_ascii_case(&provider_type)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let token_payload = match exchange_admin_provider_oauth_code(
state,
template,
code,
state_nonce,
state_data.pkce_verifier.as_deref(),
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(access_token) = json_non_empty_string(token_payload.get("access_token")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let refresh_token = json_non_empty_string(token_payload.get("refresh_token"));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let expires_at = json_u64_value(token_payload.get("expires_in"))
.map(|expires_in| now_unix_secs.saturating_add(expires_in));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!(provider_type.clone()));
auth_config.insert("updated_at".to_string(), json!(now_unix_secs));
if let Some(token_type) = token_payload.get("token_type").cloned() {
auth_config.insert("token_type".to_string(), token_type);
}
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
if let Some(expires_at) = expires_at {
auth_config.insert("expires_at".to_string(), json!(expires_at));
}
if let Some(scope) = token_payload.get("scope").cloned() {
auth_config.insert("scope".to_string(), scope);
}
enrich_admin_provider_oauth_auth_config(&provider_type, &mut auth_config, &token_payload);
let Some(encrypted_api_key) = encrypt_catalog_secret_with_fallbacks(state, &access_token)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
encrypt_catalog_secret_with_fallbacks(state, &auth_config_json)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let updated = state
.update_provider_catalog_key_oauth_credentials(
&key_id,
&encrypted_api_key,
Some(&encrypted_auth_config),
expires_at,
)
.await?;
if !updated {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
}
let mut account_state_recheck_attempted = false;
let mut account_state_recheck_error = None::<String>;
if provider_type == "codex" {
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
if let Some(endpoint) = endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}) {
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],
)
.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);
}
}
}
}
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,
}))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_complete_provider(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) =
admin_provider_oauth_complete_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(serde_json::Value::Object(map)) => map,
_ => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(callback_url) = raw_payload
.get("callback_url")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let name = raw_payload
.get("name")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let proxy_node_id = raw_payload
.get("proxy_node_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let params = parse_provider_oauth_callback_params(callback_url);
let Some(code) = params
.get("code")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let Some(state_nonce) = params
.get("state")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let state_data = match consume_provider_oauth_state(state, state_nonce).await {
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
if !state_data.key_id.trim().is_empty() || state_data.provider_id != provider_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if provider_type == "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Kiro 不支持 OAuth 授权,请使用导入授权。",
));
}
if !state_data.provider_type.trim().is_empty()
&& !state_data
.provider_type
.eq_ignore_ascii_case(&provider_type)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let token_payload = match exchange_admin_provider_oauth_code(
state,
template,
code,
state_nonce,
state_data.pkce_verifier.as_deref(),
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let (auth_config, access_token, refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(&provider_type, &token_payload);
let Some(access_token) = access_token else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id.as_deref());
let duplicate =
match find_duplicate_provider_oauth_key(state, &provider_id, &auth_config, None).await {
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let replaced = duplicate.is_some();
let persisted_key = if let Some(existing_key) = duplicate {
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} 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)
)
});
match create_provider_oauth_catalog_key(
state,
&provider_id,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
Ok(Json(json!({
"key_id": persisted_key.id,
"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),
"replaced": replaced,
}))
.into_response())
}
@@ -0,0 +1,257 @@
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
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, json_u64_value,
};
use super::shared::{
parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_key_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_complete_key(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_complete_key_id(request_context.path()) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
let payload = match parse_admin_provider_oauth_complete_request_body(request_body) {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let callback = match parse_admin_provider_oauth_complete_callback(&payload.callback_url) {
Ok(callback) => callback,
Err(response) => return Ok(response),
};
let state_data = match state
.consume_provider_oauth_state(&callback.state_nonce)
.await
{
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
if state_data.key_id != key_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next();
let Some(key) = key else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
if !key.auth_type.eq_ignore_ascii_case("oauth") {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Key 不是 oauth 认证类型",
));
}
if !state_data.provider_id.trim().is_empty() && state_data.provider_id != key.provider_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let provider_id = key.provider_id.clone();
let provider = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next();
let Some(provider) = provider else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if !state_data.provider_type.trim().is_empty()
&& !state_data
.provider_type
.eq_ignore_ascii_case(&provider_type)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let token_payload = match state
.exchange_admin_provider_oauth_code(
template,
&callback.code,
&callback.state_nonce,
state_data.pkce_verifier.as_deref(),
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(access_token) = json_non_empty_string(token_payload.get("access_token")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let refresh_token = json_non_empty_string(token_payload.get("refresh_token"));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let expires_at = json_u64_value(token_payload.get("expires_in"))
.map(|expires_in| now_unix_secs.saturating_add(expires_in));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!(provider_type.clone()));
auth_config.insert("updated_at".to_string(), json!(now_unix_secs));
if let Some(token_type) = token_payload.get("token_type").cloned() {
auth_config.insert("token_type".to_string(), token_type);
}
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
if let Some(expires_at) = expires_at {
auth_config.insert("expires_at".to_string(), json!(expires_at));
}
if let Some(scope) = token_payload.get("scope").cloned() {
auth_config.insert("scope".to_string(), scope);
}
enrich_admin_provider_oauth_auth_config(&provider_type, &mut auth_config, &token_payload);
let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(&access_token) else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
state.encrypt_catalog_secret_with_fallbacks(&auth_config_json)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let updated = state
.update_provider_catalog_key_oauth_credentials(
&key_id,
&encrypted_api_key,
Some(&encrypted_auth_config),
expires_at,
)
.await?;
if !updated {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
}
let mut account_state_recheck_attempted = false;
let mut account_state_recheck_error = None::<String>;
if provider_type == "codex" {
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
if let Some(endpoint) = endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}) {
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],
)
.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);
}
}
}
}
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,
}))
.into_response())
}
@@ -0,0 +1,27 @@
mod key;
mod provider;
mod shared;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
response::Response,
};
pub(super) async fn handle_admin_provider_oauth_complete_key(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
key::handle_admin_provider_oauth_complete_key(state, request_context, request_body).await
}
pub(super) async fn handle_admin_provider_oauth_complete_provider(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
provider::handle_admin_provider_oauth_complete_provider(state, request_context, request_body)
.await
}
@@ -0,0 +1,234 @@
use super::super::super::duplicates::find_duplicate_provider_oauth_key;
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::provisioning::{
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
update_existing_provider_oauth_catalog_key,
};
use super::super::super::runtime::refresh_provider_oauth_account_state_after_update;
use super::super::super::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
is_fixed_provider_type_for_provider_oauth,
};
use super::shared::{
parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_complete_provider(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) = admin_provider_oauth_complete_provider_id(request_context.path())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let payload = match parse_admin_provider_oauth_complete_request_body(request_body) {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let callback = match parse_admin_provider_oauth_complete_callback(&payload.callback_url) {
Ok(callback) => callback,
Err(response) => return Ok(response),
};
let state_data = match state
.consume_provider_oauth_state(&callback.state_nonce)
.await
{
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
if !state_data.key_id.trim().is_empty() || state_data.provider_id != provider_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if provider_type == "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Kiro 不支持 OAuth 授权,请使用导入授权。",
));
}
if !state_data.provider_type.trim().is_empty()
&& !state_data
.provider_type
.eq_ignore_ascii_case(&provider_type)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let token_payload = match state
.exchange_admin_provider_oauth_code(
template,
&callback.code,
&callback.state_nonce,
state_data.pkce_verifier.as_deref(),
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let (auth_config, access_token, refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(&provider_type, &token_payload);
let Some(access_token) = access_token else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(payload.proxy_node_id.as_deref());
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
.await
{
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let replaced = duplicate.is_some();
let persisted_key = if let Some(existing_key) = duplicate {
match state
.update_existing_provider_oauth_catalog_key(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let name = payload
.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)
)
});
match state
.create_provider_oauth_catalog_key(
&provider_id,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = state
.refresh_provider_oauth_account_state_after_update(&provider, &persisted_key.id)
.await;
Ok(Json(json!({
"key_id": persisted_key.id,
"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),
"replaced": replaced,
}))
.into_response())
}
@@ -0,0 +1,116 @@
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::state::parse_provider_oauth_callback_params;
use axum::{
body::{Body, Bytes},
http,
response::Response,
};
pub(super) struct AdminProviderOAuthCompleteRequest {
pub(super) callback_url: String,
pub(super) name: Option<String>,
pub(super) proxy_node_id: Option<String>,
}
pub(super) struct AdminProviderOAuthCompleteCallback {
pub(super) code: String,
pub(super) state_nonce: String,
}
pub(super) fn parse_admin_provider_oauth_callback_url(
raw_payload: &serde_json::Map<String, serde_json::Value>,
) -> Result<String, Response<Body>> {
raw_payload
.get("callback_url")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
)
})
}
pub(super) fn extract_admin_provider_oauth_code(
params: &std::collections::BTreeMap<String, String>,
) -> Result<String, Response<Body>> {
params
.get("code")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
)
})
}
pub(super) fn extract_admin_provider_oauth_state(
params: &std::collections::BTreeMap<String, String>,
) -> Result<String, Response<Body>> {
params
.get("state")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
)
})
}
pub(super) fn parse_admin_provider_oauth_complete_request_body(
request_body: Option<&Bytes>,
) -> Result<AdminProviderOAuthCompleteRequest, Response<Body>> {
let Some(request_body) = request_body else {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(serde_json::Value::Object(map)) => map,
_ => {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let callback_url = parse_admin_provider_oauth_callback_url(&raw_payload)?;
Ok(AdminProviderOAuthCompleteRequest {
callback_url,
name: raw_payload
.get("name")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
proxy_node_id: raw_payload
.get("proxy_node_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
})
}
pub(super) fn parse_admin_provider_oauth_complete_callback(
callback_url: &str,
) -> Result<AdminProviderOAuthCompleteCallback, Response<Body>> {
let params = parse_provider_oauth_callback_params(callback_url);
let code = extract_admin_provider_oauth_code(&params)?;
let state_nonce = extract_admin_provider_oauth_state(&params)?;
Ok(AdminProviderOAuthCompleteCallback { code, state_nonce })
}
@@ -1,550 +0,0 @@
use super::super::refresh::{
build_internal_control_error_response, build_provider_oauth_auth_config_from_token_payload,
create_provider_oauth_catalog_key, find_duplicate_provider_oauth_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,
};
use super::super::state::{
build_admin_provider_oauth_backend_unavailable_response, build_kiro_device_key_name,
current_unix_secs, decode_jwt_claims, default_kiro_device_region,
default_kiro_device_start_url, generate_provider_oauth_nonce, json_non_empty_string,
json_u64_value, normalize_kiro_device_region, poll_admin_kiro_device_token,
read_provider_oauth_device_session, register_admin_kiro_device_oidc_client,
save_provider_oauth_device_session, start_admin_kiro_device_authorization,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_device_authorize_provider_id, admin_provider_oauth_device_poll_provider_id,
};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use aether_data::repository::provider_oauth::{
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde::Deserialize;
use serde_json::json;
#[derive(Debug, Deserialize)]
struct AdminProviderOAuthDeviceAuthorizePayload {
#[serde(default = "default_kiro_device_start_url")]
start_url: String,
#[serde(default = "default_kiro_device_region")]
region: String,
proxy_node_id: Option<String>,
}
#[derive(Debug, Deserialize)]
struct AdminProviderOAuthDevicePollPayload {
session_id: String,
}
pub(super) async fn handle_admin_provider_oauth_device_authorize(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_device_authorize_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload =
match serde_json::from_slice::<AdminProviderOAuthDeviceAuthorizePayload>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"设备授权仅支持 Kiro provider",
));
}
let region = normalize_kiro_device_region(Some(payload.region.as_str())).ok_or_else(|| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "region 格式无效")
});
let region = match region {
Ok(region) => region,
Err(response) => return Ok(response),
};
let start_url = payload.start_url.trim();
let start_url = if start_url.is_empty() {
default_kiro_device_start_url()
} else {
start_url.to_string()
};
let client_registration =
match register_admin_kiro_device_oidc_client(state, &region, &start_url).await {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(client_id) = json_non_empty_string(client_registration.get("clientId")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let Some(client_secret) = json_non_empty_string(client_registration.get("clientSecret")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let device_authorization = match start_admin_kiro_device_authorization(
state,
&region,
&client_id,
&client_secret,
&start_url,
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(device_code) = json_non_empty_string(
device_authorization
.get("deviceCode")
.or_else(|| device_authorization.get("device_code")),
) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
));
};
let user_code = json_non_empty_string(
device_authorization
.get("userCode")
.or_else(|| device_authorization.get("user_code")),
)
.unwrap_or_default();
let verification_uri = json_non_empty_string(
device_authorization
.get("verificationUri")
.or_else(|| device_authorization.get("verification_uri"))
.or_else(|| device_authorization.get("verificationUrl")),
)
.unwrap_or_default();
let verification_uri_complete = json_non_empty_string(
device_authorization
.get("verificationUriComplete")
.or_else(|| device_authorization.get("verification_uri_complete"))
.or_else(|| device_authorization.get("verificationUrlComplete")),
)
.unwrap_or_else(|| verification_uri.clone());
let expires_in = json_u64_value(
device_authorization
.get("expiresIn")
.or_else(|| device_authorization.get("expires_in")),
)
.unwrap_or(600);
let interval = json_u64_value(device_authorization.get("interval")).unwrap_or(5);
let now_unix_secs = current_unix_secs();
let session_id = generate_provider_oauth_nonce();
let session = StoredAdminProviderOAuthDeviceSession {
provider_id: provider_id.clone(),
region,
client_id,
client_secret,
device_code,
interval,
expires_at_unix_secs: now_unix_secs.saturating_add(expires_in),
status: "pending".to_string(),
proxy_node_id: payload
.proxy_node_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
created_at_unix_secs: now_unix_secs,
key_id: None,
email: None,
replaced: false,
error_msg: None,
};
if let Err(response) = save_provider_oauth_device_session(
state,
&session_id,
&session,
expires_in.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
)
.await
{
return Ok(response);
}
Ok(Json(json!({
"session_id": session_id,
"user_code": user_code,
"verification_uri": verification_uri,
"verification_uri_complete": verification_uri_complete,
"expires_in": expires_in,
"interval": interval,
}))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_device_poll(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_device_poll_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload = match serde_json::from_slice::<AdminProviderOAuthDevicePollPayload>(request_body)
{
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let session_id = payload.session_id.trim();
if session_id.is_empty() {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"session_id 不能为空",
));
}
let Some(mut session) = read_provider_oauth_device_session(state, session_id).await? else {
return Ok(Json(json!({
"status": "expired",
"error": "会话不存在或已过期",
"replaced": false,
}))
.into_response());
};
if session.provider_id != provider_id {
return Ok(Json(json!({
"status": "error",
"error": "会话与 Provider 不匹配",
"replaced": false,
}))
.into_response());
}
if session.status == "authorized" {
return Ok(Json(json!({
"status": "authorized",
"key_id": session.key_id,
"email": session.email,
"replaced": session.replaced,
}))
.into_response());
}
if matches!(session.status.as_str(), "expired" | "error") {
return Ok(Json(json!({
"status": session.status,
"error": session.error_msg,
"replaced": session.replaced,
}))
.into_response());
}
if current_unix_secs() > session.expires_at_unix_secs {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let token_result = match poll_admin_kiro_device_token(
state,
&session.region,
&session.client_id,
&session.client_secret,
&session.device_code,
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
if token_result
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_code = json_non_empty_string(token_result.get("error")).unwrap_or_default();
if error_code == "authorization_pending" {
return Ok(Json(json!({"status": "pending", "replaced": false})).into_response());
}
if error_code == "slow_down" {
return Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response());
}
if error_code == "expired_token" {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
if error_code == "access_denied" {
session.status = "error".to_string();
session.error_msg = Some("用户拒绝授权".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"error",
Json(json!({
"status": "error",
"error": "用户拒绝授权",
"replaced": false,
}))
.into_response(),
));
}
let error_message = json_non_empty_string(token_result.get("error_description"))
.or_else(|| (!error_code.is_empty()).then_some(error_code.clone()))
.unwrap_or_else(|| "未知错误".to_string());
return Ok(Json(json!({
"status": "error",
"error": error_message,
"replaced": false,
}))
.into_response());
}
let Some(access_token) = json_non_empty_string(token_result.get("accessToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let Some(refresh_token) = json_non_empty_string(token_result.get("refreshToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let expires_at = json_u64_value(token_result.get("expiresIn"))
.map(|expires_in| current_unix_secs().saturating_add(expires_in))
.unwrap_or_else(|| current_unix_secs().saturating_add(3600));
let email = decode_jwt_claims(&access_token)
.and_then(|claims| claims.get("email").cloned())
.and_then(|value| value.as_str().map(ToOwned::to_owned));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!("kiro"));
auth_config.insert("auth_method".to_string(), json!("idc"));
auth_config.insert("refresh_token".to_string(), json!(refresh_token.clone()));
auth_config.insert("client_id".to_string(), json!(session.client_id.clone()));
auth_config.insert(
"client_secret".to_string(),
json!(session.client_secret.clone()),
);
auth_config.insert("region".to_string(), json!(session.region.clone()));
auth_config.insert("auth_region".to_string(), json!(session.region.clone()));
auth_config.insert("access_token".to_string(), json!(access_token.clone()));
auth_config.insert("expires_at".to_string(), json!(expires_at));
if let Some(email) = email.as_ref() {
auth_config.insert("email".to_string(), json!(email));
}
let duplicate =
match find_duplicate_provider_oauth_key(state, &provider_id, &auth_config, None).await {
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(Json(json!({
"status": "error",
"error": detail,
"replaced": false,
}))
.into_response());
}
};
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
let api_formats = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.filter(|endpoint| endpoint.is_active)
.map(|endpoint| endpoint.api_format)
.collect::<Vec<_>>();
let mut replaced = false;
let persisted_key = if let Some(existing_key) = duplicate {
replaced = true;
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let key_name = build_kiro_device_key_name(email.as_deref(), Some(&refresh_token));
match create_provider_oauth_catalog_key(
state,
&provider_id,
&key_name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
session.status = "authorized".to_string();
session.key_id = Some(persisted_key.id.clone());
session.email = email.clone();
session.replaced = replaced;
session.error_msg = None;
let _ = save_provider_oauth_device_session(state, session_id, &session, 60).await;
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"authorized",
Json(json!({
"status": "authorized",
"key_id": persisted_key.id,
"email": email,
"replaced": replaced,
}))
.into_response(),
))
}
fn attach_admin_provider_oauth_device_poll_terminal_response(
session_id: &str,
status: &str,
response: Response<Body>,
) -> Response<Body> {
match status {
"authorized" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_completed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"expired" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_expired",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"error" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_failed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
_ => response,
}
}
@@ -0,0 +1,194 @@
use super::session::AdminProviderOAuthDeviceAuthorizePayload;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::state::{
build_admin_provider_oauth_backend_unavailable_response, current_unix_secs,
default_kiro_device_start_url, generate_provider_oauth_nonce, json_non_empty_string,
json_u64_value, normalize_kiro_device_region,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_device_authorize_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use aether_data::repository::provider_oauth::{
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn handle_admin_provider_oauth_device_authorize(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_device_authorize_provider_id(request_context.path())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload =
match serde_json::from_slice::<AdminProviderOAuthDeviceAuthorizePayload>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"设备授权仅支持 Kiro provider",
));
}
let region = normalize_kiro_device_region(Some(payload.region.as_str())).ok_or_else(|| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "region 格式无效")
});
let region = match region {
Ok(region) => region,
Err(response) => return Ok(response),
};
let start_url = payload.start_url.trim();
let start_url = if start_url.is_empty() {
default_kiro_device_start_url()
} else {
start_url.to_string()
};
let client_registration = match state
.register_admin_kiro_device_oidc_client(&region, &start_url)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(client_id) = json_non_empty_string(client_registration.get("clientId")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let Some(client_secret) = json_non_empty_string(client_registration.get("clientSecret")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let device_authorization = match state
.start_admin_kiro_device_authorization(&region, &client_id, &client_secret, &start_url)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(device_code) = json_non_empty_string(
device_authorization
.get("deviceCode")
.or_else(|| device_authorization.get("device_code")),
) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
));
};
let user_code = json_non_empty_string(
device_authorization
.get("userCode")
.or_else(|| device_authorization.get("user_code")),
)
.unwrap_or_default();
let verification_uri = json_non_empty_string(
device_authorization
.get("verificationUri")
.or_else(|| device_authorization.get("verification_uri"))
.or_else(|| device_authorization.get("verificationUrl")),
)
.unwrap_or_default();
let verification_uri_complete = json_non_empty_string(
device_authorization
.get("verificationUriComplete")
.or_else(|| device_authorization.get("verification_uri_complete"))
.or_else(|| device_authorization.get("verificationUrlComplete")),
)
.unwrap_or_else(|| verification_uri.clone());
let expires_in = json_u64_value(
device_authorization
.get("expiresIn")
.or_else(|| device_authorization.get("expires_in")),
)
.unwrap_or(600);
let interval = json_u64_value(device_authorization.get("interval")).unwrap_or(5);
let now_unix_secs = current_unix_secs();
let session_id = generate_provider_oauth_nonce();
let session = StoredAdminProviderOAuthDeviceSession {
provider_id: provider_id.clone(),
region,
client_id,
client_secret,
device_code,
interval,
expires_at_unix_secs: now_unix_secs.saturating_add(expires_in),
status: "pending".to_string(),
proxy_node_id: payload
.proxy_node_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
created_at_unix_secs: now_unix_secs,
key_id: None,
email: None,
replaced: false,
error_msg: None,
};
if let Err(response) = state
.save_provider_oauth_device_session(
&session_id,
&session,
expires_in.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
)
.await
{
return Ok(response);
}
Ok(Json(json!({
"session_id": session_id,
"user_code": user_code,
"verification_uri": verification_uri,
"verification_uri_complete": verification_uri_complete,
"expires_in": expires_in,
"interval": interval,
}))
.into_response())
}
@@ -0,0 +1,32 @@
mod authorize;
mod poll;
mod session;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
response::Response,
};
#[cfg(any())]
pub(super) use self::authorize::handle_admin_provider_oauth_device_authorize;
#[cfg(any())]
pub(super) use self::poll::handle_admin_provider_oauth_device_poll;
pub(super) async fn handle_admin_provider_oauth_device_authorize(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
authorize::handle_admin_provider_oauth_device_authorize(state, request_context, request_body)
.await
}
pub(super) async fn handle_admin_provider_oauth_device_poll(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
poll::handle_admin_provider_oauth_device_poll(state, request_context, request_body).await
}
@@ -0,0 +1,329 @@
use super::session::{
attach_admin_provider_oauth_device_poll_terminal_response, AdminProviderOAuthDevicePollPayload,
};
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::provisioning::{
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
use crate::handlers::admin::provider::oauth::state::{
build_admin_provider_oauth_backend_unavailable_response, build_kiro_device_key_name,
current_unix_secs, decode_jwt_claims, json_non_empty_string, json_u64_value,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_device_poll_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn handle_admin_provider_oauth_device_poll(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) = admin_provider_oauth_device_poll_provider_id(request_context.path())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload = match serde_json::from_slice::<AdminProviderOAuthDevicePollPayload>(request_body)
{
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let session_id = payload.session_id.trim();
if session_id.is_empty() {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"session_id 不能为空",
));
}
let Some(mut session) = state.read_provider_oauth_device_session(session_id).await? else {
return Ok(Json(json!({
"status": "expired",
"error": "会话不存在或已过期",
"replaced": false,
}))
.into_response());
};
if session.provider_id != provider_id {
return Ok(Json(json!({
"status": "error",
"error": "会话与 Provider 不匹配",
"replaced": false,
}))
.into_response());
}
if session.status == "authorized" {
return Ok(Json(json!({
"status": "authorized",
"key_id": session.key_id,
"email": session.email,
"replaced": session.replaced,
}))
.into_response());
}
if matches!(session.status.as_str(), "expired" | "error") {
return Ok(Json(json!({
"status": session.status,
"error": session.error_msg,
"replaced": session.replaced,
}))
.into_response());
}
if current_unix_secs() > session.expires_at_unix_secs {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = state
.save_provider_oauth_device_session(session_id, &session, 30)
.await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let token_result = match state
.poll_admin_kiro_device_token(
&session.region,
&session.client_id,
&session.client_secret,
&session.device_code,
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
if token_result
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_code = json_non_empty_string(token_result.get("error")).unwrap_or_default();
if error_code == "authorization_pending" {
return Ok(Json(json!({"status": "pending", "replaced": false})).into_response());
}
if error_code == "slow_down" {
return Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response());
}
if error_code == "expired_token" {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = state
.save_provider_oauth_device_session(session_id, &session, 30)
.await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
if error_code == "access_denied" {
session.status = "error".to_string();
session.error_msg = Some("用户拒绝授权".to_string());
let _ = state
.save_provider_oauth_device_session(session_id, &session, 30)
.await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"error",
Json(json!({
"status": "error",
"error": "用户拒绝授权",
"replaced": false,
}))
.into_response(),
));
}
let error_message = json_non_empty_string(token_result.get("error_description"))
.or_else(|| (!error_code.is_empty()).then_some(error_code.clone()))
.unwrap_or_else(|| "未知错误".to_string());
return Ok(Json(json!({
"status": "error",
"error": error_message,
"replaced": false,
}))
.into_response());
}
let Some(access_token) = json_non_empty_string(token_result.get("accessToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let Some(refresh_token) = json_non_empty_string(token_result.get("refreshToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let expires_at = json_u64_value(token_result.get("expiresIn"))
.map(|expires_in| current_unix_secs().saturating_add(expires_in))
.unwrap_or_else(|| current_unix_secs().saturating_add(3600));
let email = decode_jwt_claims(&access_token)
.and_then(|claims| claims.get("email").cloned())
.and_then(|value| value.as_str().map(ToOwned::to_owned));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!("kiro"));
auth_config.insert("auth_method".to_string(), json!("idc"));
auth_config.insert("refresh_token".to_string(), json!(refresh_token.clone()));
auth_config.insert("client_id".to_string(), json!(session.client_id.clone()));
auth_config.insert(
"client_secret".to_string(),
json!(session.client_secret.clone()),
);
auth_config.insert("region".to_string(), json!(session.region.clone()));
auth_config.insert("auth_region".to_string(), json!(session.region.clone()));
auth_config.insert("access_token".to_string(), json!(access_token.clone()));
auth_config.insert("expires_at".to_string(), json!(expires_at));
if let Some(email) = email.as_ref() {
auth_config.insert("email".to_string(), json!(email));
}
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
.await
{
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(Json(json!({
"status": "error",
"error": detail,
"replaced": false,
}))
.into_response());
}
};
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
let api_formats = provider_oauth_active_api_formats(
&state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?,
);
let mut replaced = false;
let persisted_key = if let Some(existing_key) = duplicate {
replaced = true;
match state
.update_existing_provider_oauth_catalog_key(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let key_name = build_kiro_device_key_name(email.as_deref(), Some(&refresh_token));
match state
.create_provider_oauth_catalog_key(
&provider_id,
&key_name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = state
.refresh_provider_oauth_account_state_after_update(&provider, &persisted_key.id)
.await;
session.status = "authorized".to_string();
session.key_id = Some(persisted_key.id.clone());
session.email = email.clone();
session.replaced = replaced;
session.error_msg = None;
let _ = state
.save_provider_oauth_device_session(session_id, &session, 60)
.await;
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"authorized",
Json(json!({
"status": "authorized",
"key_id": persisted_key.id,
"email": email,
"replaced": replaced,
}))
.into_response(),
))
}
@@ -0,0 +1,51 @@
use crate::handlers::admin::provider::oauth::state::{
default_kiro_device_region, default_kiro_device_start_url,
};
use crate::handlers::admin::shared::attach_admin_audit_response;
use axum::{body::Body, response::Response};
use serde::Deserialize;
#[derive(Debug, Deserialize)]
pub(super) struct AdminProviderOAuthDeviceAuthorizePayload {
#[serde(default = "default_kiro_device_start_url")]
pub(super) start_url: String,
#[serde(default = "default_kiro_device_region")]
pub(super) region: String,
pub(super) proxy_node_id: Option<String>,
}
#[derive(Debug, Deserialize)]
pub(super) struct AdminProviderOAuthDevicePollPayload {
pub(super) session_id: String,
}
pub(super) fn attach_admin_provider_oauth_device_poll_terminal_response(
session_id: &str,
status: &str,
response: Response<Body>,
) -> Response<Body> {
match status {
"authorized" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_completed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"expired" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_expired",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"error" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_failed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
_ => response,
}
}
@@ -1,16 +1,18 @@
use super::super::refresh::{
build_internal_control_error_response, build_provider_oauth_auth_config_from_token_payload,
create_provider_oauth_catalog_key, find_duplicate_provider_oauth_key,
use super::super::duplicates::find_duplicate_provider_oauth_key;
use super::super::errors::build_internal_control_error_response;
use super::super::provisioning::{
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,
update_existing_provider_oauth_catalog_key,
};
use super::super::runtime::refresh_provider_oauth_account_state_after_update;
use super::super::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
exchange_admin_provider_oauth_refresh_token, is_fixed_provider_type_for_provider_oauth,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_import_provider_id;
use crate::{AppState, GatewayError};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::Body,
http,
@@ -21,15 +23,14 @@ use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
state: &AppState,
request_context: &GatewayPublicRequestContext,
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) = admin_provider_oauth_import_provider_id(&request_context.request_path)
else {
let Some(provider_id) = admin_provider_oauth_import_provider_id(request_context.path()) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
@@ -96,13 +97,13 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
return Ok(build_admin_provider_oauth_backend_unavailable_response());
};
let token_payload =
match exchange_admin_provider_oauth_refresh_token(state, template, refresh_token_input)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let token_payload = match state
.exchange_admin_provider_oauth_refresh_token(template, refresh_token_input)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let (mut auth_config, access_token, returned_refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(&provider_type, &token_payload);
@@ -124,28 +125,30 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id.as_deref());
let duplicate =
match find_duplicate_provider_oauth_key(state, &provider_id, &auth_config, None).await {
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
.await
{
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let replaced = duplicate.is_some();
let persisted_key = if let Some(existing_key) = duplicate {
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
expires_at,
)
.await?
match state
.update_existing_provider_oauth_catalog_key(
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
@@ -175,17 +178,17 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
.unwrap_or(0)
)
});
match create_provider_oauth_catalog_key(
state,
&provider_id,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
match state
.create_provider_oauth_catalog_key(
&provider_id,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
@@ -197,7 +200,8 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
}
};
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
let _ = state
.refresh_provider_oauth_account_state_after_update(&provider, &persisted_key.id)
.await;
Ok(Json(json!({
@@ -2,7 +2,6 @@ use super::state::{
build_admin_provider_oauth_backend_unavailable_response,
build_admin_provider_oauth_supported_types_payload,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_batch_import_provider_id,
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
@@ -10,7 +9,8 @@ use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_import_provider_id, admin_provider_oauth_refresh_key_id,
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
};
use crate::{AppState, GatewayError};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
http,
@@ -28,11 +28,11 @@ mod start;
mod tasks;
pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
let Some(decision) = request_context.decision() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("provider_oauth_manage") {
@@ -40,11 +40,11 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
}
let route_kind = decision.route_kind.as_deref();
let method = &request_context.request_method;
let method = &request_context.method();
if route_kind == Some("supported_types")
&& *method == http::Method::GET
&& request_context.request_path == "/api/admin/provider-oauth/supported-types"
&& request_context.path() == "/api/admin/provider-oauth/supported-types"
{
return Ok(Some(
Json(build_admin_provider_oauth_supported_types_payload()).into_response(),
@@ -58,7 +58,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_authorization_started",
"start_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_start_key_id(&request_context.request_path),
admin_provider_oauth_start_key_id(request_context.path()),
)));
}
@@ -70,7 +70,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_authorization_started",
"start_provider_oauth_for_provider",
"provider",
admin_provider_oauth_start_provider_id(&request_context.request_path),
admin_provider_oauth_start_provider_id(request_context.path()),
)));
}
@@ -93,7 +93,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_completed",
"complete_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_complete_key_id(&request_context.request_path),
admin_provider_oauth_complete_key_id(request_context.path()),
)));
}
@@ -105,7 +105,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_refreshed",
"refresh_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_refresh_key_id(&request_context.request_path),
admin_provider_oauth_refresh_key_id(request_context.path()),
)));
}
@@ -121,7 +121,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_completed",
"complete_provider_oauth_for_provider",
"provider",
admin_provider_oauth_complete_provider_id(&request_context.request_path),
admin_provider_oauth_complete_provider_id(request_context.path()),
)));
}
@@ -137,7 +137,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_refresh_token_imported",
"import_provider_oauth_refresh_token",
"provider",
admin_provider_oauth_import_provider_id(&request_context.request_path),
admin_provider_oauth_import_provider_id(request_context.path()),
)));
}
@@ -150,7 +150,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_batch_import_completed",
"batch_import_provider_oauth",
"provider",
admin_provider_oauth_batch_import_provider_id(&request_context.request_path),
admin_provider_oauth_batch_import_provider_id(request_context.path()),
)));
}
@@ -166,7 +166,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_batch_import_started",
"start_provider_oauth_batch_import",
"provider",
admin_provider_oauth_batch_import_task_provider_id(&request_context.request_path),
admin_provider_oauth_batch_import_task_provider_id(request_context.path()),
)));
}
@@ -182,7 +182,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
"admin_provider_oauth_device_authorization_started",
"start_provider_oauth_device_authorization",
"provider",
admin_provider_oauth_device_authorize_provider_id(&request_context.request_path),
admin_provider_oauth_device_authorize_provider_id(request_context.path()),
)));
}
@@ -1,227 +1,28 @@
use super::super::quota::shared::persist_provider_quota_refresh_state;
use super::super::refresh::{
build_internal_control_error_response, merge_provider_oauth_refresh_failure_reason,
normalize_provider_oauth_refresh_error_message, provider_oauth_runtime_endpoint_for_provider,
refresh_provider_oauth_account_state_after_update,
};
use super::super::state::is_fixed_provider_type_for_provider_oauth;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_refresh_key_id;
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks;
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
mod execution;
mod helpers;
mod request;
mod response;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{body::Body, response::Response};
pub(super) async fn handle_admin_provider_oauth_refresh_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_refresh_key_id(&request_context.request_path) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
if !key.auth_type.eq_ignore_ascii_case("oauth") {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Key 不是 oauth 认证类型",
));
}
let request =
match request::parse_admin_provider_oauth_refresh_request(state, request_context).await? {
helpers::RefreshDispatch::Continue(request) => request,
helpers::RefreshDispatch::Respond(response) => return Ok(response),
};
let Some(encrypted_auth_config) = key.encrypted_auth_config.as_deref() else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 auth_config,无法 refresh",
));
};
let Some(decrypted_auth_config) =
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), encrypted_auth_config)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let parsed_auth_config = serde_json::from_str::<serde_json::Value>(&decrypted_auth_config)
.ok()
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
let has_refresh_token = parsed_auth_config
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty());
if !has_refresh_token {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 refresh_token,需要重新授权",
));
}
let provider_id = key.provider_id.clone();
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"找不到有效端点,无法 refresh",
));
};
let Some(transport) = state
.read_provider_transport_snapshot(&provider_id, &endpoint.id, &key_id)
.await?
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Provider transport snapshot unavailable",
));
let refreshed = match execution::execute_admin_provider_oauth_refresh(state, request).await? {
helpers::RefreshDispatch::Continue(refreshed) => refreshed,
helpers::RefreshDispatch::Respond(response) => return Ok(response),
};
match state.force_local_oauth_refresh_entry(&transport).await {
Ok(Some(_)) => {}
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 refresh_token,需要重新授权",
));
}
Err(crate::provider_transport::LocalOAuthRefreshError::HttpStatus {
status_code,
body_excerpt,
..
}) => {
let error_reason = normalize_provider_oauth_refresh_error_message(
Some(status_code),
Some(body_excerpt.as_str()),
);
if matches!(status_code, 400 | 401 | 403) {
let merged_reason = merge_provider_oauth_refresh_failure_reason(
key.oauth_invalid_reason.as_deref(),
format!(
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 ({status_code}): {error_reason}"
)
.as_str(),
);
if let Some(merged_reason) = merged_reason {
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let _ = persist_provider_quota_refresh_state(
state,
&key_id,
None,
Some(now_unix_secs),
Some(merged_reason),
None,
)
.await?;
}
}
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Token 刷新失败:{error_reason}"),
));
}
Err(crate::provider_transport::LocalOAuthRefreshError::Transport { source, .. }) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
format!("Token 刷新失败:{}", source),
));
}
Err(crate::provider_transport::LocalOAuthRefreshError::InvalidResponse {
message, ..
}) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Token 刷新失败:{message}"),
));
}
}
if !key
.oauth_invalid_reason
.as_deref()
.map(str::trim)
.is_some_and(|value| value.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX))
{
let _ = state
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
.await?;
}
let refreshed_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
.unwrap_or(key);
let refreshed_auth_config = refreshed_key
.encrypted_auth_config
.as_deref()
.and_then(|ciphertext| {
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
})
.and_then(|plaintext| serde_json::from_str::<serde_json::Value>(&plaintext).ok())
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
let (account_state_recheck_attempted, account_state_recheck_error) =
refresh_provider_oauth_account_state_after_update(state, &provider, &key_id).await?;
Ok(Json(json!({
"provider_type": provider_type,
"expires_at": refreshed_auth_config.get("expires_at").cloned().unwrap_or(serde_json::Value::Null),
"has_refresh_token": refreshed_auth_config
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty()),
"email": refreshed_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,
}))
.into_response())
Ok(response::admin_provider_oauth_refresh_success_response(
refreshed,
))
}
@@ -0,0 +1,106 @@
use super::super::super::errors::{
merge_provider_oauth_refresh_failure_reason, normalize_provider_oauth_refresh_error_message,
};
use super::super::super::quota::shared::persist_provider_quota_refresh_state;
use super::super::super::runtime::refresh_provider_oauth_account_state_after_update;
use super::helpers::{self, RefreshDispatch, RefreshRequestContext, RefreshSuccessContext};
use super::response;
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::handlers::admin::request::{AdminAppState, AdminLocalOAuthRefreshError};
use crate::GatewayError;
use axum::http;
pub(super) async fn execute_admin_provider_oauth_refresh(
state: &AdminAppState<'_>,
request: RefreshRequestContext,
) -> Result<RefreshDispatch<RefreshSuccessContext>, GatewayError> {
let RefreshRequestContext {
key_id,
key,
provider,
provider_type,
transport,
} = request;
match state.force_local_oauth_refresh_entry(&transport).await {
Ok(Some(_)) => {}
Ok(None) => {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 refresh_token,需要重新授权",
)));
}
Err(AdminLocalOAuthRefreshError::HttpStatus {
status_code,
body_excerpt,
..
}) => {
let error_reason = normalize_provider_oauth_refresh_error_message(
Some(status_code),
Some(body_excerpt.as_str()),
);
if matches!(status_code, 400 | 401 | 403) {
let failure_reason = format!(
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 ({status_code}): {error_reason}"
);
let merged_reason = merge_provider_oauth_refresh_failure_reason(
key.oauth_invalid_reason.as_deref(),
&failure_reason,
);
if let Some(merged_reason) = merged_reason {
let _ = persist_provider_quota_refresh_state(
state,
&key_id,
None,
Some(helpers::unix_now_secs()),
Some(merged_reason),
None,
)
.await?;
}
}
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_failed_bad_request_response(&error_reason),
));
}
Err(AdminLocalOAuthRefreshError::Transport { source, .. }) => {
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_failed_service_unavailable_response(source.to_string()),
));
}
Err(AdminLocalOAuthRefreshError::InvalidResponse { message, .. }) => {
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_failed_bad_request_response(&message),
));
}
}
if !helpers::key_is_account_blocked(&key, OAUTH_ACCOUNT_BLOCK_PREFIX) {
let _ = state
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
.await?;
}
let refreshed_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
.unwrap_or(key);
let refreshed_auth_config = helpers::refreshed_auth_config_object(
state,
refreshed_key.encrypted_auth_config.as_deref(),
);
let (account_state_recheck_attempted, account_state_recheck_error) = state
.refresh_provider_oauth_account_state_after_update(&provider, &key_id)
.await?;
Ok(RefreshDispatch::Continue(RefreshSuccessContext {
provider_type,
refreshed_auth_config,
account_state_recheck_attempted,
account_state_recheck_error,
}))
}
@@ -0,0 +1,75 @@
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use axum::{body::Body, response::Response};
use serde_json::{Map, Value};
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) enum RefreshDispatch<T> {
Continue(T),
Respond(Response<Body>),
}
pub(super) struct RefreshRequestContext {
pub(super) key_id: String,
pub(super) key: StoredProviderCatalogKey,
pub(super) provider: StoredProviderCatalogProvider,
pub(super) provider_type: String,
pub(super) transport: AdminGatewayProviderTransportSnapshot,
}
pub(super) struct RefreshSuccessContext {
pub(super) provider_type: String,
pub(super) refreshed_auth_config: Map<String, Value>,
pub(super) account_state_recheck_attempted: bool,
pub(super) account_state_recheck_error: Option<String>,
}
pub(super) fn decrypt_auth_config(
state: &AdminAppState<'_>,
encrypted_auth_config: &str,
) -> Option<String> {
state.decrypt_catalog_secret_with_fallbacks(encrypted_auth_config)
}
pub(super) fn parse_auth_config_object(plaintext: &str) -> Map<String, Value> {
serde_json::from_str::<Value>(plaintext)
.ok()
.and_then(|value| value.as_object().cloned())
.unwrap_or_default()
}
pub(super) fn refreshed_auth_config_object(
state: &AdminAppState<'_>,
encrypted_auth_config: Option<&str>,
) -> Map<String, Value> {
encrypted_auth_config
.and_then(|ciphertext| decrypt_auth_config(state, ciphertext))
.map(|plaintext| parse_auth_config_object(&plaintext))
.unwrap_or_default()
}
pub(super) fn auth_config_has_refresh_token(auth_config: &Map<String, Value>) -> bool {
auth_config
.get("refresh_token")
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
pub(super) fn key_is_account_blocked(key: &StoredProviderCatalogKey, block_prefix: &str) -> bool {
key.oauth_invalid_reason
.as_deref()
.map(str::trim)
.is_some_and(|value| value.starts_with(block_prefix))
}
pub(super) fn unix_now_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
}
@@ -0,0 +1,106 @@
use super::super::super::runtime::provider_oauth_runtime_endpoint_for_provider;
use super::super::super::state::is_fixed_provider_type_for_provider_oauth;
use super::helpers::{self, RefreshDispatch, RefreshRequestContext};
use super::response;
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_refresh_key_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::http;
pub(super) async fn parse_admin_provider_oauth_refresh_request(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<RefreshDispatch<RefreshRequestContext>, GatewayError> {
let Some(key_id) = admin_provider_oauth_refresh_key_id(request_context.path()) else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
)));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
)));
};
if !key.auth_type.eq_ignore_ascii_case("oauth") {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Key 不是 oauth 认证类型",
)));
}
let Some(encrypted_auth_config) = key.encrypted_auth_config.as_deref() else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 auth_config,无法 refresh",
)));
};
let Some(decrypted_auth_config) = helpers::decrypt_auth_config(state, encrypted_auth_config)
else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
)));
};
let parsed_auth_config = helpers::parse_auth_config_object(&decrypted_auth_config);
if !helpers::auth_config_has_refresh_token(&parsed_auth_config) {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 refresh_token,需要重新授权",
)));
}
let provider_id = key.provider_id.clone();
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
)));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
)));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"找不到有效端点,无法 refresh",
)));
};
let Some(transport) = state
.read_provider_transport_snapshot(&provider_id, &endpoint.id, &key_id)
.await?
else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"Provider transport snapshot unavailable",
)));
};
Ok(RefreshDispatch::Continue(RefreshRequestContext {
key_id,
key,
provider,
provider_type,
transport,
}))
}
@@ -0,0 +1,61 @@
use super::super::super::errors::build_internal_control_error_response;
use super::helpers::RefreshSuccessContext;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::{json, Value};
pub(super) fn control_error_response(
status: http::StatusCode,
message: impl Into<String>,
) -> Response<Body> {
build_internal_control_error_response(status, message)
}
pub(super) fn oauth_refresh_failed_bad_request_response(
error_reason: impl AsRef<str>,
) -> Response<Body> {
control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Token 刷新失败:{}", error_reason.as_ref()),
)
}
pub(super) fn oauth_refresh_failed_service_unavailable_response(
error_reason: impl Into<String>,
) -> Response<Body> {
control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
format!("Token 刷新失败:{}", error_reason.into()),
)
}
pub(super) fn admin_provider_oauth_refresh_success_response(
success: RefreshSuccessContext,
) -> Response<Body> {
Json(json!({
"provider_type": success.provider_type,
"expires_at": success
.refreshed_auth_config
.get("expires_at")
.cloned()
.unwrap_or(Value::Null),
"has_refresh_token": success
.refreshed_auth_config
.get("refresh_token")
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty()),
"email": success
.refreshed_auth_config
.get("email")
.cloned()
.unwrap_or(Value::Null),
"account_state_recheck_attempted": success.account_state_recheck_attempted,
"account_state_recheck_error": success.account_state_recheck_error,
}))
.into_response()
}
@@ -1,14 +1,14 @@
use super::super::refresh::build_internal_control_error_response;
use super::super::errors::build_internal_control_error_response;
use super::super::state::{
admin_provider_oauth_template, build_provider_oauth_start_response,
generate_provider_oauth_pkce_verifier, is_fixed_provider_type_for_provider_oauth,
provider_oauth_pkce_s256, save_provider_oauth_state,
provider_oauth_pkce_s256,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
};
use crate::{AppState, GatewayError};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{
body::Body,
http,
@@ -17,10 +17,10 @@ use axum::{
};
pub(super) async fn handle_admin_provider_oauth_start_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_start_key_id(&request_context.request_path) else {
let Some(key_id) = admin_provider_oauth_start_key_id(request_context.path()) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
@@ -74,14 +74,14 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
.use_pkce
.then(generate_provider_oauth_pkce_verifier);
let code_challenge = pkce_verifier.as_deref().map(provider_oauth_pkce_s256);
let nonce = match save_provider_oauth_state(
state,
&key_id,
&provider_id,
&provider_type,
pkce_verifier.as_deref(),
)
.await
let nonce = match state
.save_provider_oauth_state(
&key_id,
&provider_id,
&provider_type,
pkce_verifier.as_deref(),
)
.await
{
Ok(nonce) => nonce,
Err(_) => {
@@ -101,11 +101,10 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
}
pub(super) async fn handle_admin_provider_oauth_start_provider(
state: &AppState,
request_context: &GatewayPublicRequestContext,
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) = admin_provider_oauth_start_provider_id(&request_context.request_path)
else {
let Some(provider_id) = admin_provider_oauth_start_provider_id(request_context.path()) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
@@ -146,14 +145,9 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
.use_pkce
.then(generate_provider_oauth_pkce_verifier);
let code_challenge = pkce_verifier.as_deref().map(provider_oauth_pkce_s256);
let nonce = match save_provider_oauth_state(
state,
"",
&provider_id,
&provider_type,
pkce_verifier.as_deref(),
)
.await
let nonce = match state
.save_provider_oauth_state("", &provider_id, &provider_type, pkce_verifier.as_deref())
.await
{
Ok(nonce) => nonce,
Err(_) => {
@@ -1,9 +1,8 @@
use super::super::refresh::build_internal_control_error_response;
use super::super::state::read_provider_oauth_batch_task_payload;
use crate::control::GatewayPublicRequestContext;
use super::super::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_path;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use crate::GatewayError;
use axum::{
body::Body,
http,
@@ -12,18 +11,20 @@ use axum::{
};
pub(super) async fn handle_admin_provider_oauth_batch_import_task_status(
state: &AppState,
request_context: &GatewayPublicRequestContext,
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let Some((provider_id, task_id)) =
admin_provider_oauth_batch_import_task_path(&request_context.request_path)
admin_provider_oauth_batch_import_task_path(request_context.path())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"批量导入任务不存在",
));
};
let payload = match read_provider_oauth_batch_task_payload(state, &provider_id, &task_id).await
let payload = match state
.read_provider_oauth_batch_task_payload(&provider_id, &task_id)
.await
{
Ok(Some(payload)) => payload,
Ok(None) => {
@@ -0,0 +1,223 @@
use crate::handlers::admin::request::AdminAppState;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
fn normalize_codex_plan_group_for_provider_oauth(
plan_type: Option<&serde_json::Value>,
) -> Option<String> {
let normalized = plan_type
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_ascii_lowercase();
match normalized.as_str() {
"free" => Some("free".to_string()),
"team" | "plus" | "enterprise" => Some("team_plus_enterprise".to_string()),
_ => None,
}
}
fn normalize_provider_oauth_identity_value(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn is_codex_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|provider_type| provider_type.eq_ignore_ascii_case("codex"))
}
fn match_codex_provider_oauth_identity(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
) -> Option<bool> {
let new_provider_type = new_auth_config.get("provider_type");
let existing_provider_type = existing_auth_config.get("provider_type");
if !is_codex_provider_oauth_provider_type(new_provider_type)
&& !is_codex_provider_oauth_provider_type(existing_provider_type)
{
return None;
}
let new_account_user_id =
normalize_provider_oauth_identity_value(new_auth_config.get("account_user_id"));
let existing_account_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("account_user_id"));
if let (Some(new_account_user_id), Some(existing_account_user_id)) =
(new_account_user_id, existing_account_user_id)
{
return Some(new_account_user_id == existing_account_user_id);
}
let new_account_id = normalize_provider_oauth_identity_value(new_auth_config.get("account_id"));
let existing_account_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("account_id"));
let new_user_id = normalize_provider_oauth_identity_value(new_auth_config.get("user_id"));
let existing_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("user_id"));
let new_email = normalize_provider_oauth_identity_value(new_auth_config.get("email"));
let existing_email = normalize_provider_oauth_identity_value(existing_auth_config.get("email"));
if let (Some(new_account_id), Some(existing_account_id)) =
(new_account_id.as_deref(), existing_account_id.as_deref())
{
if new_account_id != existing_account_id {
return Some(false);
}
}
if let (
Some(new_account_id),
Some(existing_account_id),
Some(new_user_id),
Some(existing_user_id),
) = (
new_account_id.as_deref(),
existing_account_id.as_deref(),
new_user_id.as_deref(),
existing_user_id.as_deref(),
) {
return Some(new_account_id == existing_account_id && new_user_id == existing_user_id);
}
if let (
Some(new_account_id),
Some(existing_account_id),
Some(new_email),
Some(existing_email),
) = (
new_account_id.as_deref(),
existing_account_id.as_deref(),
new_email.as_deref(),
existing_email.as_deref(),
) {
return Some(new_account_id == existing_account_id && new_email == existing_email);
}
None
}
fn is_codex_cross_plan_group_non_duplicate(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
) -> bool {
let new_provider_type = new_auth_config.get("provider_type");
let existing_provider_type = existing_auth_config.get("provider_type");
if !is_codex_provider_oauth_provider_type(new_provider_type)
&& !is_codex_provider_oauth_provider_type(existing_provider_type)
{
return false;
}
let new_group = normalize_codex_plan_group_for_provider_oauth(new_auth_config.get("plan_type"));
let existing_group =
normalize_codex_plan_group_for_provider_oauth(existing_auth_config.get("plan_type"));
matches!(
(new_group.as_deref(), existing_group.as_deref()),
(Some(left), Some(right)) if left != right
)
}
pub(crate) async fn find_duplicate_provider_oauth_key(
state: &AdminAppState<'_>,
provider_id: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
exclude_key_id: Option<&str>,
) -> Result<Option<StoredProviderCatalogKey>, String> {
let new_email = normalize_provider_oauth_identity_value(auth_config.get("email"));
let new_user_id = normalize_provider_oauth_identity_value(auth_config.get("user_id"));
let new_auth_method = normalize_provider_oauth_identity_value(auth_config.get("auth_method"));
if new_email.is_none() && new_user_id.is_none() {
return Ok(None);
}
let existing_keys = state
.list_provider_catalog_keys_by_provider_ids(&[provider_id.to_string()])
.await
.map_err(|err| format!("{err:?}"))?;
for existing_key in existing_keys.into_iter().filter(|key| {
key.auth_type.trim().eq_ignore_ascii_case("oauth")
&& exclude_key_id.is_none_or(|exclude| key.id != exclude)
}) {
let Some(existing_auth_config) = state.parse_catalog_auth_config_json(&existing_key) else {
continue;
};
let existing_email =
normalize_provider_oauth_identity_value(existing_auth_config.get("email"));
let existing_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("user_id"));
let existing_auth_method =
normalize_provider_oauth_identity_value(existing_auth_config.get("auth_method"));
let mut is_duplicate = false;
let codex_identity_match =
match_codex_provider_oauth_identity(auth_config, &existing_auth_config);
if let Some(codex_identity_match) = codex_identity_match {
is_duplicate = codex_identity_match;
}
if codex_identity_match.is_none()
&& !is_duplicate
&& new_user_id.is_some()
&& existing_user_id.is_some()
&& new_user_id == existing_user_id
&& !is_codex_cross_plan_group_non_duplicate(auth_config, &existing_auth_config)
{
is_duplicate = true;
}
if codex_identity_match.is_none()
&& !is_duplicate
&& new_email.is_some()
&& existing_email.is_some()
&& new_email == existing_email
{
let is_kiro = auth_config
.get("provider_type")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| value.eq_ignore_ascii_case("kiro"))
|| existing_auth_config
.get("provider_type")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| value.eq_ignore_ascii_case("kiro"));
if is_kiro {
if new_auth_method.is_some()
&& existing_auth_method.is_some()
&& new_auth_method
.as_deref()
.zip(existing_auth_method.as_deref())
.is_some_and(|(left, right)| left.eq_ignore_ascii_case(right))
{
is_duplicate = true;
}
} else if !is_codex_cross_plan_group_non_duplicate(auth_config, &existing_auth_config) {
is_duplicate = true;
}
}
if !is_duplicate {
continue;
}
if !existing_key.is_active {
return Ok(Some(existing_key));
}
let identifier =
normalize_provider_oauth_identity_value(auth_config.get("account_user_id"))
.or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id")))
.or_else(|| new_email.clone())
.or_else(|| new_user_id.clone())
.unwrap_or_default();
return Err(format!(
"该 OAuth 账号 ({identifier}) 已存在于当前 Provider 中(名称: {})",
existing_key.name
));
}
Ok(None)
}
@@ -0,0 +1,145 @@
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(crate) fn build_internal_control_error_response(
status: http::StatusCode,
message: impl Into<String>,
) -> Response<Body> {
(status, Json(json!({ "detail": message.into() }))).into_response()
}
pub(crate) fn normalize_provider_oauth_refresh_error_message(
status_code: Option<u16>,
body_excerpt: Option<&str>,
) -> String {
let mut message = None::<String>;
let mut error_code = None::<String>;
let mut error_type = None::<String>;
if let Some(body_excerpt) = body_excerpt {
if let Ok(value) = serde_json::from_str::<serde_json::Value>(body_excerpt) {
if let Some(object) = value.as_object() {
if let Some(error_object) =
object.get("error").and_then(serde_json::Value::as_object)
{
message = error_object
.get("message")
.or_else(|| error_object.get("error_description"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
error_code = error_object
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
error_type = error_object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if message.is_none() {
message = object
.get("message")
.or_else(|| object.get("error_description"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
}
if error_code.is_none() {
error_code = object
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if error_type.is_none() {
error_type = object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
}
}
}
let message = message
.or_else(|| {
body_excerpt
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.chars().take(300).collect::<String>())
})
.unwrap_or_default();
let lowered = message.to_ascii_lowercase();
let error_code = error_code.unwrap_or_default();
let error_type = error_type.unwrap_or_default();
if error_code == "refresh_token_reused"
|| lowered.contains("already been used to generate a new access token")
{
return "refresh_token 已被使用并轮换,请重新登录授权".to_string();
}
if error_code == "invalid_grant"
|| error_code == "invalid_refresh_token"
|| (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid"]
.iter()
.any(|keyword| lowered.contains(keyword)))
{
return "refresh_token 无效、已过期或已撤销,请重新登录授权".to_string();
}
if error_type == "invalid_request_error" && !message.is_empty() {
return message;
}
if !message.is_empty() {
return message;
}
status_code
.map(|status_code| format!("HTTP {status_code}"))
.unwrap_or_else(|| "未知错误".to_string())
}
pub(crate) fn merge_provider_oauth_refresh_failure_reason(
current_reason: Option<&str>,
refresh_reason: &str,
) -> Option<String> {
let current_reason = current_reason.map(str::trim).unwrap_or_default();
let refresh_reason = refresh_reason.trim();
if refresh_reason.is_empty() {
return (!current_reason.is_empty()).then(|| current_reason.to_string());
}
if current_reason.is_empty() {
return Some(refresh_reason.to_string());
}
if current_reason.starts_with(OAUTH_EXPIRED_PREFIX) {
return None;
}
if current_reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX) {
if let Some((head, _)) = current_reason.split_once("[REFRESH_FAILED]") {
return Some(
format!("{}\n{}", head.trim_end(), refresh_reason)
.trim()
.to_string(),
);
}
return Some(format!("{current_reason}\n{refresh_reason}"));
}
Some(refresh_reason.to_string())
}
@@ -1,6 +1,9 @@
mod dispatch;
pub(crate) mod duplicates;
pub(crate) mod errors;
pub(crate) mod provisioning;
pub(crate) mod quota;
pub(crate) mod refresh;
pub(crate) mod runtime;
pub(crate) mod state;
pub(crate) use self::dispatch::maybe_build_local_admin_provider_oauth_response;
@@ -0,0 +1,174 @@
use super::state::{
enrich_admin_provider_oauth_auth_config, json_non_empty_string, json_u64_value,
};
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
use serde_json::json;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
pub(crate) fn provider_oauth_key_proxy_value(
proxy_node_id: Option<&str>,
) -> Option<serde_json::Value> {
proxy_node_id
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| json!({ "node_id": value, "enabled": true }))
}
pub(crate) fn provider_oauth_active_api_formats(
endpoints: &[StoredProviderCatalogEndpoint],
) -> Vec<String> {
let mut formats = Vec::new();
let mut seen = BTreeSet::new();
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
let api_format = endpoint.api_format.trim();
if api_format.is_empty() || !seen.insert(api_format.to_string()) {
continue;
}
formats.push(api_format.to_string());
}
formats
}
pub(crate) fn build_provider_oauth_auth_config_from_token_payload(
provider_type: &str,
token_payload: &serde_json::Value,
) -> (
serde_json::Map<String, serde_json::Value>,
Option<String>,
Option<String>,
Option<u64>,
) {
let access_token = json_non_empty_string(token_payload.get("access_token"));
let refresh_token = json_non_empty_string(token_payload.get("refresh_token"));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let expires_at = json_u64_value(token_payload.get("expires_in"))
.map(|expires_in| now_unix_secs.saturating_add(expires_in));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!(provider_type));
auth_config.insert("updated_at".to_string(), json!(now_unix_secs));
if let Some(token_type) = token_payload.get("token_type").cloned() {
auth_config.insert("token_type".to_string(), token_type);
}
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
if let Some(expires_at) = expires_at {
auth_config.insert("expires_at".to_string(), json!(expires_at));
}
if let Some(scope) = token_payload.get("scope").cloned() {
auth_config.insert("scope".to_string(), scope);
}
enrich_admin_provider_oauth_auth_config(provider_type, &mut auth_config, token_payload);
(auth_config, access_token, refresh_token, expires_at)
}
pub(crate) async fn create_provider_oauth_catalog_key(
state: &AdminAppState<'_>,
provider_id: &str,
name: &str,
access_token: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
api_formats: &[String],
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(access_token) else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
state.encrypt_catalog_secret_with_fallbacks(&auth_config_json)
else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut record = StoredProviderCatalogKey::new(
Uuid::new_v4().to_string(),
provider_id.to_string(),
name.to_string(),
"oauth".to_string(),
None,
true,
)
.map_err(|err| GatewayError::Internal(err.to_string()))?
.with_transport_fields(
Some(json!(api_formats)),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
expires_at_unix_secs,
proxy,
None,
)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
record.internal_priority = 50;
record.cache_ttl_minutes = 5;
record.max_probe_interval_minutes = 32;
record.request_count = Some(0);
record.success_count = Some(0);
record.error_count = Some(0);
record.total_response_time_ms = Some(0);
record.health_by_format = Some(json!({}));
record.circuit_breaker_by_format = Some(json!({}));
record.created_at_unix_secs = Some(now_unix_secs);
record.updated_at_unix_secs = Some(now_unix_secs);
state.create_provider_catalog_key(&record).await
}
pub(crate) async fn update_existing_provider_oauth_catalog_key(
state: &AdminAppState<'_>,
existing_key: &StoredProviderCatalogKey,
access_token: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(access_token) else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
state.encrypt_catalog_secret_with_fallbacks(&auth_config_json)
else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut updated = existing_key.clone();
updated.encrypted_api_key = encrypted_api_key;
updated.encrypted_auth_config = Some(encrypted_auth_config);
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;
updated.health_by_format = Some(json!({}));
updated.circuit_breaker_by_format = Some(json!({}));
updated.error_count = Some(0);
if let Some(proxy) = proxy {
updated.proxy = Some(proxy);
}
updated.updated_at_unix_secs = Some(now_unix_secs);
state.update_provider_catalog_key(&updated).await
}
@@ -4,86 +4,25 @@ use super::shared::{
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::payloads::ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH;
use crate::{AppState, GatewayError};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_antigravity_usage_response;
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
fn parse_antigravity_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let models = value.get("models")?.as_object()?;
let mut quota_by_model = serde_json::Map::new();
for (model_id, model_value) in models {
let mut payload = serde_json::Map::new();
if let Some(display_name) = coerce_json_string(
model_value
.get("displayName")
.or_else(|| model_value.get("display_name")),
) {
payload.insert("display_name".to_string(), json!(display_name));
}
let quota_info = model_value
.get("quotaInfo")
.and_then(serde_json::Value::as_object);
let remaining_fraction = quota_info
.and_then(|object| object.get("remainingFraction"))
.and_then(coerce_json_f64);
let used_percent = remaining_fraction
.map(|value| ((1.0 - value).max(0.0) * 100.0).min(100.0))
.unwrap_or(100.0);
payload.insert(
"remaining_fraction".to_string(),
json!(remaining_fraction.unwrap_or(0.0)),
);
payload.insert("used_percent".to_string(), json!(used_percent));
if let Some(reset_time) = quota_info
.and_then(|object| object.get("resetTime"))
.cloned()
.filter(|value| !value.is_null())
{
payload.insert("reset_time".to_string(), reset_time);
}
quota_by_model.insert(model_id.clone(), serde_json::Value::Object(payload));
}
Some(json!({
"updated_at": updated_at_unix_secs,
"is_forbidden": false,
"forbidden_reason": serde_json::Value::Null,
"forbidden_at": serde_json::Value::Null,
"models": quota_by_model,
}))
}
async fn execute_antigravity_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
authorization: (String, String),
project_id: &str,
auth: &crate::provider_transport::antigravity::AntigravityRequestAuthSupport,
mut identity_headers: BTreeMap<String, String>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let supported_auth = match auth {
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Supported(auth) => {
auth
}
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Unsupported(_) => {
return Ok(ProviderQuotaExecutionOutcome::Failure(
"缺少 OAuth 认证信息,请先授权/刷新 Token".to_string(),
));
}
};
let mut headers =
crate::provider_transport::antigravity::build_antigravity_static_identity_headers(
supported_auth,
);
let mut headers = std::mem::take(&mut identity_headers);
headers.insert("authorization".to_string(), authorization.1);
headers.insert("content-type".to_string(), "application/json".to_string());
headers.insert("accept".to_string(), "application/json".to_string());
@@ -117,28 +56,27 @@ async fn execute_antigravity_quota_plan(
client_api_format: "gemini:chat".to_string(),
provider_api_format: "antigravity:fetch_available_models".to_string(),
model_name: Some("fetchAvailableModels".to_string()),
proxy: crate::provider_transport::resolve_transport_proxy_snapshot_with_tunnel_affinity(
state, transport,
)
.await,
tls_profile: crate::provider_transport::resolve_transport_tls_profile(transport),
timeouts: crate::provider_transport::resolve_transport_execution_timeouts(transport).or(
Some(ExecutionTimeouts {
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await,
tls_profile: state.resolve_transport_tls_profile(transport),
timeouts: state
.resolve_transport_execution_timeouts(transport)
.or(Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
}),
),
})),
};
execute_provider_quota_plan(state, transport, plan, "antigravity").await
}
pub(crate) async fn refresh_antigravity_provider_quota_locally(
state: &AppState,
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
@@ -165,11 +103,8 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
}
};
let authorization = match state.resolve_local_oauth_request_auth(&transport).await? {
Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Header {
name,
value,
}) => (name, value),
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
Some(auth) => auth,
_ => {
failed_count += 1;
results.push(json!({
@@ -182,26 +117,17 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
}
};
let antigravity_auth =
crate::provider_transport::antigravity::resolve_local_antigravity_request_auth(
&transport,
);
let project_id = match &antigravity_auth {
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Supported(
auth,
) => auth.project_id.clone(),
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Unsupported(
_,
) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}));
continue;
}
let Some((project_id, identity_headers)) =
state.resolve_local_antigravity_identity_headers(&transport)
else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}));
continue;
};
let result = match execute_antigravity_quota_plan(
@@ -209,7 +135,7 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
&transport,
authorization,
&project_id,
&antigravity_auth,
identity_headers,
)
.await?
{
@@ -1,696 +0,0 @@
use super::shared::{
coerce_json_bool, coerce_json_f64, coerce_json_string, coerce_json_u64,
execute_provider_quota_plan, extract_execution_error_message, normalize_string_id_list,
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys,
quota_refresh_success_invalid_state, should_auto_remove_structured_reason,
ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::payloads::{
CODEX_WHAM_USAGE_URL, OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
OAUTH_REQUEST_FAILED_PREFIX,
};
use crate::{AppState, GatewayError};
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::collections::{BTreeMap, BTreeSet};
use std::time::{SystemTime, UNIX_EPOCH};
fn normalize_codex_plan_type(value: Option<&str>) -> Option<String> {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase())
}
fn build_codex_quota_exhausted_fallback_metadata(
plan_type: Option<&str>,
updated_at_unix_secs: u64,
) -> serde_json::Value {
let mut object = serde_json::Map::new();
if let Some(plan_type) = normalize_codex_plan_type(plan_type) {
object.insert(
"plan_type".to_string(),
serde_json::Value::String(plan_type),
);
}
object.insert("updated_at".to_string(), json!(updated_at_unix_secs));
object.insert("primary_used_percent".to_string(), json!(100.0));
if normalize_codex_plan_type(plan_type) != Some("free".to_string()) {
object.insert("secondary_used_percent".to_string(), json!(100.0));
}
serde_json::Value::Object(object)
}
fn codex_write_window(
target: &mut serde_json::Map<String, serde_json::Value>,
source: &serde_json::Map<String, serde_json::Value>,
target_prefix: &str,
) {
if let Some(value) = source.get("used_percent").and_then(coerce_json_f64) {
target.insert(format!("{target_prefix}_used_percent"), json!(value));
}
if let Some(value) = source.get("reset_after_seconds").and_then(coerce_json_u64) {
target.insert(format!("{target_prefix}_reset_after_seconds"), json!(value));
}
if let Some(value) = source.get("reset_at").and_then(coerce_json_u64) {
target.insert(format!("{target_prefix}_reset_at"), json!(value));
}
if let Some(value) = source.get("window_minutes").and_then(coerce_json_u64) {
target.insert(format!("{target_prefix}_window_minutes"), json!(value));
}
}
fn parse_codex_wham_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let root = value.as_object()?;
if root.is_empty() {
return None;
}
let mut result = serde_json::Map::new();
let plan_type =
normalize_codex_plan_type(root.get("plan_type").and_then(serde_json::Value::as_str));
if let Some(plan_type) = plan_type.as_ref() {
result.insert("plan_type".to_string(), json!(plan_type));
}
let rate_limit = root
.get("rate_limit")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let primary_window = rate_limit
.get("primary_window")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let secondary_window = rate_limit
.get("secondary_window")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let use_paid_windows = !secondary_window.is_empty() && plan_type.as_deref() != Some("free");
if use_paid_windows {
codex_write_window(&mut result, &secondary_window, "primary");
codex_write_window(&mut result, &primary_window, "secondary");
} else {
codex_write_window(&mut result, &primary_window, "primary");
}
if let Some(credits) = root.get("credits").and_then(serde_json::Value::as_object) {
if let Some(value) = credits.get("has_credits").and_then(coerce_json_bool) {
result.insert("has_credits".to_string(), json!(value));
}
if let Some(value) = credits.get("balance").and_then(coerce_json_f64) {
result.insert("credits_balance".to_string(), json!(value));
}
if let Some(value) = credits.get("unlimited").and_then(coerce_json_bool) {
result.insert("credits_unlimited".to_string(), json!(value));
}
}
if result.is_empty() {
return None;
}
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
Some(serde_json::Value::Object(result))
}
fn parse_codex_usage_headers(
headers: &BTreeMap<String, String>,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let mut result = serde_json::Map::new();
let normalized = headers
.iter()
.map(|(key, value)| (key.trim().to_ascii_lowercase(), value.trim().to_string()))
.collect::<BTreeMap<_, _>>();
if !normalized.keys().any(|key| key.starts_with("x-codex-")) {
return None;
}
let plan_type =
normalize_codex_plan_type(normalized.get("x-codex-plan-type").map(String::as_str));
if let Some(plan_type) = plan_type.as_ref() {
result.insert("plan_type".to_string(), json!(plan_type));
}
let read_window = |prefix: &str| -> serde_json::Map<String, serde_json::Value> {
let mut object = serde_json::Map::new();
let used_key = format!("x-codex-{prefix}-used-percent");
let reset_after_key = format!("x-codex-{prefix}-reset-after-seconds");
let reset_at_key = format!("x-codex-{prefix}-reset-at");
let window_minutes_key = format!("x-codex-{prefix}-window-minutes");
if let Some(value) = normalized
.get(&used_key)
.and_then(|value| value.parse::<f64>().ok())
{
object.insert("used_percent".to_string(), json!(value));
}
if let Some(value) = normalized
.get(&reset_after_key)
.and_then(|value| value.parse::<u64>().ok())
{
object.insert("reset_after_seconds".to_string(), json!(value));
}
if let Some(value) = normalized
.get(&reset_at_key)
.and_then(|value| value.parse::<u64>().ok())
{
object.insert("reset_at".to_string(), json!(value));
}
if let Some(value) = normalized
.get(&window_minutes_key)
.and_then(|value| value.parse::<u64>().ok())
{
object.insert("window_minutes".to_string(), json!(value));
}
object
};
let primary_window = read_window("primary");
let secondary_window = read_window("secondary");
let use_paid_windows = !secondary_window.is_empty() && plan_type.as_deref() != Some("free");
if use_paid_windows {
codex_write_window(&mut result, &secondary_window, "primary");
codex_write_window(&mut result, &primary_window, "secondary");
} else {
codex_write_window(&mut result, &primary_window, "primary");
}
if let Some(value) = normalized
.get("x-codex-primary-over-secondary-limit-percent")
.and_then(|value| value.parse::<f64>().ok())
{
result.insert(
"primary_over_secondary_limit_percent".to_string(),
json!(value),
);
}
if let Some(value) = normalized
.get("x-codex-credits-has-credits")
.and_then(|value| match value.to_ascii_lowercase().as_str() {
"true" | "1" => Some(true),
"false" | "0" => Some(false),
_ => None,
})
{
result.insert("has_credits".to_string(), json!(value));
}
if let Some(value) = normalized
.get("x-codex-credits-balance")
.and_then(|value| value.parse::<f64>().ok())
{
result.insert("credits_balance".to_string(), json!(value));
}
if let Some(value) = normalized
.get("x-codex-credits-unlimited")
.and_then(|value| match value.to_ascii_lowercase().as_str() {
"true" | "1" => Some(true),
"false" | "0" => Some(false),
_ => None,
})
{
result.insert("credits_unlimited".to_string(), json!(value));
}
if result.is_empty() {
return None;
}
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
Some(serde_json::Value::Object(result))
}
fn codex_current_invalid_reason(key: &StoredProviderCatalogKey) -> String {
key.oauth_invalid_reason
.as_deref()
.map(str::trim)
.unwrap_or_default()
.to_string()
}
fn codex_merge_invalid_reason(current: &str, candidate_reason: &str) -> String {
if current.is_empty() {
return candidate_reason.to_string();
}
if current.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX) {
return current.to_string();
}
if current.starts_with(OAUTH_EXPIRED_PREFIX)
&& candidate_reason.starts_with(OAUTH_REQUEST_FAILED_PREFIX)
{
return current.to_string();
}
candidate_reason.to_string()
}
fn codex_build_invalid_state(
key: &StoredProviderCatalogKey,
candidate_reason: String,
now_unix_secs: u64,
) -> (Option<u64>, Option<String>) {
let current_reason = codex_current_invalid_reason(key);
let merged_reason = codex_merge_invalid_reason(&current_reason, &candidate_reason);
if merged_reason == current_reason {
return (key.oauth_invalid_at_unix_secs, Some(merged_reason));
}
(Some(now_unix_secs), Some(merged_reason))
}
fn codex_looks_like_token_invalidated(message: Option<&str>) -> bool {
let lowered = message.unwrap_or_default().trim().to_ascii_lowercase();
lowered.contains("token invalid")
|| lowered.contains("token invalidated")
|| lowered.contains("session has expired")
|| lowered.contains("session expired")
}
fn codex_looks_like_account_deactivated(message: Option<&str>) -> bool {
let lowered = message.unwrap_or_default().trim().to_ascii_lowercase();
lowered.contains("account has been deactivated") || lowered.contains("account deactivated")
}
fn codex_looks_like_workspace_deactivated(message: Option<&str>) -> bool {
let lowered = message.unwrap_or_default().trim().to_ascii_lowercase();
lowered.contains("deactivated_workspace")
|| (lowered.contains("workspace") && lowered.contains("deactivated"))
}
fn codex_structured_invalid_reason(status_code: u16, upstream_message: Option<&str>) -> String {
let message = upstream_message.unwrap_or_default().trim();
if status_code == 402 && codex_looks_like_workspace_deactivated(Some(message)) {
return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}工作区已停用 (deactivated_workspace)");
}
if codex_looks_like_account_deactivated(Some(message)) {
let detail = if message.is_empty() {
"OpenAI 账号已停用"
} else {
message
};
return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}");
}
if codex_looks_like_token_invalidated(Some(message)) {
let detail = if message.is_empty() {
"Codex Token 无效或已过期"
} else {
message
};
return format!("{OAUTH_EXPIRED_PREFIX}{detail}");
}
if status_code == 401 {
let detail = if message.is_empty() {
"Codex Token 无效或已过期 (401)"
} else {
message
};
return format!("{OAUTH_EXPIRED_PREFIX}{detail}");
}
if status_code == 403 {
let detail = if message.is_empty() {
"Codex 账户访问受限 (403)"
} else {
message
};
return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}");
}
message.to_string()
}
fn codex_soft_request_failure_reason(status_code: u16, upstream_message: Option<&str>) -> String {
let detail = upstream_message
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.unwrap_or_else(|| format!("Codex 请求失败 ({status_code})"));
format!("{OAUTH_REQUEST_FAILED_PREFIX}{detail}")
}
fn build_codex_refresh_headers(
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
resolved_oauth_auth: Option<(String, String)>,
) -> Result<BTreeMap<String, String>, String> {
let mut headers = BTreeMap::new();
headers.insert("accept".to_string(), "application/json".to_string());
if let Some((name, value)) = resolved_oauth_auth {
headers.insert(name.to_ascii_lowercase(), value);
} else {
let decrypted_key = transport.key.decrypted_api_key.trim();
if decrypted_key.is_empty() || decrypted_key == "__placeholder__" {
return Err("缺少 OAuth 认证信息,请先授权/刷新 Token".to_string());
}
headers.insert(
"authorization".to_string(),
format!("Bearer {decrypted_key}"),
);
}
let auth_config = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok());
let oauth_plan_type = normalize_codex_plan_type(
auth_config
.as_ref()
.and_then(|value| value.get("plan_type"))
.and_then(serde_json::Value::as_str),
);
let oauth_account_id = auth_config
.as_ref()
.and_then(|value| value.get("account_id"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if oauth_account_id.is_some() && oauth_plan_type.as_deref() != Some("free") {
headers.insert(
"chatgpt-account-id".to_string(),
oauth_account_id.unwrap_or_default().to_string(),
);
}
Ok(headers)
}
async fn execute_codex_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
headers: BTreeMap<String, String>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let plan = ExecutionPlan {
request_id: format!("codex-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("codex".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "GET".to_string(),
url: CODEX_WHAM_USAGE_URL.to_string(),
headers,
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: false,
client_api_format: "openai:cli".to_string(),
provider_api_format: "openai:cli".to_string(),
model_name: Some("codex-wham-usage".to_string()),
proxy: crate::provider_transport::resolve_transport_proxy_snapshot_with_tunnel_affinity(
state, transport,
)
.await,
tls_profile: crate::provider_transport::resolve_transport_tls_profile(transport),
timeouts: crate::provider_transport::resolve_transport_execution_timeouts(transport).or(
Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
}),
),
};
execute_provider_quota_plan(state, transport, plan, "codex").await
}
pub(crate) async fn refresh_codex_provider_quota_locally(
state: &AppState,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let auto_remove_abnormal_keys = provider_auto_remove_banned_keys(provider.config.as_ref());
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_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;
}
};
let resolved_oauth_auth = if key.auth_type.trim().eq_ignore_ascii_case("oauth") {
match state.resolve_local_oauth_request_auth(&transport).await? {
Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Header {
name,
value,
}) => Some((name, value)),
_ => None,
}
} else {
None
};
let headers = match build_codex_refresh_headers(&transport, resolved_oauth_auth) {
Ok(headers) => headers,
Err(message) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": message,
}));
continue;
}
};
let result = match execute_codex_quota_plan(state, &transport, headers).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("wham/usage 请求执行失败: {detail}"),
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = parse_codex_usage_headers(&result.headers, now_unix_secs)
.map(|metadata| json!({ "codex": metadata }));
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = (None, None);
let mut status = "error".to_string();
let mut message = None::<String>;
let mut 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(parsed) = parse_codex_wham_usage_response(body_json, now_unix_secs) {
metadata_update = Some(json!({ "codex": parsed }));
(oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
message = Some("无法解析 wham/usage API 响应".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"wham/usage API 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("wham/usage API 返回状态码 {}", result.status_code),
});
match result.status_code {
401 => {
let (at, reason) = codex_build_invalid_state(
&key,
codex_structured_invalid_reason(401, err_msg.as_deref()),
now_unix_secs,
);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "auth_invalid".to_string();
}
402 => {
if codex_looks_like_workspace_deactivated(err_msg.as_deref()) {
let mut codex_meta = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
codex_meta.insert("updated_at".to_string(), json!(now_unix_secs));
codex_meta.insert("account_disabled".to_string(), json!(true));
codex_meta.insert("reason".to_string(), json!("deactivated_workspace"));
codex_meta.insert(
"message".to_string(),
json!(err_msg
.clone()
.unwrap_or_else(|| "deactivated_workspace".to_string())),
);
let plan_type = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok())
.and_then(|value| {
value
.get("plan_type")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
});
if let Some(plan_type) = plan_type {
codex_meta
.entry("plan_type".to_string())
.or_insert_with(|| json!(plan_type.to_ascii_lowercase()));
}
metadata_update = Some(json!({ "codex": codex_meta }));
let (at, reason) = codex_build_invalid_state(
&key,
codex_structured_invalid_reason(402, err_msg.as_deref()),
now_unix_secs,
);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "workspace_deactivated".to_string();
} else {
let plan_type = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok())
.and_then(|value| {
value
.get("plan_type")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
});
metadata_update = Some(json!({
"codex": build_codex_quota_exhausted_fallback_metadata(
plan_type.as_deref(),
now_unix_secs,
)
}));
(oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
status = "quota_exhausted".to_string();
}
}
403 => {
let candidate_reason = if codex_looks_like_token_invalidated(err_msg.as_deref())
{
codex_structured_invalid_reason(403, err_msg.as_deref())
} else {
codex_soft_request_failure_reason(403, err_msg.as_deref())
};
let (at, reason) =
codex_build_invalid_state(&key, candidate_reason, now_unix_secs);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "forbidden".to_string();
}
_ => {}
}
}
let auto_removed = 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(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(),
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
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!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(status_code) = status_code.take() {
if status_code != 200 {
payload.insert("status_code".to_string(), json!(status_code));
}
}
if let Some(metadata_update) = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
.cloned()
{
payload.insert("metadata".to_string(), metadata_update);
}
if auto_removed {
payload.insert("auto_removed".to_string(), json!(true));
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": results.len(),
"results": results,
"auto_removed": auto_removed_count,
})))
}
@@ -0,0 +1,32 @@
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
pub(super) fn codex_build_invalid_state(
key: &StoredProviderCatalogKey,
candidate_reason: String,
now_unix_secs: u64,
) -> (Option<u64>, Option<String>) {
admin_provider_quota_pure::codex_build_invalid_state(key, candidate_reason, now_unix_secs)
}
pub(super) fn codex_looks_like_token_invalidated(message: Option<&str>) -> bool {
admin_provider_quota_pure::codex_looks_like_token_invalidated(message)
}
pub(super) fn codex_looks_like_workspace_deactivated(message: Option<&str>) -> bool {
admin_provider_quota_pure::codex_looks_like_workspace_deactivated(message)
}
pub(super) fn codex_structured_invalid_reason(
status_code: u16,
upstream_message: Option<&str>,
) -> String {
admin_provider_quota_pure::codex_structured_invalid_reason(status_code, upstream_message)
}
pub(super) fn codex_soft_request_failure_reason(
status_code: u16,
upstream_message: Option<&str>,
) -> String {
admin_provider_quota_pure::codex_soft_request_failure_reason(status_code, upstream_message)
}
@@ -0,0 +1,292 @@
mod invalid;
mod parse;
mod plan;
use self::invalid::{
codex_build_invalid_state, codex_looks_like_token_invalidated,
codex_looks_like_workspace_deactivated, codex_soft_request_failure_reason,
codex_structured_invalid_reason,
};
use self::parse::{
build_codex_quota_exhausted_fallback_metadata, parse_codex_usage_headers,
parse_codex_wham_usage_response,
};
use self::plan::{build_codex_refresh_headers, execute_codex_quota_plan};
use super::shared::{
extract_execution_error_message, persist_provider_quota_refresh_state,
provider_auto_remove_banned_keys, quota_refresh_success_invalid_state,
should_auto_remove_structured_reason, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(crate) async fn refresh_codex_provider_quota_locally(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let auto_remove_abnormal_keys = provider_auto_remove_banned_keys(provider.config.as_ref());
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_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;
}
};
let resolved_oauth_auth = if key.auth_type.trim().eq_ignore_ascii_case("oauth") {
state.resolve_local_oauth_header_auth(&transport).await?
} else {
None
};
let headers = match build_codex_refresh_headers(&transport, resolved_oauth_auth) {
Ok(headers) => headers,
Err(message) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": message,
}));
continue;
}
};
let result = match execute_codex_quota_plan(state, &transport, headers).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("wham/usage 请求执行失败: {detail}"),
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = parse_codex_usage_headers(&result.headers, now_unix_secs)
.map(|metadata| json!({ "codex": metadata }));
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = (None, None);
let mut status = "error".to_string();
let mut message = None::<String>;
let mut 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(parsed) = parse_codex_wham_usage_response(body_json, now_unix_secs) {
metadata_update = Some(json!({ "codex": parsed }));
(oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
message = Some("无法解析 wham/usage API 响应".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"wham/usage API 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("wham/usage API 返回状态码 {}", result.status_code),
});
match result.status_code {
401 => {
let (at, reason) = codex_build_invalid_state(
&key,
codex_structured_invalid_reason(401, err_msg.as_deref()),
now_unix_secs,
);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "auth_invalid".to_string();
}
402 => {
if codex_looks_like_workspace_deactivated(err_msg.as_deref()) {
let mut codex_meta = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
codex_meta.insert("updated_at".to_string(), json!(now_unix_secs));
codex_meta.insert("account_disabled".to_string(), json!(true));
codex_meta.insert("reason".to_string(), json!("deactivated_workspace"));
codex_meta.insert(
"message".to_string(),
json!(err_msg
.clone()
.unwrap_or_else(|| "deactivated_workspace".to_string())),
);
let plan_type = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok())
.and_then(|value| {
value
.get("plan_type")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
});
if let Some(plan_type) = plan_type {
codex_meta
.entry("plan_type".to_string())
.or_insert_with(|| json!(plan_type.to_ascii_lowercase()));
}
metadata_update = Some(json!({ "codex": codex_meta }));
let (at, reason) = codex_build_invalid_state(
&key,
codex_structured_invalid_reason(402, err_msg.as_deref()),
now_unix_secs,
);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "workspace_deactivated".to_string();
} else {
let plan_type = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok())
.and_then(|value| {
value
.get("plan_type")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
});
metadata_update = Some(json!({
"codex": build_codex_quota_exhausted_fallback_metadata(
plan_type.as_deref(),
now_unix_secs,
)
}));
(oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
status = "quota_exhausted".to_string();
}
}
403 => {
let candidate_reason = if codex_looks_like_token_invalidated(err_msg.as_deref())
{
codex_structured_invalid_reason(403, err_msg.as_deref())
} else {
codex_soft_request_failure_reason(403, err_msg.as_deref())
};
let (at, reason) =
codex_build_invalid_state(&key, candidate_reason, now_unix_secs);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "forbidden".to_string();
}
_ => {}
}
}
let auto_removed = 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(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(),
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
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!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(status_code) = status_code.take() {
if status_code != 200 {
payload.insert("status_code".to_string(), json!(status_code));
}
}
if let Some(metadata_update) = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
.cloned()
{
payload.insert("metadata".to_string(), metadata_update);
}
if auto_removed {
payload.insert("auto_removed".to_string(), json!(true));
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": results.len(),
"results": results,
"auto_removed": auto_removed_count,
})))
}
@@ -0,0 +1,30 @@
use aether_admin::provider::quota as admin_provider_quota_pure;
use std::collections::BTreeMap;
pub(super) fn normalize_codex_plan_type(value: Option<&str>) -> Option<String> {
admin_provider_quota_pure::normalize_codex_plan_type(value)
}
pub(super) fn build_codex_quota_exhausted_fallback_metadata(
plan_type: Option<&str>,
updated_at_unix_secs: u64,
) -> serde_json::Value {
admin_provider_quota_pure::build_codex_quota_exhausted_fallback_metadata(
plan_type,
updated_at_unix_secs,
)
}
pub(super) fn parse_codex_wham_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
admin_provider_quota_pure::parse_codex_wham_usage_response(value, updated_at_unix_secs)
}
pub(super) fn parse_codex_usage_headers(
headers: &BTreeMap<String, String>,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
admin_provider_quota_pure::parse_codex_usage_headers(headers, updated_at_unix_secs)
}
@@ -0,0 +1,98 @@
use super::super::shared::{execute_provider_quota_plan, ProviderQuotaExecutionOutcome};
use super::parse::normalize_codex_plan_type;
use crate::handlers::admin::provider::shared::payloads::CODEX_WHAM_USAGE_URL;
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use std::collections::BTreeMap;
pub(super) fn build_codex_refresh_headers(
transport: &AdminGatewayProviderTransportSnapshot,
resolved_oauth_auth: Option<(String, String)>,
) -> Result<BTreeMap<String, String>, String> {
let mut headers = BTreeMap::new();
headers.insert("accept".to_string(), "application/json".to_string());
if let Some((name, value)) = resolved_oauth_auth {
headers.insert(name.to_ascii_lowercase(), value);
} else {
let decrypted_key = transport.key.decrypted_api_key.trim();
if decrypted_key.is_empty() || decrypted_key == "__placeholder__" {
return Err("缺少 OAuth 认证信息,请先授权/刷新 Token".to_string());
}
headers.insert(
"authorization".to_string(),
format!("Bearer {decrypted_key}"),
);
}
let auth_config = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok());
let oauth_plan_type = normalize_codex_plan_type(
auth_config
.as_ref()
.and_then(|value| value.get("plan_type"))
.and_then(serde_json::Value::as_str),
);
let oauth_account_id = auth_config
.as_ref()
.and_then(|value| value.get("account_id"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if oauth_account_id.is_some() && oauth_plan_type.as_deref() != Some("free") {
headers.insert(
"chatgpt-account-id".to_string(),
oauth_account_id.unwrap_or_default().to_string(),
);
}
Ok(headers)
}
pub(super) async fn execute_codex_quota_plan(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
headers: BTreeMap<String, String>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let plan = ExecutionPlan {
request_id: format!("codex-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("codex".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "GET".to_string(),
url: CODEX_WHAM_USAGE_URL.to_string(),
headers,
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: false,
client_api_format: "openai:cli".to_string(),
provider_api_format: "openai:cli".to_string(),
model_name: Some("codex-wham-usage".to_string()),
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await,
tls_profile: state.resolve_transport_tls_profile(transport),
timeouts: state
.resolve_transport_execution_timeouts(transport)
.or(Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
})),
};
execute_provider_quota_plan(state, transport, plan, "codex").await
}
@@ -1,452 +0,0 @@
use super::shared::{
coerce_json_f64, 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::{
KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION,
};
use crate::handlers::admin::shared::encrypt_catalog_secret_with_fallbacks;
use crate::{AppState, GatewayError};
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use url::form_urlencoded;
use uuid::Uuid;
fn compute_kiro_total_usage_limit(breakdown: &serde_json::Value) -> f64 {
let mut total = breakdown
.get("usageLimitWithPrecision")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
if breakdown
.get("freeTrialInfo")
.and_then(serde_json::Value::as_object)
.is_some_and(|free_trial| {
free_trial
.get("freeTrialStatus")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"))
})
{
total += breakdown
.get("freeTrialInfo")
.and_then(|value| value.get("usageLimitWithPrecision"))
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
if let Some(bonuses) = breakdown
.get("bonuses")
.and_then(serde_json::Value::as_array)
{
for bonus in bonuses {
let is_active = bonus
.get("status")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"));
if is_active {
total += bonus
.get("usageLimit")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
}
}
total
}
fn compute_kiro_current_usage(breakdown: &serde_json::Value) -> f64 {
let mut total = breakdown
.get("currentUsageWithPrecision")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
if breakdown
.get("freeTrialInfo")
.and_then(serde_json::Value::as_object)
.is_some_and(|free_trial| {
free_trial
.get("freeTrialStatus")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"))
})
{
total += breakdown
.get("freeTrialInfo")
.and_then(|value| value.get("currentUsageWithPrecision"))
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
if let Some(bonuses) = breakdown
.get("bonuses")
.and_then(serde_json::Value::as_array)
{
for bonus in bonuses {
let is_active = bonus
.get("status")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"));
if is_active {
total += bonus
.get("currentUsage")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
}
}
total
}
fn parse_kiro_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let root = value.as_object()?;
let breakdown = root
.get("usageBreakdownList")
.and_then(serde_json::Value::as_array)
.and_then(|items| items.first())?;
let usage_limit = compute_kiro_total_usage_limit(breakdown);
let current_usage = compute_kiro_current_usage(breakdown);
let remaining = (usage_limit - current_usage).max(0.0);
let usage_percentage = if usage_limit > 0.0 {
((current_usage / usage_limit) * 100.0).min(100.0)
} else {
0.0
};
let mut result = serde_json::Map::new();
result.insert("current_usage".to_string(), json!(current_usage));
result.insert("usage_limit".to_string(), json!(usage_limit));
result.insert("remaining".to_string(), json!(remaining));
result.insert("usage_percentage".to_string(), json!(usage_percentage));
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
if let Some(subscription_title) = root
.get("subscriptionInfo")
.and_then(|value| value.get("subscriptionTitle"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
result.insert("subscription_title".to_string(), json!(subscription_title));
}
if let Some(next_reset_at) = root
.get("nextDateReset")
.and_then(coerce_json_f64)
.or_else(|| breakdown.get("nextDateReset").and_then(coerce_json_f64))
{
result.insert("next_reset_at".to_string(), json!(next_reset_at));
}
let email = root
.get("desktopUserInfo")
.and_then(|value| value.get("email"))
.or_else(|| root.get("userInfo").and_then(|value| value.get("email")))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if let Some(email) = email {
result.insert("email".to_string(), json!(email));
}
Some(serde_json::Value::Object(result))
}
fn build_kiro_usage_headers(
auth: &crate::provider_transport::kiro::KiroRequestAuth,
) -> BTreeMap<String, String> {
let kiro_version = auth.auth_config.effective_kiro_version();
let machine_id = auth.machine_id.trim();
let ide_tag = if machine_id.is_empty() {
format!("KiroIDE-{kiro_version}")
} else {
format!("KiroIDE-{kiro_version}-{machine_id}")
};
let host = format!(
"q.{}.amazonaws.com",
auth.auth_config.effective_api_region()
);
BTreeMap::from([
(
"x-amz-user-agent".to_string(),
format!("aws-sdk-js/{KIRO_USAGE_SDK_VERSION} {ide_tag}"),
),
(
"user-agent".to_string(),
format!(
"aws-sdk-js/{KIRO_USAGE_SDK_VERSION} ua/2.1 os/other#unknown lang/js md/nodejs#22.21.1 api/codewhispererruntime#1.0.0 m/N,E {ide_tag}"
),
),
("host".to_string(), host),
("amz-sdk-invocation-id".to_string(), Uuid::new_v4().to_string()),
("amz-sdk-request".to_string(), "attempt=1; max=1".to_string()),
("authorization".to_string(), auth.value.clone()),
("connection".to_string(), "close".to_string()),
])
}
fn build_kiro_usage_url(auth: &crate::provider_transport::kiro::KiroRequestAuth) -> String {
let host = format!(
"q.{}.amazonaws.com",
auth.auth_config.effective_api_region()
);
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("origin", "AI_EDITOR");
serializer.append_pair("resourceType", "AGENTIC_REQUEST");
serializer.append_pair("isEmailRequired", "true");
if let Some(profile_arn) = auth.auth_config.profile_arn_for_payload() {
serializer.append_pair("profileArn", profile_arn);
}
format!(
"https://{host}{KIRO_USAGE_LIMITS_PATH}?{}",
serializer.finish()
)
}
async fn execute_kiro_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
auth: &crate::provider_transport::kiro::KiroRequestAuth,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let plan = ExecutionPlan {
request_id: format!("kiro-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("kiro".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "GET".to_string(),
url: build_kiro_usage_url(auth),
headers: build_kiro_usage_headers(auth),
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: false,
client_api_format: "claude:cli".to_string(),
provider_api_format: "kiro:usage".to_string(),
model_name: Some("kiro-usage-limits".to_string()),
proxy: crate::provider_transport::resolve_transport_proxy_snapshot_with_tunnel_affinity(
state, transport,
)
.await,
tls_profile: crate::provider_transport::resolve_transport_tls_profile(transport),
timeouts: crate::provider_transport::resolve_transport_execution_timeouts(transport).or(
Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
}),
),
};
execute_provider_quota_plan(state, transport, plan, "kiro").await
}
pub(crate) async fn refresh_kiro_provider_quota_locally(
state: &AppState,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
) -> 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;
}
};
let Some(auth) = (match state.resolve_local_oauth_request_auth(&transport).await? {
Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Kiro(auth)) => {
Some(auth)
}
_ => None,
}) else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Kiro 认证配置 (auth_config)",
}));
continue;
};
let result = match execute_kiro_quota_plan(state, &transport, &auth).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("getUsageLimits 请求执行失败: {detail}"),
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = None::<serde_json::Value>;
let mut encrypted_auth_config = None::<String>;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
let mut status = "error".to_string();
let mut message = None::<String>;
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update = parse_kiro_usage_response(body_json, now_unix_secs)
.map(|metadata| json!({ "kiro": metadata }));
if metadata_update.is_some() {
let auth_config_json = auth.auth_config.to_json_value().to_string();
if let Some(auth_config_json) =
encrypt_catalog_secret_with_fallbacks(state, auth_config_json.as_str())
{
encrypted_auth_config = Some(auth_config_json);
}
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"getUsageLimits 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("getUsageLimits 返回状态码 {}", result.status_code),
});
match result.status_code {
401 => {
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string());
}
403 | 423 => {
let reason = err_msg
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户已封禁: {reason}"));
metadata_update = Some(json!({
"kiro": {
"is_banned": true,
"ban_reason": reason,
"banned_at": now_unix_secs,
"updated_at": now_unix_secs,
}
}));
status = "banned".to_string();
}
_ => {}
}
}
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
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!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("kiro"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
}
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,
})))
}
@@ -0,0 +1,199 @@
mod parse;
mod plan;
use self::parse::parse_kiro_usage_response;
use self::plan::execute_kiro_quota_plan;
use super::shared::{
extract_execution_error_message, persist_provider_quota_refresh_state,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(crate) async fn refresh_kiro_provider_quota_locally(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
) -> 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;
}
};
let Some(auth) = state
.resolve_local_oauth_kiro_request_auth(&transport)
.await?
else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Kiro 认证配置 (auth_config)",
}));
continue;
};
let result = match execute_kiro_quota_plan(state, &transport, &auth).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("getUsageLimits 请求执行失败: {detail}"),
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = None::<serde_json::Value>;
let mut encrypted_auth_config = None::<String>;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
let mut status = "error".to_string();
let mut message = None::<String>;
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update = parse_kiro_usage_response(body_json, now_unix_secs)
.map(|metadata| json!({ "kiro": metadata }));
if metadata_update.is_some() {
let auth_config_json = auth.auth_config.to_json_value().to_string();
if let Some(auth_config_json) =
state.encrypt_catalog_secret_with_fallbacks(auth_config_json.as_str())
{
encrypted_auth_config = Some(auth_config_json);
}
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"getUsageLimits 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("getUsageLimits 返回状态码 {}", result.status_code),
});
match result.status_code {
401 => {
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string());
}
403 | 423 => {
let reason = err_msg
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户已封禁: {reason}"));
metadata_update = Some(json!({
"kiro": {
"is_banned": true,
"ban_reason": reason,
"banned_at": now_unix_secs,
"updated_at": now_unix_secs,
}
}));
status = "banned".to_string();
}
_ => {}
}
}
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
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!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("kiro"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
}
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,
})))
}
@@ -0,0 +1,8 @@
use aether_admin::provider::quota as admin_provider_quota_pure;
pub(super) fn parse_kiro_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
admin_provider_quota_pure::parse_kiro_usage_response(value, updated_at_unix_secs)
}
@@ -0,0 +1,107 @@
use super::super::shared::{execute_provider_quota_plan, ProviderQuotaExecutionOutcome};
use crate::handlers::admin::provider::shared::payloads::{
KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION,
};
use crate::handlers::admin::request::{
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminKiroRequestAuth,
};
use crate::GatewayError;
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use std::collections::BTreeMap;
use url::form_urlencoded;
use uuid::Uuid;
fn build_kiro_usage_headers(auth: &AdminKiroRequestAuth) -> BTreeMap<String, String> {
let kiro_version = auth.auth_config.effective_kiro_version();
let machine_id = auth.machine_id.trim();
let ide_tag = if machine_id.is_empty() {
format!("KiroIDE-{kiro_version}")
} else {
format!("KiroIDE-{kiro_version}-{machine_id}")
};
let host = format!(
"q.{}.amazonaws.com",
auth.auth_config.effective_api_region()
);
BTreeMap::from([
(
"x-amz-user-agent".to_string(),
format!("aws-sdk-js/{KIRO_USAGE_SDK_VERSION} {ide_tag}"),
),
(
"user-agent".to_string(),
format!(
"aws-sdk-js/{KIRO_USAGE_SDK_VERSION} ua/2.1 os/other#unknown lang/js md/nodejs#22.21.1 api/codewhispererruntime#1.0.0 m/N,E {ide_tag}"
),
),
("host".to_string(), host),
("amz-sdk-invocation-id".to_string(), Uuid::new_v4().to_string()),
("amz-sdk-request".to_string(), "attempt=1; max=1".to_string()),
("authorization".to_string(), auth.value.clone()),
("connection".to_string(), "close".to_string()),
])
}
fn build_kiro_usage_url(auth: &AdminKiroRequestAuth) -> String {
let host = format!(
"q.{}.amazonaws.com",
auth.auth_config.effective_api_region()
);
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("origin", "AI_EDITOR");
serializer.append_pair("resourceType", "AGENTIC_REQUEST");
serializer.append_pair("isEmailRequired", "true");
if let Some(profile_arn) = auth.auth_config.profile_arn_for_payload() {
serializer.append_pair("profileArn", profile_arn);
}
format!(
"https://{host}{KIRO_USAGE_LIMITS_PATH}?{}",
serializer.finish()
)
}
pub(super) async fn execute_kiro_quota_plan(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
auth: &AdminKiroRequestAuth,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let plan = ExecutionPlan {
request_id: format!("kiro-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("kiro".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "GET".to_string(),
url: build_kiro_usage_url(auth),
headers: build_kiro_usage_headers(auth),
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: false,
client_api_format: "claude:cli".to_string(),
provider_api_format: "kiro:usage".to_string(),
model_name: Some("kiro-usage-limits".to_string()),
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await,
tls_profile: state.resolve_transport_tls_profile(transport),
timeouts: state
.resolve_transport_execution_timeouts(transport)
.or(Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
})),
};
execute_provider_quota_plan(state, transport, plan, "kiro").await
}
@@ -1,10 +1,11 @@
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::{AppState, GatewayError};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_contracts::{ExecutionPlan, ExecutionResult};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
@@ -14,59 +15,27 @@ pub(super) enum ProviderQuotaExecutionOutcome {
}
pub(super) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
config
.and_then(|value| value.get("pool_advanced"))
.and_then(serde_json::Value::as_object)
.and_then(|object| object.get("auto_remove_banned_keys"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
}
pub(super) fn should_auto_remove_structured_reason(reason: Option<&str>) -> bool {
reason
.map(str::trim)
.is_some_and(|value| value.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX))
admin_provider_quota_pure::should_auto_remove_structured_reason(reason)
}
pub(crate) fn normalize_string_id_list(values: Option<Vec<String>>) -> Option<Vec<String>> {
let mut out = Vec::new();
let mut seen = BTreeSet::new();
for value in values.into_iter().flatten() {
let trimmed = value.trim();
if trimmed.is_empty() || !seen.insert(trimmed.to_string()) {
continue;
}
out.push(trimmed.to_string());
}
(!out.is_empty()).then_some(out)
admin_provider_quota_pure::normalize_string_id_list(values)
}
pub(super) fn coerce_json_u64(value: &serde_json::Value) -> Option<u64> {
match value {
serde_json::Value::Number(number) => number.as_u64(),
serde_json::Value::String(text) => text.trim().parse::<u64>().ok(),
_ => None,
}
admin_provider_quota_pure::coerce_json_u64(value)
}
pub(super) fn coerce_json_f64(value: &serde_json::Value) -> Option<f64> {
match value {
serde_json::Value::Number(number) => number.as_f64(),
serde_json::Value::String(text) => text.trim().parse::<f64>().ok(),
_ => None,
}
admin_provider_quota_pure::coerce_json_f64(value)
}
pub(super) fn coerce_json_bool(value: &serde_json::Value) -> Option<bool> {
match value {
serde_json::Value::Bool(value) => Some(*value),
serde_json::Value::String(text) => match text.trim().to_ascii_lowercase().as_str() {
"true" | "1" => Some(true),
"false" | "0" => Some(false),
_ => None,
},
_ => None,
}
admin_provider_quota_pure::coerce_json_bool(value)
}
fn merge_upstream_metadata(
@@ -86,65 +55,21 @@ fn merge_upstream_metadata(
}
pub(super) fn extract_execution_error_message(result: &ExecutionResult) -> Option<String> {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(serde_json::Value::as_object)
{
if let Some(error) = body_json
.get("error")
.and_then(serde_json::Value::as_object)
{
if let Some(message) = error.get("message").and_then(serde_json::Value::as_str) {
let trimmed = message.trim();
if !trimmed.is_empty() {
return Some(trimmed.to_string());
}
}
}
if let Some(message) = body_json.get("message").and_then(serde_json::Value::as_str) {
let trimmed = message.trim();
if !trimmed.is_empty() {
return Some(trimmed.to_string());
}
}
}
result
.error
.as_ref()
.map(|error| error.message.trim().to_string())
.filter(|value| !value.is_empty())
admin_provider_quota_pure::extract_execution_error_message(result)
}
pub(super) fn quota_refresh_success_invalid_state(
key: &StoredProviderCatalogKey,
) -> (Option<u64>, Option<String>) {
let current_reason = key
.oauth_invalid_reason
.as_deref()
.map(str::trim)
.unwrap_or_default();
if current_reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX) {
return (
key.oauth_invalid_at_unix_secs,
(!current_reason.is_empty()).then_some(current_reason.to_string()),
);
}
(None, None)
admin_provider_quota_pure::quota_refresh_success_invalid_state(key)
}
pub(super) fn coerce_json_string(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
admin_provider_quota_pure::coerce_json_string(value)
}
pub(crate) async fn persist_provider_quota_refresh_state(
state: &AppState,
state: &AdminAppState<'_>,
key_id: &str,
metadata_update: Option<&serde_json::Value>,
oauth_invalid_at_unix_secs: Option<u64>,
@@ -182,12 +107,12 @@ pub(crate) async fn persist_provider_quota_refresh_state(
}
pub(super) async fn execute_provider_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
plan: ExecutionPlan,
quota_kind: &str,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
match crate::execution_runtime::execute_execution_runtime_sync_plan(state, None, &plan).await {
match state.execute_execution_runtime_sync_plan(None, &plan).await {
Ok(result) => Ok(ProviderQuotaExecutionOutcome::Response(result)),
Err(err) => {
let error = match err {
@@ -1,647 +0,0 @@
use super::quota::antigravity::refresh_antigravity_provider_quota_locally;
use super::quota::codex::refresh_codex_provider_quota_locally;
use super::quota::kiro::refresh_kiro_provider_quota_locally;
use super::quota::shared::persist_provider_quota_refresh_state;
use super::state::{
enrich_admin_provider_oauth_auth_config, json_non_empty_string, json_u64_value,
};
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
OAUTH_REQUEST_FAILED_PREFIX,
};
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks,
parse_catalog_auth_config_json,
};
use crate::{AppState, GatewayError};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
pub(crate) fn build_internal_control_error_response(
status: http::StatusCode,
message: impl Into<String>,
) -> Response<Body> {
(status, Json(json!({ "detail": message.into() }))).into_response()
}
pub(crate) fn normalize_provider_oauth_refresh_error_message(
status_code: Option<u16>,
body_excerpt: Option<&str>,
) -> String {
let mut message = None::<String>;
let mut error_code = None::<String>;
let mut error_type = None::<String>;
if let Some(body_excerpt) = body_excerpt {
if let Ok(value) = serde_json::from_str::<serde_json::Value>(body_excerpt) {
if let Some(object) = value.as_object() {
if let Some(error_object) =
object.get("error").and_then(serde_json::Value::as_object)
{
message = error_object
.get("message")
.or_else(|| error_object.get("error_description"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
error_code = error_object
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
error_type = error_object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if message.is_none() {
message = object
.get("message")
.or_else(|| object.get("error_description"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
}
if error_code.is_none() {
error_code = object
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if error_type.is_none() {
error_type = object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
}
}
}
let message = message
.or_else(|| {
body_excerpt
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.chars().take(300).collect::<String>())
})
.unwrap_or_default();
let lowered = message.to_ascii_lowercase();
let error_code = error_code.unwrap_or_default();
let error_type = error_type.unwrap_or_default();
if error_code == "refresh_token_reused"
|| lowered.contains("already been used to generate a new access token")
{
return "refresh_token 已被使用并轮换,请重新登录授权".to_string();
}
if error_code == "invalid_grant"
|| error_code == "invalid_refresh_token"
|| (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid"]
.iter()
.any(|keyword| lowered.contains(keyword)))
{
return "refresh_token 无效、已过期或已撤销,请重新登录授权".to_string();
}
if error_type == "invalid_request_error" && !message.is_empty() {
return message;
}
if !message.is_empty() {
return message;
}
status_code
.map(|status_code| format!("HTTP {status_code}"))
.unwrap_or_else(|| "未知错误".to_string())
}
pub(crate) fn merge_provider_oauth_refresh_failure_reason(
current_reason: Option<&str>,
refresh_reason: &str,
) -> Option<String> {
let current_reason = current_reason.map(str::trim).unwrap_or_default();
let refresh_reason = refresh_reason.trim();
if refresh_reason.is_empty() {
return (!current_reason.is_empty()).then(|| current_reason.to_string());
}
if current_reason.is_empty() {
return Some(refresh_reason.to_string());
}
if current_reason.starts_with(OAUTH_EXPIRED_PREFIX) {
return None;
}
if current_reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX) {
if let Some((head, _)) = current_reason.split_once("[REFRESH_FAILED]") {
return Some(
format!("{}\n{}", head.trim_end(), refresh_reason)
.trim()
.to_string(),
);
}
return Some(format!("{current_reason}\n{refresh_reason}"));
}
Some(refresh_reason.to_string())
}
pub(crate) fn provider_oauth_key_proxy_value(
proxy_node_id: Option<&str>,
) -> Option<serde_json::Value> {
proxy_node_id
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| json!({ "node_id": value, "enabled": true }))
}
pub(crate) fn provider_oauth_active_api_formats(
endpoints: &[StoredProviderCatalogEndpoint],
) -> Vec<String> {
let mut formats = Vec::new();
let mut seen = BTreeSet::new();
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
let api_format = endpoint.api_format.trim();
if api_format.is_empty() || !seen.insert(api_format.to_string()) {
continue;
}
formats.push(api_format.to_string());
}
formats
}
fn normalize_codex_plan_group_for_provider_oauth(
plan_type: Option<&serde_json::Value>,
) -> Option<String> {
let normalized = plan_type
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_ascii_lowercase();
match normalized.as_str() {
"free" => Some("free".to_string()),
"team" | "plus" | "enterprise" => Some("team_plus_enterprise".to_string()),
_ => None,
}
}
fn normalize_provider_oauth_identity_value(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn is_codex_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|provider_type| provider_type.eq_ignore_ascii_case("codex"))
}
fn match_codex_provider_oauth_identity(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
) -> Option<bool> {
let new_provider_type = new_auth_config.get("provider_type");
let existing_provider_type = existing_auth_config.get("provider_type");
if !is_codex_provider_oauth_provider_type(new_provider_type)
&& !is_codex_provider_oauth_provider_type(existing_provider_type)
{
return None;
}
let new_account_user_id =
normalize_provider_oauth_identity_value(new_auth_config.get("account_user_id"));
let existing_account_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("account_user_id"));
if let (Some(new_account_user_id), Some(existing_account_user_id)) =
(new_account_user_id, existing_account_user_id)
{
return Some(new_account_user_id == existing_account_user_id);
}
let new_account_id = normalize_provider_oauth_identity_value(new_auth_config.get("account_id"));
let existing_account_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("account_id"));
let new_user_id = normalize_provider_oauth_identity_value(new_auth_config.get("user_id"));
let existing_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("user_id"));
let new_email = normalize_provider_oauth_identity_value(new_auth_config.get("email"));
let existing_email = normalize_provider_oauth_identity_value(existing_auth_config.get("email"));
if let (Some(new_account_id), Some(existing_account_id)) =
(new_account_id.as_deref(), existing_account_id.as_deref())
{
if new_account_id != existing_account_id {
return Some(false);
}
}
if let (
Some(new_account_id),
Some(existing_account_id),
Some(new_user_id),
Some(existing_user_id),
) = (
new_account_id.as_deref(),
existing_account_id.as_deref(),
new_user_id.as_deref(),
existing_user_id.as_deref(),
) {
return Some(new_account_id == existing_account_id && new_user_id == existing_user_id);
}
if let (
Some(new_account_id),
Some(existing_account_id),
Some(new_email),
Some(existing_email),
) = (
new_account_id.as_deref(),
existing_account_id.as_deref(),
new_email.as_deref(),
existing_email.as_deref(),
) {
return Some(new_account_id == existing_account_id && new_email == existing_email);
}
None
}
fn is_codex_cross_plan_group_non_duplicate(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
) -> bool {
let new_provider_type = new_auth_config.get("provider_type");
let existing_provider_type = existing_auth_config.get("provider_type");
if !is_codex_provider_oauth_provider_type(new_provider_type)
&& !is_codex_provider_oauth_provider_type(existing_provider_type)
{
return false;
}
let new_group = normalize_codex_plan_group_for_provider_oauth(new_auth_config.get("plan_type"));
let existing_group =
normalize_codex_plan_group_for_provider_oauth(existing_auth_config.get("plan_type"));
matches!(
(new_group.as_deref(), existing_group.as_deref()),
(Some(left), Some(right)) if left != right
)
}
pub(crate) async fn find_duplicate_provider_oauth_key(
state: &AppState,
provider_id: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
exclude_key_id: Option<&str>,
) -> Result<Option<StoredProviderCatalogKey>, String> {
let new_email = normalize_provider_oauth_identity_value(auth_config.get("email"));
let new_user_id = normalize_provider_oauth_identity_value(auth_config.get("user_id"));
let new_auth_method = normalize_provider_oauth_identity_value(auth_config.get("auth_method"));
if new_email.is_none() && new_user_id.is_none() {
return Ok(None);
}
let existing_keys = state
.list_provider_catalog_keys_by_provider_ids(&[provider_id.to_string()])
.await
.map_err(|err| format!("{err:?}"))?;
for existing_key in existing_keys.into_iter().filter(|key| {
key.auth_type.trim().eq_ignore_ascii_case("oauth")
&& exclude_key_id.is_none_or(|exclude| key.id != exclude)
}) {
let Some(existing_auth_config) = parse_catalog_auth_config_json(state, &existing_key)
else {
continue;
};
let existing_email =
normalize_provider_oauth_identity_value(existing_auth_config.get("email"));
let existing_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("user_id"));
let existing_auth_method =
normalize_provider_oauth_identity_value(existing_auth_config.get("auth_method"));
let mut is_duplicate = false;
let codex_identity_match =
match_codex_provider_oauth_identity(auth_config, &existing_auth_config);
if let Some(codex_identity_match) = codex_identity_match {
is_duplicate = codex_identity_match;
}
if codex_identity_match.is_none()
&& !is_duplicate
&& new_user_id.is_some()
&& existing_user_id.is_some()
&& new_user_id == existing_user_id
&& !is_codex_cross_plan_group_non_duplicate(auth_config, &existing_auth_config)
{
is_duplicate = true;
}
if codex_identity_match.is_none()
&& !is_duplicate
&& new_email.is_some()
&& existing_email.is_some()
&& new_email == existing_email
{
let is_kiro = auth_config
.get("provider_type")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| value.eq_ignore_ascii_case("kiro"))
|| existing_auth_config
.get("provider_type")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| value.eq_ignore_ascii_case("kiro"));
if is_kiro {
if new_auth_method.is_some()
&& existing_auth_method.is_some()
&& new_auth_method
.as_deref()
.zip(existing_auth_method.as_deref())
.is_some_and(|(left, right)| left.eq_ignore_ascii_case(right))
{
is_duplicate = true;
}
} else if !is_codex_cross_plan_group_non_duplicate(auth_config, &existing_auth_config) {
is_duplicate = true;
}
}
if !is_duplicate {
continue;
}
if !existing_key.is_active {
return Ok(Some(existing_key));
}
let identifier =
normalize_provider_oauth_identity_value(auth_config.get("account_user_id"))
.or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id")))
.or_else(|| new_email.clone())
.or_else(|| new_user_id.clone())
.unwrap_or_default();
return Err(format!(
"该 OAuth 账号 ({identifier}) 已存在于当前 Provider 中(名称: {})",
existing_key.name
));
}
Ok(None)
}
pub(crate) fn build_provider_oauth_auth_config_from_token_payload(
provider_type: &str,
token_payload: &serde_json::Value,
) -> (
serde_json::Map<String, serde_json::Value>,
Option<String>,
Option<String>,
Option<u64>,
) {
let access_token = json_non_empty_string(token_payload.get("access_token"));
let refresh_token = json_non_empty_string(token_payload.get("refresh_token"));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let expires_at = json_u64_value(token_payload.get("expires_in"))
.map(|expires_in| now_unix_secs.saturating_add(expires_in));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!(provider_type));
auth_config.insert("updated_at".to_string(), json!(now_unix_secs));
if let Some(token_type) = token_payload.get("token_type").cloned() {
auth_config.insert("token_type".to_string(), token_type);
}
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
if let Some(expires_at) = expires_at {
auth_config.insert("expires_at".to_string(), json!(expires_at));
}
if let Some(scope) = token_payload.get("scope").cloned() {
auth_config.insert("scope".to_string(), scope);
}
enrich_admin_provider_oauth_auth_config(provider_type, &mut auth_config, token_payload);
(auth_config, access_token, refresh_token, expires_at)
}
pub(crate) async fn create_provider_oauth_catalog_key(
state: &AppState,
provider_id: &str,
name: &str,
access_token: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
api_formats: &[String],
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = encrypt_catalog_secret_with_fallbacks(state, access_token) else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
encrypt_catalog_secret_with_fallbacks(state, &auth_config_json)
else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut record = StoredProviderCatalogKey::new(
Uuid::new_v4().to_string(),
provider_id.to_string(),
name.to_string(),
"oauth".to_string(),
None,
true,
)
.map_err(|err| GatewayError::Internal(err.to_string()))?
.with_transport_fields(
Some(json!(api_formats)),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
expires_at_unix_secs,
proxy,
None,
)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
record.internal_priority = 50;
record.cache_ttl_minutes = 5;
record.max_probe_interval_minutes = 32;
record.request_count = Some(0);
record.success_count = Some(0);
record.error_count = Some(0);
record.total_response_time_ms = Some(0);
record.health_by_format = Some(json!({}));
record.circuit_breaker_by_format = Some(json!({}));
record.created_at_unix_secs = Some(now_unix_secs);
record.updated_at_unix_secs = Some(now_unix_secs);
state.create_provider_catalog_key(&record).await
}
pub(crate) async fn update_existing_provider_oauth_catalog_key(
state: &AppState,
existing_key: &StoredProviderCatalogKey,
access_token: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = encrypt_catalog_secret_with_fallbacks(state, access_token) else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
encrypt_catalog_secret_with_fallbacks(state, &auth_config_json)
else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut updated = existing_key.clone();
updated.encrypted_api_key = encrypted_api_key;
updated.encrypted_auth_config = Some(encrypted_auth_config);
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;
updated.health_by_format = Some(json!({}));
updated.circuit_breaker_by_format = Some(json!({}));
updated.error_count = Some(0);
if let Some(proxy) = proxy {
updated.proxy = Some(proxy);
}
updated.updated_at_unix_secs = Some(now_unix_secs);
state.update_provider_catalog_key(&updated).await
}
pub(crate) fn provider_oauth_runtime_endpoint_for_provider(
provider_type: &str,
endpoints: Vec<StoredProviderCatalogEndpoint>,
) -> Option<StoredProviderCatalogEndpoint> {
let provider_type = provider_type.trim().to_ascii_lowercase();
match provider_type.as_str() {
"codex" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}),
"antigravity" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& (endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:chat")
|| endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:cli"))
}),
"kiro" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:cli")
})
.cloned()
.or_else(|| endpoints.into_iter().find(|endpoint| endpoint.is_active)),
_ => endpoints.into_iter().find(|endpoint| endpoint.is_active),
}
}
pub(crate) async fn refresh_provider_oauth_account_state_after_update(
state: &AppState,
provider: &StoredProviderCatalogProvider,
key_id: &str,
) -> Result<(bool, Option<String>), GatewayError> {
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !matches!(provider_type.as_str(), "codex" | "kiro" | "antigravity") {
return Ok((false, None));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
else {
return Ok((false, None));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok((false, None));
};
let payload = match provider_type.as_str() {
"codex" => {
refresh_codex_provider_quota_locally(state, provider, &endpoint, vec![key]).await?
}
"kiro" => {
refresh_kiro_provider_quota_locally(state, provider, &endpoint, vec![key]).await?
}
"antigravity" => {
refresh_antigravity_provider_quota_locally(state, provider, &endpoint, vec![key])
.await?
}
_ => None,
};
let Some(payload) = payload else {
return Ok((false, None));
};
let success = payload
.get("success")
.and_then(serde_json::Value::as_u64)
.unwrap_or(0);
let error = if success == 0 {
payload
.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)
} else {
None
};
Ok((true, error))
}
@@ -0,0 +1,107 @@
use super::quota::antigravity::refresh_antigravity_provider_quota_locally;
use super::quota::codex::refresh_codex_provider_quota_locally;
use super::quota::kiro::refresh_kiro_provider_quota_locally;
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
pub(crate) fn provider_oauth_runtime_endpoint_for_provider(
provider_type: &str,
endpoints: Vec<StoredProviderCatalogEndpoint>,
) -> Option<StoredProviderCatalogEndpoint> {
let provider_type = provider_type.trim().to_ascii_lowercase();
match provider_type.as_str() {
"codex" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}),
"antigravity" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& (endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:chat")
|| endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:cli"))
}),
"kiro" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:cli")
})
.cloned()
.or_else(|| endpoints.into_iter().find(|endpoint| endpoint.is_active)),
_ => endpoints.into_iter().find(|endpoint| endpoint.is_active),
}
}
pub(crate) async fn refresh_provider_oauth_account_state_after_update(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
key_id: &str,
) -> Result<(bool, Option<String>), GatewayError> {
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !matches!(provider_type.as_str(), "codex" | "kiro" | "antigravity") {
return Ok((false, None));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
else {
return Ok((false, None));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok((false, None));
};
let payload = match provider_type.as_str() {
"codex" => {
refresh_codex_provider_quota_locally(state, provider, &endpoint, vec![key]).await?
}
"kiro" => {
refresh_kiro_provider_quota_locally(state, provider, &endpoint, vec![key]).await?
}
"antigravity" => {
refresh_antigravity_provider_quota_locally(state, provider, &endpoint, vec![key])
.await?
}
_ => None,
};
let Some(payload) = payload else {
return Ok((false, None));
};
let success = payload
.get("success")
.and_then(serde_json::Value::as_u64)
.unwrap_or(0);
let error = if success == 0 {
payload
.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)
} else {
None
};
Ok((true, error))
}
@@ -1,996 +0,0 @@
use super::refresh::{
build_internal_control_error_response, normalize_provider_oauth_refresh_error_message,
};
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
use crate::provider_transport::provider_types::{
provider_type_admin_oauth_template, provider_type_is_fixed_for_admin_oauth,
ProviderOAuthTemplate, ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES,
};
use crate::{AppState, GatewayError};
use aether_data::repository::provider_oauth::{
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key,
provider_oauth_device_session_storage_key, provider_oauth_state_storage_key,
StoredAdminProviderOAuthDeviceSession, StoredAdminProviderOAuthState,
PROVIDER_OAUTH_BATCH_TASK_TTL_SECS, PROVIDER_OAUTH_STATE_TTL_SECS,
};
use axum::{body::Body, http, response::Response};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use url::{form_urlencoded, Url};
use uuid::Uuid;
pub(crate) fn is_fixed_provider_type_for_provider_oauth(provider_type: &str) -> bool {
provider_type_is_fixed_for_admin_oauth(provider_type)
}
pub(crate) fn admin_provider_oauth_template(provider_type: &str) -> Option<ProviderOAuthTemplate> {
provider_type_admin_oauth_template(provider_type)
}
pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec<serde_json::Value> {
ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES
.into_iter()
.filter_map(|provider_type| admin_provider_oauth_template(provider_type))
.map(|template| {
json!({
"provider_type": template.provider_type,
"display_name": template.display_name,
"scopes": template.scopes,
"redirect_uri": template.redirect_uri,
"authorize_url": template.authorize_url,
"token_url": template.token_url,
"use_pkce": template.use_pkce,
})
})
.collect()
}
pub(crate) fn build_admin_provider_oauth_backend_unavailable_response() -> Response<Body> {
build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
)
}
const KIRO_DEVICE_DEFAULT_START_URL: &str = "https://view.awsapps.com/start";
const KIRO_DEVICE_DEFAULT_REGION: &str = "us-east-1";
const KIRO_IDC_AMZ_USER_AGENT: &str =
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
pub(crate) fn default_kiro_device_start_url() -> String {
KIRO_DEVICE_DEFAULT_START_URL.to_string()
}
pub(crate) fn default_kiro_device_region() -> String {
KIRO_DEVICE_DEFAULT_REGION.to_string()
}
pub(crate) fn normalize_kiro_device_region(value: Option<&str>) -> Option<String> {
let value = value
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(KIRO_DEVICE_DEFAULT_REGION);
value
.chars()
.all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
.then(|| value.to_string())
}
pub(crate) fn current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
}
pub(crate) async fn save_provider_oauth_device_session(
state: &AppState,
session_id: &str,
session: &StoredAdminProviderOAuthDeviceSession,
ttl_seconds: u64,
) -> Result<(), Response<Body>> {
let key = provider_oauth_device_session_storage_key(session_id);
let value = serde_json::to_string(session).map_err(|_| {
build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
)
})?;
if let Some(runner) = state.redis_kv_runner() {
runner
.setex(&key, &value, Some(ttl_seconds))
.await
.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
)
})?;
return Ok(());
}
if state.save_provider_oauth_device_session_for_tests(&key, &value) {
return Ok(());
}
Err(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
))
}
pub(crate) async fn read_provider_oauth_device_session(
state: &AppState,
session_id: &str,
) -> Result<Option<StoredAdminProviderOAuthDeviceSession>, GatewayError> {
let key = provider_oauth_device_session_storage_key(session_id);
let raw = if let Some(runner) = state.redis_kv_runner() {
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state.load_provider_oauth_device_session_for_tests(&key)
};
raw.map(|value| {
serde_json::from_str::<StoredAdminProviderOAuthDeviceSession>(&value)
.map_err(|err| GatewayError::Internal(err.to_string()))
})
.transpose()
}
async fn post_kiro_device_oidc_json(
state: &AppState,
endpoint_key: &str,
default_url: String,
body: serde_json::Value,
) -> Result<serde_json::Value, Response<Body>> {
let url = state.provider_oauth_token_url(endpoint_key, &default_url);
let host = Url::parse(&url)
.ok()
.and_then(|value| value.host_str().map(ToOwned::to_owned))
.unwrap_or_default();
let response = state
.client
.post(url)
.header("Content-Type", "application/json")
.header("Accept", "*/*")
.header("User-Agent", "node")
.header("x-amz-user-agent", KIRO_IDC_AMZ_USER_AGENT)
.header("Host", host)
.json(&body)
.send()
.await
.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
)
})?;
let status = response.status();
let body_text = response.text().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
)
})?;
serde_json::from_str::<serde_json::Value>(&body_text).or_else(|_| {
Ok(json!({
"_error": !status.is_success(),
"error": body_text.trim(),
}))
})
}
pub(crate) async fn register_admin_kiro_device_oidc_client(
state: &AppState,
region: &str,
start_url: &str,
) -> Result<serde_json::Value, Response<Body>> {
let payload = post_kiro_device_oidc_json(
state,
"kiro_device_register",
format!("https://oidc.{region}.amazonaws.com/client/register"),
json!({
"clientName": "Aether Gateway",
"clientType": "public",
"scopes": [
"codewhisperer:completions",
"codewhisperer:analysis",
"codewhisperer:conversations",
"codewhisperer:transformations",
"codewhisperer:taskassist"
],
"grantTypes": [
"urn:ietf:params:oauth:grant-type:device_code",
"refresh_token"
],
"issuerUrl": start_url,
}),
)
.await?;
if payload
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_desc = json_non_empty_string(payload.get("error_description"))
.or_else(|| json_non_empty_string(payload.get("error")))
.unwrap_or_else(|| "unknown".to_string());
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("注册 OIDC 客户端失败: {error_desc}"),
));
}
Ok(payload)
}
pub(crate) async fn start_admin_kiro_device_authorization(
state: &AppState,
region: &str,
client_id: &str,
client_secret: &str,
start_url: &str,
) -> Result<serde_json::Value, Response<Body>> {
let payload = post_kiro_device_oidc_json(
state,
"kiro_device_authorize",
format!("https://oidc.{region}.amazonaws.com/device_authorization"),
json!({
"clientId": client_id,
"clientSecret": client_secret,
"startUrl": start_url,
}),
)
.await?;
if payload
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_desc = json_non_empty_string(payload.get("error_description"))
.or_else(|| json_non_empty_string(payload.get("error")))
.unwrap_or_else(|| "unknown".to_string());
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("发起设备授权失败: {error_desc}"),
));
}
Ok(payload)
}
pub(crate) async fn poll_admin_kiro_device_token(
state: &AppState,
region: &str,
client_id: &str,
client_secret: &str,
device_code: &str,
) -> Result<serde_json::Value, Response<Body>> {
post_kiro_device_oidc_json(
state,
"kiro_device_poll",
format!("https://oidc.{region}.amazonaws.com/token"),
json!({
"clientId": client_id,
"clientSecret": client_secret,
"grantType": "urn:ietf:params:oauth:grant-type:device_code",
"deviceCode": device_code,
}),
)
.await
}
pub(crate) fn build_kiro_device_key_name(
email: Option<&str>,
refresh_token: Option<&str>,
) -> String {
if let Some(email) = email.map(str::trim).filter(|value| !value.is_empty()) {
return format!("{email} (idc)");
}
let fallback = refresh_token
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| {
let digest = Sha256::digest(value.as_bytes());
digest[..3]
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>()
})
.unwrap_or_else(|| "unknown".to_string());
format!("kiro_{fallback} (idc)")
}
pub(crate) fn generate_provider_oauth_nonce() -> String {
format!("{}{}", Uuid::new_v4().simple(), Uuid::new_v4().simple())
}
pub(crate) fn generate_provider_oauth_pkce_verifier() -> String {
format!(
"{}{}{}",
Uuid::new_v4().simple(),
Uuid::new_v4().simple(),
Uuid::new_v4().simple()
)
}
pub(crate) fn provider_oauth_pkce_s256(verifier: &str) -> String {
let digest = Sha256::digest(verifier.as_bytes());
URL_SAFE_NO_PAD.encode(digest)
}
pub(crate) async fn save_provider_oauth_state(
state: &AppState,
key_id: &str,
provider_id: &str,
provider_type: &str,
pkce_verifier: Option<&str>,
) -> Result<String, GatewayError> {
let nonce = generate_provider_oauth_nonce();
let payload = json!({
"nonce": nonce,
"key_id": key_id,
"provider_id": provider_id,
"provider_type": provider_type,
"pkce_verifier": pkce_verifier,
"created_at": SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs())
.unwrap_or(0),
});
let key = provider_oauth_state_storage_key(&nonce);
let value = payload.to_string();
if let Some(runner) = state.redis_kv_runner() {
runner
.setex(&key, &value, Some(PROVIDER_OAUTH_STATE_TTL_SECS))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(nonce);
}
if state.save_provider_oauth_state_for_tests(&key, &value) {
return Ok(nonce);
}
Err(GatewayError::Internal(
"provider oauth redis unavailable".to_string(),
))
}
pub(crate) fn parse_provider_oauth_callback_params(callback_url: &str) -> BTreeMap<String, String> {
let mut merged = BTreeMap::new();
let Ok(url) = Url::parse(callback_url.trim()) else {
return merged;
};
for (key, value) in url.query_pairs() {
merged.insert(key.into_owned(), value.into_owned());
}
if let Some(fragment) = url.fragment() {
for (key, value) in form_urlencoded::parse(fragment.as_bytes()) {
merged
.entry(key.into_owned())
.or_insert_with(|| value.into_owned());
}
}
if let Some(code) = merged.get("code").cloned() {
if let Some((code_part, state_part)) = code.split_once("#state=") {
merged.insert("code".to_string(), code_part.to_string());
merged
.entry("state".to_string())
.or_insert_with(|| state_part.to_string());
}
}
merged
}
pub(crate) async fn consume_provider_oauth_state(
state: &AppState,
nonce: &str,
) -> Result<Option<StoredAdminProviderOAuthState>, GatewayError> {
let key = provider_oauth_state_storage_key(nonce);
let raw = if let Some(runner) = state.redis_kv_runner() {
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GETDEL")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state.take_provider_oauth_state_for_tests(&key)
};
raw.map(|value| {
serde_json::from_str::<StoredAdminProviderOAuthState>(&value)
.map_err(|err| GatewayError::Internal(err.to_string()))
})
.transpose()
}
pub(crate) fn json_non_empty_string(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(crate) fn json_u64_value(value: Option<&serde_json::Value>) -> Option<u64> {
match value? {
serde_json::Value::Number(number) => number.as_u64(),
serde_json::Value::String(value) => value.trim().parse::<u64>().ok(),
_ => None,
}
}
pub(crate) fn decode_jwt_claims(token: &str) -> Option<serde_json::Map<String, serde_json::Value>> {
let payload = token.split('.').nth(1)?;
let bytes = URL_SAFE_NO_PAD.decode(payload.as_bytes()).ok()?;
serde_json::from_slice::<serde_json::Value>(&bytes)
.ok()?
.as_object()
.cloned()
}
fn merge_missing_auth_config_fields(
auth_config: &mut serde_json::Map<String, serde_json::Value>,
source: &serde_json::Map<String, serde_json::Value>,
fields: &[&str],
) {
for field in fields {
if auth_config.contains_key(*field) {
continue;
}
if let Some(value) = source.get(*field).cloned() {
auth_config.insert((*field).to_string(), value);
}
}
}
fn first_json_non_empty_string(
values: impl IntoIterator<Item = Option<serde_json::Value>>,
) -> Option<String> {
values.into_iter().find_map(|value| match value {
Some(serde_json::Value::String(value)) => {
let normalized = value.trim();
(!normalized.is_empty()).then(|| normalized.to_string())
}
_ => None,
})
}
fn extract_codex_auth_fields_from_object(
source: &serde_json::Map<String, serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
let auth = source
.get("https://api.openai.com/auth")
.and_then(serde_json::Value::as_object);
let mut result = serde_json::Map::new();
if let Some(email) = first_json_non_empty_string([
source.get("email").cloned(),
auth.and_then(|value| value.get("email")).cloned(),
]) {
result.insert("email".to_string(), json!(email));
}
if let Some(account_id) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_account_id"))
.cloned(),
auth.and_then(|value| value.get("chatgptAccountId"))
.cloned(),
auth.and_then(|value| value.get("account_id")).cloned(),
auth.and_then(|value| value.get("accountId")).cloned(),
source.get("chatgpt_account_id").cloned(),
source.get("chatgptAccountId").cloned(),
source.get("account_id").cloned(),
source.get("accountId").cloned(),
]) {
result.insert("account_id".to_string(), json!(account_id));
}
if let Some(account_user_id) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_account_user_id"))
.cloned(),
auth.and_then(|value| value.get("chatgptAccountUserId"))
.cloned(),
auth.and_then(|value| value.get("account_user_id")).cloned(),
auth.and_then(|value| value.get("accountUserId")).cloned(),
source.get("chatgpt_account_user_id").cloned(),
source.get("chatgptAccountUserId").cloned(),
source.get("account_user_id").cloned(),
source.get("accountUserId").cloned(),
]) {
result.insert("account_user_id".to_string(), json!(account_user_id));
}
if let Some(plan_type) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_plan_type"))
.cloned(),
auth.and_then(|value| value.get("chatgptPlanType")).cloned(),
auth.and_then(|value| value.get("plan_type")).cloned(),
auth.and_then(|value| value.get("planType")).cloned(),
source.get("chatgpt_plan_type").cloned(),
source.get("chatgptPlanType").cloned(),
source.get("plan_type").cloned(),
source.get("planType").cloned(),
]) {
result.insert("plan_type".to_string(), json!(plan_type));
}
if let Some(user_id) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_user_id")).cloned(),
auth.and_then(|value| value.get("chatgptUserId")).cloned(),
auth.and_then(|value| value.get("user_id")).cloned(),
auth.and_then(|value| value.get("userId")).cloned(),
source.get("chatgpt_user_id").cloned(),
source.get("chatgptUserId").cloned(),
source.get("user_id").cloned(),
source.get("userId").cloned(),
source.get("sub").cloned(),
]) {
result.insert("user_id".to_string(), json!(user_id));
}
if let Some(organizations) = auth
.and_then(|value| value.get("organizations"))
.and_then(serde_json::Value::as_array)
.filter(|value| !value.is_empty())
{
result.insert(
"organizations".to_string(),
serde_json::Value::Array(organizations.clone()),
);
}
result
}
pub(crate) fn enrich_admin_provider_oauth_auth_config(
provider_type: &str,
auth_config: &mut serde_json::Map<String, serde_json::Value>,
token_payload: &serde_json::Value,
) {
let Some(token_payload_object) = token_payload.as_object() else {
return;
};
merge_missing_auth_config_fields(
auth_config,
token_payload_object,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"account_name",
],
);
if !provider_type.eq_ignore_ascii_case("codex") {
return;
}
let codex_fields = extract_codex_auth_fields_from_object(token_payload_object);
merge_missing_auth_config_fields(
auth_config,
&codex_fields,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"organizations",
],
);
for token_field in ["id_token", "idToken", "access_token", "accessToken"] {
let Some(token) = json_non_empty_string(token_payload.get(token_field)) else {
continue;
};
let Some(claims) = decode_jwt_claims(&token) else {
continue;
};
merge_missing_auth_config_fields(
auth_config,
&claims,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"account_name",
],
);
let codex_claim_fields = extract_codex_auth_fields_from_object(&claims);
merge_missing_auth_config_fields(
auth_config,
&codex_claim_fields,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"organizations",
],
);
}
}
#[cfg(test)]
mod tests {
use super::enrich_admin_provider_oauth_auth_config;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(payload.to_string());
format!("{header}.{payload}.sig")
}
#[test]
fn codex_enrichment_extracts_identity_from_nested_auth_claims() {
let access_token = sample_unsigned_jwt(json!({
"email": "[email protected]",
"https://api.openai.com/auth": {
"chatgpt_account_id": "acc-1",
"chatgpt_account_user_id": "user-1__acc-1",
"chatgpt_plan_type": "team",
"chatgpt_user_id": "user-1",
"organizations": [
{"id": "org-1", "title": "Personal", "is_default": true}
],
}
}));
let token_payload = json!({
"access_token": access_token,
});
let mut auth_config = serde_json::Map::new();
enrich_admin_provider_oauth_auth_config("codex", &mut auth_config, &token_payload);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(auth_config.get("account_id"), Some(&json!("acc-1")));
assert_eq!(
auth_config.get("account_user_id"),
Some(&json!("user-1__acc-1"))
);
assert_eq!(auth_config.get("plan_type"), Some(&json!("team")));
assert_eq!(auth_config.get("user_id"), Some(&json!("user-1")));
assert_eq!(
auth_config.get("organizations"),
Some(&json!([
{"id": "org-1", "title": "Personal", "is_default": true}
]))
);
}
#[test]
fn codex_enrichment_normalizes_direct_chatgpt_alias_fields() {
let token_payload = json!({
"email": "[email protected]",
"chatgpt_account_id": "acc-2",
"chatgpt_account_user_id": "user-2__acc-2",
"chatgpt_plan_type": "plus",
"chatgpt_user_id": "user-2",
});
let mut auth_config = serde_json::Map::new();
enrich_admin_provider_oauth_auth_config("codex", &mut auth_config, &token_payload);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(auth_config.get("account_id"), Some(&json!("acc-2")));
assert_eq!(
auth_config.get("account_user_id"),
Some(&json!("user-2__acc-2"))
);
assert_eq!(auth_config.get("plan_type"), Some(&json!("plus")));
assert_eq!(auth_config.get("user_id"), Some(&json!("user-2")));
}
}
pub(crate) async fn exchange_admin_provider_oauth_code(
state: &AppState,
template: ProviderOAuthTemplate,
code: &str,
state_nonce: &str,
pkce_verifier: Option<&str>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let request = state.client.post(token_url);
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("authorization_code".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"redirect_uri".to_string(),
serde_json::Value::String(template.redirect_uri.to_string()),
),
(
"code".to_string(),
serde_json::Value::String(code.to_string()),
),
(
"state".to_string(),
serde_json::Value::String(state_nonce.to_string()),
),
]);
if let Some(verifier) = pkce_verifier {
body.insert(
"code_verifier".to_string(),
serde_json::Value::String(verifier.to_string()),
);
}
request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&serde_json::Value::Object(body))
.send()
.await
} else {
let mut form = vec![
("grant_type", "authorization_code".to_string()),
("client_id", template.client_id.to_string()),
("redirect_uri", template.redirect_uri.to_string()),
("code", code.to_string()),
];
if !template.client_secret.trim().is_empty() {
form.push(("client_secret", template.client_secret.to_string()));
}
if let Some(verifier) = pkce_verifier {
form.push(("code_verifier", verifier.to_string()));
}
request
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.form(&form)
.send()
.await
}
.map_err(|_| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "token exchange 失败")
})?;
if !response.status().is_success() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 失败",
));
}
let payload = response.json::<serde_json::Value>().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
}
Ok(payload)
}
pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
state: &AppState,
template: ProviderOAuthTemplate,
refresh_token: &str,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let request = state.client.post(token_url);
let scope = template.scopes.join(" ");
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("refresh_token".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"refresh_token".to_string(),
serde_json::Value::String(refresh_token.to_string()),
),
]);
if !scope.trim().is_empty() {
body.insert("scope".to_string(), serde_json::Value::String(scope));
}
request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&serde_json::Value::Object(body))
.send()
.await
} else {
let mut form = vec![
("grant_type", "refresh_token".to_string()),
("client_id", template.client_id.to_string()),
("refresh_token", refresh_token.to_string()),
];
if !scope.trim().is_empty() {
form.push(("scope", scope));
}
if !template.client_secret.trim().is_empty() {
form.push(("client_secret", template.client_secret.to_string()));
}
request
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.form(&form)
.send()
.await
}
.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 验证失败: token exchange 失败",
)
})?;
let status = response.status();
let body = response.text().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 验证失败: token exchange 失败",
)
})?;
if !status.is_success() {
let reason =
normalize_provider_oauth_refresh_error_message(Some(status.as_u16()), Some(&body));
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Refresh Token 验证失败: {reason}"),
));
}
let payload = serde_json::from_str::<serde_json::Value>(&body).map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
));
}
Ok(payload)
}
pub(crate) fn build_provider_oauth_start_response(
template: ProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> serde_json::Value {
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("client_id", template.client_id);
serializer.append_pair("response_type", "code");
serializer.append_pair("redirect_uri", template.redirect_uri);
serializer.append_pair("scope", &template.scopes.join(" "));
serializer.append_pair("state", nonce);
if template.provider_type == "codex" {
serializer.append_pair("prompt", "login");
serializer.append_pair("id_token_add_organizations", "true");
serializer.append_pair("codex_cli_simplified_flow", "true");
}
if template.use_pkce {
if let Some(code_challenge) = code_challenge {
serializer.append_pair("code_challenge", code_challenge);
serializer.append_pair("code_challenge_method", "S256");
}
}
json!({
"authorization_url": format!("{}?{}", template.authorize_url, serializer.finish()),
"redirect_uri": template.redirect_uri,
"provider_type": template.provider_type,
"instructions": "1) 打开 authorization_url 完成授权\n2) 授权后会跳转到 redirect_uri(localhost)\n3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url",
})
}
pub(crate) async fn save_provider_oauth_batch_task_payload(
state: &AppState,
task_id: &str,
task_state: &serde_json::Value,
) -> Result<(), GatewayError> {
let key = provider_oauth_batch_task_storage_key(task_id);
let serialized =
serde_json::to_string(task_state).map_err(|err| GatewayError::Internal(err.to_string()))?;
if let Some(runner) = state.redis_kv_runner() {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
return Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
));
};
let redis_key = runner.keyspace().key(&key);
redis::cmd("SET")
.arg(redis_key)
.arg(&serialized)
.arg("EX")
.arg(PROVIDER_OAUTH_BATCH_TASK_TTL_SECS)
.query_async::<()>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(());
}
if state.save_provider_oauth_batch_task_for_tests(&key, &serialized) {
return Ok(());
}
Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
))
}
pub(crate) async fn read_provider_oauth_batch_task_payload(
state: &AppState,
provider_id: &str,
task_id: &str,
) -> Result<Option<serde_json::Value>, GatewayError> {
let key = provider_oauth_batch_task_storage_key(task_id);
let raw = if let Some(runner) = state.redis_kv_runner() {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
return Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
));
};
let redis_key = runner.keyspace().key(&key);
redis::cmd("GET")
.arg(redis_key)
.query_async(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state.load_provider_oauth_batch_task_for_tests(&key)
};
let Some(raw) = raw else {
return Ok(None);
};
let parsed = match serde_json::from_str::<serde_json::Value>(&raw) {
Ok(value) => value,
Err(_) => return Ok(None),
};
let Some(state) = parsed.as_object() else {
return Ok(None);
};
if state
.get("provider_id")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
!= provider_id
{
return Ok(None);
}
Ok(Some(build_provider_oauth_batch_task_status_payload(
provider_id,
state,
)))
}
@@ -0,0 +1,77 @@
pub(crate) use aether_admin::provider::state::{
decode_jwt_claims, enrich_admin_provider_oauth_auth_config, json_non_empty_string,
json_u64_value,
};
#[cfg(test)]
mod tests {
use super::enrich_admin_provider_oauth_auth_config;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(payload.to_string());
format!("{header}.{payload}.sig")
}
#[test]
fn codex_enrichment_extracts_identity_from_nested_auth_claims() {
let access_token = sample_unsigned_jwt(json!({
"email": "[email protected]",
"https://api.openai.com/auth": {
"chatgpt_account_id": "acc-1",
"chatgpt_account_user_id": "user-1__acc-1",
"chatgpt_plan_type": "team",
"chatgpt_user_id": "user-1",
"organizations": [
{"id": "org-1", "title": "Personal", "is_default": true}
],
}
}));
let token_payload = json!({
"access_token": access_token,
});
let mut auth_config = serde_json::Map::new();
enrich_admin_provider_oauth_auth_config("codex", &mut auth_config, &token_payload);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(auth_config.get("account_id"), Some(&json!("acc-1")));
assert_eq!(
auth_config.get("account_user_id"),
Some(&json!("user-1__acc-1"))
);
assert_eq!(auth_config.get("plan_type"), Some(&json!("team")));
assert_eq!(auth_config.get("user_id"), Some(&json!("user-1")));
assert_eq!(
auth_config.get("organizations"),
Some(&json!([
{"id": "org-1", "title": "Personal", "is_default": true}
]))
);
}
#[test]
fn codex_enrichment_normalizes_direct_chatgpt_alias_fields() {
let token_payload = json!({
"email": "[email protected]",
"chatgpt_account_id": "acc-2",
"chatgpt_account_user_id": "user-2__acc-2",
"chatgpt_plan_type": "plus",
"chatgpt_user_id": "user-2",
});
let mut auth_config = serde_json::Map::new();
enrich_admin_provider_oauth_auth_config("codex", &mut auth_config, &token_payload);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(auth_config.get("account_id"), Some(&json!("acc-2")));
assert_eq!(
auth_config.get("account_user_id"),
Some(&json!("user-2__acc-2"))
);
assert_eq!(auth_config.get("plan_type"), Some(&json!("plus")));
assert_eq!(auth_config.get("user_id"), Some(&json!("user-2")));
}
}
@@ -0,0 +1,185 @@
use super::super::errors::{
build_internal_control_error_response, normalize_provider_oauth_refresh_error_message,
};
use super::json_non_empty_string;
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
use axum::{body::Body, http, response::Response};
pub(crate) async fn exchange_admin_provider_oauth_code(
state: &AdminAppState<'_>,
template: AdminProviderOAuthTemplate,
code: &str,
state_nonce: &str,
pkce_verifier: Option<&str>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let request = state.http_client().post(token_url);
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("authorization_code".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"redirect_uri".to_string(),
serde_json::Value::String(template.redirect_uri.to_string()),
),
(
"code".to_string(),
serde_json::Value::String(code.to_string()),
),
(
"state".to_string(),
serde_json::Value::String(state_nonce.to_string()),
),
]);
if let Some(verifier) = pkce_verifier {
body.insert(
"code_verifier".to_string(),
serde_json::Value::String(verifier.to_string()),
);
}
request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&serde_json::Value::Object(body))
.send()
.await
} else {
let mut form = vec![
("grant_type", "authorization_code".to_string()),
("client_id", template.client_id.to_string()),
("redirect_uri", template.redirect_uri.to_string()),
("code", code.to_string()),
];
if !template.client_secret.trim().is_empty() {
form.push(("client_secret", template.client_secret.to_string()));
}
if let Some(verifier) = pkce_verifier {
form.push(("code_verifier", verifier.to_string()));
}
request
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.form(&form)
.send()
.await
}
.map_err(|_| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "token exchange 失败")
})?;
if !response.status().is_success() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 失败",
));
}
let payload = response.json::<serde_json::Value>().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
}
Ok(payload)
}
pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
state: &AdminAppState<'_>,
template: AdminProviderOAuthTemplate,
refresh_token: &str,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let request = state.http_client().post(token_url);
let scope = template.scopes.join(" ");
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("refresh_token".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"refresh_token".to_string(),
serde_json::Value::String(refresh_token.to_string()),
),
]);
if !scope.trim().is_empty() {
body.insert("scope".to_string(), serde_json::Value::String(scope));
}
request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&serde_json::Value::Object(body))
.send()
.await
} else {
let mut form = vec![
("grant_type", "refresh_token".to_string()),
("client_id", template.client_id.to_string()),
("refresh_token", refresh_token.to_string()),
];
if !scope.trim().is_empty() {
form.push(("scope", scope));
}
if !template.client_secret.trim().is_empty() {
form.push(("client_secret", template.client_secret.to_string()));
}
request
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.form(&form)
.send()
.await
}
.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 验证失败: token exchange 失败",
)
})?;
let status = response.status();
let body = response.text().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 验证失败: token exchange 失败",
)
})?;
if !status.is_success() {
let reason =
normalize_provider_oauth_refresh_error_message(Some(status.as_u16()), Some(&body));
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Refresh Token 验证失败: {reason}"),
));
}
let payload = serde_json::from_str::<serde_json::Value>(&body).map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
));
}
Ok(payload)
}
@@ -0,0 +1,20 @@
mod auth_config;
mod exchange;
mod storage;
mod template;
pub(crate) use self::auth_config::enrich_admin_provider_oauth_auth_config;
pub(crate) use self::exchange::{
exchange_admin_provider_oauth_code, exchange_admin_provider_oauth_refresh_token,
};
pub(crate) use self::storage::build_provider_oauth_start_response;
pub(crate) use self::template::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
build_admin_provider_oauth_supported_types_payload, is_fixed_provider_type_for_provider_oauth,
};
pub(crate) use aether_admin::provider::state::{
build_kiro_device_key_name, current_unix_secs, decode_jwt_claims, default_kiro_device_region,
default_kiro_device_start_url, generate_provider_oauth_nonce,
generate_provider_oauth_pkce_verifier, json_non_empty_string, json_u64_value,
normalize_kiro_device_region, parse_provider_oauth_callback_params, provider_oauth_pkce_s256,
};
@@ -0,0 +1,34 @@
use crate::handlers::admin::request::AdminProviderOAuthTemplate;
use serde_json::json;
use url::form_urlencoded;
pub(crate) fn build_provider_oauth_start_response(
template: AdminProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> serde_json::Value {
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("client_id", template.client_id);
serializer.append_pair("response_type", "code");
serializer.append_pair("redirect_uri", template.redirect_uri);
serializer.append_pair("scope", &template.scopes.join(" "));
serializer.append_pair("state", nonce);
if template.provider_type == "codex" {
serializer.append_pair("prompt", "login");
serializer.append_pair("id_token_add_organizations", "true");
serializer.append_pair("codex_cli_simplified_flow", "true");
}
if template.use_pkce {
if let Some(code_challenge) = code_challenge {
serializer.append_pair("code_challenge", code_challenge);
serializer.append_pair("code_challenge_method", "S256");
}
}
json!({
"authorization_url": format!("{}?{}", template.authorize_url, serializer.finish()),
"redirect_uri": template.redirect_uri,
"provider_type": template.provider_type,
"instructions": "1) 打开 authorization_url 完成授权\n2) 授权后会跳转到 redirect_uri(localhost)\n3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url",
})
}
@@ -0,0 +1,44 @@
use super::super::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
use crate::handlers::admin::request::{
admin_provider_oauth_template as request_admin_provider_oauth_template,
admin_provider_oauth_template_types,
is_fixed_provider_type_for_admin_oauth as request_is_fixed_provider_type_for_admin_oauth,
AdminProviderOAuthTemplate,
};
use axum::{body::Body, http, response::Response};
use serde_json::json;
pub(crate) fn is_fixed_provider_type_for_provider_oauth(provider_type: &str) -> bool {
request_is_fixed_provider_type_for_admin_oauth(provider_type)
}
pub(crate) fn admin_provider_oauth_template(
provider_type: &str,
) -> Option<AdminProviderOAuthTemplate> {
request_admin_provider_oauth_template(provider_type)
}
pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec<serde_json::Value> {
admin_provider_oauth_template_types()
.filter_map(|provider_type| admin_provider_oauth_template(provider_type))
.map(|template| {
json!({
"provider_type": template.provider_type,
"display_name": template.display_name,
"scopes": template.scopes,
"redirect_uri": template.redirect_uri,
"authorize_url": template.authorize_url,
"token_url": template.token_url,
"use_pkce": template.use_pkce,
})
})
.collect()
}
pub(crate) fn build_admin_provider_oauth_backend_unavailable_response() -> Response<Body> {
build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
)
}