Skip to main content

de_mls/peer_scoring/
service.rs

1//! Reference [`PeerScoringService`] — a [`PeerScoringPlugin`] implementation
2//! over [`PeerScoreStorage`]. Threshold travels with [`ScoringConfig`];
3//! per-event score deltas are supplied at construction.
4
5use std::collections::HashMap;
6
7use crate::{
8    PeerScoreStorage, PeerScoringPlugin, ScoreEvent, ScoreOp, ScoreSnapshot, ScoringConfig,
9};
10
11/// Per-conversation, per-member score tracker. Reference [`PeerScoringPlugin`]
12/// implementation. One instance per conversation; threshold travels with
13/// [`ScoringConfig`]. Storage is abstracted via [`PeerScoreStorage`] so
14/// app-layer backends (in-memory, on-disk, …) plug in without touching
15/// this protocol logic.
16pub struct PeerScoringService<S: PeerScoreStorage> {
17    storage: S,
18    score_deltas: HashMap<ScoreEvent, i64>,
19    config: ScoringConfig,
20}
21
22impl<S: PeerScoreStorage> PeerScoringService<S> {
23    pub fn new(storage: S, score_deltas: HashMap<ScoreEvent, i64>, config: ScoringConfig) -> Self {
24        Self {
25            storage,
26            score_deltas,
27            config,
28        }
29    }
30
31    /// Signed score delta for `event`. Events not in the table contribute 0.
32    fn score_delta(&self, event: ScoreEvent) -> i64 {
33        self.score_deltas.get(&event).copied().unwrap_or(0)
34    }
35}
36
37impl<S: PeerScoreStorage> PeerScoringPlugin for PeerScoringService<S> {
38    fn add_member(&mut self, member_id: &[u8]) -> bool {
39        let default = self.config.default_score;
40        self.storage.set(member_id, default);
41        // "Untracked → tracked" treated as "above → new state": an unusual
42        // config with `default_score <= threshold` surfaces the new member
43        // as a downward cross. The standard config (default 100, threshold
44        // 0) returns false.
45        default <= self.config.threshold
46    }
47
48    fn remove_member(&mut self, member_id: &[u8]) {
49        self.storage.remove(member_id);
50    }
51
52    /// Apply an incremental delta to an already-tracked member. Unlike
53    /// `add_member` / `apply_snapshot`, this never creates an entry — a
54    /// stale op must not resurrect a removed member. The coordinator
55    /// `add_member`s first; a drop means roster and scores are out of sync.
56    fn apply_op(&mut self, op: &ScoreOp) -> bool {
57        let Some(current) = self.storage.get(&op.member_id) else {
58            tracing::debug!(
59                member = ?op.member_id,
60                event = ?op.event,
61                "score op dropped: member not tracked (add_member first)"
62            );
63            return false;
64        };
65        let delta = self.score_delta(op.event);
66        let new_score = current.saturating_add(delta);
67        self.storage.set(&op.member_id, new_score);
68        crossed_down(Some(current), new_score, self.config.threshold)
69    }
70
71    fn apply_snapshot(&mut self, snapshot: &ScoreSnapshot) -> bool {
72        let threshold = self.config.threshold;
73        let mut crossed = false;
74        for (member_id, new_score) in &snapshot.diverged {
75            let prior = self.storage.get(member_id);
76            self.storage.set(member_id, *new_score);
77            crossed |= crossed_down(prior, *new_score, threshold);
78        }
79        crossed
80    }
81
82    fn snapshot(&self) -> ScoreSnapshot {
83        let default = self.config.default_score;
84        let diverged = self
85            .storage
86            .all_scores()
87            .into_iter()
88            .filter(|(_, score)| *score != default)
89            .collect();
90        ScoreSnapshot { diverged }
91    }
92
93    fn score_for(&self, member_id: &[u8]) -> Option<i64> {
94        self.storage.get(member_id)
95    }
96
97    fn members_below_threshold(&self) -> Vec<Vec<u8>> {
98        let threshold = self.config.threshold;
99        self.storage
100            .all_scores()
101            .into_iter()
102            .filter(|(_, score)| *score <= threshold)
103            .map(|(id, _)| id)
104            .collect()
105    }
106
107    fn all_members_with_scores(&self) -> Vec<(Vec<u8>, i64)> {
108        self.storage.all_scores()
109    }
110
111    fn threshold(&self) -> i64 {
112        self.config.threshold
113    }
114
115    fn set_threshold(&mut self, threshold: i64) {
116        self.config.threshold = threshold;
117    }
118
119    fn default_score(&self) -> i64 {
120        self.config.default_score
121    }
122}
123
124/// `true` when `new_score` crosses a member *down* to at-or-below
125/// `threshold` from above. `prior == None` (untracked) counts as "above",
126/// so a fresh entry landing at-or-below threshold is a downward cross.
127/// Upward recovery is not surfaced — no coordinator consumes it.
128fn crossed_down(prior: Option<i64>, new_score: i64, threshold: i64) -> bool {
129    let was_above = prior.is_none_or(|p| p > threshold);
130    let now_below = new_score <= threshold;
131    was_above && now_below
132}
133
134#[cfg(test)]
135mod tests {
136    use std::collections::HashMap;
137
138    use super::*;
139
140    // ── Test scaffolding ────────────────────────────────────────────
141
142    /// Minimal in-memory storage for service tests. Production storage
143    /// lives in [`crate::InMemoryPeerScoreStorage`].
144    #[derive(Default)]
145    struct TestStorage(HashMap<Vec<u8>, i64>);
146
147    impl PeerScoreStorage for TestStorage {
148        fn get(&self, member_id: &[u8]) -> Option<i64> {
149            self.0.get(member_id).copied()
150        }
151        fn set(&mut self, member_id: &[u8], score: i64) {
152            self.0.insert(member_id.to_vec(), score);
153        }
154        fn remove(&mut self, member_id: &[u8]) {
155            self.0.remove(member_id);
156        }
157        fn all_scores(&self) -> Vec<(Vec<u8>, i64)> {
158            self.0.iter().map(|(k, v)| (k.clone(), *v)).collect()
159        }
160    }
161
162    fn make_service() -> PeerScoringService<TestStorage> {
163        let deltas = HashMap::from([
164            (ScoreEvent::EmergencyNoCreator, -50),
165            (ScoreEvent::EmergencyYesCreator, 20),
166            (ScoreEvent::BrokenCommit, -50),
167            (ScoreEvent::SuccessfulCommit, 10),
168            (ScoreEvent::MisbehavingCommit, -30),
169        ]);
170        PeerScoringService::new(
171            TestStorage::default(),
172            deltas,
173            ScoringConfig {
174                default_score: 100,
175                threshold: 0,
176            },
177        )
178    }
179
180    #[test]
181    fn add_member_gets_default_score() {
182        let mut svc = make_service();
183        let crossed = svc.add_member(b"alice");
184        assert!(!crossed, "default 100 > threshold 0, no cross");
185        assert_eq!(svc.score_for(b"alice"), Some(100));
186    }
187
188    #[test]
189    fn add_member_below_threshold_crosses_down() {
190        let mut svc = PeerScoringService::new(
191            TestStorage::default(),
192            HashMap::new(),
193            ScoringConfig {
194                default_score: -10,
195                threshold: 0,
196            },
197        );
198        assert!(
199            svc.add_member(b"alice"),
200            "default -10 <= threshold 0 crosses down"
201        );
202    }
203
204    #[test]
205    fn unknown_member_returns_none() {
206        let svc = make_service();
207        assert_eq!(svc.score_for(b"unknown"), None);
208    }
209
210    #[test]
211    fn remove_member_clears_score() {
212        let mut svc = make_service();
213        let _ = svc.add_member(b"alice");
214        svc.remove_member(b"alice");
215        assert_eq!(svc.score_for(b"alice"), None);
216    }
217
218    #[test]
219    fn apply_event_decreases_score() {
220        let mut svc = make_service();
221        let _ = svc.add_member(b"alice");
222        let crossed = svc.apply_op(&ScoreOp {
223            member_id: b"alice".to_vec(),
224            event: ScoreEvent::EmergencyNoCreator,
225        });
226        assert!(!crossed, "100 → 50 stays above threshold 0");
227        assert_eq!(svc.score_for(b"alice"), Some(50));
228    }
229
230    #[test]
231    fn apply_op_unknown_member_returns_false() {
232        let mut svc = make_service();
233        let crossed = svc.apply_op(&ScoreOp {
234            member_id: b"unknown".to_vec(),
235            event: ScoreEvent::EmergencyNoCreator,
236        });
237        assert!(!crossed);
238    }
239
240    #[test]
241    fn multiple_events_accumulate() {
242        let mut svc = make_service();
243        let _ = svc.add_member(b"alice");
244        for event in [
245            ScoreEvent::EmergencyNoCreator,
246            ScoreEvent::MisbehavingCommit,
247            ScoreEvent::SuccessfulCommit,
248        ] {
249            let _ = svc.apply_op(&ScoreOp {
250                member_id: b"alice".to_vec(),
251                event,
252            });
253        }
254        assert_eq!(svc.score_for(b"alice"), Some(30));
255    }
256
257    /// Recovery back above threshold emits no event — only downward
258    /// crosses are surfaced (no coordinator consumes upward recovery).
259    #[test]
260    fn apply_op_recovery_above_threshold_no_cross() {
261        let mut svc = make_service();
262        let _ = svc.add_member(b"alice");
263        // 100 → 50 → 0 (down emitted at the 0 cross), 0 → 20 (recovery).
264        for _ in 0..2 {
265            let _ = svc.apply_op(&ScoreOp {
266                member_id: b"alice".to_vec(),
267                event: ScoreEvent::EmergencyNoCreator,
268            });
269        }
270        let crossed = svc.apply_op(&ScoreOp {
271            member_id: b"alice".to_vec(),
272            event: ScoreEvent::EmergencyYesCreator,
273        });
274        assert!(!crossed, "upward recovery cross is not surfaced");
275    }
276
277    #[test]
278    fn apply_op_crosses_down_returns_true() {
279        let mut svc = make_service();
280        let _ = svc.add_member(b"alice");
281
282        // 100 → 50, still above threshold 0.
283        let crossed = svc.apply_op(&ScoreOp {
284            member_id: b"alice".to_vec(),
285            event: ScoreEvent::EmergencyNoCreator,
286        });
287        assert!(!crossed, "above threshold, no cross");
288
289        // 50 → 0, crosses to at-or-below threshold.
290        let crossed = svc.apply_op(&ScoreOp {
291            member_id: b"alice".to_vec(),
292            event: ScoreEvent::EmergencyNoCreator,
293        });
294        assert!(crossed, "0 is at-or-below threshold");
295
296        // 0 → -50, already below — no further cross.
297        let crossed = svc.apply_op(&ScoreOp {
298            member_id: b"alice".to_vec(),
299            event: ScoreEvent::BrokenCommit,
300        });
301        assert!(!crossed, "already below threshold, no cross");
302    }
303
304    #[test]
305    fn apply_ops_true_when_a_member_crosses_and_applies_all() {
306        let mut svc = make_service();
307        let _ = svc.add_member(b"alice");
308        let _ = svc.add_member(b"bob");
309        // Drop alice across threshold (-50 + -50 = 0 ≤ 0) and bob too in
310        // the same batch.
311        let ops = vec![
312            ScoreOp {
313                member_id: b"alice".to_vec(),
314                event: ScoreEvent::BrokenCommit,
315            },
316            ScoreOp {
317                member_id: b"alice".to_vec(),
318                event: ScoreEvent::BrokenCommit,
319            },
320            ScoreOp {
321                member_id: b"bob".to_vec(),
322                event: ScoreEvent::BrokenCommit,
323            },
324            ScoreOp {
325                member_id: b"bob".to_vec(),
326                event: ScoreEvent::BrokenCommit,
327            },
328        ];
329        assert!(svc.apply_ops(&ops), "a member crossed down in the batch");
330        // Every op applied (not short-circuited): both land at/below threshold.
331        let below = svc.members_below_threshold();
332        assert!(below.contains(&b"alice".to_vec()));
333        assert!(below.contains(&b"bob".to_vec()));
334    }
335
336    #[test]
337    fn snapshot_includes_only_diverged_scores() {
338        let mut svc = make_service();
339        let _ = svc.add_member(b"alice");
340        let _ = svc.add_member(b"bob");
341        let _ = svc.add_member(b"charlie");
342        let _ = svc.apply_op(&ScoreOp {
343            member_id: b"alice".to_vec(),
344            event: ScoreEvent::SuccessfulCommit,
345        });
346        let snap = svc.snapshot();
347        let ids: Vec<&[u8]> = snap.diverged.iter().map(|(id, _)| id.as_slice()).collect();
348        assert_eq!(ids, vec![b"alice".as_slice()]);
349        assert_eq!(snap.diverged[0].1, 110);
350    }
351
352    #[test]
353    fn apply_snapshot_crosses_only_on_actual_cross() {
354        let mut svc = make_service();
355        let _ = svc.add_member(b"alice");
356        let _ = svc.add_member(b"bob");
357        let snap = ScoreSnapshot {
358            diverged: vec![(b"alice".to_vec(), -10), (b"bob".to_vec(), 50)],
359        };
360        assert!(
361            svc.apply_snapshot(&snap),
362            "alice crosses down, bob stays above"
363        );
364        assert_eq!(svc.score_for(b"alice"), Some(-10));
365        assert_eq!(svc.score_for(b"bob"), Some(50));
366    }
367
368    #[test]
369    fn apply_snapshot_idempotent_on_repeat() {
370        let mut svc = make_service();
371        let _ = svc.add_member(b"alice");
372        let snap = ScoreSnapshot {
373            diverged: vec![(b"alice".to_vec(), -10)],
374        };
375        assert!(svc.apply_snapshot(&snap), "first apply crosses down");
376        assert!(
377            !svc.apply_snapshot(&snap),
378            "second apply on unchanged state does not cross"
379        );
380    }
381
382    /// A snapshot moving a member back above threshold emits no event —
383    /// only downward crosses are surfaced.
384    #[test]
385    fn apply_snapshot_recovery_above_threshold_no_cross() {
386        let mut svc = make_service();
387        let _ = svc.add_member(b"alice");
388        // First push alice below threshold.
389        let _ = svc.apply_snapshot(&ScoreSnapshot {
390            diverged: vec![(b"alice".to_vec(), -10)],
391        });
392        // Now snapshot her back above.
393        let crossed = svc.apply_snapshot(&ScoreSnapshot {
394            diverged: vec![(b"alice".to_vec(), 50)],
395        });
396        assert!(!crossed, "upward recovery cross is not surfaced");
397    }
398
399    #[test]
400    fn apply_snapshot_untracked_below_threshold_crosses_down() {
401        // Untracked → tracked (below threshold) treated as a downward
402        // cross from "above-by-default."
403        let mut svc = make_service();
404        assert!(
405            svc.apply_snapshot(&ScoreSnapshot {
406                diverged: vec![(b"newcomer".to_vec(), -10)],
407            }),
408            "untracked entry below threshold crosses down"
409        );
410    }
411
412    #[test]
413    fn members_below_threshold_filters_correctly() {
414        let mut svc = make_service();
415        let _ = svc.add_member(b"alice");
416        let _ = svc.add_member(b"bob");
417        let _ = svc.add_member(b"charlie");
418        for event in [ScoreEvent::EmergencyNoCreator, ScoreEvent::BrokenCommit] {
419            let _ = svc.apply_op(&ScoreOp {
420                member_id: b"alice".to_vec(),
421                event,
422            });
423        }
424        for _ in 0..2 {
425            let _ = svc.apply_op(&ScoreOp {
426                member_id: b"charlie".to_vec(),
427                event: ScoreEvent::EmergencyNoCreator,
428            });
429        }
430        let below = svc.members_below_threshold();
431        assert!(below.contains(&b"alice".to_vec()));
432        assert!(below.contains(&b"charlie".to_vec()));
433        assert!(!below.contains(&b"bob".to_vec()));
434    }
435
436    #[test]
437    fn set_threshold_changes_below_threshold_set() {
438        let mut svc = make_service();
439        let _ = svc.add_member(b"alice");
440        // Apply via snapshot to set an absolute score without going
441        // through the delta table.
442        let _ = svc.apply_snapshot(&ScoreSnapshot {
443            diverged: vec![(b"alice".to_vec(), -10)],
444        });
445
446        svc.set_threshold(-50);
447        assert!(!svc.members_below_threshold().contains(&b"alice".to_vec()));
448
449        svc.set_threshold(-5);
450        assert!(svc.members_below_threshold().contains(&b"alice".to_vec()));
451    }
452
453    #[test]
454    fn score_saturates_no_overflow() {
455        let mut svc = PeerScoringService::new(
456            TestStorage::default(),
457            HashMap::from([(ScoreEvent::SuccessfulCommit, i64::MAX)]),
458            ScoringConfig {
459                default_score: i64::MAX,
460                threshold: 0,
461            },
462        );
463        let _ = svc.add_member(b"alice");
464        let _ = svc.apply_op(&ScoreOp {
465            member_id: b"alice".to_vec(),
466            event: ScoreEvent::SuccessfulCommit,
467        });
468        assert_eq!(svc.score_for(b"alice"), Some(i64::MAX));
469    }
470
471    #[test]
472    fn unknown_event_yields_zero_delta() {
473        let mut svc = PeerScoringService::new(
474            TestStorage::default(),
475            HashMap::from([(ScoreEvent::EmergencyNoCreator, -50)]),
476            ScoringConfig {
477                default_score: 100,
478                threshold: 0,
479            },
480        );
481        let _ = svc.add_member(b"alice");
482        let _ = svc.apply_op(&ScoreOp {
483            member_id: b"alice".to_vec(),
484            event: ScoreEvent::SuccessfulCommit,
485        });
486        assert_eq!(svc.score_for(b"alice"), Some(100));
487    }
488}