1use std::future::Future;
4use std::time::Duration;
5
6use map_core::client::MapClient;
7use tokio::io::{AsyncRead, AsyncWrite};
8use tokio::sync::watch;
9use tokio_retry::strategy::{jitter, ExponentialBackoff};
10
11use crate::SessionError;
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
15pub enum SessionState {
16 Disconnected,
18 Connecting,
20 Active,
22 Reconnecting,
24}
25
26pub async fn run_map_session(
28 addr: bluer::Address,
29 channel: u8,
30 state: watch::Sender<SessionState>,
31 cancel: watch::Receiver<bool>,
32) {
33 let strategy = ExponentialBackoff::from_millis(2)
34 .factor(1000)
35 .max_delay(Duration::from_secs(60))
36 .map(jitter);
37 run_session_loop(move || crate::lifecycle::connect_map(addr, channel), strategy, state, cancel)
38 .await;
39}
40
41pub(crate) async fn run_session_loop<F, Fut, T, S>(
43 make_client: F,
44 strategy: S,
45 state: watch::Sender<SessionState>,
46 mut cancel: watch::Receiver<bool>,
47) where
48 F: Fn() -> Fut,
49 Fut: Future<Output = Result<MapClient<T>, SessionError>>,
50 T: AsyncRead + AsyncWrite + Unpin,
51 S: Clone + Iterator<Item = Duration>,
52{
53 let mut delays = strategy.clone();
54
55 while !*cancel.borrow() {
56 let _ = state.send(SessionState::Connecting);
57 tracing::info!("MAP: connecting");
58
59 let stop = match make_client().await {
60 Ok(client) => {
61 delays = strategy.clone();
62 serve_active(client, &state, &mut cancel).await
63 }
64 Err(e) => {
65 tracing::warn!("MAP: connect failed: {e}");
66 backoff(&mut delays, &state, &mut cancel).await
67 }
68 };
69 if stop {
70 break;
71 }
72 }
73 let _ = state.send(SessionState::Disconnected);
74}
75
76async fn serve_active<T>(
78 mut client: MapClient<T>,
79 state: &watch::Sender<SessionState>,
80 cancel: &mut watch::Receiver<bool>,
81) -> bool
82where
83 T: AsyncRead + AsyncWrite + Unpin,
84{
85 let _ = state.send(SessionState::Active);
86 tracing::info!("MAP: session active");
87 if interrupted(client.hold(), cancel).await {
88 return true;
89 }
90 let _ = state.send(SessionState::Reconnecting);
91 tracing::warn!("MAP: session dropped, reconnecting");
92 false
93}
94
95async fn backoff<S: Iterator<Item = Duration>>(
97 delays: &mut S,
98 state: &watch::Sender<SessionState>,
99 cancel: &mut watch::Receiver<bool>,
100) -> bool {
101 let Some(delay) = delays.next() else { return true };
102 let _ = state.send(SessionState::Reconnecting);
103 interrupted(tokio::time::sleep(delay), cancel).await
104}
105
106async fn interrupted<Fut: Future>(until: Fut, cancel: &mut watch::Receiver<bool>) -> bool {
108 tokio::select! {
109 _ = until => false,
110 res = cancel.changed() => res.is_err() || *cancel.borrow(),
111 }
112}
113
114#[cfg(test)]
115mod tests {
116 use bytes::Bytes;
117 use futures::{SinkExt, StreamExt};
118 use std::time::Duration;
119 use tokio::sync::watch;
120
121 use crate::{lifecycle, SessionError, SessionState};
122
123 use super::run_session_loop;
124
125 const MAP_CONNECT_RSP: &[u8] = include_bytes!("../../imsg-obex/tests/fixtures/connect_rsp.bin");
126 const NOTIF_REG_OK: &[u8] = &[0xA0, 0x00, 0x03];
127
128 #[tokio::test]
129 async fn exits_when_strategy_exhausted() {
130 let (state_tx, _state_rx) = watch::channel(SessionState::Disconnected);
131 let (cancel_tx, cancel_rx) = watch::channel(false);
132 let strategy = std::iter::once(Duration::ZERO);
133 let make_client = || async {
134 Err::<map_core::client::MapClient<tokio::io::DuplexStream>, _>(SessionError::Transport(
135 obex_core::TransportError::Io(std::io::Error::new(
136 std::io::ErrorKind::ConnectionRefused,
137 "refused",
138 )),
139 ))
140 };
141 run_session_loop(make_client, strategy, state_tx, cancel_rx).await;
142 let _ = cancel_tx.send(false);
143 }
144
145 #[tokio::test]
146 async fn exits_on_cancel_before_connect() {
147 let (state_tx, _state_rx) = watch::channel(SessionState::Disconnected);
148 let (cancel_tx, cancel_rx) = watch::channel(true);
149 let strategy = std::iter::repeat(Duration::ZERO);
150 let make_client = || async {
151 Err::<map_core::client::MapClient<tokio::io::DuplexStream>, _>(SessionError::Transport(
152 obex_core::TransportError::Io(std::io::Error::new(
153 std::io::ErrorKind::ConnectionRefused,
154 "refused",
155 )),
156 ))
157 };
158 run_session_loop(make_client, strategy, state_tx, cancel_rx).await;
159 let _ = cancel_tx;
160 }
161
162 #[tokio::test]
163 async fn enters_active_and_reconnects_on_stream_close() {
164 let (state_tx, mut state_rx) = watch::channel(SessionState::Disconnected);
165 let (cancel_tx, cancel_rx) = watch::channel(false);
166
167 let (client_io, server_io) = tokio::io::duplex(4096);
168 let client_cell = std::cell::Cell::new(Some(client_io));
169
170 let make_client = move || {
171 let io = client_cell.take().unwrap_or_else(|| {
172 let (io, _srv) = tokio::io::duplex(1);
173 io
174 });
175 async move { lifecycle::establish_map_session(io).await }
176 };
177
178 let strategy = std::iter::repeat(Duration::ZERO);
179
180 let (server_result, ()) = futures::join!(
181 async {
182 let mut srv = obex_core::wrap(server_io);
183 let _ = srv.next().await;
184 srv.send(Bytes::from_static(MAP_CONNECT_RSP))
185 .await
186 .map_err(SessionError::Transport)?;
187 let _ = srv.next().await;
188 srv.send(Bytes::from_static(NOTIF_REG_OK))
189 .await
190 .map_err(SessionError::Transport)?;
191 state_rx.wait_for(|s| *s == SessionState::Active).await.map_err(|_| {
192 SessionError::Transport(obex_core::TransportError::Io(std::io::Error::new(
193 std::io::ErrorKind::BrokenPipe,
194 "watch dropped",
195 )))
196 })?;
197 drop(srv);
198 state_rx.wait_for(|s| *s == SessionState::Reconnecting).await.map_err(|_| {
199 SessionError::Transport(obex_core::TransportError::Io(std::io::Error::new(
200 std::io::ErrorKind::BrokenPipe,
201 "watch dropped",
202 )))
203 })?;
204 cancel_tx.send(true).ok();
205 Ok::<(), SessionError>(())
206 },
207 run_session_loop(make_client, strategy, state_tx, cancel_rx),
208 );
209 assert!(server_result.is_ok());
210 }
211
212 #[tokio::test]
213 async fn state_is_connecting_before_connect_attempt() {
214 let (state_tx, mut state_rx) = watch::channel(SessionState::Disconnected);
215 let (cancel_tx, cancel_rx) = watch::channel(false);
216 let strategy = std::iter::once(Duration::ZERO);
217
218 let make_client = move || async {
219 Err::<map_core::client::MapClient<tokio::io::DuplexStream>, _>(SessionError::Transport(
220 obex_core::TransportError::Io(std::io::Error::new(
221 std::io::ErrorKind::ConnectionRefused,
222 "refused",
223 )),
224 ))
225 };
226
227 let ((), ()) = futures::join!(
228 async {
229 state_rx.wait_for(|s| *s == SessionState::Connecting).await.ok();
230 cancel_tx.send(true).ok();
231 },
232 run_session_loop(make_client, strategy, state_tx, cancel_rx),
233 );
234 }
235}