diff --git a/changelog.d/20108.bugfix b/changelog.d/20108.bugfix new file mode 100644 index 0000000000..6a842b22ea --- /dev/null +++ b/changelog.d/20108.bugfix @@ -0,0 +1 @@ +Fix a bug where read receipts could be permanently skipped and never sent to application services if more than 100 read receipts arrived in a single stream update. diff --git a/synapse/handlers/appservice.py b/synapse/handlers/appservice.py index 68b8aa71f1..8d82a98559 100644 --- a/synapse/handlers/appservice.py +++ b/synapse/handlers/appservice.py @@ -374,13 +374,21 @@ class ApplicationServicesHandler: # follow the base stream position. new_token = MultiWriterStreamToken(stream=new_token.stream) - events = await self._handle_receipts(service, new_token) - self.scheduler.enqueue_for_appservice(service, ephemeral=events) + while True: + events, reached_token = await self._handle_receipts( + service, new_token + ) + self.scheduler.enqueue_for_appservice( + service, ephemeral=events + ) - # Persist the latest handled stream token for this appservice - await self.store.set_appservice_stream_type_pos( - service, "read_receipt", new_token.stream - ) + # Persist the latest handled stream token for this appservice + await self.store.set_appservice_stream_type_pos( + service, "read_receipt", reached_token.stream + ) + + if reached_token.stream >= new_token.stream: + break elif stream_key == StreamKeyType.PRESENCE: assert isinstance(new_token, int) @@ -460,14 +468,16 @@ class ApplicationServicesHandler: async def _handle_receipts( self, service: ApplicationService, new_token: MultiWriterStreamToken - ) -> list[JsonMapping]: + ) -> tuple[list[JsonMapping], MultiWriterStreamToken]: """ - Return the latest read receipts that the given application service should receive. + Return the next batch of read receipts that the given application service + should receive. - First fetch all read receipts between the last receipt stream token that this - application service should have previously received (non-inclusive) and the - latest read receipt stream token (inclusive). Then from that set, return only - those read receipts that the given application service may be interested in. + First fetch the oldest read receipts between the last receipt stream token that + this application service should have previously received (non-inclusive) and + the latest read receipt stream token (inclusive), up to a cap. Then from that + set, return only those read receipts that the given application service may be + interested in. Args: service: The application service to check for which events it should receive. @@ -476,23 +486,40 @@ class ApplicationServicesHandler: token. Prevents accidentally duplicating work. Returns: - A list of JSON dictionaries containing data derived from the read receipts that - should be sent to the given application service. + A two-tuple containing the following: + * A list of JSON dictionaries containing data derived from the read + receipts that should be sent to the given application service. + * The receipt stream token up to which receipts were actually handled. + This is earlier than `new_token` if the fetch was truncated; callers + must call this method again to fetch the remaining receipts. """ from_key = await self.store.get_type_stream_id_for_appservice( service, "read_receipt" ) - if new_token is not None and new_token.stream <= from_key: - logger.debug("Rejecting token lower than or equal to stored: %s", new_token) - return [] + if new_token.stream <= from_key: + if new_token.stream == from_key: + logger.debug( + "Receipt token %s already handled for appservice %s", + new_token, + service.id, + ) + else: + logger.warning( + "Receipt token %s is behind stored position %s for appservice %s", + new_token, + from_key, + service.id, + ) + + return [], MultiWriterStreamToken(stream=from_key) from_token = MultiWriterStreamToken(stream=from_key) receipts_source = self.event_sources.sources.receipt - receipts, _ = await receipts_source.get_new_events_as( + receipts, reached_token = await receipts_source.get_new_events_as( service=service, from_key=from_token, to_key=new_token ) - return receipts + return receipts, reached_token async def _handle_presence( self, diff --git a/synapse/handlers/receipts.py b/synapse/handlers/receipts.py index c208b4e91a..8d41582c58 100644 --- a/synapse/handlers/receipts.py +++ b/synapse/handlers/receipts.py @@ -21,7 +21,7 @@ import logging from typing import TYPE_CHECKING, Callable, Iterable, Sequence -from synapse.api.constants import EduTypes, ReceiptTypes +from synapse.api.constants import Direction, EduTypes, ReceiptTypes from synapse.appservice import ApplicationService from synapse.streams import EventSource from synapse.types import ( @@ -355,15 +355,35 @@ class ReceiptEventSource(EventSource[MultiWriterStreamToken, JsonMapping]): A two-tuple containing the following: * A list of json dictionaries derived from read receipts that the appservice may be interested in. - * The current read receipt stream token. + * The read receipt stream token up to which receipts were actually + fetched. This is earlier than `to_key` if the fetch was + truncated; callers must call this method again from the returned + token to fetch the remaining receipts. """ if from_key == to_key: return [], to_key - # Fetch all read receipts for all rooms, up to a limit of 100. This is ordered - # by most recent. - rooms_to_events = await self.store.get_linearized_receipts_for_all_rooms( - from_key=from_key, to_key=to_key + # A stored stream position of 1 means no read receipts have ever been + # delivered to this application service: + # `get_type_stream_id_for_appservice` returns 1 when there is no stored + # position, and actual receipts occupy stream IDs 2 and upwards. + is_first_receipt_delivery = from_key.stream <= 1 + + # Fetch read receipts for all rooms, in ascending stream order. The number + # of receipts fetched is capped, in which case `reached_token` tells us how + # far we actually got. + # + # On the first delivery, don't backfill the entire receipt history: just + # send the most recent receipts and fast-forward past everything older. + ( + rooms_to_events, + reached_token, + ) = await self.store.get_linearized_receipts_for_all_rooms( + from_key=from_key, + to_key=to_key, + order=Direction.BACKWARDS + if is_first_receipt_delivery + else Direction.FORWARDS, ) # Then filter down to rooms that the AS can read @@ -379,7 +399,7 @@ class ReceiptEventSource(EventSource[MultiWriterStreamToken, JsonMapping]): # https://spec.matrix.org/v1.19/application-service-api/#pushing-ephemeral-data events = self._filter_private_receipts(events, service.is_interested_in_user) - return events, to_key + return events, reached_token def get_current_key(self) -> MultiWriterStreamToken: return self.store.get_max_receipt_stream_id() diff --git a/synapse/storage/databases/main/receipts.py b/synapse/storage/databases/main/receipts.py index ba5e07a051..3720d25c1d 100644 --- a/synapse/storage/databases/main/receipts.py +++ b/synapse/storage/databases/main/receipts.py @@ -32,7 +32,7 @@ from typing import ( import attr -from synapse.api.constants import EduTypes +from synapse.api.constants import Direction, EduTypes from synapse.replication.tcp.streams import ReceiptsStream from synapse.storage._base import SQLBaseStore, db_to_json, make_in_list_sql_clause from synapse.storage.database import ( @@ -42,7 +42,11 @@ from synapse.storage.database import ( make_tuple_in_list_sql_clause, ) from synapse.storage.engines._base import IsolationLevel -from synapse.storage.util.id_generators import MultiWriterIdGenerator +from synapse.storage.util.id_generators import ( + MultiWriterIdGenerator, + advance_multiwriter_sharded_token_after_partial_read, + make_multiwriter_sharded_token_bounds_sql, +) from synapse.types import ( JsonDict, JsonMapping, @@ -578,54 +582,81 @@ class ReceiptsWorkerStore(SQLBaseStore): self, to_key: MultiWriterStreamToken, from_key: MultiWriterStreamToken | None = None, - ) -> Mapping[str, JsonMapping]: - """Get receipts for all rooms between two stream_ids, up - to a limit of the latest 100 read receipts. + limit: int = 100, + order: Direction = Direction.FORWARDS, + ) -> tuple[Mapping[str, JsonMapping], MultiWriterStreamToken]: + """Get receipts for all rooms between two stream_ids, up to a limit of + `limit` read receipts in that range. Args: to_key: Max stream id to fetch receipts up to. from_key: Min stream id to fetch receipts from. None fetches from the start. + limit: The maximum number of receipts to fetch. + order: `Direction.FORWARDS` fetches the oldest `limit` receipts in + the range. `Direction.BACKWARDS` fetches the newest `limit` + instead; anything older is deliberately skipped, and the + returned stream token is always `to_key`. Returns: - A dictionary of roomids to a list of receipts. + A two-tuple containing the following: + * A dictionary of roomids to receipt EDUs. + * The stream token up to which receipts were actually fetched. + With `Direction.FORWARDS`, this is earlier than `to_key` (per + writer) if the limit was hit; callers must call this method + again from the returned token to fetch the remaining + receipts. With `Direction.BACKWARDS`, it is always `to_key`, + even if the limit was hit: older receipts are skipped, not + left for a later call. """ + sql_order = "DESC" if order == Direction.BACKWARDS else "ASC" - def f(txn: LoggingTransaction) -> list[tuple[str, str, str, str, str]]: - if from_key: - sql = """ - SELECT stream_id, instance_name, room_id, receipt_type, user_id, event_id, data - FROM receipts_linearized WHERE - stream_id > ? AND stream_id <= ? - ORDER BY stream_id DESC - LIMIT 100 - """ - txn.execute(sql, [from_key.stream, to_key.get_max_stream_pos()]) - else: - sql = """ - SELECT stream_id, instance_name, room_id, receipt_type, user_id, event_id, data - FROM receipts_linearized WHERE - stream_id <= ? - ORDER BY stream_id DESC - LIMIT 100 - """ + # Bound each row on its own writer's positions in the tokens, so that + # `limit` applies after the bounds and a truncated page never leaves + # in-range rows behind. + from_token = ( + from_key if from_key is not None else MultiWriterStreamToken(stream=0) + ) + bounds_clause, bounds_values = make_multiwriter_sharded_token_bounds_sql( + self.database_engine, + stream_id_column="stream_id", + instance_name_column="instance_name", + from_token_exclusive=from_token, + to_token_inclusive=to_key, + ) - txn.execute(sql, [to_key.get_max_stream_pos()]) + def f(txn: LoggingTransaction) -> list[tuple[int, str, str, str, str, str]]: + sql = f""" + SELECT stream_id, room_id, receipt_type, user_id, event_id, data + FROM receipts_linearized + WHERE {bounds_clause} + ORDER BY stream_id {sql_order} + LIMIT ? + """ + txn.execute(sql, [*bounds_values, limit]) - return [ - (room_id, receipt_type, user_id, event_id, data) - for stream_id, instance_name, room_id, receipt_type, user_id, event_id, data in txn - if MultiWriterStreamToken.is_stream_position_in_range( - from_key, to_key, instance_name, stream_id - ) - ] + return cast(list[tuple[int, str, str, str, str, str]], txn.fetchall()) txn_results = await self.db_pool.runInteraction( "get_linearized_receipts_for_all_rooms", f ) + if order == Direction.FORWARDS and len(txn_results) == limit: + # We hit the limit, so there may be more receipts in the range. + # Report how far we actually got so that the caller can fetch the + # rest, claiming each writer's position only up to the last row we + # fetched so that receipts from writers that are behind it are not + # skipped. + reached_token = advance_multiwriter_sharded_token_after_partial_read( + from_token_exclusive=from_token, + to_token_inclusive=to_key, + last_read_stream_id=txn_results[-1][0], + ) + else: + reached_token = to_key + results: JsonDict = {} - for room_id, receipt_type, user_id, event_id, data in txn_results: + for _stream_id, room_id, receipt_type, user_id, event_id, data in txn_results: # We want a single event per room, since we want to batch the # receipts by room, event and type. room_event = results.setdefault( @@ -640,7 +671,7 @@ class ReceiptsWorkerStore(SQLBaseStore): receipt_type_dict[user_id] = db_to_json(data) - return results + return results, reached_token async def get_linearized_receipts_for_user_in_rooms( self, user_id: str, room_ids: StrCollection, to_key: MultiWriterStreamToken diff --git a/synapse/storage/util/id_generators.py b/synapse/storage/util/id_generators.py index fbf75d5a2a..d370f2f815 100644 --- a/synapse/storage/util/id_generators.py +++ b/synapse/storage/util/id_generators.py @@ -36,6 +36,7 @@ from typing import ( ) import attr +from immutabledict import immutabledict from prometheus_client import Gauge from sortedcontainers import SortedList, SortedSet @@ -48,9 +49,10 @@ from synapse.storage.database import ( LoggingTransaction, make_in_list_sql_clause, ) -from synapse.storage.engines import PostgresEngine +from synapse.storage.engines import BaseDatabaseEngine, PostgresEngine, Sqlite3Engine from synapse.storage.types import Cursor from synapse.storage.util.sequence import build_sequence_generator +from synapse.types import MultiWriterStreamToken if TYPE_CHECKING: from synapse.notifier import ReplicationNotifier @@ -1019,3 +1021,198 @@ class _MultiWriterCtxManager: ) return False + + +def make_multiwriter_sharded_token_bounds_sql( + db_engine: BaseDatabaseEngine, + *, + stream_id_column: str, + instance_name_column: str, + from_token_exclusive: MultiWriterStreamToken, + to_token_inclusive: MultiWriterStreamToken, +) -> tuple[str, Sequence[str | int]]: + """ + Build an SQL clause that, in a multi-writer stream table, + matches the rows within the bounds of the `MultiWriterStreamToken` tokens provided. + + Bounds: from_token_exclusive < ... <= to_token_inclusive + + Arguments: + stream_id_column: the name of the `stream_id` column + Should be prefixed with the table name or alias + when used in a multi-table SELECT statement. + instance_name_column: the name of the `instance_name` column. + Should be prefixed with the table name or alias + when used in a multi-table SELECT statement. + from_token_exclusive: + MultiWriterStreamToken representing the highest position that has 'already been seen'. + Exclusive lower bound + to_token_inclusive: + MultiWriterStreamToken representing the highest position that has visible data. + Inclusive upper bound + + Example: + + make_multiwriter_sharded_token_bounds_sql( + stream_id_column="se.stream_id", + instance_name_column="se.instance_name", + # m5~1.8 + from_token_exclusive=MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 8}) + ), + # m10~2.14 + to_token_inclusive=MultiWriterStreamToken( + stream=10, instance_map=immutabledict({"worker2": 14}) + ), + ) + + gives: + + ( + 5 < se.stream_id AND se.stream_id <= 14 + AND NOT (se.instance_name = 'worker1' AND se.stream_id <= 8) + AND ( + se.stream_id <= 10 + OR (se.instance_name = 'worker2' AND se.stream_id <= 14) + ) + ) + + which matches, for example: + - a `worker1` row at position 9 (past the 8 we had already read from + `worker1`, and within the 10 that is visible for every writer); + - a `worker2` row at position 12 (past 5, and `worker2` is + explicitly visible up to 14); + but not: + - a `worker1` row at position 7, which the `from` token says we have + already read; + - a `worker3` row at position 12, which is beyond the position that + is visible for writers other than `worker2`. + + See also: + - `advance_multiwriter_sharded_token_after_partial_read` to create the next + `from_token_exclusive` after the result of reading a limited set of rows. + """ + + is_not_distinct_from = "IS NOT DISTINCT FROM" + if isinstance(db_engine, Sqlite3Engine): + # TODO(SQLite 3.39.0): Drop this compatibility code + # SQL standard `IS NOT DISTINCT FROM` was not introduced until 3.39.0 + # https://sqlite.org/releaselog/3_39_0.html + is_not_distinct_from = "IS" + + # The SQL we build will fundamentally consist of many clauses ANDed together. + # + # We start with an envelope delimited by the two outermost, writer-independent, bounds. + # We know everything before the lowest position of `from` and everything after the highest + # position of `to` must be out-of-bounds. + # This pair of bounds is the most likely to be index-friendly. + clauses = [f"? < {stream_id_column}", f"{stream_id_column} <= ?"] + values: list[int | str] = [ + from_token_exclusive.stream, + to_token_inclusive.get_max_stream_pos(), + ] + + # Now for every writer where we had already read further ahead than the baseline (lowest) position of `from`, + # we whittle away what we have already seen from the lower end of the envelope. + for instance_name, pos in from_token_exclusive.instance_map.items(): + # We need `IS NOT DISTINCT FROM` (analogous to `=` but treats `NULL` as a known value) + # for legacy rows where `instance_name` is `NULL`, such as before the stream was sharded. + clauses.append( + f"NOT ({instance_name_column} {is_not_distinct_from} ? AND {stream_id_column} <= ?)" + ) + values.extend((instance_name, pos)) + + # Now we need to restrict the upper end of the envelope to only include those rows where either: + # - the row is in a region where all writers' rows are visible; or + upper_or_clauses = [f"{stream_id_column} <= ?"] + upper_or_values: list[int | str] = [to_token_inclusive.stream] + # - the row's writer is explicitly visible ahead of the baseline (minimum) position + for instance_name, pos in to_token_inclusive.instance_map.items(): + upper_or_clauses.append( + f"({instance_name_column} {is_not_distinct_from} ? AND {stream_id_column} <= ?)" + ) + upper_or_values.extend((instance_name, pos)) + + clauses.append("( " + " OR ".join(upper_or_clauses) + " )") + values.extend(upper_or_values) + + return "( " + " AND ".join(clauses) + " )", values + + +def advance_multiwriter_sharded_token_after_partial_read( + *, + from_token_exclusive: MultiWriterStreamToken, + to_token_inclusive: MultiWriterStreamToken, + last_read_stream_id: int, +) -> MultiWriterStreamToken: + """ + Calculates the 'read up to' token after reading a limited number of rows + between two `MultiWriterStreamToken` bounds. + Pairs with `make_multiwriter_sharded_token_bounds_sql`. + + The read operation MUST have been ordered by the stream ID, e.g. + using `ORDER BY stream_id`. + + Caution: + For streams where one fact is represented by multiple rows with + the same `stream_id`, the ENTIRE fact corresponding to `last_read_stream_id` + MUST have been processed because this function won't help you 'pause' + in the middle of a fact. + + Arguments: + from_token_exclusive: the exclusive lower bound the rows were read with. + to_token_inclusive: the inclusive upper bound the rows were read with. + last_read_stream_id: the `stream_id` of the last row that was read. + + Returns: + A token that could be used as the `from_token_exclusive` for a subsequent + read, in order to continue reading rows in range. + + The token represents the maximum 'read up to' position. + Rows with positions strictly above the token have not yet been read. + """ + assert from_token_exclusive.stream < last_read_stream_id, "read was out of bounds" + + # Calculate the new baseline position that we have read up to, + # which applies across all workers. + new_baseline = max( + # Must be at least as far as it was before (doesn't go backwards) + from_token_exclusive.stream, + min( + # Simple case: this is just `last_read_stream_id`, + # the `stream_id` of the last row we read. + last_read_stream_id, + # Complicated case: some writers may not have advanced to `last_read_stream_id` yet, + # so we have to cap off at the baseline 'to' position (`to_token_inclusive.stream`) + to_token_inclusive.stream, + ), + ) + + # Now consider whether any workers are ahead of the baseline. + # We do this by calculating the position of each relevant worker. + # (Relevant by virtue of being in either the `from` or `to` token.) + advanced_instance_map = {} + for instance_name in ( + from_token_exclusive.instance_map.keys() + | to_token_inclusive.instance_map.keys() + ): + # Calculate the worker's position + worker_pos = max( + # Must be at least as far as it was before (doesn't go backwards) + from_token_exclusive.get_stream_pos_for_instance(instance_name), + # Then: + # - up to `last_read_stream_id` + # - unless its rows aren't visible yet, so cap off at the `to` position + # Either one of these cases will be at least as high as `new_baseline`. + min( + last_read_stream_id, + to_token_inclusive.get_stream_pos_for_instance(instance_name), + ), + ) + + if worker_pos > new_baseline: + advanced_instance_map[instance_name] = worker_pos + + return MultiWriterStreamToken( + stream=new_baseline, instance_map=immutabledict(advanced_instance_map) + ) diff --git a/tests/handlers/test_appservice.py b/tests/handlers/test_appservice.py index 78a186ee4e..43fac00296 100644 --- a/tests/handlers/test_appservice.py +++ b/tests/handlers/test_appservice.py @@ -29,6 +29,7 @@ from typing import ( ) from unittest.mock import AsyncMock, Mock +from immutabledict import immutabledict from parameterized import parameterized from twisted.internet import defer @@ -341,7 +342,7 @@ class AppServiceHandlerTestCase(unittest.TestCase): event = Mock(event_id="event_1") self.event_source.sources.receipt.get_new_events_as = AsyncMock( - return_value=([event], None) + return_value=([event], MultiWriterStreamToken(stream=580)) ) self.handler.notify_interested_services_ephemeral( @@ -371,7 +372,7 @@ class AppServiceHandlerTestCase(unittest.TestCase): event = Mock(event_id="event_1") self.event_source.sources.receipt.get_new_events_as = AsyncMock( - return_value=([event], None) + return_value=([event], MultiWriterStreamToken(stream=580)) ) self.handler.notify_interested_services_ephemeral( @@ -804,6 +805,276 @@ class ApplicationServicesHandlerSendEventsTestCase(unittest.HomeserverTestCase): }, ) + def test_sending_read_receipt_batches_with_single_token_to_application_services( + self, + ) -> None: + """Tests that a large batch of read receipts covered by a single stream + token notification (e.g. a burst arriving over federation, or an + application service catching up after downtime) is sent in full, rather + than being truncated by the per-fetch receipt limit. + """ + interested_appservice = self._register_application_service( + namespaces={ + ApplicationService.NS_USERS: [ + { + "regex": "@exclusive_as_user:.+", + "exclusive": True, + } + ], + ApplicationService.NS_ROOMS: [ + { + "regex": "!fakeroom_.*", + "exclusive": True, + } + ], + }, + ) + + # Deliver a first receipt to establish a stored read receipt stream + # position for this appservice, as an appservice without one is + # fast-forwarded to the most recent receipts instead of backfilling. + # note: stream tokens start at 2 + self.get_success( + self.hs.get_datastores().main.insert_receipt( + room_id="!fakeroom_bootstrap:test", + receipt_type="m.read", + user_id=self.local_user, + event_ids=["$eventid_bootstrap"], + thread_id=None, + data={}, + ) + ) + self.get_success( + self.hs.get_application_service_handler()._notify_interested_services_ephemeral( + services=[interested_appservice], + stream_key=StreamKeyType.RECEIPT, + new_token=MultiWriterStreamToken(stream=2), + users=[self.exclusive_as_user], + ) + ) + self.send_mock.reset_mock() + + # Insert a large burst of read receipts (300 total, past the per-fetch + # limit of 100), occupying stream IDs 3 to 302. + for i in range(300): + self.get_success( + self.hs.get_datastores().main.insert_receipt( + # We have to use unique room ID + user ID combinations here, as the db query + # is an upsert. + room_id=f"!fakeroom_{i}:test", + receipt_type="m.read", + user_id=self.local_user, + event_ids=[f"$eventid_{i}"], + thread_id=None, + data={}, + ) + ) + + # Now notify the appservice handler with a single token covering all 300 + # read receipts at once. + self.get_success( + self.hs.get_application_service_handler()._notify_interested_services_ephemeral( + services=[interested_appservice], + stream_key=StreamKeyType.RECEIPT, + new_token=MultiWriterStreamToken(stream=302), + users=[self.exclusive_as_user], + ) + ) + + # Using our txn send mock, we can see what the AS received. After iterating over every + # transaction, we'd like to see all 300 read receipts accounted for. + # No more, no less. + all_ephemeral_events = [] + for call in self.send_mock.call_args_list: + ephemeral_events = call[0][2] + all_ephemeral_events += ephemeral_events + + self.assertEqual(len(all_ephemeral_events), 300) + + # The stored stream position should have caught up with the notified token. + self.assertEqual( + self.get_success( + self.hs.get_datastores().main.get_type_stream_id_for_appservice( + interested_appservice, "read_receipt" + ) + ), + 302, + ) + + def test_read_receipts_from_lagging_writer_are_not_skipped(self) -> None: + """ + With several receipt writers, a worker can be notified with a token in + which one writer is ahead of the others, e.g. because its replication + rows arrived first. The application service handler must not treat the + leading writer's position as the point it has caught up to, or the + lagging writer's receipts would be skipped for good once they arrive. + + Scenario: writers `rw1` and `rw2` alternate stream IDs 1001 to 1200. + The worker has heard from `rw2` up to 1200 but from `rw1` only up to + 1000, so the receipt stream watermark is still 1000. + """ + interested_appservice = self._register_application_service( + namespaces={ + ApplicationService.NS_ROOMS: [ + { + "regex": "!fakeroom_.*", + "exclusive": True, + } + ], + }, + ) + + # The application service has already been sent everything up to 1000. + self.get_success( + self.hs.get_datastores().main.set_appservice_stream_type_pos( + interested_appservice, "read_receipt", 1000 + ) + ) + + # Insert the receipts as the two writers would have persisted them, each + # in its own room so they can be told apart in what the appservice gets. + self.get_success( + self.hs.get_datastores().main.db_pool.simple_insert_many( + desc="test_read_receipts_from_lagging_writer_are_not_skipped", + table="receipts_linearized", + keys=( + "stream_id", + "instance_name", + "room_id", + "receipt_type", + "user_id", + "event_id", + "data", + ), + values=[ + ( + stream_id, + "rw1" if stream_id % 2 else "rw2", + f"!fakeroom_{stream_id}:test", + "m.read", + self.local_user, + f"$eventid_{stream_id}", + "{}", + ) + for stream_id in range(1001, 1201) + ], + ) + ) + + def notify(new_token: MultiWriterStreamToken) -> None: + self.get_success( + self.hs.get_application_service_handler()._notify_interested_services_ephemeral( + services=[interested_appservice], + stream_key=StreamKeyType.RECEIPT, + new_token=new_token, + users=[], + ) + ) + + def received_stream_ids() -> list[int]: + return [ + int(event["room_id"].removeprefix("!fakeroom_").removesuffix(":test")) + for call in self.send_mock.call_args_list + for event in call[0][2] + ] + + def stored_position() -> int: + return self.get_success( + self.hs.get_datastores().main.get_type_stream_id_for_appservice( + interested_appservice, "read_receipt" + ) + ) + + # `rw2`'s rows have replicated to this worker, `rw1`'s have not. + notify( + MultiWriterStreamToken( + stream=1000, instance_map=immutabledict({"rw2": 1200}) + ) + ) + + # Nothing can be sent yet without risking `rw1`'s receipts being skipped, + # and the stored position must not move past them. + self.assertEqual(received_stream_ids(), []) + self.assertEqual(stored_position(), 1000) + + # `rw1`'s rows arrive and the watermark catches up. + notify(MultiWriterStreamToken(stream=1200)) + + # Every receipt from both writers is delivered exactly once, in order. + self.assertEqual(received_stream_ids(), list(range(1001, 1201))) + self.assertEqual(stored_position(), 1200) + + def test_application_services_are_fast_forwarded_on_first_read_receipt_delivery( + self, + ) -> None: + """Tests that an application service without a stored read receipt stream + position (i.e. one which has never been sent receipts before) is sent only + the most recent receipts, rather than the entire receipt history. + """ + interested_appservice = self._register_application_service( + namespaces={ + ApplicationService.NS_USERS: [ + { + "regex": "@exclusive_as_user:.+", + "exclusive": True, + } + ], + ApplicationService.NS_ROOMS: [ + { + "regex": "!fakeroom_.*", + "exclusive": True, + } + ], + }, + ) + + # Insert 300 read receipts, occupying stream IDs 2 to 301. + for i in range(300): + self.get_success( + self.hs.get_datastores().main.insert_receipt( + room_id=f"!fakeroom_{i}:test", + receipt_type="m.read", + user_id=self.local_user, + event_ids=[f"$eventid_{i}"], + thread_id=None, + data={}, + ) + ) + + # Notify the appservice handler, which has no stored stream position for + # this appservice yet. + self.get_success( + self.hs.get_application_service_handler()._notify_interested_services_ephemeral( + services=[interested_appservice], + stream_key=StreamKeyType.RECEIPT, + new_token=MultiWriterStreamToken(stream=301), + users=[self.exclusive_as_user], + ) + ) + + all_ephemeral_events = [] + for call in self.send_mock.call_args_list: + ephemeral_events = call[0][2] + all_ephemeral_events += ephemeral_events + + # Only the most recent 100 receipts (the per-fetch limit) should have been + # sent, not the whole history. + self.assertEqual(len(all_ephemeral_events), 100) + received_rooms = {event["room_id"] for event in all_ephemeral_events} + self.assertEqual( + received_rooms, {f"!fakeroom_{i}:test" for i in range(200, 300)} + ) + + # The stored stream position should nonetheless be at the notified token. + self.assertEqual( + self.get_success( + self.hs.get_datastores().main.get_type_stream_id_for_appservice( + interested_appservice, "read_receipt" + ) + ), + 301, + ) + @unittest.override_config( {"experimental_features": {"msc2409_to_device_messages_enabled": True}} ) diff --git a/tests/storage/test_id_generators.py b/tests/storage/test_id_generators.py index 0666eec00c..8eb2befe9c 100644 --- a/tests/storage/test_id_generators.py +++ b/tests/storage/test_id_generators.py @@ -19,8 +19,11 @@ # # +import sqlite3 from unittest import mock +from immutabledict import immutabledict + from twisted.internet.defer import CancelledError, Deferred, ensureDeferred from twisted.internet.testing import MemoryReactor @@ -31,9 +34,12 @@ from synapse.storage.database import ( LoggingDatabaseConnection, LoggingTransaction, ) +from synapse.storage.engines import Sqlite3Engine from synapse.storage.types import Cursor from synapse.storage.util.id_generators import ( MultiWriterIdGenerator, + advance_multiwriter_sharded_token_after_partial_read, + make_multiwriter_sharded_token_bounds_sql, stream_current_position_gauge, ) from synapse.storage.util.sequence import ( @@ -41,9 +47,10 @@ from synapse.storage.util.sequence import ( PostgresSequenceGenerator, SequenceGenerator, ) +from synapse.types import MultiWriterStreamToken from synapse.util.clock import Clock -from tests.unittest import HomeserverTestCase +from tests.unittest import HomeserverTestCase, TestCase from tests.utils import USE_POSTGRES_FOR_TESTS @@ -915,3 +922,333 @@ class MultiTableMultiWriterIdGeneratorTestCase(MultiWriterIdGeneratorBase): self.assertEqual(second_id_gen.get_current_token_for_writer("first"), 7) self.assertEqual(second_id_gen.get_current_token_for_writer("second"), 7) self.assertEqual(second_id_gen.get_persisted_upto_position(), 7) + + +class ShardedTokenHelpersPureTestCase(TestCase): + """ + Non-database tests for the helpers for reading multi-writer streams. + """ + + def test_bounds_sql_documented_example(self) -> None: + """ + Tests that the example in the docstring is what we actually generate. + """ + clause, values = make_multiwriter_sharded_token_bounds_sql( + Sqlite3Engine({}), + stream_id_column="se.stream_id", + instance_name_column="se.instance_name", + from_token_exclusive=MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 8}) + ), + to_token_inclusive=MultiWriterStreamToken( + stream=10, instance_map=immutabledict({"worker2": 14}) + ), + ) + self.assertEqualNormalisingWhitespace( + clause, + """ + ( + ? < se.stream_id + AND se.stream_id <= ? + AND NOT (se.instance_name IS ? AND se.stream_id <= ?) + AND ( + se.stream_id <= ? + OR (se.instance_name IS ? AND se.stream_id <= ?) + ) + ) + """, + ) + self.assertEqual(list(values), [5, 14, "worker1", 8, 10, "worker2", 14]) + + def test_token_after_partial_read_does_not_go_backwards(self) -> None: + """ + Tests that advancing the token after a partial read doesn't let it go + backwards. + + This is relevant because Synapse workers don't always advance their current + position at the same time. + """ + # As a scenario: the client has already read up to a baseline position of 60, + # but this reader worker has only caught up to 10. + # It can nonetheless see that worker2 has reached 70. + from_token = MultiWriterStreamToken(stream=60) + to_token = MultiWriterStreamToken( + stream=10, instance_map=immutabledict({"worker2": 70}) + ) + + resume_token = advance_multiwriter_sharded_token_after_partial_read( + from_token_exclusive=from_token, + to_token_inclusive=to_token, + last_read_stream_id=64, + ) + + # worker2 advances to what we read; everyone else stays where they were. + self.assertEqual( + resume_token, + MultiWriterStreamToken( + stream=60, instance_map=immutabledict({"worker2": 64}) + ), + ) + self.assertTrue( + from_token.is_before_or_eq(resume_token), + f"Expected {from_token} <= {resume_token}", + ) + + +class ShardedTokenHelpersDatabaseTestCase(TestCase): + """Tests for the helpers that read a range of a multi-writer stream. + + These don't need a homeserver: they only exercise the SQL that + `make_multiwriter_sharded_token_bounds_sql` builds, against a throwaway + SQLite table. + + NOTE: `ShardedTokenHelpersDatabaseNullTestCase` inherits all these tests + with a separate set of tables. + """ + + ROWS: list[tuple[int, str | None]] = [ + # (stream_id, instance_name) + (6, "worker1"), + (7, "worker1"), + (8, "worker2"), + (9, "worker1"), + (10, "worker3"), + (11, "worker1"), + (12, "worker2"), + (13, "worker3"), + (14, "worker2"), + ] + """ + (stream_id, instance_name) rows to insert into the example stream table. + """ + + SHARDED_TOKEN_RANGES = [ + ( + MultiWriterStreamToken(stream=5), + MultiWriterStreamToken(stream=14), + ), + ( + MultiWriterStreamToken(stream=5), + MultiWriterStreamToken( + stream=10, instance_map=immutabledict({"worker2": 14}) + ), + ), + ( + MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 9}) + ), + MultiWriterStreamToken( + stream=10, instance_map=immutabledict({"worker2": 14}) + ), + ), + ( + MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 7}) + ), + MultiWriterStreamToken( + stream=8, instance_map=immutabledict({"worker1": 11, "worker3": 13}) + ), + ), + ( + MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 9, "worker2": 8}) + ), + MultiWriterStreamToken( + stream=10, + instance_map=immutabledict( + {"worker1": 11, "worker2": 14, "worker3": 13} + ), + ), + ), + ] + """(from, to) token pairs to exercise the bounds against `ROWS`.""" + + def setUp(self) -> None: + # Just used for detecting SQL language support, + # doesn't actually connect... + self.database_engine = Sqlite3Engine({}) + self.conn = sqlite3.connect(":memory:") + self.addCleanup(self.conn.close) + self.conn.execute( + "CREATE TABLE streamtable (stream_id INTEGER PRIMARY KEY, instance_name TEXT)" + ) + self.conn.executemany("INSERT INTO streamtable VALUES (?, ?)", self.ROWS) + + def _select( + self, + from_token: MultiWriterStreamToken, + to_token: MultiWriterStreamToken, + limit: int | None = None, + ) -> list[tuple[int, str]]: + """ + Helper that selects the rows within the given bounds, in stream order, + using `make_multiwriter_sharded_token_bounds_sql`. + + Bounds: from_token < ... <= to_token + """ + clause, values = make_multiwriter_sharded_token_bounds_sql( + self.database_engine, + stream_id_column="streamtable.stream_id", + instance_name_column="streamtable.instance_name", + from_token_exclusive=from_token, + to_token_inclusive=to_token, + ) + + limit_clause = "" + if limit is not None: + limit_clause = f"LIMIT {limit}" + + return list( + self.conn.execute( + f""" + SELECT stream_id, instance_name + FROM streamtable + WHERE {clause} + ORDER BY stream_id ASC + {limit_clause} + """, + values, + ) + ) + + def _expected( + self, + from_token: MultiWriterStreamToken, + to_token: MultiWriterStreamToken, + ) -> list[tuple[int, str | None]]: + """ + Helper that returns the rows that ought to be within the given bounds. + It uses `MultiWriterStreamToken.is_stream_position_in_range` as the 'ground truth'. + + Bounds: from_token < ... <= to_token + """ + + return [ + (stream_id, instance_name) + for stream_id, instance_name in self.ROWS + if MultiWriterStreamToken.is_stream_position_in_range( + from_token, to_token, instance_name, stream_id + ) + ] + + def test_bounds_sql_selects_expected_rows(self) -> None: + """ + Tests that the `make_multiwriter_sharded_token_bounds_sql` selects the correct rows, + using `MultiWriterStreamToken.is_stream_position_in_range` as the 'ground truth'. + """ + for from_token, to_token in self.SHARDED_TOKEN_RANGES: + with self.subTest(from_token=str(from_token), to_token=str(to_token)): + self.assertEqual( + self._select(from_token, to_token), + self._expected(from_token, to_token), + ) + + def test_token_after_partial_read(self) -> None: + """ + Tests that when advancing the token after a partial read, the subsequent read returns all the remaining rows + but no more (especially not duplicates). + """ + for from_token, to_token in self.SHARDED_TOKEN_RANGES: + all_rows = self._expected(from_token, to_token) + for limit in range(1, len(all_rows) + 1): + with self.subTest( + from_token=str(from_token), to_token=str(to_token), limit=limit + ): + read = self._select(from_token, to_token, limit=limit) + resume_token = advance_multiwriter_sharded_token_after_partial_read( + from_token_exclusive=from_token, + to_token_inclusive=to_token, + last_read_stream_id=read[-1][0], + ) + + # The resumption point must be somewhere within the range we + # were asked to read. + self.assertTrue( + from_token.is_before_or_eq(resume_token), + f"Expected {from_token} <= {resume_token}", + ) + self.assertTrue( + resume_token.is_before_or_eq(to_token), + f"Expected {resume_token} <= {to_token}", + ) + + self.assertEqual( + read, + self._expected(from_token, resume_token), + f"The rows we read between {from_token} < ... <= {resume_token} seem wrong.", + ) + + self.assertEqual( + read + self._select(resume_token, to_token), + all_rows, + f"We did a limited read of {limit} rows within {from_token} < ... <= {to_token}, " + f"giving us an advanced token {resume_token}. " + f"We then read {resume_token} < ... <= {to_token} and expected to get " + "all the rows within range (in order, no duplicates), but didn't.", + ) + + +class ShardedTokenHelpersDatabaseNullTestCase(ShardedTokenHelpersDatabaseTestCase): + """ + Variant of `ShardedTokenHelpersDatabaseTestCase` where there are rows with + `instance_name` being `NULL`, corresponding to rows that existed before the stream was sharded. + """ + + ROWS = [ + # (stream_id, instance_name) + (4, None), + (5, None), + (6, None), + (7, None), + (8, "worker2"), + (9, "worker1"), + (10, "worker3"), + (11, "worker1"), + (12, "worker2"), + (13, "worker3"), + (14, "worker2"), + ] + """ + (stream_id, instance_name) rows to insert into the example stream table. + """ + + SHARDED_TOKEN_RANGES = [ + ( + MultiWriterStreamToken(stream=5), + MultiWriterStreamToken(stream=14), + ), + ( + MultiWriterStreamToken(stream=5), + MultiWriterStreamToken( + stream=10, instance_map=immutabledict({"worker2": 14}) + ), + ), + ( + MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 9}) + ), + MultiWriterStreamToken( + stream=10, instance_map=immutabledict({"worker2": 14}) + ), + ), + ( + MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 7}) + ), + MultiWriterStreamToken( + stream=8, instance_map=immutabledict({"worker1": 11, "worker3": 13}) + ), + ), + ( + MultiWriterStreamToken( + stream=5, instance_map=immutabledict({"worker1": 9, "worker2": 8}) + ), + MultiWriterStreamToken( + stream=10, + instance_map=immutabledict( + {"worker1": 11, "worker2": 14, "worker3": 13} + ), + ), + ), + ] + """(from, to) token pairs to exercise the bounds against `ROWS`.""" diff --git a/tests/storage/test_receipts.py b/tests/storage/test_receipts.py index 27875dcebb..3413978c18 100644 --- a/tests/storage/test_receipts.py +++ b/tests/storage/test_receipts.py @@ -21,11 +21,13 @@ from typing import Collection +from immutabledict import immutabledict + from twisted.internet.testing import MemoryReactor from synapse.api.constants import ReceiptTypes from synapse.server import HomeServer -from synapse.types import UserID, create_requester +from synapse.types import MultiWriterStreamToken, UserID, create_requester from synapse.util.clock import Clock from tests.test_utils.event_injection import create_event @@ -307,3 +309,67 @@ class ReceiptTestCase(HomeserverTestCase): [ReceiptTypes.READ, ReceiptTypes.READ_PRIVATE], room_id=self.room_id2 ) self.assertEqual(res, event2_1_id) + + +class MultiWriterReceiptPagingTestCase(HomeserverTestCase): + def prepare( + self, reactor: MemoryReactor, clock: Clock, homeserver: HomeServer + ) -> None: + self.store = homeserver.get_datastores().main + + def test_reached_token_does_not_cover_receipts_from_lagging_writer(self) -> None: + """ + Writers `rw1` and `rw2` alternate stream IDs 1001 to 1008. The paging + worker has heard from `rw2` up to 1008 but from `rw1` only up to 1000. + + With a limit of 4 the SQL fetches 1001 to 1004, then `rw1`'s 1001 and + 1003 are filtered out as beyond its known position. The returned token + must not claim `rw1` was handled past 1000, or those receipts are + skipped once `rw1` catches up. + """ + self.get_success( + self.store.db_pool.simple_insert_many( + table="receipts_linearized", + keys=( + "stream_id", + "instance_name", + "room_id", + "receipt_type", + "user_id", + "event_id", + "data", + ), + values=[ + ( + stream_id, + "rw1" if stream_id % 2 else "rw2", + f"!room_{stream_id}:test", + ReceiptTypes.READ, + "@user:test", + f"$event_{stream_id}", + "{}", + ) + for stream_id in range(1001, 1009) + ], + desc="insert_receipts", + ) + ) + + from_key = MultiWriterStreamToken(stream=1000) + to_key = MultiWriterStreamToken( + stream=1000, instance_map=immutabledict({"rw2": 1008}) + ) + + rooms_to_events, reached_token = self.get_success( + self.store.get_linearized_receipts_for_all_rooms( + to_key=to_key, from_key=from_key, limit=4 + ) + ) + + # `rw1`'s receipts are behind its position in `to_key`, so none of them + # are returned... + self.assertFalse({"!room_1001:test", "!room_1003:test"} & set(rooms_to_events)) + + # ...and therefore the returned token must not claim `rw1` has been + # handled beyond 1000. + self.assertLessEqual(reached_token.get_stream_pos_for_instance("rw1"), 1000) diff --git a/tests/unittest.py b/tests/unittest.py index f3dc91c58a..91adfd8318 100644 --- a/tests/unittest.py +++ b/tests/unittest.py @@ -27,6 +27,7 @@ import hmac import json import logging import os +import re import secrets import sys import time @@ -460,6 +461,32 @@ class TestCase(unittest.TestCase): self._fail_with_set_inequality(actual_items, expected_items, message, exact) + def assertEqualNormalisingWhitespace( + self, received: str, expected: str, msg: str | None = None + ) -> None: + """ + Fail the test if `received` and `expected` are not equal, + after having normalised all whitespace. + + By normalising whitespace, we mean that any run of whitespace + is replaced by a single ` `, on both sides. + The front and back of the strings are also trimmed. + + Whitespace characters considered are: `\n`, `\t`, ` `. + """ + PATTERN = r"[ \n\t]+" + REPLACEMENT = " " + normalised_received = re.sub(PATTERN, REPLACEMENT, received).strip() + normalised_expected = re.sub(PATTERN, REPLACEMENT, expected).strip() + + if normalised_received == normalised_expected: + return + + msg = "" if msg is None else msg + self.fail( + f"Expected strings to match, after normalising whitespace: {msg}\nReceived: {received}\nExpected: {expected}\nReceived (normalised): {normalised_received}\nExpected (normalised): {normalised_expected}" + ) + def DEBUG(target: TV) -> TV: """A decorator to set the .loglevel attribute to logging.DEBUG.