git.delta.rocks / fleet / refs/commits / 087938500f2e

difftreelog

source

remowt/crates/remowt-endpoints/src/subprocess.rs7.2 KiBsourcehistory
1use std::collections::HashMap;2use std::io;3use std::process::Stdio;4use std::sync::atomic::{AtomicU64, Ordering};5use std::sync::{Arc, Mutex};67use bifrostlink::declarative::endpoints;8use bifrostlink::Config;9use camino::Utf8PathBuf;10use nix::sys::signal::{self, Signal};11use nix::unistd::Pid;12use remowt_link_shared::iroh_tunnel::{TunnelAddr, TunnelDialer, TunnelStream};13use serde::{Deserialize, Serialize};14use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};15use tokio::process::{ChildStderr, ChildStdout, Command};16use tokio::sync::{mpsc, watch};17use tracing::{debug, warn};1819pub type ProcId = u64;2021#[derive(Serialize, Deserialize, Debug)]22pub enum StdioSpec {23	Null,24	Tunnel(TunnelAddr),25}2627#[derive(Serialize, Deserialize, Debug)]28pub enum StderrSpec {29	Null,30	Tunnel(TunnelAddr),31	MergeWithStdout,32}3334#[derive(Serialize, Deserialize, Debug)]35pub struct SpawnSpec {36	pub program: String,37	pub args: Vec<String>,38	pub env: Vec<(String, String)>,39	pub env_clear: bool,40	pub cwd: Option<Utf8PathBuf>,41	pub stdin: StdioSpec,42	pub stdout: StdioSpec,43	pub stderr: StderrSpec,44}4546#[derive(Serialize, Deserialize, Debug, thiserror::Error)]47pub enum Error {48	#[error("spawn failed: {0}")]49	Spawn(String),50	#[error("connect to forwarded socket failed: {0}")]51	Connect(String),52	#[error("no process with that id")]53	NoSuchProcess,54	#[error("MergeWithStdout requires stdout=Socket")]55	BadMerge,56	#[error("invalid signal: {0}")]57	BadSignal(i32),58	#[error("kill failed: {0}")]59	Kill(String),60	#[error("io error: {0}")]61	Io(String),62}6364impl From<io::Error> for Error {65	fn from(e: io::Error) -> Self {66		Error::Io(e.to_string())67	}68}6970struct ChildState {71	pid: u32,72	exit_rx: watch::Receiver<Option<Option<i32>>>,73}7475#[derive(Clone)]76pub struct Subprocess {77	children: Arc<Mutex<HashMap<ProcId, ChildState>>>,78	next_id: Arc<AtomicU64>,79	dialer: Arc<TunnelDialer>,80}8182impl Subprocess {83	pub fn new(dialer: Arc<TunnelDialer>) -> Self {84		Self {85			children: Default::default(),86			next_id: Default::default(),87			dialer,88		}89	}90}9192#[endpoints(ns = 10)]93impl Subprocess {94	#[endpoints(id = 1)]95	async fn spawn(&self, spec: SpawnSpec) -> Result<ProcId, Error> {96		let SpawnSpec {97			program,98			args,99			env,100			env_clear,101			cwd,102			stdin,103			stdout,104			stderr,105		} = spec;106107		if matches!(stderr, StderrSpec::MergeWithStdout) && !matches!(stdout, StdioSpec::Tunnel(_))108		{109			return Err(Error::BadMerge);110		}111112		let mut cmd = Command::new(&program);113		cmd.args(&args);114		if env_clear {115			cmd.env_clear();116		}117		for (k, v) in &env {118			cmd.env(k, v);119		}120		if let Some(cwd) = &cwd {121			cmd.current_dir(cwd);122		}123		cmd.stdin(match &stdin {124			StdioSpec::Tunnel(_) => Stdio::piped(),125			StdioSpec::Null => Stdio::null(),126		});127		cmd.stdout(match &stdout {128			StdioSpec::Tunnel(_) => Stdio::piped(),129			StdioSpec::Null => Stdio::null(),130		});131		cmd.stderr(match &stderr {132			StderrSpec::Tunnel(_) | StderrSpec::MergeWithStdout => Stdio::piped(),133			StderrSpec::Null => Stdio::null(),134		});135		cmd.kill_on_drop(false);136137		let mut child = cmd.spawn().map_err(|e| Error::Spawn(e.to_string()))?;138		let pid = child139			.id()140			.ok_or_else(|| Error::Spawn("child exited before pid available".to_owned()))?;141142		if let StdioSpec::Tunnel(addr) = &stdin {143			let sock = self144				.dialer145				.connect_tunnel(addr)146				.await147				.map_err(|e| Error::Connect(e.to_string()))?;148			let mut stdin_w = child.stdin.take().expect("piped");149			tokio::spawn(async move {150				let (mut sr, _) = tokio::io::split(sock);151				let _ = tokio::io::copy(&mut sr, &mut stdin_w).await;152				let _ = stdin_w.shutdown().await;153			});154		}155156		let stdout_handle = child.stdout.take();157		let stderr_handle = child.stderr.take();158159		match (&stdout, &stderr, stdout_handle, stderr_handle) {160			(StdioSpec::Tunnel(out_addr), StderrSpec::MergeWithStdout, Some(out), Some(err)) => {161				let sock = self162					.dialer163					.connect_tunnel(out_addr)164					.await165					.map_err(|e| Error::Connect(e.to_string()))?;166				tokio::spawn(merge_to_sock(out, err, sock));167			}168			(StdioSpec::Tunnel(out_addr), _, Some(out), err_opt) => {169				let sock = self170					.dialer171					.connect_tunnel(out_addr)172					.await173					.map_err(|e| Error::Connect(e.to_string()))?;174				tokio::spawn(pump_to_sock(out, sock));175				if let (StderrSpec::Tunnel(err_addr), Some(err)) = (&stderr, err_opt) {176					let err_sock = self177						.dialer178						.connect_tunnel(err_addr)179						.await180						.map_err(|e| Error::Connect(e.to_string()))?;181					tokio::spawn(pump_to_sock(err, err_sock));182				}183			}184			(StdioSpec::Null, StderrSpec::Tunnel(err_addr), _, Some(err)) => {185				let sock = self186					.dialer187					.connect_tunnel(err_addr)188					.await189					.map_err(|e| Error::Connect(e.to_string()))?;190				tokio::spawn(pump_to_sock(err, sock));191			}192			_ => {}193		}194195		let (exit_tx, exit_rx) = watch::channel(None);196		let id = self.next_id.fetch_add(1, Ordering::Relaxed);197		self.children198			.lock()199			.expect("not poisoned")200			.insert(id, ChildState { pid, exit_rx });201202		debug!(id, pid, program, "subprocess spawned");203		tokio::spawn(async move {204			let result = child.wait().await;205			let code = match result {206				Ok(status) => status.code(),207				Err(e) => {208					warn!(id, "child.wait failed: {e}");209					None210				}211			};212			let _ = exit_tx.send(Some(code));213		});214215		Ok(id)216	}217218	#[endpoints(id = 2)]219	async fn wait(&self, id: ProcId) -> Result<Option<i32>, Error> {220		let mut rx = {221			let map = self.children.lock().expect("not poisoned");222			let entry = map.get(&id).ok_or(Error::NoSuchProcess)?;223			entry.exit_rx.clone()224		};225		rx.wait_for(|v| v.is_some())226			.await227			.map_err(|_| Error::Io("exit channel closed".to_owned()))?;228		let code = rx.borrow().flatten();229		self.children.lock().expect("not poisoned").remove(&id);230		Ok(code)231	}232233	#[endpoints(id = 3)]234	async fn kill(&self, id: ProcId, signal: i32) -> Result<(), Error> {235		let pid = {236			let map = self.children.lock().expect("not poisoned");237			let entry = map.get(&id).ok_or(Error::NoSuchProcess)?;238			entry.pid239		};240		let sig = Signal::try_from(signal).map_err(|_| Error::BadSignal(signal))?;241		signal::kill(Pid::from_raw(pid as i32), sig).map_err(|e| Error::Kill(e.to_string()))?;242		Ok(())243	}244}245246async fn pump_to_sock<R>(mut from: R, sock: TunnelStream)247where248	R: tokio::io::AsyncRead + Unpin,249{250	let (_, mut sw) = tokio::io::split(sock);251	let _ = tokio::io::copy(&mut from, &mut sw).await;252	let _ = sw.shutdown().await;253}254255async fn merge_to_sock(mut stdout: ChildStdout, mut stderr: ChildStderr, sock: TunnelStream) {256	let (_, mut sw) = tokio::io::split(sock);257	let (tx, mut rx) = mpsc::channel::<Vec<u8>>(64);258	let tx_out = tx.clone();259	let out_pump = tokio::spawn(async move {260		let mut buf = vec![0u8; 4096];261		loop {262			match stdout.read(&mut buf).await {263				Ok(0) | Err(_) => break,264				Ok(n) => {265					if tx_out.send(buf[..n].to_vec()).await.is_err() {266						break;267					}268				}269			}270		}271	});272	let err_pump = tokio::spawn(async move {273		let mut buf = vec![0u8; 4096];274		loop {275			match stderr.read(&mut buf).await {276				Ok(0) | Err(_) => break,277				Ok(n) => {278					if tx.send(buf[..n].to_vec()).await.is_err() {279						break;280					}281				}282			}283		}284	});285	while let Some(chunk) = rx.recv().await {286		if sw.write_all(&chunk).await.is_err() {287			break;288		}289	}290	let _ = out_pump.await;291	let _ = err_pump.await;292	let _ = sw.shutdown().await;293}