1use std::{
2 collections::BTreeMap,
3 iter::once,
4 net::IpAddr,
5 sync::atomic::{AtomicBool, Ordering},
6 time::{Duration, Instant},
7};
8
9use axum::extract::State;
10use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt};
11use ruma::{
12 CanonicalJsonObject, CanonicalJsonValue, MilliSecondsSinceUnixEpoch, OwnedDeviceId,
13 OwnedEventId, OwnedRoomId, OwnedUserId, RoomId, ServerName, TransactionId, UserId,
14 api::{
15 error::ErrorKind,
16 federation::transactions::{
17 edu::{
18 DeviceListUpdateContent, DirectDeviceContent, Edu, PresenceContent,
19 PresenceUpdate, ReceiptContent, ReceiptData, ReceiptMap, SigningKeyUpdateContent,
20 TypingContent,
21 },
22 send_transaction_message,
23 },
24 },
25 events::receipt::{ReceiptEvent, ReceiptEventContent, ReceiptType},
26 int,
27 serde::Raw,
28 to_device::DeviceIdOrAllDevices,
29 uint,
30};
31use tuwunel_core::{
32 Err, Error, Result, debug,
33 debug::INFO_SPAN_LEVEL,
34 debug_warn, defer, err, error,
35 itertools::Itertools,
36 result::LogErr,
37 smallvec::SmallVec,
38 trace,
39 utils::{
40 debug::str_truncated,
41 future::TryExtExt,
42 millis_since_unix_epoch,
43 stream::{BroadbandExt, IterStream, ReadyExt, TryBroadbandExt, automatic_width},
44 },
45 warn,
46};
47use tuwunel_service::{
48 Services,
49 rooms::state_res::{is_topologically_sorted_in_place, topological_sort},
50 sending::{EDU_LIMIT, PDU_LIMIT},
51 users::DeviceListChange,
52};
53
54use crate::{ClientIp, Ruma};
55
56type ResolvedMap = BTreeMap<OwnedEventId, Result>;
57type RoomsPdus = SmallVec<[RoomPdus; 1]>;
58type RoomPdus = (OwnedRoomId, TxnPdus);
59type TxnPdus = SmallVec<[(usize, Pdu); 1]>;
60type Pdu = (OwnedRoomId, OwnedEventId, CanonicalJsonObject);
61
62type Deliveries = SmallVec<[(OwnedDeviceId, u64); 1]>;
65
66#[tracing::instrument(
70 name = "txn",
71 level = INFO_SPAN_LEVEL,
72 skip_all,
73 fields(
74 txn = str_truncated(body.transaction_id.as_str(), 20),
75 origin = body.origin().as_str(),
76 %client,
77 ),
78)]
79pub(crate) async fn send_transaction_message_route(
80 State(services): State<crate::State>,
81 ClientIp(client): ClientIp,
82 body: Ruma<send_transaction_message::v1::Request>,
83) -> Result<send_transaction_message::v1::Response> {
84 if body.origin() != body.body.origin {
85 return Err!(Request(Forbidden(
86 "Not allowed to send transactions on behalf of other servers"
87 )));
88 }
89
90 if body.pdus.len() > PDU_LIMIT {
91 return Err!(Request(Forbidden(
92 "Not allowed to send more than {PDU_LIMIT} PDUs in one transaction"
93 )));
94 }
95
96 if body.edus.len() > EDU_LIMIT {
97 return Err!(Request(Forbidden(
98 "Not allowed to send more than {EDU_LIMIT} EDUs in one transaction"
99 )));
100 }
101
102 services
104 .sending
105 .notify_peer_alive(body.origin())
106 .await;
107
108 let txn_start_time = Instant::now();
109 trace!(
110 pdus = body.pdus.len(),
111 edus = body.edus.len(),
112 elapsed = ?txn_start_time.elapsed(),
113 "Starting txn",
114 );
115
116 let pdus = body
117 .pdus
118 .iter()
119 .stream()
120 .enumerate()
121 .broad_filter_map(|(i, pdu)| {
122 services
123 .event_handler
124 .parse_incoming_pdu(pdu)
125 .inspect_err(move |e| debug_warn!("Could not parse PDU[{i}]: {e}"))
126 .map_ok(move |pdu| (i, pdu))
127 .ok()
128 });
129
130 let edus = body
131 .edus
132 .iter()
133 .stream()
134 .enumerate()
135 .ready_filter_map(|(i, edu)| {
136 serde_json::from_str(edu.json().get())
137 .inspect_err(|e| debug_warn!("Could not parse EDU[{i}]: {e}"))
138 .map(|edu| (i, edu))
139 .ok()
140 });
141
142 let results = handle(
143 &services,
144 &client,
145 body.origin(),
146 &body.transaction_id,
147 txn_start_time,
148 pdus,
149 edus,
150 )
151 .await?;
152
153 debug!(
154 pdus = body.pdus.len(),
155 edus = body.edus.len(),
156 elapsed = ?txn_start_time.elapsed(),
157 "Finished txn",
158 );
159
160 for (id, result) in &results {
161 if let Err(e) = result
162 && matches!(e, Error::BadRequest(ErrorKind::NotFound, _))
163 {
164 warn!("Incoming PDU failed {id}: {e:?}");
165 }
166 }
167
168 Ok(send_transaction_message::v1::Response {
169 pdus: results
170 .into_iter()
171 .map(|(e, r)| (e, r.map_err(error::sanitized_message)))
172 .collect(),
173 })
174}
175
176async fn handle(
177 services: &Services,
178 client: &IpAddr,
179 origin: &ServerName,
180 txn_id: &TransactionId,
181 started: Instant,
182 pdus: impl Stream<Item = (usize, Pdu)> + Send,
183 edus: impl Stream<Item = (usize, Edu)> + Send,
184) -> Result<ResolvedMap> {
185 let results = handle_pdus(services, client, origin, txn_id, started, pdus).await?;
186
187 handle_edus(services, client, origin, txn_id, edus).await?;
188
189 Ok(results)
190}
191
192async fn handle_pdus(
193 services: &Services,
194 client: &IpAddr,
195 origin: &ServerName,
196 txn_id: &TransactionId,
197 started: Instant,
198 pdus: impl Stream<Item = (usize, Pdu)> + Send,
199) -> Result<ResolvedMap> {
200 pdus.collect()
201 .map(Ok)
202 .map_ok(|pdus: TxnPdus| {
203 pdus.into_iter()
204 .sorted_by(|(_, (room_a, ..)), (_, (room_b, ..))| room_a.cmp(room_b))
205 .into_grouping_map_by(|(_, (room_id, ..))| room_id.clone())
206 .collect()
207 .into_iter()
208 .try_stream()
209 })
210 .try_flatten_stream()
211 .try_collect::<RoomsPdus>()
212 .map_ok(IntoIterator::into_iter)
213 .map_ok(IterStream::try_stream)
214 .try_flatten_stream()
215 .broad_and_then(async |(room_id, pdus)| {
216 handle_room(services, client, origin, txn_id, started, room_id, pdus)
217 .map_ok(ResolvedMap::into_iter)
218 .map_ok(IterStream::try_stream)
219 .await
220 })
221 .try_flatten()
222 .try_collect()
223 .await
224}
225
226#[tracing::instrument(
227 name = "room",
228 level = INFO_SPAN_LEVEL,
229 skip_all,
230 fields(%room_id)
231)]
232async fn handle_room(
233 services: &Services,
234 _client: &IpAddr,
235 origin: &ServerName,
236 txn_id: &TransactionId,
237 txn_start_time: Instant,
238 ref room_id: OwnedRoomId,
239 pdus: TxnPdus,
240) -> Result<ResolvedMap> {
241 let pdus = sort_pdus(pdus).await;
242
243 services
244 .event_handler
245 .mutex_federation
246 .lock(room_id)
247 .then(async |_lock| {
248 pdus.into_iter()
249 .enumerate()
250 .try_stream()
251 .and_then(async |pdu| {
252 services.server.check_running().map(|()| pdu) })
254 .and_then(|(ri, (ti, (room_id, event_id, value)))| {
255 let meta = (origin, txn_id, txn_start_time, ti);
256 let pdu = (ri, (room_id, event_id, value));
257 handle_pdu(services, meta, pdu).map(Ok)
258 })
259 .try_collect()
260 .await
261 })
262 .await
263}
264
265async fn sort_pdus(mut pdus: TxnPdus) -> TxnPdus {
270 if already_sorted(&pdus) {
271 return pdus;
272 }
273
274 let event_ids: BTreeMap<&str, &OwnedEventId> = pdus
275 .iter()
276 .map(|(_, (_, event_id, _))| (event_id.as_str(), event_id))
277 .collect();
278
279 let graph = pdus
280 .iter()
281 .map(|(_, (_, event_id, value))| {
282 let references = prev_event_ids(value)
283 .filter_map(|prev| event_ids.get(prev).copied())
284 .map(ToOwned::to_owned)
285 .collect();
286
287 (event_id.clone(), references)
288 })
289 .collect();
290
291 let query = async |_event_id: OwnedEventId| {
293 Ok((int!(0).into(), MilliSecondsSinceUnixEpoch(uint!(0))))
294 };
295
296 let Ok(order) = topological_sort(graph, &query).await else {
297 return pdus;
298 };
299
300 let position: BTreeMap<&str, usize> = order
301 .iter()
302 .enumerate()
303 .map(|(i, event_id)| (event_id.as_str(), i))
304 .collect();
305
306 pdus.sort_by_key(|(_, (_, event_id, _))| position.get(event_id.as_str()).copied());
307 pdus
308}
309
310fn already_sorted(pdus: &[(usize, Pdu)]) -> bool {
313 is_topologically_sorted_in_place(
314 pdus,
315 |(_, (_, id, _))| id.as_str(),
316 |(_, (_, _, value))| prev_event_ids(value),
317 )
318}
319
320fn prev_event_ids(value: &CanonicalJsonObject) -> impl Iterator<Item = &str> + '_ {
322 value
323 .get("prev_events")
324 .and_then(CanonicalJsonValue::as_array)
325 .into_iter()
326 .flatten()
327 .filter_map(CanonicalJsonValue::as_str)
328}
329
330#[tracing::instrument(
331 name = "pdu",
332 level = INFO_SPAN_LEVEL,
333 skip_all,
334 fields(%event_id, %ti, %ri)
335)]
336async fn handle_pdu(
337 services: &Services,
338 (origin, txn_id, txn_start_time, ti): (&ServerName, &TransactionId, Instant, usize),
339 (ri, (ref room_id, event_id, value)): (usize, Pdu),
340) -> (OwnedEventId, Result) {
341 let pdu_start_time = Instant::now();
342 let completed: AtomicBool = Default::default();
343 defer! {{
344 if completed.load(Ordering::Acquire) {
345 return;
346 }
347
348 if pdu_start_time.elapsed() >= Duration::from_secs(services.config.client_request_timeout) {
349 error!(
350 %origin, %txn_id, %room_id, %event_id, %ri, %ti,
351 elapsed = ?pdu_start_time.elapsed(),
352 "Incoming transaction processing timed out.",
353 );
354 } else {
355 debug_warn!(
356 %origin, %txn_id, %room_id, %event_id, %ri, %ti,
357 elapsed = ?pdu_start_time.elapsed(),
358 "Incoming transaction processing interrupted.",
359 );
360 }
361 }}
362
363 let result = services
364 .event_handler
365 .handle_incoming_pdu(origin, room_id, &event_id, value, true)
366 .map_ok(|_| ())
367 .await;
368
369 completed.store(true, Ordering::Release);
370 debug!(
371 %event_id, ri, ti,
372 pdu_elapsed = ?pdu_start_time.elapsed(),
373 txn_elapsed = ?txn_start_time.elapsed(),
374 "Finished PDU",
375 );
376
377 (event_id.clone(), result)
378}
379
380#[tracing::instrument(name = "edus", level = "debug", skip_all)]
381async fn handle_edus(
382 services: &Services,
383 client: &IpAddr,
384 origin: &ServerName,
385 txn_id: &TransactionId,
386 edus: impl Stream<Item = (usize, Edu)> + Send,
387) -> Result {
388 edus.for_each_concurrent(automatic_width(), |(i, edu)| {
389 handle_edu(services, client, origin, txn_id, i, edu)
390 })
391 .await;
392
393 Ok(())
394}
395
396#[tracing::instrument(
397 name = "edu",
398 level = "debug",
399 skip_all,
400 fields(%i),
401)]
402async fn handle_edu(
403 services: &Services,
404 client: &IpAddr,
405 origin: &ServerName,
406 _txn_id: &TransactionId,
407 i: usize,
408 edu: Edu,
409) {
410 match edu {
411 | Edu::Presence(presence) if services.server.config.allow_incoming_presence =>
412 handle_edu_presence(services, client, origin, presence).await,
413
414 | Edu::Receipt(receipt)
415 if services
416 .server
417 .config
418 .allow_incoming_read_receipts =>
419 handle_edu_receipt(services, client, origin, receipt).await,
420
421 | Edu::Typing(typing) if services.server.config.allow_incoming_typing =>
422 handle_edu_typing(services, client, origin, typing).await,
423
424 | Edu::DeviceListUpdate(content) =>
425 handle_edu_device_list_update(services, client, origin, content).await,
426
427 | Edu::DirectToDevice(content) =>
428 handle_edu_direct_to_device(services, client, origin, content).await,
429
430 | Edu::SigningKeyUpdate(content) =>
431 handle_edu_signing_key_update(services, client, origin, content).await,
432
433 | Edu::_Custom(ref _custom) => debug_warn!(?i, ?edu, "received custom/unknown EDU"),
434
435 | _ => trace!(?i, ?edu, "skipped"),
436 }
437}
438
439async fn handle_edu_presence(
440 services: &Services,
441 _client: &IpAddr,
442 origin: &ServerName,
443 presence: PresenceContent,
444) {
445 presence
446 .push
447 .into_iter()
448 .stream()
449 .for_each_concurrent(automatic_width(), |update| {
450 handle_edu_presence_update(services, origin, update)
451 })
452 .await;
453}
454
455async fn handle_edu_presence_update(
456 services: &Services,
457 origin: &ServerName,
458 update: PresenceUpdate,
459) {
460 if update.user_id.server_name() != origin {
461 debug_warn!(
462 %update.user_id, %origin,
463 "received presence EDU for user not belonging to origin"
464 );
465 return;
466 }
467
468 services
469 .presence
470 .set_presence_from_federation(
471 &update.user_id,
472 &update.presence,
473 update.currently_active,
474 update.last_active_ago,
475 update.status_msg.clone(),
476 )
477 .await
478 .log_err()
479 .ok();
480}
481
482async fn handle_edu_receipt(
483 services: &Services,
484 _client: &IpAddr,
485 origin: &ServerName,
486 receipt: ReceiptContent,
487) {
488 receipt
489 .receipts
490 .into_iter()
491 .stream()
492 .for_each_concurrent(automatic_width(), |(room_id, room_updates)| {
493 handle_edu_receipt_room(services, origin, room_id, room_updates)
494 })
495 .await;
496}
497
498async fn handle_edu_receipt_room(
499 services: &Services,
500 origin: &ServerName,
501 room_id: OwnedRoomId,
502 room_updates: ReceiptMap,
503) {
504 if services
505 .event_handler
506 .acl_check(origin, &room_id)
507 .await
508 .is_err()
509 {
510 debug_warn!(
511 %origin, %room_id,
512 "received read receipt EDU from ACL'd server"
513 );
514 return;
515 }
516
517 let room_id = &room_id;
518 room_updates
519 .read
520 .into_iter()
521 .stream()
522 .for_each_concurrent(automatic_width(), async |(user_id, user_updates)| {
523 handle_edu_receipt_room_user(services, origin, room_id, &user_id, user_updates).await;
524 })
525 .await;
526}
527
528async fn handle_edu_receipt_room_user(
529 services: &Services,
530 origin: &ServerName,
531 room_id: &RoomId,
532 user_id: &UserId,
533 user_updates: ReceiptData,
534) {
535 if user_id.server_name() != origin {
536 debug_warn!(
537 %user_id, %origin,
538 "received read receipt EDU for user not belonging to origin"
539 );
540 return;
541 }
542
543 if !services
544 .state_cache
545 .server_in_room(origin, room_id)
546 .await
547 {
548 debug_warn!(
549 %user_id, %room_id, %origin,
550 "received read receipt EDU from server who does not have a member in the room",
551 );
552 return;
553 }
554
555 let data = &user_updates.data;
556 user_updates
557 .event_ids
558 .into_iter()
559 .stream()
560 .for_each_concurrent(automatic_width(), async |event_id| {
561 let user_data = [(user_id.to_owned(), data.clone())];
562 let receipts = [(ReceiptType::Read, BTreeMap::from(user_data))];
563 let content = [(event_id.clone(), BTreeMap::from(receipts))];
564 services
565 .read_receipt
566 .readreceipt_update(user_id, room_id, &ReceiptEvent {
567 content: ReceiptEventContent(content.into()),
568 room_id: room_id.to_owned(),
569 })
570 .await;
571 })
572 .await;
573}
574
575async fn handle_edu_typing(
576 services: &Services,
577 _client: &IpAddr,
578 origin: &ServerName,
579 typing: TypingContent,
580) {
581 if typing.user_id.server_name() != origin {
582 debug_warn!(
583 %typing.user_id, %origin,
584 "received typing EDU for user not belonging to origin"
585 );
586 return;
587 }
588
589 if services
590 .event_handler
591 .acl_check(typing.user_id.server_name(), &typing.room_id)
592 .await
593 .is_err()
594 {
595 debug_warn!(
596 %typing.user_id, %typing.room_id, %origin,
597 "received typing EDU for ACL'd user's server"
598 );
599 return;
600 }
601
602 if !services
603 .state_cache
604 .is_joined(&typing.user_id, &typing.room_id)
605 .await
606 {
607 debug_warn!(
608 %typing.user_id, %typing.room_id, %origin,
609 "received typing EDU for user not in room"
610 );
611 return;
612 }
613
614 if typing.typing {
615 let secs = services.server.config.typing_federation_timeout_s;
616 let timeout = millis_since_unix_epoch().saturating_add(secs.saturating_mul(1000));
617
618 services
619 .typing
620 .typing_add(&typing.user_id, &typing.room_id, timeout)
621 .await
622 .log_err()
623 .ok();
624 } else {
625 services
626 .typing
627 .typing_remove(&typing.user_id, &typing.room_id)
628 .await
629 .log_err()
630 .ok();
631 }
632}
633
634async fn handle_edu_device_list_update(
635 services: &Services,
636 _client: &IpAddr,
637 origin: &ServerName,
638 content: DeviceListUpdateContent,
639) {
640 let DeviceListUpdateContent { user_id, .. } = content;
641
642 if user_id.server_name() != origin {
643 debug_warn!(
644 %user_id, %origin,
645 "received device list update EDU for user not belonging to origin"
646 );
647 return;
648 }
649
650 services
651 .users
652 .mark_device_key_update(&user_id, DeviceListChange::Resync)
653 .await;
654}
655
656async fn handle_edu_direct_to_device(
657 services: &Services,
658 _client: &IpAddr,
659 origin: &ServerName,
660 content: DirectDeviceContent,
661) {
662 let DirectDeviceContent {
663 ref sender,
664 ref ev_type,
665 ref message_id,
666 messages,
667 } = content;
668
669 if sender.server_name() != origin {
670 debug_warn!(
671 %sender, %origin,
672 "received direct to device EDU for user not belonging to origin"
673 );
674 return;
675 }
676
677 if services
679 .transaction_ids
680 .existing_txnid(sender, None, message_id)
681 .await
682 .is_ok()
683 {
684 return;
685 }
686
687 let ev_type = ev_type.to_string();
688
689 messages
690 .into_iter()
691 .stream()
692 .broad_filter_map(async |(target_user_id, map)| {
693 to_device_deliverable(services, &target_user_id)
694 .await
695 .then_some((target_user_id, map))
696 })
697 .for_each_concurrent(automatic_width(), |(target_user_id, map)| {
698 handle_edu_direct_to_device_user(services, target_user_id, sender, &ev_type, map)
699 })
700 .await;
701
702 services
704 .transaction_ids
705 .add_txnid(sender, None, message_id, &[]);
706}
707
708async fn to_device_deliverable(services: &Services, user_id: &UserId) -> bool {
713 services.globals.user_is_local(user_id)
714 && ((user_id == services.globals.server_user && services.users.exists(user_id).await)
715 || services.users.is_active(user_id).await
716 || services
717 .appservice
718 .is_interested_in_user(user_id)
719 .await)
720}
721
722async fn handle_edu_direct_to_device_user<Event: Send + Sync>(
723 services: &Services,
724 target_user_id: OwnedUserId,
725 sender: &UserId,
726 ev_type: &str,
727 map: BTreeMap<DeviceIdOrAllDevices, Raw<Event>>,
728) {
729 map.into_iter()
730 .stream()
731 .ready_filter_map(|(tid, raw)| {
732 raw.deserialize_as()
733 .map_err(|e| {
734 err!(Request(InvalidParam(error!("To-Device event is invalid: {e}"))))
735 })
736 .ok()
737 .map(|ev| (tid, ev))
738 })
739 .for_each_concurrent(automatic_width(), |(tid, ev)| {
740 handle_edu_direct_to_device_event(services, &target_user_id, sender, tid, ev_type, ev)
741 })
742 .await;
743}
744
745async fn handle_edu_direct_to_device_event(
746 services: &Services,
747 target_user_id: &UserId,
748 sender: &UserId,
749 target_device_id_maybe: DeviceIdOrAllDevices,
750 ev_type: &str,
751 event: serde_json::Value,
752) {
753 match target_device_id_maybe {
754 | DeviceIdOrAllDevices::DeviceId(ref target_device_id) => {
755 let count = services.users.add_to_device_event(
756 sender,
757 target_user_id,
758 target_device_id,
759 ev_type,
760 &event,
761 );
762
763 services
764 .sending
765 .send_to_device_appservices(
766 sender,
767 target_user_id,
768 once((&**target_device_id, count)),
769 ev_type,
770 &event,
771 )
772 .await
773 .log_err()
774 .ok();
775 },
776
777 | DeviceIdOrAllDevices::AllDevices => {
778 let interested = services
779 .appservice
780 .is_interested_in_user(target_user_id)
781 .await;
782
783 let deliveries: Deliveries = services
784 .users
785 .all_device_ids(target_user_id)
786 .map(|target_device_id| {
787 let count = services.users.add_to_device_event(
788 sender,
789 target_user_id,
790 target_device_id,
791 ev_type,
792 &event,
793 );
794
795 (target_device_id, count)
796 })
797 .ready_filter_map(|(target_device_id, count)| {
798 interested.then(|| (target_device_id.to_owned(), count))
799 })
800 .collect()
801 .await;
802
803 if !deliveries.is_empty() {
804 services
805 .sending
806 .send_to_device_appservices(
807 sender,
808 target_user_id,
809 deliveries
810 .iter()
811 .map(|(device_id, count)| (&**device_id, *count)),
812 ev_type,
813 &event,
814 )
815 .await
816 .log_err()
817 .ok();
818 }
819 },
820 }
821}
822
823async fn handle_edu_signing_key_update(
824 services: &Services,
825 _client: &IpAddr,
826 origin: &ServerName,
827 content: SigningKeyUpdateContent,
828) {
829 let SigningKeyUpdateContent { user_id, master_key, self_signing_key } = content;
830
831 if user_id.server_name() != origin {
832 debug_warn!(
833 %user_id, %origin,
834 "received signing key update EDU from server that does not belong to user's server"
835 );
836 return;
837 }
838
839 services
840 .users
841 .add_cross_signing_keys(&user_id, &master_key, &self_signing_key, &None, true)
842 .await
843 .log_err()
844 .ok();
845}
846
847#[cfg(test)]
848mod tests {
849 use ruma::{CanonicalJsonObject, OwnedEventId, event_id, room_id};
850 use serde_json::json;
851
852 use super::{Pdu, TxnPdus, already_sorted, prev_event_ids, sort_pdus};
853
854 fn pdu(index: usize, id: &OwnedEventId, prev: &[&OwnedEventId]) -> (usize, Pdu) {
855 let prev_events: Vec<&str> = prev.iter().map(|e| e.as_str()).collect();
856 let value: CanonicalJsonObject =
857 serde_json::from_value(json!({ "prev_events": prev_events }))
858 .expect("valid canonical json");
859
860 (index, (room_id!("!r:example.com").to_owned(), id.clone(), value))
861 }
862
863 fn ids() -> (OwnedEventId, OwnedEventId, OwnedEventId) {
864 (
865 event_id!("$a:example.com").to_owned(),
866 event_id!("$b:example.com").to_owned(),
867 event_id!("$c:example.com").to_owned(),
868 )
869 }
870
871 fn order(pdus: &[(usize, Pdu)]) -> Vec<&str> {
872 pdus.iter()
873 .map(|(_, (_, id, _))| id.as_str())
874 .collect()
875 }
876
877 #[test]
878 fn sorted_when_parents_lead() {
879 let (a, b, c) = ids();
880 let pdus = [pdu(0, &a, &[]), pdu(1, &b, &[&a]), pdu(2, &c, &[&b])];
881
882 assert!(already_sorted(&pdus));
883 }
884
885 #[test]
886 fn unsorted_when_child_leads() {
887 let (a, b, _c) = ids();
888 let pdus = [pdu(0, &b, &[&a]), pdu(1, &a, &[])];
889
890 assert!(!already_sorted(&pdus));
891 }
892
893 #[test]
894 fn sorted_ignores_out_of_batch_references() {
895 let (a, b, c) = ids();
896 let pdus = [pdu(0, &b, &[&c]), pdu(1, &a, &[&c])];
897
898 assert!(already_sorted(&pdus));
899 }
900
901 #[tokio::test]
902 async fn sort_orders_parents_before_children() {
903 let (a, b, c) = ids();
904 let pdus: TxnPdus = [pdu(0, &c, &[&b]), pdu(1, &b, &[&a]), pdu(2, &a, &[])]
905 .into_iter()
906 .collect();
907
908 let sorted = sort_pdus(pdus).await;
909
910 assert_eq!(order(&sorted), ["$a:example.com", "$b:example.com", "$c:example.com"]);
911 }
912
913 #[tokio::test]
914 async fn sort_is_noop_when_already_ordered() {
915 let (a, b, c) = ids();
916 let pdus: TxnPdus = [pdu(0, &a, &[]), pdu(1, &b, &[&a]), pdu(2, &c, &[&b])]
917 .into_iter()
918 .collect();
919
920 let sorted = sort_pdus(pdus.clone()).await;
921
922 assert_eq!(order(&sorted), order(&pdus));
923 }
924
925 #[tokio::test]
926 async fn sort_preserves_duplicates() {
927 let (a, b, _c) = ids();
928 let pdus: TxnPdus = [pdu(0, &b, &[&a]), pdu(1, &a, &[]), pdu(2, &b, &[&a])]
929 .into_iter()
930 .collect();
931
932 let sorted = sort_pdus(pdus).await;
933
934 assert_eq!(sorted.len(), 3);
935 }
936
937 #[tokio::test]
938 async fn sort_preserves_a_cycle() {
939 let (a, b, _c) = ids();
940 let pdus: TxnPdus = [pdu(0, &a, &[&b]), pdu(1, &b, &[&a])]
941 .into_iter()
942 .collect();
943
944 let sorted = sort_pdus(pdus).await;
945
946 assert_eq!(sorted.len(), 2);
947 }
948
949 #[test]
950 fn prev_event_ids_reads_the_array() {
951 let (a, b, _c) = ids();
952 let (_, (_, _, value)) = pdu(0, &a, &[&b]);
953
954 let prev: Vec<&str> = prev_event_ids(&value).collect();
955
956 assert_eq!(prev, ["$b:example.com"]);
957 }
958
959 #[test]
960 fn prev_event_ids_empty_when_absent() {
961 let value = CanonicalJsonObject::new();
962
963 assert_eq!(prev_event_ids(&value).count(), 0);
964 }
965}