Skip to main content

ruma_lean/
lib.rs

1// Copyright 2026 Shane Jaroch
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15#![no_std]
16
17extern crate alloc;
18
19use alloc::collections::{BTreeMap, BinaryHeap};
20use alloc::string::String;
21use alloc::vec::Vec;
22use core::cmp::Ordering;
23use serde::{Deserialize, Serialize};
24
25#[cfg(feature = "std")]
26extern crate std;
27
28#[cfg(feature = "std")]
29pub use std::collections::HashMap;
30
31#[cfg(not(feature = "std"))]
32pub use hashbrown::HashMap;
33
34/// The version of the Matrix State Resolution algorithm to use.
35#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
36pub enum StateResVersion {
37    V1,
38    V2,
39    V2_1,
40}
41
42/// A lightweight Matrix Event representation for Lean-equivalent resolution.
43#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
44pub struct LeanEvent {
45    pub event_id: String,
46    pub power_level: i64,
47    pub origin_server_ts: u64,
48    pub prev_events: Vec<String>,
49    pub depth: u64, // Required for V1
50}
51
52/// The core tie-breaking logic from Ruma Lean (StateRes.lean).
53/// Matches Lean model: power_level (desc) -> origin_server_ts (asc) -> event_id (asc)
54impl Ord for LeanEvent {
55    fn cmp(&self, other: &Self) -> Ordering {
56        // Higher power level comes FIRST (is "smaller" in terms of order)
57        match other.power_level.cmp(&self.power_level) {
58            Ordering::Equal => {
59                // Earlier timestamp comes FIRST
60                match self.origin_server_ts.cmp(&other.origin_server_ts) {
61                    Ordering::Equal => {
62                        // Lexicographically smaller ID comes FIRST
63                        self.event_id.cmp(&other.event_id)
64                    }
65                    ord => ord,
66                }
67            }
68            ord => ord,
69        }
70    }
71}
72
73impl PartialOrd for LeanEvent {
74    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
75        Some(self.cmp(other))
76    }
77}
78
79/// A wrapper to ensure BinaryHeap pops the "smallest" (best) event first.
80#[derive(Debug, Clone, Copy, Eq, PartialEq)]
81struct SortPriority<'a> {
82    event: &'a LeanEvent,
83    version: StateResVersion,
84}
85
86impl<'a> Ord for SortPriority<'a> {
87    fn cmp(&self, other: &Self) -> Ordering {
88        match self.version {
89            StateResVersion::V1 => {
90                // V1 tie-breaking: depth (asc) -> event_id (asc)
91                // Inverted for Max-Heap
92                match other.event.depth.cmp(&self.event.depth) {
93                    Ordering::Equal => other.event.event_id.cmp(&self.event.event_id),
94                    ord => ord,
95                }
96            }
97            StateResVersion::V2 | StateResVersion::V2_1 => {
98                // V2 tie-breaking: power_level (desc) -> origin_server_ts (asc) -> event_id (asc)
99                // Priority popping (best first)
100                match self.event.power_level.cmp(&other.event.power_level) {
101                    Ordering::Equal => {
102                        match other
103                            .event
104                            .origin_server_ts
105                            .cmp(&self.event.origin_server_ts)
106                        {
107                            Ordering::Equal => other.event.event_id.cmp(&self.event.event_id),
108                            ord => ord,
109                        }
110                    }
111                    ord => ord,
112                }
113            }
114        }
115    }
116}
117
118impl<'a> PartialOrd for SortPriority<'a> {
119    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
120        Some(self.cmp(other))
121    }
122}
123
124/// A simplified implementation of Kahn's Topological Sort.
125pub fn lean_kahn_sort(
126    events: &HashMap<String, LeanEvent>,
127    version: StateResVersion,
128) -> Vec<String> {
129    let mut in_degree: HashMap<String, usize> = HashMap::new();
130    let mut adjacency: HashMap<String, Vec<String>> = HashMap::new();
131
132    for (id, event) in events {
133        in_degree.entry(id.clone()).or_insert(0);
134        for prev in &event.prev_events {
135            if events.contains_key(prev) {
136                adjacency.entry(prev.clone()).or_default().push(id.clone());
137                *in_degree.entry(id.clone()).or_insert(0) += 1;
138            }
139        }
140    }
141
142    let mut queue: BinaryHeap<SortPriority> = BinaryHeap::new();
143    for (id, &degree) in &in_degree {
144        if degree == 0 {
145            if let Some(event) = events.get(id) {
146                queue.push(SortPriority { event, version });
147            }
148        }
149    }
150
151    let mut result = Vec::new();
152    while let Some(priority) = queue.pop() {
153        let event = priority.event;
154        result.push(event.event_id.clone());
155        if let Some(neighbors) = adjacency.get(&event.event_id) {
156            for next_id in neighbors {
157                let degree = in_degree.get_mut(next_id).unwrap();
158                *degree -= 1;
159                if *degree == 0 {
160                    queue.push(SortPriority {
161                        event: events.get(next_id).unwrap(),
162                        version,
163                    });
164                }
165            }
166        }
167    }
168    result
169}
170
171pub fn resolve_lean(
172    unconflicted_state: BTreeMap<(String, String), String>,
173    conflicted_events: HashMap<String, LeanEvent>,
174    version: StateResVersion,
175) -> BTreeMap<(String, String), String> {
176    let resolved = unconflicted_state;
177    let _sorted_ids = lean_kahn_sort(&conflicted_events, version);
178    resolved
179}
180
181#[cfg(feature = "zkvm")]
182pub fn verify_signature(_public_key: &[u8; 32], _message: &[u8], _signature: &[u8; 64]) {
183    // Verifiable signature check for ZKVM environment
184}
185
186#[cfg(all(feature = "std", not(feature = "zkvm")))]
187pub fn verify_signature(public_key: &[u8; 32], message: &[u8], signature: &[u8; 64]) {
188    use ed25519_consensus::{Signature, VerificationKey};
189    let vk = VerificationKey::try_from(*public_key).expect("Invalid public key");
190    let sig = Signature::from(*signature);
191    vk.verify(&sig, message)
192        .expect("Signature verification failed");
193}
194
195#[cfg(all(not(feature = "std"), not(feature = "zkvm")))]
196pub fn verify_signature(_public_key: &[u8; 32], _message: &[u8], _signature: &[u8; 64]) {
197    // No-op for other configurations
198}
199
200#[cfg(test)]
201mod tests {
202    use super::*;
203    use alloc::string::ToString;
204    use alloc::vec;
205
206    #[cfg(not(feature = "std"))]
207    use hashbrown::HashMap;
208    #[cfg(feature = "std")]
209    use std::collections::HashMap;
210
211    #[test]
212    fn test_v1_resolution_happy_path() {
213        let mut events = HashMap::new();
214        events.insert(
215            "A".into(),
216            LeanEvent {
217                event_id: "A".into(),
218                power_level: 0,
219                origin_server_ts: 100,
220                prev_events: vec![],
221                depth: 1,
222            },
223        );
224        events.insert(
225            "B".into(),
226            LeanEvent {
227                event_id: "B".into(),
228                power_level: 0,
229                origin_server_ts: 50,
230                prev_events: vec![],
231                depth: 2,
232            },
233        );
234        let sorted = lean_kahn_sort(&events, StateResVersion::V1);
235        assert_eq!(sorted, vec!["A", "B"]);
236    }
237
238    #[test]
239    fn test_v1_tie_break_by_id() {
240        let mut events = HashMap::new();
241        events.insert(
242            "B".into(),
243            LeanEvent {
244                event_id: "B".into(),
245                power_level: 0,
246                origin_server_ts: 100,
247                prev_events: vec![],
248                depth: 1,
249            },
250        );
251        events.insert(
252            "A".into(),
253            LeanEvent {
254                event_id: "A".into(),
255                power_level: 0,
256                origin_server_ts: 100,
257                prev_events: vec![],
258                depth: 1,
259            },
260        );
261        let sorted = lean_kahn_sort(&events, StateResVersion::V1);
262        assert_eq!(sorted, vec!["A", "B"]);
263    }
264
265    #[test]
266    fn test_v2_resolution_happy_path() {
267        let mut events = HashMap::new();
268        events.insert(
269            "A".into(),
270            LeanEvent {
271                event_id: "A".into(),
272                power_level: 100,
273                origin_server_ts: 100,
274                prev_events: vec![],
275                depth: 10,
276            },
277        );
278        events.insert(
279            "B".into(),
280            LeanEvent {
281                event_id: "B".into(),
282                power_level: 50,
283                origin_server_ts: 10,
284                prev_events: vec![],
285                depth: 1,
286            },
287        );
288        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
289        assert_eq!(sorted, vec!["A", "B"]);
290    }
291
292    #[test]
293    fn test_v2_deep_tie_break() {
294        let mut events = HashMap::new();
295        events.insert(
296            "B".into(),
297            LeanEvent {
298                event_id: "B".into(),
299                power_level: 100,
300                origin_server_ts: 10,
301                prev_events: vec![],
302                depth: 1,
303            },
304        );
305        events.insert(
306            "A".into(),
307            LeanEvent {
308                event_id: "A".into(),
309                power_level: 100,
310                origin_server_ts: 10,
311                prev_events: vec![],
312                depth: 1,
313            },
314        );
315        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
316        assert_eq!(sorted, vec!["A", "B"]);
317    }
318
319    #[test]
320    fn test_v1_v2_v2_1_comparison_determinism() {
321        let mut events = HashMap::new();
322        events.insert(
323            "A".into(),
324            LeanEvent {
325                event_id: "A".into(),
326                power_level: 10,
327                origin_server_ts: 10,
328                prev_events: vec![],
329                depth: 1,
330            },
331        );
332        events.insert(
333            "B".into(),
334            LeanEvent {
335                event_id: "B".into(),
336                power_level: 100,
337                origin_server_ts: 100,
338                prev_events: vec![],
339                depth: 10,
340            },
341        );
342        let sorted_v1 = lean_kahn_sort(&events, StateResVersion::V1);
343        let sorted_v2 = lean_kahn_sort(&events, StateResVersion::V2);
344        let sorted_v2_1 = lean_kahn_sort(&events, StateResVersion::V2_1);
345        assert_eq!(sorted_v1, vec!["A", "B"]);
346        assert_eq!(sorted_v2, vec!["B", "A"]);
347        assert_eq!(sorted_v2_1, vec!["B", "A"]);
348    }
349
350    #[test]
351    fn test_unhappy_path_cycle_detection() {
352        let mut events = HashMap::new();
353        events.insert(
354            "A".into(),
355            LeanEvent {
356                event_id: "A".into(),
357                power_level: 100,
358                origin_server_ts: 100,
359                prev_events: vec!["B".into()],
360                depth: 1,
361            },
362        );
363        events.insert(
364            "B".into(),
365            LeanEvent {
366                event_id: "B".into(),
367                power_level: 100,
368                origin_server_ts: 100,
369                prev_events: vec!["A".into()],
370                depth: 1,
371            },
372        );
373        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
374        assert!(sorted.is_empty());
375    }
376
377    #[test]
378    fn test_signature_verification_failure() {
379        #[cfg(all(feature = "std", not(feature = "zkvm")))]
380        {
381            let pk = [
382                215, 90, 152, 1, 130, 177, 10, 183, 213, 75, 254, 211, 201, 100, 7, 58, 14, 225,
383                114, 243, 218, 166, 35, 37, 175, 2, 26, 104, 247, 7, 81, 26,
384            ];
385            let sig = [0u8; 64];
386            let msg = b"test";
387            let result = std::panic::catch_unwind(|| {
388                verify_signature(&pk, msg, &sig);
389            });
390            assert!(result.is_err());
391        }
392    }
393
394    #[test]
395    fn test_serialization_roundtrip() {
396        let event = LeanEvent {
397            event_id: "$abc".into(),
398            power_level: 100,
399            origin_server_ts: 12345,
400            prev_events: vec![],
401            depth: 5,
402        };
403        let serialized = serde_json::to_string(&event).unwrap();
404        let deserialized: LeanEvent = serde_json::from_str(&serialized).unwrap();
405        assert_eq!(event, deserialized);
406    }
407
408    #[test]
409    fn test_partial_ord_implementations() {
410        let e1 = LeanEvent {
411            event_id: "a".into(),
412            power_level: 100,
413            origin_server_ts: 10,
414            prev_events: vec![],
415            depth: 1,
416        };
417        let e2 = LeanEvent {
418            event_id: "b".into(),
419            power_level: 100,
420            origin_server_ts: 10,
421            prev_events: vec![],
422            depth: 1,
423        };
424        assert!(e1.partial_cmp(&e2).is_some());
425
426        let p1 = SortPriority {
427            event: &e1,
428            version: StateResVersion::V2,
429        };
430        let p2 = SortPriority {
431            event: &e2,
432            version: StateResVersion::V2,
433        };
434        assert!(p1.partial_cmp(&p2).is_some());
435    }
436
437    #[test]
438    fn test_trait_coverage() {
439        let v = StateResVersion::V2;
440        assert_eq!(v, StateResVersion::V2);
441        let _ = alloc::format!("{:?}", v);
442
443        let e = LeanEvent {
444            event_id: "a".into(),
445            power_level: 100,
446            origin_server_ts: 10,
447            prev_events: vec![],
448            depth: 1,
449        };
450        let _ = e.clone();
451        let _ = alloc::format!("{:?}", e);
452    }
453
454    #[test]
455    fn test_complex_dag_sort() {
456        let mut events = HashMap::new();
457        events.insert(
458            "1".into(),
459            LeanEvent {
460                event_id: "1".into(),
461                power_level: 100,
462                origin_server_ts: 10,
463                prev_events: vec![],
464                depth: 1,
465            },
466        );
467        events.insert(
468            "2".into(),
469            LeanEvent {
470                event_id: "2".into(),
471                power_level: 50,
472                origin_server_ts: 20,
473                prev_events: vec!["1".into()],
474                depth: 2,
475            },
476        );
477        events.insert(
478            "3".into(),
479            LeanEvent {
480                event_id: "3".into(),
481                power_level: 50,
482                origin_server_ts: 15,
483                prev_events: vec!["1".into()],
484                depth: 2,
485            },
486        );
487        events.insert(
488            "4".into(),
489            LeanEvent {
490                event_id: "4".into(),
491                power_level: 10,
492                origin_server_ts: 30,
493                prev_events: vec!["2".into(), "3".into()],
494                depth: 3,
495            },
496        );
497        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
498        assert_eq!(sorted, vec!["1", "3", "2", "4"]);
499    }
500
501    #[test]
502    fn test_kahn_missing_parents() {
503        let mut events = HashMap::new();
504        events.insert(
505            "A".into(),
506            LeanEvent {
507                event_id: "A".into(),
508                power_level: 100,
509                origin_server_ts: 10,
510                prev_events: vec!["MISSING".into()],
511                depth: 1,
512            },
513        );
514        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
515        assert_eq!(sorted, vec!["A"]);
516    }
517
518    #[test]
519    fn test_resolve_lean_functionality() {
520        let mut unconflicted = BTreeMap::new();
521        unconflicted.insert(("type".into(), "key".into()), "id".into());
522        let conflicted = HashMap::new();
523        let resolved = resolve_lean(unconflicted.clone(), conflicted, StateResVersion::V2);
524        assert_eq!(resolved, unconflicted);
525    }
526
527    fn run_batch_test(
528        version: StateResVersion,
529        rows: &[(&str, i64, u64, u64, &[&str])],
530        expected: &[&str],
531    ) {
532        let mut events = HashMap::new();
533        for r in rows {
534            events.insert(
535                r.0.to_string(),
536                LeanEvent {
537                    event_id: r.0.to_string(),
538                    power_level: r.1,
539                    origin_server_ts: r.2,
540                    depth: r.3,
541                    prev_events: r.4.iter().map(|s| s.to_string()).collect(),
542                },
543            );
544        }
545        let result = lean_kahn_sort(&events, version);
546        assert_eq!(
547            result,
548            expected.iter().map(|s| s.to_string()).collect::<Vec<_>>()
549        );
550    }
551
552    #[test]
553    fn test_resolution_batch() {
554        run_batch_test(
555            StateResVersion::V2,
556            &[("Alice", 100, 500, 1, &[]), ("Bob", 50, 100, 1, &[])],
557            &["Alice", "Bob"],
558        );
559        run_batch_test(
560            StateResVersion::V1,
561            &[("Deep", 100, 100, 10, &[]), ("Shallow", 10, 100, 1, &[])],
562            &["Shallow", "Deep"],
563        );
564    }
565
566    #[test]
567    fn test_native_resolution_bootstrap_parity() {
568        let mut events = HashMap::new();
569        events.insert(
570            "1".into(),
571            LeanEvent {
572                event_id: "1".into(),
573                power_level: 100,
574                origin_server_ts: 10,
575                prev_events: vec![],
576                depth: 1,
577            },
578        );
579        events.insert(
580            "2".into(),
581            LeanEvent {
582                event_id: "2".into(),
583                power_level: 0,
584                origin_server_ts: 20,
585                prev_events: vec!["1".into()],
586                depth: 2,
587            },
588        );
589        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
590        let mut resolved_state = BTreeMap::new();
591        for id in sorted {
592            let ev = events.get(&id).unwrap();
593            let key = ("m.room.member".to_string(), "@user:example.com".to_string());
594            resolved_state.insert(key, ev.event_id.clone());
595        }
596        assert_eq!(
597            resolved_state.get(&("m.room.member".to_string(), "@user:example.com".to_string())),
598            Some(&"2".to_string())
599        );
600    }
601
602    #[test]
603    fn test_enum_coverage() {
604        let v = StateResVersion::V2;
605        let v2 = v;
606        assert_eq!(v, v2);
607        let debug_str = alloc::format!("{:?}", v);
608        assert!(debug_str.contains("V2"));
609    }
610
611    #[test]
612    fn test_event_traits_coverage() {
613        let e = LeanEvent {
614            event_id: "a".into(),
615            power_level: 100,
616            origin_server_ts: 10,
617            prev_events: vec![],
618            depth: 1,
619        };
620        let e2 = e.clone();
621        assert_eq!(e, e2);
622        let debug_str = alloc::format!("{:?}", e);
623        assert!(debug_str.contains("event_id"));
624    }
625
626    #[test]
627    fn test_sort_priority_traits() {
628        let e = LeanEvent {
629            event_id: "a".into(),
630            power_level: 100,
631            origin_server_ts: 10,
632            prev_events: vec![],
633            depth: 1,
634        };
635        let p = SortPriority {
636            event: &e,
637            version: StateResVersion::V2,
638        };
639        let p2 = p;
640        assert_eq!(p, p2);
641        let debug_str = alloc::format!("{:?}", p);
642        assert!(debug_str.contains("version"));
643    }
644
645    #[test]
646    fn test_v1_equal_depth_tie_break() {
647        let mut events = HashMap::new();
648        events.insert(
649            "B".into(),
650            LeanEvent {
651                event_id: "B".into(),
652                power_level: 0,
653                origin_server_ts: 10,
654                prev_events: vec![],
655                depth: 1,
656            },
657        );
658        events.insert(
659            "A".into(),
660            LeanEvent {
661                event_id: "A".into(),
662                power_level: 0,
663                origin_server_ts: 10,
664                prev_events: vec![],
665                depth: 1,
666            },
667        );
668        let sorted = lean_kahn_sort(&events, StateResVersion::V1);
669        assert_eq!(sorted, vec!["A", "B"]);
670    }
671
672    #[test]
673    fn test_kahn_no_neighbors() {
674        let mut events = HashMap::new();
675        events.insert(
676            "1".into(),
677            LeanEvent {
678                event_id: "1".into(),
679                power_level: 100,
680                origin_server_ts: 10,
681                prev_events: vec![],
682                depth: 1,
683            },
684        );
685        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
686        assert_eq!(sorted, vec!["1"]);
687    }
688
689    #[test]
690    fn test_v2_1_full_coverage() {
691        let mut events = HashMap::new();
692        events.insert(
693            "A".into(),
694            LeanEvent {
695                event_id: "A".into(),
696                power_level: 100,
697                origin_server_ts: 10,
698                prev_events: vec![],
699                depth: 1,
700            },
701        );
702        let sorted = lean_kahn_sort(&events, StateResVersion::V2_1);
703        assert_eq!(sorted, vec!["A"]);
704    }
705
706    #[test]
707    fn test_total_order_properties() {
708        let e1 = LeanEvent {
709            event_id: "a".into(),
710            power_level: 100,
711            origin_server_ts: 10,
712            prev_events: vec![],
713            depth: 1,
714        };
715        let e2 = LeanEvent {
716            event_id: "b".into(),
717            power_level: 100,
718            origin_server_ts: 10,
719            prev_events: vec![],
720            depth: 1,
721        };
722        let e3 = LeanEvent {
723            event_id: "c".into(),
724            power_level: 50,
725            origin_server_ts: 10,
726            prev_events: vec![],
727            depth: 1,
728        };
729        assert_eq!(e1.cmp(&e1), Ordering::Equal);
730        assert!(e1 <= e1);
731        assert!(e1 <= e2 || e2 <= e1);
732        if e1 <= e2 && e2 <= e3 {
733            assert!(e1 <= e3);
734        }
735        let e1_copy = e1.clone();
736        if e1 <= e1_copy && e1_copy <= e1 {
737            assert_eq!(e1, e1_copy);
738        }
739    }
740
741    #[test]
742    fn test_coverage_booster_all_branches() {
743        let e_base = LeanEvent {
744            event_id: "m".into(),
745            power_level: 50,
746            origin_server_ts: 50,
747            prev_events: vec![],
748            depth: 50,
749        };
750        let p_base = SortPriority {
751            event: &e_base,
752            version: StateResVersion::V2,
753        };
754        let e_high_power = LeanEvent {
755            power_level: 100,
756            ..e_base.clone()
757        };
758        let p_high_power = SortPriority {
759            event: &e_high_power,
760            version: StateResVersion::V2,
761        };
762        assert_eq!(p_base.cmp(&p_high_power), Ordering::Less);
763        let e_early_ts = LeanEvent {
764            origin_server_ts: 10,
765            ..e_base.clone()
766        };
767        let p_early_ts = SortPriority {
768            event: &e_early_ts,
769            version: StateResVersion::V2,
770        };
771        assert_eq!(p_base.cmp(&p_early_ts), Ordering::Less);
772        let e_early_id = LeanEvent {
773            event_id: "a".into(),
774            ..e_base.clone()
775        };
776        let p_early_id = SortPriority {
777            event: &e_early_id,
778            version: StateResVersion::V2,
779        };
780        assert_eq!(p_base.cmp(&p_early_id), Ordering::Less);
781        let p_v1_base = SortPriority {
782            event: &e_base,
783            version: StateResVersion::V1,
784        };
785        let e_shallow = LeanEvent {
786            depth: 1,
787            ..e_base.clone()
788        };
789        let p_shallow = SortPriority {
790            event: &e_shallow,
791            version: StateResVersion::V1,
792        };
793        assert_eq!(p_v1_base.cmp(&p_shallow), Ordering::Less);
794        let p_v1_early_id = SortPriority {
795            event: &e_early_id,
796            version: StateResVersion::V1,
797        };
798        assert_eq!(p_v1_base.cmp(&p_v1_early_id), Ordering::Less);
799        assert_eq!(p_v1_base.cmp(&p_v1_base), Ordering::Equal);
800    }
801}