pb_mapper_client/sdk/
handle.rs1use crate::diagnostics::{Diagnostics, TunnelDiagnostics};
8use crate::endpoint::RelayEndpoint;
9use std::time::Duration;
10
11use tokio::sync::{Mutex, watch};
12use tokio::task::JoinHandle;
13use tokio_util::sync::CancellationToken;
14
15use super::{Error, Result, TunnelStatus};
16
17pub(crate) struct LiveTunnel {
18 shutdown: CancellationToken,
19 join: Mutex<Option<JoinHandle<()>>>,
20 status: watch::Receiver<TunnelStatus>,
21 endpoint: RelayEndpoint,
22 diagnostics: Diagnostics,
23}
24
25impl LiveTunnel {
26 pub(crate) fn new(
27 shutdown: CancellationToken,
28 join: JoinHandle<()>,
29 status: watch::Receiver<TunnelStatus>,
30 endpoint: RelayEndpoint,
31 diagnostics: Diagnostics,
32 ) -> Self {
33 Self {
34 shutdown,
35 join: Mutex::new(Some(join)),
36 status,
37 endpoint,
38 diagnostics,
39 }
40 }
41
42 pub(crate) fn status(&self) -> TunnelStatus {
43 self.status.borrow().clone()
44 }
45
46 pub(crate) fn subscribe(&self) -> watch::Receiver<TunnelStatus> {
47 self.status.clone()
48 }
49
50 pub(crate) async fn wait_ready(&self) -> Result<()> {
51 wait_for_connected(&mut self.status.clone()).await
52 }
53
54 pub(crate) async fn wait_ready_timeout(&self, timeout: Duration) -> Result<()> {
55 match tokio::time::timeout(timeout, self.wait_ready()).await {
56 Ok(result) => result,
57 Err(_) => Err(Error::ReadyTimeout { timeout }),
58 }
59 }
60
61 pub(crate) async fn stop(&self) -> Result<()> {
62 self.shutdown.cancel();
63 let mut slot = self.join.lock().await;
66 if let Some(handle) = slot.as_mut() {
67 match tokio::time::timeout(Duration::from_secs(5), &mut *handle).await {
68 Ok(Ok(())) => {}
69 Ok(Err(join_error)) if join_error.is_cancelled() => {}
70 Ok(Err(join_error)) => {
71 slot.take();
72 return Err(Error::protocol(format!("tunnel task failed: {join_error}")));
73 }
74 Err(_) => {
75 handle.abort();
76 let _ = handle.await;
77 }
78 }
79 slot.take();
80 }
81 Ok(())
82 }
83}
84
85impl Drop for LiveTunnel {
86 fn drop(&mut self) {
87 self.shutdown.cancel();
88 if let Some(handle) = self.join.get_mut().take() {
89 handle.abort();
90 }
91 }
92}
93
94async fn wait_for_connected(status: &mut watch::Receiver<TunnelStatus>) -> Result<()> {
95 loop {
96 match status.borrow().clone() {
97 TunnelStatus::Connected => return Ok(()),
98 TunnelStatus::Failed(reason) => return Err(Error::TunnelFailed { reason }),
99 TunnelStatus::Stopped => return Err(Error::Stopped),
100 TunnelStatus::Starting | TunnelStatus::Retrying => {}
101 }
102 status.changed().await.map_err(|_| Error::Stopped)?;
103 }
104}
105
106macro_rules! tunnel_handle {
113 ($(#[$doc:meta])* $name:ident) => {
114 $(#[$doc])*
115 pub struct $name {
116 inner: LiveTunnel,
117 key: String,
118 }
119
120 impl $name {
121 pub(crate) fn new(inner: LiveTunnel, key: String) -> Self {
122 Self { inner, key }
123 }
124
125 pub fn key(&self) -> &str {
127 &self.key
128 }
129
130 pub fn status(&self) -> TunnelStatus {
132 self.inner.status()
133 }
134
135 pub fn diagnostics(&self) -> TunnelDiagnostics {
137 self.inner.diagnostics.snapshot(&self.inner.endpoint)
138 }
139
140 pub fn subscribe(&self) -> watch::Receiver<TunnelStatus> {
142 self.inner.subscribe()
143 }
144
145 pub async fn wait_ready(&self) -> Result<()> {
147 self.inner.wait_ready().await
148 }
149
150 pub async fn wait_ready_timeout(&self, timeout: Duration) -> Result<()> {
152 self.inner.wait_ready_timeout(timeout).await
153 }
154
155 pub async fn stop(&self) -> Result<()> {
157 self.inner.stop().await
158 }
159 }
160 };
161}
162
163tunnel_handle!(
164 Registration
166);
167
168tunnel_handle!(
169 Connection
171);
172
173#[cfg(test)]
174#[allow(clippy::unwrap_used)]
175mod tests {
176 use super::*;
177 use std::{future::Future, task::Poll};
178
179 async fn poll_pending(future: &mut std::pin::Pin<Box<impl Future>>) {
180 std::future::poll_fn(|cx| {
181 assert!(future.as_mut().poll(cx).is_pending());
182 Poll::Ready(())
183 })
184 .await;
185 }
186
187 #[tokio::test]
188 async fn cancelled_stop_keeps_worker_owned_until_handle_drop() {
189 let (tx, mut rx) = watch::channel(TunnelStatus::Starting);
190 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
191 let worker = tokio::spawn(async move {
192 let _tx = tx;
193 let _ = started_tx.send(());
194 std::future::pending::<()>().await;
195 });
196 let tunnel = LiveTunnel::new(
197 CancellationToken::new(),
198 worker,
199 rx.clone(),
200 RelayEndpoint::shared("127.0.0.1:1"),
201 Diagnostics::default(),
202 );
203 started_rx.await.unwrap();
204 let mut stopping = Box::pin(tunnel.stop());
205 poll_pending(&mut stopping).await;
206 drop(stopping);
207 drop(tunnel);
208 assert!(
209 tokio::time::timeout(Duration::from_secs(1), rx.changed())
210 .await
211 .unwrap()
212 .is_err()
213 );
214 }
215
216 #[tokio::test]
217 async fn concurrent_stop_calls_both_wait_for_cleanup() {
218 let (tx, rx) = watch::channel(TunnelStatus::Starting);
219 let (release, released) = tokio::sync::oneshot::channel();
220 let worker = tokio::spawn(async move {
221 let _tx = tx;
222 let _ = released.await;
223 });
224 let tunnel = LiveTunnel::new(
225 CancellationToken::new(),
226 worker,
227 rx,
228 RelayEndpoint::shared("127.0.0.1:1"),
229 Diagnostics::default(),
230 );
231 let mut first = Box::pin(tunnel.stop());
232 let mut second = Box::pin(tunnel.stop());
233 poll_pending(&mut first).await;
234 poll_pending(&mut second).await;
235 release.send(()).unwrap();
236 first.await.unwrap();
237 second.await.unwrap();
238 }
239}