refactor(backend): replace unwraps with typed errors
This commit is contained in:
@@ -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(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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),
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -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
|
let file_record = file_repository
|
||||||
.create_file_record(ctx.user_id(), &name, size, &key)
|
.create_file_record(ctx.user_id(), &name, size, &key)
|
||||||
.await?;
|
.await?;
|
||||||
return Ok(Json(file_record));
|
|
||||||
}
|
|
||||||
|
|
||||||
Err(LoftError::UndefinedErrorType)
|
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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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,13 +27,18 @@ 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);
|
||||||
|
|
||||||
|
if password_verification.is_err() {
|
||||||
|
return Err(LoftError::LoginFail);
|
||||||
|
}
|
||||||
|
|
||||||
let expires_at = chrono::Utc::now() + chrono::Duration::days(1);
|
let expires_at = chrono::Utc::now() + chrono::Duration::days(1);
|
||||||
let cookie = create_cookie();
|
let cookie = create_cookie();
|
||||||
let auth_token = cookie.value();
|
let auth_token = cookie.value();
|
||||||
@@ -36,9 +46,6 @@ async fn login(
|
|||||||
user_repository
|
user_repository
|
||||||
.create_session(user.id, auth_token, expires_at)
|
.create_session(user.id, auth_token, expires_at)
|
||||||
.await?;
|
.await?;
|
||||||
} else {
|
|
||||||
return Err(LoftError::LoginFail);
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user