use crate::courierust_bytes::Bytes;
use crate::courierust_error::{Error, ErrorKind};
use crate::Result;
use std::ops::Deref;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc::{Receiver, Sender, TryRecvError};
use std::sync::{Arc, Mutex};
#[derive(Default)]
pub enum Body {
#[default]
Empty,
Bytes(Bytes),
Channel(Receiver<Result<Bytes>>),
Stream(ChannelStream),
}
impl Body {
pub fn is_empty(&self) -> bool {
match self {
Self::Empty => true,
Self::Bytes(b) => b.is_empty(),
Self::Channel(_) | Self::Stream(_) => false,
}
}
pub fn is_bytes(&self) -> bool {
matches!(self, Self::Bytes(_))
}
pub fn is_stream(&self) -> bool {
matches!(self, Self::Channel(_) | Self::Stream(_))
}
pub fn into_stream(self) -> Option<ChannelStream> {
match self {
Self::Channel(rx) => Some(ChannelStream::raw(rx)),
Self::Stream(stream) => Some(stream),
Self::Empty | Self::Bytes(_) => None,
}
}
pub fn as_bytes(&self) -> Option<&[u8]> {
match self {
Self::Bytes(b) => Some(b),
_ => None,
}
}
pub fn try_next_chunk(&mut self) -> Result<Option<Bytes>> {
match self {
Self::Channel(rx) => try_next(rx),
Self::Stream(stream) => try_next(&stream.rx),
_ => Ok(None),
}
}
pub fn collect(self) -> Result<Bytes> {
self.collect_limited(usize::MAX)
}
pub fn collect_limited(self, max: usize) -> Result<Bytes> {
match self {
Self::Empty => Ok(Bytes::new()),
Self::Bytes(b) if b.len() <= max => Ok(b),
Self::Bytes(_) => Err(Error::overflow("body exceeds configured limit")),
Self::Channel(rx) => drain(rx, max),
Self::Stream(stream) => drain(stream.into_receiver(), max),
}
}
pub fn len(&self) -> Option<usize> {
match self {
Self::Empty => Some(0),
Self::Bytes(b) => Some(b.len()),
Self::Channel(_) | Self::Stream(_) => None,
}
}
}
impl crate::courierust_http::response::Response<Body> {
pub fn bytes(self) -> Result<Bytes> {
self.body.collect()
}
pub fn text(self) -> Result<String> {
let bytes = self.body.collect()?;
core::str::from_utf8(&bytes)
.map(Into::into)
.map_err(|_| Error::with_message(ErrorKind::Other, "response body is not valid UTF-8"))
}
}
fn try_next(rx: &Receiver<Result<Bytes>>) -> Result<Option<Bytes>> {
match rx.try_recv() {
Ok(chunk) => chunk.map(Some),
Err(TryRecvError::Empty) => Ok(None),
Err(TryRecvError::Disconnected) => Ok(None),
}
}
fn drain(rx: Receiver<Result<Bytes>>, max: usize) -> Result<Bytes> {
let mut out = Vec::new();
while let Ok(chunk) = rx.recv() {
let b = chunk?;
if b.len() > max.saturating_sub(out.len()) {
return Err(Error::overflow("body exceeds configured limit"));
}
out.extend_from_slice(&b);
}
Ok(Bytes::from(out))
}
pub struct ChannelStream {
rx: Receiver<Result<Bytes>>,
wake: Option<Arc<BodyWake>>,
}
impl ChannelStream {
pub fn raw(rx: Receiver<Result<Bytes>>) -> Self {
Self { rx, wake: None }
}
pub fn install_wake(&self, wake: impl Fn() + Send + Sync + 'static) {
if let Some(slot) = &self.wake {
slot.install(wake);
}
}
pub fn clear_wake(&self) {
if let Some(slot) = &self.wake {
slot.clear();
}
}
pub fn has_wake(&self) -> bool {
self.wake.as_ref().is_some_and(|slot| slot.is_installed())
}
pub fn into_receiver(self) -> Receiver<Result<Bytes>> {
self.rx
}
}
impl Deref for ChannelStream {
type Target = Receiver<Result<Bytes>>;
fn deref(&self) -> &Self::Target {
&self.rx
}
}
impl std::fmt::Debug for ChannelStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "ChannelStream(wake={})", self.has_wake())
}
}
impl From<ChannelStream> for Body {
fn from(stream: ChannelStream) -> Self {
Self::Stream(stream)
}
}
pub struct BodyWake {
installed: AtomicBool,
slot: Mutex<Option<Arc<dyn Fn() + Send + Sync>>>,
}
impl BodyWake {
fn new() -> Arc<Self> {
Arc::new(Self {
installed: AtomicBool::new(false),
slot: Mutex::new(None),
})
}
fn slot(&self) -> std::sync::MutexGuard<'_, Option<Arc<dyn Fn() + Send + Sync>>> {
self.slot
.lock()
.unwrap_or_else(|poison| poison.into_inner())
}
pub fn install(&self, wake: impl Fn() + Send + Sync + 'static) {
let wake: Arc<dyn Fn() + Send + Sync> = Arc::new(wake);
*self.slot() = Some(wake);
self.installed.store(true, Ordering::Release);
}
pub fn clear(&self) {
self.installed.store(false, Ordering::Release);
*self.slot() = None;
}
pub fn is_installed(&self) -> bool {
self.installed.load(Ordering::Acquire)
}
pub fn fire(&self) {
if !self.installed.load(Ordering::Acquire) {
return;
}
let wake = self.slot().clone();
if let Some(wake) = wake {
wake();
}
}
}
impl From<Bytes> for Body {
fn from(b: Bytes) -> Self {
if b.is_empty() {
Self::Empty
} else {
Self::Bytes(b)
}
}
}
impl From<Vec<u8>> for Body {
fn from(v: Vec<u8>) -> Self {
Self::from(Bytes::from(v))
}
}
impl From<&'static [u8]> for Body {
fn from(b: &'static [u8]) -> Self {
Self::from(Bytes::from_static(b))
}
}
impl From<&'static str> for Body {
fn from(s: &'static str) -> Self {
Self::from(Bytes::from_static(s.as_bytes()))
}
}
impl From<String> for Body {
fn from(s: String) -> Self {
Self::from(Bytes::from(s))
}
}
impl From<Receiver<Result<Bytes>>> for Body {
fn from(rx: Receiver<Result<Bytes>>) -> Self {
Self::Channel(rx)
}
}
impl std::fmt::Debug for Body {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Empty => write!(f, "Body::Empty"),
Self::Bytes(b) => write!(f, "Body::Bytes({} bytes)", b.len()),
Self::Channel(_) => write!(f, "Body::Channel"),
Self::Stream(stream) => write!(f, "Body::Stream({stream:?})"),
}
}
}
pub struct BodySender {
tx: Sender<Result<Bytes>>,
cancelled: Arc<AtomicBool>,
wake: Arc<BodyWake>,
}
impl BodySender {
pub fn from_sender(tx: std::sync::mpsc::Sender<Result<Bytes>>) -> Self {
Self {
tx,
cancelled: Arc::new(AtomicBool::new(false)),
wake: BodyWake::new(),
}
}
#[inline]
pub fn is_cancelled(&self) -> bool {
self.cancelled.load(Ordering::Acquire)
}
#[inline]
pub fn cancel(&self) {
self.cancelled.store(true, Ordering::Release);
}
#[inline]
pub fn cancel_flag(&self) -> Arc<AtomicBool> {
self.cancelled.clone()
}
pub fn send(&self, chunk: Bytes) -> Result<()> {
self.tx.send(Ok(chunk)).map_err(|_| {
self.cancel();
Error::canceled("body receiver dropped")
})?;
self.wake.fire();
Ok(())
}
pub fn send_bytes(&self, chunk: &[u8]) -> Result<()> {
self.send(Bytes::from(chunk))
}
pub fn send_result(&self, result: Result<Bytes>) -> Result<()> {
self.tx.send(result).map_err(|_| {
self.cancel();
Error::canceled("body receiver dropped")
})?;
self.wake.fire();
Ok(())
}
pub fn fail(&self, err: Error) {
if self.tx.send(Err(err)).is_ok() {
self.wake.fire();
}
}
}
pub fn channel() -> (BodySender, Body) {
let (tx, rx) = std::sync::mpsc::channel::<Result<Bytes>>();
let wake = BodyWake::new();
(
BodySender {
tx,
cancelled: Arc::new(AtomicBool::new(false)),
wake: wake.clone(),
},
Body::Stream(ChannelStream {
rx,
wake: Some(wake),
}),
)
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
#[test]
fn body_sender_wakes_only_while_installed() {
let (tx, body) = channel();
let stream = body.into_stream().expect("channel() builds a stream");
assert!(!stream.has_wake(), "no transport adopted this body yet");
let hits = Arc::new(AtomicUsize::new(0));
let counter = hits.clone();
stream.install_wake(move || {
counter.fetch_add(1, Ordering::Relaxed);
});
assert!(stream.has_wake());
tx.send(Bytes::from_static(b"one")).unwrap();
tx.send_bytes(b"two").unwrap();
assert_eq!(hits.load(Ordering::Relaxed), 2, "one wake per chunk");
tx.fail(Error::timeout("boom"));
assert_eq!(hits.load(Ordering::Relaxed), 3, "a failure wakes too");
stream.clear_wake();
assert!(!stream.has_wake());
tx.send_bytes(b"three").unwrap();
assert_eq!(
hits.load(Ordering::Relaxed),
3,
"a cleared wake must stay quiet"
);
drop(stream);
assert!(tx.send_bytes(b"four").is_err(), "the receiver is gone");
assert!(tx.is_cancelled(), "a dropped receiver cancels the producer");
}
#[test]
fn raw_channel_has_no_wake_to_fire() {
let (tx, rx) = std::sync::mpsc::channel();
let stream = Body::Channel(rx)
.into_stream()
.expect("a channel body is a stream");
assert!(!stream.has_wake());
stream.install_wake(|| panic!("a raw channel must not install a wake"));
assert!(!stream.has_wake());
tx.send(Ok(Bytes::from_static(b"x"))).unwrap();
assert_eq!(stream.try_recv().unwrap().unwrap().len(), 1);
}
}