Skip to main content

restart_manager/
tokio.rs

1//! Tokio facade backed by one dedicated standard worker thread per session.
2//!
3//! Every handle serializes commands through a FIFO channel. The blocking
4//! typestate always remains on its worker, including when an awaiting future is
5//! dropped. Progress uses a Tokio watch channel and therefore coalesces samples
6//! when a receiver is slower than Windows.
7//!
8//! Every public future returned by this module is [`Send`].
9
10use 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
461/// Error from an async consuming operation.
462///
463/// A callback-lease conflict occurs before native work and retains the
464/// reusable typestate. Worker failure or an internal state mismatch returns no
465/// state because the facade cannot prove that retrying it would be sound.
466pub 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    /// Returns the operation error.
484    #[must_use]
485    pub const fn error(&self) -> &Error {
486        &self.error
487    }
488
489    /// Returns reusable state only when native work is known not to have begun.
490    #[must_use]
491    pub const fn state(&self) -> Option<&T> {
492        self.state.as_ref()
493    }
494
495    /// Separates the optional reusable state and operation error.
496    #[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/// Cloneable, coalescing receiver for asynchronous native progress.
528#[derive(Clone)]
529pub struct ProgressReceiver {
530    receiver: watch::Receiver<Option<Progress>>,
531}
532
533impl ProgressReceiver {
534    /// Peeks at the latest progress value without marking it as received.
535    #[must_use]
536    pub fn peek(&self) -> Option<Progress> {
537        *self.receiver.borrow()
538    }
539
540    /// Receives a newer progress value, or `None` after normal operation end.
541    pub async fn recv(&mut self) -> Option<Progress> {
542        self.receiver.changed().await.ok()?;
543        self.peek()
544    }
545}
546
547/// A primary async Restart Manager session.
548pub struct RestartSession {
549    worker: Worker,
550    key: SessionKey,
551    cancellation: CancellationHandle,
552}
553
554impl RestartSession {
555    /// Starts a new session on its dedicated worker thread.
556    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    /// Returns the cross-process session key.
566    #[must_use]
567    pub fn session_key(&self) -> &SessionKey {
568        &self.key
569    }
570
571    /// Returns a cross-thread cancellation capability.
572    #[must_use]
573    pub fn cancellation_handle(&self) -> CancellationHandle {
574        self.cancellation.clone()
575    }
576
577    /// Registers a resource batch in one worker command.
578    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    /// Registers file paths.
588    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    /// Registers exact process identities.
601    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    /// Registers service short names.
614    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    /// Takes an affected-application snapshot.
627    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    /// Adds or replaces a filter.
634    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    /// Removes a filter.
645    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    /// Lists filters.
655    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    /// Starts a graceful shutdown future.
662    #[must_use]
663    pub fn shutdown(self) -> ShutdownFuture {
664        self.shutdown_with_options(ShutdownOptions::default())
665    }
666
667    /// Starts shutdown with explicit options.
668    #[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    /// Starts shutdown and returns a coalescing progress receiver.
687    #[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    /// Ends the session without shutdown.
719    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
733/// Future for an asynchronous shutdown attempt.
734///
735/// Dropping it requests best-effort cancellation. The worker still owns the
736/// session and will restart anything partially stopped before ending.
737pub 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/// Async recovery-required state after shutdown.
794#[must_use = "call restart or leave_stopped explicitly"]
795pub struct RestartPending {
796    worker: Worker,
797    shutdown: OperationOutcome,
798}
799
800impl RestartPending {
801    /// Returns the retained shutdown result.
802    #[must_use]
803    pub const fn shutdown_outcome(&self) -> &OperationOutcome {
804        &self.shutdown
805    }
806
807    /// Starts a restart future.
808    #[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    /// Starts restart and returns its coalescing progress receiver.
823    #[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    /// Explicitly opts out of restart.
845    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
854/// Future for restart without progress.
855///
856/// Dropping this future detaches the operation; the worker completes restart
857/// and ends the session.
858pub 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
886/// Future for restart with coalescing progress.
887pub 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
926/// Async completed-recovery state.
927pub struct RecoveryCompletion {
928    worker: Worker,
929    outcome: RecoveryOutcome,
930}
931
932impl RecoveryCompletion {
933    /// Returns both retained outcomes.
934    #[must_use]
935    pub const fn outcome(&self) -> &RecoveryOutcome {
936        &self.outcome
937    }
938
939    /// Takes a post-operation report.
940    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    /// Ends the worker-owned session and returns both outcomes.
947    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
955/// An async secondary-installer session.
956pub struct JoinedSession {
957    worker: Worker,
958    key: SessionKey,
959}
960
961impl JoinedSession {
962    /// Joins an existing session on a dedicated worker.
963    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    /// Returns the joined key.
970    #[must_use]
971    pub fn session_key(&self) -> &SessionKey {
972        &self.key
973    }
974
975    /// Registers a resource batch.
976    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    /// Registers file paths.
986    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    /// Registers process identities.
999    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    /// Registers service short names.
1012    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    /// Ends the joined handle.
1025    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}