use std::collections::VecDeque;
use std::net::TcpStream;
use tungstenite::Message;
use tungstenite::protocol::WebSocket;
use liminal::protocol::{Frame, encode, encoded_len};
use super::super::outbound::{DEFAULT_OUTBOUND_CAPACITY, DrainOutcome, OutboundError};
#[cfg(test)]
#[path = "outbound_tests.rs"]
mod tests;
#[derive(Debug)]
pub(in super::super) struct WebSocketOutbound {
queue: VecDeque<Vec<u8>>,
queued_bytes: usize,
capacity: usize,
in_flight: Option<usize>,
transport_flush_pending: bool,
}
impl WebSocketOutbound {
pub(in super::super) const fn new() -> Self {
Self::with_capacity(DEFAULT_OUTBOUND_CAPACITY)
}
pub(in super::super) const fn with_capacity(capacity: usize) -> Self {
Self {
queue: VecDeque::new(),
queued_bytes: 0,
capacity,
in_flight: None,
transport_flush_pending: false,
}
}
pub(in super::super) fn enqueue_frame(&mut self, frame: &Frame) -> Result<(), OutboundError> {
let needed = encoded_len(frame).map_err(OutboundError::Encode)?;
let projected = self
.queued_bytes
.checked_add(needed)
.ok_or(OutboundError::Overflow {
queued: self.queued_bytes,
needed,
capacity: self.capacity,
})?;
if projected > self.capacity {
return Err(OutboundError::Overflow {
queued: self.queued_bytes,
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.queued_bytes = self.queued_bytes.saturating_add(bytes.len());
self.queue.push_back(bytes);
Ok(())
}
pub(in super::super) const fn capacity(&self) -> usize {
self.capacity
}
pub(in super::super) fn has_room(&self, needed: usize) -> bool {
self.queued_bytes
.checked_add(needed)
.is_some_and(|projected| projected <= self.capacity)
}
pub(in super::super) const fn note_transport_write_pending(&mut self) {
self.transport_flush_pending = true;
}
pub(in super::super) fn drain(
&mut self,
socket: &mut WebSocket<TcpStream>,
budget: Option<usize>,
) -> Result<DrainOutcome, OutboundError> {
let mut written_total: usize = 0;
loop {
if self.in_flight.is_some() || self.transport_flush_pending {
match socket.flush() {
Ok(()) => {
if let Some(flushed) = self.in_flight.take() {
self.queued_bytes = self.queued_bytes.saturating_sub(flushed);
written_total = written_total.saturating_add(flushed);
}
self.transport_flush_pending = false;
}
Err(error) => return self.map_drain_error(error),
}
}
if self.queue.is_empty() {
return Ok(DrainOutcome::Drained);
}
if let Some(limit) = budget {
if written_total >= limit {
return Ok(DrainOutcome::Progress);
}
}
let Some(message) = self.queue.pop_front() else {
return Ok(DrainOutcome::Drained);
};
self.in_flight = Some(message.len());
match socket.write(Message::Binary(message.into())) {
Ok(()) => {}
Err(error) => return self.map_drain_error(error),
}
}
}
fn map_drain_error(
&mut self,
error: tungstenite::Error,
) -> Result<DrainOutcome, OutboundError> {
match error {
tungstenite::Error::Io(io_error)
if io_error.kind() == std::io::ErrorKind::WouldBlock
|| io_error.kind() == std::io::ErrorKind::Interrupted =>
{
Ok(DrainOutcome::WouldBlockWithResidue)
}
tungstenite::Error::Io(io_error) => Err(OutboundError::Write(io_error)),
tungstenite::Error::WriteBufferFull(message) => {
if let Message::Binary(bytes) = *message {
self.in_flight = None;
self.queue.push_front(bytes.to_vec());
}
Ok(DrainOutcome::WouldBlockWithResidue)
}
other => Err(OutboundError::Write(std::io::Error::other(
other.to_string(),
))),
}
}
#[cfg(test)]
pub(in super::super) const fn queued_len(&self) -> usize {
self.queued_bytes
}
#[cfg(test)]
pub(in super::super) fn take_messages(&mut self) -> Vec<Vec<u8>> {
let messages: Vec<Vec<u8>> = self.queue.drain(..).collect();
let drained: usize = messages.iter().map(Vec::len).sum();
self.queued_bytes = self.queued_bytes.saturating_sub(drained);
messages
}
}
impl super::super::delivery::DeliverySink for WebSocketOutbound {
fn capacity(&self) -> usize {
Self::capacity(self)
}
fn has_room(&self, needed: usize) -> bool {
Self::has_room(self, needed)
}
fn enqueue_frame(&mut self, frame: &Frame) -> Result<(), OutboundError> {
Self::enqueue_frame(self, frame)
}
}