diff --git a/synapse/storage/controllers/persist_events.py b/synapse/storage/controllers/persist_events.py index 1ad124421a..3601dc843c 100644 --- a/synapse/storage/controllers/persist_events.py +++ b/synapse/storage/controllers/persist_events.py @@ -617,7 +617,7 @@ class EventsPersistenceStorageController: if not events_and_contexts: return replaced_events - await assign_stitched_orders(room_id, events_and_contexts, self.main_store) + await assign_stitched_orders(room_id, [ev for (ev, _) in events_and_contexts], self.main_store) chunks = [ events_and_contexts[x : x + 100] @@ -1246,9 +1246,7 @@ class EventsPersistenceStorageController: return False -def find_predecessors( - event_ids: Iterable[str], batch: List[EventPersistencePair] -) -> Set[str]: +def find_predecessors(event_ids: Iterable[str], batch: List[EventBase]) -> Set[str]: """ Walk the tree of dependencies (in batch), and return every event that is in batch, and is an ancestor of one of the supplied events. @@ -1261,7 +1259,7 @@ def find_predecessors( # Iterate through the incoming events, looking for events in our "unexplored" # set. For each matching event, add it to the "found" set, and add its # "prev_events" to the "unexplored" set for the next pass. - for event, _ in batch: + for event in batch: if event.event_id in unexplored: found.add(event.event_id) next_unexplored.update( @@ -1279,20 +1277,16 @@ def find_predecessors( async def assign_stitched_orders( room_id: str, - events_and_contexts: List[EventPersistencePair], + events: List[EventBase], store: DataStore, ) -> None: """ - Updates the EventContexts within `events_and_contexts`, to assign a + Updates the events within `events`, to assign a stitched_ordering to each event. """ # Take a copy of the events we have to process # TODO find a better way to exclude outliers - remaining_batch = list( - (ev, ctx) - for ev, ctx in events_and_contexts - if not ev.internal_metadata.is_outlier() - ) + remaining_batch = list(ev for ev in events if not ev.internal_metadata.is_outlier()) # Find all events in the current batch which are in a timeline gap gap_events = await store.db_pool.simple_select_many_batch( @@ -1339,9 +1333,9 @@ async def assign_stitched_orders( return still_remaining_batch = [] - for event, context in remaining_batch: + for event in remaining_batch: if event.event_id not in to_insert: - still_remaining_batch.append((event, context)) + still_remaining_batch.append(event) continue # TODO we may need to reorder existing events @@ -1354,18 +1348,18 @@ async def assign_stitched_orders( remaining_batch = still_remaining_batch logger.debug( - "Remaining events: %s", [ev.event_id for (ev, _) in remaining_batch] + "Remaining events: %s", [ev.event_id for ev in remaining_batch] ) logger.debug( "Remaining events after processing gap matches: %s", - [ev.event_id for (ev, _) in remaining_batch], + [ev.event_id for ev in remaining_batch], ) current_max_stream_ordering = ( await store.get_room_max_stitched_ordering(room_id) or 0 ) - for event, _ in remaining_batch: + for event in remaining_batch: current_max_stream_ordering += 2**16 event.assign_stitched_ordering(current_max_stream_ordering) diff --git a/tests/storage/controllers/test_persist_events.py b/tests/storage/controllers/test_persist_events.py index ecc1f66f42..c3ae09b355 100644 --- a/tests/storage/controllers/test_persist_events.py +++ b/tests/storage/controllers/test_persist_events.py @@ -22,7 +22,7 @@ from synapse.storage.controllers.persist_events import find_predecessors class FindPredecessorsTestCase(TestCase): def test_predecessors_finds_nothing_if_event_is_not_in_batch(self) -> None: batch = [ - (FakeEvent(event_id="B", prev_event_ids=["C"]), None), + FakeEvent(event_id="B", prev_event_ids=["C"]), ] predecessors = find_predecessors({"A"}, batch) # type: ignore[arg-type] @@ -30,8 +30,8 @@ class FindPredecessorsTestCase(TestCase): def test_predecessors_finds_only_event_if_it_has_no_predecessors(self) -> None: batch = [ - (FakeEvent(event_id="E1", prev_event_ids=[]), None), - (FakeEvent(event_id="E2", prev_event_ids=["E3"]), None), + FakeEvent(event_id="E1", prev_event_ids=[]), + FakeEvent(event_id="E2", prev_event_ids=["E3"]), ] predecessors = find_predecessors({"E1"}, batch) # type: ignore[arg-type] @@ -39,21 +39,21 @@ class FindPredecessorsTestCase(TestCase): def test_predecessors_finds_all_ancestors(self) -> None: batch = [ - (FakeEvent(event_id="A", prev_event_ids=["B", "C"]), None), - (FakeEvent(event_id="B", prev_event_ids=["D"]), None), - (FakeEvent(event_id="C", prev_event_ids=["D"]), None), - (FakeEvent(event_id="D", prev_event_ids=["E"]), None), - (FakeEvent(event_id="E", prev_event_ids=[]), None), - (FakeEvent(event_id="F", prev_event_ids=["G", "H"]), None), - (FakeEvent(event_id="G", prev_event_ids=[]), None), + FakeEvent(event_id="A", prev_event_ids=["B", "C"]), + FakeEvent(event_id="B", prev_event_ids=["D"]), + FakeEvent(event_id="C", prev_event_ids=["D"]), + FakeEvent(event_id="D", prev_event_ids=["E"]), + FakeEvent(event_id="E", prev_event_ids=[]), + FakeEvent(event_id="F", prev_event_ids=["G", "H"]), + FakeEvent(event_id="G", prev_event_ids=[]), ] predecessors = find_predecessors({"A"}, batch) # type: ignore[arg-type] self.assertEqual(predecessors, {"A", "B", "C", "D", "E"}) def test_predecessors_ignores_cycles(self) -> None: batch = [ - (FakeEvent(event_id="E1", prev_event_ids=["E2"]), None), - (FakeEvent(event_id="E2", prev_event_ids=["E1"]), None), + FakeEvent(event_id="E1", prev_event_ids=["E2"]), + FakeEvent(event_id="E2", prev_event_ids=["E1"]), ] predecessors = find_predecessors({"E1"}, batch) # type: ignore[arg-type] @@ -61,8 +61,8 @@ class FindPredecessorsTestCase(TestCase): def test_predecessors_ignores_self_reference_cycles(self) -> None: batch = [ - (FakeEvent(event_id="E1", prev_event_ids=["E2"]), None), - (FakeEvent(event_id="E2", prev_event_ids=["E2"]), None), + FakeEvent(event_id="E1", prev_event_ids=["E2"]), + FakeEvent(event_id="E2", prev_event_ids=["E2"]), ] predecessors = find_predecessors({"E1"}, batch) # type: ignore[arg-type] @@ -70,13 +70,13 @@ class FindPredecessorsTestCase(TestCase): def test_predecessors_finds_ancestors_of_multiple_starting_events(self) -> None: batch = [ - (FakeEvent(event_id="A", prev_event_ids=["B"]), None), - (FakeEvent(event_id="B", prev_event_ids=[]), None), - (FakeEvent(event_id="C", prev_event_ids=["D"]), None), - (FakeEvent(event_id="D", prev_event_ids=["E"]), None), - (FakeEvent(event_id="E", prev_event_ids=[]), None), - (FakeEvent(event_id="F", prev_event_ids=["G"]), None), - (FakeEvent(event_id="G", prev_event_ids=[]), None), + FakeEvent(event_id="A", prev_event_ids=["B"]), + FakeEvent(event_id="B", prev_event_ids=[]), + FakeEvent(event_id="C", prev_event_ids=["D"]), + FakeEvent(event_id="D", prev_event_ids=["E"]), + FakeEvent(event_id="E", prev_event_ids=[]), + FakeEvent(event_id="F", prev_event_ids=["G"]), + FakeEvent(event_id="G", prev_event_ids=[]), ] predecessors = find_predecessors({"A", "C"}, batch) # type: ignore[arg-type] self.assertEqual(predecessors, {"A", "B", "C", "D", "E"}) diff --git a/tests/storage/test_events.py b/tests/storage/test_events.py index 2a4e1bdc2c..084a0e1317 100644 --- a/tests/storage/test_events.py +++ b/tests/storage/test_events.py @@ -561,9 +561,8 @@ class AssignStitchedOrderingTestCase(HomeserverTestCase): RoomVersions.V12, ) - context = self.get_success(self.state.compute_event_context(test_event)) self.get_success( - assign_stitched_orders(room_id, [(test_event, context)], self.store) + assign_stitched_orders(room_id, [test_event], self.store) ) self.assertEqual(test_event.stitched_ordering, 6 * 2**16)