auth.rs (12581B)
1 use askama::Template; 2 use axum::{ 3 extract::{Query, State}, 4 response::{Html, Redirect}, 5 Form, 6 }; 7 use axum_extra::extract::cookie::{Cookie, CookieJar}; 8 use serde::Deserialize; 9 10 use crate::auth; 11 use crate::error::AppError; 12 use crate::state::AppState; 13 use crate::templates::{ApplicationTemplate, LoginTemplate, RegisterTemplate, ResendVerificationTemplate}; 14 15 const AUTH_COOKIE: &str = "session"; 16 17 /// Extract the session token from cookies. The caller must look it up in the DB. 18 pub fn get_session_token(jar: &CookieJar) -> Option<String> { 19 jar.get(AUTH_COOKIE).map(|c| c.value().to_string()) 20 } 21 22 fn set_auth_cookie(jar: CookieJar, token: &str) -> CookieJar { 23 let cookie = Cookie::build((AUTH_COOKIE, token.to_string())) 24 .path("/") 25 .http_only(true) 26 .secure(true) 27 .same_site(axum_extra::extract::cookie::SameSite::Lax); 28 jar.add(cookie) 29 } 30 31 fn clear_auth_cookie(jar: CookieJar) -> CookieJar { 32 jar.remove(Cookie::from(AUTH_COOKIE)) 33 } 34 35 fn send_verification_email(state: &AppState, email: &str, token: &str) -> Result<(), String> { 36 use std::io::{BufRead, BufReader, Write}; 37 use std::net::TcpStream; 38 39 let verify_url = format!("{}/verify-email?token={}", state.config.site_url, token); 40 41 let mut stream = TcpStream::connect(&*state.config.smtp_addr) 42 .map_err(|e| format!("SMTP connect failed: {e}"))?; 43 stream 44 .set_read_timeout(Some(std::time::Duration::from_secs(10))) 45 .ok(); 46 47 let mut reader = BufReader::new(stream.try_clone().map_err(|e| format!("Clone failed: {e}"))?); 48 49 fn read_reply(reader: &mut BufReader<TcpStream>, expect: &str) -> Result<(), String> { 50 let mut line = String::new(); 51 reader 52 .read_line(&mut line) 53 .map_err(|e| format!("SMTP read failed: {e}"))?; 54 if !line.starts_with(expect) { 55 return Err(format!("SMTP unexpected reply: {}", line.trim())); 56 } 57 Ok(()) 58 } 59 60 read_reply(&mut reader, "220")?; // greeting 61 62 write!(stream, "EHLO localhost\r\n").map_err(|e| format!("SMTP write: {e}"))?; 63 // EHLO can return multiple lines; read until we get a non-continuation line 64 loop { 65 let mut line = String::new(); 66 reader 67 .read_line(&mut line) 68 .map_err(|e| format!("SMTP read: {e}"))?; 69 if line.starts_with("250 ") { 70 break; 71 } 72 if !line.starts_with("250-") { 73 return Err(format!("SMTP EHLO failed: {}", line.trim())); 74 } 75 } 76 77 write!(stream, "MAIL FROM:<{}>\r\n", state.config.smtp_from).map_err(|e| format!("SMTP write: {e}"))?; 78 read_reply(&mut reader, "250")?; 79 80 write!(stream, "RCPT TO:<{email}>\r\n").map_err(|e| format!("SMTP write: {e}"))?; 81 read_reply(&mut reader, "250")?; 82 83 write!(stream, "DATA\r\n").map_err(|e| format!("SMTP write: {e}"))?; 84 read_reply(&mut reader, "354")?; 85 86 write!( 87 stream, 88 "From: <{}>\r\nTo: <{}>\r\nSubject: Verify your email\r\nContent-Type: text/plain; charset=utf-8\r\n\r\nClick to verify your email: {}\r\n.\r\n", 89 state.config.smtp_from, email, verify_url 90 ) 91 .map_err(|e| format!("SMTP write: {e}"))?; 92 read_reply(&mut reader, "250")?; 93 94 write!(stream, "QUIT\r\n").map_err(|e| format!("SMTP write: {e}"))?; 95 Ok(()) 96 } 97 98 pub async fn show_login() -> Result<Html<String>, AppError> { 99 let login = LoginTemplate { error: None }; 100 let app = ApplicationTemplate { 101 content: login.render()?, 102 }; 103 Ok(Html(app.render()?)) 104 } 105 106 pub async fn show_register() -> Result<Html<String>, AppError> { 107 let register = RegisterTemplate { 108 error: None, 109 username: String::new(), 110 email: String::new(), 111 }; 112 let app = ApplicationTemplate { 113 content: register.render()?, 114 }; 115 Ok(Html(app.render()?)) 116 } 117 118 #[derive(Debug, Deserialize)] 119 pub struct LoginForm { 120 pub email: String, 121 pub password: String, 122 } 123 124 pub async fn login( 125 State(state): State<AppState>, 126 jar: CookieJar, 127 Form(form): Form<LoginForm>, 128 ) -> Result<(CookieJar, Redirect), Html<String>> { 129 let email = form.email.clone(); 130 let password = form.password.clone(); 131 132 // Look up user (DB read — uses read pool) 133 let user = state.db().get_user_by_email(&email).await 134 .and_then(|opt| opt.ok_or_else(|| AppError::Unauthorized("Invalid email or password".to_string()))) 135 .map_err(|e| render_login_error(&format!("{e}")))?; 136 137 if !user.email_verified { 138 return Err(render_login_error( 139 "Please verify your email before logging in. Check your inbox or visit /resend-verification to get a new link." 140 )); 141 } 142 143 // Verify password (CPU-bound — runs on tokio's blocking pool, 144 // which is separate from the dedicated DB thread pool) 145 let hash = user.password_hash.clone(); 146 let valid = tokio::task::spawn_blocking(move || { 147 auth::verify_password(&password, &hash) 148 }).await 149 .map_err(|e| render_login_error(&format!("{e}")))? 150 .map_err(|e| render_login_error(&format!("{e}")))?; 151 152 if !valid { 153 return Err(render_login_error("Invalid email or password")); 154 } 155 156 // Create session (DB write) 157 let result = state.db().create_session(user.id).await; 158 159 match result { 160 Ok(token) => { 161 let jar = set_auth_cookie(jar, &token); 162 Ok((jar, Redirect::to("/"))) 163 } 164 Err(e) => Err(render_login_error(&format!("{}", e))), 165 } 166 } 167 168 #[derive(Debug, Deserialize)] 169 pub struct RegisterForm { 170 pub username: String, 171 pub email: String, 172 pub password: String, 173 pub password_confirm: String, 174 } 175 176 pub async fn register( 177 State(state): State<AppState>, 178 Form(form): Form<RegisterForm>, 179 ) -> Result<Html<String>, Html<String>> { 180 // Validate username 181 if form.username.trim().is_empty() { 182 return Err(render_register_error(&form, "Username is required")); 183 } 184 185 if form.username.len() < 2 || form.username.len() > 30 { 186 return Err(render_register_error( 187 &form, 188 "Username must be between 2 and 30 characters", 189 )); 190 } 191 192 if !form 193 .username 194 .chars() 195 .all(|c| c.is_alphanumeric() || c == '_') 196 { 197 return Err(render_register_error( 198 &form, 199 "Username can only contain letters, numbers, and underscores", 200 )); 201 } 202 203 if form.password != form.password_confirm { 204 return Err(render_register_error(&form, "Passwords do not match")); 205 } 206 207 let username = form.username.clone(); 208 let email = form.email.clone(); 209 let password = form.password.clone(); 210 let token = { 211 use rand::Rng; 212 let bytes: [u8; 16] = rand::thread_rng().gen(); 213 bytes.iter().map(|b| format!("{:02x}", b)).collect::<String>() 214 }; 215 let token_clone = token.clone(); 216 217 // Check uniqueness (DB reads) 218 let db = state.db(); 219 if db.get_user_by_username(&username).await 220 .map_err(|e| render_register_error(&form, &format!("{e}")))? 221 .is_some() 222 { 223 return Err(render_register_error(&form, "Username is already taken")); 224 } 225 if db.get_user_by_email(&email).await 226 .map_err(|e| render_register_error(&form, &format!("{e}")))? 227 .is_some() 228 { 229 return Err(render_register_error(&form, "Email is already registered")); 230 } 231 232 // Hash password (CPU-bound — uses semaphore) 233 // Hash password (CPU-bound — runs on tokio's blocking pool, 234 // separate from DB thread pool) 235 let password_hash = tokio::task::spawn_blocking(move || { 236 auth::hash_password(&password) 237 }).await 238 .map_err(|e| render_register_error(&form, &format!("{e}")))? 239 .map_err(|e| render_register_error(&form, &format!("{e}")))?; 240 241 // Create user (DB write) 242 let result = state.db().create_user(&email, &password_hash, &username, Some(&token_clone)) 243 .await.map(|_| ()); 244 245 match result { 246 Ok(()) => { 247 // Send verification email 248 let email_addr = form.email.clone(); 249 if let Err(e) = send_verification_email(&state, &email_addr, &token) { 250 eprintln!("Failed to send verification email: {}", e); 251 } 252 253 // Show "check your email" message instead of auto-login 254 let template = ResendVerificationTemplate { 255 error: None, 256 message: Some("Registration successful! Please check your email to verify your account.".to_string()), 257 }; 258 let app = ApplicationTemplate { 259 content: template.render().unwrap_or_default(), 260 }; 261 Ok(Html(app.render().unwrap_or_default())) 262 } 263 Err(e) => Err(render_register_error(&form, &format!("{}", e))), 264 } 265 } 266 267 #[derive(Deserialize)] 268 pub struct VerifyEmailQuery { 269 pub token: String, 270 } 271 272 pub async fn verify_email( 273 State(state): State<AppState>, 274 Query(query): Query<VerifyEmailQuery>, 275 ) -> Result<Redirect, Html<String>> { 276 let token = query.token; 277 278 let result = state.db().verify_email_token(&token).await; 279 280 match result { 281 Ok(true) => Ok(Redirect::to("/login")), 282 Ok(false) => { 283 let template = ResendVerificationTemplate { 284 error: Some("Invalid or expired verification token.".to_string()), 285 message: None, 286 }; 287 let app = ApplicationTemplate { 288 content: template.render().unwrap_or_default(), 289 }; 290 Err(Html(app.render().unwrap_or_default())) 291 } 292 Err(e) => { 293 let template = ResendVerificationTemplate { 294 error: Some(format!("Verification failed: {}", e)), 295 message: None, 296 }; 297 let app = ApplicationTemplate { 298 content: template.render().unwrap_or_default(), 299 }; 300 Err(Html(app.render().unwrap_or_default())) 301 } 302 } 303 } 304 305 pub async fn show_resend_verification() -> Result<Html<String>, AppError> { 306 let template = ResendVerificationTemplate { 307 error: None, 308 message: None, 309 }; 310 let app = ApplicationTemplate { 311 content: template.render()?, 312 }; 313 Ok(Html(app.render()?)) 314 } 315 316 #[derive(Deserialize)] 317 pub struct ResendVerificationForm { 318 pub email: String, 319 } 320 321 pub async fn resend_verification( 322 State(state): State<AppState>, 323 Form(form): Form<ResendVerificationForm>, 324 ) -> Result<Html<String>, AppError> { 325 let email = form.email.clone(); 326 let new_token = { 327 use rand::Rng; 328 let bytes: [u8; 16] = rand::thread_rng().gen(); 329 bytes.iter().map(|b| format!("{:02x}", b)).collect::<String>() 330 }; 331 let new_token_clone = new_token.clone(); 332 333 let db = state.db(); 334 let user = db.get_user_by_email(&email).await?; 335 let user_found = match user { 336 Some(u) if !u.email_verified => { 337 db.update_verification_token(u.id, &new_token_clone).await?; 338 true 339 } 340 Some(_) => false, // already verified 341 None => false, // user not found 342 }; 343 344 if user_found { 345 if let Err(e) = send_verification_email(&state, &form.email, &new_token) { 346 eprintln!("Failed to send verification email: {}", e); 347 } 348 } 349 350 // Always show the same message to avoid leaking whether the email exists 351 let template = ResendVerificationTemplate { 352 error: None, 353 message: Some("If an account exists with that email, a new verification link has been sent.".to_string()), 354 }; 355 let app = ApplicationTemplate { 356 content: template.render()?, 357 }; 358 Ok(Html(app.render()?)) 359 } 360 361 pub async fn logout( 362 State(state): State<AppState>, 363 jar: CookieJar, 364 ) -> (CookieJar, Redirect) { 365 if let Some(token) = get_session_token(&jar) { 366 let _ = state.db().delete_session(&token).await; 367 } 368 let jar = clear_auth_cookie(jar); 369 (jar, Redirect::to("/")) 370 } 371 372 fn render_login_error(error: &str) -> Html<String> { 373 let login = LoginTemplate { 374 error: Some(error.to_string()), 375 }; 376 let app = ApplicationTemplate { 377 content: login.render().unwrap_or_default(), 378 }; 379 Html(app.render().unwrap_or_default()) 380 } 381 382 fn render_register_error(form: &RegisterForm, error: &str) -> Html<String> { 383 let register = RegisterTemplate { 384 error: Some(error.to_string()), 385 username: form.username.clone(), 386 email: form.email.clone(), 387 }; 388 let app = ApplicationTemplate { 389 content: register.render().unwrap_or_default(), 390 }; 391 Html(app.render().unwrap_or_default()) 392 }