use std::io::Read;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex, OnceLock};
use crate::protocol::{PACK_BODY_LIMIT_USIZE, PackKey, TransportError, TransportResult};
use super::{Sequential, Shard, ShardSet, decode_shard_iter, default_parallel_strategy_for_len};
pub const MAX_BUFFERED_SHARD_BYTES: usize = PACK_BODY_LIMIT_USIZE.saturating_add(1024 * 1024);
pub const MAX_SHARD_WORKERS: usize = 32;
const READ_BYTES: usize = 16 * 1024;
#[derive(Debug)]
struct Budget {
used: AtomicUsize,
limit: usize,
}
impl Budget {
fn reserve(&self, bytes: usize) -> TransportResult<()> {
self.used
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |used| {
used.checked_add(bytes).filter(|&next| next <= self.limit)
})
.map(|_| ())
.map_err(|_| TransportError::PayloadTooLarge(self.limit.saturating_add(1)))
}
}
#[derive(Debug)]
struct Reservation {
budget: Arc<Budget>,
bytes: usize,
}
impl Reservation {
fn grow_to(&mut self, capacity: usize) -> TransportResult<()> {
if capacity > self.bytes {
self.budget.reserve(capacity - self.bytes)?;
self.bytes = capacity;
}
Ok(())
}
}
impl Drop for Reservation {
fn drop(&mut self) {
self.budget.used.fetch_sub(self.bytes, Ordering::AcqRel);
}
}
#[derive(Debug, Default)]
pub struct DownloadGroup {
cancelled: Arc<AtomicBool>,
}
impl DownloadGroup {
#[must_use]
pub fn token(&self) -> Cancellation {
Cancellation(Arc::clone(&self.cancelled))
}
pub fn cancel(&self) {
self.cancelled.store(true, Ordering::Release);
}
}
impl Drop for DownloadGroup {
fn drop(&mut self) {
self.cancel();
}
}
#[derive(Debug, Clone)]
pub struct Cancellation(Arc<AtomicBool>);
impl Cancellation {
pub fn check(&self) -> TransportResult<()> {
if self.0.load(Ordering::Acquire) {
return Err(TransportError::ProtocolError);
}
Ok(())
}
}
#[derive(Debug)]
pub struct DownloadedShard {
shard: Shard,
_reservation: Reservation,
}
impl DownloadedShard {
pub fn read(index: u16, reader: impl Read, cancel: &Cancellation) -> TransportResult<Self> {
static BUDGET: OnceLock<Arc<Budget>> = OnceLock::new();
let budget = BUDGET.get_or_init(|| {
Arc::new(Budget {
used: AtomicUsize::new(0),
limit: MAX_BUFFERED_SHARD_BYTES,
})
});
Self::read_with_budget(index, reader, cancel, Arc::clone(budget))
}
fn read_with_budget(
index: u16,
mut reader: impl Read,
cancel: &Cancellation,
budget: Arc<Budget>,
) -> TransportResult<Self> {
let mut reservation = Reservation { budget, bytes: 0 };
let mut bytes = Vec::new();
let mut chunk = [0u8; READ_BYTES];
loop {
cancel.check()?;
let n = reader
.read(&mut chunk)
.map_err(|_| TransportError::ConnectionFailed)?;
if n == 0 {
break;
}
let next = bytes
.len()
.checked_add(n)
.filter(|&n| n <= PACK_BODY_LIMIT_USIZE)
.ok_or(TransportError::PayloadTooLarge(
PACK_BODY_LIMIT_USIZE.saturating_add(1),
))?;
if next > bytes.capacity() {
reservation.grow_to(next)?;
bytes
.try_reserve_exact(next - bytes.len())
.map_err(|_| TransportError::PayloadTooLarge(next))?;
reservation.grow_to(bytes.capacity())?;
}
bytes.extend_from_slice(&chunk[..n]);
}
Ok(Self {
shard: Shard { index, bytes },
_reservation: reservation,
})
}
}
pub fn decode_downloaded_pack(
shards: &[DownloadedShard],
manifest: &ShardSet,
key: &PackKey,
) -> TransportResult<Vec<u8>> {
if manifest.pack_hash != *key.as_bytes() {
return Err(TransportError::InvalidResponse);
}
let size_hint = shards.iter().map(|shard| shard.shard.bytes.len()).sum();
let pack = match default_parallel_strategy_for_len(size_hint) {
Some(strategy) => {
decode_shard_iter(shards.iter().map(|shard| &shard.shard), manifest, &strategy)
}
None => decode_shard_iter(
shards.iter().map(|shard| &shard.shard),
manifest,
&Sequential,
),
}
.map_err(|_| TransportError::InvalidResponse)?;
key.verify_bytes(&pack)?;
Ok(pack)
}
pub fn download_shards(
config: super::Config,
fetch: impl Fn(u16, &Cancellation) -> TransportResult<DownloadedShard> + Send + Sync + 'static,
) -> TransportResult<Vec<DownloadedShard>> {
use std::sync::mpsc::{self, TryRecvError};
use std::time::Duration;
let total = config.total_shards();
if total > 256 {
return Err(TransportError::InvalidResponse);
}
let minimum = usize::from(config.minimum_shards.get());
let group = DownloadGroup::default();
let fetch = Arc::new(fetch);
let (tx, rx) = mpsc::channel();
let mut sender = Some(tx);
let mut next = 0u16;
let mut shards = Vec::with_capacity(minimum);
let mut failures = 0u16;
loop {
let result = match rx.try_recv() {
Ok(result) => result,
Err(TryRecvError::Disconnected) => return Err(TransportError::ConnectionFailed),
Err(TryRecvError::Empty) => {
if u32::from(next) < total
&& let Some(slot) = WorkerSlot::try_acquire()?
{
let index = next;
let fetch = Arc::clone(&fetch);
let tx = sender
.as_ref()
.ok_or(TransportError::ConnectionFailed)?
.clone();
let cancel = group.token();
std::thread::Builder::new()
.spawn(move || {
let _slot = slot;
let result = fetch(index, &cancel);
let _ = tx.send(result);
})
.map_err(|_| TransportError::ConnectionFailed)?;
next += 1;
if u32::from(next) == total {
drop(sender.take());
}
continue;
}
match rx.recv_timeout(Duration::from_millis(10)) {
Ok(result) => result,
Err(mpsc::RecvTimeoutError::Timeout) => continue,
Err(mpsc::RecvTimeoutError::Disconnected) => {
return Err(TransportError::ConnectionFailed);
}
}
}
};
if let Ok(shard) = result {
shards.push(shard);
if shards.len() == minimum {
return Ok(shards);
}
} else {
failures += 1;
if failures > config.extra_shards.get() {
return Err(TransportError::PackNotFound);
}
}
}
}
fn workers() -> &'static (Mutex<usize>, Condvar) {
static WORKERS: OnceLock<(Mutex<usize>, Condvar)> = OnceLock::new();
WORKERS.get_or_init(|| (Mutex::new(0), Condvar::new()))
}
#[derive(Debug)]
pub struct WorkerSlot;
impl WorkerSlot {
pub fn try_acquire() -> TransportResult<Option<Self>> {
let (lock, _) = workers();
let mut active = lock.lock().map_err(|_| TransportError::ConnectionFailed)?;
if *active >= MAX_SHARD_WORKERS {
return Ok(None);
}
*active += 1;
Ok(Some(Self))
}
pub fn acquire() -> TransportResult<Self> {
let (lock, ready) = workers();
let mut active = lock.lock().map_err(|_| TransportError::ConnectionFailed)?;
while *active >= MAX_SHARD_WORKERS {
active = ready
.wait(active)
.map_err(|_| TransportError::ConnectionFailed)?;
}
*active += 1;
Ok(Self)
}
}
impl Drop for WorkerSlot {
fn drop(&mut self) {
let (lock, ready) = workers();
let mut active = lock
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*active -= 1;
ready.notify_one();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn exited_shard_workers_report_failure() {
use std::num::NonZeroU16;
use std::sync::mpsc;
use std::time::Duration;
let config = super::super::Config {
minimum_shards: NonZeroU16::new(1).unwrap(),
extra_shards: NonZeroU16::new(1).unwrap(),
};
let (done, result) = mpsc::channel();
std::thread::spawn(move || {
let result = download_shards(config, |_, _| panic!("worker exited without a result"));
let _ = done.send(result);
});
assert!(matches!(
result
.recv_timeout(Duration::from_secs(2))
.expect("closed workers must not hang collection"),
Err(TransportError::ConnectionFailed)
));
}
#[test]
fn aggregate_budget_follows_live_buffers_and_releases_on_drop() {
let budget = Arc::new(Budget {
used: AtomicUsize::new(0),
limit: 96,
});
let group = DownloadGroup::default();
let first =
DownloadedShard::read_with_budget(0, &[0; 64][..], &group.token(), Arc::clone(&budget))
.unwrap();
assert!(matches!(
DownloadedShard::read_with_budget(1, &[1; 64][..], &group.token(), Arc::clone(&budget)),
Err(TransportError::PayloadTooLarge(_))
));
assert_eq!(budget.used.load(Ordering::Acquire), 64);
drop(first);
let second =
DownloadedShard::read_with_budget(1, &[1; 64][..], &group.token(), Arc::clone(&budget))
.unwrap();
drop(second);
assert_eq!(budget.used.load(Ordering::Acquire), 0);
}
#[test]
fn cancellation_stops_before_reading() {
struct MustNotRead;
impl Read for MustNotRead {
fn read(&mut self, _: &mut [u8]) -> std::io::Result<usize> {
panic!("cancelled request read the body");
}
}
let group = DownloadGroup::default();
let token = group.token();
drop(group);
assert!(matches!(
DownloadedShard::read(0, MustNotRead, &token),
Err(TransportError::ProtocolError)
));
}
}