use axum::{ Json, Router, body::Body, extract::{Multipart, Path, State}, response::IntoResponse, routing::get, }; use sqlx::types::uuid; use tracing::info; use crate::{ ctx::Ctx, error::LoftError, model::{FileRecord, FileRepository} }; pub fn routes_file(file_repository: FileRepository) -> Router { Router::new() .route("/files", get(list_files).post(upload_file)) .route("/files/{id}", get(get_file).delete(delete_file)) .route("/files/{id}/download", get(download_file)) .with_state(file_repository) } async fn upload_file( State(file_repository): State, ctx: Ctx, mut multipart: Multipart, ) -> Result, LoftError> { info!("handler: upload_file"); 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()); let size = file_repository.upload_file(field, &key).await?; file_name = Some(name); 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 file_record = file_repository .create_file_record(ctx.user_id(), &name, size, &key) .await?; return Ok(Json(file_record)); } Err(LoftError::UndefinedErrorType) } #[axum::debug_handler] async fn get_file( State(file_repository): State, Path(file_id): Path, ) -> Result, LoftError> { let record = file_repository.get_file(file_id as i64).await?; Ok(Json(record)) } #[axum::debug_handler] async fn download_file( State(file_repository): State, Path(file_id): Path, ) -> Result { let stream = file_repository.download_file(file_id as i64).await?; Ok(Body::from_stream(stream)) } async fn delete_file( State(file_repository): State, Path(file_id): Path, ) -> Result, LoftError> { info!("handler: delete_file"); let file = file_repository.delete_file(file_id as i64).await?; Ok(Json(file)) } 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)) } #[cfg(test)] mod tests { use axum::{Router, middleware}; use axum_test::{ TestServer, multipart::{MultipartForm, Part}, }; use rand::RngExt; use serde_json::json; use sqlx::PgPool; use tower_cookies::CookieManagerLayer; use crate::{ model::{FileRepository, UserRepository}, web::{ mw_auth::{mw_ctx_resolver, mw_require_auth}, routes_file::routes_file, }, }; const BAD_AUTH_COOKIE: &str = "auth-token=user-0123456789"; async fn file_repository() -> FileRepository { dotenvy::dotenv().ok(); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let pool = PgPool::connect(&database_url).await.unwrap(); let file_repository = FileRepository::new(pool).unwrap(); file_repository } async fn user_repository() -> UserRepository { dotenvy::dotenv().ok(); let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); let pool = PgPool::connect(&database_url).await.unwrap(); let user_repository = UserRepository::new(pool).unwrap(); user_repository } async fn test_server() -> TestServer { let user_repository = user_repository().await; let file_repository = file_repository().await; let routes_file = routes_file(file_repository.clone()).route_layer(middleware::from_fn(mw_require_auth)); let app = Router::new() .nest("/api", routes_file) .layer(middleware::from_fn_with_state( user_repository, mw_ctx_resolver, )) .layer(CookieManagerLayer::new()); TestServer::new(app) } async fn create_test_session(user_repository: &UserRepository) -> String { let user = user_repository .create_user("testuser".to_string(), "hash".to_string()) .await .unwrap(); let token: String = rand::rng() .sample_iter(&rand::distr::Alphanumeric) .take(64) .map(char::from) .collect(); let expires_at = chrono::Utc::now() + chrono::Duration::days(1); user_repository .create_session(user.id, token.clone(), expires_at) .await .unwrap(); token } async fn truncate(pool: &PgPool) { sqlx::query!("TRUNCATE TABLE file_records") .execute(pool) .await .unwrap(); } #[tokio::test] async fn test_requires_auth() { let server = test_server().await; server.get("/api/files").await.assert_status_unauthorized(); } #[tokio::test] async fn test_requires_auth_invalid_cookie() { let server = test_server().await; server .get("/api/files") .add_header(axum::http::header::COOKIE, BAD_AUTH_COOKIE) .await .assert_status_unauthorized(); } #[tokio::test] #[serial_test::serial] async fn test_requires_auth_post() { let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; server .post("/api/files") .multipart(MultipartForm::new().add_part( "file", Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"), )) .await .assert_status_unauthorized(); truncate(&file_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_list_files_empty() { let user_repository = user_repository().await; let token = create_test_session(&user_repository).await; let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; server .get("/api/files") .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_ok() .assert_json(&json!([])); } #[tokio::test] #[serial_test::serial] async fn test_upload_and_list_files() { let user_repository = user_repository().await; let token = create_test_session(&user_repository).await; let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; let res = server .post("/api/files") .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .multipart(MultipartForm::new().add_part( "file", Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"), )) .await; res.assert_status_ok(); let file = res.json::(); assert_eq!(file["name"], "a.jpg"); let list = server .get("/api/files") .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .json::(); assert_eq!(list.as_array().unwrap().len(), 1); truncate(&file_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_download_file() { let user_repository = user_repository().await; let token = create_test_session(&user_repository).await; let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; let post_res = server .post("/api/files") .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .multipart(MultipartForm::new().add_part( "file", Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"), )) .await; let id = post_res.json::()["id"].as_i64().unwrap(); let res = server .get(&format!("/api/files/{id}/download")) .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await; res.assert_status_ok(); assert_eq!(res.as_bytes(), b"fake_bytes".as_ref()); truncate(&file_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_download_file_not_found() { let user_repository = user_repository().await; let token = create_test_session(&user_repository).await; let server = test_server().await; server .get("/api/files/99/download") .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_not_found(); } #[tokio::test] #[serial_test::serial] async fn test_delete_file() { let user_repository = user_repository().await; let token = create_test_session(&user_repository).await; let file_repository = file_repository().await; truncate(&file_repository.pool).await; let server = test_server().await; let post_res = server .post("/api/files") .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .multipart(MultipartForm::new().add_part( "file", Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"), )) .await; let id = post_res.json::()["id"].as_i64().unwrap(); server .delete(&format!("/api/files/{id}")) .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_ok(); server .get(&format!("/api/files/{id}")) .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_not_found(); truncate(&file_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_delete_file_not_found() { let user_repository = user_repository().await; let token = create_test_session(&user_repository).await; let server = test_server().await; server .delete("/api/files/99") .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .await .assert_status_not_found(); } }