use std::collections::BTreeSet;
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)
}
}
#[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>;
}
#[derive(Debug, Default)]
pub struct InMemoryStateStore {
inner: tokio::sync::Mutex<std::collections::HashMap<String, DownloadState>>,
}
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(())
}
}
#[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 {
let mut name = String::with_capacity(key.len() * 2 + 5);
for b in key.as_bytes() {
name.push(char::from_digit((b >> 4) as u32, 16).unwrap());
name.push(char::from_digit((b & 0x0f) as u32, 16).unwrap());
}
name.push_str(".json");
self.dir.join(name)
}
}
#[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)),
}
}
}
#[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 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() };
}
}