use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use async_trait::async_trait;
use dig_dht::ProviderRecord;
use dig_nat::{AvailabilityItem, AvailabilityResponse, RangeFrame, RangeRequest};
use dig_peer::DigPeer;
use tokio::io::{AsyncRead, AsyncReadExt};
use crate::error::DownloadError;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RangeMeta {
pub total_length: Option<u64>,
pub chunk_lens: Option<Vec<u64>>,
pub chunk_index: Option<u64>,
pub root: Option<String>,
pub inclusion_proof: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FetchedRange {
pub request_offset: u64,
pub bytes: Vec<u8>,
pub meta: RangeMeta,
}
#[async_trait]
pub trait RangeTransport: Send + Sync {
async fn query_availability(
&self,
provider: &ProviderRecord,
items: Vec<AvailabilityItem>,
) -> Result<AvailabilityResponse, DownloadError>;
async fn fetch_range(
&self,
provider: &ProviderRecord,
req: &RangeRequest,
) -> Result<FetchedRange, DownloadError>;
}
#[derive(Debug, Clone, Default)]
pub struct SourceHealth {
pub failures: u32,
pub served: u64,
pub backoff_until: Option<Instant>,
}
#[derive(Debug, Default)]
pub struct SourceTracker {
health: HashMap<String, SourceHealth>,
base_backoff: Duration,
max_backoff: Duration,
}
impl SourceTracker {
pub fn new(base_backoff: Duration, max_backoff: Duration) -> Self {
SourceTracker {
health: HashMap::new(),
base_backoff,
max_backoff,
}
}
pub fn is_available(&self, peer_id: &str, now: Instant) -> bool {
match self.health.get(peer_id) {
Some(h) => match h.backoff_until {
Some(until) => now >= until,
None => true,
},
None => true,
}
}
pub fn record_success(&mut self, peer_id: &str) {
let h = self.health.entry(peer_id.to_string()).or_default();
h.failures = 0;
h.served += 1;
h.backoff_until = None;
}
pub fn record_failure(&mut self, peer_id: &str, now: Instant) {
let base = self.base_backoff;
let max = self.max_backoff;
let h = self.health.entry(peer_id.to_string()).or_default();
h.failures = h.failures.saturating_add(1);
let shift = h.failures.saturating_sub(1).min(16);
let backoff = base.checked_mul(1u32 << shift).unwrap_or(max).min(max);
h.backoff_until = Some(now + backoff);
}
pub fn served(&self, peer_id: &str) -> u64 {
self.health.get(peer_id).map(|h| h.served).unwrap_or(0)
}
pub fn failures(&self, peer_id: &str) -> u32 {
self.health.get(peer_id).map(|h| h.failures).unwrap_or(0)
}
}
pub async fn assemble_range_stream<R: AsyncRead + Unpin>(
reader: &mut R,
max_len: u64,
) -> Result<(Vec<u8>, RangeMeta), DownloadError> {
let mut buf: Vec<u8> = Vec::new();
let mut meta = RangeMeta::default();
let mut first = true;
loop {
let frame = RangeFrame::decode(reader)
.await
.map_err(|e| DownloadError::Transport {
provider: String::new(),
reason: format!("range frame decode: {e}"),
})?;
let Some(frame) = frame else {
break; };
if first {
meta = RangeMeta {
total_length: frame.total_length,
chunk_lens: frame.chunk_lens.clone(),
chunk_index: frame.chunk_index,
root: frame.root.clone(),
inclusion_proof: frame.inclusion_proof.clone(),
};
first = false;
}
let start = frame.offset as usize;
let end = start + frame.bytes.len();
if end as u64 > max_len {
return Err(DownloadError::Transport {
provider: String::new(),
reason: format!("range frame overflows expected length {max_len}"),
});
}
if buf.len() < end {
buf.resize(end, 0);
}
buf[start..end].copy_from_slice(&frame.bytes);
if frame.complete {
break;
}
}
Ok((buf, meta))
}
const MAX_TRAILER_DRAIN: u64 = 64 * 1024;
pub async fn drain_trailer_bounded<R: AsyncRead + Unpin>(reader: &mut R, cap: u64) -> u64 {
let mut scratch = [0u8; 4096];
let mut drained: u64 = 0;
while drained < cap {
let want = ((cap - drained) as usize).min(scratch.len());
match reader.read(&mut scratch[..want]).await {
Ok(0) => break, Ok(n) => drained += n as u64,
Err(_) => break, }
}
drained
}
type PooledConn = Arc<tokio::sync::Mutex<DigPeer>>;
pub struct NatRangeTransport {
node: std::sync::Arc<dig_nat::NodeCert>,
config: dig_nat::NatConfig,
network_id: String,
runtime: Arc<dig_nat::NatRuntime>,
pool: tokio::sync::Mutex<HashMap<String, PooledConn>>,
}
impl NatRangeTransport {
pub fn new(
node: std::sync::Arc<dig_nat::NodeCert>,
config: dig_nat::NatConfig,
network_id: impl Into<String>,
) -> Self {
Self::new_with_runtime(
node,
config,
network_id,
Arc::new(dig_nat::NatRuntime::default()),
)
}
pub fn new_with_runtime(
node: std::sync::Arc<dig_nat::NodeCert>,
config: dig_nat::NatConfig,
network_id: impl Into<String>,
runtime: Arc<dig_nat::NatRuntime>,
) -> Self {
NatRangeTransport {
node,
config,
network_id: network_id.into(),
runtime,
pool: tokio::sync::Mutex::new(HashMap::new()),
}
}
pub fn provider_to_target(
&self,
provider: &ProviderRecord,
) -> Result<dig_nat::PeerTarget, DownloadError> {
let peer_id = provider.provider_peer_id().ok_or_else(|| {
DownloadError::transport(&provider.provider_peer_id, "malformed provider peer_id")
})?;
match provider.best_address() {
Some(addr) => {
let socket = format!("{}:{}", addr.host, addr.port)
.parse::<std::net::SocketAddr>()
.map_err(|e| {
DownloadError::transport(&provider.provider_peer_id, format!("addr: {e}"))
})?;
Ok(dig_nat::PeerTarget::with_addr(
peer_id,
socket,
self.network_id.clone(),
))
}
None => Ok(dig_nat::PeerTarget::relay_only(
peer_id,
self.network_id.clone(),
)),
}
}
async fn connect(&self, provider: &ProviderRecord) -> Result<DigPeer, DownloadError> {
let target = self.provider_to_target(provider)?;
DigPeer::connect_with_runtime(&target, &self.node, &self.config, &self.runtime)
.await
.map_err(|e| DownloadError::transport(&provider.provider_peer_id, e))
}
async fn pooled_conn(&self, provider: &ProviderRecord) -> Result<PooledConn, DownloadError> {
let key = provider.provider_peer_id.clone();
if let Some(conn) = self.pool.lock().await.get(&key).cloned() {
return Ok(conn);
}
let fresh = Arc::new(tokio::sync::Mutex::new(self.connect(provider).await?));
let mut pool = self.pool.lock().await;
Ok(pool.entry(key).or_insert(fresh).clone())
}
async fn evict(&self, provider: &ProviderRecord) {
self.pool.lock().await.remove(&provider.provider_peer_id);
}
}
#[async_trait]
impl RangeTransport for NatRangeTransport {
async fn query_availability(
&self,
provider: &ProviderRecord,
items: Vec<AvailabilityItem>,
) -> Result<AvailabilityResponse, DownloadError> {
let conn = self.pooled_conn(provider).await?;
let res = {
let mut guard = conn.lock().await;
guard.get_availability(items).await
};
match res {
Ok(resp) => Ok(resp),
Err(e) => {
self.evict(provider).await;
Err(DownloadError::transport(&provider.provider_peer_id, e))
}
}
}
async fn fetch_range(
&self,
provider: &ProviderRecord,
req: &RangeRequest,
) -> Result<FetchedRange, DownloadError> {
let conn = self.pooled_conn(provider).await?;
let stream = {
let mut guard = conn.lock().await;
guard.fetch_range(req).await
};
let mut stream = match stream {
Ok(s) => s,
Err(e) => {
self.evict(provider).await;
return Err(DownloadError::transport(&provider.provider_peer_id, e));
}
};
let (bytes, meta) = assemble_range_stream(&mut stream, req.length)
.await
.map_err(|e| {
DownloadError::transport(&provider.provider_peer_id, e)
})?;
let _ = drain_trailer_bounded(&mut stream, MAX_TRAILER_DRAIN).await;
Ok(FetchedRange {
request_offset: req.offset,
bytes,
meta,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use dig_dht::{CandidateAddr, ProviderRecord};
use dig_nat::PeerId;
fn provider(peer: u8, host: &str, port: u16) -> ProviderRecord {
ProviderRecord::new(
&dig_dht::Key::from_bytes([0xAB; 32]),
&PeerId::from_bytes([peer; 32]),
vec![CandidateAddr::direct(host, port)],
u64::MAX,
)
}
#[test]
fn provider_to_target_uses_direct_address() {
let t = NatRangeTransport::new(
fake_node_cert(),
dig_nat::NatConfig::default(),
"DIG_MAINNET",
);
let p = provider(1, "203.0.113.7", 9444);
let target = t.provider_to_target(&p).unwrap();
assert_eq!(
target.direct_addr().unwrap().to_string(),
"203.0.113.7:9444"
);
assert_eq!(target.network_id, "DIG_MAINNET");
}
#[test]
fn new_with_runtime_builds_a_full_ladder_transport() {
let runtime = std::sync::Arc::new(dig_nat::NatRuntime::default());
let t = NatRangeTransport::new_with_runtime(
fake_node_cert(),
dig_nat::NatConfig::default(),
"DIG_MAINNET",
runtime,
);
let p = provider(1, "203.0.113.7", 9444);
let target = t.provider_to_target(&p).unwrap();
assert_eq!(
target.direct_addr().unwrap().to_string(),
"203.0.113.7:9444"
);
assert_eq!(target.network_id, "DIG_MAINNET");
}
#[test]
fn provider_to_target_relay_only_without_address() {
let t = NatRangeTransport::new(
fake_node_cert(),
dig_nat::NatConfig::default(),
"DIG_MAINNET",
);
let p = ProviderRecord::new(
&dig_dht::Key::from_bytes([0xAB; 32]),
&PeerId::from_bytes([2; 32]),
vec![CandidateAddr::relay_marker()],
u64::MAX,
);
let target = t.provider_to_target(&p).unwrap();
assert!(target.direct_addr().is_none());
}
#[tokio::test]
async fn assemble_reassembles_ordered_frames() {
let f0 = RangeFrame {
offset: 0,
length: 3,
bytes: b"ABC".to_vec(),
complete: false,
total_length: Some(6),
chunk_lens: Some(vec![3, 3]),
chunk_index: Some(0),
inclusion_proof: Some("proof".into()),
root: Some("aa".repeat(32)),
};
let f1 = RangeFrame {
offset: 3,
length: 3,
bytes: b"DEF".to_vec(),
complete: true,
total_length: None,
chunk_lens: None,
chunk_index: None,
inclusion_proof: None,
root: None,
};
let mut wire = f0.encode();
wire.extend_from_slice(&f1.encode());
let mut cur = std::io::Cursor::new(wire);
let (bytes, meta) = assemble_range_stream(&mut cur, 6).await.unwrap();
assert_eq!(bytes, b"ABCDEF");
assert_eq!(meta.total_length, Some(6));
assert_eq!(meta.chunk_lens, Some(vec![3, 3]));
assert_eq!(meta.chunk_index, Some(0));
assert_eq!(meta.root, Some("aa".repeat(32)));
assert_eq!(meta.inclusion_proof, Some("proof".into()));
}
#[tokio::test]
async fn assemble_rejects_overflowing_frame() {
let f = RangeFrame {
offset: 0,
length: 10,
bytes: vec![0u8; 10],
complete: true,
total_length: None,
chunk_lens: None,
chunk_index: None,
inclusion_proof: None,
root: None,
};
let mut cur = std::io::Cursor::new(f.encode());
let err = assemble_range_stream(&mut cur, 5).await;
assert!(matches!(err, Err(DownloadError::Transport { .. })));
}
#[tokio::test]
async fn drain_trailer_is_bounded_by_cap() {
let flood = vec![0u8; 1_000_000];
let mut cur = std::io::Cursor::new(flood);
let drained = drain_trailer_bounded(&mut cur, 64 * 1024).await;
assert_eq!(drained, 64 * 1024, "drain must stop exactly at the cap");
assert!((cur.position() as usize) < 1_000_000);
}
#[tokio::test]
async fn drain_trailer_stops_at_eof_below_cap() {
let mut cur = std::io::Cursor::new(vec![0u8; 100]);
assert_eq!(drain_trailer_bounded(&mut cur, 64 * 1024).await, 100);
let mut empty = std::io::Cursor::new(Vec::<u8>::new());
assert_eq!(drain_trailer_bounded(&mut empty, 64 * 1024).await, 0);
}
#[tokio::test]
async fn assemble_stops_on_clean_eof() {
let f = RangeFrame {
offset: 0,
length: 2,
bytes: b"hi".to_vec(),
complete: false,
total_length: Some(2),
chunk_lens: Some(vec![2]),
chunk_index: Some(0),
inclusion_proof: None,
root: None,
};
let mut cur = std::io::Cursor::new(f.encode());
let (bytes, meta) = assemble_range_stream(&mut cur, 2).await.unwrap();
assert_eq!(bytes, b"hi");
assert_eq!(meta.total_length, Some(2));
}
#[test]
fn source_tracker_backoff_and_recovery() {
let mut t = SourceTracker::new(Duration::from_millis(100), Duration::from_secs(10));
let now = Instant::now();
assert!(t.is_available("p", now));
t.record_failure("p", now);
assert!(!t.is_available("p", now)); assert_eq!(t.failures("p"), 1);
assert!(t.is_available("p", now + Duration::from_millis(101)));
t.record_success("p");
assert!(t.is_available("p", now));
assert_eq!(t.failures("p"), 0);
assert_eq!(t.served("p"), 1);
}
#[test]
fn source_tracker_backoff_is_exponential_and_capped() {
let mut t = SourceTracker::new(Duration::from_millis(100), Duration::from_millis(250));
let now = Instant::now();
t.record_failure("p", now); assert!(t.is_available("p", now + Duration::from_millis(150)));
t.record_failure("p", now); assert!(!t.is_available("p", now + Duration::from_millis(150)));
t.record_failure("p", now); assert!(t.is_available("p", now + Duration::from_millis(260)));
}
fn fake_node_cert() -> std::sync::Arc<dig_nat::NodeCert> {
use sha2::{Digest, Sha256};
let seed: [u8; 32] = Sha256::digest(b"dig-download/tests/fake-node-cert").into();
let bls_sk = dig_tls::bls::SecretKey::from_seed(&seed);
std::sync::Arc::new(dig_nat::NodeCert::generate_signed(&bls_sk).unwrap())
}
}