Andrew Ferrazzutti
2026-08-28 02:05:59 -04:00
committed by GitHub
co-authored by Olivier 'reivilibre'
parent 305faa2ab5
commit c3bec60936
5 changed files with 238 additions and 20 deletions
+1
View File
@@ -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.
+26 -2
View File
@@ -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
+28 -3
View File
@@ -134,6 +134,28 @@ class SendDelayedEventServlet(RestServlet):
return 200, {}
class DelayedEventServlet(RestServlet):
PATTERNS = client_patterns(
r"/org\.matrix\.msc4140/delayed_events/(?P<delay_id>[^/]+)$",
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)
@@ -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
+114
View File
@@ -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"