use std::{
fmt::Display,
ops::{Range, RangeBounds},
};
use super::Result;
use bytes::Bytes;
use futures_util::{Stream, TryStreamExt, stream::StreamExt};
#[cfg(any(feature = "azure-base", feature = "http-base"))]
pub(crate) static RFC1123_FMT: &str = "%a, %d %h %Y %T GMT";
#[cfg(any(feature = "azure-base", feature = "http-base"))]
pub(crate) fn deserialize_rfc1123<'de, D>(
deserializer: D,
) -> Result<chrono::DateTime<chrono::Utc>, D::Error>
where
D: serde::Deserializer<'de>,
{
let s: String = serde::Deserialize::deserialize(deserializer)?;
let naive =
chrono::NaiveDateTime::parse_from_str(&s, RFC1123_FMT).map_err(serde::de::Error::custom)?;
Ok(chrono::TimeZone::from_utc_datetime(&chrono::Utc, &naive))
}
pub async fn collect_bytes<S, E>(mut stream: S, size_hint: Option<u64>) -> Result<Bytes, E>
where
E: Send,
S: Stream<Item = Result<Bytes, E>> + Send + Unpin,
{
let first = stream.next().await.transpose()?.unwrap_or_default();
match stream.next().await.transpose()? {
None => Ok(first),
Some(second) => {
let size_hint = size_hint.unwrap_or_else(|| first.len() as u64 + second.len() as u64);
let mut buf = Vec::with_capacity(size_hint as usize);
buf.extend_from_slice(&first);
buf.extend_from_slice(&second);
while let Some(maybe_bytes) = stream.next().await {
buf.extend_from_slice(&maybe_bytes?);
}
Ok(buf.into())
}
}
}
#[cfg(all(feature = "fs", not(target_arch = "wasm32")))]
pub(crate) async fn maybe_spawn_blocking<F, T>(f: F) -> Result<T>
where
F: FnOnce() -> Result<T> + Send + 'static,
T: Send + 'static,
{
match tokio::runtime::Handle::try_current() {
Ok(runtime) => runtime.spawn_blocking(f).await?,
Err(_) => f(),
}
}
pub const OBJECT_STORE_COALESCE_DEFAULT: u64 = 1024 * 1024;
pub(crate) const OBJECT_STORE_COALESCE_PARALLEL: usize = 10;
pub async fn coalesce_ranges<F, E, Fut>(
ranges: &[Range<u64>],
fetch: F,
coalesce: u64,
) -> Result<Vec<Bytes>, E>
where
F: Send + FnMut(Range<u64>) -> Fut,
E: Send,
Fut: std::future::Future<Output = Result<Bytes, E>> + Send,
{
let fetch_ranges = merge_ranges(ranges, coalesce);
let fetched: Vec<_> = futures_util::stream::iter(fetch_ranges.iter().cloned())
.map(fetch)
.buffered(OBJECT_STORE_COALESCE_PARALLEL)
.try_collect()
.await?;
Ok(ranges
.iter()
.map(|range| {
let idx = fetch_ranges.partition_point(|v| v.start <= range.start) - 1;
let fetch_range = &fetch_ranges[idx];
let fetch_bytes = &fetched[idx];
let start = range.start - fetch_range.start;
let end = range.end - fetch_range.start;
let range = (start as usize)..(end as usize).min(fetch_bytes.len());
fetch_bytes.slice(range)
})
.collect())
}
pub(crate) fn merge_ranges(ranges: &[Range<u64>], coalesce: u64) -> Vec<Range<u64>> {
if ranges.is_empty() {
return vec![];
}
let mut ranges = ranges.to_vec();
ranges.sort_unstable_by_key(|range| range.start);
let mut ret = Vec::with_capacity(ranges.len());
let mut start_idx = 0;
let mut end_idx = 1;
while start_idx != ranges.len() {
let mut range_end = ranges[start_idx].end;
while end_idx != ranges.len()
&& ranges[end_idx]
.start
.checked_sub(range_end)
.map(|delta| delta <= coalesce)
.unwrap_or(true)
{
range_end = range_end.max(ranges[end_idx].end);
end_idx += 1;
}
let start = ranges[start_idx].start;
let end = range_end;
ret.push(start..end);
start_idx = end_idx;
end_idx += 1;
}
ret
}
#[derive(Debug, PartialEq, Eq, Clone)]
pub enum GetRange {
Bounded(Range<u64>),
Offset(u64),
Suffix(u64),
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum InvalidGetRange {
#[error("Wanted range starting at {requested}, but object was only {length} bytes long")]
StartTooLarge { requested: u64, length: u64 },
#[error("Range started at {start} and ended at {end}")]
Inconsistent { start: u64, end: u64 },
#[error("Range {requested} is larger than system memory limit {max}")]
TooLarge { requested: u64, max: u64 },
}
impl GetRange {
pub fn is_valid(&self) -> Result<(), InvalidGetRange> {
if let Self::Bounded(r) = self {
if r.end <= r.start {
return Err(InvalidGetRange::Inconsistent {
start: r.start,
end: r.end,
});
}
if (r.end - r.start) > usize::MAX as u64 {
return Err(InvalidGetRange::TooLarge {
requested: r.start,
max: usize::MAX as u64,
});
}
}
Ok(())
}
pub fn as_range(&self, len: u64) -> Result<Range<u64>, InvalidGetRange> {
self.is_valid()?;
match self {
Self::Bounded(r) => {
if r.start >= len {
Err(InvalidGetRange::StartTooLarge {
requested: r.start,
length: len,
})
} else if r.end > len {
Ok(r.start..len)
} else {
Ok(r.clone())
}
}
Self::Offset(o) => {
if *o >= len {
Err(InvalidGetRange::StartTooLarge {
requested: *o,
length: len,
})
} else {
Ok(*o..len)
}
}
Self::Suffix(n) => Ok(len.saturating_sub(*n)..len),
}
}
}
impl Display for GetRange {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Bounded(r) => write!(f, "bytes={}-{}", r.start, r.end - 1),
Self::Offset(o) => write!(f, "bytes={o}-"),
Self::Suffix(n) => write!(f, "bytes=-{n}"),
}
}
}
impl<T: RangeBounds<u64>> From<T> for GetRange {
fn from(value: T) -> Self {
use std::ops::Bound::*;
let first = match value.start_bound() {
Included(i) => *i,
Excluded(i) => i + 1,
Unbounded => 0,
};
match value.end_bound() {
Included(i) => Self::Bounded(first..(i + 1)),
Excluded(i) => Self::Bounded(first..*i),
Unbounded => Self::Offset(first),
}
}
}
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
pub(crate) const STRICT_ENCODE_SET: percent_encoding::AsciiSet = percent_encoding::NON_ALPHANUMERIC
.remove(b'-')
.remove(b'.')
.remove(b'_')
.remove(b'~');
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
pub(crate) fn append_strict_query_pairs(url: &mut url::Url, pairs: &[(String, String)]) {
use percent_encoding::utf8_percent_encode;
use std::fmt::Write;
if pairs.is_empty() {
return;
}
let mut query = url.query().unwrap_or_default().to_owned();
for (key, value) in pairs {
if !query.is_empty() {
query.push('&');
}
let _ = write!(
query,
"{}={}",
utf8_percent_encode(key, &STRICT_ENCODE_SET),
utf8_percent_encode(value, &STRICT_ENCODE_SET),
);
}
url.set_query(Some(&query));
}
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
const RESERVED_SIGNED_HEADERS: [&str; 4] =
["host", "authorization", "content-length", "user-agent"];
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
pub(crate) fn validate_signed_url_extras(
store: &'static str,
extra_query: &[(String, String)],
signed_headers: &http::HeaderMap,
reserved_query_prefix: &str,
) -> Result<()> {
let err = |source: String| crate::Error::Generic {
store,
source: source.into(),
};
for (name, _) in extra_query {
if name.is_empty() {
return Err(err("query parameter name must not be empty".to_string()));
}
if name.to_ascii_lowercase().starts_with(reserved_query_prefix) {
return Err(err(format!(
"query parameter {name:?} is reserved for request signing and cannot be set via SignedUrlOptions"
)));
}
}
for (name, value) in signed_headers {
if RESERVED_SIGNED_HEADERS.contains(&name.as_str()) {
return Err(err(format!(
"header {name:?} is controlled by the signer and cannot be signed via SignedUrlOptions"
)));
}
if std::str::from_utf8(value.as_bytes()).is_err() {
return Err(err(format!(
"value of header {name:?} is not valid UTF-8 and cannot be signed via SignedUrlOptions"
)));
}
}
Ok(())
}
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
pub(crate) fn hex_digest(
crypto: &dyn crate::client::CryptoProvider,
bytes: &[u8],
) -> Result<String> {
let mut ctx = crypto.digest(crate::client::DigestAlgorithm::Sha256)?;
ctx.update(bytes);
Ok(hex_encode(ctx.finish()?))
}
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
pub(crate) fn hex_encode(bytes: &[u8]) -> String {
use std::fmt::Write;
let mut out = String::with_capacity(bytes.len() * 2);
for byte in bytes {
let _ = write!(out, "{byte:02x}");
}
out
}
#[cfg(test)]
mod tests {
use crate::Error;
use super::*;
use rand::{RngExt, rng};
use std::ops::Range;
async fn do_fetch(ranges: Vec<Range<u64>>, coalesce: u64) -> Vec<Range<u64>> {
let max = ranges.iter().map(|x| x.end).max().unwrap_or(0);
let src: Vec<_> = (0..max).map(|x| x as u8).collect();
let mut fetches = vec![];
let coalesced = coalesce_ranges::<_, Error, _>(
&ranges,
|range| {
fetches.push(range.clone());
let start = usize::try_from(range.start).unwrap();
let end = usize::try_from(range.end).unwrap();
futures_util::future::ready(Ok(Bytes::from(src[start..end].to_vec())))
},
coalesce,
)
.await
.unwrap();
assert_eq!(ranges.len(), coalesced.len());
for (range, bytes) in ranges.iter().zip(coalesced) {
assert_eq!(
bytes.as_ref(),
&src[usize::try_from(range.start).unwrap()..usize::try_from(range.end).unwrap()]
);
}
fetches
}
#[tokio::test]
async fn test_coalesce_ranges() {
let fetches = do_fetch(vec![], 0).await;
assert!(fetches.is_empty());
let fetches = do_fetch(vec![0..3; 1], 0).await;
assert_eq!(fetches, vec![0..3]);
let fetches = do_fetch(vec![0..2, 3..5], 0).await;
assert_eq!(fetches, vec![0..2, 3..5]);
let fetches = do_fetch(vec![0..1, 1..2], 0).await;
assert_eq!(fetches, vec![0..2]);
let fetches = do_fetch(vec![0..1, 2..72], 1).await;
assert_eq!(fetches, vec![0..72]);
let fetches = do_fetch(vec![0..1, 56..72, 73..75], 1).await;
assert_eq!(fetches, vec![0..1, 56..75]);
let fetches = do_fetch(vec![0..1, 5..6, 7..9, 2..3, 4..6], 1).await;
assert_eq!(fetches, vec![0..9]);
let fetches = do_fetch(vec![0..1, 5..6, 7..9, 2..3, 4..6], 1).await;
assert_eq!(fetches, vec![0..9]);
let fetches = do_fetch(vec![0..1, 6..7, 8..9, 10..14, 9..10], 4).await;
assert_eq!(fetches, vec![0..1, 6..14]);
}
#[tokio::test]
async fn test_coalesce_fuzz() {
let mut rand = rng();
for _ in 0..100 {
let object_len = rand.random_range(10..250);
let range_count = rand.random_range(0..10);
let ranges: Vec<_> = (0..range_count)
.map(|_| {
let start = rand.random_range(0..object_len);
let max_len = 20.min(object_len - start);
let len = rand.random_range(0..max_len);
start..start + len
})
.collect();
let coalesce = rand.random_range(1..5);
let fetches = do_fetch(ranges.clone(), coalesce).await;
for fetch in fetches.windows(2) {
assert!(
fetch[0].start <= fetch[1].start,
"fetches should be sorted, {:?} vs {:?}",
fetch[0],
fetch[1]
);
let delta = fetch[1].end - fetch[0].end;
assert!(
delta > coalesce,
"fetches should not overlap by {}, {:?} vs {:?} for {:?}",
coalesce,
fetch[0],
fetch[1],
ranges
);
}
}
}
#[test]
fn getrange_str() {
assert_eq!(GetRange::Offset(0).to_string(), "bytes=0-");
assert_eq!(GetRange::Bounded(10..19).to_string(), "bytes=10-18");
assert_eq!(GetRange::Suffix(10).to_string(), "bytes=-10");
}
#[test]
fn getrange_from() {
assert_eq!(Into::<GetRange>::into(10..15), GetRange::Bounded(10..15),);
assert_eq!(Into::<GetRange>::into(10..=15), GetRange::Bounded(10..16),);
assert_eq!(Into::<GetRange>::into(10..), GetRange::Offset(10),);
assert_eq!(Into::<GetRange>::into(..=15), GetRange::Bounded(0..16));
}
#[test]
fn test_as_range() {
let range = GetRange::Bounded(2..5);
assert_eq!(range.as_range(5).unwrap(), 2..5);
let range = range.as_range(4).unwrap();
assert_eq!(range, 2..4);
let range = GetRange::Bounded(3..3);
let err = range.as_range(2).unwrap_err().to_string();
assert_eq!(err, "Range started at 3 and ended at 3");
let range = GetRange::Bounded(2..2);
let err = range.as_range(3).unwrap_err().to_string();
assert_eq!(err, "Range started at 2 and ended at 2");
let range = GetRange::Suffix(3);
assert_eq!(range.as_range(3).unwrap(), 0..3);
assert_eq!(range.as_range(2).unwrap(), 0..2);
let range = GetRange::Suffix(0);
assert_eq!(range.as_range(0).unwrap(), 0..0);
let range = GetRange::Offset(2);
let err = range.as_range(2).unwrap_err().to_string();
assert_eq!(
err,
"Wanted range starting at 2, but object was only 2 bytes long"
);
let err = range.as_range(1).unwrap_err().to_string();
assert_eq!(
err,
"Wanted range starting at 2, but object was only 1 bytes long"
);
let range = GetRange::Offset(1);
assert_eq!(range.as_range(2).unwrap(), 1..2);
}
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
fn owned_pairs(pairs: &[(&str, &str)]) -> Vec<(String, String)> {
pairs
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect()
}
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
#[test]
fn append_strict_query_pairs_uses_percent_encoding() {
let mut url = url::Url::parse("https://example.com/object").unwrap();
append_strict_query_pairs(
&mut url,
&owned_pairs(&[("a key", "a value"), ("plus", "a+b/c=d")]),
);
assert_eq!(url.query().unwrap(), "a%20key=a%20value&plus=a%2Bb%2Fc%3Dd");
append_strict_query_pairs(&mut url, &owned_pairs(&[("x", "y")]));
assert!(url.query().unwrap().ends_with("&x=y"));
let before = url.query().unwrap().to_owned();
append_strict_query_pairs(&mut url, &[]);
assert_eq!(url.query().unwrap(), before);
}
#[cfg(any(feature = "aws-base", feature = "gcp-base"))]
#[test]
fn validate_signed_url_extras_accepts_and_rejects() {
use http::{HeaderMap, HeaderName, HeaderValue};
let header = |name: &'static str| {
let mut h = HeaderMap::new();
h.insert(HeaderName::from_static(name), HeaderValue::from_static("v"));
h
};
validate_signed_url_extras(
"S3",
&owned_pairs(&[("partNumber", "1"), ("uploadId", "abc")]),
&header("content-type"),
"x-amz-",
)
.unwrap();
for key in ["X-Amz-Signature", "x-amz-expires", "X-Amz-Security-Token"] {
let err = validate_signed_url_extras(
"S3",
&owned_pairs(&[(key, "x")]),
&HeaderMap::new(),
"x-amz-",
)
.unwrap_err();
assert!(
matches!(err, Error::Generic { .. }),
"{key} should be rejected"
);
}
assert!(
validate_signed_url_extras(
"GCS",
&owned_pairs(&[("X-Goog-Signature", "x")]),
&HeaderMap::new(),
"x-goog-"
)
.is_err()
);
assert!(
validate_signed_url_extras(
"S3",
&owned_pairs(&[("", "x")]),
&HeaderMap::new(),
"x-amz-"
)
.is_err()
);
for name in ["host", "authorization", "content-length", "user-agent"] {
let err = validate_signed_url_extras("S3", &[], &header(name), "x-amz-").unwrap_err();
assert!(
matches!(err, Error::Generic { .. }),
"{name} should be rejected"
);
}
let mut non_utf8 = HeaderMap::new();
non_utf8.insert(
HeaderName::from_static("x-custom"),
HeaderValue::from_bytes(&[0xFF]).unwrap(),
);
let err = validate_signed_url_extras("S3", &[], &non_utf8, "x-amz-").unwrap_err();
assert!(
matches!(err, Error::Generic { .. }),
"non-UTF-8 header value should be rejected"
);
}
}