use std::collections::{HashMap, HashSet};
use std::future::Future;
use crate::key::Key;
use crate::record::ProviderRecord;
use crate::routing::Contact;
pub const MAX_LOOKUP_ROUNDS: usize = 64;
pub const MAX_PROVIDERS_PER_RESPONSE: usize = 64;
pub const MAX_CLOSER_PER_RESPONSE: usize = 64;
#[derive(Debug, Default, Clone)]
pub struct QueryOutcome {
pub closer: Vec<Contact>,
pub providers: Vec<ProviderRecord>,
}
#[derive(Debug, Default, Clone)]
pub struct LookupResult {
pub closest: Vec<Contact>,
pub providers: Vec<ProviderRecord>,
}
struct ShortlistEntry {
contact: Contact,
distance: crate::key::Distance,
queried: bool,
failed: bool,
}
pub async fn iterative_find<F, Fut>(
target: Key,
seeds: Vec<Contact>,
k: usize,
alpha: usize,
stop_on_providers: bool,
query: F,
) -> LookupResult
where
F: Fn(Contact) -> Fut + Clone + Send + 'static,
Fut: Future<Output = Result<QueryOutcome, ()>> + Send + 'static,
{
let mut shortlist: Vec<ShortlistEntry> = Vec::new();
let mut seen: HashSet<String> = HashSet::new();
let mut providers: HashMap<String, ProviderRecord> = HashMap::new();
let max_providers_per_response = MAX_PROVIDERS_PER_RESPONSE.max(k.saturating_mul(2));
let max_closer_per_response = MAX_CLOSER_PER_RESPONSE.max(k.saturating_mul(2));
for c in seeds {
merge_contact(&mut shortlist, &mut seen, &target, c);
}
sort_and_cap(&mut shortlist, k, alpha);
let mut rounds = 0;
loop {
rounds += 1;
if rounds > MAX_LOOKUP_ROUNDS {
break;
}
let batch: Vec<Contact> = shortlist
.iter()
.filter(|e| !e.queried && !e.failed)
.take(alpha)
.map(|e| e.contact.clone())
.collect();
if batch.is_empty() {
break; }
let mut set = tokio::task::JoinSet::new();
for c in &batch {
let q = query.clone();
let c2 = c.clone();
set.spawn(async move { (c2.peer_id.clone(), q(c2).await) });
}
let mut results = Vec::with_capacity(batch.len());
while let Some(joined) = set.join_next().await {
match joined {
Ok(pair) => results.push(pair),
Err(_join_err) => {
}
}
}
for (peer_id, res) in results {
mark_queried(&mut shortlist, &peer_id, res.is_err());
if let Ok(mut outcome) = res {
outcome.providers.truncate(max_providers_per_response);
outcome.closer.truncate(max_closer_per_response);
for p in outcome.providers {
providers.entry(p.provider_peer_id.clone()).or_insert(p);
}
for c in outcome.closer {
merge_contact(&mut shortlist, &mut seen, &target, c);
}
}
}
sort_and_cap(&mut shortlist, k, alpha);
if stop_on_providers && !providers.is_empty() {
break;
}
let any_unqueried_in_top_k = shortlist.iter().take(k).any(|e| !e.queried && !e.failed);
if !any_unqueried_in_top_k {
break;
}
}
let closest = shortlist
.into_iter()
.filter(|e| !e.failed)
.take(k)
.map(|e| e.contact)
.collect();
LookupResult {
closest,
providers: providers.into_values().collect(),
}
}
fn merge_contact(
shortlist: &mut Vec<ShortlistEntry>,
seen: &mut HashSet<String>,
target: &Key,
contact: Contact,
) {
if seen.contains(&contact.peer_id) {
return;
}
let Some(key) = contact.key() else {
return;
};
seen.insert(contact.peer_id.clone());
let distance = target.distance(&key);
shortlist.push(ShortlistEntry {
contact,
distance,
queried: false,
failed: false,
});
}
fn mark_queried(shortlist: &mut [ShortlistEntry], peer_id: &str, failed: bool) {
if let Some(e) = shortlist.iter_mut().find(|e| e.contact.peer_id == peer_id) {
e.queried = true;
e.failed = failed;
}
}
fn sort_and_cap(shortlist: &mut Vec<ShortlistEntry>, k: usize, alpha: usize) {
shortlist.sort_by_key(|e| e.distance);
let cap = (k * 3).max(k + alpha * 2);
if shortlist.len() > cap {
shortlist.truncate(cap);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::record::CandidateAddr;
use dig_nat::PeerId;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
fn contact_from_key(key_bytes: [u8; 32]) -> Contact {
Contact::new(
&PeerId::from_bytes(key_bytes),
vec![CandidateAddr::direct("h", 1)],
)
}
fn oracle_query(
all_ids: Vec<[u8; 32]>,
target: Key,
k: usize,
counter: Arc<AtomicUsize>,
) -> impl Fn(Contact) -> std::pin::Pin<Box<dyn Future<Output = Result<QueryOutcome, ()>> + Send>>
+ Clone {
move |_c: Contact| {
counter.fetch_add(1, Ordering::SeqCst);
let all_ids = all_ids.clone();
Box::pin(async move {
let mut sorted: Vec<[u8; 32]> = all_ids;
sorted.sort_by_key(|id| *target.distance(&Key::from_bytes(*id)).as_bytes());
let closer = sorted.into_iter().take(k).map(contact_from_key).collect();
Ok(QueryOutcome {
closer,
providers: vec![],
})
})
}
}
#[tokio::test]
async fn converges_to_k_closest_in_simulated_network() {
let all_ids: Vec<[u8; 32]> = (0u8..50)
.map(|i| {
let mut b = [0u8; 32];
b[0] = i.wrapping_mul(5);
b[1] = i;
b
})
.collect();
let target = Key::from_bytes([0u8; 32]);
let k = 20;
let counter = Arc::new(AtomicUsize::new(0));
let seeds = vec![contact_from_key(all_ids[40]), contact_from_key(all_ids[45])];
let query = oracle_query(all_ids.clone(), target, k, counter.clone());
let result = iterative_find(target, seeds, k, 3, false, query).await;
let mut expected = all_ids.clone();
expected.sort_by_key(|id| *target.distance(&Key::from_bytes(*id)).as_bytes());
let expected_top: Vec<Contact> =
expected.into_iter().take(k).map(contact_from_key).collect();
assert_eq!(result.closest.len(), k);
let got: HashSet<String> = result.closest.iter().map(|c| c.peer_id.clone()).collect();
let exp: HashSet<String> = expected_top.iter().map(|c| c.peer_id.clone()).collect();
assert_eq!(got, exp, "lookup must converge on the true k-closest");
assert_eq!(result.closest[0].peer_id, expected_top[0].peer_id);
}
#[tokio::test]
async fn empty_seeds_returns_empty() {
let target = Key::from_bytes([0u8; 32]);
let result = iterative_find(target, vec![], 20, 3, false, |_c: Contact| async {
Ok(QueryOutcome::default())
})
.await;
assert!(result.closest.is_empty());
assert!(result.providers.is_empty());
}
#[tokio::test]
async fn stop_on_providers_ends_early() {
let target = Key::from_bytes([0u8; 32]);
let seed = contact_from_key([0x10; 32]);
let provider = ProviderRecord::new(
&target,
&PeerId::from_bytes([0xAB; 32]),
vec![CandidateAddr::direct("h", 9444)],
u64::MAX,
);
let p2 = provider.clone();
let result = iterative_find(target, vec![seed], 20, 3, true, move |_c: Contact| {
let p = p2.clone();
async move {
Ok(QueryOutcome {
closer: vec![],
providers: vec![p],
})
}
})
.await;
assert_eq!(result.providers.len(), 1);
assert_eq!(
result.providers[0].provider_peer_id,
provider.provider_peer_id
);
}
#[tokio::test]
async fn failed_peers_do_not_abort_lookup() {
let target = Key::from_bytes([0u8; 32]);
let seeds = vec![contact_from_key([0x01; 32]), contact_from_key([0x02; 32])];
let result = iterative_find(target, seeds, 20, 3, false, |_c: Contact| async {
Err::<QueryOutcome, ()>(())
})
.await;
assert!(result.closest.is_empty());
}
#[tokio::test]
async fn dedups_providers_by_peer_id() {
let target = Key::from_bytes([0u8; 32]);
let seeds = vec![contact_from_key([0x01; 32]), contact_from_key([0x02; 32])];
let provider =
ProviderRecord::new(&target, &PeerId::from_bytes([0xAB; 32]), vec![], u64::MAX);
let p = provider.clone();
let result = iterative_find(target, seeds, 20, 3, false, move |_c: Contact| {
let p = p.clone();
async move {
Ok(QueryOutcome {
closer: vec![],
providers: vec![p],
})
}
})
.await;
assert_eq!(result.providers.len(), 1);
}
#[tokio::test]
async fn providers_per_response_is_capped() {
let target = Key::from_bytes([0u8; 32]);
let seed = contact_from_key([0x10; 32]);
let flood: Vec<ProviderRecord> = (0u32..10_000)
.map(|i| {
let mut b = [0u8; 32];
b[0..4].copy_from_slice(&i.to_be_bytes());
ProviderRecord::new(&target, &PeerId::from_bytes(b), vec![], u64::MAX)
})
.collect();
let result = iterative_find(target, vec![seed], 20, 3, false, move |_c: Contact| {
let flood = flood.clone();
async move {
Ok(QueryOutcome {
closer: vec![],
providers: flood,
})
}
})
.await;
let cap = MAX_PROVIDERS_PER_RESPONSE.max(20 * 2);
assert!(
result.providers.len() <= cap,
"providers-per-response must be capped at {cap}, got {}",
result.providers.len()
);
}
#[tokio::test]
async fn round_guard_terminates_a_non_converging_lookup() {
let target = Key::from_bytes([0xFF; 32]);
let seed = contact_from_key([0x00; 32]);
let calls = Arc::new(AtomicUsize::new(0));
let counter = calls.clone();
const ALPHA: usize = 3;
let result = iterative_find(target, vec![seed], 20, ALPHA, false, move |_c: Contact| {
let n = counter.fetch_add(1, Ordering::SeqCst);
let mut b = [0xFFu8; 32];
b[24..32].copy_from_slice(&(n as u64).to_be_bytes());
async move {
Ok(QueryOutcome {
closer: vec![contact_from_key(b)],
providers: vec![],
})
}
})
.await;
let total = calls.load(Ordering::SeqCst);
assert!(
total <= MAX_LOOKUP_ROUNDS * ALPHA + 1,
"round guard must bound total queries to ~MAX_LOOKUP_ROUNDS*alpha, got {total}"
);
assert!(result.closest.len() <= 20);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn alpha_batch_is_queried_concurrently_not_sequentially() {
const ALPHA: usize = 4;
const DELAY: std::time::Duration = std::time::Duration::from_millis(150);
let target = Key::from_bytes([0u8; 32]);
let seeds: Vec<Contact> = (1u8..=ALPHA as u8)
.map(|i| contact_from_key([i; 32]))
.collect();
let start = std::time::Instant::now();
let result = iterative_find(target, seeds, 20, ALPHA, false, |_c: Contact| async move {
tokio::time::sleep(DELAY).await;
Ok(QueryOutcome::default())
})
.await;
let elapsed = start.elapsed();
assert_eq!(result.closest.len(), ALPHA, "all alpha peers answered");
assert!(
elapsed < DELAY * (ALPHA as u32) / 2,
"batch must run concurrently: expected well under {:?} (alpha * delay), got {:?}",
DELAY * (ALPHA as u32),
elapsed
);
}
}