Add sticky events store and stream

This commit is contained in:
Olivier 'reivilibre
2026-01-16 08:59:42 +00:00
parent cf5ab85fb8
commit a2b0da6c0a
14 changed files with 249 additions and 7 deletions
+2
View File
@@ -102,6 +102,7 @@ from synapse.storage.databases.main.signatures import SignatureWorkerStore
from synapse.storage.databases.main.sliding_sync import SlidingSyncStore
from synapse.storage.databases.main.state import StateGroupWorkerStore
from synapse.storage.databases.main.stats import StatsStore
from synapse.storage.databases.main.sticky_events import StickyEventsWorkerStore
from synapse.storage.databases.main.stream import StreamWorkerStore
from synapse.storage.databases.main.tags import TagsWorkerStore
from synapse.storage.databases.main.task_scheduler import TaskSchedulerWorkerStore
@@ -137,6 +138,7 @@ class GenericWorkerStore(
RoomWorkerStore,
DirectoryWorkerStore,
ThreadSubscriptionsWorkerStore,
StickyEventsWorkerStore,
PushRulesWorkerStore,
ApplicationServiceTransactionWorkerStore,
ApplicationServiceWorkerStore,
+1 -1
View File
@@ -127,7 +127,7 @@ class WriterLocations:
"""Specifies the instances that write various streams.
Attributes:
events: The instances that write to the event and backfill streams.
events: The instances that write to the event, backfill and sticky events streams.
typing: The instances that write to the typing stream. Currently
can only be a single instance.
to_device: The instances that write to the to_device stream. Currently
+1
View File
@@ -526,6 +526,7 @@ class Notifier:
StreamKeyType.TYPING,
StreamKeyType.UN_PARTIAL_STATED_ROOMS,
StreamKeyType.THREAD_SUBSCRIPTIONS,
StreamKeyType.STICKY_EVENTS,
],
new_token: int,
users: Collection[str | UserID] | None = None,
+10 -1
View File
@@ -43,7 +43,10 @@ from synapse.replication.tcp.streams import (
UnPartialStatedEventStream,
UnPartialStatedRoomStream,
)
from synapse.replication.tcp.streams._base import ThreadSubscriptionsStream
from synapse.replication.tcp.streams._base import (
StickyEventsStream,
ThreadSubscriptionsStream,
)
from synapse.replication.tcp.streams.events import (
EventsStream,
EventsStreamEventRow,
@@ -262,6 +265,12 @@ class ReplicationDataHandler:
token,
users=[row.user_id for row in rows],
)
elif stream_name == StickyEventsStream.NAME:
self.notifier.on_new_event(
StreamKeyType.STICKY_EVENTS,
token,
rooms=[row.room_id for row in rows],
)
await self._presence_handler.process_replication_rows(
stream_name, instance_name, token, rows
+7
View File
@@ -66,6 +66,7 @@ from synapse.replication.tcp.streams import (
)
from synapse.replication.tcp.streams._base import (
DeviceListsStream,
StickyEventsStream,
ThreadSubscriptionsStream,
)
from synapse.util.background_queue import BackgroundQueue
@@ -216,6 +217,12 @@ class ReplicationCommandHandler:
continue
if isinstance(stream, StickyEventsStream):
if hs.get_instance_name() in hs.config.worker.writers.events:
self._streams_to_replicate.append(stream)
continue
if isinstance(stream, DeviceListsStream):
if hs.get_instance_name() in hs.config.worker.writers.device_lists:
self._streams_to_replicate.append(stream)
@@ -40,6 +40,7 @@ from synapse.replication.tcp.streams._base import (
PushersStream,
PushRulesStream,
ReceiptsStream,
StickyEventsStream,
Stream,
ThreadSubscriptionsStream,
ToDeviceStream,
@@ -68,6 +69,7 @@ STREAMS_MAP = {
ToDeviceStream,
FederationStream,
AccountDataStream,
StickyEventsStream,
ThreadSubscriptionsStream,
UnPartialStatedRoomStream,
UnPartialStatedEventStream,
@@ -90,6 +92,7 @@ __all__ = [
"ToDeviceStream",
"FederationStream",
"AccountDataStream",
"StickyEventsStream",
"ThreadSubscriptionsStream",
"UnPartialStatedRoomStream",
"UnPartialStatedEventStream",
+45
View File
@@ -763,3 +763,48 @@ class ThreadSubscriptionsStream(_StreamFromIdGen):
return [], to_token, False
return rows, rows[-1][0], len(updates) == limit
@attr.s(slots=True, auto_attribs=True)
class StickyEventsStreamRow:
"""Stream to inform workers about changes to sticky events."""
room_id: str
event_id: str
"""The sticky event ID"""
class StickyEventsStream(_StreamFromIdGen):
"""A sticky event was changed."""
NAME = "sticky_events"
ROW_TYPE = StickyEventsStreamRow
def __init__(self, hs: "HomeServer"):
self.store = hs.get_datastores().main
super().__init__(
hs.get_instance_name(),
self._update_function,
self.store._sticky_events_id_gen,
)
async def _update_function(
self, instance_name: str, from_token: int, to_token: int, limit: int
) -> StreamUpdateResult:
updates = await self.store.get_updated_sticky_events(
from_id=from_token, to_id=to_token, limit=limit
)
rows = [
(
stream_id,
# These are the args to `StickyEventsStreamRow`
(room_id, event_id),
)
for stream_id, room_id, event_id, _ in updates
]
if not rows:
return [], to_token, False
return rows, rows[-1][0], len(updates) == limit
@@ -34,6 +34,7 @@ from synapse.storage.database import (
)
from synapse.storage.databases.main.sliding_sync import SlidingSyncStore
from synapse.storage.databases.main.stats import UserSortOrder
from synapse.storage.databases.main.sticky_events import StickyEventsWorkerStore
from synapse.storage.databases.main.thread_subscriptions import (
ThreadSubscriptionsWorkerStore,
)
@@ -144,6 +145,7 @@ class DataStore(
TagsStore,
AccountDataStore,
ThreadSubscriptionsWorkerStore,
StickyEventsWorkerStore,
PushRulesWorkerStore,
StreamWorkerStore,
OpenIdStore,
@@ -68,6 +68,10 @@ from synapse.metrics.background_process_metrics import (
wrap_as_background_process,
)
from synapse.replication.tcp.streams import BackfillStream, UnPartialStatedEventStream
from synapse.replication.tcp.streams._base import (
StickyEventsStream,
StickyEventsStreamRow,
)
from synapse.replication.tcp.streams.events import EventsStream
from synapse.replication.tcp.streams.partial_state import UnPartialStatedEventStreamRow
from synapse.storage._base import SQLBaseStore, db_to_json, make_in_list_sql_clause
@@ -459,6 +463,11 @@ class EventsWorkerStore(SQLBaseStore):
# If the partial-stated event became rejected or unrejected
# when it wasn't before, we need to invalidate this cache.
self._invalidate_local_get_event_cache(row.event_id)
elif stream_name == StickyEventsStream.NAME:
for row in rows:
assert isinstance(row, StickyEventsStreamRow)
# In case soft-failure status changed, invalidate the cache.
self._invalidate_local_get_event_cache(row.event_id)
super().process_replication_rows(stream_name, instance_name, token, rows)
@@ -0,0 +1,153 @@
#
# This file is licensed under the Affero General Public License (AGPL) version 3.
#
# Copyright (C) 2025 New Vector, Ltd
#
# This program is free software: you can redistribute it and/or modify
# it under the terms of the GNU Affero General Public License as
# published by the Free Software Foundation, either version 3 of the
# License, or (at your option) any later version.
#
# See the GNU Affero General Public License for more details:
# <https://www.gnu.org/licenses/agpl-3.0.html>.
import logging
from typing import (
TYPE_CHECKING,
cast,
)
from twisted.internet.defer import Deferred
from synapse.events import EventBase
from synapse.replication.tcp.streams._base import StickyEventsStream
from synapse.storage.database import (
DatabasePool,
LoggingDatabaseConnection,
LoggingTransaction,
)
from synapse.storage.databases.main.cache import CacheInvalidationWorkerStore
from synapse.storage.databases.main.state import StateGroupWorkerStore
from synapse.storage.util.id_generators import MultiWriterIdGenerator
from synapse.util.duration import Duration
if TYPE_CHECKING:
from synapse.server import HomeServer
logger = logging.getLogger(__name__)
# Remove entries from the sticky_events table at this frequency.
# Note: this does NOT mean we don't honour shorter expiration timeouts.
# Consumers call 'get_sticky_events_in_rooms' which has `WHERE expires_at > ?`
# to filter out expired sticky events that have yet to be deleted.
DELETE_EXPIRED_STICKY_EVENTS_INTERVAL = Duration(hours=1)
class StickyEventsWorkerStore(StateGroupWorkerStore, CacheInvalidationWorkerStore):
def __init__(
self,
database: DatabasePool,
db_conn: LoggingDatabaseConnection,
hs: "HomeServer",
):
super().__init__(database, db_conn, hs)
self._can_write_to_sticky_events = (
self._instance_name in hs.config.worker.writers.events
)
# Technically this means we will cleanup N times, once per event persister, maybe put on master?
if self._can_write_to_sticky_events:
self.clock.looping_call(
self._run_background_cleanup, DELETE_EXPIRED_STICKY_EVENTS_INTERVAL
)
self._sticky_events_id_gen: MultiWriterIdGenerator = MultiWriterIdGenerator(
db_conn=db_conn,
db=database,
notifier=hs.get_replication_notifier(),
stream_name="sticky_events",
server_name=self.server_name,
instance_name=self._instance_name,
tables=[
("sticky_events", "instance_name", "stream_id"),
],
sequence_name="sticky_events_sequence",
writers=hs.config.worker.writers.events,
)
def process_replication_position(
self, stream_name: str, instance_name: str, token: int
) -> None:
if stream_name == StickyEventsStream.NAME:
self._sticky_events_id_gen.advance(instance_name, token)
super().process_replication_position(stream_name, instance_name, token)
def get_max_sticky_events_stream_id(self) -> int:
"""Get the current maximum stream_id for thread subscriptions.
Returns:
The maximum stream_id
"""
return self._sticky_events_id_gen.get_current_token()
def get_sticky_events_stream_id_generator(self) -> MultiWriterIdGenerator:
return self._sticky_events_id_gen
async def get_updated_sticky_events(
self, from_id: int, to_id: int, limit: int
) -> list[tuple[int, str, str, bool]]:
"""Get updates to sticky events between two stream IDs.
Args:
from_id: The starting stream ID (exclusive)
to_id: The ending stream ID (inclusive)
limit: The maximum number of rows to return
Returns:
list of (stream_id, room_id, event_id, soft_failed) tuples
"""
return await self.db_pool.runInteraction(
"get_updated_sticky_events",
self._get_updated_sticky_events_txn,
from_id,
to_id,
limit,
)
def _get_updated_sticky_events_txn(
self, txn: LoggingTransaction, from_id: int, to_id: int, limit: int
) -> list[tuple[int, str, str, bool]]:
txn.execute(
"""
SELECT stream_id, room_id, event_id, soft_failed
FROM sticky_events
WHERE ? < stream_id AND stream_id <= ?
LIMIT ?
""",
(from_id, to_id, limit),
)
return cast(list[tuple[int, str, str, bool]], txn.fetchall())
async def _delete_expired_sticky_events(self) -> None:
logger.info("delete_expired_sticky_events")
await self.db_pool.runInteraction(
"_delete_expired_sticky_events",
self._delete_expired_sticky_events_txn,
self.clock.time_msec(),
)
def _delete_expired_sticky_events_txn(
self, txn: LoggingTransaction, now: int
) -> None:
txn.execute(
"""
DELETE FROM sticky_events WHERE expires_at < ?
""",
(now,),
)
def _run_background_cleanup(self) -> Deferred:
return self.hs.run_as_background_process(
"delete_expired_sticky_events",
self._delete_expired_sticky_events,
)
+3
View File
@@ -84,6 +84,7 @@ class EventSources:
self._instance_name
)
thread_subscriptions_key = self.store.get_max_thread_subscriptions_stream_id()
sticky_events_key = self.store.get_max_sticky_events_stream_id()
token = StreamToken(
room_key=self.sources.room.get_current_key(),
@@ -98,6 +99,7 @@ class EventSources:
groups_key=0,
un_partial_stated_rooms_key=un_partial_stated_rooms_key,
thread_subscriptions_key=thread_subscriptions_key,
sticky_events_key=sticky_events_key,
)
return token
@@ -125,6 +127,7 @@ class EventSources:
StreamKeyType.DEVICE_LIST: self.store.get_device_stream_id_generator(),
StreamKeyType.UN_PARTIAL_STATED_ROOMS: self.store.get_un_partial_stated_rooms_id_generator(),
StreamKeyType.THREAD_SUBSCRIPTIONS: self.store.get_thread_subscriptions_stream_id_generator(),
StreamKeyType.STICKY_EVENTS: self.store.get_sticky_events_stream_id_generator(),
}
for _, key in StreamKeyType.__members__.items():
+9 -1
View File
@@ -1006,6 +1006,7 @@ class StreamKeyType(Enum):
DEVICE_LIST = "device_list_key"
UN_PARTIAL_STATED_ROOMS = "un_partial_stated_rooms_key"
THREAD_SUBSCRIPTIONS = "thread_subscriptions_key"
STICKY_EVENTS = "sticky_events_key"
@attr.s(slots=True, frozen=True, auto_attribs=True)
@@ -1027,6 +1028,7 @@ class StreamToken:
9. `groups_key`: `1` (note that this key is now unused)
10. `un_partial_stated_rooms_key`: `379`
11. `thread_subscriptions_key`: 4242
12. `sticky_events_key`: 4141
You can see how many of these keys correspond to the various
fields in a "/sync" response:
@@ -1086,6 +1088,7 @@ class StreamToken:
groups_key: int
un_partial_stated_rooms_key: int
thread_subscriptions_key: int
sticky_events_key: int
_SEPARATOR = "_"
START: ClassVar["StreamToken"]
@@ -1114,6 +1117,7 @@ class StreamToken:
groups_key,
un_partial_stated_rooms_key,
thread_subscriptions_key,
sticky_events_key,
) = keys
return cls(
@@ -1130,6 +1134,7 @@ class StreamToken:
groups_key=int(groups_key),
un_partial_stated_rooms_key=int(un_partial_stated_rooms_key),
thread_subscriptions_key=int(thread_subscriptions_key),
sticky_events_key=int(sticky_events_key),
)
except CancelledError:
raise
@@ -1153,6 +1158,7 @@ class StreamToken:
str(self.groups_key),
str(self.un_partial_stated_rooms_key),
str(self.thread_subscriptions_key),
str(self.sticky_events_key),
]
)
@@ -1218,6 +1224,7 @@ class StreamToken:
StreamKeyType.TYPING,
StreamKeyType.UN_PARTIAL_STATED_ROOMS,
StreamKeyType.THREAD_SUBSCRIPTIONS,
StreamKeyType.STICKY_EVENTS,
],
) -> int: ...
@@ -1274,7 +1281,7 @@ class StreamToken:
f"account_data: {self.account_data_key}, push_rules: {self.push_rules_key}, "
f"to_device: {self.to_device_key}, device_list: {self.device_list_key}, "
f"groups: {self.groups_key}, un_partial_stated_rooms: {self.un_partial_stated_rooms_key},"
f"thread_subscriptions: {self.thread_subscriptions_key})"
f"thread_subscriptions: {self.thread_subscriptions_key}, sticky_events: {self.sticky_events_key})"
)
@@ -1290,6 +1297,7 @@ StreamToken.START = StreamToken(
groups_key=0,
un_partial_stated_rooms_key=0,
thread_subscriptions_key=0,
sticky_events_key=0,
)
+2 -2
View File
@@ -2545,7 +2545,7 @@ class RoomMessagesTestCase(unittest.HomeserverTestCase):
def test_topo_token_is_accepted(self) -> None:
"""Test Topo Token is accepted."""
token = "t1-0_0_0_0_0_0_0_0_0_0_0"
token = "t1-0_0_0_0_0_0_0_0_0_0_0_0"
channel = self.make_request(
"GET",
"/_synapse/admin/v1/rooms/%s/messages?from=%s" % (self.room_id, token),
@@ -2559,7 +2559,7 @@ class RoomMessagesTestCase(unittest.HomeserverTestCase):
def test_stream_token_is_accepted_for_fwd_pagianation(self) -> None:
"""Test that stream token is accepted for forward pagination."""
token = "s0_0_0_0_0_0_0_0_0_0_0"
token = "s0_0_0_0_0_0_0_0_0_0_0_0"
channel = self.make_request(
"GET",
"/_synapse/admin/v1/rooms/%s/messages?from=%s" % (self.room_id, token),
+2 -2
View File
@@ -2245,7 +2245,7 @@ class RoomMessageListTestCase(RoomBase):
self.room_id = self.helper.create_room_as(self.user_id)
def test_topo_token_is_accepted(self) -> None:
token = "t1-0_0_0_0_0_0_0_0_0_0_0"
token = "t1-0_0_0_0_0_0_0_0_0_0_0_0"
channel = self.make_request(
"GET", "/rooms/%s/messages?access_token=x&from=%s" % (self.room_id, token)
)
@@ -2256,7 +2256,7 @@ class RoomMessageListTestCase(RoomBase):
self.assertTrue("end" in channel.json_body)
def test_stream_token_is_accepted_for_fwd_pagianation(self) -> None:
token = "s0_0_0_0_0_0_0_0_0_0_0"
token = "s0_0_0_0_0_0_0_0_0_0_0_0"
channel = self.make_request(
"GET", "/rooms/%s/messages?access_token=x&from=%s" % (self.room_id, token)
)