Skip to content
Open
Show file tree
Hide file tree
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
150 changes: 133 additions & 17 deletions crates/blockchain/src/key_manager.rs
Original file line number Diff line number Diff line change
Expand Up @@ -169,6 +169,35 @@ impl KeyManager {
self.sign_with_attestation_key(validator_id, slot, &message_hash)
}

/// Signs `attestation_data` with each of `validator_ids`' attestation keys,
/// spread over at most one thread per core, and returns one result per id
/// in `validator_ids` order.
///
/// Every validator a node runs votes at the same instant, and one XMSS
/// signature takes tens of milliseconds, so signing them one after another
/// held a node's last vote back by most of a second, and the aggregators
/// waited on that last vote before their first proof covered the subnet.
/// The keys are independent and `sign` needs only `&self`, so the
/// signatures run concurrently instead.
pub fn sign_attestations(
&self,
validator_ids: &[u64],
attestation_data: &AttestationData,
) -> Vec<Result<XmssSignature, KeyManagerError>> {
let message_hash = attestation_data.hash_tree_root();
let slot = attestation_data.slot as u32;
let max_threads = thread::available_parallelism().map_or(1, NonZeroUsize::get);
// One key per batch at minimum: a signature costs far more than the
// thread spawn it would amortize.
map_batched(
validator_ids,
max_threads,
1,
"xmss-sign",
|&validator_id| self.sign_with_attestation_key(validator_id, slot, &message_hash),
)
}

