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
19#[cfg(not(feature = "zkvm"))]
20use alloc::collections::BTreeSet;
21use alloc::collections::{BTreeMap, BinaryHeap};
22
23use alloc::string::String;
24use alloc::vec::Vec;
25use core::cmp::Ordering;
26use serde::{Deserialize, Serialize};
27
28use serde_json::Value;
29
30#[cfg(feature = "std")]
31extern crate std;
32
33#[cfg(feature = "std")]
34pub use std::collections::HashMap;
35
36#[cfg(not(feature = "std"))]
37pub use hashbrown::HashMap;
38
39/// The version of the Matrix State Resolution algorithm to use.
40#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
41#[cfg_attr(feature = "cli", derive(clap::ValueEnum))]
42pub enum StateResVersion {
43    V1,
44    V2,
45    V2_1,
46}
47
48/// A lightweight Matrix Event representation for Lean-equivalent resolution.
49#[derive(Debug, Clone, Serialize, Deserialize, Default)]
50pub struct LeanEvent {
51    pub event_id: String,
52    #[serde(rename = "type")]
53    pub event_type: String,
54    #[serde(default)]
55    pub state_key: String,
56    #[serde(default)]
57    pub power_level: i64,
58    pub origin_server_ts: u64,
59    #[serde(default)]
60    pub sender: String,
61    #[serde(default)]
62    pub content: Value,
63    #[serde(default)]
64    pub prev_events: Vec<String>,
65    #[serde(default)]
66    pub auth_events: Vec<String>,
67    #[serde(default)]
68    pub depth: u64, // Required for V1
69}
70
71impl PartialEq for LeanEvent {
72    fn eq(&self, other: &Self) -> bool {
73        self.event_id == other.event_id
74    }
75}
76
77impl Eq for LeanEvent {}
78
79impl Ord for LeanEvent {
80    fn cmp(&self, other: &Self) -> Ordering {
81        self.event_id.cmp(&other.event_id)
82    }
83}
84
85impl PartialOrd for LeanEvent {
86    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
87        Some(self.cmp(other))
88    }
89}
90
91/// A wrapper to ensure BinaryHeap pops the "smallest" (best) event first.
92#[derive(Debug, Clone, Copy)]
93struct SortPriority<'a> {
94    event: &'a LeanEvent,
95    version: StateResVersion,
96}
97
98impl<'a> PartialEq for SortPriority<'a> {
99    fn eq(&self, other: &Self) -> bool {
100        self.cmp(other) == Ordering::Equal
101    }
102}
103
104impl<'a> Eq for SortPriority<'a> {}
105
106impl<'a> Ord for SortPriority<'a> {
107    fn cmp(&self, other: &Self) -> Ordering {
108        match self.version {
109            StateResVersion::V1 => {
110                // V1 tie-breaking: depth (asc) -> event_id (asc)
111                // Inverted for Max-Heap
112                match other.event.depth.cmp(&self.event.depth) {
113                    Ordering::Equal => other.event.event_id.cmp(&self.event.event_id),
114                    ord => ord,
115                }
116            }
117            StateResVersion::V2 | StateResVersion::V2_1 => {
118                // V2 tie-breaking: power_level (desc) -> origin_server_ts (asc) -> event_id (asc)
119                // To have "best" events come LAST in the sorted list, we must pop "worst" events FIRST.
120                // In Rust's Max-Heap BinaryHeap, "greater" elements are popped first.
121                // So "worst" must be "greater" than "best".
122
123                // Lower power level is WORSE (pops first, overwritten).
124                match other.event.power_level.cmp(&self.event.power_level) {
125                    Ordering::Equal => {
126                        // Earlier timestamp is WORSE (pops first, overwritten).
127                        match other
128                            .event
129                            .origin_server_ts
130                            .cmp(&self.event.origin_server_ts)
131                        {
132                            Ordering::Equal => {
133                                // Lexicographically SMALLER ID is WORSE (pops first, overwritten).
134                                other.event.event_id.cmp(&self.event.event_id)
135                            }
136                            ord => ord,
137                        }
138                    }
139                    ord => ord,
140                }
141            }
142        }
143    }
144}
145
146impl<'a> PartialOrd for SortPriority<'a> {
147    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
148        Some(self.cmp(other))
149    }
150}
151
152/// A simplified implementation of Kahn's Topological Sort.
153pub fn lean_kahn_sort(
154    events: &HashMap<String, LeanEvent>,
155    version: StateResVersion,
156) -> Vec<String> {
157    let mut in_degree: HashMap<String, usize> = HashMap::new();
158    let mut adjacency: HashMap<String, Vec<String>> = HashMap::new();
159
160    for (id, event) in events {
161        in_degree.entry(id.clone()).or_insert(0);
162        for auth in &event.auth_events {
163            if events.contains_key(auth) {
164                adjacency.entry(auth.clone()).or_default().push(id.clone());
165                *in_degree.entry(id.clone()).or_insert(0) += 1;
166            }
167        }
168    }
169
170    let mut queue: BinaryHeap<SortPriority> = BinaryHeap::new();
171    for (id, &degree) in &in_degree {
172        if degree == 0 {
173            if let Some(event) = events.get(id) {
174                queue.push(SortPriority { event, version });
175            }
176        }
177    }
178
179    let mut result = Vec::new();
180    while let Some(priority) = queue.pop() {
181        let event = priority.event;
182        result.push(event.event_id.clone());
183        if let Some(neighbors) = adjacency.get(&event.event_id) {
184            for next_id in neighbors {
185                let degree = in_degree.get_mut(next_id).unwrap();
186                *degree -= 1;
187                if *degree == 0 {
188                    queue.push(SortPriority {
189                        event: events.get(next_id).unwrap(),
190                        version,
191                    });
192                }
193            }
194        }
195    }
196
197    // Failsafe: If the result length doesn't match the input length, there is a cycle in the DAG.
198    // Matrix spec mandates failure.
199    if result.len() != events.len() {
200        return Vec::new();
201    }
202
203    result
204}
205
206pub fn resolve_lean(
207    unconflicted_state: BTreeMap<(String, String), String>,
208    conflicted_events: HashMap<String, LeanEvent>,
209    version: StateResVersion,
210) -> BTreeMap<(String, String), String> {
211    // MSC4297 (v2.1): The algorithm starts from an empty set of state.
212    // Unconflicted state events are added to the conflicted events set and sorted together.
213    let (mut resolved, sort_set) = match version {
214        StateResVersion::V2_1 => {
215            // MSC4297: We assume all necessary event objects (conflicted and unconflicted)
216            // are provided in conflicted_events for the sorting process.
217            (BTreeMap::new(), conflicted_events.clone())
218        }
219        _ => (unconflicted_state, conflicted_events),
220    };
221
222    let sorted_ids = lean_kahn_sort(&sort_set, version);
223
224    for id in sorted_ids {
225        if let Some(event) = sort_set.get(&id) {
226            resolved.insert(
227                (event.event_type.clone(), event.state_key.clone()),
228                event.event_id.clone(),
229            );
230        }
231    }
232
233    resolved
234}
235
236#[cfg(not(feature = "zkvm"))] // ONLY run this on the Host!
237pub fn compute_v2_1_conflicted_subgraph(
238    auth_graph: &HashMap<String, LeanEvent>,
239    conflicted_set: &[String],
240) -> HashMap<String, LeanEvent> {
241    let mut backwards_reachable = BTreeSet::new();
242    let mut forwards_reachable = BTreeSet::new();
243
244    // 1. Calculate Backwards Reachable (Ancestors up the auth chain)
245    let mut b_stack: Vec<String> = conflicted_set.to_vec();
246    while let Some(node) = b_stack.pop() {
247        if backwards_reachable.insert(node.clone()) {
248            if let Some(event) = auth_graph.get(&node) {
249                b_stack.extend(event.auth_events.clone());
250            }
251        }
252    }
253
254    // 2. Build Reverse Adjacency for Forwards Search
255    let mut children_map: HashMap<String, Vec<String>> = HashMap::new();
256    for (id, event) in auth_graph {
257        for prev in &event.auth_events {
258            children_map
259                .entry(prev.clone())
260                .or_default()
261                .push(id.clone());
262        }
263    }
264
265    // 3. Calculate Forwards Reachable (Descendants down the auth chain)
266    let mut f_stack: Vec<String> = conflicted_set.to_vec();
267    while let Some(node) = f_stack.pop() {
268        if forwards_reachable.insert(node.clone()) {
269            if let Some(children) = children_map.get(&node) {
270                f_stack.extend(children.clone());
271            }
272        }
273    }
274
275    // 4. Intersect and build the final Conflicted Subgraph
276    let mut subgraph = HashMap::new();
277    let backwards_ids: BTreeSet<String> = backwards_reachable.iter().cloned().collect();
278    let forwards_ids: BTreeSet<String> = forwards_reachable.iter().cloned().collect();
279
280    for id in backwards_ids.intersection(&forwards_ids) {
281        if let Some(event) = auth_graph.get(id) {
282            subgraph.insert(id.clone(), event.clone());
283        }
284    }
285    subgraph
286}
287
288#[cfg(feature = "zkvm")]
289pub fn verify_signature(_public_key: &[u8; 32], _message: &[u8], _signature: &[u8; 64]) {
290    // Verifiable signature check for ZKVM environment
291}
292
293#[cfg(all(feature = "std", not(feature = "zkvm")))]
294pub fn verify_signature(public_key: &[u8; 32], message: &[u8], signature: &[u8; 64]) {
295    use ed25519_consensus::{Signature, VerificationKey};
296    let vk = VerificationKey::try_from(*public_key).expect("Invalid public key");
297    let sig = Signature::from(*signature);
298    vk.verify(&sig, message)
299        .expect("Signature verification failed");
300}
301
302#[cfg(all(not(feature = "std"), not(feature = "zkvm")))]
303pub fn verify_signature(_public_key: &[u8; 32], _message: &[u8], _signature: &[u8; 64]) {
304    // No-op for other configurations
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310    use alloc::string::ToString;
311    use alloc::vec;
312
313    #[cfg(not(feature = "std"))]
314    use hashbrown::HashMap;
315    #[cfg(feature = "std")]
316    use std::collections::HashMap;
317
318    #[test]
319    fn test_leanevent_deserialization_defaults() {
320        let json = r#"{
321            "event_id": "$test",
322            "type": "m.room.message",
323            "origin_server_ts": 12345
324        }"#;
325        let ev: LeanEvent = serde_json::from_str(json).unwrap();
326        assert_eq!(ev.event_id, "$test");
327        assert_eq!(ev.event_type, "m.room.message");
328        assert_eq!(ev.origin_server_ts, 12345);
329        assert_eq!(ev.state_key, "");
330        assert_eq!(ev.power_level, 0);
331        assert_eq!(ev.sender, "");
332        assert_eq!(ev.prev_events.len(), 0);
333        assert_eq!(ev.auth_events.len(), 0);
334        assert_eq!(ev.depth, 0);
335    }
336
337    #[test]
338    fn test_sort_priority_v2_tie_break() {
339        let e_base = LeanEvent {
340            event_id: "$1".into(),
341            power_level: 100,
342            origin_server_ts: 10,
343            ..Default::default()
344        };
345        let e_worst_pl = LeanEvent {
346            event_id: "$2".into(),
347            power_level: 50,
348            origin_server_ts: 10,
349            ..Default::default()
350        };
351        let p_base = SortPriority {
352            event: &e_base,
353            version: StateResVersion::V2,
354        };
355        let p_worst_pl = SortPriority {
356            event: &e_worst_pl,
357            version: StateResVersion::V2,
358        };
359
360        // Earlier ts / smaller id events should be GREATER so they pop FIRST from Max-Heap.
361        assert_eq!(p_base.cmp(&p_worst_pl), Ordering::Less); // p_worst_pl has power 50, p_base 100. Lower pl gets popped first, so it is Greater. p_worst_pl > p_base => p_base < p_worst_pl => Less
362
363        let e_earlier_ts = LeanEvent {
364            event_id: "$3".into(),
365            power_level: 100,
366            origin_server_ts: 5,
367            ..Default::default()
368        };
369        let p_earlier_ts = SortPriority {
370            event: &e_earlier_ts,
371            version: StateResVersion::V2,
372        };
373        // p_earlier_ts has ts 5, p_base has ts 10. Earlier TS gets popped first, so it is Greater. p_earlier_ts > p_base => p_base < p_earlier_ts => Less
374        assert_eq!(p_base.cmp(&p_earlier_ts), Ordering::Less);
375
376        let e_smaller_id = LeanEvent {
377            event_id: "$0".into(),
378            power_level: 100,
379            origin_server_ts: 10,
380            ..Default::default()
381        };
382        let p_smaller_id = SortPriority {
383            event: &e_smaller_id,
384            version: StateResVersion::V2,
385        };
386        // p_smaller_id has id "$0", p_base has id "$1". Smaller ID gets popped first, so it is Greater. p_smaller_id > p_base => p_base < p_smaller_id => Less
387        assert_eq!(p_base.cmp(&p_smaller_id), Ordering::Less);
388    }
389
390    #[test]
391    fn test_v1_resolution_happy_path() {
392        let mut events = HashMap::new();
393        events.insert(
394            "A".into(),
395            LeanEvent {
396                event_id: "A".into(),
397                event_type: "m.room.member".into(),
398                state_key: "@alice:example.com".into(),
399                power_level: 0,
400                origin_server_ts: 100,
401                prev_events: vec![],
402                auth_events: vec![],
403                depth: 1,
404                ..Default::default()
405            },
406        );
407        events.insert(
408            "B".into(),
409            LeanEvent {
410                event_id: "B".into(),
411                event_type: "m.room.member".into(),
412                state_key: "@alice:example.com".into(),
413                power_level: 0,
414                origin_server_ts: 50,
415                prev_events: vec![],
416                auth_events: vec!["A".into()],
417                depth: 2,
418                ..Default::default()
419            },
420        );
421        let sorted = lean_kahn_sort(&events, StateResVersion::V1);
422        assert_eq!(sorted, vec!["A", "B"]);
423    }
424
425    #[test]
426    fn test_v2_1_strict_resolution() {
427        let mut unconflicted = BTreeMap::new();
428        unconflicted.insert(
429            ("m.room.member".into(), "@alice:example.com".into()),
430            "A".into(),
431        );
432
433        let mut conflicted = HashMap::new();
434        conflicted.insert(
435            "A".into(),
436            LeanEvent {
437                event_id: "A".into(),
438                event_type: "m.room.member".into(),
439                state_key: "@alice:example.com".into(),
440                power_level: 50,
441                origin_server_ts: 100,
442                prev_events: vec![],
443                auth_events: vec![],
444                depth: 1,
445                ..Default::default()
446            },
447        );
448        conflicted.insert(
449            "B".into(),
450            LeanEvent {
451                event_id: "B".into(),
452                event_type: "m.room.member".into(),
453                state_key: "@alice:example.com".into(),
454                power_level: 100,
455                origin_server_ts: 50,
456                prev_events: vec![],
457                auth_events: vec![],
458                depth: 1,
459                ..Default::default()
460            },
461        );
462
463        // In V2, A would win because it's unconflicted.
464        // In V2.1, B should win because it has a higher power level (100 > 50) and it's sorted together with A.
465        let resolved = resolve_lean(unconflicted, conflicted, StateResVersion::V2_1);
466        assert_eq!(
467            resolved.get(&("m.room.member".into(), "@alice:example.com".into())),
468            Some(&"B".into())
469        );
470    }
471
472    #[test]
473    fn test_v1_tie_break_by_id() {
474        let mut events = HashMap::new();
475        events.insert(
476            "B".into(),
477            LeanEvent {
478                event_id: "B".into(),
479                event_type: "m.room.member".into(),
480                state_key: "@alice:example.com".into(),
481                power_level: 0,
482                origin_server_ts: 100,
483                prev_events: vec![],
484                auth_events: vec![],
485                depth: 1,
486                ..Default::default()
487            },
488        );
489        events.insert(
490            "A".into(),
491            LeanEvent {
492                event_id: "A".into(),
493                event_type: "m.room.member".into(),
494                state_key: "@alice:example.com".into(),
495                power_level: 0,
496                origin_server_ts: 100,
497                prev_events: vec![],
498                auth_events: vec![],
499                depth: 1,
500                ..Default::default()
501            },
502        );
503        let sorted = lean_kahn_sort(&events, StateResVersion::V1);
504        assert_eq!(sorted, vec!["A", "B"]);
505    }
506
507    #[test]
508    fn test_v2_resolution_happy_path() {
509        let mut events = HashMap::new();
510        events.insert(
511            "A".into(),
512            LeanEvent {
513                event_id: "A".into(),
514                event_type: "m.room.member".into(),
515                state_key: "@alice:example.com".into(),
516                power_level: 100,
517                origin_server_ts: 100,
518                prev_events: vec![],
519                auth_events: vec![],
520                depth: 10,
521                ..Default::default()
522            },
523        );
524        events.insert(
525            "B".into(),
526            LeanEvent {
527                event_id: "B".into(),
528                event_type: "m.room.member".into(),
529                state_key: "@alice:example.com".into(),
530                power_level: 50,
531                origin_server_ts: 10,
532                prev_events: vec![],
533                auth_events: vec![],
534                depth: 1,
535                ..Default::default()
536            },
537        );
538        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
539        // Best (A) comes LAST.
540        assert_eq!(sorted, vec!["B", "A"]);
541    }
542
543    #[test]
544    fn test_v2_deep_tie_break() {
545        let mut events = HashMap::new();
546        events.insert(
547            "B".into(),
548            LeanEvent {
549                event_id: "B".into(),
550                event_type: "m.room.member".into(),
551                state_key: "@alice:example.com".into(),
552                power_level: 100,
553                origin_server_ts: 10,
554                prev_events: vec![],
555                auth_events: vec![],
556                depth: 1,
557                ..Default::default()
558            },
559        );
560        events.insert(
561            "A".into(),
562            LeanEvent {
563                event_id: "A".into(),
564                event_type: "m.room.member".into(),
565                state_key: "@alice:example.com".into(),
566                power_level: 100,
567                origin_server_ts: 10,
568                prev_events: vec![],
569                auth_events: vec![],
570                depth: 1,
571                ..Default::default()
572            },
573        );
574        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
575        // Best (B, larger ID) comes LAST.
576        assert_eq!(sorted, vec!["A", "B"]);
577    }
578
579    #[test]
580    fn test_v1_v2_v2_1_comparison_determinism() {
581        let mut events = HashMap::new();
582        events.insert(
583            "A".into(),
584            LeanEvent {
585                event_id: "A".into(),
586                event_type: "m.room.member".into(),
587                state_key: "@alice:example.com".into(),
588                power_level: 10,
589                origin_server_ts: 10,
590                prev_events: vec![],
591                auth_events: vec![],
592                depth: 1,
593                ..Default::default()
594            },
595        );
596        events.insert(
597            "B".into(),
598            LeanEvent {
599                event_id: "B".into(),
600                event_type: "m.room.member".into(),
601                state_key: "@alice:example.com".into(),
602                power_level: 100,
603                origin_server_ts: 100,
604                prev_events: vec![],
605                auth_events: vec![],
606                depth: 10,
607                ..Default::default()
608            },
609        );
610        let sorted_v1 = lean_kahn_sort(&events, StateResVersion::V1);
611        let sorted_v2 = lean_kahn_sort(&events, StateResVersion::V2);
612        let sorted_v2_1 = lean_kahn_sort(&events, StateResVersion::V2_1);
613        assert_eq!(sorted_v1, vec!["A", "B"]);
614        // B is better (higher power level), so it comes LAST in V2 and V2.1
615        assert_eq!(sorted_v2, vec!["A", "B"]);
616        assert_eq!(sorted_v2_1, vec!["A", "B"]);
617    }
618
619    #[test]
620    fn test_unhappy_path_cycle_detection() {
621        let mut events = HashMap::new();
622        events.insert(
623            "A".into(),
624            LeanEvent {
625                event_id: "A".into(),
626                event_type: "m.room.member".into(),
627                state_key: "@alice:example.com".into(),
628                power_level: 100,
629                origin_server_ts: 100,
630                prev_events: vec!["B".into()],
631                auth_events: vec!["B".into()],
632                depth: 1,
633                ..Default::default()
634            },
635        );
636        events.insert(
637            "B".into(),
638            LeanEvent {
639                event_id: "B".into(),
640                event_type: "m.room.member".into(),
641                state_key: "@alice:example.com".into(),
642                power_level: 100,
643                origin_server_ts: 100,
644                prev_events: vec!["A".into()],
645                auth_events: vec!["A".into()],
646                depth: 1,
647                ..Default::default()
648            },
649        );
650        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
651        assert!(sorted.is_empty());
652    }
653
654    #[test]
655    #[cfg(all(feature = "std", not(feature = "zkvm")))]
656    #[should_panic(expected = "Signature verification failed")]
657    fn test_signature_verification_failure() {
658        let pk = [
659            215, 90, 152, 1, 130, 177, 10, 183, 213, 75, 254, 211, 201, 100, 7, 58, 14, 225, 114,
660            243, 218, 166, 35, 37, 175, 2, 26, 104, 247, 7, 81, 26,
661        ];
662        let sig = [0u8; 64];
663        let msg = b"test";
664        verify_signature(&pk, msg, &sig);
665    }
666
667    #[test]
668    fn test_serialization_roundtrip() {
669        let event = LeanEvent {
670            event_id: "$abc".into(),
671            event_type: "m.room.member".into(),
672            state_key: "@alice:example.com".into(),
673            power_level: 100,
674            origin_server_ts: 12345,
675            prev_events: vec![],
676            auth_events: vec![],
677            depth: 5,
678            ..Default::default()
679        };
680        let serialized = serde_json::to_string(&event).unwrap();
681        let deserialized: LeanEvent = serde_json::from_str(&serialized).unwrap();
682        assert_eq!(event, deserialized);
683    }
684
685    #[test]
686    fn test_partial_ord_implementations() {
687        let e1 = LeanEvent {
688            event_id: "a".into(),
689            event_type: "m.room.member".into(),
690            state_key: "@alice:example.com".into(),
691            power_level: 100,
692            origin_server_ts: 10,
693            prev_events: vec![],
694            auth_events: vec![],
695            depth: 1,
696            ..Default::default()
697        };
698        let e2 = LeanEvent {
699            event_id: "b".into(),
700            event_type: "m.room.member".into(),
701            state_key: "@alice:example.com".into(),
702            power_level: 100,
703            origin_server_ts: 10,
704            prev_events: vec![],
705            auth_events: vec![],
706            depth: 1,
707            ..Default::default()
708        };
709        assert!(e1.partial_cmp(&e2).is_some());
710
711        let p1 = SortPriority {
712            event: &e1,
713            version: StateResVersion::V2,
714        };
715        let p2 = SortPriority {
716            event: &e2,
717            version: StateResVersion::V2,
718        };
719        assert!(p1.partial_cmp(&p2).is_some());
720    }
721
722    #[test]
723    fn test_trait_coverage() {
724        let v = StateResVersion::V2;
725        assert_eq!(v, StateResVersion::V2);
726        let _ = alloc::format!("{:?}", v);
727
728        let e = LeanEvent {
729            event_id: "a".into(),
730            event_type: "m.room.member".into(),
731            state_key: "@alice:example.com".into(),
732            power_level: 100,
733            origin_server_ts: 10,
734            prev_events: vec![],
735            auth_events: vec![],
736            depth: 1,
737            ..Default::default()
738        };
739        let _ = e.clone();
740        let _ = alloc::format!("{:?}", e);
741    }
742
743    #[test]
744    fn test_complex_dag_sort() {
745        let mut events = HashMap::new();
746        events.insert(
747            "1".into(),
748            LeanEvent {
749                event_id: "1".into(),
750                event_type: "m.room.member".into(),
751                state_key: "@alice:example.com".into(),
752                power_level: 100,
753                origin_server_ts: 10,
754                prev_events: vec![],
755                auth_events: vec![],
756                depth: 1,
757                ..Default::default()
758            },
759        );
760        events.insert(
761            "2".into(),
762            LeanEvent {
763                event_id: "2".into(),
764                event_type: "m.room.member".into(),
765                state_key: "@alice:example.com".into(),
766                power_level: 50,
767                origin_server_ts: 20,
768                prev_events: vec!["1".into()],
769                auth_events: vec!["1".into()],
770                depth: 2,
771                ..Default::default()
772            },
773        );
774        events.insert(
775            "3".into(),
776            LeanEvent {
777                event_id: "3".into(),
778                event_type: "m.room.member".into(),
779                state_key: "@alice:example.com".into(),
780                power_level: 50,
781                origin_server_ts: 15,
782                prev_events: vec!["1".into()],
783                auth_events: vec!["1".into()],
784                depth: 2,
785                ..Default::default()
786            },
787        );
788        events.insert(
789            "4".into(),
790            LeanEvent {
791                event_id: "4".into(),
792                event_type: "m.room.member".into(),
793                state_key: "@alice:example.com".into(),
794                power_level: 10,
795                origin_server_ts: 30,
796                prev_events: vec!["2".into(), "3".into()],
797                auth_events: vec!["2".into(), "3".into()],
798                depth: 3,
799                ..Default::default()
800            },
801        );
802        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
803        // 1 pops first (only one with in-degree 0).
804        // Then 2 and 3 are in queue. 3 is earlier (TS 15 < 20), so 3 pops first.
805        // Then 2 pops.
806        // Then 4 pops.
807        assert_eq!(sorted, vec!["1", "3", "2", "4"]);
808    }
809
810    #[test]
811    fn test_kahn_missing_parents() {
812        let mut events = HashMap::new();
813        events.insert(
814            "A".into(),
815            LeanEvent {
816                event_id: "A".into(),
817                event_type: "m.room.member".into(),
818                state_key: "@alice:example.com".into(),
819                power_level: 100,
820                origin_server_ts: 10,
821                prev_events: vec!["MISSING".into()],
822                auth_events: vec!["MISSING".into()],
823                depth: 1,
824                ..Default::default()
825            },
826        );
827        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
828        assert_eq!(sorted, vec!["A"]);
829    }
830
831    #[test]
832    fn test_resolve_lean_functionality() {
833        let mut unconflicted = BTreeMap::new();
834        unconflicted.insert(("type".into(), "key".into()), "id".into());
835        let conflicted = HashMap::new();
836        let resolved = resolve_lean(unconflicted.clone(), conflicted, StateResVersion::V2);
837        assert_eq!(resolved, unconflicted);
838    }
839
840    #[test]
841    fn test_resolve_lean_v2_1_overlay() {
842        let mut unconflicted = BTreeMap::new();
843        unconflicted.insert(("type1".into(), "key1".into()), "id1".into());
844        unconflicted.insert(("type2".into(), "key2".into()), "id2".into());
845
846        let mut conflicted = HashMap::new();
847        // Provide objects for all events to be sorted in V2.1
848        conflicted.insert(
849            "id1".into(),
850            LeanEvent {
851                event_id: "id1".into(),
852                event_type: "type1".into(),
853                state_key: "key1".into(),
854                power_level: 50,
855                origin_server_ts: 500,
856                prev_events: vec![],
857                auth_events: vec![],
858                depth: 1,
859                ..Default::default()
860            },
861        );
862        conflicted.insert(
863            "id2".into(),
864            LeanEvent {
865                event_id: "id2".into(),
866                event_type: "type2".into(),
867                state_key: "key2".into(),
868                power_level: 50,
869                origin_server_ts: 500,
870                prev_events: vec![],
871                auth_events: vec![],
872                depth: 1,
873                ..Default::default()
874            },
875        );
876        conflicted.insert(
877            "id2_new".into(),
878            LeanEvent {
879                event_id: "id2_new".into(),
880                event_type: "type2".into(),
881                state_key: "key2".into(),
882                power_level: 100,
883                origin_server_ts: 1000,
884                prev_events: vec![],
885                auth_events: vec![],
886                depth: 1,
887                ..Default::default()
888            },
889        );
890
891        let resolved = resolve_lean(unconflicted.clone(), conflicted, StateResVersion::V2_1);
892
893        assert_eq!(
894            resolved.get(&("type1".into(), "key1".into())),
895            Some(&"id1".into())
896        );
897        assert_eq!(
898            resolved.get(&("type2".into(), "key2".into())),
899            Some(&"id2_new".into())
900        );
901    }
902
903    fn run_batch_test(
904        version: StateResVersion,
905        rows: &[(&str, i64, u64, u64, &[&str])],
906        expected: &[&str],
907    ) {
908        let mut events = HashMap::new();
909        for r in rows {
910            events.insert(
911                r.0.to_string(),
912                LeanEvent {
913                    event_id: r.0.to_string(),
914                    event_type: "m.room.member".into(),
915                    state_key: "@alice:example.com".into(),
916                    power_level: r.1,
917                    origin_server_ts: r.2,
918                    depth: r.3,
919                    prev_events: r.4.iter().map(|s| s.to_string()).collect(),
920                    auth_events: r.4.iter().map(|s| s.to_string()).collect(),
921                    ..Default::default()
922                },
923            );
924        }
925        let result = lean_kahn_sort(&events, version);
926        assert_eq!(
927            result,
928            expected.iter().map(|s| s.to_string()).collect::<Vec<_>>()
929        );
930    }
931
932    #[test]
933    fn test_resolution_batch() {
934        run_batch_test(
935            StateResVersion::V2,
936            &[("Alice", 100, 500, 1, &[]), ("Bob", 50, 100, 1, &[])],
937            &["Bob", "Alice"], // Bob is worse (PL 50), pops first.
938        );
939        run_batch_test(
940            StateResVersion::V1,
941            &[("Deep", 100, 100, 10, &[]), ("Shallow", 10, 100, 1, &[])],
942            &["Shallow", "Deep"],
943        );
944    }
945
946    #[test]
947    fn test_native_resolution_bootstrap_parity() {
948        let mut events = HashMap::new();
949        events.insert(
950            "1".into(),
951            LeanEvent {
952                event_id: "1".into(),
953                event_type: "m.room.member".into(),
954                state_key: "@user:example.com".into(),
955                power_level: 100,
956                origin_server_ts: 10,
957                prev_events: vec![],
958                auth_events: vec![],
959                depth: 1,
960                ..Default::default()
961            },
962        );
963        events.insert(
964            "2".into(),
965            LeanEvent {
966                event_id: "2".into(),
967                event_type: "m.room.member".into(),
968                state_key: "@user:example.com".into(),
969                power_level: 0,
970                origin_server_ts: 20,
971                prev_events: vec!["1".into()],
972                auth_events: vec!["1".into()],
973                depth: 2,
974                ..Default::default()
975            },
976        );
977        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
978        let mut resolved_state = BTreeMap::new();
979        for id in sorted {
980            let ev = events.get(&id).unwrap();
981            let key = (ev.event_type.clone(), ev.state_key.clone());
982            resolved_state.insert(key, ev.event_id.clone());
983        }
984        assert_eq!(
985            resolved_state.get(&("m.room.member".to_string(), "@user:example.com".to_string())),
986            Some(&"2".to_string())
987        );
988    }
989
990    #[test]
991    fn test_enum_coverage() {
992        let v = StateResVersion::V2;
993        let v2 = v;
994        assert_eq!(v, v2);
995        let debug_str = alloc::format!("{:?}", v);
996        assert!(debug_str.contains("V2"));
997    }
998
999    #[test]
1000    fn test_event_traits_coverage() {
1001        let e = LeanEvent {
1002            event_id: "a".into(),
1003            event_type: "m.room.member".into(),
1004            state_key: "@alice:example.com".into(),
1005            power_level: 100,
1006            origin_server_ts: 10,
1007            prev_events: vec![],
1008            auth_events: vec![],
1009            depth: 1,
1010            ..Default::default()
1011        };
1012        let e2 = e.clone();
1013        assert_eq!(e, e2);
1014        let debug_str = alloc::format!("{:?}", e);
1015        assert!(debug_str.contains("event_id"));
1016    }
1017
1018    #[test]
1019    fn test_sort_priority_traits() {
1020        let e = LeanEvent {
1021            event_id: "a".into(),
1022            event_type: "m.room.member".into(),
1023            state_key: "@alice:example.com".into(),
1024            power_level: 100,
1025            origin_server_ts: 10,
1026            prev_events: vec![],
1027            auth_events: vec![],
1028            depth: 1,
1029            ..Default::default()
1030        };
1031        let p = SortPriority {
1032            event: &e,
1033            version: StateResVersion::V2,
1034        };
1035        let p2 = p;
1036        assert_eq!(p, p2);
1037        let debug_str = alloc::format!("{:?}", p);
1038        assert!(debug_str.contains("version"));
1039    }
1040
1041    #[test]
1042    fn test_v1_equal_depth_tie_break() {
1043        let mut events = HashMap::new();
1044        events.insert(
1045            "B".into(),
1046            LeanEvent {
1047                event_id: "B".into(),
1048                event_type: "m.room.member".into(),
1049                state_key: "@alice:example.com".into(),
1050                power_level: 0,
1051                origin_server_ts: 10,
1052                prev_events: vec![],
1053                auth_events: vec![],
1054                depth: 1,
1055                ..Default::default()
1056            },
1057        );
1058        events.insert(
1059            "A".into(),
1060            LeanEvent {
1061                event_id: "A".into(),
1062                event_type: "m.room.member".into(),
1063                state_key: "@alice:example.com".into(),
1064                power_level: 0,
1065                origin_server_ts: 10,
1066                prev_events: vec![],
1067                auth_events: vec![],
1068                depth: 1,
1069                ..Default::default()
1070            },
1071        );
1072        let sorted = lean_kahn_sort(&events, StateResVersion::V1);
1073        assert_eq!(sorted, vec!["A", "B"]);
1074    }
1075
1076    #[test]
1077    fn test_kahn_no_neighbors() {
1078        let mut events = HashMap::new();
1079        events.insert(
1080            "1".into(),
1081            LeanEvent {
1082                event_id: "1".into(),
1083                event_type: "m.room.member".into(),
1084                state_key: "@alice:example.com".into(),
1085                power_level: 100,
1086                origin_server_ts: 10,
1087                prev_events: vec![],
1088                auth_events: vec![],
1089                depth: 1,
1090                ..Default::default()
1091            },
1092        );
1093        let sorted = lean_kahn_sort(&events, StateResVersion::V2);
1094        assert_eq!(sorted, vec!["1"]);
1095    }
1096
1097    #[test]
1098    fn test_v2_1_full_coverage() {
1099        let mut events = HashMap::new();
1100        events.insert(
1101            "A".into(),
1102            LeanEvent {
1103                event_id: "A".into(),
1104                event_type: "m.room.member".into(),
1105                state_key: "@alice:example.com".into(),
1106                power_level: 100,
1107                origin_server_ts: 10,
1108                prev_events: vec![],
1109                auth_events: vec![],
1110                depth: 1,
1111                ..Default::default()
1112            },
1113        );
1114        let sorted = lean_kahn_sort(&events, StateResVersion::V2_1);
1115        assert_eq!(sorted, vec!["A"]);
1116    }
1117
1118    #[test]
1119    fn test_total_order_properties() {
1120        let e1 = LeanEvent {
1121            event_id: "a".into(),
1122            event_type: "m.room.member".into(),
1123            state_key: "@alice:example.com".into(),
1124            power_level: 100,
1125            origin_server_ts: 10,
1126            prev_events: vec![],
1127            auth_events: vec![],
1128            depth: 1,
1129            ..Default::default()
1130        };
1131        let e2 = LeanEvent {
1132            event_id: "b".into(),
1133            event_type: "m.room.member".into(),
1134            state_key: "@alice:example.com".into(),
1135            power_level: 100,
1136            origin_server_ts: 10,
1137            prev_events: vec![],
1138            auth_events: vec![],
1139            depth: 1,
1140            ..Default::default()
1141        };
1142        let e3 = LeanEvent {
1143            event_id: "c".into(),
1144            event_type: "m.room.member".into(),
1145            state_key: "@alice:example.com".into(),
1146            power_level: 50,
1147            origin_server_ts: 10,
1148            prev_events: vec![],
1149            auth_events: vec![],
1150            depth: 1,
1151            ..Default::default()
1152        };
1153        assert_eq!(e1.cmp(&e1), Ordering::Equal);
1154        assert!(e1 <= e1);
1155        assert!(e1 <= e2 || e2 <= e1);
1156        if e1 <= e2 && e2 <= e3 {
1157            assert!(e1 <= e3);
1158        }
1159        let e1_copy = e1.clone();
1160        if e1 <= e1_copy && e1_copy <= e1 {
1161            assert_eq!(e1, e1_copy);
1162        }
1163    }
1164
1165    #[test]
1166    fn test_coverage_booster_all_branches() {
1167        let e_base = LeanEvent {
1168            event_id: "m".into(),
1169            event_type: "m.room.member".into(),
1170            state_key: "@alice:example.com".into(),
1171            power_level: 50,
1172            origin_server_ts: 50,
1173            prev_events: vec![],
1174            auth_events: vec![],
1175            depth: 50,
1176            ..Default::default()
1177        };
1178        let p_base = SortPriority {
1179            event: &e_base,
1180            version: StateResVersion::V2,
1181        };
1182        let e_high_power = LeanEvent {
1183            power_level: 100,
1184            ..e_base.clone()
1185        };
1186        let p_high_power = SortPriority {
1187            event: &e_high_power,
1188            version: StateResVersion::V2,
1189        };
1190        // p_base is WORSE (PL 50 < 100), so it should be GREATER.
1191        assert_eq!(p_base.cmp(&p_high_power), Ordering::Greater);
1192        let e_early_ts = LeanEvent {
1193            origin_server_ts: 10,
1194            ..e_base.clone()
1195        };
1196        let p_early_ts = SortPriority {
1197            event: &e_early_ts,
1198            version: StateResVersion::V2,
1199        };
1200        // p_early_ts has TS 10, p_base has TS 50. Earlier TS pops first, so p_early_ts is GREATER. p_base < p_early_ts => Less.
1201        assert_eq!(p_base.cmp(&p_early_ts), Ordering::Less);
1202        let e_early_id = LeanEvent {
1203            event_id: "a".into(),
1204            ..e_base.clone()
1205        };
1206        let p_early_id = SortPriority {
1207            event: &e_early_id,
1208            version: StateResVersion::V2,
1209        };
1210        // p_early_id has ID "a", p_base has ID "m". Smaller ID pops first, so p_early_id is GREATER. p_base < p_early_id => Less.
1211        assert_eq!(p_base.cmp(&p_early_id), Ordering::Less);
1212        let p_v1_base = SortPriority {
1213            event: &e_base,
1214            version: StateResVersion::V1,
1215        };
1216        let e_shallow = LeanEvent {
1217            depth: 1,
1218            ..e_base.clone()
1219        };
1220        let p_shallow = SortPriority {
1221            event: &e_shallow,
1222            version: StateResVersion::V1,
1223        };
1224        assert_eq!(p_v1_base.cmp(&p_shallow), Ordering::Less);
1225        let p_v1_early_id = SortPriority {
1226            event: &e_early_id,
1227            version: StateResVersion::V1,
1228        };
1229        assert_eq!(p_v1_base.cmp(&p_v1_early_id), Ordering::Less);
1230        assert_eq!(p_v1_base.cmp(&p_v1_base), Ordering::Equal);
1231    }
1232}