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 }