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_dial_targets(
&self,
provider: &ProviderRecord,
) -> Result<Vec<(String, dig_nat::PeerTarget)>, DownloadError> {
let peer_id = provider.provider_peer_id().ok_or_else(|| {
DownloadError::transport(&provider.provider_peer_id, "malformed provider peer_id")
})?;
let mut targets = Vec::new();
for candidate in crate::addr::dial_candidates(provider) {
match crate::addr::candidate_socket(candidate) {
Ok(socket) => targets.push((
socket.to_string(),
dig_nat::PeerTarget::with_addr(peer_id, socket, self.network_id.clone()),
)),
Err(e) => tracing::warn!(
peer = %provider.provider_peer_id,
candidate = %crate::addr::display(candidate),
error = %e,
"skipping unusable provider candidate address"
),
}
}
targets.push((
"relay-only".to_string(),
dig_nat::PeerTarget::relay_only(peer_id, self.network_id.clone()),
));
Ok(targets)
}
pub fn provider_to_target(
&self,
provider: &ProviderRecord,
) -> Result<dig_nat::PeerTarget, DownloadError> {
let (_, target) = self
.provider_dial_targets(provider)?
.into_iter()
.next()
.expect("dial targets always include the relay-only fallback");
Ok(target)
}
async fn connect(&self, provider: &ProviderRecord) -> Result<DigPeer, DownloadError> {
let mut last_error = None;
for (addr, target) in self.provider_dial_targets(provider)? {
match DigPeer::connect_with_runtime(&target, &self.node, &self.config, &self.runtime)
.await
{
Ok(peer) => return Ok(peer),
Err(e) => {
tracing::debug!(
peer = %provider.provider_peer_id,
candidate = %addr,
error = %e,
"provider dial candidate failed; trying the next address"
);
last_error = Some(format!("dial {addr}: {e}"));
}
}
}
Err(DownloadError::transport(
&provider.provider_peer_id,
last_error.unwrap_or_else(|| "no dialable candidate address".to_string()),
))
}
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_accepts_v4_mapped_v6_host() {
let t = NatRangeTransport::new(
fake_node_cert(),
dig_nat::NatConfig::default(),
"DIG_MAINNET",
);
let p = provider(1, "::ffff:172.31.79.22", 9444);
let target = t
.provider_to_target(&p)
.expect("v4-mapped v6 host must resolve");
assert_eq!(
target.direct_addr().unwrap(),
std::net::SocketAddr::new("::ffff:172.31.79.22".parse().unwrap(), 9444)
);
}
#[test]
fn provider_to_target_accepts_plain_v6_host() {
let t = NatRangeTransport::new(
fake_node_cert(),
dig_nat::NatConfig::default(),
"DIG_MAINNET",
);
let p = provider(1, "2001:db8::1", 9444);
let target = t.provider_to_target(&p).expect("v6 host must resolve");
assert_eq!(
target.direct_addr().unwrap(),
std::net::SocketAddr::new("2001:db8::1".parse().unwrap(), 9444)
);
}
#[test]
fn unusable_first_candidate_falls_through_to_the_ipv4_one() {
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([3; 32]),
vec![
CandidateAddr::direct("not-an-ip-literal", 9444),
CandidateAddr::direct("10.0.0.1", 9444),
],
u64::MAX,
);
let target = t
.provider_to_target(&p)
.expect("the v4 candidate is dialable");
assert_eq!(
target.direct_addr().unwrap(),
"10.0.0.1:9444".parse::<std::net::SocketAddr>().unwrap()
);
}
#[test]
fn dial_targets_order_v6_then_v4_then_relay() {
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([4; 32]),
vec![
CandidateAddr::direct("172.31.79.22", 9444),
CandidateAddr::direct("::ffff:172.31.79.22", 9444),
],
u64::MAX,
);
let addrs: Vec<String> = t
.provider_dial_targets(&p)
.unwrap()
.into_iter()
.map(|(addr, _)| addr)
.collect();
assert_eq!(
addrs,
vec![
"[::ffff:172.31.79.22]:9444",
"172.31.79.22:9444",
"relay-only"
]
);
}
#[tokio::test]
async fn connect_tries_every_candidate_before_failing() {
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([5; 32]),
vec![
CandidateAddr::direct("::1", 1),
CandidateAddr::direct("127.0.0.1", 1),
],
u64::MAX,
);
let err = t.connect(&p).await.expect_err("no listener is up");
let reason = err.to_string();
assert!(
reason.contains("relay-only"),
"the last attempt must be named: {reason}"
);
}
#[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())
}
}