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)]
11pub struct MultiRoundState {
12 pub round: u64,
13 pub participants: Vec<String>,
14 pub contributions: BTreeMap<String, String>,
15 #[serde(default)]
16 pub convergence_type: String,
17 #[serde(default)]
18 pub converged: bool,
19}
20
21#[derive(Debug, Clone, Deserialize)]
23struct ContributePayload {
24 value: String,
25}
26
27#[derive(Debug, Serialize)]
29struct ResolutionPayload {
30 converged_value: String,
31 round: u64,
32 #[serde(rename = "final")]
33 final_values: BTreeMap<String, String>,
34}
35
36pub struct MultiRoundMode;
37
38impl MultiRoundMode {
39 fn encode_state(state: &MultiRoundState) -> Vec<u8> {
40 serde_json::to_vec(state).expect("MultiRoundState is always serializable")
41 }
42
43 fn decode_state(data: &[u8]) -> Result<MultiRoundState, MacpError> {
44 serde_json::from_slice(data).map_err(|_| MacpError::InvalidModeState)
45 }
46
47 fn check_convergence(state: &MultiRoundState) -> bool {
48 let all_contributed = state
49 .participants
50 .iter()
51 .all(|p| state.contributions.contains_key(p));
52
53 if !all_contributed {
54 return false;
55 }
56
57 let values: Vec<&String> = state.contributions.values().collect();
58 values.windows(2).all(|w| w[0] == w[1])
59 }
60}
61
62impl Mode for MultiRoundMode {
63 fn on_session_start(
64 &self,
65 session: &Session,
66 _env: &Envelope,
67 ) -> Result<ModeResponse, MacpError> {
68 let participants = session.participants.clone();
69
70 if participants.is_empty() {
71 return Err(MacpError::InvalidPayload);
72 }
73
74 let state = MultiRoundState {
75 round: 0,
76 participants,
77 contributions: BTreeMap::new(),
78 convergence_type: "all_equal".into(),
79 converged: false,
80 };
81
82 Ok(ModeResponse::PersistState(Self::encode_state(&state)))
83 }
84
85 fn on_message(&self, session: &Session, env: &Envelope) -> Result<ModeResponse, MacpError> {
86 match env.message_type.as_str() {
87 "Contribute" => self.handle_contribute(session, env),
88 "Commitment" => self.handle_commitment(session, env),
89 _ => Err(MacpError::InvalidPayload),
90 }
91 }
92
93 fn authorize_sender(&self, session: &Session, env: &Envelope) -> Result<(), MacpError> {
94 if env.message_type == "Commitment" {
95 if env.sender != session.initiator_sender {
97 return Err(MacpError::Forbidden);
98 }
99 return Ok(());
100 }
101 if !session.participants.is_empty() && !session.participants.contains(&env.sender) {
103 return Err(MacpError::Forbidden);
104 }
105 Ok(())
106 }
107}
108
109impl MultiRoundMode {
110 fn handle_contribute(
111 &self,
112 session: &Session,
113 env: &Envelope,
114 ) -> Result<ModeResponse, MacpError> {
115 let mut state = Self::decode_state(&session.mode_state)?;
116
117 if state.converged {
118 return Err(MacpError::InvalidPayload);
119 }
120
121 let text = std::str::from_utf8(&env.payload).map_err(|_| MacpError::InvalidPayload)?;
122 let contribute: ContributePayload =
123 serde_json::from_str(text).map_err(|_| MacpError::InvalidPayload)?;
124
125 let previous = state.contributions.get(&env.sender);
126 let value_changed = previous.is_none_or(|prev| *prev != contribute.value);
127
128 if value_changed {
129 state.round += 1;
130 state
131 .contributions
132 .insert(env.sender.clone(), contribute.value);
133 }
134
135 if Self::check_convergence(&state) {
136 state.converged = true;
137 }
138
139 Ok(ModeResponse::PersistState(Self::encode_state(&state)))
140 }
141
142 fn handle_commitment(
143 &self,
144 session: &Session,
145 env: &Envelope,
146 ) -> Result<ModeResponse, MacpError> {
147 let state = Self::decode_state(&session.mode_state)?;
148
149 if !state.converged {
150 return Err(MacpError::InvalidPayload);
151 }
152
153 validate_commitment_payload_for_session(session, &env.payload)?;
154
155 let converged_value = state
156 .contributions
157 .values()
158 .next()
159 .cloned()
160 .unwrap_or_default();
161 let resolution = ResolutionPayload {
162 converged_value,
163 round: state.round,
164 final_values: state.contributions.clone(),
165 };
166 let resolution_bytes =
167 serde_json::to_vec(&resolution).expect("ResolutionPayload is always serializable");
168
169 Ok(ModeResponse::PersistAndResolve {
170 state: Self::encode_state(&state),
171 resolution: resolution_bytes,
172 })
173 }
174}
175
176#[cfg(test)]
177mod tests {
178 use super::*;
179 use macp_core::session::SessionState;
180 use macp_pb::pb::CommitmentPayload;
181 use prost::Message;
182 use std::collections::HashSet;
183
184 fn base_session() -> Session {
185 Session {
186 session_id: "s1".into(),
187 state: SessionState::Open,
188 ttl_expiry: i64::MAX,
189 ttl_ms: 60_000,
190 started_at_unix_ms: 0,
191 resolution: None,
192 mode: "ext.multi_round.v1".into(),
193 mode_state: vec![],
194 participants: vec![],
195 seen_message_ids: HashSet::new(),
196 intent: String::new(),
197 mode_version: "1.0.0".into(),
198 configuration_version: "cfg-1".into(),
199 policy_version: String::new(),
200 context_id: String::new(),
201 extensions: std::collections::HashMap::new(),
202 roots: vec![],
203 initiator_sender: "coordinator".into(),
204 participant_message_counts: std::collections::HashMap::new(),
205 participant_last_seen: std::collections::HashMap::new(),
206 policy_definition: None,
207 suspended_at_ms: None,
208 accumulated_suspended_ms: 0,
209 }
210 }
211
212 fn session_start_env() -> Envelope {
213 Envelope {
214 macp_version: "1.0".into(),
215 mode: "ext.multi_round.v1".into(),
216 message_type: "SessionStart".into(),
217 message_id: "m0".into(),
218 session_id: "s1".into(),
219 sender: "coordinator".into(),
220 timestamp_unix_ms: 1_700_000_000_000,
221 payload: vec![],
222 }
223 }
224
225 fn contribute_env(sender: &str, value: &str) -> Envelope {
226 let payload = serde_json::json!({"value": value}).to_string();
227 Envelope {
228 macp_version: "1.0".into(),
229 mode: "ext.multi_round.v1".into(),
230 message_type: "Contribute".into(),
231 message_id: format!("m_{}", sender),
232 session_id: "s1".into(),
233 sender: sender.into(),
234 timestamp_unix_ms: 1_700_000_000_000,
235 payload: payload.into_bytes(),
236 }
237 }
238
239 fn commitment_env(sender: &str) -> Envelope {
240 let payload = CommitmentPayload {
241 commitment_id: "c1".into(),
242 action: "multi_round.converged".into(),
243 authority_scope: "test".into(),
244 reason: "converged".into(),
245 mode_version: "1.0.0".into(),
246 policy_version: String::new(),
247 configuration_version: "cfg-1".into(),
248 outcome_positive: true,
249 supersedes: None,
250 }
251 .encode_to_vec();
252 Envelope {
253 macp_version: "1.0".into(),
254 mode: "ext.multi_round.v1".into(),
255 message_type: "Commitment".into(),
256 message_id: "m_commit".into(),
257 session_id: "s1".into(),
258 sender: sender.into(),
259 timestamp_unix_ms: 1_700_000_000_000,
260 payload,
261 }
262 }
263
264 fn session_with_state(state: &MultiRoundState) -> Session {
265 let mut s = base_session();
266 s.mode_state = MultiRoundMode::encode_state(state);
267 s.participants = state.participants.clone();
268 s
269 }
270
271 #[test]
272 fn session_start_parses_valid_config() {
273 let mode = MultiRoundMode;
274 let mut session = base_session();
275 session.participants = vec!["alice".into(), "bob".into()];
276 let env = session_start_env();
277
278 let result = mode.on_session_start(&session, &env).unwrap();
279 match result {
280 ModeResponse::PersistState(data) => {
281 let state: MultiRoundState = serde_json::from_slice(&data).unwrap();
282 assert_eq!(state.round, 0);
283 assert_eq!(state.participants, vec!["alice", "bob"]);
284 assert!(state.contributions.is_empty());
285 assert!(!state.converged);
286 }
287 _ => panic!("Expected PersistState"),
288 }
289 }
290
291 #[test]
292 fn session_start_rejects_empty_participants() {
293 let mode = MultiRoundMode;
294 let session = base_session();
295 let env = session_start_env();
296
297 let err = mode.on_session_start(&session, &env).unwrap_err();
298 assert_eq!(err.to_string(), "InvalidPayload");
299 }
300
301 #[test]
302 fn contribute_first_value_increments_round() {
303 let mode = MultiRoundMode;
304 let state = MultiRoundState {
305 round: 0,
306 participants: vec!["alice".into(), "bob".into()],
307 contributions: BTreeMap::new(),
308 convergence_type: "all_equal".into(),
309 converged: false,
310 };
311 let session = session_with_state(&state);
312 let env = contribute_env("alice", "option_a");
313
314 let result = mode.on_message(&session, &env).unwrap();
315 match result {
316 ModeResponse::PersistState(data) => {
317 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
318 assert_eq!(new_state.round, 1);
319 assert_eq!(new_state.contributions.get("alice").unwrap(), "option_a");
320 assert!(!new_state.converged);
321 }
322 _ => panic!("Expected PersistState"),
323 }
324 }
325
326 #[test]
327 fn resubmit_same_value_does_not_increment_round() {
328 let mode = MultiRoundMode;
329 let mut contributions = BTreeMap::new();
330 contributions.insert("alice".to_string(), "option_a".to_string());
331 let state = MultiRoundState {
332 round: 1,
333 participants: vec!["alice".into(), "bob".into()],
334 contributions,
335 convergence_type: "all_equal".into(),
336 converged: false,
337 };
338 let session = session_with_state(&state);
339 let env = contribute_env("alice", "option_a");
340
341 let result = mode.on_message(&session, &env).unwrap();
342 match result {
343 ModeResponse::PersistState(data) => {
344 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
345 assert_eq!(new_state.round, 1);
346 }
347 _ => panic!("Expected PersistState"),
348 }
349 }
350
351 #[test]
352 fn revise_value_increments_round() {
353 let mode = MultiRoundMode;
354 let mut contributions = BTreeMap::new();
355 contributions.insert("alice".to_string(), "option_a".to_string());
356 let state = MultiRoundState {
357 round: 1,
358 participants: vec!["alice".into(), "bob".into()],
359 contributions,
360 convergence_type: "all_equal".into(),
361 converged: false,
362 };
363 let session = session_with_state(&state);
364 let env = contribute_env("alice", "option_b");
365
366 let result = mode.on_message(&session, &env).unwrap();
367 match result {
368 ModeResponse::PersistState(data) => {
369 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
370 assert_eq!(new_state.round, 2);
371 assert_eq!(new_state.contributions.get("alice").unwrap(), "option_b");
372 }
373 _ => panic!("Expected PersistState"),
374 }
375 }
376
377 #[test]
378 fn convergence_sets_converged_flag() {
379 let mode = MultiRoundMode;
380 let mut contributions = BTreeMap::new();
381 contributions.insert("alice".to_string(), "option_a".to_string());
382 let state = MultiRoundState {
383 round: 1,
384 participants: vec!["alice".into(), "bob".into()],
385 contributions,
386 convergence_type: "all_equal".into(),
387 converged: false,
388 };
389 let session = session_with_state(&state);
390 let env = contribute_env("bob", "option_a");
391
392 let result = mode.on_message(&session, &env).unwrap();
393 match result {
394 ModeResponse::PersistState(data) => {
395 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
396 assert_eq!(new_state.round, 2);
397 assert!(new_state.converged);
398 }
399 _ => panic!("Expected PersistState (convergence tracked, not auto-resolved)"),
400 }
401 }
402
403 #[test]
404 fn commitment_after_convergence_resolves() {
405 let mode = MultiRoundMode;
406 let mut contributions = BTreeMap::new();
407 contributions.insert("alice".to_string(), "option_a".to_string());
408 contributions.insert("bob".to_string(), "option_a".to_string());
409 let state = MultiRoundState {
410 round: 2,
411 participants: vec!["alice".into(), "bob".into()],
412 contributions,
413 convergence_type: "all_equal".into(),
414 converged: true,
415 };
416 let session = session_with_state(&state);
417 let env = commitment_env("coordinator");
418
419 let result = mode.on_message(&session, &env).unwrap();
420 match result {
421 ModeResponse::PersistAndResolve { resolution, .. } => {
422 let res: serde_json::Value = serde_json::from_slice(&resolution).unwrap();
423 assert_eq!(res["converged_value"], "option_a");
424 assert_eq!(res["round"], 2);
425 }
426 _ => panic!("Expected PersistAndResolve"),
427 }
428 }
429
430 #[test]
431 fn commitment_before_convergence_rejected() {
432 let mode = MultiRoundMode;
433 let state = MultiRoundState {
434 round: 0,
435 participants: vec!["alice".into(), "bob".into()],
436 contributions: BTreeMap::new(),
437 convergence_type: "all_equal".into(),
438 converged: false,
439 };
440 let session = session_with_state(&state);
441 let env = commitment_env("coordinator");
442
443 let err = mode.on_message(&session, &env).unwrap_err();
444 assert_eq!(err.to_string(), "InvalidPayload");
445 }
446
447 #[test]
448 fn contribute_after_convergence_rejected() {
449 let mode = MultiRoundMode;
450 let mut contributions = BTreeMap::new();
451 contributions.insert("alice".to_string(), "option_a".to_string());
452 contributions.insert("bob".to_string(), "option_a".to_string());
453 let state = MultiRoundState {
454 round: 2,
455 participants: vec!["alice".into(), "bob".into()],
456 contributions,
457 convergence_type: "all_equal".into(),
458 converged: true,
459 };
460 let session = session_with_state(&state);
461 let env = contribute_env("alice", "option_b");
462
463 let err = mode.on_message(&session, &env).unwrap_err();
464 assert_eq!(err.to_string(), "InvalidPayload");
465 }
466
467 #[test]
468 fn non_initiator_commitment_rejected() {
469 let mode = MultiRoundMode;
470 let mut contributions = BTreeMap::new();
471 contributions.insert("alice".to_string(), "option_a".to_string());
472 contributions.insert("bob".to_string(), "option_a".to_string());
473 let state = MultiRoundState {
474 round: 2,
475 participants: vec!["alice".into(), "bob".into()],
476 contributions,
477 convergence_type: "all_equal".into(),
478 converged: true,
479 };
480 let session = session_with_state(&state);
481 let env = commitment_env("alice"); let err = mode.authorize_sender(&session, &env).unwrap_err();
484 assert_eq!(err.to_string(), "Forbidden");
485 }
486
487 #[test]
488 fn no_convergence_when_values_differ() {
489 let mode = MultiRoundMode;
490 let mut contributions = BTreeMap::new();
491 contributions.insert("alice".to_string(), "option_a".to_string());
492 let state = MultiRoundState {
493 round: 1,
494 participants: vec!["alice".into(), "bob".into()],
495 contributions,
496 convergence_type: "all_equal".into(),
497 converged: false,
498 };
499 let session = session_with_state(&state);
500 let env = contribute_env("bob", "option_b");
501
502 let result = mode.on_message(&session, &env).unwrap();
503 match result {
504 ModeResponse::PersistState(data) => {
505 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
506 assert!(!new_state.converged);
507 }
508 _ => panic!("Expected PersistState"),
509 }
510 }
511
512 #[test]
513 fn no_convergence_when_not_all_contributed() {
514 let mode = MultiRoundMode;
515 let state = MultiRoundState {
516 round: 0,
517 participants: vec!["alice".into(), "bob".into(), "carol".into()],
518 contributions: BTreeMap::new(),
519 convergence_type: "all_equal".into(),
520 converged: false,
521 };
522 let session = session_with_state(&state);
523 let env = contribute_env("alice", "option_a");
524
525 let result = mode.on_message(&session, &env).unwrap();
526 assert!(matches!(result, ModeResponse::PersistState(_)));
527 }
528
529 #[test]
530 fn non_contribute_message_rejected() {
531 let mode = MultiRoundMode;
532 let state = MultiRoundState {
533 round: 0,
534 participants: vec!["alice".into()],
535 contributions: BTreeMap::new(),
536 convergence_type: "all_equal".into(),
537 converged: false,
538 };
539 let session = session_with_state(&state);
540 let env = Envelope {
541 macp_version: "1.0".into(),
542 mode: "ext.multi_round.v1".into(),
543 message_type: "Message".into(),
544 message_id: "m1".into(),
545 session_id: "s1".into(),
546 sender: "alice".into(),
547 timestamp_unix_ms: 1_700_000_000_000,
548 payload: b"hello".to_vec(),
549 };
550
551 let err = mode.on_message(&session, &env).unwrap_err();
552 assert_eq!(err.error_code(), "INVALID_ENVELOPE");
553 }
554
555 #[test]
556 fn contribute_invalid_payload_returns_error() {
557 let mode = MultiRoundMode;
558 let state = MultiRoundState {
559 round: 0,
560 participants: vec!["alice".into()],
561 contributions: BTreeMap::new(),
562 convergence_type: "all_equal".into(),
563 converged: false,
564 };
565 let session = session_with_state(&state);
566 let env = Envelope {
567 macp_version: "1.0".into(),
568 mode: "ext.multi_round.v1".into(),
569 message_type: "Contribute".into(),
570 message_id: "m1".into(),
571 session_id: "s1".into(),
572 sender: "alice".into(),
573 timestamp_unix_ms: 1_700_000_000_000,
574 payload: b"not json".to_vec(),
575 };
576
577 let err = mode.on_message(&session, &env).unwrap_err();
578 assert_eq!(err.to_string(), "InvalidPayload");
579 }
580
581 #[test]
582 fn encode_decode_round_trip() {
583 let mut contributions = BTreeMap::new();
584 contributions.insert("alice".into(), "value_a".into());
585 let original = MultiRoundState {
586 round: 5,
587 participants: vec!["alice".into(), "bob".into()],
588 contributions,
589 convergence_type: "all_equal".into(),
590 converged: true,
591 };
592
593 let encoded = MultiRoundMode::encode_state(&original);
594 let decoded = MultiRoundMode::decode_state(&encoded).unwrap();
595
596 assert_eq!(decoded.round, original.round);
597 assert_eq!(decoded.participants, original.participants);
598 assert_eq!(decoded.contributions, original.contributions);
599 assert_eq!(decoded.converged, original.converged);
600 }
601
602 #[test]
603 fn decode_invalid_state_returns_error() {
604 let err = MultiRoundMode::decode_state(b"garbage").unwrap_err();
605 assert_eq!(err.to_string(), "InvalidModeState");
606 }
607
608 #[test]
609 fn three_participant_convergence() {
610 let mode = MultiRoundMode;
611
612 let mut contributions = BTreeMap::new();
613 contributions.insert("alice".to_string(), "option_a".to_string());
614 contributions.insert("bob".to_string(), "option_a".to_string());
615 let state = MultiRoundState {
616 round: 2,
617 participants: vec!["alice".into(), "bob".into(), "carol".into()],
618 contributions,
619 convergence_type: "all_equal".into(),
620 converged: false,
621 };
622 let session = session_with_state(&state);
623 let env = contribute_env("carol", "option_a");
624
625 let result = mode.on_message(&session, &env).unwrap();
626 match result {
627 ModeResponse::PersistState(data) => {
628 let new_state: MultiRoundState = serde_json::from_slice(&data).unwrap();
629 assert!(new_state.converged);
630 }
631 _ => panic!("Expected PersistState with converged=true"),
632 }
633 }
634
635 #[test]
636 fn unknown_message_type_rejected() {
637 let mode = MultiRoundMode;
638 let state = MultiRoundState {
639 round: 0,
640 participants: vec!["alice".into(), "bob".into()],
641 contributions: BTreeMap::new(),
642 convergence_type: "all_equal".into(),
643 converged: false,
644 };
645 let session = session_with_state(&state);
646 let env = Envelope {
647 macp_version: "1.0".into(),
648 mode: "ext.multi_round.v1".into(),
649 message_type: "UnknownType".into(),
650 message_id: "msg-unknown".into(),
651 session_id: "s1".into(),
652 sender: "alice".into(),
653 timestamp_unix_ms: 0,
654 payload: vec![],
655 };
656 let err = mode.on_message(&session, &env).unwrap_err();
657 assert_eq!(err.error_code(), "INVALID_ENVELOPE");
658 }
659}