use std::cell::Cell;
use std::future::Future;
use std::pin::pin;
use std::sync::Arc;
use std::task::{Context, Poll, Wake, Waker};
use std::time::Duration;
use tokio::runtime::Handle;
use tokio::sync::oneshot;
use crate::{Error, Reply};
thread_local! {
static CLIENT_THREAD: Cell<bool> = const { Cell::new(false) };
}
pub(crate) fn current() -> Result<Handle, Error> {
Handle::try_current().map_err(|_| Error::invalid("this async call needs a tokio runtime (or use net_backend_client::blocking)"))
}
pub(crate) fn refuse_on_client_thread() -> Result<(), Error> {
if CLIENT_THREAD.with(Cell::get) {
return Err(Error::invalid("a blocking call made on the client's own runtime thread (e.g. from an SSH prompt responder) would wait for itself"));
}
Ok(())
}
struct ThreadWaker(std::thread::Thread);
impl Wake for ThreadWaker {
fn wake(self: Arc<Self>) {
self.0.unpark();
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.unpark();
}
}
pub(crate) fn park_on<F: Future>(future: F) -> F::Output {
let waker = Waker::from(Arc::new(ThreadWaker(std::thread::current())));
let mut context = Context::from_waker(&waker);
let mut future = pin!(future);
loop {
match future.as_mut().poll(&mut context) {
Poll::Ready(output) => return output,
Poll::Pending => std::thread::park(),
}
}
}
pub(crate) struct RuntimeThread {
handle: Handle,
stop: Option<oneshot::Sender<()>>,
}
impl RuntimeThread {
pub(crate) fn start() -> Result<Arc<Self>, Error> {
let (ready_sender, ready) = std::sync::mpsc::channel::<Result<Handle, String>>();
let (stop, stopped) = oneshot::channel::<()>();
std::thread::Builder::new()
.name("net-backend-client".into())
.spawn(move || {
CLIENT_THREAD.with(|flag| flag.set(true));
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.max_blocking_threads(2)
.thread_name("net-backend-client-io")
.on_thread_start(|| CLIENT_THREAD.with(|flag| flag.set(true)))
.build();
match runtime {
Ok(runtime) => {
let _ = ready_sender.send(Ok(runtime.handle().clone()));
runtime.block_on(async {
let _ = stopped.await;
});
runtime.shutdown_timeout(Duration::from_secs(1));
}
Err(e) => {
let _ = ready_sender.send(Err(e.to_string()));
}
}
})
.map_err(|e| Error::network(format!("could not start the client thread: {e}"), Some(false)))?;
let handle = ready
.recv()
.map_err(|_| Error::network("the client thread stopped at once", Some(false)))?
.map_err(|e| Error::network(format!("could not start the client runtime: {e}"), Some(false)))?;
Ok(Arc::new(Self { handle, stop: Some(stop) }))
}
pub(crate) fn spawn<T: Send + 'static>(&self, future: impl Future<Output = Result<T, Error>> + Send + 'static) -> Reply<T> {
let (sender, reply) = Reply::channel();
self.handle.spawn(async move {
let _ = sender.send(future.await);
});
reply
}
pub(crate) fn block<T: Send + 'static>(&self, future: impl Future<Output = Result<T, Error>> + Send + 'static) -> Result<T, Error> {
self.spawn(future).wait()
}
#[cfg(feature = "ssh")]
pub(crate) fn enter<T>(&self, work: impl FnOnce() -> T) -> T {
let _guard = self.handle.enter();
work()
}
}
impl Drop for RuntimeThread {
fn drop(&mut self) {
if let Some(stop) = self.stop.take() {
let _ = stop.send(());
}
}
}
impl std::fmt::Debug for RuntimeThread {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("RuntimeThread")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn only_the_client_thread_refuses_and_waiting_never_panics() {
assert!(refuse_on_client_thread().is_ok());
let refused = std::thread::spawn(|| {
CLIENT_THREAD.with(|flag| flag.set(true));
refuse_on_client_thread()
})
.join()
.expect("join");
assert!(matches!(refused, Err(Error::InvalidRequest(_))));
let runtime = tokio::runtime::Builder::new_current_thread().enable_all().build().expect("runtime");
runtime.block_on(async {
assert!(Reply::ready(Ok(5)).wait().ok() == Some(5));
});
}
}