use std::collections::{HashMap, HashSet};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, Instant};
use dig_dht::{ContentId, ProviderRecord};
use dig_nat::{AvailabilityItem, RangeRequest};
use futures::stream::{FuturesUnordered, StreamExt};
use tokio::sync::mpsc;
use crate::error::DownloadError;
use crate::gc::ActiveDownloads;
use crate::locate::ProviderLocator;
use crate::plan::{plan_ranges, Range, RangeState};
use crate::progress::{DownloadEvent, DownloadProgress, DownloadState, StateStore};
use crate::sink::Sink;
use crate::source::{FetchedRange, RangeTransport, SourceTracker};
use crate::verify::{ResourceCommitment, ResourceHasher, Verifier};
#[derive(Debug, Clone)]
pub struct DownloadConfig {
pub window: u64,
pub max_concurrency: usize,
pub max_inflight_per_source: usize,
pub base_backoff: Duration,
pub max_backoff: Duration,
pub max_relocate_attempts: usize,
pub max_range_attempts: usize,
pub verify_whole_resource: bool,
}
impl Default for DownloadConfig {
fn default() -> Self {
DownloadConfig {
window: 3 * 1024 * 1024,
max_concurrency: 8,
max_inflight_per_source: 4,
base_backoff: Duration::from_millis(200),
max_backoff: Duration::from_secs(10),
max_relocate_attempts: 4,
max_range_attempts: 6,
verify_whole_resource: true,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct DownloadOptions {
pub start_paused: bool,
pub resume_key: Option<String>,
}
#[derive(Debug)]
enum Control {
Pause,
Resume,
Cancel,
}
pub fn download_key(content: &ContentId) -> String {
content.to_key().to_hex()
}
pub struct Downloader {
locator: Arc<dyn ProviderLocator>,
transport: Arc<dyn RangeTransport>,
verifier: Arc<dyn Verifier>,
state_store: Arc<dyn StateStore>,
registry: Arc<ActiveDownloads>,
config: DownloadConfig,
}
impl Downloader {
pub fn new(
locator: Arc<dyn ProviderLocator>,
transport: Arc<dyn RangeTransport>,
verifier: Arc<dyn Verifier>,
state_store: Arc<dyn StateStore>,
config: DownloadConfig,
) -> Self {
Downloader {
locator,
transport,
verifier,
state_store,
registry: Arc::new(ActiveDownloads::new()),
config,
}
}
pub fn active_downloads(&self) -> Arc<ActiveDownloads> {
self.registry.clone()
}
pub fn download(
&self,
content: ContentId,
sink: Arc<dyn Sink>,
opts: DownloadOptions,
) -> DownloadHandle {
let key = opts
.resume_key
.clone()
.unwrap_or_else(|| download_key(&content));
let (control_tx, control_rx) = mpsc::channel(16);
let (events_tx, events_rx) = mpsc::channel(256);
let job = Job {
content,
key,
sink,
verifier: self.verifier.clone(),
transport: self.transport.clone(),
locator: self.locator.clone(),
state_store: self.state_store.clone(),
registry: self.registry.clone(),
config: self.config.clone(),
events: events_tx,
control: control_rx,
providers: Vec::new(),
commitment: None,
ranges: Vec::new(),
range_state: Vec::new(),
tracker: SourceTracker::new(self.config.base_backoff, self.config.max_backoff),
inflight_per_source: HashMap::new(),
resume: DownloadState::new(String::new()),
paused: opts.start_paused,
bytes_done: 0,
hasher: None,
relocate_attempts: 0,
relocated_since_progress: false,
total_failures: 0,
};
let task = tokio::spawn(job.run());
DownloadHandle {
control: control_tx,
events: events_rx,
task,
}
}
pub async fn gc(
&self,
dir: impl Into<std::path::PathBuf>,
ttl: Duration,
) -> Result<usize, DownloadError> {
crate::gc::TmpGc::new(dir, ttl, self.registry.clone())
.sweep()
.await
}
}
pub struct DownloadHandle {
control: mpsc::Sender<Control>,
events: mpsc::Receiver<DownloadEvent>,
task: tokio::task::JoinHandle<Result<u64, DownloadError>>,
}
impl DownloadHandle {
pub fn pause(&self) {
let _ = self.control.try_send(Control::Pause);
}
pub fn resume(&self) {
let _ = self.control.try_send(Control::Resume);
}
pub fn cancel(&self) {
let _ = self.control.try_send(Control::Cancel);
}
pub async fn next_event(&mut self) -> Option<DownloadEvent> {
self.events.recv().await
}
pub fn events(&mut self) -> &mut mpsc::Receiver<DownloadEvent> {
&mut self.events
}
pub async fn join(self) -> Result<u64, DownloadError> {
match self.task.await {
Ok(res) => res,
Err(_) => Err(DownloadError::TaskEnded),
}
}
}
type FetchOutput = (usize, String, Result<FetchedRange, DownloadError>);
struct Job {
content: ContentId,
key: String,
sink: Arc<dyn Sink>,
verifier: Arc<dyn Verifier>,
transport: Arc<dyn RangeTransport>,
locator: Arc<dyn ProviderLocator>,
state_store: Arc<dyn StateStore>,
registry: Arc<ActiveDownloads>,
config: DownloadConfig,
events: mpsc::Sender<DownloadEvent>,
control: mpsc::Receiver<Control>,
providers: Vec<ProviderRecord>,
commitment: Option<ResourceCommitment>,
ranges: Vec<Range>,
range_state: Vec<RangeState>,
tracker: SourceTracker,
inflight_per_source: HashMap<String, usize>,
resume: DownloadState,
paused: bool,
bytes_done: u64,
hasher: Option<ResourceHasher>,
relocate_attempts: usize,
relocated_since_progress: bool,
total_failures: usize,
}
impl Job {
async fn run(mut self) -> Result<u64, DownloadError> {
self.resume = match self.state_store.load(&self.key).await {
Ok(Some(state)) => state,
Ok(None) => DownloadState::new(self.key.clone()),
Err(e) => {
self.emit(DownloadEvent::Failed {
reason: e.to_string(),
})
.await;
return Err(e);
}
};
if !self.resume.chunk_lens.is_empty() {
match ResourceCommitment::from_first_frame(
self.resume.total_length,
self.resume.chunk_lens.clone(),
self.resume.root.clone(),
self.resume.inclusion_proof.clone(),
) {
Ok(c) => self.commitment = Some(c),
Err(_) => self.commitment = None,
}
}
let staging = self.sink.staging_path().map(|p| p.to_path_buf());
if let Some(path) = &staging {
self.registry.register(path.clone()).await;
}
let result = self.run_inner().await;
if let Some(path) = &staging {
self.registry.unregister(path).await;
}
result
}
async fn run_inner(&mut self) -> Result<u64, DownloadError> {
self.availability_item()?;
self.providers = self.locate_and_confirm().await?;
if self.providers.is_empty() {
let reason = format!("{:?}", self.content);
self.emit(DownloadEvent::Failed {
reason: format!("no providers for {reason}"),
})
.await;
return Err(DownloadError::NotFound { content: reason });
}
if self.commitment.is_none() {
self.establish_commitment().await?;
}
self.persist_commitment().await?;
let commitment = self.commitment.clone().expect("commitment established");
self.ranges = plan_ranges(&commitment.layout, self.config.window);
self.range_state = self
.ranges
.iter()
.map(|r| {
if self.resume.is_done(r.index) {
RangeState::Done
} else {
RangeState::Pending
}
})
.collect();
let resumed_ranges = self
.ranges
.iter()
.filter(|r| self.resume.is_done(r.index))
.count();
self.hasher = if self.config.verify_whole_resource && resumed_ranges == 0 {
Some(ResourceHasher::new())
} else {
None
};
self.bytes_done = self
.ranges
.iter()
.filter(|r| self.resume.is_done(r.index))
.map(|r| r.length)
.sum();
self.emit(DownloadEvent::Planned {
ranges_total: self.ranges.len(),
total_length: commitment.total_length,
})
.await;
self.schedule_loop().await?;
if let Some(hasher) = self.hasher.take() {
let hashed_len = hasher.hashed_len();
let leaf = hasher.finalize();
if let Err(e) = self
.verifier
.verify_resource_leaf(&commitment, &leaf, hashed_len)
{
self.emit(DownloadEvent::Failed {
reason: e.to_string(),
})
.await;
return Err(e.into());
}
}
self.sink.finalize().await?;
let _ = self.state_store.clear(&self.key).await;
self.emit(DownloadEvent::Completed {
total_length: commitment.total_length,
})
.await;
Ok(commitment.total_length)
}
async fn schedule_loop(&mut self) -> Result<(), DownloadError> {
let mut inflight: FuturesUnordered<Pin<Box<dyn Future<Output = FetchOutput> + Send>>> =
FuturesUnordered::new();
loop {
if !self.paused {
self.fill(&mut inflight);
}
if self.all_done() && inflight.is_empty() {
return Ok(());
}
let budget = self
.ranges
.len()
.saturating_mul(self.config.max_range_attempts)
.max(self.config.max_range_attempts);
if self.total_failures > budget {
let needed = self.pending_count();
self.emit(DownloadEvent::Failed {
reason: format!("provider set exhausted ({needed} range(s) unmet)"),
})
.await;
return Err(DownloadError::NoProviders { needed });
}
let mut wakeup: Option<Instant> = None;
if !self.paused && inflight.is_empty() && !self.all_done() {
if !self.relocated_since_progress
&& self.relocate_attempts < self.config.max_relocate_attempts
{
let added = self.relocate().await?;
self.relocated_since_progress = true;
if added > 0 {
continue; }
}
match self.earliest_backoff() {
Some(t) => wakeup = Some(t),
None => {
let needed = self.pending_count();
self.emit(DownloadEvent::Failed {
reason: format!("no live providers ({needed} range(s) unmet)"),
})
.await;
return Err(DownloadError::NoProviders { needed });
}
}
}
let sleep = wakeup.map(|t| {
let now = Instant::now();
tokio::time::sleep(t.saturating_duration_since(now))
});
tokio::select! {
ctrl = self.control.recv() => {
match ctrl {
Some(Control::Pause) => {
if !self.paused {
self.paused = true;
let _ = self.checkpoint().await;
self.emit(DownloadEvent::Paused).await;
}
}
Some(Control::Resume) => {
if self.paused {
self.paused = false;
self.emit(DownloadEvent::Resumed).await;
}
}
Some(Control::Cancel) | None => {
let _ = self.checkpoint().await;
self.emit(DownloadEvent::Failed { reason: "cancelled".into() }).await;
return Err(DownloadError::Cancelled);
}
}
}
Some((idx, peer, res)) = inflight.next(), if !inflight.is_empty() => {
self.handle_result(idx, peer, res).await?;
}
_ = async { sleep.unwrap().await }, if wakeup.is_some() => {
}
}
}
}
fn fill(
&mut self,
inflight: &mut FuturesUnordered<Pin<Box<dyn Future<Output = FetchOutput> + Send>>>,
) {
let now = Instant::now();
loop {
if inflight.len() >= self.config.max_concurrency {
break;
}
let Some(range_idx) = self.next_pending() else {
break;
};
let Some(peer) = self.pick_source(now) else {
break; };
self.range_state[range_idx] = RangeState::InFlight(peer.clone());
*self.inflight_per_source.entry(peer.clone()).or_insert(0) += 1;
inflight.push(self.fetch_future(range_idx, peer));
}
}
fn fetch_future(
&self,
range_idx: usize,
peer: String,
) -> Pin<Box<dyn Future<Output = FetchOutput> + Send>> {
let range = self.ranges[range_idx];
let provider = self
.providers
.iter()
.find(|p| p.provider_peer_id == peer)
.cloned();
let transport = self.transport.clone();
let req = self.range_request(range.offset, range.length);
Box::pin(async move {
let provider = match provider {
Some(p) => p,
None => {
return (
range_idx,
peer.clone(),
Err(DownloadError::transport(&peer, "provider vanished")),
)
}
};
let req = match req {
Ok(r) => r,
Err(e) => return (range_idx, peer, Err(e)),
};
let res = transport.fetch_range(&provider, &req).await;
(range_idx, peer, res)
})
}
async fn handle_result(
&mut self,
idx: usize,
peer: String,
res: Result<FetchedRange, DownloadError>,
) -> Result<(), DownloadError> {
if let Some(n) = self.inflight_per_source.get_mut(&peer) {
*n = n.saturating_sub(1);
}
let commitment = self.commitment.clone().expect("commitment established");
let range = self.ranges[idx];
let outcome = match res {
Ok(fetched) => self.verify_fetched(&commitment, &range, fetched),
Err(e) => Err(e),
};
match outcome {
Ok(bytes) => {
self.sink.write_at(range.offset, &bytes).await?;
if let Some(hasher) = self.hasher.as_mut() {
hasher.feed(range.offset, bytes);
}
self.range_state[idx] = RangeState::Done;
self.resume.mark_done(idx);
self.bytes_done = self.bytes_done.saturating_add(range.length);
self.tracker.record_success(&peer);
self.relocated_since_progress = false;
self.checkpoint().await?;
let progress = self.snapshot();
self.emit(DownloadEvent::RangeCompleted {
range: idx,
provider: peer,
progress,
})
.await;
}
Err(e) => {
if !e.is_recoverable() {
self.emit(DownloadEvent::Failed {
reason: e.to_string(),
})
.await;
return Err(e);
}
self.range_state[idx] = RangeState::Pending;
self.tracker.record_failure(&peer, Instant::now());
self.total_failures = self.total_failures.saturating_add(1);
self.emit(DownloadEvent::RangeFailed {
range: idx,
provider: peer,
reason: e.to_string(),
})
.await;
}
}
Ok(())
}
fn verify_fetched(
&self,
commitment: &ResourceCommitment,
range: &Range,
fetched: FetchedRange,
) -> Result<Vec<u8>, DownloadError> {
commitment.check_consistent(
fetched.meta.total_length,
fetched.meta.chunk_lens.as_deref(),
fetched.meta.root.as_deref(),
)?;
self.verifier.verify_range(
commitment,
range.chunk_start as u64,
range.length,
&fetched.bytes,
)?;
Ok(fetched.bytes)
}
async fn locate_and_confirm(&self) -> Result<Vec<ProviderRecord>, DownloadError> {
let found = self.locator.find_providers(&self.content).await?;
let item = self.availability_item()?;
let mut confirmed = Vec::new();
for p in found {
match self
.transport
.query_availability(&p, vec![item.clone()])
.await
{
Ok(resp) if resp.items.first().map(|a| a.available).unwrap_or(false) => {
confirmed.push(p)
}
_ => {}
}
}
Ok(confirmed)
}
async fn relocate(&mut self) -> Result<usize, DownloadError> {
self.relocate_attempts += 1;
let more = self.locate_and_confirm().await?;
let known: HashSet<String> = self
.providers
.iter()
.map(|p| p.provider_peer_id.clone())
.collect();
let mut added = 0;
for p in more {
if !known.contains(&p.provider_peer_id) {
self.providers.push(p);
added += 1;
}
}
if added > 0 {
self.emit(DownloadEvent::ProvidersRefreshed {
providers: self.providers.len(),
})
.await;
}
Ok(added)
}
async fn establish_commitment(&mut self) -> Result<(), DownloadError> {
let providers = self.providers.clone();
let want_root = self.content_root_hex();
for provider in &providers {
let req = self.range_request(0, 1)?;
if let Ok(f) = self.transport.fetch_range(provider, &req).await {
if let (Some(want), Some(got)) = (&want_root, &f.meta.root) {
if got != want {
continue;
}
}
if let (Some(tl), Some(cl)) = (f.meta.total_length, f.meta.chunk_lens.clone()) {
match ResourceCommitment::from_first_frame(
tl,
cl,
f.meta.root.clone(),
f.meta.inclusion_proof.clone(),
) {
Ok(c) => {
self.commitment = Some(c);
return Ok(());
}
Err(_) => continue,
}
}
}
}
let reason = format!("{:?}", self.content);
self.emit(DownloadEvent::Failed {
reason: format!("could not read resource metadata for {reason}"),
})
.await;
Err(DownloadError::NotFound { content: reason })
}
async fn persist_commitment(&mut self) -> Result<(), DownloadError> {
if let Some(c) = &self.commitment {
self.resume.total_length = c.total_length;
self.resume.chunk_lens = c.layout.chunk_lens().to_vec();
self.resume.root = c.root.clone();
self.resume.inclusion_proof = c.inclusion_proof.clone();
self.checkpoint().await?;
}
Ok(())
}
fn next_pending(&self) -> Option<usize> {
self.range_state
.iter()
.position(|s| matches!(s, RangeState::Pending))
}
fn pick_source(&self, now: Instant) -> Option<String> {
self.providers
.iter()
.map(|p| p.provider_peer_id.clone())
.filter(|peer| self.tracker.is_available(peer, now))
.filter(|peer| {
self.inflight_per_source.get(peer).copied().unwrap_or(0)
< self.config.max_inflight_per_source
})
.min_by_key(|peer| self.inflight_per_source.get(peer).copied().unwrap_or(0))
}
fn earliest_backoff(&self) -> Option<Instant> {
let now = Instant::now();
self.providers
.iter()
.filter_map(|p| {
if self.tracker.is_available(&p.provider_peer_id, now) {
None
} else {
self.next_available_at(&p.provider_peer_id, now)
}
})
.min()
}
fn next_available_at(&self, peer: &str, now: Instant) -> Option<Instant> {
let step = self.config.base_backoff.max(Duration::from_millis(1));
let mut t = now;
let limit = now + self.config.max_backoff + step;
while t <= limit {
if self.tracker.is_available(peer, t) {
return Some(t);
}
t += step;
}
Some(limit)
}
fn all_done(&self) -> bool {
!self.range_state.is_empty()
&& self
.range_state
.iter()
.all(|s| matches!(s, RangeState::Done))
}
fn pending_count(&self) -> usize {
self.range_state
.iter()
.filter(|s| s.is_incomplete())
.count()
}
fn snapshot(&self) -> DownloadProgress {
let ranges_done = self
.range_state
.iter()
.filter(|s| matches!(s, RangeState::Done))
.count();
let active_sources = self
.inflight_per_source
.values()
.filter(|&&n| n > 0)
.count();
DownloadProgress {
bytes_done: self.bytes_done,
total_length: self
.commitment
.as_ref()
.map(|c| c.total_length)
.unwrap_or(0),
ranges_done,
ranges_total: self.ranges.len(),
active_sources,
}
}
async fn checkpoint(&self) -> Result<(), DownloadError> {
self.state_store.save(&self.resume).await
}
async fn emit(&self, event: DownloadEvent) {
let _ = self.events.send(event).await;
}
fn availability_item(&self) -> Result<AvailabilityItem, DownloadError> {
match &self.content {
ContentId::Store { .. } => Err(DownloadError::NotDownloadable),
ContentId::Root { store_id, root } => Ok(AvailabilityItem {
store_id: hex32(store_id),
root: Some(hex32(root)),
retrieval_key: None,
}),
ContentId::Resource {
store_id,
root,
retrieval_key,
} => Ok(AvailabilityItem {
store_id: hex32(store_id),
root: Some(hex32(root)),
retrieval_key: Some(hex32(retrieval_key)),
}),
}
}
fn content_root_hex(&self) -> Option<String> {
match &self.content {
ContentId::Store { .. } => None,
ContentId::Root { root, .. } | ContentId::Resource { root, .. } => Some(hex32(root)),
}
}
fn range_request(&self, offset: u64, length: u64) -> Result<RangeRequest, DownloadError> {
match &self.content {
ContentId::Store { .. } => Err(DownloadError::NotDownloadable),
ContentId::Root { store_id, root } => Ok(RangeRequest {
store_id: hex32(store_id),
retrieval_key: None,
root: Some(hex32(root)),
capsule: true,
offset,
length,
}),
ContentId::Resource {
store_id,
root,
retrieval_key,
} => Ok(RangeRequest {
store_id: hex32(store_id),
retrieval_key: Some(hex32(retrieval_key)),
root: Some(hex32(root)),
capsule: false,
offset,
length,
}),
}
}
}
fn hex32(b: &[u8; 32]) -> String {
let mut s = String::with_capacity(64);
for x in b {
s.push(char::from_digit((x >> 4) as u32, 16).unwrap());
s.push(char::from_digit((x & 0x0f) as u32, 16).unwrap());
}
s
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn download_key_is_content_key_hex() {
let c = ContentId::resource([1; 32], [2; 32], [3; 32]);
assert_eq!(download_key(&c), c.to_key().to_hex());
assert_eq!(download_key(&c).len(), 64);
}
#[test]
fn hex32_round_trips_length() {
assert_eq!(hex32(&[0xAB; 32]), "ab".repeat(32));
}
}