use std::collections::BTreeSet;
use std::time::Duration;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use crate::error::DownloadError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct DownloadProgress {
pub bytes_done: u64,
pub total_length: u64,
pub ranges_done: usize,
pub ranges_total: usize,
pub active_sources: usize,
}
impl DownloadProgress {
pub fn fraction(&self) -> f64 {
if self.total_length == 0 {
0.0
} else {
self.bytes_done as f64 / self.total_length as f64
}
}
pub fn is_complete(&self) -> bool {
self.ranges_total > 0 && self.ranges_done == self.ranges_total
}
}
#[derive(Debug, Clone)]
pub enum DownloadEvent {
Planned {
ranges_total: usize,
total_length: u64,
},
RangeCompleted {
range: usize,
provider: String,
progress: DownloadProgress,
},
RangeFailed {
range: usize,
provider: String,
reason: String,
},
ProvidersRefreshed {
providers: usize,
},
Paused,
Resumed,
Completed {
total_length: u64,
},
Failed {
reason: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DownloadState {
pub key: String,
pub total_length: u64,
pub chunk_lens: Vec<u64>,
pub root: Option<String>,
pub inclusion_proof: Option<String>,
pub done_ranges: BTreeSet<usize>,
}
impl DownloadState {
pub fn new(key: impl Into<String>) -> Self {
DownloadState {
key: key.into(),
total_length: 0,
chunk_lens: Vec::new(),
root: None,
inclusion_proof: None,
done_ranges: BTreeSet::new(),
}
}
pub fn has_commitment(&self) -> bool {
!self.chunk_lens.is_empty()
}
pub fn mark_done(&mut self, index: usize) {
self.done_ranges.insert(index);
}
pub fn is_done(&self, index: usize) -> bool {
self.done_ranges.contains(&index)
}
}
pub const BAD_DESCRIPTOR_TTL: Duration = Duration::from_secs(24 * 60 * 60);
pub const MAX_BAD_DESCRIPTOR_PEERS: usize = 32;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct BadDescriptorVerdict {
pub peer_id: String,
pub recorded_at_unix: u64,
}
pub fn unix_now() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
fn is_hex_peer_id(peer_id: &str) -> bool {
peer_id.len() == 64 && peer_id.bytes().all(|b| b.is_ascii_hexdigit())
}
fn record_verdict(verdicts: &mut Vec<BadDescriptorVerdict>, peer_id: &str, now: u64) {
if !is_hex_peer_id(peer_id) {
return;
}
verdicts.retain(|v| !is_expired(v, now) && v.peer_id != peer_id);
verdicts.push(BadDescriptorVerdict {
peer_id: peer_id.to_string(),
recorded_at_unix: now,
});
while verdicts.len() > MAX_BAD_DESCRIPTOR_PEERS {
verdicts.remove(0); }
}
fn live_peers(verdicts: &[BadDescriptorVerdict], now: u64) -> Vec<String> {
verdicts
.iter()
.filter(|v| !is_expired(v, now))
.map(|v| v.peer_id.clone())
.collect()
}
fn is_expired(verdict: &BadDescriptorVerdict, now: u64) -> bool {
now.saturating_sub(verdict.recorded_at_unix) > BAD_DESCRIPTOR_TTL.as_secs()
}
#[async_trait]
pub trait StateStore: Send + Sync {
async fn load(&self, key: &str) -> Result<Option<DownloadState>, DownloadError>;
async fn save(&self, state: &DownloadState) -> Result<(), DownloadError>;
async fn clear(&self, key: &str) -> Result<(), DownloadError>;
async fn record_bad_descriptor(
&self,
target_key: &str,
peer_id: &str,
) -> Result<(), DownloadError> {
let _ = (target_key, peer_id);
Ok(())
}
async fn bad_descriptor_peers(&self, target_key: &str) -> Result<Vec<String>, DownloadError> {
let _ = target_key;
Ok(Vec::new())
}
}
#[derive(Debug, Default)]
pub struct InMemoryStateStore {
inner: tokio::sync::Mutex<std::collections::HashMap<String, DownloadState>>,
reputation: tokio::sync::Mutex<std::collections::HashMap<String, Vec<BadDescriptorVerdict>>>,
}
impl InMemoryStateStore {
pub fn new() -> Self {
InMemoryStateStore::default()
}
}
#[async_trait]
impl StateStore for InMemoryStateStore {
async fn load(&self, key: &str) -> Result<Option<DownloadState>, DownloadError> {
Ok(self.inner.lock().await.get(key).cloned())
}
async fn save(&self, state: &DownloadState) -> Result<(), DownloadError> {
self.inner
.lock()
.await
.insert(state.key.clone(), state.clone());
Ok(())
}
async fn clear(&self, key: &str) -> Result<(), DownloadError> {
self.inner.lock().await.remove(key);
Ok(())
}
async fn record_bad_descriptor(
&self,
target_key: &str,
peer_id: &str,
) -> Result<(), DownloadError> {
let mut reputation = self.reputation.lock().await;
let verdicts = reputation.entry(target_key.to_string()).or_default();
record_verdict(verdicts, peer_id, unix_now());
Ok(())
}
async fn bad_descriptor_peers(&self, target_key: &str) -> Result<Vec<String>, DownloadError> {
let reputation = self.reputation.lock().await;
Ok(reputation
.get(target_key)
.map(|v| live_peers(v, unix_now()))
.unwrap_or_default())
}
}
fn checkpoint_file_stem(key: &str) -> String {
use sha2::Digest;
crate::module::hex_of(sha2::Sha256::digest(key.as_bytes()))
}
const CHECKPOINT_STEM_LEN: usize = 64;
#[derive(Debug, Clone)]
pub struct FileStateStore {
dir: std::path::PathBuf,
}
impl FileStateStore {
pub fn new(dir: impl Into<std::path::PathBuf>) -> Self {
FileStateStore { dir: dir.into() }
}
fn path_for(&self, key: &str) -> std::path::PathBuf {
self.file_for(key, ".json")
}
fn reputation_path_for(&self, key: &str) -> std::path::PathBuf {
self.file_for(key, ".holders.json")
}
fn file_for(&self, key: &str, suffix: &str) -> std::path::PathBuf {
let mut name = checkpoint_file_stem(key);
debug_assert_eq!(name.len(), CHECKPOINT_STEM_LEN);
name.reserve_exact(suffix.len());
name.push_str(suffix);
self.dir.join(name)
}
fn read_verdicts(&self, key: &str) -> Vec<BadDescriptorVerdict> {
std::fs::read(self.reputation_path_for(key))
.ok()
.and_then(|bytes| serde_json::from_slice(&bytes).ok())
.unwrap_or_default()
}
}
#[async_trait]
impl StateStore for FileStateStore {
async fn load(&self, key: &str) -> Result<Option<DownloadState>, DownloadError> {
let path = self.path_for(key);
match std::fs::read(&path) {
Ok(bytes) => {
let state = serde_json::from_slice(&bytes).map_err(DownloadError::state)?;
Ok(Some(state))
}
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None),
Err(e) => Err(DownloadError::state(e)),
}
}
async fn save(&self, state: &DownloadState) -> Result<(), DownloadError> {
std::fs::create_dir_all(&self.dir).map_err(DownloadError::state)?;
let bytes = serde_json::to_vec(state).map_err(DownloadError::state)?;
std::fs::write(self.path_for(&state.key), bytes).map_err(DownloadError::state)
}
async fn clear(&self, key: &str) -> Result<(), DownloadError> {
match std::fs::remove_file(self.path_for(key)) {
Ok(()) => Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(e) => Err(DownloadError::state(e)),
}
}
async fn record_bad_descriptor(
&self,
target_key: &str,
peer_id: &str,
) -> Result<(), DownloadError> {
let mut verdicts = self.read_verdicts(target_key);
record_verdict(&mut verdicts, peer_id, unix_now());
std::fs::create_dir_all(&self.dir).map_err(DownloadError::state)?;
let bytes = serde_json::to_vec(&verdicts).map_err(DownloadError::state)?;
std::fs::write(self.reputation_path_for(target_key), bytes).map_err(DownloadError::state)
}
async fn bad_descriptor_peers(&self, target_key: &str) -> Result<Vec<String>, DownloadError> {
Ok(live_peers(&self.read_verdicts(target_key), unix_now()))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn progress_fraction_and_complete() {
let mut p = DownloadProgress {
total_length: 100,
bytes_done: 25,
ranges_total: 4,
ranges_done: 1,
active_sources: 2,
};
assert!((p.fraction() - 0.25).abs() < 1e-9);
assert!(!p.is_complete());
p.ranges_done = 4;
p.bytes_done = 100;
assert!(p.is_complete());
assert_eq!(DownloadProgress::default().fraction(), 0.0);
}
#[test]
fn state_marks_and_queries_done() {
let mut s = DownloadState::new("k");
assert!(!s.is_done(2));
s.mark_done(2);
assert!(s.is_done(2));
assert_eq!(s.done_ranges.len(), 1);
}
#[tokio::test]
async fn in_memory_store_round_trips() {
let store = InMemoryStateStore::new();
assert!(store.load("k").await.unwrap().is_none());
let mut s = DownloadState::new("k");
s.mark_done(1);
s.total_length = 42;
store.save(&s).await.unwrap();
assert_eq!(store.load("k").await.unwrap().unwrap(), s);
store.clear("k").await.unwrap();
assert!(store.load("k").await.unwrap().is_none());
}
#[tokio::test]
async fn file_store_round_trips_and_survives_reload() {
let dir = std::env::temp_dir().join(format!(
"dig-download-test-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let store = FileStateStore::new(&dir);
assert!(store.load("abc").await.unwrap().is_none());
let mut s = DownloadState::new("abc");
s.total_length = 100;
s.chunk_lens = vec![10, 20];
s.root = Some("aa".repeat(32));
s.mark_done(0);
store.save(&s).await.unwrap();
let reloaded = FileStateStore::new(&dir);
assert_eq!(reloaded.load("abc").await.unwrap().unwrap(), s);
store.clear("abc").await.unwrap();
assert!(store.load("abc").await.unwrap().is_none());
store.clear("abc").await.unwrap();
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn digest_name_is_bounded() {
const NAME_MAX: usize = 255;
let long_key = crate::module::module_download_key(&"ab".repeat(32), &"cd".repeat(32));
assert_eq!(
long_key.len(),
136,
"the production key really is this long"
);
let store = FileStateStore::new("dir");
for produced in [
store.path_for(&long_key),
store.reputation_path_for(&long_key),
] {
let name = produced.file_name().unwrap().to_string_lossy().into_owned();
assert!(
name.chars().count() < NAME_MAX,
"{} chars exceeds NAME_MAX: {name}",
name.chars().count()
);
}
assert_eq!(
store
.reputation_path_for(&long_key)
.file_name()
.unwrap()
.to_string_lossy()
.chars()
.count(),
CHECKPOINT_STEM_LEN + ".holders.json".len(),
);
assert_eq!(checkpoint_file_stem(&long_key).len(), CHECKPOINT_STEM_LEN);
assert_eq!(checkpoint_file_stem("abc").len(), CHECKPOINT_STEM_LEN);
assert_ne!(
checkpoint_file_stem(&long_key),
checkpoint_file_stem(&crate::module::module_download_key(
&"ab".repeat(32),
&"ce".repeat(32)
))
);
}
#[tokio::test]
async fn file_store_round_trips_a_real_module_download_key() {
let dir = std::env::temp_dir().join(format!(
"dig-download-modkey-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let key = crate::module::module_download_key(&"ab".repeat(32), &"cd".repeat(32));
let store = FileStateStore::new(&dir);
let mut s = DownloadState::new(&key);
s.total_length = 100;
s.chunk_lens = vec![10, 20];
s.mark_done(0);
store.save(&s).await.unwrap();
assert_eq!(
FileStateStore::new(&dir).load(&key).await.unwrap().unwrap(),
s
);
let peer = "ef".repeat(32);
store.record_bad_descriptor(&key, &peer).await.unwrap();
assert_eq!(
FileStateStore::new(&dir)
.bad_descriptor_peers(&key)
.await
.unwrap(),
vec![peer]
);
for entry in std::fs::read_dir(&dir).unwrap() {
let name = entry.unwrap().file_name();
assert!(
name.to_string_lossy().chars().count() < 255,
"checkpoint filename must fit NAME_MAX: {name:?}"
);
}
store.clear(&key).await.unwrap();
assert!(store.load(&key).await.unwrap().is_none());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn download_event_variants_construct() {
let _ = DownloadEvent::Planned {
ranges_total: 3,
total_length: 30,
};
let _ = DownloadEvent::Paused;
let _ = DownloadEvent::Resumed;
let _ = DownloadEvent::ProvidersRefreshed { providers: 2 };
let _ = DownloadEvent::Completed { total_length: 30 };
let _ = DownloadEvent::Failed { reason: "x".into() };
}
#[tokio::test]
async fn in_memory_store_remembers_a_bad_descriptor_verdict() {
let store = InMemoryStateStore::new();
let peer = "ab".repeat(32);
assert!(store.bad_descriptor_peers("k").await.unwrap().is_empty());
store.record_bad_descriptor("k", &peer).await.unwrap();
assert_eq!(store.bad_descriptor_peers("k").await.unwrap(), vec![peer]);
assert!(store
.bad_descriptor_peers("other")
.await
.unwrap()
.is_empty());
}
#[tokio::test]
async fn file_store_reputation_survives_a_process_restart_and_outlives_the_checkpoint() {
let dir = std::env::temp_dir().join(format!(
"dig-download-rep-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let store = FileStateStore::new(&dir);
let peer = "cd".repeat(32);
store
.record_bad_descriptor("module:x", &peer)
.await
.unwrap();
let restarted = FileStateStore::new(&dir);
assert_eq!(
restarted.bad_descriptor_peers("module:x").await.unwrap(),
vec![peer.clone()]
);
restarted.clear("module:x").await.unwrap();
assert_eq!(
restarted.bad_descriptor_peers("module:x").await.unwrap(),
vec![peer]
);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn a_verdict_expires_after_its_ttl() {
let now = 10_000_000u64;
let fresh = BadDescriptorVerdict {
peer_id: "ab".repeat(32),
recorded_at_unix: now - 60,
};
let stale = BadDescriptorVerdict {
peer_id: "cd".repeat(32),
recorded_at_unix: now - BAD_DESCRIPTOR_TTL.as_secs() - 1,
};
assert_eq!(
live_peers(&[fresh.clone(), stale], now),
vec![fresh.peer_id],
"only the un-expired verdict is live"
);
}
#[test]
fn the_verdict_record_is_bounded_and_deduplicated() {
let now = 10_000_000u64;
let mut verdicts = Vec::new();
for i in 0..(MAX_BAD_DESCRIPTOR_PEERS + 10) {
record_verdict(&mut verdicts, &format!("{i:064x}"), now);
}
assert_eq!(verdicts.len(), MAX_BAD_DESCRIPTOR_PEERS, "capped");
assert!(
!verdicts.iter().any(|v| v.peer_id == format!("{:064x}", 0)),
"the oldest verdicts were evicted first"
);
let repeat = format!("{:064x}", MAX_BAD_DESCRIPTOR_PEERS + 9);
record_verdict(&mut verdicts, &repeat, now + 5);
assert_eq!(
verdicts.iter().filter(|v| v.peer_id == repeat).count(),
1,
"a repeat verdict refreshes the entry instead of duplicating it"
);
}
#[test]
fn a_malformed_peer_id_is_never_recorded() {
let mut verdicts = Vec::new();
record_verdict(&mut verdicts, "../../etc/passwd", 1);
record_verdict(&mut verdicts, "not-hex", 1);
record_verdict(&mut verdicts, &"ab".repeat(31), 1); assert!(
verdicts.is_empty(),
"only 64-hex ids are stored: {verdicts:?}"
);
}
}