#![cfg(all(unix, feature = "client"))]
use std::io::{self, Read, Write};
use std::os::fd::AsFd;
use std::process::{Child, ChildStdin, ChildStdout, Command, Stdio};
use std::time::{Duration, Instant};
use nix::fcntl::{FcntlArg, OFlag, fcntl};
use nix::poll::{PollFd, PollFlags, PollTimeout, poll};
use crate::client::Transport;
pub fn expand_tokens(s: &str, host: &str, port: u16, user: &str) -> String {
let mut out = String::with_capacity(s.len());
let mut chars = s.chars();
while let Some(c) = chars.next() {
if c != '%' {
out.push(c);
continue;
}
match chars.next() {
Some('h') => out.push_str(host),
Some('p') => {
use core::fmt::Write as _;
let _ = write!(out, "{port}");
}
Some('r') => out.push_str(user),
Some('%') => out.push('%'),
Some(other) => {
out.push('%');
out.push(other);
}
None => out.push('%'),
}
}
out
}
pub struct ProcTransport {
child: Child,
stdin: ChildStdin,
stdout: ChildStdout,
read_timeout: Option<Duration>,
}
impl ProcTransport {
pub fn spawn(command: &str) -> io::Result<Self> {
let mut child = Command::new("/bin/sh")
.arg("-c")
.arg(command)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::inherit())
.spawn()?;
let stdin = child
.stdin
.take()
.ok_or_else(|| io::Error::other("ProxyCommand: child stdin not captured"))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| io::Error::other("ProxyCommand: child stdout not captured"))?;
set_nonblocking(&stdout)?;
Ok(Self {
child,
stdin,
stdout,
read_timeout: None,
})
}
}
fn set_nonblocking<F: AsFd>(fd: &F) -> io::Result<()> {
let borrowed = fd.as_fd();
let cur = fcntl(borrowed, FcntlArg::F_GETFL).map_err(io::Error::from)?;
let flags = OFlag::from_bits_truncate(cur) | OFlag::O_NONBLOCK;
fcntl(borrowed, FcntlArg::F_SETFL(flags)).map_err(io::Error::from)?;
Ok(())
}
impl Read for ProcTransport {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let deadline = self.read_timeout.map(|d| Instant::now() + d);
loop {
let timeout: PollTimeout = match deadline {
None => PollTimeout::NONE,
Some(end) => {
let now = Instant::now();
if now >= end {
return Err(io::Error::from(io::ErrorKind::WouldBlock));
}
let remaining = end - now;
let ms = remaining.as_millis().min(u16::MAX as u128) as u16;
PollTimeout::from(ms)
}
};
let mut fds = [PollFd::new(self.stdout.as_fd(), PollFlags::POLLIN)];
match poll(&mut fds, timeout) {
Ok(0) => {
return Err(io::Error::from(io::ErrorKind::WouldBlock));
}
Ok(_) => {
match self.stdout.read(buf) {
Err(e) if e.kind() == io::ErrorKind::WouldBlock => continue,
other => return other,
}
}
Err(nix::errno::Errno::EINTR) => continue,
Err(e) => return Err(io::Error::from(e)),
}
}
}
}
impl Write for ProcTransport {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.stdin.write(buf)
}
fn flush(&mut self) -> io::Result<()> {
self.stdin.flush()
}
}
impl Transport for ProcTransport {
fn set_read_timeout(&mut self, t: Option<Duration>) -> io::Result<()> {
self.read_timeout = t;
Ok(())
}
}
impl Drop for ProcTransport {
fn drop(&mut self) {
let _ = self.child.kill();
let _ = self.child.wait();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn expand_basic_tokens() {
assert_eq!(
expand_tokens("nc %h %p", "example.com", 2222, "alice"),
"nc example.com 2222"
);
assert_eq!(expand_tokens("login=%r", "h", 22, "bob"), "login=bob");
assert_eq!(expand_tokens("100%%done", "h", 22, "u"), "100%done");
}
#[test]
fn expand_unknown_token_passthrough() {
assert_eq!(expand_tokens("a%zb", "h", 22, "u"), "a%zb");
assert_eq!(expand_tokens("end%", "h", 22, "u"), "end%");
}
#[test]
fn spawn_failure_is_strict_error() {
let mut t = ProcTransport::spawn("exit 0").expect("sh -c spawns");
let mut buf = [0u8; 16];
let n = t.read(&mut buf).expect("read after child exit");
assert_eq!(n, 0, "closed child stdout should read as EOF");
}
#[test]
fn round_trip_through_cat() {
let mut t = ProcTransport::spawn("cat").expect("spawn cat");
t.write_all(b"hello pipe").expect("write");
t.flush().expect("flush");
let mut buf = [0u8; 10];
t.read_exact(&mut buf).expect("read echo");
assert_eq!(&buf, b"hello pipe");
}
#[test]
fn read_timeout_ticks_as_wouldblock() {
let mut t = ProcTransport::spawn("sleep 5").expect("spawn sleep");
t.set_read_timeout(Some(Duration::from_millis(50)))
.expect("set timeout");
let mut buf = [0u8; 16];
let start = Instant::now();
let err = t.read(&mut buf).expect_err("should time out");
assert_eq!(err.kind(), io::ErrorKind::WouldBlock);
assert!(
start.elapsed() >= Duration::from_millis(40),
"poll should have waited out the deadline, took {:?}",
start.elapsed()
);
}
#[test]
fn read_timeout_still_delivers_data() {
let mut t = ProcTransport::spawn("printf abc; sleep 5").expect("spawn");
t.set_read_timeout(Some(Duration::from_millis(500)))
.expect("set timeout");
let mut buf = [0u8; 3];
t.read_exact(&mut buf).expect("read the printf output");
assert_eq!(&buf, b"abc");
}
#[test]
fn blocking_read_after_clearing_timeout() {
let mut t = ProcTransport::spawn("sleep 0.1; printf ok").expect("spawn");
t.set_read_timeout(None).expect("clear timeout");
let mut buf = [0u8; 2];
t.read_exact(&mut buf)
.expect("blocking read waits for data");
assert_eq!(&buf, b"ok");
}
}