#![allow(clippy::duration_suboptimal_units)]
use core::fmt;
use core::time::Duration;
use crate::hash::{FromHexError, Hash, to_hex};
use crate::refs::Ref;
pub use crate::refs::RefWriteCondition;
#[derive(Debug, thiserror::Error)]
pub enum TransportError {
#[error("pack not found on remote")]
PackNotFound,
#[error("access denied by remote")]
AccessDenied,
#[error("remote error: {0}")]
RemoteError(String),
#[error("ref CAS precondition failed")]
RefConflict,
#[error("invalid ref name: {0}")]
InvalidRef(String),
#[error("connection to remote failed")]
ConnectionFailed,
#[error("server error (status {status})")]
ServerError {
status: u16,
},
#[error("invalid response from remote")]
InvalidResponse,
#[error("protocol error")]
ProtocolError,
#[error("payload too large: {0} bytes")]
PayloadTooLarge(usize),
#[error("insecure scheme: plain http:// is allowed only for loopback hosts")]
InsecureScheme,
}
pub type TransportResult<T> = Result<T, TransportError>;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct PackKey(pub [u8; 32]);
impl PackKey {
#[must_use]
pub const fn new(bytes: [u8; 32]) -> Self {
Self(bytes)
}
#[must_use]
pub const fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
#[must_use]
pub fn to_hex(&self) -> String {
to_hex(&self.0)
}
#[must_use]
pub const fn from_hash(h: Hash) -> Self {
Self(h)
}
#[must_use]
pub const fn into_hash(self) -> Hash {
self.0
}
}
impl From<Hash> for PackKey {
fn from(h: Hash) -> Self {
Self(h)
}
}
impl From<PackKey> for Hash {
fn from(k: PackKey) -> Hash {
k.0
}
}
impl fmt::Display for PackKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.to_hex())
}
}
pub fn pack_key_from_hex(s: &str) -> Result<PackKey, FromHexError> {
let h = crate::hash::from_hex(s)?;
Ok(PackKey(h))
}
#[must_use]
pub fn is_retryable(err: &TransportError) -> bool {
match err {
TransportError::ConnectionFailed => true,
TransportError::ServerError { status } => *status >= 500 || *status == 429,
_ => false,
}
}
pub const BACKOFF_MAX_ATTEMPTS: u32 = 5;
pub const BACKOFF_INITIAL: Duration = Duration::from_secs(1);
pub const BACKOFF_CAP: Duration = Duration::from_secs(300);
#[cfg(target_pointer_width = "64")]
pub const PACK_BODY_LIMIT: u64 = 4 * 1024 * 1024 * 1024;
#[cfg(not(target_pointer_width = "64"))]
pub const PACK_BODY_LIMIT: u64 = usize::MAX as u64;
#[allow(clippy::cast_possible_truncation)]
pub const PACK_BODY_LIMIT_USIZE: usize = PACK_BODY_LIMIT as usize;
const _: () = assert!(
(PACK_BODY_LIMIT_USIZE as u64) == PACK_BODY_LIMIT,
"PACK_BODY_LIMIT does not fit in usize on this target",
);
#[derive(Debug, Clone)]
pub struct BackoffIterator {
next_delay: Duration,
attempts_remaining: u32,
cap: Duration,
}
impl BackoffIterator {
#[must_use]
pub const fn new() -> Self {
Self {
next_delay: BACKOFF_INITIAL,
attempts_remaining: BACKOFF_MAX_ATTEMPTS,
cap: BACKOFF_CAP,
}
}
#[must_use]
pub const fn with(initial: Duration, cap: Duration, attempts: u32) -> Self {
Self {
next_delay: initial,
attempts_remaining: attempts,
cap,
}
}
}
impl Default for BackoffIterator {
fn default() -> Self {
Self::new()
}
}
impl Iterator for BackoffIterator {
type Item = Duration;
fn next(&mut self) -> Option<Self::Item> {
if self.attempts_remaining == 0 {
return None;
}
self.attempts_remaining -= 1;
let current = self.next_delay;
let doubled = current.saturating_mul(2);
self.next_delay = if doubled > self.cap {
self.cap
} else {
doubled
};
Some(current)
}
}
pub fn retrying<T>(
mut op: impl FnMut() -> TransportResult<T>,
backoff: fn() -> BackoffIterator,
sleep: fn(Duration),
) -> TransportResult<T> {
let mut ladder = backoff();
loop {
match op() {
Ok(v) => return Ok(v),
Err(err) => {
if is_retryable(&err)
&& let Some(delay) = ladder.next()
{
sleep(delay);
continue;
}
return Err(err);
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PackChunk {
pub offset: u64,
pub data: Vec<u8>,
pub last: bool,
}
pub trait Transport: Send + Sync {
fn upload_pack(&self, bytes: &[u8], key: &PackKey) -> TransportResult<()>;
fn download_pack(&self, key: &PackKey) -> TransportResult<Vec<u8>>;
fn upload_pack_streaming(
&self,
key: &PackKey,
total_bytes: u64,
chunks: &mut dyn Iterator<Item = TransportResult<PackChunk>>,
) -> TransportResult<()> {
if total_bytes > PACK_BODY_LIMIT {
return Err(TransportError::PayloadTooLarge(PACK_BODY_LIMIT_USIZE));
}
let initial = usize::try_from(total_bytes).unwrap_or(PACK_BODY_LIMIT_USIZE);
let mut buf = Vec::with_capacity(initial);
let mut saw_last = false;
for chunk in chunks {
let c = chunk?;
if buf.len().saturating_add(c.data.len()) > PACK_BODY_LIMIT_USIZE {
return Err(TransportError::PayloadTooLarge(PACK_BODY_LIMIT_USIZE));
}
buf.extend_from_slice(&c.data);
if c.last {
saw_last = true;
break;
}
}
if !saw_last || buf.len() as u64 != total_bytes {
return Err(TransportError::ProtocolError);
}
self.upload_pack(&buf, key)
}
fn download_pack_streaming(
&self,
key: &PackKey,
) -> TransportResult<Box<dyn Iterator<Item = TransportResult<PackChunk>> + '_>> {
let bytes = self.download_pack(key)?;
Ok(Box::new(core::iter::once(Ok(PackChunk {
offset: 0,
data: bytes,
last: true,
}))))
}
fn pack_exists(&self, key: &PackKey) -> TransportResult<bool>;
fn upload_blob(&self, bytes: &[u8], key: &PackKey) -> TransportResult<()> {
self.upload_pack(bytes, key)
}
fn download_blob(&self, key: &PackKey) -> TransportResult<Vec<u8>> {
self.download_pack(key)
}
fn write_ref(&self, name: &str, hash: &Hash) -> TransportResult<()> {
self.update_ref(name, RefWriteCondition::Any, hash)
}
fn update_ref(
&self,
name: &str,
condition: RefWriteCondition,
hash: &Hash,
) -> TransportResult<()>;
fn read_ref(&self, name: &str) -> TransportResult<Option<Hash>>;
fn list_refs(&self, prefix: &str) -> TransportResult<Vec<Ref>>;
fn advance_refs(
&self,
head_ref: &str,
head_condition: RefWriteCondition,
head_value: &Hash,
packmap_ref: &str,
packmap_condition: RefWriteCondition,
packmap_value: &Hash,
) -> TransportResult<AdvanceOutcome> {
match self.update_ref(packmap_ref, packmap_condition, packmap_value) {
Ok(()) => {}
Err(TransportError::RefConflict) => return Ok(AdvanceOutcome::PackmapConflict),
Err(e) => return Err(e),
}
match self.update_ref(head_ref, head_condition, head_value) {
Ok(()) => Ok(AdvanceOutcome::Committed),
Err(TransportError::RefConflict) => Ok(AdvanceOutcome::HeadConflict),
Err(e) => Err(e),
}
}
fn supports_atomic_advance(&self) -> bool {
false
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AdvanceOutcome {
Committed,
HeadConflict,
PackmapConflict,
}
pub mod async_shim {
pub trait Executor: Send + Sync {
fn block_on<F, T>(&self, fut: F) -> T
where
F: core::future::Future<Output = T> + Send,
T: Send;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pack_key_hex_roundtrip() {
let bytes = [0x42u8; 32];
let pk = PackKey::new(bytes);
let hex = pk.to_hex();
assert_eq!(hex.len(), 64);
let pk2 = pack_key_from_hex(&hex).unwrap();
assert_eq!(pk, pk2);
}
#[test]
fn is_retryable_matches_spec() {
assert!(is_retryable(&TransportError::ConnectionFailed));
assert!(is_retryable(&TransportError::ServerError { status: 500 }));
assert!(is_retryable(&TransportError::ServerError { status: 503 }));
assert!(is_retryable(&TransportError::ServerError { status: 429 }));
assert!(!is_retryable(&TransportError::ServerError { status: 404 }));
assert!(!is_retryable(&TransportError::ServerError { status: 401 }));
assert!(!is_retryable(&TransportError::PackNotFound));
assert!(!is_retryable(&TransportError::AccessDenied));
assert!(!is_retryable(&TransportError::RefConflict));
}
#[test]
fn backoff_default_ladder_is_1_2_4_8_16() {
let delays: Vec<Duration> = BackoffIterator::new().collect();
assert_eq!(
delays,
vec![
Duration::from_secs(1),
Duration::from_secs(2),
Duration::from_secs(4),
Duration::from_secs(8),
Duration::from_secs(16),
]
);
}
#[test]
fn backoff_caps_at_max() {
let cap = Duration::from_secs(10);
let delays: Vec<Duration> = BackoffIterator::with(Duration::from_secs(8), cap, 5).collect();
assert_eq!(delays[0], Duration::from_secs(8));
for d in &delays[1..] {
assert!(*d <= cap);
}
}
#[derive(Default)]
struct RecordingTransport {
uploaded: std::sync::Mutex<Option<(Vec<u8>, PackKey)>>,
stored: std::sync::Mutex<std::collections::HashMap<[u8; 32], Vec<u8>>>,
}
impl Transport for RecordingTransport {
fn upload_pack(&self, bytes: &[u8], key: &PackKey) -> TransportResult<()> {
*self.uploaded.lock().unwrap() = Some((bytes.to_vec(), *key));
self.stored
.lock()
.unwrap()
.insert(*key.as_bytes(), bytes.to_vec());
Ok(())
}
fn download_pack(&self, key: &PackKey) -> TransportResult<Vec<u8>> {
self.stored
.lock()
.unwrap()
.get(key.as_bytes())
.cloned()
.ok_or(TransportError::PackNotFound)
}
fn pack_exists(&self, _key: &PackKey) -> TransportResult<bool> {
unimplemented!("not exercised by these tests")
}
fn update_ref(
&self,
_name: &str,
_condition: RefWriteCondition,
_hash: &Hash,
) -> TransportResult<()> {
unimplemented!("not exercised by these tests")
}
fn read_ref(&self, _name: &str) -> TransportResult<Option<Hash>> {
unimplemented!("not exercised by these tests")
}
fn list_refs(&self, _prefix: &str) -> TransportResult<Vec<Ref>> {
unimplemented!("not exercised by these tests")
}
}
fn chunks_of(data: &[u8], chunk_len: usize) -> Vec<PackChunk> {
if data.is_empty() {
return vec![PackChunk {
offset: 0,
data: Vec::new(),
last: true,
}];
}
let mut out = Vec::new();
let mut offset = 0usize;
while offset < data.len() {
let end = core::cmp::min(offset + chunk_len, data.len());
out.push(PackChunk {
offset: offset as u64,
data: data[offset..end].to_vec(),
last: end == data.len(),
});
offset = end;
}
out
}
#[test]
fn upload_pack_streaming_default_delegates_to_upload_pack() {
let t = RecordingTransport::default();
let payload = b"hello mkit pack bytes".repeat(100);
let key = PackKey::new([0x11; 32]);
let mut it = chunks_of(&payload, 7).into_iter().map(Ok);
t.upload_pack_streaming(&key, payload.len() as u64, &mut it)
.expect("streaming upload via default impl");
let (got_bytes, got_key) = t.uploaded.lock().unwrap().clone().expect("upload recorded");
assert_eq!(got_bytes, payload);
assert_eq!(got_key, key);
}
#[test]
fn upload_pack_streaming_default_rejects_missing_last_chunk() {
let t = RecordingTransport::default();
let key = PackKey::new([0x22; 32]);
let mut it = core::iter::empty();
let err = t
.upload_pack_streaming(&key, 0, &mut it)
.expect_err("must reject a stream with no last=true chunk");
assert!(matches!(err, TransportError::ProtocolError));
}
#[test]
fn upload_pack_streaming_default_rejects_total_bytes_mismatch() {
let t = RecordingTransport::default();
let key = PackKey::new([0x33; 32]);
let mut it = core::iter::once(Ok(PackChunk {
offset: 0,
data: vec![1, 2, 3],
last: true,
}));
let err = t
.upload_pack_streaming(&key, 10, &mut it)
.expect_err("must reject a total_bytes/accumulated-length mismatch");
assert!(matches!(err, TransportError::ProtocolError));
}
#[test]
fn upload_pack_streaming_default_propagates_chunk_error() {
let t = RecordingTransport::default();
let key = PackKey::new([0x44; 32]);
let mut it = core::iter::once(Err(TransportError::ConnectionFailed));
let err = t
.upload_pack_streaming(&key, 0, &mut it)
.expect_err("must propagate an error yielded mid-stream");
assert!(matches!(err, TransportError::ConnectionFailed));
}
#[test]
fn upload_pack_streaming_default_rejects_oversize_total() {
let t = RecordingTransport::default();
let key = PackKey::new([0x55; 32]);
let mut it = core::iter::empty();
let err = t
.upload_pack_streaming(&key, PACK_BODY_LIMIT + 1, &mut it)
.expect_err("must reject total_bytes above PACK_BODY_LIMIT");
assert!(matches!(err, TransportError::PayloadTooLarge(_)));
}
#[test]
fn download_pack_streaming_default_wraps_whole_pack() {
let t = RecordingTransport::default();
let key = PackKey::new([0x66; 32]);
let payload = vec![9u8; 4096];
t.upload_pack(&payload, &key).unwrap();
let mut stream = t.download_pack_streaming(&key).expect("stream opens");
let first = stream.next().expect("one chunk").expect("no error");
assert_eq!(first.data, payload);
assert!(first.last);
assert!(
stream.next().is_none(),
"default impl yields exactly one chunk"
);
}
#[test]
fn download_pack_streaming_default_propagates_not_found() {
let t = RecordingTransport::default();
let key = PackKey::new([0x77; 32]);
match t.download_pack_streaming(&key) {
Err(TransportError::PackNotFound) => {}
Err(other) => panic!("expected PackNotFound, got {other:?}"),
Ok(_) => panic!("missing pack must fail before any chunk is produced"),
}
}
fn test_backoff() -> BackoffIterator {
BackoffIterator::with(Duration::from_millis(1), Duration::from_millis(1), 5)
}
fn no_sleep(_delay: Duration) {}
#[test]
fn retrying_succeeds_on_first_try_without_sleeping() {
use core::sync::atomic::{AtomicUsize, Ordering};
fn record_sleep(_delay: Duration) {
SLEEPS.fetch_add(1, Ordering::SeqCst);
}
static CALLS: AtomicUsize = AtomicUsize::new(0);
static SLEEPS: AtomicUsize = AtomicUsize::new(0);
CALLS.store(0, Ordering::SeqCst);
SLEEPS.store(0, Ordering::SeqCst);
let result = retrying::<u32>(
|| {
CALLS.fetch_add(1, Ordering::SeqCst);
Ok(7)
},
test_backoff,
record_sleep,
);
assert_eq!(result.unwrap(), 7);
assert_eq!(CALLS.load(Ordering::SeqCst), 1);
assert_eq!(SLEEPS.load(Ordering::SeqCst), 0);
}
#[test]
fn retrying_recovers_after_transient_connection_failures() {
use core::sync::atomic::{AtomicUsize, Ordering};
static CALLS: AtomicUsize = AtomicUsize::new(0);
CALLS.store(0, Ordering::SeqCst);
let result = retrying::<u32>(
|| {
let n = CALLS.fetch_add(1, Ordering::SeqCst);
if n < 3 {
Err(TransportError::ConnectionFailed)
} else {
Ok(42)
}
},
test_backoff,
no_sleep,
);
assert_eq!(result.unwrap(), 42);
assert_eq!(CALLS.load(Ordering::SeqCst), 4);
}
#[test]
fn retrying_gives_up_after_ladder_exhausts() {
use core::sync::atomic::{AtomicUsize, Ordering};
static CALLS: AtomicUsize = AtomicUsize::new(0);
CALLS.store(0, Ordering::SeqCst);
let result = retrying::<u32>(
|| {
CALLS.fetch_add(1, Ordering::SeqCst);
Err(TransportError::ConnectionFailed)
},
test_backoff,
no_sleep,
);
assert!(matches!(result, Err(TransportError::ConnectionFailed)));
assert_eq!(CALLS.load(Ordering::SeqCst), 6);
}
#[test]
fn retrying_does_not_retry_non_retryable_errors() {
use core::sync::atomic::{AtomicUsize, Ordering};
static CALLS: AtomicUsize = AtomicUsize::new(0);
CALLS.store(0, Ordering::SeqCst);
let result = retrying::<u32>(
|| {
CALLS.fetch_add(1, Ordering::SeqCst);
Err(TransportError::PackNotFound)
},
test_backoff,
no_sleep,
);
assert!(matches!(result, Err(TransportError::PackNotFound)));
assert_eq!(CALLS.load(Ordering::SeqCst), 1);
}
}