Skip to main content

backbone_auth/
auth_service.rs

1//! Authentication service implementation
2
3use anyhow::Result;
4use uuid::Uuid;
5use std::sync::OnceLock;
6use std::time::{SystemTime, UNIX_EPOCH, Duration};
7use regex::Regex;
8use crate::jwt::{JwtService, Claims, KeyRotationConfig, JwtAlgorithm};
9use crate::traits::{UserRepository, SecurityService, SecurityFlags, RefreshTokenClaims, User, AuthRequest, AuthResultEnhanced};
10
11/// Lazily-compiled email validation regex (compiled once, reused forever)
12fn email_regex() -> &'static Regex {
13    static RE: OnceLock<Regex> = OnceLock::new();
14    RE.get_or_init(|| {
15        Regex::new(r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$")
16            .expect("email regex is a valid constant pattern")
17    })
18}
19
20/// Authentication service configuration
21#[derive(Debug, Clone)]
22pub struct AuthServiceConfig {
23    pub jwt_secret: String,
24    pub token_expiry_hours: u64,
25    /// Optional key rotation configuration. When set, enables JWT key rotation support.
26    pub key_rotation: Option<KeyRotationConfig>,
27    /// JWT algorithm to use. Defaults to HS256 if not specified.
28    pub jwt_algorithm: Option<JwtAlgorithm>,
29    /// PEM-encoded RSA private key (required for RS256)
30    pub rsa_private_key_pem: Option<String>,
31    /// PEM-encoded RSA public key (required for RS256)
32    pub rsa_public_key_pem: Option<String>,
33}
34
35impl Default for AuthServiceConfig {
36    fn default() -> Self {
37        Self {
38            jwt_secret: "default_secret_change_in_production".to_string(),
39            token_expiry_hours: 24,
40            key_rotation: None,
41            jwt_algorithm: None,
42            rsa_private_key_pem: None,
43            rsa_public_key_pem: None,
44        }
45    }
46}
47
48/// Authentication service
49pub struct AuthService {
50    jwt_service: JwtService,
51    config: AuthServiceConfig,
52}
53
54impl AuthService {
55    pub fn new(config: AuthServiceConfig) -> anyhow::Result<Self> {
56        let jwt_service = match config.jwt_algorithm.unwrap_or(JwtAlgorithm::HS256) {
57            JwtAlgorithm::RS256 => {
58                let private_pem = config.rsa_private_key_pem.as_deref()
59                    .ok_or_else(|| anyhow::anyhow!("rsa_private_key_pem is required for RS256"))?;
60                let public_pem = config.rsa_public_key_pem.as_deref()
61                    .ok_or_else(|| anyhow::anyhow!("rsa_public_key_pem is required for RS256"))?;
62                match &config.key_rotation {
63                    Some(rotation_config) => {
64                        JwtService::with_rs256_rotation(private_pem, public_pem, rotation_config.clone())?
65                    }
66                    None => JwtService::new_rs256(private_pem, public_pem)?,
67                }
68            }
69            JwtAlgorithm::HS256 => {
70                match &config.key_rotation {
71                    Some(rotation_config) => {
72                        JwtService::with_rotation(&config.jwt_secret, rotation_config.clone())
73                    }
74                    None => JwtService::new(&config.jwt_secret),
75                }
76            }
77        };
78        Ok(Self {
79            jwt_service,
80            config,
81        })
82    }
83
84    /// Create an AuthService with HS256 and the given secret (infallible).
85    pub fn with_secret(jwt_secret: &str) -> Self {
86        Self {
87            jwt_service: JwtService::new(jwt_secret),
88            config: AuthServiceConfig {
89                jwt_secret: jwt_secret.to_string(),
90                token_expiry_hours: 24,
91                key_rotation: None,
92                jwt_algorithm: None,
93                rsa_private_key_pem: None,
94                rsa_public_key_pem: None,
95            },
96        }
97    }
98
99    /// Authenticate user with enhanced security.
100    ///
101    /// Orchestrates: validation, rate-limiting, credential verification,
102    /// security analysis, token generation, and audit logging.
103    pub async fn authenticate_enhanced(
104        &self,
105        request: AuthRequest,
106        user_repository: &dyn UserRepository,
107        security_service: &dyn SecurityService,
108    ) -> Result<AuthResultEnhanced> {
109        tracing::info!(
110            event = "auth.attempt_started",
111            email = %request.email,
112            ip_address = ?request.ip_address,
113            "Authentication attempt started"
114        );
115
116        if let Err(e) = self.validate_auth_request(&request) {
117            tracing::warn!(event = "auth.validation_failed", email = %request.email, reason = %e, "Validation failed");
118            return Err(e);
119        }
120
121        security_service.check_rate_limit(&request.email, request.ip_address.as_deref()).await?;
122
123        let user = self.verify_user_credentials(&request, user_repository, security_service).await?;
124        let (security_flags, requires_2fa) = self.evaluate_security_context(&user, &request, security_service).await?;
125        self.finalize_auth(&user, &request, security_flags, requires_2fa, security_service).await
126    }
127
128    /// Steps 3-5: look up user, check account status, verify password.
129    async fn verify_user_credentials(
130        &self,
131        request: &AuthRequest,
132        user_repository: &dyn UserRepository,
133        security_service: &dyn SecurityService,
134    ) -> Result<User> {
135        let user = user_repository.find_by_email(&request.email).await?.ok_or_else(|| {
136            tracing::warn!(event = "auth.user_not_found", email = %request.email, "User not found");
137            anyhow::anyhow!("Invalid credentials")
138        })?;
139
140        if let Err(e) = self.check_account_status(&user) {
141            tracing::warn!(event = "auth.account_status_failed", user_id = %user.id, reason = %e, "Account status check failed");
142            return Err(e);
143        }
144
145        if !self.verify_password(&request.password, &user.password_hash)? {
146            tracing::warn!(event = "auth.password_mismatch", user_id = %user.id, "Invalid password");
147            security_service.log_failed_auth_attempt(&user.id, request.ip_address.as_deref()).await?;
148            return Err(anyhow::anyhow!("Invalid credentials"));
149        }
150
151        Ok(user)
152    }
153
154    /// Steps 6-7: run security analysis and determine 2FA requirement.
155    async fn evaluate_security_context(
156        &self,
157        user: &User,
158        request: &AuthRequest,
159        security_service: &dyn SecurityService,
160    ) -> Result<(SecurityFlags, bool)> {
161        let security_flags = security_service
162            .analyze_login_attempt(&user.id, &request.device_info, request.ip_address.as_deref())
163            .await?;
164
165        if security_flags.new_device {
166            tracing::info!(event = "auth.new_device_detected", user_id = %user.id, "New device detected");
167        }
168
169        let requires_2fa = user.two_factor_enabled && !user.two_factor_methods.is_empty();
170        if requires_2fa {
171            tracing::info!(event = "auth.2fa_required", user_id = %user.id, "2FA required");
172        }
173
174        Ok((security_flags, requires_2fa))
175    }
176
177    /// Steps 8-9: generate tokens, log success, build result.
178    async fn finalize_auth(
179        &self,
180        user: &User,
181        request: &AuthRequest,
182        security_flags: SecurityFlags,
183        requires_2fa: bool,
184        security_service: &dyn SecurityService,
185    ) -> Result<AuthResultEnhanced> {
186        let remember_me = request.remember_me.unwrap_or(false);
187        let token = self.generate_access_token(&user.id)?;
188        let refresh_token = self.generate_refresh_token(&user.id, remember_me)?;
189
190        tracing::info!(event = "auth.tokens_generated", user_id = %user.id, has_refresh_token = refresh_token.is_some(), "Tokens generated");
191
192        security_service.log_successful_auth(&user.id, request.ip_address.as_deref()).await?;
193
194        tracing::info!(
195            event = "auth.success", user_id = %user.id, ip_address = ?request.ip_address,
196            requires_2fa = requires_2fa, risk_score = security_flags.risk_score, "Authentication successful"
197        );
198
199        Ok(AuthResultEnhanced {
200            user_id: user.id,
201            token,
202            refresh_token,
203            expires_at: chrono::Utc::now() + chrono::Duration::hours(self.config.token_expiry_hours as i64),
204            requires_2fa,
205            security_flags,
206        })
207    }
208
209    /// Validate authentication request with comprehensive checks
210    fn validate_auth_request(&self, request: &AuthRequest) -> Result<()> {
211        // Enhanced email validation
212        if !self.is_valid_email(&request.email) {
213            return Err(anyhow::anyhow!("Invalid email format"));
214        }
215
216        // Enhanced password validation
217        if !self.is_valid_password_format(&request.password) {
218            return Err(anyhow::anyhow!(
219                "Password must be at least 8 characters with uppercase, lowercase, and number"
220            ));
221        }
222
223        // Check for common passwords
224        if self.is_common_password(&request.password) {
225            return Err(anyhow::anyhow!("Password is too common"));
226        }
227
228        Ok(())
229    }
230
231    /// Enhanced email validation with regex
232    fn is_valid_email(&self, email: &str) -> bool {
233        !email.is_empty() && email_regex().is_match(email) && email.len() <= 254
234    }
235
236    /// Enhanced password validation
237    fn is_valid_password_format(&self, password: &str) -> bool {
238        password.len() >= 8
239            && password.len() <= 128
240            && password.chars().any(|c| c.is_uppercase())
241            && password.chars().any(|c| c.is_lowercase())
242            && password.chars().any(|c| c.is_numeric())
243    }
244
245    /// Check against common passwords list
246    fn is_common_password(&self, password: &str) -> bool {
247        let common_passwords = vec![
248            "password", "123456", "123456789", "qwerty", "abc123",
249            "password123", "admin", "letmein", "welcome", "monkey"
250        ];
251        common_passwords.contains(&password.to_lowercase().as_str())
252    }
253
254    /// Verify password against stored hash
255    fn verify_password(&self, password: &str, hash: &str) -> Result<bool> {
256        use argon2::{Argon2, PasswordHash, PasswordVerifier};
257
258        let parsed_hash = PasswordHash::new(hash)
259            .map_err(|e| anyhow::anyhow!("Failed to parse password hash: {}", e))?;
260
261        let argon2 = Argon2::default();
262
263        Ok(argon2.verify_password(password.as_bytes(), &parsed_hash).is_ok())
264    }
265
266    /// Check account status (active, suspended, locked, etc.)
267    fn check_account_status(&self, user: &User) -> Result<()> {
268        if !user.is_active {
269            return Err(anyhow::anyhow!("Account is disabled"));
270        }
271
272        if user.is_locked {
273            return Err(anyhow::anyhow!("Account is locked. Please contact support."));
274        }
275
276        if let Some(expires_at) = user.account_expires_at {
277            if chrono::Utc::now() > expires_at {
278                return Err(anyhow::anyhow!("Account has expired"));
279            }
280        }
281
282        if user.requires_password_change {
283            return Err(anyhow::anyhow!("Password change required"));
284        }
285
286        Ok(())
287    }
288
289    /// Generate an access token for the given user.
290    fn generate_access_token(&self, user_id: &Uuid) -> Result<String> {
291        let now = SystemTime::now()
292            .duration_since(UNIX_EPOCH)
293            .map_err(|e| anyhow::anyhow!("System time error: {}", e))?;
294
295        let exp = now + Duration::from_secs(self.config.token_expiry_hours * 3600);
296
297        let claims = Claims {
298            sub: user_id.to_string(),
299            exp: exp.as_secs() as usize,
300            iat: now.as_secs() as usize,
301            iss: "backbone".to_string(),
302        };
303
304        self.jwt_service.create_token(&claims)
305    }
306
307    /// Generate a refresh token if `remember_me` is true; returns `None` otherwise.
308    fn generate_refresh_token(&self, user_id: &Uuid, remember_me: bool) -> Result<Option<String>> {
309        if !remember_me {
310            return Ok(None);
311        }
312
313        let now = SystemTime::now()
314            .duration_since(UNIX_EPOCH)
315            .map_err(|e| anyhow::anyhow!("System time error: {}", e))?;
316
317        let refresh_exp = now + Duration::from_secs(30 * 24 * 3600); // 30 days
318
319        let claims = RefreshTokenClaims {
320            sub: user_id.to_string(),
321            exp: refresh_exp.as_secs() as usize,
322            iat: now.as_secs() as usize,
323            iss: "backbone".to_string(),
324            token_type: "refresh".to_string(),
325        };
326
327        Ok(Some(self.jwt_service.create_refresh_token(&claims)?))
328    }
329
330    /// Generate JWT token
331    pub async fn generate_token(&self, user_id: &Uuid) -> Result<String> {
332        let token = self.generate_token_internal(user_id)?;
333        tracing::info!(
334            event = "auth.token_generated",
335            user_id = %user_id,
336            token_type = "access",
337            "Access token generated"
338        );
339        Ok(token)
340    }
341
342    /// Internal method to generate JWT token
343    fn generate_token_internal(&self, user_id: &Uuid) -> Result<String> {
344        let now = SystemTime::now()
345            .duration_since(UNIX_EPOCH)
346            .map_err(|e| anyhow::anyhow!("System time error: {}", e))?;
347
348        let exp = now + Duration::from_secs(self.config.token_expiry_hours * 3600);
349
350        let claims = Claims {
351            sub: user_id.to_string(),
352            exp: exp.as_secs() as usize,
353            iat: now.as_secs() as usize,
354            iss: "backbone".to_string(),
355        };
356
357        self.jwt_service.create_token(&claims)
358    }
359
360    /// Validate JWT token
361    pub async fn validate_token(&self, token: &str) -> Result<TokenValidation> {
362        match self.jwt_service.validate_token(token) {
363            Ok(claims) => {
364                let user_id = Uuid::parse_str(&claims.sub)
365                    .map_err(|e| anyhow::anyhow!("Invalid user ID in token: {}", e))?;
366
367                tracing::debug!(
368                    event = "auth.token_validated",
369                    user_id = %user_id,
370                    "Token validated successfully"
371                );
372
373                Ok(TokenValidation {
374                    valid: true,
375                    user_id: Some(user_id),
376                })
377            }
378            Err(e) => {
379                tracing::warn!(
380                    event = "auth.token_validation_failed",
381                    reason = %e,
382                    "Token validation failed"
383                );
384                Ok(TokenValidation {
385                    valid: false,
386                    user_id: None,
387                })
388            }
389        }
390    }
391}
392
393/// Authentication result
394#[derive(Debug, Clone)]
395pub struct AuthResult {
396    pub user_id: Uuid,
397    pub token: Option<String>,
398}
399
400impl AuthResult {
401    pub fn new(user_id: Uuid) -> Self {
402        Self {
403            user_id,
404            token: None,
405        }
406    }
407}
408
409/// Token validation result
410#[derive(Debug, Clone)]
411pub struct TokenValidation {
412    pub valid: bool,
413    pub user_id: Option<Uuid>,
414}
415
416impl TokenValidation {
417    pub fn new(valid: bool) -> Self {
418        Self {
419            valid,
420            user_id: None,
421        }
422    }
423}