#![allow(dead_code)]
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 trusty_common::uds::server::RpcError;
use trusty_common::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 = trusty_common::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),
}
}
pub fn tools_call_envelope(inner: &Value) -> Value {
serde_json::json!({ "content": [{ "type": "text", "text": inner.to_string() }] })
}
pub fn always(result: Value) -> impl Fn(&str, Value) -> MockFuture + Send + Sync + 'static {
move |_method, _params| {
let result = result.clone();
Box::pin(async move { Ok(result) })
}
}
pub struct BlockingMockDaemon {
socket: PathBuf,
shutdown: Option<tokio::sync::oneshot::Sender<()>>,
thread: Option<std::thread::JoinHandle<()>>,
}
impl BlockingMockDaemon {
pub fn socket(&self) -> &Path {
&self.socket
}
}
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();
}
}
}
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 =
trusty_common::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),
}
}