git.delta.rocks / fleet / refs/heads / push-kyumtlkprzyo

difftreelog

source

remowt/crates/remowt-endpoints/src/pty.rs6.4 KiBsourcehistory
1use std::collections::HashMap;2use std::io;3use std::os::fd::{AsRawFd, OwnedFd};4use std::pin::Pin;5use std::process::Stdio;6use std::sync::atomic::{AtomicU64, Ordering};7use std::sync::{Arc, Mutex};8use std::task::{Context, Poll};910use bifrostlink::declarative::endpoints;11use bifrostlink::Config;12use nix::libc;13use nix::pty::{openpty, OpenptyResult, Winsize};14use remowt_link_shared::iroh_tunnel::{TunnelAddr, TunnelDialer};15use serde::{Deserialize, Serialize};16use tokio::io::unix::AsyncFd;17use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};18use tracing::{debug, info, warn};1920pub type ShellId = u64;2122#[derive(Serialize, Deserialize, Debug, thiserror::Error)]23pub enum Error {24	#[error("openpty failed: {0}")]25	Open(String),26	#[error("failed to spawn shell: {0}")]27	Spawn(String),28	#[error("failed to connect to forwarded socket: {0}")]29	Connect(String),30	#[error("no shell with that id")]31	NoSuchShell,32	#[error("resize failed: {0}")]33	Resize(String),34	#[error("io error: {0}")]35	Io(String),36}3738impl From<io::Error> for Error {39	fn from(e: io::Error) -> Self {40		Error::Io(e.to_string())41	}42}4344#[derive(Clone)]45pub struct Pty {46	shells: Arc<Mutex<HashMap<ShellId, OwnedFd>>>,47	next_id: Arc<AtomicU64>,48	dialer: Arc<TunnelDialer>,49}5051impl Pty {52	pub fn new(dialer: Arc<TunnelDialer>) -> Self {53		Self {54			shells: Default::default(),55			next_id: Default::default(),56			dialer,57		}58	}59}6061#[endpoints(ns = 7)]62impl Pty {63	#[endpoints(id = 1)]64	async fn open_shell(65		&self,66		tunnel: TunnelAddr,67		term: String,68		cols: u16,69		rows: u16,70	) -> Result<ShellId, Error> {71		let ws = Winsize {72			ws_row: rows,73			ws_col: cols,74			ws_xpixel: 0,75			ws_ypixel: 0,76		};77		let OpenptyResult { master, slave } =78			openpty(Some(&ws), None).map_err(|e| Error::Open(e.to_string()))?;7980		let shell = std::env::var("SHELL").unwrap_or_else(|_| "/bin/sh".to_owned());8182		let slave_in = slave.try_clone()?;83		let slave_out = slave.try_clone()?;84		let slave_err = slave;8586		let mut cmd = tokio::process::Command::new(&shell);87		cmd.env("TERM", &term);88		if let Ok(home) = std::env::var("HOME") {89			cmd.current_dir(home);90		}91		cmd.stdin(Stdio::from(slave_in));92		cmd.stdout(Stdio::from(slave_out));93		cmd.stderr(Stdio::from(slave_err));94		// SAFETY: only async-signal-safe calls (setsid, ioctl) before exec.95		unsafe {96			cmd.pre_exec(|| {97				nix::unistd::setsid().map_err(|e| io::Error::from_raw_os_error(e as i32))?;98				if libc::ioctl(0, libc::TIOCSCTTY as _, 0) < 0 {99					return Err(io::Error::last_os_error());100				}101				Ok(())102			});103		}104105		let mut child = cmd.spawn().map_err(|e| Error::Spawn(e.to_string()))?;106107		let resize_fd = master.try_clone()?;108		let id = self.next_id.fetch_add(1, Ordering::Relaxed);109		self.shells110			.lock()111			.expect("not poisoned")112			.insert(id, resize_fd);113114		let sock = match self.dialer.connect_tunnel(&tunnel).await {115			Ok(s) => s,116			Err(e) => {117				self.shells.lock().expect("not poisoned").remove(&id);118				let _ = child.kill().await;119				return Err(Error::Connect(e.to_string()));120			}121		};122		let pty = AsyncPty::new(master)?;123124		debug!(id, shell, "shell opened");125		let shells = self.shells.clone();126		tokio::spawn(async move {127			let mut pty = pty;128			let mut sock = sock;129			if let Err(e) = tokio::io::copy_bidirectional(&mut pty, &mut sock).await {130				warn!(id, "shell pump ended: {e}");131			}132			let _ = child.kill().await;133			shells.lock().expect("not poisoned").remove(&id);134			info!(id, "shell closed");135		});136137		Ok(id)138	}139140	#[endpoints(id = 2)]141	async fn resize(&self, id: ShellId, cols: u16, rows: u16) -> Result<(), Error> {142		let ws = libc::winsize {143			ws_row: rows,144			ws_col: cols,145			ws_xpixel: 0,146			ws_ypixel: 0,147		};148		let shells = self.shells.lock().expect("not poisoned");149		let fd = shells.get(&id).ok_or(Error::NoSuchShell)?;150		// SAFETY: `fd` is a live PTY master151		let rc = unsafe { libc::ioctl(fd.as_raw_fd(), libc::TIOCSWINSZ as _, &ws) };152		if rc < 0 {153			return Err(Error::Resize(io::Error::last_os_error().to_string()));154		}155		Ok(())156	}157}158159struct AsyncPty {160	fd: AsyncFd<OwnedFd>,161}162163impl AsyncPty {164	fn new(fd: OwnedFd) -> io::Result<Self> {165		let raw = fd.as_raw_fd();166		// SAFETY: standard F_GETFL/F_SETFL round-trip on a valid fd.167		unsafe {168			let flags = libc::fcntl(raw, libc::F_GETFL);169			if flags < 0 {170				return Err(io::Error::last_os_error());171			}172			if libc::fcntl(raw, libc::F_SETFL, flags | libc::O_NONBLOCK) < 0 {173				return Err(io::Error::last_os_error());174			}175		}176		Ok(Self {177			fd: AsyncFd::new(fd)?,178		})179	}180}181182impl AsyncRead for AsyncPty {183	fn poll_read(184		self: Pin<&mut Self>,185		cx: &mut Context<'_>,186		buf: &mut ReadBuf<'_>,187	) -> Poll<io::Result<()>> {188		let this = self.get_mut();189		loop {190			let mut guard = match this.fd.poll_read_ready(cx) {191				Poll::Ready(Ok(g)) => g,192				Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),193				Poll::Pending => return Poll::Pending,194			};195			let unfilled = buf.initialize_unfilled();196			let res = guard.try_io(|inner| {197				let fd = inner.get_ref().as_raw_fd();198				// SAFETY: writing into `unfilled`'s own backing storage.199				let n = unsafe { libc::read(fd, unfilled.as_mut_ptr().cast(), unfilled.len()) };200				if n < 0 {201					let err = io::Error::last_os_error();202					if err.raw_os_error() == Some(libc::EIO) {203						Ok(0)204					} else {205						Err(err)206					}207				} else {208					Ok(n as usize)209				}210			});211			match res {212				Ok(Ok(n)) => {213					buf.advance(n);214					return Poll::Ready(Ok(()));215				}216				Ok(Err(e)) => return Poll::Ready(Err(e)),217				Err(_would_block) => continue,218			}219		}220	}221}222223impl AsyncWrite for AsyncPty {224	fn poll_write(225		self: Pin<&mut Self>,226		cx: &mut Context<'_>,227		buf: &[u8],228	) -> Poll<io::Result<usize>> {229		let this = self.get_mut();230		loop {231			let mut guard = match this.fd.poll_write_ready(cx) {232				Poll::Ready(Ok(g)) => g,233				Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),234				Poll::Pending => return Poll::Pending,235			};236			let res = guard.try_io(|inner| {237				let fd = inner.get_ref().as_raw_fd();238				// SAFETY: reading from `buf` for `buf.len()` bytes.239				let n = unsafe { libc::write(fd, buf.as_ptr().cast(), buf.len()) };240				if n < 0 {241					Err(io::Error::last_os_error())242				} else {243					Ok(n as usize)244				}245			});246			match res {247				Ok(r) => return Poll::Ready(r),248				Err(_would_block) => continue,249			}250		}251	}252253	fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {254		Poll::Ready(Ok(()))255	}256257	fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {258		Poll::Ready(Ok(()))259	}260}