1use std::collections::HashMap;
6
7use crate::{
8 PeerScoreStorage, PeerScoringPlugin, ScoreEvent, ScoreOp, ScoreSnapshot, ScoringConfig,
9};
10
11pub 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 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 default <= self.config.threshold
46 }
47
48 fn remove_member(&mut self, member_id: &[u8]) {
49 self.storage.remove(member_id);
50 }
51
52 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
124fn 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 #[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 #[test]
260 fn apply_op_recovery_above_threshold_no_cross() {
261 let mut svc = make_service();
262 let _ = svc.add_member(b"alice");
263 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 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 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 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 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 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 #[test]
385 fn apply_snapshot_recovery_above_threshold_no_cross() {
386 let mut svc = make_service();
387 let _ = svc.add_member(b"alice");
388 let _ = svc.apply_snapshot(&ScoreSnapshot {
390 diverged: vec![(b"alice".to_vec(), -10)],
391 });
392 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 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 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}