use axum::{ Json, Router, body::Body, extract::{Multipart, Path, State}, http::{HeaderMap, Response, StatusCode, header}, response::IntoResponse, routing::get, }; use sqlx::types::uuid; use tokio::sync::mpsc::Sender; use tracing::warn; use utoipa::OpenApi; use crate::{ ctx::Ctx, error::{LoftError, Result}, model::{FileRecord, FileRepository}, tasks::ThumbnailTask, }; #[derive(OpenApi)] #[openapi( paths( list_files, upload_file, get_file, delete_file, download_file, thumbnail, stream_part ), components(schemas(FileRecord)) )] pub struct FileApi; #[derive(Clone)] pub struct FileState { pub file_repository: FileRepository, pub tx: Sender, } impl FileState { pub fn new(file_repository: FileRepository, tx: Sender) -> Self { Self { file_repository, tx, } } } pub fn routes_file(file_service: FileState) -> 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)) .route("/files/{id}/thumbnail", get(thumbnail)) .route("/files/{id}/stream_part", get(stream_part)) .with_state(file_service) } /// Upload a file #[utoipa::path( post, path = "/api/files", tag = "files", security(("cookie_auth" = [])), request_body(content_type = "multipart/form-data", description = "Multipart form with a `file` field"), responses( (status = 201, description = "Uploaded file metadata", body = FileRecord), (status = 400, description = "No file provided or malformed upload"), (status = 401, description = "Unauthorized"), (status = 500, description = "Internal server error") ) )] async fn upload_file( State(service): State, ctx: Ctx, mut multipart: Multipart, ) -> Result<(StatusCode, Json)> { let mut uploaded: Option<(String, String, usize)> = None; 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 = service.file_repository.upload_file(field, &key).await?; uploaded = Some((name, key, size)); } } let (name, key, size) = uploaded.ok_or(LoftError::NoFileProvided)?; let record = service .file_repository .create_file_record(ctx.user_id(), &name, size, &key) .await?; if let Err(e) = service .tx .send(ThumbnailTask::new(record.id, ctx.user_id())) .await { warn!("failed to send thumbnail task for file {}: {e}", record.id) } Ok((StatusCode::CREATED, Json(record))) } /// Get a file's metadata #[utoipa::path( get, path = "/api/files/{id}", tag = "files", security(("cookie_auth" = [])), params(("id" = u64, Path, description = "File id")), responses( (status = 200, description = "File metadata", body = FileRecord), (status = 404, description = "File not found"), (status = 401, description = "Unauthorized"), (status = 500, description = "Internal server error") ) )] async fn get_file( State(service): State, ctx: Ctx, Path(file_id): Path, ) -> Result> { let record = service .file_repository .get_file(file_id as i64, ctx.user_id()) .await?; Ok(Json(record)) } /// Download a file #[utoipa::path( get, path = "/api/files/{id}/download", tag = "files", security(("cookie_auth" = [])), params(("id" = u64, Path, description = "File id")), responses( (status = 200, description = "File bytes", content_type = "application/octet-stream"), (status = 404, description = "File not found"), (status = 401, description = "Unauthorized"), (status = 500, description = "Internal server error") ) )] async fn download_file( State(service): State, ctx: Ctx, Path(file_id): Path, ) -> Result { let stream = service .file_repository .download_file(file_id as i64, ctx.user_id()) .await?; Ok(Body::from_stream(stream)) } /// Download a file thumbnail #[utoipa::path( get, path = "/api/files/{id}/thumbnail", tag = "files", security(("cookie_auth" = [])), params(("id" = u64, Path, description = "File id")), responses( (status = 200, description = "File bytes", content_type = "application/octet-stream"), (status = 404, description = "File not found"), (status = 401, description = "Unauthorized"), (status = 500, description = "Internal server error") ) )] async fn thumbnail( State(service): State, ctx: Ctx, Path(file_id): Path, ) -> Result { let stream = service .file_repository .thumbnail(file_id as i64, ctx.user_id()) .await?; Ok(Body::from_stream(stream)) } /// Stream a byte range of a file (HTTP 206) #[utoipa::path( get, path = "/api/files/{id}/stream_part", tag = "files", security(("cookie_auth" = [])), params( ("id" = u64, Path, description = "File id"), ("Range" = Option, Header, description = "Byte range, e.g. bytes=0-1023") ), responses( (status = 206, description = "Partial content"), (status = 400, description = "Invalid range"), (status = 404, description = "File not found"), (status = 401, description = "Unauthorized"), (status = 500, description = "Internal server error") ) )] async fn stream_part( State(service): State, ctx: Ctx, headers: HeaderMap, Path(file_id): Path, ) -> Result { let record = service .file_repository .get_file(file_id as i64, ctx.user_id()) .await?; let file_size = record.size as u64; let (start, end): (u64, u64) = parse_range(&headers, file_size)?; let stream = service .file_repository .stream_part(file_id as i64, ctx.user_id(), start, end + 1) .await?; let respones = Response::builder() .status(StatusCode::PARTIAL_CONTENT) .header(header::CONTENT_TYPE, record.file_type) .header(header::CONTENT_LENGTH, end - start + 1) .header( header::CONTENT_RANGE, format!("bytes {start}-{end}/{file_size}"), ) .header(header::ACCEPT_RANGES, "bytes") .body(Body::from_stream(stream)) .expect("Failed to build response with valid headers"); Ok(respones) } 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); Ok((start, end)) } /// Delete a file #[utoipa::path( delete, path = "/api/files/{id}", tag = "files", security(("cookie_auth" = [])), params(("id" = u64, Path, description = "File id")), responses( (status = 200, description = "Deleted file metadata", body = FileRecord), (status = 404, description = "File not found"), (status = 401, description = "Unauthorized"), (status = 500, description = "Internal server error") ) )] async fn delete_file( State(service): State, ctx: Ctx, Path(file_id): Path, ) -> Result> { let file = service .file_repository .delete_file(file_id as i64, ctx.user_id()) .await?; Ok(Json(file)) } /// List the current user's files #[utoipa::path( get, path = "/api/files", tag = "files", security(("cookie_auth" = [])), responses( (status = 200, description = "List user's files", body = Vec), (status = 401, description = "Unauthorized"), (status = 500, description = "Internal server error") ) )] async fn list_files(State(service): State, ctx: Ctx) -> Result>> { let files = service.file_repository.list_files(ctx.user_id()).await?; Ok(Json(files)) } #[cfg(test)] mod tests { use axum::{ Router, http::{StatusCode, header}, middleware, }; use axum_test::{ TestServer, multipart::{MultipartForm, Part}, }; use rand::RngExt; use serde_json::json; use sqlx::PgPool; use tokio::sync::mpsc; use tower_cookies::CookieManagerLayer; use crate::{ model::{FileRepository, UserRepository}, web::{ mw_auth::{mw_ctx_resolver, mw_require_auth}, routes_file::{FileState, routes_file}, }, }; const BAD_AUTH_COOKIE: &str = "auth-token=user-0123456789"; async fn file_repository() -> FileRepository { 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 file_repository = FileRepository::new(pool).unwrap(); file_repository } async fn user_repository() -> UserRepository { 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); user_repository } async fn test_server() -> TestServer { let user_repository = user_repository().await; let file_repository = file_repository().await; let (tx, _) = mpsc::channel(100); let file_service = FileState::new(file_repository.clone(), tx); let routes_file = routes_file(file_service).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", "hash") .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, 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(StatusCode::CREATED); 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_stream_part_range() { 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"0123456789".to_vec()).file_name("a.txt"), )) .await; let id = post_res.json::()["id"].as_i64().unwrap(); let res = server .get(&format!("/api/files/{id}/stream_part")) .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .add_header(axum::http::header::RANGE, "bytes=2-5") .await; res.assert_status(axum::http::StatusCode::PARTIAL_CONTENT); assert_eq!(res.as_bytes(), b"2345".as_ref()); assert_eq!(res.header(header::CONTENT_RANGE), "bytes 2-5/10"); assert_eq!(res.header(header::CONTENT_LENGTH), "4"); assert_eq!(res.header(header::ACCEPT_RANGES), "bytes"); truncate(&file_repository.pool).await; } #[tokio::test] #[serial_test::serial] async fn test_stream_part_open_ended() { 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"0123456789".to_vec()).file_name("a.txt"), )) .await; let id = post_res.json::()["id"].as_i64().unwrap(); let res = server .get(&format!("/api/files/{id}/stream_part")) .add_header(axum::http::header::COOKIE, format!("auth-token={token}")) .add_header(axum::http::header::RANGE, "bytes=0-") .await; res.assert_status(axum::http::StatusCode::PARTIAL_CONTENT); assert_eq!(res.as_bytes(), b"0123456789".as_ref()); assert_eq!(res.header(header::CONTENT_RANGE), "bytes 0-9/10"); assert_eq!(res.header(header::CONTENT_LENGTH), "10"); 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(); } }