1use crate::auth::session::{Session, SessionConfig, SessionStore};
2use crate::auth::{AuthError, AuthInput, AuthMethod, AuthResult, ErasedOAuthFlow, Identity};
3#[cfg(feature = "token")]
4use crate::token::TokenManager;
5use std::collections::HashMap;
6use std::sync::Arc;
7
8#[derive(Clone, Default, Debug)]
10pub struct Missing;
11
12#[derive(Clone, Debug)]
14pub struct Configured<T>(pub T);
15
16pub trait SessionStoreState: Send + Sync + Clone {
18 fn get_store(&self) -> Arc<dyn SessionStore>;
20}
21
22impl SessionStoreState for Configured<Arc<dyn SessionStore>> {
23 fn get_store(&self) -> Arc<dyn SessionStore> {
24 self.0.clone()
25 }
26}
27
28pub trait TokenManagerState: Send + Sync + Clone {
30 #[cfg(feature = "token")]
32 fn get_manager(&self) -> Arc<TokenManager>;
33}
34
35#[cfg(feature = "token")]
36impl TokenManagerState for Configured<Arc<TokenManager>> {
37 fn get_manager(&self) -> Arc<TokenManager> {
38 self.0.clone()
39 }
40}
41
42pub struct Engine<S = Missing, T = Missing> {
48 pub providers: HashMap<String, Arc<dyn ErasedOAuthFlow>>,
50 pub auth_methods: HashMap<String, Arc<dyn AuthMethod>>,
52 pub mfa_methods: HashMap<String, Arc<dyn AuthMethod>>,
54 pub mfa_jwt_secret: [u8; 32],
56 pub session_store: S,
58 pub session_config: SessionConfig,
60 #[cfg(feature = "token")]
62 pub token_manager: T,
63}
64
65impl<S, T> Clone for Engine<S, T>
66where
67 S: Clone,
68 T: Clone,
69{
70 fn clone(&self) -> Self {
71 Self {
72 providers: self.providers.clone(),
73 auth_methods: self.auth_methods.clone(),
74 mfa_methods: self.mfa_methods.clone(),
75 mfa_jwt_secret: self.mfa_jwt_secret,
76 session_store: self.session_store.clone(),
77 session_config: self.session_config.clone(),
78 #[cfg(feature = "token")]
79 token_manager: self.token_manager.clone(),
80 }
81 }
82}
83
84impl Engine<Missing, Missing> {
85 pub fn builder() -> EngineBuilder<Missing, Missing> {
87 let mut secret = [0u8; 32];
88 rand::RngCore::fill_bytes(&mut rand::rng(), &mut secret);
89
90 EngineBuilder {
91 providers: HashMap::new(),
92 auth_methods: HashMap::new(),
93 mfa_methods: HashMap::new(),
94 mfa_jwt_secret: secret,
95 session_store: Missing,
96 session_config: SessionConfig::default(),
97 #[cfg(feature = "token")]
98 token_manager: Missing,
99 }
100 }
101}
102
103pub struct EngineBuilder<S = Missing, T = Missing> {
105 providers: HashMap<String, Arc<dyn ErasedOAuthFlow>>,
106 auth_methods: HashMap<String, Arc<dyn AuthMethod>>,
107 mfa_methods: HashMap<String, Arc<dyn AuthMethod>>,
108 mfa_jwt_secret: [u8; 32],
109 session_store: S,
110 session_config: SessionConfig,
111 #[cfg(feature = "token")]
112 token_manager: T,
113}
114
115impl<S, T> EngineBuilder<S, T> {
116 pub fn provider<F>(mut self, flow: F) -> Self
118 where
119 F: ErasedOAuthFlow + 'static,
120 {
121 let id = flow.provider_id();
122 self.providers.insert(id, Arc::new(flow));
123 self
124 }
125
126 pub fn with_auth_method<M>(mut self, method: M) -> Self
128 where
129 M: AuthMethod + 'static,
130 {
131 self.auth_methods
132 .insert(method.name().to_string(), Arc::new(method));
133 self
134 }
135
136 pub fn with_mfa_method<M>(mut self, method: M) -> Self
138 where
139 M: AuthMethod + 'static,
140 {
141 self.mfa_methods
142 .insert(method.name().to_string(), Arc::new(method));
143 self
144 }
145
146 #[cfg(feature = "totp")]
148 pub fn with_totp<C>(self, store: C) -> Self
149 where
150 C: crate::CredentialStore + 'static,
151 {
152 self.with_auth_method(crate::auth::totp::TotpAuthMethod::new(store))
153 }
154
155 #[cfg(feature = "webauthn")]
157 pub fn with_webauthn<C>(self, webauthn: Arc<webauthn_rs::prelude::Webauthn>, store: C) -> Self
158 where
159 C: crate::CredentialStore + 'static,
160 {
161 self.with_auth_method(crate::auth::webauthn::WebAuthnAuthMethod::new(
162 webauthn, store,
163 ))
164 }
165
166 pub fn session_store(
168 self,
169 store: Arc<dyn SessionStore>,
170 ) -> EngineBuilder<Configured<Arc<dyn SessionStore>>, T> {
171 EngineBuilder {
172 providers: self.providers,
173 auth_methods: self.auth_methods,
174 mfa_methods: self.mfa_methods,
175 mfa_jwt_secret: self.mfa_jwt_secret,
176 session_store: Configured(store),
177 session_config: self.session_config,
178 #[cfg(feature = "token")]
179 token_manager: self.token_manager,
180 }
181 }
182
183 #[cfg(feature = "token")]
185 pub fn token_manager(
186 self,
187 manager: Arc<TokenManager>,
188 ) -> EngineBuilder<S, Configured<Arc<TokenManager>>> {
189 EngineBuilder {
190 providers: self.providers,
191 auth_methods: self.auth_methods,
192 mfa_methods: self.mfa_methods,
193 mfa_jwt_secret: self.mfa_jwt_secret,
194 session_store: self.session_store,
195 session_config: self.session_config,
196 token_manager: Configured(manager),
197 }
198 }
199
200 #[cfg(feature = "token")]
202 pub fn jwt_secret(self, secret: &[u8]) -> EngineBuilder<S, Configured<Arc<TokenManager>>> {
203 self.token_manager(Arc::new(TokenManager::new(secret, None)))
204 }
205
206 pub fn session_config(mut self, config: SessionConfig) -> Self {
208 self.session_config = config;
209 self
210 }
211
212 pub fn build(self) -> Engine<S, T> {
214 Engine {
215 providers: self.providers,
216 auth_methods: self.auth_methods,
217 mfa_methods: self.mfa_methods,
218 mfa_jwt_secret: self.mfa_jwt_secret,
219 session_store: self.session_store,
220 session_config: self.session_config,
221 #[cfg(feature = "token")]
222 token_manager: self.token_manager,
223 }
224 }
225}
226
227impl<S, T> Engine<S, T> {
228 pub async fn authenticate(&self, input: AuthInput) -> Result<AuthResult, AuthError> {
232 if let AuthInput::MfaChallenge {
234 mfa_token,
235 challenge_input,
236 } = input
237 {
238 let token_data = jsonwebtoken::decode::<crate::auth::state::MfaTokenClaims>(
240 &mfa_token,
241 &jsonwebtoken::DecodingKey::from_secret(&self.mfa_jwt_secret),
242 &jsonwebtoken::Validation::new(jsonwebtoken::Algorithm::HS256),
243 )
244 .map_err(|_| AuthError::InvalidInput)?;
245
246 if !token_data.claims.mfa_pending {
247 return Err(AuthError::InvalidInput);
248 }
249
250 let method_name = match &*challenge_input {
251 #[cfg(feature = "totp")]
252 AuthInput::Totp { .. } => "totp",
253 #[cfg(feature = "webauthn")]
254 AuthInput::WebAuthnAuthentication { .. } => "webauthn",
255 _ => "",
256 };
257
258 if method_name.is_empty() {
259 return Err(AuthError::InvalidInput);
260 }
261
262 let method = self
263 .auth_methods
264 .get(method_name)
265 .or_else(|| self.mfa_methods.get(method_name))
266 .ok_or_else(|| {
267 AuthError::Internal(format!("MFA method {} not registered", method_name))
268 })?;
269
270 let identity = method.authenticate(*challenge_input).await?;
271
272 if identity.external_id != token_data.claims.sub {
273 return Err(AuthError::Credentials("MFA token user mismatch".into()));
274 }
275
276 return Ok(AuthResult::Success(identity));
277 }
278
279 let method_name = match &input {
281 AuthInput::Password { .. } => "password",
282 #[cfg(feature = "totp")]
283 AuthInput::Totp { .. } => "totp", #[cfg(feature = "webauthn")]
285 AuthInput::WebAuthnAuthentication { .. } => "webauthn",
286 _ => "",
287 };
288
289 if method_name.is_empty() {
290 return Err(AuthError::InvalidInput);
291 }
292
293 let method = self.auth_methods.get(method_name).ok_or_else(|| {
294 AuthError::Internal(format!(
295 "Primary auth method {} not registered or is step-up only",
296 method_name
297 ))
298 })?;
299
300 let identity = method.authenticate(input).await?;
301
302 let mut enrolled_methods = Vec::new();
304 for (name, m) in self.auth_methods.iter().chain(self.mfa_methods.iter()) {
305 if name == "password" || name == method_name {
306 continue;
307 }
308 if !enrolled_methods.contains(name)
309 && m.has_enrolled(&identity.external_id).await.unwrap_or(false)
310 {
311 enrolled_methods.push(name.clone());
312 }
313 }
314
315 if enrolled_methods.is_empty() || method.is_mfa_equivalent() {
318 Ok(AuthResult::Success(identity))
319 } else {
320 let exp = chrono::Utc::now() + chrono::Duration::minutes(15);
322 let claims = crate::auth::state::MfaTokenClaims {
323 sub: identity.external_id.clone(),
324 mfa_pending: true,
325 exp: exp.timestamp() as usize,
326 };
327
328 let mfa_token = jsonwebtoken::encode(
329 &jsonwebtoken::Header::default(),
330 &claims,
331 &jsonwebtoken::EncodingKey::from_secret(&self.mfa_jwt_secret),
332 )
333 .map_err(|e| AuthError::Internal(e.to_string()))?;
334
335 Ok(AuthResult::MfaRequired {
336 mfa_token,
337 user_id: identity.external_id,
338 allowed_methods: enrolled_methods,
339 })
340 }
341 }
342}
343
344impl<T> Engine<Configured<Arc<dyn SessionStore>>, T> {
346 pub fn session_store(&self) -> Arc<dyn SessionStore> {
348 self.session_store.0.clone()
349 }
350
351 #[tracing::instrument(skip(self, identity), fields(user_id = %identity.external_id))]
353 pub async fn create_session(&self, identity: Identity) -> Result<Session, AuthError> {
354 let session_duration = self
355 .session_config
356 .max_age
357 .unwrap_or(chrono::Duration::hours(24));
358 let session = Session {
359 id: uuid::Uuid::new_v4().to_string(),
360 identity,
361 expires_at: chrono::Utc::now() + session_duration,
362 };
363
364 tracing::debug!(session_id = %session.id, "creating new session");
365
366 self.session_store
367 .0
368 .save_session(&session)
369 .await
370 .map_err(|e| {
371 tracing::error!(error = %e, "failed to save session");
372 AuthError::Session(e.to_string())
373 })?;
374
375 tracing::info!(session_id = %session.id, "session created successfully");
376 Ok(session)
377 }
378}
379
380#[cfg(feature = "token")]
381impl<S> Engine<S, Configured<Arc<TokenManager>>> {
382 pub fn token_manager(&self) -> Arc<TokenManager> {
384 self.token_manager.0.clone()
385 }
386
387 #[tracing::instrument(skip(self, identity), fields(user_id = %identity.external_id))]
389 pub fn issue_token(
390 &self,
391 identity: Identity,
392 expires_in_secs: u64,
393 ) -> Result<String, AuthError> {
394 tracing::debug!("issuing token for user");
395 self.token_manager
396 .0
397 .issue_user_token(identity, expires_in_secs, None, None)
398 .map_err(|e| {
399 tracing::error!(error = %e, "failed to issue token");
400 AuthError::Token(e.to_string())
401 })
402 .inspect(|_| {
403 tracing::info!("token issued successfully");
404 })
405 }
406}
407
408pub trait HasSessionStore {
410 fn session_store(&self) -> Arc<dyn SessionStore>;
412}
413
414impl<T> HasSessionStore for Engine<Configured<Arc<dyn SessionStore>>, T> {
415 fn session_store(&self) -> Arc<dyn SessionStore> {
416 self.session_store.0.clone()
417 }
418}
419
420#[cfg(feature = "token")]
422pub trait HasTokenManager {
423 fn token_manager(&self) -> Arc<TokenManager>;
425}
426
427#[cfg(feature = "token")]
428impl<S> HasTokenManager for Engine<S, Configured<Arc<TokenManager>>> {
429 fn token_manager(&self) -> Arc<TokenManager> {
430 self.token_manager.0.clone()
431 }
432}