use std::future::Future;
use std::path::{Path, PathBuf};
use std::pin::Pin;
use std::sync::Arc;
use async_trait::async_trait;
use serde_json::Value;
use tempfile::TempDir;
pub use crate::uds::server::RpcError;
use crate::uds::server::{RpcFallback, RpcRouter, RpcServeOptions, serve_until};
pub type MockFuture = Pin<Box<dyn Future<Output = Result<Value, RpcError>> + Send>>;
pub struct MockUdsDaemon {
socket: PathBuf,
_dir: Option<TempDir>,
shutdown: Option<tokio::sync::oneshot::Sender<()>>,
}
impl MockUdsDaemon {
pub fn socket(&self) -> &Path {
&self.socket
}
}
impl Drop for MockUdsDaemon {
fn drop(&mut self) {
if let Some(tx) = self.shutdown.take() {
let _ = tx.send(());
}
}
}
struct MockFallback<F> {
handler: F,
}
#[async_trait]
impl<F> RpcFallback for MockFallback<F>
where
F: Fn(&str, Value) -> MockFuture + Send + Sync + 'static,
{
async fn call(&self, method: &str, params: Value) -> Result<Value, RpcError> {
(self.handler)(method, params).await
}
}
pub async fn spawn<F>(handler: F) -> MockUdsDaemon
where
F: Fn(&str, Value) -> MockFuture + Send + Sync + 'static,
{
let dir = TempDir::new().expect("tempdir for the mock socket");
let socket = dir.path().join("daemon.sock");
let mut daemon = spawn_at(socket, handler).await;
daemon._dir = Some(dir);
daemon
}
pub async fn spawn_at<F>(socket: PathBuf, handler: F) -> MockUdsDaemon
where
F: Fn(&str, Value) -> MockFuture + Send + Sync + 'static,
{
let listener = crate::uds::bind_hardened(&socket).expect("bind the mock socket");
let router = Arc::new(RpcRouter::new().fallback(MockFallback { handler }));
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
tokio::spawn(async move {
serve_until(&listener, router, RpcServeOptions::default(), async {
let _ = rx.await;
})
.await;
});
MockUdsDaemon {
socket,
_dir: None,
shutdown: Some(tx),
}
}
#[cfg(feature = "search-index")]
pub struct BlockingMockDaemon {
socket: PathBuf,
shutdown: Option<tokio::sync::oneshot::Sender<()>>,
thread: Option<std::thread::JoinHandle<()>>,
}
#[cfg(feature = "search-index")]
impl BlockingMockDaemon {
pub fn socket(&self) -> &Path {
&self.socket
}
}
#[cfg(feature = "search-index")]
impl Drop for BlockingMockDaemon {
fn drop(&mut self) {
if let Some(tx) = self.shutdown.take() {
let _ = tx.send(());
}
if let Some(thread) = self.thread.take() {
let _ = thread.join();
}
}
}
#[cfg(feature = "search-index")]
pub fn spawn_blocking_at<F>(socket: PathBuf, handler: F) -> BlockingMockDaemon
where
F: Fn(&str, Value) -> MockFuture + Send + Sync + 'static,
{
let (ready_tx, ready_rx) = std::sync::mpsc::channel::<()>();
let (stop_tx, stop_rx) = tokio::sync::oneshot::channel::<()>();
let bind_at = socket.clone();
let thread = std::thread::spawn(move || {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("runtime for the mock daemon");
runtime.block_on(async move {
let listener = crate::uds::bind_hardened(&bind_at).expect("bind the mock socket");
let router = Arc::new(RpcRouter::new().fallback(MockFallback { handler }));
let _ = ready_tx.send(());
serve_until(&listener, router, RpcServeOptions::default(), async {
let _ = stop_rx.await;
})
.await;
});
});
ready_rx
.recv()
.expect("the mock daemon must bind before the rig proceeds");
BlockingMockDaemon {
socket,
shutdown: Some(stop_tx),
thread: Some(thread),
}
}