Skip to main content

fraiseql_auth/oauth/
provider.rs

1//! External OAuth provider registry and session management.
2
3use std::{collections::HashMap, sync::Arc};
4
5use chrono::{DateTime, Duration, Utc};
6use serde::{Deserialize, Serialize};
7
8use super::{super::error::AuthError, client::OIDCProviderConfig};
9
10/// External authentication provider type
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
12#[non_exhaustive]
13pub enum ProviderType {
14    /// OAuth2 provider
15    OAuth2,
16    /// OIDC provider
17    OIDC,
18}
19
20impl std::fmt::Display for ProviderType {
21    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
22        match self {
23            Self::OAuth2 => write!(f, "oauth2"),
24            Self::OIDC => write!(f, "oidc"),
25        }
26    }
27}
28
29/// OAuth session stored in database
30#[derive(Debug, Clone, Serialize, Deserialize)]
31pub struct OAuthSession {
32    /// Session ID
33    pub id:               String,
34    /// User ID (local system)
35    pub user_id:          String,
36    /// Provider type (oauth2, oidc)
37    pub provider_type:    ProviderType,
38    /// Provider name (Auth0, Google, etc.)
39    pub provider_name:    String,
40    /// Provider's user ID (sub claim)
41    pub provider_user_id: String,
42    /// Access token (encrypted)
43    pub access_token:     String,
44    /// Refresh token (encrypted), if available
45    pub refresh_token:    Option<String>,
46    /// When access token expires
47    pub token_expiry:     DateTime<Utc>,
48    /// Session creation time
49    pub created_at:       DateTime<Utc>,
50    /// Last time token was refreshed
51    pub last_refreshed:   Option<DateTime<Utc>>,
52}
53
54impl OAuthSession {
55    /// Create new OAuth session
56    #[must_use]
57    pub fn new(
58        user_id: String,
59        provider_type: ProviderType,
60        provider_name: String,
61        provider_user_id: String,
62        access_token: String,
63        token_expiry: DateTime<Utc>,
64    ) -> Self {
65        Self {
66            id: uuid::Uuid::new_v4().to_string(),
67            user_id,
68            provider_type,
69            provider_name,
70            provider_user_id,
71            access_token,
72            refresh_token: None,
73            token_expiry,
74            created_at: Utc::now(),
75            last_refreshed: None,
76        }
77    }
78
79    /// Check if session is expired
80    #[must_use]
81    pub fn is_expired(&self) -> bool {
82        self.token_expiry <= Utc::now()
83    }
84
85    /// Check if session will be expired within grace period
86    #[must_use]
87    pub fn is_expiring_soon(&self, grace_seconds: i64) -> bool {
88        self.token_expiry <= (Utc::now() + Duration::seconds(grace_seconds))
89    }
90
91    /// Update tokens after refresh
92    pub fn refresh_tokens(&mut self, access_token: String, token_expiry: DateTime<Utc>) {
93        self.access_token = access_token;
94        self.token_expiry = token_expiry;
95        self.last_refreshed = Some(Utc::now());
96    }
97}
98
99/// External auth provider configuration
100#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
101pub struct ExternalAuthProvider {
102    /// Provider ID
103    pub id: String,
104    /// Provider type (oauth2, oidc)
105    pub provider_type: ProviderType,
106    /// Provider name (Auth0, Google, Microsoft, Okta)
107    pub provider_name: String,
108    /// Client ID
109    pub client_id: String,
110    /// Client secret (should be fetched from vault)
111    pub client_secret_vault_path: String,
112    /// Provider configuration (OIDC)
113    pub oidc_config: Option<OIDCProviderConfig>,
114    /// OAuth2 configuration
115    pub oauth2_config: Option<OAuth2ClientConfig>,
116    /// Enabled flag
117    pub enabled: bool,
118    /// Requested scopes
119    pub scopes: Vec<String>,
120}
121
122/// OAuth2 client configuration
123#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
124pub struct OAuth2ClientConfig {
125    /// Authorization endpoint
126    pub authorization_endpoint: String,
127    /// Token endpoint
128    pub token_endpoint:         String,
129    /// Use PKCE
130    pub use_pkce:               bool,
131}
132
133impl ExternalAuthProvider {
134    /// Create new external auth provider
135    pub fn new(
136        provider_type: ProviderType,
137        provider_name: impl Into<String>,
138        client_id: impl Into<String>,
139        client_secret_vault_path: impl Into<String>,
140    ) -> Self {
141        Self {
142            id: uuid::Uuid::new_v4().to_string(),
143            provider_type,
144            provider_name: provider_name.into(),
145            client_id: client_id.into(),
146            client_secret_vault_path: client_secret_vault_path.into(),
147            oidc_config: None,
148            oauth2_config: None,
149            enabled: true,
150            scopes: vec![
151                "openid".to_string(),
152                "profile".to_string(),
153                "email".to_string(),
154            ],
155        }
156    }
157
158    /// Enable or disable provider
159    pub const fn set_enabled(&mut self, enabled: bool) {
160        self.enabled = enabled;
161    }
162
163    /// Set requested scopes
164    pub fn set_scopes(&mut self, scopes: Vec<String>) {
165        self.scopes = scopes;
166    }
167}
168
169/// Provider registry managing multiple OAuth providers
170#[derive(Debug, Clone)]
171pub struct ProviderRegistry {
172    /// Map of providers by name
173    // std::sync::Mutex is intentional: this lock is never held across .await.
174    // Switch to tokio::sync::Mutex if that constraint ever changes.
175    providers: Arc<std::sync::Mutex<HashMap<String, ExternalAuthProvider>>>,
176}
177
178impl ProviderRegistry {
179    /// Create new provider registry
180    #[must_use]
181    pub fn new() -> Self {
182        Self {
183            providers: Arc::new(std::sync::Mutex::new(HashMap::new())),
184        }
185    }
186
187    /// Register provider
188    ///
189    /// # Errors
190    ///
191    /// Returns `AuthError::Internal` if the mutex is poisoned.
192    pub fn register(&self, provider: ExternalAuthProvider) -> std::result::Result<(), AuthError> {
193        let mut providers = self.providers.lock().map_err(|_| AuthError::Internal {
194            message: "provider registry mutex poisoned".to_string(),
195        })?;
196        providers.insert(provider.provider_name.clone(), provider);
197        Ok(())
198    }
199
200    /// Get provider by name
201    ///
202    /// # Errors
203    ///
204    /// Returns `AuthError::Internal` if the mutex is poisoned.
205    pub fn get(&self, name: &str) -> std::result::Result<Option<ExternalAuthProvider>, AuthError> {
206        let providers = self.providers.lock().map_err(|_| AuthError::Internal {
207            message: "provider registry mutex poisoned".to_string(),
208        })?;
209        Ok(providers.get(name).cloned())
210    }
211
212    /// List all enabled providers
213    ///
214    /// # Errors
215    ///
216    /// Returns `AuthError::Internal` if the mutex is poisoned.
217    pub fn list_enabled(&self) -> std::result::Result<Vec<ExternalAuthProvider>, AuthError> {
218        let providers = self.providers.lock().map_err(|_| AuthError::Internal {
219            message: "provider registry mutex poisoned".to_string(),
220        })?;
221        Ok(providers.values().filter(|p| p.enabled).cloned().collect())
222    }
223
224    /// Disable provider
225    ///
226    /// # Errors
227    ///
228    /// Returns `AuthError::Internal` if the mutex is poisoned.
229    pub fn disable(&self, name: &str) -> std::result::Result<bool, AuthError> {
230        let mut providers = self.providers.lock().map_err(|_| AuthError::Internal {
231            message: "provider registry mutex poisoned".to_string(),
232        })?;
233        if let Some(provider) = providers.get_mut(name) {
234            provider.set_enabled(false);
235            Ok(true)
236        } else {
237            Ok(false)
238        }
239    }
240
241    /// Enable provider
242    ///
243    /// # Errors
244    ///
245    /// Returns `AuthError::Internal` if the mutex is poisoned.
246    pub fn enable(&self, name: &str) -> std::result::Result<bool, AuthError> {
247        let mut providers = self.providers.lock().map_err(|_| AuthError::Internal {
248            message: "provider registry mutex poisoned".to_string(),
249        })?;
250        if let Some(provider) = providers.get_mut(name) {
251            provider.set_enabled(true);
252            Ok(true)
253        } else {
254            Ok(false)
255        }
256    }
257}
258
259impl Default for ProviderRegistry {
260    fn default() -> Self {
261        Self::new()
262    }
263}