use dig_nat::TraversalKind;
use crate::quality::PeerQuality;
use crate::registry::PeerEntry;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum PeerClass {
DirectPath,
RelayedPath,
Unknown,
}
impl PeerClass {
pub fn of(class: Option<TraversalKind>) -> Self {
match class {
Some(TraversalKind::Relayed) => PeerClass::RelayedPath,
Some(_) => PeerClass::DirectPath,
None => PeerClass::Unknown,
}
}
pub fn is_relayed(self) -> bool {
matches!(self, PeerClass::RelayedPath)
}
}
#[derive(Debug, Clone)]
pub struct SaturationModel {
direct: SaturationEstimate,
relayed: SaturationEstimate,
unknown: SaturationEstimate,
}
impl Default for SaturationModel {
fn default() -> Self {
SaturationModel {
direct: SaturationEstimate::seeded(4.0),
relayed: SaturationEstimate::seeded(2.0),
unknown: SaturationEstimate::seeded(3.0),
}
}
}
impl SaturationModel {
fn slot(&self, class: PeerClass) -> &SaturationEstimate {
match class {
PeerClass::DirectPath => &self.direct,
PeerClass::RelayedPath => &self.relayed,
PeerClass::Unknown => &self.unknown,
}
}
fn slot_mut(&mut self, class: PeerClass) -> &mut SaturationEstimate {
match class {
PeerClass::DirectPath => &mut self.direct,
PeerClass::RelayedPath => &mut self.relayed,
PeerClass::Unknown => &mut self.unknown,
}
}
pub fn saturation_point(&self, class: PeerClass) -> u32 {
self.slot(class).point()
}
pub fn observe(&mut self, class: PeerClass, in_flight_at_dispatch: u32, throughput_bps: f64) {
self.slot_mut(class)
.observe(in_flight_at_dispatch, throughput_bps);
}
}
#[derive(Debug, Clone)]
struct SaturationEstimate {
point: f64,
best_throughput: f64,
best_concurrency: f64,
samples: u64,
}
impl SaturationEstimate {
fn seeded(point: f64) -> Self {
SaturationEstimate {
point,
best_throughput: 0.0,
best_concurrency: 1.0,
samples: 0,
}
}
fn point(&self) -> u32 {
(self.point.round() as i64).max(1) as u32
}
fn observe(&mut self, in_flight_at_dispatch: u32, throughput_bps: f64) {
let conc = (in_flight_at_dispatch.max(1)) as f64;
self.samples = self.samples.saturating_add(1);
if throughput_bps <= 0.0 {
return;
}
let alpha = (1.0 / (self.samples as f64)).max(0.1);
if throughput_bps >= self.best_throughput {
self.best_throughput = throughput_bps;
self.best_concurrency = conc;
let target = conc.max(self.point);
self.point = (1.0 - alpha) * self.point + alpha * target;
} else if conc > self.best_concurrency && throughput_bps < 0.8 * self.best_throughput {
let target = self.best_concurrency;
self.point = (1.0 - alpha) * self.point + alpha * target;
}
self.point = self.point.clamp(1.0, 64.0);
}
}
#[derive(Debug, Clone, Default)]
pub struct RelayModel {
relayed_mean: Option<f64>,
direct_mean: Option<f64>,
relayed_samples: u64,
direct_samples: u64,
}
impl RelayModel {
pub fn observe(&mut self, relayed: bool, throughput_bps: f64) {
if throughput_bps <= 0.0 {
return;
}
if relayed {
self.relayed_samples += 1;
let a = (1.0 / self.relayed_samples as f64).max(0.1);
self.relayed_mean = Some(match self.relayed_mean {
None => throughput_bps,
Some(m) => (1.0 - a) * m + a * throughput_bps,
});
} else {
self.direct_samples += 1;
let a = (1.0 / self.direct_samples as f64).max(0.1);
self.direct_mean = Some(match self.direct_mean {
None => throughput_bps,
Some(m) => (1.0 - a) * m + a * throughput_bps,
});
}
}
pub fn penalty(&self) -> f64 {
const FLOOR: f64 = 0.25; match (self.relayed_mean, self.direct_mean) {
(Some(r), Some(d)) if d > 0.0 => (r / d).clamp(FLOOR, 1.0),
_ => 0.85,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct ScoredPeer {
pub effective_score: f64,
pub headroom: u32,
pub exploratory: bool,
pub tie_break: u64,
}
#[cfg(test)]
pub(crate) static SCORE_PEER_CALLS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
#[cfg(test)]
pub(crate) static SCORE_PEER_CALLS_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
pub fn score_peer(
entry: &PeerEntry,
saturation: &SaturationModel,
relay: &RelayModel,
exploration_bonus: f64,
) -> ScoredPeer {
#[cfg(test)]
SCORE_PEER_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let q = &entry.quality;
let class = PeerClass::of(entry.connection_class);
let sat_point = saturation.saturation_point(class);
let headroom = sat_point.saturating_sub(q.in_flight);
let tie_break = tie_break_salt(entry);
if q.is_cold() {
return ScoredPeer {
effective_score: exploration_bonus,
headroom: headroom.max(1),
exploratory: true,
tie_break,
};
}
if is_bad_source(q) {
return ScoredPeer {
effective_score: -1.0e12 + tie_break as f64 * 1e-6,
headroom: headroom.max(1),
exploratory: false,
tie_break,
};
}
let tput = q.throughput.value().unwrap_or(0.0);
let rel = q.reliability.rate().unwrap_or(0.5);
let conf = q.confidence();
let base = tput * rel;
let volatility = q.throughput.relative_volatility();
let saturation_pressure = if sat_point == 0 {
1.0
} else {
(q.in_flight as f64 / sat_point as f64).clamp(0.0, 1.0)
};
let tail_discount =
(1.0 - 0.5 * volatility) * (0.5 + 0.5 * rel) * (1.0 - 0.4 * saturation_pressure);
let confident_value = base * (0.3 + 0.7 * conf);
let relay_factor = if class.is_relayed() {
relay.penalty()
} else {
1.0
};
let effective_score = confident_value * tail_discount * relay_factor;
ScoredPeer {
effective_score,
headroom,
exploratory: false,
tie_break,
}
}
fn is_bad_source(q: &PeerQuality) -> bool {
q.reliability.hard_failures() > 0 && q.reliability.rate().unwrap_or(1.0) < 0.5
}
fn tie_break_salt(entry: &PeerEntry) -> u64 {
let b = entry.peer_id.as_bytes();
let mut acc = 0u64;
for &x in b.iter().take(8) {
acc = (acc << 8) | x as u64;
}
acc
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::Provenance;
use dig_nat::PeerId;
fn entry(b: u8) -> PeerEntry {
PeerEntry::cold(PeerId::from_bytes([b; 32]), Provenance::Dht, 0)
}
#[test]
fn peer_class_maps_relayed_vs_direct() {
assert_eq!(
PeerClass::of(Some(TraversalKind::Relayed)),
PeerClass::RelayedPath
);
assert_eq!(
PeerClass::of(Some(TraversalKind::Direct)),
PeerClass::DirectPath
);
assert_eq!(
PeerClass::of(Some(TraversalKind::HolePunch)),
PeerClass::DirectPath
);
assert_eq!(PeerClass::of(None), PeerClass::Unknown);
}
#[test]
fn faster_reliable_peer_scores_higher() {
let sat = SaturationModel::default();
let relay = RelayModel::default();
let mut fast = entry(1);
for _ in 0..10 {
fast.quality.observe_throughput(1000.0);
fast.quality.observe_result(true, false);
fast.quality.bump_samples();
}
let mut slow = entry(2);
for _ in 0..10 {
slow.quality.observe_throughput(100.0);
slow.quality.observe_result(true, false);
slow.quality.bump_samples();
}
let sf = score_peer(&fast, &sat, &relay, 0.0);
let ss = score_peer(&slow, &sat, &relay, 0.0);
assert!(sf.effective_score > ss.effective_score);
}
#[test]
fn cold_peer_is_exploratory_and_bounded() {
let sat = SaturationModel::default();
let relay = RelayModel::default();
let cold = entry(1);
let s = score_peer(&cold, &sat, &relay, 50.0);
assert!(s.exploratory);
assert_eq!(s.effective_score, 50.0);
assert!(s.headroom >= 1);
}
#[test]
fn bad_source_sinks_below_cold_peers() {
let sat = SaturationModel::default();
let relay = RelayModel::default();
let mut bad = entry(1);
for _ in 0..5 {
bad.quality.observe_throughput(1000.0);
bad.quality.observe_result(true, false);
bad.quality.bump_samples();
}
for _ in 0..8 {
bad.quality.observe_result(false, true);
bad.quality.bump_samples();
}
let sbad = score_peer(&bad, &sat, &relay, 0.0);
let cold = entry(2);
let scold = score_peer(&cold, &sat, &relay, 10.0);
assert!(
sbad.effective_score < scold.effective_score,
"a verification-failing source must rank below a cold peer"
);
}
#[test]
fn saturation_model_lowers_point_on_oversubscription_degradation() {
let mut m = SaturationModel::default();
for _ in 0..5 {
m.observe(PeerClass::DirectPath, 2, 1000.0);
}
let p_before = m.saturation_point(PeerClass::DirectPath);
for _ in 0..8 {
m.observe(PeerClass::DirectPath, 8, 200.0);
}
let p_after = m.saturation_point(PeerClass::DirectPath);
assert!(
p_after <= p_before,
"oversubscription with degraded throughput must not raise the saturation point (before {p_before}, after {p_after})"
);
assert!(p_after >= 1);
}
#[test]
fn relay_penalty_shrinks_when_relayed_measures_as_well_as_direct() {
let mut m = RelayModel::default();
for _ in 0..10 {
m.observe(false, 1000.0);
m.observe(true, 950.0);
}
assert!(
m.penalty() > 0.9,
"near-parity relayed links => small penalty"
);
let mut m2 = RelayModel::default();
for _ in 0..10 {
m2.observe(false, 1000.0);
m2.observe(true, 300.0);
}
assert!(
m2.penalty() < 0.5,
"much-worse relayed links => large penalty"
);
}
}