diff --git a/server/src/blob.rs b/server/src/blob.rs index 891e77d990..f1d9fc12d2 100644 --- a/server/src/blob.rs +++ b/server/src/blob.rs @@ -1,25 +1,23 @@ -use actix_web::dev::ServiceRequest; -use actix_web::http::header::ContentLength; -use actix_web::web::{Header, Payload}; +use std::error::Error as StdError; + use aws_sdk_s3::primitives::ByteStream; -use aws_sdk_s3::types::{CompletedMultipartUpload, CompletedPart}; use blake3::Hasher; -use bytes::BytesMut; -use futures_util::StreamExt; +use bytes::{Bytes, BytesMut}; +use futures::stream::StreamExt; +use futures_util::Stream; use size::Size; use tracing::*; +use crate::handlers::ApiError; use crate::s3::S3Client; use crate::{ config::CONFIG, postgres::{self, Pool}, }; -use crate::handlers::{ApiError, HandlerResult}; - pub struct Blob { pub s3_key: String, - pub length: u64, + pub length: usize, pub inline: Option>, } @@ -31,213 +29,105 @@ fn random_key() -> String { } #[instrument(level = "debug", skip_all, fields(s3_bucket))] -pub async fn upload( +pub async fn upload( s3: &S3Client, pool: &Pool, - request: &mut ServiceRequest, - payload: Payload, -) -> Result { + size: Option, + mut source: S, +) -> Result +where + S: Stream> + Unpin, + E: StdError + Send + Sync + 'static, +{ let span = Span::current(); let s3_bucket = &CONFIG.s3_bucket; - span.record("s3_bucket", &s3_bucket); - if let Ok(length) = request.extract::>().await - && length.0 < Size::from_megabytes(MULTIPART_THRESHOLD).bytes() as usize + let blob = if let Some(length) = size + && length < Size::from_megabytes(MULTIPART_THRESHOLD) { - upload_regular( - s3, - pool, - &s3_bucket, - length.0 < Size::from_kilobytes(INLINE_THRESHHOLD).bytes() as usize, - payload, - ) - .await - } else { - upload_multipart(s3, pool, &s3_bucket, payload).await - } -} + let mut hash = Hasher::new(); -#[instrument(level = "debug", skip_all, fields(s3_key))] -async fn upload_regular( - s3: &S3Client, - pool: &Pool, - s3_bucket: &str, - require_inline: bool, - payload: Payload, -) -> Result { - let span = Span::current(); + let mut buffer = BytesMut::new(); - let payload = payload - .to_bytes_limited(Size::from_megabytes(MULTIPART_THRESHOLD).bytes() as usize) - .await - .map_err(|_| actix_web::error::ErrorPayloadTooLarge("payload too large"))??; + while let Some(Ok(chunk)) = source.next().await { + buffer.extend_from_slice(&chunk); - let length = payload.len() as u64; - let inline = if require_inline { - Some(payload.to_vec()) - } else { - None - }; + if buffer.len() > length.bytes() as usize { + return Err(actix_web::error::ErrorPayloadTooLarge("payload too large").into()); + } + } - let mut hash = Hasher::new(); - hash.update(&payload); + if buffer.len() != length.bytes() as usize { + return Err(actix_web::error::ErrorBadRequest("payload size mismatch").into()); + } - let hash = hash.finalize().to_hex().to_string(); + let hash = hash.update(&buffer).finalize().to_hex(); + let length = buffer.len(); + let inline = if length < Size::from_kilobytes(INLINE_THRESHHOLD).bytes() as usize { + Some(buffer.to_vec()) + } else { + None + }; - let s3_key = if let Some(s3_key_found) = postgres::find_blob_by_hash(&pool, &hash).await? { - span.record("s3_key", &s3_key_found); - debug!(s3_key_found, "blob deduplicated"); - s3_key_found + let s3_key = if let Some(s3_key_found) = postgres::find_blob_by_hash(&pool, &hash).await? { + span.record("s3_key", &s3_key_found); + debug!(s3_key_found, "blob deduplicated"); + s3_key_found + } else { + let s3_key = random_key(); + span.record("s3_key", &s3_key); + + s3.put_object() + .bucket(s3_bucket) + .key(&s3_key) + .body(ByteStream::from(buffer.freeze())) + .send() + .await?; + + postgres::insert_blob(&pool, &s3_key, &hash).await?; + + debug!("blob created"); + s3_key + }; + + Blob { + s3_key, + length, + inline, + } } else { let s3_key = random_key(); span.record("s3_key", &s3_key); - s3.put_object() - .bucket(s3_bucket) - .key(&s3_key) - .body(ByteStream::from(payload)) - .send() - .await?; + let upload = crate::s3::multipart_upload(&s3, &s3_bucket, &s3_key, source).await?; - postgres::insert_blob(&pool, &s3_key, &hash).await?; + let hash = upload.hash.to_hex().to_string(); - debug!("blob created"); + let s3_key = if let Some(s3_key_found) = postgres::find_blob_by_hash(&pool, &hash).await? { + debug!(s3_key_found, "blob deduplicated"); - s3_key - }; + // delete uploaded + s3.delete_object() + .bucket(s3_bucket) + .key(s3_key) + .send() + .await?; - Ok(Blob { - s3_key, - length, - inline, - }) -} - -#[instrument(level = "debug", skip_all, fields(upload, s3_key))] -async fn upload_multipart( - s3: &S3Client, - pool: &Pool, - s3_bucket: &str, - mut payload: Payload, -) -> Result { - let span = Span::current(); - - let s3_key = random_key(); - - span.record("s3_key", &s3_key); - - let create_multipart = s3 - .create_multipart_upload() - .bucket(s3_bucket) - .key(&s3_key) - .send() - .await?; - - let upload_id = create_multipart.upload_id().unwrap(); - - span.record("upload", &upload_id[upload_id.len().saturating_sub(16)..]); - - debug!("upload start"); - - let upload_part = async |number, buffer: BytesMut| -> HandlerResult { - let upload = s3 - .upload_part() - .bucket(s3_bucket) - .key(&s3_key) - .upload_id(upload_id) - .body(buffer.freeze().into()) - .part_number(number) - .send() - .await?; - - let part = CompletedPart::builder() - .e_tag(upload.e_tag.unwrap()) - .part_number(number) - .build(); - - Ok(part) - }; - - let mut buffer = BytesMut::with_capacity(1024 * 1024 * 6); - let mut complete = CompletedMultipartUpload::builder(); - let mut part_number = 1; - let mut hash = Hasher::new(); - let mut total_in = 0; - let mut total_uploaded = 0; - - while let Some(part) = payload.next().await { - if let Ok(part) = part { - hash.update(&part); - - total_in += part.len(); - - buffer.extend_from_slice(&part); - - // each part must be at least 5MB - if buffer.len() > 1024 * 1024 * 5 { - trace!(length = buffer.len(), part_number, "upload part"); - - total_uploaded += buffer.len(); - - let uploaded = upload_part(part_number, buffer).await?; - - complete = complete.parts(uploaded); - - buffer = BytesMut::new(); - part_number += 1; - } + s3_key_found } else { - // TODO: cleanup incomplete upload - panic!("read error") + debug!("blob created"); + postgres::insert_blob(&pool, &s3_key, &hash).await?; + s3_key + }; + + Blob { + s3_key, + length: upload.length, + inline: None, } - } - - // the last part - if buffer.len() > 0 { - total_uploaded += buffer.len(); - - trace!(length = buffer.len(), part_number, "upload part"); - let uploaded = upload_part(part_number, buffer).await?; - complete = complete.parts(uploaded); - } - - assert_eq!(total_in, total_uploaded); - - let _ = s3 - .complete_multipart_upload() - .bucket(s3_bucket) - .key(&s3_key) - .multipart_upload(complete.build()) - .upload_id(upload_id) - .send() - .await?; - - let hash = hash.finalize().to_hex().to_string(); - - debug!(hash, "upload complete"); - - let s3_key = if let Some(s3_key_found) = postgres::find_blob_by_hash(&pool, &hash).await? { - debug!(s3_key_found, "blob deduplicated"); - - // delete uploaded - s3.delete_object() - .bucket(s3_bucket) - .key(s3_key) - .send() - .await?; - - s3_key_found - } else { - debug!("blob created"); - postgres::insert_blob(&pool, &s3_key, &hash).await?; - s3_key }; - Ok(Blob { - s3_key, - length: total_uploaded as u64, - inline: None, - }) + Ok(blob) } diff --git a/server/src/handlers.rs b/server/src/handlers.rs index 8a0b979929..5822820c63 100644 --- a/server/src/handlers.rs +++ b/server/src/handlers.rs @@ -6,7 +6,7 @@ use actix_web::{ dev::ServiceRequest, http::{ self, - header::{self, ContentType}, + header::{self, ContentLength, ContentType}, }, web::{Data, Header, Path, Payload}, }; @@ -14,6 +14,7 @@ use aws_sdk_s3::error::SdkError; use bytes::Bytes; use futures_util::Stream; use serde::{Deserialize, Serialize}; +use size::Size; use tracing::*; use uuid::Uuid; @@ -70,12 +71,25 @@ impl From Option; +} + +impl ServiceRequestExt for ServiceRequest { + async fn content_length(&mut self) -> Option { + self.extract::>() + .await + .map(|header| Size::from_bytes(*header.0)) + .ok() + } +} + #[derive(Serialize, serde::Deserialize, Debug)] struct PartData { workspace: Uuid, key: String, part: u32, - size: u64, + size: usize, blob: String, etag: String, @@ -129,7 +143,9 @@ pub async fn put(request: HttpRequest, payload: Payload) -> HandlerResult HandlerResult>().unwrap().to_owned(); let s3 = request.app_data::>().unwrap().to_owned(); - let uploaded = upload(&s3, &pool, &mut request, payload).await?; + let content_length = request.content_length().await; + + let uploaded = upload(&s3, &pool, content_length, payload).await?; let parts = postgres::find_parts::(&pool, path.workspace, &path.key).await?; @@ -291,7 +309,7 @@ pub async fn get(request: HttpRequest) -> HandlerResult { } } - response.body(SizedStream::new(content_length, stream(parts, s3))) + response.body(SizedStream::new(content_length as u64, stream(parts, s3))) } else { HttpResponse::NotFound().finish() }; diff --git a/server/src/s3.rs b/server/src/s3.rs index e9ce9c3ce7..0ca106f43c 100644 --- a/server/src/s3.rs +++ b/server/src/s3.rs @@ -1,3 +1,5 @@ +use std::error::Error as StdError; + use anyhow::Result; use aws_config::BehaviorVersion; use aws_sdk_s3::{ @@ -32,13 +34,17 @@ pub struct Upload { pub length: usize, } -async fn multipart_upload_stream( +async fn multipart_upload_stream( s3: &S3Client, bucket: &str, key: &str, upload_id: &str, - mut source: impl Stream> + Unpin, -) -> Result<(CompletedMultipartUpload, Upload)> { + mut source: S, +) -> Result<(CompletedMultipartUpload, Upload)> +where + S: Stream> + Unpin, + E: StdError + Send + Sync + 'static, +{ debug!("upload start"); let upload_part = async |number, buffer: Bytes| -> Result { @@ -68,7 +74,7 @@ async fn multipart_upload_stream( let mut length = 0; while let Some(part) = source.next().await { - let part = part?; + let part = part.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e))?; hash.update(&part); @@ -108,12 +114,16 @@ async fn multipart_upload_stream( } #[tracing::instrument(level = "debug", skip_all)] -pub async fn multipart_upload( +pub async fn multipart_upload( s3: &S3Client, bucket: &str, key: &str, - source: impl Stream> + Unpin, -) -> Result { + source: S, +) -> Result +where + S: Stream> + Unpin, + E: StdError + Send + Sync + 'static, +{ let span = Span::current(); let create_multipart = s3