1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
12#[non_exhaustive]
13pub enum ProviderType {
14 OAuth2,
16 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#[derive(Debug, Clone, Serialize, Deserialize)]
31pub struct OAuthSession {
32 pub id: String,
34 pub user_id: String,
36 pub provider_type: ProviderType,
38 pub provider_name: String,
40 pub provider_user_id: String,
42 pub access_token: String,
44 pub refresh_token: Option<String>,
46 pub token_expiry: DateTime<Utc>,
48 pub created_at: DateTime<Utc>,
50 pub last_refreshed: Option<DateTime<Utc>>,
52}
53
54impl OAuthSession {
55 #[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 #[must_use]
81 pub fn is_expired(&self) -> bool {
82 self.token_expiry <= Utc::now()
83 }
84
85 #[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 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#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
101pub struct ExternalAuthProvider {
102 pub id: String,
104 pub provider_type: ProviderType,
106 pub provider_name: String,
108 pub client_id: String,
110 pub client_secret_vault_path: String,
112 pub oidc_config: Option<OIDCProviderConfig>,
114 pub oauth2_config: Option<OAuth2ClientConfig>,
116 pub enabled: bool,
118 pub scopes: Vec<String>,
120}
121
122#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
124pub struct OAuth2ClientConfig {
125 pub authorization_endpoint: String,
127 pub token_endpoint: String,
129 pub use_pkce: bool,
131}
132
133impl ExternalAuthProvider {
134 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 pub const fn set_enabled(&mut self, enabled: bool) {
160 self.enabled = enabled;
161 }
162
163 pub fn set_scopes(&mut self, scopes: Vec<String>) {
165 self.scopes = scopes;
166 }
167}
168
169#[derive(Debug, Clone)]
171pub struct ProviderRegistry {
172 providers: Arc<std::sync::Mutex<HashMap<String, ExternalAuthProvider>>>,
176}
177
178impl ProviderRegistry {
179 #[must_use]
181 pub fn new() -> Self {
182 Self {
183 providers: Arc::new(std::sync::Mutex::new(HashMap::new())),
184 }
185 }
186
187 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 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 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 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 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}