use std::collections::HashMap;
use std::io::{BufRead, BufReader, Write};
use std::path::PathBuf;
use std::process::{Child, ChildStdin, ChildStdout};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{mpsc, Arc, Mutex};
use serde_json::Value;
pub type RuntimeHandler = Arc<dyn Fn(&str, Value) -> Result<Value, String> + Send + Sync>;
type Response = Result<Value, String>;
type Pending = Arc<Mutex<HashMap<u64, mpsc::Sender<Response>>>>;
pub type ToolUpdateHandler = Arc<dyn Fn(Value) + Send + Sync>;
const NODE_STOPPED: &str = "Node extension host stopped";
struct Inner {
child: Arc<Mutex<Child>>,
stdin: Arc<Mutex<ChildStdin>>,
pending: Pending,
tool_update_handlers: Arc<Mutex<HashMap<String, ToolUpdateHandler>>>,
runtime_handlers: Arc<Mutex<Vec<RuntimeHandler>>>,
next_id: AtomicU64,
shutdown: Arc<AtomicBool>,
reader: Mutex<Option<std::thread::JoinHandle<()>>>,
cleanup_path: Option<PathBuf>,
}
impl Inner {
fn is_shutdown(&self) -> bool {
self.shutdown.load(Ordering::Acquire)
}
fn mark_failed(&self, error: &str) {
self.shutdown.store(true, Ordering::Release);
fail_pending(&self.pending, error);
terminate_child(&self.child);
}
fn shutdown(&self) {
let first_shutdown = !self.shutdown.swap(true, Ordering::AcqRel);
if first_shutdown {
fail_pending(&self.pending, NODE_STOPPED);
}
terminate_child(&self.child);
}
}
impl Drop for Inner {
fn drop(&mut self) {
self.shutdown();
if let Ok(reader) = self.reader.get_mut() {
if let Some(reader) = reader.take() {
let _ = reader.join();
}
}
if let Some(path) = self.cleanup_path.as_ref() {
let _ = std::fs::remove_file(path);
}
}
}
#[derive(Clone)]
pub struct NodeTransport {
inner: Arc<Inner>,
}
pub struct PendingRequest {
id: u64,
receiver: mpsc::Receiver<Response>,
}
pub struct NodeTransportStartup {
transport: NodeTransport,
init_receiver: mpsc::Receiver<Response>,
}
impl NodeTransportStartup {
pub fn transport(&self) -> NodeTransport {
self.transport.clone()
}
pub fn wait(self) -> Result<(NodeTransport, Value), String> {
let result = match self.init_receiver.recv() {
Ok(Ok(init)) => Ok(init),
Ok(Err(error)) => Err(error),
Err(_) => Err("Node extension host exited before initialization".into()),
};
match result {
Ok(init) => Ok((self.transport, init)),
Err(error) => {
self.transport.shutdown();
Err(error)
}
}
}
}
impl PendingRequest {
pub fn id(&self) -> u64 {
self.id
}
pub fn wait(self) -> Response {
self.receiver
.recv()
.unwrap_or_else(|_| Err("Node extension response channel closed".into()))
}
}
impl NodeTransport {
pub fn start(
child: Child,
stdin: ChildStdin,
stdout: ChildStdout,
) -> Result<(Self, Value), String> {
Self::start_with_cleanup_and_handlers(child, stdin, stdout, None, Vec::new())
}
pub fn start_with_cleanup(
child: Child,
stdin: ChildStdin,
stdout: ChildStdout,
cleanup_path: Option<PathBuf>,
) -> Result<(Self, Value), String> {
Self::start_with_cleanup_and_handlers(child, stdin, stdout, cleanup_path, Vec::new())
}
pub fn start_with_cleanup_and_handlers(
child: Child,
stdin: ChildStdin,
stdout: ChildStdout,
cleanup_path: Option<PathBuf>,
initial_handlers: Vec<RuntimeHandler>,
) -> Result<(Self, Value), String> {
Self::start_pending_with_cleanup_and_handlers(
child,
stdin,
stdout,
cleanup_path,
initial_handlers,
)?
.wait()
}
pub fn start_pending_with_cleanup_and_handlers(
child: Child,
stdin: ChildStdin,
stdout: ChildStdout,
cleanup_path: Option<PathBuf>,
initial_handlers: Vec<RuntimeHandler>,
) -> Result<NodeTransportStartup, String> {
let stdout = BufReader::new(stdout);
let stdin = Arc::new(Mutex::new(stdin));
let pending = Arc::new(Mutex::new(HashMap::new()));
let tool_update_handlers = Arc::new(Mutex::new(HashMap::new()));
let runtime_handlers = Arc::new(Mutex::new(initial_handlers));
let shutdown = Arc::new(AtomicBool::new(false));
let (init_sender, init_receiver) = mpsc::channel();
let child = Arc::new(Mutex::new(child));
let inner = Arc::new(Inner {
child: child.clone(),
stdin: stdin.clone(),
pending: pending.clone(),
tool_update_handlers: tool_update_handlers.clone(),
runtime_handlers: runtime_handlers.clone(),
next_id: AtomicU64::new(1),
shutdown: shutdown.clone(),
reader: Mutex::new(None),
cleanup_path,
});
let reader = std::thread::Builder::new()
.name("rpi-node-transport".into())
.spawn(move || {
read_loop(
stdout,
stdin,
child,
pending,
tool_update_handlers,
runtime_handlers,
Some(init_sender),
shutdown,
)
})
.map_err(|error| {
terminate_child(&inner.child);
if let Some(path) = inner.cleanup_path.as_ref() {
let _ = std::fs::remove_file(path);
}
format!("could not start Node response reader: {error}")
})?;
if let Err(error) = inner.reader.lock().map(|mut slot| *slot = Some(reader)) {
terminate_child(&inner.child);
if let Some(path) = inner.cleanup_path.as_ref() {
let _ = std::fs::remove_file(path);
}
return Err(format!("Node reader lock poisoned: {error}"));
}
Ok(NodeTransportStartup {
transport: Self { inner },
init_receiver,
})
}
pub fn shutdown(&self) {
self.inner.shutdown();
}
pub fn is_shutdown(&self) -> bool {
self.inner.is_shutdown()
}
pub fn same_instance(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.inner, &other.inner)
}
pub fn begin_request(&self, method: &str, payload: Value) -> Result<PendingRequest, String> {
if self.inner.is_shutdown() {
return Err(NODE_STOPPED.into());
}
let id = self.inner.next_id.fetch_add(1, Ordering::Relaxed);
let (sender, receiver) = mpsc::channel();
let mut pending = self
.inner
.pending
.lock()
.map_err(|_| "Node pending-request lock poisoned")?;
if self.inner.is_shutdown() {
return Err(NODE_STOPPED.into());
}
pending.insert(id, sender);
if self.inner.is_shutdown() {
let sender = pending.remove(&id);
drop(pending);
if let Some(sender) = sender {
let _ = sender.send(Err(NODE_STOPPED.into()));
}
return Err(NODE_STOPPED.into());
}
drop(pending);
let mut request = serde_json::json!({"id": id, "method": method});
if let (Some(object), Some(values)) = (request.as_object_mut(), payload.as_object()) {
object.extend(values.clone());
}
if let Err(error) = write_value(&self.inner.stdin, &request) {
if let Ok(mut pending) = self.inner.pending.lock() {
pending.remove(&id);
}
self.inner.mark_failed(&error);
return Err(error);
}
Ok(PendingRequest { id, receiver })
}
pub fn request(&self, method: &str, payload: Value) -> Response {
self.begin_request(method, payload)?.wait()
}
pub fn send_event(&self, event: Value) -> Result<(), String> {
if self.inner.is_shutdown() {
return Err(NODE_STOPPED.into());
}
if let Err(error) = write_value(&self.inner.stdin, &event) {
self.inner.mark_failed(&error);
return Err(error);
}
Ok(())
}
pub fn request_event(&self, mut event: Value) -> Response {
if self.inner.is_shutdown() {
return Err(NODE_STOPPED.into());
}
let supported = event
.get("event")
.and_then(Value::as_str)
.is_some_and(|name| matches!(name, "custom_input" | "custom_resize"));
if !supported {
return Err("host event does not support request/response mode".into());
}
let id = self.inner.next_id.fetch_add(1, Ordering::Relaxed);
let (sender, receiver) = mpsc::channel();
let mut pending = self
.inner
.pending
.lock()
.map_err(|_| "Node pending-request lock poisoned")?;
if self.inner.is_shutdown() {
return Err(NODE_STOPPED.into());
}
pending.insert(id, sender);
if self.inner.is_shutdown() {
let sender = pending.remove(&id);
drop(pending);
if let Some(sender) = sender {
let _ = sender.send(Err(NODE_STOPPED.into()));
}
return Err(NODE_STOPPED.into());
}
drop(pending);
let Some(object) = event.as_object_mut() else {
if let Ok(mut pending) = self.inner.pending.lock() {
pending.remove(&id);
}
return Err("Node host event must be a JSON object".into());
};
object.insert("id".into(), Value::from(id));
if let Err(error) = write_value(&self.inner.stdin, &event) {
if let Ok(mut pending) = self.inner.pending.lock() {
pending.remove(&id);
}
self.inner.mark_failed(&error);
return Err(error);
}
receiver
.recv()
.unwrap_or_else(|_| Err("Node extension response channel closed".into()))
}
pub fn cancel(&self, id: u64) -> Result<(), String> {
self.send_event(serde_json::json!({
"type": "host_event",
"event": "cancel_request",
"id": id,
}))
}
pub fn register_tool_update_handler(
&self,
tool_call_id: impl Into<String>,
handler: ToolUpdateHandler,
) -> Result<(), String> {
self.inner
.tool_update_handlers
.lock()
.map_err(|_| "Node tool-update lock poisoned")?
.insert(tool_call_id.into(), handler);
Ok(())
}
pub fn unregister_tool_update_handler(&self, tool_call_id: &str) {
if let Ok(mut handlers) = self.inner.tool_update_handlers.lock() {
handlers.remove(tool_call_id);
}
}
pub fn replace_runtime_handlers(&self, handler: RuntimeHandler) -> Result<(), String> {
*self
.inner
.runtime_handlers
.lock()
.map_err(|_| "Node runtime handler lock poisoned")? = vec![handler];
Ok(())
}
pub fn add_runtime_handler(&self, handler: RuntimeHandler) -> Result<(), String> {
self.inner
.runtime_handlers
.lock()
.map_err(|_| "Node runtime handler lock poisoned")?
.push(handler);
Ok(())
}
}
fn read_loop(
mut stdout: BufReader<ChildStdout>,
stdin: Arc<Mutex<ChildStdin>>,
child: Arc<Mutex<Child>>,
pending: Pending,
tool_update_handlers: Arc<Mutex<HashMap<String, ToolUpdateHandler>>>,
handlers: Arc<Mutex<Vec<RuntimeHandler>>>,
init_sender: Option<mpsc::Sender<Response>>,
shutdown: Arc<AtomicBool>,
) {
let mut init_sender = init_sender;
loop {
let message = match read_value(&mut stdout) {
Ok(message) => message,
Err(error) => {
shutdown.store(true, Ordering::Release);
fail_pending(&pending, &error);
terminate_child(&child);
if let Some(sender) = init_sender.take() {
let _ = sender.send(Err(error));
}
return;
}
};
if init_sender.is_some() && message.get("id").and_then(Value::as_u64) == Some(0) {
if let Some(sender) = init_sender.take() {
let _ = sender.send(Ok(message));
}
continue;
}
if message.get("type").and_then(Value::as_str) == Some("runtime_request") {
let background = message
.get("action")
.and_then(Value::as_str)
.is_some_and(runtime_action_may_block);
if background {
let stdin = stdin.clone();
let handlers = handlers.clone();
let pending = pending.clone();
let child = child.clone();
let shutdown = shutdown.clone();
std::thread::spawn(move || {
handle_runtime_request(message, &stdin, &handlers, &pending, &child, &shutdown)
});
} else {
handle_runtime_request(message, &stdin, &handlers, &pending, &child, &shutdown);
}
continue;
}
if message.get("type").and_then(Value::as_str) == Some("host_event")
&& message.get("event").and_then(Value::as_str) == Some("tool_update")
{
let Some(tool_call_id) = message.get("toolCallId").and_then(Value::as_str) else {
continue;
};
let handler = tool_update_handlers
.lock()
.ok()
.and_then(|handlers| handlers.get(tool_call_id).cloned());
if let Some(handler) = handler {
handler(message.get("partialResult").cloned().unwrap_or_default());
}
continue;
}
let Some(id) = message.get("id").and_then(Value::as_u64) else {
continue;
};
let sender = pending
.lock()
.ok()
.and_then(|mut values| values.remove(&id));
let Some(sender) = sender else {
continue;
};
let response = response_from_message(&message);
let _ = sender.send(response);
}
}
fn terminate_child(child: &Arc<Mutex<Child>>) {
if let Ok(mut child) = child.lock() {
let _ = child.kill();
let _ = child.wait();
}
}
fn response_from_message(message: &Value) -> Response {
if message.get("ok").and_then(Value::as_bool) == Some(true) {
Ok(message.get("result").cloned().unwrap_or_default())
} else {
Err(message
.get("error")
.and_then(Value::as_str)
.unwrap_or("JS extension failed")
.to_string())
}
}
fn runtime_action_may_block(action: &str) -> bool {
action.starts_with("provider.")
|| action.starts_with("tools.")
|| action.starts_with("session.")
|| action == "ui.dialog"
}
fn handle_runtime_request(
message: Value,
stdin: &Arc<Mutex<ChildStdin>>,
handlers: &Arc<Mutex<Vec<RuntimeHandler>>>,
pending: &Pending,
child: &Arc<Mutex<Child>>,
shutdown: &Arc<AtomicBool>,
) {
let Some(request_id) = message.get("requestId").and_then(Value::as_u64) else {
return;
};
let Some(action) = message.get("action").and_then(Value::as_str) else {
return;
};
let args = message.get("args").cloned().unwrap_or_default();
let handlers = handlers.lock().map(|values| values.clone());
let mut result = Err(format!("unsupported capability: {action}"));
match handlers {
Ok(handlers) => {
for handler in handlers {
match handler(action, args.clone()) {
Ok(value) => {
result = Ok(value);
break;
}
Err(error) if error.starts_with("unsupported capability:") => {}
Err(error) => {
result = Err(error);
break;
}
}
}
}
Err(_) => result = Err("Node runtime handler lock poisoned".into()),
}
let response = match result {
Ok(result) => serde_json::json!({
"type": "runtime_response",
"requestId": request_id,
"ok": true,
"result": result,
}),
Err(error) => serde_json::json!({
"type": "runtime_response",
"requestId": request_id,
"ok": false,
"error": error,
}),
};
if let Err(error) = write_value(stdin, &response) {
shutdown.store(true, Ordering::Release);
fail_pending(pending, &error);
terminate_child(child);
}
}
fn write_value(stdin: &Arc<Mutex<ChildStdin>>, value: &Value) -> Result<(), String> {
let mut stdin = stdin.lock().map_err(|_| "Node host stdin lock poisoned")?;
writeln!(stdin, "{value}")
.map_err(|error| format!("could not write to Node extension host: {error}"))?;
stdin
.flush()
.map_err(|error| format!("could not flush Node extension host: {error}"))
}
fn read_value(reader: &mut BufReader<ChildStdout>) -> Result<Value, String> {
let mut line = String::new();
reader
.read_line(&mut line)
.map_err(|error| format!("could not read Node extension host: {error}"))?;
if line.trim().is_empty() {
return Err("Node extension host exited without a response".into());
}
serde_json::from_str(&line).map_err(|error| format!("invalid Node extension response: {error}"))
}
fn fail_pending(pending: &Pending, error: &str) {
let values = pending
.lock()
.map(|mut values| values.drain().map(|(_, sender)| sender).collect::<Vec<_>>())
.unwrap_or_default();
for sender in values {
let _ = sender.send(Err(error.to_string()));
}
}