1use crate::{ManagedService, ServiceRuntimeState, SupervisorError, SupervisorResult};
14use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
15use std::sync::mpsc::{self, Receiver, RecvTimeoutError, SyncSender, TrySendError};
16use std::sync::{Arc, Mutex};
17use std::thread::JoinHandle;
18use std::time::{Duration, Instant};
19
20pub(crate) const DEFAULT_RESTART_QUEUE_CAPACITY: usize = 64;
21pub(crate) const DEFAULT_RESTART_WORKERS: usize = 2;
22const RESTART_THREAD_STACK_BYTES: usize = 1024 * 1024;
23
24#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum RestartState {
27 None,
29 Scheduled {
31 execute_at_ms: u64,
33 },
34 Stopping,
36 Starting,
38 Backoff,
40 Failed,
42}
43
44pub(crate) struct RestartCommand {
45 pub service: Arc<dyn ManagedService>,
46 pub attempt: u64,
47}
48
49pub(crate) enum RestartOutcome {
50 Restarted,
51 Orphaned,
52 Failed,
53 Cancelled,
54}
55
56pub(crate) struct RestartCompletion {
57 pub service_id: String,
58 pub attempt: u64,
59 pub outcome: RestartOutcome,
60}
61
62pub(crate) struct RestartExecutor {
63 sender: Mutex<Option<SyncSender<RestartCommand>>>,
64 commands: Arc<Mutex<Receiver<RestartCommand>>>,
65 completions: Mutex<Receiver<RestartCompletion>>,
66 workers: Mutex<Vec<JoinHandle<()>>>,
67 cancellation: Arc<AtomicBool>,
68 healthy: Arc<AtomicBool>,
69 pending: Arc<AtomicU64>,
70 queue_capacity: usize,
71 worker_count: usize,
72}
73
74impl RestartExecutor {
75 pub fn new(queue_capacity: usize, worker_count: usize) -> Self {
76 let queue_capacity = queue_capacity.max(1);
77 let worker_count = worker_count.max(1);
78 let (sender, receiver) = mpsc::sync_channel::<RestartCommand>(queue_capacity);
79 let completion_capacity = queue_capacity.saturating_add(worker_count);
80 let (completion_sender, completions) =
81 mpsc::sync_channel::<RestartCompletion>(completion_capacity);
82 let receiver = Arc::new(Mutex::new(receiver));
83 let cancellation = Arc::new(AtomicBool::new(false));
84 let healthy = Arc::new(AtomicBool::new(true));
85 let pending = Arc::new(AtomicU64::new(0));
86 let workers = spawn_workers(
87 worker_count,
88 Arc::clone(&receiver),
89 completion_sender,
90 Arc::clone(&cancellation),
91 Arc::clone(&healthy),
92 Arc::clone(&pending),
93 );
94 Self {
95 sender: Mutex::new(Some(sender)),
96 commands: receiver,
97 completions: Mutex::new(completions),
98 workers: Mutex::new(workers),
99 cancellation,
100 healthy,
101 pending,
102 queue_capacity,
103 worker_count,
104 }
105 }
106
107 pub fn schedule(&self, command: RestartCommand) -> SupervisorResult<()> {
108 if self.cancellation.load(Ordering::Acquire) {
109 return Err(SupervisorError::RestartExecutorStopped);
110 }
111 let sender = self
112 .sender
113 .lock()
114 .map_err(|_| SupervisorError::RestartExecutorStopped)?;
115 let sender = sender
116 .as_ref()
117 .ok_or(SupervisorError::RestartExecutorStopped)?;
118 self.pending.fetch_add(1, Ordering::AcqRel);
119 match sender.try_send(command) {
120 Ok(()) => Ok(()),
121 Err(TrySendError::Full(_)) => {
122 self.pending.fetch_sub(1, Ordering::AcqRel);
123 Err(SupervisorError::RestartQueueFull)
124 }
125 Err(TrySendError::Disconnected(_)) => {
126 self.pending.fetch_sub(1, Ordering::AcqRel);
127 Err(SupervisorError::RestartExecutorStopped)
128 }
129 }
130 }
131
132 pub fn drain_completions(&self) -> Vec<RestartCompletion> {
133 let Ok(receiver) = self.completions.lock() else {
134 self.healthy.store(false, Ordering::Release);
135 return Vec::new();
136 };
137 receiver.try_iter().collect()
138 }
139
140 pub fn snapshot(&self) -> crate::RestartExecutorSnapshot {
141 let workers_healthy = self
142 .workers
143 .lock()
144 .map(|workers| workers.iter().all(|worker| !worker.is_finished()))
145 .unwrap_or(false);
146 crate::RestartExecutorSnapshot {
147 healthy: self.healthy.load(Ordering::Acquire)
148 && workers_healthy
149 && !self.cancellation.load(Ordering::Acquire),
150 pending: self.pending.load(Ordering::Acquire),
151 queue_capacity: self.queue_capacity,
152 worker_count: self.worker_count,
153 }
154 }
155
156 pub fn shutdown(&self, timeout: Duration) -> bool {
157 self.cancellation.store(true, Ordering::Release);
158 if let Ok(mut sender) = self.sender.lock() {
159 sender.take();
160 } else {
161 self.healthy.store(false, Ordering::Release);
162 }
163 let deadline = Instant::now().checked_add(timeout);
164 while deadline.is_none_or(|deadline| Instant::now() < deadline) {
165 let complete = self
166 .workers
167 .lock()
168 .map(|workers| workers.iter().all(JoinHandle::is_finished))
169 .unwrap_or(false);
170 if complete {
171 break;
172 }
173 std::thread::sleep(Duration::from_millis(5));
174 }
175 let Ok(mut workers) = self.workers.lock() else {
176 self.healthy.store(false, Ordering::Release);
177 return false;
178 };
179 let all_finished = workers.iter().all(JoinHandle::is_finished);
180 for worker in workers.drain(..).filter(JoinHandle::is_finished) {
181 if worker.join().is_err() {
182 self.healthy.store(false, Ordering::Release);
183 }
184 }
185 drain_retained(&self.commands, &self.healthy);
186 drain_retained(&self.completions, &self.healthy);
187 self.pending.store(0, Ordering::Release);
188 self.healthy.store(false, Ordering::Release);
189 all_finished
190 }
191}
192
193fn spawn_workers(
194 count: usize,
195 receiver: Arc<Mutex<Receiver<RestartCommand>>>,
196 completions: SyncSender<RestartCompletion>,
197 cancellation: Arc<AtomicBool>,
198 healthy: Arc<AtomicBool>,
199 pending: Arc<AtomicU64>,
200) -> Vec<JoinHandle<()>> {
201 (0..count)
202 .filter_map(|index| {
203 let receiver = Arc::clone(&receiver);
204 let completions = completions.clone();
205 let cancellation = Arc::clone(&cancellation);
206 let healthy = Arc::clone(&healthy);
207 let worker_health = Arc::clone(&healthy);
208 let pending = Arc::clone(&pending);
209 std::thread::Builder::new()
210 .name(format!("appcore-restart-{index}"))
211 .stack_size(RESTART_THREAD_STACK_BYTES)
212 .spawn(move || {
213 restart_worker(receiver, completions, cancellation, worker_health, pending)
214 })
215 .map_err(|_| healthy.store(false, Ordering::Release))
216 .ok()
217 })
218 .collect()
219}
220
221fn restart_worker(
222 receiver: Arc<Mutex<Receiver<RestartCommand>>>,
223 completions: SyncSender<RestartCompletion>,
224 cancellation: Arc<AtomicBool>,
225 healthy: Arc<AtomicBool>,
226 pending: Arc<AtomicU64>,
227) {
228 loop {
229 if cancellation.load(Ordering::Acquire) {
230 return;
231 }
232 let command = match receiver
233 .lock()
234 .map(|receiver| receiver.recv_timeout(Duration::from_millis(25)))
235 {
236 Ok(Ok(command)) => command,
237 Ok(Err(RecvTimeoutError::Timeout)) => continue,
238 Ok(Err(RecvTimeoutError::Disconnected)) | Err(_) => return,
239 };
240 let service_id = command.service.descriptor().name().to_string();
241 let attempt = command.attempt;
242 let completion = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
243 execute_restart(command, &cancellation)
244 }))
245 .unwrap_or(RestartCompletion {
246 service_id,
247 attempt,
248 outcome: RestartOutcome::Failed,
249 });
250 send_completion(&completions, completion, &cancellation, &healthy);
251 decrement_pending(&pending);
252 }
253}
254
255fn send_completion(
256 completions: &SyncSender<RestartCompletion>,
257 mut completion: RestartCompletion,
258 cancellation: &AtomicBool,
259 healthy: &AtomicBool,
260) {
261 loop {
262 if cancellation.load(Ordering::Acquire) {
263 return;
264 }
265 match completions.try_send(completion) {
266 Ok(()) => return,
267 Err(TrySendError::Full(retained)) => {
268 completion = retained;
269 std::thread::sleep(Duration::from_millis(1));
270 }
271 Err(TrySendError::Disconnected(_)) => {
272 healthy.store(false, Ordering::Release);
273 return;
274 }
275 }
276 }
277}
278
279fn drain_retained<T>(receiver: &Mutex<Receiver<T>>, healthy: &AtomicBool) {
280 match receiver.lock() {
281 Ok(receiver) => receiver.try_iter().for_each(drop),
282 Err(_) => healthy.store(false, Ordering::Release),
283 }
284}
285
286fn decrement_pending(pending: &AtomicU64) {
287 let _ = pending.fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| {
288 Some(value.saturating_sub(1))
289 });
290}
291
292fn execute_restart(command: RestartCommand, cancellation: &AtomicBool) -> RestartCompletion {
293 let service_id = command.service.descriptor().name().to_string();
294 let timeout = command
295 .service
296 .descriptor()
297 .restart_policy()
298 .shutdown_timeout;
299 let outcome = match command.service.stop(timeout) {
300 Err(_) if command.service.runtime_state() == ServiceRuntimeState::Orphaned => {
301 RestartOutcome::Orphaned
302 }
303 Err(_) => RestartOutcome::Failed,
304 Ok(()) if cancellation.load(Ordering::Acquire) => RestartOutcome::Cancelled,
305 Ok(()) => match command.service.start() {
306 Ok(()) => RestartOutcome::Restarted,
307 Err(_) => RestartOutcome::Failed,
308 },
309 };
310 RestartCompletion {
311 service_id,
312 attempt: command.attempt,
313 outcome,
314 }
315}
316
317#[cfg(test)]
318#[path = "restart_executor_tests.rs"]
319mod tests;