Compare commits

...

2 Commits

4 changed files with 61 additions and 65 deletions

View File

@@ -61,7 +61,7 @@ impl FileRepository {
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, LoftError> {
let mut writer = self.op.writer(&file_storage_key).await.unwrap(); let mut writer = self.op.writer(file_storage_key).await.unwrap();
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.unwrap();
@@ -90,7 +90,9 @@ impl FileRepository {
"#, "#,
user_id, user_id,
file_name, file_name,
mime_guess::from_path(file_name).first_or_octet_stream().to_string(), mime_guess::from_path(file_name)
.first_or_octet_stream()
.to_string(),
file_size as i64, file_size as i64,
file_storage_key file_storage_key
) )
@@ -210,8 +212,8 @@ impl UserRepository {
pub async fn create_user( pub async fn create_user(
&self, &self,
username: String, username: &str,
password_hash: String, password_hash: &str,
) -> Result<User, LoftError> { ) -> Result<User, LoftError> {
let user = sqlx::query_as!( let user = sqlx::query_as!(
User, User,
@@ -232,7 +234,7 @@ impl UserRepository {
Ok(user) Ok(user)
} }
pub async fn find_by_username(&self, username: String) -> Result<User, LoftError> { pub async fn find_by_username(&self, username: &str) -> Result<User, LoftError> {
info!("Fetching username \"{}\" from users", username); info!("Fetching username \"{}\" from users", username);
let user = sqlx::query_as!( let user = sqlx::query_as!(
User, User,
@@ -254,7 +256,7 @@ impl UserRepository {
pub async fn create_session( pub async fn create_session(
&self, &self,
user_id: i64, user_id: i64,
token: String, token: &str,
expires_at: chrono::DateTime<chrono::Utc>, expires_at: chrono::DateTime<chrono::Utc>,
) -> Result<Session, LoftError> { ) -> Result<Session, LoftError> {
let session = sqlx::query_as!( let session = sqlx::query_as!(
@@ -280,7 +282,7 @@ impl UserRepository {
Ok(session) Ok(session)
} }
pub async fn get_session(&self, token: String) -> Result<Session, LoftError> { pub async fn get_session(&self, token: &str) -> Result<Session, LoftError> {
info!("Fetching session \"{}\" from sessions", token); info!("Fetching session \"{}\" from sessions", token);
let session = sqlx::query_as!( let session = sqlx::query_as!(
Session, Session,
@@ -300,7 +302,7 @@ impl UserRepository {
Ok(session) Ok(session)
} }
pub async fn delete_session(&self, token: String) -> Result<(), LoftError> { pub async fn delete_session(&self, token: &str) -> Result<(), LoftError> {
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
@@ -338,8 +340,14 @@ mod tests {
#[serial_test::serial] #[serial_test::serial]
async fn test_upload_and_list() { async fn test_upload_and_list() {
let user_repository = user_repository().await.unwrap(); let user_repository = user_repository().await.unwrap();
let user1 = user_repository.create_user("username1".to_string(), "password_hash".to_string()).await.unwrap(); let user1 = user_repository
let user2 = user_repository.create_user("username2".to_string(), "password_hash".to_string()).await.unwrap(); .create_user("username1", "password_hash")
.await
.unwrap();
let user2 = user_repository
.create_user("username2", "password_hash")
.await
.unwrap();
let file_repository = file_repository().await.unwrap(); let file_repository = file_repository().await.unwrap();
truncate_file_records(&file_repository.pool).await; truncate_file_records(&file_repository.pool).await;
@@ -377,7 +385,10 @@ mod tests {
#[serial_test::serial] #[serial_test::serial]
async fn test_download() { async fn test_download() {
let user_repository = user_repository().await.unwrap(); let user_repository = user_repository().await.unwrap();
let user = user_repository.create_user("username".to_string(), "password_hash".to_string()).await.unwrap(); let user = user_repository
.create_user("username", "password_hash")
.await
.unwrap();
let file_repository = file_repository().await.unwrap(); let file_repository = file_repository().await.unwrap();
truncate_file_records(&file_repository.pool).await; truncate_file_records(&file_repository.pool).await;
@@ -420,7 +431,10 @@ mod tests {
#[serial_test::serial] #[serial_test::serial]
async fn test_delete() { async fn test_delete() {
let user_repository = user_repository().await.unwrap(); let user_repository = user_repository().await.unwrap();
let user = user_repository.create_user("username".to_string(), "password_hash".to_string()).await.unwrap(); let user = user_repository
.create_user("username", "password_hash")
.await
.unwrap();
let file_repository = file_repository().await.unwrap(); let file_repository = file_repository().await.unwrap();
truncate_file_records(&file_repository.pool).await; truncate_file_records(&file_repository.pool).await;
@@ -487,11 +501,8 @@ mod tests {
let user_repository = user_repository().await.unwrap(); let user_repository = user_repository().await.unwrap();
truncate_users(&user_repository.pool).await; truncate_users(&user_repository.pool).await;
let username = "picolo".to_string(); let username = "picolo";
let user = user_repository let user = user_repository.create_user(&username, "pw").await.unwrap();
.create_user(username.clone(), "pw".to_string())
.await
.unwrap();
assert_eq!(user.username, username); assert_eq!(user.username, username);
truncate_users(&user_repository.pool).await; truncate_users(&user_repository.pool).await;
@@ -503,16 +514,10 @@ mod tests {
let user_repository = user_repository().await.unwrap(); let user_repository = user_repository().await.unwrap();
truncate_users(&user_repository.pool).await; truncate_users(&user_repository.pool).await;
let username = "picolo".to_string(); let username = "picolo";
user_repository user_repository.create_user(username, "pw").await.unwrap();
.create_user(username.clone(), "pw".to_string())
.await
.unwrap();
let fetched_user = user_repository let fetched_user = user_repository.find_by_username(username).await.unwrap();
.find_by_username(username.clone())
.await
.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;
@@ -533,16 +538,13 @@ mod tests {
truncate_sessions(&user_repository.pool).await; truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await; truncate_users(&user_repository.pool).await;
let user = user_repository let user = user_repository.create_user("picolo", "pw").await.unwrap();
.create_user("picolo".to_string(), "pw".to_string())
.await
.unwrap();
let token = token(); let token = token();
let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2); let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2);
let session = user_repository let session = user_repository
.create_session(user.id, token.clone(), expires_at) .create_session(user.id, &token, expires_at)
.await .await
.unwrap(); .unwrap();
@@ -560,20 +562,17 @@ mod tests {
truncate_sessions(&user_repository.pool).await; truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await; truncate_users(&user_repository.pool).await;
let user = user_repository let user = user_repository.create_user("picolo", "pw").await.unwrap();
.create_user("picolo".to_string(), "pw".to_string())
.await
.unwrap();
let token = token(); let token = token();
let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2); let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2);
user_repository user_repository
.create_session(user.id, token.clone(), expires_at) .create_session(user.id, &token, expires_at)
.await .await
.unwrap(); .unwrap();
let fetched_session = user_repository.get_session(token.clone()).await.unwrap(); let fetched_session = user_repository.get_session(&token).await.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);
@@ -589,16 +588,13 @@ mod tests {
truncate_sessions(&user_repository.pool).await; truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await; truncate_users(&user_repository.pool).await;
let user = user_repository let user = user_repository.create_user("picolo", "pw").await.unwrap();
.create_user("picolo".to_string(), "pw".to_string())
.await
.unwrap();
let token = token(); let token = token();
let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2); let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2);
let session = user_repository let session = user_repository
.create_session(user.id, token.clone(), expires_at) .create_session(user.id, &token, expires_at)
.await .await
.unwrap(); .unwrap();
@@ -606,10 +602,10 @@ mod tests {
assert_eq!(session.user_id, user.id); assert_eq!(session.user_id, user.id);
assert!(expires_at > chrono::Utc::now()); assert!(expires_at > chrono::Utc::now());
user_repository.delete_session(token.clone()).await.unwrap(); user_repository.delete_session(&token).await.unwrap();
assert!(matches!( assert!(matches!(
user_repository.get_session(token.clone()).await, user_repository.get_session(&token).await,
Err(LoftError::LoginFail) Err(LoftError::LoginFail)
)); ));

View File

@@ -27,7 +27,7 @@ pub async fn mw_ctx_resolver(
let result_ctx = match auth_token { let result_ctx = match auth_token {
Some(token) => user_repository Some(token) => user_repository
.get_session(token) .get_session(&token)
.await .await
.map(|s| Ctx::new(s.user_id)), .map(|s| Ctx::new(s.user_id)),
None => Err(LoftError::AuthFailNoAuthTokenCookie), None => Err(LoftError::AuthFailNoAuthTokenCookie),
@@ -39,11 +39,10 @@ pub async fn mw_ctx_resolver(
impl<S: Send + Sync> FromRequestParts<S> for Ctx { impl<S: Send + Sync> FromRequestParts<S> for Ctx {
type Rejection = LoftError; type Rejection = LoftError;
fn from_request_parts( async fn from_request_parts(
parts: &mut axum::http::request::Parts, parts: &mut axum::http::request::Parts,
_: &S, _: &S,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send { ) -> Result<Self, Self::Rejection> {
async move {
parts parts
.extensions .extensions
.get::<Result<Ctx, LoftError>>() .get::<Result<Ctx, LoftError>>()
@@ -51,4 +50,3 @@ impl<S: Send + Sync> FromRequestParts<S> for Ctx {
.clone() .clone()
} }
} }
}

View File

@@ -9,7 +9,9 @@ use sqlx::types::uuid;
use tracing::info; use tracing::info;
use crate::{ use crate::{
ctx::Ctx, error::LoftError, model::{FileRecord, FileRepository} ctx::Ctx,
error::LoftError,
model::{FileRecord, FileRepository},
}; };
pub fn routes_file(file_repository: FileRepository) -> Router { pub fn routes_file(file_repository: FileRepository) -> Router {
@@ -82,7 +84,7 @@ async fn delete_file(
async fn list_files( 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"); info!("handler: list_files");
@@ -146,7 +148,7 @@ mod tests {
async fn create_test_session(user_repository: &UserRepository) -> String { async fn create_test_session(user_repository: &UserRepository) -> String {
let user = user_repository let user = user_repository
.create_user("testuser".to_string(), "hash".to_string()) .create_user("testuser", "hash")
.await .await
.unwrap(); .unwrap();
let token: String = rand::rng() let token: String = rand::rng()
@@ -156,7 +158,7 @@ mod tests {
.collect(); .collect();
let expires_at = chrono::Utc::now() + chrono::Duration::days(1); let expires_at = chrono::Utc::now() + chrono::Duration::days(1);
user_repository user_repository
.create_session(user.id, token.clone(), expires_at) .create_session(user.id, &token, expires_at)
.await .await
.unwrap(); .unwrap();
token token

View File

@@ -22,7 +22,7 @@ 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.find_by_username(&payload.username).await?;
// TODO: replace unwrap with ? // TODO: replace unwrap with ?
let parsed_hash = PasswordHash::new(&user.password_hash).unwrap(); let parsed_hash = PasswordHash::new(&user.password_hash).unwrap();
if Argon2::default() if Argon2::default()
@@ -31,8 +31,8 @@ async fn login(
{ {
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().to_string(); let auth_token = cookie.value();
cookies.add(cookie); cookies.add(cookie.clone());
user_repository user_repository
.create_session(user.id, auth_token, expires_at) .create_session(user.id, auth_token, expires_at)
.await?; .await?;
@@ -50,7 +50,7 @@ async fn logout(
let auth_token = cookies.get(AUTH_TOKEN).map(|c| c.value().to_string()); let auth_token = cookies.get(AUTH_TOKEN).map(|c| c.value().to_string());
if let Some(auth_token) = auth_token { if let Some(auth_token) = auth_token {
user_repository.delete_session(auth_token.clone()).await?; user_repository.delete_session(&auth_token).await?;
cookies.remove(Cookie::build(AUTH_TOKEN).path("/").build()); cookies.remove(Cookie::build(AUTH_TOKEN).path("/").build());
} }
@@ -62,7 +62,7 @@ async fn register(
Json(payload): Json<LoginPayload>, Json(payload): Json<LoginPayload>,
) -> Result<StatusCode, LoftError> { ) -> Result<StatusCode, LoftError> {
if user_repository if user_repository
.find_by_username(payload.username.clone()) .find_by_username(&payload.username)
.await .await
.is_ok() .is_ok()
{ {
@@ -78,7 +78,7 @@ async fn register(
.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)