Skip to main content

webserver_base/webserver/
shutdown.rs

1//! Coordinated graceful shutdown.
2
3use std::sync::Arc;
4use std::time::Duration;
5
6use tokio::sync::watch;
7use tracing::{info, instrument, warn};
8
9/// The longest in-flight work may take once shutdown begins.
10///
11/// A ceiling, never a wait: a server holding no connection exits the instant it
12/// is signalled. The ceiling is what makes graceful shutdown terminate at all —
13/// `with_graceful_shutdown` waits for every connection, and a WebSocket never
14/// closes on its own, so without a deadline a socket-holding server hangs until
15/// it is killed.
16pub const DEFAULT_DRAIN_TIMEOUT: Duration = Duration::from_secs(10);
17
18/// A cloneable handle that resolves when the process should stop.
19///
20/// One listener, many holders — which is what lets a binary run several servers
21/// and drain them together.
22#[derive(Clone, Debug)]
23pub struct Shutdown {
24    sender: Arc<watch::Sender<bool>>,
25    receiver: watch::Receiver<bool>,
26}
27
28impl Shutdown {
29    /// A handle fired by hand, for tests and custom signal handling.
30    #[must_use]
31    pub fn manual() -> Self {
32        let (sender, receiver) = watch::channel(false);
33        Self {
34            sender: Arc::new(sender),
35            receiver,
36        }
37    }
38
39    /// A handle wired to this platform's termination signals: `SIGTERM`,
40    /// `SIGINT` and `SIGQUIT` on unix, ctrl-c elsewhere.
41    ///
42    /// The first signal starts the drain; a second exits immediately. Must be
43    /// called from inside a Tokio runtime.
44    #[must_use]
45    #[instrument(skip_all)]
46    pub fn listen() -> Self {
47        let shutdown: Self = Self::manual();
48        let trigger: Self = shutdown.clone();
49
50        tokio::spawn(async move {
51            wait_for_signal().await;
52            info!("shutdown signal received; draining");
53            trigger.trigger();
54
55            wait_for_signal().await;
56            warn!("second shutdown signal received; exiting immediately");
57            std::process::exit(130);
58        });
59
60        shutdown
61    }
62
63    /// Starts the drain.
64    pub fn trigger(&self) {
65        // A send failure means every receiver is already gone.
66        let _ = self.sender.send(true);
67    }
68
69    /// Whether the drain has started.
70    #[must_use]
71    pub fn is_shutting_down(&self) -> bool {
72        *self.receiver.borrow()
73    }
74
75    /// Resolves when the drain starts, immediately if it already has.
76    ///
77    /// Consumes the handle so the future is `'static`; clone first if the
78    /// handle is still needed.
79    pub async fn recv(mut self) {
80        if *self.receiver.borrow_and_update() {
81            return;
82        }
83        let _ = self.receiver.changed().await;
84    }
85}
86
87/// Waits for whichever termination signal arrives first.
88#[cfg(unix)]
89async fn wait_for_signal() {
90    use tokio::signal::unix::{SignalKind, signal};
91
92    let mut terminate = signal(SignalKind::terminate()).expect("failed to install SIGTERM handler");
93    let mut interrupt = signal(SignalKind::interrupt()).expect("failed to install SIGINT handler");
94    let mut quit = signal(SignalKind::quit()).expect("failed to install SIGQUIT handler");
95
96    tokio::select! {
97        _ = terminate.recv() => {}
98        _ = interrupt.recv() => {}
99        _ = quit.recv() => {}
100    }
101}
102
103/// Waits for whichever termination signal arrives first.
104#[cfg(not(unix))]
105async fn wait_for_signal() {
106    let _ = tokio::signal::ctrl_c().await;
107}
108
109#[cfg(test)]
110mod tests {
111    use std::time::Duration;
112
113    use tokio::time::timeout;
114
115    use super::Shutdown;
116
117    #[tokio::test]
118    async fn a_fresh_handle_is_not_shutting_down_and_does_not_resolve() {
119        let shutdown: Shutdown = Shutdown::manual();
120
121        let expected: bool = false;
122        let actual: bool = shutdown.is_shutting_down();
123        assert_eq!(expected, actual);
124
125        let resolved: bool = timeout(Duration::from_millis(20), shutdown.clone().recv())
126            .await
127            .is_ok();
128        assert!(!resolved, "recv resolved before anything triggered it");
129    }
130
131    #[tokio::test]
132    async fn every_clone_hears_one_trigger() {
133        let shutdown: Shutdown = Shutdown::manual();
134        let first: Shutdown = shutdown.clone();
135        let second: Shutdown = shutdown.clone();
136
137        shutdown.trigger();
138
139        timeout(Duration::from_millis(200), first.recv())
140            .await
141            .expect("the first clone resolved");
142        timeout(Duration::from_millis(200), second.recv())
143            .await
144            .expect("the second clone resolved");
145    }
146
147    #[tokio::test]
148    async fn recv_after_the_fact_resolves_immediately() {
149        let shutdown: Shutdown = Shutdown::manual();
150        shutdown.trigger();
151
152        let expected: bool = true;
153        let actual: bool = shutdown.is_shutting_down();
154        assert_eq!(expected, actual);
155
156        timeout(Duration::from_millis(50), shutdown.clone().recv())
157            .await
158            .expect("a handle created before the trigger still resolves after it");
159    }
160
161    #[tokio::test]
162    async fn a_handle_cloned_after_the_trigger_still_resolves() {
163        let shutdown: Shutdown = Shutdown::manual();
164        shutdown.trigger();
165
166        let late: Shutdown = shutdown.clone();
167        timeout(Duration::from_millis(50), late.recv())
168            .await
169            .expect("a late clone sees the state, not just the transition");
170    }
171}