streamline normal/multipart uploads

Signed-off-by: Alexey Aristov <aav@acm.org>
This commit is contained in:
Alexey Aristov
2025-09-01 20:52:05 +02:00
parent d7bf59c34a
commit bd41fd1d04
3 changed files with 122 additions and 204 deletions
+82 -192
View File
@@ -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<Vec<u8>>,
}
@@ -31,213 +29,105 @@ fn random_key() -> String {
}
#[instrument(level = "debug", skip_all, fields(s3_bucket))]
pub async fn upload(
pub async fn upload<S, E>(
s3: &S3Client,
pool: &Pool,
request: &mut ServiceRequest,
payload: Payload,
) -> Result<Blob, ApiError> {
size: Option<Size>,
mut source: S,
) -> Result<Blob, ApiError>
where
S: Stream<Item = Result<Bytes, E>> + 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::<Header<ContentLength>>().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<Blob, ApiError> {
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<Blob, ApiError> {
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<CompletedPart> {
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)
}
+23 -5
View File
@@ -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<E: Display + std::error::Error + 'static, B: std::fmt::Debug> From<SdkError
}
}
trait ServiceRequestExt {
async fn content_length(&mut self) -> Option<Size>;
}
impl ServiceRequestExt for ServiceRequest {
async fn content_length(&mut self) -> Option<Size> {
self.extract::<Header<ContentLength>>()
.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<HttpRe
}
}
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 part_data = PartData {
workspace: path.workspace,
@@ -175,7 +191,9 @@ pub async fn patch(request: HttpRequest, payload: Payload) -> HandlerResult<Http
let pool = request.app_data::<Data<Pool>>().unwrap().to_owned();
let s3 = request.app_data::<Data<S3Client>>().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::<PartData>(&pool, path.workspace, &path.key).await?;
@@ -291,7 +309,7 @@ pub async fn get(request: HttpRequest) -> HandlerResult<HttpResponse> {
}
}
response.body(SizedStream::new(content_length, stream(parts, s3)))
response.body(SizedStream::new(content_length as u64, stream(parts, s3)))
} else {
HttpResponse::NotFound().finish()
};
+17 -7
View File
@@ -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<S, E>(
s3: &S3Client,
bucket: &str,
key: &str,
upload_id: &str,
mut source: impl Stream<Item = Result<Bytes, std::io::Error>> + Unpin,
) -> Result<(CompletedMultipartUpload, Upload)> {
mut source: S,
) -> Result<(CompletedMultipartUpload, Upload)>
where
S: Stream<Item = Result<Bytes, E>> + Unpin,
E: StdError + Send + Sync + 'static,
{
debug!("upload start");
let upload_part = async |number, buffer: Bytes| -> Result<CompletedPart> {
@@ -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<S>(
pub async fn multipart_upload<S, E>(
s3: &S3Client,
bucket: &str,
key: &str,
source: impl Stream<Item = Result<Bytes, std::io::Error>> + Unpin,
) -> Result<Upload> {
source: S,
) -> Result<Upload>
where
S: Stream<Item = Result<Bytes, E>> + Unpin,
E: StdError + Send + Sync + 'static,
{
let span = Span::current();
let create_multipart = s3