use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use tokio::sync::Mutex;
use dig_nat::PeerId;
use crate::config::DhtConfig;
use crate::content::ContentId;
use crate::error::DhtError;
use crate::key::Key;
use crate::lookup::{iterative_find, QueryOutcome};
use crate::provider_store::{ProviderStore, PutOutcome};
use crate::record::{CandidateAddr, ProviderRecord};
use crate::routing::{Contact, InsertOutcome, RoutingTable};
use crate::transport::DhtTransport;
use crate::wire::{DhtRequest, DhtResponse};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BootstrapPeer {
pub peer_id: PeerId,
pub addresses: Vec<CandidateAddr>,
}
impl BootstrapPeer {
pub fn direct(peer_id: PeerId, host: impl Into<String>, port: u16) -> Self {
BootstrapPeer {
peer_id,
addresses: vec![CandidateAddr::direct(host, port)],
}
}
fn to_contact(&self) -> Contact {
Contact::new(&self.peer_id, self.addresses.clone())
}
}
pub struct DhtService {
local_id: PeerId,
local_addresses: Vec<CandidateAddr>,
config: DhtConfig,
routing: Arc<Mutex<RoutingTable>>,
providers: Arc<Mutex<ProviderStore>>,
transport: Arc<dyn DhtTransport>,
}
impl DhtService {
pub fn new(
local_id: PeerId,
local_addresses: Vec<CandidateAddr>,
config: DhtConfig,
transport: Arc<dyn DhtTransport>,
) -> Self {
let routing = RoutingTable::new(&local_id, config.k);
let providers = ProviderStore::with_limits(config.provider_store_limits);
DhtService {
local_id,
local_addresses,
config,
routing: Arc::new(Mutex::new(routing)),
providers: Arc::new(Mutex::new(providers)),
transport,
}
}
pub fn local_id(&self) -> &PeerId {
&self.local_id
}
fn local_contact(&self) -> Contact {
Contact::new(&self.local_id, self.local_addresses.clone())
}
pub async fn bootstrap(&self, peers: &[BootstrapPeer]) -> Result<usize, DhtError> {
{
let mut rt = self.routing.lock().await;
for p in peers {
let _ = rt.insert(p.to_contact());
}
}
let self_key = Key::from_peer_id(&self.local_id);
let seeds: Vec<Contact> = peers.iter().map(|p| p.to_contact()).collect();
let result = self.run_lookup(self_key, seeds, false).await;
self.absorb_contacts(&result.closest).await;
Ok(self.routing.lock().await.len())
}
pub async fn find_node(&self, peer_id: &PeerId) -> Result<Vec<Contact>, DhtError> {
let target = Key::from_peer_id(peer_id);
let seeds = self.seed_contacts(&target).await;
if seeds.is_empty() {
return Err(DhtError::NoPeers);
}
let result = self.run_lookup(target, seeds, false).await;
self.absorb_contacts(&result.closest).await;
Ok(result.closest)
}
pub async fn find_providers(
&self,
content: &ContentId,
) -> Result<Vec<ProviderRecord>, DhtError> {
let target = content.to_key();
let now = now_secs();
let mut local = self.providers.lock().await.get(&target.to_hex(), now);
let seeds = self.seed_contacts(&target).await;
if seeds.is_empty() {
return Ok(local);
}
let result = self.run_lookup(target, seeds, true).await;
self.absorb_contacts(&result.closest).await;
let mut discovered = result.providers;
for r in &mut discovered {
crate::record::sort_and_cap_addresses(&mut r.addresses);
}
local.extend(discovered);
let now = now_secs();
let mut seen = std::collections::HashSet::new();
local.retain(|r| !r.is_expired(now) && seen.insert(r.provider_peer_id.clone()));
Ok(local)
}
pub async fn announce_provider(&self, content: &ContentId) -> Result<usize, DhtError> {
let target = content.to_key();
let record = self.build_local_record(&target);
{
let mut ps = self.providers.lock().await;
ps.put(record.clone());
ps.mark_announced(target.to_hex());
}
let seeds = self.seed_contacts(&target).await;
if seeds.is_empty() {
return Ok(0);
}
let result = self.run_lookup(target, seeds, false).await;
self.absorb_contacts(&result.closest).await;
Ok(self.put_record_at(&result.closest, &record).await)
}
pub async fn withdraw_provider(&self, content: &ContentId) -> bool {
let key = content.to_key().to_hex();
self.providers.lock().await.unmark_announced(&key)
}
pub async fn ingest_verified_provider(&self, record: ProviderRecord) -> PutOutcome {
self.admit_verified_record(record).await
}
pub async fn remove_provider_record(&self, content_key: &str, provider_peer_id: &str) -> bool {
self.providers
.lock()
.await
.remove(content_key, provider_peer_id)
}
pub async fn retract_own_provider(&self, content: &ContentId) -> bool {
let key = content.to_key().to_hex();
let self_id = self.local_id.to_hex();
let mut ps = self.providers.lock().await;
let removed_record = ps.remove(&key, &self_id);
let was_announced = ps.unmark_announced(&key);
removed_record || was_announced
}
pub async fn holders_of(&self, content: &ContentId) -> Result<Vec<PeerId>, DhtError> {
let records = self.find_providers(content).await?;
Ok(records
.iter()
.filter_map(|r| r.provider_peer_id())
.collect())
}
pub async fn republish(&self) -> usize {
let keys = self.providers.lock().await.local_announcements();
let count = keys.len();
for hex in keys {
let Some(bytes) = hex64_to_bytes(&hex) else {
continue;
};
let target = Key::from_bytes(bytes);
let record = self.build_local_record(&target);
self.providers.lock().await.put(record.clone());
let seeds = self.seed_contacts(&target).await;
if !seeds.is_empty() {
let result = self.run_lookup(target, seeds, false).await;
self.absorb_contacts(&result.closest).await;
self.put_record_at(&result.closest, &record).await;
}
}
count
}
pub async fn refresh_buckets(&self) -> usize {
let indices = self.routing.lock().await.non_empty_bucket_indices();
let count = indices.len();
for idx in indices {
let target = self.random_key_in_bucket(idx);
let seeds = self.seed_contacts(&target).await;
if !seeds.is_empty() {
let result = self.run_lookup(target, seeds, false).await;
self.absorb_contacts(&result.closest).await;
}
}
count
}
pub async fn gc(&self) -> usize {
self.providers.lock().await.gc(now_secs())
}
pub async fn ping(&self, peer: &Contact) -> bool {
let nonce = rand::random::<u64>();
let from = self.local_contact();
match self
.transport
.rpc(&from, peer, &DhtRequest::Ping { nonce })
.await
{
Ok(DhtResponse::Pong { nonce: got }) if got == nonce => true,
_ => {
self.routing.lock().await.remove(&peer.peer_id);
false
}
}
}
pub async fn handle_request(&self, request: DhtRequest) -> DhtResponse {
self.handle_request_from(None, request).await
}
pub async fn handle_request_from(
&self,
caller: Option<Contact>,
request: DhtRequest,
) -> DhtResponse {
let caller_peer_id = caller.as_ref().map(|c| c.peer_id.clone());
if let Some(mut c) = caller {
if c.peer_id != self.local_id.to_hex() {
crate::record::sort_and_cap_addresses(&mut c.addresses);
let _ = self.routing.lock().await.insert(c);
}
}
match request {
DhtRequest::Ping { nonce } => DhtResponse::Pong { nonce },
DhtRequest::FindNode { target } => {
let Some(key) = parse_key(&target) else {
return DhtResponse::Error {
code: 2,
message: "bad target key".into(),
};
};
let nodes = self.routing.lock().await.closest(&key);
DhtResponse::Nodes { nodes }
}
DhtRequest::FindProviders { content_key } => {
let Some(key) = parse_key(&content_key) else {
return DhtResponse::Error {
code: 2,
message: "bad content key".into(),
};
};
let now = now_secs();
let providers = self.providers.lock().await.get(&key.to_hex(), now);
let closer = self.routing.lock().await.closest(&key);
DhtResponse::Providers { providers, closer }
}
DhtRequest::AddProvider { record } => {
if let Some(caller_id) = &caller_peer_id {
if *caller_id != record.provider_peer_id {
return DhtResponse::Error {
code: 4,
message:
"add_provider: provider_peer_id must match the authenticated caller"
.into(),
};
}
}
match self.admit_verified_record(record).await {
PutOutcome::Accepted => DhtResponse::AddProviderOk,
PutOutcome::RejectedOverCapacity => DhtResponse::Error {
code: 3,
message: "provider store over capacity".into(),
},
}
}
}
}
async fn admit_verified_record(&self, mut record: ProviderRecord) -> PutOutcome {
crate::record::sort_and_cap_addresses(&mut record.addresses);
let clamp_ceiling = now_secs().saturating_add(self.config.provider_ttl_secs());
record.expires_at = record.expires_at.min(clamp_ceiling);
let outcome = self.providers.lock().await.put(record.clone());
if outcome == PutOutcome::Accepted {
if let Some(pid) = record.provider_peer_id() {
let contact = Contact::new(&pid, record.addresses.clone());
let _ = self.routing.lock().await.insert(contact);
}
}
outcome
}
fn build_local_record(&self, target: &Key) -> ProviderRecord {
let expires_at = now_secs().saturating_add(self.config.provider_ttl_secs());
ProviderRecord::new(
target,
&self.local_id,
self.local_addresses.clone(),
expires_at,
)
}
async fn seed_contacts(&self, target: &Key) -> Vec<Contact> {
self.routing.lock().await.closest(target)
}
async fn run_lookup(
&self,
target: Key,
seeds: Vec<Contact>,
stop_on_providers: bool,
) -> crate::lookup::LookupResult {
let transport = self.transport.clone();
let content_key = target.to_hex();
let from = self.local_contact();
let query = move |contact: Contact| {
let transport = transport.clone();
let content_key = content_key.clone();
let from = from.clone();
async move {
let req = DhtRequest::FindProviders { content_key };
match transport.rpc(&from, &contact, &req).await {
Ok(DhtResponse::Providers { providers, closer }) => {
Ok(QueryOutcome { closer, providers })
}
Ok(DhtResponse::Nodes { nodes }) => Ok(QueryOutcome {
closer: nodes,
providers: vec![],
}),
_ => Err(()),
}
}
};
iterative_find(
target,
seeds,
self.config.k,
self.config.alpha,
stop_on_providers,
query,
)
.await
}
async fn absorb_contacts(&self, contacts: &[Contact]) {
let mut rt = self.routing.lock().await;
for c in contacts {
let mut c = c.clone();
crate::record::sort_and_cap_addresses(&mut c.addresses);
match rt.insert(c) {
InsertOutcome::Inserted => {}
InsertOutcome::Full { .. } => {
}
}
}
}
async fn put_record_at(&self, peers: &[Contact], record: &ProviderRecord) -> usize {
let req = DhtRequest::AddProvider {
record: record.clone(),
};
let from = self.local_contact();
let mut accepted = 0;
for p in peers {
if p.peer_id == self.local_id.to_hex() {
continue; }
if let Ok(DhtResponse::AddProviderOk) = self.transport.rpc(&from, p, &req).await {
accepted += 1;
}
}
accepted
}
fn random_key_in_bucket(&self, idx: usize) -> Key {
let local = *self.local_id.as_bytes();
let mut distance = [0u8; 32];
let bit = 255 - idx; let byte = bit / 8;
let bit_in_byte = 7 - (bit % 8);
distance[byte] = 1 << bit_in_byte;
for b in distance.iter_mut().skip(byte + 1) {
*b = rand::random::<u8>();
}
let mut target = [0u8; 32];
for i in 0..32 {
target[i] = local[i] ^ distance[i];
}
Key::from_bytes(target)
}
pub async fn known_closest(&self, target: &Key) -> Vec<Contact> {
self.routing.lock().await.closest(target)
}
pub async fn routing_len(&self) -> usize {
self.routing.lock().await.len()
}
}
fn now_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
fn parse_key(hex: &str) -> Option<Key> {
hex64_to_bytes(hex).map(Key::from_bytes)
}
fn hex64_to_bytes(hex: &str) -> Option<[u8; 32]> {
if hex.len() != 64 {
return None;
}
let mut out = [0u8; 32];
let bytes = hex.as_bytes();
for (i, chunk) in bytes.chunks(2).enumerate() {
let hi = (chunk[0] as char).to_digit(16)?;
let lo = (chunk[1] as char).to_digit(16)?;
out[i] = ((hi << 4) | lo) as u8;
}
Some(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn key_hex_round_trips() {
}
#[test]
fn hex64_round_trip() {
let bytes = [0xABu8; 32];
let hex = Key::from_bytes(bytes).to_hex();
assert_eq!(hex64_to_bytes(&hex).unwrap(), bytes);
assert!(hex64_to_bytes("short").is_none());
assert!(hex64_to_bytes(&"zz".repeat(32)).is_none());
key_hex_round_trips();
}
#[test]
fn parse_key_rejects_bad_hex() {
assert!(parse_key("nothex").is_none());
assert!(parse_key(&"00".repeat(32)).is_some());
}
}