Skip to main content

reachability_handler/
lib.rs

1// Copyright 2023 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use async_utils::hanging_get::server as hanging_get;
6
7use fidl_fuchsia_net_reachability as freachability;
8use fuchsia_async as fasync;
9use fuchsia_component::server::{ServiceFsDir, ServiceObjLocal};
10use futures::lock::Mutex;
11use futures::{TryFutureExt as _, TryStreamExt as _};
12use log::error;
13use std::sync::Arc;
14
15type WatchResponder = freachability::MonitorWatchResponder;
16type NotifyFn = Box<dyn Fn(&freachability::Snapshot, WatchResponder) -> bool>;
17type ReachabilityBroker =
18    hanging_get::HangingGet<freachability::Snapshot, WatchResponder, NotifyFn>;
19type ReachabilityPublisher =
20    hanging_get::Publisher<freachability::Snapshot, WatchResponder, NotifyFn>;
21
22pub struct ReachabilityHandler {
23    state: Arc<Mutex<ReachabilityState>>,
24    broker: Arc<Mutex<ReachabilityBroker>>,
25    publisher: Arc<Mutex<ReachabilityPublisher>>,
26}
27
28#[derive(Clone, Debug, PartialEq)]
29pub struct ReachabilityState {
30    pub internet_available: bool,
31    pub gateway_reachable: bool,
32    pub dns_active: bool,
33    pub http_active: bool,
34}
35
36impl From<ReachabilityState> for freachability::Snapshot {
37    fn from(state: ReachabilityState) -> Self {
38        Self {
39            internet_available: Some(state.internet_available),
40            gateway_reachable: Some(state.gateway_reachable),
41            dns_active: Some(state.dns_active),
42            http_active: Some(state.http_active),
43            ..Default::default()
44        }
45    }
46}
47
48impl ReachabilityHandler {
49    pub fn new() -> Self {
50        let notify_fn: NotifyFn = Box::new(|state, responder| match responder.send(&state) {
51            Ok(()) => true,
52            Err(e) => {
53                error!("Failed to send reachability state to client: {}", e);
54                false
55            }
56        });
57        let state = ReachabilityState {
58            internet_available: false,
59            gateway_reachable: false,
60            dns_active: false,
61            http_active: false,
62        };
63        let broker = hanging_get::HangingGet::new(state.clone().into(), notify_fn);
64        let publisher = broker.new_publisher();
65        Self {
66            state: Arc::new(Mutex::new(state)),
67            broker: Arc::new(Mutex::new(broker)),
68            publisher: Arc::new(Mutex::new(publisher)),
69        }
70    }
71
72    pub async fn replace_state(&mut self, new_state: ReachabilityState) {
73        self.update_state(|state| *state = new_state).await;
74    }
75
76    async fn update_state(&mut self, update_callback: impl FnOnce(&mut ReachabilityState)) {
77        let mut current_state_guard = self.state.lock().await;
78        let previous_state = current_state_guard.clone();
79
80        update_callback(&mut current_state_guard);
81
82        if *current_state_guard != previous_state {
83            self.publisher
84                .lock()
85                .await
86                .set(freachability::Snapshot::from(current_state_guard.clone()));
87        }
88    }
89
90    pub fn publish_service<'a, 'b>(
91        &mut self,
92        mut svc_dir: ServiceFsDir<'a, ServiceObjLocal<'b, ()>>,
93    ) {
94        let _ = svc_dir.add_fidl_service({
95            let broker = self.broker.clone();
96            move |mut stream: freachability::MonitorRequestStream| {
97                let broker = broker.clone();
98                fasync::Task::local(
99                    async move {
100                        let subscriber = broker.lock().await.new_subscriber();
101                        // Keep track of whether SetOptions or Watch were already called. Calling
102                        // SetOptions after either it or Watch have already been called will result in us
103                        // closing the request stream.
104                        let mut set_options_called = false;
105                        let mut watch_called = false;
106                        while let Some(req) = stream.try_next().await? {
107                            match req {
108                                freachability::MonitorRequest::Watch { responder } => {
109                                    watch_called = true;
110                                    subscriber.register(responder)?
111                                }
112                                freachability::MonitorRequest::SetOptions {
113                                    payload: _,
114                                    control_handle,
115                                } => {
116                                    if watch_called || set_options_called {
117                                        control_handle.shutdown_with_epitaph(
118                                            fidl::Status::CONNECTION_ABORTED,
119                                        );
120                                        break;
121                                    }
122                                    set_options_called = true;
123                                }
124                            }
125                        }
126
127                        Ok(())
128                    }
129                    .unwrap_or_else(|e: anyhow::Error| error!("{:?}", e)),
130                )
131                .detach()
132            }
133        });
134    }
135}
136
137#[cfg(test)]
138mod tests {
139    use super::*;
140    use anyhow::Error;
141    use assert_matches::assert_matches;
142    use fidl::endpoints::Proxy;
143    use fuchsia_component::server::ServiceFs;
144    use futures::StreamExt as _;
145    use std::cell::RefCell;
146    use std::task::Poll;
147
148    struct TestEnv {
149        connector: fuchsia_component::server::ProtocolConnector,
150    }
151
152    impl TestEnv {
153        fn new(mut service_fs: ServiceFs<ServiceObjLocal<'static, ()>>) -> Self {
154            let connector = service_fs.create_protocol_connector().unwrap();
155            fasync::Task::local(service_fs.collect()).detach();
156            Self { connector }
157        }
158
159        fn connect_client(&self) -> FakeClient {
160            let watcher_proxy =
161                self.connector.connect_to_protocol::<freachability::MonitorMarker>().unwrap();
162            FakeClient { watcher_proxy, hanging_watcher_request: RefCell::new(None) }
163        }
164    }
165
166    struct FakeClient {
167        watcher_proxy: freachability::MonitorProxy,
168        hanging_watcher_request:
169            RefCell<Option<fidl::client::QueryResponseFut<freachability::Snapshot>>>,
170    }
171
172    impl FakeClient {
173        fn get_reachability_state(
174            &self,
175            executor: &mut fasync::TestExecutor,
176        ) -> Result<Option<freachability::Snapshot>, Error> {
177            let mut watch_request = self
178                .hanging_watcher_request
179                .take()
180                .take()
181                .unwrap_or_else(|| self.watcher_proxy.watch());
182
183            match executor.run_until_stalled(&mut watch_request) {
184                Poll::Pending => {
185                    let _: Option<fidl::client::QueryResponseFut<freachability::Snapshot>> =
186                        self.hanging_watcher_request.replace(Some(watch_request));
187                    Ok(None)
188                }
189                Poll::Ready(Ok(state)) => Ok(Some(state)),
190                Poll::Ready(Err(e)) => Err(e.into()),
191            }
192        }
193    }
194
195    // Tests that the handler correctly implements the hanging-get pattern.
196    #[test]
197    fn test_hanging_get() {
198        let mut executor = fasync::TestExecutor::new();
199        let mut service_fs = ServiceFs::new_local();
200        let mut handler = ReachabilityHandler::new();
201        handler.publish_service(service_fs.root_dir());
202        let test_env = TestEnv::new(service_fs);
203        let client = test_env.connect_client();
204
205        assert_matches!(
206            client.get_reachability_state(&mut executor),
207            Ok(Some(freachability::Snapshot {
208                internet_available: Some(false),
209                gateway_reachable: Some(false),
210                dns_active: Some(false),
211                ..
212            }))
213        );
214
215        // Verify no response as state hasn't changed.
216        assert_matches!(client.get_reachability_state(&mut executor), Ok(None));
217
218        executor.run_singlethreaded(handler.replace_state(ReachabilityState {
219            internet_available: true,
220            gateway_reachable: true,
221            dns_active: true,
222            http_active: true,
223        }));
224
225        assert_matches!(
226            client.get_reachability_state(&mut executor),
227            Ok(Some(freachability::Snapshot {
228                internet_available: Some(true),
229                gateway_reachable: Some(true),
230                dns_active: Some(true),
231                http_active: Some(true),
232                ..
233            }))
234        );
235    }
236
237    #[test]
238    fn test_hanging_get_multiple_clients() {
239        let mut executor = fasync::TestExecutor::new();
240        let mut service_fs = ServiceFs::new_local();
241        let mut handler = ReachabilityHandler::new();
242        handler.publish_service(service_fs.root_dir());
243        let test_env = TestEnv::new(service_fs);
244
245        let client1 = test_env.connect_client();
246        let client2 = test_env.connect_client();
247
248        assert_matches!(
249            client1.get_reachability_state(&mut executor),
250            Ok(Some(freachability::Snapshot {
251                internet_available: Some(false),
252                gateway_reachable: Some(false),
253                dns_active: Some(false),
254                ..
255            }))
256        );
257        assert_matches!(
258            client2.get_reachability_state(&mut executor),
259            Ok(Some(freachability::Snapshot {
260                internet_available: Some(false),
261                gateway_reachable: Some(false),
262                dns_active: Some(false),
263                ..
264            }))
265        );
266
267        assert_matches!(client1.get_reachability_state(&mut executor), Ok(None));
268        assert_matches!(client2.get_reachability_state(&mut executor), Ok(None));
269
270        executor.run_singlethreaded(handler.update_state(|state| {
271            state.internet_available = true;
272            state.gateway_reachable = true;
273        }));
274
275        assert_matches!(
276            client1.get_reachability_state(&mut executor),
277            Ok(Some(freachability::Snapshot {
278                internet_available: Some(true),
279                gateway_reachable: Some(true),
280                dns_active: Some(false),
281                ..
282            }))
283        );
284        assert_matches!(
285            client2.get_reachability_state(&mut executor),
286            Ok(Some(freachability::Snapshot {
287                internet_available: Some(true),
288                gateway_reachable: Some(true),
289                dns_active: Some(false),
290                ..
291            }))
292        );
293
294        // An update that does not change the current state should not be published.
295        executor.run_singlethreaded(handler.update_state(|state| {
296            state.internet_available = true;
297            state.gateway_reachable = true;
298            state.dns_active = false;
299        }));
300
301        assert_matches!(client1.get_reachability_state(&mut executor), Ok(None));
302        assert_matches!(client2.get_reachability_state(&mut executor), Ok(None));
303    }
304
305    // Tests that the handler closes the request stream if the client calls SetOptions after having
306    // already called Watch.
307    #[test]
308    fn test_cannot_call_set_options_after_watch() {
309        let mut executor = fasync::TestExecutor::new();
310        let mut service_fs = ServiceFs::new_local();
311        let mut handler = ReachabilityHandler::new();
312        handler.publish_service(service_fs.root_dir());
313        let test_env = TestEnv::new(service_fs);
314        let client = test_env.connect_client();
315
316        assert_matches!(client.get_reachability_state(&mut executor), Ok(_));
317        assert_matches!(
318            client.watcher_proxy.set_options(&freachability::MonitorOptions::default()),
319            Ok(())
320        );
321        assert_matches!(executor.run_singlethreaded(client.watcher_proxy.on_closed()), Ok(_));
322    }
323
324    // Tests that the handler closes the request stream if the client calls SetOptions after having
325    // already called it before.
326    #[test]
327    fn test_cannot_call_set_options_twice() {
328        let mut executor = fasync::TestExecutor::new();
329        let mut service_fs = ServiceFs::new_local();
330        let mut handler = ReachabilityHandler::new();
331        handler.publish_service(service_fs.root_dir());
332        let test_env = TestEnv::new(service_fs);
333        let client = test_env.connect_client();
334
335        assert_matches!(
336            client.watcher_proxy.set_options(&freachability::MonitorOptions::default()),
337            Ok(())
338        );
339        assert_matches!(
340            client.watcher_proxy.set_options(&freachability::MonitorOptions::default()),
341            Ok(())
342        );
343        assert_matches!(executor.run_singlethreaded(client.watcher_proxy.on_closed()), Ok(_));
344    }
345}