use crate::{Error, Framer, encode};
use std::io::{self, Read, Write};
pub trait Transport {
fn send(&mut self, payload: &[u8]) -> io::Result<()>;
fn receive(&mut self) -> io::Result<Option<Vec<u8>>>;
fn send_str(&mut self, message: &str) -> io::Result<()> {
self.send(message.as_bytes())
}
}
#[derive(Debug)]
pub struct IoTransport<S> {
stream: S,
framer: Framer,
chunk: Vec<u8>,
}
const CHUNK: usize = 8 * 1024;
impl<S> IoTransport<S> {
pub fn new(stream: S) -> IoTransport<S> {
IoTransport::with_framer(stream, Framer::new())
}
pub fn with_framer(stream: S, framer: Framer) -> IoTransport<S> {
IoTransport {
stream,
framer,
chunk: vec![0; CHUNK],
}
}
pub fn stream(&self) -> &S {
&self.stream
}
pub fn stream_mut(&mut self) -> &mut S {
&mut self.stream
}
pub fn into_stream(self) -> S {
self.stream
}
pub fn framer(&self) -> &Framer {
&self.framer
}
}
impl<S: Read + Write> Transport for IoTransport<S> {
fn send(&mut self, payload: &[u8]) -> io::Result<()> {
self.stream.write_all(&encode(payload))?;
self.stream.flush()
}
fn receive(&mut self) -> io::Result<Option<Vec<u8>>> {
loop {
if let Some(frame) = self.framer.next_frame()? {
return Ok(Some(frame));
}
let read = self.stream.read(&mut self.chunk)?;
if read == 0 {
return if self.framer.is_empty() {
Ok(None)
} else {
Err(io::Error::from(Error::Incomplete))
};
}
self.framer.push(&self.chunk[..read]);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug)]
struct Pipe {
incoming: Vec<u8>,
position: usize,
bite: usize,
outgoing: Vec<u8>,
}
impl Pipe {
fn new(incoming: &[u8], bite: usize) -> Pipe {
Pipe {
incoming: incoming.to_vec(),
position: 0,
bite,
outgoing: Vec::new(),
}
}
}
impl Read for Pipe {
fn read(&mut self, out: &mut [u8]) -> io::Result<usize> {
let remaining = self.incoming.len() - self.position;
let count = remaining.min(self.bite).min(out.len());
out[..count].copy_from_slice(&self.incoming[self.position..self.position + count]);
self.position += count;
Ok(count)
}
}
impl Write for Pipe {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
self.outgoing.extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[test]
fn receives_messages_however_the_stream_chops_them() {
let wire = b"\x0bMSH|one\x1c\r\x0bMSH|two\x1c\r";
for bite in [1, 2, 7, 1024] {
let mut transport = IoTransport::new(Pipe::new(wire, bite));
assert_eq!(
transport.receive().unwrap().unwrap(),
b"MSH|one",
"bite {bite}"
);
assert_eq!(
transport.receive().unwrap().unwrap(),
b"MSH|two",
"bite {bite}"
);
assert_eq!(transport.receive().unwrap(), None, "bite {bite}");
}
}
#[test]
fn sends_one_framed_message() {
let mut transport = IoTransport::new(Pipe::new(b"", 64));
transport.send_str("MSH|^~\\&|LAB").unwrap();
assert_eq!(transport.stream().outgoing, b"\x0bMSH|^~\\&|LAB\x1c\r");
}
#[test]
fn a_peer_that_hangs_up_mid_message_is_an_error_not_a_message() {
let mut transport = IoTransport::new(Pipe::new(b"\x0bMSH|trunc", 4));
let error = transport.receive().unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(error.to_string().contains("incomplete"), "{error}");
}
#[test]
fn a_clean_close_between_frames_is_the_end_of_the_stream() {
let mut transport = IoTransport::new(Pipe::new(b"\x0bonly\x1c\r", 3));
assert_eq!(transport.receive().unwrap().unwrap(), b"only");
assert_eq!(transport.receive().unwrap(), None);
assert_eq!(transport.receive().unwrap(), None);
}
#[test]
fn framing_violations_surface_as_invalid_data() {
let mut transport = IoTransport::with_framer(
Pipe::new(b"garbage\x0bMSH|\x1c\r", 64),
Framer::new().with_tolerance(crate::Tolerance::Strict),
);
let error = transport.receive().unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
}