mirror of
https://github.com/hcengineering/platform.git
synced 2026-09-08 10:47:42 +02:00
streamline normal/multipart uploads
Signed-off-by: Alexey Aristov <aav@acm.org>
This commit is contained in:
+82
-192
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user