1#![allow(rustdoc::private_intra_doc_links)]
16#![doc = include_str!("README.md")]
17
18use std::{fmt, time::Duration};
19
20use futures_util::{StreamExt, pin_mut};
21use matrix_sdk_common::executor::spawn;
22use ruma::api::client::delayed_events::DelayParameters;
23use serde::de::{self, Deserialize, Deserializer, Visitor};
24use tokio::sync::{
25 Mutex,
26 mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel},
27};
28use tokio_stream::wrappers::UnboundedReceiverStream;
29use tokio_util::sync::{CancellationToken, DropGuard};
30
31use self::{
32 machine::{
33 Action, IncomingMessage, MatrixDriverRequestData, MatrixDriverResponse, SendEventRequest,
34 WidgetMachine,
35 },
36 matrix::MatrixDriver,
37};
38use crate::{Result, room::Room, widget::machine::DownloadFileResponse};
39
40mod capabilities;
41mod filter;
42mod machine;
43mod matrix;
44mod settings;
45
46pub use self::{
47 capabilities::{Capabilities, CapabilitiesProvider},
48 filter::{Filter, MessageLikeEventFilter, StateEventFilter, ToDeviceEventFilter},
49 settings::{
50 ClientProperties, EncryptionSystem, Intent, VirtualElementCallWidgetConfig,
51 VirtualElementCallWidgetProperties, WidgetSettings,
52 },
53};
54
55#[derive(Debug)]
58pub struct WidgetDriver {
59 settings: WidgetSettings,
60
61 from_widget_rx: UnboundedReceiver<String>,
65
66 to_widget_tx: UnboundedSender<String>,
71
72 event_forwarding_guard: Option<DropGuard>,
77}
78
79#[derive(Debug)]
82pub struct WidgetDriverHandle {
83 to_widget_rx: Mutex<UnboundedReceiver<String>>,
90
91 from_widget_tx: UnboundedSender<String>,
98}
99
100impl WidgetDriverHandle {
101 pub async fn recv(&self) -> Option<String> {
114 self.to_widget_rx.lock().await.recv().await
115 }
116
117 pub fn send(&self, message: String) -> bool {
121 self.from_widget_tx.send(message).is_ok()
122 }
123}
124
125impl WidgetDriver {
126 pub fn new(settings: WidgetSettings) -> (Self, WidgetDriverHandle) {
129 let (from_widget_tx, from_widget_rx) = unbounded_channel();
130 let (to_widget_tx, to_widget_rx) = unbounded_channel();
131
132 let driver = Self { settings, from_widget_rx, to_widget_tx, event_forwarding_guard: None };
133 let channels =
134 WidgetDriverHandle { from_widget_tx, to_widget_rx: Mutex::new(to_widget_rx) };
135
136 (driver, channels)
137 }
138
139 #[expect(clippy::result_unit_err)]
144 pub async fn run(
145 self,
146 room: Room,
147 capabilities_provider: impl CapabilitiesProvider,
148 ) -> Result<(), ()> {
149 let (incoming_msg_tx, incoming_msg_rx) = unbounded_channel();
157
158 spawn({
165 let incoming_msg_tx = incoming_msg_tx.clone();
166 let mut from_widget_rx = self.from_widget_rx;
167
168 async move {
169 while let Some(msg) = from_widget_rx.recv().await {
170 let _ = incoming_msg_tx.send(IncomingMessage::WidgetMessage(msg));
171 }
172 }
173 });
174
175 let (mut widget_machine, initial_actions) = WidgetMachine::new(
179 self.settings.widget_id().to_owned(),
180 room.room_id().to_owned(),
181 self.settings.init_on_content_load(),
182 );
183
184 let matrix_driver = MatrixDriver::new(room.clone());
185
186 let stream = UnboundedReceiverStream::new(incoming_msg_rx)
188 .flat_map(|message| tokio_stream::iter(widget_machine.process(message)));
189
190 let mut combined = tokio_stream::iter(initial_actions).chain(stream);
193
194 let to_widget_tx = self.to_widget_tx;
195 let mut event_forwarding_guard = self.event_forwarding_guard;
196
197 while let Some(action) = combined.next().await {
199 Self::process_action(
200 &to_widget_tx,
201 &mut event_forwarding_guard,
202 &matrix_driver,
203 &incoming_msg_tx,
204 &capabilities_provider,
205 action,
206 )
207 .await?;
208 }
209
210 Ok(())
211 }
212
213 async fn process_action(
215 to_widget_tx: &UnboundedSender<String>,
216 event_forwarding_guard: &mut Option<DropGuard>,
217 matrix_driver: &MatrixDriver,
218 incoming_msg_tx: &UnboundedSender<IncomingMessage>,
219 capabilities_provider: &impl CapabilitiesProvider,
220 action: Action,
221 ) -> Result<(), ()> {
222 match action {
223 Action::SendToWidget(msg) => {
224 to_widget_tx.send(msg).map_err(|_| ())?;
225 }
226
227 Action::MatrixDriverRequest { request_id, data } => {
228 let response = match data {
229 MatrixDriverRequestData::AcquireCapabilities(cmd) => {
230 let obtained = capabilities_provider
231 .acquire_capabilities(cmd.desired_capabilities)
232 .await;
233 Ok(MatrixDriverResponse::CapabilitiesAcquired(obtained))
234 }
235
236 MatrixDriverRequestData::GetOpenId => {
237 matrix_driver.get_open_id().await.map(MatrixDriverResponse::OpenIdReceived)
238 }
239
240 MatrixDriverRequestData::ReadEvents(cmd) => matrix_driver
241 .read_events(cmd.event_type.into(), cmd.state_key, cmd.limit)
242 .await
243 .map(MatrixDriverResponse::EventsRead),
244
245 MatrixDriverRequestData::ReadState(cmd) => matrix_driver
246 .read_state(cmd.event_type.into(), &cmd.state_key)
247 .await
248 .map(MatrixDriverResponse::StateRead),
249
250 MatrixDriverRequestData::SendEvent(req) => {
251 let SendEventRequest { event_type, state_key, content, delay } = req;
252 let delay_event_parameter = delay.map(|d| DelayParameters::Timeout {
257 timeout: Duration::from_millis(d),
258 });
259 matrix_driver
260 .send(event_type.into(), state_key, content, delay_event_parameter)
261 .await
262 .map(MatrixDriverResponse::EventSent)
263 }
264
265 MatrixDriverRequestData::UpdateDelayedEvent(req) => matrix_driver
266 .update_delayed_event(req.delay_id, req.action)
267 .await
268 .map(MatrixDriverResponse::DelayedEventUpdated),
269
270 MatrixDriverRequestData::SendToDeviceEvent(send_to_device_request) => {
271 matrix_driver
272 .send_to_device(
273 send_to_device_request.event_type.into(),
274 send_to_device_request.messages,
275 )
276 .await
277 .map(MatrixDriverResponse::ToDeviceSent)
278 }
279 MatrixDriverRequestData::DownloadFile(req) => matrix_driver
280 .download_attachment(req.content_uri)
281 .await
282 .map(|file_data_base64| {
283 MatrixDriverResponse::FileDownloaded(DownloadFileResponse {
284 file_data_base64,
285 })
286 }),
287
288 MatrixDriverRequestData::GetRtcTransports => matrix_driver
289 .get_rtc_transports()
290 .await
291 .map(MatrixDriverResponse::RtcTransportsReceived),
292 };
293
294 incoming_msg_tx
297 .send(IncomingMessage::MatrixDriverResponse { request_id, response })
298 .map_err(|_| ())?;
299 }
300
301 Action::Subscribe => {
302 if event_forwarding_guard.is_some() {
304 return Ok(());
305 }
306
307 let (stop_forwarding, guard) = {
308 let token = CancellationToken::new();
309 (token.child_token(), token.drop_guard())
310 };
311
312 event_forwarding_guard.replace(guard);
313
314 let mut events = matrix_driver.events();
315 let mut state_updates = matrix_driver.state_updates();
316 let to_device_events = matrix_driver.to_device_events();
317 let incoming_msg_tx = incoming_msg_tx.clone();
318
319 spawn(async move {
320 pin_mut!(to_device_events);
321
322 loop {
323 tokio::select! {
324 _ = stop_forwarding.cancelled() => {
325 return;
327 }
328
329 Some(event) = events.recv() => {
330 let _ = incoming_msg_tx.send(IncomingMessage::MatrixEventReceived(event));
332 }
333
334 Ok(state) = state_updates.recv() => {
335 let _ = incoming_msg_tx.send(IncomingMessage::StateUpdateReceived(state));
337 }
338
339 Some(event) = to_device_events.next() => {
340 let _ = incoming_msg_tx.send(IncomingMessage::ToDeviceReceived(event));
342 }
343 }
344 }
345 });
346 }
347
348 Action::Unsubscribe => {
349 event_forwarding_guard.take();
350 }
351 }
352
353 Ok(())
354 }
355}
356
357#[derive(Clone, Debug)]
359pub(crate) enum StateKeySelector {
360 Key(String),
361 Any,
362}
363
364impl<'de> Deserialize<'de> for StateKeySelector {
365 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
366 where
367 D: Deserializer<'de>,
368 {
369 struct StateKeySelectorVisitor;
370
371 impl Visitor<'_> for StateKeySelectorVisitor {
372 type Value = StateKeySelector;
373
374 fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
375 write!(f, "a string or `true`")
376 }
377
378 fn visit_bool<E>(self, v: bool) -> Result<Self::Value, E>
379 where
380 E: de::Error,
381 {
382 if v {
383 Ok(StateKeySelector::Any)
384 } else {
385 Err(E::invalid_value(de::Unexpected::Bool(v), &self))
386 }
387 }
388
389 fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
390 where
391 E: de::Error,
392 {
393 self.visit_string(v.to_owned())
394 }
395
396 fn visit_string<E>(self, v: String) -> Result<Self::Value, E>
397 where
398 E: de::Error,
399 {
400 Ok(StateKeySelector::Key(v))
401 }
402 }
403
404 deserializer.deserialize_any(StateKeySelectorVisitor)
405 }
406}
407
408#[cfg(test)]
409mod tests {
410 use assert_matches::assert_matches;
411 use serde_json::json;
412
413 use super::StateKeySelector;
414
415 #[test]
416 fn state_key_selector_from_true() {
417 let state_key = serde_json::from_value(json!(true)).unwrap();
418 assert_matches!(state_key, StateKeySelector::Any);
419 }
420
421 #[test]
422 fn state_key_selector_from_string() {
423 let state_key = serde_json::from_value(json!("test")).unwrap();
424 assert_matches!(state_key, StateKeySelector::Key(k) if k == "test");
425 }
426
427 #[test]
428 fn state_key_selector_from_false() {
429 let result = serde_json::from_value::<StateKeySelector>(json!(false));
430 assert_matches!(result, Err(e) if e.is_data());
431 }
432
433 #[test]
434 fn state_key_selector_from_number() {
435 let result = serde_json::from_value::<StateKeySelector>(json!(5));
436 assert_matches!(result, Err(e) if e.is_data());
437 }
438}