From 2616896ecbac7846b399ed01d56e335e7532dcfd Mon Sep 17 00:00:00 2001 From: stefiosif Date: Sat, 13 Jun 2026 12:38:44 +0300 Subject: [PATCH] refactor(backend): replace unwraps with typed errors --- backend/src/error.rs | 59 +++++++++-- backend/src/main.rs | 18 ++-- backend/src/model.rs | 179 ++++++++++---------------------- backend/src/web/mw_auth.rs | 1 + backend/src/web/routes_file.rs | 68 ++++++------ backend/src/web/routes_login.rs | 57 +++++----- 6 files changed, 181 insertions(+), 201 deletions(-) diff --git a/backend/src/error.rs b/backend/src/error.rs index c92cddd..650f628 100644 --- a/backend/src/error.rs +++ b/backend/src/error.rs @@ -1,16 +1,24 @@ use std::fmt; -use axum::{http::StatusCode, response::IntoResponse}; -use tracing::info; +use axum::{extract::multipart::MultipartError, http::StatusCode, response::IntoResponse}; +use tracing::{error, info}; + +pub type Result = std::result::Result; #[derive(Debug, Clone)] pub enum LoftError { LoginFail, RegisterFail, AuthFailNoAuthTokenCookie, + AuthFailSessionNotFound, AuthFailCtxNotInRequestExt, FileIdNotFound, - UndefinedErrorType, + DatabaseError(String), + ArgonError(String), + OpenDalError(String), + NoFileProvided, + MultipartError(String), + InvalidRange, } impl fmt::Display for LoftError { @@ -19,6 +27,30 @@ impl fmt::Display for LoftError { } } +impl From for LoftError { + fn from(value: sqlx::Error) -> Self { + LoftError::DatabaseError(value.to_string()) + } +} + +impl From for LoftError { + fn from(value: MultipartError) -> Self { + LoftError::MultipartError(value.to_string()) + } +} + +impl From for LoftError { + fn from(value: opendal::Error) -> Self { + LoftError::OpenDalError(value.to_string()) + } +} + +impl From for LoftError { + fn from(value: argon2::password_hash::Error) -> Self { + LoftError::ArgonError(value.to_string()) + } +} + impl std::error::Error for LoftError {} impl IntoResponse for LoftError { @@ -27,7 +59,8 @@ impl IntoResponse for LoftError { Self::LoginFail | Self::RegisterFail | Self::AuthFailNoAuthTokenCookie - | Self::AuthFailCtxNotInRequestExt => { + | Self::AuthFailCtxNotInRequestExt + | Self::AuthFailSessionNotFound => { info!("UNAUTHORIZED"); StatusCode::UNAUTHORIZED.into_response() } @@ -35,10 +68,24 @@ impl IntoResponse for LoftError { info!("NOT_FOUND"); StatusCode::NOT_FOUND.into_response() } - Self::UndefinedErrorType => { - info!("INTERNAL_SERVER_ERROR"); + Self::DatabaseError(e) => { + error!("database error: {e}"); StatusCode::INTERNAL_SERVER_ERROR.into_response() } + Self::ArgonError(e) => { + error!("argon2 error: {e}"); + StatusCode::INTERNAL_SERVER_ERROR.into_response() + } + Self::OpenDalError(e) => { + error!("opendal storage error: {e}"); + StatusCode::INTERNAL_SERVER_ERROR.into_response() + } + Self::NoFileProvided => StatusCode::BAD_REQUEST.into_response(), + Self::MultipartError(e) => { + info!("bad request: {e}"); + StatusCode::BAD_REQUEST.into_response() + } + Self::InvalidRange => StatusCode::BAD_REQUEST.into_response(), } } } diff --git a/backend/src/main.rs b/backend/src/main.rs index 622285a..b34b3a4 100644 --- a/backend/src/main.rs +++ b/backend/src/main.rs @@ -54,31 +54,31 @@ async fn main() -> Result<()> { .burst_size(2) .key_extractor(SmartIpKeyExtractor) .finish() - .unwrap(); + .expect("failed to initialize rate limiter configurations"); let governor_auth_limiter = governor_conf_auth.limiter().clone(); let interval = Duration::from_secs(60); std::thread::spawn(move || { loop { std::thread::sleep(interval); - info!( - "rate limiting auth storage size: {}", - governor_auth_limiter.len() - ); + let len = governor_auth_limiter.len(); + if len > 0 { + info!("rate limiting auth storage size: {len}"); + } governor_auth_limiter.retain_recent(); } }); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); const BODY_LIMIT: usize = 1000 * 1000 * 1000 * 5; - let pool = PgPool::connect(&database_url).await.unwrap(); + let pool = PgPool::connect(&database_url).await?; sqlx::migrate!().run(&pool).await?; 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::max(BODY_LIMIT)); - let user_repository = UserRepository::new(pool)?; + let user_repository = UserRepository::new(pool); let routes_auth = routes_auth(user_repository.clone()).layer(GovernorLayer::new(governor_conf_auth)); @@ -95,7 +95,7 @@ async fn main() -> Result<()> { .layer(CookieManagerLayer::new()) .layer( CorsLayer::new() - .allow_origin("http://localhost:5173".parse::().unwrap()) + .allow_origin("http://localhost:5173".parse::()?) .allow_methods([Method::GET, Method::POST, Method::DELETE]) .allow_credentials(true) .allow_headers([header::CONTENT_TYPE]), @@ -105,7 +105,7 @@ async fn main() -> Result<()> { ); let listener = tokio::net::TcpListener::bind("0.0.0.0:3000").await?; - info!("listening on {}", listener.local_addr().unwrap()); + info!("listening on {}", listener.local_addr()?); axum::serve( listener, diff --git a/backend/src/model.rs b/backend/src/model.rs index 69fd35e..16fbba3 100644 --- a/backend/src/model.rs +++ b/backend/src/model.rs @@ -1,12 +1,10 @@ use axum::{body::Bytes, extract::multipart::MultipartError}; use futures_util::{Stream, StreamExt}; use opendal::{Operator, layers::LoggingLayer, services}; -use serde::{Deserialize, Serialize}; +use serde::Serialize; use sqlx::{PgPool, prelude::FromRow}; -use std::fmt::Display; -use tracing::info; -use crate::error::LoftError; +use crate::error::{LoftError, Result}; #[derive(Clone, Debug, Serialize, FromRow)] #[serde(rename_all = "camelCase")] @@ -21,23 +19,6 @@ pub struct FileRecord { pub uploaded_at: chrono::DateTime, } -#[derive(Clone, Debug, Serialize, Deserialize)] -pub enum FileType { - Image, - Video, - Document, -} - -impl Display for FileType { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - FileType::Image => write!(f, "Image"), - FileType::Video => write!(f, "Video"), - FileType::Document => write!(f, "Document"), - } - } -} - #[derive(Clone)] pub struct FileRepository { pub pool: PgPool, @@ -45,13 +26,11 @@ pub struct FileRepository { } impl FileRepository { - pub fn new(pool: PgPool) -> Result { + 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() + let op = Operator::new(services::Fs::default().root(&storage_path))? .layer(LoggingLayer::default()) .finish(); - //.map_err(|x| LoftError::customerror)?; Ok(Self { pool, op }) } @@ -60,16 +39,15 @@ impl FileRepository { &self, mut file_byte_stream: impl Stream> + Unpin, file_storage_key: &str, - ) -> Result { - let mut writer = self.op.writer(file_storage_key).await.unwrap(); + ) -> Result { + let mut writer = self.op.writer(file_storage_key).await?; let mut total_size = 0; while let Some(chunk) = file_byte_stream.next().await { - let chunk = chunk.unwrap(); + let chunk = chunk?; total_size += chunk.len(); - writer.write(chunk).await.unwrap(); + writer.write(chunk).await?; } - // must - writer.close().await.unwrap(); + writer.close().await?; Ok(total_size) } @@ -79,9 +57,8 @@ impl FileRepository { file_name: &str, file_size: usize, file_storage_key: &str, - ) -> Result { - info!("Saving metadata of file \"{}\" in file_records", file_name); - let file_record = sqlx::query_as!( + ) -> Result { + let record = sqlx::query_as!( FileRecord, r#" INSERT INTO file_records (user_id, name, file_type, size, storage_key) @@ -97,26 +74,19 @@ impl FileRepository { file_storage_key ) .fetch_one(&self.pool) - .await - .unwrap(); + .await?; - Ok(file_record) + Ok(record) } pub async fn download_file( &self, file_id: i64, user_id: i64, - ) -> Result> + use<>, LoftError> { - info!( - "Fetching metadata of file \"{}\" from file_records", - file_id - ); + ) -> Result> + use<>> { let record = self.get_file(file_id, user_id).await?; - info!("Downloading file \"{}\"", file_id); - let reader = self.op.reader(&record.storage_key).await.unwrap(); - let stream = reader.into_bytes_stream(0..).await.unwrap(); - + let reader = self.op.reader(&record.storage_key).await?; + let stream = reader.into_bytes_stream(0..).await?; Ok(stream) } @@ -126,25 +96,15 @@ impl FileRepository { user_id: i64, from: u64, to: u64, - ) -> Result> + use<>, LoftError> { - info!( - "Fetching metadata of file \"{}\" from file_records", - file_id - ); + ) -> Result> + use<>> { let record = self.get_file(file_id, user_id).await?; - info!("Streaming chunk of file \"{}\"", file_id); - let reader = self.op.reader(&record.storage_key).await.unwrap(); - let stream = reader.into_bytes_stream(from..to).await.unwrap(); - + let reader = self.op.reader(&record.storage_key).await?; + let stream = reader.into_bytes_stream(from..to).await?; Ok(stream) } - pub async fn get_file(&self, file_id: i64, user_id: i64) -> Result { - info!( - "Fetching metadata of file \"{}\" from file_records", - file_id - ); - let record = sqlx::query_as!( + pub async fn get_file(&self, file_id: i64, user_id: i64) -> Result { + sqlx::query_as!( FileRecord, r#" SELECT * @@ -156,23 +116,14 @@ impl FileRepository { user_id ) .fetch_optional(&self.pool) - .await - .unwrap() - .ok_or(LoftError::FileIdNotFound)?; - - Ok(record) + .await? + .ok_or(LoftError::FileIdNotFound) } - pub async fn delete_file(&self, file_id: i64, user_id: i64) -> Result { - info!( - "Fetching metadata of file \"{}\" from file_records", - file_id - ); + pub async fn delete_file(&self, file_id: i64, user_id: i64) -> Result { let record = self.get_file(file_id, user_id).await?; - info!("Deleting file bytes \"{}\"", file_id); - self.op.delete(&record.storage_key).await.unwrap(); + self.op.delete(&record.storage_key).await?; - info!("Deleting file record \"{}\"", file_id); sqlx::query_as!( FileRecord, r#" @@ -185,13 +136,12 @@ impl FileRepository { user_id ) .fetch_optional(&self.pool) - .await - .unwrap() + .await? .ok_or(LoftError::FileIdNotFound) } - pub async fn list_files(&self, user_id: i64) -> Result, LoftError> { - let files = sqlx::query_as!( + pub async fn list_files(&self, user_id: i64) -> Result> { + let records = sqlx::query_as!( FileRecord, r#" SELECT * @@ -201,10 +151,9 @@ impl FileRepository { user_id ) .fetch_all(&self.pool) - .await - .unwrap(); + .await?; - Ok(files) + Ok(records) } } @@ -230,15 +179,11 @@ pub struct UserRepository { } impl UserRepository { - pub fn new(pool: PgPool) -> Result { - Ok(Self { pool }) + pub fn new(pool: PgPool) -> Self { + Self { pool } } - pub async fn create_user( - &self, - username: &str, - password_hash: &str, - ) -> Result { + pub async fn create_user(&self, username: &str, password_hash: &str) -> Result { let user = sqlx::query_as!( User, r#" @@ -250,16 +195,12 @@ impl UserRepository { password_hash ) .fetch_one(&self.pool) - .await - .unwrap(); - - info!("Persisted user: {}", username); + .await?; Ok(user) } - pub async fn find_by_username(&self, username: &str) -> Result { - info!("Fetching username \"{}\" from users", username); + pub async fn find_by_username(&self, username: &str) -> Result> { let user = sqlx::query_as!( User, r#" @@ -270,9 +211,7 @@ impl UserRepository { username ) .fetch_optional(&self.pool) - .await - .unwrap() - .ok_or(LoftError::LoginFail)?; + .await?; Ok(user) } @@ -282,8 +221,8 @@ impl UserRepository { user_id: i64, token: &str, expires_at: chrono::DateTime, - ) -> Result { - let session = sqlx::query_as!( + ) -> Result { + let sessions = sqlx::query_as!( Session, r#" INSERT INTO sessions (id, user_id, expires_at) @@ -295,20 +234,13 @@ impl UserRepository { expires_at ) .fetch_one(&self.pool) - .await - .unwrap(); + .await?; - info!( - "Persisted session: {} with expiration date: {}", - token, expires_at - ); - - Ok(session) + Ok(sessions) } - pub async fn get_session(&self, token: &str) -> Result { - info!("Fetching session \"{}\" from sessions", token); - let session = sqlx::query_as!( + pub async fn get_session(&self, token: &str) -> Result> { + let sessions = sqlx::query_as!( Session, r#" SELECT * @@ -319,20 +251,15 @@ impl UserRepository { token ) .fetch_optional(&self.pool) - .await - .unwrap() - .ok_or(LoftError::LoginFail)?; + .await?; - Ok(session) + Ok(sessions) } - pub async fn delete_session(&self, token: &str) -> Result<(), LoftError> { + pub async fn delete_session(&self, token: &str) -> Result<()> { sqlx::query!(r#"DELETE FROM sessions WHERE id = $1"#, token) .execute(&self.pool) - .await - .unwrap(); - - info!("Deleted session \"{}\"", token); + .await?; Ok(()) } @@ -353,7 +280,7 @@ mod tests { .unwrap(); } - async fn file_repository() -> Result { + async fn file_repository() -> Result { 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(); @@ -563,7 +490,7 @@ mod tests { 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(); - Ok(UserRepository::new(pool)?) + Ok(UserRepository::new(pool)) } #[tokio::test] @@ -588,7 +515,11 @@ mod tests { let username = "picolo"; user_repository.create_user(username, "pw").await.unwrap(); - let fetched_user = user_repository.find_by_username(username).await.unwrap(); + let fetched_user = user_repository + .find_by_username(username) + .await + .unwrap() + .unwrap(); assert_eq!(fetched_user.username, username); truncate_users(&user_repository.pool).await; @@ -643,7 +574,7 @@ mod tests { .await .unwrap(); - let fetched_session = user_repository.get_session(&token).await.unwrap(); + let fetched_session = user_repository.get_session(&token).await.unwrap().unwrap(); assert_eq!(fetched_session.id, token); assert_eq!(fetched_session.user_id, user.id); @@ -677,7 +608,7 @@ mod tests { assert!(matches!( user_repository.get_session(&token).await, - Err(LoftError::LoginFail) + Ok(None) )); truncate_sessions(&user_repository.pool).await; diff --git a/backend/src/web/mw_auth.rs b/backend/src/web/mw_auth.rs index a87b755..4df22ae 100644 --- a/backend/src/web/mw_auth.rs +++ b/backend/src/web/mw_auth.rs @@ -29,6 +29,7 @@ pub async fn mw_ctx_resolver( Some(token) => user_repository .get_session(&token) .await + .and_then(|s| s.ok_or(LoftError::AuthFailSessionNotFound)) .map(|s| Ctx::new(s.user_id)), None => Err(LoftError::AuthFailNoAuthTokenCookie), }; diff --git a/backend/src/web/routes_file.rs b/backend/src/web/routes_file.rs index 3625dfe..bda0f54 100644 --- a/backend/src/web/routes_file.rs +++ b/backend/src/web/routes_file.rs @@ -7,11 +7,10 @@ use axum::{ routing::get, }; use sqlx::types::uuid; -use tracing::info; use crate::{ ctx::Ctx, - error::LoftError, + error::{LoftError, Result}, model::{FileRecord, FileRepository}, }; @@ -29,31 +28,24 @@ async fn upload_file( ctx: Ctx, mut multipart: Multipart, ) -> Result, LoftError> { - info!("handler: upload_file"); + let mut uploaded: Option<(String, String, usize)> = None; - let mut file_name = None; - let mut file_storage_key = None; - let mut file_size = None; - - while let Some(field) = multipart.next_field().await.unwrap() { - if field.name().unwrap() == "file" { - let name = field.file_name().map(|s| s.to_string()).unwrap_or_default(); - let key = format!("{}-{}", name, uuid::Uuid::new_v4()); + while let Some(field) = multipart.next_field().await? { + if field.name() == Some("file") { + let name = field.file_name().map(str::to_string).unwrap_or_default(); + let key = format!("{name}-{}", uuid::Uuid::new_v4()); let size = file_repository.upload_file(field, &key).await?; - file_name = Some(name); - file_storage_key = Some(key); - file_size = Some(size); + uploaded = Some((name, key, size)); } } - if let (Some(name), Some(key), Some(size)) = (file_name, file_storage_key, file_size) { - let file_record = file_repository - .create_file_record(ctx.user_id(), &name, size, &key) - .await?; - return Ok(Json(file_record)); - } + let (name, key, size) = uploaded.ok_or(LoftError::NoFileProvided)?; - Err(LoftError::UndefinedErrorType) + let file_record = file_repository + .create_file_record(ctx.user_id(), &name, size, &key) + .await?; + + Ok(Json(file_record)) } async fn get_file( @@ -71,7 +63,7 @@ async fn download_file( State(file_repository): State, ctx: Ctx, Path(file_id): Path, -) -> Result { +) -> Result { let stream = file_repository .download_file(file_id as i64, ctx.user_id()) .await?; @@ -83,13 +75,12 @@ async fn stream_part( ctx: Ctx, headers: HeaderMap, Path(file_id): Path, -) -> Result { - info!("stream_part"); +) -> Result { let file_record = file_repository .get_file(file_id as i64, ctx.user_id()) .await?; let file_size = file_record.size as u64; - let (start, end): (u64, u64) = parse_range(&headers, file_size); + let (start, end): (u64, u64) = parse_range(&headers, file_size)?; let stream = file_repository .stream_part(file_id as i64, ctx.user_id(), start, end + 1) @@ -105,20 +96,25 @@ async fn stream_part( ) .header(header::ACCEPT_RANGES, "bytes") .body(Body::from_stream(stream)) - .unwrap(); + .expect("Failed to build response with valid headers"); Ok(respones) } -fn parse_range(headers: &HeaderMap, file_size: u64) -> (u64, u64) { - let range = headers.get(header::RANGE); - let str = range.unwrap().to_str().unwrap(); - let strip = str.strip_prefix("bytes=").unwrap(); - let tuple = strip.split_once("-").unwrap(); - let start = tuple.0.parse::().unwrap(); +fn parse_range(headers: &HeaderMap, file_size: u64) -> Result<(u64, u64)> { + let range = headers.get(header::RANGE).ok_or(LoftError::InvalidRange)?; + let str = range.to_str().map_err(|_| LoftError::InvalidRange)?; + let tuple = str + .strip_prefix("bytes=") + .and_then(|s| s.split_once("-")) + .ok_or(LoftError::InvalidRange)?; + let start = tuple + .0 + .parse::() + .map_err(|_| LoftError::InvalidRange)?; let end = tuple.1.parse::().unwrap_or(file_size - 1); - (start, end) + Ok((start, end)) } async fn delete_file( @@ -126,8 +122,6 @@ async fn delete_file( ctx: Ctx, Path(file_id): Path, ) -> Result, LoftError> { - info!("handler: delete_file"); - let file = file_repository .delete_file(file_id as i64, ctx.user_id()) .await?; @@ -138,8 +132,6 @@ async fn list_files( State(file_repository): State, ctx: Ctx, ) -> Result>, LoftError> { - info!("handler: list_files"); - let files = file_repository.list_files(ctx.user_id()).await?; Ok(Json(files)) } @@ -178,7 +170,7 @@ mod tests { 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).unwrap(); + let user_repository = UserRepository::new(pool); user_repository } diff --git a/backend/src/web/routes_login.rs b/backend/src/web/routes_login.rs index 07c6d04..a2faf31 100644 --- a/backend/src/web/routes_login.rs +++ b/backend/src/web/routes_login.rs @@ -6,8 +6,13 @@ 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, model::UserRepository, web::AUTH_TOKEN}; +use crate::{ + error::{LoftError, Result}, + model::UserRepository, + web::AUTH_TOKEN, +}; pub fn routes_auth(user_repository: UserRepository) -> Router { Router::new() @@ -22,24 +27,26 @@ async fn login( cookies: Cookies, Json(payload): Json, ) -> 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(); - cookies.add(cookie.clone()); - user_repository - .create_session(user.id, auth_token, expires_at) - .await?; - } else { + 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) } @@ -63,22 +70,24 @@ async fn register( ) -> Result { if user_repository .find_by_username(&payload.username) - .await - .is_ok() + .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(); - // TODO: replace unwrap - let password_hash = argon2 - .hash_password(payload.password.as_bytes(), &salt) - .unwrap() + let password_hash = &argon2 + .hash_password(payload.password.as_bytes(), &salt)? .to_string(); user_repository - .create_user(&payload.username, &password_hash) + .create_user(&payload.username, password_hash) .await?; Ok(StatusCode::CREATED) @@ -135,7 +144,7 @@ mod tests { 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).unwrap(); + let user_repository = UserRepository::new(pool); user_repository }