1use ruma::{
19 MilliSecondsSinceUnixEpoch, OwnedEventId, UInt,
20 events::{
21 AnyMessageLikeEventContent, AnySyncMessageLikeEvent, AnySyncTimelineEvent,
22 MessageLikeEventType,
23 relation::{BundledThread, RelationType},
24 },
25 room_version_rules::RedactionRules,
26 serde::Raw,
27};
28use serde::Deserialize;
29use serde_json::value::RawValue;
30
31#[derive(Deserialize)]
32struct RelatesTo {
33 #[serde(rename = "rel_type")]
34 rel_type: RelationType,
35 #[serde(rename = "event_id")]
36 event_id: Option<OwnedEventId>,
37}
38
39#[allow(missing_debug_implementations)]
40#[derive(Deserialize)]
41struct SimplifiedContent {
42 #[serde(rename = "m.relates_to")]
43 relates_to: Option<RelatesTo>,
44}
45
46pub fn extract_thread_root_from_content(
54 content: Raw<AnyMessageLikeEventContent>,
55) -> Option<OwnedEventId> {
56 let relates_to = content.deserialize_as_unchecked::<SimplifiedContent>().ok()?.relates_to?;
57 match relates_to.rel_type {
58 RelationType::Thread => relates_to.event_id,
59 _ => None,
60 }
61}
62
63pub fn extract_thread_root(event: &Raw<AnySyncTimelineEvent>) -> Option<OwnedEventId> {
72 extract_thread_root_from_content(event.get_field("content").ok().flatten()?)
73}
74
75pub fn extract_relation(event: &Raw<AnySyncTimelineEvent>) -> Option<(RelationType, OwnedEventId)> {
78 let relates_to = event.get_field::<SimplifiedContent>("content").ok().flatten()?.relates_to?;
79 Some((relates_to.rel_type, relates_to.event_id?))
80}
81
82pub fn extract_redaction_target(
85 event: &Raw<AnySyncTimelineEvent>,
86 redaction_rules: &RedactionRules,
87) -> Option<OwnedEventId> {
88 let Ok(Some(MessageLikeEventType::RoomRedaction)) =
90 event.get_field::<MessageLikeEventType>("type")
91 else {
92 return None;
94 };
95
96 let Ok(AnySyncTimelineEvent::MessageLike(AnySyncMessageLikeEvent::RoomRedaction(redaction))) =
99 event.deserialize()
100 else {
101 return None;
103 };
104
105 redaction.redacts(redaction_rules).map(ToOwned::to_owned)
106}
107
108#[derive(Deserialize)]
109struct UnsignedRelations<Relations> {
110 #[serde(rename = "m.relations")]
111 relations: Option<Relations>,
112}
113
114#[derive(Deserialize)]
115struct ThreadRelation<BundledThread> {
116 #[serde(rename = "m.thread")]
117 thread: Option<BundledThread>,
118}
119
120pub fn extract_bundled_thread(event: &Raw<AnySyncTimelineEvent>) -> Option<BundledThread> {
122 extract_thread_relation::<BundledThread>(event)
123}
124
125pub fn extract_is_thread_root(event: &Raw<AnySyncTimelineEvent>) -> bool {
133 #[derive(Deserialize)]
134 struct LightBundledThread<'a> {
135 #[allow(unused)]
138 #[serde(borrow)]
139 latest_event: &'a RawValue,
140 #[allow(unused)]
141 count: UInt,
142 #[allow(unused)]
143 current_user_participated: bool,
144 }
145
146 extract_thread_relation::<LightBundledThread<'_>>(event).is_some()
147}
148
149fn extract_thread_relation<'de, O>(event: &'de Raw<AnySyncTimelineEvent>) -> Option<O>
150where
151 O: Deserialize<'de>,
152{
153 match event.get_field::<UnsignedRelations<ThreadRelation<O>>>("unsigned") {
154 Ok(Some(UnsignedRelations {
155 relations: Some(ThreadRelation { thread: Some(bundled_thread) }),
156 })) => Some(bundled_thread),
157 Ok(_) | Err(_) => None,
158 }
159}
160
161pub fn extract_timestamp(
167 event: &Raw<AnySyncTimelineEvent>,
168 max_value: MilliSecondsSinceUnixEpoch,
169) -> Option<MilliSecondsSinceUnixEpoch> {
170 let mut origin_server_ts = event.get_field("origin_server_ts").ok().flatten()?;
171
172 if origin_server_ts > max_value {
173 origin_server_ts = max_value;
174 }
175
176 Some(origin_server_ts)
177}
178
179#[cfg(test)]
180mod tests {
181 use std::ops::Not;
182
183 use assert_matches::assert_matches;
184 use ruma::{UInt, event_id, owned_event_id};
185 use serde_json::json;
186
187 use super::{
188 MilliSecondsSinceUnixEpoch, Raw, RelationType, extract_bundled_thread,
189 extract_is_thread_root, extract_relation, extract_thread_root, extract_timestamp,
190 };
191
192 #[test]
193 fn test_extract_thread_root() {
194 let thread_root = event_id!("$thread_root_event_id:example.com");
200 let event = Raw::new(&json!({
201 "event_id": "$eid:example.com",
202 "type": "m.room.message",
203 "sender": "@alice:example.com",
204 "origin_server_ts": 42,
205 "content": {
206 "body": "Hello, world!",
207 "m.relates_to": {
208 "rel_type": "m.thread",
209 "event_id": thread_root,
210 }
211 }
212 }))
213 .unwrap()
214 .cast_unchecked();
215
216 let observed_thread_root = extract_thread_root(&event);
217 assert_eq!(observed_thread_root.as_deref(), Some(thread_root));
218 let observed_relation = extract_relation(&event).unwrap();
219 assert_eq!(observed_relation, (RelationType::Thread, thread_root.to_owned()));
220
221 let event = Raw::new(&json!({
224 "event_id": "$eid:example.com",
225 "type": "m.room.message",
226 "sender": "@alice:example.com",
227 "origin_server_ts": 42,
228 }))
229 .unwrap()
230 .cast_unchecked();
231
232 let observed_thread_root = extract_thread_root(&event);
233 assert_matches!(observed_thread_root, None);
234 assert_matches!(extract_relation(&event), None);
235
236 let event = Raw::new(&json!({
239 "event_id": "$eid:example.com",
240 "type": "m.room.message",
241 "sender": "@alice:example.com",
242 "origin_server_ts": 42,
243 "content": {
244 "body": "Hello, world!",
245 }
246 }))
247 .unwrap()
248 .cast_unchecked();
249
250 let observed_thread_root = extract_thread_root(&event);
251 assert_matches!(observed_thread_root, None);
252 assert_matches!(extract_relation(&event), None);
253
254 let event = Raw::new(&json!({
257 "event_id": "$eid:example.com",
258 "type": "m.room.message",
259 "sender": "@alice:example.com",
260 "origin_server_ts": 42,
261 "content": {
262 "body": "Hello, world!",
263 "m.relates_to": {
264 "rel_type": "m.reference",
265 "event_id": "$referenced_event_id:example.com",
266 }
267 }
268 }))
269 .unwrap()
270 .cast_unchecked();
271
272 let observed_thread_root = extract_thread_root(&event);
273 assert_matches!(observed_thread_root, None);
274 let observed_relation = extract_relation(&event).unwrap();
275 assert_eq!(
276 observed_relation,
277 (RelationType::Reference, owned_event_id!("$referenced_event_id:example.com"))
278 );
279 }
280
281 #[test]
282 fn test_extract_bundled_thread_and_is_thread_root() {
283 let event = Raw::new(&json!({
285 "event_id": "$eid:example.com",
286 "type": "m.room.message",
287 "sender": "@alice:example.com",
288 "origin_server_ts": 42,
289 "content": {
290 "body": "Hello, world!",
291 },
292 "unsigned": {
293 "m.relations": {
294 "m.thread": {
295 "latest_event": {
296 "event_id": "$latest_event:example.com",
297 "type": "m.room.message",
298 "sender": "@bob:example.com",
299 "origin_server_ts": 42,
300 "content": {
301 "body": "Hello to you too!",
302 }
303 },
304 "count": 2,
305 "current_user_participated": true,
306 }
307 }
308 }
309 }))
310 .unwrap()
311 .cast_unchecked();
312
313 assert!(extract_bundled_thread(&event).is_some());
314 assert!(extract_is_thread_root(&event));
315
316 let event = Raw::new(&json!({
319 "event_id": "$eid:example.com",
320 "type": "m.room.message",
321 "sender": "@alice:example.com",
322 "origin_server_ts": 42,
323 }))
324 .unwrap()
325 .cast_unchecked();
326
327 assert!(extract_bundled_thread(&event).is_none());
328 assert!(extract_is_thread_root(&event).not());
329
330 let event = Raw::new(&json!({
333 "event_id": "$eid:example.com",
334 "type": "m.room.message",
335 "sender": "@alice:example.com",
336 "origin_server_ts": 42,
337 "content": {
338 "body": "Bonjour, monde!",
339 },
340 "unsigned": {
341 "m.relations": {
342 "m.replace":
343 {
344 "event_id": "$update:example.com",
345 "type": "m.room.message",
346 "sender": "@alice:example.com",
347 "origin_server_ts": 43,
348 "content": {
349 "body": "* Hello, world!",
350 }
351 },
352 }
353 }
354 }))
355 .unwrap()
356 .cast_unchecked();
357
358 assert!(extract_bundled_thread(&event).is_none());
359 assert!(extract_is_thread_root(&event).not());
360
361 let event = Raw::new(&json!({
365 "event_id": "$eid:example.com",
366 "type": "m.room.message",
367 "sender": "@alice:example.com",
368 "origin_server_ts": 42,
369 "unsigned": {
370 "m.relations": {
371 "m.thread": {
372 }
374 }
375 }
376 }))
377 .unwrap()
378 .cast_unchecked();
379
380 assert!(extract_bundled_thread(&event).is_none());
381 assert!(extract_is_thread_root(&event).not());
382 }
383
384 #[test]
385 fn test_extract_timestamp() {
386 let event = Raw::new(&json!({
387 "event_id": "$ev0",
388 "type": "m.room.message",
389 "sender": "@mnt_io:matrix.org",
390 "origin_server_ts": 42,
391 "content": {
392 "body": "Le gras, c'est la vie",
393 }
394 }))
395 .unwrap()
396 .cast_unchecked();
397
398 let timestamp = extract_timestamp(&event, MilliSecondsSinceUnixEpoch(UInt::from(100u32)));
399
400 assert_eq!(timestamp, Some(MilliSecondsSinceUnixEpoch(UInt::from(42u32))));
401 }
402
403 #[test]
404 fn test_extract_timestamp_no_origin_server_ts() {
405 let event = Raw::new(&json!({
406 "event_id": "$ev0",
407 "type": "m.room.message",
408 "sender": "@mnt_io:matrix.org",
409 "content": {
410 "body": "Le gras, c'est la vie",
411 }
412 }))
413 .unwrap()
414 .cast_unchecked();
415
416 let timestamp = extract_timestamp(&event, MilliSecondsSinceUnixEpoch(UInt::from(100u32)));
417
418 assert!(timestamp.is_none());
419 }
420
421 #[test]
422 fn test_extract_timestamp_invalid_origin_server_ts() {
423 let event = Raw::new(&json!({
424 "event_id": "$ev0",
425 "type": "m.room.message",
426 "sender": "@mnt_io:matrix.org",
427 "origin_server_ts": "saucisse",
428 "content": {
429 "body": "Le gras, c'est la vie",
430 }
431 }))
432 .unwrap()
433 .cast_unchecked();
434
435 let timestamp = extract_timestamp(&event, MilliSecondsSinceUnixEpoch(UInt::from(100u32)));
436
437 assert!(timestamp.is_none());
438 }
439
440 #[test]
441 fn test_extract_timestamp_malicious_origin_server_ts() {
442 let event = Raw::new(&json!({
443 "event_id": "$ev0",
444 "type": "m.room.message",
445 "sender": "@mnt_io:matrix.org",
446 "origin_server_ts": 101,
447 "content": {
448 "body": "Le gras, c'est la vie",
449 }
450 }))
451 .unwrap()
452 .cast_unchecked();
453
454 let timestamp = extract_timestamp(&event, MilliSecondsSinceUnixEpoch(UInt::from(100u32)));
455
456 assert_eq!(timestamp, Some(MilliSecondsSinceUnixEpoch(UInt::from(100u32))));
457 }
458}