use std::borrow::Cow;
use std::collections::VecDeque;
use std::io::{IoSlice, Write};
use std::mem::MaybeUninit;
use std::os::fd::{AsFd, AsRawFd, BorrowedFd, OwnedFd};
use std::os::unix::io::RawFd;
use std::os::unix::net::UnixStream;
use std::time::Instant;
use gnitz_wire::{Deframer, FrameLenError, HelloError};
use super::error::ProtocolError;
use crate::ClientError;
mod tls;
pub struct ClientTransport {
inner: Inner,
deframer: Deframer,
window: Box<[MaybeUninit<u8>]>,
torn: bool,
queue: OutQueue,
}
const WINDOW_BYTES: usize = 64 * 1024;
enum Inner {
Unix(UnixStream),
Tls(Box<tls::TlsInner>),
}
enum WriteOutcome {
Written(usize),
WouldBlock,
}
impl Inner {
fn write_slices(&mut self, slices: &[IoSlice<'_>]) -> Result<WriteOutcome, ProtocolError> {
match self {
Inner::Unix(s) => write_nonblocking(|| (&*s).write_vectored(slices)),
Inner::Tls(t) => t.write_slices(slices),
}
}
fn has_pending_ciphertext(&self) -> bool {
match self {
Inner::Unix(_) => false,
Inner::Tls(t) => t.wants_write(),
}
}
fn ship_ciphertext(&mut self) -> Result<(), ProtocolError> {
match self {
Inner::Unix(_) => Ok(()),
Inner::Tls(t) => t.ship(),
}
}
fn as_fd(&self) -> BorrowedFd<'_> {
match self {
Inner::Unix(s) => s.as_fd(),
Inner::Tls(t) => t.as_fd(),
}
}
}
fn write_nonblocking(mut write: impl FnMut() -> std::io::Result<usize>) -> Result<WriteOutcome, ProtocolError> {
loop {
match write() {
Ok(0) => return Err(std::io::Error::from(std::io::ErrorKind::WriteZero).into()),
Ok(n) => return Ok(WriteOutcome::Written(n)),
Err(e) if e.kind() == std::io::ErrorKind::Interrupted => {}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => return Ok(WriteOutcome::WouldBlock),
Err(e) => return Err(e.into()),
}
}
}
fn recv(fd: RawFd, buf: &mut [MaybeUninit<u8>]) -> Result<Option<usize>, ProtocolError> {
loop {
let n = unsafe { libc::recv(fd, buf.as_mut_ptr() as *mut libc::c_void, buf.len(), 0) };
if n >= 0 {
return Ok(Some(n as usize));
}
let e = std::io::Error::last_os_error();
match e.kind() {
std::io::ErrorKind::Interrupted => {}
std::io::ErrorKind::WouldBlock => return Ok(None),
_ => return Err(e.into()),
}
}
}
fn timed_out() -> std::io::Error {
std::io::Error::new(std::io::ErrorKind::TimedOut, "socket operation timed out")
}
pub(crate) fn poll_fd(fd: RawFd, events: libc::c_short, until: Option<Instant>) -> std::io::Result<libc::c_short> {
let timeout_ms: libc::c_int = match until {
None => -1,
Some(t) => {
let left = t.saturating_duration_since(Instant::now());
if left.is_zero() {
return Err(timed_out());
}
left.as_nanos().div_ceil(1_000_000).min(i32::MAX as u128) as libc::c_int
}
};
let mut pfd = libc::pollfd { fd, events, revents: 0 };
match unsafe { libc::poll(&mut pfd, 1, timeout_ms) } {
rc if rc < 0 => Err(std::io::Error::last_os_error()),
0 => Err(timed_out()),
_ => Ok(pfd.revents),
}
}
impl ClientTransport {
fn new(inner: Inner) -> Self {
ClientTransport {
inner,
deframer: Deframer::default(),
window: Box::new_uninit_slice(WINDOW_BYTES),
torn: false,
queue: OutQueue::default(),
}
}
pub fn try_clone_fd(&self) -> Result<OwnedFd, ProtocolError> {
Ok(self.inner.as_fd().try_clone_to_owned()?)
}
pub fn connect(target: &str, until: Instant) -> Result<Self, ClientError> {
let mut t = match target.strip_prefix("tls://") {
Some(rest) => tls::connect_tls(rest, until)?,
None => ClientTransport::unix(UnixStream::connect(target)?)?,
};
t.hello(until)?;
Ok(t)
}
pub(crate) fn unix(stream: UnixStream) -> Result<Self, ProtocolError> {
stream.set_nonblocking(true)?;
Ok(ClientTransport::new(Inner::Unix(stream)))
}
pub fn as_raw_fd(&self) -> RawFd {
self.inner.as_fd().as_raw_fd()
}
fn hello(&mut self, until: Instant) -> Result<(), ClientError> {
self.enqueue(gnitz_wire::HELLO.to_vec());
loop {
self.flush()?;
let events = if self.wants_write() {
libc::POLLIN | libc::POLLOUT
} else {
libc::POLLIN
};
match poll_fd(self.as_raw_fd(), events, Some(until)) {
Err(e) if e.kind() == std::io::ErrorKind::Interrupted => continue,
polled => polled.map_err(ProtocolError::from)?,
};
let mut frames = Vec::new();
let read = self.read(|frame| {
frames.push(frame.into_owned());
Ok(())
});
let reply = match &frames[..] {
[] => {
read?;
continue;
}
[reply] => reply,
_ => return Err(ProtocolError::DecodeError("a frame behind the server's HELLO".into()).into()),
};
gnitz_wire::check_hello(reply).map_err(|e| match e {
HelloError::Malformed => ClientError::from(ProtocolError::DecodeError(e.to_string())),
HelloError::Version { .. } => ClientError::from(e.to_string()),
})?;
read?;
return Ok(());
}
}
pub(crate) fn enqueue(&mut self, payload: Vec<u8>) {
self.queue.push(gnitz_wire::frame_len_prefix(payload.len()), payload);
}
pub(crate) fn flush(&mut self) -> Result<(), ProtocolError> {
loop {
self.inner.ship_ciphertext()?;
if self.inner.has_pending_ciphertext() || self.queue.is_empty() {
return Ok(());
}
let ClientTransport { inner, queue, .. } = self;
let mut slices: Vec<IoSlice<'_>> = Vec::with_capacity(IOV_MAX_CHUNK.min(2 * queue.len()));
queue.build_slices(&mut slices);
match inner.write_slices(&slices)? {
WriteOutcome::Written(n) => queue.advance(n),
WriteOutcome::WouldBlock => return Ok(()),
}
}
}
pub(crate) fn queued_bytes(&self) -> usize {
self.queue.bytes
}
pub(crate) fn wants_write(&self) -> bool {
!self.queue.is_empty() || self.inner.has_pending_ciphertext()
}
pub(crate) fn close(&mut self) {
let _ = unsafe { libc::shutdown(self.as_raw_fd(), libc::SHUT_RDWR) };
self.queue = OutQueue::default();
}
pub(crate) fn read(
&mut self,
mut on_frame: impl FnMut(Cow<'_, [u8]>) -> Result<(), ProtocolError>,
) -> Result<bool, ProtocolError> {
let ClientTransport { inner, deframer, window, torn, .. } = self;
if *torn {
return Err(std::io::Error::other("an earlier read was abandoned with bytes unfed").into());
}
let fd = inner.as_fd().as_raw_fd();
let into = match inner {
Inner::Unix(_) => deframer.window(window),
Inner::Tls(_) => &mut window[..],
};
let len = into.len();
let Some(n) = recv(fd, into)? else { return Ok(false) };
let mut feed = |deframer: &mut Deframer, mut src: &[u8]| {
while let Some((frame, ())) = deframer.feed(&mut src, |_| Ok::<_, ProtocolError>(()))? {
on_frame(frame)?;
}
Ok::<_, ProtocolError>(())
};
*torn = true;
let open = n > 0
&& match inner {
Inner::Unix(_) => {
let src = unsafe { deframer.landed(window, n) };
feed(deframer, src)?;
true
}
Inner::Tls(t) => t.ingest(unsafe { window[..n].assume_init_ref() }, |plain| feed(deframer, plain))?,
};
*torn = false;
if !open {
let msg = if deframer.is_mid_frame() {
"connection closed mid-frame"
} else {
"connection closed"
};
return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, msg).into());
}
Ok(n == len)
}
}
const IOV_MAX_CHUNK: usize = 1024;
struct QueuedFrame {
prefix: [u8; gnitz_wire::FRAME_LEN_PREFIX_BYTES],
payload: Vec<u8>,
}
impl QueuedFrame {
fn segments(&self) -> [&[u8]; 2] {
[&self.prefix, &self.payload]
}
}
#[derive(Default)]
struct OutQueue {
frames: VecDeque<QueuedFrame>,
off: usize,
bytes: usize,
}
impl OutQueue {
fn push(&mut self, prefix: [u8; gnitz_wire::FRAME_LEN_PREFIX_BYTES], payload: Vec<u8>) {
self.bytes += gnitz_wire::FRAME_LEN_PREFIX_BYTES + payload.len();
self.frames.push_back(QueuedFrame { prefix, payload });
}
fn is_empty(&self) -> bool {
self.frames.is_empty()
}
fn len(&self) -> usize {
self.frames.len()
}
fn build_slices<'a>(&'a self, out: &mut Vec<IoSlice<'a>>) {
let mut skip = self.off;
for f in &self.frames {
for s in f.segments() {
if skip >= s.len() {
skip -= s.len();
continue;
}
out.push(IoSlice::new(&s[skip..]));
skip = 0;
if out.len() == IOV_MAX_CHUNK {
return;
}
}
}
}
fn advance(&mut self, n: usize) {
debug_assert!(n <= self.bytes);
self.bytes -= n;
self.off += n;
while let Some(f) = self.frames.front() {
let total = gnitz_wire::FRAME_LEN_PREFIX_BYTES + f.payload.len();
if self.off < total {
break;
}
self.off -= total;
self.frames.pop_front();
}
}
}
impl From<FrameLenError> for ProtocolError {
fn from(e: FrameLenError) -> Self {
match e {
FrameLenError::Zero => ProtocolError::DecodeError("zero-length frame".into()),
FrameLenError::Oversize { len } => ProtocolError::DecodeError(format!(
"payload length {len} exceeds maximum {} bytes",
gnitz_wire::MAX_FRAME_PAYLOAD
)),
FrameLenError::Alloc { .. } => std::io::Error::from(std::io::ErrorKind::OutOfMemory).into(),
}
}
}
#[cfg(test)]
#[path = "tests/transport.rs"]
mod tests;
#[cfg(test)]
#[path = "benches/transport.rs"]
mod bench;