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

View File

@@ -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)?;
}
}
}

View File

@@ -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)]

View File

@@ -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,

View File

@@ -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)
)
})
}

View File

@@ -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_urilocalhost\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_urilocalhost\n3) 复制浏览器地址栏完整 URL调用 complete 接口粘贴 callback_url",
})
format!("{}?{}", template.authorize_url, serializer.finish())
}

View File

@@ -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,
})
}
}

View File

@@ -32,6 +32,8 @@ mod support_dashboard;
mod support_models;
#[path = "support/monitoring.rs"]
mod support_monitoring;
#[path = "support/oauth.rs"]
mod support_oauth;
#[path = "support/payment.rs"]
mod support_payment;
#[path = "support/test_connection.rs"]
@@ -56,13 +58,14 @@ use self::support_auth::auth_session::{
};
use self::support_auth::{
build_auth_error_response, build_auth_json_response, build_auth_registration_settings_payload,
build_auth_settings_payload, maybe_build_local_auth_response,
build_auth_settings_payload, extract_client_device_id, maybe_build_local_auth_response,
};
use self::support_dashboard::maybe_build_local_dashboard_response;
use self::support_models::{
build_models_auth_error_response, maybe_build_local_models_response, models_api_format,
};
use self::support_monitoring::maybe_build_local_user_monitoring_response;
use self::support_oauth::maybe_build_local_oauth_response;
use self::support_payment::maybe_build_local_payment_callback_response;
use self::support_test_connection::maybe_build_local_test_connection_response;
use self::support_user_me::maybe_build_local_users_me_response;
@@ -106,6 +109,11 @@ pub(crate) async fn maybe_build_local_public_support_response(
.await;
}
if decision.route_family.as_deref() == Some("oauth") {
return maybe_build_local_oauth_response(state, request_context, headers, request_body)
.await;
}
if decision.route_family.as_deref() == Some("dashboard") {
return Some(maybe_build_local_dashboard_response(state, request_context, headers).await);
}

View File

