use std::collections::HashMap;
use std::collections::hash_map::Entry;
use std::path::Path;
use std::sync::Arc;
use std::sync::RwLock;
use aws_config::BehaviorVersion;
use aws_credential_types::Credentials;
use aws_credential_types::provider::ProvideCredentials;
use aws_credential_types::provider::error::CredentialsError;
use aws_credential_types::provider::future;
use aws_sdk_s3::error::ProvideErrorMetadata;
use aws_sdk_s3::error::SdkError;
use aws_sdk_s3::primitives::ByteStream;
use aws_types::region::Region;
use tracing::debug;
use tracing::info;
use tracing::warn;
use crate::Error;
use crate::Res;
use crate::auth;
use crate::auth::OAuthParams;
use crate::auth::RoleInfo;
use crate::error::LoginError;
use crate::error::RemoteCatalogError;
use crate::error::S3Error;
use crate::error::S3ErrorKind;
use crate::io::remote::HostChecksums;
use crate::io::remote::HostConfig;
use crate::io::remote::HttpClient;
use crate::io::remote::Remote;
use crate::io::remote::describe_sdk_error;
use crate::io::remote::host::fetch_host_config;
use crate::io::remote::object::multipart_upload_and_sha256_chunksum;
use crate::io::remote::object::put_and_request_checksum;
use crate::io::storage::LocalStorage;
use crate::io::storage::auth::OAuthClient;
use crate::object_hash::ObjectHash;
use crate::paths::DomainPaths;
use quilt_uri::Host;
use quilt_uri::S3Uri;
use crate::io::remote::RemoteObjectStream;
async fn find_bucket_region(client: &impl HttpClient, bucket: &str) -> Res<String> {
let headers = client
.head(&format!("https://s3.amazonaws.com/{bucket}"))
.await?;
let region = headers
.get("x-amz-bucket-region")
.ok_or_else(|| RemoteCatalogError::BucketUnreachable(bucket.to_string()))?;
Ok(region.to_str()?.into())
}
pub(super) fn classify_s3_error(
code: Option<&str>,
status: Option<u16>,
described: &str,
fallback: fn(String) -> S3ErrorKind,
) -> S3ErrorKind {
match (code, status) {
(Some("AccessDenied"), _) | (None, Some(403)) => {
S3ErrorKind::AccessDenied(described.to_string())
}
_ => fallback(described.to_string()),
}
}
pub(super) fn classify_sdk_error<E>(
err: SdkError<E>,
fallback: fn(String) -> S3ErrorKind,
) -> S3ErrorKind
where
E: ProvideErrorMetadata + std::error::Error + Send + Sync + 'static,
{
let code = err
.as_service_error()
.and_then(ProvideErrorMetadata::code)
.map(str::to_owned);
let status = err.raw_response().map(|raw| raw.status().as_u16());
classify_s3_error(code.as_deref(), status, &describe_sdk_error(err), fallback)
}
async fn get_object_stream(client: &aws_sdk_s3::Client, s3_uri: &S3Uri) -> Res<RemoteObjectStream> {
let result = client.get_object().bucket(&s3_uri.bucket).key(&s3_uri.key);
let result = match &s3_uri.version {
Some(version) => result.version_id(version),
None => result,
};
let result = result.send().await.map_err(|err| match &err {
SdkError::ServiceError(svc) if svc.err().is_no_such_key() => {
Error::S3(S3Error::new(S3ErrorKind::NotFound(s3_uri.to_string())))
}
_ => Error::S3(S3Error::new(classify_sdk_error(err, S3ErrorKind::Raw))),
})?;
let uri_versioned = S3Uri {
version: result.version_id,
..s3_uri.clone()
};
Ok(RemoteObjectStream {
body: result.body,
uri: uri_versioned,
})
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
struct CredsRef {
region: Region,
host: Option<Host>,
}
#[derive(Clone, Debug)]
struct QuiltCredentialsProvider<H> {
auth: auth::Auth,
http: H,
host: Host,
}
impl<H> ProvideCredentials for QuiltCredentialsProvider<H>
where
H: HttpClient + Clone + std::fmt::Debug + Send + Sync + 'static,
{
fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a>
where
Self: 'a,
{
future::ProvideCredentials::new(async move {
let c = self
.auth
.get_credentials_or_refresh(&self.http, &self.host)
.await
.map_err(CredentialsError::provider_error)?;
Ok(Credentials::new(
c.access_key,
c.secret_key,
Some(c.token),
Some(c.expires_at.into()),
"quilt-registry",
))
})
}
}
#[derive(Debug)]
pub struct RemoteS3 {
auth: auth::Auth,
http: crate::io::remote::client::ReqwestClient,
s3: RwLock<HashMap<CredsRef, aws_sdk_s3::Client>>,
regions: RwLock<HashMap<String, Region>>,
}
impl RemoteS3 {
#[must_use]
pub fn new(paths: DomainPaths, storage: LocalStorage) -> Self {
RemoteS3 {
http: crate::io::remote::client::ReqwestClient::new(),
s3: RwLock::new(HashMap::new()),
regions: RwLock::new(HashMap::new()),
auth: auth::Auth::new(paths, Arc::new(storage)),
}
}
pub fn try_clone(&self) -> Res<Self> {
let s3 = match self.s3.read() {
Ok(s3) => s3.clone(),
Err(_) => return Err(Error::S3(S3Error::new(S3ErrorKind::RemoteInit))),
};
let regions = match self.regions.read() {
Ok(regions) => regions.clone(),
Err(_) => return Err(Error::S3(S3Error::new(S3ErrorKind::RemoteInit))),
};
Ok(RemoteS3 {
http: self.http.clone(),
s3: RwLock::new(s3),
regions: RwLock::new(regions),
auth: self.auth.clone(),
})
}
pub async fn login(&self, host: &Host, refresh_token: String) -> Res {
self.auth.login(&self.http, host, refresh_token).await
}
pub async fn login_oauth(&self, host: &Host, params: OAuthParams) -> Res {
self.auth.login_oauth(&self.http, host, params).await
}
pub async fn get_or_register_client(
&self,
host: &Host,
redirect_uri: &str,
) -> Res<OAuthClient> {
self.auth
.get_or_register_client(&self.http, host, redirect_uri)
.await
}
pub async fn refresh_roles(&self, host: &Host) -> Res<RoleInfo> {
self.auth.refresh_roles(&self.http, host).await
}
pub async fn switch_role(&self, host: &Host, role_name: &str) -> Res<RoleInfo> {
self.auth.switch_role(&self.http, host, role_name).await
}
pub async fn readable_buckets(&self, host: &Host) -> Res<Vec<String>> {
self.auth.readable_buckets(&self.http, host).await
}
pub async fn expire_credentials(&self, host: &Host) -> Res {
self.auth.expire_credentials(host).await
}
async fn get_region_for_bucket(&self, bucket: &str) -> Res<Region> {
{
if let Some(region) = self
.regions
.read()
.map_err(|e| S3Error::new(S3ErrorKind::PoisonLock(e.to_string())))?
.get(bucket)
{
return Ok(region.clone());
}
}
let region = find_bucket_region(&self.http, bucket).await?;
let mut map = self
.regions
.write()
.map_err(|e| S3Error::new(S3ErrorKind::PoisonLock(e.to_string())))?;
match map.entry(bucket.to_owned()) {
Entry::Occupied(entry) => Ok(entry.get().clone()),
Entry::Vacant(entry) => Ok(entry.insert(Region::new(region)).clone()),
}
}
async fn get_client_for_region(
&self,
host: Option<&Host>,
region: aws_types::region::Region,
) -> Res<aws_sdk_s3::Client> {
let creds_ref = CredsRef {
region: region.clone(),
host: host.cloned(),
};
let cached_client = {
let map = self
.s3
.read()
.map_err(|e| S3Error::new(S3ErrorKind::PoisonLock(e.to_string())))?;
map.get(&creds_ref).cloned()
};
if let Some(client) = cached_client {
info!("✔️ Using cached S3 client for region {:?}", region);
return Ok(client);
}
info!("⏳ Creating new S3 client for region {:?}", region);
let config = match host {
None => {
info!("⏳ No `&catalog=`, so we use credentials in ~/.aws");
let config = aws_config::defaults(BehaviorVersion::latest())
.region(region.clone())
.load()
.await;
if config.credentials_provider().is_none() {
return Err(Error::Login(LoginError::Required(None)));
}
config
}
Some(host) => {
self.auth
.get_credentials_or_refresh(&self.http, host)
.await?;
debug!("✔️ Got credentials for host {:?}", host);
aws_config::defaults(BehaviorVersion::latest())
.region(region.clone())
.credentials_provider(QuiltCredentialsProvider {
auth: self.auth.clone(),
http: self.http.clone(),
host: host.clone(),
})
.load()
.await
}
};
let client = aws_sdk_s3::Client::new(&config);
debug!("✔️ created new S3 client for region {:?}", region);
let mut map = self
.s3
.write()
.map_err(|e| S3Error::new(S3ErrorKind::PoisonLock(e.to_string())))?;
match map.entry(creds_ref) {
Entry::Occupied(mut entry) => {
entry.insert(client.clone());
Ok(client)
}
Entry::Vacant(entry) => Ok(entry.insert(client).clone()),
}
}
pub fn clear_client_cache(&self, host: Option<&Host>) {
let mut map = match self.s3.write() {
Ok(map) => map,
Err(poisoned) => poisoned.into_inner(),
};
match host {
Some(host) => map.retain(|creds_ref, _| creds_ref.host.as_ref() != Some(host)),
None => map.clear(),
}
}
async fn get_client_for_bucket(
&self,
host: Option<&Host>,
bucket: &str,
) -> Res<aws_sdk_s3::Client> {
let region = self.get_region_for_bucket(bucket).await?.clone();
self.get_client_for_region(host, region)
.await
.map_err(|e| match e {
Error::Login(LoginError::Required(_)) | Error::S3(_) => e,
_ => Error::S3(S3Error {
host: host.cloned(),
kind: S3ErrorKind::Client(e.to_string()),
}),
})
}
}
impl Remote for RemoteS3 {
async fn exists(&self, host: Option<&Host>, s3_uri: &S3Uri) -> Res<bool> {
debug!(
"⏳ Checking if object exists - host: {:?}, uri: {}",
host, s3_uri
);
let client = self.get_client_for_bucket(host, &s3_uri.bucket).await?;
let result = client.head_object().bucket(&s3_uri.bucket).key(&s3_uri.key);
let result = match &s3_uri.version {
Some(version) => result.version_id(version),
None => result,
};
match result.send().await {
Ok(_) => {
info!("✔️ Object exists at {}", s3_uri);
Ok(true)
}
Err(SdkError::ServiceError(err)) if err.err().is_not_found() => {
info!("ℹ️ Object does not exist at {}", s3_uri);
Ok(false)
}
Err(err) => {
warn!("❌ Failed to check object existence at {}: {}", s3_uri, err);
Err(Error::S3(S3Error {
host: host.cloned(),
kind: classify_sdk_error(err, S3ErrorKind::Exists),
}))
}
}
}
async fn get_object_stream(
&self,
host: Option<&Host>,
s3_uri: &S3Uri,
) -> Res<RemoteObjectStream> {
debug!(
"⏳ Getting object stream - host: {:?}, uri: {}",
host, s3_uri
);
let client = self.get_client_for_bucket(host, &s3_uri.bucket).await?;
match get_object_stream(&client, s3_uri).await {
Ok(stream) => {
info!("✔️ Created stream for object {}", s3_uri);
Ok(stream)
}
Err(e) if e.is_not_found() => {
info!("ℹ️ Object not found: {}", s3_uri);
Err(e)
}
Err(e) if e.is_access_denied() => {
warn!("❌ Access denied reading {}: {}", s3_uri, e);
Err(e)
}
Err(e) => {
warn!("❌ Failed to create stream for {}: {}", s3_uri, e);
Err(Error::S3(S3Error {
host: host.cloned(),
kind: S3ErrorKind::GetObjectStream(e.to_string()),
}))
}
}
}
async fn put_object(
&self,
host: Option<&Host>,
s3_uri: &S3Uri,
contents: impl Into<ByteStream>,
) -> Res {
self.get_client_for_bucket(host, &s3_uri.bucket)
.await?
.put_object()
.bucket(&s3_uri.bucket)
.key(&s3_uri.key)
.body(contents.into())
.send()
.await
.map_err(|err| {
Error::S3(S3Error {
host: host.cloned(),
kind: classify_sdk_error(err, S3ErrorKind::PutObject),
})
})?;
Ok(())
}
async fn resolve_url(&self, host: Option<&Host>, s3_uri: &S3Uri) -> Res<S3Uri> {
let client = self.get_client_for_bucket(host, &s3_uri.bucket).await?;
let result = client.head_object().bucket(&s3_uri.bucket).key(&s3_uri.key);
let result = match &s3_uri.version {
Some(version) => result.version_id(version),
None => result,
};
match result.send().await {
Ok(head) => Ok(S3Uri {
version: head.version_id,
..s3_uri.clone()
}),
Err(err) => Err(Error::S3(S3Error {
host: host.cloned(),
kind: classify_sdk_error(err, S3ErrorKind::ResolveUrl),
})),
}
}
async fn upload_file(
&self,
host_config: &HostConfig,
source_path: impl AsRef<Path>,
dest_uri: &S3Uri,
size: u64,
) -> Res<(S3Uri, ObjectHash)> {
let client = self
.get_client_for_bucket(host_config.host.as_ref(), &dest_uri.bucket)
.await?;
if host_config.checksums == HostChecksums::Sha256Chunked && size != 0 {
multipart_upload_and_sha256_chunksum(client, source_path, dest_uri, size).await
} else {
put_and_request_checksum(client, source_path, dest_uri, host_config).await
}
}
async fn host_config(&self, host: Option<&Host>) -> Res<HostConfig> {
fetch_host_config(&self.http, host).await
}
async fn verify_bucket(&self, bucket: &str) -> Res {
self.get_region_for_bucket(bucket).await?;
Ok(())
}
fn clear_client_cache(&self, host: Option<&Host>) {
RemoteS3::clear_client_cache(self, host);
}
}
#[cfg(test)]
mod tests {
use super::*;
use test_log::test;
use std::io::Write;
use async_trait::async_trait;
use reqwest::header::HeaderMap;
use tempfile::NamedTempFile;
use crate::fixtures::objects::LESS_THAN_8MB_HASH_B64;
use crate::fixtures::objects::ZERO_HASH_B64;
use crate::fixtures::objects::less_than_8mb;
use crate::fixtures::objects::zero_bytes;
use crate::io::storage::LocalStorage;
use crate::paths::DomainPaths;
#[test(tokio::test)]
async fn test_multipart_upload() -> Res<()> {
let mut temp_file = NamedTempFile::new()?;
temp_file.write_all(less_than_8mb())?;
let temp_path = temp_file.path();
let paths = DomainPaths::default();
let storage = LocalStorage::new();
let remote = RemoteS3::new(paths, storage);
let host_config = HostConfig {
checksums: HostChecksums::Sha256Chunked,
host: None,
};
let s3_uri =
S3Uri::try_from("s3://data-yaml-spec-tests/test_quilt_rs/multipart-upload.txt")?;
let size = less_than_8mb().len() as u64;
let result = remote
.upload_file(&host_config, temp_path, &s3_uri, size)
.await;
assert!(result.is_ok());
let (uploaded_uri, object_hash) = result?;
assert!(uploaded_uri.version.is_some());
assert_eq!(object_hash.to_string(), LESS_THAN_8MB_HASH_B64);
Ok(())
}
#[test(tokio::test)]
async fn test_zero_bytes_upload() -> Res<()> {
let mut temp_file = NamedTempFile::new()?;
temp_file.write_all(zero_bytes())?;
let temp_path = temp_file.path();
let paths = DomainPaths::default();
let storage = LocalStorage::new();
let remote = RemoteS3::new(paths, storage);
let host_config = HostConfig {
checksums: HostChecksums::Sha256Chunked,
host: None,
};
let s3_uri =
S3Uri::try_from("s3://data-yaml-spec-tests/test_quilt_rs/zero-bytes-file.txt")?;
let size = zero_bytes().len() as u64;
assert_eq!(size, 0);
let result = remote
.upload_file(&host_config, temp_path, &s3_uri, size)
.await;
assert!(result.is_ok());
let (uploaded_uri, object_hash) = result?;
assert!(uploaded_uri.version.is_some());
assert_eq!(object_hash.to_string(), ZERO_HASH_B64);
Ok(())
}
#[test(tokio::test)]
async fn test_crc64_upload() -> Res<()> {
let fixture_path = std::path::Path::new("fixtures/user-settings.mkfg");
let file_content = std::fs::read(fixture_path)?;
let mut temp_file = NamedTempFile::new()?;
temp_file.write_all(&file_content)?;
let temp_path = temp_file.path();
let paths = DomainPaths::default();
let storage = LocalStorage::new();
let remote = RemoteS3::new(paths, storage);
let host_config = HostConfig {
checksums: HostChecksums::Crc64,
host: None,
};
let s3_uri = S3Uri::try_from("s3://data-yaml-spec-tests/test_quilt_rs/crc64.txt")?;
let size = file_content.len() as u64;
let result = remote
.upload_file(&host_config, temp_path, &s3_uri, size)
.await;
assert!(result.is_ok());
let (uploaded_uri, object_hash) = result?;
assert!(uploaded_uri.version.is_some());
assert_eq!(object_hash.to_string(), "LZmmpqbBItw=");
Ok(())
}
#[test]
fn access_denied_code_classifies_as_access_denied() {
let err = classify_s3_error(
Some("AccessDenied"),
Some(403),
"AccessDenied: forbidden",
S3ErrorKind::Raw,
);
assert_eq!(
err,
S3ErrorKind::AccessDenied("AccessDenied: forbidden".to_string())
);
assert!(S3Error::new(err).is_access_denied());
}
#[test]
fn upload_denial_classifies_as_access_denied() {
let err = classify_s3_error(
Some("AccessDenied"),
Some(403),
"AccessDenied: forbidden",
S3ErrorKind::PutObject,
);
assert_eq!(
err,
S3ErrorKind::AccessDenied("AccessDenied: forbidden".to_string())
);
}
#[test]
fn non_denial_codes_keep_the_callers_fallback_kind() {
assert_eq!(
classify_s3_error(
Some("SlowDown"),
Some(503),
"SlowDown: throttled",
S3ErrorKind::PutObject
),
S3ErrorKind::PutObject("SlowDown: throttled".to_string())
);
assert_eq!(
classify_s3_error(
None,
Some(500),
"HTTP 500 (no error body)",
S3ErrorKind::UploadFile
),
S3ErrorKind::UploadFile("HTTP 500 (no error body)".to_string())
);
}
#[test]
fn expired_credential_codes_do_not_classify_as_access_denied() {
for code in ["ExpiredToken", "InvalidAccessKeyId"] {
let err = classify_s3_error(
Some(code),
Some(403),
&format!("{code}: nope"),
S3ErrorKind::Raw,
);
assert!(
!S3Error::new(err).is_access_denied(),
"{code} must not be read as a role denial"
);
}
}
#[test]
fn unknown_codes_stay_raw() {
let err = classify_s3_error(
Some("SlowDown"),
Some(503),
"SlowDown: throttled",
S3ErrorKind::Raw,
);
assert_eq!(err, S3ErrorKind::Raw("SlowDown: throttled".to_string()));
}
#[test]
fn missing_code_stays_raw() {
let err = classify_s3_error(
None,
Some(500),
"HTTP 500 (no error body)",
S3ErrorKind::Raw,
);
assert_eq!(
err,
S3ErrorKind::Raw("HTTP 500 (no error body)".to_string())
);
}
#[test]
fn a_bodyless_403_classifies_as_access_denied() {
let err = classify_s3_error(
None,
Some(403),
"HTTP 403 (no error body)",
S3ErrorKind::Exists,
);
assert_eq!(
err,
S3ErrorKind::AccessDenied("HTTP 403 (no error body)".to_string())
);
}
#[test]
fn other_bodyless_statuses_keep_the_callers_fallback_kind() {
for status in [404u16, 500, 503] {
let described = format!("HTTP {status} (no error body)");
assert_eq!(
classify_s3_error(None, Some(status), &described, S3ErrorKind::ResolveUrl),
S3ErrorKind::ResolveUrl(described.clone()),
"HTTP {status} is not a denial"
);
}
}
async fn spawn_canned_s3_endpoint(response: Vec<u8>) -> std::net::SocketAddr {
use tokio::io::AsyncReadExt;
use tokio::io::AsyncWriteExt;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((mut stream, _)) = listener.accept().await {
let response = response.clone();
tokio::spawn(async move {
let mut buf = [0u8; 8192];
let _ = stream.read(&mut buf).await;
let _ = stream.write_all(&response).await;
let _ = stream.shutdown().await;
});
}
});
addr
}
const DENIED_BUCKET: &str = "locked";
async fn access_denied_remote() -> RemoteS3 {
let body = "<?xml version=\"1.0\" encoding=\"UTF-8\"?>\
<Error><Code>AccessDenied</Code><Message>Access Denied</Message>\
<RequestId>REQ</RequestId><HostId>HID</HostId></Error>";
let addr = spawn_canned_s3_endpoint(
format!(
"HTTP/1.1 403 Forbidden\r\n\
Content-Type: application/xml\r\n\
Content-Length: {}\r\n\
Connection: close\r\n\r\n{body}",
body.len()
)
.into_bytes(),
)
.await;
let region = Region::new("us-east-1");
let denied_client = aws_sdk_s3::Client::from_conf(
aws_sdk_s3::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(region.clone())
.credentials_provider(Credentials::new("AK", "SK", None, None, "test"))
.endpoint_url(format!("http://{addr}"))
.force_path_style(true)
.build(),
);
let remote = RemoteS3::new(DomainPaths::default(), LocalStorage::new());
remote
.regions
.write()
.unwrap()
.insert(DENIED_BUCKET.to_string(), region.clone());
remote
.s3
.write()
.unwrap()
.insert(CredsRef { region, host: None }, denied_client);
remote
}
#[test(tokio::test)]
async fn get_object_stream_method_preserves_access_denied() -> Res<()> {
let remote = access_denied_remote().await;
let Err(err) =
Remote::get_object_stream(&remote, None, &S3Uri::try_from("s3://locked/x")?).await
else {
panic!("the stub endpoint always refuses");
};
assert!(
err.is_access_denied(),
"the method must not re-wrap the denial, got: {err}"
);
Ok(())
}
#[test(tokio::test)]
async fn exists_method_preserves_access_denied() -> Res<()> {
let remote = access_denied_remote().await;
let Err(err) = Remote::exists(&remote, None, &S3Uri::try_from("s3://locked/x")?).await
else {
panic!("the stub endpoint always refuses");
};
assert!(
err.is_access_denied(),
"the existence check must not re-wrap the denial, got: {err}"
);
Ok(())
}
#[test(tokio::test)]
async fn resolve_url_method_preserves_access_denied() -> Res<()> {
let remote = access_denied_remote().await;
let Err(err) = Remote::resolve_url(&remote, None, &S3Uri::try_from("s3://locked/x")?).await
else {
panic!("the stub endpoint always refuses");
};
assert!(
err.is_access_denied(),
"the resolve path must not re-wrap the denial, got: {err}"
);
Ok(())
}
#[test(tokio::test)]
async fn put_object_method_preserves_access_denied() -> Res<()> {
let remote = access_denied_remote().await;
let Err(err) = Remote::put_object(
&remote,
None,
&S3Uri::try_from("s3://locked/x")?,
ByteStream::from_static(b"payload"),
)
.await
else {
panic!("the stub endpoint always refuses");
};
assert!(
err.is_access_denied(),
"the upload path must not re-wrap the denial, got: {err}"
);
Ok(())
}
#[test(tokio::test)]
async fn upload_file_method_preserves_access_denied() -> Res<()> {
let remote = access_denied_remote().await;
let mut temp_file = NamedTempFile::new()?;
temp_file.write_all(b"payload")?;
let host_config = HostConfig {
checksums: HostChecksums::Crc64,
host: None,
};
let Err(err) = remote
.upload_file(
&host_config,
temp_file.path(),
&S3Uri::try_from("s3://locked/x")?,
7,
)
.await
else {
panic!("the stub endpoint always refuses");
};
assert!(
err.is_access_denied(),
"the upload path must not re-wrap the denial, got: {err}"
);
Ok(())
}
#[test(tokio::test)]
async fn multipart_upload_preserves_access_denied() -> Res<()> {
let remote = access_denied_remote().await;
let mut temp_file = NamedTempFile::new()?;
temp_file.write_all(b"payload")?;
let host_config = HostConfig {
checksums: HostChecksums::Sha256Chunked,
host: None,
};
let Err(err) = remote
.upload_file(
&host_config,
temp_file.path(),
&S3Uri::try_from("s3://locked/x")?,
7,
)
.await
else {
panic!("the stub endpoint always refuses");
};
assert!(
err.is_access_denied(),
"the multipart path must not re-wrap the denial, got: {err}"
);
Ok(())
}
fn dummy_client(region: &str) -> aws_sdk_s3::Client {
let conf = aws_sdk_s3::Config::builder()
.behavior_version(BehaviorVersion::latest())
.region(Region::new(region.to_string()))
.build();
aws_sdk_s3::Client::from_conf(conf)
}
#[test]
fn test_clear_client_cache_filters_by_host() {
use std::str::FromStr;
let host_a = Host::from_str("a.example.com").unwrap();
let host_b = Host::from_str("b.example.com").unwrap();
let remote = RemoteS3::new(DomainPaths::default(), LocalStorage::new());
{
let mut map = remote.s3.write().unwrap();
map.insert(
CredsRef {
region: Region::new("us-east-1"),
host: Some(host_a.clone()),
},
dummy_client("us-east-1"),
);
map.insert(
CredsRef {
region: Region::new("us-west-2"),
host: Some(host_b.clone()),
},
dummy_client("us-west-2"),
);
map.insert(
CredsRef {
region: Region::new("eu-west-1"),
host: None,
},
dummy_client("eu-west-1"),
);
}
remote.clear_client_cache(Some(&host_a));
{
let map = remote.s3.read().unwrap();
assert_eq!(map.len(), 2);
assert!(!map.keys().any(|k| k.host.as_ref() == Some(&host_a)));
assert!(map.keys().any(|k| k.host.as_ref() == Some(&host_b)));
}
remote.clear_client_cache(None);
assert!(remote.s3.read().unwrap().is_empty());
}
async fn role_api_shapes(remote: &RemoteS3, host: &Host) -> Res<()> {
let _: RoleInfo = remote.refresh_roles(host).await?;
let _: RoleInfo = remote.switch_role(host, "ReadOnly").await?;
let _: Vec<String> = remote.readable_buckets(host).await?;
Ok(())
}
#[test(tokio::test)]
async fn remote_exposes_the_role_api() -> Res<()> {
use std::str::FromStr;
use tempfile::TempDir;
let _ = role_api_shapes;
let temp = TempDir::new()?;
let remote = RemoteS3::new(
DomainPaths::new(temp.path().to_path_buf()),
LocalStorage::new(),
);
let host = Host::from_str("catalog.example.com").unwrap();
remote.expire_credentials(&host).await?;
Ok(())
}
#[test(tokio::test)]
async fn test_quilt_credentials_provider_returns_stored_creds() -> Res<()> {
use std::str::FromStr;
use tempfile::TempDir;
use crate::io::storage::auth::AuthIo;
use crate::io::storage::auth::Credentials as QuiltCreds;
let temp = TempDir::new()?;
let paths = DomainPaths::new(temp.path().to_path_buf());
let storage = Arc::new(LocalStorage::new());
let host = Host::from_str("catalog.example.com").unwrap();
let stored = QuiltCreds {
access_key: "AKIAEXAMPLE".to_string(),
secret_key: "secret".to_string(),
token: "session-token".to_string(),
expires_at: chrono::Utc::now() + chrono::Duration::hours(1),
};
let auth_io = AuthIo::new(Arc::clone(&storage), paths.auth_host(&host));
auth_io.write_credentials(&stored).await?;
let provider = QuiltCredentialsProvider {
auth: auth::Auth::new(paths, storage),
http: crate::io::remote::client::ReqwestClient::new(),
host,
};
let sdk_creds = provider.provide_credentials().await.unwrap();
assert_eq!(sdk_creds.access_key_id(), stored.access_key);
assert_eq!(sdk_creds.secret_access_key(), stored.secret_key);
assert_eq!(sdk_creds.session_token(), Some(stored.token.as_str()));
Ok(())
}
#[derive(Clone, Debug)]
struct RefreshMock {
refreshed_access_key: String,
}
#[async_trait]
impl HttpClient for RefreshMock {
async fn get<T: serde::de::DeserializeOwned>(
&self,
url: &str,
auth_token: Option<&str>,
) -> Res<T> {
if url.ends_with("/config.json") {
let body = serde_json::json!({
"registryUrl": "https://registry.example.com",
});
return Ok(serde_json::from_value(body)?);
}
if url.contains("/api/auth/get_credentials") {
assert_eq!(auth_token, Some("fresh-access-token"));
let body = serde_json::json!({
"AccessKeyId": self.refreshed_access_key,
"SecretAccessKey": "refreshed-secret",
"SessionToken": "refreshed-session",
"Expiration": (chrono::Utc::now() + chrono::Duration::hours(1))
.to_rfc3339(),
});
return Ok(serde_json::from_value(body)?);
}
panic!("unexpected GET: {url}");
}
async fn head(&self, _url: &str) -> Res<HeaderMap> {
unimplemented!("head not used")
}
async fn post<T: serde::de::DeserializeOwned>(
&self,
_url: &str,
_form_data: &HashMap<String, String>,
) -> Res<T> {
unimplemented!("fresh tokens → no token refresh")
}
async fn post_json<T: serde::de::DeserializeOwned, B: serde::Serialize + Send + Sync>(
&self,
_url: &str,
_body: &B,
) -> Res<T> {
unimplemented!("post_json not used")
}
async fn post_json_auth<
T: serde::de::DeserializeOwned,
B: serde::Serialize + Send + Sync,
>(
&self,
_url: &str,
_body: &B,
_auth_token: &str,
) -> Res<T> {
unimplemented!("post_json_auth not used")
}
}
#[test(tokio::test)]
async fn test_quilt_credentials_provider_refreshes_when_expired() -> Res<()> {
use std::str::FromStr;
use tempfile::TempDir;
use crate::io::storage::auth::AuthIo;
use crate::io::storage::auth::Credentials as QuiltCreds;
use crate::io::storage::auth::Tokens;
let temp = TempDir::new()?;
let paths = DomainPaths::new(temp.path().to_path_buf());
let storage = Arc::new(LocalStorage::new());
let host = Host::from_str("catalog.example.com").unwrap();
let auth_io = AuthIo::new(Arc::clone(&storage), paths.auth_host(&host));
auth_io
.write_credentials(&QuiltCreds {
access_key: "STALE".to_string(),
secret_key: "stale-secret".to_string(),
token: "stale-session".to_string(),
expires_at: chrono::Utc::now() - chrono::Duration::hours(1),
})
.await?;
auth_io
.write_tokens(&Tokens {
access_token: "fresh-access-token".to_string(),
refresh_token: "refresh-token".to_string(),
expires_at: chrono::Utc::now() + chrono::Duration::hours(1),
})
.await?;
let provider = QuiltCredentialsProvider {
auth: auth::Auth::new(paths, storage),
http: RefreshMock {
refreshed_access_key: "REFRESHED".to_string(),
},
host,
};
let sdk_creds = provider.provide_credentials().await.unwrap();
assert_eq!(sdk_creds.access_key_id(), "REFRESHED");
assert_eq!(sdk_creds.secret_access_key(), "refreshed-secret");
assert_eq!(sdk_creds.session_token(), Some("refreshed-session"));
Ok(())
}
}