1use 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
11fn 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#[derive(Debug, Clone)]
22pub struct AuthServiceConfig {
23 pub jwt_secret: String,
24 pub token_expiry_hours: u64,
25 pub key_rotation: Option<KeyRotationConfig>,
27 pub jwt_algorithm: Option<JwtAlgorithm>,
29 pub rsa_private_key_pem: Option<String>,
31 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
48pub 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 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 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 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 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 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 fn validate_auth_request(&self, request: &AuthRequest) -> Result<()> {
211 if !self.is_valid_email(&request.email) {
213 return Err(anyhow::anyhow!("Invalid email format"));
214 }
215
216 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 if self.is_common_password(&request.password) {
225 return Err(anyhow::anyhow!("Password is too common"));
226 }
227
228 Ok(())
229 }
230
231 fn is_valid_email(&self, email: &str) -> bool {
233 !email.is_empty() && email_regex().is_match(email) && email.len() <= 254
234 }
235
236 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 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 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 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 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 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); 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 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 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 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#[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#[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}