use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct ServeToken(u64);
impl ServeToken {
pub(crate) fn new(value: u64) -> Self {
Self(value)
}
pub(crate) fn get(self) -> u64 {
self.0
}
}
struct ServeHandleInner {
token: ServeToken,
join: Mutex<Option<tokio::task::JoinHandle<()>>>,
abort_requested: AtomicBool,
shutdown_requested: AtomicBool,
shutdown_notify: Arc<tokio::sync::Notify>,
close_flag: Arc<AtomicBool>,
close_connections: Arc<tokio::sync::Notify>,
drain_timeout: std::time::Duration,
done_rx: tokio::sync::watch::Receiver<bool>,
}
#[derive(Clone)]
pub struct ServeHandle {
inner: Arc<ServeHandleInner>,
}
impl ServeHandle {
pub(super) fn pending(
token: ServeToken,
shutdown_notify: Arc<tokio::sync::Notify>,
close_flag: Arc<AtomicBool>,
close_connections: Arc<tokio::sync::Notify>,
drain_timeout: std::time::Duration,
done_rx: tokio::sync::watch::Receiver<bool>,
) -> Self {
Self {
inner: Arc::new(ServeHandleInner {
token,
join: Mutex::new(None),
abort_requested: AtomicBool::new(false),
shutdown_requested: AtomicBool::new(false),
shutdown_notify,
close_flag,
close_connections,
drain_timeout,
done_rx,
}),
}
}
pub(super) fn attach_join(&self, join: tokio::task::JoinHandle<()>) {
let mut slot = self
.inner
.join
.lock()
.unwrap_or_else(|error| error.into_inner());
if self.inner.abort_requested.load(Ordering::Acquire) {
join.abort();
} else {
*slot = Some(join);
}
}
pub(crate) fn token(&self) -> ServeToken {
self.inner.token
}
pub(crate) fn is_same_cycle(&self, other: &Self) -> bool {
self.token() == other.token()
}
pub(crate) fn is_shutdown_requested(&self) -> bool {
self.inner.shutdown_requested.load(Ordering::Acquire)
}
pub fn shutdown(&self) {
self.inner.shutdown_requested.store(true, Ordering::Release);
self.inner.shutdown_notify.notify_one();
}
pub fn shutdown_and_close(&self) {
self.inner.shutdown_requested.store(true, Ordering::Release);
self.inner.close_flag.store(true, Ordering::Release);
self.inner.shutdown_notify.notify_one();
self.inner.close_connections.notify_waiters();
}
pub async fn drain(self) {
self.shutdown();
let join = self
.inner
.join
.lock()
.unwrap_or_else(|error| error.into_inner())
.take();
if let Some(join) = join {
let _ = join.await;
} else {
let mut done = self.subscribe_done();
let _ = done.wait_for(|value| *value).await;
}
}
pub fn abort(&self) {
self.inner.abort_requested.store(true, Ordering::Release);
if let Some(join) = self
.inner
.join
.lock()
.unwrap_or_else(|error| error.into_inner())
.as_ref()
{
join.abort();
}
}
pub fn drain_timeout(&self) -> std::time::Duration {
self.inner.drain_timeout
}
pub fn subscribe_done(&self) -> tokio::sync::watch::Receiver<bool> {
self.inner.done_rx.clone()
}
}