use bytes::Buf;
use bytes::BufMut;
use bytes::Bytes;
use bytes::BytesMut;
use std::sync::Arc;
use crate::compression::CompressionPolicy;
use crate::compression::CompressionRegistry;
use crate::error::ConnectError;
pub mod flags {
pub const DATA: u8 = 0x00;
pub const COMPRESSED: u8 = 0x01;
pub const END_STREAM: u8 = 0x02;
pub const GRPC_WEB_TRAILER: u8 = 0x80;
}
pub const HEADER_SIZE: usize = 5;
pub(crate) const MIN_CHAIN_SIZE: usize = 16 * 1024;
#[derive(Debug, Clone)]
pub struct Envelope {
pub flags: u8,
pub data: Bytes,
}
impl Envelope {
pub fn data(data: Bytes) -> Self {
Self {
flags: flags::DATA,
data,
}
}
pub fn compressed(data: Bytes) -> Self {
Self {
flags: flags::COMPRESSED,
data,
}
}
pub fn end_stream(data: Bytes) -> Self {
Self {
flags: flags::END_STREAM,
data,
}
}
pub fn is_compressed(&self) -> bool {
self.flags & flags::COMPRESSED != 0
}
pub fn is_end_stream(&self) -> bool {
self.flags & flags::END_STREAM != 0
}
pub fn encode(&self) -> Bytes {
let mut buf = BytesMut::with_capacity(HEADER_SIZE + self.data.len());
write_envelope(self.flags, &self.data, &mut buf)
.expect("envelope payload exceeds u32::MAX");
buf.freeze()
}
pub(crate) fn encode_body_parts(
flags: u8,
body: crate::response::EncodedBody,
min_chain: usize,
) -> (Bytes, Vec<Bytes>) {
let mut head = BytesMut::new();
let mut segments = Vec::new();
write_envelope_chained(flags, body, &mut head, &mut segments, min_chain)
.expect("envelope payload exceeds u32::MAX");
(head.freeze(), segments)
}
pub fn decode(buf: &mut BytesMut) -> Result<Option<Self>, ConnectError> {
Self::decode_with_limit(buf, usize::MAX)
}
pub fn decode_with_limit(
buf: &mut BytesMut,
max_size: usize,
) -> Result<Option<Self>, ConnectError> {
if buf.len() < HEADER_SIZE {
return Ok(None);
}
let flags = buf[0];
let length = u32::from_be_bytes([buf[1], buf[2], buf[3], buf[4]]) as usize;
if length > max_size {
return Err(ConnectError::resource_exhausted(format!(
"message size {length} exceeds limit {max_size}"
)));
}
if buf.len() < HEADER_SIZE.saturating_add(length) {
return Ok(None);
}
buf.advance(HEADER_SIZE);
let data = buf.split_to(length).freeze();
Ok(Some(Self { flags, data }))
}
}
pub(crate) struct EnvelopeDecoder {
max_message_size: usize,
streaming_encoding: Option<String>,
compression: Arc<CompressionRegistry>,
done: bool,
}
impl EnvelopeDecoder {
pub(crate) fn new(
max_message_size: usize,
streaming_encoding: Option<String>,
compression: Arc<CompressionRegistry>,
) -> Self {
Self {
max_message_size,
streaming_encoding,
compression,
done: false,
}
}
pub(crate) fn is_done(&self) -> bool {
self.done
}
}
impl tokio_util::codec::Decoder for EnvelopeDecoder {
type Item = Bytes;
type Error = ConnectError;
fn decode(&mut self, buf: &mut BytesMut) -> Result<Option<Bytes>, ConnectError> {
if self.done {
return Ok(None);
}
let envelope = match Envelope::decode_with_limit(buf, self.max_message_size)? {
Some(envelope) => envelope,
None => return Ok(None), };
if envelope.is_end_stream() {
tracing::trace!("client stream: received end-stream envelope");
self.done = true;
return Ok(None);
}
let data = if envelope.is_compressed() {
let encoding = match self.streaming_encoding.as_deref() {
Some(enc) if enc != "identity" => enc,
_ => {
return Err(ConnectError::internal(
"received compressed message without connect-content-encoding header",
));
}
};
self.compression.decompress_with_limit(
encoding,
envelope.data,
self.max_message_size,
)?
} else {
envelope.data
};
tracing::trace!(
size = data.len(),
"client stream: dispatching message to handler"
);
Ok(Some(data))
}
fn decode_eof(&mut self, buf: &mut BytesMut) -> Result<Option<Bytes>, ConnectError> {
match self.decode(buf)? {
some @ Some(_) => Ok(some),
None => {
if !buf.is_empty() {
tracing::debug!(
remaining_bytes = buf.len(),
"client stream: body ended with incomplete envelope"
);
Err(ConnectError::invalid_argument(
"incomplete request envelope",
))
} else {
Ok(None)
}
}
}
}
}
pub(crate) struct EnvelopeEncoder {
compression: Option<(Arc<CompressionRegistry>, String)>,
policy: CompressionPolicy,
}
impl EnvelopeEncoder {
pub(crate) fn new(
compression: Option<(Arc<CompressionRegistry>, impl Into<String>)>,
policy: CompressionPolicy,
) -> Self {
Self {
compression: compression.map(|(reg, enc)| (reg, enc.into())),
policy,
}
}
pub(crate) fn uncompressed() -> Self {
Self {
compression: None,
policy: CompressionPolicy::disabled(),
}
}
pub(crate) fn encode_end_stream(
&mut self,
data: Bytes,
dst: &mut BytesMut,
) -> Result<(), ConnectError> {
write_envelope(flags::END_STREAM, &data, dst)
}
pub(crate) fn encode_chained(
&mut self,
body: crate::response::EncodedBody,
dst: &mut BytesMut,
chained: &mut impl Extend<Bytes>,
min_chain: usize,
) -> Result<(), ConnectError> {
let (flag, body) = if let Some((ref comp, ref encoding)) = self.compression
&& self.policy.should_compress(body.len())
{
let compressed = comp.compress(encoding, &body.into_contiguous())?;
(flags::COMPRESSED, compressed.into())
} else {
(flags::DATA, body)
};
write_envelope_chained(flag, body, dst, chained, min_chain)
}
}
impl tokio_util::codec::Encoder<Bytes> for EnvelopeEncoder {
type Error = ConnectError;
fn encode(&mut self, data: Bytes, dst: &mut BytesMut) -> Result<(), ConnectError> {
let mut chained = Vec::new();
self.encode_chained(data.into(), dst, &mut chained, usize::MAX)?;
debug_assert!(chained.is_empty(), "usize::MAX threshold cannot chain");
Ok(())
}
}
fn write_envelope(flag: u8, data: &[u8], dst: &mut BytesMut) -> Result<(), ConnectError> {
put_envelope_header(flag, envelope_length(data.len())?, dst);
dst.put_slice(data);
Ok(())
}
fn write_envelope_chained(
flag: u8,
body: crate::response::EncodedBody,
dst: &mut BytesMut,
chained: &mut impl Extend<Bytes>,
min_chain: usize,
) -> Result<(), ConnectError> {
let total = body.len();
let declared = envelope_length(total)?;
if total < min_chain {
dst.reserve(HEADER_SIZE + total);
put_envelope_header(flag, declared, dst);
for segment in body.segments() {
dst.put_slice(segment);
}
return Ok(());
}
put_envelope_header(flag, declared, dst);
match body {
crate::response::EncodedBody::Contiguous(bytes) => chained.extend([bytes]),
crate::response::EncodedBody::Segmented(segments) => chained.extend(segments),
}
Ok(())
}
fn envelope_length(len: usize) -> Result<u32, ConnectError> {
u32::try_from(len).map_err(|_| {
ConnectError::resource_exhausted(format!("envelope payload {len} bytes exceeds u32::MAX"))
})
}
fn put_envelope_header(flag: u8, len: u32, dst: &mut BytesMut) {
dst.reserve(HEADER_SIZE);
dst.put_u8(flag);
dst.put_u32(len);
}
#[cfg(test)]
mod tests {
use super::*;
use tokio_util::codec::{Decoder, Encoder};
fn decoder(max_message_size: usize) -> EnvelopeDecoder {
EnvelopeDecoder::new(
max_message_size,
None,
Arc::new(CompressionRegistry::default()),
)
}
#[test]
fn encode_body_parts_chains_large_payload_by_refcount() {
let payload = Bytes::from(vec![7u8; 64]);
let ptr = payload.as_ptr();
let (head, chained) = Envelope::encode_body_parts(flags::DATA, payload.clone().into(), 64);
assert_eq!(head.len(), HEADER_SIZE);
let [chained] = &chained[..] else {
panic!("payload at threshold must chain as one segment");
};
assert!(std::ptr::eq(chained.as_ptr(), ptr), "must not copy");
let mut reassembled = BytesMut::from(&head[..]);
reassembled.put_slice(chained);
assert_eq!(
reassembled.freeze(),
Envelope::data(payload).encode(),
"chained wire bytes must match contiguous encoding"
);
}
#[test]
#[cfg(feature = "gzip")]
fn encode_chained_chains_large_compressed_payload() {
let registry = Arc::new(CompressionRegistry::default());
let mut enc = EnvelopeEncoder::new(
Some((Arc::clone(®istry), "gzip")),
CompressionPolicy::default().with_min_size(0),
);
let data: Vec<u8> = (0..64 * 1024u32)
.map(|i| (i.wrapping_mul(2654435761) >> 13) as u8)
.collect();
let mut dst = BytesMut::new();
let mut chained = Vec::new();
enc.encode_chained(Bytes::from(data).into(), &mut dst, &mut chained, 1024)
.unwrap();
let [chained] = &chained[..] else {
panic!("large compressed payload must chain as one segment");
};
assert_eq!(dst.len(), HEADER_SIZE);
assert_eq!(dst[0], flags::COMPRESSED);
assert_eq!(
u32::from_be_bytes([dst[1], dst[2], dst[3], dst[4]]) as usize,
chained.len()
);
let mut wire = dst;
wire.put_slice(chained);
let mut dec = EnvelopeDecoder::new(1024 * 1024, Some("gzip".to_owned()), registry);
let decoded = Decoder::decode(&mut dec, &mut wire).unwrap().unwrap();
assert_eq!(decoded.len(), 64 * 1024);
}
#[test]
fn encode_chained_keeps_segments_of_a_large_body() {
use crate::response::EncodedBody;
let lead = Bytes::from_static(b"tag");
let big = Bytes::from(vec![3u8; 64]);
let tail = Bytes::from_static(b"end");
let body = EncodedBody::Segmented(vec![lead.clone(), big.clone(), tail.clone()]);
let total = body.len();
let mut enc = EnvelopeEncoder::uncompressed();
let mut dst = BytesMut::new();
let mut segments = Vec::new();
enc.encode_chained(body, &mut dst, &mut segments, 32)
.unwrap();
assert_eq!(dst.len(), HEADER_SIZE, "only the header is copied");
assert_eq!(
u32::from_be_bytes([dst[1], dst[2], dst[3], dst[4]]) as usize,
total,
"header declares the total across segments"
);
let [s0, s1, s2] = &segments[..] else {
panic!("expected three segments, got {}", segments.len());
};
assert!(std::ptr::eq(s0.as_ptr(), lead.as_ptr()));
assert!(std::ptr::eq(s1.as_ptr(), big.as_ptr()));
assert!(std::ptr::eq(s2.as_ptr(), tail.as_ptr()));
}
#[test]
fn encode_chained_copies_a_small_segmented_body() {
use crate::response::EncodedBody;
let body =
EncodedBody::Segmented(vec![Bytes::from_static(b"ab"), Bytes::from_static(b"cd")]);
let mut enc = EnvelopeEncoder::uncompressed();
let mut dst = BytesMut::new();
let mut segments = Vec::new();
enc.encode_chained(body, &mut dst, &mut segments, 32)
.unwrap();
assert!(segments.is_empty());
assert_eq!(
dst.freeze(),
Envelope::data(Bytes::from_static(b"abcd")).encode()
);
}
#[test]
#[cfg(feature = "gzip")]
fn encode_chained_compresses_a_segmented_body_as_one() {
use crate::response::EncodedBody;
let registry = Arc::new(CompressionRegistry::default());
let mut enc = EnvelopeEncoder::new(
Some((Arc::clone(®istry), "gzip")),
CompressionPolicy::default().with_min_size(0),
);
let body = EncodedBody::Segmented(vec![
Bytes::from(vec![b'a'; 4096]),
Bytes::from(vec![b'b'; 4096]),
]);
let expected = body.clone().into_contiguous();
let mut wire = BytesMut::new();
let mut segments = Vec::new();
enc.encode_chained(body, &mut wire, &mut segments, usize::MAX)
.unwrap();
assert!(segments.is_empty());
assert_eq!(wire[0], flags::COMPRESSED);
let mut dec = EnvelopeDecoder::new(1024 * 1024, Some("gzip".to_owned()), registry);
let decoded = Decoder::decode(&mut dec, &mut wire).unwrap().unwrap();
assert_eq!(decoded, expected);
}
#[test]
fn encode_body_parts_small_payload_stays_contiguous() {
let payload = Bytes::from_static(b"tiny");
let (head, chained) = Envelope::encode_body_parts(flags::DATA, payload.clone().into(), 64);
assert!(chained.is_empty());
assert_eq!(head, Envelope::data(payload).encode());
}
#[test]
fn encode_body_parts_declares_the_length_it_emits() {
use crate::response::EncodedBody;
let cases: Vec<(&str, EncodedBody)> = vec![
("empty", EncodedBody::Contiguous(Bytes::new())),
(
"sub-threshold contiguous",
EncodedBody::Contiguous(Bytes::from_static(b"small")),
),
(
"sub-threshold segmented",
EncodedBody::Segmented(vec![Bytes::from_static(b"ab"), Bytes::from_static(b"cd")]),
),
(
"over-threshold contiguous",
EncodedBody::Contiguous(Bytes::from(vec![9u8; 128])),
),
(
"over-threshold two segments",
EncodedBody::Segmented(vec![
Bytes::from(vec![1u8; 64]),
Bytes::from(vec![2u8; 64]),
]),
),
(
"over-threshold many segments",
EncodedBody::Segmented((0..5).map(|i| Bytes::from(vec![i as u8; 40])).collect()),
),
];
for (name, body) in cases {
let total = body.len();
let expected = Envelope::data(body.clone().into_contiguous()).encode();
let (head, segments) = Envelope::encode_body_parts(flags::DATA, body, 64);
let declared = u32::from_be_bytes([head[1], head[2], head[3], head[4]]) as usize;
assert_eq!(declared, total, "{name}: header must declare the total");
let emitted: usize =
head.len() - HEADER_SIZE + segments.iter().map(Bytes::len).sum::<usize>();
assert_eq!(
emitted, total,
"{name}: emitted payload bytes must match the declared length"
);
let mut reassembled = BytesMut::from(&head[..]);
for segment in &segments {
assert!(!segment.is_empty(), "{name}: no empty segments");
reassembled.put_slice(segment);
}
assert_eq!(
reassembled.freeze(),
expected,
"{name}: must reassemble to the contiguous envelope"
);
}
}
#[test]
fn test_envelope_roundtrip() {
let original = Envelope::data(Bytes::from_static(b"hello world"));
let encoded = original.encode();
let mut buf = BytesMut::from(&encoded[..]);
let decoded = Envelope::decode(&mut buf).unwrap().unwrap();
assert_eq!(decoded.flags, original.flags);
assert_eq!(decoded.data, original.data);
}
#[test]
fn test_envelope_partial() {
let mut buf = BytesMut::from(&[0u8, 0, 0, 0][..]);
assert!(Envelope::decode(&mut buf).unwrap().is_none());
}
#[test]
fn test_envelope_size_limit() {
let mut buf = BytesMut::new();
buf.put_u8(0); buf.put_u32(1024 * 1024);
let result = Envelope::decode_with_limit(&mut buf, 512 * 1024);
assert!(result.is_err());
let err = result.unwrap_err();
assert_eq!(err.code, crate::error::ErrorCode::ResourceExhausted);
}
#[test]
fn test_envelope_size_limit_ok() {
let original = Envelope::data(Bytes::from_static(b"small"));
let encoded = original.encode();
let mut buf = BytesMut::from(&encoded[..]);
let result = Envelope::decode_with_limit(&mut buf, 1024 * 1024);
assert!(result.is_ok());
assert!(result.unwrap().is_some());
}
#[test]
fn test_envelope_unlimited_decode_huge_length_no_panic() {
let mut buf = BytesMut::new();
buf.put_u8(0); buf.put_u32(u32::MAX); let result = Envelope::decode(&mut buf);
assert!(matches!(result, Ok(None)));
}
#[test]
fn test_decoder_complete_message() {
let mut dec = decoder(1024);
let envelope = Envelope::data(Bytes::from_static(b"hello"));
let mut buf = BytesMut::from(&envelope.encode()[..]);
let result = dec.decode(&mut buf).unwrap();
assert_eq!(result.unwrap(), Bytes::from_static(b"hello"));
assert!(buf.is_empty());
}
#[test]
fn test_decoder_incomplete_header() {
let mut dec = decoder(1024);
let mut buf = BytesMut::from(&[0u8, 0, 0][..]);
assert!(dec.decode(&mut buf).unwrap().is_none());
assert_eq!(buf.len(), 3, "buffer should be untouched");
}
#[test]
fn test_decoder_incomplete_payload() {
let mut dec = decoder(1024);
let mut buf = BytesMut::new();
buf.put_u8(flags::DATA);
buf.put_u32(10);
buf.put_slice(&[1, 2, 3]);
assert!(dec.decode(&mut buf).unwrap().is_none());
assert_eq!(buf.len(), 8, "buffer should be untouched");
}
#[test]
fn test_decoder_end_stream_signals_eof() {
let mut dec = decoder(1024);
let envelope = Envelope::end_stream(Bytes::from_static(b"{}"));
let mut buf = BytesMut::from(&envelope.encode()[..]);
assert!(dec.decode(&mut buf).unwrap().is_none());
assert!(dec.decode(&mut buf).unwrap().is_none());
}
#[test]
fn test_decoder_message_exceeds_size_limit() {
let mut dec = decoder(4); let envelope = Envelope::data(Bytes::from_static(b"too long"));
let mut buf = BytesMut::from(&envelope.encode()[..]);
let err = dec.decode(&mut buf).unwrap_err();
assert_eq!(err.code, crate::error::ErrorCode::ResourceExhausted);
}
#[test]
fn test_decoder_multiple_envelopes_in_buffer() {
let mut dec = decoder(1024);
let e1 = Envelope::data(Bytes::from_static(b"first"));
let e2 = Envelope::data(Bytes::from_static(b"second"));
let mut buf = BytesMut::new();
buf.extend_from_slice(&e1.encode());
buf.extend_from_slice(&e2.encode());
let r1 = dec.decode(&mut buf).unwrap().unwrap();
assert_eq!(r1, Bytes::from_static(b"first"));
let r2 = dec.decode(&mut buf).unwrap().unwrap();
assert_eq!(r2, Bytes::from_static(b"second"));
assert!(buf.is_empty());
}
#[test]
fn test_decoder_data_then_end_stream() {
let mut dec = decoder(1024);
let data_env = Envelope::data(Bytes::from_static(b"msg"));
let end_env = Envelope::end_stream(Bytes::from_static(b"{}"));
let mut buf = BytesMut::new();
buf.extend_from_slice(&data_env.encode());
buf.extend_from_slice(&end_env.encode());
let r1 = dec.decode(&mut buf).unwrap().unwrap();
assert_eq!(r1, Bytes::from_static(b"msg"));
assert!(dec.decode(&mut buf).unwrap().is_none());
}
#[test]
fn test_decode_eof_empty_buffer() {
let mut dec = decoder(1024);
let mut buf = BytesMut::new();
assert!(dec.decode_eof(&mut buf).unwrap().is_none());
}
#[test]
fn test_decode_eof_with_complete_envelope() {
let mut dec = decoder(1024);
let envelope = Envelope::data(Bytes::from_static(b"final"));
let mut buf = BytesMut::from(&envelope.encode()[..]);
let result = dec.decode_eof(&mut buf).unwrap();
assert_eq!(result.unwrap(), Bytes::from_static(b"final"));
}
#[test]
fn test_decode_eof_with_leftover_bytes() {
let mut dec = decoder(1024);
let mut buf = BytesMut::from(&[0u8, 0, 0][..]);
let err = dec.decode_eof(&mut buf).unwrap_err();
assert_eq!(err.code, crate::error::ErrorCode::InvalidArgument);
}
#[test]
fn test_decoder_compressed_without_encoding_header() {
let mut dec = decoder(1024);
let envelope = Envelope::compressed(Bytes::from_static(b"data"));
let mut buf = BytesMut::from(&envelope.encode()[..]);
let err = dec.decode(&mut buf).unwrap_err();
assert_eq!(err.code, crate::error::ErrorCode::Internal);
}
#[test]
fn test_encoder_uncompressed() {
let mut enc = EnvelopeEncoder::uncompressed();
let mut buf = BytesMut::new();
enc.encode(Bytes::from_static(b"hello"), &mut buf).unwrap();
assert_eq!(buf.len(), HEADER_SIZE + 5);
assert_eq!(buf[0], flags::DATA);
assert_eq!(u32::from_be_bytes([buf[1], buf[2], buf[3], buf[4]]), 5);
assert_eq!(&buf[HEADER_SIZE..], b"hello");
}
#[test]
#[cfg(feature = "gzip")]
fn test_encoder_empty_payload_skips_compression() {
let registry = Arc::new(CompressionRegistry::default());
let mut enc = EnvelopeEncoder::new(Some((registry, "gzip")), CompressionPolicy::default());
let mut buf = BytesMut::new();
enc.encode(Bytes::new(), &mut buf).unwrap();
assert_eq!(buf[0], flags::DATA, "empty payload should use DATA flag");
assert_eq!(u32::from_be_bytes([buf[1], buf[2], buf[3], buf[4]]), 0);
}
#[test]
#[cfg(feature = "gzip")]
fn test_encoder_with_compression() {
let registry = Arc::new(CompressionRegistry::default());
let mut enc = EnvelopeEncoder::new(
Some((registry, "gzip")),
CompressionPolicy::default().with_min_size(0),
);
let mut buf = BytesMut::new();
enc.encode(Bytes::from_static(b"compress me"), &mut buf)
.unwrap();
assert_eq!(buf[0], flags::COMPRESSED, "should use COMPRESSED flag");
let payload_len = u32::from_be_bytes([buf[1], buf[2], buf[3], buf[4]]) as usize;
assert!(payload_len > 0);
assert_eq!(buf.len(), HEADER_SIZE + payload_len);
}
#[test]
fn test_encoder_end_stream() {
let mut enc = EnvelopeEncoder::uncompressed();
let mut buf = BytesMut::new();
enc.encode_end_stream(Bytes::from_static(b"{}"), &mut buf)
.unwrap();
assert_eq!(buf[0], flags::END_STREAM);
assert_eq!(u32::from_be_bytes([buf[1], buf[2], buf[3], buf[4]]), 2);
assert_eq!(&buf[HEADER_SIZE..], b"{}");
}
#[test]
#[cfg(feature = "gzip")]
fn test_encoder_decoder_roundtrip() {
let registry = Arc::new(CompressionRegistry::default());
let mut enc = EnvelopeEncoder::new(
Some((Arc::clone(®istry), "gzip")),
CompressionPolicy::default(),
);
let mut dec = EnvelopeDecoder::new(1024, Some("gzip".to_owned()), registry);
let original = Bytes::from_static(b"roundtrip test data");
let mut buf = BytesMut::new();
enc.encode(original.clone(), &mut buf).unwrap();
let decoded = dec.decode(&mut buf).unwrap().unwrap();
assert_eq!(decoded, original);
assert!(buf.is_empty());
}
#[test]
fn test_encoder_multiple_messages() {
let mut enc = EnvelopeEncoder::uncompressed();
let mut buf = BytesMut::new();
enc.encode(Bytes::from_static(b"one"), &mut buf).unwrap();
enc.encode(Bytes::from_static(b"two"), &mut buf).unwrap();
assert_eq!(buf.len(), 2 * HEADER_SIZE + 3 + 3);
let mut dec = decoder(1024);
let r1 = dec.decode(&mut buf).unwrap().unwrap();
assert_eq!(r1, Bytes::from_static(b"one"));
let r2 = dec.decode(&mut buf).unwrap().unwrap();
assert_eq!(r2, Bytes::from_static(b"two"));
assert!(buf.is_empty());
}
}