559 lines
18 KiB
Rust
559 lines
18 KiB
Rust
use axum::{
|
|
Json, Router,
|
|
body::Body,
|
|
extract::{Multipart, Path, State},
|
|
http::{HeaderMap, Response, StatusCode, header},
|
|
response::IntoResponse,
|
|
routing::get,
|
|
};
|
|
use sqlx::types::uuid;
|
|
use utoipa::OpenApi;
|
|
|
|
use crate::{
|
|
ctx::Ctx,
|
|
error::{LoftError, Result},
|
|
model::{FileRecord, FileRepository},
|
|
};
|
|
|
|
#[derive(OpenApi)]
|
|
#[openapi(
|
|
paths(
|
|
list_files,
|
|
upload_file,
|
|
get_file,
|
|
delete_file,
|
|
download_file,
|
|
stream_part
|
|
),
|
|
components(schemas(FileRecord))
|
|
)]
|
|
pub struct FileApi;
|
|
|
|
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))
|
|
.route("/files/{id}/stream_part", get(stream_part))
|
|
.with_state(file_repository)
|
|
}
|
|
|
|
/// 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(file_repository): State<FileRepository>,
|
|
ctx: Ctx,
|
|
mut multipart: Multipart,
|
|
) -> Result<(StatusCode, Json<FileRecord>)> {
|
|
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 = file_repository.upload_file(field, &key).await?;
|
|
uploaded = Some((name, key, size));
|
|
}
|
|
}
|
|
|
|
let (name, key, size) = uploaded.ok_or(LoftError::NoFileProvided)?;
|
|
|
|
let file_record = file_repository
|
|
.create_file_record(ctx.user_id(), &name, size, &key)
|
|
.await?;
|
|
|
|
Ok((StatusCode::CREATED, Json(file_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(file_repository): State<FileRepository>,
|
|
ctx: Ctx,
|
|
Path(file_id): Path<u64>,
|
|
) -> Result<Json<FileRecord>> {
|
|
let record = 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(file_repository): State<FileRepository>,
|
|
ctx: Ctx,
|
|
Path(file_id): Path<u64>,
|
|
) -> Result<impl IntoResponse> {
|
|
let stream = file_repository
|
|
.download_file(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<String>, 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(file_repository): State<FileRepository>,
|
|
ctx: Ctx,
|
|
headers: HeaderMap,
|
|
Path(file_id): Path<u64>,
|
|
) -> Result<impl IntoResponse> {
|
|
let file_record = file_repository
|
|
.get_file(file_id as i64, ctx.user_id())
|
|
.await?;
|
|
let file_size = file_record.size as u64;
|
|
let (start, end): (u64, u64) = parse_range(&headers, file_size)?;
|
|
|
|
let stream = 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, file_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::<u64>()
|
|
.map_err(|_| LoftError::InvalidRange)?;
|
|
let end = tuple.1.parse::<u64>().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(file_repository): State<FileRepository>,
|
|
ctx: Ctx,
|
|
Path(file_id): Path<u64>,
|
|
) -> Result<Json<FileRecord>> {
|
|
let file = 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<FileRecord>),
|
|
(status = 401, description = "Unauthorized"),
|
|
(status = 500, description = "Internal server error")
|
|
)
|
|
)]
|
|
async fn list_files(
|
|
State(file_repository): State<FileRepository>,
|
|
ctx: Ctx,
|
|
) -> Result<Json<Vec<FileRecord>>> {
|
|
let files = 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 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);
|
|
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(StatusCode::CREATED);
|
|
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_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::<serde_json::Value>()["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::<serde_json::Value>()["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::<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();
|
|
}
|
|
}
|