use crate::diagnostics::{Diagnostics, TunnelDiagnostics};
use crate::endpoint::RelayEndpoint;
use std::time::Duration;
use tokio::sync::{Mutex, watch};
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use super::{Error, Result, TunnelStatus};
pub(crate) struct LiveTunnel {
shutdown: CancellationToken,
join: Mutex<Option<JoinHandle<()>>>,
status: watch::Receiver<TunnelStatus>,
endpoint: RelayEndpoint,
diagnostics: Diagnostics,
}
impl LiveTunnel {
pub(crate) fn new(
shutdown: CancellationToken,
join: JoinHandle<()>,
status: watch::Receiver<TunnelStatus>,
endpoint: RelayEndpoint,
diagnostics: Diagnostics,
) -> Self {
Self {
shutdown,
join: Mutex::new(Some(join)),
status,
endpoint,
diagnostics,
}
}
pub(crate) fn status(&self) -> TunnelStatus {
self.status.borrow().clone()
}
pub(crate) fn subscribe(&self) -> watch::Receiver<TunnelStatus> {
self.status.clone()
}
pub(crate) async fn wait_ready(&self) -> Result<()> {
wait_for_connected(&mut self.status.clone()).await
}
pub(crate) async fn wait_ready_timeout(&self, timeout: Duration) -> Result<()> {
match tokio::time::timeout(timeout, self.wait_ready()).await {
Ok(result) => result,
Err(_) => Err(Error::ReadyTimeout { timeout }),
}
}
pub(crate) async fn stop(&self) -> Result<()> {
self.shutdown.cancel();
let mut slot = self.join.lock().await;
if let Some(handle) = slot.as_mut() {
match tokio::time::timeout(Duration::from_secs(5), &mut *handle).await {
Ok(Ok(())) => {}
Ok(Err(join_error)) if join_error.is_cancelled() => {}
Ok(Err(join_error)) => {
slot.take();
return Err(Error::protocol(format!("tunnel task failed: {join_error}")));
}
Err(_) => {
handle.abort();
let _ = handle.await;
}
}
slot.take();
}
Ok(())
}
}
impl Drop for LiveTunnel {
fn drop(&mut self) {
self.shutdown.cancel();
if let Some(handle) = self.join.get_mut().take() {
handle.abort();
}
}
}
async fn wait_for_connected(status: &mut watch::Receiver<TunnelStatus>) -> Result<()> {
loop {
match status.borrow().clone() {
TunnelStatus::Connected => return Ok(()),
TunnelStatus::Failed(reason) => return Err(Error::TunnelFailed { reason }),
TunnelStatus::Stopped => return Err(Error::Stopped),
TunnelStatus::Starting | TunnelStatus::Retrying => {}
}
status.changed().await.map_err(|_| Error::Stopped)?;
}
}
macro_rules! tunnel_handle {
($(#[$doc:meta])* $name:ident) => {
$(#[$doc])*
pub struct $name {
inner: LiveTunnel,
key: String,
}
impl $name {
pub(crate) fn new(inner: LiveTunnel, key: String) -> Self {
Self { inner, key }
}
pub fn key(&self) -> &str {
&self.key
}
pub fn status(&self) -> TunnelStatus {
self.inner.status()
}
pub fn diagnostics(&self) -> TunnelDiagnostics {
self.inner.diagnostics.snapshot(&self.inner.endpoint)
}
pub fn subscribe(&self) -> watch::Receiver<TunnelStatus> {
self.inner.subscribe()
}
pub async fn wait_ready(&self) -> Result<()> {
self.inner.wait_ready().await
}
pub async fn wait_ready_timeout(&self, timeout: Duration) -> Result<()> {
self.inner.wait_ready_timeout(timeout).await
}
pub async fn stop(&self) -> Result<()> {
self.inner.stop().await
}
}
};
}
tunnel_handle!(
Registration
);
tunnel_handle!(
Connection
);
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
use std::{future::Future, task::Poll};
async fn poll_pending(future: &mut std::pin::Pin<Box<impl Future>>) {
std::future::poll_fn(|cx| {
assert!(future.as_mut().poll(cx).is_pending());
Poll::Ready(())
})
.await;
}
#[tokio::test]
async fn cancelled_stop_keeps_worker_owned_until_handle_drop() {
let (tx, mut rx) = watch::channel(TunnelStatus::Starting);
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let worker = tokio::spawn(async move {
let _tx = tx;
let _ = started_tx.send(());
std::future::pending::<()>().await;
});
let tunnel = LiveTunnel::new(
CancellationToken::new(),
worker,
rx.clone(),
RelayEndpoint::shared("127.0.0.1:1"),
Diagnostics::default(),
);
started_rx.await.unwrap();
let mut stopping = Box::pin(tunnel.stop());
poll_pending(&mut stopping).await;
drop(stopping);
drop(tunnel);
assert!(
tokio::time::timeout(Duration::from_secs(1), rx.changed())
.await
.unwrap()
.is_err()
);
}
#[tokio::test]
async fn concurrent_stop_calls_both_wait_for_cleanup() {
let (tx, rx) = watch::channel(TunnelStatus::Starting);
let (release, released) = tokio::sync::oneshot::channel();
let worker = tokio::spawn(async move {
let _tx = tx;
let _ = released.await;
});
let tunnel = LiveTunnel::new(
CancellationToken::new(),
worker,
rx,
RelayEndpoint::shared("127.0.0.1:1"),
Diagnostics::default(),
);
let mut first = Box::pin(tunnel.stop());
let mut second = Box::pin(tunnel.stop());
poll_pending(&mut first).await;
poll_pending(&mut second).await;
release.send(()).unwrap();
first.await.unwrap();
second.await.unwrap();
}
}