use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::RwLock;
use dig_dht::routing::Contact;
use dig_dht::transport::DhtTransport;
use dig_dht::wire::{DhtRequest, DhtResponse};
use dig_dht::{BootstrapPeer, CandidateAddr, ContentId, DhtConfig, DhtError, DhtService, PeerId};
#[derive(Clone, Default)]
struct SwarmRouter {
nodes: Arc<RwLock<HashMap<String, Arc<DhtService>>>>,
offline: Arc<RwLock<HashMap<String, ()>>>,
}
impl SwarmRouter {
fn new() -> Self {
SwarmRouter::default()
}
async fn add(&self, service: Arc<DhtService>) {
self.nodes
.write()
.await
.insert(service.local_id().to_hex(), service);
}
async fn set_offline(&self, peer_id: &str) {
self.offline.write().await.insert(peer_id.to_string(), ());
}
fn transport(&self) -> Arc<dyn DhtTransport> {
Arc::new(RouterTransport {
router: self.clone(),
})
}
}
struct RouterTransport {
router: SwarmRouter,
}
#[async_trait]
impl DhtTransport for RouterTransport {
async fn rpc(
&self,
from: &Contact,
peer: &Contact,
request: &DhtRequest,
) -> Result<DhtResponse, DhtError> {
if self.router.offline.read().await.contains_key(&peer.peer_id) {
return Err(DhtError::transport("offline"));
}
let encoded = request.encode();
let mut cur = std::io::Cursor::new(encoded);
let decoded = DhtRequest::decode(&mut cur)
.await
.map_err(DhtError::transport)?;
let service = {
let nodes = self.router.nodes.read().await;
nodes.get(&peer.peer_id).cloned()
};
match service {
Some(s) => {
let resp = s.handle_request_from(Some(from.clone()), decoded).await;
let mut rcur = std::io::Cursor::new(resp.encode());
DhtResponse::decode(&mut rcur)
.await
.map_err(DhtError::transport)
}
None => Err(DhtError::transport("no route")),
}
}
}
fn pid(hi: u8, lo: u8) -> PeerId {
let mut b = [0u8; 32];
b[0] = hi;
b[1] = lo;
PeerId::from_bytes(b)
}
fn addr() -> Vec<CandidateAddr> {
vec![CandidateAddr::direct("203.0.113.1", 9444)]
}
async fn make_node(router: &SwarmRouter, id: PeerId, config: DhtConfig) -> Arc<DhtService> {
let svc = Arc::new(DhtService::new(id, addr(), config, router.transport()));
router.add(svc.clone()).await;
svc
}
fn bootstrap_of(svc: &DhtService) -> BootstrapPeer {
BootstrapPeer {
peer_id: *svc.local_id(),
addresses: addr(),
}
}
async fn build_swarm(router: &SwarmRouter, n: u8, config: &DhtConfig) -> Vec<Arc<DhtService>> {
let mut nodes = Vec::new();
for i in 0..n {
let svc = make_node(
router,
pid(i.wrapping_mul(7).wrapping_add(1), i),
config.clone(),
)
.await;
nodes.push(svc);
}
let seeds: Vec<BootstrapPeer> = nodes.iter().take(3).map(|s| bootstrap_of(s)).collect();
for svc in &nodes {
svc.bootstrap(&seeds).await.unwrap();
}
for svc in &nodes {
svc.bootstrap(&seeds).await.unwrap();
}
nodes
}
#[tokio::test]
async fn announce_then_find_providers_roundtrip() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 15, &config).await;
let holder = &nodes[7];
let content = ContentId::capsule([0x42; 32], [0x24; 32]);
let accepted = holder.announce_provider(&content).await.unwrap();
assert!(accepted > 0, "announce must PUT the record at some peers");
let seeker = &nodes[2];
let providers = seeker.find_providers(&content).await.unwrap();
assert_eq!(providers.len(), 1, "exactly the one holder");
assert_eq!(
providers[0].provider_peer_id,
holder.local_id().to_hex(),
"the provider must be the announcing node"
);
assert!(providers[0].best_address().is_some());
}
#[tokio::test]
async fn find_providers_returns_all_distinct_holders() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 15, &config).await;
let content = ContentId::store([0x99; 32]);
nodes[3].announce_provider(&content).await.unwrap();
nodes[8].announce_provider(&content).await.unwrap();
nodes[11].announce_provider(&content).await.unwrap();
let providers = nodes[1].find_providers(&content).await.unwrap();
let holder_ids: std::collections::HashSet<String> = providers
.iter()
.map(|p| p.provider_peer_id.clone())
.collect();
assert!(holder_ids.contains(&nodes[3].local_id().to_hex()));
assert!(holder_ids.contains(&nodes[8].local_id().to_hex()));
assert!(holder_ids.contains(&nodes[11].local_id().to_hex()));
assert_eq!(holder_ids.len(), 3, "all three distinct holders, deduped");
}
#[tokio::test]
async fn find_node_converges_across_the_swarm() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 20, &config).await;
let target = *nodes[17].local_id();
let found = nodes[0].find_node(&target).await.unwrap();
let ids: std::collections::HashSet<String> = found.iter().map(|c| c.peer_id.clone()).collect();
assert!(
ids.contains(&target.to_hex()),
"iterative find_node must locate the target peer across hops"
);
}
#[tokio::test]
async fn no_providers_returns_empty_not_error() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 10, &config).await;
let content = ContentId::resource([0x01; 32], [0x02; 32], [0x03; 32]);
let providers = nodes[4].find_providers(&content).await.unwrap();
assert!(
providers.is_empty(),
"unknown content → empty provider set (not an error)"
);
}
#[tokio::test]
async fn find_providers_with_no_peers_is_not_an_error() {
let router = SwarmRouter::new();
let solo = make_node(&router, pid(0xAA, 0), DhtConfig::default()).await;
let content = ContentId::store([0x77; 32]);
let providers = solo.find_providers(&content).await.unwrap();
assert!(providers.is_empty());
}
#[tokio::test]
async fn find_node_with_no_peers_errors() {
let router = SwarmRouter::new();
let solo = make_node(&router, pid(0xBB, 0), DhtConfig::default()).await;
let err = solo.find_node(&pid(0xCC, 0)).await;
assert!(matches!(err, Err(DhtError::NoPeers)));
}
#[tokio::test]
async fn provider_record_ttl_expires_then_republish_restores() {
let router = SwarmRouter::new();
let config = DhtConfig {
provider_ttl: std::time::Duration::from_secs(0),
..Default::default()
};
let nodes = build_swarm(&router, 12, &config).await;
let content = ContentId::capsule([0x55; 32], [0x66; 32]);
nodes[5].announce_provider(&content).await.unwrap();
let providers = nodes[1].find_providers(&content).await.unwrap();
assert!(
providers.is_empty(),
"expired provider records must not be returned"
);
for svc in &nodes {
svc.gc().await;
}
let router2 = SwarmRouter::new();
let good_config = DhtConfig::default();
let nodes2 = build_swarm(&router2, 12, &good_config).await;
nodes2[5].announce_provider(&content).await.unwrap();
let republished = nodes2[5].republish().await;
assert_eq!(republished, 1, "the one announced key is republished");
let providers2 = nodes2[1].find_providers(&content).await.unwrap();
assert_eq!(providers2.len(), 1, "republished record is findable");
}
#[tokio::test]
async fn ping_liveness_evicts_dead_peer() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 8, &config).await;
let pinger = &nodes[0];
let target_key = dig_dht::Key::from_peer_id(nodes[3].local_id());
let known = pinger.known_closest(&target_key).await;
let victim = known
.iter()
.find(|c| c.peer_id == nodes[3].local_id().to_hex())
.cloned()
.expect("pinger should know node 3 after bootstrap");
assert!(pinger.ping(&victim).await, "live peer answers ping");
let before = pinger.routing_len().await;
router.set_offline(&victim.peer_id).await;
assert!(!pinger.ping(&victim).await, "offline peer fails ping");
let after = pinger.routing_len().await;
assert_eq!(
after,
before - 1,
"failed-ping peer is evicted from the routing table"
);
}
#[tokio::test]
async fn serving_side_find_providers_returns_closer_when_no_providers() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 10, &config).await;
let key = ContentId::store([0xEE; 32]).to_key();
let resp = nodes[0]
.handle_request(DhtRequest::FindProviders {
content_key: key.to_hex(),
})
.await;
match resp {
DhtResponse::Providers { providers, closer } => {
assert!(providers.is_empty(), "no providers announced for this key");
assert!(
!closer.is_empty(),
"must return closer peers to continue the walk"
);
}
other => panic!("expected Providers, got {other:?}"),
}
}
#[tokio::test]
async fn refresh_buckets_runs_over_populated_buckets() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 12, &config).await;
let refreshed = nodes[0].refresh_buckets().await;
assert!(refreshed > 0, "at least one populated bucket to refresh");
}
#[tokio::test]
async fn add_provider_over_global_capacity_is_rejected_not_stored() {
let router = SwarmRouter::new();
let config = DhtConfig {
provider_store_limits: dig_dht::provider_store::ProviderStoreLimits {
max_providers_per_key: 20,
max_total_records: 2,
},
..Default::default()
};
let victim = make_node(&router, pid(0x10, 0), config).await;
let mk_record = |tag: u8| {
dig_dht::ProviderRecord::new(
&dig_dht::ContentId::store([tag; 32]).to_key(),
&pid(0x20, tag),
addr(),
u64::MAX,
)
};
let ok1 = victim
.handle_request(DhtRequest::AddProvider {
record: mk_record(1),
})
.await;
assert_eq!(ok1, DhtResponse::AddProviderOk);
let ok2 = victim
.handle_request(DhtRequest::AddProvider {
record: mk_record(2),
})
.await;
assert_eq!(ok2, DhtResponse::AddProviderOk);
let rejected = victim
.handle_request(DhtRequest::AddProvider {
record: mk_record(3),
})
.await;
match rejected {
DhtResponse::Error { .. } => {}
other => panic!("expected an over-capacity Error response, got {other:?}"),
}
let key3 = dig_dht::ContentId::store([3u8; 32]).to_key();
let resp = victim
.handle_request(DhtRequest::FindProviders {
content_key: key3.to_hex(),
})
.await;
match resp {
DhtResponse::Providers { providers, .. } => {
assert!(providers.is_empty(), "rejected record must not be stored")
}
other => panic!("expected Providers, got {other:?}"),
}
}
#[tokio::test]
async fn add_provider_with_malicious_expiry_is_clamped_to_local_ttl() {
let router = SwarmRouter::new();
let short_ttl = std::time::Duration::from_secs(60);
let config = DhtConfig {
provider_ttl: short_ttl,
..Default::default()
};
let victim = make_node(&router, pid(0x11, 0), config).await;
let content = ContentId::store([0x77; 32]);
let malicious = dig_dht::ProviderRecord::new(
&content.to_key(),
&pid(0x22, 0),
addr(),
u64::MAX, );
let resp = victim
.handle_request(DhtRequest::AddProvider { record: malicious })
.await;
assert_eq!(resp, DhtResponse::AddProviderOk);
let key = content.to_key();
let resp = victim
.handle_request(DhtRequest::FindProviders {
content_key: key.to_hex(),
})
.await;
let stored = match resp {
DhtResponse::Providers { providers, .. } => providers,
other => panic!("expected Providers, got {other:?}"),
};
assert_eq!(stored.len(), 1);
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
assert!(
stored[0].expires_at <= now + short_ttl.as_secs() + 5, "malicious expires_at must be clamped to local TTL, got {}",
stored[0].expires_at
);
assert!(
stored[0].expires_at < u64::MAX / 2,
"clamp must actually bound the value, not just leave it near u64::MAX"
);
}
#[tokio::test]
async fn add_provider_with_unbounded_addresses_is_capped_at_the_boundary() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let victim = make_node(&router, pid(0x13, 0), config).await;
let content = ContentId::store([0x88; 32]);
let flood: Vec<CandidateAddr> = (0..2000)
.map(|i| CandidateAddr::direct(format!("203.0.113.{}", i % 255), 9444))
.collect();
let malicious = dig_dht::ProviderRecord {
content_key: content.to_key().to_hex(),
provider_peer_id: pid(0x24, 0).to_hex(),
addresses: flood,
expires_at: u64::MAX,
};
let resp = victim
.handle_request(DhtRequest::AddProvider { record: malicious })
.await;
assert_eq!(resp, DhtResponse::AddProviderOk);
let key = content.to_key();
let resp = victim
.handle_request(DhtRequest::FindProviders {
content_key: key.to_hex(),
})
.await;
let stored = match resp {
DhtResponse::Providers { providers, .. } => providers,
other => panic!("expected Providers, got {other:?}"),
};
assert_eq!(stored.len(), 1);
assert_eq!(
stored[0].addresses.len(),
dig_dht::record::MAX_ADDRESSES_PER_RECORD,
"stored record's address list must be capped, not the raw flood"
);
}
#[tokio::test]
async fn add_provider_naming_a_third_party_provider_is_rejected() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let victim = make_node(&router, pid(0x30, 0), config).await;
let caller = Contact::new(&pid(0x31, 0), addr());
let third_party = pid(0x32, 0); let content = ContentId::store([0x55; 32]);
let record = dig_dht::ProviderRecord::new(&content.to_key(), &third_party, addr(), u64::MAX);
let resp = victim
.handle_request_from(Some(caller), DhtRequest::AddProvider { record })
.await;
match resp {
DhtResponse::Error { .. } => {}
other => panic!("expected an Error response for third-party announce, got {other:?}"),
}
let key = content.to_key();
let resp = victim
.handle_request(DhtRequest::FindProviders {
content_key: key.to_hex(),
})
.await;
match resp {
DhtResponse::Providers { providers, .. } => assert!(
providers.is_empty(),
"third-party-named record must not be stored"
),
other => panic!("expected Providers, got {other:?}"),
}
}
#[tokio::test]
async fn add_provider_self_announce_is_accepted() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let victim = make_node(&router, pid(0x33, 0), config).await;
let announcer_id = pid(0x34, 0);
let caller = Contact::new(&announcer_id, addr());
let content = ContentId::store([0x56; 32]);
let record = dig_dht::ProviderRecord::new(&content.to_key(), &announcer_id, addr(), u64::MAX);
let resp = victim
.handle_request_from(Some(caller), DhtRequest::AddProvider { record })
.await;
assert_eq!(resp, DhtResponse::AddProviderOk);
}
#[tokio::test]
async fn add_provider_with_no_authenticated_caller_is_still_accepted() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let victim = make_node(&router, pid(0x35, 0), config).await;
let content = ContentId::store([0x57; 32]);
let record = dig_dht::ProviderRecord::new(&content.to_key(), &pid(0x36, 0), addr(), u64::MAX);
let resp = victim
.handle_request(DhtRequest::AddProvider { record })
.await;
assert_eq!(resp, DhtResponse::AddProviderOk);
}
#[tokio::test]
async fn withdraw_provider_stops_republish() {
let router = SwarmRouter::new();
let config = DhtConfig::default();
let nodes = build_swarm(&router, 10, &config).await;
let content = ContentId::store([0x33; 32]);
nodes[6].announce_provider(&content).await.unwrap();
assert!(nodes[6].withdraw_provider(&content).await, "was announced");
assert!(
!nodes[6].withdraw_provider(&content).await,
"no longer announced"
);
assert_eq!(nodes[6].republish().await, 0);
}