simple-web-app

Unnamed repository; edit this file 'description' to name the repository.
Log | Files | Refs | README

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 }