mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
Add shared OAuth flows
This commit is contained in:
@@ -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())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user