refactor: 大规模模块拆分与代码精简,新增 ai-pipeline/data-contracts 独立 crate

- 新增 aether-ai-pipeline 和 aether-data-contracts crate,将 pipeline 逻辑与数据契约从 gateway 中解耦
- 重构 admin handlers:拆分单体模块为 auth/billing/endpoint/features/model/observability/provider/system 等独立子模块
- 合并 chat/cli 重复代码路径:精简 conversion、finalize、planner 中的 sync/chat/cli 分支
- 重构 scheduler/executor/data 层,引入 facade 模式降低模块间耦合
- 移除冗余的 intent 模块,将 plan_fallback/policy/stream_path/sync_path 迁移至 executor
- 前端适配:调整 admin API 调用和 provider 模型测试对话框
This commit is contained in:
fawney19
2026-04-07 02:50:19 +08:00
parent 763ff03a7b
commit 5d96d6673b
732 changed files with 28589 additions and 20662 deletions
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,562 @@
use super::super::provider_oauth_quota::refresh_codex_provider_quota_locally;
use super::super::provider_oauth_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::provider_oauth_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::{
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,550 @@
use super::super::provider_oauth_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::provider_oauth_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::{
admin_provider_oauth_device_authorize_provider_id, admin_provider_oauth_device_poll_provider_id,
};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use aether_data::repository::provider_oauth::{
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde::Deserialize;
use serde_json::json;
#[derive(Debug, Deserialize)]
struct AdminProviderOAuthDeviceAuthorizePayload {
#[serde(default = "default_kiro_device_start_url")]
start_url: String,
#[serde(default = "default_kiro_device_region")]
region: String,
proxy_node_id: Option<String>,
}
#[derive(Debug, Deserialize)]
struct AdminProviderOAuthDevicePollPayload {
session_id: String,
}
pub(super) async fn handle_admin_provider_oauth_device_authorize(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_device_authorize_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload =
match serde_json::from_slice::<AdminProviderOAuthDeviceAuthorizePayload>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"设备授权仅支持 Kiro provider",
));
}
let region = normalize_kiro_device_region(Some(payload.region.as_str())).ok_or_else(|| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "region 格式无效")
});
let region = match region {
Ok(region) => region,
Err(response) => return Ok(response),
};
let start_url = payload.start_url.trim();
let start_url = if start_url.is_empty() {
default_kiro_device_start_url()
} else {
start_url.to_string()
};
let client_registration =
match register_admin_kiro_device_oidc_client(state, &region, &start_url).await {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(client_id) = json_non_empty_string(client_registration.get("clientId")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let Some(client_secret) = json_non_empty_string(client_registration.get("clientSecret")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let device_authorization = match start_admin_kiro_device_authorization(
state,
&region,
&client_id,
&client_secret,
&start_url,
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(device_code) = json_non_empty_string(
device_authorization
.get("deviceCode")
.or_else(|| device_authorization.get("device_code")),
) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
));
};
let user_code = json_non_empty_string(
device_authorization
.get("userCode")
.or_else(|| device_authorization.get("user_code")),
)
.unwrap_or_default();
let verification_uri = json_non_empty_string(
device_authorization
.get("verificationUri")
.or_else(|| device_authorization.get("verification_uri"))
.or_else(|| device_authorization.get("verificationUrl")),
)
.unwrap_or_default();
let verification_uri_complete = json_non_empty_string(
device_authorization
.get("verificationUriComplete")
.or_else(|| device_authorization.get("verification_uri_complete"))
.or_else(|| device_authorization.get("verificationUrlComplete")),
)
.unwrap_or_else(|| verification_uri.clone());
let expires_in = json_u64_value(
device_authorization
.get("expiresIn")
.or_else(|| device_authorization.get("expires_in")),
)
.unwrap_or(600);
let interval = json_u64_value(device_authorization.get("interval")).unwrap_or(5);
let now_unix_secs = current_unix_secs();
let session_id = generate_provider_oauth_nonce();
let session = StoredAdminProviderOAuthDeviceSession {
provider_id: provider_id.clone(),
region,
client_id,
client_secret,
device_code,
interval,
expires_at_unix_secs: now_unix_secs.saturating_add(expires_in),
status: "pending".to_string(),
proxy_node_id: payload
.proxy_node_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
created_at_unix_secs: now_unix_secs,
key_id: None,
email: None,
replaced: false,
error_msg: None,
};
if let Err(response) = save_provider_oauth_device_session(
state,
&session_id,
&session,
expires_in.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
)
.await
{
return Ok(response);
}
Ok(Json(json!({
"session_id": session_id,
"user_code": user_code,
"verification_uri": verification_uri,
"verification_uri_complete": verification_uri_complete,
"expires_in": expires_in,
"interval": interval,
}))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_device_poll(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_device_poll_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload = match serde_json::from_slice::<AdminProviderOAuthDevicePollPayload>(request_body)
{
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let session_id = payload.session_id.trim();
if session_id.is_empty() {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"session_id 不能为空",
));
}
let Some(mut session) = read_provider_oauth_device_session(state, session_id).await? else {
return Ok(Json(json!({
"status": "expired",
"error": "会话不存在或已过期",
"replaced": false,
}))
.into_response());
};
if session.provider_id != provider_id {
return Ok(Json(json!({
"status": "error",
"error": "会话与 Provider 不匹配",
"replaced": false,
}))
.into_response());
}
if session.status == "authorized" {
return Ok(Json(json!({
"status": "authorized",
"key_id": session.key_id,
"email": session.email,
"replaced": session.replaced,
}))
.into_response());
}
if matches!(session.status.as_str(), "expired" | "error") {
return Ok(Json(json!({
"status": session.status,
"error": session.error_msg,
"replaced": session.replaced,
}))
.into_response());
}
if current_unix_secs() > session.expires_at_unix_secs {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let token_result = match poll_admin_kiro_device_token(
state,
&session.region,
&session.client_id,
&session.client_secret,
&session.device_code,
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
if token_result
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_code = json_non_empty_string(token_result.get("error")).unwrap_or_default();
if error_code == "authorization_pending" {
return Ok(Json(json!({"status": "pending", "replaced": false})).into_response());
}
if error_code == "slow_down" {
return Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response());
}
if error_code == "expired_token" {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
if error_code == "access_denied" {
session.status = "error".to_string();
session.error_msg = Some("用户拒绝授权".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"error",
Json(json!({
"status": "error",
"error": "用户拒绝授权",
"replaced": false,
}))
.into_response(),
));
}
let error_message = json_non_empty_string(token_result.get("error_description"))
.or_else(|| (!error_code.is_empty()).then_some(error_code.clone()))
.unwrap_or_else(|| "未知错误".to_string());
return Ok(Json(json!({
"status": "error",
"error": error_message,
"replaced": false,
}))
.into_response());
}
let Some(access_token) = json_non_empty_string(token_result.get("accessToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let Some(refresh_token) = json_non_empty_string(token_result.get("refreshToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let expires_at = json_u64_value(token_result.get("expiresIn"))
.map(|expires_in| current_unix_secs().saturating_add(expires_in))
.unwrap_or_else(|| current_unix_secs().saturating_add(3600));
let email = decode_jwt_claims(&access_token)
.and_then(|claims| claims.get("email").cloned())
.and_then(|value| value.as_str().map(ToOwned::to_owned));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!("kiro"));
auth_config.insert("auth_method".to_string(), json!("idc"));
auth_config.insert("refresh_token".to_string(), json!(refresh_token.clone()));
auth_config.insert("client_id".to_string(), json!(session.client_id.clone()));
auth_config.insert(
"client_secret".to_string(),
json!(session.client_secret.clone()),
);
auth_config.insert("region".to_string(), json!(session.region.clone()));
auth_config.insert("auth_region".to_string(), json!(session.region.clone()));
auth_config.insert("access_token".to_string(), json!(access_token.clone()));
auth_config.insert("expires_at".to_string(), json!(expires_at));
if let Some(email) = email.as_ref() {
auth_config.insert("email".to_string(), json!(email));
}
let duplicate =
match find_duplicate_provider_oauth_key(state, &provider_id, &auth_config, None).await {
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(Json(json!({
"status": "error",
"error": detail,
"replaced": false,
}))
.into_response());
}
};
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
let api_formats = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.filter(|endpoint| endpoint.is_active)
.map(|endpoint| endpoint.api_format)
.collect::<Vec<_>>();
let mut replaced = false;
let persisted_key = if let Some(existing_key) = duplicate {
replaced = true;
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let key_name = build_kiro_device_key_name(email.as_deref(), Some(&refresh_token));
match create_provider_oauth_catalog_key(
state,
&provider_id,
&key_name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
session.status = "authorized".to_string();
session.key_id = Some(persisted_key.id.clone());
session.email = email.clone();
session.replaced = replaced;
session.error_msg = None;
let _ = save_provider_oauth_device_session(state, session_id, &session, 60).await;
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"authorized",
Json(json!({
"status": "authorized",
"key_id": persisted_key.id,
"email": email,
"replaced": replaced,
}))
.into_response(),
))
}
fn attach_admin_provider_oauth_device_poll_terminal_response(
session_id: &str,
status: &str,
response: Response<Body>,
) -> Response<Body> {
match status {
"authorized" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_completed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"expired" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_expired",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"error" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_failed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
_ => response,
}
}
@@ -0,0 +1,212 @@
use super::super::provider_oauth_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::provider_oauth_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::admin_provider_oauth_import_provider_id;
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
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,
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 {
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 对象",
));
}
};
let Some(refresh_token_input) = raw_payload
.get("refresh_token")
.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,
"Refresh Token 不能为空",
));
};
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 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 Some(template) = admin_provider_oauth_template(&provider_type) else {
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 (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 {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
));
};
let refresh_token = returned_refresh_token
.or_else(|| Some(refresh_token_input.to_string()))
.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));
}
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,222 @@
use super::provider_oauth_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::{
admin_provider_oauth_batch_import_provider_id,
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id,
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::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
mod batch;
mod complete;
mod device;
mod import;
mod refresh;
mod start;
mod tasks;
pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("provider_oauth_manage") {
return Ok(None);
}
let route_kind = decision.route_kind.as_deref();
let method = &request_context.request_method;
if route_kind == Some("supported_types")
&& *method == http::Method::GET
&& request_context.request_path == "/api/admin/provider-oauth/supported-types"
{
return Ok(Some(
Json(build_admin_provider_oauth_supported_types_payload()).into_response(),
));
}
if route_kind == Some("start_key_oauth") && *method == http::Method::POST {
let response = start::handle_admin_provider_oauth_start_key(state, request_context).await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_authorization_started",
"start_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_start_key_id(&request_context.request_path),
)));
}
if route_kind == Some("start_provider_oauth") && *method == http::Method::POST {
let response =
start::handle_admin_provider_oauth_start_provider(state, request_context).await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_authorization_started",
"start_provider_oauth_for_provider",
"provider",
admin_provider_oauth_start_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("get_batch_import_task_status") && *method == http::Method::GET {
return Ok(Some(
tasks::handle_admin_provider_oauth_batch_import_task_status(state, request_context)
.await?,
));
}
if route_kind == Some("complete_key_oauth") && *method == http::Method::POST {
let response = complete::handle_admin_provider_oauth_complete_key(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_completed",
"complete_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_complete_key_id(&request_context.request_path),
)));
}
if route_kind == Some("refresh_key_oauth") && *method == http::Method::POST {
let response =
refresh::handle_admin_provider_oauth_refresh_key(state, request_context).await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_refreshed",
"refresh_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_refresh_key_id(&request_context.request_path),
)));
}
if route_kind == Some("complete_provider_oauth") && *method == http::Method::POST {
let response = complete::handle_admin_provider_oauth_complete_provider(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_completed",
"complete_provider_oauth_for_provider",
"provider",
admin_provider_oauth_complete_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("import_refresh_token") && *method == http::Method::POST {
let response = import::handle_admin_provider_oauth_import_refresh_token(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_refresh_token_imported",
"import_provider_oauth_refresh_token",
"provider",
admin_provider_oauth_import_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("batch_import_oauth") && *method == http::Method::POST {
let response =
batch::handle_admin_provider_oauth_batch_import(state, request_context, request_body)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_batch_import_completed",
"batch_import_provider_oauth",
"provider",
admin_provider_oauth_batch_import_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("start_batch_import_oauth_task") && *method == http::Method::POST {
let response = batch::handle_admin_provider_oauth_start_batch_import_task(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
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),
)));
}
if route_kind == Some("device_authorize") && *method == http::Method::POST {
let response = device::handle_admin_provider_oauth_device_authorize(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_device_authorization_started",
"start_provider_oauth_device_authorization",
"provider",
admin_provider_oauth_device_authorize_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("device_poll") && *method == http::Method::POST {
return Ok(Some(
device::handle_admin_provider_oauth_device_poll(state, request_context, request_body)
.await?,
));
}
if matches!(
route_kind,
Some("refresh_key_oauth" | "import_refresh_token")
) {
return Ok(Some(
build_admin_provider_oauth_backend_unavailable_response(),
));
}
Ok(None)
}
fn attach_admin_provider_oauth_audit_response(
response: Response<Body>,
event_name: &'static str,
action: &'static str,
target_type: &'static str,
target_id: Option<String>,
) -> Response<Body> {
if !response.status().is_success() {
return response;
}
let Some(target_id) = target_id else {
return response;
};
attach_admin_audit_response(response, event_name, action, target_type, &target_id)
}
@@ -0,0 +1,226 @@
use super::super::provider_oauth_quota::persist_provider_quota_refresh_state;
use super::super::provider_oauth_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::provider_oauth_state::is_fixed_provider_type_for_provider_oauth;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_oauth_refresh_key_id, 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};
pub(super) async fn handle_admin_provider_oauth_refresh_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> 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 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",
));
};
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())
}
@@ -0,0 +1,173 @@
use super::super::provider_oauth_refresh::build_internal_control_error_response;
use super::super::provider_oauth_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,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
};
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
pub(super) async fn handle_admin_provider_oauth_start_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_start_key_id(&request_context.request_path) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
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 认证类型",
));
}
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",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let pkce_verifier = template
.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
{
Ok(nonce) => nonce,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
Ok(Json(build_provider_oauth_start_response(
template,
&nonce,
code_challenge.as_deref(),
))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_start_provider(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) = admin_provider_oauth_start_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
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 provider_type == "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Kiro 不支持 OAuth 授权,请使用导入授权。",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let pkce_verifier = template
.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
{
Ok(nonce) => nonce,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
Ok(Json(build_provider_oauth_start_response(
template,
&nonce,
code_challenge.as_deref(),
))
.into_response())
}
@@ -0,0 +1,65 @@
use super::super::provider_oauth_refresh::build_internal_control_error_response;
use super::super::provider_oauth_state::read_provider_oauth_batch_task_payload;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::admin_provider_oauth_batch_import_task_path;
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
pub(super) async fn handle_admin_provider_oauth_batch_import_task_status(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some((provider_id, task_id)) =
admin_provider_oauth_batch_import_task_path(&request_context.request_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
{
Ok(Some(payload)) => payload,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"批量导入任务不存在或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth batch task redis unavailable",
));
}
};
let status = payload
.get("status")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
.unwrap_or_default();
let response = Json(payload).into_response();
Ok(match status.as_str() {
"completed" => attach_admin_audit_response(
response,
"admin_provider_oauth_batch_task_completed_viewed",
"view_provider_oauth_batch_task_terminal_state",
"provider_oauth_batch_task",
&format!("{provider_id}:{task_id}"),
),
"failed" => attach_admin_audit_response(
response,
"admin_provider_oauth_batch_task_failed_viewed",
"view_provider_oauth_batch_task_terminal_state",
"provider_oauth_batch_task",
&format!("{provider_id}:{task_id}"),
),
_ => response,
})
}