use std::collections::VecDeque;
use std::io::Write;
use std::net::TcpStream;
use liminal::protocol::{Frame, ProtocolError, encode, encoded_len};
pub(super) const DEFAULT_OUTBOUND_CAPACITY: usize = 4 * 1024 * 1024;
#[derive(Debug)]
pub(super) enum OutboundError {
Overflow {
queued: usize,
needed: usize,
capacity: usize,
},
Encode(ProtocolError),
Write(std::io::Error),
}
impl std::fmt::Display for OutboundError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Overflow {
queued,
needed,
capacity,
} => write!(
formatter,
"outbound buffer overflow: {queued} queued + {needed} needed exceeds \
capacity {capacity}"
),
Self::Encode(error) => write!(formatter, "outbound frame encode failed: {error}"),
Self::Write(error) => write!(formatter, "outbound socket write failed: {error}"),
}
}
}
impl std::error::Error for OutboundError {}
#[derive(Debug)]
pub(super) struct OutboundWriter {
buffer: VecDeque<u8>,
capacity: usize,
}
impl OutboundWriter {
pub(super) const fn new() -> Self {
Self::with_capacity(DEFAULT_OUTBOUND_CAPACITY)
}
pub(super) const fn with_capacity(capacity: usize) -> Self {
Self {
buffer: VecDeque::new(),
capacity,
}
}
pub(super) fn enqueue_frame(&mut self, frame: &Frame) -> Result<(), OutboundError> {
let needed = encoded_len(frame).map_err(OutboundError::Encode)?;
let queued = self.buffer.len();
let projected = queued.checked_add(needed).ok_or(OutboundError::Overflow {
queued,
needed,
capacity: self.capacity,
})?;
if projected > self.capacity {
return Err(OutboundError::Overflow {
queued,
needed,
capacity: self.capacity,
});
}
let mut bytes = vec![0_u8; needed];
let written = encode(frame, &mut bytes).map_err(OutboundError::Encode)?;
bytes.truncate(written);
self.buffer.extend(bytes);
Ok(())
}
pub(super) const fn capacity(&self) -> usize {
self.capacity
}
pub(super) fn has_room(&self, needed: usize) -> bool {
self.buffer
.len()
.checked_add(needed)
.is_some_and(|projected| projected <= self.capacity)
}
pub(super) fn drain(&mut self, stream: &mut TcpStream) -> Result<(), OutboundError> {
while !self.buffer.is_empty() {
let front = self.buffer.as_slices().0;
match stream.write(front) {
Ok(0) => {
return Err(OutboundError::Write(std::io::Error::new(
std::io::ErrorKind::WriteZero,
"connection peer accepted zero bytes",
)));
}
Ok(written) => {
self.buffer.drain(..written);
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => return Ok(()),
Err(error) if error.kind() == std::io::ErrorKind::Interrupted => {}
Err(error) => return Err(OutboundError::Write(error)),
}
}
Ok(())
}
#[cfg(test)]
pub(super) fn queued_len(&self) -> usize {
self.buffer.len()
}
#[cfg(test)]
pub(super) fn take_bytes(&mut self) -> Vec<u8> {
self.buffer.drain(..).collect()
}
}
#[cfg(test)]
#[path = "outbound_tests.rs"]
mod tests;