From c595869fb70b38a69cd3243c85ce4f71572fcdad Mon Sep 17 00:00:00 2001 From: Erik Johnston Date: Fri, 18 Sep 2026 12:11:09 +0100 Subject: [PATCH] Add cache to `get_partial_filtered_current_state_ids` (#20160) For state filters that ask for concrete types. This allows us to cache the common case of asking for a specific type/state key. I noticed a bunch of queries in the jaeger traces that could be cached. --- changelog.d/20160.misc | 1 + synapse/replication/tcp/client.py | 21 +++-- synapse/storage/_base.py | 2 + synapse/storage/controllers/state.py | 4 +- synapse/storage/databases/main/state.py | 105 +++++++++++++++++++++-- tests/storage/test_state.py | 108 ++++++++++++++++++++++++ 6 files changed, 227 insertions(+), 14 deletions(-) create mode 100644 changelog.d/20160.misc diff --git a/changelog.d/20160.misc b/changelog.d/20160.misc new file mode 100644 index 0000000000..476608a767 --- /dev/null +++ b/changelog.d/20160.misc @@ -0,0 +1 @@ +Add a cache for looking up individual pieces of current room state. diff --git a/synapse/replication/tcp/client.py b/synapse/replication/tcp/client.py index c0896b83e7..8a60f8df6f 100644 --- a/synapse/replication/tcp/client.py +++ b/synapse/replication/tcp/client.py @@ -50,6 +50,8 @@ from synapse.replication.tcp.streams._base import ( ) from synapse.replication.tcp.streams.events import ( EventsStream, + EventsStreamAllStateRow, + EventsStreamCurrentStateRow, EventsStreamEventRow, EventsStreamRow, ) @@ -191,6 +193,20 @@ class ReplicationDataHandler: # We shouldn't get multiple rows per token for events stream, so # we don't need to optimise this for multiple rows. for row in rows: + # If this is a server ACL event, clear the cache in the storage controller. + if row.type in ( + EventsStreamEventRow.TypeId, + EventsStreamCurrentStateRow.TypeId, + ): + if row.data.type == EventTypes.ServerACL: + self._state_storage_controller.get_server_acl_for_room.invalidate( + (row.data.room_id,) + ) + elif row.type == EventsStreamAllStateRow.TypeId: + self._state_storage_controller.get_server_acl_for_room.invalidate( + (row.data.room_id,) + ) + if row.type != EventsStreamEventRow.TypeId: # The row's data is an `EventsStreamCurrentStateRow`. # When we recompute the current state of a room based on forward @@ -238,11 +254,6 @@ class ReplicationDataHandler: row.data.event_id, row.data.room_id ) - # If this is a server ACL event, clear the cache in the storage controller. - if row.data.type == EventTypes.ServerACL: - self._state_storage_controller.get_server_acl_for_room.invalidate( - (row.data.room_id,) - ) elif stream_name == UnPartialStatedRoomStream.NAME: for row in rows: assert isinstance(row, UnPartialStatedRoomStreamRow) diff --git a/synapse/storage/_base.py b/synapse/storage/_base.py index 1df7f70b71..bc13a54dd1 100644 --- a/synapse/storage/_base.py +++ b/synapse/storage/_base.py @@ -137,6 +137,7 @@ class SQLBaseStore(metaclass=ABCMeta): # Purge other caches based on room state. self._attempt_to_invalidate_cache("get_room_summary", (room_id,)) self._attempt_to_invalidate_cache("get_partial_current_state_ids", (room_id,)) + self._attempt_to_invalidate_cache("_get_current_state_event_id", (room_id,)) self._attempt_to_invalidate_cache("get_room_type", (room_id,)) self._attempt_to_invalidate_cache("get_room_encryption", (room_id,)) self._attempt_to_invalidate_cache( @@ -154,6 +155,7 @@ class SQLBaseStore(metaclass=ABCMeta): room_id: Room where state changed """ self._attempt_to_invalidate_cache("get_partial_current_state_ids", (room_id,)) + self._attempt_to_invalidate_cache("_get_current_state_event_id", (room_id,)) self._attempt_to_invalidate_cache("get_users_in_room", (room_id,)) self._attempt_to_invalidate_cache("is_host_invited", None) self._attempt_to_invalidate_cache("is_host_joined", None) diff --git a/synapse/storage/controllers/state.py b/synapse/storage/controllers/state.py index 4885268305..b2e1d61a33 100644 --- a/synapse/storage/controllers/state.py +++ b/synapse/storage/controllers/state.py @@ -580,8 +580,8 @@ class StateStorageController: """Get the current state event ids for a room based on the current_state_events table. - If a state filter is given (that is not `StateFilter.all()`) the query - result is *not* cached. + If a wildcard state filter is given (that is not `StateFilter.all()`) + the query result is *not* cached. Args: room_id: The room to get the state IDs of. state_filter: The state diff --git a/synapse/storage/databases/main/state.py b/synapse/storage/databases/main/state.py index 87523e6f18..3e42d46084 100644 --- a/synapse/storage/databases/main/state.py +++ b/synapse/storage/databases/main/state.py @@ -49,6 +49,7 @@ from synapse.storage.database import ( LoggingDatabaseConnection, LoggingTransaction, make_in_list_sql_clause, + make_tuple_in_list_sql_clause, ) from synapse.storage.databases.main.events_worker import EventsWorkerStore from synapse.storage.databases.main.roommember import RoomMemberWorkerStore @@ -536,7 +537,84 @@ class StateGroupWorkerStore(EventsWorkerStore, SQLBaseStore): return frozenset(event_id for (event_id,) in rows) - # FIXME: how should this be cached? + @cached(max_entries=100000, tree=True) + async def _get_current_state_event_id( + self, room_id: str, event_type_and_state_key: tuple[str, str] + ) -> str | None: + """Get the event ID of the given piece of current state in the room. + + Returns None if there is no such event in the current state. + """ + return await self.db_pool.simple_select_one_onecol( + table="current_state_events", + keyvalues={ + "room_id": room_id, + "type": event_type_and_state_key[0], + "state_key": event_type_and_state_key[1], + }, + retcol="event_id", + allow_none=True, + desc="_get_current_state_event_id", + ) + + @cachedList( + cached_method_name="_get_current_state_event_id", + list_name="event_types_and_state_keys", + num_args=2, + ) + async def _get_current_state_event_ids( + self, room_id: str, event_types_and_state_keys: Collection[tuple[str, str]] + ) -> Mapping[tuple[str, str], str | None]: + """Bulk version of `_get_current_state_event_id`. + + Types/state keys that aren't in the room's current state map to None, so + that their absence gets cached too. + """ + if not event_types_and_state_keys: + return {} + + # Check if the room_id is in `get_partial_current_state_ids` cache, if + # so, we can use that to avoid a DB query. + room_state = self.get_partial_current_state_ids.cache.get_immediate( + room_id, None, update_metrics=False + ) + if room_state is not None: + return { + (intern_string(typ), intern_string(state_key)): room_state.get( + (typ, state_key) + ) + for typ, state_key in event_types_and_state_keys + } + + def _get_current_state_event_ids_txn( + txn: LoggingTransaction, + ) -> dict[tuple[str, str], str | None]: + results: dict[tuple[str, str], str | None] = { + (intern_string(typ), intern_string(state_key)): None + for typ, state_key in event_types_and_state_keys + } + + for batch in batch_iter(event_types_and_state_keys, 500): + clause, args = make_tuple_in_list_sql_clause( + self.database_engine, ("type", "state_key"), batch + ) + + sql = f""" + SELECT type, state_key, event_id FROM current_state_events + WHERE room_id = ? AND {clause} + """ + + txn.execute(sql, [room_id, *args]) + + for typ, state_key, event_id in txn: + results[(intern_string(typ), intern_string(state_key))] = event_id + + return results + + return await self.db_pool.runInteraction( + "_get_current_state_event_ids", _get_current_state_event_ids_txn + ) + @cancellable async def get_partial_filtered_current_state_ids( self, room_id: str, state_filter: StateFilter | None = None @@ -555,15 +633,28 @@ class StateGroupWorkerStore(EventsWorkerStore, SQLBaseStore): Returns: Map from type/state_key to event ID. """ - if state_filter is None: - state_filter = StateFilter.all() + # First we check if we can delegate to one of the cached functions. + if state_filter is None or state_filter.is_full(): + return await self.get_partial_current_state_ids(room_id) + + if not state_filter.has_wildcards(): + results = StateMapWrapper(state_filter=state_filter) + + concrete_types = state_filter.concrete_types() + if not concrete_types: + # The filter matches nothing. + return results + + ids = await self._get_current_state_event_ids(room_id, concrete_types) + results.update( + (type_and_state_key, event_id) + for type_and_state_key, event_id in ids.items() + if event_id is not None + ) + return results where_clause, where_args = (state_filter).make_sql_filter_clause() - if not where_clause: - # We delegate to the cached version - return await self.get_partial_current_state_ids(room_id) - def _get_filtered_current_state_ids_txn( txn: LoggingTransaction, ) -> StateMap[str]: diff --git a/tests/storage/test_state.py b/tests/storage/test_state.py index 6c3506f348..46019c2e7b 100644 --- a/tests/storage/test_state.py +++ b/tests/storage/test_state.py @@ -30,6 +30,7 @@ from twisted.internet.testing import MemoryReactor from synapse.api.constants import EventTypes, Membership from synapse.api.room_versions import RoomVersions from synapse.events import EventBase +from synapse.logging.context import LoggingContext from synapse.server import HomeServer from synapse.types import JsonDict, RoomID, StateMap, UserID from synapse.types.state import StateFilter @@ -644,6 +645,113 @@ class StateStoreTestCase(HomeserverTestCase): ) self.assertEqual(context.state_group_before_event, groups[0][0]) + def test_get_partial_filtered_current_state_ids_concrete(self) -> None: + """A filter with no wildcards is served from the per-key cache, and the + absence of a key is cached too.""" + room_id = self.room.to_string() + + create = self.inject_state_event( + self.room, self.u_alice, EventTypes.Create, "", {} + ) + name = self.inject_state_event( + self.room, self.u_alice, EventTypes.Name, "", {"name": "test room"} + ) + + state_filter = StateFilter.from_types( + [(EventTypes.Create, ""), (EventTypes.Name, ""), (EventTypes.Topic, "")] + ) + + state = self.get_success( + self.store.get_partial_filtered_current_state_ids(room_id, state_filter) + ) + + self.assertEqual( + dict(state), + { + (EventTypes.Create, ""): create.event_id, + (EventTypes.Name, ""): name.event_id, + }, + ) + + # The room has no topic, and that fact is cached, so asking again does + # not go back to the database. + sentinel = object() + cache = self.store._get_current_state_event_id.cache + self.assertIsNone( + cache.get_immediate((room_id, (EventTypes.Topic, "")), sentinel) + ) + self.assertEqual( + cache.get_immediate((room_id, (EventTypes.Name, "")), sentinel), + name.event_id, + ) + + def test_get_partial_filtered_current_state_ids_invalidation(self) -> None: + """Persisting a new state event invalidates the cached entries for the + room.""" + room_id = self.room.to_string() + + self.inject_state_event(self.room, self.u_alice, EventTypes.Create, "", {}) + + state_filter = StateFilter.from_types([(EventTypes.Name, "")]) + + state = self.get_success( + self.store.get_partial_filtered_current_state_ids(room_id, state_filter) + ) + self.assertEqual(dict(state), {}) + + name = self.inject_state_event( + self.room, self.u_alice, EventTypes.Name, "", {"name": "test room"} + ) + + state = self.get_success( + self.store.get_partial_filtered_current_state_ids(room_id, state_filter) + ) + self.assertEqual(dict(state), {(EventTypes.Name, ""): name.event_id}) + + name2 = self.inject_state_event( + self.room, self.u_alice, EventTypes.Name, "", {"name": "renamed"} + ) + + state = self.get_success( + self.store.get_partial_filtered_current_state_ids(room_id, state_filter) + ) + self.assertEqual(dict(state), {(EventTypes.Name, ""): name2.event_id}) + + def test_get_partial_filtered_current_state_ids_uses_full_cache(self) -> None: + """Test that fetching a single state key from a room with a full cache + hits the full cache and does not go to the database.""" + + room_id = self.room.to_string() + + create = self.inject_state_event( + self.room, self.u_alice, EventTypes.Create, "", {} + ) + name = self.inject_state_event( + self.room, self.u_alice, EventTypes.Name, "", {"name": "test room"} + ) + + # prime the full cache + state_filter = StateFilter.all() + state = self.get_success( + self.store.get_partial_filtered_current_state_ids(room_id, state_filter) + ) + self.assertEqual( + dict(state), + { + (EventTypes.Create, ""): create.event_id, + (EventTypes.Name, ""): name.event_id, + }, + ) + + # now fetch a single key and check that it hits the full cache + with LoggingContext(name="test", server_name=self.hs.hostname) as ctx: + state_filter = StateFilter.from_types([(EventTypes.Name, "")]) + state = self.get_success( + self.store.get_partial_filtered_current_state_ids(room_id, state_filter) + ) + self.assertEqual(dict(state), {(EventTypes.Name, ""): name.event_id}) + self.assertEqual(ctx.get_resource_usage().db_txn_count, 0) + class CurrentStateDeltaStreamTestCase(HomeserverTestCase): def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None: