1use std::{
2 any::type_name,
3 collections::{BTreeMap, BTreeSet, HashMap},
4 fmt::Display,
5 hash::Hash,
6 mem,
7 ops::Deref,
8};
9
10use alloy::primitives::U256;
11use hotshot_types::{
12 data::{EpochNumber, ViewNumber},
13 epoch_membership::EpochMembershipCoordinator,
14 message::UpgradeLock,
15 simple_certificate::{
16 Certificate1, Certificate2, SimpleCertificate, Threshold, TimeoutCertificate2,
17 },
18 simple_vote::{HasEpoch, Voteable},
19 stake_table::StakeTableEntries,
20 traits::{node_implementation::NodeType, signature_key::SignatureKey},
21 vote::{Certificate, HasViewNumber},
22};
23use hotshot_utils::anytrace::Result;
24use tokio_util::task::JoinMap;
25use tracing::{error, warn};
26
27use crate::message::{EpochChangeMessage, Unchecked, Validated};
28
29#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
30pub struct ValidCert<C> {
31 cert: C,
32 epoch: EpochNumber,
33}
34
35impl<C> ValidCert<C> {
36 pub(crate) fn new(cert: C, epoch: EpochNumber) -> Self {
37 Self { cert, epoch }
38 }
39
40 pub fn cert(&self) -> &C {
41 &self.cert
42 }
43
44 pub fn epoch(&self) -> EpochNumber {
45 self.epoch
46 }
47
48 pub fn into_cert(self) -> C {
49 self.cert
50 }
51}
52
53impl<C> Deref for ValidCert<C> {
54 type Target = C;
55
56 fn deref(&self) -> &Self::Target {
57 self.cert()
58 }
59}
60
61impl<C: HasViewNumber> HasViewNumber for ValidCert<C> {
62 fn view_number(&self) -> ViewNumber {
63 self.cert.view_number()
64 }
65}
66
67pub trait Verifiable<T: NodeType>: HasViewNumber + HasEpoch + Sized {
68 type Key: Copy + Ord + Hash + Display + Send + Sync + 'static;
70
71 type Output: Send + 'static;
72
73 fn key(&self) -> Option<Self::Key>;
74
75 fn check(
76 self,
77 stake_table: &[<T::SignatureKey as SignatureKey>::StakeTableEntry],
78 threshold: U256,
79 upgrade_lock: &UpgradeLock<T>,
80 ) -> Result<Self::Output>;
81}
82
83impl<T, D, V> Verifiable<T> for SimpleCertificate<T, D, V>
84where
85 T: NodeType,
86 D: Voteable<T> + HasEpoch + 'static,
87 V: Threshold<T>,
88 Self: Certificate<T, D> + Send + 'static,
89{
90 type Key = ViewNumber;
91 type Output = Self;
92
93 fn key(&self) -> Option<ViewNumber> {
94 Some(self.view_number())
95 }
96
97 fn check(
98 self,
99 stake_table: &[<T::SignatureKey as SignatureKey>::StakeTableEntry],
100 threshold: U256,
101 upgrade_lock: &UpgradeLock<T>,
102 ) -> Result<Self> {
103 self.is_valid_cert(stake_table, threshold, upgrade_lock)?;
104 Ok(self)
105 }
106}
107
108impl<T: NodeType> Verifiable<T> for EpochChangeMessage<T, Unchecked> {
109 type Key = EpochNumber;
110 type Output = EpochChangeMessage<T, Validated>;
111
112 fn key(&self) -> Option<EpochNumber> {
113 self.epoch()
114 }
115
116 fn check(
117 self,
118 stake_table: &[<T::SignatureKey as SignatureKey>::StakeTableEntry],
119 threshold: U256,
120 upgrade_lock: &UpgradeLock<T>,
121 ) -> Result<Self::Output> {
122 self.cert1
123 .is_valid_cert(stake_table, threshold, upgrade_lock)?;
124 self.cert2
125 .is_valid_cert(stake_table, threshold, upgrade_lock)?;
126 Ok(self.into_validated())
127 }
128}
129
130pub struct CertVerifier<T: NodeType, C: Verifiable<T>> {
144 tasks: JoinMap<C::Key, Option<ValidCert<C::Output>>>,
145 pending_task: BTreeMap<C::Key, HashMap<T::SignatureKey, C>>,
146 pending_membership: BTreeMap<C::Key, HashMap<T::SignatureKey, C>>,
147 completed: BTreeSet<C::Key>,
148 lower_bound: Option<C::Key>,
149 membership: EpochMembershipCoordinator<T>,
150 upgrade_lock: UpgradeLock<T>,
151 invalid_certs: u64,
152}
153
154impl<T: NodeType, C: Verifiable<T> + Send + 'static> CertVerifier<T, C> {
155 pub fn new(membership: EpochMembershipCoordinator<T>, upgrade_lock: UpgradeLock<T>) -> Self {
156 Self {
157 tasks: JoinMap::new(),
158 pending_task: BTreeMap::new(),
159 pending_membership: BTreeMap::new(),
160 completed: BTreeSet::new(),
161 lower_bound: None,
162 membership,
163 upgrade_lock,
164 invalid_certs: 0,
165 }
166 }
167
168 pub fn verify(&mut self, sender: T::SignatureKey, cert: C) -> Option<EpochNumber> {
172 let Some(key) = cert.key() else {
173 warn!(cert = type_name::<C>(), "certificate has no key");
174 return None;
175 };
176
177 let Some(epoch) = cert.epoch() else {
178 warn!(%key, cert = type_name::<C>(), "certificate has no epoch number");
179 return None;
180 };
181
182 if self.is_stale(key) || self.completed.contains(&key) {
183 return None;
184 }
185
186 if let Some(senders) = self.pending_task.get(&key)
187 && senders.contains_key(&sender)
188 {
189 return None;
190 }
191
192 if let Some(senders) = self.pending_membership.get(&key)
193 && senders.contains_key(&sender)
194 {
195 return None;
196 }
197
198 if self.tasks.contains_key(&key) {
199 self.pending_task
200 .entry(key)
201 .or_default()
202 .insert(sender, cert);
203 return None;
204 }
205
206 let Ok(membership) = self.membership.membership_for_epoch(Some(epoch)) else {
207 self.pending_membership
208 .entry(key)
209 .or_default()
210 .insert(sender, cert);
211 return Some(epoch);
212 };
213
214 let lock = self.upgrade_lock.clone();
215
216 self.tasks.spawn_blocking(key, move || {
217 let entries = StakeTableEntries::from_iter(membership.stake_table()).0;
218 let threshold = membership.success_threshold();
219 match cert.check(&entries, threshold, &lock) {
220 Ok(valid) => Some(ValidCert::new(valid, epoch)),
221 Err(err) => {
222 warn!(%key, %epoch, %err, cert = type_name::<C>(), "invalid certificate");
223 None
224 },
225 }
226 });
227
228 None
229 }
230
231 pub fn mark_completed(&mut self, key: C::Key) {
235 if self.is_stale(key) {
236 return;
237 }
238 self.completed.insert(key);
239 self.pending_task.remove(&key);
240 self.pending_membership.remove(&key);
241 self.tasks.abort(&key);
242 }
243
244 pub fn retry_pending(&mut self) -> Vec<EpochNumber> {
249 mem::take(&mut self.pending_membership)
250 .into_values()
251 .flatten()
252 .filter_map(|(sender, cert)| self.verify(sender, cert))
253 .collect()
254 }
255
256 pub async fn next(&mut self) -> Option<ValidCert<C::Output>> {
257 loop {
258 match self.tasks.join_next().await? {
259 (key, Ok(Some(cert))) => {
260 if !self.is_stale(key) {
261 self.completed.insert(key);
262 self.pending_task.remove(&key);
263 self.pending_membership.remove(&key);
264 return Some(cert);
265 }
266 },
267 (key, Ok(None)) => {
268 self.invalid_certs += 1;
269 if !self.is_stale(key)
270 && let Some((sender, cert)) = self.next_pending_sender(key)
271 {
272 self.verify(sender, cert);
273 }
274 },
275 (key, Err(err)) => {
276 if err.is_panic() {
277 error!(%key, %err, cert = type_name::<C>(), "cert verification task panic");
278 }
279 if !self.is_stale(key)
280 && let Some((sender, cert)) = self.next_pending_sender(key)
281 {
282 self.verify(sender, cert);
283 }
284 },
285 }
286 }
287 }
288
289 pub fn gc(&mut self, key: C::Key) {
290 self.completed = self.completed.split_off(&key);
291 self.pending_task = self.pending_task.split_off(&key);
292 self.pending_membership = self.pending_membership.split_off(&key);
293 self.lower_bound = Some(key);
294 self.tasks.abort_matching(|k| *k < key);
295 }
296
297 pub fn num_invalid_certs(&self) -> u64 {
298 self.invalid_certs
299 }
300
301 fn next_pending_sender(&mut self, k: C::Key) -> Option<(T::SignatureKey, C)> {
302 let map = self.pending_task.get_mut(&k)?;
303 let sender = map.keys().next().cloned()?;
304 let cert = map.remove(&sender)?;
305 if map.is_empty() {
306 self.pending_task.remove(&k);
307 }
308 Some((sender, cert))
309 }
310
311 fn is_stale(&self, key: C::Key) -> bool {
312 self.lower_bound.is_some_and(|lb| key < lb)
313 }
314}
315
316pub struct CertBySenderVerifier<T: NodeType, C: Verifiable<T>> {
323 tasks: JoinMap<T::SignatureKey, Option<ValidCert<C::Output>>>,
324 pending: HashMap<T::SignatureKey, C>,
325 completed: BTreeSet<ViewNumber>,
326 lower_bound: ViewNumber,
327 membership: EpochMembershipCoordinator<T>,
328 upgrade_lock: UpgradeLock<T>,
329 invalid_certs: u64,
330}
331
332impl<T: NodeType, C: Verifiable<T> + Send + 'static> CertBySenderVerifier<T, C>
333where
334 C::Output: HasViewNumber,
335{
336 pub fn new(membership: EpochMembershipCoordinator<T>, upgrade_lock: UpgradeLock<T>) -> Self {
337 Self {
338 tasks: JoinMap::new(),
339 pending: HashMap::new(),
340 completed: BTreeSet::new(),
341 lower_bound: ViewNumber::genesis(),
342 membership,
343 upgrade_lock,
344 invalid_certs: 0,
345 }
346 }
347
348 pub fn verify(&mut self, sender: T::SignatureKey, cert: C) -> Option<EpochNumber> {
354 let view = cert.view_number();
355
356 if view < self.lower_bound
357 || self.completed.contains(&view)
358 || self.tasks.contains_key(&sender)
359 {
360 return None;
361 }
362
363 let Some(epoch) = cert.epoch() else {
364 warn!(%view, cert = type_name::<C>(), "received certificate has no epoch number");
365 return None;
366 };
367
368 let Ok(membership) = self.membership.membership_for_epoch(Some(epoch)) else {
369 self.pending.insert(sender, cert);
370 return Some(epoch);
371 };
372
373 let lock = self.upgrade_lock.clone();
374
375 self.tasks.spawn_blocking(sender, move || {
376 let entries = StakeTableEntries::from_iter(membership.stake_table()).0;
377 let threshold = membership.success_threshold();
378 match cert.check(&entries, threshold, &lock) {
379 Ok(valid) => Some(ValidCert::new(valid, epoch)),
380 Err(err) => {
381 warn!(%view, %epoch, %err, cert = type_name::<C>(), "invalid certificate");
382 None
383 },
384 }
385 });
386
387 None
388 }
389
390 pub fn mark_completed(&mut self, view: ViewNumber) {
394 if view < self.lower_bound {
395 return;
396 }
397 self.completed.insert(view);
398 self.pending.retain(|_, c| c.view_number() != view);
399 }
400
401 pub fn retry_pending(&mut self) -> Vec<EpochNumber> {
405 mem::take(&mut self.pending)
406 .into_iter()
407 .filter_map(|(sender, cert)| self.verify(sender, cert))
408 .collect()
409 }
410
411 pub async fn next(&mut self) -> Option<ValidCert<C::Output>> {
412 loop {
413 match self.tasks.join_next().await? {
414 (_, Ok(Some(cert))) => {
415 let view = cert.view_number();
416 if view >= self.lower_bound && self.completed.insert(view) {
417 return Some(cert);
418 }
419 },
420 (_, Ok(None)) => {
421 self.invalid_certs += 1;
422 },
423 (sender, Err(err)) => {
424 if err.is_panic() {
425 error!(?sender, %err, cert = type_name::<C>(), "cert verification task panic");
426 }
427 },
428 }
429 }
430 }
431
432 pub fn gc(&mut self, view: ViewNumber) {
433 self.completed = self.completed.split_off(&view);
434 self.pending.retain(|_, c| c.view_number() >= view);
435 self.lower_bound = view;
436 }
437
438 pub fn num_invalid_certs(&self) -> u64 {
439 self.invalid_certs
440 }
441}
442
443pub struct CertVerifiers<T: NodeType> {
445 pub cert1: CertVerifier<T, Certificate1<T>>,
446 pub cert2: CertVerifier<T, Certificate2<T>>,
447 pub timeout: CertBySenderVerifier<T, TimeoutCertificate2<T>>,
448 pub advance: CertBySenderVerifier<T, Certificate1<T>>,
449 pub epoch_change: CertVerifier<T, EpochChangeMessage<T, Unchecked>>,
450}
451
452impl<T: NodeType> CertVerifiers<T> {
453 pub fn new(membership: EpochMembershipCoordinator<T>, upgrade_lock: UpgradeLock<T>) -> Self {
454 Self {
455 cert1: CertVerifier::new(membership.clone(), upgrade_lock.clone()),
456 cert2: CertVerifier::new(membership.clone(), upgrade_lock.clone()),
457 timeout: CertBySenderVerifier::new(membership.clone(), upgrade_lock.clone()),
458 advance: CertBySenderVerifier::new(membership.clone(), upgrade_lock.clone()),
459 epoch_change: CertVerifier::new(membership, upgrade_lock),
460 }
461 }
462
463 pub fn retry_pending<F>(&mut self, mut request: F)
464 where
465 F: FnMut(EpochNumber),
466 {
467 for epoch in self.cert1.retry_pending() {
468 request(epoch);
469 }
470 for epoch in self.cert2.retry_pending() {
471 request(epoch);
472 }
473 for epoch in self.timeout.retry_pending() {
474 request(epoch);
475 }
476 for epoch in self.advance.retry_pending() {
477 request(epoch);
478 }
479 for epoch in self.epoch_change.retry_pending() {
480 request(epoch);
481 }
482 }
483
484 pub fn gc(&mut self, view: ViewNumber, epoch: EpochNumber) {
485 self.cert1.gc(view);
486 self.cert2.gc(view);
487 self.timeout.gc(view);
488 self.advance.gc(view);
489 self.epoch_change.gc(epoch);
490 }
491
492 pub fn num_invalid_certs(&self) -> u64 {
493 self.cert1
494 .num_invalid_certs()
495 .saturating_add(self.cert2.num_invalid_certs())
496 .saturating_add(self.timeout.num_invalid_certs())
497 .saturating_add(self.advance.num_invalid_certs())
498 .saturating_add(self.epoch_change.num_invalid_certs())
499 }
500}