mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
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:
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(¶ms)?;
|
||||
let state_nonce = extract_admin_provider_oauth_state(¶ms)?;
|
||||
|
||||
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, ®ion, &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,
|
||||
®ion,
|
||||
&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(®ion, &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(®ion, &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) => {
|
||||
|
||||
Reference in New Issue
Block a user