use bytes::{Bytes, BytesMut};
use flume::{Receiver, Sender};
use monocoque_core::options::SocketOptions;
use monocoque_core::rt::{OwnedReadHalf, OwnedWriteHalf, TcpListener};
use std::collections::HashMap;
use std::io;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use tracing::{debug, trace, warn};
type RoutingId = Bytes;
type InboundMsg = Vec<Bytes>;
async fn peer_reader(
routing_id: RoutingId,
mut reader: OwnedReadHalf,
inbound: Sender<InboundMsg>,
read_buffer_size: usize,
shutdown: Receiver<()>,
) {
use compio_buf::BufResult;
use compio_io::AsyncRead;
use futures::{FutureExt, select_biased};
use monocoque_core::io::take_read_buffer;
let _ = inbound
.send_async(vec![routing_id.clone(), Bytes::new(), Bytes::new()])
.await;
let mut read_buf = BytesMut::new();
loop {
let buf = unsafe { take_read_buffer(&mut read_buf, read_buffer_size) };
let BufResult(result, mut buf) = select_biased! {
_ = shutdown.recv_async().fuse() => {
debug!("[STREAM] Peer {:?} reader cancelled", routing_id);
break;
}
res = reader.read(buf).fuse() => res,
};
match result {
Ok(0) => {
debug!("[STREAM] Peer {:?} disconnected (EOF)", routing_id);
break;
}
Ok(n) => {
debug_assert!(n <= read_buffer_size);
buf.truncate(n);
let data = buf.freeze();
trace!("[STREAM] Received {} bytes from peer {:?}", n, routing_id);
let msg = vec![routing_id.clone(), Bytes::new(), data];
if inbound.send_async(msg).await.is_err() {
break; }
}
Err(e) => {
debug!("[STREAM] Peer {:?} read error: {}", routing_id, e);
break;
}
}
}
let _ = inbound.try_send(vec![routing_id, Bytes::new(), Bytes::new()]);
}
struct PeerHandle {
out_tx: Sender<Bytes>,
_shutdown: Sender<()>,
}
async fn peer_writer(mut writer: OwnedWriteHalf, outbound: Receiver<Bytes>) {
use compio_buf::BufResult;
use compio_io::AsyncWriteExt;
while let Ok(data) = outbound.recv_async().await {
let BufResult(res, _) = writer.write_all(data).await;
if res.is_err() {
break;
}
}
}
pub struct StreamSocket {
listener: TcpListener,
inbound_rx: Receiver<InboundMsg>,
inbound_tx: Sender<InboundMsg>,
peers: HashMap<RoutingId, PeerHandle>,
next_id: Arc<AtomicU64>,
options: SocketOptions,
}
impl StreamSocket {
pub async fn bind(addr: impl monocoque_core::rt::ToSocketAddrs) -> io::Result<Self> {
let listener = TcpListener::bind(addr).await?;
debug!("[STREAM] Bound to {}", listener.local_addr()?);
let options = SocketOptions::default();
let (tx, rx) = if options.recv_hwm == 0 {
flume::unbounded()
} else {
flume::bounded(options.recv_hwm)
};
Ok(Self {
listener,
inbound_rx: rx,
inbound_tx: tx,
peers: HashMap::new(),
next_id: Arc::new(AtomicU64::new(1)),
options,
})
}
pub async fn accept_raw(&mut self) -> io::Result<RoutingId> {
let (stream, addr) = match self.listener.accept().await {
Ok(pair) => pair,
Err(e) => {
crate::utils::backoff_on_fd_exhaustion(&e).await;
return Err(e);
}
};
crate::utils::configure_tcp_stream(&stream, &self.options, "STREAM")?;
debug!("[STREAM] Accepted raw connection from {}", addr);
let id_u64 = self.next_id.fetch_add(1, Ordering::Relaxed);
let routing_id = Bytes::copy_from_slice(&id_u64.to_be_bytes());
let (read_half, write_half) = stream.into_split();
let (out_tx, out_rx) = if self.options.send_hwm == 0 {
flume::unbounded::<Bytes>()
} else {
flume::bounded::<Bytes>(self.options.send_hwm)
};
let (shutdown_tx, shutdown_rx) = flume::bounded::<()>(1);
self.peers.insert(
routing_id.clone(),
PeerHandle {
out_tx,
_shutdown: shutdown_tx,
},
);
let inbound = self.inbound_tx.clone();
let rid = routing_id.clone();
monocoque_core::rt::spawn_detached(peer_reader(
rid,
read_half,
inbound,
self.options.read_buffer_size,
shutdown_rx,
));
monocoque_core::rt::spawn_detached(peer_writer(write_half, out_rx));
debug!("[STREAM] Peer {:?} registered", routing_id);
Ok(routing_id)
}
pub async fn recv(&mut self) -> io::Result<Option<InboundMsg>> {
match self.inbound_rx.recv_async().await {
Ok(msg) => {
trace!("[STREAM] Dequeued message from peer {:?}", msg[0]);
Ok(Some(msg))
}
Err(_) => Ok(None),
}
}
pub async fn send(&mut self, msg: Vec<Bytes>) -> io::Result<()> {
if msg.is_empty() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"STREAM send requires at least a routing-id frame",
));
}
let routing_id = msg[0].clone();
let data: Bytes = msg
.iter()
.skip(1)
.find(|f| !f.is_empty())
.cloned()
.unwrap_or_default();
if data.is_empty() {
self.disconnect(&routing_id);
return Ok(());
}
match self.peers.get(&routing_id) {
Some(peer) => {
peer.out_tx.try_send(data).map_err(|e| match e {
flume::TrySendError::Full(_) => io::Error::new(
io::ErrorKind::WouldBlock,
format!("Peer {:?} send queue reached send_hwm", routing_id),
),
flume::TrySendError::Disconnected(_) => io::Error::new(
io::ErrorKind::BrokenPipe,
format!("Peer {:?} send channel disconnected", routing_id),
),
})?;
trace!("[STREAM] Queued data for peer {:?}", routing_id);
}
None => {
warn!(
"[STREAM] Unknown routing-id {:?}, dropping message",
routing_id
);
}
}
Ok(())
}
pub fn close_peer(&mut self, routing_id: &Bytes) -> bool {
if self.peers.remove(routing_id).is_some() {
debug!("[STREAM] Peer {:?} closed and removed", routing_id);
true
} else {
false
}
}
pub fn disconnect(&mut self, routing_id: &Bytes) {
self.close_peer(routing_id);
}
#[inline]
pub fn peer_count(&self) -> usize {
self.peers.len()
}
pub fn local_addr(&self) -> io::Result<std::net::SocketAddr> {
self.listener.local_addr()
}
#[inline]
pub const fn options(&self) -> &SocketOptions {
&self.options
}
#[inline]
pub fn options_mut(&mut self) -> &mut SocketOptions {
&mut self.options
}
}