Skip to main content

tuwunel_api/server/
send.rs

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
62/// Recipient devices of one `AllDevices` to-device send paired with their
63/// inbox counts.
64type Deliveries = SmallVec<[(OwnedDeviceId, u64); 1]>;
65
66/// # `PUT /_matrix/federation/v1/send/{txnId}`
67///
68/// Push EDUs and PDUs to this server.
69#[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	// Clear any failure bucket before processing consults the peer gate.
103	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) // interruption point
253				})
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
265/// Reorder a room's transaction PDUs so each event follows the in-batch events
266/// it references. An already-ordered batch is returned unchanged; references to
267/// events outside the batch are non-edges. The sort is an optimization, so a
268/// failure falls back to the arrival order.
269async 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	// Causal order alone matters here, so the tie-break inputs are constant.
292	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
310/// Whether the batch is already in causal order, in which case the sort can be
311/// skipped.
312fn 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
320/// The `prev_events` of a PDU held as canonical JSON.
321fn 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	// Check if this is a new transaction id
678	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	// Save transaction id with empty data
703	services
704		.transaction_ids
705		.add_txnid(sender, None, message_id, &[]);
706}
707
708/// A local account we store or forward to-device events for.
709///
710/// Qualifying accounts are active, the server user once its account exists, or
711/// claimed by an appservice namespace so its puppet events reach the bridge.
712async 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}