up_rust/
local_transport.rs1use std::{collections::HashSet, sync::Arc};
19
20use tokio::sync::RwLock;
21
22use crate::{ComparableListener, UListener, UMessage, UStatus, UTransport, UUri};
23
24#[derive(Eq, PartialEq, Hash)]
25struct RegisteredListener {
26 source_filter: UUri,
27 sink_filter: Option<UUri>,
28 listener: ComparableListener,
29}
30
31impl RegisteredListener {
32 fn matches(&self, source: &UUri, sink: Option<&UUri>) -> bool {
33 if !self.source_filter.matches(source) {
34 return false;
35 }
36
37 if let Some(pattern) = &self.sink_filter {
38 sink.is_some_and(|candidate_sink| pattern.matches(candidate_sink))
39 } else {
40 sink.is_none()
41 }
42 }
43 fn matches_msg(&self, msg: &UMessage) -> bool {
44 self.matches(msg.source(), msg.sink())
45 }
46 async fn on_receive(&self, msg: UMessage) {
47 self.listener.on_receive(msg).await
48 }
49}
50
51#[derive(Default)]
56pub struct LocalTransport {
57 listeners: RwLock<HashSet<RegisteredListener>>,
58}
59
60impl core::fmt::Debug for LocalTransport {
61 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
62 f.debug_struct("LocalTransport").finish_non_exhaustive()
63 }
64}
65
66impl LocalTransport {
67 async fn dispatch(&self, message: UMessage) {
68 let listeners = self.listeners.read().await;
69 for listener in listeners.iter() {
70 if listener.matches_msg(&message) {
71 listener.on_receive(message.clone()).await;
72 }
73 }
74 }
75}
76
77#[async_trait::async_trait]
78impl UTransport for LocalTransport {
79 async fn send(&self, message: UMessage) -> Result<(), UStatus> {
80 self.dispatch(message).await;
81 Ok(())
82 }
83
84 async fn register_listener(
85 &self,
86 source_filter: &UUri,
87 sink_filter: Option<&UUri>,
88 listener: Arc<dyn UListener>,
89 ) -> Result<(), UStatus> {
90 let registered_listener = RegisteredListener {
91 source_filter: source_filter.to_owned(),
92 sink_filter: sink_filter.map(|u| u.to_owned()),
93 listener: ComparableListener::new(listener),
94 };
95 let mut listeners = self.listeners.write().await;
96 if listeners.contains(®istered_listener) {
97 Err(UStatus::fail_with_code(
98 crate::UCode::AlreadyExists,
99 "listener already registered for filters",
100 ))
101 } else {
102 listeners.insert(registered_listener);
103 Ok(())
104 }
105 }
106
107 async fn unregister_listener(
108 &self,
109 source_filter: &UUri,
110 sink_filter: Option<&UUri>,
111 listener: Arc<dyn UListener>,
112 ) -> Result<(), UStatus> {
113 let registered_listener = RegisteredListener {
114 source_filter: source_filter.to_owned(),
115 sink_filter: sink_filter.map(|u| u.to_owned()),
116 listener: ComparableListener::new(listener),
117 };
118 let mut listeners = self.listeners.write().await;
119 if listeners.remove(®istered_listener) {
120 Ok(())
121 } else {
122 Err(UStatus::fail_with_code(
123 crate::UCode::NotFound,
124 "no such listener registered for filters",
125 ))
126 }
127 }
128}
129
130#[cfg(test)]
131mod tests {
132 use super::*;
133 use crate::{utransport::MockUListener, LocalUriProvider, StaticUriProvider, UMessageBuilder};
134
135 #[tokio::test]
136 async fn test_send_dispatches_to_matching_listener() {
137 const RESOURCE_ID: u16 = 0xa1b3;
138 let mut listener = MockUListener::new();
139 listener.expect_on_receive().once().return_const(());
140 let listener_ref = Arc::new(listener);
141 let uri_provider = StaticUriProvider::new("my-vehicle", 0x100d, 0x02)
142 .expect("failed to create URI provider");
143 let transport = LocalTransport::default();
144
145 transport
146 .register_listener(
147 &uri_provider.get_resource_uri(RESOURCE_ID),
148 None,
149 listener_ref.clone(),
150 )
151 .await
152 .unwrap();
153 let _ = transport
154 .send(
155 UMessageBuilder::publish(uri_provider.get_resource_uri(RESOURCE_ID))
156 .build()
157 .unwrap(),
158 )
159 .await;
160
161 transport
162 .unregister_listener(
163 &uri_provider.get_resource_uri(RESOURCE_ID),
164 None,
165 listener_ref,
166 )
167 .await
168 .unwrap();
169 let _ = transport
170 .send(
171 UMessageBuilder::publish(uri_provider.get_resource_uri(RESOURCE_ID))
172 .build()
173 .unwrap(),
174 )
175 .await;
176 }
177
178 #[tokio::test]
179 async fn test_send_does_not_dispatch_to_non_matching_listener() {
180 const RESOURCE_ID: u16 = 0xa1b3;
181 let mut listener = MockUListener::new();
182 listener.expect_on_receive().never().return_const(());
183 let listener_ref = Arc::new(listener);
184 let uri_provider = StaticUriProvider::new("my-vehicle", 0x100d, 0x02)
185 .expect("failed to create URI provider");
186 let transport = LocalTransport::default();
187
188 transport
189 .register_listener(
190 &uri_provider.get_resource_uri(RESOURCE_ID + 10),
191 None,
192 listener_ref.clone(),
193 )
194 .await
195 .unwrap();
196 let _ = transport
197 .send(
198 UMessageBuilder::publish(uri_provider.get_resource_uri(RESOURCE_ID))
199 .build()
200 .unwrap(),
201 )
202 .await;
203
204 transport
205 .unregister_listener(
206 &uri_provider.get_resource_uri(RESOURCE_ID + 10),
207 None,
208 listener_ref,
209 )
210 .await
211 .unwrap();
212 let _ = transport
213 .send(
214 UMessageBuilder::publish(uri_provider.get_resource_uri(RESOURCE_ID))
215 .build()
216 .unwrap(),
217 )
218 .await;
219 }
220}