use std::{collections::HashMap, sync::Arc};
use chrono::{DateTime, Duration, Utc};
use serde::{Deserialize, Serialize};
use super::{super::error::AuthError, client::OIDCProviderConfig};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum ProviderType {
OAuth2,
OIDC,
}
impl std::fmt::Display for ProviderType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::OAuth2 => write!(f, "oauth2"),
Self::OIDC => write!(f, "oidc"),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OAuthSession {
pub id: String,
pub user_id: String,
pub provider_type: ProviderType,
pub provider_name: String,
pub provider_user_id: String,
pub access_token: String,
pub refresh_token: Option<String>,
pub token_expiry: DateTime<Utc>,
pub created_at: DateTime<Utc>,
pub last_refreshed: Option<DateTime<Utc>>,
}
impl OAuthSession {
#[must_use]
pub fn new(
user_id: String,
provider_type: ProviderType,
provider_name: String,
provider_user_id: String,
access_token: String,
token_expiry: DateTime<Utc>,
) -> Self {
Self {
id: uuid::Uuid::new_v4().to_string(),
user_id,
provider_type,
provider_name,
provider_user_id,
access_token,
refresh_token: None,
token_expiry,
created_at: Utc::now(),
last_refreshed: None,
}
}
#[must_use]
pub fn is_expired(&self) -> bool {
self.token_expiry <= Utc::now()
}
#[must_use]
pub fn is_expiring_soon(&self, grace_seconds: i64) -> bool {
self.token_expiry <= (Utc::now() + Duration::seconds(grace_seconds))
}
pub fn refresh_tokens(&mut self, access_token: String, token_expiry: DateTime<Utc>) {
self.access_token = access_token;
self.token_expiry = token_expiry;
self.last_refreshed = Some(Utc::now());
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ExternalAuthProvider {
pub id: String,
pub provider_type: ProviderType,
pub provider_name: String,
pub client_id: String,
pub client_secret_vault_path: String,
pub oidc_config: Option<OIDCProviderConfig>,
pub oauth2_config: Option<OAuth2ClientConfig>,
pub enabled: bool,
pub scopes: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct OAuth2ClientConfig {
pub authorization_endpoint: String,
pub token_endpoint: String,
pub use_pkce: bool,
}
impl ExternalAuthProvider {
pub fn new(
provider_type: ProviderType,
provider_name: impl Into<String>,
client_id: impl Into<String>,
client_secret_vault_path: impl Into<String>,
) -> Self {
Self {
id: uuid::Uuid::new_v4().to_string(),
provider_type,
provider_name: provider_name.into(),
client_id: client_id.into(),
client_secret_vault_path: client_secret_vault_path.into(),
oidc_config: None,
oauth2_config: None,
enabled: true,
scopes: vec![
"openid".to_string(),
"profile".to_string(),
"email".to_string(),
],
}
}
pub const fn set_enabled(&mut self, enabled: bool) {
self.enabled = enabled;
}
pub fn set_scopes(&mut self, scopes: Vec<String>) {
self.scopes = scopes;
}
}
#[derive(Debug, Clone)]
pub struct ProviderRegistry {
providers: Arc<std::sync::Mutex<HashMap<String, ExternalAuthProvider>>>,
}
impl ProviderRegistry {
#[must_use]
pub fn new() -> Self {
Self {
providers: Arc::new(std::sync::Mutex::new(HashMap::new())),
}
}
pub fn register(&self, provider: ExternalAuthProvider) -> std::result::Result<(), AuthError> {
let mut providers = self.providers.lock().map_err(|_| AuthError::Internal {
message: "provider registry mutex poisoned".to_string(),
})?;
providers.insert(provider.provider_name.clone(), provider);
Ok(())
}
pub fn get(&self, name: &str) -> std::result::Result<Option<ExternalAuthProvider>, AuthError> {
let providers = self.providers.lock().map_err(|_| AuthError::Internal {
message: "provider registry mutex poisoned".to_string(),
})?;
Ok(providers.get(name).cloned())
}
pub fn list_enabled(&self) -> std::result::Result<Vec<ExternalAuthProvider>, AuthError> {
let providers = self.providers.lock().map_err(|_| AuthError::Internal {
message: "provider registry mutex poisoned".to_string(),
})?;
Ok(providers.values().filter(|p| p.enabled).cloned().collect())
}
pub fn disable(&self, name: &str) -> std::result::Result<bool, AuthError> {
let mut providers = self.providers.lock().map_err(|_| AuthError::Internal {
message: "provider registry mutex poisoned".to_string(),
})?;
if let Some(provider) = providers.get_mut(name) {
provider.set_enabled(false);
Ok(true)
} else {
Ok(false)
}
}
pub fn enable(&self, name: &str) -> std::result::Result<bool, AuthError> {
let mut providers = self.providers.lock().map_err(|_| AuthError::Internal {
message: "provider registry mutex poisoned".to_string(),
})?;
if let Some(provider) = providers.get_mut(name) {
provider.set_enabled(true);
Ok(true)
} else {
Ok(false)
}
}
}
impl Default for ProviderRegistry {
fn default() -> Self {
Self::new()
}
}