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 => {
631 if Self::is_coordinator_claim(&c.args) {
632 self.note_coordinator_activity();
634 if self.is_coordinator && c.sender_id < self.member_id {
638 self.is_coordinator = false;
639 }
640 self.events.push(Event::CoordinatorClaim {
641 claimant: c.sender_id,
642 });
643 }
644 }
645 _ => {}
646 }
647 self.events.push(Event::Control {
648 from: c.sender_id,
649 opcode,
650 transition_id: c.transition_id,
651 request_id: c.request_id,
652 args: c.args.to_vec(),
653 });
654 }
655
656 pub fn apply_transition(&mut self, tid: TransitionId) {
659 self.current_epoch += 1;
660 self.last_transition_id = tid;
661 self.pending_transition_id = 0;
662 self.pending_commit_sender = None;
663 self.transition_state = TransitionState::TExecuted;
664 self.out_seq.clear();
665 self.in_hw.clear();
666 self.events.push(Event::EpochAdvanced {
667 epoch: self.current_epoch,
668 transition_id: tid,
669 });
670 }
671
672 pub fn trigger_resync(&mut self) {
674 if self.state != NodeState::Resyncing {
675 self.transition(NodeState::Resyncing);
676 }
677 }
678
679 pub fn check_timeouts(&mut self) -> Vec<Event> {
686 let now = Instant::now();
687
688 if self.prepare_deadline.is_some_and(|d| now >= d) {
689 self.prepare_deadline = None;
690 self.execute_deadline = None;
691 self.pending_transition_id = 0;
692 self.transition_state = TransitionState::TAborted;
693 self.emit_err_spec(codes::PREPARE_TIMEOUT, "T_prepare_max exceeded");
694 }
695
696 if self.execute_deadline.is_some_and(|d| now >= d) {
697 self.execute_deadline = None;
698 self.emit_err_spec(codes::EXECUTE_TIMEOUT, "T_execute_max exceeded");
699 }
700
701 if self.coordinator_last_seen.is_some_and(|t| {
702 now.duration_since(t).as_millis() as u64 >= timeouts::T_COORDINATOR_GRACE_MS
703 }) {
704 self.coordinator_last_seen = None;
705 self.is_coordinator = false;
706 self.emit_err_spec(
707 codes::COORDINATOR_GONE,
708 "coordinator silence exceeded T_coordinator_grace",
709 );
710 self.events.push(Event::CoordinatorElectionNeeded);
711 }
712
713 self.drain_events()
714 }
715
716 pub fn note_coordinator_activity(&mut self) {
723 self.coordinator_last_seen = Some(Instant::now());
724 }
725
726 pub fn claim_coordinator<S: Sealer>(
737 &mut self,
738 seal: &mut S,
739 target: MemberId,
740 ) -> Result<OutboundFrame, NodeError> {
741 let args = vec![0xA1u8, 0x00, 0xF5];
743 self.is_coordinator = true;
744 self.coordinator_last_seen = Some(Instant::now());
745 self.events.push(Event::BecameCoordinator);
746 self.send_control(
747 seal,
748 target,
749 ControlOpcode::CapabilitiesAdvertise,
750 self.last_transition_id,
751 0,
752 args,
753 )
754 }
755
756 fn is_coordinator_claim(args: &[u8]) -> bool {
759 if args == [0xA1, 0x00, 0xF5] {
763 return true;
764 }
765 args.windows(2).any(|w| w == [0x00, 0xF5])
769 }
770
771 fn transition(&mut self, next: NodeState) {
772 if self.state == next {
773 return;
774 }
775 if !self.state.can_transition_to(next) {
776 let from = self.state;
777 self.state = NodeState::Failed;
778 self.events.push(Event::StateChanged {
779 from,
780 to: NodeState::Failed,
781 });
782 return;
783 }
784 let from = self.state;
785 self.state = next;
786 self.events.push(Event::StateChanged { from, to: next });
787 }
788
789 fn assert_can_send(&self) -> Result<(), NodeError> {
790 if matches!(
791 self.state,
792 NodeState::Active | NodeState::Resyncing | NodeState::EstablishingGroup
793 ) {
794 Ok(())
795 } else {
796 Err(NodeError::InvalidState(format!(
797 "cannot send in state {}",
798 self.state
799 )))
800 }
801 }
802
803 fn next_seq(&mut self, st: StreamType, sid: StreamId) -> SequenceNo {
804 let entry = self.out_seq.entry((st, sid)).or_insert(0);
805 *entry += 1;
806 *entry
807 }
808
809 fn emit_err_spec(&mut self, code: u16, reason: impl Into<String>) {
810 if let Some(spec) = ErrorSpec::lookup(code) {
811 self.emit_err_named(spec.code, spec.class, spec.retryable, spec.fatal, reason);
812 } else {
813 self.emit_err_named(code, ErrorClass::Policy, false, false, reason);
814 }
815 }
816
817 fn emit_err_named(
818 &mut self,
819 code: u16,
820 class: ErrorClass,
821 retryable: bool,
822 fatal: bool,
823 reason: impl Into<String>,
824 ) {
825 let reason = reason.into();
826 let (class, retryable, fatal) = if let Some(spec) = ErrorSpec::lookup(code) {
829 (spec.class, spec.retryable, spec.fatal)
830 } else {
831 (class, retryable, fatal)
832 };
833 let _ = ErrorObject::new(code, class, retryable, fatal, reason.clone()).to_cbor();
834 self.events.push(Event::Error {
835 code,
836 class,
837 retryable,
838 fatal,
839 reason,
840 });
841 if fatal {
842 let from = self.state;
843 self.state = NodeState::Failed;
844 self.events.push(Event::StateChanged {
845 from,
846 to: NodeState::Failed,
847 });
848 }
849 }
850}
851
852pub trait Sealer {
858 fn seal(&mut self, st: StreamType, seq: SequenceNo, pt: &[u8]) -> Result<Vec<u8>, MlsError>;
860 fn open(&mut self, st: StreamType, seq: SequenceNo, ct: &[u8]) -> Result<Vec<u8>, MlsError>;
862}
863
864impl Sealer for gbp_mls::MlsContext {
865 fn seal(&mut self, st: StreamType, seq: SequenceNo, pt: &[u8]) -> Result<Vec<u8>, MlsError> {
866 gbp_mls::MlsContext::seal(self, label_for(st), seq, pt)
867 }
868 fn open(&mut self, st: StreamType, seq: SequenceNo, ct: &[u8]) -> Result<Vec<u8>, MlsError> {
869 gbp_mls::MlsContext::open(self, label_for(st), seq, ct)
870 }
871}
872
873#[cfg(test)]
874mod tests {
875 use super::*;
876
877 struct PlainSealer;
878 impl Sealer for PlainSealer {
879 fn seal(
880 &mut self,
881 _st: StreamType,
882 _seq: SequenceNo,
883 pt: &[u8],
884 ) -> Result<Vec<u8>, MlsError> {
885 Ok(pt.to_vec())
886 }
887 fn open(
888 &mut self,
889 _st: StreamType,
890 _seq: SequenceNo,
891 ct: &[u8],
892 ) -> Result<Vec<u8>, MlsError> {
893 Ok(ct.to_vec())
894 }
895 }
896
897 fn group_id() -> GroupId {
898 let mut g = [0u8; 16];
899 g[..3].copy_from_slice(b"GBP");
900 g
901 }
902
903 #[test]
904 fn replay_window_rejects_repeat() {
905 let mut alice = GroupNode::new(1, group_id());
906 let mut bob = GroupNode::new(2, group_id());
907 alice.bootstrap_as_creator(1);
908 bob.bootstrap_as_joiner(1, 0);
909 let mut s = PlainSealer;
910 let sid = alice.member_stream_id(2);
911 let f = alice
912 .send_payload(
913 &mut s,
914 2,
915 StreamType::Text,
916 sid,
917 GbpFlags::ordered_reliable_ack(),
918 b"hi",
919 PayloadCodec::Cbor,
920 )
921 .unwrap();
922 let _ = bob.on_wire(&mut s, &f.wire).unwrap();
923 let evs = bob.on_wire(&mut s, &f.wire).unwrap();
924 assert!(evs.iter().any(|e| matches!(
925 e,
926 Event::Error {
927 code: codes::REPLAY_DETECTED,
928 ..
929 }
930 )));
931 }
932
933 #[test]
934 fn epoch_mismatch_triggers_resync() {
935 let mut alice = GroupNode::new(1, group_id());
936 let mut bob = GroupNode::new(2, group_id());
937 alice.bootstrap_as_creator(1);
938 bob.bootstrap_as_joiner(1, 0);
939 alice.current_epoch = 2;
940 let mut s = PlainSealer;
941 let sid = alice.member_stream_id(2);
942 let f = alice
943 .send_payload(
944 &mut s,
945 2,
946 StreamType::Text,
947 sid,
948 GbpFlags::ordered_reliable_ack(),
949 b"x",
950 PayloadCodec::Cbor,
951 )
952 .unwrap();
953 let _ = bob.on_wire(&mut s, &f.wire).unwrap();
954 assert_eq!(bob.state, NodeState::Resyncing);
955 }
956
957 #[test]
958 fn payload_emits_received_event() {
959 let mut alice = GroupNode::new(1, group_id());
960 let mut bob = GroupNode::new(2, group_id());
961 alice.bootstrap_as_creator(1);
962 bob.bootstrap_as_joiner(1, 0);
963 let mut s = PlainSealer;
964 let sid = alice.member_stream_id(2);
965 let f = alice
966 .send_payload(
967 &mut s,
968 2,
969 StreamType::Text,
970 sid,
971 GbpFlags::ordered_reliable_ack(),
972 b"payload",
973 PayloadCodec::Cbor,
974 )
975 .unwrap();
976 let evs = bob.on_wire(&mut s, &f.wire).unwrap();
977 let pr = evs
978 .into_iter()
979 .find_map(|e| match e {
980 Event::PayloadReceived(p) => Some(p),
981 _ => None,
982 })
983 .expect("payload");
984 assert_eq!(pr.stream_type, StreamType::Text);
985 assert_eq!(pr.plaintext, b"payload");
986 }
987
988 fn drain_errs(events: &[Event]) -> Vec<u16> {
991 events
992 .iter()
993 .filter_map(|e| match e {
994 Event::Error { code, .. } => Some(*code),
995 _ => None,
996 })
997 .collect()
998 }
999
1000 fn drain_controls(events: &[Event]) -> Vec<(ControlOpcode, TransitionId)> {
1001 events
1002 .iter()
1003 .filter_map(|e| match e {
1004 Event::Control {
1005 opcode,
1006 transition_id,
1007 ..
1008 } => Some((*opcode, *transition_id)),
1009 _ => None,
1010 })
1011 .collect()
1012 }
1013
1014 #[test]
1015 fn prepare_transition_sets_pending_on_sender_and_receiver() {
1016 let mut coord = GroupNode::new(1, group_id());
1017 let mut peer = GroupNode::new(2, group_id());
1018 coord.bootstrap_as_creator(0);
1019 peer.bootstrap_as_joiner(0, 0);
1020 let mut s = PlainSealer;
1021 let f = coord
1023 .send_control(
1024 &mut s,
1025 0,
1026 ControlOpcode::PrepareTransition,
1027 1,
1028 100,
1029 b"commit-blob".to_vec(),
1030 )
1031 .unwrap();
1032 assert_eq!(coord.pending_transition_id, 1, "sender mirrors pending");
1033 assert_eq!(coord.transition_state, TransitionState::TPrepared);
1034 let evs = peer.on_wire(&mut s, &f.wire).unwrap();
1035 assert_eq!(peer.pending_transition_id, 1, "receiver records pending");
1036 assert!(
1037 drain_errs(&evs).is_empty(),
1038 "no error: {:?}",
1039 drain_errs(&evs)
1040 );
1041 let ctls = drain_controls(&evs);
1042 assert_eq!(ctls, vec![(ControlOpcode::PrepareTransition, 1)]);
1043 }
1044
1045 #[test]
1046 fn ready_with_wrong_tid_is_rejected() {
1047 let mut coord = GroupNode::new(1, group_id());
1048 let mut peer = GroupNode::new(2, group_id());
1049 coord.bootstrap_as_creator(0);
1050 peer.bootstrap_as_joiner(0, 0);
1051 let mut s = PlainSealer;
1052 let f = coord
1053 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1054 .unwrap();
1055 peer.on_wire(&mut s, &f.wire).unwrap();
1056 let bogus = peer
1058 .send_control(&mut s, 1, ControlOpcode::ReadyForTransition, 7, 1, vec![])
1059 .unwrap();
1060 let evs = coord.on_wire(&mut s, &bogus.wire).unwrap();
1061 let errs = drain_errs(&evs);
1062 assert!(errs.contains(&codes::TRANSITION_MISMATCH), "got {:?}", errs);
1063 }
1064
1065 #[test]
1066 fn execute_advances_epoch_and_clears_pending() {
1067 let mut coord = GroupNode::new(1, group_id());
1068 let mut peer = GroupNode::new(2, group_id());
1069 coord.bootstrap_as_creator(0);
1070 peer.bootstrap_as_joiner(0, 0);
1071 let mut s = PlainSealer;
1072 let prep = coord
1073 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1074 .unwrap();
1075 peer.on_wire(&mut s, &prep.wire).unwrap();
1076 let exec = coord
1078 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1079 .unwrap();
1080 coord.apply_transition(1);
1081 let evs = peer.on_wire(&mut s, &exec.wire).unwrap();
1082 assert_eq!(coord.last_transition_id, 1);
1083 assert_eq!(coord.current_epoch, 1);
1084 assert_eq!(peer.last_transition_id, 1);
1085 assert_eq!(peer.current_epoch, 1);
1086 assert_eq!(peer.pending_transition_id, 0);
1087 assert!(evs.iter().any(|e| matches!(
1088 e,
1089 Event::EpochAdvanced {
1090 transition_id: 1,
1091 ..
1092 }
1093 )));
1094 }
1095
1096 #[test]
1097 fn abort_clears_pending_no_advance() {
1098 let mut coord = GroupNode::new(1, group_id());
1099 let mut peer = GroupNode::new(2, group_id());
1100 coord.bootstrap_as_creator(0);
1101 peer.bootstrap_as_joiner(0, 0);
1102 let mut s = PlainSealer;
1103 let prep = coord
1104 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1105 .unwrap();
1106 peer.on_wire(&mut s, &prep.wire).unwrap();
1107 let abort = coord
1108 .send_control(&mut s, 0, ControlOpcode::AbortTransition, 1, 2, vec![])
1109 .unwrap();
1110 peer.on_wire(&mut s, &abort.wire).unwrap();
1111 assert_eq!(peer.pending_transition_id, 0);
1112 assert_eq!(peer.current_epoch, 0);
1113 assert_eq!(peer.transition_state, TransitionState::TAborted);
1114 assert_eq!(coord.transition_state, TransitionState::TAborted);
1115 }
1116
1117 #[test]
1118 fn bootstrap_as_joiner_with_expected_tid_accepts_first_execute() {
1119 let mut coord = GroupNode::new(1, group_id());
1120 let mut joiner = GroupNode::new(2, group_id());
1122 coord.bootstrap_as_creator(0);
1123 joiner.bootstrap_as_joiner(0, 1);
1124 assert_eq!(joiner.pending_transition_id, 1);
1125 let mut s = PlainSealer;
1126 let _ = coord
1128 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1129 .unwrap();
1130 let exec = coord
1132 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1133 .unwrap();
1134 let evs = joiner.on_wire(&mut s, &exec.wire).unwrap();
1135 let errs = drain_errs(&evs);
1136 assert!(
1137 errs.is_empty(),
1138 "expected clean apply, got errors {:?}",
1139 errs
1140 );
1141 assert_eq!(joiner.last_transition_id, 1);
1142 assert_eq!(joiner.current_epoch, 1);
1143 }
1144
1145 #[test]
1148 fn claim_coordinator_sets_flag_and_emits_event() {
1149 let mut node = GroupNode::new(1, group_id());
1150 node.bootstrap_as_creator(0);
1151 node.drain_events();
1152 let mut s = PlainSealer;
1153 let _ = node.claim_coordinator(&mut s, 0).unwrap();
1154 assert!(node.is_coordinator);
1155 let evs = node.drain_events();
1156 assert!(evs.iter().any(|e| matches!(e, Event::BecameCoordinator)));
1157 }
1158
1159 #[test]
1160 fn coordinator_gone_emits_election_needed() {
1161 let mut member = GroupNode::new(2, group_id());
1162 member.bootstrap_as_joiner(0, 0);
1163 member.coordinator_last_seen = Some(Instant::now() - Duration::from_millis(11_000));
1164 let evs = member.check_timeouts();
1165 assert!(
1166 evs.iter()
1167 .any(|e| matches!(e, Event::CoordinatorElectionNeeded))
1168 );
1169 assert!(!member.is_coordinator, "flag cleared on silence");
1170 }
1171
1172 #[test]
1173 fn capabilities_advertise_with_claim_resets_silence_timer() {
1174 let mut member = GroupNode::new(2, group_id());
1175 let mut coord = GroupNode::new(1, group_id());
1176 member.bootstrap_as_joiner(0, 0);
1177 coord.bootstrap_as_creator(0);
1178 let mut s = PlainSealer;
1179 let f = coord.claim_coordinator(&mut s, 2).unwrap();
1181 let evs = member.on_wire(&mut s, &f.wire).unwrap();
1183 assert!(
1184 member.coordinator_last_seen.is_some(),
1185 "silence timer reset"
1186 );
1187 assert!(
1188 evs.iter()
1189 .any(|e| matches!(e, Event::CoordinatorClaim { claimant: 1 }))
1190 );
1191 }
1192
1193 #[test]
1194 fn higher_id_yields_to_lower_claimant() {
1195 let mut node5 = GroupNode::new(5, group_id());
1197 let mut node2 = GroupNode::new(2, group_id());
1198 node5.bootstrap_as_joiner(0, 0);
1199 node2.bootstrap_as_creator(0);
1200 let mut s = PlainSealer;
1201 node5.is_coordinator = true;
1203 let f = node2.claim_coordinator(&mut s, 5).unwrap();
1205 node5.on_wire(&mut s, &f.wire).unwrap();
1206 assert!(!node5.is_coordinator, "node5 yielded to node2");
1207 }
1208
1209 #[test]
1210 fn lower_id_keeps_coordinator_against_higher_claimant() {
1211 let mut node1 = GroupNode::new(1, group_id());
1212 let mut node5 = GroupNode::new(5, group_id());
1213 node1.bootstrap_as_creator(0);
1214 node5.bootstrap_as_joiner(0, 0);
1215 let mut s = PlainSealer;
1216 node1.is_coordinator = true;
1217 let f = node5.claim_coordinator(&mut s, 1).unwrap();
1218 node1.on_wire(&mut s, &f.wire).unwrap();
1219 assert!(node1.is_coordinator, "node1 keeps role — it has lower id");
1220 }
1221
1222 #[test]
1225 fn competing_prepare_lower_member_id_wins() {
1226 let mut node = GroupNode::new(10, group_id());
1229 node.bootstrap_as_joiner(0, 0);
1230 let mut s = PlainSealer;
1231
1232 let mut sender1 = GroupNode::new(1, group_id());
1234 sender1.bootstrap_as_creator(0);
1235 let f1 = sender1
1236 .send_control(
1237 &mut s,
1238 10,
1239 ControlOpcode::PrepareTransition,
1240 1,
1241 1,
1242 b"commit-A".to_vec(),
1243 )
1244 .unwrap();
1245 node.on_wire(&mut s, &f1.wire).unwrap();
1246 assert_eq!(
1247 node.pending_commit_sender,
1248 Some(1),
1249 "member 1 is initial winner"
1250 );
1251
1252 let mut sender3 = GroupNode::new(3, group_id());
1254 sender3.bootstrap_as_creator(0);
1255 let f3 = sender3
1256 .send_control(
1257 &mut s,
1258 10,
1259 ControlOpcode::PrepareTransition,
1260 1,
1261 2,
1262 b"commit-B".to_vec(),
1263 )
1264 .unwrap();
1265 node.on_wire(&mut s, &f3.wire).unwrap();
1266 assert_eq!(node.pending_commit_sender, Some(1), "member 1 still wins");
1268 assert_eq!(node.pending_transition_id, 1);
1269 }
1270
1271 #[test]
1272 fn competing_prepare_later_lower_id_displaces_winner() {
1273 let mut node = GroupNode::new(10, group_id());
1275 node.bootstrap_as_joiner(0, 0);
1276 let mut s = PlainSealer;
1277
1278 let mut sender5 = GroupNode::new(5, group_id());
1279 sender5.bootstrap_as_creator(0);
1280 let f5 = sender5
1281 .send_control(
1282 &mut s,
1283 10,
1284 ControlOpcode::PrepareTransition,
1285 1,
1286 1,
1287 b"commit-X".to_vec(),
1288 )
1289 .unwrap();
1290 node.on_wire(&mut s, &f5.wire).unwrap();
1291 assert_eq!(node.pending_commit_sender, Some(5));
1292
1293 let mut sender2 = GroupNode::new(2, group_id());
1294 sender2.bootstrap_as_creator(0);
1295 let f2 = sender2
1296 .send_control(
1297 &mut s,
1298 10,
1299 ControlOpcode::PrepareTransition,
1300 1,
1301 2,
1302 b"commit-Y".to_vec(),
1303 )
1304 .unwrap();
1305 node.on_wire(&mut s, &f2.wire).unwrap();
1306 assert_eq!(
1307 node.pending_commit_sender,
1308 Some(2),
1309 "member 2 displaces member 5"
1310 );
1311 }
1312
1313 #[test]
1314 fn apply_transition_clears_commit_sender() {
1315 let mut coord = GroupNode::new(1, group_id());
1316 coord.bootstrap_as_creator(0);
1317 let mut s = PlainSealer;
1318 coord
1319 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1320 .unwrap();
1321 coord.apply_transition(1);
1322 assert_eq!(coord.pending_commit_sender, None);
1323 }
1324
1325 #[test]
1328 fn prepare_timeout_fires_when_deadline_exceeded() {
1329 let mut coord = GroupNode::new(1, group_id());
1330 coord.bootstrap_as_creator(0);
1331 let mut s = PlainSealer;
1332 coord
1333 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1334 .unwrap();
1335 coord.prepare_deadline = Some(Instant::now() - Duration::from_millis(1));
1337 let evs = coord.check_timeouts();
1338 assert!(
1339 evs.iter().any(|e| matches!(
1340 e,
1341 Event::Error {
1342 code: codes::PREPARE_TIMEOUT,
1343 ..
1344 }
1345 )),
1346 "expected PREPARE_TIMEOUT, got {:?}",
1347 evs
1348 );
1349 assert_eq!(
1350 coord.transition_state,
1351 TransitionState::TAborted,
1352 "transition aborted"
1353 );
1354 assert_eq!(coord.prepare_deadline, None, "deadline cleared");
1355 }
1356
1357 #[test]
1358 fn execute_timeout_fires_when_deadline_exceeded() {
1359 let mut member = GroupNode::new(2, group_id());
1360 member.bootstrap_as_joiner(0, 0);
1361 let mut s = PlainSealer;
1362 member.pending_transition_id = 1;
1364 member.transition_state = TransitionState::TPrepared;
1365 member
1366 .send_control(&mut s, 1, ControlOpcode::ReadyForTransition, 1, 1, vec![])
1367 .unwrap();
1368 member.execute_deadline = Some(Instant::now() - Duration::from_millis(1));
1370 let evs = member.check_timeouts();
1371 assert!(
1372 evs.iter().any(|e| matches!(
1373 e,
1374 Event::Error {
1375 code: codes::EXECUTE_TIMEOUT,
1376 ..
1377 }
1378 )),
1379 "expected EXECUTE_TIMEOUT, got {:?}",
1380 evs
1381 );
1382 assert_eq!(member.execute_deadline, None, "deadline cleared");
1383 }
1384
1385 #[test]
1386 fn coordinator_gone_fires_after_silence() {
1387 let mut member = GroupNode::new(2, group_id());
1388 member.bootstrap_as_joiner(0, 0);
1389 member.coordinator_last_seen = Some(Instant::now() - Duration::from_millis(11_000));
1391 let evs = member.check_timeouts();
1392 assert!(
1393 evs.iter().any(|e| matches!(
1394 e,
1395 Event::Error {
1396 code: codes::COORDINATOR_GONE,
1397 ..
1398 }
1399 )),
1400 "expected COORDINATOR_GONE, got {:?}",
1401 evs
1402 );
1403 assert_eq!(member.coordinator_last_seen, None, "timer cleared");
1404 }
1405
1406 #[test]
1407 fn note_coordinator_activity_resets_silence_timer() {
1408 let mut member = GroupNode::new(2, group_id());
1409 member.bootstrap_as_joiner(0, 0);
1410 member.coordinator_last_seen = Some(Instant::now() - Duration::from_millis(11_000));
1412 member.note_coordinator_activity();
1414 let evs = member.check_timeouts();
1415 assert!(
1416 !evs.iter().any(|e| matches!(
1417 e,
1418 Event::Error {
1419 code: codes::COORDINATOR_GONE,
1420 ..
1421 }
1422 )),
1423 "should NOT fire after reset"
1424 );
1425 }
1426
1427 #[test]
1428 fn execute_clears_prepare_deadline() {
1429 let mut coord = GroupNode::new(1, group_id());
1430 coord.bootstrap_as_creator(0);
1431 let mut s = PlainSealer;
1432 coord
1433 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1434 .unwrap();
1435 assert!(coord.prepare_deadline.is_some(), "deadline armed");
1436 coord
1437 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1438 .unwrap();
1439 assert_eq!(coord.prepare_deadline, None, "deadline cleared on EXECUTE");
1440 assert_eq!(
1441 coord.execute_deadline, None,
1442 "execute_deadline also cleared"
1443 );
1444 }
1445
1446 #[test]
1447 fn receive_prepare_arms_execute_deadline() {
1448 let mut coord = GroupNode::new(1, group_id());
1449 let mut member = GroupNode::new(2, group_id());
1450 coord.bootstrap_as_creator(0);
1451 member.bootstrap_as_joiner(0, 0);
1452 let mut s = PlainSealer;
1453 let f = coord
1454 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1455 .unwrap();
1456 member.on_wire(&mut s, &f.wire).unwrap();
1457 assert!(
1458 member.execute_deadline.is_some(),
1459 "execute_deadline armed on receiving PREPARE"
1460 );
1461 }
1462
1463 #[test]
1464 fn receive_execute_clears_execute_deadline() {
1465 let mut coord = GroupNode::new(1, group_id());
1466 let mut member = GroupNode::new(2, group_id());
1467 coord.bootstrap_as_creator(0);
1468 member.bootstrap_as_joiner(0, 0);
1469 let mut s = PlainSealer;
1470 let prep = coord
1471 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1472 .unwrap();
1473 member.on_wire(&mut s, &prep.wire).unwrap();
1474 let exec = coord
1475 .send_control(&mut s, 0, ControlOpcode::ExecuteTransition, 1, 2, vec![])
1476 .unwrap();
1477 member.on_wire(&mut s, &exec.wire).unwrap();
1478 assert_eq!(member.execute_deadline, None, "cleared on EXECUTE");
1479 }
1480
1481 #[test]
1482 fn no_timeout_when_deadlines_not_set() {
1483 let mut node = GroupNode::new(1, group_id());
1484 node.bootstrap_as_creator(0);
1485 node.drain_events(); let evs = node.check_timeouts();
1487 assert!(evs.is_empty(), "no events without armed deadlines");
1488 }
1489
1490 #[test]
1491 fn prepare_with_already_applied_tid_is_rejected() {
1492 let mut coord = GroupNode::new(1, group_id());
1495 coord.bootstrap_as_creator(0);
1496 let mut s = PlainSealer;
1497 let _ = coord
1498 .send_control(&mut s, 0, ControlOpcode::PrepareTransition, 1, 1, vec![])
1499 .unwrap();
1500 coord.apply_transition(1);
1501 assert_eq!(coord.last_transition_id, 1);
1502 assert_eq!(coord.pending_transition_id, 0);
1503 let mut peer = GroupNode::new(2, group_id());
1507 peer.bootstrap_as_joiner(coord.current_epoch, 0);
1508 let stale = peer
1509 .send_control(&mut s, 1, ControlOpcode::PrepareTransition, 1, 9, vec![])
1510 .unwrap();
1511 let evs = coord.on_wire(&mut s, &stale.wire).unwrap();
1512 let errs = drain_errs(&evs);
1513 assert!(
1514 errs.contains(&codes::TRANSITION_MISMATCH),
1515 "expected TRANSITION_MISMATCH, got {:?}",
1516 errs
1517 );
1518 }
1519
1520 #[test]
1521 fn decrypt_failed_is_non_fatal() {
1522 struct OpenFailSealer;
1524 impl Sealer for OpenFailSealer {
1525 fn seal(
1526 &mut self,
1527 _: StreamType,
1528 _: SequenceNo,
1529 p: &[u8],
1530 ) -> Result<Vec<u8>, MlsError> {
1531 Ok(p.to_vec())
1532 }
1533 fn open(
1534 &mut self,
1535 _: StreamType,
1536 _: SequenceNo,
1537 _: &[u8],
1538 ) -> Result<Vec<u8>, MlsError> {
1539 Err(MlsError::Aead("simulated".into()))
1540 }
1541 }
1542 let mut alice = GroupNode::new(1, group_id());
1543 let mut bob = GroupNode::new(2, group_id());
1544 alice.bootstrap_as_creator(1);
1545 bob.bootstrap_as_joiner(1, 0);
1546 let mut s = PlainSealer;
1547 let sid = alice.member_stream_id(2);
1548 let f = alice
1549 .send_payload(
1550 &mut s,
1551 2,
1552 StreamType::Text,
1553 sid,
1554 GbpFlags::ordered_reliable_ack(),
1555 b"x",
1556 PayloadCodec::Cbor,
1557 )
1558 .unwrap();
1559 let mut fail = OpenFailSealer;
1560 let evs = bob.on_wire(&mut fail, &f.wire).unwrap();
1561 let err = evs
1562 .iter()
1563 .find_map(|e| match e {
1564 Event::Error {
1565 code,
1566 fatal,
1567 retryable,
1568 ..
1569 } => Some((*code, *fatal, *retryable)),
1570 _ => None,
1571 })
1572 .expect("error event");
1573 assert_eq!(err.0, codes::DECRYPT_FAILED);
1574 assert!(!err.1, "must be non-fatal");
1575 assert!(err.2, "must be retryable");
1576 assert_eq!(bob.state, NodeState::Active, "bob stays Active");
1577 }
1578}