use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use async_trait::async_trait;
use dig_dht::{CandidateAddr, ContentId, PeerId, ProviderRecord};
use dig_nat::{AvailabilityAnswer, AvailabilityItem, AvailabilityResponse};
use tokio::sync::Mutex;
use crate::error::DownloadError;
use crate::locate::ProviderLocator;
use crate::select::{RangeOutcome, SelectPlan, SelectRequest, SourceSelector};
use crate::source::{FetchedRange, RangeMeta, RangeTransport};
#[derive(Debug, Clone)]
pub struct MockContent {
pub bytes: Vec<u8>,
pub chunk_lens: Vec<u64>,
pub root: String,
pub inclusion_proof: Option<String>,
offsets: Vec<u64>,
}
impl MockContent {
pub fn new(bytes: Vec<u8>, chunk_lens: Vec<u64>) -> Self {
assert_eq!(
bytes.len() as u64,
chunk_lens.iter().sum::<u64>(),
"chunk_lens must sum to bytes.len()"
);
let mut offsets = Vec::with_capacity(chunk_lens.len() + 1);
let mut acc = 0u64;
offsets.push(0);
for &l in &chunk_lens {
acc += l;
offsets.push(acc);
}
MockContent {
bytes,
chunk_lens,
root: "ab".repeat(32),
inclusion_proof: Some("mock-proof".into()),
offsets,
}
}
pub fn even(n: usize, chunks: usize) -> Self {
let chunks = chunks.max(1);
let base = n / chunks;
let mut lens = vec![base as u64; chunks];
let assigned: u64 = lens.iter().sum();
if let Some(last) = lens.last_mut() {
*last += n as u64 - assigned;
}
let bytes: Vec<u8> = (0..n).map(|i| (i % 251) as u8).collect();
MockContent::new(bytes, lens)
}
fn chunk_index_at(&self, offset: u64) -> u64 {
self.offsets.iter().position(|&o| o == offset).unwrap_or(0) as u64
}
fn meta(&self, offset: u64) -> RangeMeta {
RangeMeta {
total_length: Some(self.bytes.len() as u64),
chunk_lens: Some(self.chunk_lens.clone()),
chunk_index: Some(self.chunk_index_at(offset)),
root: Some(self.root.clone()),
inclusion_proof: self.inclusion_proof.clone(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Behavior {
Honest,
Corrupt,
Truncate,
ShortAligned,
Unavailable,
DropAfter(usize),
AlwaysFail,
WrongRoot,
}
pub struct MockRangeTransport {
content: MockContent,
behaviors: Mutex<HashMap<String, Behavior>>,
provider_attempts: Mutex<HashMap<String, usize>>,
offset_attempts: Mutex<HashMap<u64, usize>>,
delay: Mutex<Option<std::time::Duration>>,
}
impl MockRangeTransport {
pub fn new(content: MockContent) -> Self {
MockRangeTransport {
content,
behaviors: Mutex::new(HashMap::new()),
provider_attempts: Mutex::new(HashMap::new()),
offset_attempts: Mutex::new(HashMap::new()),
delay: Mutex::new(None),
}
}
pub async fn set_delay(&self, delay: std::time::Duration) {
*self.delay.lock().await = Some(delay);
}
pub async fn set_behavior(&self, peer_id: &str, behavior: Behavior) {
self.behaviors
.lock()
.await
.insert(peer_id.to_string(), behavior);
}
pub async fn attempts_for(&self, peer_id: &str) -> usize {
self.provider_attempts
.lock()
.await
.get(peer_id)
.copied()
.unwrap_or(0)
}
pub async fn attempts_at(&self, offset: u64) -> usize {
self.offset_attempts
.lock()
.await
.get(&offset)
.copied()
.unwrap_or(0)
}
async fn behavior(&self, peer_id: &str) -> Behavior {
self.behaviors
.lock()
.await
.get(peer_id)
.cloned()
.unwrap_or(Behavior::Honest)
}
}
#[async_trait]
impl RangeTransport for MockRangeTransport {
async fn query_availability(
&self,
provider: &ProviderRecord,
items: Vec<AvailabilityItem>,
) -> Result<AvailabilityResponse, DownloadError> {
let behavior = self.behavior(&provider.provider_peer_id).await;
let held = !matches!(behavior, Behavior::Unavailable | Behavior::AlwaysFail);
let answers = items
.iter()
.map(|_| AvailabilityAnswer {
available: held,
roots: Some(vec![self.content.root.clone()]),
total_length: Some(self.content.bytes.len() as u64),
chunk_count: Some(self.content.chunk_lens.len() as u64),
complete: Some(true),
})
.collect();
Ok(AvailabilityResponse { items: answers })
}
async fn fetch_range(
&self,
provider: &ProviderRecord,
req: &dig_nat::RangeRequest,
) -> Result<FetchedRange, DownloadError> {
let peer = provider.provider_peer_id.clone();
let attempts = {
let mut a = self.provider_attempts.lock().await;
let n = a.entry(peer.clone()).or_insert(0);
*n += 1;
*n
};
*self
.offset_attempts
.lock()
.await
.entry(req.offset)
.or_insert(0) += 1;
if let Some(d) = *self.delay.lock().await {
tokio::time::sleep(d).await;
}
let behavior = self.behavior(&peer).await;
let fail = || DownloadError::transport(&peer, "mock: source failed");
match behavior {
Behavior::Unavailable | Behavior::AlwaysFail => return Err(fail()),
Behavior::DropAfter(n) if attempts > n => return Err(fail()),
_ => {}
}
let start = req.offset as usize;
let end = (req.offset + req.length).min(self.content.bytes.len() as u64) as usize;
let mut bytes = self.content.bytes[start..end].to_vec();
match behavior {
Behavior::Truncate => {
bytes.pop(); }
Behavior::Corrupt => {
for b in bytes.iter_mut() {
*b ^= 0xFF; }
}
Behavior::ShortAligned => {
let first_chunk_idx = self.content.chunk_index_at(req.offset) as usize;
if let Some(&first_len) = self.content.chunk_lens.get(first_chunk_idx) {
let keep = (first_len as usize).min(bytes.len());
if keep < bytes.len() {
bytes.truncate(keep);
}
}
}
_ => {}
}
let mut meta = self.content.meta(req.offset);
if matches!(behavior, Behavior::WrongRoot) {
meta.root = Some("cd".repeat(32));
}
Ok(FetchedRange {
request_offset: req.offset,
bytes,
meta,
})
}
}
pub struct MockProviderLocator {
batches: Vec<Vec<ProviderRecord>>,
calls: Mutex<usize>,
}
impl MockProviderLocator {
pub fn fixed(providers: Vec<ProviderRecord>) -> Self {
MockProviderLocator {
batches: vec![providers],
calls: Mutex::new(0),
}
}
pub fn scripted(batches: Vec<Vec<ProviderRecord>>) -> Self {
MockProviderLocator {
batches: if batches.is_empty() {
vec![vec![]]
} else {
batches
},
calls: Mutex::new(0),
}
}
pub async fn call_count(&self) -> usize {
*self.calls.lock().await
}
}
#[async_trait]
impl ProviderLocator for MockProviderLocator {
async fn find_providers(
&self,
_content: &ContentId,
) -> Result<Vec<ProviderRecord>, DownloadError> {
let mut calls = self.calls.lock().await;
let idx = (*calls).min(self.batches.len() - 1);
*calls += 1;
Ok(self.batches[idx].clone())
}
}
#[derive(Default)]
pub struct MockSelector {
select_calls: AtomicUsize,
recorded: std::sync::Mutex<Vec<RangeOutcome>>,
forced_order: std::sync::Mutex<Option<Vec<String>>>,
}
impl MockSelector {
pub fn new() -> Arc<Self> {
Arc::new(MockSelector::default())
}
pub fn with_order(order: Vec<String>) -> Arc<Self> {
let sel = MockSelector::default();
*sel.forced_order.lock().unwrap() = Some(order);
Arc::new(sel)
}
pub fn select_call_count(&self) -> usize {
self.select_calls.load(Ordering::Relaxed)
}
pub fn outcomes(&self) -> Vec<RangeOutcome> {
self.recorded.lock().unwrap().clone()
}
}
impl SourceSelector for MockSelector {
fn select(&self, req: &SelectRequest) -> SelectPlan {
self.select_calls.fetch_add(1, Ordering::Relaxed);
let offered: Vec<String> = req.candidates.iter().map(|c| c.peer_id.clone()).collect();
match self.forced_order.lock().unwrap().as_ref() {
Some(order) => {
let mut plan: Vec<String> = order
.iter()
.filter(|p| offered.contains(p))
.cloned()
.collect();
for p in &offered {
if !plan.contains(p) {
plan.push(p.clone());
}
}
SelectPlan::ordered(plan)
}
None => SelectPlan::ordered(offered),
}
}
fn record(&self, outcome: &RangeOutcome) {
self.recorded.lock().unwrap().push(outcome.clone());
}
}
pub fn mock_provider(n: u8, content: &ContentId) -> ProviderRecord {
ProviderRecord::new(
&content.to_key(),
&PeerId::from_bytes([n; 32]),
vec![CandidateAddr::direct(format!("10.0.0.{n}"), 9444)],
u64::MAX,
)
}
pub fn mock_peer_hex(n: u8) -> String {
PeerId::from_bytes([n; 32]).to_hex()
}
pub fn mock_content_id() -> ContentId {
ContentId::resource([1; 32], [0xAB; 32], [3; 32])
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn honest_transport_serves_correct_slice() {
let content = MockContent::even(30, 3);
let t = MockRangeTransport::new(content.clone());
let cid = mock_content_id();
let p = mock_provider(1, &cid);
let req = dig_nat::RangeRequest::resource("s", "r", 10, 10);
let got = t.fetch_range(&p, &req).await.unwrap();
assert_eq!(got.bytes, content.bytes[10..20]);
assert_eq!(got.meta.chunk_lens, Some(content.chunk_lens.clone()));
assert_eq!(t.attempts_at(10).await, 1);
assert_eq!(t.attempts_for(&mock_peer_hex(1)).await, 1);
}
#[tokio::test]
async fn behaviors_corrupt_and_truncate() {
let content = MockContent::even(30, 3);
let cid = mock_content_id();
let p = mock_provider(2, &cid);
let hex = mock_peer_hex(2);
let t = MockRangeTransport::new(content.clone());
t.set_behavior(&hex, Behavior::Truncate).await;
let req = dig_nat::RangeRequest::resource("s", "r", 0, 10);
let got = t.fetch_range(&p, &req).await.unwrap();
assert_eq!(got.bytes.len(), 9);
let t2 = MockRangeTransport::new(content.clone());
t2.set_behavior(&hex, Behavior::Corrupt).await;
let got2 = t2.fetch_range(&p, &req).await.unwrap();
assert_eq!(got2.bytes.len(), 10);
assert_ne!(got2.bytes, content.bytes[0..10]);
}
#[tokio::test]
async fn drop_after_fails_late() {
let content = MockContent::even(30, 3);
let cid = mock_content_id();
let p = mock_provider(3, &cid);
let hex = mock_peer_hex(3);
let t = MockRangeTransport::new(content);
t.set_behavior(&hex, Behavior::DropAfter(1)).await;
let req = dig_nat::RangeRequest::resource("s", "r", 0, 10);
assert!(t.fetch_range(&p, &req).await.is_ok()); assert!(t.fetch_range(&p, &req).await.is_err()); }
#[tokio::test]
async fn unavailable_source_reports_not_held() {
let content = MockContent::even(30, 3);
let cid = mock_content_id();
let p = mock_provider(4, &cid);
let hex = mock_peer_hex(4);
let t = MockRangeTransport::new(content);
t.set_behavior(&hex, Behavior::Unavailable).await;
let resp = t
.query_availability(
&p,
vec![AvailabilityItem {
store_id: "s".into(),
root: None,
retrieval_key: None,
}],
)
.await
.unwrap();
assert!(!resp.items[0].available);
}
#[tokio::test]
async fn scripted_locator_advances_batches() {
let cid = mock_content_id();
let loc = MockProviderLocator::scripted(vec![
vec![mock_provider(1, &cid)],
vec![mock_provider(1, &cid), mock_provider(2, &cid)],
]);
assert_eq!(loc.find_providers(&cid).await.unwrap().len(), 1);
assert_eq!(loc.find_providers(&cid).await.unwrap().len(), 2);
assert_eq!(loc.find_providers(&cid).await.unwrap().len(), 2); assert_eq!(loc.call_count().await, 3);
}
}