use crate::error::{Result, SZipError};
use aws_sdk_s3::primitives::ByteStream;
use aws_sdk_s3::types::{CompletedMultipartUpload, CompletedPart};
use aws_sdk_s3::Client;
use std::future::Future;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::io::{AsyncSeek, AsyncWrite};
use tokio::sync::mpsc;
pub const DEFAULT_PART_SIZE: usize = 5 * 1024 * 1024;
pub const MAX_PART_SIZE: usize = 5 * 1024 * 1024 * 1024;
pub const MAX_PARTS: usize = 10_000;
pub struct S3ZipWriter {
upload_tx: mpsc::UnboundedSender<UploadCommand>,
upload_task: Option<tokio::task::JoinHandle<Result<()>>>,
buffer: Vec<u8>,
part_size: usize,
position: u64,
current_part_number: usize,
shutdown_initiated: bool,
#[allow(dead_code)]
max_concurrent_uploads: usize,
}
enum UploadCommand {
UploadPart { part_number: usize, data: Vec<u8> },
Complete { final_data: Option<Vec<u8>> },
}
pub struct S3ZipWriterBuilder {
client: Option<Client>,
bucket: String,
key: String,
part_size: usize,
endpoint_url: Option<String>,
region: Option<String>,
force_path_style: bool,
max_concurrent_uploads: usize,
}
impl S3ZipWriter {
pub async fn new(
client: Client,
bucket: impl Into<String>,
key: impl Into<String>,
) -> Result<Self> {
Self::builder()
.client(client)
.bucket(bucket)
.key(key)
.build()
.await
}
pub fn builder() -> S3ZipWriterBuilder {
S3ZipWriterBuilder {
client: None,
bucket: String::new(),
key: String::new(),
part_size: DEFAULT_PART_SIZE,
endpoint_url: None,
region: None,
force_path_style: false,
max_concurrent_uploads: 4, }
}
}
impl S3ZipWriterBuilder {
pub fn client(mut self, client: Client) -> Self {
self.client = Some(client);
self
}
pub fn bucket(mut self, bucket: impl Into<String>) -> Self {
self.bucket = bucket.into();
self
}
pub fn key(mut self, key: impl Into<String>) -> Self {
self.key = key.into();
self
}
pub fn endpoint_url(mut self, url: impl Into<String>) -> Self {
self.endpoint_url = Some(url.into());
self
}
pub fn region(mut self, region: impl Into<String>) -> Self {
self.region = Some(region.into());
self
}
pub fn force_path_style(mut self, force: bool) -> Self {
self.force_path_style = force;
self
}
pub fn part_size(mut self, part_size: usize) -> Self {
assert!(
part_size >= DEFAULT_PART_SIZE,
"Part size must be at least 5MB"
);
assert!(part_size <= MAX_PART_SIZE, "Part size must not exceed 5GB");
self.part_size = part_size;
self
}
pub fn max_concurrent_uploads(mut self, max: usize) -> Self {
assert!(max > 0, "max_concurrent_uploads must be at least 1");
assert!(max <= 20, "max_concurrent_uploads should not exceed 20");
self.max_concurrent_uploads = max;
self
}
pub async fn build(self) -> Result<S3ZipWriter> {
let client = match self.client {
Some(c) => c,
None => {
let mut config_loader = aws_config::from_env();
if let Some(ref endpoint) = self.endpoint_url {
config_loader = config_loader.endpoint_url(endpoint);
}
if let Some(ref region) = self.region {
config_loader = config_loader.region(aws_config::Region::new(region.clone()));
}
let sdk_config = config_loader.load().await;
let mut s3_config = aws_sdk_s3::config::Builder::from(&sdk_config);
if self.force_path_style {
s3_config = s3_config.force_path_style(true);
}
Client::from_conf(s3_config.build())
}
};
let (tx, rx) = mpsc::unbounded_channel();
let max_concurrent = self.max_concurrent_uploads;
let upload_task = tokio::spawn(upload_worker_concurrent(
client,
self.bucket,
self.key,
rx,
max_concurrent,
));
Ok(S3ZipWriter {
upload_tx: tx,
upload_task: Some(upload_task),
buffer: Vec::with_capacity(self.part_size),
part_size: self.part_size,
position: 0,
current_part_number: 0,
shutdown_initiated: false,
max_concurrent_uploads: self.max_concurrent_uploads,
})
}
}
impl AsyncWrite for S3ZipWriter {
fn poll_write(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
self.buffer.extend_from_slice(buf);
self.position += buf.len() as u64;
if self.buffer.len() >= self.part_size {
let part_size = self.part_size;
let data = std::mem::replace(&mut self.buffer, Vec::with_capacity(part_size));
self.current_part_number += 1;
if self
.upload_tx
.send(UploadCommand::UploadPart {
part_number: self.current_part_number,
data,
})
.is_err()
{
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"Upload task terminated unexpectedly",
)));
}
}
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
if !self.shutdown_initiated {
self.shutdown_initiated = true;
let final_data = if !self.buffer.is_empty() {
Some(std::mem::take(&mut self.buffer))
} else {
None
};
if self
.upload_tx
.send(UploadCommand::Complete { final_data })
.is_err()
{
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"Upload task terminated unexpectedly",
)));
}
}
if let Some(task) = self.upload_task.as_mut() {
match Pin::new(task).poll(cx) {
Poll::Ready(Ok(Ok(()))) => Poll::Ready(Ok(())),
Poll::Ready(Ok(Err(e))) => {
Poll::Ready(Err(io::Error::other(format!("S3 upload failed: {}", e))))
}
Poll::Ready(Err(e)) => Poll::Ready(Err(io::Error::other(format!(
"Upload task panicked: {}",
e
)))),
Poll::Pending => Poll::Pending,
}
} else {
Poll::Ready(Ok(()))
}
}
}
impl AsyncSeek for S3ZipWriter {
fn start_seek(self: Pin<&mut Self>, position: io::SeekFrom) -> io::Result<()> {
match position {
io::SeekFrom::Current(0) => Ok(()), _ => Err(io::Error::new(
io::ErrorKind::Unsupported,
"S3 writer does not support seeking",
)),
}
}
fn poll_complete(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<u64>> {
Poll::Ready(Ok(self.position))
}
}
impl Unpin for S3ZipWriter {}
#[allow(dead_code)]
async fn upload_worker(
client: Client,
bucket: String,
key: String,
mut rx: mpsc::UnboundedReceiver<UploadCommand>,
) -> Result<()> {
let mut upload_id: Option<String> = None;
let mut parts: Vec<CompletedPart> = Vec::new();
while let Some(cmd) = rx.recv().await {
match cmd {
UploadCommand::UploadPart { part_number, data } => {
if upload_id.is_none() {
let response = client
.create_multipart_upload()
.bucket(&bucket)
.key(&key)
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to create multipart upload: {}",
e
)))
})?;
upload_id = Some(
response
.upload_id()
.ok_or_else(|| {
SZipError::Io(io::Error::other("No upload_id returned from S3"))
})?
.to_string(),
);
}
let response = client
.upload_part()
.bucket(&bucket)
.key(&key)
.upload_id(upload_id.as_ref().unwrap())
.part_number(part_number as i32)
.body(ByteStream::from(data))
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to upload part {}: {}",
part_number, e
)))
})?;
let etag = response
.e_tag()
.ok_or_else(|| {
SZipError::Io(io::Error::other(format!(
"No ETag returned for part {}",
part_number
)))
})?
.to_string();
parts.push(
CompletedPart::builder()
.part_number(part_number as i32)
.e_tag(etag)
.build(),
);
}
UploadCommand::Complete { final_data } => {
if let Some(data) = final_data {
if !data.is_empty() {
if upload_id.is_none() {
let response = client
.create_multipart_upload()
.bucket(&bucket)
.key(&key)
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to create multipart upload: {}",
e
)))
})?;
upload_id = Some(
response
.upload_id()
.ok_or_else(|| {
SZipError::Io(io::Error::other(
"No upload_id returned from S3",
))
})?
.to_string(),
);
}
let part_number = parts.len() + 1;
let response = client
.upload_part()
.bucket(&bucket)
.key(&key)
.upload_id(upload_id.as_ref().unwrap())
.part_number(part_number as i32)
.body(ByteStream::from(data))
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to upload final part: {}",
e
)))
})?;
let etag = response
.e_tag()
.ok_or_else(|| {
SZipError::Io(io::Error::other("No ETag returned for final part"))
})?
.to_string();
parts.push(
CompletedPart::builder()
.part_number(part_number as i32)
.e_tag(etag)
.build(),
);
}
}
if let Some(id) = upload_id {
client
.complete_multipart_upload()
.bucket(&bucket)
.key(&key)
.upload_id(&id)
.multipart_upload(
CompletedMultipartUpload::builder()
.set_parts(Some(parts))
.build(),
)
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to complete multipart upload: {}",
e
)))
})?;
}
break;
}
}
}
Ok(())
}
async fn upload_worker_concurrent(
client: Client,
bucket: String,
key: String,
mut rx: mpsc::UnboundedReceiver<UploadCommand>,
max_concurrent: usize,
) -> Result<()> {
use futures_util::stream::{FuturesUnordered, StreamExt};
let client = Arc::new(client);
let bucket = Arc::new(bucket);
let key = Arc::new(key);
let mut upload_id: Option<String> = None;
let mut completed_parts: Vec<(usize, CompletedPart)> = Vec::new();
let mut upload_futures = FuturesUnordered::new();
let mut pending_parts: Vec<(usize, Vec<u8>)> = Vec::new();
while let Some(cmd) = rx.recv().await {
match cmd {
UploadCommand::UploadPart { part_number, data } => {
if upload_id.is_none() {
let response = client
.create_multipart_upload()
.bucket(bucket.as_ref())
.key(key.as_ref())
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to create multipart upload: {}",
e
)))
})?;
upload_id = Some(
response
.upload_id()
.ok_or_else(|| {
SZipError::Io(io::Error::other("No upload_id returned from S3"))
})?
.to_string(),
);
}
let upload_id_clone = upload_id.clone().unwrap();
if upload_futures.len() < max_concurrent {
let fut = upload_part_with_retry(
client.clone(),
bucket.clone(),
key.clone(),
upload_id_clone,
part_number,
data,
);
upload_futures.push(fut);
} else {
pending_parts.push((part_number, data));
}
while let Some(result) = upload_futures.next().await {
let (part_num, completed_part) = result?;
completed_parts.push((part_num, completed_part));
if let Some((pn, pdata)) = pending_parts.pop() {
let fut = upload_part_with_retry(
client.clone(),
bucket.clone(),
key.clone(),
upload_id.clone().unwrap(),
pn,
pdata,
);
upload_futures.push(fut);
}
if upload_futures.len() < max_concurrent {
break;
}
}
}
UploadCommand::Complete { final_data } => {
if let Some(data) = final_data {
if !data.is_empty() {
if upload_id.is_none() {
let response = client
.create_multipart_upload()
.bucket(bucket.as_ref())
.key(key.as_ref())
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to create multipart upload: {}",
e
)))
})?;
upload_id = Some(
response
.upload_id()
.ok_or_else(|| {
SZipError::Io(io::Error::other(
"No upload_id returned from S3",
))
})?
.to_string(),
);
}
let part_number =
completed_parts.len() + upload_futures.len() + pending_parts.len() + 1;
let fut = upload_part_with_retry(
client.clone(),
bucket.clone(),
key.clone(),
upload_id.clone().unwrap(),
part_number,
data,
);
upload_futures.push(fut);
}
}
while let Some(result) = upload_futures.next().await {
let (part_num, completed_part) = result?;
completed_parts.push((part_num, completed_part));
}
completed_parts.sort_by_key(|(part_num, _)| *part_num);
let parts: Vec<_> = completed_parts.into_iter().map(|(_, p)| p).collect();
if let Some(id) = upload_id {
client
.complete_multipart_upload()
.bucket(bucket.as_ref())
.key(key.as_ref())
.upload_id(&id)
.multipart_upload(
CompletedMultipartUpload::builder()
.set_parts(Some(parts))
.build(),
)
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to complete multipart upload: {}",
e
)))
})?;
}
break;
}
}
}
Ok(())
}
async fn upload_part_with_retry(
client: Arc<Client>,
bucket: Arc<String>,
key: Arc<String>,
upload_id: String,
part_number: usize,
data: Vec<u8>,
) -> Result<(usize, CompletedPart)> {
const MAX_RETRIES: u32 = 3;
const BASE_DELAY_MS: u64 = 100;
let mut retries = 0;
loop {
match client
.upload_part()
.bucket(bucket.as_ref())
.key(key.as_ref())
.upload_id(&upload_id)
.part_number(part_number as i32)
.body(ByteStream::from(data.clone()))
.send()
.await
{
Ok(response) => {
let etag = response
.e_tag()
.ok_or_else(|| {
SZipError::Io(io::Error::other(format!(
"No ETag returned for part {}",
part_number
)))
})?
.to_string();
let completed_part = CompletedPart::builder()
.part_number(part_number as i32)
.e_tag(etag)
.build();
return Ok((part_number, completed_part));
}
Err(_e) if retries < MAX_RETRIES => {
retries += 1;
let delay = BASE_DELAY_MS * 2_u64.pow(retries - 1);
tokio::time::sleep(Duration::from_millis(delay)).await;
}
Err(e) => {
return Err(SZipError::Io(io::Error::other(format!(
"Failed to upload part {} after {} retries: {}",
part_number, MAX_RETRIES, e
))));
}
}
}
}
impl Drop for S3ZipWriter {
fn drop(&mut self) {
}
}
use tokio::io::AsyncRead;
pub struct S3ZipReader {
client: Client,
bucket: String,
key: String,
position: u64,
size: u64,
#[allow(clippy::type_complexity)]
read_future: Option<Pin<Box<dyn Future<Output = io::Result<Vec<u8>>> + Send>>>,
}
pub struct S3ZipReaderBuilder {
client: Option<Client>,
bucket: String,
key: String,
endpoint_url: Option<String>,
region: Option<String>,
force_path_style: bool,
}
impl S3ZipReader {
pub async fn new(
client: Client,
bucket: impl Into<String>,
key: impl Into<String>,
) -> Result<Self> {
Self::builder()
.client(client)
.bucket(bucket)
.key(key)
.build()
.await
}
pub fn builder() -> S3ZipReaderBuilder {
S3ZipReaderBuilder {
client: None,
bucket: String::new(),
key: String::new(),
endpoint_url: None,
region: None,
force_path_style: false,
}
}
pub fn size(&self) -> u64 {
self.size
}
}
impl S3ZipReaderBuilder {
pub fn client(mut self, client: Client) -> Self {
self.client = Some(client);
self
}
pub fn bucket(mut self, bucket: impl Into<String>) -> Self {
self.bucket = bucket.into();
self
}
pub fn key(mut self, key: impl Into<String>) -> Self {
self.key = key.into();
self
}
pub fn endpoint_url(mut self, url: impl Into<String>) -> Self {
self.endpoint_url = Some(url.into());
self
}
pub fn region(mut self, region: impl Into<String>) -> Self {
self.region = Some(region.into());
self
}
pub fn force_path_style(mut self, force: bool) -> Self {
self.force_path_style = force;
self
}
pub async fn build(self) -> Result<S3ZipReader> {
let client = match self.client {
Some(c) => c,
None => {
let mut config_loader = aws_config::from_env();
if let Some(ref endpoint) = self.endpoint_url {
config_loader = config_loader.endpoint_url(endpoint);
}
if let Some(ref region) = self.region {
config_loader = config_loader.region(aws_config::Region::new(region.clone()));
}
let sdk_config = config_loader.load().await;
let mut s3_config = aws_sdk_s3::config::Builder::from(&sdk_config);
if self.force_path_style {
s3_config = s3_config.force_path_style(true);
}
Client::from_conf(s3_config.build())
}
};
let head = client
.head_object()
.bucket(&self.bucket)
.key(&self.key)
.send()
.await
.map_err(|e| {
SZipError::Io(io::Error::other(format!(
"Failed to get S3 object metadata: {}",
e
)))
})?;
let size = head
.content_length()
.ok_or_else(|| SZipError::Io(io::Error::other("S3 object has no content length")))?
as u64;
Ok(S3ZipReader {
client,
bucket: self.bucket,
key: self.key,
position: 0,
size,
read_future: None,
})
}
}
impl AsyncRead for S3ZipReader {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if let Some(fut) = self.read_future.as_mut() {
match fut.as_mut().poll(cx) {
Poll::Ready(Ok(bytes)) => {
let n = bytes.len().min(buf.remaining());
buf.put_slice(&bytes[..n]);
self.position += n as u64;
self.read_future = None;
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(e)) => {
self.read_future = None;
return Poll::Ready(Err(e));
}
Poll::Pending => return Poll::Pending,
}
}
let start = self.position;
let end = (start + buf.remaining() as u64 - 1).min(self.size - 1);
if start >= self.size {
return Poll::Ready(Ok(())); }
let range = format!("bytes={}-{}", start, end);
let client = self.client.clone();
let bucket = self.bucket.clone();
let key = self.key.clone();
let fut = Box::pin(async move {
let response = client
.get_object()
.bucket(&bucket)
.key(&key)
.range(range)
.send()
.await
.map_err(|e| io::Error::other(format!("S3 GetObject failed: {}", e)))?;
let bytes = response
.body
.collect()
.await
.map_err(|e| io::Error::other(format!("Failed to read S3 body: {}", e)))?;
Ok::<_, io::Error>(bytes.into_bytes().to_vec())
});
self.read_future = Some(fut);
self.poll_read(cx, buf)
}
}
impl AsyncSeek for S3ZipReader {
fn start_seek(mut self: Pin<&mut Self>, position: io::SeekFrom) -> io::Result<()> {
let new_pos = match position {
io::SeekFrom::Start(pos) => pos as i64,
io::SeekFrom::End(offset) => self.size as i64 + offset,
io::SeekFrom::Current(offset) => self.position as i64 + offset,
};
if new_pos < 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"Invalid seek position",
));
}
self.position = new_pos as u64;
Ok(())
}
fn poll_complete(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<u64>> {
Poll::Ready(Ok(self.position))
}
}
impl Unpin for S3ZipReader {}
unsafe impl Send for S3ZipReader {}