Skip to content
File

Blob: src/rust/cxx-integration/tokio.rs

rust83 lines
1use std::future::Future;
2use std::sync::OnceLock;
3 
4use tokio::task::JoinHandle;
5use tracing::Instrument;
6use tracing::info;
7 
8static TOKIO_RUNTIME: OnceLock<tokio::runtime::Runtime> = OnceLock::new();
9 
10/// Initialize tokio runtime.
11/// Must be called after all forking and sandbox setup is finished.
12pub(crate) fn init(worker_threads: Option<usize>) {
13 assert!(TOKIO_RUNTIME.get().is_none());
14 
15 let mut builder = tokio::runtime::Builder::new_multi_thread();
16 
17 if let Some(worker_threads) = worker_threads {
18 builder.worker_threads(worker_threads);
19 }
20 
21 let runtime = builder
22 .enable_time()
23 .enable_io()
24 .build()
25 .expect("failed to build tokio runtime");
26 
27 TOKIO_RUNTIME
28 .set(runtime)
29 .expect("failed to set tokio runtime");
30 spawn(async {
31 info!(nosentry = true, "tokio runtime is online");
32 });
33}
34 
35/// Obtain a handle to the shared tokio runtime.
36/// Requires calling [`init_tokio`] first.
37/// # Panics
38/// if tokio runtime is not available yet.
39pub fn runtime_handle() -> tokio::runtime::Handle {
40 TOKIO_RUNTIME
41 .get()
42 .expect("tokio runtime is not initialized")
43 .handle()
44 .clone()
45}
46 
47/// This is helper to set the spawn and stuff duplicating the signature of
48/// `https://docs.rs/tokio/latest/tokio/task/fn.spawn.html`.
49pub fn spawn<F>(future: F) -> JoinHandle<F::Output>
50where
51 F: Future + Send + 'static,
52 F::Output: Send + 'static,
53{
54 let handle = runtime_handle();
55 let _guard = handle.enter();
56 
57 tokio::spawn(future.in_current_span())
58}
59 
60/// Exposes `Runtime::block_on`
61/// `https://docs.rs/tokio/latest/tokio/runtime/struct.Runtime.html#method.block_on`.
62pub fn block_on<F: Future>(f: F) -> F::Output {
63 runtime_handle().block_on(f)
64}
65 
66#[cfg(test)]
67mod test {
68 use std::time::Duration;
69 
70 use super::*;
71 
72 #[test]
73 fn test_tokio_init() {
74 init(None);
75 let join = spawn(async move {
76 tokio::time::sleep(Duration::from_millis(1)).await;
77 42
78 });
79 let result = runtime_handle().block_on(join).unwrap();
80 assert_eq!(42, result);
81 }
82}