diff --git a/src/body/chan.rs b/src/body/chan.rs new file mode 100644 index 0000000000..b45bea866b --- /dev/null +++ b/src/body/chan.rs @@ -0,0 +1,264 @@ +use std::fmt; +use std::sync::{Arc, Mutex}; +use std::task::{Context, Poll}; + +use atomic_waker::AtomicWaker; +use bytes::Bytes; +use http::HeaderMap; + +use crate::common::lock::LockResultExt; + +pub(crate) fn channel(wanter: bool) -> (Sender, Receiver) { + let shared = Arc::new(Shared { + state: Mutex::new(State { + item: None, + pending_error: None, + trailers: None, + sender_open: true, + receiver_open: true, + want: !wanter, + }), + sender_waker: AtomicWaker::new(), + receiver_waker: AtomicWaker::new(), + }); + + ( + Sender { + shared: Arc::clone(&shared), + trailers_sent: false, + }, + Receiver { + shared, + terminated: false, + }, + ) +} + +#[must_use = "Sender does nothing unless sent on"] +pub(crate) struct Sender { + shared: Arc, + trailers_sent: bool, +} + +pub(crate) struct Receiver { + shared: Arc, + terminated: bool, +} + +struct Shared { + state: Mutex, + sender_waker: AtomicWaker, + receiver_waker: AtomicWaker, +} + +struct State { + item: Option>, + // An error must not displace data which was already accepted. The old + // mpsc channel achieved this by sending the error from a cloned sender. + pending_error: Option, + trailers: Option, + sender_open: bool, + receiver_open: bool, + want: bool, +} + +impl Sender { + pub(crate) fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.shared.sender_waker.register(cx.waker()); + let state = self.shared.state.lock().panic_if_poisoned(); + if !state.receiver_open { + Poll::Ready(Err(crate::Error::new_closed())) + } else if state.want && state.item.is_none() && state.pending_error.is_none() { + Poll::Ready(Ok(())) + } else { + Poll::Pending + } + } + + #[cfg(test)] + pub(crate) async fn ready(&mut self) -> crate::Result<()> { + futures_util::future::poll_fn(|cx| self.poll_ready(cx)).await + } + + #[cfg(test)] + #[allow(dead_code)] + pub(crate) async fn send_data(&mut self, chunk: Bytes) -> crate::Result<()> { + self.ready().await?; + self.try_send_data(chunk) + .map_err(|_| crate::Error::new_closed()) + } + + pub(crate) fn try_send_data(&mut self, chunk: Bytes) -> Result<(), Bytes> { + let mut state = self.shared.state.lock().panic_if_poisoned(); + if !state.receiver_open + || !state.want + || state.item.is_some() + || state.pending_error.is_some() + { + return Err(chunk); + } + state.item = Some(Ok(chunk)); + drop(state); + self.shared.receiver_waker.wake(); + Ok(()) + } + + pub(crate) fn try_send_trailers( + &mut self, + trailers: HeaderMap, + ) -> Result<(), Option> { + if self.trailers_sent { + return Err(None); + } + self.trailers_sent = true; + + let mut state = self.shared.state.lock().panic_if_poisoned(); + if !state.receiver_open { + return Err(Some(trailers)); + } + state.trailers = Some(trailers); + drop(state); + self.shared.receiver_waker.wake(); + Ok(()) + } + + #[allow(dead_code)] + #[allow(clippy::unused_async_trait_impl)] + pub(crate) async fn send_trailers(&mut self, trailers: HeaderMap) -> crate::Result<()> { + self.try_send_trailers(trailers) + .map_err(|_| crate::Error::new_closed()) + } + + pub(crate) fn send_error(&mut self, err: crate::Error) { + let mut state = self.shared.state.lock().panic_if_poisoned(); + if !state.receiver_open { + return; + } + if state.item.is_none() { + state.item = Some(Err(err)); + } else if state.pending_error.is_none() { + state.pending_error = Some(err); + } + drop(state); + self.shared.receiver_waker.wake(); + } + + #[cfg(test)] + pub(crate) fn abort(mut self) { + self.send_error(crate::Error::new_body_write_aborted()); + } + + fn is_closed(&self) -> bool { + !self.shared.state.lock().panic_if_poisoned().receiver_open + } +} + +impl Receiver { + pub(crate) fn poll_next( + &mut self, + cx: &mut Context<'_>, + ) -> Poll>> { + if self.terminated { + return Poll::Ready(None); + } + + self.shared.receiver_waker.register(cx.waker()); + let mut state = self.shared.state.lock().panic_if_poisoned(); + let wake_sender = if !state.want { + state.want = true; + true + } else { + false + }; + if let Some(item) = state.item.take() { + drop(state); + self.shared.sender_waker.wake(); + return Poll::Ready(Some(item)); + } + if let Some(err) = state.pending_error.take() { + drop(state); + if wake_sender { + self.shared.sender_waker.wake(); + } + return Poll::Ready(Some(Err(err))); + } + let sender_open = state.sender_open; + drop(state); + if wake_sender { + self.shared.sender_waker.wake(); + } + if sender_open { + Poll::Pending + } else { + self.terminated = true; + Poll::Ready(None) + } + } + + pub(crate) fn take_trailers(&mut self) -> Option { + debug_assert!(self.terminated, "data channel still open before trailers"); + let mut state = self.shared.state.lock().panic_if_poisoned(); + state.trailers.take() + } +} + +impl Drop for Sender { + fn drop(&mut self) { + let mut state = self.shared.state.lock().panic_if_poisoned(); + state.sender_open = false; + drop(state); + self.shared.receiver_waker.wake(); + } +} + +impl Drop for Receiver { + fn drop(&mut self) { + let mut state = self.shared.state.lock().panic_if_poisoned(); + state.receiver_open = false; + drop(state); + self.shared.sender_waker.wake(); + } +} + +impl fmt::Debug for Sender { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + if self.is_closed() { + f.debug_tuple("Sender").field(&"Closed").finish() + } else { + f.debug_tuple("Sender").field(&"Open").finish() + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + async fn recv(rx: &mut Receiver) -> Option> { + futures_util::future::poll_fn(|cx| rx.poll_next(cx)).await + } + + #[tokio::test] + async fn error_queued_behind_accepted_data() { + let (mut tx, mut rx) = channel(false); + tx.try_send_data(Bytes::from_static(b"data")).unwrap(); + tx.send_error(crate::Error::new_incomplete()); + drop(tx); + + assert_eq!(recv(&mut rx).await.unwrap().unwrap(), "data"); + assert!(recv(&mut rx).await.unwrap().is_err()); + assert!(recv(&mut rx).await.is_none()); + } + + #[tokio::test] + async fn trailers_follow_data_close() { + let (mut tx, mut rx) = channel(false); + let mut trailers = HeaderMap::new(); + trailers.insert("x-trailer", "value".parse().unwrap()); + tx.try_send_trailers(trailers).unwrap(); + drop(tx); + + assert!(recv(&mut rx).await.is_none()); + assert_eq!(rx.take_trailers().unwrap()["x-trailer"], "value"); + } +} diff --git a/src/body/incoming.rs b/src/body/incoming.rs index 6ef6c0963a..b026a24d6b 100644 --- a/src/body/incoming.rs +++ b/src/body/incoming.rs @@ -1,38 +1,25 @@ use std::fmt; -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -use std::future::Future; use std::pin::Pin; use std::task::{Context, Poll}; use bytes::Bytes; -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -use futures_channel::{mpsc, oneshot}; #[cfg(all( any(feature = "http1", feature = "http2"), any(feature = "client", feature = "server") ))] use futures_core::ready; -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -use futures_core::{stream::FusedStream, Stream}; // for mpsc::Receiver -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -use http::HeaderMap; use http_body::{Body, Frame, SizeHint}; +#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] +use super::chan; #[cfg(all( any(feature = "http1", feature = "http2"), any(feature = "client", feature = "server") ))] use super::DecodedLength; -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -use crate::common::watch; #[cfg(all(feature = "http2", any(feature = "client", feature = "server")))] use crate::proto::h2::ping; -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -type BodySender = mpsc::Sender>; -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -type TrailersSender = oneshot::Sender; - /// A stream of `Bytes`, used when receiving bodies from the network. /// /// Note that Users should not instantiate this struct directly. When working with the hyper client, @@ -58,9 +45,7 @@ enum Kind { #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] Chan { content_length: DecodedLength, - want_tx: watch::Sender, - data_rx: mpsc::Receiver>, - trailers_rx: oneshot::Receiver, + rx: chan::Receiver, }, #[cfg(all(feature = "http2", any(feature = "client", feature = "server")))] H2 { @@ -86,18 +71,8 @@ enum Kind { /// /// [`Body::channel()`]: struct.Body.html#method.channel /// [`Sender::abort()`]: struct.Sender.html#method.abort -#[must_use = "Sender does nothing unless sent on"] #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -pub(crate) struct Sender { - want_rx: watch::Receiver, - data_tx: BodySender, - trailers_tx: Option, -} - -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -const WANT_PENDING: usize = 1; -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -const WANT_READY: usize = 2; +pub(crate) use super::chan::Sender; impl Incoming { /// Create a `Body` stream with an associated sender half. @@ -112,26 +87,8 @@ impl Incoming { #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] pub(crate) fn new_channel(content_length: DecodedLength, wanter: bool) -> (Sender, Incoming) { - let (data_tx, data_rx) = mpsc::channel(0); - let (trailers_tx, trailers_rx) = oneshot::channel(); - - // If wanter is true, `Sender::poll_ready()` won't becoming ready - // until the `Body` has been polled for data once. - let want = if wanter { WANT_PENDING } else { WANT_READY }; - - let (want_tx, want_rx) = watch::channel(want); - - let tx = Sender { - want_rx, - data_tx, - trailers_tx: Some(trailers_tx), - }; - let rx = Incoming::new(Kind::Chan { - content_length, - want_tx, - data_rx, - trailers_rx, - }); + let (tx, rx) = chan::channel(wanter); + let rx = Incoming::new(Kind::Chan { content_length, rx }); (tx, rx) } @@ -210,24 +167,13 @@ impl Body for Incoming { #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] Kind::Chan { content_length: len, - data_rx, - want_tx, - trailers_rx, + rx, } => { - want_tx.send(WANT_READY); - - if !data_rx.is_terminated() { - if let Some(chunk) = ready!(Pin::new(data_rx).poll_next(cx)?) { - len.sub_if(chunk.len() as u64); - return Poll::Ready(Some(Ok(Frame::data(chunk)))); - } - } - - // check trailers after data is terminated - match ready!(Pin::new(trailers_rx).poll(cx)) { - Ok(t) => Poll::Ready(Some(Ok(Frame::trailers(t)))), - Err(_) => Poll::Ready(None), + if let Some(chunk) = ready!(rx.poll_next(cx)?) { + len.sub_if(chunk.len() as u64); + return Poll::Ready(Some(Ok(Frame::data(chunk)))); } + Poll::Ready(rx.take_trailers().map(Frame::trailers).map(Ok)) } #[cfg(all(feature = "http2", any(feature = "client", feature = "server")))] Kind::H2 { @@ -353,120 +299,12 @@ impl fmt::Debug for Incoming { } } -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -impl Sender { - /// Check to see if this `Sender` can send more data. - pub(crate) fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { - // Check if the receiver end has tried polling for the body yet - ready!(self.poll_want(cx)?); - self.data_tx - .poll_ready(cx) - .map_err(|_| crate::Error::new_closed()) - } - - fn poll_want(&mut self, cx: &mut Context<'_>) -> Poll> { - match self.want_rx.load(cx) { - WANT_READY => Poll::Ready(Ok(())), - WANT_PENDING => Poll::Pending, - watch::CLOSED => Poll::Ready(Err(crate::Error::new_closed())), - unexpected => unreachable!("want_rx value: {}", unexpected), - } - } - - #[cfg(test)] - async fn ready(&mut self) -> crate::Result<()> { - futures_util::future::poll_fn(|cx| self.poll_ready(cx)).await - } - - /// Send data on data channel when it is ready. - #[cfg(test)] - #[allow(unused)] - pub(crate) async fn send_data(&mut self, chunk: Bytes) -> crate::Result<()> { - self.ready().await?; - self.data_tx - .try_send(Ok(chunk)) - .map_err(|_| crate::Error::new_closed()) - } - - /// Send trailers on trailers channel. - #[allow(unused)] - #[allow(clippy::unused_async_trait_impl)] - pub(crate) async fn send_trailers(&mut self, trailers: HeaderMap) -> crate::Result<()> { - let tx = match self.trailers_tx.take() { - Some(tx) => tx, - None => return Err(crate::Error::new_closed()), - }; - tx.send(trailers).map_err(|_| crate::Error::new_closed()) - } - - /// Try to send data on this channel. - /// - /// # Errors - /// - /// Returns `Err(Bytes)` if the channel could not (currently) accept - /// another `Bytes`. - /// - /// # Note - /// - /// This is mostly useful for when trying to send from some other thread - /// that doesn't have an async context. If in an async context, prefer - /// `send_data()` instead. - #[cfg(feature = "http1")] - pub(crate) fn try_send_data(&mut self, chunk: Bytes) -> Result<(), Bytes> { - self.data_tx - .try_send(Ok(chunk)) - .map_err(|err| err.into_inner().expect("just sent Ok")) - } - - #[cfg(feature = "http1")] - pub(crate) fn try_send_trailers( - &mut self, - trailers: HeaderMap, - ) -> Result<(), Option> { - let tx = match self.trailers_tx.take() { - Some(tx) => tx, - None => return Err(None), - }; - - tx.send(trailers).map_err(Some) - } - - #[cfg(test)] - pub(crate) fn abort(mut self) { - self.send_error(crate::Error::new_body_write_aborted()); - } - - pub(crate) fn send_error(&mut self, err: crate::Error) { - let _ = self - .data_tx - // clone so the send works even if buffer is full - .clone() - .try_send(Err(err)); - } -} - -#[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] -impl fmt::Debug for Sender { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - #[derive(Debug)] - struct Open; - #[derive(Debug)] - struct Closed; - - let mut builder = f.debug_tuple("Sender"); - match self.want_rx.peek() { - watch::CLOSED => builder.field(&Closed), - _ => builder.field(&Open), - }; - - builder.finish() - } -} - #[cfg(test)] mod tests { #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] use std::mem; + #[cfg(all(feature = "nightly", not(miri)))] + use std::pin::Pin; #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] use std::task::Poll; @@ -474,6 +312,8 @@ mod tests { use super::{Body, Incoming, SizeHint}; #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] use super::{DecodedLength, Sender}; + #[cfg(all(feature = "nightly", not(miri)))] + use bytes::Bytes; #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] use http_body_util::BodyExt; @@ -494,7 +334,7 @@ mod tests { assert_eq!( mem::size_of::(), - mem::size_of::() * 5, + mem::size_of::() * 2, "Sender" ); @@ -505,6 +345,51 @@ mod tests { ); } + #[cfg(all(feature = "nightly", not(miri)))] + #[bench] + fn bench_channel_create_and_drop(b: &mut test::Bencher) { + b.iter(|| { + let _ = test::black_box(Incoming::new_channel( + DecodedLength::CHUNKED, + /* wanter = */ false, + )); + }); + } + + #[cfg(all(feature = "nightly", not(miri)))] + #[bench] + fn bench_channel_data_handoff(b: &mut test::Bencher) { + let (mut tx, mut body) = + Incoming::new_channel(DecodedLength::CHUNKED, /* wanter = */ false); + let mut cx = std::task::Context::from_waker(std::task::Waker::noop()); + + b.iter(|| { + assert!(tx.poll_ready(&mut cx).is_ready()); + tx.try_send_data(Bytes::from_static(b"hello world")) + .unwrap(); + let frame = match Pin::new(&mut body).poll_frame(&mut cx) { + Poll::Ready(Some(Ok(frame))) => frame, + unexpected => panic!("unexpected body poll: {unexpected:?}"), + }; + test::black_box(frame); + }); + } + + #[cfg(all(feature = "nightly", not(miri)))] + #[bench] + fn bench_channel_want_transition(b: &mut test::Bencher) { + let mut cx = std::task::Context::from_waker(std::task::Waker::noop()); + + b.iter(|| { + let (mut tx, mut body) = + Incoming::new_channel(DecodedLength::CHUNKED, /* wanter = */ true); + assert!(tx.poll_ready(&mut cx).is_pending()); + assert!(Pin::new(&mut body).poll_frame(&mut cx).is_pending()); + assert!(tx.poll_ready(&mut cx).is_ready()); + let _ = test::black_box((tx, body)); + }); + } + #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] #[test] fn size_hint() { diff --git a/src/body/mod.rs b/src/body/mod.rs index c34d019d21..42678b63d1 100644 --- a/src/body/mod.rs +++ b/src/body/mod.rs @@ -79,6 +79,8 @@ pub(crate) use self::incoming::Sender; ))] pub(crate) use self::length::DecodedLength; +#[cfg(all(any(feature = "client", feature = "server"), feature = "http1"))] +mod chan; mod incoming; #[cfg(all( any(feature = "http1", feature = "http2"), diff --git a/src/common/mod.rs b/src/common/mod.rs index 5be740b000..da1ad9eafb 100644 --- a/src/common/mod.rs +++ b/src/common/mod.rs @@ -18,5 +18,3 @@ pub(crate) mod task; all(any(feature = "client", feature = "server"), feature = "http2"), ))] pub(crate) mod time; -#[cfg(all(any(feature = "client", feature = "server"), feature = "http1"))] -pub(crate) mod watch; diff --git a/src/common/watch.rs b/src/common/watch.rs deleted file mode 100644 index 81acd5df70..0000000000 --- a/src/common/watch.rs +++ /dev/null @@ -1,73 +0,0 @@ -//! An SPSC broadcast channel. -//! -//! - The value can only be a `usize`. -//! - The consumer is only notified if the value is different. -//! - The value `0` is reserved for closed. - -use atomic_waker::AtomicWaker; -use std::sync::{ - atomic::{AtomicUsize, Ordering}, - Arc, -}; -use std::task; - -type Value = usize; - -pub(crate) const CLOSED: usize = 0; - -pub(crate) fn channel(initial: Value) -> (Sender, Receiver) { - debug_assert!( - initial != CLOSED, - "watch::channel initial state of 0 is reserved" - ); - - let shared = Arc::new(Shared { - value: AtomicUsize::new(initial), - waker: AtomicWaker::new(), - }); - - ( - Sender { - shared: shared.clone(), - }, - Receiver { shared }, - ) -} - -pub(crate) struct Sender { - shared: Arc, -} - -pub(crate) struct Receiver { - shared: Arc, -} - -struct Shared { - value: AtomicUsize, - waker: AtomicWaker, -} - -impl Sender { - pub(crate) fn send(&mut self, value: Value) { - if self.shared.value.swap(value, Ordering::SeqCst) != value { - self.shared.waker.wake(); - } - } -} - -impl Drop for Sender { - fn drop(&mut self) { - self.send(CLOSED); - } -} - -impl Receiver { - pub(crate) fn load(&mut self, cx: &mut task::Context<'_>) -> Value { - self.shared.waker.register(cx.waker()); - self.shared.value.load(Ordering::SeqCst) - } - - pub(crate) fn peek(&self) -> Value { - self.shared.value.load(Ordering::Relaxed) - } -}