From 06c8deae629ebf0a0f1799d94ab01e69c3e1a566 Mon Sep 17 00:00:00 2001 From: Kegan Dougal <7190048+kegsay@users.noreply.github.com> Date: Wed, 19 Aug 2026 09:07:45 +0100 Subject: [PATCH] Hand-roll some randomised DAGs, add more edge case tests --- tests/storage/test_event_federation.py | 232 +++++++++++++++++++++---- 1 file changed, 203 insertions(+), 29 deletions(-) diff --git a/tests/storage/test_event_federation.py b/tests/storage/test_event_federation.py index 976c7b1b49..c785235202 100644 --- a/tests/storage/test_event_federation.py +++ b/tests/storage/test_event_federation.py @@ -19,6 +19,7 @@ # import datetime +import random from typing import ( Collection, Iterable, @@ -1519,7 +1520,15 @@ class EventFederationGetMissingEventsStateDAGTestCase( "T": ["R"], "E": ["W", "D", "T"], } - creator = "@test_get_missing_events_state_dag:localhost" + self.graph = graph + (self.room_id, self.graph_events) = self._persist_state_dag( + "@test_get_missing_events_state_dag:localhost", graph + ) + + def _persist_state_dag( + self, creator: str, graph: dict[str, list[str]] + ) -> tuple[str, dict[str, FrozenEventVMSC4242]]: + """Build and persist a state DAG in its own room, as `build_state_dag` returns it.""" (room_id, graph_events) = build_state_dag(creator, graph) def insert(txn: LoggingTransaction) -> None: @@ -1546,8 +1555,32 @@ class EventFederationGetMissingEventsStateDAGTestCase( insert, ) ) - self.room_id = room_id - self.graph_events = graph_events + return room_id, graph_events + + def _get_missing_events( + self, + latest: list[str], + earliest: list[str], + limit: int, + room_id: str | None = None, + graph_events: dict[str, FrozenEventVMSC4242] | None = None, + ) -> list[str]: + """Run `get_missing_events_state_dag` in terms of fake event IDs, returning the fake + event IDs which came back. Queries the room built in `prepare` unless told otherwise. + """ + if room_id is None or graph_events is None: + room_id = self.room_id + graph_events = self.graph_events + fake_event_ids = {ev.event_id: fake for fake, ev in graph_events.items()} + got = self.get_success( + self.store.get_missing_events_state_dag( + room_id=room_id, + earliest_event_ids=[graph_events[fake].event_id for fake in earliest], + latest_event_ids=[graph_events[fake].event_id for fake in latest], + limit=limit, + ), + ) + return [fake_event_ids[ev.event_id] for ev in got] @parameterized.expand( [ @@ -1578,22 +1611,133 @@ class EventFederationGetMissingEventsStateDAGTestCase( # A <- B E # `- R -- W --` # `-- T -` - got = self.get_success( - self.store.get_missing_events_state_dag( - room_id=self.room_id, - earliest_event_ids=[], - latest_event_ids=[ - self.graph_events[graph_event_id].event_id - for graph_event_id in latest - ], - limit=limit, - ), - ) self.assertEquals( - [ev.event_id for ev in got], - [self.graph_events[graph_event_id].event_id for graph_event_id in want], + self._get_missing_events(latest=latest, earliest=[], limit=limit), + want, f"latest={latest} want={want} limit={limit}", ) + # These expectations are written by hand, so use them to check `walk_state_dag`, which + # the randomly generated graphs below are compared against. + self.assertEquals( + walk_state_dag(self.graph, self.graph_events, latest, [], limit), + want, + f"latest={latest} want={want} limit={limit}", + ) + + @tests.unittest.override_config( + {"experimental_features": {"msc4242_enabled": True}} + ) + def test_get_missing_events_state_dag_limit_truncates(self) -> None: + """Lowering `limit` must truncate the result rather than change it. + + `limit` also bounds how many hops the walk takes, so this checks the two uses agree: + as every extra hop contributes at least one event, a shallower walk cannot miss an + event that a `limit`-sized response should have contained. + """ + full = self._get_missing_events( + latest=["E"], earliest=[], limit=len(self.graph_events) + ) + self.assertEquals(full, ["D", "T", "W", "C", "R", "B", "A"]) + for limit in range(1, len(full) + 1): + self.assertEquals( + self._get_missing_events(latest=["E"], earliest=[], limit=limit), + full[:limit], + f"limit={limit}", + ) + + @tests.unittest.override_config( + {"experimental_features": {"msc4242_enabled": True}} + ) + def test_get_missing_events_state_dag_returns_nothing(self) -> None: + """The cases where there is nothing to walk back to.""" + # Nothing was asked for. + self.assertEquals( + self._get_missing_events(latest=["E"], earliest=[], limit=0), [] + ) + # Every event in the room has already been seen. + self.assertEquals( + self._get_missing_events( + latest=["E"], earliest=list(self.graph), limit=100 + ), + [], + ) + # Every event walked back from has already been seen. + self.assertEquals( + self._get_missing_events(latest=["E", "C"], earliest=["C", "E"], limit=100), + [], + ) + + @tests.unittest.override_config( + {"experimental_features": {"msc4242_enabled": True}} + ) + def test_get_missing_events_state_dag_random_graphs(self) -> None: + """Check the query against a direct implementation of the MSC4242 ordering. + + Uses randomly generated DAGs, as the hand-written cases above can only cover the + shapes we thought to write down. The seed is fixed so that a failure is reproducible + and cannot appear on an unrelated change. + """ + rand = random.Random(42) + # Single letters keep `build_state_dag`'s hash mining cheap. At least 4 events, as + # smaller graphs have too little to order for the walk to be interesting. + names = "ABCDEFGH" + for room_number in range(20): + # Each event picks its prev_state_events from the events before it, which keeps + # the graph acyclic and in causal order. "A" is the create event. + graph: dict[str, list[str]] = {names[0]: []} + for index in range(1, rand.randint(4, len(names))): + graph[names[index]] = sorted( + rand.sample(names[:index], rand.randint(1, index)) + ) + built = list(graph) + (room_id, graph_events) = self._persist_state_dag( + f"@test_random_state_dag_{room_number}:localhost", graph + ) + + for _ in range(3): + # `latest` is sampled with replacement, as callers can repeat an event ID. + latest = [ + rand.choice(built) for _ in range(rand.randint(1, len(built))) + ] + earliest = rand.sample(built, rand.randint(0, len(built) // 2)) + limit = rand.randint(0, len(built) + 1) + message = ( + f"graph={graph} latest={latest} earliest={earliest} limit={limit}" + ) + self.assertEquals( + self._get_missing_events( + latest=latest, + earliest=earliest, + limit=limit, + room_id=room_id, + graph_events=graph_events, + ), + walk_state_dag(graph, graph_events, latest, earliest, limit), + message, + ) + + # Lowering `limit` must truncate the result rather than change it. Take the + # limit from the events actually available rather than from the random one, + # which is often high enough that nothing is dropped. + full = self._get_missing_events( + latest=latest, + earliest=earliest, + limit=len(built), + room_id=room_id, + graph_events=graph_events, + ) + if len(full) > 1: + self.assertEquals( + self._get_missing_events( + latest=latest, + earliest=earliest, + limit=len(full) - 1, + room_id=room_id, + graph_events=graph_events, + ), + full[:-1], + message, + ) @tests.unittest.override_config( {"experimental_features": {"msc4242_enabled": True}} @@ -1619,22 +1763,11 @@ class EventFederationGetMissingEventsStateDAGTestCase( ) ) - got = self.get_success( - self.store.get_missing_events_state_dag( - room_id=self.room_id, - earliest_event_ids=[], - latest_event_ids=[self.graph_events["E"].event_id], - limit=100, - ), - ) # Same as the acyclic walk from E, with E itself now reachable via A at the end. Each # event appears exactly once: the CTE groups by event ID and keeps the fewest hops. self.assertEquals( - [ev.event_id for ev in got], - [ - self.graph_events[graph_event_id].event_id - for graph_event_id in ["D", "T", "W", "C", "R", "B", "A", "E"] - ], + self._get_missing_events(latest=["E"], earliest=[], limit=100), + ["D", "T", "W", "C", "R", "B", "A", "E"], ) @@ -1656,6 +1789,45 @@ class FakeEvent: return True +def walk_state_dag( + graph: dict[str, list[str]], + graph_events: dict[str, FrozenEventVMSC4242], + latest: list[str], + earliest: list[str], + limit: int, +) -> list[str]: + """Work out what `get_missing_events_state_dag` should return, from the MSC4242 rules. + + Walks back from `latest` via prev_state_events, then orders the events found by how many + hops they are from `latest`, breaking ties on the real event ID. An `earliest` event is + never returned and is never walked through, though its predecessors are still returned if + some other path reaches them. + + Takes the graph and events as `build_state_dag` returns them, in terms of fake event IDs, + and returns the fake event IDs which should come back, in order. Ties break on the real + event ID rather than the fake one because the create event is the one event whose real ID + does not start with its fake ID. + """ + hops: dict[str, int] = {} + # An event `limit` hops away can only be reached by returning an event at every hop + # before it, so a walk deeper than `limit` cannot contribute to the response. + frontier = sorted(set(latest) - set(earliest)) + for hop in range(1, limit + 1): + next_frontier = set() + for fake_event_id in frontier: + for prev_fake_event_id in graph[fake_event_id]: + if prev_fake_event_id in earliest or prev_fake_event_id in hops: + continue + # The first time we see an event is by definition its fewest hops. + hops[prev_fake_event_id] = hop + next_frontier.add(prev_fake_event_id) + frontier = sorted(next_frontier) + return sorted( + hops, + key=lambda fake: (hops[fake], graph_events[fake].event_id), + )[:limit] + + def build_state_dag( creator: str, graph: dict[str, list[str]] ) -> tuple[str, dict[str, FrozenEventVMSC4242]]: @@ -1670,6 +1842,8 @@ def build_state_dag( A tuple of the room ID and a map from fake event ID e.g. "B" to real event which you can use .event_id to extract the real event ID. Guarantees that the real event IDs start with the fake event ID e.g. the real event for "B" is guarantees to start "$B...." which makes sorting tests much easier to reason about. + The create event is the exception: its ID is fixed by the room ID, so it does not start with its + fake event ID. """ graph_events: dict[str, FrozenEventVMSC4242] = {} # graph ID => built event create_event = FrozenEventVMSC4242(