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;
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
25pub enum RestartState {
26 None,
28 Scheduled {
30 execute_at_ms: u64,
32 },
33 Stopping,
35 Starting,
37 Backoff,
39 Failed,
41}
42
43pub(crate) struct RestartCommand {
44 pub service: Arc<dyn ManagedService>,
45 pub attempt: u64,
46}
47
48pub(crate) enum RestartOutcome {
49 Restarted,
50 Orphaned,
51 Failed,
52 Cancelled,
53}
54
55pub(crate) struct RestartCompletion {
56 pub service_id: String,
57 pub attempt: u64,
58 pub outcome: RestartOutcome,
59}
60
61pub(crate) struct RestartExecutor {
62 sender: SyncSender<RestartCommand>,
63 completions: Mutex<Receiver<RestartCompletion>>,
64 workers: Mutex<Vec<JoinHandle<()>>>,
65 cancellation: Arc<AtomicBool>,
66 healthy: Arc<AtomicBool>,
67 pending: Arc<AtomicU64>,
68 queue_capacity: usize,
69 worker_count: usize,
70}
71
72impl RestartExecutor {
73 pub fn new(queue_capacity: usize, worker_count: usize) -> Self {
74 let queue_capacity = queue_capacity.max(1);
75 let worker_count = worker_count.max(1);
76 let (sender, receiver) = mpsc::sync_channel::<RestartCommand>(queue_capacity);
77 let (completion_sender, completions) = mpsc::channel();
78 let receiver = Arc::new(Mutex::new(receiver));
79 let cancellation = Arc::new(AtomicBool::new(false));
80 let healthy = Arc::new(AtomicBool::new(true));
81 let pending = Arc::new(AtomicU64::new(0));
82 let workers = spawn_workers(
83 worker_count,
84 receiver,
85 completion_sender,
86 Arc::clone(&cancellation),
87 Arc::clone(&healthy),
88 Arc::clone(&pending),
89 );
90 Self {
91 sender,
92 completions: Mutex::new(completions),
93 workers: Mutex::new(workers),
94 cancellation,
95 healthy,
96 pending,
97 queue_capacity,
98 worker_count,
99 }
100 }
101
102 pub fn schedule(&self, command: RestartCommand) -> SupervisorResult<()> {
103 self.pending.fetch_add(1, Ordering::AcqRel);
104 match self.sender.try_send(command) {
105 Ok(()) => Ok(()),
106 Err(TrySendError::Full(_)) => {
107 self.pending.fetch_sub(1, Ordering::AcqRel);
108 Err(SupervisorError::RestartQueueFull)
109 }
110 Err(TrySendError::Disconnected(_)) => {
111 self.pending.fetch_sub(1, Ordering::AcqRel);
112 Err(SupervisorError::RestartExecutorStopped)
113 }
114 }
115 }
116
117 pub fn drain_completions(&self) -> Vec<RestartCompletion> {
118 let Ok(receiver) = self.completions.lock() else {
119 self.healthy.store(false, Ordering::Release);
120 return Vec::new();
121 };
122 receiver.try_iter().collect()
123 }
124
125 pub fn snapshot(&self) -> crate::RestartExecutorSnapshot {
126 let workers_healthy = self
127 .workers
128 .lock()
129 .map(|workers| workers.iter().all(|worker| !worker.is_finished()))
130 .unwrap_or(false);
131 crate::RestartExecutorSnapshot {
132 healthy: self.healthy.load(Ordering::Acquire)
133 && workers_healthy
134 && !self.cancellation.load(Ordering::Acquire),
135 pending: self.pending.load(Ordering::Acquire),
136 queue_capacity: self.queue_capacity,
137 worker_count: self.worker_count,
138 }
139 }
140
141 pub fn shutdown(&self, timeout: Duration) -> bool {
142 self.cancellation.store(true, Ordering::Release);
143 let deadline = Instant::now().checked_add(timeout);
144 while deadline.is_none_or(|deadline| Instant::now() < deadline) {
145 let complete = self
146 .workers
147 .lock()
148 .map(|workers| workers.iter().all(JoinHandle::is_finished))
149 .unwrap_or(false);
150 if complete {
151 break;
152 }
153 std::thread::sleep(Duration::from_millis(5));
154 }
155 let Ok(mut workers) = self.workers.lock() else {
156 self.healthy.store(false, Ordering::Release);
157 return false;
158 };
159 let all_finished = workers.iter().all(JoinHandle::is_finished);
160 for worker in workers.drain(..).filter(JoinHandle::is_finished) {
161 if worker.join().is_err() {
162 self.healthy.store(false, Ordering::Release);
163 }
164 }
165 self.pending.store(0, Ordering::Release);
166 self.healthy.store(false, Ordering::Release);
167 all_finished
168 }
169}
170
171fn spawn_workers(
172 count: usize,
173 receiver: Arc<Mutex<Receiver<RestartCommand>>>,
174 completions: mpsc::Sender<RestartCompletion>,
175 cancellation: Arc<AtomicBool>,
176 healthy: Arc<AtomicBool>,
177 pending: Arc<AtomicU64>,
178) -> Vec<JoinHandle<()>> {
179 (0..count)
180 .filter_map(|index| {
181 let receiver = Arc::clone(&receiver);
182 let completions = completions.clone();
183 let cancellation = Arc::clone(&cancellation);
184 let healthy = Arc::clone(&healthy);
185 let pending = Arc::clone(&pending);
186 std::thread::Builder::new()
187 .name(format!("appcore-restart-{index}"))
188 .spawn(move || restart_worker(receiver, completions, cancellation, pending))
189 .map_err(|_| healthy.store(false, Ordering::Release))
190 .ok()
191 })
192 .collect()
193}
194
195fn restart_worker(
196 receiver: Arc<Mutex<Receiver<RestartCommand>>>,
197 completions: mpsc::Sender<RestartCompletion>,
198 cancellation: Arc<AtomicBool>,
199 pending: Arc<AtomicU64>,
200) {
201 loop {
202 if cancellation.load(Ordering::Acquire) {
203 return;
204 }
205 let command = match receiver
206 .lock()
207 .map(|receiver| receiver.recv_timeout(Duration::from_millis(25)))
208 {
209 Ok(Ok(command)) => command,
210 Ok(Err(RecvTimeoutError::Timeout)) => continue,
211 Ok(Err(RecvTimeoutError::Disconnected)) | Err(_) => return,
212 };
213 let service_id = command.service.descriptor().name().to_string();
214 let attempt = command.attempt;
215 let completion = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
216 execute_restart(command, &cancellation)
217 }))
218 .unwrap_or(RestartCompletion {
219 service_id,
220 attempt,
221 outcome: RestartOutcome::Failed,
222 });
223 decrement_pending(&pending);
224 let _ = completions.send(completion);
225 }
226}
227
228fn decrement_pending(pending: &AtomicU64) {
229 let _ = pending.fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| {
230 Some(value.saturating_sub(1))
231 });
232}
233
234fn execute_restart(command: RestartCommand, cancellation: &AtomicBool) -> RestartCompletion {
235 let service_id = command.service.descriptor().name().to_string();
236 let timeout = command
237 .service
238 .descriptor()
239 .restart_policy()
240 .shutdown_timeout;
241 let outcome = match command.service.stop(timeout) {
242 Err(_) if command.service.runtime_state() == ServiceRuntimeState::Orphaned => {
243 RestartOutcome::Orphaned
244 }
245 Err(_) => RestartOutcome::Failed,
246 Ok(()) if cancellation.load(Ordering::Acquire) => RestartOutcome::Cancelled,
247 Ok(()) => match command.service.start() {
248 Ok(()) => RestartOutcome::Restarted,
249 Err(_) => RestartOutcome::Failed,
250 },
251 };
252 RestartCompletion {
253 service_id,
254 attempt: command.attempt,
255 outcome,
256 }
257}
258
259#[cfg(test)]
260mod tests {
261 use super::*;
262 use crate::{CallbackManagedService, ManagedResource, RestartPolicy, ServiceDescriptor};
263
264 #[test]
265 fn saturated_queue_does_not_block_executor_shutdown() {
266 let stop_started = Arc::new(AtomicBool::new(false));
267 let stop_signal = Arc::clone(&stop_started);
268 let descriptor =
269 ServiceDescriptor::new("worker", ManagedResource::Worker, RestartPolicy::never())
270 .unwrap();
271 let service: Arc<dyn ManagedService> = Arc::new(CallbackManagedService::new(
272 descriptor,
273 || Ok(()),
274 move |_| {
275 stop_signal.store(true, Ordering::Release);
276 std::thread::sleep(Duration::from_millis(100));
277 Ok(())
278 },
279 || crate::ServiceHealth::Healthy,
280 ));
281 service.start().unwrap();
282 let executor = RestartExecutor::new(1, 1);
283 executor
284 .schedule(RestartCommand {
285 service: Arc::clone(&service),
286 attempt: 1,
287 })
288 .unwrap();
289 let deadline = Instant::now() + Duration::from_secs(1);
290 while !stop_started.load(Ordering::Acquire) && Instant::now() < deadline {
291 std::thread::sleep(Duration::from_millis(1));
292 }
293 assert!(stop_started.load(Ordering::Acquire));
294 executor
295 .schedule(RestartCommand {
296 service: Arc::clone(&service),
297 attempt: 2,
298 })
299 .unwrap();
300 assert!(matches!(
301 executor.schedule(RestartCommand {
302 service,
303 attempt: 3,
304 }),
305 Err(SupervisorError::RestartQueueFull)
306 ));
307
308 assert!(executor.shutdown(Duration::from_secs(1)));
309 assert_eq!(executor.snapshot().pending, 0);
310 }
311
312 #[test]
313 fn pending_counter_saturates_after_late_worker_completion() {
314 let pending = AtomicU64::new(0);
315 decrement_pending(&pending);
316 assert_eq!(pending.load(Ordering::Acquire), 0);
317 }
318
319 #[test]
320 fn managed_service_panic_does_not_kill_restart_worker() {
321 struct PanicService {
322 descriptor: ServiceDescriptor,
323 }
324
325 impl ManagedService for PanicService {
326 fn descriptor(&self) -> &ServiceDescriptor {
327 &self.descriptor
328 }
329
330 fn start(&self) -> SupervisorResult<()> {
331 Ok(())
332 }
333
334 fn stop(&self, _timeout: Duration) -> SupervisorResult<()> {
335 panic!("injected managed-service panic");
336 }
337
338 fn health(&self) -> crate::ServiceHealth {
339 crate::ServiceHealth::Failed
340 }
341 }
342
343 let service: Arc<dyn ManagedService> = Arc::new(PanicService {
344 descriptor: ServiceDescriptor::new(
345 "panic-worker",
346 ManagedResource::Worker,
347 RestartPolicy::never(),
348 )
349 .unwrap(),
350 });
351 let executor = RestartExecutor::new(1, 1);
352 executor
353 .schedule(RestartCommand {
354 service,
355 attempt: 1,
356 })
357 .unwrap();
358 let deadline = Instant::now() + Duration::from_secs(1);
359 let completion = loop {
360 if let Some(completion) = executor.drain_completions().pop() {
361 break completion;
362 }
363 assert!(Instant::now() < deadline);
364 std::thread::sleep(Duration::from_millis(1));
365 };
366 assert!(matches!(completion.outcome, RestartOutcome::Failed));
367 assert_eq!(executor.snapshot().pending, 0);
368 assert!(executor.shutdown(Duration::from_secs(1)));
369 }
370}