mod bounds;
use std::future::Future;
pub use crate::bounds::{MaybeSend, MaybeSync};
use bytes::{Buf, BufMut, Bytes, BytesMut};
pub trait Error: std::error::Error + MaybeSend + MaybeSync + 'static {
fn session_error(&self) -> Option<(u32, String)>;
fn stream_error(&self) -> Option<u32> {
None
}
}
pub trait Session: Clone + MaybeSend + MaybeSync + 'static {
type SendStream: SendStream;
type RecvStream: RecvStream;
type Error: Error;
fn accept_uni(&self)
-> impl Future<Output = Result<Self::RecvStream, Self::Error>> + MaybeSend;
fn accept_bi(
&self,
) -> impl Future<Output = Result<(Self::SendStream, Self::RecvStream), Self::Error>> + MaybeSend;
fn open_bi(
&self,
) -> impl Future<Output = Result<(Self::SendStream, Self::RecvStream), Self::Error>> + MaybeSend;
fn open_uni(&self) -> impl Future<Output = Result<Self::SendStream, Self::Error>> + MaybeSend;
fn send_datagram(
&self,
payload: Bytes,
) -> impl Future<Output = Result<(), Self::Error>> + MaybeSend;
fn recv_datagram(&self) -> impl Future<Output = Result<Bytes, Self::Error>> + MaybeSend;
fn max_datagram_size(&self) -> usize;
fn close(&self, code: u32, reason: &str);
fn closed(&self) -> impl Future<Output = Self::Error> + MaybeSend;
}
pub trait SendStream: MaybeSend {
type Error: Error;
fn write(&mut self, buf: &[u8])
-> impl Future<Output = Result<usize, Self::Error>> + MaybeSend;
fn write_buf<B: Buf + MaybeSend>(
&mut self,
buf: &mut B,
) -> impl Future<Output = Result<usize, Self::Error>> + MaybeSend {
async move {
let chunk = buf.chunk();
let size = self.write(chunk).await?;
assert!(
size > 0 || chunk.is_empty(),
"SendStream::write returned zero for a non-empty buffer"
);
assert!(
size <= chunk.len(),
"SendStream::write returned more bytes than provided"
);
buf.advance(size);
Ok(size)
}
}
fn write_chunk(
&mut self,
chunk: Bytes,
) -> impl Future<Output = Result<(), Self::Error>> + MaybeSend {
async move {
let mut c = chunk;
self.write_all_buf(&mut c).await
}
}
fn write_all(
&mut self,
buf: &[u8],
) -> impl Future<Output = Result<(), Self::Error>> + MaybeSend {
async move {
let mut pos = 0;
while pos < buf.len() {
let written = self.write(&buf[pos..]).await?;
assert!(
written > 0,
"SendStream::write returned zero for a non-empty buffer"
);
assert!(
written <= buf.len() - pos,
"SendStream::write returned more bytes than provided"
);
pos += written;
}
Ok(())
}
}
fn write_all_buf<B: Buf + MaybeSend>(
&mut self,
buf: &mut B,
) -> impl Future<Output = Result<(), Self::Error>> + MaybeSend {
async move {
while buf.has_remaining() {
let written = self.write_buf(buf).await?;
assert!(
written > 0,
"SendStream::write returned zero for a non-empty buffer"
);
}
Ok(())
}
}
fn set_priority(&mut self, order: u8);
fn finish(&mut self) -> Result<(), Self::Error>;
fn reset(&mut self, code: u32);
fn closed(&mut self) -> impl Future<Output = Result<(), Self::Error>> + MaybeSend;
}
pub trait RecvStream: MaybeSend {
type Error: Error;
fn read(
&mut self,
dst: &mut [u8],
) -> impl Future<Output = Result<Option<usize>, Self::Error>> + MaybeSend;
fn read_buf<B: BufMut + MaybeSend>(
&mut self,
buf: &mut B,
) -> impl Future<Output = Result<Option<usize>, Self::Error>> + MaybeSend {
async move {
let capacity = buf.chunk_mut().len().min(8 * 1024);
if capacity == 0 {
return Ok(Some(0));
}
let mut dst = vec![0; capacity];
let size = match self.read(&mut dst).await? {
Some(size) => size,
None => return Ok(None),
};
assert!(
size <= dst.len(),
"RecvStream::read returned more bytes than the provided buffer"
);
buf.put_slice(&dst[..size]);
Ok(Some(size))
}
}
fn read_chunk(
&mut self,
max: usize,
) -> impl Future<Output = Result<Option<Bytes>, Self::Error>> + MaybeSend {
async move {
let mut buf = BytesMut::with_capacity(max.min(8 * 1024));
Ok(self.read_buf(&mut buf).await?.map(|_| buf.freeze()))
}
}
fn stop(&mut self, code: u32);
fn closed(&mut self) -> impl Future<Output = Result<(), Self::Error>> + MaybeSend;
fn read_all(&mut self) -> impl Future<Output = Result<Bytes, Self::Error>> + MaybeSend {
async move {
let mut buf = BytesMut::new();
self.read_all_buf(&mut buf).await?;
Ok(buf.freeze())
}
}
fn read_all_buf<B: BufMut + MaybeSend>(
&mut self,
buf: &mut B,
) -> impl Future<Output = Result<usize, Self::Error>> + MaybeSend {
async move {
let mut size = 0;
while buf.has_remaining_mut() {
match self.read_buf(buf).await? {
Some(n) => size += n,
None => break,
}
}
Ok(size)
}
}
}
#[cfg(test)]
mod tests {
use super::{Error, RecvStream, SendStream};
use bytes::{Bytes, BytesMut};
use futures::executor::block_on;
use std::fmt;
#[derive(Debug)]
struct TestError;
impl fmt::Display for TestError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "test error")
}
}
impl std::error::Error for TestError {}
impl Error for TestError {
fn session_error(&self) -> Option<(u32, String)> {
None
}
}
struct TestRecvStream {
data: Vec<u8>,
pos: usize,
}
impl TestRecvStream {
fn new(data: &[u8]) -> Self {
Self {
data: data.to_vec(),
pos: 0,
}
}
}
impl RecvStream for TestRecvStream {
type Error = TestError;
async fn read(&mut self, dst: &mut [u8]) -> Result<Option<usize>, Self::Error> {
let available = self.data.len().saturating_sub(self.pos);
if available == 0 {
return Ok(None);
}
let size = available.min(dst.len());
let end = self.pos + size;
dst[..size].copy_from_slice(&self.data[self.pos..end]);
self.pos = end;
Ok(Some(size))
}
fn stop(&mut self, _code: u32) {}
async fn closed(&mut self) -> Result<(), Self::Error> {
Ok(())
}
}
struct PartialSendStream {
data: Vec<u8>,
max_write: usize,
}
impl SendStream for PartialSendStream {
type Error = TestError;
async fn write(&mut self, buf: &[u8]) -> Result<usize, Self::Error> {
let size = buf.len().min(self.max_write);
self.data.extend_from_slice(&buf[..size]);
Ok(size)
}
fn set_priority(&mut self, _order: u8) {}
fn finish(&mut self) -> Result<(), Self::Error> {
Ok(())
}
fn reset(&mut self, _code: u32) {}
async fn closed(&mut self) -> Result<(), Self::Error> {
Ok(())
}
}
#[test]
fn read_chunk_respects_max_and_eof() {
let mut stream = TestRecvStream::new(b"hello world");
let first = block_on(stream.read_chunk(5)).unwrap().unwrap();
assert_eq!(first, Bytes::from_static(b"hello"));
let second = block_on(stream.read_chunk(1024)).unwrap().unwrap();
assert_eq!(second, Bytes::from_static(b" world"));
let end = block_on(stream.read_chunk(1)).unwrap();
assert!(end.is_none());
}
#[test]
fn read_buf_advances_buffer() {
let mut stream = TestRecvStream::new(b"test");
let mut buf = BytesMut::with_capacity(4);
let size = block_on(stream.read_buf(&mut buf)).unwrap().unwrap();
assert_eq!(size, 4);
assert_eq!(&buf[..], b"test");
}
#[test]
fn write_chunk_retries_partial_writes() {
let mut stream = PartialSendStream {
data: Vec::new(),
max_write: 2,
};
block_on(stream.write_chunk(Bytes::from_static(b"hello"))).unwrap();
assert_eq!(stream.data, b"hello");
}
}