Files
loft/backend/src/web/routes_file.rs

354 lines
11 KiB
Rust

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<FileRepository>,
ctx: Ctx,
mut multipart: Multipart,
) -> Result<Json<FileRecord>, 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<FileRepository>,
ctx: Ctx,
Path(file_id): Path<u64>,
) -> Result<Json<FileRecord>, LoftError> {
let record = file_repository
.get_file(file_id as i64, ctx.user_id())
.await?;
Ok(Json(record))
}
#[axum::debug_handler]
async fn download_file(
State(file_repository): State<FileRepository>,
ctx: Ctx,
Path(file_id): Path<u64>,
) -> Result<impl IntoResponse, LoftError> {
let stream = file_repository
.download_file(file_id as i64, ctx.user_id())
.await?;
Ok(Body::from_stream(stream))
}
async fn delete_file(
State(file_repository): State<FileRepository>,
ctx: Ctx,
Path(file_id): Path<u64>,
) -> Result<Json<FileRecord>, LoftError> {
info!("handler: delete_file");
let file = file_repository
.delete_file(file_id as i64, ctx.user_id())
.await?;
Ok(Json(file))
}
async fn list_files(
State(file_repository): State<FileRepository>,
ctx: Ctx,
) -> Result<Json<Vec<FileRecord>>, 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::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).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", "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_ok();
let file = res.json::<serde_json::Value>();
assert_eq!(file["name"], "a.jpg");
let list = server
.get("/api/files")
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await
.json::<serde_json::Value>();
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::<serde_json::Value>()["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::<serde_json::Value>()["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();
}
}