mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
Add shared OAuth flows
This commit is contained in:
@@ -34,14 +34,24 @@ pub(crate) struct AdminOAuthProviderUpsertRequest {
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_oauth_supported_types_payload() -> Vec<serde_json::Value> {
|
||||
vec![json!({
|
||||
"provider_type": "linuxdo",
|
||||
"display_name": "Linux Do",
|
||||
"default_authorization_url": "https://connect.linux.do/oauth2/authorize",
|
||||
"default_token_url": "https://connect.linux.do/oauth2/token",
|
||||
"default_userinfo_url": "https://connect.linux.do/api/user",
|
||||
"default_scopes": [],
|
||||
})]
|
||||
vec![
|
||||
json!({
|
||||
"provider_type": "linuxdo",
|
||||
"display_name": "Linux Do",
|
||||
"default_authorization_url": "https://connect.linux.do/oauth2/authorize",
|
||||
"default_token_url": "https://connect.linux.do/oauth2/token",
|
||||
"default_userinfo_url": "https://connect.linux.do/api/user",
|
||||
"default_scopes": [],
|
||||
}),
|
||||
json!({
|
||||
"provider_type": "custom_oidc",
|
||||
"display_name": "Custom OIDC",
|
||||
"default_authorization_url": "",
|
||||
"default_token_url": "",
|
||||
"default_userinfo_url": "",
|
||||
"default_scopes": ["openid", "profile", "email"],
|
||||
}),
|
||||
]
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_oauth_provider_payload(
|
||||
@@ -78,10 +88,13 @@ pub(crate) fn admin_oauth_test_provider_type_from_path(request_path: &str) -> Op
|
||||
}
|
||||
|
||||
fn admin_oauth_is_supported_provider(provider_type: &str) -> bool {
|
||||
provider_type.eq_ignore_ascii_case("linuxdo")
|
||||
matches!(
|
||||
provider_type.to_ascii_lowercase().as_str(),
|
||||
"linuxdo" | "custom_oidc"
|
||||
)
|
||||
}
|
||||
|
||||
fn admin_oauth_allowed_domains(provider_type: &str) -> Option<&'static [&'static str]> {
|
||||
fn admin_oauth_builtin_allowed_domains(provider_type: &str) -> Option<&'static [&'static str]> {
|
||||
if provider_type.eq_ignore_ascii_case("linuxdo") {
|
||||
Some(&["linux.do", "connect.linux.do", "connect.linuxdo.org"])
|
||||
} else {
|
||||
@@ -89,6 +102,27 @@ fn admin_oauth_allowed_domains(provider_type: &str) -> Option<&'static [&'static
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_oauth_custom_allowed_domains(extra_config: Option<&serde_json::Value>) -> Vec<String> {
|
||||
extra_config
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|object| {
|
||||
object
|
||||
.get("allowed_domains")
|
||||
.or_else(|| object.get("oauth_allowed_domains"))
|
||||
})
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.map(|items| {
|
||||
items
|
||||
.iter()
|
||||
.filter_map(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| value.trim_end_matches('.').to_ascii_lowercase())
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn validate_admin_oauth_frontend_callback_url(url: &str) -> Result<(), String> {
|
||||
let parsed = Url::parse(url).map_err(|_| "frontend_callback_url 必须是绝对 URL".to_string())?;
|
||||
if !matches!(parsed.scheme(), "http" | "https") {
|
||||
@@ -134,6 +168,17 @@ fn validate_admin_oauth_url_override(url: &str, allowed_domains: &[&str]) -> Res
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_admin_oauth_url_override_for_domains(
|
||||
url: &str,
|
||||
allowed_domains: &[String],
|
||||
) -> Result<(), String> {
|
||||
let allowed = allowed_domains
|
||||
.iter()
|
||||
.map(String::as_str)
|
||||
.collect::<Vec<_>>();
|
||||
validate_admin_oauth_url_override(url, &allowed)
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_oauth_upsert_record(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_type: &str,
|
||||
@@ -163,21 +208,64 @@ pub(super) fn build_admin_oauth_upsert_record(
|
||||
validate_admin_oauth_frontend_callback_url(frontend_callback_url)?;
|
||||
validate_admin_oauth_redirect_uri(redirect_uri)?;
|
||||
|
||||
let allowed_domains = admin_oauth_allowed_domains(provider_type)
|
||||
.ok_or_else(|| "不支持的 provider_type".to_string())?;
|
||||
let is_custom_oidc = provider_type.eq_ignore_ascii_case("custom_oidc");
|
||||
let custom_allowed_domains = if is_custom_oidc {
|
||||
let domains = admin_oauth_custom_allowed_domains(payload.extra_config.as_ref());
|
||||
if domains.is_empty() {
|
||||
return Err(
|
||||
"custom_oidc 必须在 extra_config.allowed_domains 配置域名白名单".to_string(),
|
||||
);
|
||||
}
|
||||
domains
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
let builtin_allowed_domains = admin_oauth_builtin_allowed_domains(provider_type);
|
||||
|
||||
if is_custom_oidc {
|
||||
for (field_name, value) in [
|
||||
(
|
||||
"authorization_url_override",
|
||||
payload.authorization_url_override.as_deref(),
|
||||
),
|
||||
("token_url_override", payload.token_url_override.as_deref()),
|
||||
(
|
||||
"userinfo_url_override",
|
||||
payload.userinfo_url_override.as_deref(),
|
||||
),
|
||||
] {
|
||||
let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||
return Err(format!("custom_oidc 必须配置 {field_name}"));
|
||||
};
|
||||
validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(value) = payload.authorization_url_override.as_deref().map(str::trim) {
|
||||
if !value.is_empty() {
|
||||
validate_admin_oauth_url_override(value, allowed_domains)?;
|
||||
if let Some(allowed_domains) = builtin_allowed_domains {
|
||||
validate_admin_oauth_url_override(value, allowed_domains)?;
|
||||
} else {
|
||||
validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?;
|
||||
}
|
||||
}
|
||||
}
|
||||
if let Some(value) = payload.token_url_override.as_deref().map(str::trim) {
|
||||
if !value.is_empty() {
|
||||
validate_admin_oauth_url_override(value, allowed_domains)?;
|
||||
if let Some(allowed_domains) = builtin_allowed_domains {
|
||||
validate_admin_oauth_url_override(value, allowed_domains)?;
|
||||
} else {
|
||||
validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?;
|
||||
}
|
||||
}
|
||||
}
|
||||
if let Some(value) = payload.userinfo_url_override.as_deref().map(str::trim) {
|
||||
if !value.is_empty() {
|
||||
validate_admin_oauth_url_override(value, allowed_domains)?;
|
||||
if let Some(allowed_domains) = builtin_allowed_domains {
|
||||
validate_admin_oauth_url_override(value, allowed_domains)?;
|
||||
} else {
|
||||
validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -20,11 +20,17 @@ pub(crate) use self::observability::{
|
||||
admin_stats_bad_request_response, maybe_build_local_admin_usage_response, parse_bounded_u32,
|
||||
round_to, AdminStatsTimeRange, AdminStatsUsageFilter,
|
||||
};
|
||||
pub(crate) use self::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
|
||||
pub(crate) use self::provider::oauth::errors::build_internal_control_error_response;
|
||||
pub(crate) use self::provider::oauth::provisioning::{
|
||||
create_provider_oauth_catalog_key, update_existing_provider_oauth_catalog_key,
|
||||
};
|
||||
pub(crate) use self::provider::oauth::quota::antigravity::refresh_antigravity_provider_quota_locally;
|
||||
pub(crate) use self::provider::oauth::quota::codex::refresh_codex_provider_quota_locally;
|
||||
pub(crate) use self::provider::oauth::quota::kiro::refresh_kiro_provider_quota_locally;
|
||||
pub(crate) use self::provider::oauth::runtime::provider_oauth_runtime_endpoint_for_provider;
|
||||
pub(crate) use self::provider::oauth::runtime::{
|
||||
provider_oauth_runtime_endpoint_for_provider, refresh_provider_oauth_account_state_after_update,
|
||||
};
|
||||
pub(crate) use self::provider::ops::providers::actions::admin_provider_ops_local_action_response;
|
||||
pub(crate) use self::provider::pool::config::admin_provider_pool_config;
|
||||
pub(crate) use self::provider::pool_admin::maybe_build_local_admin_pool_response;
|
||||
@@ -32,7 +38,8 @@ pub(crate) use self::provider::{
|
||||
maybe_build_local_admin_provider_oauth_response, maybe_build_local_admin_providers_response,
|
||||
};
|
||||
pub(crate) use self::request::{
|
||||
AdminAppState, AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult,
|
||||
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminLocalOAuthRefreshError,
|
||||
AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult,
|
||||
};
|
||||
pub(crate) use self::routes::maybe_build_local_admin_response;
|
||||
#[cfg(test)]
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -37,23 +37,24 @@ impl<'a> AdminAppState<'a> {
|
||||
encrypted_auth_config: Option<&str>,
|
||||
expires_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app
|
||||
.update_provider_catalog_key_oauth_credentials(
|
||||
key_id,
|
||||
encrypted_api_key,
|
||||
encrypted_auth_config,
|
||||
expires_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
crate::oauth::ProviderOAuthRepository::update_provider_catalog_key_oauth_credentials(
|
||||
self,
|
||||
key_id,
|
||||
encrypted_api_key,
|
||||
encrypted_auth_config,
|
||||
expires_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn clear_provider_catalog_key_oauth_invalid_marker(
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(key_id)
|
||||
.await
|
||||
crate::oauth::ProviderOAuthRepository::clear_provider_catalog_key_oauth_invalid_marker(
|
||||
self, key_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn force_local_oauth_refresh_entry(
|
||||
@@ -61,7 +62,8 @@ impl<'a> AdminAppState<'a> {
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
) -> Result<Option<crate::provider_transport::CachedOAuthEntry>, AdminLocalOAuthRefreshError>
|
||||
{
|
||||
self.app.force_local_oauth_refresh_entry(transport).await
|
||||
crate::oauth::ProviderOAuthRepository::force_local_oauth_refresh_entry(self, transport)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn save_provider_oauth_state(
|
||||
@@ -425,23 +427,12 @@ impl<'a> AdminAppState<'a> {
|
||||
temporary_proxy_node_id: Option<&str>,
|
||||
configured_proxies: &[Option<&serde_json::Value>],
|
||||
) -> Option<ProxySnapshot> {
|
||||
if let Some(snapshot) = self
|
||||
.resolve_admin_proxy_node_snapshot(temporary_proxy_node_id)
|
||||
.await
|
||||
{
|
||||
return Some(snapshot);
|
||||
}
|
||||
|
||||
for proxy in configured_proxies {
|
||||
if let Some(snapshot) = self
|
||||
.app
|
||||
.resolve_configured_proxy_snapshot_with_tunnel_affinity(*proxy)
|
||||
.await
|
||||
{
|
||||
return Some(snapshot);
|
||||
}
|
||||
}
|
||||
self.app.resolve_system_proxy_snapshot().await
|
||||
crate::oauth::resolve_provider_oauth_operation_proxy_snapshot(
|
||||
self,
|
||||
temporary_proxy_node_id,
|
||||
configured_proxies,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
@@ -453,7 +444,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
String,
|
||||
> {
|
||||
crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key(
|
||||
crate::oauth::ProviderOAuthRepository::find_duplicate_provider_oauth_key(
|
||||
self,
|
||||
provider_id,
|
||||
auth_config,
|
||||
@@ -476,7 +467,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
GatewayError,
|
||||
> {
|
||||
crate::handlers::admin::provider::oauth::provisioning::create_provider_oauth_catalog_key(
|
||||
crate::oauth::ProviderOAuthRepository::create_provider_oauth_catalog_key(
|
||||
self,
|
||||
provider_id,
|
||||
provider_type,
|
||||
@@ -503,7 +494,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
GatewayError,
|
||||
> {
|
||||
crate::handlers::admin::provider::oauth::provisioning::update_existing_provider_oauth_catalog_key(
|
||||
crate::oauth::ProviderOAuthRepository::update_existing_provider_oauth_catalog_key(
|
||||
self,
|
||||
existing_key,
|
||||
provider_type,
|
||||
@@ -522,7 +513,7 @@ impl<'a> AdminAppState<'a> {
|
||||
key_id: &str,
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
) -> Result<(bool, Option<String>), GatewayError> {
|
||||
crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update(
|
||||
crate::oauth::ProviderOAuthRepository::refresh_provider_oauth_account_state_after_update(
|
||||
self,
|
||||
provider,
|
||||
key_id,
|
||||
@@ -618,56 +609,31 @@ impl<'a> AdminAppState<'a> {
|
||||
body_bytes: Option<Vec<u8>>,
|
||||
proxy: Option<ProxySnapshot>,
|
||||
) -> Result<AdminProviderOAuthHttpResponse, String> {
|
||||
let body = if let Some(json_body) = json_body {
|
||||
RequestBody::from_json(json_body)
|
||||
} else {
|
||||
RequestBody {
|
||||
json_body: None,
|
||||
body_bytes_b64: body_bytes.map(|bytes| STANDARD.encode(bytes)),
|
||||
body_ref: None,
|
||||
}
|
||||
};
|
||||
let timeout_ms = admin_provider_oauth_timeout_ms(proxy.as_ref());
|
||||
let plan = ExecutionPlan {
|
||||
let network = aether_oauth::network::OAuthNetworkContext::provider_operation(proxy);
|
||||
let request = aether_oauth::network::OAuthHttpRequest {
|
||||
request_id: request_id.to_string(),
|
||||
candidate_id: None,
|
||||
provider_name: Some("provider_oauth".to_string()),
|
||||
provider_id: String::new(),
|
||||
endpoint_id: String::new(),
|
||||
key_id: String::new(),
|
||||
method: method.as_str().to_string(),
|
||||
method,
|
||||
url: url.to_string(),
|
||||
headers: admin_provider_oauth_execution_headers(headers),
|
||||
content_type: content_type
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
content_encoding: None,
|
||||
body,
|
||||
stream: false,
|
||||
client_api_format: "provider_oauth:exchange".to_string(),
|
||||
provider_api_format: "provider_oauth:exchange".to_string(),
|
||||
model_name: Some("oauth-exchange".to_string()),
|
||||
proxy,
|
||||
tls_profile: None,
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
connect_ms: Some(timeout_ms),
|
||||
read_ms: Some(timeout_ms),
|
||||
write_ms: Some(timeout_ms),
|
||||
pool_ms: Some(timeout_ms),
|
||||
total_ms: Some(timeout_ms),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
json_body,
|
||||
body_bytes,
|
||||
network,
|
||||
};
|
||||
let result = self
|
||||
.execute_execution_runtime_sync_plan(None, &plan)
|
||||
.await
|
||||
.map_err(admin_provider_oauth_gateway_error_message)?;
|
||||
let response = aether_oauth::network::OAuthHttpExecutor::execute(
|
||||
&crate::oauth::GatewayOAuthHttpExecutor::new(*self),
|
||||
request,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| err.to_string())?;
|
||||
Ok(AdminProviderOAuthHttpResponse {
|
||||
status: http::StatusCode::from_u16(result.status_code)
|
||||
status: http::StatusCode::from_u16(response.status_code)
|
||||
.unwrap_or(http::StatusCode::BAD_GATEWAY),
|
||||
body_text: admin_provider_oauth_execution_body_text(&result),
|
||||
json_body: admin_provider_oauth_execution_json_body(&result),
|
||||
body_text: response.body_text,
|
||||
json_body: response.json_body,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user