mirror of
https://github.com/element-hq/synapse.git
synced 2026-09-01 20:18:19 +00:00
Hand-roll some randomised DAGs, add more edge case tests
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user