1use crate::mode::util::validate_commitment_payload_for_session;
2use crate::mode::{Mode, ModeResponse};
3use macp_core::error::MacpError;
4use macp_core::session::Session;
5use macp_pb::pb::Envelope;
6use serde::{Deserialize, Serialize};
7use std::collections::BTreeMap;
8
9#[derive(Debug, Clone, Serialize, Deserialize)]
29#[non_exhaustive]
30pub struct MultiRoundState {
31 pub round: u64,
32 pub participants: Vec<String>,
33 pub contributions: BTreeMap<String, String>,
34 #[serde(default)]
35 pub convergence_type: String,
36 #[serde(default)]
37 pub converged: bool,
38}
39
40#[derive(Debug, Clone, Deserialize)]
42struct ContributeJson {
43 value: String,
44}
45
46#[doc(hidden)]
63pub fn parse_contribute_value(payload: &[u8]) -> Result<String, MacpError> {
64 if payload.is_empty() {
68 return Err(MacpError::InvalidPayload);
69 }
70 if let Ok(text) = std::str::from_utf8(payload) {
71 if let Ok(c) = serde_json::from_str::<ContributeJson>(text) {
72 return Ok(c.value);
73 }
74 }
75 <macp_pb::multi_round_pb::ContributePayload as prost::Message>::decode(payload)
76 .map(|c| c.value)
77 .map_err(|_| MacpError::InvalidPayload)
78}
79
80#[derive(Debug, Serialize)]
82struct ResolutionPayload {
83 converged_value: String,
84 round: u64,
85 #[serde(rename = "final")]
86 final_values: BTreeMap<String, String>,
87}
88
89pub struct MultiRoundMode;
90
91impl MultiRoundMode {
92 fn encode_state(state: &MultiRoundState) -> Vec<u8> {
93 crate::mode::util::encode_mode_state(state)
94 }
95
96 fn decode_state(data: &[u8]) -> Result<MultiRoundState, MacpError> {
97 crate::mode::util::decode_mode_state(data)
98 }
99
100 fn check_convergence(state: &MultiRoundState) -> bool {
101 let all_contributed = state
102 .participants
103 .iter()
104 .all(|p| state.contributions.contains_key(p));
105
106 if !all_contributed {
107 return false;
108 }
109
110 let values: Vec<&String> = state.contributions.values().collect();
111 values.windows(2).all(|w| w[0] == w[1])
112 }
113}
114
115impl Mode for MultiRoundMode {
116 fn on_session_start(
117 &self,
118 session: &Session,
119 _env: &Envelope,
120 ) -> Result<ModeResponse, MacpError> {
121 let participants = session.participants.clone();
122
123 if participants.is_empty() {
124 return Err(MacpError::InvalidPayload);
125 }
126
127 let state = MultiRoundState {
128 round: 0,
129 participants,
130 contributions: BTreeMap::new(),
131 convergence_type: "all_equal".into(),
132 converged: false,
133 };
134
135 Ok(ModeResponse::PersistState(Self::encode_state(&state)))
136 }
137
138 fn on_message(&self, session: &Session, env: &Envelope) -> Result<ModeResponse, MacpError> {
139 match env.message_type.as_str() {
140 "Contribute" => self.handle_contribute(session, env),
141 "Commitment" => self.handle_commitment(session, env),
142 _ => Err(MacpError::InvalidPayload),
143 }
144 }
145
146 fn authorize_sender(&self, session: &Session, env: &Envelope) -> Result<(), MacpError> {
147 if env.message_type == "Commitment" {
148 if env.sender != session.initiator_sender {
150 return Err(MacpError::Forbidden);
151 }
152 return Ok(());
153 }
154 if !session.participants.is_empty() && !session.participants.contains(&env.sender) {
156 return Err(MacpError::Forbidden);
157 }
158 Ok(())
159 }
160}
161
162impl MultiRoundMode {
163 fn handle_contribute(
164 &self,
165 session: &Session,
166 env: &Envelope,
167 ) -> Result<ModeResponse, MacpError> {
168 let mut state = Self::decode_state(&session.mode_state)?;
169
170 if state.converged {
171 return Err(MacpError::InvalidPayload);
172 }
173
174 let value = parse_contribute_value(&env.payload)?;
175
176 let previous = state.contributions.get(&env.sender);
177 let value_changed = previous.is_none_or(|prev| *prev != value);
178
179 if value_changed {
180 state.round += 1;
181 state.contributions.insert(env.sender.clone(), value);
182 }
183
184 if Self::check_convergence(&state) {
185 state.converged = true;
186 }
187
188 Ok(ModeResponse::PersistState(Self::encode_state(&state)))
189 }
190
191 fn handle_commitment(
192 &self,
193 session: &Session,
194 env: &Envelope,
195 ) -> Result<ModeResponse, MacpError> {
196 let state = Self::decode_state(&session.mode_state)?;
197
198 if !state.converged {
199 return Err(MacpError::InvalidPayload);
200 }
201
202 validate_commitment_payload_for_session(session, &env.payload)?;
203
204 let converged_value = state
205 .contributions
206 .values()
207 .next()
208 .cloned()
209 .unwrap_or_default();
210 let resolution = ResolutionPayload {
211 converged_value,
212 round: state.round,
213 final_values: state.contributions.clone(),
214 };
215 let resolution_bytes =
216 serde_json::to_vec(&resolution).expect("ResolutionPayload is always serializable");
217
218 Ok(ModeResponse::PersistAndResolve {
219 state: Self::encode_state(&state),
220 resolution: resolution_bytes,
221 })
222 }
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228
229 use macp_pb::pb::CommitmentPayload;
230 use prost::Message;
231
232 fn base_session() -> Session {
233 Session::builder("s1", "ext.multi_round.v1", "coordinator")
234 .ttl_ms(60_000)
235 .mode_version("1.0.0")
236 .configuration_version("cfg-1")
237 .build()
238 }
239
240 fn session_start_env() -> Envelope {
241 Envelope {
242 macp_version: "1.0".into(),
243 mode: "ext.multi_round.v1".into(),
244 message_type: "SessionStart".into(),
245 message_id: "m0".into(),
246 session_id: "s1".into(),
247 sender: "coordinator".into(),
248 timestamp_unix_ms: 1_700_000_000_000,
249 payload: vec![],
250 }
251 }
252
253 fn contribute_env_with_payload(sender: &str, payload: Vec<u8>) -> Envelope {
254 Envelope {
255 macp_version: "1.0".into(),
256 mode: "ext.multi_round.v1".into(),
257 message_type: "Contribute".into(),
258 message_id: format!("m_{}", sender),
259 session_id: "s1".into(),
260 sender: sender.into(),
261 timestamp_unix_ms: 1_700_000_000_000,
262 payload,
263 }
264 }
265
266 fn contribute_env(sender: &str, value: &str) -> Envelope {
268 let payload = macp_pb::multi_round_pb::ContributePayload {
269 value: value.into(),
270 }
271 .encode_to_vec();
272 contribute_env_with_payload(sender, payload)
273 }
274
275 fn contribute_env_json(sender: &str, value: &str) -> Envelope {
277 let payload = serde_json::json!({"value": value}).to_string();
278 contribute_env_with_payload(sender, payload.into_bytes())
279 }
280
281 fn commitment_env(sender: &str) -> Envelope {
282 let payload = CommitmentPayload {
283 commitment_id: "c1".into(),
284 action: "multi_round.converged".into(),
285 authority_scope: "test".into(),
286 reason: "converged".into(),
287 mode_version: "1.0.0".into(),
288 policy_version: String::new(),
289 configuration_version: "cfg-1".into(),
290 outcome_positive: true,
291 supersedes: None,
292 }
293 .encode_to_vec();
294 Envelope {
295 macp_version: "1.0".into(),
296 mode: "ext.multi_round.v1".into(),
297 message_type: "Commitment".into(),
298 message_id: "m_commit".into(),
299 session_id: "s1".into(),
300 sender: sender.into(),
301 timestamp_unix_ms: 1_700_000_000_000,
302 payload,
303 }
304 }
305
306 fn session_with_state(state: &MultiRoundState) -> Session {
307 let mut s = base_session();
308 s.mode_state = MultiRoundMode::encode_state(state);
309 s.participants = state.participants.clone();
310 s
311 }
312
313 #[test]
314 fn session_start_parses_valid_config() {
315 let mode = MultiRoundMode;
316 let mut session = base_session();
317 session.participants = vec!["alice".into(), "bob".into()];
318 let env = session_start_env();
319
320 let result = mode.on_session_start(&session, &env).unwrap();
321 match result {
322 ModeResponse::PersistState(data) => {
323 let state: MultiRoundState = serde_json::from_slice(&data).unwrap();
324 assert_eq!(state.round, 0);
325 assert_eq!(state.participants, vec!["alice", "bob"]);
326 assert!(state.contributions.is_empty());
327 assert!(!state.converged);
328 }
329 _ => panic!("Expected PersistState"),
330 }
331 }
332
333 #[test]
334 fn session_start_rejects_empty_participants() {
335 let mode = MultiRoundMode;
336 let session = base_session();
337 let env = session_start_env();
338
339 let err = mode.on_session_start(&session, &env).unwrap_err();
340 assert_eq!(err.to_string(), "InvalidPayload");
341 }
342
343 #[test]
344 fn contribute_first_value_increments_round() {
345 let mode = MultiRoundMode;
346 let state = MultiRoundState {
347 round: 0,
348 participants: vec!["alice".into(), "bob".into()],
349 contributions: BTreeMap::new(),
350 convergence_type: "all_equal".into(),
351 converged: false,
352 };
353 let session = session_with_state(&state);
354 let env = contribute_env("alice", "option_a");
355
356 let result = mode.on_message(&session, &env).unwrap();
357 match result {
358 ModeResponse::PersistState(data) => {
359 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
360 assert_eq!(new_state.round, 1);
361 assert_eq!(new_state.contributions.get("alice").unwrap(), "option_a");
362 assert!(!new_state.converged);
363 }
364 _ => panic!("Expected PersistState"),
365 }
366 }
367
368 #[test]
369 fn resubmit_same_value_does_not_increment_round() {
370 let mode = MultiRoundMode;
371 let mut contributions = BTreeMap::new();
372 contributions.insert("alice".to_string(), "option_a".to_string());
373 let state = MultiRoundState {
374 round: 1,
375 participants: vec!["alice".into(), "bob".into()],
376 contributions,
377 convergence_type: "all_equal".into(),
378 converged: false,
379 };
380 let session = session_with_state(&state);
381 let env = contribute_env("alice", "option_a");
382
383 let result = mode.on_message(&session, &env).unwrap();
384 match result {
385 ModeResponse::PersistState(data) => {
386 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
387 assert_eq!(new_state.round, 1);
388 }
389 _ => panic!("Expected PersistState"),
390 }
391 }
392
393 #[test]
394 fn revise_value_increments_round() {
395 let mode = MultiRoundMode;
396 let mut contributions = BTreeMap::new();
397 contributions.insert("alice".to_string(), "option_a".to_string());
398 let state = MultiRoundState {
399 round: 1,
400 participants: vec!["alice".into(), "bob".into()],
401 contributions,
402 convergence_type: "all_equal".into(),
403 converged: false,
404 };
405 let session = session_with_state(&state);
406 let env = contribute_env("alice", "option_b");
407
408 let result = mode.on_message(&session, &env).unwrap();
409 match result {
410 ModeResponse::PersistState(data) => {
411 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
412 assert_eq!(new_state.round, 2);
413 assert_eq!(new_state.contributions.get("alice").unwrap(), "option_b");
414 }
415 _ => panic!("Expected PersistState"),
416 }
417 }
418
419 #[test]
420 fn convergence_sets_converged_flag() {
421 let mode = MultiRoundMode;
422 let mut contributions = BTreeMap::new();
423 contributions.insert("alice".to_string(), "option_a".to_string());
424 let state = MultiRoundState {
425 round: 1,
426 participants: vec!["alice".into(), "bob".into()],
427 contributions,
428 convergence_type: "all_equal".into(),
429 converged: false,
430 };
431 let session = session_with_state(&state);
432 let env = contribute_env("bob", "option_a");
433
434 let result = mode.on_message(&session, &env).unwrap();
435 match result {
436 ModeResponse::PersistState(data) => {
437 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
438 assert_eq!(new_state.round, 2);
439 assert!(new_state.converged);
440 }
441 _ => panic!("Expected PersistState (convergence tracked, not auto-resolved)"),
442 }
443 }
444
445 #[test]
446 fn commitment_after_convergence_resolves() {
447 let mode = MultiRoundMode;
448 let mut contributions = BTreeMap::new();
449 contributions.insert("alice".to_string(), "option_a".to_string());
450 contributions.insert("bob".to_string(), "option_a".to_string());
451 let state = MultiRoundState {
452 round: 2,
453 participants: vec!["alice".into(), "bob".into()],
454 contributions,
455 convergence_type: "all_equal".into(),
456 converged: true,
457 };
458 let session = session_with_state(&state);
459 let env = commitment_env("coordinator");
460
461 let result = mode.on_message(&session, &env).unwrap();
462 match result {
463 ModeResponse::PersistAndResolve { resolution, .. } => {
464 let res: serde_json::Value = serde_json::from_slice(&resolution).unwrap();
465 assert_eq!(res["converged_value"], "option_a");
466 assert_eq!(res["round"], 2);
467 }
468 _ => panic!("Expected PersistAndResolve"),
469 }
470 }
471
472 #[test]
473 fn commitment_before_convergence_rejected() {
474 let mode = MultiRoundMode;
475 let state = MultiRoundState {
476 round: 0,
477 participants: vec!["alice".into(), "bob".into()],
478 contributions: BTreeMap::new(),
479 convergence_type: "all_equal".into(),
480 converged: false,
481 };
482 let session = session_with_state(&state);
483 let env = commitment_env("coordinator");
484
485 let err = mode.on_message(&session, &env).unwrap_err();
486 assert_eq!(err.to_string(), "InvalidPayload");
487 }
488
489 #[test]
490 fn contribute_after_convergence_rejected() {
491 let mode = MultiRoundMode;
492 let mut contributions = BTreeMap::new();
493 contributions.insert("alice".to_string(), "option_a".to_string());
494 contributions.insert("bob".to_string(), "option_a".to_string());
495 let state = MultiRoundState {
496 round: 2,
497 participants: vec!["alice".into(), "bob".into()],
498 contributions,
499 convergence_type: "all_equal".into(),
500 converged: true,
501 };
502 let session = session_with_state(&state);
503 let env = contribute_env("alice", "option_b");
504
505 let err = mode.on_message(&session, &env).unwrap_err();
506 assert_eq!(err.to_string(), "InvalidPayload");
507 }
508
509 #[test]
510 fn non_initiator_commitment_rejected() {
511 let mode = MultiRoundMode;
512 let mut contributions = BTreeMap::new();
513 contributions.insert("alice".to_string(), "option_a".to_string());
514 contributions.insert("bob".to_string(), "option_a".to_string());
515 let state = MultiRoundState {
516 round: 2,
517 participants: vec!["alice".into(), "bob".into()],
518 contributions,
519 convergence_type: "all_equal".into(),
520 converged: true,
521 };
522 let session = session_with_state(&state);
523 let env = commitment_env("alice"); let err = mode.authorize_sender(&session, &env).unwrap_err();
526 assert_eq!(err.to_string(), "Forbidden");
527 }
528
529 #[test]
530 fn no_convergence_when_values_differ() {
531 let mode = MultiRoundMode;
532 let mut contributions = BTreeMap::new();
533 contributions.insert("alice".to_string(), "option_a".to_string());
534 let state = MultiRoundState {
535 round: 1,
536 participants: vec!["alice".into(), "bob".into()],
537 contributions,
538 convergence_type: "all_equal".into(),
539 converged: false,
540 };
541 let session = session_with_state(&state);
542 let env = contribute_env("bob", "option_b");
543
544 let result = mode.on_message(&session, &env).unwrap();
545 match result {
546 ModeResponse::PersistState(data) => {
547 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
548 assert!(!new_state.converged);
549 }
550 _ => panic!("Expected PersistState"),
551 }
552 }
553
554 #[test]
555 fn no_convergence_when_not_all_contributed() {
556 let mode = MultiRoundMode;
557 let state = MultiRoundState {
558 round: 0,
559 participants: vec!["alice".into(), "bob".into(), "carol".into()],
560 contributions: BTreeMap::new(),
561 convergence_type: "all_equal".into(),
562 converged: false,
563 };
564 let session = session_with_state(&state);
565 let env = contribute_env("alice", "option_a");
566
567 let result = mode.on_message(&session, &env).unwrap();
568 assert!(matches!(result, ModeResponse::PersistState(_)));
569 }
570
571 #[test]
572 fn non_contribute_message_rejected() {
573 let mode = MultiRoundMode;
574 let state = MultiRoundState {
575 round: 0,
576 participants: vec!["alice".into()],
577 contributions: BTreeMap::new(),
578 convergence_type: "all_equal".into(),
579 converged: false,
580 };
581 let session = session_with_state(&state);
582 let env = Envelope {
583 macp_version: "1.0".into(),
584 mode: "ext.multi_round.v1".into(),
585 message_type: "Message".into(),
586 message_id: "m1".into(),
587 session_id: "s1".into(),
588 sender: "alice".into(),
589 timestamp_unix_ms: 1_700_000_000_000,
590 payload: b"hello".to_vec(),
591 };
592
593 let err = mode.on_message(&session, &env).unwrap_err();
594 assert_eq!(err.error_code(), "INVALID_ENVELOPE");
595 }
596
597 #[test]
598 fn contribute_invalid_payload_returns_error() {
599 let mode = MultiRoundMode;
600 let state = MultiRoundState {
601 round: 0,
602 participants: vec!["alice".into()],
603 contributions: BTreeMap::new(),
604 convergence_type: "all_equal".into(),
605 converged: false,
606 };
607 let session = session_with_state(&state);
608 let env = Envelope {
609 macp_version: "1.0".into(),
610 mode: "ext.multi_round.v1".into(),
611 message_type: "Contribute".into(),
612 message_id: "m1".into(),
613 session_id: "s1".into(),
614 sender: "alice".into(),
615 timestamp_unix_ms: 1_700_000_000_000,
616 payload: b"not json".to_vec(),
617 };
618
619 let err = mode.on_message(&session, &env).unwrap_err();
620 assert_eq!(err.to_string(), "InvalidPayload");
621 }
622
623 #[test]
627 fn contribute_json_fallback_still_accepted() {
628 let mode = MultiRoundMode;
629 let state = MultiRoundState {
630 round: 0,
631 participants: vec!["alice".into(), "bob".into()],
632 contributions: BTreeMap::new(),
633 convergence_type: "all_equal".into(),
634 converged: false,
635 };
636 let session = session_with_state(&state);
637
638 let result = mode
639 .on_message(&session, &contribute_env_json("alice", "option_a"))
640 .unwrap();
641 match result {
642 ModeResponse::PersistState(data) => {
643 let state: MultiRoundState = serde_json::from_slice(&data).unwrap();
644 assert_eq!(state.contributions["alice"], "option_a");
645 assert_eq!(state.round, 1);
646 }
647 _ => panic!("Expected PersistState"),
648 }
649 }
650
651 #[test]
654 fn proto_and_json_contributions_are_equivalent() {
655 let mode = MultiRoundMode;
656 let state = MultiRoundState {
657 round: 0,
658 participants: vec!["alice".into(), "bob".into()],
659 contributions: BTreeMap::new(),
660 convergence_type: "all_equal".into(),
661 converged: false,
662 };
663 let session = session_with_state(&state);
664
665 let after_json = match mode
666 .on_message(&session, &contribute_env_json("alice", "option_a"))
667 .unwrap()
668 {
669 ModeResponse::PersistState(data) => data,
670 _ => panic!("Expected PersistState"),
671 };
672 let session = {
673 let state: MultiRoundState = serde_json::from_slice(&after_json).unwrap();
674 session_with_state(&state)
675 };
676
677 match mode
679 .on_message(&session, &contribute_env("alice", "option_a"))
680 .unwrap()
681 {
682 ModeResponse::PersistState(data) => {
683 let state: MultiRoundState = serde_json::from_slice(&data).unwrap();
684 assert_eq!(state.round, 1, "unchanged value must not advance the round");
685 assert_eq!(state.contributions["alice"], "option_a");
686 }
687 _ => panic!("Expected PersistState"),
688 }
689 }
690
691 #[test]
694 fn contribute_empty_payload_rejected() {
695 let mode = MultiRoundMode;
696 let state = MultiRoundState {
697 round: 0,
698 participants: vec!["alice".into()],
699 contributions: BTreeMap::new(),
700 convergence_type: "all_equal".into(),
701 converged: false,
702 };
703 let session = session_with_state(&state);
704 let env = contribute_env_with_payload("alice", vec![]);
705
706 let err = mode.on_message(&session, &env).unwrap_err();
707 assert_eq!(err.to_string(), "InvalidPayload");
708 }
709
710 #[test]
711 fn encode_decode_round_trip() {
712 let mut contributions = BTreeMap::new();
713 contributions.insert("alice".into(), "value_a".into());
714 let original = MultiRoundState {
715 round: 5,
716 participants: vec!["alice".into(), "bob".into()],
717 contributions,
718 convergence_type: "all_equal".into(),
719 converged: true,
720 };
721
722 let encoded = MultiRoundMode::encode_state(&original);
723 let decoded = MultiRoundMode::decode_state(&encoded).unwrap();
724
725 assert_eq!(decoded.round, original.round);
726 assert_eq!(decoded.participants, original.participants);
727 assert_eq!(decoded.contributions, original.contributions);
728 assert_eq!(decoded.converged, original.converged);
729 }
730
731 #[test]
732 fn decode_invalid_state_returns_error() {
733 let err = MultiRoundMode::decode_state(b"garbage").unwrap_err();
734 assert_eq!(err.to_string(), "InvalidModeState");
735 }
736
737 #[test]
738 fn three_participant_convergence() {
739 let mode = MultiRoundMode;
740
741 let mut contributions = BTreeMap::new();
742 contributions.insert("alice".to_string(), "option_a".to_string());
743 contributions.insert("bob".to_string(), "option_a".to_string());
744 let state = MultiRoundState {
745 round: 2,
746 participants: vec!["alice".into(), "bob".into(), "carol".into()],
747 contributions,
748 convergence_type: "all_equal".into(),
749 converged: false,
750 };
751 let session = session_with_state(&state);
752 let env = contribute_env("carol", "option_a");
753
754 let result = mode.on_message(&session, &env).unwrap();
755 match result {
756 ModeResponse::PersistState(data) => {
757 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
758 assert!(new_state.converged);
759 }
760 _ => panic!("Expected PersistState with converged=true"),
761 }
762 }
763
764 #[test]
765 fn unknown_message_type_rejected() {
766 let mode = MultiRoundMode;
767 let state = MultiRoundState {
768 round: 0,
769 participants: vec!["alice".into(), "bob".into()],
770 contributions: BTreeMap::new(),
771 convergence_type: "all_equal".into(),
772 converged: false,
773 };
774 let session = session_with_state(&state);
775 let env = Envelope {
776 macp_version: "1.0".into(),
777 mode: "ext.multi_round.v1".into(),
778 message_type: "UnknownType".into(),
779 message_id: "msg-unknown".into(),
780 session_id: "s1".into(),
781 sender: "alice".into(),
782 timestamp_unix_ms: 0,
783 payload: vec![],
784 };
785 let err = mode.on_message(&session, &env).unwrap_err();
786 assert_eq!(err.error_code(), "INVALID_ENVELOPE");
787 }
788}