From e7eaa420641b157747c713fcd3a180389c22dea6 Mon Sep 17 00:00:00 2001 From: stefiosif Date: Tue, 26 May 2026 02:34:23 +0300 Subject: [PATCH] feat(backend): replace auth stub with real session, create users and sessions tables --- backend/Cargo.lock | 40 ++- backend/Cargo.toml | 2 + .../20260523200650_create_users.sql | 6 + .../20260523200652_create_sessions.sql | 6 + backend/src/ctx.rs | 6 +- backend/src/error.rs | 4 +- backend/src/main.rs | 18 +- backend/src/model.rs | 320 +++++++++++++++++- backend/src/web/mw_auth.rs | 47 +-- backend/src/web/routes_file.rs | 89 +++-- backend/src/web/routes_login.rs | 233 ++++++++++--- 11 files changed, 636 insertions(+), 135 deletions(-) create mode 100644 backend/migrations/20260523200650_create_users.sql create mode 100644 backend/migrations/20260523200652_create_sessions.sql diff --git a/backend/Cargo.lock b/backend/Cargo.lock index fc12f9a..a94d377 100644 --- a/backend/Cargo.lock +++ b/backend/Cargo.lock @@ -32,6 +32,18 @@ version = "1.0.102" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" +[[package]] +name = "argon2" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3c3610892ee6e0cbce8ae2700349fcf8f98adb0dbfbee85aec3c9179d29cc072" +dependencies = [ + "base64ct", + "blake2", + "cpufeatures 0.2.17", + "password-hash", +] + [[package]] name = "atoi" version = "2.0.0" @@ -173,12 +185,14 @@ name = "backend" version = "0.1.0" dependencies = [ "anyhow", + "argon2", "axum", "axum-test", "chrono", "dotenvy", "lazy-regex", "opendal", + "rand 0.10.1", "serde", "serde_json", "serial_test", @@ -222,6 +236,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "blake2" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "46502ad458c9a52b69d4d4d32775c788b7a1b85e8bc9d482d92250fc0e3f8efe" +dependencies = [ + "digest", +] + [[package]] name = "block-buffer" version = "0.10.4" @@ -1735,6 +1758,17 @@ dependencies = [ "windows-link", ] +[[package]] +name = "password-hash" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "346f04948ba92c43e8469c1ee6736c7563d71012b17d40745260fe106aac2166" +dependencies = [ + "base64ct", + "rand_core 0.6.4", + "subtle", +] + [[package]] name = "pem-rfc7468" version = "0.7.0" @@ -1967,9 +2001,9 @@ dependencies = [ [[package]] name = "rand" -version = "0.10.0" +version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bc266eb313df6c5c09c1c7b1fbe2510961e5bcd3add930c1e31f7ed9da0feff8" +checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207" dependencies = [ "chacha20", "getrandom 0.4.2", @@ -2181,7 +2215,7 @@ dependencies = [ "futures-util", "http", "mime", - "rand 0.10.0", + "rand 0.10.1", "thiserror", ] diff --git a/backend/Cargo.toml b/backend/Cargo.toml index 78ecad1..8dd0d86 100644 --- a/backend/Cargo.toml +++ b/backend/Cargo.toml @@ -18,6 +18,8 @@ sqlx = { version = "0.8", features = ["runtime-tokio-rustls", "postgres", "uuid" chrono = { version = "0.4.44", features = ["serde"] } dotenvy = "0.15.7" opendal = { version = "0.56.0", features = ["tests", "services-fs"] } +argon2 = "0.5.3" +rand = "0.10.1" [dev-dependencies] axum-test = "20.0.0" diff --git a/backend/migrations/20260523200650_create_users.sql b/backend/migrations/20260523200650_create_users.sql new file mode 100644 index 0000000..8c1ff18 --- /dev/null +++ b/backend/migrations/20260523200650_create_users.sql @@ -0,0 +1,6 @@ +CREATE TABLE users ( + id BIGSERIAL PRIMARY KEY, + username TEXT NOT NULL, + password_hash TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); diff --git a/backend/migrations/20260523200652_create_sessions.sql b/backend/migrations/20260523200652_create_sessions.sql new file mode 100644 index 0000000..891523a --- /dev/null +++ b/backend/migrations/20260523200652_create_sessions.sql @@ -0,0 +1,6 @@ +CREATE TABLE sessions ( + id TEXT PRIMARY KEY, + user_id BIGINT NOT NULL REFERENCES users(id), + expires_at TIMESTAMPTZ NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); diff --git a/backend/src/ctx.rs b/backend/src/ctx.rs index 7034526..14f4a29 100644 --- a/backend/src/ctx.rs +++ b/backend/src/ctx.rs @@ -1,14 +1,14 @@ #[derive(Clone)] pub struct Ctx { - user_id: u64, + user_id: i64, } impl Ctx { - pub fn new(user_id: u64) -> Self { + pub fn new(user_id: i64) -> Self { Self { user_id } } - pub fn user_id(&self) -> u64 { + pub fn user_id(&self) -> i64 { self.user_id } } diff --git a/backend/src/error.rs b/backend/src/error.rs index afd02b3..c92cddd 100644 --- a/backend/src/error.rs +++ b/backend/src/error.rs @@ -6,8 +6,8 @@ use tracing::info; #[derive(Debug, Clone)] pub enum LoftError { LoginFail, + RegisterFail, AuthFailNoAuthTokenCookie, - AuthFailTokenWrongFormat, AuthFailCtxNotInRequestExt, FileIdNotFound, UndefinedErrorType, @@ -25,8 +25,8 @@ impl IntoResponse for LoftError { fn into_response(self) -> axum::response::Response { match self { Self::LoginFail + | Self::RegisterFail | Self::AuthFailNoAuthTokenCookie - | Self::AuthFailTokenWrongFormat | Self::AuthFailCtxNotInRequestExt => { info!("UNAUTHORIZED"); StatusCode::UNAUTHORIZED.into_response() diff --git a/backend/src/main.rs b/backend/src/main.rs index e14143f..c662fe3 100644 --- a/backend/src/main.rs +++ b/backend/src/main.rs @@ -11,18 +11,19 @@ use axum::{ middleware, response::Response, }; +use sqlx::PgPool; use tower_cookies::CookieManagerLayer; use tower_http::{cors::CorsLayer, services::ServeDir, trace::TraceLayer}; use tracing::info; use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt}; use crate::{ - model::FileRepository, + model::{FileRepository, UserRepository}, web::{ mw_auth::{mw_ctx_resolver, mw_require_auth}, routes_file::routes_file, routes_health::routes_health, - routes_login::routes_login, + routes_login::routes_auth, }, }; @@ -37,20 +38,25 @@ async fn main() -> Result<()> { .with(tracing_subscriber::fmt::layer()) .init(); - let file_repository = FileRepository::new().await?; - + dotenvy::dotenv().ok(); + let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); + let pool = PgPool::connect(&database_url).await.unwrap(); + let file_repository = FileRepository::new(pool.clone())?; let routes_file = routes_file(file_repository.clone()) .route_layer(middleware::from_fn(mw_require_auth)) .layer(DefaultBodyLimit::disable()); + let user_repository = UserRepository::new(pool)?; + let routes_auth = routes_auth(user_repository.clone()); + let app = Router::new() .nest("/api", routes_file) + .nest("/api/auth", routes_auth) .merge(routes_health()) - .merge(routes_login()) .layer(TraceLayer::new_for_http()) .layer(middleware::map_response(main_response_mapper)) .layer(middleware::from_fn_with_state( - file_repository, + user_repository, mw_ctx_resolver, )) .layer(CookieManagerLayer::new()) diff --git a/backend/src/model.rs b/backend/src/model.rs index cbbf4ab..ada7880 100644 --- a/backend/src/model.rs +++ b/backend/src/model.rs @@ -43,11 +43,7 @@ pub struct FileRepository { } impl FileRepository { - pub async fn new() -> Result { - dotenvy::dotenv().ok(); - let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); - let pool = PgPool::connect(&database_url).await.unwrap(); - + pub fn new(pool: PgPool) -> Result { let storage_path = std::env::var("STORAGE_PATH").expect("STORAGE_PATH must be set"); let op = Operator::new(services::Fs::default().root(&storage_path)) .unwrap() @@ -129,10 +125,10 @@ impl FileRepository { file_id ); let record = self.get_file(file_id).await?; - info!("Downloading file bytes \"{}\"", file_id); + info!("Deleting file bytes \"{}\"", file_id); self.op.delete(&record.storage_path).await.unwrap(); - info!("Downloading file record \"{}\"", file_id); + info!("Deleting file record \"{}\"", file_id); sqlx::query_as!( FileRecord, r#" @@ -164,11 +160,143 @@ impl FileRepository { } } +#[derive(Clone, Debug, FromRow)] +pub struct User { + pub id: i64, + pub username: String, + pub password_hash: String, + pub created_at: chrono::DateTime, +} + +#[derive(Clone, Debug, FromRow)] +pub struct Session { + pub id: String, + pub user_id: i64, + pub expires_at: chrono::DateTime, + pub created_at: chrono::DateTime, +} + +#[derive(Clone)] +pub struct UserRepository { + pub pool: PgPool, +} + +impl UserRepository { + pub fn new(pool: PgPool) -> Result { + Ok(Self { pool }) + } + + pub async fn create_user( + &self, + username: String, + password_hash: String, + ) -> Result { + let user = sqlx::query_as!( + User, + r#" + INSERT INTO users (username, password_hash) + VALUES ($1, $2) + RETURNING * + "#, + username, + password_hash + ) + .fetch_one(&self.pool) + .await + .unwrap(); + + info!("Persisted user: {}", username); + + Ok(user) + } + + pub async fn find_by_username(&self, username: String) -> Result { + info!("Fetching username \"{}\" from users", username); + let user = sqlx::query_as!( + User, + r#" + SELECT * + FROM users u + WHERE u.username = $1 + "#, + username + ) + .fetch_optional(&self.pool) + .await + .unwrap() + .ok_or(LoftError::LoginFail)?; + + Ok(user) + } + + pub async fn create_session( + &self, + user_id: i64, + token: String, + expires_at: chrono::DateTime, + ) -> Result { + let session = sqlx::query_as!( + Session, + r#" + INSERT INTO sessions (id, user_id, expires_at) + VALUES ($1, $2, $3) + RETURNING * + "#, + token, + user_id, + expires_at + ) + .fetch_one(&self.pool) + .await + .unwrap(); + + info!( + "Persisted session: {} with expiration date: {}", + token, expires_at + ); + + Ok(session) + } + + pub async fn get_session(&self, token: String) -> Result { + info!("Fetching session \"{}\" from sessions", token); + let session = sqlx::query_as!( + Session, + r#" + SELECT * + FROM sessions s + WHERE s.id = $1 + AND s.expires_at > NOW() + "#, + token + ) + .fetch_optional(&self.pool) + .await + .unwrap() + .ok_or(LoftError::LoginFail)?; + + Ok(session) + } + + pub async fn delete_session(&self, token: String) -> Result<(), LoftError> { + sqlx::query!(r#"DELETE FROM sessions WHERE id = $1"#, token) + .execute(&self.pool) + .await + .unwrap(); + + info!("Deleted session \"{}\"", token); + + Ok(()) + } +} + #[cfg(test)] mod tests { + use rand::RngExt; + use super::*; - async fn truncate(pool: &PgPool) { + async fn truncate_file_records(pool: &PgPool) { sqlx::query!("TRUNCATE TABLE file_records") .execute(pool) .await @@ -176,14 +304,17 @@ mod tests { } async fn file_repository() -> Result { - Ok(FileRepository::new().await.unwrap()) + dotenvy::dotenv().ok(); + let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); + let pool = PgPool::connect(&database_url).await.unwrap(); + Ok(FileRepository::new(pool)?) } #[tokio::test] #[serial_test::serial] async fn test_upload_and_list() { let file_repository = file_repository().await.unwrap(); - truncate(&file_repository.pool).await; + truncate_file_records(&file_repository.pool).await; file_repository .upload_file(vec![0u8; 10], "a.jpg".to_string()) @@ -196,14 +327,14 @@ mod tests { let files = file_repository.list_files().await.unwrap(); assert_eq!(files.len(), 2); - truncate(&file_repository.pool).await; + truncate_file_records(&file_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_download() { let file_repository = file_repository().await.unwrap(); - truncate(&file_repository.pool).await; + truncate_file_records(&file_repository.pool).await; let uploaded = file_repository .upload_file(vec![0u8; 10], "a.jpg".to_string()) @@ -212,14 +343,14 @@ mod tests { let downloaded = file_repository.download_file(uploaded.id).await.unwrap(); assert!(!downloaded.is_empty()); - truncate(&file_repository.pool).await; + truncate_file_records(&file_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_download_not_found() { let file_repository = file_repository().await.unwrap(); - truncate(&file_repository.pool).await; + truncate_file_records(&file_repository.pool).await; assert!(matches!( file_repository.download_file(i64::MAX).await, @@ -231,7 +362,7 @@ mod tests { #[serial_test::serial] async fn test_delete() { let file_repository = file_repository().await.unwrap(); - truncate(&file_repository.pool).await; + truncate_file_records(&file_repository.pool).await; let uploaded = file_repository .upload_file(vec![0u8; 10], "a.jpg".to_string()) @@ -243,7 +374,7 @@ mod tests { file_repository.download_file(uploaded.id).await, Err(LoftError::FileIdNotFound) )); - truncate(&file_repository.pool).await; + truncate_file_records(&file_repository.pool).await; } #[tokio::test] @@ -256,4 +387,161 @@ mod tests { Err(LoftError::FileIdNotFound) )); } + + 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() -> Result { + dotenvy::dotenv().ok(); + let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); + let pool = PgPool::connect(&database_url).await.unwrap(); + Ok(UserRepository::new(pool)?) + } + + #[tokio::test] + #[serial_test::serial] + async fn test_create_user() { + let user_repository = user_repository().await.unwrap(); + truncate_users(&user_repository.pool).await; + + let username = "picolo".to_string(); + let user = user_repository + .create_user(username.clone(), "pw".to_string()) + .await + .unwrap(); + + assert_eq!(user.username, username); + truncate_users(&user_repository.pool).await; + } + + #[tokio::test] + #[serial_test::serial] + async fn test_find_by_username() { + let user_repository = user_repository().await.unwrap(); + truncate_users(&user_repository.pool).await; + + let username = "picolo".to_string(); + user_repository + .create_user(username.clone(), "pw".to_string()) + .await + .unwrap(); + + let fetched_user = user_repository + .find_by_username(username.clone()) + .await + .unwrap(); + + assert_eq!(fetched_user.username, username); + truncate_users(&user_repository.pool).await; + } + + fn token() -> String { + rand::rng() + .sample_iter(&rand::distr::Alphanumeric) + .take(64) + .map(char::from) + .collect() + } + + #[tokio::test] + #[serial_test::serial] + async fn test_create_session() { + let user_repository = user_repository().await.unwrap(); + truncate_sessions(&user_repository.pool).await; + truncate_users(&user_repository.pool).await; + + let user = user_repository + .create_user("picolo".to_string(), "pw".to_string()) + .await + .unwrap(); + + let token = token(); + + let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2); + let session = user_repository + .create_session(user.id, token.clone(), expires_at) + .await + .unwrap(); + + assert_eq!(session.id, token); + assert_eq!(session.user_id, user.id); + assert!(expires_at > chrono::Utc::now()); + truncate_sessions(&user_repository.pool).await; + truncate_users(&user_repository.pool).await; + } + + #[tokio::test] + #[serial_test::serial] + async fn test_get_session() { + let user_repository = user_repository().await.unwrap(); + truncate_sessions(&user_repository.pool).await; + truncate_users(&user_repository.pool).await; + + let user = user_repository + .create_user("picolo".to_string(), "pw".to_string()) + .await + .unwrap(); + + let token = token(); + + let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2); + user_repository + .create_session(user.id, token.clone(), expires_at) + .await + .unwrap(); + + let fetched_session = user_repository.get_session(token.clone()).await.unwrap(); + + assert_eq!(fetched_session.id, token); + assert_eq!(fetched_session.user_id, user.id); + assert!(fetched_session.expires_at > chrono::Utc::now()); + truncate_sessions(&user_repository.pool).await; + truncate_users(&user_repository.pool).await; + } + + #[tokio::test] + #[serial_test::serial] + async fn test_delete_session() { + let user_repository = user_repository().await.unwrap(); + truncate_sessions(&user_repository.pool).await; + truncate_users(&user_repository.pool).await; + + let user = user_repository + .create_user("picolo".to_string(), "pw".to_string()) + .await + .unwrap(); + + let token = token(); + + let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2); + let session = user_repository + .create_session(user.id, token.clone(), expires_at) + .await + .unwrap(); + + assert_eq!(session.id, token); + assert_eq!(session.user_id, user.id); + assert!(expires_at > chrono::Utc::now()); + + user_repository.delete_session(token.clone()).await.unwrap(); + + assert!(matches!( + user_repository.get_session(token.clone()).await, + Err(LoftError::LoginFail) + )); + + truncate_sessions(&user_repository.pool).await; + truncate_users(&user_repository.pool).await; + } } diff --git a/backend/src/web/mw_auth.rs b/backend/src/web/mw_auth.rs index 9de7c94..9d26532 100644 --- a/backend/src/web/mw_auth.rs +++ b/backend/src/web/mw_auth.rs @@ -1,14 +1,12 @@ use axum::{ - extract::{FromRequestParts, Request}, + extract::{FromRequestParts, Request, State}, middleware::Next, response::Response, }; -use lazy_regex::regex_captures; -use tower_cookies::{Cookie, Cookies}; +use tower_cookies::Cookies; -use crate::{ctx::Ctx, error::LoftError, web::AUTH_TOKEN}; +use crate::{ctx::Ctx, error::LoftError, model::UserRepository, web::AUTH_TOKEN}; -/// validates the cookie exists and is well-formed (3-part format) pub async fn mw_require_auth( ctx: Result, req: Request, @@ -20,35 +18,27 @@ pub async fn mw_require_auth( } pub async fn mw_ctx_resolver( + State(user_repository): State, cookies: Cookies, mut req: Request, next: Next, ) -> Result { let auth_token = cookies.get(AUTH_TOKEN).map(|c| c.value().to_string()); - let result_ctx = match auth_token - .ok_or(LoftError::AuthFailNoAuthTokenCookie) - .and_then(parse_auth_token) - { - Ok((user_id, _, _)) => { - //TODO: add validation - Ok(Ctx::new(user_id)) - } - Err(e) => Err(e), + + let result_ctx = match auth_token { + Some(token) => user_repository + .get_session(token) + .await + .map(|s| Ctx::new(s.user_id)), + None => Err(LoftError::AuthFailNoAuthTokenCookie), }; - - if result_ctx.is_err() && !matches!(result_ctx, Err(LoftError::AuthFailNoAuthTokenCookie)) { - cookies.remove(Cookie::from(AUTH_TOKEN)) - } - req.extensions_mut().insert(result_ctx); - Ok(next.run(req).await) } impl FromRequestParts for Ctx { type Rejection = LoftError; - // extracts user_id from the token and makes it available to handlers as an extractor fn from_request_parts( parts: &mut axum::http::request::Parts, _: &S, @@ -62,18 +52,3 @@ impl FromRequestParts for Ctx { } } } - -fn parse_auth_token(auth_token: String) -> Result<(u64, u64, String), LoftError> { - let (_, user_id, expiration, signature) = - regex_captures!(r"^user-(\d+)\.(\d+)\.([a-f0-9]+)$", &auth_token) - .ok_or(LoftError::AuthFailTokenWrongFormat)?; - - let user_id: u64 = user_id - .parse() - .map_err(|_| LoftError::AuthFailTokenWrongFormat)?; - let expiration: u64 = expiration - .parse() - .map_err(|_| LoftError::AuthFailTokenWrongFormat)?; - - Ok((user_id, expiration, signature.to_string())) -} diff --git a/backend/src/web/routes_file.rs b/backend/src/web/routes_file.rs index 75b932a..ea10aa3 100644 --- a/backend/src/web/routes_file.rs +++ b/backend/src/web/routes_file.rs @@ -93,36 +93,71 @@ mod tests { TestServer, multipart::{MultipartForm, Part}, }; + use rand::RngExt; use serde_json::json; use sqlx::PgPool; use tower_cookies::CookieManagerLayer; use crate::{ - model::FileRepository, + model::{FileRepository, UserRepository}, web::{ mw_auth::{mw_ctx_resolver, mw_require_auth}, routes_file::routes_file, }, }; - // Cookie format: user-[user-id].[expiration].[signature] - const AUTH_COOKIE: &str = "auth-token=user-1.0123456789.a1b2c3d4e5f6"; - const BAD_AUTH_COOKIE: &str = "auth-token=user-1.0123456789"; + const BAD_AUTH_COOKIE: &str = "auth-token=user-0123456789"; + + async fn file_repository() -> FileRepository { + dotenvy::dotenv().ok(); + let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); + let pool = PgPool::connect(&database_url).await.unwrap(); + let file_repository = FileRepository::new(pool).unwrap(); + file_repository + } + + async fn user_repository() -> UserRepository { + dotenvy::dotenv().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).unwrap(); + user_repository + } async fn test_server() -> TestServer { - let file_repository = FileRepository::new().await.unwrap(); + let user_repository = user_repository().await; + let file_repository = file_repository().await; let routes_file = routes_file(file_repository.clone()).route_layer(middleware::from_fn(mw_require_auth)); let app = Router::new() .nest("/api", routes_file) .layer(middleware::from_fn_with_state( - file_repository, + user_repository, mw_ctx_resolver, )) .layer(CookieManagerLayer::new()); + TestServer::new(app) } + async fn create_test_session(user_repository: &UserRepository) -> String { + let user = user_repository + .create_user("testuser".to_string(), "hash".to_string()) + .await + .unwrap(); + let token: String = rand::rng() + .sample_iter(&rand::distr::Alphanumeric) + .take(64) + .map(char::from) + .collect(); + let expires_at = chrono::Utc::now() + chrono::Duration::days(1); + user_repository + .create_session(user.id, token.clone(), expires_at) + .await + .unwrap(); + token + } + async fn truncate(pool: &PgPool) { sqlx::query!("TRUNCATE TABLE file_records") .execute(pool) @@ -149,7 +184,7 @@ mod tests { #[tokio::test] #[serial_test::serial] async fn test_requires_auth_post() { - let file_repository = FileRepository::new().await.unwrap(); + let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; @@ -167,13 +202,15 @@ mod tests { #[tokio::test] #[serial_test::serial] async fn test_list_files_empty() { - let file_repository = FileRepository::new().await.unwrap(); + let user_repository = user_repository().await; + let token = create_test_session(&user_repository).await; + let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; server .get("/api/files") - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_ok() .assert_json(&json!([])); @@ -182,13 +219,15 @@ mod tests { #[tokio::test] #[serial_test::serial] async fn test_upload_and_list_files() { - let file_repository = FileRepository::new().await.unwrap(); + let user_repository = user_repository().await; + let token = create_test_session(&user_repository).await; + let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; let res = server .post("/api/files") - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .multipart(MultipartForm::new().add_part( "file", Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"), @@ -201,7 +240,7 @@ mod tests { let list = server .get("/api/files") - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .json::(); @@ -212,13 +251,15 @@ mod tests { #[tokio::test] #[serial_test::serial] async fn test_download_file() { - let file_repository = FileRepository::new().await.unwrap(); + let user_repository = user_repository().await; + let token = create_test_session(&user_repository).await; + let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; let post_res = server .post("/api/files") - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .multipart(MultipartForm::new().add_part( "file", Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"), @@ -227,7 +268,7 @@ mod tests { let id = post_res.json::()["id"].as_i64().unwrap(); let res = server .get(&format!("/api/files/{id}/download")) - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await; res.assert_status_ok(); @@ -238,10 +279,12 @@ mod tests { #[tokio::test] #[serial_test::serial] async fn test_download_file_not_found() { + let user_repository = user_repository().await; + let token = create_test_session(&user_repository).await; let server = test_server().await; server .get("/api/files/99/download") - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_not_found(); } @@ -249,13 +292,15 @@ mod tests { #[tokio::test] #[serial_test::serial] async fn test_delete_file() { - let file_repository = FileRepository::new().await.unwrap(); + let user_repository = user_repository().await; + let token = create_test_session(&user_repository).await; + let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; let post_res = server .post("/api/files") - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .multipart(MultipartForm::new().add_part( "file", Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"), @@ -265,13 +310,13 @@ mod tests { server .delete(&format!("/api/files/{id}")) - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_ok(); server .get(&format!("/api/files/{id}")) - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_not_found(); @@ -281,10 +326,12 @@ mod tests { #[tokio::test] #[serial_test::serial] async fn test_delete_file_not_found() { + let user_repository = user_repository().await; + let token = create_test_session(&user_repository).await; let server = test_server().await; server .delete("/api/files/99") - .add_header(axum::http::header::COOKIE, AUTH_COOKIE) + .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_not_found(); } diff --git a/backend/src/web/routes_login.rs b/backend/src/web/routes_login.rs index 6c41260..50e6b32 100644 --- a/backend/src/web/routes_login.rs +++ b/backend/src/web/routes_login.rs @@ -1,41 +1,103 @@ -use axum::{ - Json, Router, - routing::{get, post}, +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 serde_json::{Value, json}; +use sqlx::PgPool; use tower_cookies::{Cookie, Cookies}; -use crate::{error::LoftError, web::AUTH_TOKEN}; - -pub fn routes_login() -> Router { +use crate::{error::LoftError, model::UserRepository, web::AUTH_TOKEN}; + +pub fn routes_auth(user_repository: UserRepository) -> Router { Router::new() .route("/login", post(login)) - .route("/register", get(register)) + .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, LoftError> { - //TODO: real db/auth logic - if payload.username != "x" || payload.password != "y" { +) -> Result { + let user = user_repository.find_by_username(payload.username).await?; + // TODO: replace unwrap with ? + let parsed_hash = PasswordHash::new(&user.password_hash).unwrap(); + if Argon2::default() + .verify_password(payload.password.as_bytes(), &parsed_hash) + .is_ok() + { + let expires_at = chrono::Utc::now() + chrono::Duration::days(1); + let cookie = create_cookie(); + let auth_token = cookie.value().to_string(); + cookies.add(cookie); + user_repository + .create_session(user.id, auth_token, expires_at) + .await?; + } else { return Err(LoftError::LoginFail); } - // FIXME: real auth-token generation-signature - cookies.add(Cookie::new(AUTH_TOKEN, "user-1.exp.sign")); - - let body = Json(json!({ - "result": { - "success": true - } - })); - Ok(body) + Ok(StatusCode::OK) } -async fn register() -> &'static str { - "register" +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.clone()).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.clone()) + .await + .is_ok() + { + return Err(LoftError::RegisterFail); + } + + let salt = SaltString::generate(&mut OsRng); + let argon2 = Argon2::default(); + // TODO: replace unwrap + let password_hash = argon2 + .hash_password(payload.password.as_bytes(), &salt) + .unwrap() + .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)] @@ -46,23 +108,54 @@ struct LoginPayload { #[cfg(test)] mod tests { - use axum::{ - Router, - routing::{get, post}, - }; + use axum::{Router, http::StatusCode}; use axum_test::TestServer; use serde_json::json; + use sqlx::PgPool; - use crate::web::routes_login::{login, register}; + 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::dotenv().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).unwrap(); + 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 app = Router::new() - .route(&"/login", post(login)) - .layer(tower_cookies::CookieManagerLayer::new()); - let server = TestServer::new(app); + let server = test_server().await; let response = server - .post("/login") + .post("/api/auth/login") .json(&json!({ "username": "wrong", "password": "wrong", @@ -72,30 +165,74 @@ mod tests { } #[tokio::test] + #[serial_test::serial] async fn test_routes_login() { - let app = Router::new() - .route(&"/login", post(login)) - .layer(tower_cookies::CookieManagerLayer::new()); - let server = TestServer::new(app); + 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("/login") + .post("/api/auth/login") .json(&json!({ - "username": "x", - "password": "y", + "username": "picolo", + "password": "picolo", })) .await; - response.assert_status_ok().assert_json(&json!({ - "result": { - "success": true - } - })); + + response.assert_status_ok(); } #[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(); + } + + #[tokio::test] + #[serial_test::serial] async fn test_routes_register() { - let app = Router::new().route(&"/register", get(register)); - let server = TestServer::new(app); - let response = server.get("/register").await; - response.assert_status_ok().assert_text("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); } }