mirror of
https://github.com/element-hq/synapse.git
synced 2026-10-05 21:17:20 +00:00
Fix appservice ephemeral read receipts being permanently skipped on large receipt bursts (#20108)
Fixes https://github.com/element-hq/synapse/issues/20096. Part of https://github.com/element-hq/synapse/issues/18118 The spec requires read receipts to be delivered to interested application services ([Pushing ephemeral data](https://spec.matrix.org/v1.19/application-service-api/#pushing-ephemeral-data)): > If the `receive_ephemeral` settings is enabled in the registration file, homeservers MUST send ephemeral data that is relevant to the application service via the transaction API, using the `ephemeral` property of the request's body. > > `m.receipt`: MUST be sent to the application service under the same rules as regular events, meaning that the application service must have registered interest in the room itself, or in a user that is in the room. Ephemeral data delivery to application services was added in Matrix v1.13 from [MSC2409](https://github.com/matrix-org/matrix-spec-proposals/pull/2409) (see the [v1.13 changelog](https://spec.matrix.org/v1.19/changelog/v1.13/#application-service-api)). ## Problem When more than 100 read receipts are covered by a single receipt stream update, only the last 100 are sent to appservices. The rest is permanently skipped. ## Proposed Fix - Looping over the receipts in batches of 100 items and updating the cursor on every iteration - Keep the fast-forward behavior for new appservices as intended in https://github.com/matrix-org/synapse/pull/8744 --------- Co-authored-by: Olivier 'reivilibre <oliverw@element.io>
This commit is contained in:
co-authored by
Olivier 'reivilibre
parent
373fa7f542
commit
db618e276d
@@ -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.
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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}}
|
||||
)
|
||||
|
||||
@@ -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`."""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user