diff --git a/changelog.d/19926.feature b/changelog.d/19926.feature new file mode 100644 index 0000000000..88f3111079 --- /dev/null +++ b/changelog.d/19926.feature @@ -0,0 +1 @@ +[MSC4140: Cancellable delayed events](https://github.com/matrix-org/matrix-spec-proposals/pull/4140): Add an endpoint for getting a single delayed event. diff --git a/synapse/handlers/delayed_events.py b/synapse/handlers/delayed_events.py index 13d6a54de2..cfaf16a30c 100644 --- a/synapse/handlers/delayed_events.py +++ b/synapse/handlers/delayed_events.py @@ -31,6 +31,8 @@ from synapse.replication.http.delayed_events import ( ) from synapse.storage.databases.main.delayed_events import ( DelayedEventDetails, + DelayedEventResponse, + DelayedEventResponseLegacyCompat, EventType, StateKey, Timestamp, @@ -549,8 +551,30 @@ class DelayedEventsHandler: else: self._next_delayed_event_call.reset(delay_duration.as_secs()) - async def get_all_for_user(self, requester: Requester) -> list[JsonDict]: - """Return all pending delayed events requested by the given user.""" + async def get_for_user( + self, requester: Requester, delay_id: str + ) -> DelayedEventResponse: + """ + Return the specified pending delayed event requested by the given user. + + Raises: + NotFoundError: if no matching delayed event could be found. + """ + await self._delayed_event_mgmt_ratelimiter.ratelimit(requester) + return await self._store.get_delayed_event_for_user( + delay_id, + requester.user.localpart, + ) + + async def get_all_for_user( + self, requester: Requester + ) -> list[DelayedEventResponseLegacyCompat]: + """ + Return all pending delayed events owned by the given user. + Includes fields from earlier revisions of MSC4140 for + compatibility with clients that still expect them. + """ + # TODO: Remove legacy fields once stable await self._delayed_event_mgmt_ratelimiter.ratelimit(requester) return await self._store.get_all_delayed_events_for_user( requester.user.localpart diff --git a/synapse/rest/client/delayed_events.py b/synapse/rest/client/delayed_events.py index 7afecffe2d..07ca07b7b9 100644 --- a/synapse/rest/client/delayed_events.py +++ b/synapse/rest/client/delayed_events.py @@ -134,6 +134,28 @@ class SendDelayedEventServlet(RestServlet): return 200, {} +class DelayedEventServlet(RestServlet): + PATTERNS = client_patterns( + r"/org\.matrix\.msc4140/delayed_events/(?P[^/]+)$", + releases=(), + ) + CATEGORY = "Delayed event management requests" + + def __init__(self, hs: "HomeServer"): + super().__init__() + self.auth = hs.get_auth() + self.delayed_events_handler = hs.get_delayed_events_handler() + + async def on_GET( + self, request: SynapseRequest, delay_id: str + ) -> tuple[int, JsonDict]: + requester = await self.auth.get_user_by_req(request) + delayed_event = await self.delayed_events_handler.get_for_user( + requester, delay_id + ) + return 200, delayed_event.asdict() + + class DelayedEventsServlet(RestServlet): PATTERNS = client_patterns( r"/org\.matrix\.msc4140/delayed_events$", @@ -150,9 +172,11 @@ class DelayedEventsServlet(RestServlet): requester = await self.auth.get_user_by_req(request) # TODO: Support Pagination stream API ("from" query parameter) delayed_events = await self.delayed_events_handler.get_all_for_user(requester) - - ret = {"delayed_events": delayed_events} - return 200, ret + return 200, { + "delayed_events": [ + delayed_event.asdict() for delayed_event in delayed_events + ] + } def register_servlets(hs: "HomeServer", http_server: HttpServer) -> None: @@ -162,4 +186,5 @@ def register_servlets(hs: "HomeServer", http_server: HttpServer) -> None: CancelDelayedEventServlet(hs).register(http_server) SendDelayedEventServlet(hs).register(http_server) RestartDelayedEventServlet(hs).register(http_server) + DelayedEventServlet(hs).register(http_server) DelayedEventsServlet(hs).register(http_server) diff --git a/synapse/storage/databases/main/delayed_events.py b/synapse/storage/databases/main/delayed_events.py index bb512611e4..35f78e3f97 100644 --- a/synapse/storage/databases/main/delayed_events.py +++ b/synapse/storage/databases/main/delayed_events.py @@ -64,6 +64,33 @@ class DelayedEventDetails(EventDetails): user_localpart: UserLocalpart +@attr.s(slots=True, frozen=True, auto_attribs=True) +class DelayedEventResponse: + """The representation of a delayed event in API format.""" + + delay_id: str + room_id: str + type: str + state_key: str | None + delay_ms: int + delayed_since_ts: int + content: JsonDict = attr.ib(converter=db_to_json) + + def asdict(self) -> JsonDict: + return attr.asdict(self, filter=lambda _attr, v: v is not None) + + +# TODO: Remove this class once the response format is stable +class DelayedEventResponseLegacyCompat(DelayedEventResponse): + """For backwards compatibility with field names from earlier revisions of MSC4140.""" + + def asdict(self) -> JsonDict: + return super().asdict() | { + "delay": self.delay_ms, + "running_since": self.delayed_since_ts, + } + + class DelayedEventsStore(SQLBaseStore): def __init__( self, @@ -290,11 +317,49 @@ class DelayedEventsStore(SQLBaseStore): _get_count_of_delayed_events, ) + async def get_delayed_event_for_user( + self, + delay_id: str, + user_localpart: str, + ) -> DelayedEventResponse: + """ + Returns the specified pending delayed event owned by the given user. + + Raises: + NotFoundError: if there is no matching delayed event. + """ + row = await self.db_pool.simple_select_one( + table="delayed_events", + keyvalues={ + "delay_id": delay_id, + "user_localpart": user_localpart, + "is_processed": False, + }, + retcols=( + "room_id", + "event_type", + "state_key", + "delay", + "send_ts - delay", + "content", + ), + allow_none=True, + desc="get_delayed_event_for_user", + ) + if row is None: + raise NotFoundError("Delayed event not found") + return DelayedEventResponse(delay_id, *row) + async def get_all_delayed_events_for_user( self, user_localpart: str, - ) -> list[JsonDict]: - """Returns all pending delayed events owned by the given user.""" + ) -> list[DelayedEventResponseLegacyCompat]: + """ + Return all pending delayed events owned by the given user. + Includes fields from earlier revisions of MSC4140 for + compatibility with clients that still expect them. + """ + # TODO: Remove legacy fields once stable # TODO: Support Pagination stream API ("next_batch" field) rows = await self.db_pool.execute( "get_all_delayed_events_for_user", @@ -305,7 +370,7 @@ class DelayedEventsStore(SQLBaseStore): event_type, state_key, delay, - send_ts, + send_ts - delay, content FROM delayed_events WHERE user_localpart = ? AND NOT is_processed @@ -313,18 +378,7 @@ class DelayedEventsStore(SQLBaseStore): """, user_localpart, ) - return [ - { - "delay_id": DelayID(row[0]), - "room_id": str(RoomID.from_string(row[1])), - "type": EventType(row[2]), - **({"state_key": StateKey(row[3])} if row[3] is not None else {}), - "delay": Delay(row[4]), - "running_since": Timestamp(row[5] - row[4]), - "content": db_to_json(row[6]), - } - for row in rows - ] + return [DelayedEventResponseLegacyCompat(*row) for row in rows] async def process_timeout_delayed_events( self, current_ts: Timestamp, reprocess_events: bool = False diff --git a/tests/rest/client/test_delayed_events.py b/tests/rest/client/test_delayed_events.py index 75d716244a..3af7d13858 100644 --- a/tests/rest/client/test_delayed_events.py +++ b/tests/rest/client/test_delayed_events.py @@ -127,6 +127,120 @@ class DelayedEventsTestCase(HomeserverTestCase): def test_delayed_events_empty_on_startup(self) -> None: self.assertListEqual([], self._get_delayed_events()) + def test_delayed_event_lookup(self) -> None: + # Schedule a message event + delay_ms = 100000 + content: JsonDict = {"message": "hello"} + delayed_since_ts = self.hs.get_clock().time_msec() + channel = self.make_request( + "POST", + _get_path_for_delayed_send(self.room_id, _EVENT_TYPE, delay_ms), + content, + self.user1_access_token, + ) + self.assertEqual(channel.code, HTTPStatus.OK, channel.result) + delay_id = channel.json_body["delay_id"] + + # Test that the scheduled delayed event can be retrieved + channel = self.make_request( + "GET", + f"{PATH_PREFIX}/{delay_id}", + access_token=self.user1_access_token, + ) + self.assertEqual(channel.code, HTTPStatus.OK, channel.result) + + # Assert the stored properties of the delayed event + event = channel.json_body + self.assertDictEqual( + event, + { + "delay_id": delay_id, + "room_id": self.room_id, + "type": _EVENT_TYPE, + "delay_ms": delay_ms, + "delayed_since_ts": delayed_since_ts, + "content": content, + }, + ) + + # Test that a non-existent delayed event cannot be found + channel = self.make_request( + "GET", + f"{PATH_PREFIX}/{delay_id}-fake", + access_token=self.user1_access_token, + ) + self.assertEqual(channel.code, HTTPStatus.NOT_FOUND, channel.result) + + # Test that other users cannot access this delayed event + channel = self.make_request( + "GET", + f"{PATH_PREFIX}/{delay_id}", + access_token=self.user2_access_token, + ) + self.assertEqual(channel.code, HTTPStatus.NOT_FOUND, channel.result) + + # Now schedule a state event. + # Do it in this test, as opposed to a new one, to confirm that the correct delayed event + # is retrieved when multiple delayed events have been scheduled. + delay_ms += 2000 + state_key = "" + state_event_type = _EVENT_TYPE + "_state" + content = {"state_message": "greetings"} + delayed_since_ts = self.hs.get_clock().time_msec() + channel = self.make_request( + "PUT", + _get_path_for_delayed_state( + self.room_id, state_event_type, state_key, delay_ms + ), + content, + self.user1_access_token, + ) + self.assertEqual(channel.code, HTTPStatus.OK, channel.result) + delay_id_2 = channel.json_body["delay_id"] + + # Test that the new delayed event has a different ID from the previous one + self.assertNotEqual(delay_id, delay_id_2) + + # Test that the scheduled delayed event can be retrieved + channel = self.make_request( + "GET", + f"{PATH_PREFIX}/{delay_id_2}", + access_token=self.user1_access_token, + ) + self.assertEqual(channel.code, HTTPStatus.OK, channel.result) + + # Assert the stored properties of the delayed event + state_event = channel.json_body + self.assertDictEqual( + state_event, + { + "delay_id": delay_id_2, + "room_id": self.room_id, + "type": state_event_type, + "state_key": state_key, + "delay_ms": delay_ms, + "delayed_since_ts": delayed_since_ts, + "content": content, + }, + ) + + # Test that the list lookup retrieves the same items (with legacy fields included) + self.assertEqual( + self._get_delayed_events(), + [ + event + | { + "delay": event["delay_ms"], + "running_since": event["delayed_since_ts"], + }, + state_event + | { + "delay": state_event["delay_ms"], + "running_since": state_event["delayed_since_ts"], + }, + ], + ) + def test_delayed_state_events_are_sent_on_timeout(self) -> None: state_key = "to_send_on_timeout"