use super::{Connection, MAX_FRAME_BYTES};
use std::{
collections::VecDeque,
io,
os::fd::RawFd,
time::{Duration, Instant},
};
const MAX_PENDING: usize = 32;
const OUTPUT_DEADLINE: Duration = Duration::from_secs(2);
struct PipeFlags {
fd: RawFd,
original: i32,
}
impl PipeFlags {
fn nonblocking(fd: RawFd) -> io::Result<Self> {
let mut stat = std::mem::MaybeUninit::<libc::stat>::uninit();
if unsafe { libc::fstat(fd, stat.as_mut_ptr()) } != 0 {
return Err(io::Error::last_os_error());
}
if unsafe { stat.assume_init() }.st_mode & libc::S_IFMT != libc::S_IFIFO {
return Err(io::Error::new(
io::ErrorKind::Unsupported,
"stdio must be pipes",
));
}
let original = unsafe { libc::fcntl(fd, libc::F_GETFL) };
if original < 0
|| unsafe { libc::fcntl(fd, libc::F_SETFL, original | libc::O_NONBLOCK) } < 0
{
return Err(io::Error::last_os_error());
}
Ok(Self { fd, original })
}
}
impl Drop for PipeFlags {
fn drop(&mut self) {
unsafe { libc::fcntl(self.fd, libc::F_SETFL, self.original) };
}
}
struct Pending {
bytes: Vec<u8>,
offset: usize,
deadline: Instant,
}
fn enqueue(queue: &mut VecDeque<Pending>, bytes: Vec<u8>) -> io::Result<()> {
if queue.len() == MAX_PENDING || bytes.len() > MAX_FRAME_BYTES {
return Err(io::Error::other("output overflow"));
}
queue.push_back(Pending {
bytes,
offset: 0,
deadline: Instant::now() + OUTPUT_DEADLINE,
});
Ok(())
}
pub(super) fn run() -> anyhow::Result<()> {
let _input_flags = PipeFlags::nonblocking(libc::STDIN_FILENO)?;
let _output_flags = PipeFlags::nonblocking(libc::STDOUT_FILENO)?;
let mut connection = Connection::new()?;
let mut pending = VecDeque::<Pending>::new();
let mut input = Vec::with_capacity(MAX_FRAME_BYTES);
let mut closing = false;
let mut closing_deadline = None;
loop {
if !closing {
if let Some(response) = connection.poll_startup() {
enqueue(&mut pending, response)?;
}
if let Some(response) = connection.poll_session() {
enqueue(&mut pending, response)?;
}
} else {
connection.cancel_startup();
}
if closing {
let deadline = closing_deadline.get_or_insert_with(|| Instant::now() + OUTPUT_DEADLINE);
if Instant::now() >= *deadline {
anyhow::bail!("closing deadline");
}
}
if pending.len() < MAX_PENDING
&& let Some(event) = connection.poll_turn()
{
enqueue(&mut pending, event)?;
}
if closing && pending.is_empty() && connection.turn_worker.is_none() {
return Ok(());
}
if pending
.front()
.is_some_and(|frame| Instant::now() >= frame.deadline)
{
anyhow::bail!("output deadline");
}
let mut descriptors = [
libc::pollfd {
fd: if closing { -1 } else { libc::STDIN_FILENO },
events: libc::POLLIN,
revents: 0,
},
libc::pollfd {
fd: libc::STDOUT_FILENO,
events: if pending.is_empty() { 0 } else { libc::POLLOUT },
revents: 0,
},
];
if unsafe { libc::poll(descriptors.as_mut_ptr(), 2, 20) } < 0 {
if io::Error::last_os_error().kind() == io::ErrorKind::Interrupted {
continue;
}
return Err(io::Error::last_os_error().into());
}
if pending.is_empty() {
descriptors[1].events = libc::POLLOUT;
if unsafe { libc::poll(&mut descriptors[1], 1, 0) } < 0 {
if io::Error::last_os_error().kind() == io::ErrorKind::Interrupted {
continue;
}
return Err(io::Error::last_os_error().into());
}
}
if descriptors[1].revents & (libc::POLLERR | libc::POLLHUP | libc::POLLNVAL) != 0 {
anyhow::bail!("output disconnected");
}
if !pending.is_empty() && descriptors[1].revents & libc::POLLOUT != 0 {
let frame = pending
.front_mut()
.expect("only poll output with pending frame");
let remaining = &frame.bytes[frame.offset..];
let written = unsafe {
libc::write(
libc::STDOUT_FILENO,
remaining.as_ptr().cast(),
remaining.len(),
)
};
if written > 0 {
frame.offset += written as usize;
if frame.offset == frame.bytes.len() {
pending.pop_front();
}
} else if written == 0 {
anyhow::bail!("output closed");
} else {
let error = io::Error::last_os_error();
if !matches!(
error.kind(),
io::ErrorKind::WouldBlock | io::ErrorKind::Interrupted
) {
return Err(error.into());
}
}
}
if descriptors[0].revents & (libc::POLLERR | libc::POLLNVAL) != 0 {
anyhow::bail!("input disconnected");
}
if descriptors[0].revents & (libc::POLLIN | libc::POLLHUP) != 0 {
let mut buffer = [0u8; 4096];
let count =
unsafe { libc::read(libc::STDIN_FILENO, buffer.as_mut_ptr().cast(), buffer.len()) };
if count == 0 {
closing = true;
connection.cancel_startup();
if !input.is_empty() {
enqueue(
&mut pending,
connection.response(None, Err("incomplete_frame")),
)?;
}
} else if count < 0 {
let error = io::Error::last_os_error();
if !matches!(
error.kind(),
io::ErrorKind::WouldBlock | io::ErrorKind::Interrupted
) {
return Err(error.into());
}
} else {
for &byte in &buffer[..count as usize] {
if byte == b'\n' {
let (response, stop) = connection.receive(&input);
input.clear();
if let Some(response) = response {
enqueue(&mut pending, response)?;
}
if stop {
if let Some(response) = connection.begin_shutdown() {
enqueue(&mut pending, response)?;
}
closing = true;
break;
}
} else if input.len() == MAX_FRAME_BYTES {
enqueue(
&mut pending,
connection.response(None, Err("frame_too_large")),
)?;
closing = true;
break;
} else {
input.push(byte);
}
}
}
}
}
}