1use gbp::{CodecError, ControlMessage, ErrorObject, GbpFrame};
23use gbp_core::{
24 ControlOpcode, ErrorClass, GbpFlags, GroupId, MemberId, NodeState, PayloadCodec, SequenceNo,
25 StreamId, StreamType, TransitionId, TransitionState, codes, errors::ErrorSpec, timeouts,
26};
27use gbp_mls::{MlsError, label_for};
28use std::collections::HashMap;
29use std::time::Duration;
30#[cfg(not(target_arch = "wasm32"))]
31use std::time::Instant;
32#[cfg(target_arch = "wasm32")]
33use web_time::Instant;
34
35#[derive(Debug, thiserror::Error)]
37pub enum NodeError {
38 #[error("codec: {0}")]
40 Codec(#[from] CodecError),
41 #[error("mls: {0}")]
43 Mls(#[from] MlsError),
44 #[error("invalid state: {0}")]
46 InvalidState(String),
47}
48
49pub struct OutboundFrame {
51 pub to: MemberId,
53 pub wire: Vec<u8>,
55}
56
57#[derive(Debug, Clone)]
59pub struct DeliveredPayload {
60 pub stream_type: StreamType,
62 pub stream_id: StreamId,
65 pub sequence_no: SequenceNo,
67 pub flags: u16,
69 pub plaintext: Vec<u8>,
71 pub codec: PayloadCodec,
73}
74
75#[derive(Debug, Clone)]
77pub enum Event {
78 StateChanged {
80 from: NodeState,
82 to: NodeState,
84 },
85 PayloadReceived(DeliveredPayload),
89 Control {
91 from: MemberId,
93 opcode: ControlOpcode,
95 transition_id: TransitionId,
97 request_id: u32,
99 args: Vec<u8>,
102 },
103 Error {
105 code: u16,
107 class: ErrorClass,
109 retryable: bool,
111 fatal: bool,
113 reason: String,
115 },
116 EpochAdvanced {
118 epoch: u64,
120 transition_id: TransitionId,
122 },
123 CoordinatorElectionNeeded,
127 BecameCoordinator,
130 CoordinatorClaim {
132 claimant: MemberId,
134 },
135}
136
137pub struct GroupNode {
144 pub member_id: MemberId,
146 pub is_coordinator: bool,
148 pub group_id: GroupId,
150 pub current_epoch: u64,
153 pub last_transition_id: TransitionId,
155 pub pending_transition_id: TransitionId,
157 pub state: NodeState,
159 pub transition_state: TransitionState,
161
162 out_seq: HashMap<(StreamType, StreamId), SequenceNo>,
163 in_hw: HashMap<(StreamType, StreamId), SequenceNo>,
164 events: Vec<Event>,
165
166 pending_commit_sender: Option<MemberId>,
170 prepare_deadline: Option<Instant>,
173 execute_deadline: Option<Instant>,
176 coordinator_last_seen: Option<Instant>,
179}
180
181impl GroupNode {
182 pub fn new(member_id: MemberId, group_id: GroupId) -> Self {
184 Self {
185 member_id,
186 group_id,
187 is_coordinator: false,
188 current_epoch: 0,
189 last_transition_id: 0,
190 pending_transition_id: 0,
191 state: NodeState::Idle,
192 transition_state: TransitionState::TIdle,
193 out_seq: HashMap::new(),
194 in_hw: HashMap::new(),
195 events: Vec::new(),
196 pending_commit_sender: None,
197 prepare_deadline: None,
198 execute_deadline: None,
199 coordinator_last_seen: None,
200 }
201 }
202
203 pub fn bootstrap_as_creator(&mut self, epoch: u64) {
205 self.transition(NodeState::Connecting);
206 self.transition(NodeState::EstablishingGroup);
207 self.current_epoch = epoch;
208 self.transition(NodeState::Active);
209 }
210
211 pub fn bootstrap_as_joiner(&mut self, epoch: u64, expected_first_tid: u32) {
220 self.transition(NodeState::Connecting);
221 self.transition(NodeState::EstablishingGroup);
222 self.current_epoch = epoch;
223 if expected_first_tid > 0 {
224 self.pending_transition_id = expected_first_tid;
225 self.transition_state = TransitionState::TPrepared;
226 }
227 self.transition(NodeState::Active);
228 }
229
230 pub fn drain_events(&mut self) -> Vec<Event> {
232 std::mem::take(&mut self.events)
233 }
234
235 pub fn member_stream_id(&self, base: u32) -> StreamId {
240 debug_assert!(
241 self.member_id < 1_000_000,
242 "member_id overflow: {0}",
243 self.member_id
244 );
245 base + self.member_id * 100
246 }
247
248 pub fn export_out_seq(&self) -> Vec<u8> {
259 let mut out = Vec::with_capacity(4 + self.out_seq.len() * 9);
260 out.extend_from_slice(&(self.out_seq.len() as u32).to_le_bytes());
261 for ((st, sid), seq) in &self.out_seq {
262 out.push(*st as u8);
263 out.extend_from_slice(&sid.to_le_bytes());
264 out.extend_from_slice(&seq.to_le_bytes());
265 }
266 out
267 }
268
269 pub fn restore_out_seq(&mut self, bytes: &[u8]) {
272 if bytes.len() < 4 {
273 return;
274 }
275 let n = u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]) as usize;
276 let mut cur = &bytes[4..];
277 for _ in 0..n {
278 if cur.len() < 9 {
279 break;
280 }
281 let st = match StreamType::try_from(cur[0]) {
282 Ok(s) => s,
283 Err(_) => break,
284 };
285 let sid = u32::from_le_bytes([cur[1], cur[2], cur[3], cur[4]]);
286 let seq = u32::from_le_bytes([cur[5], cur[6], cur[7], cur[8]]);
287 self.out_seq.insert((st, sid), seq);
288 cur = &cur[9..];
289 }
290 }
291
292 #[allow(clippy::too_many_arguments)]
301 pub fn send_payload<S: Sealer>(
302 &mut self,
303 seal: &mut S,
304 target: MemberId,
305 stream_type: StreamType,
306 stream_id: StreamId,
307 flags: u16,
308 plaintext: &[u8],
309 codec: PayloadCodec,
310 ) -> Result<OutboundFrame, NodeError> {
311 self.assert_can_send()?;
312 let seq = self.next_seq(stream_type, stream_id);
313 let ciphertext = seal.seal(stream_type, seq, plaintext)?;
314 let frame = GbpFrame::new(
315 self.group_id,
316 self.current_epoch,
317 self.last_transition_id,
318 stream_type,
319 stream_id,
320 flags,
321 seq,
322 ciphertext,
323 codec.as_u8(),
324 );
325 Ok(OutboundFrame {
326 to: target,
327 wire: frame.to_cbor(),
328 })
329 }
330
331 pub fn send_control<S: Sealer>(
341 &mut self,
342 seal: &mut S,
343 target: MemberId,
344 opcode: ControlOpcode,
345 transition_id: TransitionId,
346 request_id: u32,
347 args: Vec<u8>,
348 ) -> Result<OutboundFrame, NodeError> {
349 let ctl = ControlMessage::with_args(
350 opcode as u16,
351 request_id,
352 self.member_id,
353 transition_id,
354 args,
355 );
356 let mut flags = GbpFlags::ordered_reliable_system();
357 if matches!(
358 opcode,
359 ControlOpcode::PrepareTransition
360 | ControlOpcode::ReadyForTransition
361 | ControlOpcode::ExecuteTransition
362 ) {
363 flags |= GbpFlags::CRITICAL;
364 }
365 match opcode {
369 ControlOpcode::PrepareTransition => {
370 self.pending_transition_id = transition_id;
371 self.transition_state = TransitionState::TPrepared;
372 self.prepare_deadline =
373 Some(Instant::now() + Duration::from_millis(timeouts::T_PREPARE_MAX_MS));
374 self.execute_deadline = None;
375 }
376 ControlOpcode::ReadyForTransition => {
377 self.execute_deadline =
378 Some(Instant::now() + Duration::from_millis(timeouts::T_EXECUTE_MAX_MS));
379 }
380 ControlOpcode::ExecuteTransition | ControlOpcode::AbortTransition => {
381 self.prepare_deadline = None;
382 self.execute_deadline = None;
383 if opcode == ControlOpcode::AbortTransition {
384 self.pending_transition_id = 0;
385 self.transition_state = TransitionState::TAborted;
386 }
387 }
388 _ => {}
389 }
390 let stream_id = self.member_stream_id(0);
391 self.send_payload(
392 seal,
393 target,
394 StreamType::Control,
395 stream_id,
396 flags,
397 &ctl.to_cbor(),
398 PayloadCodec::Cbor,
399 )
400 }
401
402 pub fn on_wire<S: Sealer>(
413 &mut self,
414 seal: &mut S,
415 wire: &[u8],
416 ) -> Result<Vec<Event>, NodeError> {
417 let frame = match GbpFrame::decode(wire) {
422 Ok(f) => f,
423 Err(e) => {
424 self.emit_err_spec(codes::STREAM_POLICY_VIOLATION, format!("frame decode: {e}"));
425 return Ok(self.drain_events());
426 }
427 };
428 self.deliver_frame(seal, frame)?;
429 Ok(self.drain_events())
430 }
431
432 fn deliver_frame<S: Sealer>(&mut self, seal: &mut S, frame: GbpFrame) -> Result<(), NodeError> {
433 if frame.version != 1 {
436 self.emit_err_spec(codes::UNSUPPORTED_VERSION, "version != 1");
437 return Ok(());
438 }
439 if frame.group_id_array() != self.group_id {
440 self.emit_err_spec(codes::UNKNOWN_GROUP, "group_id");
441 return Ok(());
442 }
443 if frame.epoch != self.current_epoch {
444 self.emit_err_spec(
445 codes::EPOCH_MISMATCH,
446 format!("got {}, expected {}", frame.epoch, self.current_epoch),
447 );
448 self.trigger_resync();
449 return Ok(());
450 }
451 if let Err(e) = frame.validate_payload_size() {
452 self.emit_err_spec(codes::STREAM_POLICY_VIOLATION, format!("payload size: {e}"));
453 return Ok(());
454 }
455 let flags = GbpFlags::from_bits(frame.flags);
456 let st = match frame.stream_type_typed() {
457 Ok(st) => st,
458 Err(_) => {
459 self.emit_err_spec(codes::STREAM_POLICY_VIOLATION, "unknown stream_type");
460 return Ok(());
461 }
462 };
463
464 if st != StreamType::Control
470 && flags.has(GbpFlags::CRITICAL)
471 && frame.transition_id != self.last_transition_id
472 {
473 self.emit_err_spec(
474 codes::TRANSITION_MISMATCH,
475 format!(
476 "got tid={}, expected {}",
477 frame.transition_id, self.last_transition_id
478 ),
479 );
480 return Ok(());
481 }
482
483 let key = (st, frame.stream_id);
484 let hw = self.in_hw.get(&key).copied().unwrap_or(0);
485 if frame.sequence_no <= hw {
486 self.emit_err_spec(
487 codes::REPLAY_DETECTED,
488 format!(
489 "st={} sid={} seq={} hw={}",
490 st, frame.stream_id, frame.sequence_no, hw
491 ),
492 );
493 return Ok(());
494 }
495 self.in_hw.insert(key, frame.sequence_no);
496
497 let plain = match seal.open(st, frame.sequence_no, &frame.encrypted_payload) {
498 Ok(p) => p,
499 Err(e) => {
500 self.emit_err_named(
507 codes::DECRYPT_FAILED,
508 ErrorClass::Crypto,
509 true, false, format!("aead open: {e}"),
512 );
513 return Ok(());
514 }
515 };
516
517 match st {
518 StreamType::Control => self.handle_control(plain),
519 other => self.events.push(Event::PayloadReceived(DeliveredPayload {
520 stream_type: other,
521 stream_id: frame.stream_id,
522 sequence_no: frame.sequence_no,
523 flags: frame.flags,
524 plaintext: plain,
525 codec: frame.payload_codec(),
526 })),
527 }
528 Ok(())
529 }
530
531 fn handle_control(&mut self, plain: Vec<u8>) {
532 let c = match ControlMessage::from_cbor(&plain) {
533 Ok(c) => c,
534 Err(_) => {
535 self.emit_err_spec(codes::STREAM_POLICY_VIOLATION, "control decode");
536 return;
537 }
538 };
539 let opcode = match ControlOpcode::try_from(c.opcode) {
540 Ok(op) => op,
541 Err(_) => {
542 self.emit_err_spec(codes::STREAM_POLICY_VIOLATION, "unknown opcode");
543 return;
544 }
545 };
546 let tid_ok = match opcode {
548 ControlOpcode::PrepareTransition => {
552 c.transition_id > self.last_transition_id
553 && (self.pending_transition_id == 0
554 || self.pending_transition_id == c.transition_id)
555 }
556 ControlOpcode::ReadyForTransition
558 | ControlOpcode::ExecuteTransition
559 | ControlOpcode::AbortTransition => {
560 self.pending_transition_id != 0 && c.transition_id == self.pending_transition_id
561 }
562 _ => true,
565 };
566 if !tid_ok {
567 self.emit_err_spec(
568 codes::TRANSITION_MISMATCH,
569 format!(
570 "control tid={} not valid for {:?} (last={}, pending={})",
571 c.transition_id, opcode, self.last_transition_id, self.pending_transition_id
572 ),
573 );
574 return;
575 }
576 match opcode {
577 ControlOpcode::PrepareTransition => {
578 if self.pending_transition_id == c.transition_id {
582 let current_winner = self.pending_commit_sender.unwrap_or(MemberId::MAX);
583 if c.sender_id >= current_winner {
584 self.events.push(Event::Control {
588 from: c.sender_id,
589 opcode,
590 transition_id: c.transition_id,
591 request_id: c.request_id,
592 args: c.args.to_vec(),
593 });
594 return;
595 }
596 }
598 self.pending_transition_id = c.transition_id;
599 self.pending_commit_sender = Some(c.sender_id);
600 self.transition_state = TransitionState::TPrepared;
601 self.note_coordinator_activity();
603 self.execute_deadline =
605 Some(Instant::now() + Duration::from_millis(timeouts::T_EXECUTE_MAX_MS));
606 }
607 ControlOpcode::ReadyForTransition => {
608 self.transition_state = TransitionState::TReady;
609 self.prepare_deadline = None;
611 }
612 ControlOpcode::ExecuteTransition => {
613 self.execute_deadline = None;
614 self.pending_commit_sender = None;
615 self.apply_transition(c.transition_id);
616 self.note_coordinator_activity();
617 }
618 ControlOpcode::AbortTransition => {
619 self.prepare_deadline = None;
620 self.execute_deadline = None;
621 self.pending_commit_sender = None;
622 self.transition_state = TransitionState::TAborted;
623 self.pending_transition_id = 0;
624 }
625 ControlOpcode::GroupStateDigestResponse => {
626 if self.state == NodeState::Resyncing {
627 self.transition(NodeState::Active);
628 }
629 }
630 ControlOpcode::CapabilitiesAdvertise if Self::is_coordinator_claim(&c.args) => {
631 self.note_coordinator_activity();
633 if self.is_coordinator && c.sender_id < self.member_id {
637 self.is_coordinator = false;
638 }
639 self.events.push(Event::CoordinatorClaim {
640 claimant: c.sender_id,
641 });
642 }
643 _ => {}
644 }
645 self.events.push(Event::Control {
646 from: c.sender_id,
647 opcode,
648 transition_id: c.transition_id,
649 request_id: c.request_id,
650 args: c.args.to_vec(),
651 });
652 }
653
654 pub fn apply_transition(&mut self, tid: TransitionId) {
657 self.current_epoch += 1;
658 self.last_transition_id = tid;
659 self.pending_transition_id = 0;
660 self.pending_commit_sender = None;
661 self.transition_state = TransitionState::TExecuted;
662 self.out_seq.clear();
663 self.in_hw.clear();
664 self.events.push(Event::EpochAdvanced {
665 epoch: self.current_epoch,
666 transition_id: tid,
667 });
668 }
669
670 pub fn trigger_resync(&mut self) {
672 if self.state != NodeState::Resyncing {
673 self.transition(NodeState::Resyncing);
674 }
675 }
676
677 pub fn check_timeouts(&mut self) -> Vec<Event> {
684 let now = Instant::now();
685
686 if self.prepare_deadline.is_some_and(|d| now >= d) {
687 self.prepare_deadline = None;
688 self.execute_deadline = None;
689 self.pending_transition_id = 0;
690 self.transition_state = TransitionState::TAborted;
691 self.emit_err_spec(codes::PREPARE_TIMEOUT, "T_prepare_max exceeded");
692 }
693
694 if self.execute_deadline.is_some_and(|d| now >= d) {
695 self.execute_deadline = None;
696 self.emit_err_spec(codes::EXECUTE_TIMEOUT, "T_execute_max exceeded");
697 }
698
699 if self.coordinator_last_seen.is_some_and(|t| {
700 now.duration_since(t).as_millis() as u64 >= timeouts::T_COORDINATOR_GRACE_MS
701 }) {
702 self.coordinator_last_seen = None;
703 self.is_coordinator = false;
704 self.emit_err_spec(
705 codes::COORDINATOR_GONE,
706 "coordinator silence exceeded T_coordinator_grace",
707 );
708 self.events.push(Event::CoordinatorElectionNeeded);
709 }
710
711 self.drain_events()
712 }
713
714 pub fn note_coordinator_activity(&mut self) {
721 self.coordinator_last_seen = Some(Instant::now());
722 }
723
724 pub fn claim_coordinator<S: Sealer>(
735 &mut self,
736 seal: &mut S,
737 target: MemberId,
738 ) -> Result<OutboundFrame, NodeError> {
739 let args = vec![0xA1u8, 0x00, 0xF5];
741 self.is_coordinator = true;
742 self.coordinator_last_seen = Some(Instant::now());
743 self.events.push(Event::BecameCoordinator);
744 self.send_control(
745 seal,
746 target,
747 ControlOpcode::CapabilitiesAdvertise,
748 self.last_transition_id,
749 0,
750 args,
751 )
752 }
753
754 fn is_coordinator_claim(args: &[u8]) -> bool {
757 if args == [0xA1, 0x00, 0xF5] {
761 return true;
762 }
763 args.windows(2).any(|w| w == [0x00, 0xF5])
767 }
768
769 fn transition(&mut self, next: NodeState) {
770 if self.state == next {
771 return;
772 }
773 if !self.state.can_transition_to(next) {
774 let from = self.state;
775 self.state = NodeState::Failed;
776 self.events.push(Event::StateChanged {
777 from,
778 to: NodeState::Failed,
779 });
780 return;
781 }
782 let from = self.state;
783 self.state = next;
784 self.events.push(Event::StateChanged { from, to: next });
785 }
786
787 fn assert_can_send(&self) -> Result<(), NodeError> {
788 if matches!(
789 self.state,
790 NodeState::Active | NodeState::Resyncing | NodeState::EstablishingGroup
791 ) {
792 Ok(())
793 } else {
794 Err(NodeError::InvalidState(format!(
795 "cannot send in state {}",
796 self.state
797 )))
798 }
799 }
800
801 fn next_seq(&mut self, st: StreamType, sid: StreamId) -> SequenceNo {
802 let entry = self.out_seq.entry((st, sid)).or_insert(0);
803 *entry += 1;
804 *entry
805 }
806
807 fn emit_err_spec(&mut self, code: u16, reason: impl Into<String>) {
808 if let Some(spec) = ErrorSpec::lookup(code) {
809 self.emit_err_named(spec.code, spec.class, spec.retryable, spec.fatal, reason);
810 } else {
811 self.emit_err_named(code, ErrorClass::Policy, false, false, reason);
812 }
813 }
814
815 fn emit_err_named(
816 &mut self,
817 code: u16,
818 class: ErrorClass,
819 retryable: bool,
820 fatal: bool,
821 reason: impl Into<String>,
822 ) {
823 let reason = reason.into();
824 let (class, retryable, fatal) = if let Some(spec) = ErrorSpec::lookup(code) {
827 (spec.class, spec.retryable, spec.fatal)
828 } else {
829 (class, retryable, fatal)
830 };
831 let _ = ErrorObject::new(code, class, retryable, fatal, reason.clone()).to_cbor();
832 self.events.push(Event::Error {
833 code,
834 class,
835 retryable,
836 fatal,
837 reason,
838 });
839 if fatal {
840 let from = self.state;
841 self.state = NodeState::Failed;
842 self.events.push(Event::StateChanged {
843 from,
844 to: NodeState::Failed,
845 });
846 }
847 }
848}
849
850pub trait Sealer {
856 fn seal(&mut self, st: StreamType, seq: SequenceNo, pt: &[u8]) -> Result<Vec<u8>, MlsError>;
858 fn open(&mut self, st: StreamType, seq: SequenceNo, ct: &[u8]) -> Result<Vec<u8>, MlsError>;
860}
861
862impl Sealer for gbp_mls::MlsContext {
863 fn seal(&mut self, st: StreamType, seq: SequenceNo, pt: &[u8]) -> Result<Vec<u8>, MlsError> {
864 gbp_mls::MlsContext::seal(self, label_for(st), seq, pt)
865 }
866 fn open(&mut self, st: StreamType, seq: SequenceNo, ct: &[u8]) -> Result<Vec<u8>, MlsError> {
867 gbp_mls::MlsContext::open(self, label_for(st), seq, ct)
868 }
869}
870
871#[cfg(test)]
872mod tests {
873 use super::*;
874
875 struct PlainSealer;
876 impl Sealer for PlainSealer {
877 fn seal(
878 &mut self,
879 _st: StreamType,
880 _seq: SequenceNo,
881 pt: &[u8],
882 ) -> Result<Vec<u8>, MlsError> {
883 Ok(pt.to_vec())
884 }
885 fn open(
886 &mut self,
887 _st: StreamType,
888 _seq: SequenceNo,
889 ct: &[u8],
890 ) -> Result<Vec<u8>, MlsError> {
891 Ok(ct.to_vec())
892 }
893 }
894
895 fn group_id() -> GroupId {
896 let mut g = [0u8; 16];
897 g[..3].copy_from_slice(b"GBP");
898 g
899 }
900
901 #[test]
902 fn replay_window_rejects_repeat() {
903 let mut alice = GroupNode::new(1, group_id());
904 let mut bob = GroupNode::new(2, group_id());
905 alice.bootstrap_as_creator(1);
906 bob.bootstrap_as_joiner(1, 0);
907 let mut s = PlainSealer;
908 let sid = alice.member_stream_id(2);
909 let f = alice
910 .send_payload(
911 &mut s,
912 2,
913 StreamType::Text,
914 sid,
915 GbpFlags::ordered_reliable_ack(),
916 b"hi",
917 PayloadCodec::Cbor,
918 )
919 .unwrap();
920 let _ = bob.on_wire(&mut s, &f.wire).unwrap();
921 let evs = bob.on_wire(&mut s, &f.wire).unwrap();
922 assert!(evs.iter().any(|e| matches!(
923 e,
924 Event::Error {
925 code: codes::REPLAY_DETECTED,
926 ..
927 }
928 )));
929 }
930
931 #[test]
932 fn epoch_mismatch_triggers_resync() {
933 let mut alice = GroupNode::new(1, group_id());
934 let mut bob = GroupNode::new(2, group_id());
935 alice.bootstrap_as_creator(1);
936 bob.bootstrap_as_joiner(1, 0);
937 alice.current_epoch = 2;
938 let mut s = PlainSealer;
939 let sid = alice.member_stream_id(2);
940 let f = alice
941 .send_payload(
942 &mut s,
943 2,
944 StreamType::Text,
945 sid,
946 GbpFlags::ordered_reliable_ack(),
947 b"x",
948 PayloadCodec::Cbor,
949 )
950 .unwrap();
951 let _ = bob.on_wire(&mut s, &f.wire).unwrap();
952 assert_eq!(bob.state, NodeState::Resyncing);
953 }
954
955 #[test]
956 fn payload_emits_received_event() {
957 let mut alice = GroupNode::new(1, group_id());
958 let mut bob = GroupNode::new(2, group_id());
959 alice.bootstrap_as_creator(1);
960 bob.bootstrap_as_joiner(1, 0);
961 let mut s = PlainSealer;
962 let sid = alice.member_stream_id(2);
963 let f = alice
964 .send_payload(
965 &mut s,
966 2,
967 StreamType::Text,
968 sid,
969 GbpFlags::ordered_reliable_ack(),
970 b"payload",
971 PayloadCodec::Cbor,
972 )
973 .unwrap();
974 let evs = bob.on_wire(&mut s, &f.wire).unwrap();
975 let pr = evs
976 .into_iter()
977 .find_map(|e| match e {
978 Event::PayloadReceived(p) => Some(p),
979 _ => None,
980 })
981 .expect("payload");
982 assert_eq!(pr.stream_type, StreamType::Text);
983 assert_eq!(pr.plaintext, b"payload");
984 }
985
986 fn drain_errs(events: &[Event]) -> Vec<u16> {
989 events
990 .iter()
991 .filter_map(|e| match e {
992 Event::Error { code, .. } => Some(*code),
993 _ => None,
994 })
995 .collect()
996 }
997
998 fn drain_controls(events: &[Event]) -> Vec<(ControlOpcode, TransitionId)> {
999 events
1000 .iter()
1001 .filter_map(|e| match e {
1002 Event::Control {
1003 opcode,
1004 transition_id,
1005 ..
1006 } => Some((*opcode, *transition_id)),
1007 _ => None,
1008 })
1009 .collect()
1010 }
1011
1012 #[test]
1013 fn prepare_transition_sets_pending_on_sender_and_receiver() {
1014 let mut coord = GroupNode::new(1, group_id());
1015 let mut peer = GroupNode::new(2, group_id());
1016 coord.bootstrap_as_creator(0);
1017 peer.bootstrap_as_joiner(0, 0);
1018 let mut s = PlainSealer;
1019 let f = coord
1021 .send_control(
1022 &mut s,
1023 0,
1024 ControlOpcode::PrepareTransition,
1025 1,
1026 100,
1027 b"commit-blob".to_vec(),
1028 )
1029 .unwrap();
1030 assert_eq!(coord.pending_transition_id, 1, "sender mirrors pending");
1031 assert_eq!(coord.transition_state, TransitionState::TPrepared);
1032 let evs = peer.on_wire(&mut s, &f.wire).unwrap();
1033 assert_eq!(peer.pending_transition_id, 1, "receiver records pending");
1034 assert!(
1035 drain_errs(&evs).is_empty(),
1036 "no error: {:?}",
1037 drain_errs(&evs)
1038 );
1039 let ctls = drain_controls(&evs);
1040 assert_eq!(ctls, vec![(ControlOpcode::PrepareTransition, 1)]);
1041 }
1042
1043 #[test]
1044 fn ready_with_wrong_tid_is_rejected() {
1045 let mut coord = GroupNode::new(1, group_id());
1046 let mut peer = GroupNode::new(2, group_id());
1047 coord.bootstrap_as_creator(0);
1048 peer.bootstrap_as_joiner(0, 0);
1049 let mut s = PlainSealer;
1050 let f = coord
1051 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1052 .unwrap();
1053 peer.on_wire(&mut s, &f.wire).unwrap();
1054 let bogus = peer
1056 .send_control(&mut s, 1, ControlOpcode::ReadyForTransition, 7, 1, vec![])
1057 .unwrap();
1058 let evs = coord.on_wire(&mut s, &bogus.wire).unwrap();
1059 let errs = drain_errs(&evs);
1060 assert!(errs.contains(&codes::TRANSITION_MISMATCH), "got {:?}", errs);
1061 }
1062
1063 #[test]
1064 fn execute_advances_epoch_and_clears_pending() {
1065 let mut coord = GroupNode::new(1, group_id());
1066 let mut peer = GroupNode::new(2, group_id());
1067 coord.bootstrap_as_creator(0);
1068 peer.bootstrap_as_joiner(0, 0);
1069 let mut s = PlainSealer;
1070 let prep = coord
1071 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1072 .unwrap();
1073 peer.on_wire(&mut s, &prep.wire).unwrap();
1074 let exec = coord
1076 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1077 .unwrap();
1078 coord.apply_transition(1);
1079 let evs = peer.on_wire(&mut s, &exec.wire).unwrap();
1080 assert_eq!(coord.last_transition_id, 1);
1081 assert_eq!(coord.current_epoch, 1);
1082 assert_eq!(peer.last_transition_id, 1);
1083 assert_eq!(peer.current_epoch, 1);
1084 assert_eq!(peer.pending_transition_id, 0);
1085 assert!(evs.iter().any(|e| matches!(
1086 e,
1087 Event::EpochAdvanced {
1088 transition_id: 1,
1089 ..
1090 }
1091 )));
1092 }
1093
1094 #[test]
1095 fn abort_clears_pending_no_advance() {
1096 let mut coord = GroupNode::new(1, group_id());
1097 let mut peer = GroupNode::new(2, group_id());
1098 coord.bootstrap_as_creator(0);
1099 peer.bootstrap_as_joiner(0, 0);
1100 let mut s = PlainSealer;
1101 let prep = coord
1102 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1103 .unwrap();
1104 peer.on_wire(&mut s, &prep.wire).unwrap();
1105 let abort = coord
1106 .send_control(&mut s, 0, ControlOpcode::AbortTransition, 1, 2, vec![])
1107 .unwrap();
1108 peer.on_wire(&mut s, &abort.wire).unwrap();
1109 assert_eq!(peer.pending_transition_id, 0);
1110 assert_eq!(peer.current_epoch, 0);
1111 assert_eq!(peer.transition_state, TransitionState::TAborted);
1112 assert_eq!(coord.transition_state, TransitionState::TAborted);
1113 }
1114
1115 #[test]
1116 fn bootstrap_as_joiner_with_expected_tid_accepts_first_execute() {
1117 let mut coord = GroupNode::new(1, group_id());
1118 let mut joiner = GroupNode::new(2, group_id());
1120 coord.bootstrap_as_creator(0);
1121 joiner.bootstrap_as_joiner(0, 1);
1122 assert_eq!(joiner.pending_transition_id, 1);
1123 let mut s = PlainSealer;
1124 let _ = coord
1126 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1127 .unwrap();
1128 let exec = coord
1130 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1131 .unwrap();
1132 let evs = joiner.on_wire(&mut s, &exec.wire).unwrap();
1133 let errs = drain_errs(&evs);
1134 assert!(
1135 errs.is_empty(),
1136 "expected clean apply, got errors {:?}",
1137 errs
1138 );
1139 assert_eq!(joiner.last_transition_id, 1);
1140 assert_eq!(joiner.current_epoch, 1);
1141 }
1142
1143 #[test]
1146 fn claim_coordinator_sets_flag_and_emits_event() {
1147 let mut node = GroupNode::new(1, group_id());
1148 node.bootstrap_as_creator(0);
1149 node.drain_events();
1150 let mut s = PlainSealer;
1151 let _ = node.claim_coordinator(&mut s, 0).unwrap();
1152 assert!(node.is_coordinator);
1153 let evs = node.drain_events();
1154 assert!(evs.iter().any(|e| matches!(e, Event::BecameCoordinator)));
1155 }
1156
1157 #[test]
1158 fn coordinator_gone_emits_election_needed() {
1159 let mut member = GroupNode::new(2, group_id());
1160 member.bootstrap_as_joiner(0, 0);
1161 member.coordinator_last_seen = Some(Instant::now() - Duration::from_millis(11_000));
1162 let evs = member.check_timeouts();
1163 assert!(
1164 evs.iter()
1165 .any(|e| matches!(e, Event::CoordinatorElectionNeeded))
1166 );
1167 assert!(!member.is_coordinator, "flag cleared on silence");
1168 }
1169
1170 #[test]
1171 fn capabilities_advertise_with_claim_resets_silence_timer() {
1172 let mut member = GroupNode::new(2, group_id());
1173 let mut coord = GroupNode::new(1, group_id());
1174 member.bootstrap_as_joiner(0, 0);
1175 coord.bootstrap_as_creator(0);
1176 let mut s = PlainSealer;
1177 let f = coord.claim_coordinator(&mut s, 2).unwrap();
1179 let evs = member.on_wire(&mut s, &f.wire).unwrap();
1181 assert!(
1182 member.coordinator_last_seen.is_some(),
1183 "silence timer reset"
1184 );
1185 assert!(
1186 evs.iter()
1187 .any(|e| matches!(e, Event::CoordinatorClaim { claimant: 1 }))
1188 );
1189 }
1190
1191 #[test]
1192 fn higher_id_yields_to_lower_claimant() {
1193 let mut node5 = GroupNode::new(5, group_id());
1195 let mut node2 = GroupNode::new(2, group_id());
1196 node5.bootstrap_as_joiner(0, 0);
1197 node2.bootstrap_as_creator(0);
1198 let mut s = PlainSealer;
1199 node5.is_coordinator = true;
1201 let f = node2.claim_coordinator(&mut s, 5).unwrap();
1203 node5.on_wire(&mut s, &f.wire).unwrap();
1204 assert!(!node5.is_coordinator, "node5 yielded to node2");
1205 }
1206
1207 #[test]
1208 fn lower_id_keeps_coordinator_against_higher_claimant() {
1209 let mut node1 = GroupNode::new(1, group_id());
1210 let mut node5 = GroupNode::new(5, group_id());
1211 node1.bootstrap_as_creator(0);
1212 node5.bootstrap_as_joiner(0, 0);
1213 let mut s = PlainSealer;
1214 node1.is_coordinator = true;
1215 let f = node5.claim_coordinator(&mut s, 1).unwrap();
1216 node1.on_wire(&mut s, &f.wire).unwrap();
1217 assert!(node1.is_coordinator, "node1 keeps role — it has lower id");
1218 }
1219
1220 #[test]
1223 fn competing_prepare_lower_member_id_wins() {
1224 let mut node = GroupNode::new(10, group_id());
1227 node.bootstrap_as_joiner(0, 0);
1228 let mut s = PlainSealer;
1229
1230 let mut sender1 = GroupNode::new(1, group_id());
1232 sender1.bootstrap_as_creator(0);
1233 let f1 = sender1
1234 .send_control(
1235 &mut s,
1236 10,
1237 ControlOpcode::PrepareTransition,
1238 1,
1239 1,
1240 b"commit-A".to_vec(),
1241 )
1242 .unwrap();
1243 node.on_wire(&mut s, &f1.wire).unwrap();
1244 assert_eq!(
1245 node.pending_commit_sender,
1246 Some(1),
1247 "member 1 is initial winner"
1248 );
1249
1250 let mut sender3 = GroupNode::new(3, group_id());
1252 sender3.bootstrap_as_creator(0);
1253 let f3 = sender3
1254 .send_control(
1255 &mut s,
1256 10,
1257 ControlOpcode::PrepareTransition,
1258 1,
1259 2,
1260 b"commit-B".to_vec(),
1261 )
1262 .unwrap();
1263 node.on_wire(&mut s, &f3.wire).unwrap();
1264 assert_eq!(node.pending_commit_sender, Some(1), "member 1 still wins");
1266 assert_eq!(node.pending_transition_id, 1);
1267 }
1268
1269 #[test]
1270 fn competing_prepare_later_lower_id_displaces_winner() {
1271 let mut node = GroupNode::new(10, group_id());
1273 node.bootstrap_as_joiner(0, 0);
1274 let mut s = PlainSealer;
1275
1276 let mut sender5 = GroupNode::new(5, group_id());
1277 sender5.bootstrap_as_creator(0);
1278 let f5 = sender5
1279 .send_control(
1280 &mut s,
1281 10,
1282 ControlOpcode::PrepareTransition,
1283 1,
1284 1,
1285 b"commit-X".to_vec(),
1286 )
1287 .unwrap();
1288 node.on_wire(&mut s, &f5.wire).unwrap();
1289 assert_eq!(node.pending_commit_sender, Some(5));
1290
1291 let mut sender2 = GroupNode::new(2, group_id());
1292 sender2.bootstrap_as_creator(0);
1293 let f2 = sender2
1294 .send_control(
1295 &mut s,
1296 10,
1297 ControlOpcode::PrepareTransition,
1298 1,
1299 2,
1300 b"commit-Y".to_vec(),
1301 )
1302 .unwrap();
1303 node.on_wire(&mut s, &f2.wire).unwrap();
1304 assert_eq!(
1305 node.pending_commit_sender,
1306 Some(2),
1307 "member 2 displaces member 5"
1308 );
1309 }
1310
1311 #[test]
1312 fn apply_transition_clears_commit_sender() {
1313 let mut coord = GroupNode::new(1, group_id());
1314 coord.bootstrap_as_creator(0);
1315 let mut s = PlainSealer;
1316 coord
1317 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1318 .unwrap();
1319 coord.apply_transition(1);
1320 assert_eq!(coord.pending_commit_sender, None);
1321 }
1322
1323 #[test]
1326 fn prepare_timeout_fires_when_deadline_exceeded() {
1327 let mut coord = GroupNode::new(1, group_id());
1328 coord.bootstrap_as_creator(0);
1329 let mut s = PlainSealer;
1330 coord
1331 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1332 .unwrap();
1333 coord.prepare_deadline = Some(Instant::now() - Duration::from_millis(1));
1335 let evs = coord.check_timeouts();
1336 assert!(
1337 evs.iter().any(|e| matches!(
1338 e,
1339 Event::Error {
1340 code: codes::PREPARE_TIMEOUT,
1341 ..
1342 }
1343 )),
1344 "expected PREPARE_TIMEOUT, got {:?}",
1345 evs
1346 );
1347 assert_eq!(
1348 coord.transition_state,
1349 TransitionState::TAborted,
1350 "transition aborted"
1351 );
1352 assert_eq!(coord.prepare_deadline, None, "deadline cleared");
1353 }
1354
1355 #[test]
1356 fn execute_timeout_fires_when_deadline_exceeded() {
1357 let mut member = GroupNode::new(2, group_id());
1358 member.bootstrap_as_joiner(0, 0);
1359 let mut s = PlainSealer;
1360 member.pending_transition_id = 1;
1362 member.transition_state = TransitionState::TPrepared;
1363 member
1364 .send_control(&mut s, 1, ControlOpcode::ReadyForTransition, 1, 1, vec![])
1365 .unwrap();
1366 member.execute_deadline = Some(Instant::now() - Duration::from_millis(1));
1368 let evs = member.check_timeouts();
1369 assert!(
1370 evs.iter().any(|e| matches!(
1371 e,
1372 Event::Error {
1373 code: codes::EXECUTE_TIMEOUT,
1374 ..
1375 }
1376 )),
1377 "expected EXECUTE_TIMEOUT, got {:?}",
1378 evs
1379 );
1380 assert_eq!(member.execute_deadline, None, "deadline cleared");
1381 }
1382
1383 #[test]
1384 fn coordinator_gone_fires_after_silence() {
1385 let mut member = GroupNode::new(2, group_id());
1386 member.bootstrap_as_joiner(0, 0);
1387 member.coordinator_last_seen = Some(Instant::now() - Duration::from_millis(11_000));
1389 let evs = member.check_timeouts();
1390 assert!(
1391 evs.iter().any(|e| matches!(
1392 e,
1393 Event::Error {
1394 code: codes::COORDINATOR_GONE,
1395 ..
1396 }
1397 )),
1398 "expected COORDINATOR_GONE, got {:?}",
1399 evs
1400 );
1401 assert_eq!(member.coordinator_last_seen, None, "timer cleared");
1402 }
1403
1404 #[test]
1405 fn note_coordinator_activity_resets_silence_timer() {
1406 let mut member = GroupNode::new(2, group_id());
1407 member.bootstrap_as_joiner(0, 0);
1408 member.coordinator_last_seen = Some(Instant::now() - Duration::from_millis(11_000));
1410 member.note_coordinator_activity();
1412 let evs = member.check_timeouts();
1413 assert!(
1414 !evs.iter().any(|e| matches!(
1415 e,
1416 Event::Error {
1417 code: codes::COORDINATOR_GONE,
1418 ..
1419 }
1420 )),
1421 "should NOT fire after reset"
1422 );
1423 }
1424
1425 #[test]
1426 fn execute_clears_prepare_deadline() {
1427 let mut coord = GroupNode::new(1, group_id());
1428 coord.bootstrap_as_creator(0);
1429 let mut s = PlainSealer;
1430 coord
1431 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1432 .unwrap();
1433 assert!(coord.prepare_deadline.is_some(), "deadline armed");
1434 coord
1435 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1436 .unwrap();
1437 assert_eq!(coord.prepare_deadline, None, "deadline cleared on EXECUTE");
1438 assert_eq!(
1439 coord.execute_deadline, None,
1440 "execute_deadline also cleared"
1441 );
1442 }
1443
1444 #[test]
1445 fn receive_prepare_arms_execute_deadline() {
1446 let mut coord = GroupNode::new(1, group_id());
1447 let mut member = GroupNode::new(2, group_id());
1448 coord.bootstrap_as_creator(0);
1449 member.bootstrap_as_joiner(0, 0);
1450 let mut s = PlainSealer;
1451 let f = coord
1452 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1453 .unwrap();
1454 member.on_wire(&mut s, &f.wire).unwrap();
1455 assert!(
1456 member.execute_deadline.is_some(),
1457 "execute_deadline armed on receiving PREPARE"
1458 );
1459 }
1460
1461 #[test]
1462 fn receive_execute_clears_execute_deadline() {
1463 let mut coord = GroupNode::new(1, group_id());
1464 let mut member = GroupNode::new(2, group_id());
1465 coord.bootstrap_as_creator(0);
1466 member.bootstrap_as_joiner(0, 0);
1467 let mut s = PlainSealer;
1468 let prep = coord
1469 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1470 .unwrap();
1471 member.on_wire(&mut s, &prep.wire).unwrap();
1472 let exec = coord
1473 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1474 .unwrap();
1475 member.on_wire(&mut s, &exec.wire).unwrap();
1476 assert_eq!(member.execute_deadline, None, "cleared on EXECUTE");
1477 }
1478
1479 #[test]
1480 fn no_timeout_when_deadlines_not_set() {
1481 let mut node = GroupNode::new(1, group_id());
1482 node.bootstrap_as_creator(0);
1483 node.drain_events(); let evs = node.check_timeouts();
1485 assert!(evs.is_empty(), "no events without armed deadlines");
1486 }
1487
1488 #[test]
1489 fn prepare_with_already_applied_tid_is_rejected() {
1490 let mut coord = GroupNode::new(1, group_id());
1493 coord.bootstrap_as_creator(0);
1494 let mut s = PlainSealer;
1495 let _ = coord
1496 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1497 .unwrap();
1498 coord.apply_transition(1);
1499 assert_eq!(coord.last_transition_id, 1);
1500 assert_eq!(coord.pending_transition_id, 0);
1501 let mut peer = GroupNode::new(2, group_id());
1505 peer.bootstrap_as_joiner(coord.current_epoch, 0);
1506 let stale = peer
1507 .send_control(&mut s, 1, ControlOpcode::PrepareTransition, 1, 9, vec![])
1508 .unwrap();
1509 let evs = coord.on_wire(&mut s, &stale.wire).unwrap();
1510 let errs = drain_errs(&evs);
1511 assert!(
1512 errs.contains(&codes::TRANSITION_MISMATCH),
1513 "expected TRANSITION_MISMATCH, got {:?}",
1514 errs
1515 );
1516 }
1517
1518 #[test]
1519 fn decrypt_failed_is_non_fatal() {
1520 struct OpenFailSealer;
1522 impl Sealer for OpenFailSealer {
1523 fn seal(
1524 &mut self,
1525 _: StreamType,
1526 _: SequenceNo,
1527 p: &[u8],
1528 ) -> Result<Vec<u8>, MlsError> {
1529 Ok(p.to_vec())
1530 }
1531 fn open(
1532 &mut self,
1533 _: StreamType,
1534 _: SequenceNo,
1535 _: &[u8],
1536 ) -> Result<Vec<u8>, MlsError> {
1537 Err(MlsError::Aead("simulated".into()))
1538 }
1539 }
1540 let mut alice = GroupNode::new(1, group_id());
1541 let mut bob = GroupNode::new(2, group_id());
1542 alice.bootstrap_as_creator(1);
1543 bob.bootstrap_as_joiner(1, 0);
1544 let mut s = PlainSealer;
1545 let sid = alice.member_stream_id(2);
1546 let f = alice
1547 .send_payload(
1548 &mut s,
1549 2,
1550 StreamType::Text,
1551 sid,
1552 GbpFlags::ordered_reliable_ack(),
1553 b"x",
1554 PayloadCodec::Cbor,
1555 )
1556 .unwrap();
1557 let mut fail = OpenFailSealer;
1558 let evs = bob.on_wire(&mut fail, &f.wire).unwrap();
1559 let err = evs
1560 .iter()
1561 .find_map(|e| match e {
1562 Event::Error {
1563 code,
1564 fatal,
1565 retryable,
1566 ..
1567 } => Some((*code, *fatal, *retryable)),
1568 _ => None,
1569 })
1570 .expect("error event");
1571 assert_eq!(err.0, codes::DECRYPT_FAILED);
1572 assert!(!err.1, "must be non-fatal");
1573 assert!(err.2, "must be retryable");
1574 assert_eq!(bob.state, NodeState::Active, "bob stays Active");
1575 }
1576}