use argon2::{ Argon2, PasswordHash, PasswordHasher, PasswordVerifier, password_hash::{SaltString, rand_core::OsRng}, }; use axum::{Json, Router, extract::State, http::StatusCode, routing::post}; use rand::RngExt; use serde::Deserialize; use tower_cookies::{Cookie, Cookies}; use tracing::warn; use crate::{ error::{LoftError, Result}, model::UserRepository, web::AUTH_TOKEN, }; pub fn routes_auth(user_repository: UserRepository) -> Router { Router::new() .route("/login", post(login)) .route("/logout", post(logout)) .route("/register", post(register)) .with_state(user_repository) } async fn login( State(user_repository): State, cookies: Cookies, Json(payload): Json, ) -> Result { let user = user_repository .find_by_username(&payload.username) .await? .ok_or(LoftError::LoginFail)?; let parsed_hash = PasswordHash::new(&user.password_hash)?; let password_verification = Argon2::default().verify_password(payload.password.as_bytes(), &parsed_hash); if password_verification.is_err() { return Err(LoftError::LoginFail); } let expires_at = chrono::Utc::now() + chrono::Duration::days(1); let cookie = create_cookie(); let auth_token = cookie.value(); cookies.add(cookie.clone()); user_repository .create_session(user.id, auth_token, expires_at) .await?; Ok(StatusCode::OK) } async fn logout( State(user_repository): State, cookies: Cookies, ) -> Result { let auth_token = cookies.get(AUTH_TOKEN).map(|c| c.value().to_string()); if let Some(auth_token) = auth_token { user_repository.delete_session(&auth_token).await?; cookies.remove(Cookie::build(AUTH_TOKEN).path("/").build()); } Ok(StatusCode::OK) } async fn register( State(user_repository): State, Json(payload): Json, ) -> Result { if user_repository .find_by_username(&payload.username) .await? .is_some() { warn!( "Register fail, username {} already exists", &payload.username ); // also fix "Login fail" typo return Err(LoftError::RegisterFail); } let salt = SaltString::generate(&mut OsRng); let argon2 = Argon2::default(); let password_hash = &argon2 .hash_password(payload.password.as_bytes(), &salt)? .to_string(); user_repository .create_user(&payload.username, password_hash) .await?; Ok(StatusCode::CREATED) } fn create_cookie() -> Cookie<'static> { let auth_token: String = rand::rng() .sample_iter(&rand::distr::Alphanumeric) .take(64) .map(char::from) .collect(); Cookie::build((AUTH_TOKEN, auth_token)) .http_only(true) .same_site(tower_cookies::cookie::SameSite::Lax) .path("/") .secure(std::env::var("ENVIRONMENT").unwrap_or_default() == "prod") .build() } #[derive(Debug, Deserialize)] struct LoginPayload { username: String, password: String, } #[cfg(test)] mod tests { use axum::{Router, http::StatusCode}; use axum_test::TestServer; use serde_json::json; use sqlx::PgPool; use crate::{ model::UserRepository, web::{AUTH_TOKEN, routes_login::routes_auth}, }; async fn truncate_users(pool: &PgPool) { sqlx::query!("TRUNCATE TABLE users CASCADE") .execute(pool) .await .unwrap(); } async fn truncate_sessions(pool: &PgPool) { sqlx::query!("TRUNCATE TABLE sessions") .execute(pool) .await .unwrap(); } async fn user_repository() -> UserRepository { dotenvy::from_filename(".env.test").ok(); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let pool = PgPool::connect(&database_url).await.unwrap(); let user_repository = UserRepository::new(pool); user_repository } async fn test_server() -> TestServer { let user_repository = user_repository().await; let routes_auth = routes_auth(user_repository); let app = Router::new() .nest("/api/auth", routes_auth) .layer(tower_cookies::CookieManagerLayer::new()); TestServer::new(app) } #[tokio::test] #[serial_test::serial] async fn test_routes_login_wrong_credentials() { let user_repository = user_repository().await; truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; let server = test_server().await; let response = server .post("/api/auth/login") .json(&json!({ "username": "wrong", "password": "wrong", })) .await; response.assert_status_unauthorized(); truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_routes_login() { let user_repository = user_repository().await; truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; let server = test_server().await; server .post("/api/auth/register") .json(&json!({ "username": "picolo", "password": "picolo" })) .await; let response = server .post("/api/auth/login") .json(&json!({ "username": "picolo", "password": "picolo", })) .await; response.assert_status_ok(); truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_routes_logout() { let user_repository = user_repository().await; truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; let server = test_server().await; let register_response = server .post("/api/auth/register") .json(&json!({ "username": "picolo", "password": "picolo" })) .await; register_response.assert_status(StatusCode::CREATED); let login_response = server .post("/api/auth/login") .json(&json!({ "username": "picolo", "password": "picolo", })) .await; login_response.assert_status_ok(); let cookie = login_response.cookie(AUTH_TOKEN); let logout_response = server.post("/api/auth/logout").add_cookie(cookie).await; logout_response.assert_status_ok(); truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_routes_register() { let user_repository = user_repository().await; truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; let server = test_server().await; let register_response = server .post("/api/auth/register") .json(&json!({ "username": "picolo", "password": "picolo" })) .await; register_response.assert_status(StatusCode::CREATED); truncate_sessions(&user_repository.pool).await; truncate_users(&user_repository.pool).await; } }