refactor(backend): replace unwraps with typed errors

This commit is contained in:
2026-06-13 12:38:44 +03:00
parent 21d65af470
commit 2616896ecb
6 changed files with 181 additions and 201 deletions

View File

@@ -1,16 +1,24 @@
use std::fmt; use std::fmt;
use axum::{http::StatusCode, response::IntoResponse}; use axum::{extract::multipart::MultipartError, http::StatusCode, response::IntoResponse};
use tracing::info; use tracing::{error, info};
pub type Result<T, E = LoftError> = std::result::Result<T, E>;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub enum LoftError { pub enum LoftError {
LoginFail, LoginFail,
RegisterFail, RegisterFail,
AuthFailNoAuthTokenCookie, AuthFailNoAuthTokenCookie,
AuthFailSessionNotFound,
AuthFailCtxNotInRequestExt, AuthFailCtxNotInRequestExt,
FileIdNotFound, FileIdNotFound,
UndefinedErrorType, DatabaseError(String),
ArgonError(String),
OpenDalError(String),
NoFileProvided,
MultipartError(String),
InvalidRange,
} }
impl fmt::Display for LoftError { impl fmt::Display for LoftError {
@@ -19,6 +27,30 @@ impl fmt::Display for LoftError {
} }
} }
impl From<sqlx::Error> for LoftError {
fn from(value: sqlx::Error) -> Self {
LoftError::DatabaseError(value.to_string())
}
}
impl From<MultipartError> for LoftError {
fn from(value: MultipartError) -> Self {
LoftError::MultipartError(value.to_string())
}
}
impl From<opendal::Error> for LoftError {
fn from(value: opendal::Error) -> Self {
LoftError::OpenDalError(value.to_string())
}
}
impl From<argon2::password_hash::Error> for LoftError {
fn from(value: argon2::password_hash::Error) -> Self {
LoftError::ArgonError(value.to_string())
}
}
impl std::error::Error for LoftError {} impl std::error::Error for LoftError {}
impl IntoResponse for LoftError { impl IntoResponse for LoftError {
@@ -27,7 +59,8 @@ impl IntoResponse for LoftError {
Self::LoginFail Self::LoginFail
| Self::RegisterFail | Self::RegisterFail
| Self::AuthFailNoAuthTokenCookie | Self::AuthFailNoAuthTokenCookie
| Self::AuthFailCtxNotInRequestExt => { | Self::AuthFailCtxNotInRequestExt
| Self::AuthFailSessionNotFound => {
info!("UNAUTHORIZED"); info!("UNAUTHORIZED");
StatusCode::UNAUTHORIZED.into_response() StatusCode::UNAUTHORIZED.into_response()
} }
@@ -35,10 +68,24 @@ impl IntoResponse for LoftError {
info!("NOT_FOUND"); info!("NOT_FOUND");
StatusCode::NOT_FOUND.into_response() StatusCode::NOT_FOUND.into_response()
} }
Self::UndefinedErrorType => { Self::DatabaseError(e) => {
info!("INTERNAL_SERVER_ERROR"); error!("database error: {e}");
StatusCode::INTERNAL_SERVER_ERROR.into_response() 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(),
} }
} }
} }

View File

