diff --git a/crates/matrix-sdk/src/widget/machine/driver_req.rs b/crates/matrix-sdk/src/widget/machine/driver_req.rs index dfa6c4e20..c175eb16b 100644 --- a/crates/matrix-sdk/src/widget/machine/driver_req.rs +++ b/crates/matrix-sdk/src/widget/machine/driver_req.rs @@ -26,10 +26,10 @@ use serde::Deserialize; use serde_json::Value as JsonValue; use tracing::error; -use super::{incoming::MatrixDriverResponse, MatrixDriverRequestMeta, WidgetMachine}; +use super::{incoming::MatrixDriverResponse, Action, MatrixDriverRequestMeta, WidgetMachine}; use crate::widget::{Capabilities, StateKeySelector}; -#[derive(Debug)] +#[derive(Clone, Debug)] pub(crate) enum MatrixDriverRequestData { /// Acquire capabilities from the user given the set of desired /// capabilities. @@ -65,18 +65,18 @@ where Self { request_meta: Some(request_meta), _phantom: PhantomData } } - pub(crate) fn null() -> Self { - Self { request_meta: None, _phantom: PhantomData } - } - pub(crate) fn then( self, - response_handler: impl FnOnce(Result, &mut WidgetMachine) + Send + 'static, + response_handler: impl FnOnce(Result, &mut WidgetMachine) -> Vec + + Send + + 'static, ) { if let Some(request_meta) = self.request_meta { request_meta.response_fn = Some(Box::new(move |response, machine| { if let Some(response_data) = response.map(T::from_response).transpose() { response_handler(response_data, machine) + } else { + Vec::new() } })); } @@ -94,7 +94,7 @@ pub(crate) trait FromMatrixDriverResponse: Sized { /// Ask the client (capability provider) to acquire given capabilities /// from the user. The client must eventually respond with granted capabilities. -#[derive(Debug)] +#[derive(Clone, Debug)] pub(crate) struct AcquireCapabilities { pub(crate) desired_capabilities: Capabilities, } @@ -149,7 +149,7 @@ impl FromMatrixDriverResponse for request_openid_token::v3::Response { /// Ask the client to read matrix event(s) that corresponds to the given /// description and return a list of events as a response. -#[derive(Debug)] +#[derive(Clone, Debug)] pub(crate) struct ReadMessageLikeEventRequest { /// The event type to read. pub(crate) event_type: MessageLikeEventType, @@ -182,7 +182,7 @@ impl FromMatrixDriverResponse for Vec> { /// Ask the client to read matrix event(s) that corresponds to the given /// description and return a list of events as a response. -#[derive(Debug)] +#[derive(Clone, Debug)] pub(crate) struct ReadStateEventRequest { /// The event type to read. pub(crate) event_type: StateEventType, @@ -204,7 +204,7 @@ impl MatrixDriverRequest for ReadStateEventRequest { /// Ask the client to send matrix event that corresponds to the given /// description and return an event ID as a response. -#[derive(Debug, Deserialize)] +#[derive(Clone, Debug, Deserialize)] pub(crate) struct SendEventRequest { /// The type of the event. #[serde(rename = "type")] diff --git a/crates/matrix-sdk/src/widget/machine/mod.rs b/crates/matrix-sdk/src/widget/machine/mod.rs index 2b8ff46da..3cb8f8e91 100644 --- a/crates/matrix-sdk/src/widget/machine/mod.rs +++ b/crates/matrix-sdk/src/widget/machine/mod.rs @@ -25,8 +25,7 @@ use ruma::{ }; use serde::Serialize; use serde_json::value::RawValue as RawJsonValue; -use tokio::sync::mpsc::{unbounded_channel, UnboundedReceiver, UnboundedSender}; -use tracing::{debug, error, info, instrument, trace, warn}; +use tracing::{debug, error, info, instrument, warn}; use uuid::Uuid; use self::{ @@ -65,7 +64,7 @@ pub(crate) use self::{ }; /// Action (a command) that client (driver) must perform. -#[derive(Debug)] +#[derive(Clone, Debug)] pub(crate) enum Action { /// Send a raw message to the widget. SendToWidget(String), @@ -102,7 +101,6 @@ pub(crate) enum Action { pub(crate) struct WidgetMachine { widget_id: String, room_id: OwnedRoomId, - actions_sender: UnboundedSender, pending_to_widget_requests: IndexMap, pending_matrix_driver_requests: IndexMap, capabilities: CapabilitiesState, @@ -116,76 +114,72 @@ impl WidgetMachine { widget_id: String, room_id: OwnedRoomId, init_on_content_load: bool, - ) -> (Self, UnboundedReceiver) { - let (actions_sender, actions_receiver) = unbounded_channel(); + ) -> (Self, Vec) { let mut machine = Self { widget_id, room_id, - actions_sender, pending_to_widget_requests: IndexMap::new(), pending_matrix_driver_requests: IndexMap::new(), capabilities: CapabilitiesState::Unset, }; - if !init_on_content_load { - machine.negotiate_capabilities(); - } - - (machine, actions_receiver) + let actions = (!init_on_content_load).then(|| machine.negotiate_capabilities()); + (machine, actions.unwrap_or_default()) } /// Main entry point to drive the state machine. - pub(crate) fn process(&mut self, event: IncomingMessage) { + pub(crate) fn process(&mut self, event: IncomingMessage) -> Vec { match event { - IncomingMessage::WidgetMessage(raw) => { - self.process_widget_message(&raw); - } + IncomingMessage::WidgetMessage(raw) => self.process_widget_message(&raw), IncomingMessage::MatrixDriverResponse { request_id, response } => { - self.process_matrix_driver_response(request_id, response); + self.process_matrix_driver_response(request_id, response) } IncomingMessage::MatrixEventReceived(event) => { let CapabilitiesState::Negotiated(capabilities) = &self.capabilities else { error!("Received matrix event before capabilities negotiation"); - return; + return Vec::new(); }; let filter_in = match event.deserialize_as::() { Ok(i) => i, Err(e) => { error!("Failed to deserialize event: {e}"); - return; + return Vec::new(); } }; - if capabilities.read.iter().any(|f| f.matches(&filter_in)) { - self.send_to_widget_request(NotifyNewMatrixEvent(event)); - } + capabilities + .read + .iter() + .any(|f| f.matches(&filter_in)) + .then(|| vec![self.send_to_widget_request(NotifyNewMatrixEvent(event)).1]) + .unwrap_or_default() } } } - fn process_widget_message(&mut self, raw: &str) { + fn process_widget_message(&mut self, raw: &str) -> Vec { let message = match serde_json::from_str::(raw) { Ok(msg) => msg, Err(e) => { // TODO: There is a special error handling required for the invalid // messages. Refer to the `widget-api-poc` for implementation notes. error!("Failed to parse incoming message: {e}"); - return; + return Vec::new(); } }; if message.widget_id != self.widget_id { error!("Received a message from a wrong widget, ignoring"); - return; + return Vec::new(); } match message.kind { IncomingWidgetMessageKind::Request(request) => { - self.process_from_widget_request(message.request_id, request); + self.process_from_widget_request(message.request_id, request) } IncomingWidgetMessageKind::Response(response) => { - self.process_to_widget_response(message.request_id, response); + self.process_to_widget_response(message.request_id, response) } } } @@ -195,38 +189,38 @@ impl WidgetMachine { &mut self, request_id: String, raw_request: Raw, - ) { + ) -> Vec { let request = match raw_request.deserialize() { Ok(r) => r, - Err(e) => { - self.send_from_widget_error_response(raw_request, e); - return; - } + Err(e) => return vec![self.send_from_widget_error_response(raw_request, e)], }; match request { FromWidgetRequest::SupportedApiVersions {} => { - self.send_from_widget_response(raw_request, SupportedApiVersionsResponse::new()); + let response = SupportedApiVersionsResponse::new(); + vec![self.send_from_widget_response(raw_request, response)] } FromWidgetRequest::ContentLoaded {} => { - self.send_from_widget_response(raw_request, JsonObject::new()); - if self.capabilities.is_unset() { - self.negotiate_capabilities(); - } + let response = vec![self.send_from_widget_response(raw_request, JsonObject::new())]; + self.capabilities + .is_unset() + .then(|| [&response, self.negotiate_capabilities().as_slice()].concat()) + .unwrap_or(response) } FromWidgetRequest::ReadEvent(req) => { - self.process_read_event_request(req, raw_request); + vec![self.process_read_event_request(req, raw_request)] } - FromWidgetRequest::SendEvent(req) => { - self.process_send_event_request(req, raw_request); - } + FromWidgetRequest::SendEvent(req) => self + .process_send_event_request(req, raw_request) + .map(|a| vec![a]) + .unwrap_or_default(), FromWidgetRequest::GetOpenId {} => { - self.send_from_widget_response(raw_request, OpenIdResponse::Pending); - self.send_matrix_driver_request(RequestOpenId).then(|res, machine| { + let (request, request_action) = self.send_matrix_driver_request(RequestOpenId); + request.then(|res, machine| { let response = match res { Ok(res) => OpenIdResponse::Allowed(OpenIdState::new(request_id, res)), Err(msg) => { @@ -235,8 +229,11 @@ impl WidgetMachine { } }; - machine.send_to_widget_request(NotifyOpenIdChanged(response)); + vec![machine.send_to_widget_request(NotifyOpenIdChanged(response)).1] }); + + let response = self.send_from_widget_response(raw_request, OpenIdResponse::Pending); + vec![response, request_action] } } } @@ -245,21 +242,16 @@ impl WidgetMachine { &mut self, request: ReadEventRequest, raw_request: Raw, - ) { + ) -> Action { let CapabilitiesState::Negotiated(capabilities) = &self.capabilities else { - self.send_from_widget_error_response( - raw_request, - "Received read event request before capabilities were negotiated", - ); - return; + let text = "Received read event request before capabilities were negotiated"; + return self.send_from_widget_error_response(raw_request, text); }; match request { ReadEventRequest::ReadMessageLikeEvent { .. } => { - self.send_from_widget_error_response( - raw_request, - "Reading of message events is not yet supported", - ); + let text = "Reading of message events is not yet supported"; + self.send_from_widget_error_response(raw_request, text) } ReadEventRequest::ReadStateEvent { event_type, state_key } => { let allowed = match &state_key { @@ -282,12 +274,14 @@ impl WidgetMachine { if allowed { let request = ReadStateEventRequest { event_type, state_key }; - self.send_matrix_driver_request(request).then(|result, machine| { + let (request, action) = self.send_matrix_driver_request(request); + request.then(|result, machine| { let response = result.map(|events| ReadEventResponse { events }); - machine.send_from_widget_result_response(raw_request, response); + vec![machine.send_from_widget_result_response(raw_request, response)] }); + action } else { - self.send_from_widget_error_response(raw_request, "Not allowed"); + self.send_from_widget_error_response(raw_request, "Not allowed") } } } @@ -297,10 +291,10 @@ impl WidgetMachine { &mut self, request: SendEventRequest, raw_request: Raw, - ) { + ) -> Option { let CapabilitiesState::Negotiated(capabilities) = &self.capabilities else { error!("Received send event request before capabilities negotiation"); - return; + return None; }; let filter_in = MatrixEventFilterInput { @@ -314,27 +308,35 @@ impl WidgetMachine { }), }; - if capabilities.send.iter().any(|filter| filter.matches(&filter_in)) { - self.send_matrix_driver_request(request).then(|result, machine| { - let response = result - .map(|event_id| SendEventResponse { event_id, room_id: &machine.room_id }); - machine.send_from_widget_result_response(raw_request, response); + let action = if capabilities.send.iter().any(|filter| filter.matches(&filter_in)) { + let (request, action) = self.send_matrix_driver_request(request); + request.then(|result, machine| { + let room_id = &machine.room_id; + let response = result.map(|event_id| SendEventResponse { event_id, room_id }); + vec![machine.send_from_widget_result_response(raw_request, response)] }); + action } else { - self.send_from_widget_error_response(raw_request, "Not allowed"); - } + self.send_from_widget_error_response(raw_request, "Not allowed") + }; + + Some(action) } #[instrument(skip_all, fields(?request_id))] - fn process_to_widget_response(&mut self, request_id: String, response: ToWidgetResponse) { + fn process_to_widget_response( + &mut self, + request_id: String, + response: ToWidgetResponse, + ) -> Vec { let Ok(request_id) = Uuid::parse_str(&request_id) else { error!("Response's request_id is not a valid UUID"); - return; + return Vec::new(); }; let Some(request) = self.pending_to_widget_requests.remove(&request_id) else { warn!("Received response for an unknown request"); - return; + return Vec::new(); }; if response.action != request.action { @@ -342,14 +344,13 @@ impl WidgetMachine { ?request.action, ?response.action, "Received response with different `action` than request" ); + return Vec::new(); } - if let Some(response_fn) = request.response_fn { - trace!("Calling response_fn"); - response_fn(response.response_data, self); - } else { - trace!("No response_fn registered"); - } + request + .response_fn + .map(|response_fn| response_fn(response.response_data, self)) + .unwrap_or_default() } #[instrument(skip_all, fields(?request_id))] @@ -357,18 +358,13 @@ impl WidgetMachine { &mut self, request_id: Uuid, response: Result, - ) { + ) -> Vec { let Some(request) = self.pending_matrix_driver_requests.remove(&request_id) else { error!("Received response for an unknown request"); - return; + return Vec::new(); }; - if let Some(response_fn) = request.response_fn { - trace!("Calling response_fn"); - response_fn(response, self); - } else { - trace!("No response_fn registered"); - } + request.response_fn.map(|response_fn| response_fn(response, self)).unwrap_or_default() } #[instrument(skip_all, fields(request_id))] @@ -376,41 +372,22 @@ impl WidgetMachine { &self, raw_request: Raw, response_data: impl Serialize, - ) { - let mut object = match raw_request.deserialize_as::>>() { - Ok(o) => o, - Err(e) => { - error!("Failed to converted FromWidgetRequest to object representation: {e}"); - return; - } - }; - let response_data = match serde_json::value::to_raw_value(&response_data) { - Ok(d) => d, - Err(e) => { - error!("Failed to serialize response data: {e}"); - return; - } - }; + ) -> Action { + let mut object = raw_request + .deserialize_as::>>() + .expect("Failed to converted FromWidgetRequest to object representation"); + let response_data = serde_json::value::to_raw_value(&response_data) + .expect("Failed to serialize response data"); object.insert("response".to_owned(), response_data); - - let serialized = match serde_json::to_string(&object) { - Ok(s) => s, - Err(e) => { - error!("Failed to serialize response: {e}"); - return; - } - }; - - if let Err(e) = self.actions_sender.send(Action::SendToWidget(serialized)) { - error!("Failed to send action: {e}"); - } + let serialized = serde_json::to_string(&object).expect("Failed to serialize response"); + Action::SendToWidget(serialized) } fn send_from_widget_error_response( &self, raw_request: Raw, error: impl fmt::Display, - ) { + ) -> Action { self.send_from_widget_response(raw_request, FromWidgetErrorResponse::new(error)) } @@ -418,7 +395,7 @@ impl WidgetMachine { &self, raw_request: Raw, result: Result, - ) { + ) -> Action { match result { Ok(res) => self.send_from_widget_response(raw_request, res), Err(msg) => self.send_from_widget_error_response(raw_request, msg), @@ -429,7 +406,7 @@ impl WidgetMachine { fn send_to_widget_request( &mut self, to_widget_request: T, - ) -> ToWidgetRequestHandle<'_, T::ResponseData> { + ) -> (ToWidgetRequestHandle<'_, T::ResponseData>, Action) { #[derive(Serialize)] #[serde(tag = "api", rename = "toWidget", rename_all = "camelCase")] struct ToWidgetRequestSerHelper<'a, T> { @@ -447,93 +424,70 @@ impl WidgetMachine { data: to_widget_request, }; - let serialized = match serde_json::to_string(&full_request) { - Ok(msg) => msg, - Err(e) => { - error!("Failed to serialize outgoing message: {e}"); - return ToWidgetRequestHandle::null(); - } - }; - - if let Err(e) = self.actions_sender.send(Action::SendToWidget(serialized)) { - error!("Failed to send action: {e}"); - return ToWidgetRequestHandle::null(); - } - let request_meta = ToWidgetRequestMeta::new(T::ACTION); let Entry::Vacant(entry) = self.pending_to_widget_requests.entry(request_id) else { panic!("uuid collision"); }; let meta = entry.insert(request_meta); - ToWidgetRequestHandle::new(meta) + let serialized = serde_json::to_string(&full_request).expect("Failed to serialize request"); + (ToWidgetRequestHandle::new(meta), Action::SendToWidget(serialized)) } #[instrument(skip_all)] fn send_matrix_driver_request( &mut self, - matrix_driver_request: T, - ) -> MatrixDriverRequestHandle<'_, T::Response> { + request: T, + ) -> (MatrixDriverRequestHandle<'_, T::Response>, Action) { let request_id = Uuid::new_v4(); - if let Err(e) = self - .actions_sender - .send(Action::MatrixDriverRequest { request_id, data: matrix_driver_request.into() }) - { - error!("Failed to send action: {e}"); - return MatrixDriverRequestHandle::null(); - } - - let request_meta = MatrixDriverRequestMeta::new(); let Entry::Vacant(entry) = self.pending_matrix_driver_requests.entry(request_id) else { panic!("uuid collision"); }; - let meta = entry.insert(request_meta); - - MatrixDriverRequestHandle::new(meta) + let meta = entry.insert(MatrixDriverRequestMeta::new()); + let action = Action::MatrixDriverRequest { request_id, data: request.into() }; + (MatrixDriverRequestHandle::new(meta), action) } - fn negotiate_capabilities(&mut self) { - if let CapabilitiesState::Negotiated(capabilities) = &self.capabilities { - if !capabilities.read.is_empty() { - if let Err(err) = self.actions_sender.send(Action::Unsubscribe) { - error!("Failed to send action: {err}"); - } - } - } - + fn negotiate_capabilities(&mut self) -> Vec { + let unsubscribe_required = + matches!(&self.capabilities, CapabilitiesState::Negotiated(c) if !c.read.is_empty()); self.capabilities = CapabilitiesState::Negotiating; - self.send_to_widget_request(RequestCapabilities {}) - // TODO: Each request can actually fail here, take this into an account. - .then(|response, machine| { - let requested = response.capabilities; - machine - .send_matrix_driver_request(AcquireCapabilities { - desired_capabilities: requested.clone(), - }) - .then(|result, machine| { - let approved = result.unwrap_or_else(|e| { - error!("Acquiring capabilities failed: {e}"); - Capabilities::default() - }); - - if !approved.read.is_empty() { - if let Err(err) = machine.actions_sender.send(Action::Subscribe) { - error!("Failed to send action: {err}"); - } - } - - machine.capabilities = CapabilitiesState::Negotiated(approved.clone()); - machine.send_to_widget_request(NotifyCapabilitiesChanged { - approved, - requested, - }); - }) + let (request, action) = self.send_to_widget_request(RequestCapabilities {}); + request.then(|response, machine| { + let requested = response.capabilities; + let (request, action) = machine.send_matrix_driver_request(AcquireCapabilities { + desired_capabilities: requested.clone(), }); + + request.then(|result, machine| { + let approved = result.unwrap_or_else(|e| { + error!("Acquiring capabilities failed: {e}"); + Capabilities::default() + }); + + let subscribe_required = !approved.read.is_empty(); + machine.capabilities = CapabilitiesState::Negotiated(approved.clone()); + + let update = NotifyCapabilitiesChanged { approved, requested }; + let (_request, action) = machine.send_to_widget_request(update); + + (subscribe_required) + .then(|| Action::Subscribe) + .into_iter() + .chain(Some(action)) + .collect() + }); + + vec![action] + }); + + unsubscribe_required.then(|| Action::Unsubscribe).into_iter().chain(Some(action)).collect() } } -type ToWidgetResponseFn = Box, &mut WidgetMachine) + Send>; +type ToWidgetResponseFn = + Box, &mut WidgetMachine) -> Vec + Send>; pub(crate) struct ToWidgetRequestMeta { action: &'static str, @@ -547,7 +501,7 @@ impl ToWidgetRequestMeta { } type MatrixDriverResponseFn = - Box, &mut WidgetMachine) + Send>; + Box, &mut WidgetMachine) -> Vec + Send>; pub(crate) struct MatrixDriverRequestMeta { response_fn: Option, diff --git a/crates/matrix-sdk/src/widget/machine/tests/api_versions.rs b/crates/matrix-sdk/src/widget/machine/tests/api_versions.rs index 5476f0972..a409cf654 100644 --- a/crates/matrix-sdk/src/widget/machine/tests/api_versions.rs +++ b/crates/matrix-sdk/src/widget/machine/tests/api_versions.rs @@ -21,10 +21,10 @@ use crate::widget::machine::{Action, IncomingMessage, WidgetMachine}; #[test] fn get_supported_api_versions() { - let (mut machine, mut actions_recv) = + let (mut machine, _) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), true); - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "fromWidget", "widgetId": WIDGET_ID, "requestId": "S2ixNhjaC0kd0jJn", @@ -32,7 +32,7 @@ fn get_supported_api_versions() { "data": {}, }))); - let action = actions_recv.try_recv().unwrap(); + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let msg: JsonValue = serde_json::from_str(&msg).unwrap(); assert_eq!( diff --git a/crates/matrix-sdk/src/widget/machine/tests/capabilities.rs b/crates/matrix-sdk/src/widget/machine/tests/capabilities.rs index 080e80dab..7be991dc1 100644 --- a/crates/matrix-sdk/src/widget/machine/tests/capabilities.rs +++ b/crates/matrix-sdk/src/widget/machine/tests/capabilities.rs @@ -15,7 +15,6 @@ use assert_matches::assert_matches; use ruma::owned_room_id; use serde_json::{from_value, json}; -use tokio::sync::mpsc::UnboundedReceiver; use super::{parse_msg, WIDGET_ID}; use crate::widget::machine::{ @@ -24,21 +23,20 @@ use crate::widget::machine::{ #[test] fn machine_can_negotiate_capabilities_immediately() { - let (mut machine, mut actions_recv) = + let (mut machine, initial_actions) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), false); - assert_capabilities_dance(&mut machine, &mut actions_recv); - assert_matches!(actions_recv.try_recv(), Err(_)); + assert_capabilities_dance(&mut machine, initial_actions); } #[test] fn machine_can_request_capabilities_on_content_load() { - let (mut machine, mut actions_recv) = + let (mut machine, initial_actions) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), true); - assert_matches!(actions_recv.try_recv(), Err(_)); + assert!(initial_actions.is_empty()); // Content loaded event processed. - { - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let actions = { + let mut actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "fromWidget", "widgetId": WIDGET_ID, "requestId": "content-loaded-request-id", @@ -46,7 +44,7 @@ fn machine_can_request_capabilities_on_content_load() { "data": {}, }))); - let action = actions_recv.try_recv().unwrap(); + let action = actions.remove(0); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!(request_id, "content-loaded-request-id"); @@ -60,19 +58,21 @@ fn machine_can_request_capabilities_on_content_load() { "response": {}, }), ); - } - assert_capabilities_dance(&mut machine, &mut actions_recv); + actions + }; + + assert_capabilities_dance(&mut machine, actions); } #[test] fn capabilities_failure_results_into_empty_capabilities() { - let (mut machine, mut actions_recv) = + let (mut machine, actions) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), false); // Ask widget to provide desired capabilities. - { - let action = actions_recv.try_recv().unwrap(); + let actions = { + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!( @@ -94,12 +94,12 @@ fn capabilities_failure_results_into_empty_capabilities() { "response": { "capabilities": ["org.matrix.msc2762.receive.state_event:m.room.member"], }, - }))); - } + }))) + }; // Try to acquire capabilities by sending a request to a matrix driver. - { - let action = actions_recv.try_recv().unwrap(); + let actions = { + let [action]: [Action; 1] = actions.try_into().unwrap(); let (request_id, capabilities) = assert_matches!( action, Action::MatrixDriverRequest { @@ -115,11 +115,11 @@ fn capabilities_failure_results_into_empty_capabilities() { machine.process(IncomingMessage::MatrixDriverResponse { request_id, response: Err("OHMG!".into()), - }); - } + }) + }; // Inform the widget about the new capabilities, or lack of thereof :) - let action = actions_recv.try_recv().unwrap(); + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, _request_id) = parse_msg(&msg); assert_eq!( @@ -134,17 +134,12 @@ fn capabilities_failure_results_into_empty_capabilities() { }, }), ); - - assert_matches!(actions_recv.try_recv(), Err(_)); } -pub(super) fn assert_capabilities_dance( - machine: &mut WidgetMachine, - actions_recv: &mut UnboundedReceiver, -) { +pub(super) fn assert_capabilities_dance(machine: &mut WidgetMachine, actions: Vec) { // Ask widget to provide desired capabilities. - { - let action = actions_recv.try_recv().unwrap(); + let actions = { + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!( @@ -166,12 +161,12 @@ pub(super) fn assert_capabilities_dance( "response": { "capabilities": ["org.matrix.msc2762.receive.state_event:m.room.member"], }, - }))); - } + }))) + }; // Try to acquire capabilities by sending a request to a matrix driver. - { - let action = actions_recv.try_recv().unwrap(); + let mut actions = { + let [action]: [Action; 1] = actions.try_into().unwrap(); let (request_id, capabilities) = assert_matches!( action, Action::MatrixDriverRequest { @@ -184,20 +179,20 @@ pub(super) fn assert_capabilities_dance( from_value(json!(["org.matrix.msc2762.receive.state_event:m.room.member"])).unwrap() ); - let response = MatrixDriverResponse::CapabilitiesAcquired(capabilities); - machine - .process(IncomingMessage::MatrixDriverResponse { request_id, response: Ok(response) }); - } + let response = Ok(MatrixDriverResponse::CapabilitiesAcquired(capabilities)); + let message = IncomingMessage::MatrixDriverResponse { request_id, response }; + machine.process(message) + }; // We get the `Subscribe` command since we requested some reading capabilities. { - let action = actions_recv.try_recv().unwrap(); + let action = actions.remove(0); assert_matches!(action, Action::Subscribe); } // Inform the widget about the acquired capabilities. { - let action = actions_recv.try_recv().unwrap(); + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!( @@ -213,7 +208,7 @@ pub(super) fn assert_capabilities_dance( }), ); - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "toWidget", "widgetId": WIDGET_ID, "requestId": request_id, @@ -224,5 +219,7 @@ pub(super) fn assert_capabilities_dance( }, "response": {}, }))); + + assert!(actions.is_empty()); } } diff --git a/crates/matrix-sdk/src/widget/machine/tests/error.rs b/crates/matrix-sdk/src/widget/machine/tests/error.rs index 8ca6c362e..d206165f1 100644 --- a/crates/matrix-sdk/src/widget/machine/tests/error.rs +++ b/crates/matrix-sdk/src/widget/machine/tests/error.rs @@ -21,13 +21,10 @@ use crate::widget::machine::{Action, IncomingMessage, WidgetMachine}; #[test] fn machine_sends_error_for_unknown_request() { - let (mut machine, mut actions_recv) = + let (mut machine, _) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), true); - // No messages from the machine at first - assert_matches!(actions_recv.try_recv(), Err(_)); - - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "fromWidget", "widgetId": WIDGET_ID, "requestId": "invalid-req", @@ -37,7 +34,7 @@ fn machine_sends_error_for_unknown_request() { }, }))); - let action = actions_recv.try_recv().unwrap(); + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!(request_id, "invalid-req"); @@ -50,10 +47,10 @@ fn machine_sends_error_for_unknown_request() { #[test] fn read_messages_without_capabilities() { - let (mut machine, mut actions_recv) = + let (mut machine, _) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), true); - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "fromWidget", "widgetId": WIDGET_ID, "requestId": "get-me-some-messages", @@ -63,7 +60,7 @@ fn read_messages_without_capabilities() { }, }))); - let action = actions_recv.try_recv().unwrap(); + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!(request_id, "get-me-some-messages"); @@ -77,12 +74,11 @@ fn read_messages_without_capabilities() { #[test] fn read_messages_not_yet_supported() { - let (mut machine, mut actions_recv) = + let (mut machine, actions) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), false); + assert_capabilities_dance(&mut machine, actions); - assert_capabilities_dance(&mut machine, &mut actions_recv); - - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "fromWidget", "widgetId": WIDGET_ID, "requestId": "get-me-some-messages", @@ -92,7 +88,7 @@ fn read_messages_not_yet_supported() { }, }))); - let action = actions_recv.try_recv().unwrap(); + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!(request_id, "get-me-some-messages"); diff --git a/crates/matrix-sdk/src/widget/machine/tests/openid.rs b/crates/matrix-sdk/src/widget/machine/tests/openid.rs index bbd523d09..73842ce4a 100644 --- a/crates/matrix-sdk/src/widget/machine/tests/openid.rs +++ b/crates/matrix-sdk/src/widget/machine/tests/openid.rs @@ -25,13 +25,13 @@ use crate::widget::machine::{ #[test] fn openid_request_handling_works() { - let (mut machine, mut actions_recv) = + let (mut machine, _) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), true); // Widget requests an open ID token, since we don't have any caching yet, // we reply with a pending response right away. - { - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let actions = { + let mut actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "fromWidget", "widgetId": WIDGET_ID, "requestId": "openid-request-id", @@ -39,7 +39,7 @@ fn openid_request_handling_works() { "data": {}, }))); - let action = actions_recv.try_recv().unwrap(); + let action = actions.remove(0); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!(request_id, "openid-request-id"); @@ -55,11 +55,13 @@ fn openid_request_handling_works() { }, }), ); - } + + actions + }; // Then we send an OpenID request to the driver and expect an answer. - { - let action = actions_recv.try_recv().unwrap(); + let actions = { + let [action]: [Action; 1] = actions.try_into().unwrap(); let request_id = assert_matches!( action, Action::MatrixDriverRequest { @@ -78,12 +80,12 @@ fn openid_request_handling_works() { Duration::from_secs(3600), ), )), - }); - } + }) + }; // We inform the widget about the new OpenID token. { - let action = actions_recv.try_recv().unwrap(); + let [action]: [Action; 1] = actions.try_into().unwrap(); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, _request_id) = parse_msg(&msg); assert_eq!( @@ -103,20 +105,17 @@ fn openid_request_handling_works() { }), ); } - - // No further actions expected. - assert_matches!(actions_recv.try_recv(), Err(_)); } #[test] fn openid_fail_results_in_response_blocked() { - let (mut machine, mut actions_recv) = + let (mut machine, _) = WidgetMachine::new(WIDGET_ID.to_owned(), owned_room_id!("!a98sd12bjh:example.org"), true); // Widget requests an open ID token, since we don't have any caching yet, // we reply with a pending response right away. - { - machine.process(IncomingMessage::WidgetMessage(json_string!({ + let mut actions = { + let mut actions = machine.process(IncomingMessage::WidgetMessage(json_string!({ "api": "fromWidget", "widgetId": WIDGET_ID, "requestId": "openid-request-id", @@ -124,7 +123,7 @@ fn openid_fail_results_in_response_blocked() { "data": {}, }))); - let action = actions_recv.try_recv().unwrap(); + let action = actions.remove(0); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, request_id) = parse_msg(&msg); assert_eq!(request_id, "openid-request-id"); @@ -140,11 +139,14 @@ fn openid_fail_results_in_response_blocked() { }, }), ); - } + + actions + }; // Then we send an OpenID request to the driver and expect a fail. - { - let action = actions_recv.try_recv().unwrap(); + let mut actions = { + let action = actions.remove(0); + assert!(actions.is_empty()); let request_id = assert_matches!( action, Action::MatrixDriverRequest { @@ -156,12 +158,12 @@ fn openid_fail_results_in_response_blocked() { machine.process(IncomingMessage::MatrixDriverResponse { request_id, response: Err("Unlucky one".into()), - }); - } + }) + }; // We inform the widget about the new OpenID token. { - let action = actions_recv.try_recv().unwrap(); + let action = actions.remove(0); let msg = assert_matches!(action, Action::SendToWidget(msg) => msg); let (msg, _request_id) = parse_msg(&msg); assert_eq!( @@ -179,5 +181,5 @@ fn openid_fail_results_in_response_blocked() { } // No further actions expected. - assert_matches!(actions_recv.try_recv(), Err(_)); + assert!(actions.is_empty()); } diff --git a/crates/matrix-sdk/src/widget/machine/to_widget.rs b/crates/matrix-sdk/src/widget/machine/to_widget.rs index e499f5080..842aad968 100644 --- a/crates/matrix-sdk/src/widget/machine/to_widget.rs +++ b/crates/matrix-sdk/src/widget/machine/to_widget.rs @@ -19,7 +19,7 @@ use serde::{de::DeserializeOwned, Deserialize, Serialize}; use serde_json::value::RawValue as RawJsonValue; use tracing::error; -use super::{openid::OpenIdResponse, ToWidgetRequestMeta, WidgetMachine}; +use super::{openid::OpenIdResponse, Action, ToWidgetRequestMeta, WidgetMachine}; use crate::widget::Capabilities; /// A handle to a pending `toWidget` request. @@ -36,19 +36,18 @@ where Self { request_meta: Some(request_meta), _phantom: PhantomData } } - pub(crate) fn null() -> Self { - Self { request_meta: None, _phantom: PhantomData } - } - pub(crate) fn then( self, - response_handler: impl FnOnce(T, &mut WidgetMachine) + Send + 'static, + response_handler: impl FnOnce(T, &mut WidgetMachine) -> Vec + Send + 'static, ) { if let Some(request_meta) = self.request_meta { request_meta.response_fn = Some(Box::new(move |raw_response_data, machine| { match serde_json::from_str(raw_response_data.get()) { Ok(response_data) => response_handler(response_data, machine), - Err(e) => error!("Failed to deserialize toWidget response: {e}"), + Err(e) => { + error!("Failed to deserialize toWidget response: {e}"); + Vec::new() + } } })); } diff --git a/crates/matrix-sdk/src/widget/mod.rs b/crates/matrix-sdk/src/widget/mod.rs index d071c13f8..2024a380d 100644 --- a/crates/matrix-sdk/src/widget/mod.rs +++ b/crates/matrix-sdk/src/widget/mod.rs @@ -18,7 +18,7 @@ use std::fmt; use async_channel::{Receiver, Sender}; use serde::de::{self, Deserialize, Deserializer, Visitor}; -use tokio::sync::mpsc::unbounded_channel; +use tokio::sync::mpsc::{unbounded_channel, UnboundedSender}; use tokio_util::sync::{CancellationToken, DropGuard}; use self::{ @@ -124,12 +124,6 @@ impl WidgetDriver { room: Room, capabilities_provider: impl CapabilitiesProvider, ) -> Result<(), ()> { - let (mut client_api, mut actions) = WidgetMachine::new( - self.settings.widget_id().to_owned(), - room.room_id().to_owned(), - self.settings.init_on_content_load(), - ); - // Create a channel so that we can conveniently send all events to it. let (events_tx, mut events_rx) = unbounded_channel(); @@ -141,89 +135,132 @@ impl WidgetDriver { } }); - // Forward all of the incoming events to the `ClientApi` implementation. - tokio::spawn(async move { - while let Some(event) = events_rx.recv().await { - client_api.process(event); + // Create widget API machine. + let (client_api, initial_actions) = WidgetMachine::new( + self.settings.widget_id().to_owned(), + room.room_id().to_owned(), + self.settings.init_on_content_load(), + ); + + // The environment for the processing of actions from the widget machine. + let mut ctx = ProcessingContext { + widget_machine: client_api, + matrix_driver: MatrixDriver::new(room.clone()), + event_forwarding_guard: None, + to_widget_tx: self.to_widget_tx, + events_tx, + capabilities_provider, + }; + + // Process initial actions that "initialise" the widget api machine. + for action in initial_actions { + ctx.process_action(action).await?; + } + + // Process incoming events. + while let Some(event) = events_rx.recv().await { + ctx.process_event(event).await?; + } + + Ok(()) + } +} + +/// A small wrapper of all the data that we need to process an incoming event. +struct ProcessingContext { + widget_machine: WidgetMachine, + matrix_driver: MatrixDriver, + event_forwarding_guard: Option, + to_widget_tx: Sender, + events_tx: UnboundedSender, + capabilities_provider: T, +} + +impl ProcessingContext { + async fn process_event(&mut self, event: IncomingMessage) -> Result<(), ()> { + for action in self.widget_machine.process(event) { + self.process_action(action).await?; + } + + Ok(()) + } + + async fn process_action(&mut self, action: Action) -> Result<(), ()> { + match action { + Action::SendToWidget(msg) => { + self.to_widget_tx.send(msg).await.map_err(|_| ())?; } - }); + Action::MatrixDriverRequest { request_id, data } => { + let response = match data { + MatrixDriverRequestData::AcquireCapabilities(cmd) => { + let obtained = self + .capabilities_provider + .acquire_capabilities(cmd.desired_capabilities) + .await; + Ok(MatrixDriverResponse::CapabilitiesAcquired(obtained)) + } - // Process events that we receive **from** the client api implementation, - // i.e. the commands (actions) that the client sends to us. - let matrix_driver = MatrixDriver::new(room); - let mut event_forwarding_guard: Option = None; - while let Some(action) = actions.recv().await { - match action { - Action::SendToWidget(msg) => self.to_widget_tx.send(msg).await.map_err(|_| ())?, - Action::MatrixDriverRequest { request_id, data } => { - let response = match data { - MatrixDriverRequestData::AcquireCapabilities(cmd) => { - let obtained = capabilities_provider - .acquire_capabilities(cmd.desired_capabilities.clone()) - .await; + MatrixDriverRequestData::GetOpenId => self + .matrix_driver + .get_open_id() + .await + .map(MatrixDriverResponse::OpenIdReceived) + .map_err(|e| e.to_string()), - Ok(MatrixDriverResponse::CapabilitiesAcquired(obtained)) - } + MatrixDriverRequestData::ReadMessageLikeEvent(cmd) => self + .matrix_driver + .read_message_like_events(cmd.event_type.clone(), cmd.limit) + .await + .map(MatrixDriverResponse::MatrixEventRead) + .map_err(|e| e.to_string()), - MatrixDriverRequestData::GetOpenId => matrix_driver - .get_open_id() + MatrixDriverRequestData::ReadStateEvent(cmd) => self + .matrix_driver + .read_state_events(cmd.event_type.clone(), &cmd.state_key) + .await + .map(MatrixDriverResponse::MatrixEventRead) + .map_err(|e| e.to_string()), + + MatrixDriverRequestData::SendMatrixEvent(req) => { + let SendEventRequest { event_type, state_key, content } = req; + self.matrix_driver + .send(event_type, state_key, content) .await - .map(MatrixDriverResponse::OpenIdReceived) - .map_err(|e| e.to_string()), + .map(MatrixDriverResponse::MatrixEventSent) + .map_err(|e| e.to_string()) + } + }; - MatrixDriverRequestData::ReadMessageLikeEvent(cmd) => matrix_driver - .read_message_like_events(cmd.event_type.clone(), cmd.limit) - .await - .map(MatrixDriverResponse::MatrixEventRead) - .map_err(|e| e.to_string()), - - MatrixDriverRequestData::ReadStateEvent(cmd) => matrix_driver - .read_state_events(cmd.event_type.clone(), &cmd.state_key) - .await - .map(MatrixDriverResponse::MatrixEventRead) - .map_err(|e| e.to_string()), - - MatrixDriverRequestData::SendMatrixEvent(req) => { - let SendEventRequest { event_type, state_key, content } = req; - - matrix_driver - .send(event_type, state_key, content) - .await - .map(MatrixDriverResponse::MatrixEventSent) - .map_err(|e| e.to_string()) - } + self.events_tx + .send(IncomingMessage::MatrixDriverResponse { request_id, response }) + .map_err(|_| ())?; + } + Action::Subscribe => { + // Only subscribe if we are not already subscribed. + if self.event_forwarding_guard.is_none() { + let (stop_forwarding, guard) = { + let token = CancellationToken::new(); + (token.child_token(), token.drop_guard()) }; - events_tx - .send(IncomingMessage::MatrixDriverResponse { request_id, response }) - .map_err(|_| ())?; - } - Action::Subscribe => { - // Only subscribe if we are not already subscribed. - if event_forwarding_guard.is_none() { - let (stop_forwarding, guard) = { - let token = CancellationToken::new(); - (token.child_token(), token.drop_guard()) - }; - - event_forwarding_guard = Some(guard); - let (mut matrix, events_tx) = (matrix_driver.events(), events_tx.clone()); - tokio::spawn(async move { - loop { - tokio::select! { - _ = stop_forwarding.cancelled() => { return } - Some(event) = matrix.recv() => { - let _ = events_tx.send(IncomingMessage::MatrixEventReceived(event)); - } + self.event_forwarding_guard = Some(guard); + let (mut matrix, events_tx) = + (self.matrix_driver.events(), self.events_tx.clone()); + tokio::spawn(async move { + loop { + tokio::select! { + _ = stop_forwarding.cancelled() => { return } + Some(event) = matrix.recv() => { + let _ = events_tx.send(IncomingMessage::MatrixEventReceived(event)); } } - }); - } - } - Action::Unsubscribe => { - event_forwarding_guard = None; + } + }); } } + Action::Unsubscribe => { + self.event_forwarding_guard = None; + } } Ok(())