1use std::borrow::Borrow;
11use std::ffi::OsStr;
12use std::fmt;
13use std::future::Future;
14use std::path::Path;
15use std::pin::Pin;
16use std::sync::mpsc;
17use std::task::{Context, Poll};
18
19use ::tokio::sync::{oneshot, watch};
20
21use crate::{
22 AffectedApplications, CancellationHandle, Error, ErrorKind, Filter, FilterAction, FilterTarget,
23 OperationOutcome, ProcessIdentity, Progress, RecoveryOutcome, ResourceBatch, Result,
24 SessionKey, ShutdownOptions,
25};
26
27enum WorkerState {
28 Primary(crate::RestartSession),
29 Joined(crate::JoinedSession),
30 Pending(crate::RestartPending),
31 Completion(crate::RecoveryCompletion),
32 Empty,
33}
34
35enum Command {
36 PrimaryRegister {
37 resources: ResourceBatch,
38 reply: oneshot::Sender<Result<()>>,
39 },
40 PrimaryAffected {
41 reply: oneshot::Sender<Result<AffectedApplications>>,
42 },
43 SetFilter {
44 target: FilterTarget,
45 action: FilterAction,
46 reply: oneshot::Sender<Result<()>>,
47 },
48 RemoveFilter {
49 target: FilterTarget,
50 reply: oneshot::Sender<Result<()>>,
51 },
52 Filters {
53 reply: oneshot::Sender<Result<Vec<Filter>>>,
54 },
55 Shutdown {
56 options: ShutdownOptions,
57 reply: oneshot::Sender<ShutdownReply>,
58 },
59 ShutdownWithProgress {
60 options: ShutdownOptions,
61 progress: watch::Sender<Option<Progress>>,
62 reply: oneshot::Sender<ShutdownReply>,
63 },
64 EndPrimary {
65 reply: oneshot::Sender<Result<()>>,
66 },
67 Restart {
68 reply: oneshot::Sender<Result<RecoveryOutcome>>,
69 },
70 RestartWithProgress {
71 progress: watch::Sender<Option<Progress>>,
72 reply: oneshot::Sender<RestartProgressReply>,
73 },
74 LeaveStopped {
75 reply: oneshot::Sender<Result<RecoveryOutcome>>,
76 },
77 CompletionAffected {
78 reply: oneshot::Sender<Result<AffectedApplications>>,
79 },
80 EndCompletion {
81 reply: oneshot::Sender<Result<RecoveryOutcome>>,
82 },
83 JoinedRegister {
84 resources: ResourceBatch,
85 reply: oneshot::Sender<Result<()>>,
86 },
87 EndJoined {
88 reply: oneshot::Sender<Result<()>>,
89 },
90}
91
92#[cfg(all(test, windows))]
93#[derive(Clone)]
94struct WorkerExitNotification {
95 state: std::sync::Arc<(std::sync::Mutex<bool>, std::sync::Condvar)>,
96}
97
98#[cfg(all(test, windows))]
99impl WorkerExitNotification {
100 fn new() -> Self {
101 Self {
102 state: std::sync::Arc::new((std::sync::Mutex::new(false), std::sync::Condvar::new())),
103 }
104 }
105
106 fn notify(&self) {
107 let (stopped, changed) = &*self.state;
108 *stopped.lock().unwrap_or_else(|error| error.into_inner()) = true;
109 changed.notify_all();
110 }
111
112 fn wait(&self) {
113 let (stopped, changed) = &*self.state;
114 let mut stopped = stopped.lock().unwrap_or_else(|error| error.into_inner());
115 while !*stopped {
116 stopped = changed
117 .wait(stopped)
118 .unwrap_or_else(|error| error.into_inner());
119 }
120 }
121}
122
123#[cfg(all(test, windows))]
124struct WorkerExitGuard(WorkerExitNotification);
125
126#[cfg(all(test, windows))]
127impl Drop for WorkerExitGuard {
128 fn drop(&mut self) {
129 self.0.notify();
130 }
131}
132
133struct Worker {
134 sender: mpsc::Sender<Command>,
135 #[cfg(all(test, windows))]
136 exit: WorkerExitNotification,
137}
138
139impl Worker {
140 fn send(&self, command: Command) -> Result<()> {
141 self.sender.send(command).map_err(|_| worker_unavailable())
142 }
143
144 #[cfg(all(test, windows))]
145 fn exit_notification(&self) -> WorkerExitNotification {
146 self.exit.clone()
147 }
148}
149
150struct PrimaryInit {
151 key: SessionKey,
152 cancellation: CancellationHandle,
153}
154
155async fn start_primary_worker() -> Result<(Worker, PrimaryInit)> {
156 let (command_sender, command_receiver) = mpsc::channel::<Command>();
157 let (init_sender, init_receiver) = oneshot::channel();
158 #[cfg(all(test, windows))]
159 let exit = WorkerExitNotification::new();
160 #[cfg(all(test, windows))]
161 let thread_exit = exit.clone();
162 std::thread::Builder::new()
163 .name("restart-manager".to_owned())
164 .spawn(move || {
165 #[cfg(all(test, windows))]
166 let _exit_guard = WorkerExitGuard(thread_exit);
167 let session = match crate::RestartSession::new() {
168 Ok(session) => session,
169 Err(error) => {
170 let _ = init_sender.send(Err(error));
171 return;
172 }
173 };
174 let init = PrimaryInit {
175 key: session.session_key().clone(),
176 cancellation: session.cancellation_handle(),
177 };
178 if init_sender.send(Ok(init)).is_err() {
179 return;
180 }
181 run_worker(command_receiver, WorkerState::Primary(session));
182 })
183 .map_err(|error| {
184 Error::new(
185 ErrorKind::AsyncWorkerUnavailable,
186 error.raw_os_error().map(|code| code as u32),
187 format!("the Restart Manager worker thread could not be created: {error}"),
188 )
189 })?;
190 let init = init_receiver.await.map_err(|_| worker_unavailable())??;
191 Ok((
192 Worker {
193 sender: command_sender,
194 #[cfg(all(test, windows))]
195 exit,
196 },
197 init,
198 ))
199}
200
201async fn start_joined_worker(key: SessionKey) -> Result<Worker> {
202 let (command_sender, command_receiver) = mpsc::channel::<Command>();
203 let (init_sender, init_receiver) = oneshot::channel();
204 #[cfg(all(test, windows))]
205 let exit = WorkerExitNotification::new();
206 #[cfg(all(test, windows))]
207 let thread_exit = exit.clone();
208 std::thread::Builder::new()
209 .name("restart-manager-joined".to_owned())
210 .spawn(move || {
211 #[cfg(all(test, windows))]
212 let _exit_guard = WorkerExitGuard(thread_exit);
213 let session = match crate::JoinedSession::join(&key) {
214 Ok(session) => session,
215 Err(error) => {
216 let _ = init_sender.send(Err(error));
217 return;
218 }
219 };
220 if init_sender.send(Ok(())).is_err() {
221 return;
222 }
223 run_worker(command_receiver, WorkerState::Joined(session));
224 })
225 .map_err(|error| {
226 Error::new(
227 ErrorKind::AsyncWorkerUnavailable,
228 error.raw_os_error().map(|code| code as u32),
229 format!("the Restart Manager worker thread could not be created: {error}"),
230 )
231 })?;
232 init_receiver.await.map_err(|_| worker_unavailable())??;
233 Ok(Worker {
234 sender: command_sender,
235 #[cfg(all(test, windows))]
236 exit,
237 })
238}
239
240fn run_worker(receiver: mpsc::Receiver<Command>, mut state: WorkerState) {
241 while let Ok(command) = receiver.recv() {
242 execute_command(command, &mut state);
243 }
244}
245
246fn execute_command(command: Command, state: &mut WorkerState) {
247 match command {
248 Command::PrimaryRegister { resources, reply } => {
249 let result = match state {
250 WorkerState::Primary(session) => session.register_resources(&resources),
251 _ => Err(wrong_state()),
252 };
253 let _ = reply.send(result);
254 }
255 Command::PrimaryAffected { reply } => {
256 let result = match state {
257 WorkerState::Primary(session) => session.affected_applications(),
258 _ => Err(wrong_state()),
259 };
260 let _ = reply.send(result);
261 }
262 Command::SetFilter {
263 target,
264 action,
265 reply,
266 } => {
267 let result = match state {
268 WorkerState::Primary(session) => session.set_filter(&target, action),
269 _ => Err(wrong_state()),
270 };
271 let _ = reply.send(result);
272 }
273 Command::RemoveFilter { target, reply } => {
274 let result = match state {
275 WorkerState::Primary(session) => session.remove_filter(&target),
276 _ => Err(wrong_state()),
277 };
278 let _ = reply.send(result);
279 }
280 Command::Filters { reply } => {
281 let result = match state {
282 WorkerState::Primary(session) => session.filters(),
283 _ => Err(wrong_state()),
284 };
285 let _ = reply.send(result);
286 }
287 Command::Shutdown { options, reply } => {
288 let previous = std::mem::replace(state, WorkerState::Empty);
289 let result = match previous {
290 WorkerState::Primary(session) => {
291 let pending = session.shutdown_with_options(options);
292 let outcome = pending.shutdown_outcome().clone();
293 *state = WorkerState::Pending(pending);
294 ShutdownReply::Pending(outcome)
295 }
296 other => {
297 *state = other;
298 ShutdownReply::Failed(wrong_state())
299 }
300 };
301 let _ = reply.send(result);
302 }
303 Command::ShutdownWithProgress {
304 options,
305 progress,
306 reply,
307 } => {
308 let previous = std::mem::replace(state, WorkerState::Empty);
309 let result = match previous {
310 WorkerState::Primary(session) => {
311 match session.shutdown_with_progress(options, |value| {
312 progress.send_replace(Some(value));
313 }) {
314 Ok(pending) => {
315 let outcome = pending.shutdown_outcome().clone();
316 *state = WorkerState::Pending(pending);
317 ShutdownReply::Pending(outcome)
318 }
319 Err(not_started) => {
320 let (session, error) = not_started.into_parts();
321 *state = WorkerState::Primary(session);
322 ShutdownReply::Recoverable(error)
323 }
324 }
325 }
326 other => {
327 *state = other;
328 ShutdownReply::Failed(wrong_state())
329 }
330 };
331 let _ = reply.send(result);
332 }
333 Command::EndPrimary { reply } => {
334 let previous = std::mem::replace(state, WorkerState::Empty);
335 let result = match previous {
336 WorkerState::Primary(session) => session.end(),
337 other => {
338 *state = other;
339 Err(wrong_state())
340 }
341 };
342 let _ = reply.send(result);
343 }
344 Command::Restart { reply } => {
345 let previous = std::mem::replace(state, WorkerState::Empty);
346 let result = match previous {
347 WorkerState::Pending(pending) => {
348 let completion = pending.restart();
349 let outcome = completion.outcome().clone();
350 *state = WorkerState::Completion(completion);
351 Ok(outcome)
352 }
353 other => {
354 *state = other;
355 Err(wrong_state())
356 }
357 };
358 let _ = reply.send(result);
359 }
360 Command::RestartWithProgress { progress, reply } => {
361 let previous = std::mem::replace(state, WorkerState::Empty);
362 let result = match previous {
363 WorkerState::Pending(pending) => {
364 match pending.restart_with_progress(|value| {
365 progress.send_replace(Some(value));
366 }) {
367 Ok(completion) => {
368 let outcome = completion.outcome().clone();
369 *state = WorkerState::Completion(completion);
370 RestartProgressReply::Completed(outcome)
371 }
372 Err(not_started) => {
373 let (pending, error) = not_started.into_parts();
374 *state = WorkerState::Pending(pending);
375 RestartProgressReply::Recoverable(error)
376 }
377 }
378 }
379 other => {
380 *state = other;
381 RestartProgressReply::Failed(wrong_state())
382 }
383 };
384 let _ = reply.send(result);
385 }
386 Command::LeaveStopped { reply } => {
387 let previous = std::mem::replace(state, WorkerState::Empty);
388 let result = match previous {
389 WorkerState::Pending(pending) => {
390 let completion = pending.leave_stopped();
391 let outcome = completion.outcome().clone();
392 *state = WorkerState::Completion(completion);
393 Ok(outcome)
394 }
395 other => {
396 *state = other;
397 Err(wrong_state())
398 }
399 };
400 let _ = reply.send(result);
401 }
402 Command::CompletionAffected { reply } => {
403 let result = match state {
404 WorkerState::Completion(completion) => completion.affected_applications(),
405 _ => Err(wrong_state()),
406 };
407 let _ = reply.send(result);
408 }
409 Command::EndCompletion { reply } => {
410 let previous = std::mem::replace(state, WorkerState::Empty);
411 let result = match previous {
412 WorkerState::Completion(completion) => completion.end(),
413 other => {
414 *state = other;
415 Err(wrong_state())
416 }
417 };
418 let _ = reply.send(result);
419 }
420 Command::JoinedRegister { resources, reply } => {
421 let result = match state {
422 WorkerState::Joined(session) => session.register_resources(&resources),
423 _ => Err(wrong_state()),
424 };
425 let _ = reply.send(result);
426 }
427 Command::EndJoined { reply } => {
428 let previous = std::mem::replace(state, WorkerState::Empty);
429 let result = match previous {
430 WorkerState::Joined(session) => session.end(),
431 other => {
432 *state = other;
433 Err(wrong_state())
434 }
435 };
436 let _ = reply.send(result);
437 }
438 }
439}
440
441async fn receive<T>(receiver: oneshot::Receiver<T>) -> Result<T> {
442 receiver.await.map_err(|_| worker_unavailable())
443}
444
445fn worker_unavailable() -> Error {
446 Error::new(
447 ErrorKind::AsyncWorkerUnavailable,
448 None,
449 "the dedicated Restart Manager worker is unavailable",
450 )
451}
452
453fn wrong_state() -> Error {
454 Error::new(
455 ErrorKind::OperationOutOfSequence,
456 None,
457 "the asynchronous worker received an operation for the wrong typestate",
458 )
459}
460
461pub struct AsyncOperationError<T> {
467 state: Option<T>,
468 error: Error,
469}
470
471impl<T> AsyncOperationError<T> {
472 fn recoverable(state: T, error: Error) -> Self {
473 Self {
474 state: Some(state),
475 error,
476 }
477 }
478
479 fn unavailable(error: Error) -> Self {
480 Self { state: None, error }
481 }
482
483 #[must_use]
485 pub const fn error(&self) -> &Error {
486 &self.error
487 }
488
489 #[must_use]
491 pub const fn state(&self) -> Option<&T> {
492 self.state.as_ref()
493 }
494
495 #[must_use]
497 pub fn into_parts(self) -> (Option<T>, Error) {
498 (self.state, self.error)
499 }
500}
501
502impl<T> fmt::Debug for AsyncOperationError<T> {
503 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
504 formatter
505 .debug_struct("AsyncOperationError")
506 .field(
507 "state",
508 &self.state.as_ref().map(|_| std::any::type_name::<T>()),
509 )
510 .field("error", &self.error)
511 .finish()
512 }
513}
514
515impl<T> fmt::Display for AsyncOperationError<T> {
516 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
517 self.error.fmt(formatter)
518 }
519}
520
521impl<T> std::error::Error for AsyncOperationError<T> {
522 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
523 Some(&self.error)
524 }
525}
526
527#[derive(Clone)]
529pub struct ProgressReceiver {
530 receiver: watch::Receiver<Option<Progress>>,
531}
532
533impl ProgressReceiver {
534 #[must_use]
536 pub fn peek(&self) -> Option<Progress> {
537 *self.receiver.borrow()
538 }
539
540 pub async fn recv(&mut self) -> Option<Progress> {
542 self.receiver.changed().await.ok()?;
543 self.peek()
544 }
545}
546
547pub struct RestartSession {
549 worker: Worker,
550 key: SessionKey,
551 cancellation: CancellationHandle,
552}
553
554impl RestartSession {
555 pub async fn new() -> Result<Self> {
557 let (worker, init) = start_primary_worker().await?;
558 Ok(Self {
559 worker,
560 key: init.key,
561 cancellation: init.cancellation,
562 })
563 }
564
565 #[must_use]
567 pub fn session_key(&self) -> &SessionKey {
568 &self.key
569 }
570
571 #[must_use]
573 pub fn cancellation_handle(&self) -> CancellationHandle {
574 self.cancellation.clone()
575 }
576
577 pub async fn register_resources(&mut self, resources: &ResourceBatch) -> Result<()> {
579 let (reply, receiver) = oneshot::channel();
580 self.worker.send(Command::PrimaryRegister {
581 resources: resources.clone(),
582 reply,
583 })?;
584 receive(receiver).await?
585 }
586
587 pub async fn register_files<I, P>(&mut self, files: I) -> Result<()>
589 where
590 I: IntoIterator<Item = P>,
591 P: AsRef<Path>,
592 {
593 let mut resources = ResourceBatch::new();
594 for file in files {
595 resources.add_file(file.as_ref().to_path_buf());
596 }
597 self.register_resources(&resources).await
598 }
599
600 pub async fn register_processes<I, P>(&mut self, processes: I) -> Result<()>
602 where
603 I: IntoIterator<Item = P>,
604 P: Borrow<ProcessIdentity>,
605 {
606 let mut resources = ResourceBatch::new();
607 for process in processes {
608 resources.add_process(*process.borrow());
609 }
610 self.register_resources(&resources).await
611 }
612
613 pub async fn register_services<I, S>(&mut self, services: I) -> Result<()>
615 where
616 I: IntoIterator<Item = S>,
617 S: AsRef<OsStr>,
618 {
619 let mut resources = ResourceBatch::new();
620 for service in services {
621 resources.add_service(service.as_ref().to_os_string());
622 }
623 self.register_resources(&resources).await
624 }
625
626 pub async fn affected_applications(&mut self) -> Result<AffectedApplications> {
628 let (reply, receiver) = oneshot::channel();
629 self.worker.send(Command::PrimaryAffected { reply })?;
630 receive(receiver).await?
631 }
632
633 pub async fn set_filter(&mut self, target: &FilterTarget, action: FilterAction) -> Result<()> {
635 let (reply, receiver) = oneshot::channel();
636 self.worker.send(Command::SetFilter {
637 target: target.clone(),
638 action,
639 reply,
640 })?;
641 receive(receiver).await?
642 }
643
644 pub async fn remove_filter(&mut self, target: &FilterTarget) -> Result<()> {
646 let (reply, receiver) = oneshot::channel();
647 self.worker.send(Command::RemoveFilter {
648 target: target.clone(),
649 reply,
650 })?;
651 receive(receiver).await?
652 }
653
654 pub async fn filters(&mut self) -> Result<Vec<Filter>> {
656 let (reply, receiver) = oneshot::channel();
657 self.worker.send(Command::Filters { reply })?;
658 receive(receiver).await?
659 }
660
661 #[must_use]
663 pub fn shutdown(self) -> ShutdownFuture {
664 self.shutdown_with_options(ShutdownOptions::default())
665 }
666
667 #[must_use]
669 pub fn shutdown_with_options(self, options: ShutdownOptions) -> ShutdownFuture {
670 let Self {
671 worker,
672 key,
673 cancellation,
674 } = self;
675 let (reply, receiver) = oneshot::channel();
676 let _ = worker.send(Command::Shutdown { options, reply });
677 ShutdownFuture {
678 worker: Some(worker),
679 key: Some(key),
680 cancellation: Some(cancellation),
681 receiver,
682 cancel_on_drop: true,
683 }
684 }
685
686 #[must_use]
688 pub fn shutdown_with_progress(
689 self,
690 options: ShutdownOptions,
691 ) -> (ShutdownFuture, ProgressReceiver) {
692 let Self {
693 worker,
694 key,
695 cancellation,
696 } = self;
697 let (progress_sender, progress_receiver) = watch::channel(None);
698 let (reply, receiver) = oneshot::channel();
699 let _ = worker.send(Command::ShutdownWithProgress {
700 options,
701 progress: progress_sender,
702 reply,
703 });
704 (
705 ShutdownFuture {
706 worker: Some(worker),
707 key: Some(key),
708 cancellation: Some(cancellation),
709 receiver,
710 cancel_on_drop: true,
711 },
712 ProgressReceiver {
713 receiver: progress_receiver,
714 },
715 )
716 }
717
718 pub async fn end(self) -> Result<()> {
720 let Self { worker, .. } = self;
721 let (reply, receiver) = oneshot::channel();
722 worker.send(Command::EndPrimary { reply })?;
723 receive(receiver).await?
724 }
725}
726
727enum ShutdownReply {
728 Pending(OperationOutcome),
729 Recoverable(Error),
730 Failed(Error),
731}
732
733pub struct ShutdownFuture {
738 worker: Option<Worker>,
739 key: Option<SessionKey>,
740 cancellation: Option<CancellationHandle>,
741 receiver: oneshot::Receiver<ShutdownReply>,
742 cancel_on_drop: bool,
743}
744
745impl Future for ShutdownFuture {
746 type Output = std::result::Result<RestartPending, AsyncOperationError<RestartSession>>;
747
748 fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
749 let this = self.get_mut();
750 match Pin::new(&mut this.receiver).poll(context) {
751 Poll::Pending => Poll::Pending,
752 Poll::Ready(Ok(ShutdownReply::Pending(shutdown))) => {
753 this.cancel_on_drop = false;
754 Poll::Ready(Ok(RestartPending {
755 worker: this.worker.take().expect("future owns its worker"),
756 shutdown,
757 }))
758 }
759 Poll::Ready(Ok(ShutdownReply::Recoverable(error))) => {
760 this.cancel_on_drop = false;
761 let state = RestartSession {
762 worker: this.worker.take().expect("future owns its worker"),
763 key: this.key.take().expect("future owns its session key"),
764 cancellation: this
765 .cancellation
766 .take()
767 .expect("future owns its cancellation handle"),
768 };
769 Poll::Ready(Err(AsyncOperationError::recoverable(state, error)))
770 }
771 Poll::Ready(Ok(ShutdownReply::Failed(error))) => {
772 this.cancel_on_drop = false;
773 Poll::Ready(Err(AsyncOperationError::unavailable(error)))
774 }
775 Poll::Ready(Err(_)) => {
776 this.cancel_on_drop = false;
777 Poll::Ready(Err(AsyncOperationError::unavailable(worker_unavailable())))
778 }
779 }
780 }
781}
782
783impl Drop for ShutdownFuture {
784 fn drop(&mut self) {
785 if self.cancel_on_drop
786 && let Some(cancellation) = &self.cancellation
787 {
788 let _ = cancellation.cancel();
789 }
790 }
791}
792
793#[must_use = "call restart or leave_stopped explicitly"]
795pub struct RestartPending {
796 worker: Worker,
797 shutdown: OperationOutcome,
798}
799
800impl RestartPending {
801 #[must_use]
803 pub const fn shutdown_outcome(&self) -> &OperationOutcome {
804 &self.shutdown
805 }
806
807 #[must_use]
809 pub fn restart(self) -> RestartFuture {
810 let Self {
811 worker,
812 shutdown: _,
813 } = self;
814 let (reply, receiver) = oneshot::channel();
815 let _ = worker.send(Command::Restart { reply });
816 RestartFuture {
817 worker: Some(worker),
818 receiver,
819 }
820 }
821
822 #[must_use]
824 pub fn restart_with_progress(self) -> (RestartWithProgressFuture, ProgressReceiver) {
825 let Self { worker, shutdown } = self;
826 let (progress_sender, progress_receiver) = watch::channel(None);
827 let (reply, receiver) = oneshot::channel();
828 let _ = worker.send(Command::RestartWithProgress {
829 progress: progress_sender,
830 reply,
831 });
832 (
833 RestartWithProgressFuture {
834 worker: Some(worker),
835 shutdown: Some(shutdown),
836 receiver,
837 },
838 ProgressReceiver {
839 receiver: progress_receiver,
840 },
841 )
842 }
843
844 pub async fn leave_stopped(self) -> Result<RecoveryCompletion> {
846 let Self { worker, .. } = self;
847 let (reply, receiver) = oneshot::channel();
848 worker.send(Command::LeaveStopped { reply })?;
849 let outcome = receive(receiver).await??;
850 Ok(RecoveryCompletion { worker, outcome })
851 }
852}
853
854pub struct RestartFuture {
859 worker: Option<Worker>,
860 receiver: oneshot::Receiver<Result<RecoveryOutcome>>,
861}
862
863impl Future for RestartFuture {
864 type Output = Result<RecoveryCompletion>;
865
866 fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
867 let this = self.get_mut();
868 match Pin::new(&mut this.receiver).poll(context) {
869 Poll::Pending => Poll::Pending,
870 Poll::Ready(Ok(Ok(outcome))) => Poll::Ready(Ok(RecoveryCompletion {
871 worker: this.worker.take().expect("future owns its worker"),
872 outcome,
873 })),
874 Poll::Ready(Ok(Err(error))) => Poll::Ready(Err(error)),
875 Poll::Ready(Err(_)) => Poll::Ready(Err(worker_unavailable())),
876 }
877 }
878}
879
880enum RestartProgressReply {
881 Completed(RecoveryOutcome),
882 Recoverable(Error),
883 Failed(Error),
884}
885
886pub struct RestartWithProgressFuture {
888 worker: Option<Worker>,
889 shutdown: Option<OperationOutcome>,
890 receiver: oneshot::Receiver<RestartProgressReply>,
891}
892
893impl Future for RestartWithProgressFuture {
894 type Output = std::result::Result<RecoveryCompletion, AsyncOperationError<RestartPending>>;
895
896 fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
897 let this = self.get_mut();
898 match Pin::new(&mut this.receiver).poll(context) {
899 Poll::Pending => Poll::Pending,
900 Poll::Ready(Ok(RestartProgressReply::Completed(outcome))) => {
901 Poll::Ready(Ok(RecoveryCompletion {
902 worker: this.worker.take().expect("future owns its worker"),
903 outcome,
904 }))
905 }
906 Poll::Ready(Ok(RestartProgressReply::Recoverable(error))) => {
907 let state = RestartPending {
908 worker: this.worker.take().expect("future owns its worker"),
909 shutdown: this
910 .shutdown
911 .take()
912 .expect("future owns its shutdown outcome"),
913 };
914 Poll::Ready(Err(AsyncOperationError::recoverable(state, error)))
915 }
916 Poll::Ready(Ok(RestartProgressReply::Failed(error))) => {
917 Poll::Ready(Err(AsyncOperationError::unavailable(error)))
918 }
919 Poll::Ready(Err(_)) => {
920 Poll::Ready(Err(AsyncOperationError::unavailable(worker_unavailable())))
921 }
922 }
923 }
924}
925
926pub struct RecoveryCompletion {
928 worker: Worker,
929 outcome: RecoveryOutcome,
930}
931
932impl RecoveryCompletion {
933 #[must_use]
935 pub const fn outcome(&self) -> &RecoveryOutcome {
936 &self.outcome
937 }
938
939 pub async fn affected_applications(&mut self) -> Result<AffectedApplications> {
941 let (reply, receiver) = oneshot::channel();
942 self.worker.send(Command::CompletionAffected { reply })?;
943 receive(receiver).await?
944 }
945
946 pub async fn end(self) -> Result<RecoveryOutcome> {
948 let Self { worker, .. } = self;
949 let (reply, receiver) = oneshot::channel();
950 worker.send(Command::EndCompletion { reply })?;
951 receive(receiver).await?
952 }
953}
954
955pub struct JoinedSession {
957 worker: Worker,
958 key: SessionKey,
959}
960
961impl JoinedSession {
962 pub async fn join(key: &SessionKey) -> Result<Self> {
964 let key = key.clone();
965 let worker = start_joined_worker(key.clone()).await?;
966 Ok(Self { worker, key })
967 }
968
969 #[must_use]
971 pub fn session_key(&self) -> &SessionKey {
972 &self.key
973 }
974
975 pub async fn register_resources(&mut self, resources: &ResourceBatch) -> Result<()> {
977 let (reply, receiver) = oneshot::channel();
978 self.worker.send(Command::JoinedRegister {
979 resources: resources.clone(),
980 reply,
981 })?;
982 receive(receiver).await?
983 }
984
985 pub async fn register_files<I, P>(&mut self, files: I) -> Result<()>
987 where
988 I: IntoIterator<Item = P>,
989 P: AsRef<Path>,
990 {
991 let mut resources = ResourceBatch::new();
992 for file in files {
993 resources.add_file(file.as_ref().to_path_buf());
994 }
995 self.register_resources(&resources).await
996 }
997
998 pub async fn register_processes<I, P>(&mut self, processes: I) -> Result<()>
1000 where
1001 I: IntoIterator<Item = P>,
1002 P: Borrow<ProcessIdentity>,
1003 {
1004 let mut resources = ResourceBatch::new();
1005 for process in processes {
1006 resources.add_process(*process.borrow());
1007 }
1008 self.register_resources(&resources).await
1009 }
1010
1011 pub async fn register_services<I, S>(&mut self, services: I) -> Result<()>
1013 where
1014 I: IntoIterator<Item = S>,
1015 S: AsRef<OsStr>,
1016 {
1017 let mut resources = ResourceBatch::new();
1018 for service in services {
1019 resources.add_service(service.as_ref().to_os_string());
1020 }
1021 self.register_resources(&resources).await
1022 }
1023
1024 pub async fn end(self) -> Result<()> {
1026 let Self { worker, .. } = self;
1027 let (reply, receiver) = oneshot::channel();
1028 worker.send(Command::EndJoined { reply })?;
1029 receive(receiver).await?
1030 }
1031}
1032
1033#[cfg(all(test, windows))]
1034mod tests {
1035 use std::sync::Arc;
1036 use std::task::{Wake, Waker};
1037
1038 use super::*;
1039
1040 struct ThreadWake(std::thread::Thread);
1041
1042 impl Wake for ThreadWake {
1043 fn wake(self: Arc<Self>) {
1044 self.0.unpark();
1045 }
1046
1047 fn wake_by_ref(self: &Arc<Self>) {
1048 self.0.unpark();
1049 }
1050 }
1051
1052 fn block_on<F: Future>(future: F) -> F::Output {
1053 let waker = Waker::from(Arc::new(ThreadWake(std::thread::current())));
1054 let mut context = Context::from_waker(&waker);
1055 let mut future = std::pin::pin!(future);
1056 loop {
1057 match future.as_mut().poll(&mut context) {
1058 Poll::Ready(output) => return output,
1059 Poll::Pending => std::thread::park(),
1060 }
1061 }
1062 }
1063
1064 fn disconnected_worker() -> Worker {
1065 let (sender, receiver) = mpsc::channel();
1066 drop(receiver);
1067 Worker {
1068 sender,
1069 exit: WorkerExitNotification::new(),
1070 }
1071 }
1072
1073 #[test]
1074 fn worker_owns_the_complete_typestate_sequence() {
1075 let mut session = block_on(RestartSession::new()).unwrap();
1076 block_on(session.register_resources(&ResourceBatch::new())).unwrap();
1077 let pending = block_on(session.shutdown()).unwrap();
1078 assert!(pending.shutdown_outcome().is_success());
1079 let completion = block_on(pending.restart()).unwrap();
1080 assert!(completion.outcome().shutdown_outcome().is_success());
1081 assert!(
1082 completion
1083 .outcome()
1084 .restart_outcome()
1085 .is_some_and(OperationOutcome::is_success)
1086 );
1087 block_on(completion.end()).unwrap();
1088 }
1089
1090 #[test]
1091 fn progress_receiver_coalesces_and_peek_does_not_consume() {
1092 let (sender, receiver) = watch::channel(None);
1093 let mut receiver = ProgressReceiver { receiver };
1094 sender.send_replace(Progress::try_from_native(10));
1095 sender.send_replace(Progress::try_from_native(20));
1096 assert_eq!(receiver.peek().unwrap().percent_complete(), 20);
1097 assert_eq!(receiver.peek().unwrap().percent_complete(), 20);
1098 assert_eq!(block_on(receiver.recv()).unwrap().percent_complete(), 20);
1099 }
1100
1101 #[test]
1102 fn progress_receiver_recv_wakes_and_closes_normally() {
1103 let (sender, receiver) = watch::channel(None);
1104 let mut receiver = ProgressReceiver { receiver };
1105 let producer = std::thread::spawn(move || {
1106 sender.send_replace(Progress::try_from_native(30));
1107 });
1108 assert_eq!(block_on(receiver.recv()).unwrap().percent_complete(), 30);
1109 producer.join().unwrap();
1110 assert_eq!(block_on(receiver.recv()), None);
1111 }
1112
1113 #[test]
1114 fn worker_failure_remains_on_the_operation_future() {
1115 let session = crate::RestartSession::new().unwrap();
1116 let key = session.session_key().clone();
1117 let cancellation = session.cancellation_handle();
1118 session.end().unwrap();
1119
1120 let (reply, _receiver) = oneshot::channel();
1121 assert_eq!(
1122 disconnected_worker()
1123 .send(Command::Filters { reply })
1124 .unwrap_err()
1125 .kind(),
1126 ErrorKind::AsyncWorkerUnavailable
1127 );
1128
1129 let worker = disconnected_worker();
1130 let (reply_sender, receiver) = oneshot::channel();
1131 drop(reply_sender);
1132 let future = ShutdownFuture {
1133 worker: Some(worker),
1134 key: Some(key),
1135 cancellation: Some(cancellation),
1136 receiver,
1137 cancel_on_drop: true,
1138 };
1139 let error = match block_on(future) {
1140 Ok(_) => panic!("a disconnected shutdown reply unexpectedly succeeded"),
1141 Err(error) => error,
1142 };
1143 assert_eq!(error.error().kind(), ErrorKind::AsyncWorkerUnavailable);
1144 assert!(error.state().is_none());
1145 assert!(format!("{error:?}").contains("AsyncOperationError"));
1146 assert_eq!(error.to_string(), error.error().to_string());
1147 assert!(std::error::Error::source(&error).is_some());
1148
1149 let worker = disconnected_worker();
1150 let (reply_sender, receiver) = oneshot::channel();
1151 drop(reply_sender);
1152 let result = block_on(RestartFuture {
1153 worker: Some(worker),
1154 receiver,
1155 });
1156 let error = match result {
1157 Ok(_) => panic!("a disconnected restart reply unexpectedly succeeded"),
1158 Err(error) => error,
1159 };
1160 assert_eq!(error.kind(), ErrorKind::AsyncWorkerUnavailable);
1161
1162 let worker = disconnected_worker();
1163 let (reply_sender, receiver) = oneshot::channel();
1164 drop(reply_sender);
1165 let result = block_on(RestartWithProgressFuture {
1166 worker: Some(worker),
1167 shutdown: Some(OperationOutcome::Succeeded),
1168 receiver,
1169 });
1170 let error = match result {
1171 Ok(_) => panic!("a disconnected restart-progress reply unexpectedly succeeded"),
1172 Err(error) => error,
1173 };
1174 assert_eq!(error.error().kind(), ErrorKind::AsyncWorkerUnavailable);
1175 assert!(error.state().is_none());
1176 }
1177
1178 #[test]
1179 fn callback_conflicts_retain_only_provably_reusable_async_state() {
1180 let session = block_on(RestartSession::new()).unwrap();
1181 let error = crate::sys::with_callback_lease_for_test(|| {
1182 let (shutdown, _progress) = session.shutdown_with_progress(ShutdownOptions::default());
1183 match block_on(shutdown) {
1184 Err(error) => error,
1185 Ok(_) => panic!("shutdown unexpectedly acquired the callback lease"),
1186 }
1187 });
1188 assert_eq!(error.error().kind(), ErrorKind::CallbackInUse);
1189 assert!(error.state().is_some());
1190 let (session, _) = error.into_parts();
1191
1192 let pending = block_on(
1193 session
1194 .expect("callback conflict retains the session")
1195 .shutdown(),
1196 )
1197 .unwrap();
1198 let error = crate::sys::with_callback_lease_for_test(|| {
1199 let (restart, _progress) = pending.restart_with_progress();
1200 match block_on(restart) {
1201 Err(error) => error,
1202 Ok(_) => panic!("restart unexpectedly acquired the callback lease"),
1203 }
1204 });
1205 assert_eq!(error.error().kind(), ErrorKind::CallbackInUse);
1206 assert!(error.state().is_some());
1207 let (pending, _) = error.into_parts();
1208 let completion = block_on(
1209 pending
1210 .expect("callback conflict retains the pending state")
1211 .leave_stopped(),
1212 )
1213 .unwrap();
1214 block_on(completion.end()).unwrap();
1215 }
1216
1217 #[test]
1218 fn wrong_command_state_is_reported_without_changing_the_worker_state() {
1219 let primary = block_on(RestartSession::new()).unwrap();
1220 let (reply, receiver) = oneshot::channel();
1221 primary
1222 .worker
1223 .send(Command::JoinedRegister {
1224 resources: ResourceBatch::new(),
1225 reply,
1226 })
1227 .unwrap();
1228 let error = block_on(receive(receiver)).unwrap().unwrap_err();
1229 assert_eq!(error.kind(), ErrorKind::OperationOutOfSequence);
1230 block_on(primary.end()).unwrap();
1231 }
1232
1233 #[test]
1234 fn internal_shutdown_state_mismatch_does_not_expose_retryable_state() {
1235 let session = block_on(RestartSession::new()).unwrap();
1236 let key = session.session_key().clone();
1237 let cancellation = session.cancellation_handle();
1238 let pending = block_on(session.shutdown()).unwrap();
1239 let RestartPending {
1240 worker,
1241 shutdown: _,
1242 } = pending;
1243 let (reply, receiver) = oneshot::channel();
1244 worker
1245 .send(Command::Shutdown {
1246 options: ShutdownOptions::default(),
1247 reply,
1248 })
1249 .unwrap();
1250 let error = match block_on(ShutdownFuture {
1251 worker: Some(worker),
1252 key: Some(key),
1253 cancellation: Some(cancellation),
1254 receiver,
1255 cancel_on_drop: true,
1256 }) {
1257 Err(error) => error,
1258 Ok(_) => panic!("shutdown unexpectedly accepted the pending worker state"),
1259 };
1260 assert_eq!(error.error().kind(), ErrorKind::OperationOutOfSequence);
1261 assert!(error.state().is_none());
1262 }
1263
1264 #[test]
1265 fn primary_and_joined_workers_cover_every_registration_and_filter_command() {
1266 let mut primary = block_on(RestartSession::new()).unwrap();
1267 let key = primary.session_key().clone();
1268 let cancellation = primary.cancellation_handle();
1269 let process = ProcessIdentity::current().unwrap();
1270 let executable = std::env::current_exe().unwrap();
1271
1272 let mut joined = block_on(JoinedSession::join(&key)).unwrap();
1273 assert_eq!(joined.session_key(), &key);
1274 block_on(joined.register_resources(&ResourceBatch::new())).unwrap();
1275 block_on(joined.register_files([&executable])).unwrap();
1276 block_on(joined.register_processes([process])).unwrap();
1277 block_on(joined.register_services(["EventLog"])).unwrap();
1278 block_on(joined.end()).unwrap();
1279
1280 block_on(primary.register_files([&executable])).unwrap();
1281 block_on(primary.register_processes([process])).unwrap();
1282 block_on(primary.register_services(["EventLog"])).unwrap();
1283 assert!(
1284 !block_on(primary.affected_applications())
1285 .unwrap()
1286 .is_empty()
1287 );
1288
1289 let process_target = FilterTarget::process(process);
1290 let service_target = FilterTarget::service("EventLog").unwrap();
1291 block_on(primary.set_filter(&process_target, FilterAction::PreventRestart)).unwrap();
1292 block_on(primary.set_filter(&service_target, FilterAction::PreventShutdown)).unwrap();
1293 let filters = block_on(primary.filters()).unwrap();
1294 assert!(
1295 filters
1296 .iter()
1297 .any(|filter| filter.target() == &process_target)
1298 );
1299 assert!(
1300 filters
1301 .iter()
1302 .any(|filter| filter.target() == &service_target)
1303 );
1304 block_on(primary.remove_filter(&process_target)).unwrap();
1305 block_on(primary.remove_filter(&service_target)).unwrap();
1306 block_on(primary.end()).unwrap();
1307 assert_eq!(
1308 cancellation.cancel().unwrap_err().kind(),
1309 ErrorKind::SessionEnded
1310 );
1311 }
1312
1313 #[test]
1314 fn progress_futures_and_leave_stopped_preserve_outcomes() {
1315 let _test_guard = crate::sys::serialize_callback_test();
1316 let mut session = block_on(RestartSession::new()).unwrap();
1317 let process = ProcessIdentity::current().unwrap();
1318 block_on(session.register_processes([process])).unwrap();
1319 block_on(session.set_filter(
1320 &FilterTarget::process(process),
1321 FilterAction::PreventShutdown,
1322 ))
1323 .unwrap();
1324 let (shutdown, progress) = session.shutdown_with_progress(ShutdownOptions::default());
1325 let pending = block_on(shutdown).unwrap();
1326 assert!(progress.peek().is_some());
1327 let (restart, progress) = pending.restart_with_progress();
1328 let mut completion = block_on(restart).unwrap();
1329 let _ = progress.peek();
1330 block_on(completion.affected_applications()).unwrap();
1331 block_on(completion.end()).unwrap();
1332
1333 let session = block_on(RestartSession::new()).unwrap();
1334 let pending = block_on(session.shutdown()).unwrap();
1335 let completion = block_on(pending.leave_stopped()).unwrap();
1336 assert!(completion.outcome().restart_outcome().is_none());
1337 block_on(completion.end()).unwrap();
1338 }
1339
1340 #[test]
1341 fn dropping_restart_futures_detaches_worker_cleanup() {
1342 let _test_guard = crate::sys::serialize_callback_test();
1343 let session = block_on(RestartSession::new()).unwrap();
1344 let pending = block_on(session.shutdown()).unwrap();
1345 let exit = pending.worker.exit_notification();
1346 drop(pending.restart());
1347 exit.wait();
1348
1349 let session = block_on(RestartSession::new()).unwrap();
1350 let pending = block_on(session.shutdown()).unwrap();
1351 let exit = pending.worker.exit_notification();
1352 let (restart, _progress) = pending.restart_with_progress();
1353 drop(restart);
1354 exit.wait();
1355
1356 crate::RestartSession::new().unwrap().end().unwrap();
1357 }
1358
1359 #[test]
1360 fn dropping_shutdown_future_leaves_cleanup_with_the_worker() {
1361 let session = block_on(RestartSession::new()).unwrap();
1362 let exit = session.worker.exit_notification();
1363 drop(session.shutdown());
1364 exit.wait();
1365 crate::RestartSession::new().unwrap().end().unwrap();
1366 }
1367}