/// Signs a block root using the validator's proposal key.
pub fn sign_block_root(
&mut self,
Expand All @@ -180,14 +209,14 @@ impl KeyManager {
}

fn sign_with_attestation_key(
&mut self,
&self,
validator_id: u64,
slot: u32,
message: &H256,
) -> Result<XmssSignature, KeyManagerError> {
let key_pair = self
.keys
.get_mut(&validator_id)
.get(&validator_id)
.ok_or(KeyManagerError::ValidatorKeyNotFound(validator_id))?;

// A slot outside the key's range can never be signed, however long the
Expand Down Expand Up @@ -285,41 +314,71 @@ fn prepare_key(validator_id: u64, role: KeyRole, key: &ValidatorSecretKey, slot:

/// Runs `f` on every item, in batches of at least [`MIN_KEYS_PER_WARM_THREAD`]
/// items spread over at most `max_threads` threads, and returns once all are
/// done.
/// done. See [`map_batched`].
fn for_each_batched<T: Sync>(items: &[T], max_threads: usize, f: impl Fn(&T) + Sync) {
map_batched(items, max_threads, MIN_KEYS_PER_WARM_THREAD, "xmss-warm", f);
}

/// Maps `f` over `items`, in batches of at least `min_batch` items spread over
/// at most `max_threads` threads named `thread_name`, and returns the results
/// in `items` order once all are done.
///
/// The calling thread takes the first batch, so a single batch spawns nothing.
/// A batch whose thread cannot be spawned also runs on the calling thread,
/// which costs latency instead of panicking the caller.
fn for_each_batched<T: Sync>(items: &[T], max_threads: usize, f: impl Fn(&T) + Sync) {
fn map_batched<T: Sync, R: Send>(
items: &[T],
max_threads: usize,
min_batch: usize,
thread_name: &str,
f: impl Fn(&T) -> R + Sync,
) -> Vec<R> {
let batch_len = items
.len()
.div_ceil(max_threads.max(1))
.max(MIN_KEYS_PER_WARM_THREAD);
.max(min_batch.max(1));
let mut batches = items.chunks(batch_len);
let Some(first) = batches.next() else {
return;
return Vec::new();
};
let f = &f;
let run = move |batch: &[T]| batch.iter().map(f).collect::<Vec<R>>();
thread::scope(|scope| {
let mut inline = vec![first];
for batch in batches {
let spawned = thread::Builder::new()
.name("xmss-warm".into())
.spawn_scoped(scope, move || batch.iter().for_each(f));
if let Err(err) = spawned {
warn!(%err, "Failed to spawn an XMSS warm thread, warming its keys inline");
inline.push(batch);
// Spawned before the first batch runs, so every batch is in flight
// while the calling thread works through its own.
let rest: Vec<_> = batches
.map(|batch| {
thread::Builder::new()
.name(thread_name.into())
.spawn_scoped(scope, move || run(batch))
.map_err(|err| {
warn!(%err, thread_name, "Failed to spawn a key thread, running its batch inline");
batch
})
})
.collect();
let mut results = run(first);
for batch in rest {
match batch {
Ok(handle) => results.extend(
handle
.join()
.unwrap_or_else(|panic| std::panic::resume_unwind(panic)),
),
Err(batch) => results.extend(run(batch)),
}
}
inline.into_iter().flatten().for_each(f);
});
results
})
}

#[cfg(test)]
mod tests {
use std::collections::HashSet;
use std::sync::Mutex;

use ethlambda_types::checkpoint::Checkpoint;

use super::*;

#[test]
Expand All @@ -332,7 +391,7 @@ mod tests {
#[test]
fn test_sign_attestation_validator_not_found() {
let keys = HashMap::new();
let mut key_manager = KeyManager::new(keys);
let key_manager = KeyManager::new(keys);
let message = H256::default();

let result = key_manager.sign_with_attestation_key(123, 0, &message);
Expand Down Expand Up @@ -443,6 +502,63 @@ mod tests {
assert_eq!(threads_used(3 * batch, 64).len(), 3);
}

#[test]
fn map_batched_returns_results_in_item_order() {
for len in [0, 1, 2, 7, 40] {
for max_threads in [0, 1, 3, 64] {
for min_batch in [0, 1, MIN_KEYS_PER_WARM_THREAD] {
let items: Vec<usize> = (0..len).collect();
let doubled = map_batched(&items, max_threads, min_batch, "test", |&i| 2 * i);
let expected: Vec<usize> = items.iter().map(|i| 2 * i).collect();
assert_eq!(
doubled, expected,
"len {len}, max_threads {max_threads}, min_batch {min_batch}"
);
}
}
}
}

#[test]
fn map_batched_with_unit_batches_uses_one_thread_per_item() {
let items: Vec<usize> = (0..5).collect();
let threads: HashSet<_> = map_batched(&items, 64, 1, "test", |_| thread::current().id())
.into_iter()
.collect();
assert_eq!(threads.len(), items.len());
}

#[test]
fn sign_attestations_signs_each_id_in_order() {
let count = 5;
let key_manager = tiny_key_manager(count);
let data = AttestationData {
slot: 1,
head: Checkpoint::default(),
target: Checkpoint::default(),
source: Checkpoint::default(),
};
let message = data.hash_tree_root();
// An unknown id in the middle keeps its place and reports itself.
let ids = [3, 0, 99, 4, 1, 2];

let results = key_manager.sign_attestations(&ids, &data);

assert_eq!(results.len(), ids.len());
for (&id, result) in ids.iter().zip(&results) {
let Some(pair) = key_manager.keys.get(&id) else {
assert!(
matches!(result, Err(KeyManagerError::ValidatorKeyNotFound(missing)) if *missing == id),
"id {id}: {result:?}"
);
continue;
};
let signature = ValidatorSignature::from_bytes(result.as_ref().unwrap()).unwrap();
let pubkey = pair.attestation_key.public_key();
assert!(signature.is_valid(&pubkey, 1, &message), "id {id}");
}
}

#[test]
fn prepare_keys_for_leaves_every_key_signable() {
// More keys than one batch holds, so the warm spans several threads.
Expand Down
21 changes: 11 additions & 10 deletions crates/blockchain/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -736,16 +736,17 @@ impl BlockChainServer {
// Produce attestation data once for all validators
let attestation_data = store::produce_attestation_data(&self.store, slot);

// For each registered validator, produce and publish attestation
for validator_id in self.key_manager.validator_ids() {
// Sign the attestation
let Ok(signature) = self
.key_manager
.sign_attestation(validator_id, &attestation_data)
.inspect_err(
|err| error!(%slot, %validator_id, %err, "Failed to sign attestation"),
)
else {
// Sign every validator's attestation at once, then deliver and publish
// each. Signing them in parallel is what keeps a node's last vote from
// trailing its first by the sum of every signature before it.
let validator_ids = self.key_manager.validator_ids();
let signatures = self
.key_manager
.sign_attestations(&validator_ids, &attestation_data);
for (validator_id, signature) in validator_ids.into_iter().zip(signatures) {
let Ok(signature) = signature.inspect_err(
|err| error!(%slot, %validator_id, %err, "Failed to sign attestation"),
) else {
continue;
};

Expand Down
Loading