@@ -236,7 +236,7 @@ pub(super) fn extract_cookie_value(headers: &http::HeaderMap, cookie_name: &str)
None
}
pub(super) fn extract_client_device_id(
pub(crate) fn extract_client_device_id(
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Result<String, Response<Body>> {

View File

@@ -22,7 +22,7 @@ fn base64url_decode(value: &str) -> Result<Vec<u8>, String> {
.map_err(|_| "无效的Token".to_string())
}
pub(super) fn create_auth_token(
pub(crate) fn create_auth_token(
token_type: &str,
mut payload: serde_json::Map<String, serde_json::Value>,
expires_at: chrono::DateTime<chrono::Utc>,
@@ -54,7 +54,7 @@ pub(super) fn create_auth_token(
))
}
pub(super) fn decode_auth_token(
pub(crate) fn decode_auth_token(
token: &str,
expected_type: &str,
) -> Result<serde_json::Map<String, serde_json::Value>, String> {
@@ -525,7 +525,7 @@ pub(super) async fn handle_auth_refresh(
)
}
pub(super) async fn build_auth_login_success_response(
pub(crate) async fn build_auth_login_success_response(
state: &AppState,
headers: &http::HeaderMap,
client_device_id: String,

View File

@@ -0,0 +1,728 @@
use super::support_auth::auth_session::{
build_auth_login_success_response, create_auth_token, decode_auth_token,
};
use super::{
build_auth_error_response, build_auth_json_response, extract_client_device_id, http, json,
query_param_value, resolve_authenticated_local_user, AppState, Body, Bytes,
GatewayPublicRequestContext, IntoResponse, Json, Response,
};
use aether_oauth::core::{generate_pkce_verifier, pkce_s256, OAuthError};
use aether_oauth::identity::{
IdentityClaims, IdentityOAuthExchangeContext, IdentityOAuthService, IdentityOAuthStartContext,
};
use axum::body::to_bytes;
use axum::http::header::{LOCATION, SET_COOKIE};
use axum::http::HeaderValue;
use url::form_urlencoded;
pub(super) async fn maybe_build_local_oauth_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
_request_body: Option<&Bytes>,
) -> Option<Response<Body>> {
let decision = request_context.control_decision.as_ref()?;
if decision.route_family.as_deref() != Some("oauth") {
return None;
}
match decision.route_kind.as_deref() {
Some("list_providers")
if request_context.request_method == http::Method::GET
&& request_context.request_path == "/api/oauth/providers" =>
{
Some(handle_oauth_list_providers(state).await)
}
Some("authorize") if request_context.request_method == http::Method::GET => {
Some(handle_oauth_authorize(state, request_context, headers).await)
}
Some("callback") if request_context.request_method == http::Method::GET => {
Some(handle_oauth_callback(state, request_context, headers).await)
}
Some("bindable_providers")
if request_context.request_method == http::Method::GET
&& request_context.request_path == "/api/user/oauth/bindable-providers" =>
{
Some(handle_oauth_bindable_providers(state, request_context, headers).await)
}
Some("links")
if request_context.request_method == http::Method::GET
&& request_context.request_path == "/api/user/oauth/links" =>
{
Some(handle_oauth_links(state, request_context, headers).await)
}
Some("bind_token") if request_context.request_method == http::Method::POST => {
Some(handle_oauth_bind_token(state, request_context, headers).await)
}
Some("bind") if request_context.request_method == http::Method::GET => {
Some(handle_oauth_bind_start(state, request_context, headers).await)
}
Some("unbind") if request_context.request_method == http::Method::DELETE => {
Some(handle_oauth_unbind(state, request_context, headers).await)
}
_ => Some(super::build_unhandled_public_support_response(
request_context,
)),
}
}
async fn handle_oauth_list_providers(state: &AppState) -> Response<Body> {
match crate::oauth::list_enabled_identity_oauth_providers(state).await {
Ok(providers) => Json(json!({ "providers": providers })).into_response(),
Err(err) => build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("oauth provider lookup failed: {err:?}"),
false,
),
}
}
async fn handle_oauth_bindable_providers(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Response<Body> {
let auth = match resolve_authenticated_local_user(state, request_context, headers).await {
Ok(value) => value,
Err(response) => return response,
};
if auth.user.auth_source.eq_ignore_ascii_case("ldap") {
return Json(json!({ "providers": [] })).into_response();
}
match crate::oauth::list_bindable_identity_oauth_providers(state, &auth.user.id).await {
Ok(providers) => Json(json!({ "providers": providers })).into_response(),
Err(err) => build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("oauth provider lookup failed: {err:?}"),
false,
),
}
}
async fn handle_oauth_links(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Response<Body> {
let auth = match resolve_authenticated_local_user(state, request_context, headers).await {
Ok(value) => value,
Err(response) => return response,
};
match crate::oauth::list_identity_oauth_links(state, &auth.user.id).await {
Ok(links) => Json(json!({ "links": links })).into_response(),
Err(err) => build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("oauth link lookup failed: {err:?}"),
false,
),
}
}
async fn handle_oauth_authorize(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Response<Body> {
let Some(provider_type) =
public_oauth_provider_from_path(&request_context.request_path, "authorize")
else {
return build_auth_error_response(
http::StatusCode::NOT_FOUND,
"OAuth Provider 不存在",
false,
);
};
let client_device_id = match extract_client_device_id(request_context, headers) {
Ok(value) => value,
Err(response) => return response,
};
start_identity_oauth(
state,
&provider_type,
client_device_id,
crate::oauth::IdentityOAuthStateMode::Login,
None,
None,
)
.await
}
async fn handle_oauth_bind_token(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Response<Body> {
let Some(provider_type) =
user_oauth_provider_from_path(&request_context.request_path, "bind-token")
else {
return build_auth_error_response(
http::StatusCode::NOT_FOUND,
"OAuth Provider 不存在",
false,
);
};
let auth = match resolve_authenticated_local_user(state, request_context, headers).await {
Ok(value) => value,
Err(response) => return response,
};
if auth.user.auth_source.eq_ignore_ascii_case("ldap") {
return build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"LDAP 用户不支持 OAuth 绑定",
false,
);
}
match crate::oauth::get_enabled_identity_oauth_provider_config(state, &provider_type).await {
Ok(Some(_)) => {}
Ok(None) => {
return build_auth_error_response(
http::StatusCode::NOT_FOUND,
"OAuth Provider 不存在或已禁用",
false,
)
}
Err(err) => return oauth_account_error_response(err),
}
let token = match create_auth_token(
"oauth_bind",
serde_json::Map::from_iter([
("user_id".to_string(), json!(auth.user.id)),
("session_id".to_string(), json!(auth.session_id)),
("provider_type".to_string(), json!(provider_type)),
]),
chrono::Utc::now() + chrono::Duration::minutes(10),
) {
Ok(value) => value,
Err(detail) => {
return build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
detail,
false,
)
}
};
build_auth_json_response(http::StatusCode::OK, json!({ "bind_token": token }), None)
}
async fn handle_oauth_bind_start(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Response<Body> {
let Some(provider_type) = user_oauth_provider_from_path(&request_context.request_path, "bind")
else {
return build_auth_error_response(
http::StatusCode::NOT_FOUND,
"OAuth Provider 不存在",
false,
);
};
let client_device_id = match extract_client_device_id(request_context, headers) {
Ok(value) => value,
Err(response) => return response,
};
let Some(bind_token) = query_param_value(
request_context.request_query_string.as_deref(),
"bind_token",
) else {
return build_auth_error_response(http::StatusCode::BAD_REQUEST, "缺少绑定令牌", false);
};
let bind =
match validate_bind_token(state, &provider_type, &client_device_id, &bind_token).await {
Ok(value) => value,
Err(response) => return response,
};
start_identity_oauth(
state,
&provider_type,
client_device_id,
crate::oauth::IdentityOAuthStateMode::Bind,
Some(bind.user_id),
Some(bind.session_id),
)
.await
}
async fn handle_oauth_callback(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Response<Body> {
let Some(provider_type) =
public_oauth_provider_from_path(&request_context.request_path, "callback")
else {
return redirect_oauth_error(None, "provider_unavailable");
};
let params = callback_params(request_context);
if params
.get("error")
.is_some_and(|value| value.eq_ignore_ascii_case("access_denied"))
{
return redirect_oauth_error(None, "authorization_denied");
}
let Some(code) = params
.get("code")
.map(String::as_str)
.filter(|value| !value.is_empty())
else {
return redirect_oauth_error(None, "invalid_callback");
};
let Some(nonce) = params
.get("state")
.map(String::as_str)
.filter(|value| !value.is_empty())
else {
return redirect_oauth_error(None, "invalid_state");
};
let stored = match crate::oauth::consume_identity_oauth_state(state, nonce).await {
Ok(Some(value)) => value,
Ok(None) => return redirect_oauth_error(None, "invalid_state"),
Err(_) => return redirect_oauth_error(None, "invalid_state"),
};
if stored.provider_type != provider_type {
return redirect_oauth_error(None, "invalid_state");
}
let config =
match crate::oauth::get_enabled_identity_oauth_provider_config(state, &provider_type).await
{
Ok(Some(value)) => value,
Ok(None) => return redirect_oauth_error(None, "provider_disabled"),
Err(err) => return redirect_oauth_error(None, err.code()),
};
let network = crate::oauth::resolve_identity_oauth_network_context(state).await;
let exchange_ctx = IdentityOAuthExchangeContext {
code: code.to_string(),
state: nonce.to_string(),
pkce_verifier: stored.pkce_verifier.clone(),
network,
};
let executor = crate::oauth::GatewayOAuthHttpExecutor::from_app(state);
let service = IdentityOAuthService::with_builtin_providers();
let claims = match service.login(&executor, &config, &exchange_ctx).await {
Ok(outcome) => outcome.claims,
Err(err) => {
return redirect_oauth_error(
Some(&config.frontend_callback_url),
oauth_error_code(&err),
)
}
};
match stored.mode {
crate::oauth::IdentityOAuthStateMode::Login => {
complete_oauth_login(
state,
headers,
&config.frontend_callback_url,
stored.client_device_id,
claims,
)
.await
}
crate::oauth::IdentityOAuthStateMode::Bind => {
complete_oauth_bind(state, &config.frontend_callback_url, stored, claims).await
}
}
}
async fn handle_oauth_unbind(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
) -> Response<Body> {
let Some(provider_type) =
user_oauth_provider_from_path_without_suffix(&request_context.request_path)
else {
return build_auth_error_response(
http::StatusCode::NOT_FOUND,
"OAuth Provider 不存在",
false,
);
};
let auth = match resolve_authenticated_local_user(state, request_context, headers).await {
Ok(value) => value,
Err(response) => return response,
};
match crate::oauth::unbind_identity_oauth(state, &auth.user, &provider_type).await {
Ok(true) => Json(json!({ "message": "解绑成功" })).into_response(),
Ok(false) => {
build_auth_error_response(http::StatusCode::NOT_FOUND, "OAuth 绑定不存在", false)
}
Err(err) => oauth_account_error_response(err),
}
}
async fn start_identity_oauth(
state: &AppState,
provider_type: &str,
client_device_id: String,
mode: crate::oauth::IdentityOAuthStateMode,
bind_user_id: Option<String>,
bind_session_id: Option<String>,
) -> Response<Body> {
let config = match crate::oauth::get_enabled_identity_oauth_provider_config(
state,
provider_type,
)
.await
{
Ok(Some(value)) => value,
Ok(None) => {
return build_auth_error_response(
http::StatusCode::NOT_FOUND,
"OAuth Provider 不存在或已禁用",
false,
)
}
Err(err) => return oauth_account_error_response(err),
};
let pkce_verifier = generate_pkce_verifier();
let code_challenge = pkce_s256(&pkce_verifier);
let stored = match mode {
crate::oauth::IdentityOAuthStateMode::Login => {
crate::oauth::StoredIdentityOAuthState::login(
provider_type,
client_device_id,
Some(pkce_verifier),
)
}
crate::oauth::IdentityOAuthStateMode::Bind => {
let Some(user_id) = bind_user_id else {
return build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"缺少绑定用户",
false,
);
};
let Some(session_id) = bind_session_id else {
return build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"缺少绑定会话",
false,
);
};
crate::oauth::StoredIdentityOAuthState::bind(
provider_type,
client_device_id,
Some(pkce_verifier),
user_id,
session_id,
)
}
};
if crate::oauth::save_identity_oauth_state(state, &stored)
.await
.is_err()
{
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"OAuth 状态存储不可用",
false,
);
}
let network = crate::oauth::resolve_identity_oauth_network_context(state).await;
let start_ctx = IdentityOAuthStartContext {
state: stored.nonce,
code_challenge: Some(code_challenge),
network,
};
let authorize = match IdentityOAuthService::with_builtin_providers().start(&config, &start_ctx)
{
Ok(value) => value,
Err(_) => {
return build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"OAuth Provider 不可用",
false,
)
}
};
redirect_to(&authorize.authorize_url, None)
}
async fn complete_oauth_login(
state: &AppState,
headers: &http::HeaderMap,
frontend_callback_url: &str,
client_device_id: String,
claims: IdentityClaims,
) -> Response<Body> {
let user = match crate::oauth::resolve_identity_oauth_login_user(state, &claims).await {
Ok(user) if user.is_active && !user.is_deleted => user,
Ok(_) => return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"),
Err(err) => return redirect_oauth_error(Some(frontend_callback_url), err.code()),
};
let login_response =
build_auth_login_success_response(state, headers, client_device_id, user).await;
if login_response.status() != http::StatusCode::OK {
return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable");
}
let set_cookies = login_response
.headers()
.get_all(SET_COOKIE)
.iter()
.cloned()
.collect::<Vec<_>>();
let body = login_response.into_body();
let body = match to_bytes(body, usize::MAX).await {
Ok(value) => value,
Err(_) => return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"),
};
let payload = match serde_json::from_slice::<serde_json::Value>(&body) {
Ok(value) => value,
Err(_) => return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"),
};
let Some(access_token) = payload
.get("access_token")
.and_then(serde_json::Value::as_str)
else {
return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable");
};
let expires_in = payload
.get("expires_in")
.and_then(serde_json::Value::as_i64)
.unwrap_or(24 * 60 * 60)
.to_string();
let mut response = redirect_to(
frontend_callback_url,
Some(RedirectParams::Fragment(vec![
("access_token", access_token.to_string()),
("token_type", "bearer".to_string()),
("expires_in", expires_in),
])),
);
for cookie in set_cookies {
response.headers_mut().append(SET_COOKIE, cookie);
}
response
}
async fn complete_oauth_bind(
state: &AppState,
frontend_callback_url: &str,
stored: crate::oauth::StoredIdentityOAuthState,
claims: IdentityClaims,
) -> Response<Body> {
let Some(user_id) = stored.bind_user_id.as_deref() else {
return redirect_oauth_error(Some(frontend_callback_url), "invalid_state");
};
let Some(session_id) = stored.bind_session_id.as_deref() else {
return redirect_oauth_error(Some(frontend_callback_url), "invalid_state");
};
let user = match state.find_user_auth_by_id(user_id).await {
Ok(Some(user)) if user.is_active && !user.is_deleted => user,
_ => return redirect_oauth_error(Some(frontend_callback_url), "invalid_state"),
};
let session = match state.find_user_session(user_id, session_id).await {
Ok(Some(session)) => session,
_ => return redirect_oauth_error(Some(frontend_callback_url), "invalid_state"),
};
let now = chrono::Utc::now();
if session.is_revoked()
|| session.is_expired(now)
|| session.client_device_id != stored.client_device_id
{
return redirect_oauth_error(Some(frontend_callback_url), "invalid_state");
}
if let Err(err) = crate::oauth::bind_identity_oauth_to_user(state, &user, &claims).await {
return redirect_oauth_error(Some(frontend_callback_url), err.code());
}
redirect_to(
frontend_callback_url,
Some(RedirectParams::Query(vec![(
"oauth_bound",
claims.provider_type,
)])),
)
}
#[derive(Debug, Clone)]
struct ValidatedBindToken {
user_id: String,
session_id: String,
}
async fn validate_bind_token(
state: &AppState,
provider_type: &str,
client_device_id: &str,
bind_token: &str,
) -> Result<ValidatedBindToken, Response<Body>> {
let payload = decode_auth_token(bind_token, "oauth_bind").map_err(|detail| {
build_auth_error_response(http::StatusCode::UNAUTHORIZED, detail, false)
})?;
let token_provider = payload
.get("provider_type")
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
if token_provider != provider_type {
return Err(build_auth_error_response(
http::StatusCode::UNAUTHORIZED,
"绑定令牌不匹配",
false,
));
}
let Some(user_id) = payload.get("user_id").and_then(serde_json::Value::as_str) else {
return Err(build_auth_error_response(
http::StatusCode::UNAUTHORIZED,
"绑定令牌无效",
false,
));
};
let Some(session_id) = payload
.get("session_id")
.and_then(serde_json::Value::as_str)
else {
return Err(build_auth_error_response(
http::StatusCode::UNAUTHORIZED,
"绑定令牌无效",
false,
));
};
let session = state
.find_user_session(user_id, session_id)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("auth session lookup failed: {err:?}"),
false,
)
})?
.ok_or_else(|| {
build_auth_error_response(http::StatusCode::UNAUTHORIZED, "绑定会话已失效", false)
})?;
if session.is_revoked()
|| session.is_expired(chrono::Utc::now())
|| session.client_device_id != client_device_id
{
return Err(build_auth_error_response(
http::StatusCode::UNAUTHORIZED,
"绑定会话已失效",
false,
));
}
Ok(ValidatedBindToken {
user_id: user_id.to_string(),
session_id: session_id.to_string(),
})
}
fn callback_params(
request_context: &GatewayPublicRequestContext,
) -> std::collections::BTreeMap<String, String> {
request_context
.request_query_string
.as_deref()
.map(|query| {
form_urlencoded::parse(query.as_bytes())
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect()
})
.unwrap_or_default()
}
fn public_oauth_provider_from_path(path: &str, suffix: &str) -> Option<String> {
path.strip_prefix("/api/oauth/")?
.strip_suffix(&format!("/{suffix}"))?
.split('/')
.next()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_ascii_lowercase)
}
fn user_oauth_provider_from_path(path: &str, suffix: &str) -> Option<String> {
path.strip_prefix("/api/user/oauth/")?
.strip_suffix(&format!("/{suffix}"))?
.split('/')
.next()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_ascii_lowercase)
}
fn user_oauth_provider_from_path_without_suffix(path: &str) -> Option<String> {
let provider_type = path.strip_prefix("/api/user/oauth/")?;
(!provider_type.is_empty() && !provider_type.contains('/'))
.then(|| provider_type.trim().to_ascii_lowercase())
}
fn oauth_error_code(error: &OAuthError) -> &'static str {
match error {
OAuthError::InvalidState => "invalid_state",
OAuthError::UnsupportedProvider(_) | OAuthError::InvalidRequest(_) => {
"provider_unavailable"
}
OAuthError::HttpStatus { .. }
| OAuthError::InvalidResponse(_)
| OAuthError::Transport(_) => "token_exchange_failed",
OAuthError::Storage(_) | OAuthError::EncryptionUnavailable => "provider_unavailable",
}
}
fn oauth_account_error_response(error: crate::oauth::IdentityOAuthAccountError) -> Response<Body> {
let status = match error {
crate::oauth::IdentityOAuthAccountError::ProviderUnavailable
| crate::oauth::IdentityOAuthAccountError::Storage(_) => {
http::StatusCode::SERVICE_UNAVAILABLE
}
crate::oauth::IdentityOAuthAccountError::OAuthAlreadyBound
| crate::oauth::IdentityOAuthAccountError::AlreadyBoundProvider
| crate::oauth::IdentityOAuthAccountError::LastOAuthBinding
| crate::oauth::IdentityOAuthAccountError::LastLoginMethod => http::StatusCode::CONFLICT,
_ => http::StatusCode::BAD_REQUEST,
};
build_auth_error_response(status, error.detail(), false)
}
enum RedirectParams {
Query(Vec<(&'static str, String)>),
Fragment(Vec<(&'static str, String)>),
}
fn redirect_oauth_error(frontend_callback_url: Option<&str>, code: &str) -> Response<Body> {
redirect_to(
frontend_callback_url.unwrap_or("/auth/callback"),
Some(RedirectParams::Query(vec![(
"error_code",
code.to_string(),
)])),
)
}
fn redirect_to(target: &str, params: Option<RedirectParams>) -> Response<Body> {
let location = build_redirect_location(target, params);
let mut response = Response::new(Body::empty());
*response.status_mut() = http::StatusCode::FOUND;
if let Ok(value) = HeaderValue::from_str(&location) {
response.headers_mut().insert(LOCATION, value);
}
response
}
fn build_redirect_location(target: &str, params: Option<RedirectParams>) -> String {
let Ok(mut url) = url::Url::parse(target) else {
return target.to_string();
};
match params {
Some(RedirectParams::Query(items)) => {
{
let mut query = url.query_pairs_mut();
for (key, value) in items {
query.append_pair(key, &value);
}
}
url.to_string()
}
Some(RedirectParams::Fragment(items)) => {
let mut serializer = form_urlencoded::Serializer::new(String::new());
for (key, value) in items {
serializer.append_pair(key, &value);
}
url.set_fragment(Some(&serializer.finish()));
url.to_string()
}
None => url.to_string(),
}
}