mirror of
https://github.com/element-hq/synapse.git
synced 2026-09-18 16:34:44 +00:00
MSC4140: support getting a single delayed event (#19926)
See https://github.com/matrix-org/matrix-spec-proposals/blob/toger5/expiring-events-keep-alive/proposals/4140-delayed-events-futures.md#getting-a-single-delayed-event Co-authored-by: Olivier 'reivilibre' <olivier@librepush.net>
This commit is contained in:
co-authored by
Olivier 'reivilibre'
parent
305faa2ab5
commit
c3bec60936
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user