use std::net::SocketAddr;
use std::time::Duration;
use tokio::sync::broadcast;
use crate::server::errors::{ServerError, ShutdownResult};
use crate::server::lifecycle::Lifecycle;
pub struct ServerHandle {
local_addr: SocketAddr,
shutdown_tx: broadcast::Sender<()>,
join: Option<tokio::task::JoinHandle<ShutdownResult>>,
lifecycle: std::sync::Arc<Lifecycle>,
}
impl std::fmt::Debug for ServerHandle {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ServerHandle")
.field("local_addr", &self.local_addr)
.field("state", &self.lifecycle.state())
.finish()
}
}
impl ServerHandle {
pub(crate) fn new(
local_addr: SocketAddr,
shutdown_tx: broadcast::Sender<()>,
join: tokio::task::JoinHandle<ShutdownResult>,
lifecycle: std::sync::Arc<Lifecycle>,
) -> Self {
Self {
local_addr,
shutdown_tx,
join: Some(join),
lifecycle,
}
}
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub fn state(&self) -> crate::server::lifecycle::LifecycleState {
self.lifecycle.state()
}
pub async fn ready(&self) -> Result<(), ServerError> {
let state = self.lifecycle.state();
match state {
crate::server::lifecycle::LifecycleState::Running => Ok(()),
crate::server::lifecycle::LifecycleState::Starting => {
self.lifecycle.wait_ready().await;
let state = self.lifecycle.state();
match state {
crate::server::lifecycle::LifecycleState::Running => Ok(()),
crate::server::lifecycle::LifecycleState::Failed => {
Err(ServerError::Startup("server failed during startup".into()))
}
other => Err(ServerError::Config(format!(
"unexpected state after ready: {other}"
))),
}
}
crate::server::lifecycle::LifecycleState::Failed => {
Err(ServerError::Startup("server failed during startup".into()))
}
other => Err(ServerError::Config(format!(
"server not ready: in {other} state"
))),
}
}
pub fn shutdown(&self) {
let _ = self.lifecycle.drain();
let _ = self.shutdown_tx.send(());
}
pub async fn force_shutdown(
mut self,
deadline: Duration,
) -> Result<ShutdownResult, ServerError> {
self.shutdown();
match tokio::time::timeout(deadline, self.wait_internal()).await {
Ok(()) => {
if let Some(join) = self.join.take() {
match join.await {
Ok(result) => Ok(result),
Err(e) => Err(ServerError::Accept(std::io::Error::other(format!(
"server task panicked: {}",
e
)))),
}
} else {
Ok(ShutdownResult::Clean)
}
}
Err(_deadline_exceeded) => {
if let Some(join) = self.join.take() {
join.abort();
let _ = join.await;
}
let _ = self.lifecycle.mark_stopped();
Ok(ShutdownResult::Forced)
}
}
}
pub async fn wait(mut self) -> Result<ShutdownResult, ServerError> {
let state = self.lifecycle.state();
if !state.is_terminal() {
self.shutdown();
}
self.wait_internal().await;
if let Some(join) = self.join.take() {
match join.await {
Ok(result) => Ok(result),
Err(e) => Err(ServerError::Accept(std::io::Error::other(format!(
"server task panicked: {}",
e
)))),
}
} else {
Ok(ShutdownResult::Clean)
}
}
async fn wait_internal(&self) {
let mut terminal_rx = self.lifecycle.subscribe_terminal();
let state = self.lifecycle.state();
if state.is_terminal() {
return;
}
let _ = terminal_rx.recv().await;
}
}
impl Drop for ServerHandle {
fn drop(&mut self) {
if self.join.is_some() {
let _ = self.lifecycle.drain();
let _ = self.shutdown_tx.send(());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::server::lifecycle::Lifecycle;
use std::sync::Arc;
async fn make_test_handle() -> ServerHandle {
let lifecycle = Arc::new(Lifecycle::new());
let (tx, _rx) = broadcast::channel(1);
let join = tokio::spawn(async { ShutdownResult::Clean });
ServerHandle::new("127.0.0.1:8000".parse().unwrap(), tx, join, lifecycle)
}
fn make_handle_with_state(state: crate::server::lifecycle::LifecycleState) -> ServerHandle {
let lifecycle = Arc::new(Lifecycle::new());
match state {
crate::server::lifecycle::LifecycleState::Created => {}
crate::server::lifecycle::LifecycleState::Starting => {
lifecycle.start().unwrap();
}
crate::server::lifecycle::LifecycleState::Running => {
lifecycle.start().unwrap();
lifecycle.mark_running().unwrap();
}
crate::server::lifecycle::LifecycleState::Failed => {
lifecycle.mark_failed().unwrap();
}
crate::server::lifecycle::LifecycleState::Draining => {
lifecycle.start().unwrap();
lifecycle.mark_running().unwrap();
lifecycle.drain().unwrap();
}
crate::server::lifecycle::LifecycleState::Stopped => {
lifecycle.start().unwrap();
lifecycle.mark_running().unwrap();
lifecycle.drain().unwrap();
lifecycle.mark_stopped().unwrap();
}
}
let (shutdown_tx, _) = broadcast::channel(1);
let join = tokio::spawn(async { ShutdownResult::Clean });
ServerHandle::new("127.0.0.1:0".parse().unwrap(), shutdown_tx, join, lifecycle)
}
#[tokio::test]
async fn handle_local_addr() {
let handle = make_test_handle().await;
assert_eq!(
handle.local_addr(),
"127.0.0.1:8000".parse::<SocketAddr>().unwrap()
);
}
#[tokio::test]
async fn handle_state_initial() {
let handle = make_test_handle().await;
assert_eq!(
handle.state(),
crate::server::lifecycle::LifecycleState::Created
);
}
#[tokio::test]
async fn handle_shutdown_sends_signal() {
let lifecycle = Arc::new(Lifecycle::new());
lifecycle.start().unwrap();
lifecycle.mark_running().unwrap();
let (tx, mut rx) = broadcast::channel(1);
let join = tokio::spawn(async move {
let _ = rx.recv().await;
ShutdownResult::Clean
});
let handle = ServerHandle::new("127.0.0.1:0".parse().unwrap(), tx, join, lifecycle);
handle.shutdown();
}
#[tokio::test]
async fn handle_ready_returns_error_for_failed() {
let lifecycle = Arc::new(Lifecycle::new());
lifecycle.mark_failed().unwrap();
let (tx, _rx) = broadcast::channel(1);
let join = tokio::spawn(async { ShutdownResult::Clean });
let handle = ServerHandle::new("127.0.0.1:0".parse().unwrap(), tx, join, lifecycle);
let result = handle.ready().await;
assert!(result.is_err());
}
#[tokio::test]
async fn handle_debug_format() {
let handle = make_test_handle().await;
let debug = format!("{:?}", handle);
assert!(debug.contains("ServerHandle"));
assert!(debug.contains("127.0.0.1:8000"));
}
#[tokio::test]
async fn ready_already_running_returns_ok() {
let lifecycle = Arc::new(Lifecycle::new());
lifecycle.start().unwrap();
lifecycle.mark_running().unwrap();
assert_eq!(
lifecycle.state(),
crate::server::lifecycle::LifecycleState::Running
);
let (tx, _rx) = broadcast::channel(1);
let join = tokio::spawn(async { ShutdownResult::Clean });
let handle = ServerHandle::new("127.0.0.1:0".parse().unwrap(), tx, join, lifecycle);
let result = handle.ready().await;
assert!(
result.is_ok(),
"ready() on already-Running server: {:?}",
result.err()
);
}
#[tokio::test]
async fn ready_failed_returns_error() {
let handle = make_handle_with_state(crate::server::lifecycle::LifecycleState::Failed);
let result = handle.ready().await;
assert!(result.is_err());
}
#[tokio::test]
async fn ready_starting_then_running_succeeds() {
let lifecycle = Arc::new(Lifecycle::new());
lifecycle.start().unwrap();
let (tx, _) = broadcast::channel(1);
let join = tokio::spawn(async { ShutdownResult::Clean });
let handle = ServerHandle::new("127.0.0.1:0".parse().unwrap(), tx, join, lifecycle.clone());
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
lifecycle.mark_running().unwrap();
});
let result = tokio::time::timeout(Duration::from_secs(5), handle.ready()).await;
assert!(result.is_ok());
assert!(result.unwrap().is_ok());
}
#[tokio::test]
async fn ready_starting_then_failed_returns_error() {
let lifecycle = Arc::new(Lifecycle::new());
lifecycle.start().unwrap();
let (tx, _) = broadcast::channel(1);
let join = tokio::spawn(async { ShutdownResult::Clean });
let handle = ServerHandle::new("127.0.0.1:0".parse().unwrap(), tx, join, lifecycle.clone());
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
lifecycle.mark_failed().unwrap();
});
let result = handle.ready().await;
assert!(result.is_err());
}
#[tokio::test]
async fn ready_stuck_starting_times_out() {
let handle = make_handle_with_state(crate::server::lifecycle::LifecycleState::Starting);
let result = tokio::time::timeout(Duration::from_millis(50), handle.ready()).await;
assert!(result.is_err());
}
#[tokio::test]
async fn ready_draining_is_error() {
let handle = make_handle_with_state(crate::server::lifecycle::LifecycleState::Draining);
let result = handle.ready().await;
assert!(result.is_err());
}
#[tokio::test]
async fn ready_stopped_is_error() {
let handle = make_handle_with_state(crate::server::lifecycle::LifecycleState::Stopped);
let result = handle.ready().await;
assert!(result.is_err());
}
#[tokio::test]
async fn ready_idempotent_on_running() {
let handle = make_handle_with_state(crate::server::lifecycle::LifecycleState::Running);
let r1 = tokio::time::timeout(Duration::from_millis(50), handle.ready()).await;
assert!(r1.is_ok() && r1.unwrap().is_ok());
let r2 = tokio::time::timeout(Duration::from_millis(50), handle.ready()).await;
assert!(r2.is_ok() && r2.unwrap().is_ok());
}
}