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:
Paul Chobert
2026-09-29 12:06:15 +01:00
committed by GitHub
co-authored by Olivier 'reivilibre
parent 373fa7f542
commit db618e276d
9 changed files with 1042 additions and 65 deletions
+1
View File
@@ -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.
+46 -19
View File
@@ -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,
+27 -7
View File
@@ -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()
+65 -34
View File
@@ -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
+198 -1
View File
@@ -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)
)
+273 -2
View File
@@ -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}}
)
+338 -1
View File
@@ -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`."""
+67 -1
View File
@@ -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
View File
@@ -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.