Skip to main content

macp_modes/mode/
multi_round.rs

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/// Internal state tracked across rounds.
10#[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/// Payload for Contribute messages.
22#[derive(Debug, Clone, Deserialize)]
23struct ContributePayload {
24    value: String,
25}
26
27/// Resolution payload emitted on convergence.
28#[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            // Only the initiator can emit Commitment
96            if env.sender != session.initiator_sender {
97                return Err(MacpError::Forbidden);
98            }
99            return Ok(());
100        }
101        // Default: must be a declared participant
102        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"); // not the initiator
482
483        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}