diff --git a/Cargo.lock b/Cargo.lock index 87f50ea5dfa..04ae3152360 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -52,13 +52,19 @@ version = "2.9.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1b8e56985ec62d17e9c1001dc89c88ecd7dc08e47eba5ec7c29c7b5eeecde967" +[[package]] +name = "bitmaps" +version = "3.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1d084b0137aaa901caf9f1e8b21daa6aa24d41cd806e111335541eff9683bd6" + [[package]] name = "blake2" version = "0.10.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "46502ad458c9a52b69d4d4d32775c788b7a1b85e8bc9d482d92250fc0e3f8efe" dependencies = [ - "digest", + "digest 0.10.7", ] [[package]] @@ -70,12 +76,33 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-buffer" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2f6c7dbe95a6ed67ad9f18e57daf93a2f034c524b99fd2b76d18fdfeb6660aa" +dependencies = [ + "hybrid-array", +] + [[package]] name = "bumpalo" version = "3.19.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "46c5e41b57b8bba42a04676d81cb89e9ee8e859a1a66f80a5a72e1cb76b34d43" +[[package]] +name = "bytemuck" +version = "1.25.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec" + +[[package]] +name = "byteorder" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" + [[package]] name = "bytes" version = "1.11.1" @@ -137,6 +164,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + [[package]] name = "crypto-common" version = "0.1.6" @@ -147,17 +183,36 @@ dependencies = [ "typenum", ] +[[package]] +name = "crypto-common" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce6e4c961d6cd6c9a86db418387425e8bdeaf05b3c8bc1411e6dca4c252f1453" +dependencies = [ + "hybrid-array", +] + [[package]] name = "digest" version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer", - "crypto-common", + "block-buffer 0.10.4", + "crypto-common 0.1.6", "subtle", ] +[[package]] +name = "digest" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" +dependencies = [ + "block-buffer 0.12.1", + "crypto-common 0.2.2", +] + [[package]] name = "displaydoc" version = "0.2.5" @@ -187,6 +242,12 @@ version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + [[package]] name = "form_urlencoded" version = "1.2.1" @@ -345,6 +406,9 @@ name = "hashbrown" version = "0.15.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5971ac85611da7067dbfcabef3c70ebb5606018acd9e2a3903a0da507521e0d5" +dependencies = [ + "foldhash", +] [[package]] name = "headers" @@ -427,6 +491,15 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" +[[package]] +name = "hybrid-array" +version = "0.4.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3944cf8cf766b40e2a1a333ee5e9b563f854d5fa49d6a8ca2764e97c6eddb214" +dependencies = [ + "typenum", +] + [[package]] name = "hyper" version = "1.6.0" @@ -641,6 +714,29 @@ dependencies = [ "icu_properties", ] +[[package]] +name = "imbl" +version = "4.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ae128b3bc67ed43ec0a7bb1c337a9f026717628b3c4033f07ded1da3e854951" +dependencies = [ + "bitmaps", + "imbl-sized-chunks", + "rand_core 0.6.4", + "rand_xoshiro", + "serde", + "version_check", +] + +[[package]] +name = "imbl-sized-chunks" +version = "0.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f4241005618a62f8d57b2febd02510fb96e0137304728543dfc5fd6f052c22d" +dependencies = [ + "bitmaps", +] + [[package]] name = "indexmap" version = "2.10.0" @@ -692,6 +788,16 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "keccak" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e24a010dd405bd7ed803e5253182815b41bf2e6a80cc3bfc066658e03a198aa" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", +] + [[package]] name = "lazy_static" version = "1.5.0" @@ -968,7 +1074,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea" dependencies = [ "rand_chacha", - "rand_core", + "rand_core 0.9.3", ] [[package]] @@ -978,9 +1084,15 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" dependencies = [ "ppv-lite86", - "rand_core", + "rand_core 0.9.3", ] +[[package]] +name = "rand_core" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c" + [[package]] name = "rand_core" version = "0.9.3" @@ -990,6 +1102,15 @@ dependencies = [ "getrandom 0.3.3", ] +[[package]] +name = "rand_xoshiro" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f97cdb2a36ed4183de61b2f824cc45c9f1037f28afe0a322e9fff4c108b5aaa" +dependencies = [ + "rand_core 0.6.4", +] + [[package]] name = "regex" version = "1.12.4" @@ -1061,6 +1182,21 @@ dependencies = [ "web-sys", ] +[[package]] +name = "rezzy" +version = "0.5.2" +source = "git+https://github.com/gamesguru/rezzy.git?branch=dev#92e31a3a041980ee14f35046ee517308b024f5cd" +dependencies = [ + "blake2", + "hashbrown", + "imbl", + "roaring", + "serde", + "serde_json", + "sha2 0.11.0", + "sha3", +] + [[package]] name = "ring" version = "0.17.14" @@ -1075,6 +1211,16 @@ dependencies = [ "windows-sys 0.52.0", ] +[[package]] +name = "roaring" +version = "0.10.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "19e8d2cfa184d94d0726d650a9f4a1be7f9b76ac9fdb954219878dc00c1c1e7b" +dependencies = [ + "bytemuck", + "byteorder", +] + [[package]] name = "rustc-hash" version = "2.1.1" @@ -1249,8 +1395,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", ] [[package]] @@ -1260,8 +1406,29 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", +] + +[[package]] +name = "sha2" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.3", +] + +[[package]] +name = "sha3" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be176f1a57ce4e3d31c1a166222d9768de5954f811601fb7ca06fc8203905ce1" +dependencies = [ + "digest 0.11.3", + "keccak", ] [[package]] @@ -1350,10 +1517,11 @@ dependencies = [ "pythonize", "regex", "reqwest", + "rezzy", "rustc_version", "serde", "serde_json", - "sha2", + "sha2 0.10.9", "tokio", "ulid", ] diff --git a/changelog.d/15.feature b/changelog.d/15.feature new file mode 120000 index 00000000000..4f27f8c591a --- /dev/null +++ b/changelog.d/15.feature @@ -0,0 +1 @@ +19945.feature \ No newline at end of file diff --git a/changelog.d/19945.feature b/changelog.d/19945.feature new file mode 100644 index 00000000000..1d0bbab8e7b --- /dev/null +++ b/changelog.d/19945.feature @@ -0,0 +1 @@ +Add preliminary Rust-backed state res (v2.0) and auth chain processing. diff --git a/rust/Cargo.toml b/rust/Cargo.toml index 612ab09f6d0..a7755ad981d 100644 --- a/rust/Cargo.toml +++ b/rust/Cargo.toml @@ -64,6 +64,10 @@ futures = "0.3.31" tokio = { version = "1.44.2", features = ["rt", "rt-multi-thread"] } once_cell = "1.18.0" itertools = "0.14.0" +# rezzy = "0.5.1" +# rezzy = { path = "../../rezzy" } +# rezzy = { git = "https://github.com/gamesguru/rezzy.git", rev = "50ac86688fb0e978daa3c233429faa8341662b4d" } +rezzy = { git = "https://github.com/gamesguru/rezzy.git", branch = "dev" } [build-dependencies] blake2 = "0.10.4" diff --git a/rust/src/events/internal_metadata.rs b/rust/src/events/internal_metadata.rs index 0778fbfeaa2..78b55c2831c 100644 --- a/rust/src/events/internal_metadata.rs +++ b/rust/src/events/internal_metadata.rs @@ -620,7 +620,7 @@ impl EventInternalMetadata { Ok(self.read_inner()?.need_to_check_redaction()) } - fn is_soft_failed(&self) -> PyResult { + pub fn is_soft_failed(&self) -> PyResult { Ok(self.read_inner()?.is_soft_failed()) } diff --git a/rust/src/events/mod.rs b/rust/src/events/mod.rs index 0d746a2ec81..1742d647342 100644 --- a/rust/src/events/mod.rs +++ b/rust/src/events/mod.rs @@ -96,6 +96,21 @@ pub mod utils; use json_object::JsonObject; +#[derive(Clone)] +pub(crate) struct EventResolverData { + pub(crate) event_id: String, + pub(crate) event_type: String, + pub(crate) state_key: Option, + pub(crate) sender: String, + pub(crate) origin_server_ts: u64, + pub(crate) depth: u64, + pub(crate) prev_events: Vec, + pub(crate) auth_events: Vec, + pub(crate) content: Value, + pub(crate) rejected: bool, + pub(crate) soft_failed: bool, +} + /// Called when registering modules with python. pub fn register_module(py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { // Register the `JsonObject` class as a `Mapping` so that `isinstance` works. @@ -344,7 +359,7 @@ impl Event { /// Returns the list of auth event IDs. The order matches the order /// specified in the event, though there is no meaning to it. - fn auth_event_ids(&self) -> PyResult> { + pub(crate) fn auth_event_ids(&self) -> PyResult> { match &*self.parsed_event.specific_fields { EventFormatEnum::V1(format) => Ok(format.auth_event_ids()), EventFormatEnum::V2V3(format) => Ok(format.auth_event_ids()), @@ -631,6 +646,38 @@ impl Event { } } +impl Event { + pub(crate) fn resolver_data(&self) -> PyResult { + let origin_server_ts = u64::try_from(self.origin_server_ts()).map_err(|_| { + PyValueError::new_err(format!( + "event {} has a negative origin_server_ts", + self.event_id + )) + })?; + let depth = u64::try_from(self.depth()).map_err(|_| { + PyValueError::new_err(format!("event {} has a negative depth", self.event_id)) + })?; + let content = + serde_json::to_value(&self.parsed_event.common_fields.content).map_err(|err| { + PyValueError::new_err(format!("Failed to serialize event content: {err}")) + })?; + + Ok(EventResolverData { + event_id: self.event_id.to_string(), + event_type: self.r#type().to_string(), + state_key: self.get_state_key().map(str::to_owned), + sender: self.sender().to_string(), + origin_server_ts, + depth, + prev_events: self.prev_event_ids(), + auth_events: self.auth_event_ids()?, + content, + rejected: self.rejected_reason.is_some(), + soft_failed: self.internal_metadata.is_soft_failed()?, + }) + } +} + /// Parses a JSON string into a [`FormattedEvent`] for the given room version. fn event_dict_from_json_str( room_version: &RoomVersion, diff --git a/rust/src/lib.rs b/rust/src/lib.rs index 28783afbbac..7d0087b6c61 100644 --- a/rust/src/lib.rs +++ b/rust/src/lib.rs @@ -22,6 +22,7 @@ pub mod push; pub mod rendezvous; pub mod room_versions; pub mod segmenter; +pub mod state_res; pub mod storage; pub mod tokio_runtime; pub mod types; @@ -79,6 +80,7 @@ fn synapse_rust(py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { msc4388_rendezvous::register_module(py, m)?; segmenter::register_module(py, m)?; room_versions::register_module(py, m)?; + state_res::register_module(py, m)?; types::register_module(py, m)?; Ok(()) diff --git a/rust/src/state_res.rs b/rust/src/state_res.rs new file mode 100644 index 00000000000..47ae331f12c --- /dev/null +++ b/rust/src/state_res.rs @@ -0,0 +1,218 @@ +/* + * This file is licensed under the Affero General Public License (AGPL) version 3. + * + * Copyright (C) 2026 Element Creations Ltd. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as + * published by the Free Software Foundation, either version 3 of the + * License, or (at your option) any later version. + * + * See the GNU Affero General Public License for more details: + * . + */ + +use std::collections::{HashMap, HashSet}; + +use pyo3::prelude::*; +use pyo3::types::{PyAny, PyDict, PySet, PyTuple}; +use pythonize::depythonize; +use rezzy::{ + auth::roaring::AuthGraph, resolve_lattice_fold, LeanEvent, SharedState, StateResVersion, +}; +use serde_json::Value; + +use crate::events::{Event, EventResolverData}; + +#[pyfunction] +#[pyo3(text_signature = "(state_sets, event_map, /)")] +pub fn get_auth_chain_difference_from_event_graph<'py>( + py: Python<'py>, + state_sets: Bound<'py, PyAny>, + event_map: Bound<'py, PyDict>, +) -> PyResult> { + let mut auth_graph_events: HashMap> = + HashMap::with_capacity(event_map.len()); + for (k, v) in event_map.iter() { + let event_id: String = k.extract()?; + let auth_ids: Vec = if let Ok(event) = v.extract::>() { + event.auth_event_ids()? + } else { + v.call_method0("auth_event_ids")?.extract()? + }; + auth_graph_events.insert( + event_id.clone(), + LeanEvent { + event_id, + event_type: String::new(), + state_key: None, + power_level: 0, + origin_server_ts: 0, + sender: String::new(), + content: (), + prev_events: Vec::new(), + auth_events: auth_ids, + depth: 0, + rejected: false, + soft_fail: false, + }, + ); + } + let auth_graph = AuthGraph::build(&auth_graph_events); + + let mut union: Option> = None; + let mut intersection: HashSet = HashSet::new(); + + for state_set in state_sets.try_iter()? { + let state_set = state_set?; + let values = state_set.call_method0("values")?; + let mut state_set_ids = Vec::with_capacity(values.len()?); + for value in values.try_iter()? { + state_set_ids.push(value?.extract()?); + } + let closure: HashSet = auth_graph + .auth_difference(&[], &state_set_ids) + .into_iter() + .collect(); + + match &mut union { + None => { + intersection = closure.clone(); + union = Some(closure); + } + Some(union) => { + union.extend(closure.iter().cloned()); + intersection = intersection.intersection(&closure).cloned().collect(); + } + } + } + + let Some(union) = union else { + return PySet::empty(py); + }; + + let result: HashSet = union.difference(&intersection).cloned().collect(); + PySet::new(py, result) +} + +fn resolver_data_to_lean_event(data: EventResolverData) -> LeanEvent { + LeanEvent { + event_id: data.event_id, + event_type: data.event_type, + state_key: data.state_key, + power_level: 0, + origin_server_ts: data.origin_server_ts, + sender: data.sender, + content: data.content, + prev_events: data.prev_events, + auth_events: data.auth_events, + depth: data.depth, + rejected: data.rejected, + soft_fail: data.soft_failed, + } +} + +fn py_to_lean_event(py_ev: &Bound<'_, PyAny>) -> PyResult> { + let event_id: String = py_ev.getattr("event_id")?.extract()?; + let event_type: String = py_ev.getattr("type")?.extract()?; + let state_key: Option = py_ev.call_method0("get_state_key")?.extract()?; + let sender: String = py_ev.getattr("sender")?.extract()?; + let origin_server_ts: u64 = py_ev.getattr("origin_server_ts")?.extract()?; + let depth: u64 = py_ev.getattr("depth")?.extract()?; + + let prev_events: Vec = py_ev.call_method0("prev_event_ids")?.extract()?; + let auth_events: Vec = py_ev.call_method0("auth_event_ids")?.extract()?; + let rejected_reason: Option = py_ev.getattr("rejected_reason")?.extract()?; + let soft_failed: bool = py_ev + .getattr("internal_metadata")? + .call_method0("is_soft_failed")? + .extract()?; + + let py_content = py_ev.getattr("content")?; + let content: Value = depythonize(&py_content)?; + + let power_level: i64 = 0; + + Ok(LeanEvent { + event_id, + event_type, + state_key, + power_level, + origin_server_ts, + sender, + content, + prev_events, + auth_events, + depth, + rejected: rejected_reason.is_some(), + soft_fail: soft_failed, + }) +} + +#[pyfunction] +#[pyo3(text_signature = "(unconflicted_state, conflicted_event_ids, event_map, /)")] +pub fn resolve_v2_via_lattice_fold<'py>( + py: Python<'py>, + unconflicted_state: Bound<'py, PyDict>, + conflicted_event_ids: Bound<'py, PyAny>, + event_map: Bound<'py, PyDict>, +) -> PyResult> { + let version = StateResVersion::V2; + + let mut unconf_state = SharedState::new(); + for (k, v) in unconflicted_state.iter() { + let key: (String, String) = k.extract()?; + let val: String = v.extract()?; + unconf_state.insert(key, val); + } + + let conflicted_ids: Vec = conflicted_event_ids.extract()?; + + let mut parsed_events: HashMap> = + HashMap::with_capacity(event_map.len()); + for (k, v) in event_map.iter() { + let event_id: String = k.extract()?; + let lean_ev = if let Ok(event) = v.extract::>() { + resolver_data_to_lean_event(event.resolver_data()?) + } else { + py_to_lean_event(&v)? + }; + parsed_events.insert(event_id, lean_ev); + } + + let mut conflicted_events = HashMap::with_capacity(conflicted_ids.len()); + for id in conflicted_ids { + if let Some(ev) = parsed_events.get(&id) { + conflicted_events.insert(id.clone(), ev.clone()); + } + } + + let resolved = resolve_lattice_fold(unconf_state, conflicted_events, &parsed_events, version); + + let py_resolved = PyDict::new(py); + for ((type_, state_key), event_id) in resolved { + let py_key = PyTuple::new(py, &[type_, state_key])?; + py_resolved.set_item(py_key, event_id)?; + } + + Ok(py_resolved) +} + +pub fn register_module(py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { + let child_module = PyModule::new(py, "state_res")?; + child_module.add_function(wrap_pyfunction!( + get_auth_chain_difference_from_event_graph, + &child_module + )?)?; + child_module.add_function(wrap_pyfunction!( + resolve_v2_via_lattice_fold, + &child_module + )?)?; + m.add_submodule(&child_module)?; + + py.import("sys")? + .getattr("modules")? + .set_item("synapse.synapse_rust.state_res", child_module)?; + + Ok(()) +} diff --git a/scripts-dev/benchmark_state_res.py b/scripts-dev/benchmark_state_res.py new file mode 100755 index 00000000000..40ef2224a2d --- /dev/null +++ b/scripts-dev/benchmark_state_res.py @@ -0,0 +1,766 @@ +#!/usr/bin/env python +# +# This file is licensed under the Affero General Public License (AGPL) version 3. +# +# Copyright (C) 2026 Element Creations Ltd. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as +# published by the Free Software Foundation, either version 3 of the +# License, or (at your option) any later version. +# +# See the GNU Affero General Public License for more details: +# . + +import argparse +import asyncio +import cProfile +import hashlib +import io +import json +import pstats +import re +import sys +import time +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from textwrap import dedent +from typing import Any, Collection, Iterable, Iterator, Mapping, cast + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +import synapse.event_auth # noqa: E402 # Import first to resolve circular import dependency +import synapse.state # noqa: F401,E402 +import synapse.state.v2 as v2 # noqa: E402 +from synapse.api.room_versions import ( # noqa: E402 + RoomVersion, + RoomVersions, +) +from synapse.events import EventBase, make_event_from_dict # noqa: E402 +from synapse.state import StateDifference # noqa: E402 +from synapse.util.duration import Duration # noqa: E402 + + +class _RustResolverDisabledForBenchmark(RuntimeError): + pass + + +@contextmanager +def _disable_rust_lattice_fold_resolver() -> Iterator[None]: + import synapse.synapse_rust.state_res as rust_res + + original = rust_res.resolve_v2_via_lattice_fold + original_logger_exception = v2.logger.exception + + def disabled_resolver(*args: Any, **kwargs: Any) -> Any: + raise _RustResolverDisabledForBenchmark( + "Rust lattice-fold state resolver disabled for benchmark" + ) + + def logger_exception(message: object, *args: Any, **kwargs: Any) -> None: + if isinstance(kwargs.get("exc_info"), _RustResolverDisabledForBenchmark): + return + original_logger_exception(message, *args, **kwargs) + + cast(Any, rust_res).resolve_v2_via_lattice_fold = disabled_resolver + cast(Any, v2.logger).exception = logger_exception + try: + yield + finally: + cast(Any, rust_res).resolve_v2_via_lattice_fold = original + cast(Any, v2.logger).exception = original_logger_exception + + +# Mock Clock +class MockClock: + def time_msec(self) -> int: + return int(time.time() * 1000) + + async def sleep(self, duration: Duration) -> None: + await asyncio.sleep(0) + + +# Mock Room Version +class MockRoomVersion: + def __init__(self, real_version: RoomVersion, state_res: int | None) -> None: + self.state_res = state_res + for attr in dir(real_version): + if not attr.startswith("_"): + try: + setattr(self, attr, getattr(real_version, attr)) + except (AttributeError, TypeError): + pass + self.state_res = state_res + + +# Mock Event compatible with PyO3 translation layer +class MockEvent: + room_version: MockRoomVersion | None = None + + class _InternalMetadata: + def is_soft_failed(self) -> bool: + return False + + def __init__( + self, + event_id: str, + sender: str, + event_type: str, + state_key: str | None, + content: dict, + origin_server_ts: int = 0, + depth: int = 1, + auth_event_ids: list[str] | None = None, + prev_event_ids: list[str] | None = None, + room_id: str = "!room:example.com", + ): + self.event_id = event_id + self.sender = sender + self.type = event_type + self.state_key = state_key + self.content = content + self.room_id = room_id + self._auth_event_ids: list[str] = auth_event_ids or [] + self._prev_event_ids: list[str] = prev_event_ids or [] + self.depth = depth + self.origin_server_ts = origin_server_ts + self.rejected_reason = None + self.internal_metadata = MockEvent._InternalMetadata() + + @property + def membership(self) -> str | None: + return self.content.get("membership") + + def auth_event_ids(self) -> list[str]: + return self._auth_event_ids + + def prev_event_ids(self) -> list[str]: + return self._prev_event_ids + + def get_state_key(self) -> str | None: + return self.state_key + + +# Mock Store matching Synapse's StateResolutionStore API +class MockStateResolutionStore: + def __init__(self, event_map: dict[str, Any]): + self.event_map = event_map + self.auth_chains: dict[str, set[str]] = {} + + async def get_events( + self, event_ids: Collection[str], allow_rejected: bool = False + ) -> dict[str, EventBase]: + missing = [eid for eid in event_ids if eid not in self.event_map] + if missing: + raise KeyError(f"Missing benchmark events: {missing}") + + return cast( + dict[str, EventBase], + {eid: self.event_map[eid] for eid in event_ids}, + ) + + def _get_auth_chain(self, event_ids: Iterable[str]) -> list[str]: + event_ids = list(event_ids) + if self.auth_chains and all(eid in self.auth_chains for eid in event_ids): + result = set() + for eid in event_ids: + result.update(self.auth_chains[eid]) + return list(result) + + result = set() + stack = list(event_ids) + while stack: + event_id = stack.pop() + if event_id in result: + continue + result.add(event_id) + event = self.event_map[event_id] + for aid in event.auth_event_ids(): + stack.append(aid) + return list(result) + + async def get_auth_chain_difference( + self, + room_id: str, + auth_sets: list[set[str]], + conflicted_state: set[str] | None, + additional_backwards_reachable_conflicted_events: set[str] | None, + ) -> StateDifference: + chains = [frozenset(self._get_auth_chain(a)) for a in auth_sets] + common = set(chains[0]).intersection(*chains[1:]) + return StateDifference( + auth_difference=set().union(*chains) - common, + conflicted_subgraph=None, + ) + + +def _make_benchmark_event( + *, + room_version: RoomVersion, + event_id: str, + sender: str, + event_type: str, + state_key: str | None, + content: dict[str, Any], + origin_server_ts: int, + depth: int, + auth_event_ids: list[str] | None = None, + prev_event_ids: list[str] | None = None, + room_id: str = "!room:example.com", +) -> EventBase: + reference_hashes = {"sha256": "benchmark"} + event_dict: dict[str, Any] = { + "room_id": room_id, + "type": event_type, + "sender": sender, + "content": content, + "depth": depth, + "origin_server_ts": origin_server_ts, + "hashes": {"sha256": "aGVsbG8="}, + "signatures": {}, + "auth_events": [ + (event_id, reference_hashes) for event_id in (auth_event_ids or []) + ], + "prev_events": [ + (event_id, reference_hashes) for event_id in (prev_event_ids or []) + ], + "event_id": event_id, + } + if state_key is not None: + event_dict["state_key"] = state_key + return make_event_from_dict(event_dict, room_version=room_version) + + +@dataclass(slots=True) +class RunStats: + total_s: float + bookkeeping_s: float + resolve_s: float + merge_points: int + + +def _print_profile(profile: cProfile.Profile, title: str, limit: int) -> None: + stream = io.StringIO() + stats = pstats.Stats(profile, stream=stream) + stats.sort_stats("cumtime") + stats.print_stats(limit) + print(f"\n{title}") + print(stream.getvalue().rstrip()) + + +def _resolved_state_checksum(state: Mapping[tuple[str, str], str]) -> str: + encoded_state = json.dumps( + [ + [event_type, state_key, event_id] + for (event_type, state_key), event_id in sorted(state.items()) + ], + separators=(",", ":"), + ) + return hashlib.sha256(encoded_state.encode()).hexdigest() + + +def _print_results_table( + stats_py: RunStats, + stats_rust: RunStats, + checksum_py: str, + checksum_rust: str, + title: str, +) -> None: + speedup_py = "1.0x (Baseline)" + speedup_rust = f"{stats_py.total_s / stats_rust.total_s:.1f}x" + speedup_width = max(len("Speedup"), len(speedup_py), len(speedup_rust)) + checksum_width = len(checksum_py) + + print( + dedent( + f""" + {title} + +--------------------+---------------+{"-" * (speedup_width + 2)}+{"-" * (checksum_width + 2)}+ + | {"Implementation":<18} | {"Duration (s)":<13} | {"Speedup":<{speedup_width}} | {"State SHA-256":<{checksum_width}} | + +--------------------+---------------+{"-" * (speedup_width + 2)}+{"-" * (checksum_width + 2)}+ + | {"Python V2":<18} | {stats_py.total_s:<13.5f} | {speedup_py:<{speedup_width}} | {checksum_py:<{checksum_width}} | + | {"Rust (rezzy)":<18} | {stats_rust.total_s:<13.5f} | {speedup_rust:<{speedup_width}} | {checksum_rust:<{checksum_width}} | + +--------------------+---------------+{"-" * (speedup_width + 2)}+{"-" * (checksum_width + 2)}+ + """ + ).rstrip() + ) + + +def _print_stage_breakdown(stats_py: RunStats, stats_rust: RunStats) -> None: + def format_stage(label: str, stats: RunStats) -> str: + parts = [f"{label:<6} total: {stats.total_s:.5f}s"] + if stats.bookkeeping_s: + parts.append(f"bookkeeping: {stats.bookkeeping_s:.5f}s") + if stats.resolve_s != stats.total_s: + parts.append(f"resolver: {stats.resolve_s:.5f}s") + parts.append(f"merge points: {stats.merge_points}") + return " - " + ", ".join(parts) + + print("\nStage breakdown:") + print(format_stage("Python", stats_py)) + print(format_stage("Rust", stats_rust)) + + +def _load_jsonl_events(path: str) -> tuple[dict[str, Any], list[MockEvent]]: + print(f"Loading DAG from {path}...") + event_map: dict[str, Any] = {} + events_list: list[MockEvent] = [] + room_id_hint = None + filename = Path(path).name + match = re.match( + r"^(?:local|remote)-dag-(.+)-v\d+-.+\.jsonl$", + filename, + ) + if match: + room_id_hint = f"!{match.group(1)}" + + with open(path, "r") as f: + for line_no, line in enumerate(f, start=1): + if not line.strip(): + continue + + d = json.loads(line) + if "event_id" not in d: + keys = ", ".join(sorted(d.keys())) + raise ValueError( + "The --jsonl input must include a top-level 'event_id' field for each " + f"event. The first parsed row at line {line_no} had keys: {keys}. " + "This file looks like a DAG export that omits event IDs, so the " + "benchmark cannot build its event_map or resolve prev/auth chains from it." + ) + + ev = MockEvent( + event_id=d["event_id"], + sender=d["sender"], + event_type=d["type"], + state_key=d.get("state_key"), + content=d["content"], + origin_server_ts=d["origin_server_ts"], + depth=d["depth"], + auth_event_ids=d["auth_events"], + prev_event_ids=d["prev_events"], + room_id=d.get("room_id", room_id_hint) or "", + ) + if not ev.room_id: + raise ValueError( + "The --jsonl input must include a top-level 'room_id' field, or use " + "a file name that encodes the room id (for example local-dag--v*.jsonl). " + f"The row at line {line_no} had no room_id and the filename {filename!r} " + "did not match the expected pattern." + ) + event_map[ev.event_id] = ev + events_list.append(ev) + + return event_map, events_list + + +async def main() -> None: + import logging + + logging.basicConfig(level=logging.INFO, force=True) + parser = argparse.ArgumentParser( + description="Benchmark state resolution V2 (Rust vs Python)" + ) + parser.add_argument( + "-p", "--partitions", type=int, default=50, help="Number of partitions (P)" + ) + parser.add_argument( + "-n", + "--events", + type=int, + default=100, + help="Conflicting events per partition (N)", + ) + parser.add_argument( + "--jsonl", + type=str, + default=None, + help="Path to JSONL DAG file to resolve", + ) + parser.add_argument( + "--profile", + action="store_true", + help="Print a cProfile summary for each benchmarked run", + ) + parser.add_argument( + "--profile-limit", + type=int, + default=20, + help="Number of cProfile rows to print per run", + ) + args = parser.parse_args() + P = args.partitions + N = args.events + + room_id = "!room:example.com" + real_version = RoomVersions.V2 + event_map: dict[str, Any] = {} + events_list: list[MockEvent] = [] + auth_chains: dict[str, set[str]] = {} + if args.jsonl: + if "v11" in args.jsonl: + real_version = getattr( + RoomVersions, "V11", RoomVersions.V10 + ) # Fallback to V10 if V11 not found? + elif "v6" in args.jsonl: + real_version = RoomVersions.V6 + else: + real_version = RoomVersions.V6 + room_version_rust = MockRoomVersion(real_version, real_version.state_res) + room_version_py = MockRoomVersion(real_version, real_version.state_res) + + # All MockEvent instances will share room_version_rust initially + MockEvent.room_version = room_version_rust + + if args.jsonl: + event_map, events_list = _load_jsonl_events(args.jsonl) + + # Sort events topologically by depth to simulate chronological ordering + events_list.sort(key=lambda e: (e.depth, e.origin_server_ts)) + + # Precompute auth chains for all events + for ev in events_list: + chain = {ev.event_id} + for aid in ev.auth_event_ids(): + if aid in auth_chains: + chain.update(auth_chains[aid]) + auth_chains[ev.event_id] = chain + + async def run_simulation( + room_version_to_use: MockRoomVersion, + ) -> tuple[RunStats, dict]: + # Set the room version for events + MockEvent.room_version = room_version_to_use + + # Initialize store + store = MockStateResolutionStore(event_map) + store.auth_chains = auth_chains # attach precomputed chains + + clock = MockClock() + if not events_list: + raise ValueError("JSONL DAG contains no events") + room_id = events_list[0].room_id + + # Map from event_id to the state *after* that event + event_states: dict[str, dict[tuple[str, str], str]] = {} + start_time = time.perf_counter() + bookkeeping_s = 0.0 + resolve_s = 0.0 + merge_points = 0 + try: + # Iterate through events and construct/resolve state + for ev in events_list: + loop_start = time.perf_counter() + prev_ids = ev.prev_event_ids() + + # Compute state before this event + if not prev_ids: + state_before: dict[tuple[str, str], str] = {} + elif len(prev_ids) == 1: + prev_id = prev_ids[0] + state_before = dict(event_states.get(prev_id, {})) + else: + # Merge point! We must resolve the states after the prev events + state_sets = [] + for pid in prev_ids: + if pid in event_states: + state_sets.append(event_states[pid]) + + if not state_sets: + state_before = {} + elif len(state_sets) == 1: + state_before = dict(state_sets[0]) + else: + merge_points += 1 + print(f"Resolving {len(state_sets)} states at {ev.event_id}") + resolve_start = time.perf_counter() + state_before = dict( + await v2.resolve_events_with_store( + cast(Any, clock), + room_id, + cast(RoomVersion, room_version_to_use), + state_sets, + None, + cast(Any, store), + ) + ) + resolve_s += time.perf_counter() - resolve_start + + # Compute state after this event + state_after = dict(state_before) + if ev.state_key is not None: + state_after[(ev.type, ev.state_key)] = ev.event_id + + event_states[ev.event_id] = state_after + bookkeeping_s += time.perf_counter() - loop_start + + duration = time.perf_counter() - start_time + # Get state of the last event + final_state = event_states[events_list[-1].event_id] + return ( + RunStats( + total_s=duration, + bookkeeping_s=bookkeeping_s - resolve_s, + resolve_s=resolve_s, + merge_points=merge_points, + ), + final_state, + ) + finally: + pass + + async def run_profiled_simulation( + room_version_to_use: MockRoomVersion, + title: str, + ) -> tuple[RunStats, dict]: + if not args.profile: + return await run_simulation(room_version_to_use) + + profiler = cProfile.Profile() + profiler.enable() + try: + result = await run_simulation(room_version_to_use) + finally: + profiler.disable() + _print_profile(profiler, title, args.profile_limit) + return result + + if args.jsonl: + print("Simulating resolution using Rust (rezzy)...") + stats_rust, res_rust = await run_profiled_simulation( + room_version_rust, "cProfile: Rust run" + ) + + print("Simulating resolution using Python fallback...") + try: + with _disable_rust_lattice_fold_resolver(): + stats_py, res_py = await run_profiled_simulation( + room_version_py, "cProfile: Python run" + ) + except Exception as e: + print( + "Python fallback benchmark failed for this DAG. " + f"The Rust run completed, but the Python path raised: {type(e).__name__}: {e}" + ) + return + + # Restore + MockEvent.room_version = room_version_rust + + assert res_rust == res_py, ( + "Error: Resolved states differ between Rust and Python!" + ) + + _print_results_table( + stats_py, + stats_rust, + _resolved_state_checksum(res_py), + _resolved_state_checksum(res_rust), + "Simulation Results:", + ) + _print_stage_breakdown(stats_py, stats_rust) + return + + if not args.jsonl: + # Baseline Events + bench_event_map: dict[str, EventBase] = {} + + # 1. CREATE + create = _make_benchmark_event( + room_version=real_version, + event_id="$CREATE", + sender="@alice:example.com", + event_type="m.room.create", + state_key="", + content={"creator": "@alice:example.com"}, + origin_server_ts=1000, + depth=1, + ) + bench_event_map[create.event_id] = create + + # 2. MEMBERS + alice_join = _make_benchmark_event( + room_version=real_version, + event_id="$IMA", + sender="@alice:example.com", + event_type="m.room.member", + state_key="@alice:example.com", + content={"membership": "join"}, + origin_server_ts=1001, + depth=2, + auth_event_ids=[create.event_id], + ) + bench_event_map[alice_join.event_id] = alice_join + + pl = _make_benchmark_event( + room_version=real_version, + event_id="$IPOWER", + sender="@alice:example.com", + event_type="m.room.power_levels", + state_key="", + content={"users": {"@alice:example.com": 100}, "users_default": 0}, + origin_server_ts=1002, + depth=3, + auth_event_ids=[create.event_id, alice_join.event_id], + ) + bench_event_map[pl.event_id] = pl + + # Join rules + jr = _make_benchmark_event( + room_version=real_version, + event_id="$IJR", + sender="@alice:example.com", + event_type="m.room.join_rules", + state_key="", + content={"join_rule": "public"}, + origin_server_ts=1003, + depth=4, + auth_event_ids=[create.event_id, alice_join.event_id, pl.event_id], + ) + bench_event_map[jr.event_id] = jr + + baseline_state = { + ("m.room.create", ""): create.event_id, + ("m.room.member", "@alice:example.com"): alice_join.event_id, + ("m.room.power_levels", ""): pl.event_id, + ("m.room.join_rules", ""): jr.event_id, + } + + # Generate P parallel partitions, each with N conflicting events + state_sets: list[dict[tuple[str, str], str]] = [] + + for p in range(P): + sender = f"@user_{p}:example.com" + # Join user to the room first + join_ev = _make_benchmark_event( + room_version=real_version, + event_id=f"$JOIN_{p}", + sender=sender, + event_type="m.room.member", + state_key=sender, + content={"membership": "join"}, + origin_server_ts=2000 + p, + depth=5 + p, + auth_event_ids=[create.event_id, jr.event_id, pl.event_id], + ) + bench_event_map[join_ev.event_id] = join_ev + + part_state = dict(baseline_state) + part_state[("m.room.member", sender)] = join_ev.event_id + + prev_id = join_ev.event_id + for i in range(N): + # Each event changes a topic or custom type to create conflicts + ev_id = f"$EV_{p}_{i}" + bench_ev: EventBase = _make_benchmark_event( + room_version=real_version, + event_id=ev_id, + sender=sender, + event_type=f"org.example.test_{i}", + state_key=f"state_key_{i}", + content={"value": f"val_{p}_{i}"}, + origin_server_ts=3000 + p * N + i, + depth=6 + p * N + i, + auth_event_ids=[create.event_id, join_ev.event_id, pl.event_id], + prev_event_ids=[prev_id], + ) + bench_event_map[ev_id] = bench_ev + + assert bench_ev.state_key is not None + part_state[(bench_ev.type, bench_ev.state_key)] = ev_id + prev_id = ev_id + + state_sets.append(part_state) + + clock = MockClock() + store = MockStateResolutionStore(bench_event_map) + + print("Benchmark Configuration:") + print(f" - Partitions: {P}") + print(f" - Conflicting events per partition: {N}") + print(f" - Total events in map: {len(bench_event_map)}") + print(" - Warm-up resolution...") + + async def run_resolution( + room_version_to_use: MockRoomVersion, + title: str, + ) -> tuple[RunStats, dict]: + MockEvent.room_version = room_version_to_use + + profiler = cProfile.Profile() if args.profile else None + if profiler is not None: + profiler.enable() + + try: + start = time.perf_counter() + resolved_state = dict( + await v2.resolve_events_with_store( + cast(Any, clock), + room_id, + cast(RoomVersion, room_version_to_use), + state_sets, + bench_event_map, + cast(Any, store), + ) + ) + duration = time.perf_counter() - start + finally: + if profiler is not None: + profiler.disable() + _print_profile(profiler, title, args.profile_limit) + + return ( + RunStats( + total_s=duration, + bookkeeping_s=0.0, + resolve_s=duration, + merge_points=1, + ), + resolved_state, + ) + + # Warmup + await v2.resolve_events_with_store( + cast(Any, clock), + room_id, + cast(RoomVersion, room_version_rust), + state_sets, + bench_event_map, + cast(Any, store), + ) + + # 1. Benchmark Rust (rezzy) + stats_rust, res_rust = await run_resolution( + room_version_rust, "cProfile: Rust run" + ) + + # 2. Benchmark Python (fallback) + # Disable only the Rust lattice-fold resolver while keeping the room + # version's state resolution algorithm unchanged. + with _disable_rust_lattice_fold_resolver(): + stats_py, res_py = await run_resolution( + room_version_py, "cProfile: Python run" + ) + + # Restore + MockEvent.room_version = room_version_rust + + assert res_rust == res_py, ( + "Error: Resolved states differ between Rust and Python!" + ) + + _print_results_table( + stats_py, + stats_rust, + _resolved_state_checksum(res_py), + _resolved_state_checksum(res_rust), + "Benchmark Results:", + ) + _print_stage_breakdown(stats_py, stats_rust) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/stubs/hiredis.pyi b/stubs/hiredis.pyi new file mode 100644 index 00000000000..64b2fa5f4bf --- /dev/null +++ b/stubs/hiredis.pyi @@ -0,0 +1,15 @@ +from __future__ import annotations + +from typing import Any + +class Reader: + def __init__( + self, + encoding: str | None = ..., + errors: str = ..., + notEnoughData: Any = ..., + protocolError: Any = ..., + replyError: Any = ..., + ) -> None: ... + def feed(self, data: bytes) -> None: ... + def gets(self) -> object: ... diff --git a/stubs/jaeger_client/__init__.pyi b/stubs/jaeger_client/__init__.pyi new file mode 100644 index 00000000000..029133f4242 --- /dev/null +++ b/stubs/jaeger_client/__init__.pyi @@ -0,0 +1,14 @@ +from __future__ import annotations + +from .config import Config, ConstSampler, Span, SpanContext, Tracer +from .reporter import BaseReporter, InMemoryReporter + +__all__ = [ + "BaseReporter", + "Config", + "ConstSampler", + "InMemoryReporter", + "Span", + "SpanContext", + "Tracer", +] diff --git a/stubs/jaeger_client/config.pyi b/stubs/jaeger_client/config.pyi new file mode 100644 index 00000000000..0feaf7aae9b --- /dev/null +++ b/stubs/jaeger_client/config.pyi @@ -0,0 +1,33 @@ +from __future__ import annotations + +from typing import Any + +from opentracing import Span as OpenTracingSpan, Tracer as OpenTracingTracer + +class SpanContext: ... + +class Span(OpenTracingSpan): + context: SpanContext + start_time: float | None + end_time: float | None + +class Tracer(OpenTracingTracer): + active_span: Span | None + +class Sampler: ... + +class ConstSampler(Sampler): + def __init__(self, decision: bool) -> None: ... + +class Config: + sampler: Any + + def __init__( + self, + config: Any, + service_name: str, + scope_manager: Any, + metrics_factory: Any = ..., + ) -> None: ... + def create_tracer(self, sampler: Any, reporter: Any = ...) -> Any: ... + def initialize_tracer(self, io_loop: Any = ...) -> Tracer | None: ... diff --git a/stubs/jaeger_client/metrics/__init__.pyi b/stubs/jaeger_client/metrics/__init__.pyi new file mode 100644 index 00000000000..9d48db4f9f8 --- /dev/null +++ b/stubs/jaeger_client/metrics/__init__.pyi @@ -0,0 +1 @@ +from __future__ import annotations diff --git a/stubs/jaeger_client/metrics/prometheus.pyi b/stubs/jaeger_client/metrics/prometheus.pyi new file mode 100644 index 00000000000..21846be11b6 --- /dev/null +++ b/stubs/jaeger_client/metrics/prometheus.pyi @@ -0,0 +1,4 @@ +from __future__ import annotations + +class PrometheusMetricsFactory: + def __init__(self) -> None: ... diff --git a/stubs/jaeger_client/reporter.pyi b/stubs/jaeger_client/reporter.pyi new file mode 100644 index 00000000000..e1eaa94553f --- /dev/null +++ b/stubs/jaeger_client/reporter.pyi @@ -0,0 +1,13 @@ +from __future__ import annotations + +from typing import Any + +from .config import Span + +class BaseReporter: + def set_process(self, service_name: str, tags: Any, max_length: int) -> None: ... + def report_span(self, span: Span) -> None: ... + +class InMemoryReporter(BaseReporter): + def __init__(self) -> None: ... + def get_spans(self) -> list[Span]: ... diff --git a/stubs/sentry_sdk.pyi b/stubs/sentry_sdk.pyi new file mode 100644 index 00000000000..f3bd1095d79 --- /dev/null +++ b/stubs/sentry_sdk.pyi @@ -0,0 +1,10 @@ +from __future__ import annotations + +from typing import Any + +def init(*args: Any, **kwargs: Any) -> None: ... + +class Scope: + @staticmethod + def get_global_scope() -> Scope: ... + def set_tag(self, key: str, value: Any) -> None: ... diff --git a/synapse/logging/opentracing.py b/synapse/logging/opentracing.py index 6e4e029163a..1789c302a03 100644 --- a/synapse/logging/opentracing.py +++ b/synapse/logging/opentracing.py @@ -253,13 +253,19 @@ class _DummyTagNames: except ImportError: opentracing = None # type: ignore[assignment] tags = _DummyTagNames # type: ignore[assignment] +JaegerConfig: Any = None +LogContextScopeManager: Any = None try: - from jaeger_client import Config as JaegerConfig + from jaeger_client import Config as _JaegerConfig - from synapse.logging.scopecontextmanager import LogContextScopeManager + from synapse.logging.scopecontextmanager import ( + LogContextScopeManager as _LogContextScopeManager, + ) except ImportError: - JaegerConfig = None # type: ignore - LogContextScopeManager = None # type: ignore + pass +else: + JaegerConfig = cast(Any, _JaegerConfig) + LogContextScopeManager = cast(Any, _LogContextScopeManager) try: diff --git a/synapse/state/v2.py b/synapse/state/v2.py index 1241a4d66e5..992c788b20f 100644 --- a/synapse/state/v2.py +++ b/synapse/state/v2.py @@ -30,6 +30,7 @@ Literal, Protocol, Sequence, + cast, overload, ) @@ -40,6 +41,7 @@ from synapse.events import EventBase, is_creator from synapse.storage.databases.main.event_federation import StateDifference from synapse.types import MutableStateMap, StateMap, StrCollection +from synapse.util.async_helpers import yieldable_gather_results from synapse.util.duration import Duration logger = logging.getLogger(__name__) @@ -141,11 +143,23 @@ async def resolve_events_with_store( ) ) - events = await state_res_store.get_events( - [eid for eid in full_conflicted_set if eid not in event_map], - allow_rejected=True, - ) - event_map.update(events) + # Prefetch the conflicted set and its auth chain concurrently so later + # auth checks can stay on the in-memory fast path. + to_fetch = {eid for eid in full_conflicted_set if eid not in event_map} + to_fetch.update(eid for eid in unconflicted_state.values() if eid not in event_map) + + while to_fetch: + fetched = await state_res_store.get_events(list(to_fetch), allow_rejected=True) + event_map.update(fetched) + + new_to_fetch = set() + for eid in to_fetch: + ev = event_map.get(eid) + if ev: + for aid in ev.auth_event_ids(): + if aid not in event_map: + new_to_fetch.add(aid) + to_fetch = new_to_fetch # everything in the event map should be in the right room for event in event_map.values(): @@ -159,6 +173,27 @@ async def resolve_events_with_store( ) ) + # Attempt to run high-performance state resolution in Rust via rezzy's lattice fold + if room_version.state_res == StateResolutionVersions.V2: + try: + import synapse.synapse_rust.state_res as rust_res + + resolve_v2_via_lattice_fold = rust_res.resolve_v2_via_lattice_fold + + logger.debug("Resolving state v2 via Rust rezzy lattice fold") + + resolved_state_rust: StateMap[str] = resolve_v2_via_lattice_fold( + dict(unconflicted_state), + list(full_conflicted_set), + event_map, + ) + return resolved_state_rust + except Exception as e: + logger.exception( + "Failed to run Rust state resolution via lattice fold, falling back to python", + exc_info=e, + ) + full_conflicted_set = {eid for eid in full_conflicted_set if eid in event_map} logger.debug("%d full_conflicted_set entries", len(full_conflicted_set)) @@ -335,6 +370,22 @@ async def _get_auth_chain_difference( len(conflicted_state) if conflicted_state is not None else None ) + if not is_state_res_v21: + try: + import synapse.synapse_rust.state_res as rust_res + + return cast( + set[str], + cast(Any, rust_res).get_auth_chain_difference_from_event_graph( + state_sets, + unpersisted_events, + ), + ) + except Exception: + # Fall back to the Python/store path if the Rust fast path cannot + # handle the graph we were given. + pass + # The `StateResolutionStore.get_auth_chain_difference` function assumes that # all events passed to it (and their auth chains) have been persisted # previously. We need to manually handle any other events that are yet to be @@ -348,27 +399,62 @@ async def _get_auth_chain_difference( # the set of persisted events belonging to the auth difference. # 3. Adding the results of 1 and 2 together. - # Map from event ID in `unpersisted_events` to their auth event IDs, and their auth - # event IDs if they appear in the `unpersisted_events`. This is the intersection of - # the event's auth chain with the events in `unpersisted_events` *plus* their - # auth event IDs. + # Map from event ID in `unpersisted_events` to the auth chain reachable from it. + # + # We memoize the transitive closure so shared auth chains are only traversed once. events_to_auth_chain: dict[str, set[str]] = {} - # remember the forward links when doing the graph traversal, we'll need it for v2.1 checks - # This is a map from an event to the set of events that contain it as an auth event. + + # Forward links are needed later for v2.1 conflicted-subgraph expansion. + # This maps an event to the set of events that reference it directly in auth_events. event_to_next_event: dict[str, set[str]] = {} + direct_auth_events: dict[str, list[str]] = {} for event in unpersisted_events.values(): - chain = {event.event_id} - events_to_auth_chain[event.event_id] = chain - - to_search = [event] - while to_search: - next_event = to_search.pop() - for auth_id in next_event.auth_event_ids(): + auth_ids = list(event.auth_event_ids()) + direct_auth_events[event.event_id] = auth_ids + for auth_id in auth_ids: + event_to_next_event.setdefault(auth_id, set()).add(event.event_id) + + def _build_auth_chain(root_event_id: str) -> set[str]: + """Return the full auth closure for an unpersisted event.""" + + cached = events_to_auth_chain.get(root_event_id) + if cached is not None: + return cached + + postorder: list[str] = [] + stack: list[tuple[str, bool]] = [(root_event_id, False)] + + while stack: + event_id, expanded = stack.pop() + + if expanded: + postorder.append(event_id) + continue + + if event_id in events_to_auth_chain: + continue + + stack.append((event_id, True)) + for auth_id in direct_auth_events.get(event_id, []): + if ( + auth_id in unpersisted_events + and auth_id not in events_to_auth_chain + ): + stack.append((auth_id, False)) + + for event_id in postorder: + chain = {event_id} + for auth_id in direct_auth_events.get(event_id, []): chain.add(auth_id) - event_to_next_event.setdefault(auth_id, set()).add(next_event.event_id) - auth_event = unpersisted_events.get(auth_id) - if auth_event: - to_search.append(auth_event) + auth_chain = events_to_auth_chain.get(auth_id) + if auth_chain is not None: + chain.update(auth_chain) + events_to_auth_chain[event_id] = chain + + return events_to_auth_chain[root_event_id] + + for event_id in unpersisted_events: + _build_auth_chain(event_id) # We now 1) calculate the auth chain difference for the unpersisted events # and 2) work out the state sets to pass to the store. @@ -698,9 +784,11 @@ async def _iterative_auth_checks( auth_events = {} for aid in event.auth_event_ids(): - ev = await _get_event( - room_id, aid, event_map, state_res_store, allow_none=True - ) + ev = event_map.get(aid) + if not ev: + ev = await _get_event( + room_id, aid, event_map, state_res_store, allow_none=True + ) if not ev: logger.warning( @@ -713,7 +801,9 @@ async def _iterative_auth_checks( for key in event_auth.auth_types_for_event(room_version, event): if key in resolved_state: ev_id = resolved_state[key] - ev = await _get_event(room_id, ev_id, event_map, state_res_store) + ev = event_map.get(ev_id) + if not ev: + ev = await _get_event(room_id, ev_id, event_map, state_res_store) if ev.rejected_reason is None: auth_events[key] = event_map[ev_id] @@ -806,16 +896,14 @@ async def _mainline_sort( event_ids = list(event_ids) order_map = {} - for idx, ev_id in enumerate(event_ids, start=1): + + async def get_depth(ev_id: str) -> None: depth = await _get_mainline_depth_for_event( clock, event_map[ev_id], mainline_map, event_map, state_res_store ) order_map[ev_id] = (depth, event_map[ev_id].origin_server_ts, ev_id) - # We await occasionally when we're working with large data sets to - # ensure that we don't block the reactor loop for too long. - if idx % _AWAIT_AFTER_ITERATIONS == 0: - await clock.sleep(Duration(seconds=0)) + await yieldable_gather_results(get_depth, event_ids) event_ids.sort(key=lambda ev_id: order_map[ev_id]) diff --git a/synapse/synapse_rust/state_res.pyi b/synapse/synapse_rust/state_res.pyi new file mode 100644 index 00000000000..c47c661c6e3 --- /dev/null +++ b/synapse/synapse_rust/state_res.pyi @@ -0,0 +1,24 @@ +# This file is licensed under the Affero General Public License (AGPL) version 3. +# +# Copyright (C) 2026 Element Creations Ltd. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as +# published by the Free Software Foundation, either version 3 of the +# License, or (at your option) any later version. +# +# See the GNU Affero General Public License for more details: +# . + +from collections.abc import Iterable +from typing import Any + +def get_auth_chain_difference_from_event_graph( + state_sets: Iterable[Any], + event_map: dict[str, Any], +) -> set[str]: ... +def resolve_v2_via_lattice_fold( + unconflicted_state: dict[tuple[str, str], str], + conflicted_event_ids: Iterable[str], + event_map: dict[str, Any], +) -> dict[tuple[str, str], str]: ... diff --git a/test_profile.rs b/test_profile.rs new file mode 100644 index 00000000000..483ef9515d4 --- /dev/null +++ b/test_profile.rs @@ -0,0 +1,2 @@ +use std::time::Instant; +// just testing how to profile... diff --git a/tests/config/test_load.py b/tests/config/test_load.py index 8d94390acf9..51ecc226d1e 100644 --- a/tests/config/test_load.py +++ b/tests/config/test_load.py @@ -20,7 +20,7 @@ # # import tempfile -from typing import Callable +from typing import Any, Callable, cast from unittest import mock import yaml @@ -37,10 +37,13 @@ except ImportError: authlib = None +hiredis: Any = None try: - import hiredis + import hiredis as _hiredis except ImportError: - hiredis = None # type: ignore + pass +else: + hiredis = cast(Any, _hiredis) class ConfigLoadingFileTestCase(ConfigFileTestCase): diff --git a/tests/logging/test_opentracing.py b/tests/logging/test_opentracing.py index d5e643585d5..8dc5d681bcf 100644 --- a/tests/logging/test_opentracing.py +++ b/tests/logging/test_opentracing.py @@ -19,7 +19,8 @@ # # -from typing import Awaitable, cast +import logging +from typing import Any, Awaitable, cast from twisted.internet import defer from twisted.internet.testing import MemoryReactorClock @@ -40,24 +41,29 @@ from synapse.util.duration import Duration from tests.server import get_clock +from tests.unittest import TestCase +jaeger_client: Any = None try: - import jaeger_client + import jaeger_client as _jaeger_client except ImportError: - jaeger_client = None # type: ignore - + pass +else: + jaeger_client = cast(Any, _jaeger_client) +opentracing: Any = None +LogContextScopeManager: Any = None try: - import opentracing + import opentracing as _opentracing - from synapse.logging.scopecontextmanager import LogContextScopeManager + from synapse.logging.scopecontextmanager import ( + LogContextScopeManager as _LogContextScopeManager, + ) except ImportError: - opentracing = None # type: ignore - LogContextScopeManager = None # type: ignore - -import logging - -from tests.unittest import TestCase + pass +else: + opentracing = cast(Any, _opentracing) + LogContextScopeManager = cast(Any, _LogContextScopeManager) logger = logging.getLogger(__name__) @@ -74,9 +80,9 @@ class LogContextScopeManagerTestCase(TestCase): """ if opentracing is None or LogContextScopeManager is None: - skip = "Requires opentracing" # type: ignore[unreachable] + skip = "Requires opentracing" if jaeger_client is None: - skip = "Requires jaeger_client" # type: ignore[unreachable] + skip = "Requires jaeger_client" def setUp(self) -> None: # since this is a unit test, we don't really want to mess around with the @@ -164,15 +170,8 @@ def test_nested_spans(self) -> None: def test_overlapping_spans(self) -> None: """Overlapping spans which are not neatly nested should work""" reactor = MemoryReactorClock() - # type-ignore: mypy-zope doesn't seem to recognise that `MemoryReactorClock` - # implements `ISynapseThreadlessReactor` (combination of the normal Twisted - # Reactor/Clock interfaces), via inheritance from - # `twisted.internet.testing.MemoryReactor` and `twisted.internet.testing.Clock` - # Ignore `multiple-internal-clocks` linter error here since we are creating a `Clock` - # for testing purposes. clock = Clock( # type: ignore[multiple-internal-clocks] - reactor, # type: ignore[arg-type] - server_name="test_server", + cast(Any, reactor), server_name="test_server" ) scopes = [] @@ -342,7 +341,6 @@ async def bg_task() -> None: # so that the test can complete and we see the underlying error. callback_finished = True - # type-ignore: We ignore because the point is to test the bare function run_as_background_process( # type: ignore[untracked-background-process] desc="some-bg-task", server_name="test_server", @@ -409,7 +407,6 @@ async def bg_task() -> None: "some-request", tracer=self._tracer, ): - # type-ignore: We ignore because the point is to test the bare function run_as_background_process( # type: ignore[untracked-background-process] desc="some-bg-task", server_name="test_server", diff --git a/tests/replication/_base.py b/tests/replication/_base.py index b23696668f3..9bb4c912ecf 100644 --- a/tests/replication/_base.py +++ b/tests/replication/_base.py @@ -19,7 +19,7 @@ # import logging from collections import defaultdict -from typing import Any +from typing import Any, cast from twisted.internet.address import IPv4Address from twisted.internet.protocol import Protocol, connectionDone @@ -44,10 +44,13 @@ from tests.server import FakeTransport from tests.utils import USE_POSTGRES_FOR_TESTS +hiredis: Any = None try: - import hiredis + import hiredis as _hiredis except ImportError: - hiredis = None # type: ignore + pass +else: + hiredis = cast(Any, _hiredis) logger = logging.getLogger(__name__)