use std::collections::{HashMap, HashSet};
use crate::record::ProviderRecord;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProviderStoreLimits {
pub max_providers_per_key: usize,
pub max_total_records: usize,
}
impl Default for ProviderStoreLimits {
fn default() -> Self {
ProviderStoreLimits {
max_providers_per_key: 20,
max_total_records: 100_000,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PutOutcome {
Accepted,
RejectedOverCapacity,
}
#[derive(Debug)]
pub struct ProviderStore {
by_key: HashMap<String, HashMap<String, ProviderRecord>>,
announced: HashSet<String>,
limits: ProviderStoreLimits,
}
impl Default for ProviderStore {
fn default() -> Self {
ProviderStore::new()
}
}
impl ProviderStore {
pub fn new() -> Self {
ProviderStore::with_limits(ProviderStoreLimits::default())
}
pub fn with_limits(limits: ProviderStoreLimits) -> Self {
ProviderStore {
by_key: HashMap::new(),
announced: HashSet::new(),
limits,
}
}
pub fn put(&mut self, record: ProviderRecord) -> PutOutcome {
let is_new_provider = !self
.by_key
.get(&record.content_key)
.is_some_and(|providers| providers.contains_key(&record.provider_peer_id));
if is_new_provider {
if self.len() >= self.limits.max_total_records {
return PutOutcome::RejectedOverCapacity;
}
if let Some(providers) = self.by_key.get_mut(&record.content_key) {
if providers.len() >= self.limits.max_providers_per_key {
if let Some(evict_id) = providers
.iter()
.min_by_key(|(_, r)| r.expires_at)
.map(|(pid, _)| pid.clone())
{
providers.remove(&evict_id);
}
}
}
}
self.by_key
.entry(record.content_key.clone())
.or_default()
.insert(record.provider_peer_id.clone(), record);
PutOutcome::Accepted
}
pub fn remove(&mut self, content_key: &str, provider_peer_id: &str) -> bool {
let Some(providers) = self.by_key.get_mut(content_key) else {
return false;
};
let removed = providers.remove(provider_peer_id).is_some();
if providers.is_empty() {
self.by_key.remove(content_key);
}
removed
}
pub fn get(&self, content_key: &str, now: u64) -> Vec<ProviderRecord> {
self.by_key
.get(content_key)
.map(|providers| {
providers
.values()
.filter(|r| !r.is_expired(now))
.cloned()
.collect()
})
.unwrap_or_default()
}
pub fn gc(&mut self, now: u64) -> usize {
let mut removed = 0;
self.by_key.retain(|_key, providers| {
let before = providers.len();
providers.retain(|_pid, r| !r.is_expired(now));
removed += before - providers.len();
!providers.is_empty()
});
removed
}
pub fn mark_announced(&mut self, content_key: String) {
self.announced.insert(content_key);
}
pub fn unmark_announced(&mut self, content_key: &str) -> bool {
self.announced.remove(content_key)
}
pub fn local_announcements(&self) -> Vec<String> {
self.announced.iter().cloned().collect()
}
pub fn len(&self) -> usize {
self.by_key.values().map(|p| p.len()).sum()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::key::Key;
use crate::record::CandidateAddr;
use dig_nat::PeerId;
fn rec(content: &Key, provider: u8, expires_at: u64) -> ProviderRecord {
ProviderRecord::new(
content,
&PeerId::from_bytes([provider; 32]),
vec![CandidateAddr::direct("h", 9444)],
expires_at,
)
}
#[test]
fn put_then_get_returns_live_record() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
s.put(rec(&key, 1, 100));
let got = s.get(&key.to_hex(), 50);
assert_eq!(got.len(), 1);
assert_eq!(
got[0].provider_peer_id,
PeerId::from_bytes([1u8; 32]).to_hex()
);
}
#[test]
fn get_hides_expired_records() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
s.put(rec(&key, 1, 100));
assert!(
s.get(&key.to_hex(), 100).is_empty(),
"expired at exactly TTL"
);
assert!(s.get(&key.to_hex(), 200).is_empty());
}
#[test]
fn same_provider_dedups_and_refreshes() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
s.put(rec(&key, 1, 100));
s.put(rec(&key, 1, 500)); assert_eq!(s.len(), 1, "same provider must not duplicate");
assert_eq!(s.get(&key.to_hex(), 300).len(), 1);
}
#[test]
fn distinct_providers_for_same_key_coexist() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
s.put(rec(&key, 1, 100));
s.put(rec(&key, 2, 100));
assert_eq!(s.get(&key.to_hex(), 50).len(), 2);
}
#[test]
fn put_returns_accepted_under_capacity() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
assert_eq!(s.put(rec(&key, 1, 100)), PutOutcome::Accepted);
}
#[test]
fn refreshing_same_provider_always_succeeds_even_at_per_key_cap() {
let mut s = ProviderStore::with_limits(ProviderStoreLimits {
max_providers_per_key: 1,
max_total_records: 1000,
});
let key = Key::from_bytes([0xAA; 32]);
assert_eq!(s.put(rec(&key, 1, 100)), PutOutcome::Accepted);
assert_eq!(s.put(rec(&key, 1, 999)), PutOutcome::Accepted, "refresh");
assert_eq!(s.len(), 1);
}
#[test]
fn per_key_cap_evicts_soonest_to_expire_to_make_room() {
let mut s = ProviderStore::with_limits(ProviderStoreLimits {
max_providers_per_key: 2,
max_total_records: 1000,
});
let key = Key::from_bytes([0xAA; 32]);
assert_eq!(s.put(rec(&key, 1, 100)), PutOutcome::Accepted); assert_eq!(s.put(rec(&key, 2, 500)), PutOutcome::Accepted);
assert_eq!(s.put(rec(&key, 3, 900)), PutOutcome::Accepted);
assert_eq!(
s.get(&key.to_hex(), 0).len(),
2,
"per-key cap must not be exceeded"
);
let ids: std::collections::HashSet<String> = s
.get(&key.to_hex(), 0)
.into_iter()
.map(|r| r.provider_peer_id)
.collect();
assert!(
!ids.contains(&PeerId::from_bytes([1u8; 32]).to_hex()),
"soonest-to-expire provider must be the one evicted"
);
}
#[test]
fn global_cap_rejects_new_content_keys_over_ceiling() {
let mut s = ProviderStore::with_limits(ProviderStoreLimits {
max_providers_per_key: 20,
max_total_records: 2,
});
let k1 = Key::from_bytes([0x01; 32]);
let k2 = Key::from_bytes([0x02; 32]);
let k3 = Key::from_bytes([0x03; 32]);
assert_eq!(s.put(rec(&k1, 1, 100)), PutOutcome::Accepted);
assert_eq!(s.put(rec(&k2, 1, 100)), PutOutcome::Accepted);
assert_eq!(
s.put(rec(&k3, 1, 100)),
PutOutcome::RejectedOverCapacity,
"third distinct record must be rejected once the global ceiling is hit"
);
assert_eq!(s.len(), 2, "rejected record must not be stored");
assert!(
s.get(&k3.to_hex(), 0).is_empty(),
"rejected key must not appear in the store at all"
);
}
#[test]
fn global_cap_does_not_evict_a_different_key_to_make_room() {
let mut s = ProviderStore::with_limits(ProviderStoreLimits {
max_providers_per_key: 20,
max_total_records: 1,
});
let legit = Key::from_bytes([0xAA; 32]);
s.put(rec(&legit, 1, 100));
let attacker_key = Key::from_bytes([0xBB; 32]);
assert_eq!(
s.put(rec(&attacker_key, 2, 100)),
PutOutcome::RejectedOverCapacity
);
assert_eq!(
s.get(&legit.to_hex(), 0).len(),
1,
"the legitimate key's record must survive"
);
}
#[test]
fn remove_deletes_only_the_named_provider_record() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
s.put(rec(&key, 1, 100));
s.put(rec(&key, 2, 100));
let pid1 = PeerId::from_bytes([1u8; 32]).to_hex();
let pid2 = PeerId::from_bytes([2u8; 32]).to_hex();
assert!(
s.remove(&key.to_hex(), &pid1),
"the named record was removed"
);
let survivors: std::collections::HashSet<String> = s
.get(&key.to_hex(), 0)
.into_iter()
.map(|r| r.provider_peer_id)
.collect();
assert_eq!(survivors.len(), 1, "the other provider must survive");
assert!(survivors.contains(&pid2));
assert!(!survivors.contains(&pid1));
}
#[test]
fn remove_of_absent_record_returns_false() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
s.put(rec(&key, 1, 100));
let absent = PeerId::from_bytes([9u8; 32]).to_hex();
assert!(!s.remove(&key.to_hex(), &absent), "no such provider");
assert!(!s.remove(&"00".repeat(32), &absent), "no such content key");
assert_eq!(s.len(), 1, "nothing removed");
}
#[test]
fn remove_drops_content_key_when_last_provider_leaves() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0xAA; 32]);
s.put(rec(&key, 1, 100));
let pid1 = PeerId::from_bytes([1u8; 32]).to_hex();
assert!(s.remove(&key.to_hex(), &pid1));
assert!(
s.is_empty(),
"the now-empty content key must be dropped entirely"
);
}
#[test]
fn gc_removes_expired_and_empty_keys() {
let mut s = ProviderStore::new();
let k1 = Key::from_bytes([0x01; 32]);
let k2 = Key::from_bytes([0x02; 32]);
s.put(rec(&k1, 1, 100)); s.put(rec(&k2, 1, 500)); let removed = s.gc(200);
assert_eq!(removed, 1);
assert!(s.get(&k1.to_hex(), 200).is_empty());
assert_eq!(s.get(&k2.to_hex(), 200).len(), 1);
}
#[test]
fn announcements_track_and_untrack() {
let mut s = ProviderStore::new();
let key = Key::from_bytes([0x07; 32]).to_hex();
s.mark_announced(key.clone());
s.mark_announced(key.clone()); assert_eq!(s.local_announcements(), vec![key.clone()]);
assert!(s.unmark_announced(&key));
assert!(!s.unmark_announced(&key));
assert!(s.local_announcements().is_empty());
}
}