Skip to main content

imsg_session/
reconnect.rs

1//! Reconnect policy: back-off timing and retry logic for dropped RFCOMM sessions.
2
3use 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/// Disconnected → Connecting → Active → Reconnecting.
14#[derive(Debug, Clone, Copy, PartialEq, Eq)]
15pub enum SessionState {
16    /// Initial state; no connection attempted yet.
17    Disconnected,
18    /// RFCOMM connect + OBEX handshake ongoing.
19    Connecting,
20    /// OBEX CONNECT + `SetNotificationRegistration` done; session live.
21    Active,
22    /// Session dropped; backoff running before retry.
23    Reconnecting,
24}
25
26/// Exponential backoff (2 s → 60 s, ×2, with jitter). Retries indefinitely.
27pub 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
41/// Inner reconnect loop. Resets backoff to start after each successful session.
42pub(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
76/// Hold session until drop or cancel. Returns `true` to stop loop, `false` to retry.
77async 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
95/// Wait one backoff delay. Returns `true` when strategy exhausted or cancelled.
96async 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
106/// Race `until` against cancellation. Returns `true` when cancel fired.
107async 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}