use std::collections::VecDeque;
use std::io::{self, Read, Write};
use std::sync::{Arc, Condvar, Mutex, MutexGuard, PoisonError};
use std::time::{Duration, Instant};
use crate::ServerError;
use super::super::process::InboundPending;
const MIN_RING_CAPACITY_BYTES: usize = 1;
fn recover<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(PoisonError::into_inner)
}
#[derive(Debug)]
struct RingState {
bytes: VecDeque<u8>,
writer_gone: bool,
reader_gone: bool,
}
#[derive(Debug)]
struct Ring {
capacity: usize,
state: Mutex<RingState>,
changed: Condvar,
}
impl Ring {
const fn new(capacity: usize) -> Self {
Self {
capacity,
state: Mutex::new(RingState {
bytes: VecDeque::new(),
writer_gone: false,
reader_gone: false,
}),
changed: Condvar::new(),
}
}
fn readable_bytes(&self) -> usize {
recover(&self.state).bytes.len()
}
fn read_would_answer(&self) -> bool {
let state = recover(&self.state);
!state.bytes.is_empty() || state.writer_gone
}
fn write(&self, buf: &[u8]) -> io::Result<(usize, bool)> {
let mut state = recover(&self.state);
if state.reader_gone {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"loopback peer end was dropped",
));
}
if buf.is_empty() {
return Ok((0, false));
}
let free = self.capacity.saturating_sub(state.bytes.len());
if free == 0 {
return Err(io::Error::new(
io::ErrorKind::WouldBlock,
"loopback ring is full",
));
}
let accepted = free.min(buf.len());
let was_empty = state.bytes.is_empty();
state.bytes.extend(buf.get(..accepted).unwrap_or(buf));
drop(state);
self.changed.notify_all();
Ok((accepted, was_empty))
}
fn write_until(&self, buf: &[u8], timeout: Option<Duration>) -> io::Result<(usize, bool)> {
if buf.is_empty() {
return Ok((0, false));
}
let deadline = timeout.and_then(|window| Instant::now().checked_add(window));
let mut state = recover(&self.state);
loop {
if state.reader_gone {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"loopback peer end was dropped",
));
}
let free = self.capacity.saturating_sub(state.bytes.len());
if free > 0 {
let accepted = free.min(buf.len());
let was_empty = state.bytes.is_empty();
state.bytes.extend(buf.get(..accepted).unwrap_or(buf));
drop(state);
self.changed.notify_all();
return Ok((accepted, was_empty));
}
let Some(deadline) = deadline else {
state = self
.changed
.wait(state)
.unwrap_or_else(PoisonError::into_inner);
continue;
};
let Some(remaining) = deadline.checked_duration_since(Instant::now()) else {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"loopback write deadline expired",
));
};
let (next, _) = self
.changed
.wait_timeout(state, remaining)
.unwrap_or_else(PoisonError::into_inner);
state = next;
}
}
fn read(&self, buf: &mut [u8]) -> io::Result<usize> {
let (taken, writer_gone) = {
let mut state = recover(&self.state);
(take_from(&mut state.bytes, buf), state.writer_gone)
};
if taken > 0 {
self.changed.notify_all();
return Ok(taken);
}
if buf.is_empty() || writer_gone {
return Ok(taken);
}
Err(io::Error::new(
io::ErrorKind::WouldBlock,
"loopback ring is empty",
))
}
fn read_until(&self, buf: &mut [u8], timeout: Option<Duration>) -> io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let deadline = timeout.and_then(|window| Instant::now().checked_add(window));
let mut state = recover(&self.state);
loop {
let taken = take_from(&mut state.bytes, buf);
if taken > 0 {
drop(state);
self.changed.notify_all();
return Ok(taken);
}
if state.writer_gone {
return Ok(0);
}
let Some(deadline) = deadline else {
state = self
.changed
.wait(state)
.unwrap_or_else(PoisonError::into_inner);
continue;
};
let Some(remaining) = deadline.checked_duration_since(Instant::now()) else {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"loopback read deadline expired",
));
};
let (next, _) = self
.changed
.wait_timeout(state, remaining)
.unwrap_or_else(PoisonError::into_inner);
state = next;
}
}
fn close_writer(&self) {
recover(&self.state).writer_gone = true;
self.changed.notify_all();
}
fn close_reader(&self) {
recover(&self.state).reader_gone = true;
self.changed.notify_all();
}
}
fn take_from(bytes: &mut VecDeque<u8>, buf: &mut [u8]) -> usize {
let taken = bytes.len().min(buf.len());
for (slot, byte) in buf.iter_mut().zip(bytes.drain(..taken)) {
*slot = byte;
}
taken
}
#[derive(Default)]
struct WakerSlot {
callback: Mutex<Option<Arc<dyn Fn() + Send + Sync>>>,
}
impl WakerSlot {
fn set(&self, callback: Arc<dyn Fn() + Send + Sync>) {
*recover(&self.callback) = Some(callback);
}
fn fire(&self) {
let callback = recover(&self.callback).as_ref().map(Arc::clone);
if let Some(callback) = callback {
callback();
}
}
}
impl std::fmt::Debug for WakerSlot {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("WakerSlot")
.field("registered", &recover(&self.callback).is_some())
.finish()
}
}
#[derive(Debug)]
pub enum LoopbackDuplex {}
impl LoopbackDuplex {
#[must_use]
pub fn bounded(capacity_per_ring: usize) -> (LoopbackClientEnd, LoopbackServerEnd) {
let capacity = capacity_per_ring.max(MIN_RING_CAPACITY_BYTES);
let to_server = Arc::new(Ring::new(capacity));
let to_client = Arc::new(Ring::new(capacity));
let waker = Arc::new(WakerSlot::default());
let client = LoopbackClientEnd {
to_server: Arc::clone(&to_server),
from_server: Arc::clone(&to_client),
server_waker: Arc::clone(&waker),
};
let server = LoopbackServerEnd {
from_client: to_server,
to_client,
waker,
};
(client, server)
}
}
#[derive(Debug)]
pub struct LoopbackClientEnd {
to_server: Arc<Ring>,
from_server: Arc<Ring>,
server_waker: Arc<WakerSlot>,
}
impl LoopbackClientEnd {
#[must_use]
pub fn readable_bytes(&self) -> usize {
self.from_server.readable_bytes()
}
pub fn read_timeout(&mut self, buf: &mut [u8], timeout: Option<Duration>) -> io::Result<usize> {
self.from_server.read_until(buf, timeout)
}
pub fn write_timeout(&mut self, buf: &[u8], timeout: Option<Duration>) -> io::Result<usize> {
let (written, opened_the_ring) = self.to_server.write_until(buf, timeout)?;
if opened_the_ring {
self.server_waker.fire();
}
Ok(written)
}
}
impl Read for LoopbackClientEnd {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
self.from_server.read_until(buf, None)
}
}
impl Write for LoopbackClientEnd {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let (written, opened_the_ring) = self.to_server.write(buf)?;
if opened_the_ring {
self.server_waker.fire();
}
Ok(written)
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl Drop for LoopbackClientEnd {
fn drop(&mut self) {
self.to_server.close_writer();
self.from_server.close_reader();
self.server_waker.fire();
}
}
#[derive(Debug)]
pub struct LoopbackServerEnd {
from_client: Arc<Ring>,
to_client: Arc<Ring>,
waker: Arc<WakerSlot>,
}
impl LoopbackServerEnd {
pub fn set_waker(&self, waker: Box<dyn Fn() + Send + Sync>) {
self.waker.set(Arc::from(waker));
}
#[must_use]
pub fn readable_bytes(&self) -> usize {
self.from_client.readable_bytes()
}
}
impl Read for LoopbackServerEnd {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
self.from_client.read(buf)
}
}
impl Write for LoopbackServerEnd {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let (written, _) = self.to_client.write(buf)?;
Ok(written)
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl InboundPending for LoopbackServerEnd {
fn inbound_pending(&self) -> Result<bool, ServerError> {
Ok(self.from_client.read_would_answer())
}
}
impl Drop for LoopbackServerEnd {
fn drop(&mut self) {
self.to_client.close_writer();
self.from_client.close_reader();
}
}