@@ -54,31 +54,31 @@ async fn main() -> Result<()> {
.burst_size(2) .burst_size(2)
.key_extractor(SmartIpKeyExtractor) .key_extractor(SmartIpKeyExtractor)
.finish() .finish()
.unwrap(); .expect("failed to initialize rate limiter configurations");
let governor_auth_limiter = governor_conf_auth.limiter().clone(); let governor_auth_limiter = governor_conf_auth.limiter().clone();
let interval = Duration::from_secs(60); let interval = Duration::from_secs(60);
std::thread::spawn(move || { std::thread::spawn(move || {
loop { loop {
std::thread::sleep(interval); std::thread::sleep(interval);
info!( let len = governor_auth_limiter.len();
"rate limiting auth storage size: {}", if len > 0 {
governor_auth_limiter.len() info!("rate limiting auth storage size: {len}");
); }
governor_auth_limiter.retain_recent(); governor_auth_limiter.retain_recent();
} }
}); });
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
const BODY_LIMIT: usize = 1000 * 1000 * 1000 * 5; 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?; sqlx::migrate!().run(&pool).await?;
let file_repository = FileRepository::new(pool.clone())?; let file_repository = FileRepository::new(pool.clone())?;
let routes_file = routes_file(file_repository.clone()) let routes_file = routes_file(file_repository.clone())
.route_layer(middleware::from_fn(mw_require_auth)) .route_layer(middleware::from_fn(mw_require_auth))
.layer(DefaultBodyLimit::max(BODY_LIMIT)); .layer(DefaultBodyLimit::max(BODY_LIMIT));
let user_repository = UserRepository::new(pool)?; let user_repository = UserRepository::new(pool);
let routes_auth = let routes_auth =
routes_auth(user_repository.clone()).layer(GovernorLayer::new(governor_conf_auth)); routes_auth(user_repository.clone()).layer(GovernorLayer::new(governor_conf_auth));
@@ -95,7 +95,7 @@ async fn main() -> Result<()> {
.layer(CookieManagerLayer::new()) .layer(CookieManagerLayer::new())
.layer( .layer(
CorsLayer::new() CorsLayer::new()
.allow_origin("http://localhost:5173".parse::<HeaderValue>().unwrap()) .allow_origin("http://localhost:5173".parse::<HeaderValue>()?)
.allow_methods([Method::GET, Method::POST, Method::DELETE]) .allow_methods([Method::GET, Method::POST, Method::DELETE])
.allow_credentials(true) .allow_credentials(true)
.allow_headers([header::CONTENT_TYPE]), .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?; 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( axum::serve(
listener, listener,

View File

@@ -1,12 +1,10 @@
use axum::{body::Bytes, extract::multipart::MultipartError}; use axum::{body::Bytes, extract::multipart::MultipartError};
use futures_util::{Stream, StreamExt}; use futures_util::{Stream, StreamExt};
use opendal::{Operator, layers::LoggingLayer, services}; use opendal::{Operator, layers::LoggingLayer, services};
use serde::{Deserialize, Serialize}; use serde::Serialize;
use sqlx::{PgPool, prelude::FromRow}; 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)] #[derive(Clone, Debug, Serialize, FromRow)]
#[serde(rename_all = "camelCase")] #[serde(rename_all = "camelCase")]
@@ -21,23 +19,6 @@ pub struct FileRecord {
pub uploaded_at: chrono::DateTime<chrono::Utc>, pub uploaded_at: chrono::DateTime<chrono::Utc>,
} }
#[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)] #[derive(Clone)]
pub struct FileRepository { pub struct FileRepository {
pub pool: PgPool, pub pool: PgPool,
@@ -45,13 +26,11 @@ pub struct FileRepository {
} }
impl FileRepository { impl FileRepository {
pub fn new(pool: PgPool) -> Result<Self, LoftError> { pub fn new(pool: PgPool) -> Result<Self> {
let storage_path = std::env::var("STORAGE_PATH").expect("STORAGE_PATH must be set"); let storage_path = std::env::var("STORAGE_PATH").expect("STORAGE_PATH must be set");
let op = Operator::new(services::Fs::default().root(&storage_path)) let op = Operator::new(services::Fs::default().root(&storage_path))?
.unwrap()
.layer(LoggingLayer::default()) .layer(LoggingLayer::default())
.finish(); .finish();
//.map_err(|x| LoftError::customerror)?;
Ok(Self { pool, op }) Ok(Self { pool, op })
} }
@@ -60,16 +39,15 @@ impl FileRepository {
&self, &self,
mut file_byte_stream: impl Stream<Item = Result<Bytes, MultipartError>> + Unpin, mut file_byte_stream: impl Stream<Item = Result<Bytes, MultipartError>> + Unpin,
file_storage_key: &str, file_storage_key: &str,
) -> Result<usize, LoftError> { ) -> Result<usize> {
let mut writer = self.op.writer(file_storage_key).await.unwrap(); let mut writer = self.op.writer(file_storage_key).await?;
let mut total_size = 0; let mut total_size = 0;
while let Some(chunk) = file_byte_stream.next().await { while let Some(chunk) = file_byte_stream.next().await {
let chunk = chunk.unwrap(); let chunk = chunk?;
total_size += chunk.len(); total_size += chunk.len();
writer.write(chunk).await.unwrap(); writer.write(chunk).await?;
} }
// must writer.close().await?;
writer.close().await.unwrap();
Ok(total_size) Ok(total_size)
} }
@@ -79,9 +57,8 @@ impl FileRepository {
file_name: &str, file_name: &str,
file_size: usize, file_size: usize,
file_storage_key: &str, file_storage_key: &str,
) -> Result<FileRecord, LoftError> { ) -> Result<FileRecord> {
info!("Saving metadata of file \"{}\" in file_records", file_name); let record = sqlx::query_as!(
let file_record = sqlx::query_as!(
FileRecord, FileRecord,
r#" r#"
INSERT INTO file_records (user_id, name, file_type, size, storage_key) INSERT INTO file_records (user_id, name, file_type, size, storage_key)
@@ -97,26 +74,19 @@ impl FileRepository {
file_storage_key file_storage_key
) )
.fetch_one(&self.pool) .fetch_one(&self.pool)
.await .await?;
.unwrap();
Ok(file_record) Ok(record)
} }
pub async fn download_file( pub async fn download_file(
&self, &self,
file_id: i64, file_id: i64,
user_id: i64, user_id: i64,
) -> Result<impl Stream<Item = std::io::Result<Bytes>> + use<>, LoftError> { ) -> Result<impl Stream<Item = std::io::Result<Bytes>> + use<>> {
info!(
"Fetching metadata of file \"{}\" from file_records",
file_id
);
let record = self.get_file(file_id, user_id).await?; let record = self.get_file(file_id, user_id).await?;
info!("Downloading file \"{}\"", file_id); let reader = self.op.reader(&record.storage_key).await?;
let reader = self.op.reader(&record.storage_key).await.unwrap(); let stream = reader.into_bytes_stream(0..).await?;
let stream = reader.into_bytes_stream(0..).await.unwrap();
Ok(stream) Ok(stream)
} }
@@ -126,25 +96,15 @@ impl FileRepository {
user_id: i64, user_id: i64,
from: u64, from: u64,
to: u64, to: u64,
) -> Result<impl Stream<Item = std::io::Result<Bytes>> + use<>, LoftError> { ) -> Result<impl Stream<Item = std::io::Result<Bytes>> + use<>> {
info!(
"Fetching metadata of file \"{}\" from file_records",
file_id
);
let record = self.get_file(file_id, user_id).await?; 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?;
let reader = self.op.reader(&record.storage_key).await.unwrap(); let stream = reader.into_bytes_stream(from..to).await?;
let stream = reader.into_bytes_stream(from..to).await.unwrap();
Ok(stream) Ok(stream)
} }
pub async fn get_file(&self, file_id: i64, user_id: i64) -> Result<FileRecord, LoftError> { pub async fn get_file(&self, file_id: i64, user_id: i64) -> Result<FileRecord> {
info!( sqlx::query_as!(
"Fetching metadata of file \"{}\" from file_records",
file_id
);
let record = sqlx::query_as!(
FileRecord, FileRecord,
r#" r#"
SELECT * SELECT *
@@ -156,23 +116,14 @@ impl FileRepository {
user_id user_id
) )
.fetch_optional(&self.pool) .fetch_optional(&self.pool)
.await .await?
.unwrap() .ok_or(LoftError::FileIdNotFound)
.ok_or(LoftError::FileIdNotFound)?;
Ok(record)
} }
pub async fn delete_file(&self, file_id: i64, user_id: i64) -> Result<FileRecord, LoftError> { pub async fn delete_file(&self, file_id: i64, user_id: i64) -> Result<FileRecord> {
info!(
"Fetching metadata of file \"{}\" from file_records",
file_id
);
let record = self.get_file(file_id, user_id).await?; let record = self.get_file(file_id, user_id).await?;
info!("Deleting file bytes \"{}\"", file_id); self.op.delete(&record.storage_key).await?;
self.op.delete(&record.storage_key).await.unwrap();
info!("Deleting file record \"{}\"", file_id);
sqlx::query_as!( sqlx::query_as!(
FileRecord, FileRecord,
r#" r#"
@@ -185,13 +136,12 @@ impl FileRepository {
user_id user_id
) )
.fetch_optional(&self.pool) .fetch_optional(&self.pool)
.await .await?
.unwrap()
.ok_or(LoftError::FileIdNotFound) .ok_or(LoftError::FileIdNotFound)
} }
pub async fn list_files(&self, user_id: i64) -> Result<Vec<FileRecord>, LoftError> { pub async fn list_files(&self, user_id: i64) -> Result<Vec<FileRecord>> {
let files = sqlx::query_as!( let records = sqlx::query_as!(
FileRecord, FileRecord,
r#" r#"
SELECT * SELECT *
@@ -201,10 +151,9 @@ impl FileRepository {
user_id user_id
) )
.fetch_all(&self.pool) .fetch_all(&self.pool)
.await .await?;
.unwrap();
Ok(files) Ok(records)
} }
} }
@@ -230,15 +179,11 @@ pub struct UserRepository {
} }
impl UserRepository { impl UserRepository {
pub fn new(pool: PgPool) -> Result<Self, LoftError> { pub fn new(pool: PgPool) -> Self {
Ok(Self { pool }) Self { pool }
} }
pub async fn create_user( pub async fn create_user(&self, username: &str, password_hash: &str) -> Result<User> {
&self,
username: &str,
password_hash: &str,
) -> Result<User, LoftError> {
let user = sqlx::query_as!( let user = sqlx::query_as!(
User, User,
r#" r#"
@@ -250,16 +195,12 @@ impl UserRepository {
password_hash password_hash
) )
.fetch_one(&self.pool) .fetch_one(&self.pool)
.await .await?;
.unwrap();
info!("Persisted user: {}", username);
Ok(user) Ok(user)
} }
pub async fn find_by_username(&self, username: &str) -> Result<User, LoftError> { pub async fn find_by_username(&self, username: &str) -> Result<Option<User>> {
info!("Fetching username \"{}\" from users", username);
let user = sqlx::query_as!( let user = sqlx::query_as!(
User, User,
r#" r#"
@@ -270,9 +211,7 @@ impl UserRepository {
username username
) )
.fetch_optional(&self.pool) .fetch_optional(&self.pool)
.await .await?;
.unwrap()
.ok_or(LoftError::LoginFail)?;
Ok(user) Ok(user)
} }
@@ -282,8 +221,8 @@ impl UserRepository {
user_id: i64, user_id: i64,
token: &str, token: &str,
expires_at: chrono::DateTime<chrono::Utc>, expires_at: chrono::DateTime<chrono::Utc>,
) -> Result<Session, LoftError> { ) -> Result<Session> {
let session = sqlx::query_as!( let sessions = sqlx::query_as!(
Session, Session,
r#" r#"
INSERT INTO sessions (id, user_id, expires_at) INSERT INTO sessions (id, user_id, expires_at)
@@ -295,20 +234,13 @@ impl UserRepository {
expires_at expires_at
) )
.fetch_one(&self.pool) .fetch_one(&self.pool)
.await .await?;
.unwrap();
info!( Ok(sessions)
"Persisted session: {} with expiration date: {}",
token, expires_at
);
Ok(session)
} }
pub async fn get_session(&self, token: &str) -> Result<Session, LoftError> { pub async fn get_session(&self, token: &str) -> Result<Option<Session>> {
info!("Fetching session \"{}\" from sessions", token); let sessions = sqlx::query_as!(
let session = sqlx::query_as!(
Session, Session,
r#" r#"
SELECT * SELECT *
@@ -319,20 +251,15 @@ impl UserRepository {
token token
) )
.fetch_optional(&self.pool) .fetch_optional(&self.pool)
.await .await?;
.unwrap()
.ok_or(LoftError::LoginFail)?;
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) sqlx::query!(r#"DELETE FROM sessions WHERE id = $1"#, token)
.execute(&self.pool) .execute(&self.pool)
.await .await?;
.unwrap();
info!("Deleted session \"{}\"", token);
Ok(()) Ok(())
} }
@@ -353,7 +280,7 @@ mod tests {
.unwrap(); .unwrap();
} }
async fn file_repository() -> Result<FileRepository, LoftError> { async fn file_repository() -> Result<FileRepository> {
dotenvy::from_filename(".env.test").ok(); dotenvy::from_filename(".env.test").ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap(); let pool = PgPool::connect(&database_url).await.unwrap();
@@ -563,7 +490,7 @@ mod tests {
dotenvy::from_filename(".env.test").ok(); dotenvy::from_filename(".env.test").ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap(); let pool = PgPool::connect(&database_url).await.unwrap();
Ok(UserRepository::new(pool)?) Ok(UserRepository::new(pool))
} }
#[tokio::test] #[tokio::test]
@@ -588,7 +515,11 @@ mod tests {
let username = "picolo"; let username = "picolo";
user_repository.create_user(username, "pw").await.unwrap(); 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); assert_eq!(fetched_user.username, username);
truncate_users(&user_repository.pool).await; truncate_users(&user_repository.pool).await;
@@ -643,7 +574,7 @@ mod tests {
.await .await
.unwrap(); .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.id, token);
assert_eq!(fetched_session.user_id, user.id); assert_eq!(fetched_session.user_id, user.id);
@@ -677,7 +608,7 @@ mod tests {
assert!(matches!( assert!(matches!(
user_repository.get_session(&token).await, user_repository.get_session(&token).await,
Err(LoftError::LoginFail) Ok(None)
)); ));
truncate_sessions(&user_repository.pool).await; truncate_sessions(&user_repository.pool).await;

View File

@@ -29,6 +29,7 @@ pub async fn mw_ctx_resolver(
Some(token) => user_repository Some(token) => user_repository
.get_session(&token) .get_session(&token)
.await .await
.and_then(|s| s.ok_or(LoftError::AuthFailSessionNotFound))
.map(|s| Ctx::new(s.user_id)), .map(|s| Ctx::new(s.user_id)),
None => Err(LoftError::AuthFailNoAuthTokenCookie), None => Err(LoftError::AuthFailNoAuthTokenCookie),
}; };

View File

@@ -7,11 +7,10 @@ use axum::{
routing::get, routing::get,
}; };
use sqlx::types::uuid; use sqlx::types::uuid;
use tracing::info;
use crate::{ use crate::{
ctx::Ctx, ctx::Ctx,
error::LoftError, error::{LoftError, Result},
model::{FileRecord, FileRepository}, model::{FileRecord, FileRepository},
}; };
@@ -29,31 +28,24 @@ async fn upload_file(
ctx: Ctx, ctx: Ctx,
mut multipart: Multipart, mut multipart: Multipart,
) -> Result<Json<FileRecord>, LoftError> { ) -> Result<Json<FileRecord>, LoftError> {
info!("handler: upload_file"); let mut uploaded: Option<(String, String, usize)> = None;
let mut file_name = None; while let Some(field) = multipart.next_field().await? {
let mut file_storage_key = None; if field.name() == Some("file") {
let mut file_size = None; let name = field.file_name().map(str::to_string).unwrap_or_default();
let key = format!("{name}-{}", uuid::Uuid::new_v4());
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());
let size = file_repository.upload_file(field, &key).await?; let size = file_repository.upload_file(field, &key).await?;
file_name = Some(name); uploaded = Some((name, key, size));
file_storage_key = Some(key);
file_size = Some(size);
} }
} }
if let (Some(name), Some(key), Some(size)) = (file_name, file_storage_key, file_size) { let (name, key, size) = uploaded.ok_or(LoftError::NoFileProvided)?;
let file_record = file_repository
.create_file_record(ctx.user_id(), &name, size, &key)
.await?;
return Ok(Json(file_record));
}
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( async fn get_file(
@@ -71,7 +63,7 @@ async fn download_file(
State(file_repository): State<FileRepository>, State(file_repository): State<FileRepository>,
ctx: Ctx, ctx: Ctx,
Path(file_id): Path<u64>, Path(file_id): Path<u64>,
) -> Result<impl IntoResponse, LoftError> { ) -> Result<impl IntoResponse> {
let stream = file_repository let stream = file_repository
.download_file(file_id as i64, ctx.user_id()) .download_file(file_id as i64, ctx.user_id())
.await?; .await?;
@@ -83,13 +75,12 @@ async fn stream_part(
ctx: Ctx, ctx: Ctx,
headers: HeaderMap, headers: HeaderMap,
Path(file_id): Path<u64>, Path(file_id): Path<u64>,
) -> Result<impl IntoResponse, LoftError> { ) -> Result<impl IntoResponse> {
info!("stream_part");
let file_record = file_repository let file_record = file_repository
.get_file(file_id as i64, ctx.user_id()) .get_file(file_id as i64, ctx.user_id())
.await?; .await?;
let file_size = file_record.size as u64; 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 let stream = file_repository
.stream_part(file_id as i64, ctx.user_id(), start, end + 1) .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") .header(header::ACCEPT_RANGES, "bytes")
.body(Body::from_stream(stream)) .body(Body::from_stream(stream))
.unwrap(); .expect("Failed to build response with valid headers");
Ok(respones) Ok(respones)
} }
fn parse_range(headers: &HeaderMap, file_size: u64) -> (u64, u64) { fn parse_range(headers: &HeaderMap, file_size: u64) -> Result<(u64, u64)> {
let range = headers.get(header::RANGE); let range = headers.get(header::RANGE).ok_or(LoftError::InvalidRange)?;
let str = range.unwrap().to_str().unwrap(); let str = range.to_str().map_err(|_| LoftError::InvalidRange)?;
let strip = str.strip_prefix("bytes=").unwrap(); let tuple = str
let tuple = strip.split_once("-").unwrap(); .strip_prefix("bytes=")
let start = tuple.0.parse::<u64>().unwrap(); .and_then(|s| s.split_once("-"))
.ok_or(LoftError::InvalidRange)?;
let start = tuple
.0
.parse::<u64>()
.map_err(|_| LoftError::InvalidRange)?;
let end = tuple.1.parse::<u64>().unwrap_or(file_size - 1); let end = tuple.1.parse::<u64>().unwrap_or(file_size - 1);
(start, end) Ok((start, end))
} }
async fn delete_file( async fn delete_file(
@@ -126,8 +122,6 @@ async fn delete_file(
ctx: Ctx, ctx: Ctx,
Path(file_id): Path<u64>, Path(file_id): Path<u64>,
) -> Result<Json<FileRecord>, LoftError> { ) -> Result<Json<FileRecord>, LoftError> {
info!("handler: delete_file");
let file = file_repository let file = file_repository
.delete_file(file_id as i64, ctx.user_id()) .delete_file(file_id as i64, ctx.user_id())
.await?; .await?;
@@ -138,8 +132,6 @@ async fn list_files(
State(file_repository): State<FileRepository>, State(file_repository): State<FileRepository>,
ctx: Ctx, ctx: Ctx,
) -> Result<Json<Vec<FileRecord>>, LoftError> { ) -> Result<Json<Vec<FileRecord>>, LoftError> {
info!("handler: list_files");
let files = file_repository.list_files(ctx.user_id()).await?; let files = file_repository.list_files(ctx.user_id()).await?;
Ok(Json(files)) Ok(Json(files))
} }
@@ -178,7 +170,7 @@ mod tests {
dotenvy::from_filename(".env.test").ok(); dotenvy::from_filename(".env.test").ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap(); let pool = PgPool::connect(&database_url).await.unwrap();
let user_repository = UserRepository::new(pool).unwrap(); let user_repository = UserRepository::new(pool);
user_repository user_repository
} }

View File

@@ -6,8 +6,13 @@ use axum::{Json, Router, extract::State, http::StatusCode, routing::post};
use rand::RngExt; use rand::RngExt;
use serde::Deserialize; use serde::Deserialize;
use tower_cookies::{Cookie, Cookies}; 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 { pub fn routes_auth(user_repository: UserRepository) -> Router {
Router::new() Router::new()
@@ -22,24 +27,26 @@ async fn login(
cookies: Cookies, cookies: Cookies,
Json(payload): Json<LoginPayload>, Json(payload): Json<LoginPayload>,
) -> Result<StatusCode, LoftError> { ) -> Result<StatusCode, LoftError> {
let user = user_repository.find_by_username(&payload.username).await?; let user = user_repository
// TODO: replace unwrap with ? .find_by_username(&payload.username)
let parsed_hash = PasswordHash::new(&user.password_hash).unwrap(); .await?
if Argon2::default() .ok_or(LoftError::LoginFail)?;
.verify_password(payload.password.as_bytes(), &parsed_hash) let parsed_hash = PasswordHash::new(&user.password_hash)?;
.is_ok() let password_verification =
{ Argon2::default().verify_password(payload.password.as_bytes(), &parsed_hash);
let expires_at = chrono::Utc::now() + chrono::Duration::days(1);
let cookie = create_cookie(); if password_verification.is_err() {
let auth_token = cookie.value();
cookies.add(cookie.clone());
user_repository
.create_session(user.id, auth_token, expires_at)
.await?;
} else {
return Err(LoftError::LoginFail); 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) Ok(StatusCode::OK)
} }
@@ -63,22 +70,24 @@ async fn register(
) -> Result<StatusCode, LoftError> { ) -> Result<StatusCode, LoftError> {
if user_repository if user_repository
.find_by_username(&payload.username) .find_by_username(&payload.username)
.await .await?
.is_ok() .is_some()
{ {
warn!(
"Register fail, username {} already exists",
&payload.username
); // also fix "Login fail" typo
return Err(LoftError::RegisterFail); return Err(LoftError::RegisterFail);
} }
let salt = SaltString::generate(&mut OsRng); let salt = SaltString::generate(&mut OsRng);
let argon2 = Argon2::default(); let argon2 = Argon2::default();
// TODO: replace unwrap let password_hash = &argon2
let password_hash = argon2 .hash_password(payload.password.as_bytes(), &salt)?
.hash_password(payload.password.as_bytes(), &salt)
.unwrap()
.to_string(); .to_string();
user_repository user_repository
.create_user(&payload.username, &password_hash) .create_user(&payload.username, password_hash)
.await?; .await?;
Ok(StatusCode::CREATED) Ok(StatusCode::CREATED)
@@ -135,7 +144,7 @@ mod tests {
dotenvy::from_filename(".env.test").ok(); dotenvy::from_filename(".env.test").ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap(); let pool = PgPool::connect(&database_url).await.unwrap();
let user_repository = UserRepository::new(pool).unwrap(); let user_repository = UserRepository::new(pool);
user_repository user_repository
} }