258 lines
7.5 KiB
Rust
258 lines
7.5 KiB
Rust
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<UserRepository>,
|
|
cookies: Cookies,
|
|
Json(payload): Json<LoginPayload>,
|
|
) -> Result<StatusCode, LoftError> {
|
|
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<UserRepository>,
|
|
cookies: Cookies,
|
|
) -> Result<StatusCode, LoftError> {
|
|
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<UserRepository>,
|
|
Json(payload): Json<LoginPayload>,
|
|
) -> Result<StatusCode, LoftError> {
|
|
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;
|
|
}
|
|
}
|