use crate::{
CreateBucketOutcome, FileStream, StorageLayerError, StorageLayerImpl, UploadFileOptions,
UploadFileTag,
};
use aws_config::SdkConfig;
use aws_sdk_s3::{
config::Credentials,
error::SdkError,
operation::{
create_bucket::CreateBucketError, delete_bucket::DeleteBucketError,
delete_object::DeleteObjectError,
get_bucket_lifecycle_configuration::GetBucketLifecycleConfigurationError,
get_object::GetObjectError, head_bucket::HeadBucketError,
put_bucket_cors::PutBucketCorsError,
put_bucket_lifecycle_configuration::PutBucketLifecycleConfigurationError,
put_bucket_notification_configuration::PutBucketNotificationConfigurationError,
put_object::PutObjectError,
},
presigning::{PresignedRequest, PresigningConfig},
primitives::ByteStream,
types::{
BucketLifecycleConfiguration, BucketLocationConstraint, CorsConfiguration, CorsRule,
CreateBucketConfiguration, LifecycleExpiration, LifecycleRule, LifecycleRuleFilter,
NotificationConfiguration, QueueConfiguration, Tag,
},
};
use bytes::Bytes;
use chrono::{DateTime, TimeDelta, Utc};
use futures::Stream;
use serde::{Deserialize, Serialize};
use std::{error::Error, fmt::Debug, time::Duration};
use thiserror::Error;
type S3Client = aws_sdk_s3::Client;
#[derive(Debug, Default, Clone, Deserialize, Serialize)]
#[serde(default)]
pub struct S3StorageLayerFactoryConfig {
pub endpoint: S3Endpoint,
}
#[derive(Debug, Error)]
pub enum S3StorageLayerFactoryConfigError {
#[error("cannot use DOCBOX_S3_ENDPOINT without specifying DOCBOX_S3_ACCESS_KEY_ID")]
MissingAccessKeyId,
#[error("cannot use DOCBOX_S3_ENDPOINT without specifying DOCBOX_S3_ACCESS_KEY_SECRET")]
MissingAccessKeySecret,
}
impl S3StorageLayerFactoryConfig {
pub fn from_env() -> Result<Self, S3StorageLayerFactoryConfigError> {
let endpoint = S3Endpoint::from_env()?;
Ok(Self { endpoint })
}
}
#[derive(Default, Clone, Deserialize, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum S3Endpoint {
#[default]
Aws,
Custom {
endpoint: String,
external_endpoint: Option<String>,
access_key_id: String,
access_key_secret: String,
},
}
impl Debug for S3Endpoint {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Aws => write!(f, "Aws"),
Self::Custom { endpoint, .. } => f
.debug_struct("Custom")
.field("endpoint", endpoint)
.finish(),
}
}
}
impl S3Endpoint {
pub fn from_env() -> Result<Self, S3StorageLayerFactoryConfigError> {
match std::env::var("DOCBOX_S3_ENDPOINT") {
Ok(endpoint_url) => {
let access_key_id = std::env::var("DOCBOX_S3_ACCESS_KEY_ID")
.map_err(|_| S3StorageLayerFactoryConfigError::MissingAccessKeyId)?;
let access_key_secret = std::env::var("DOCBOX_S3_ACCESS_KEY_SECRET")
.map_err(|_| S3StorageLayerFactoryConfigError::MissingAccessKeySecret)?;
let external_endpoint = std::env::var("DOCBOX_S3_EXTERNAL_ENDPOINT").ok();
Ok(S3Endpoint::Custom {
endpoint: endpoint_url,
external_endpoint,
access_key_id,
access_key_secret,
})
}
Err(_) => Ok(S3Endpoint::Aws),
}
}
}
#[derive(Clone)]
pub struct S3StorageLayerFactory {
client: S3Client,
external_client: Option<S3Client>,
}
impl S3StorageLayerFactory {
pub fn from_config(aws_config: &SdkConfig, config: S3StorageLayerFactoryConfig) -> Self {
let (client, external_client) = match config.endpoint {
S3Endpoint::Aws => {
tracing::debug!("using aws s3 storage layer");
(S3Client::new(aws_config), None)
}
S3Endpoint::Custom {
endpoint,
external_endpoint,
access_key_id,
access_key_secret,
} => {
tracing::debug!("using custom s3 storage layer");
let credentials = Credentials::new(
access_key_id,
access_key_secret,
None,
None,
"docbox_key_provider",
);
let config_builder = aws_sdk_s3::config::Builder::from(aws_config)
.force_path_style(true)
.endpoint_url(endpoint)
.credentials_provider(credentials);
let external_client = match external_endpoint {
Some(external_endpoint) => {
let config = config_builder
.clone()
.endpoint_url(external_endpoint)
.build();
let client = S3Client::from_conf(config);
Some(client)
}
None => None,
};
let config = config_builder.build();
let client = S3Client::from_conf(config);
(client, external_client)
}
};
Self {
client,
external_client,
}
}
pub fn create_storage_layer(&self, bucket_name: String) -> S3StorageLayer {
S3StorageLayer::new(
self.client.clone(),
self.external_client.clone(),
bucket_name,
)
}
}
#[derive(Clone)]
pub struct S3StorageLayer {
bucket_name: String,
client: S3Client,
external_client: Option<S3Client>,
}
impl S3StorageLayer {
fn new(client: S3Client, external_client: Option<S3Client>, bucket_name: String) -> Self {
Self {
bucket_name,
client,
external_client,
}
}
async fn m1_storage_lifecycle_rules(&self) -> Result<(), StorageLayerError> {
let existing_lifecycle_configuration_rules = match self
.client
.get_bucket_lifecycle_configuration()
.bucket(&self.bucket_name)
.send()
.await
.inspect_err(|error| {
tracing::error!(
?error,
"failed to get existing bucket lifecycle configuration"
)
}) {
Ok(value) => value.rules,
Err(error) => match error.as_service_error() {
Some(error)
if error
.meta()
.code()
.is_some_and(|code| code == "NoSuchLifecycleConfiguration") =>
{
None
}
_ => return Err(S3StorageError::GetBucketLifecycleConfiguration(error).into()),
},
};
self.client
.put_bucket_lifecycle_configuration()
.bucket(&self.bucket_name)
.lifecycle_configuration(
BucketLifecycleConfiguration::builder()
.set_rules(existing_lifecycle_configuration_rules)
.rules(
LifecycleRule::builder()
.id("expire-1d")
.status(aws_sdk_s3::types::ExpirationStatus::Enabled)
.filter(
LifecycleRuleFilter::builder()
.tag(
Tag::builder()
.key("expire")
.value("1d")
.build()
.expect("invalid tag"),
)
.build(),
)
.expiration(LifecycleExpiration::builder().days(1).build())
.build()
.expect("invalid lifecycle rule configuration"),
)
.rules(
LifecycleRule::builder()
.id("expire-30d")
.status(aws_sdk_s3::types::ExpirationStatus::Enabled)
.filter(
LifecycleRuleFilter::builder()
.tag(
Tag::builder()
.key("expire")
.value("30d")
.build()
.expect("invalid tag"),
)
.build(),
)
.expiration(LifecycleExpiration::builder().days(30).build())
.build()
.expect("invalid lifecycle rule configuration"),
)
.build()
.expect("invalid lifecycle configuration"),
)
.send()
.await
.inspect_err(|error| {
tracing::error!(?error, "failed to put bucket lifecycle configuration")
})
.map_err(S3StorageError::PutBucketLifecycleConfiguration)?;
Ok(())
}
}
#[derive(Debug, Error)]
pub enum S3StorageError {
#[error("invalid server configuration (region)")]
MissingRegion,
#[error("failed to create storage bucket")]
CreateBucket(SdkError<CreateBucketError>),
#[error("failed to delete storage bucket")]
DeleteBucket(SdkError<DeleteBucketError>),
#[error("failed to get storage bucket")]
HeadBucket(SdkError<HeadBucketError>),
#[error("failed to store file object")]
PutObject(SdkError<PutObjectError>),
#[error("failed to calculate expiry timestamp")]
UnixTimeCalculation,
#[error("failed to create presigned store file object")]
PutObjectPresigned(SdkError<PutObjectError>),
#[error("failed to create presigned config")]
PresignedConfig,
#[error("failed to get presigned store file object")]
GetObjectPresigned(SdkError<GetObjectError>),
#[error("failed to create bucket notification queue config")]
QueueConfig,
#[error("failed to add bucket notification queue: {0}")]
PutBucketNotification(SdkError<PutBucketNotificationConfigurationError>),
#[error("failed to create bucket cors config")]
CreateCorsConfig,
#[error("failed to set bucket cors rules: {0}")]
PutBucketCors(SdkError<PutBucketCorsError>),
#[error("failed to delete file object")]
DeleteObject(SdkError<DeleteObjectError>),
#[error("failed to get file storage object")]
GetObject(SdkError<GetObjectError>),
#[error("failed to get bucket lifecycle configuration: {0}")]
GetBucketLifecycleConfiguration(SdkError<GetBucketLifecycleConfigurationError>),
#[error("failed to put bucket lifecycle configuration: {0}")]
PutBucketLifecycleConfiguration(SdkError<PutBucketLifecycleConfigurationError>),
}
const MIGRATION_NAMES: &[&str] = &["m1_storage_lifecycle_rules"];
impl StorageLayerImpl for S3StorageLayer {
fn bucket_name(&self) -> String {
self.bucket_name.clone()
}
async fn create_bucket(&self) -> Result<CreateBucketOutcome, StorageLayerError> {
let bucket_region = self
.client
.config()
.region()
.ok_or(S3StorageError::MissingRegion)?
.to_string();
let mut builder = self.client.create_bucket().bucket(&self.bucket_name);
if bucket_region != "us-east-1" {
let constraint = BucketLocationConstraint::from(bucket_region.as_str());
let cfg = CreateBucketConfiguration::builder()
.location_constraint(constraint)
.build();
builder = builder.create_bucket_configuration(cfg)
}
if let Err(error) = builder.send().await {
let already_exists = error
.as_service_error()
.is_some_and(|value| value.is_bucket_already_owned_by_you());
if already_exists {
tracing::debug!("bucket already exists");
return Ok(CreateBucketOutcome::Existing);
}
tracing::error!(?error, "failed to create bucket");
return Err(S3StorageError::CreateBucket(error).into());
}
Ok(CreateBucketOutcome::New)
}
async fn bucket_exists(&self) -> Result<bool, StorageLayerError> {
if let Err(error) = self
.client
.head_bucket()
.bucket(&self.bucket_name)
.send()
.await
{
if error
.as_service_error()
.is_some_and(|error| error.is_not_found())
{
return Ok(false);
}
return Err(S3StorageError::HeadBucket(error).into());
}
Ok(true)
}
async fn delete_bucket(&self) -> Result<(), StorageLayerError> {
if let Err(error) = self
.client
.delete_bucket()
.bucket(&self.bucket_name)
.send()
.await
{
if error
.as_service_error()
.and_then(|err| err.meta().code())
.is_some_and(|code| code == "NoSuchBucket")
{
tracing::debug!("bucket did not exist");
return Ok(());
}
tracing::error!(?error, "failed to delete bucket");
return Err(S3StorageError::DeleteBucket(error).into());
}
Ok(())
}
async fn upload_file(
&self,
key: &str,
body: Bytes,
options: UploadFileOptions,
) -> Result<(), StorageLayerError> {
let tagging = options.tags.map(|tags| {
use itertools::Itertools;
tags.into_iter()
.map(|tag| match tag {
UploadFileTag::ExpireDays1 => "expire=1d",
UploadFileTag::ExpireDays30 => "expire=30d",
})
.join("&")
});
self.client
.put_object()
.bucket(&self.bucket_name)
.content_type(options.content_type)
.key(key)
.set_tagging(tagging)
.body(body.into())
.send()
.await
.map_err(|error| {
tracing::error!(?error, "failed to store file object");
S3StorageError::PutObject(error)
})?;
Ok(())
}
async fn create_presigned(
&self,
key: &str,
size: i64,
) -> Result<(PresignedRequest, DateTime<Utc>), StorageLayerError> {
let expiry_time_minutes = 30;
let expires_at = Utc::now()
.checked_add_signed(TimeDelta::minutes(expiry_time_minutes))
.ok_or(S3StorageError::UnixTimeCalculation)?;
let client = match self.external_client.as_ref() {
Some(external_client) => external_client,
None => &self.client,
};
let result = client
.put_object()
.bucket(&self.bucket_name)
.key(key)
.content_length(size)
.presigned(
PresigningConfig::builder()
.expires_in(Duration::from_secs(60 * expiry_time_minutes as u64))
.build()
.map_err(|error| {
tracing::error!(?error, "Failed to create presigned store config");
S3StorageError::PresignedConfig
})?,
)
.await
.map_err(|error| {
tracing::error!(?error, "failed to create presigned store file object");
S3StorageError::PutObjectPresigned(error)
})?;
Ok((result, expires_at))
}
async fn create_presigned_download(
&self,
key: &str,
expires_in: Duration,
) -> Result<(PresignedRequest, DateTime<Utc>), StorageLayerError> {
let expires_at = Utc::now()
.checked_add_signed(TimeDelta::seconds(expires_in.as_secs() as i64))
.ok_or(S3StorageError::UnixTimeCalculation)?;
let client = match self.external_client.as_ref() {
Some(external_client) => external_client,
None => &self.client,
};
let result = client
.get_object()
.bucket(&self.bucket_name)
.key(key)
.presigned(PresigningConfig::expires_in(expires_in).map_err(|error| {
tracing::error!(?error, "failed to create presigned download config");
S3StorageError::PresignedConfig
})?)
.await
.map_err(|error| {
tracing::error!(?error, "failed to create presigned download");
S3StorageError::GetObjectPresigned(error)
})?;
Ok((result, expires_at))
}
async fn add_bucket_notifications(&self, sqs_arn: &str) -> Result<(), StorageLayerError> {
self.client
.put_bucket_notification_configuration()
.bucket(&self.bucket_name)
.notification_configuration(
NotificationConfiguration::builder()
.set_queue_configurations(Some(vec![
QueueConfiguration::builder()
.queue_arn(sqs_arn)
.events(aws_sdk_s3::types::Event::S3ObjectCreated)
.build()
.map_err(|error| {
tracing::error!(
?error,
"failed to create bucket notification queue config"
);
S3StorageError::QueueConfig
})?,
]))
.build(),
)
.send()
.await
.map_err(|error| {
tracing::error!(?error, "failed to add bucket notification queue");
S3StorageError::PutBucketNotification(error)
})?;
Ok(())
}
async fn set_bucket_cors_origins(&self, origins: Vec<String>) -> Result<(), StorageLayerError> {
if let Err(error) = self
.client
.put_bucket_cors()
.bucket(&self.bucket_name)
.cors_configuration(
CorsConfiguration::builder()
.cors_rules(
CorsRule::builder()
.allowed_headers("*")
.allowed_methods("PUT")
.set_allowed_origins(Some(origins))
.set_expose_headers(Some(Vec::new()))
.build()
.map_err(|error| {
tracing::error!(?error, "failed to create cors rule");
S3StorageError::CreateCorsConfig
})?,
)
.build()
.map_err(|error| {
tracing::error!(?error, "failed to create cors config");
S3StorageError::CreateCorsConfig
})?,
)
.send()
.await
{
if error
.raw_response()
.is_some_and(|response| response.status().as_u16() == 501)
{
tracing::warn!("storage s3 backend does not support PutBucketCors.. skipping..");
return Ok(());
}
tracing::error!(?error, "failed to add bucket cors");
return Err(S3StorageError::PutBucketCors(error).into());
};
Ok(())
}
async fn delete_file(&self, key: &str) -> Result<(), StorageLayerError> {
if let Err(error) = self
.client
.delete_object()
.bucket(&self.bucket_name)
.key(key)
.send()
.await
{
if error
.as_service_error()
.and_then(|err| err.source())
.and_then(|source| source.downcast_ref::<aws_sdk_s3::Error>())
.is_some_and(|err| matches!(err, aws_sdk_s3::Error::NoSuchKey(_)))
{
return Ok(());
}
tracing::error!(?error, "failed to delete file object");
return Err(S3StorageError::DeleteObject(error).into());
}
Ok(())
}
async fn get_file(&self, key: &str) -> Result<FileStream, StorageLayerError> {
let object = self
.client
.get_object()
.bucket(&self.bucket_name)
.key(key)
.send()
.await
.map_err(|error| {
tracing::error!(?error, "failed to get file storage object");
S3StorageError::GetObject(error)
})?;
let stream = FileStream {
stream: Box::pin(AwsFileStream { inner: object.body }),
};
Ok(stream)
}
async fn get_pending_migrations(
&self,
applied_names: Vec<String>,
) -> Result<Vec<String>, StorageLayerError> {
Ok(MIGRATION_NAMES
.iter()
.map(|name| name.to_string())
.filter(|name| !applied_names.contains(name))
.collect())
}
async fn apply_migration(&self, name: &str) -> Result<(), StorageLayerError> {
#[allow(clippy::single_match)]
match name {
"m1_storage_lifecycle_rules" => self.m1_storage_lifecycle_rules().await,
_ => Ok(()),
}
}
}
pub struct AwsFileStream {
inner: ByteStream,
}
impl AwsFileStream {
pub fn into_inner(self) -> ByteStream {
self.inner
}
}
impl Stream for AwsFileStream {
type Item = std::io::Result<Bytes>;
fn poll_next(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
let this = self.get_mut();
let inner = std::pin::Pin::new(&mut this.inner);
inner.poll_next(cx).map_err(std::io::Error::other)
}
}