use std::{collections::HashMap, sync::Arc};
use axum::{
Form, Json,
extract::{Query, State},
http::StatusCode,
response::{IntoResponse, Redirect, Response},
};
use serde::{Deserialize, Serialize};
use url::Url;
use crate::{
account_linking::{AccountStore, TrustedEmailProviders},
audit::logger::{AuditEventType, SecretType, get_audit_logger},
handlers::generate_secure_state,
provider::OAuthProvider,
session::SessionStore,
state_store::StateStore,
};
const MAX_REDIRECT_URI_BYTES: usize = 2_048;
const MAX_PROVIDER_NAME_BYTES: usize = 128;
const STATE_VALUE_SEPARATOR: char = '\n';
#[must_use]
pub fn is_redirect_uri_allowed(candidate: &str, allowlist: &[String]) -> bool {
let Ok(candidate_url) = Url::parse(candidate) else {
return false;
};
allowlist.iter().any(|entry| {
Url::parse(entry).is_ok_and(|entry_url| redirect_uri_matches(&candidate_url, &entry_url))
})
}
fn redirect_uri_matches(candidate: &Url, entry: &Url) -> bool {
if candidate.scheme() != entry.scheme()
|| candidate.host_str() != entry.host_str()
|| candidate.port_or_known_default() != entry.port_or_known_default()
{
return false;
}
let (candidate_path, entry_path) = (candidate.path(), entry.path());
candidate_path == entry_path
|| candidate_path
.strip_prefix(entry_path)
.is_some_and(|rest| entry_path.ends_with('/') || rest.starts_with('/'))
}
fn encode_state_value(provider: &str, redirect_uri: Option<&str>) -> String {
match redirect_uri {
Some(uri) => format!("{provider}{STATE_VALUE_SEPARATOR}{uri}"),
None => provider.to_string(),
}
}
fn decode_state_value(value: &str) -> (String, Option<String>) {
match value.split_once(STATE_VALUE_SEPARATOR) {
Some((provider, uri)) => (provider.to_string(), Some(uri.to_string())),
None => (value.to_string(), None),
}
}
fn build_redirect_with_tokens(
redirect_uri: &str,
access_token: &str,
refresh_token: &str,
expires_in: u64,
provider: &str,
) -> String {
format!(
"{redirect_uri}#access_token={}&token_type=Bearer&expires_in={expires_in}&refresh_token={}&provider={}",
urlencoding::encode(access_token),
urlencoding::encode(refresh_token),
urlencoding::encode(provider),
)
}
#[derive(Clone)]
pub struct MultiProviderAuthState {
providers: HashMap<String, Arc<dyn OAuthProvider>>,
state_store: Arc<dyn StateStore>,
session_store: Arc<dyn SessionStore>,
user_store: Option<Arc<dyn AccountStore>>,
trusted_email_providers: TrustedEmailProviders,
redirect_uri_allowlist: Vec<String>,
}
impl MultiProviderAuthState {
pub fn new(state_store: Arc<dyn StateStore>, session_store: Arc<dyn SessionStore>) -> Self {
Self {
providers: HashMap::new(),
state_store,
session_store,
user_store: None,
trusted_email_providers: TrustedEmailProviders::default(),
redirect_uri_allowlist: Vec::new(),
}
}
pub fn with_user_store(mut self, user_store: Arc<dyn AccountStore>) -> Self {
self.user_store = Some(user_store);
self
}
#[must_use = "builder method returns the modified state"]
pub fn with_trusted_email_providers(mut self, trusted: TrustedEmailProviders) -> Self {
self.trusted_email_providers = trusted;
self
}
#[must_use = "builder method returns the modified state"]
pub fn with_redirect_uri_allowlist(mut self, allowlist: Vec<String>) -> Self {
self.redirect_uri_allowlist = allowlist;
self
}
pub fn register_provider(&mut self, name: impl Into<String>, provider: Arc<dyn OAuthProvider>) {
self.providers.insert(name.into(), provider);
}
#[must_use]
pub fn provider_names(&self) -> Vec<String> {
let mut names: Vec<String> = self.providers.keys().cloned().collect();
names.sort();
names
}
#[must_use]
pub fn get_provider(&self, name: &str) -> Option<&Arc<dyn OAuthProvider>> {
self.providers.get(name)
}
}
#[derive(Debug, Deserialize)]
pub struct AuthorizeQuery {
pub provider: String,
pub redirect_uri: String,
}
#[derive(Debug, Deserialize)]
pub struct CallbackQuery {
pub code: Option<String>,
pub state: Option<String>,
pub error: Option<String>,
pub error_description: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct CallbackForm {
pub code: Option<String>,
pub state: Option<String>,
pub error: Option<String>,
pub error_description: Option<String>,
pub user: Option<String>,
}
#[derive(Debug, Serialize)]
pub struct ProvidersResponse {
pub providers: Vec<String>,
}
#[derive(Debug, Serialize)]
pub struct AuthTokenResponse {
pub access_token: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub refresh_token: Option<String>,
pub token_type: String,
pub expires_in: u64,
pub provider: String,
}
impl AuthTokenResponse {
#[must_use = "builder does nothing until .build() is called"]
pub fn builder() -> AuthTokenResponseBuilder {
AuthTokenResponseBuilder::default()
}
}
#[derive(Debug, Default)]
pub struct AuthTokenResponseBuilder {
access_token: Option<String>,
refresh_token: Option<String>,
token_type: Option<String>,
expires_in: Option<u64>,
provider: Option<String>,
}
impl AuthTokenResponseBuilder {
pub fn access_token(mut self, access_token: impl Into<String>) -> Self {
self.access_token = Some(access_token.into());
self
}
pub fn refresh_token(mut self, refresh_token: impl Into<String>) -> Self {
self.refresh_token = Some(refresh_token.into());
self
}
pub fn token_type(mut self, token_type: impl Into<String>) -> Self {
self.token_type = Some(token_type.into());
self
}
#[must_use = "builder method returns modified builder"]
pub const fn expires_in(mut self, expires_in: u64) -> Self {
self.expires_in = Some(expires_in);
self
}
pub fn provider(mut self, provider: impl Into<String>) -> Self {
self.provider = Some(provider.into());
self
}
pub fn build(self) -> Result<AuthTokenResponse, String> {
Ok(AuthTokenResponse {
access_token: self
.access_token
.ok_or("AuthTokenResponse: access_token is required")?,
refresh_token: self.refresh_token,
token_type: self.token_type.ok_or("AuthTokenResponse: token_type is required")?,
expires_in: self.expires_in.ok_or("AuthTokenResponse: expires_in is required")?,
provider: self.provider.ok_or("AuthTokenResponse: provider is required")?,
})
}
}
fn json_error(status: StatusCode, message: &str) -> Response {
(status, Json(serde_json::json!({ "error": message }))).into_response()
}
fn effective_email_verified(
trusted: &TrustedEmailProviders,
provider: &str,
claimed_verified: bool,
) -> bool {
claimed_verified && trusted.is_trusted(provider)
}
pub async fn list_providers(
State(state): State<Arc<MultiProviderAuthState>>,
) -> Json<ProvidersResponse> {
Json(ProvidersResponse {
providers: state.provider_names(),
})
}
pub async fn authorize(
State(state): State<Arc<MultiProviderAuthState>>,
Query(q): Query<AuthorizeQuery>,
) -> Response {
if q.provider.len() > MAX_PROVIDER_NAME_BYTES {
return json_error(StatusCode::BAD_REQUEST, "provider name exceeds maximum length");
}
if q.redirect_uri.is_empty() {
return json_error(StatusCode::BAD_REQUEST, "redirect_uri is required");
}
if q.redirect_uri.len() > MAX_REDIRECT_URI_BYTES {
return json_error(StatusCode::BAD_REQUEST, "redirect_uri exceeds maximum length");
}
let bound_redirect_uri = if state.redirect_uri_allowlist.is_empty() {
None
} else if is_redirect_uri_allowed(&q.redirect_uri, &state.redirect_uri_allowlist) {
Some(q.redirect_uri.clone())
} else {
return json_error(StatusCode::BAD_REQUEST, "redirect_uri is not allow-listed");
};
let Some(provider) = state.get_provider(&q.provider) else {
return json_error(StatusCode::BAD_REQUEST, &format!("unknown provider: {}", q.provider));
};
let state_value = generate_secure_state();
let Ok(now) = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
else {
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "system clock error");
};
let expiry = now + 600;
let state_payload = encode_state_value(&q.provider, bound_redirect_uri.as_deref());
if let Err(e) = state.state_store.store(state_value.clone(), state_payload, expiry).await {
tracing::error!("state store failed: {e}");
return json_error(
StatusCode::INTERNAL_SERVER_ERROR,
"authorization flow could not be started",
);
}
let authorization_url = provider.authorization_url(&state_value);
Redirect::to(&authorization_url).into_response()
}
pub async fn callback(
State(state): State<Arc<MultiProviderAuthState>>,
Query(q): Query<CallbackQuery>,
) -> Response {
complete_callback(&state, q, None).await
}
pub async fn callback_form_post(
State(state): State<Arc<MultiProviderAuthState>>,
Form(form): Form<CallbackForm>,
) -> Response {
let display_name = form
.user
.as_deref()
.and_then(crate::providers::apple::AppleFirstAuthUser::parse)
.and_then(|u| u.display_name());
complete_callback(
&state,
CallbackQuery {
code: form.code,
state: form.state,
error: form.error,
error_description: form.error_description,
},
display_name,
)
.await
}
#[allow(clippy::cognitive_complexity)] async fn complete_callback(
state: &Arc<MultiProviderAuthState>,
q: CallbackQuery,
first_auth_name: Option<String>,
) -> Response {
if let Some(err) = q.error {
let desc = q.error_description.as_deref().unwrap_or("(no description)");
tracing::warn!(provider_error = %err, description = %desc, "OAuth provider returned error");
let client_message = match err.as_str() {
"access_denied" => "Access was denied",
"login_required" => "Authentication is required",
"invalid_request" | "invalid_scope" => "Invalid authorization request",
"server_error" | "temporarily_unavailable" => "Authorization server error",
_ => "Authorization failed",
};
return json_error(StatusCode::BAD_REQUEST, client_message);
}
let (Some(code), Some(state_token)) = (q.code, q.state) else {
return json_error(StatusCode::BAD_REQUEST, "missing code or state parameter");
};
let Ok((state_payload, expiry)) = state.state_store.retrieve(&state_token).await else {
return json_error(StatusCode::BAD_REQUEST, "invalid or expired state token");
};
let (provider_name, bound_redirect_uri) = decode_state_value(&state_payload);
let Ok(now) = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
else {
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "system clock error");
};
if now > expiry {
return json_error(StatusCode::BAD_REQUEST, "state token expired");
}
let Some(provider) = state.get_provider(&provider_name) else {
tracing::error!(provider = %provider_name, "provider from state not found in registry");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "provider configuration error");
};
let token_response = match provider.exchange_code(&code).await {
Ok(t) => t,
Err(e) => {
tracing::error!(error = %e, "token exchange failed");
return json_error(StatusCode::BAD_GATEWAY, "token exchange with provider failed");
},
};
let mut user_info = match provider.user_info_from_tokens(&token_response).await {
Ok(u) => u,
Err(e) => {
tracing::error!(error = %e, "user info fetch failed");
return json_error(StatusCode::BAD_GATEWAY, "failed to retrieve user information");
},
};
if user_info.name.is_none() {
user_info.name = first_auth_name;
}
let provider_trusted = state.trusted_email_providers.is_trusted(&provider_name);
let email_verified = effective_email_verified(
&state.trusted_email_providers,
&provider_name,
user_info.email_verified,
);
if user_info.email_verified && !provider_trusted {
get_audit_logger().log_failure(
AuditEventType::AuthFailure,
SecretType::StateToken,
None,
"social_callback",
&format!(
"untrusted_provider_email_downgraded:{provider_name} — email_verified claim not \
honored for account linking"
),
);
tracing::warn!(
provider = %provider_name,
"provider asserted email_verified but is not in the trusted-email set; treating email \
as unverified for account linking (#368)"
);
}
let local_user_id = if let Some(account_store) = &state.user_store {
match account_store
.link_or_create_user(
None,
user_info.email.as_deref(),
email_verified,
&provider_name,
&user_info.id,
)
.await
{
Ok(result) => result.user_id,
Err(e) => {
tracing::error!(error = %e, "account store lookup failed");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "user resolution failed");
},
}
} else {
user_info.id.clone()
};
let session_expiry = now + (7 * 24 * 60 * 60);
let session_tokens = match state
.session_store
.create_session(&local_user_id, session_expiry)
.await
{
Ok(t) => t,
Err(e) => {
tracing::error!(error = %e, "session creation failed");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "session could not be created");
},
};
if let Some(redirect_uri) = bound_redirect_uri {
let location = build_redirect_with_tokens(
&redirect_uri,
&session_tokens.access_token,
&session_tokens.refresh_token,
session_tokens.expires_in,
&provider_name,
);
return Redirect::to(&location).into_response();
}
Json(AuthTokenResponse {
access_token: session_tokens.access_token,
refresh_token: Some(session_tokens.refresh_token),
token_type: "Bearer".to_string(),
expires_in: session_tokens.expires_in,
provider: provider_name,
})
.into_response()
}
#[cfg(test)]
mod tests;