Skip to main content

tuwunel_api/client/sync/v5/extensions/
e2ee.rs

1use std::collections::HashSet;
2
3use futures::{
4	FutureExt, StreamExt, TryFutureExt,
5	future::{join, join3},
6	stream::once,
7};
8use ruma::{
9	OwnedUserId, RoomId,
10	api::client::sync::sync_events::{DeviceLists, v5::response},
11	events::{
12		StateEventType, TimelineEventType,
13		room::member::{MembershipState, RoomMemberEventContent},
14	},
15};
16use tuwunel_core::{
17	Result, error,
18	matrix::{Event, pdu::PduCount},
19	pair_of,
20	utils::{
21		BoolExt, FutureBoolExt, IterStream, ReadyExt, TryFutureExtExt,
22		future::{OptionFutureExt, OptionStream, ReadyBoolExt},
23		stream::BroadbandExt,
24	},
25};
26use tuwunel_service::sync::Connection;
27
28use super::{SyncInfo, share_encrypted_room};
29
30#[tracing::instrument(name = "e2ee", level = "trace", skip_all)]
31pub(super) async fn collect(
32	sync_info: SyncInfo<'_>,
33	conn: &Connection,
34) -> Result<response::E2EE> {
35	let SyncInfo { services, sender_user, sender_device, .. } = sync_info;
36	let Some(sender_device) = sender_device else {
37		return Ok(response::E2EE::default());
38	};
39
40	let keys_changed = services
41		.users
42		.keys_changed(sender_user, conn.globalsince, Some(conn.next_batch))
43		.map(ToOwned::to_owned)
44		.collect::<HashSet<_>>()
45		.map(|changed| (changed, HashSet::new()));
46
47	let (changed, left) = (HashSet::new(), HashSet::new());
48	let (changed, left) = services
49		.state_cache
50		.rooms_joined(sender_user)
51		.map(ToOwned::to_owned)
52		.broad_filter_map(async |room_id| collect_room(sync_info, conn, &room_id).await.ok())
53		.chain(once(keys_changed))
54		.ready_fold((changed, left), |(mut changed, mut left), room| {
55			changed.extend(room.0);
56			left.extend(room.1);
57			(changed, left)
58		})
59		.await;
60
61	let left = left
62		.into_iter()
63		.stream()
64		.filter_map(async |user_id| {
65			share_encrypted_room(services, sender_user, &user_id, None)
66				.await
67				.is_false()
68				.then_some(user_id)
69		})
70		.collect();
71
72	let device_one_time_keys_count = services
73		.users
74		.last_one_time_keys_update(sender_user)
75		.then(|since| {
76			since.gt(&conn.globalsince).then_async(|| {
77				services
78					.users
79					.count_one_time_keys(sender_user, sender_device)
80			})
81		})
82		.map(Option::unwrap_or_default);
83
84	let device_unused_fallback_key_types = services
85		.users
86		.unused_fallback_key_algorithms(sender_user, sender_device)
87		.collect::<Vec<_>>()
88		.map(Some);
89
90	let (left, device_one_time_keys_count, device_unused_fallback_key_types) =
91		join3(left, device_one_time_keys_count, device_unused_fallback_key_types)
92			.boxed()
93			.await;
94
95	Ok(response::E2EE {
96		device_one_time_keys_count,
97		device_unused_fallback_key_types,
98		device_lists: DeviceLists {
99			changed: changed.into_iter().collect(),
100			left,
101		},
102	})
103}
104
105#[tracing::instrument(level = "trace", skip_all, fields(room_id), ret)]
106async fn collect_room(
107	SyncInfo { services, sender_user, .. }: SyncInfo<'_>,
108	conn: &Connection,
109	room_id: &RoomId,
110) -> Result<pair_of!(HashSet<OwnedUserId>)> {
111	let current_shortstatehash = services
112		.state
113		.get_room_shortstatehash(room_id)
114		.inspect_err(|e| error!("Room {room_id} has no state: {e}"));
115
116	let room_keys_changed = services
117		.users
118		.room_keys_changed(room_id, conn.globalsince, Some(conn.next_batch))
119		.map(|(user_id, _)| user_id)
120		.map(ToOwned::to_owned)
121		.collect::<HashSet<_>>();
122
123	let (current_shortstatehash, device_list_changed) =
124		join(current_shortstatehash, room_keys_changed)
125			.boxed()
126			.await;
127
128	let lists = (device_list_changed, HashSet::new());
129	let Ok(current_shortstatehash) = current_shortstatehash else {
130		return Ok(lists);
131	};
132
133	if current_shortstatehash <= conn.globalsince {
134		return Ok(lists);
135	}
136
137	let Ok(since_shortstatehash) = services
138		.timeline
139		.prev_shortstatehash(room_id, PduCount::Normal(conn.globalsince).saturating_add(1))
140		.await
141	else {
142		return Ok(lists);
143	};
144
145	if since_shortstatehash == current_shortstatehash {
146		return Ok(lists);
147	}
148
149	let encrypted_at = |shortstatehash| {
150		services
151			.state_accessor
152			.state_get_shortid(shortstatehash, &StateEventType::RoomEncryption, "")
153			.is_ok()
154	};
155
156	let current_encrypted = encrypted_at(current_shortstatehash).await;
157
158	if !current_encrypted
159		&& services
160			.config
161			.device_key_update_encrypted_rooms_only
162	{
163		return Ok(lists);
164	}
165
166	let joined_since_last_sync = services
167		.state_cache
168		.get_joined_count(room_id, sender_user)
169		.map_ok_or(false, |count| count > conn.globalsince);
170
171	// A plaintext room bursts only on the sender's own join.
172	let newly_encrypted = current_encrypted
173		.then_async(|| encrypted_at(since_shortstatehash).is_false())
174		.unwrap_or_default();
175
176	// The keyed membership read leads; the state lookup trails it.
177	let members_burst = joined_since_last_sync
178		.is_false()
179		.and(newly_encrypted.is_false())
180		.is_false()
181		.await;
182
183	let joined_members_burst = members_burst.then_async(|| {
184		services
185			.state_cache
186			.room_members(room_id)
187			.ready_filter(|&user_id| user_id != sender_user)
188			.map(ToOwned::to_owned)
189			.map(|user_id| (MembershipState::Join, user_id))
190			.boxed()
191			.into_future()
192	});
193
194	services
195		.state_accessor
196		.state_added((since_shortstatehash, current_shortstatehash))
197		.broad_filter_map(async |(_shortstatekey, shorteventid)| {
198			services
199				.timeline
200				.get_pdu_from_shorteventid(shorteventid)
201				.ok()
202				.await
203		})
204		.ready_filter(|event| *event.kind() == TimelineEventType::RoomMember)
205		.ready_filter(|event| {
206			event
207				.state_key()
208				.is_some_and(|state_key| state_key != sender_user)
209		})
210		.ready_filter_map(|event| {
211			let content: RoomMemberEventContent = event.get_content().ok()?;
212			let user_id: OwnedUserId = event.state_key()?.parse().ok()?;
213
214			Some((content.membership, user_id))
215		})
216		.chain(joined_members_burst.stream())
217		.fold(lists, async |(mut changed, mut left), (membership, user_id)| {
218			use MembershipState::*;
219
220			let should_add = async |user_id| {
221				!share_encrypted_room(services, sender_user, user_id, Some(room_id)).await
222			};
223
224			match membership {
225				| Join if should_add(&user_id).await => changed.insert(user_id),
226				| Leave => left.insert(user_id),
227				| _ => false,
228			};
229
230			(changed, left)
231		})
232		.map(Ok)
233		.boxed()
234		.await
235}