use std::{collections::HashMap, sync::Arc};
use axum::{
Json,
extract::{Query, State},
http::StatusCode,
response::{IntoResponse, Redirect, Response},
};
use serde::{Deserialize, Serialize};
use crate::{
account_linking::AccountStore, 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;
#[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>>,
}
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,
}
}
pub fn with_user_store(mut self, user_store: Arc<dyn AccountStore>) -> Self {
self.user_store = Some(user_store);
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, 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()
}
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 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;
if let Err(e) = state.state_store.store(state_value.clone(), q.provider.clone(), 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()
}
#[allow(clippy::cognitive_complexity)] pub async fn callback(
State(state): State<Arc<MultiProviderAuthState>>,
Query(q): Query<CallbackQuery>,
) -> 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((provider_name, expiry)) = state.state_store.retrieve(&state_token).await else {
return json_error(StatusCode::BAD_REQUEST, "invalid or expired state token");
};
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
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 user_info = match provider.user_info(&token_response.access_token).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");
},
};
let local_user_id = if let Some(account_store) = &state.user_store {
match account_store
.link_or_create_user(&user_info.email, &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");
},
};
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()
}