use crate::error::Result;
use aws_sdk_s3 as s3;
use aws_sdk_s3::config::Builder as S3ConfigBuilder;
use aws_sdk_s3::config::Credentials;
use aws_sdk_s3::error::ProvideErrorMetadata;
use aws_types::region::Region;
#[derive(Debug, Clone)]
pub struct S3Client {
inner: s3::Client,
}
#[derive(Debug, Clone)]
pub struct S3ClientOptions {
pub access_key: String,
pub secret_key: String,
pub region: Option<String>,
pub endpoint: Option<String>,
pub force_path_style: bool,
}
impl Default for S3ClientOptions {
fn default() -> Self {
Self {
access_key: "".to_string(),
secret_key: "".to_string(),
region: None,
endpoint: None,
force_path_style: false,
}
}
}
impl S3ClientOptions {
pub fn with_access_key<S: Into<String>>(mut self, access_key: S) -> Self {
self.access_key = access_key.into();
self
}
pub fn with_secret_key<S: Into<String>>(mut self, secret_key: S) -> Self {
self.secret_key = secret_key.into();
self
}
pub fn with_region<S: Into<String>>(mut self, region: S) -> Self {
self.region = Some(region.into());
self
}
pub fn with_endpoint<S: Into<String>>(mut self, endpoint: S) -> Self {
self.endpoint = Some(endpoint.into());
self
}
pub fn with_force_path_style(mut self, force_path_style: bool) -> Self {
self.force_path_style = force_path_style;
self
}
}
impl S3Client {
pub async fn new(opts: S3ClientOptions) -> Result<Self> {
let credentials = Credentials::new(
opts.access_key,
opts.secret_key,
None, None, "user-supplied", );
let cred_provider = s3::config::SharedCredentialsProvider::new(credentials);
let mut cfg_loader = aws_config::defaults(aws_config::BehaviorVersion::latest())
.credentials_provider(cred_provider);
let region = opts.region.unwrap_or_else(|| "us-east-1".to_string());
cfg_loader = cfg_loader.region(Region::new(region));
let shared_cfg = cfg_loader.load().await;
let mut s3_cfg = S3ConfigBuilder::from(&shared_cfg);
if let Some(url) = opts.endpoint {
s3_cfg = s3_cfg
.endpoint_url(url)
.force_path_style(opts.force_path_style);
}
let client = s3::Client::from_conf(s3_cfg.build());
Ok(Self { inner: client })
}
pub async fn check_bucket_access(&self, bucket: &str) -> Result<bool> {
match self.inner.head_bucket().bucket(bucket).send().await {
Ok(_) => Ok(true),
Err(sdk_err) => {
if matches!(sdk_err.code(), Some("AccessDenied") | Some("Forbidden")) {
return Ok(false);
}
let aws_err: aws_sdk_s3::Error = sdk_err.into();
Err(aws_err.into())
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use wiremock::matchers::{method, path_regex};
use wiremock::{Mock, MockServer, ResponseTemplate};
const AK: &str = "TEST_AK";
const SK: &str = "TEST_SK";
async fn client(server: &MockServer) -> S3Client {
S3Client::new(S3ClientOptions {
access_key: AK.to_string(),
secret_key: SK.to_string(),
region: None, endpoint: Some(server.uri()), force_path_style: true,
})
.await
.expect("client init")
}
#[tokio::test]
async fn bucket_exists_returns_true() {
let server = MockServer::start().await;
let bucket = "mybucket";
Mock::given(method("HEAD"))
.and(path_regex(r"^/mybucket(/)?$")) .respond_with(ResponseTemplate::new(200))
.mount(&server)
.await;
let ok = client(&server)
.await
.check_bucket_access(bucket)
.await
.unwrap();
assert!(ok);
}
#[tokio::test]
async fn not_found_propagates_error() {
let server = MockServer::start().await;
let bucket = "missing";
Mock::given(method("HEAD"))
.and(path_regex(r"^/missing(/)?$"))
.respond_with(ResponseTemplate::new(404))
.mount(&server)
.await;
let err = client(&server)
.await
.check_bucket_access(bucket)
.await
.expect_err("should be Err");
assert!(format!("{err:?}").contains("NotFound"));
}
#[tokio::test]
async fn default_region_is_us_east_1() {
let server = MockServer::start().await;
let cli = client(&server).await;
let region = cli.inner.config().region().unwrap().as_ref();
assert_eq!(region, "us-east-1");
}
#[tokio::test]
async fn custom_region_is_set() {
let server = MockServer::start().await;
let region_name = "eu-west-1";
let cli = S3Client::new(S3ClientOptions {
access_key: AK.to_string(),
secret_key: SK.to_string(),
region: Some(region_name.to_string()),
endpoint: Some(server.uri()),
force_path_style: true,
})
.await
.expect("client init");
let region = cli.inner.config().region().unwrap().as_ref();
assert_eq!(region, region_name);
}
}