use crate::net::types::NetError;
use crate::types::PeekBuf;
use bytes::Bytes;
use futures_core::stream::BoxStream;
use futures_core::Stream;
use futures_util::{stream, StreamExt, TryStreamExt};
use parking_lot::Mutex;
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::atomic::AtomicU64;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncReadExt};
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tokio_util::io::StreamReader;
use tokio_util::sync::CancellationToken;
#[derive(Clone)]
pub struct SharedBody {
inner: Arc<Mutex<State>>,
}
struct State {
subs: HashMap<u64, mpsc::Sender<Result<Bytes, NetError>>>,
next_id: AtomicU64,
max_queue: usize,
closed: bool,
}
impl SharedBody {
pub fn new(max_queue: usize) -> Self {
Self {
inner: Arc::new(Mutex::new(State {
subs: HashMap::new(),
next_id: AtomicU64::new(1),
max_queue,
closed: false,
})),
}
}
pub fn push(&self, chunk: Bytes) {
let (subs, mut to_remove) = {
let st = self.inner.lock();
if st.closed {
return;
}
let subs: Vec<(u64, mpsc::Sender<_>)> =
st.subs.iter().map(|(id, tx)| (*id, tx.clone())).collect();
(subs, Vec::new())
};
for (id, tx) in subs {
match tx.try_send(Ok(chunk.clone())) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
to_remove.push(id);
}
Err(mpsc::error::TrySendError::Closed(_)) => {
to_remove.push(id);
}
}
}
if !to_remove.is_empty() {
let mut st = self.inner.lock();
for id in &to_remove {
st.subs.remove(id);
}
}
}
pub fn error(&self, e: NetError) {
let senders: Vec<mpsc::Sender<Result<Bytes, NetError>>> = {
let mut st = self.inner.lock();
if st.closed {
return;
}
st.closed = true;
st.subs.drain().map(|(_, tx)| tx).collect()
};
for tx in senders {
let _ = tx.try_send(Err(e.clone()));
}
}
pub fn finish(&self) {
let _dropped: Vec<mpsc::Sender<Result<Bytes, NetError>>> = {
let mut st = self.inner.lock();
if st.closed {
return;
}
st.closed = true;
st.subs.drain().map(|(_, tx)| tx).collect()
};
}
pub fn subscribe_with_cap(
&self,
max_queue: usize,
) -> BoxStream<'static, Result<Bytes, NetError>> {
let (rx, id) = {
let mut st = self.inner.lock();
if st.closed {
return stream::empty::<Result<Bytes, NetError>>().boxed();
}
let (tx, rx) = mpsc::channel(max_queue);
let id = st
.next_id
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
st.subs.insert(id, tx);
(rx, id)
};
SubStream {
id,
parent: self.inner.clone(),
inner: ReceiverStream::new(rx),
}
.boxed()
}
pub fn subscribe_stream(&self) -> BoxStream<'static, Result<Bytes, NetError>> {
let cap = {
let st = self.inner.lock();
st.max_queue
};
self.subscribe_with_cap(cap)
}
pub fn combined_reader(
peek_buf: PeekBuf,
shared: Arc<SharedBody>,
) -> Pin<Box<dyn AsyncRead + Send>> {
let head = stream::iter([Ok::<Bytes, std::io::Error>(peek_buf.into_bytes())]);
let rest_stream = shared.subscribe_stream().map_err(|e: NetError| e.to_io());
let combined = head.chain(rest_stream);
Box::pin(StreamReader::new(combined))
}
}
pub struct ReaderOptions {
pub capacity: usize,
pub buf_size: usize,
pub cancel: Option<CancellationToken>,
pub idle_timeout: Option<Duration>,
pub total_timeout: Option<Duration>,
pub max_size: Option<u64>,
}
impl Default for ReaderOptions {
fn default() -> Self {
Self {
capacity: 32,
buf_size: 16 * 1024,
cancel: None,
idle_timeout: None,
total_timeout: None,
max_size: None,
}
}
}
impl SharedBody {
pub fn from_reader<R>(mut reader: R, opts: ReaderOptions) -> Arc<Self>
where
R: AsyncRead + Send + 'static + Unpin,
{
let sb = Arc::new(SharedBody::new(opts.capacity));
let sb_clone = sb.clone();
tokio::spawn(async move {
let ReaderOptions {
capacity: _,
buf_size,
cancel,
idle_timeout,
total_timeout,
max_size,
} = opts;
let deadline = total_timeout.map(|d| tokio::time::Instant::now() + d);
let cancel = cancel.unwrap_or_else(CancellationToken::new);
let mut buf = vec![0u8; buf_size];
let mut total_read: u64 = 0;
let check_total_deadline = |now: tokio::time::Instant| -> Result<(), NetError> {
if let Some(dl) = deadline {
if now >= dl {
return Err(NetError::Timeout("total read timeout".to_string()));
}
}
Ok(())
};
if let Err(e) = check_total_deadline(tokio::time::Instant::now()) {
sb_clone.error(e);
return;
}
loop {
if cancel.is_cancelled() {
sb_clone.error(NetError::Cancelled("read cancelled".to_string()));
return;
}
let (read_cap, probing) = if let Some(max) = max_size {
let remaining = max.saturating_sub(total_read);
if remaining == 0 {
(1, true)
} else {
(remaining.min(buf.len() as u64) as usize, false)
}
} else {
(buf.len(), false)
};
let idle_sleep = async {
match idle_timeout {
Some(d) => tokio::time::sleep(d).await,
None => std::future::pending().await,
}
};
let deadline_sleep = async {
match deadline {
Some(dl) => tokio::time::sleep_until(dl).await,
None => std::future::pending().await,
}
};
let read_res = tokio::select! {
r = reader.read(&mut buf[..read_cap]) => r.map_err(|e| NetError::Io(Arc::new(e))),
_ = cancel.cancelled() => Err(NetError::Cancelled("read cancelled".to_string())),
_ = idle_sleep => Err(NetError::Timeout("read idle timeout".to_string())),
_ = deadline_sleep => Err(NetError::Timeout("total read timeout".to_string())),
};
match read_res {
Ok(0) => {
sb_clone.finish();
return;
}
Ok(_) if probing => {
sb_clone.error(NetError::Io(Arc::new(std::io::Error::other(
"max size exceeded during read",
))));
return;
}
Ok(n) => {
total_read = total_read.saturating_add(n as u64);
sb_clone.push(Bytes::copy_from_slice(&buf[..n]));
if let Err(e) = check_total_deadline(tokio::time::Instant::now()) {
sb_clone.error(e);
return;
}
}
Err(e) => {
sb_clone.error(e);
return;
}
}
}
});
sb
}
}
struct SubStream {
id: u64,
parent: Arc<Mutex<State>>,
inner: ReceiverStream<Result<Bytes, NetError>>,
}
impl Stream for SubStream {
type Item = Result<Bytes, NetError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
Pin::new(&mut this.inner).poll_next(cx)
}
}
impl Drop for SubStream {
fn drop(&mut self) {
self.parent.lock().subs.remove(&self.id);
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::AsyncReadExt;
#[tokio::test(flavor = "current_thread")]
async fn shared_body_broadcasts_and_finishes() {
let sb = SharedBody::new(8);
let mut s1 = sb.subscribe_stream();
let mut s2 = sb.subscribe_stream();
sb.push(Bytes::from_static(b"hello"));
sb.push(Bytes::from_static(b" world"));
sb.finish();
let a1 = s1.next().await.unwrap().unwrap();
let a2 = s1.next().await.unwrap().unwrap();
let expected1: &[u8] = b"hello";
let expected2: &[u8] = b" world";
assert_eq!((&a1[..], &a2[..]), (expected1, expected2));
assert!(s1.next().await.is_none());
let b1 = s2.next().await.unwrap().unwrap();
let b2 = s2.next().await.unwrap().unwrap();
let expected1: &[u8] = b"hello";
let expected2: &[u8] = b" world";
assert_eq!((&b1[..], &b2[..]), (expected1, expected2));
assert!(s2.next().await.is_none());
}
#[tokio::test(flavor = "current_thread")]
async fn shared_body_drops_slow_subscriber() {
let sb = SharedBody::new(1);
let mut slow = sb.subscribe_stream();
let mut fast = sb.subscribe_stream();
sb.push(Bytes::from_static(b"A"));
let fa = fast.next().await.unwrap().unwrap();
sb.push(Bytes::from_static(b"B"));
let fb = fast.next().await.unwrap().unwrap();
let first = slow.next().await.unwrap();
assert!(first.is_ok());
let tail = slow.next().await; assert!(tail.is_none() || tail.unwrap().is_err());
sb.push(Bytes::from_static(b"C"));
let fc = fast.next().await.unwrap().unwrap();
let exp1: &[u8] = b"A";
let exp2: &[u8] = b"B";
let exp3: &[u8] = b"C";
assert_eq!((&fa[..], &fb[..], &fc[..]), (exp1, exp2, exp3));
}
#[tokio::test(flavor = "current_thread")]
async fn combined_reader_yields_peek_then_tail() {
let sb = SharedBody::new(8);
let sb2 = sb.clone();
let peek_buf = PeekBuf::from_slice(b"PEEK-");
tokio::spawn(async move {
sb2.push(Bytes::from_static(b"TAIL1"));
sb2.push(Bytes::from_static(b"TAIL2"));
sb2.finish();
});
let mut reader = SharedBody::combined_reader(peek_buf, Arc::new(sb));
let mut out = Vec::new();
reader.read_to_end(&mut out).await.unwrap();
assert_eq!(&out[..], b"PEEK-TAIL1TAIL2");
}
struct BlockingReader;
impl tokio::io::AsyncRead for BlockingReader {
fn poll_read(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
_buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::task::Poll::Pending
}
}
impl Unpin for BlockingReader {}
struct ErrorReader;
impl tokio::io::AsyncRead for ErrorReader {
fn poll_read(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
_buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::task::Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"test",
)))
}
}
impl Unpin for ErrorReader {}
async fn drain_result(sb: &Arc<SharedBody>) -> Result<Vec<u8>, NetError> {
let mut stream = sb.subscribe_stream();
let mut out = Vec::new();
while let Some(chunk) = stream.next().await {
out.extend_from_slice(&chunk?);
}
Ok(out)
}
#[tokio::test(flavor = "current_thread")]
async fn push_after_finish_is_noop() {
let sb = SharedBody::new(8);
let mut s = sb.subscribe_stream();
sb.push(Bytes::from_static(b"before"));
sb.finish();
sb.push(Bytes::from_static(b"after"));
let first = s.next().await.unwrap().unwrap();
assert_eq!(&first[..], b"before");
assert!(s.next().await.is_none(), "post-finish push must not appear");
}
#[tokio::test(flavor = "current_thread")]
async fn error_after_finish_is_noop() {
let sb = SharedBody::new(8);
let mut s = sb.subscribe_stream();
sb.finish();
sb.error(NetError::Cancelled("ignored".into()));
assert!(
s.next().await.is_none(),
"error after finish must not reopen stream"
);
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_delivers_all_data() {
let sb = SharedBody::from_reader(
std::io::Cursor::new(b"hello reader".to_vec()),
ReaderOptions::default(),
);
assert_eq!(drain_result(&sb).await.unwrap(), b"hello reader");
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_cancellation_errors_subscribers() {
let cancel = CancellationToken::new();
let sb = SharedBody::from_reader(
BlockingReader,
ReaderOptions {
cancel: Some(cancel.clone()),
..ReaderOptions::default()
},
);
let mut stream = sb.subscribe_stream();
cancel.cancel();
assert!(stream.next().await.unwrap().is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_idle_timeout_errors_subscribers() {
let sb = SharedBody::from_reader(
BlockingReader,
ReaderOptions {
idle_timeout: Some(Duration::from_millis(50)),
..ReaderOptions::default()
},
);
let err = drain_result(&sb).await.unwrap_err();
assert!(err.to_string().contains("idle"), "got: {err}");
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_total_timeout_errors_subscribers() {
let sb = SharedBody::from_reader(
BlockingReader,
ReaderOptions {
total_timeout: Some(Duration::ZERO),
..ReaderOptions::default()
},
);
let err = drain_result(&sb).await.unwrap_err();
assert!(err.to_string().contains("total"), "got: {err}");
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_io_error_errors_subscribers() {
let sb = SharedBody::from_reader(ErrorReader, ReaderOptions::default());
assert!(drain_result(&sb).await.is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_exact_max_size_succeeds() {
let sb = SharedBody::from_reader(
std::io::Cursor::new(vec![7u8; 10]),
ReaderOptions {
max_size: Some(10),
..ReaderOptions::default()
},
);
let body = drain_result(&sb).await.unwrap();
assert_eq!(body, vec![7u8; 10]);
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_total_timeout_fires_during_stalled_read() {
let sb = SharedBody::from_reader(
BlockingReader,
ReaderOptions {
total_timeout: Some(Duration::from_millis(50)),
..ReaderOptions::default()
},
);
let err = tokio::time::timeout(Duration::from_secs(2), drain_result(&sb))
.await
.expect("total timeout did not interrupt the blocked read")
.unwrap_err();
assert!(err.to_string().contains("total"), "got: {err}");
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_cancel_interrupts_blocked_read() {
let cancel = CancellationToken::new();
let sb = SharedBody::from_reader(
BlockingReader,
ReaderOptions {
cancel: Some(cancel.clone()),
..ReaderOptions::default()
},
);
let mut stream = sb.subscribe_stream();
tokio::time::sleep(Duration::from_millis(20)).await;
cancel.cancel();
let item = tokio::time::timeout(Duration::from_secs(2), stream.next())
.await
.expect("cancel did not interrupt the blocked read")
.unwrap();
assert!(item.is_err());
}
#[tokio::test(flavor = "current_thread")]
async fn from_reader_max_size_exceeded_errors_subscribers() {
let sb = SharedBody::from_reader(
std::io::Cursor::new(vec![0u8; 200]),
ReaderOptions {
max_size: Some(10),
..ReaderOptions::default()
},
);
let mut stream = sb.subscribe_stream();
let mut got_err = false;
while let Some(r) = stream.next().await {
if r.is_err() {
got_err = true;
break;
}
}
assert!(got_err);
}
}