Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
66 changes: 49 additions & 17 deletions crates/types/src/validator_set.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ use crate::signing::PublicKey;
use crate::{Address, ArcContext};

/// A validator is a public key and voting power
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub struct Validator {
pub address: Address,
pub public_key: PublicKey,
Expand All @@ -42,18 +42,6 @@ impl Validator {
}
}

impl PartialOrd for Validator {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}

impl Ord for Validator {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.address.cmp(&other.address)
}
}

impl malachitebft_core_types::Validator<ArcContext> for Validator {
fn address(&self) -> &Address {
&self.address
Expand Down Expand Up @@ -81,6 +69,12 @@ impl ValidatorSet {

assert!(!validators.is_empty());

// Verify that total voting power does not overflow u64
validators
.iter()
.try_fold(0u64, |acc, v| acc.checked_add(v.voting_power))
.expect("total voting power overflow");

Self {
validators: Arc::new(validators),
}
Expand Down Expand Up @@ -134,7 +128,8 @@ impl ValidatorSet {
a.cmp(&b)
});

vals.dedup();
let mut seen = std::collections::HashSet::new();
vals.retain(|v| seen.insert(v.address));
}

pub fn get_keys(&self) -> Vec<PublicKey> {
Expand Down Expand Up @@ -187,7 +182,7 @@ mod tests {

#[test]
#[should_panic(expected = "total voting power overflow")]
fn total_voting_power_overflow_panics() {
fn total_voting_power_overflow_panics_on_construction() {
let mut rng = StdRng::seed_from_u64(0x42);

let sk1 = PrivateKey::generate(&mut rng);
Expand All @@ -196,7 +191,44 @@ mod tests {
let v1 = Validator::new(sk1.public_key(), u64::MAX);
let v2 = Validator::new(sk2.public_key(), 1);

let vs = ValidatorSet::new(vec![v1, v2]);
let _ = vs.total_voting_power();
// Panics immediately upon construction rather than deferring to consensus runtime
let _ = ValidatorSet::new(vec![v1, v2]);
}

#[test]
fn validator_ord_consistent_with_partialeq() {
let mut rng = StdRng::seed_from_u64(0x42);
let sk = PrivateKey::generate(&mut rng);

let v1 = Validator::new(sk.public_key(), 10);
let v2 = Validator::new(sk.public_key(), 20);
let v3 = Validator::new(sk.public_key(), 10);

// Consistent ordering and equality:
assert_ne!(v1.cmp(&v2), std::cmp::Ordering::Equal);
assert_ne!(v1, v2);

assert_eq!(v1.cmp(&v3), std::cmp::Ordering::Equal);
assert_eq!(v1, v3);
}

#[test]
fn sort_validators_deduplicates_by_address() {
let mut rng = StdRng::seed_from_u64(0x42);
let sk = PrivateKey::generate(&mut rng);

let v1 = Validator::new(sk.public_key(), 10);
let v2 = Validator::new(sk.public_key(), 20);

assert_eq!(v1.address, v2.address);

let vs = ValidatorSet::new(vec![v1.clone(), v2.clone()]);

// Retains highest voting power entry and deduplicates by address:
assert_eq!(vs.len(), 1);
assert_eq!(vs.total_voting_power(), 20);

let fetched = vs.get_by_address(&v1.address).unwrap();
assert_eq!(fetched.voting_power, 20);
}
}
Loading