use std::{collections::HashMap, sync::Arc};
use axum::{
Json,
extract::{Query, State},
http::StatusCode,
response::{IntoResponse, Redirect, Response},
};
use base64::Engine as _;
use serde::Deserialize;
use crate::{
audit::logger::{AuditEventType, SecretType, get_audit_logger},
provider::OAuthProvider,
rate_limiting::RateLimiters,
session::unix_now,
state_store::StateStore,
};
pub struct SocialProviderRegistry {
providers: HashMap<String, Arc<dyn OAuthProvider>>,
}
impl SocialProviderRegistry {
#[must_use]
pub fn new() -> Self {
Self {
providers: HashMap::new(),
}
}
pub fn register(&mut self, name: impl Into<String>, provider: Arc<dyn OAuthProvider>) {
self.providers.insert(name.into(), provider);
}
#[must_use]
pub fn get(&self, name: &str) -> Option<Arc<dyn OAuthProvider>> {
self.providers.get(name).cloned()
}
#[must_use]
pub fn names(&self) -> Vec<&str> {
self.providers.keys().map(String::as_str).collect()
}
#[must_use]
pub fn len(&self) -> usize {
self.providers.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.providers.is_empty()
}
}
impl Default for SocialProviderRegistry {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone)]
pub struct SocialLoginState {
pub registry: Arc<SocialProviderRegistry>,
pub state_store: Arc<dyn StateStore>,
pub rate_limiters: Arc<RateLimiters>,
}
#[derive(Debug, Deserialize)]
pub struct SocialAuthorizeParams {
pub provider: String,
}
const STATE_TTL_SECS: u64 = 600;
pub async fn social_authorize(
State(state): State<Arc<SocialLoginState>>,
Query(params): Query<SocialAuthorizeParams>,
) -> Response {
let logger = get_audit_logger();
let Some(provider) = state.registry.get(¶ms.provider) else {
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::StateToken,
None,
"social_authorize",
&format!("unknown provider: {}", params.provider),
);
return (
StatusCode::BAD_REQUEST,
Json(serde_json::json!({
"error": "unknown_provider",
"message": format!("Provider '{}' is not configured", params.provider)
})),
)
.into_response();
};
let csrf_state = generate_state_token();
let expiry = match unix_now() {
Ok(now) => now + STATE_TTL_SECS,
Err(_) => {
return (StatusCode::INTERNAL_SERVER_ERROR, "internal error").into_response();
},
};
if let Err(e) = state
.state_store
.store(csrf_state.clone(), params.provider.clone(), expiry)
.await
{
tracing::error!(error = %e, "Failed to store OAuth CSRF state");
return (StatusCode::INTERNAL_SERVER_ERROR, "failed to store auth state").into_response();
}
logger.log_success(
AuditEventType::OauthStart,
SecretType::StateToken,
None,
&format!("social_authorize:{}", params.provider),
);
let auth_url = provider.authorization_url(&csrf_state);
Redirect::to(&auth_url).into_response()
}
fn generate_state_token() -> String {
use rand::RngCore as _;
let mut bytes = [0u8; 32];
rand::rng().fill_bytes(&mut bytes);
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
}
#[cfg(test)]
mod tests;