use crate::{BitBuf, BitDecode, BitEncode, BitError, BitReader, BitWriter, ErrorKind};
use alloc::vec;
use alloc::vec::Vec;
use std::io::{self, Read, Write};
use std::net::{SocketAddr, UdpSocket};
#[cfg(feature = "mock")]
use std::{
cell::{Cell, RefCell},
collections::VecDeque,
};
fn encode<T: BitEncode>(msg: &T) -> Result<Vec<u8>, BitError> {
let mut w = BitWriter::with_layout(<T as BitEncode>::LAYOUT);
msg.bit_encode(&mut w)?;
Ok(w.into_bytes())
}
#[derive(Debug)]
pub struct MessageStream<S> {
inner: S,
buf: BitBuf,
}
impl<S> MessageStream<S> {
pub fn new(inner: S) -> Self {
Self {
inner,
buf: BitBuf::new(),
}
}
pub fn with_cap(inner: S, cap: usize) -> Self {
Self {
inner,
buf: BitBuf::bounded(cap),
}
}
pub fn get_mut(&mut self) -> &mut S {
&mut self.inner
}
pub fn into_inner(self) -> S {
self.inner
}
}
impl<S: Read> MessageStream<S> {
pub fn read_message<T: BitDecode + BitEncode>(&mut self) -> Result<T, BitError> {
loop {
if let Some(msg) = self.buf.pull::<T>()? {
return Ok(msg);
}
let mut chunk = [0u8; 4096];
let n = self.inner.read(&mut chunk)?;
if n == 0 {
return Err(
io::Error::new(io::ErrorKind::UnexpectedEof, "connection closed").into(),
);
}
self.buf.try_push(&chunk[..n]).map_err(|e| {
BitError::new(ErrorKind::BufferFull { cap: e.cap }, self.buf.bit_len())
})?;
}
}
}
impl<S: Write> MessageStream<S> {
pub fn write_message<T: BitEncode>(&mut self, msg: &T) -> Result<(), BitError> {
self.inner.write_all(&encode(msg)?)?;
Ok(())
}
}
pub trait DatagramSocket: sealed::Sealed {
type Addr;
fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, Self::Addr)>;
fn send_to(&self, buf: &[u8], addr: &Self::Addr) -> io::Result<usize>;
}
mod sealed {
pub trait Sealed {}
}
impl sealed::Sealed for UdpSocket {}
impl DatagramSocket for UdpSocket {
type Addr = SocketAddr;
fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
UdpSocket::recv_from(self, buf)
}
fn send_to(&self, buf: &[u8], addr: &SocketAddr) -> io::Result<usize> {
UdpSocket::send_to(self, buf, *addr)
}
}
#[cfg(unix)]
impl sealed::Sealed for std::os::unix::net::UnixDatagram {}
#[cfg(unix)]
impl DatagramSocket for std::os::unix::net::UnixDatagram {
type Addr = std::os::unix::net::SocketAddr;
fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, Self::Addr)> {
std::os::unix::net::UnixDatagram::recv_from(self, buf)
}
fn send_to(&self, buf: &[u8], addr: &Self::Addr) -> io::Result<usize> {
self.send_to_addr(buf, addr)
}
}
#[derive(Debug)]
pub struct MessageDatagram<D> {
sock: D,
buf: Vec<u8>,
}
impl<D> MessageDatagram<D> {
pub fn new(sock: D) -> Self {
Self::with_capacity(sock, 65_536)
}
pub fn with_capacity(sock: D, capacity: usize) -> Self {
Self {
sock,
buf: vec![0u8; capacity],
}
}
pub fn get_ref(&self) -> &D {
&self.sock
}
pub fn get_mut(&mut self) -> &mut D {
&mut self.sock
}
pub fn into_inner(self) -> D {
self.sock
}
}
impl<D: DatagramSocket> MessageDatagram<D> {
pub fn send_message<T: BitEncode>(&self, msg: &T, addr: &D::Addr) -> Result<usize, BitError> {
Ok(self.sock.send_to(&encode(msg)?, addr)?)
}
pub fn recv_message<T: BitDecode + BitEncode>(&mut self) -> Result<(T, D::Addr), BitError> {
let (n, from) = self.sock.recv_from(&mut self.buf)?;
let mut r = BitReader::with_layout(&self.buf[..n], <T as BitEncode>::LAYOUT);
let msg = <T as BitDecode>::bit_decode(&mut r)?;
Ok((msg, from))
}
}
#[cfg(feature = "mock")]
#[derive(Debug, Default)]
pub struct MockDatagramSocket {
inbound: RefCell<VecDeque<(Vec<u8>, SocketAddr)>>,
sent: RefCell<Vec<(Vec<u8>, SocketAddr)>>,
fail_recv: Cell<bool>,
}
#[cfg(feature = "mock")]
impl MockDatagramSocket {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn push_inbound(&self, bytes: &[u8], from: SocketAddr) {
self.inbound.borrow_mut().push_back((bytes.to_vec(), from));
}
#[must_use]
pub fn sent(&self) -> Vec<(Vec<u8>, SocketAddr)> {
self.sent.borrow().clone()
}
#[must_use]
pub fn fail_next_recv(self) -> Self {
self.fail_recv.set(true);
self
}
}
#[cfg(feature = "mock")]
impl sealed::Sealed for MockDatagramSocket {}
#[cfg(feature = "mock")]
impl DatagramSocket for MockDatagramSocket {
type Addr = SocketAddr;
fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
if self.fail_recv.replace(false) {
return Err(io::Error::new(
io::ErrorKind::ConnectionReset,
"mock: recv failed",
));
}
let (data, from) = self
.inbound
.borrow_mut()
.pop_front()
.ok_or_else(|| io::Error::new(io::ErrorKind::WouldBlock, "no queued datagram"))?;
let n = data.len().min(buf.len());
buf[..n].copy_from_slice(&data[..n]);
Ok((n, from))
}
fn send_to(&self, buf: &[u8], addr: &SocketAddr) -> io::Result<usize> {
self.sent.borrow_mut().push((buf.to_vec(), *addr));
Ok(buf.len())
}
}
#[cfg(feature = "mock")]
#[derive(Debug, Default, Clone)]
pub struct MockStream {
inbound: VecDeque<u8>,
outbound: Vec<u8>,
chunk: usize,
fail_after: Option<usize>,
read_total: usize,
}
#[cfg(feature = "mock")]
impl MockStream {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_chunk_size(n: usize) -> Self {
Self {
chunk: n,
..Self::default()
}
}
pub fn push_inbound(&mut self, bytes: &[u8]) {
self.inbound.extend(bytes.iter().copied());
}
#[must_use]
pub fn written(&self) -> &[u8] {
&self.outbound
}
#[must_use]
pub fn fail_after(mut self, n: usize) -> Self {
self.fail_after = Some(n);
self
}
}
#[cfg(feature = "mock")]
impl Read for MockStream {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
if let Some(at) = self.fail_after {
if self.read_total >= at {
return Err(io::Error::new(
io::ErrorKind::ConnectionReset,
"mock: connection reset",
));
}
}
if self.inbound.is_empty() || buf.is_empty() {
return Ok(0); }
let mut cap = if self.chunk == 0 {
buf.len()
} else {
buf.len().min(self.chunk)
};
cap = cap.min(self.inbound.len());
if let Some(at) = self.fail_after {
cap = cap.min(at - self.read_total); }
for slot in buf.iter_mut().take(cap) {
*slot = self.inbound.pop_front().unwrap();
}
self.read_total += cap;
Ok(cap)
}
}
#[cfg(feature = "mock")]
impl Write for MockStream {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.outbound.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[cfg(all(test, feature = "mock"))]
mod component {
use bnb::{MessageDatagram, MessageStream, MockDatagramSocket, MockStream, bin};
#[bin(big)]
#[derive(Debug, Clone, PartialEq, Eq)]
struct Msg {
seq: u16,
}
#[test]
fn stream_write_message_is_captured() {
let mut conn = MessageStream::new(MockStream::new());
conn.write_message(&Msg { seq: 7 }).unwrap();
assert_eq!(
conn.get_mut().written(),
&Msg { seq: 7 }.to_bytes().unwrap()[..]
);
}
#[test]
fn stream_reads_a_queued_message() {
let mut conn = MessageStream::new(MockStream::new());
conn.get_mut()
.push_inbound(&Msg { seq: 0xABCD }.to_bytes().unwrap());
assert_eq!(conn.read_message::<Msg>().unwrap(), Msg { seq: 0xABCD });
}
#[test]
fn stream_reassembles_a_message_split_across_reads() {
let mut conn = MessageStream::new(MockStream::with_chunk_size(1));
conn.get_mut()
.push_inbound(&Msg { seq: 0x1234 }.to_bytes().unwrap());
assert_eq!(conn.read_message::<Msg>().unwrap(), Msg { seq: 0x1234 });
}
#[test]
fn stream_eof_mid_message_is_an_error() {
let mut conn = MessageStream::new(MockStream::new());
conn.get_mut().push_inbound(&[0x12]);
assert!(conn.read_message::<Msg>().is_err());
}
#[test]
fn stream_connection_reset_surfaces_as_an_error() {
let mut conn = MessageStream::new(MockStream::new().fail_after(0));
conn.get_mut()
.push_inbound(&Msg { seq: 1 }.to_bytes().unwrap());
assert!(conn.read_message::<Msg>().is_err());
}
#[test]
fn stream_into_inner_recovers_the_transport() {
let conn = MessageStream::new(MockStream::new());
let _inner: MockStream = conn.into_inner();
}
#[test]
fn datagram_recv_then_send_to_the_sender() {
let mut peer = MessageDatagram::new(MockDatagramSocket::new());
let from = "127.0.0.1:5000".parse().unwrap();
peer.get_ref()
.push_inbound(&Msg { seq: 7 }.to_bytes().unwrap(), from);
let (msg, who) = peer.recv_message::<Msg>().unwrap();
assert_eq!(msg, Msg { seq: 7 });
assert_eq!(who, from);
let n = peer.send_message(&Msg { seq: 8 }, &who).unwrap();
assert_eq!(n, 2, "send_message returns the byte count");
assert_eq!(
peer.get_ref().sent()[0].0,
Msg { seq: 8 }.to_bytes().unwrap()
);
assert_eq!(
peer.get_ref().sent()[0].1,
who,
"sent to the original sender"
);
}
#[test]
fn datagram_recv_error_is_injected() {
let mut peer = MessageDatagram::new(MockDatagramSocket::new().fail_next_recv());
assert!(peer.recv_message::<Msg>().is_err());
}
#[test]
fn datagram_recv_malformed_is_a_codec_error() {
#[bin(big, magic = 0xCAFEu16)]
#[derive(Debug, PartialEq, Eq)]
struct M {
v: u8,
}
let mut peer = MessageDatagram::new(MockDatagramSocket::new());
let from = "127.0.0.1:1".parse().unwrap();
peer.get_ref().push_inbound(&[0x00, 0x00, 0x09], from); assert!(peer.recv_message::<M>().is_err());
}
#[test]
fn datagram_with_capacity_truncates_an_oversized_datagram() {
let mut peer = MessageDatagram::with_capacity(MockDatagramSocket::new(), 2);
let from = "127.0.0.1:2".parse().unwrap();
peer.get_ref().push_inbound(&[0x00, 0x05, 0xFF, 0xFF], from);
let (msg, _) = peer.recv_message::<Msg>().unwrap();
assert_eq!(msg, Msg { seq: 5 });
}
#[test]
fn datagram_get_mut_and_into_inner() {
let mut peer = MessageDatagram::new(MockDatagramSocket::new());
let _m: &mut MockDatagramSocket = peer.get_mut();
let _inner: MockDatagramSocket = peer.into_inner();
}
}