app

Local-first trade for farms and co-ops
git clone https://radroots.dev/git/app.git
Log | Files | Refs | README | LICENSE

host_runtime.rs (5826B)


      1 use std::future::Future;
      2 use std::sync::{Arc, Mutex};
      3 use std::thread::JoinHandle;
      4 use std::time::Duration;
      5 
      6 use tokio::runtime::{Builder, Handle};
      7 use tokio::sync::{oneshot, watch};
      8 
      9 pub(crate) struct HostRuntime {
     10     handle: Handle,
     11     shutdown: Mutex<Option<oneshot::Sender<()>>>,
     12     completion: watch::Receiver<bool>,
     13     thread: Mutex<Option<JoinHandle<()>>>,
     14 }
     15 
     16 #[cfg(test)]
     17 pub(crate) struct CompletionGatedHostRuntime {
     18     pub(crate) runtime: Arc<HostRuntime>,
     19     pub(crate) entered: std::sync::mpsc::Receiver<()>,
     20     pub(crate) release: std::sync::mpsc::SyncSender<()>,
     21 }
     22 
     23 impl HostRuntime {
     24     pub(crate) fn new() -> Result<Arc<Self>, ()> {
     25         Self::new_inner(None)
     26     }
     27 
     28     fn new_inner(
     29         completion_gate: Option<(
     30             std::sync::mpsc::SyncSender<()>,
     31             std::sync::mpsc::Receiver<()>,
     32         )>,
     33     ) -> Result<Arc<Self>, ()> {
     34         let (startup_sender, startup_receiver) = std::sync::mpsc::sync_channel(1);
     35         let (shutdown_sender, shutdown_receiver) = oneshot::channel();
     36         let (completion_sender, completion_receiver) = watch::channel(false);
     37         let thread = std::thread::Builder::new()
     38             .name("harvestcircle-host-runtime".to_owned())
     39             .spawn(move || {
     40                 let Ok(runtime) = Builder::new_multi_thread()
     41                     .enable_all()
     42                     .thread_name("harvestcircle-runtime-worker")
     43                     .build()
     44                 else {
     45                     let _ = startup_sender.send(Err(()));
     46                     let _ = completion_sender.send(true);
     47                     return;
     48                 };
     49                 if startup_sender.send(Ok(runtime.handle().clone())).is_err() {
     50                     let _ = completion_sender.send(true);
     51                     return;
     52                 }
     53                 runtime.block_on(async {
     54                     let _ = shutdown_receiver.await;
     55                 });
     56                 runtime.shutdown_timeout(Duration::from_secs(5));
     57                 if let Some((entered, release)) = completion_gate {
     58                     let _ = entered.send(());
     59                     let _ = release.recv();
     60                 }
     61                 let _ = completion_sender.send(true);
     62             })
     63             .map_err(|_| ())?;
     64         let handle = startup_receiver.recv().map_err(|_| ())??;
     65         Ok(Arc::new(Self {
     66             handle,
     67             shutdown: Mutex::new(Some(shutdown_sender)),
     68             completion: completion_receiver,
     69             thread: Mutex::new(Some(thread)),
     70         }))
     71     }
     72 
     73     #[cfg(test)]
     74     pub(crate) fn new_completion_gated_for_test() -> Result<CompletionGatedHostRuntime, ()> {
     75         let (entered_sender, entered_receiver) = std::sync::mpsc::sync_channel(1);
     76         let (release_sender, release_receiver) = std::sync::mpsc::sync_channel(1);
     77         let runtime = Self::new_inner(Some((entered_sender, release_receiver)))?;
     78         Ok(CompletionGatedHostRuntime {
     79             runtime,
     80             entered: entered_receiver,
     81             release: release_sender,
     82         })
     83     }
     84 
     85     pub(crate) fn handle(&self) -> &Handle {
     86         &self.handle
     87     }
     88 
     89     pub(crate) fn block_on<F>(&self, future: F) -> Result<F::Output, ()>
     90     where
     91         F: Future + Send + 'static,
     92         F::Output: Send + 'static,
     93     {
     94         let (sender, receiver) = std::sync::mpsc::sync_channel(1);
     95         self.handle.spawn(async move {
     96             let _ = sender.send(future.await);
     97         });
     98         receiver.recv().map_err(|_| ())
     99     }
    100 
    101     pub(crate) async fn run<F>(&self, future: F) -> Result<F::Output, ()>
    102     where
    103         F: Future + Send + 'static,
    104         F::Output: Send + 'static,
    105     {
    106         self.handle.spawn(future).await.map_err(|_| ())
    107     }
    108 
    109     pub(crate) async fn shutdown(&self) -> Result<(), ()> {
    110         let sender = self.shutdown.lock().map_err(|_| ())?.take();
    111         if let Some(sender) = sender {
    112             let _ = sender.send(());
    113         }
    114 
    115         let mut completion = self.completion.clone();
    116         while !*completion.borrow() {
    117             completion.changed().await.map_err(|_| ())?;
    118         }
    119 
    120         let thread = self.thread.lock().map_err(|_| ())?.take();
    121         if let Some(thread) = thread {
    122             thread.join().map_err(|_| ())?;
    123         }
    124         Ok(())
    125     }
    126 }
    127 
    128 impl Drop for HostRuntime {
    129     fn drop(&mut self) {
    130         if let Ok(sender) = self.shutdown.get_mut()
    131             && let Some(sender) = sender.take()
    132         {
    133             let _ = sender.send(());
    134         }
    135     }
    136 }
    137 
    138 #[cfg(test)]
    139 mod tests {
    140     use std::sync::Arc;
    141 
    142     use super::HostRuntime;
    143 
    144     #[test]
    145     fn host_runtime_executes_work_and_shuts_down_explicitly() {
    146         let host = HostRuntime::new().expect("host runtime");
    147         assert_eq!(host.block_on(async { 7 }).expect("runtime result"), 7);
    148         let test_runtime = tokio::runtime::Runtime::new().expect("test runtime");
    149         test_runtime.block_on(host.shutdown()).expect("shutdown");
    150         test_runtime
    151             .block_on(host.shutdown())
    152             .expect("idempotent shutdown");
    153     }
    154 
    155     #[tokio::test]
    156     async fn cancelled_shutdown_is_resumable() {
    157         let gated = HostRuntime::new_completion_gated_for_test().expect("host runtime");
    158         let host = gated.runtime;
    159         let entered_receiver = gated.entered;
    160         let release_sender = gated.release;
    161         let first_host = Arc::clone(&host);
    162         let first = tokio::spawn(async move { first_host.shutdown().await });
    163         tokio::task::spawn_blocking(move || entered_receiver.recv())
    164             .await
    165             .expect("entered join")
    166             .expect("shutdown entered completion gate");
    167         first.abort();
    168         assert!(first.await.is_err());
    169         release_sender.send(()).expect("release shutdown");
    170         host.shutdown().await.expect("resumed shutdown");
    171     }
    172 }