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, Clone, Copy, PartialEq, Eq)]
pub(super) enum DrainOutcome {
Drained,
Progress,
WouldBlockWithResidue,
}
#[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,
budget: Option<usize>,
) -> Result<DrainOutcome, OutboundError> {
let mut written_total: usize = 0;
while !self.buffer.is_empty() {
let remaining_budget = match budget {
Some(limit) if written_total >= limit => return Ok(DrainOutcome::Progress),
Some(limit) => Some(limit - written_total),
None => None,
};
let front = self.buffer.as_slices().0;
let front = remaining_budget.map_or(front, |limit| {
front.get(..limit.min(front.len())).unwrap_or(front)
});
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);
written_total = written_total.saturating_add(written);
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
return Ok(DrainOutcome::WouldBlockWithResidue);
}
Err(error) if error.kind() == std::io::ErrorKind::Interrupted => {}
Err(error) => return Err(OutboundError::Write(error)),
}
}
Ok(DrainOutcome::Drained)
}
#[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;