1use crate::cx::Cx;
31use crate::error::{Error, ErrorKind, Result};
32use crate::time::timeout;
33use crate::types::{Outcome, Time};
34use serde::{Deserialize, Serialize};
35use std::collections::{HashMap, VecDeque};
36use std::sync::{Arc, Mutex};
37use std::time::Duration;
38
39use super::types::{
40 ConsensusBatch, ConsensusRequest, ConsensusResponse, MessageCertificate, MessageDigest,
41 PhaseKind, ReplicaId, SequenceNumber, ViewNumber,
42};
43
44#[derive(Debug, Clone)]
46pub struct PbftConfig {
47 pub replica_count: usize,
49 pub fault_tolerance: usize,
51 pub preprepare_timeout: Duration,
53 pub prepare_timeout: Duration,
55 pub commit_timeout: Duration,
57 pub view_change_timeout: Duration,
59 pub max_batch_size: usize,
61 pub batch_timeout: Duration,
63}
64
65impl PbftConfig {
66 pub fn new(replica_count: usize, fault_tolerance: usize) -> Result<Self> {
68 if replica_count < 3 * fault_tolerance + 1 {
69 return Err(Error::new(ErrorKind::InvalidInput));
70 }
71
72 Ok(Self {
73 replica_count,
74 fault_tolerance,
75 preprepare_timeout: Duration::from_secs(5),
76 prepare_timeout: Duration::from_secs(5),
77 commit_timeout: Duration::from_secs(5),
78 view_change_timeout: Duration::from_secs(10),
79 max_batch_size: 100,
80 batch_timeout: Duration::from_millis(10),
81 })
82 }
83
84 pub fn is_valid(&self) -> bool {
86 self.replica_count > 3 * self.fault_tolerance
87 }
88
89 pub fn quorum_size(&self) -> usize {
91 2 * self.fault_tolerance + 1
92 }
93}
94
95#[derive(Debug, Clone, Serialize, Deserialize)]
97pub enum PbftMessage {
98 Request(ConsensusRequest),
100 PrePrepare {
102 view: ViewNumber,
103 sequence: SequenceNumber,
104 digest: MessageDigest,
105 batch: ConsensusBatch,
106 replica_id: ReplicaId,
107 },
108 Prepare {
110 view: ViewNumber,
111 sequence: SequenceNumber,
112 digest: MessageDigest,
113 replica_id: ReplicaId,
114 },
115 Commit {
117 view: ViewNumber,
118 sequence: SequenceNumber,
119 digest: MessageDigest,
120 replica_id: ReplicaId,
121 },
122 ViewChange {
124 new_view: ViewNumber,
125 replica_id: ReplicaId,
126 certificates: Vec<MessageCertificate>,
127 },
128 NewView {
130 view: ViewNumber,
131 view_change_msgs: Vec<PbftMessage>,
132 preprepare_msgs: Vec<PbftMessage>,
133 },
134}
135
136impl PbftMessage {
137 pub fn digest(&self) -> Result<MessageDigest> {
139 MessageDigest::of(self)
140 }
141
142 pub fn phase(&self) -> PhaseKind {
144 match self {
145 PbftMessage::PrePrepare { .. } => PhaseKind::PrePrepare,
146 PbftMessage::Prepare { .. } => PhaseKind::Prepare,
147 PbftMessage::Commit { .. } => PhaseKind::Commit,
148 PbftMessage::ViewChange { .. } => PhaseKind::ViewChange,
149 PbftMessage::NewView { .. } => PhaseKind::NewView,
150 PbftMessage::Request(_) => PhaseKind::PrePrepare, }
152 }
153}
154
155#[derive(Debug, Clone)]
157pub struct PbftState {
158 pub view: ViewNumber,
160 pub sequence: SequenceNumber,
162 pub log: HashMap<SequenceNumber, LogEntry>,
164 pub pending_requests: VecDeque<ConsensusRequest>,
166 pub last_executed: SequenceNumber,
168 pub view_change_state: Option<ViewChangeState>,
170}
171
172#[derive(Debug, Clone)]
174pub struct LogEntry {
175 pub batch: ConsensusBatch,
177 pub digest: MessageDigest,
179 pub view: ViewNumber,
181 pub preprepared: bool,
183 pub prepare_msgs: HashMap<ReplicaId, PbftMessage>,
185 pub commit_msgs: HashMap<ReplicaId, PbftMessage>,
187 pub result: Option<Outcome<Vec<u8>, String>>,
189}
190
191#[derive(Debug, Clone)]
193pub struct ViewChangeState {
194 pub target_view: ViewNumber,
196 pub view_change_msgs: HashMap<ReplicaId, PbftMessage>,
198 pub sent_view_change: bool,
200 pub started_at: Time,
202}
203
204pub trait PbftTransport: Send + Sync {
206 fn send_to_replica(
208 &self,
209 replica_id: &ReplicaId,
210 message: PbftMessage,
211 ) -> impl std::future::Future<Output = Result<()>> + Send;
212
213 fn broadcast(
215 &self,
216 message: PbftMessage,
217 ) -> impl std::future::Future<Output = Result<()>> + Send;
218
219 fn receive(&self) -> impl std::future::Future<Output = Result<PbftMessage>> + Send;
221}
222
223pub struct PbftNode<T: PbftTransport> {
225 replica_id: ReplicaId,
227 replica_index: usize,
229 config: PbftConfig,
231 state: Arc<Mutex<PbftState>>,
233 transport: T,
235}
236
237impl<T: PbftTransport> PbftNode<T> {
238 pub fn new(replica_id: ReplicaId, config: PbftConfig, transport: T) -> Result<Self> {
240 if !config.is_valid() {
241 return Err(Error::new(ErrorKind::InvalidInput));
242 }
243 let replica_index = parse_replica_index(&replica_id, config.replica_count)?;
244
245 let state = PbftState {
246 view: ViewNumber::new(0),
247 sequence: SequenceNumber::new(1),
255 log: HashMap::new(),
256 pending_requests: VecDeque::new(),
257 last_executed: SequenceNumber::new(0),
258 view_change_state: None,
259 };
260
261 Ok(Self {
262 replica_id,
263 replica_index,
264 config,
265 state: Arc::new(Mutex::new(state)),
266 transport,
267 })
268 }
269
270 pub fn is_primary(&self) -> bool {
272 let state = self.state.lock().unwrap();
273 let primary_idx = state.view.primary(self.config.replica_count);
274 self.replica_index == primary_idx
275 }
276
277 pub fn last_executed(&self) -> SequenceNumber {
280 self.state.lock().unwrap().last_executed
281 }
282
283 pub async fn submit_request(&self, cx: &Cx, request: ConsensusRequest) -> Result<()> {
285 {
286 let mut state = self.state.lock().unwrap();
287 state.pending_requests.push_back(request);
288 }
289
290 if self.is_primary() {
292 self.try_create_batch(cx).await?;
293 }
294
295 Ok(())
296 }
297
298 async fn try_create_batch(&self, cx: &Cx) -> Result<()> {
300 let (batch, sequence, view) = {
301 let mut state = self.state.lock().unwrap();
302
303 if state.pending_requests.is_empty() {
304 return Ok(()); }
306
307 let mut requests = Vec::new();
309 while requests.len() < self.config.max_batch_size && !state.pending_requests.is_empty()
310 {
311 if let Some(request) = state.pending_requests.pop_front() {
312 requests.push(request);
313 }
314 }
315
316 let batch = ConsensusBatch::new(requests);
317 let sequence = state.sequence;
318 let view = state.view;
319
320 (batch, sequence, view)
321 };
322
323 let result = self
324 .send_preprepare(cx, view, sequence, batch.clone())
325 .await;
326 let mut state = self.state.lock().unwrap();
327 match result {
328 Ok(()) => {
329 if state.sequence == sequence {
330 state.sequence = state.sequence.next();
331 }
332 Ok(())
333 }
334 Err(err) => {
335 if state.sequence == sequence.next() {
336 state.sequence = sequence;
337 }
338 if let Ok(digest) = MessageDigest::of(&batch) {
339 if state
340 .log
341 .get(&sequence)
342 .is_some_and(|entry| entry.view == view && entry.digest == digest)
343 {
344 state.log.remove(&sequence);
345 }
346 }
347 for request in batch.requests.iter().rev() {
348 state.pending_requests.push_front(request.clone());
349 }
350 Err(err)
351 }
352 }
353 }
354
355 async fn send_preprepare(
357 &self,
358 _cx: &Cx,
359 view: ViewNumber,
360 sequence: SequenceNumber,
361 batch: ConsensusBatch,
362 ) -> Result<()> {
363 let digest = MessageDigest::of(&batch)?;
364
365 {
367 let mut state = self.state.lock().unwrap();
368 if state.log.contains_key(&sequence) {
369 return Err(
370 Error::new(ErrorKind::InvalidStateTransition).with_message(format!(
371 "PBFT pre-prepare sequence {sequence} already has a log entry"
372 )),
373 );
374 }
375 let entry = LogEntry {
376 batch: batch.clone(),
377 digest: digest.clone(),
378 view,
379 preprepared: true,
380 prepare_msgs: HashMap::new(),
381 commit_msgs: HashMap::new(),
382 result: None,
383 };
384 state.log.insert(sequence, entry);
385 }
386
387 let message = PbftMessage::PrePrepare {
388 view,
389 sequence,
390 digest,
391 batch,
392 replica_id: self.replica_id.clone(),
393 };
394
395 timeout(
397 Time::from_millis(0),
398 self.config.preprepare_timeout,
399 self.transport.broadcast(message),
400 )
401 .await
402 .map_err(|_| Error::new(ErrorKind::DeadlineExceeded))?
403 }
404
405 pub async fn process_message(&self, cx: &Cx, message: PbftMessage) -> Result<()> {
407 match message {
408 PbftMessage::Request(request) => self.submit_request(cx, request).await,
409 PbftMessage::PrePrepare {
410 view,
411 sequence,
412 digest,
413 batch,
414 replica_id,
415 } => {
416 self.handle_preprepare(cx, view, sequence, digest, batch, replica_id)
417 .await
418 }
419 PbftMessage::Prepare {
420 view,
421 sequence,
422 digest,
423 replica_id,
424 } => {
425 self.handle_prepare(cx, view, sequence, digest, replica_id)
426 .await
427 }
428 PbftMessage::Commit {
429 view,
430 sequence,
431 digest,
432 replica_id,
433 } => {
434 self.handle_commit(cx, view, sequence, digest, replica_id)
435 .await
436 }
437 PbftMessage::ViewChange {
438 new_view,
439 replica_id,
440 certificates,
441 } => {
442 self.handle_view_change(cx, new_view, replica_id, certificates)
443 .await
444 }
445 PbftMessage::NewView {
446 view,
447 view_change_msgs,
448 preprepare_msgs,
449 } => {
450 self.handle_new_view(cx, view, view_change_msgs, preprepare_msgs)
451 .await
452 }
453 }
454 }
455
456 async fn handle_preprepare(
458 &self,
459 _cx: &Cx,
460 view: ViewNumber,
461 sequence: SequenceNumber,
462 digest: MessageDigest,
463 batch: ConsensusBatch,
464 replica_id: ReplicaId,
465 ) -> Result<()> {
466 self.validate_preprepare_primary(view, &replica_id)?;
467 {
469 let mut state = self.state.lock().unwrap();
470 if view != state.view {
471 return Err(Error::new(ErrorKind::InvalidInput));
472 }
473 if sequence <= state.last_executed {
474 return Err(
475 Error::new(ErrorKind::InvalidStateTransition).with_message(format!(
476 "PBFT pre-prepare sequence {sequence} is at or below executed watermark {}",
477 state.last_executed
478 )),
479 );
480 }
481
482 if let Some(entry) = state.log.get_mut(&sequence) {
483 if entry.view != view || entry.digest != digest {
484 return Err(Error::new(ErrorKind::InvalidStateTransition).with_message(
485 format!("PBFT pre-prepare equivocation for {sequence} in {view}"),
486 ));
487 }
488 if entry.preprepared {
489 return Ok(());
490 }
491 }
492 }
493
494 let computed_digest = MessageDigest::of(&batch)?;
496 if digest != computed_digest {
497 return Err(Error::new(ErrorKind::InvalidInput));
498 }
499
500 {
502 let mut state = self.state.lock().unwrap();
503 if let Some(entry) = state.log.get_mut(&sequence) {
504 entry.batch = batch;
505 entry.preprepared = true;
506 } else {
507 let entry = LogEntry {
508 batch,
509 digest: digest.clone(),
510 view,
511 preprepared: true,
512 prepare_msgs: HashMap::new(),
513 commit_msgs: HashMap::new(),
514 result: None,
515 };
516 state.log.insert(sequence, entry);
517 }
518 }
519
520 let prepare_msg = PbftMessage::Prepare {
522 view,
523 sequence,
524 digest,
525 replica_id: self.replica_id.clone(),
526 };
527
528 timeout(
529 Time::from_millis(0),
530 self.config.prepare_timeout,
531 self.transport.broadcast(prepare_msg),
532 )
533 .await
534 .map_err(|_| Error::new(ErrorKind::DeadlineExceeded))?
535 }
536
537 async fn handle_prepare(
539 &self,
540 _cx: &Cx,
541 view: ViewNumber,
542 sequence: SequenceNumber,
543 digest: MessageDigest,
544 replica_id: ReplicaId,
545 ) -> Result<()> {
546 self.validate_remote_replica(&replica_id)?;
547 let should_commit = {
548 let mut state = self.state.lock().unwrap();
549
550 let entry = match state.log.get_mut(&sequence) {
552 Some(entry) if entry.view == view && entry.digest == digest => entry,
553 _ => return Ok(()), };
555
556 let msg = PbftMessage::Prepare {
558 view,
559 sequence,
560 digest: digest.clone(),
561 replica_id: replica_id.clone(),
562 };
563 entry.prepare_msgs.insert(replica_id, msg);
564
565 entry.preprepared && entry.prepare_msgs.len() + 1 >= self.config.quorum_size()
567 };
568
569 if should_commit {
571 let commit_msg = PbftMessage::Commit {
572 view,
573 sequence,
574 digest,
575 replica_id: self.replica_id.clone(),
576 };
577
578 timeout(
579 Time::from_millis(0),
580 self.config.commit_timeout,
581 self.transport.broadcast(commit_msg),
582 )
583 .await
584 .map_err(|_| Error::new(ErrorKind::DeadlineExceeded))??;
585 }
586
587 Ok(())
588 }
589
590 async fn handle_commit(
592 &self,
593 _cx: &Cx,
594 view: ViewNumber,
595 sequence: SequenceNumber,
596 digest: MessageDigest,
597 replica_id: ReplicaId,
598 ) -> Result<()> {
599 self.validate_remote_replica(&replica_id)?;
600 let should_execute = {
601 let mut state = self.state.lock().unwrap();
602 let next_to_execute = state.last_executed.next();
603
604 let entry = match state.log.get_mut(&sequence) {
606 Some(entry) if entry.view == view && entry.digest == digest => entry,
607 _ => return Ok(()), };
609
610 let msg = PbftMessage::Commit {
612 view,
613 sequence,
614 digest: digest.clone(),
615 replica_id: replica_id.clone(),
616 };
617 entry.commit_msgs.insert(replica_id, msg);
618
619 let prepared =
620 entry.preprepared && entry.prepare_msgs.len() + 1 >= self.config.quorum_size();
621 let committed = entry.commit_msgs.len() + 1 >= self.config.quorum_size();
622
623 prepared && committed && sequence == next_to_execute && entry.result.is_none()
624 };
625
626 if should_execute {
633 self.execute_batch(sequence).await?;
634 while let Some(next) = self.next_executable_sequence() {
635 self.execute_batch(next).await?;
636 }
637 }
638
639 Ok(())
640 }
641
642 fn next_executable_sequence(&self) -> Option<SequenceNumber> {
648 let state = self.state.lock().unwrap();
649 let next = state.last_executed.next();
650 let entry = state.log.get(&next)?;
651 let prepared =
652 entry.preprepared && entry.prepare_msgs.len() + 1 >= self.config.quorum_size();
653 let committed = entry.commit_msgs.len() + 1 >= self.config.quorum_size();
654 (prepared && committed && entry.result.is_none()).then_some(next)
655 }
656
657 async fn execute_batch(&self, sequence: SequenceNumber) -> Result<()> {
659 let batch = {
660 let mut state = self.state.lock().unwrap();
661
662 if sequence != state.last_executed.next() {
663 return Ok(());
664 }
665
666 let batch = {
667 let entry = state.log.get_mut(&sequence).ok_or_else(|| {
668 Error::new(ErrorKind::InvalidStateTransition).with_message(format!(
669 "PBFT cannot execute missing log entry for {sequence}"
670 ))
671 })?;
672 if entry.result.is_some() {
673 return Ok(());
674 }
675 let batch = entry.batch.clone();
676
677 let result = Outcome::Ok(b"executed".to_vec());
679 entry.result = Some(result);
680 batch
681 };
682
683 state.last_executed = sequence;
684 batch
685 };
686
687 let batch_size = batch.len();
688
689 #[cfg(feature = "tracing-integration")]
692 tracing::info!(
693 replica_id = %self.replica_id,
694 sequence = %sequence,
695 batch_size,
696 "Executed consensus batch"
697 );
698 #[cfg(not(feature = "tracing-integration"))]
699 let _ = batch_size;
700
701 Ok(())
702 }
703
704 fn validate_remote_replica(&self, replica_id: &ReplicaId) -> Result<()> {
705 let index = parse_replica_index(replica_id, self.config.replica_count)?;
706 if index == self.replica_index {
707 return Err(Error::new(ErrorKind::InvalidInput).with_message(format!(
708 "PBFT rejected self-authored remote quorum message from {replica_id}"
709 )));
710 }
711 Ok(())
712 }
713
714 fn validate_preprepare_primary(&self, view: ViewNumber, replica_id: &ReplicaId) -> Result<()> {
715 let index = parse_replica_index(replica_id, self.config.replica_count)?;
716 let expected = view.primary(self.config.replica_count);
717 if index != expected {
718 return Err(Error::new(ErrorKind::InvalidInput).with_message(format!(
719 "PBFT rejected pre-prepare from {replica_id}; primary for {view} is replica:{expected}"
720 )));
721 }
722 Ok(())
723 }
724
725 async fn handle_view_change(
734 &self,
735 _cx: &Cx,
736 _new_view: ViewNumber,
737 _replica_id: ReplicaId,
738 _certificates: Vec<MessageCertificate>,
739 ) -> Result<()> {
740 Err(Error::new(ErrorKind::InvalidStateTransition).with_message(
741 "PBFT view-change is not implemented (experimental consensus; no Byzantine \
742 fault tolerance under primary failure)",
743 ))
744 }
745
746 async fn handle_new_view(
751 &self,
752 _cx: &Cx,
753 _view: ViewNumber,
754 _view_change_msgs: Vec<PbftMessage>,
755 _preprepare_msgs: Vec<PbftMessage>,
756 ) -> Result<()> {
757 Err(Error::new(ErrorKind::InvalidStateTransition).with_message(
758 "PBFT new-view is not implemented (experimental consensus; no Byzantine \
759 fault tolerance under primary failure)",
760 ))
761 }
762}
763
764fn parse_replica_index(replica_id: &ReplicaId, replica_count: usize) -> Result<usize> {
765 let index = replica_id.as_str().parse::<usize>().map_err(|_| {
766 Error::new(ErrorKind::InvalidInput).with_message(format!(
767 "PBFT replica id {replica_id} must be a numeric index"
768 ))
769 })?;
770 if index >= replica_count {
771 return Err(Error::new(ErrorKind::InvalidInput).with_message(format!(
772 "PBFT replica id {replica_id} is outside configured replica set size {replica_count}"
773 )));
774 }
775 Ok(index)
776}
777
778pub struct PbftConsensus<T: PbftTransport> {
780 node: PbftNode<T>,
781}
782
783impl<T: PbftTransport> PbftConsensus<T> {
784 pub fn new(replica_id: ReplicaId, config: PbftConfig, transport: T) -> Result<Self> {
786 let node = PbftNode::new(replica_id, config, transport)?;
787 Ok(Self { node })
788 }
789
790 pub async fn submit(&self, cx: &Cx, request: ConsensusRequest) -> Result<ConsensusResponse> {
792 self.node.submit_request(cx, request.clone()).await?;
793
794 Ok(ConsensusResponse {
797 view: ViewNumber::new(0),
798 sequence: SequenceNumber::new(0),
799 result: Outcome::Ok(b"consensus result".to_vec()),
800 replica_id: self.node.replica_id.clone(),
801 timestamp: Time::from_millis(0),
802 })
803 }
804
805 pub async fn run(&self, cx: &Cx) -> Result<()> {
807 loop {
808 let message = self.node.transport.receive().await?;
810 self.node.process_message(cx, message).await?;
811 }
812 }
813}