Add shared OAuth flows

This commit is contained in:
fawney19
2026-04-28 15:46:21 +08:00
parent 712b484bc8
commit 70f747d406
56 changed files with 5927 additions and 977 deletions
@@ -95,6 +95,19 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
response::oauth_refresh_failed_service_unavailable_response(source.to_string()),
));
}
Err(AdminLocalOAuthRefreshError::TransportMessage { message, .. }) => {
tracing::warn!(
trace_id = %trace_id,
key_id = %key_id,
provider_id = %provider.id,
provider_type = %provider_type,
error = %message,
"gateway manual provider oauth refresh transport failed"
);
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_failed_service_unavailable_response(message),
));
}
Err(AdminLocalOAuthRefreshError::InvalidResponse { message, .. }) => {
tracing::warn!(
trace_id = %trace_id,
@@ -1,11 +1,12 @@
use super::super::errors::{
build_internal_control_error_response, normalize_provider_oauth_refresh_error_message,
};
use super::json_non_empty_string;
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
use aether_contracts::ProxySnapshot;
use aether_oauth::provider::providers::GenericProviderOAuthAdapter;
use aether_oauth::provider::{ProviderOAuthService, ProviderOAuthTransportContext};
use axum::{body::Body, http, response::Response};
use url::form_urlencoded;
use std::sync::Arc;
fn provider_oauth_transport_error_detail(prefix: &str, error: &str) -> String {
let error = error.trim();
@@ -15,6 +16,51 @@ fn provider_oauth_transport_error_detail(prefix: &str, error: &str) -> String {
format!("{prefix}: {error}")
}
fn provider_oauth_exchange_context(
provider_type: &str,
proxy: Option<ProxySnapshot>,
) -> ProviderOAuthTransportContext {
ProviderOAuthTransportContext {
provider_id: String::new(),
provider_type: provider_type.to_string(),
endpoint_id: None,
key_id: None,
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: None,
endpoint_config: None,
key_config: None,
network: aether_oauth::network::OAuthNetworkContext::provider_operation(proxy),
}
}
fn provider_oauth_service_for_template(
template: AdminProviderOAuthTemplate,
token_url: String,
) -> Result<ProviderOAuthService, Response<Body>> {
GenericProviderOAuthAdapter::for_provider_type(template.provider_type)
.map(|adapter| adapter.with_token_url_override(token_url))
.map(|adapter| ProviderOAuthService::new().with_adapter(Arc::new(adapter)))
.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
)
})
}
fn token_payload_from_provider_oauth_result(
result: aether_oauth::provider::ProviderOAuthTokenSet,
) -> Result<serde_json::Value, Response<Body>> {
result.token_set.raw_payload.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
)
})
}
pub(crate) async fn exchange_admin_provider_oauth_code(
state: &AdminAppState<'_>,
template: AdminProviderOAuthTemplate,
@@ -24,122 +70,25 @@ pub(crate) async fn exchange_admin_provider_oauth_code(
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("authorization_code".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"redirect_uri".to_string(),
serde_json::Value::String(template.redirect_uri.to_string()),
),
(
"code".to_string(),
serde_json::Value::String(code.to_string()),
),
(
"state".to_string(),
serde_json::Value::String(state_nonce.to_string()),
),
]);
if let Some(verifier) = pkce_verifier {
body.insert(
"code_verifier".to_string(),
serde_json::Value::String(verifier.to_string()),
);
}
let headers = reqwest::header::HeaderMap::from_iter([
(
reqwest::header::CONTENT_TYPE,
reqwest::header::HeaderValue::from_static("application/json"),
),
(
reqwest::header::ACCEPT,
reqwest::header::HeaderValue::from_static("application/json"),
),
]);
state
.execute_admin_provider_oauth_http_request(
"provider-oauth:exchange-code",
reqwest::Method::POST,
&token_url,
&headers,
Some("application/json"),
Some(serde_json::Value::Object(body)),
None,
proxy.clone(),
)
.await
} else {
let form_body = {
let mut form = form_urlencoded::Serializer::new(String::new());
form.append_pair("grant_type", "authorization_code");
form.append_pair("client_id", template.client_id);
form.append_pair("redirect_uri", template.redirect_uri);
form.append_pair("code", code);
if !template.client_secret.trim().is_empty() {
form.append_pair("client_secret", template.client_secret);
let service = provider_oauth_service_for_template(template, token_url)?;
let ctx = provider_oauth_exchange_context(template.provider_type, proxy);
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let result = service
.exchange_code(&executor, &ctx, code, state_nonce, pkce_verifier)
.await
.map_err(|error| match error {
aether_oauth::core::OAuthError::HttpStatus { .. } => {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 失败",
)
}
if let Some(verifier) = pkce_verifier {
form.append_pair("code_verifier", verifier);
}
form.finish().into_bytes()
};
let headers = reqwest::header::HeaderMap::from_iter([
(
reqwest::header::CONTENT_TYPE,
reqwest::header::HeaderValue::from_static("application/x-www-form-urlencoded"),
error => build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
provider_oauth_transport_error_detail("token exchange 失败", &error.to_string()),
),
(
reqwest::header::ACCEPT,
reqwest::header::HeaderValue::from_static("application/json"),
),
]);
state
.execute_admin_provider_oauth_http_request(
"provider-oauth:exchange-code",
reqwest::Method::POST,
&token_url,
&headers,
Some("application/x-www-form-urlencoded"),
None,
Some(form_body),
proxy.clone(),
)
.await
}
.map_err(|error| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
provider_oauth_transport_error_detail("token exchange 失败", &error),
)
})?;
if !response.status.is_success() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 失败",
));
}
let payload = response.json_body.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
}
Ok(payload)
})?;
token_payload_from_provider_oauth_result(result)
}
pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
@@ -149,119 +98,45 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let scope = template.scopes.join(" ");
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("refresh_token".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"refresh_token".to_string(),
serde_json::Value::String(refresh_token.to_string()),
),
]);
if !scope.trim().is_empty() {
body.insert("scope".to_string(), serde_json::Value::String(scope));
}
let headers = reqwest::header::HeaderMap::from_iter([
(
reqwest::header::CONTENT_TYPE,
reqwest::header::HeaderValue::from_static("application/json"),
),
(
reqwest::header::ACCEPT,
reqwest::header::HeaderValue::from_static("application/json"),
),
]);
state
.execute_admin_provider_oauth_http_request(
"provider-oauth:refresh-token",
reqwest::Method::POST,
&token_url,
&headers,
Some("application/json"),
Some(serde_json::Value::Object(body)),
None,
proxy.clone(),
)
.await
} else {
let form_body = {
let mut form = form_urlencoded::Serializer::new(String::new());
form.append_pair("grant_type", "refresh_token");
form.append_pair("client_id", template.client_id);
form.append_pair("refresh_token", refresh_token);
if !scope.trim().is_empty() {
form.append_pair("scope", &scope);
let service = provider_oauth_service_for_template(template, token_url)?;
let ctx = provider_oauth_exchange_context(template.provider_type, proxy);
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let input = aether_oauth::provider::ProviderOAuthImportInput {
provider_type: template.provider_type.to_string(),
name: None,
refresh_token: Some(refresh_token.to_string()),
raw_credentials: None,
network: ctx.network.clone(),
};
let result = service
.import_credentials(&executor, &ctx, input)
.await
.map_err(|error| match error {
aether_oauth::core::OAuthError::HttpStatus {
status_code,
body_excerpt,
} => {
let reason = normalize_provider_oauth_refresh_error_message(
Some(status_code),
Some(&body_excerpt),
);
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Refresh Token 验证失败: {reason}"),
)
}
if !template.client_secret.trim().is_empty() {
form.append_pair("client_secret", template.client_secret);
}
form.finish().into_bytes()
};
let headers = reqwest::header::HeaderMap::from_iter([
(
reqwest::header::CONTENT_TYPE,
reqwest::header::HeaderValue::from_static("application/x-www-form-urlencoded"),
error => build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
provider_oauth_transport_error_detail(
"Refresh Token 验证失败: token exchange 失败",
&error.to_string(),
),
),
(
reqwest::header::ACCEPT,
reqwest::header::HeaderValue::from_static("application/json"),
),
]);
state
.execute_admin_provider_oauth_http_request(
"provider-oauth:refresh-token",
reqwest::Method::POST,
&token_url,
&headers,
Some("application/x-www-form-urlencoded"),
None,
Some(form_body),
proxy.clone(),
)
.await
}
.map_err(|error| {
})?;
token_payload_from_provider_oauth_result(result).map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
provider_oauth_transport_error_detail(
"Refresh Token 验证失败: token exchange 失败",
&error,
),
)
})?;
let status = response.status;
let body = response.body_text;
if !status.is_success() {
let reason =
normalize_provider_oauth_refresh_error_message(Some(status.as_u16()), Some(&body));
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Refresh Token 验证失败: {reason}"),
));
}
let payload = response
.json_body
.or_else(|| serde_json::from_str::<serde_json::Value>(&body).ok())
.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
));
}
Ok(payload)
)
})
}
@@ -1,4 +1,5 @@
use crate::handlers::admin::request::AdminProviderOAuthTemplate;
use aether_oauth::provider::{ProviderOAuthService, ProviderOAuthTransportContext};
use serde_json::json;
use url::form_urlencoded;
@@ -7,6 +8,48 @@ pub(crate) fn build_provider_oauth_start_response(
nonce: &str,
code_challenge: Option<&str>,
) -> serde_json::Value {
let authorization_url = build_provider_oauth_authorization_url(template, nonce, code_challenge)
.unwrap_or_else(|| {
build_provider_oauth_authorization_url_legacy(template, nonce, code_challenge)
});
json!({
"authorization_url": authorization_url,
"redirect_uri": template.redirect_uri,
"provider_type": template.provider_type,
"instructions": "1) 打开 authorization_url 完成授权\n2) 授权后会跳转到 redirect_uri(localhost)\n3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url",
})
}
fn build_provider_oauth_authorization_url(
template: AdminProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> Option<String> {
let ctx = ProviderOAuthTransportContext {
provider_id: String::new(),
provider_type: template.provider_type.to_string(),
endpoint_id: None,
key_id: None,
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: None,
endpoint_config: None,
key_config: None,
network: aether_oauth::network::OAuthNetworkContext::provider_operation(None),
};
ProviderOAuthService::with_builtin_adapters()
.build_authorize_url(&ctx, nonce, code_challenge)
.ok()
.map(|response| response.authorize_url)
}
fn build_provider_oauth_authorization_url_legacy(
template: AdminProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> String {
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("client_id", template.client_id);
serializer.append_pair("response_type", "code");
@@ -25,10 +68,5 @@ pub(crate) fn build_provider_oauth_start_response(
}
}
json!({
"authorization_url": format!("{}?{}", template.authorize_url, serializer.finish()),
"redirect_uri": template.redirect_uri,
"provider_type": template.provider_type,
"instructions": "1) 打开 authorization_url 完成授权\n2) 授权后会跳转到 redirect_uri(localhost)\n3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url",
})
format!("{}?{}", template.authorize_url, serializer.finish())
}