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