use bytes::{Buf, BufMut, Bytes, BytesMut};
use tokio_util::codec::{Decoder, Encoder};
use super::TransportError;
pub(crate) const DEFAULT_MAX_SIZE: usize = 16 * 1024 * 1024;
#[derive(Debug)]
pub struct ContentLengthCodec {
state: State,
max_size: usize,
}
#[derive(Debug)]
enum State {
Headers,
Body { length: usize },
}
impl Default for ContentLengthCodec {
fn default() -> Self {
Self::new(DEFAULT_MAX_SIZE)
}
}
impl ContentLengthCodec {
pub fn new(max_size: usize) -> Self {
Self {
state: State::Headers,
max_size,
}
}
pub(crate) fn validate_size(&self, length: usize) -> Result<(), TransportError> {
if length > self.max_size {
return Err(TransportError::OversizedMessage {
length,
limit: self.max_size,
});
}
Ok(())
}
pub(crate) fn header_for(&self, length: usize) -> Result<String, TransportError> {
self.validate_size(length)?;
Ok(format!("Content-Length: {length}\r\n\r\n"))
}
}
impl Decoder for ContentLengthCodec {
type Item = Bytes;
type Error = TransportError;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Bytes>, TransportError> {
loop {
match self.state {
State::Headers => {
let Some(headers_end) = find_header_end(src) else {
return Ok(None);
};
let headers = src.split_to(headers_end);
src.advance(4);
let length = parse_content_length(&headers)?;
self.validate_size(length)?;
self.state = State::Body { length };
}
State::Body { length } => {
if src.len() < length {
src.reserve(length - src.len());
return Ok(None);
}
let body = src.split_to(length).freeze();
self.state = State::Headers;
return Ok(Some(body));
}
}
}
}
}
impl Encoder<Bytes> for ContentLengthCodec {
type Error = TransportError;
fn encode(&mut self, item: Bytes, dst: &mut BytesMut) -> Result<(), TransportError> {
let header = self.header_for(item.len())?;
dst.reserve(header.len() + item.len());
dst.put_slice(header.as_bytes());
dst.put_slice(&item);
Ok(())
}
}
fn find_header_end(src: &[u8]) -> Option<usize> {
src.windows(4).position(|w| w == b"\r\n\r\n")
}
fn parse_content_length(headers: &[u8]) -> Result<usize, TransportError> {
let s = std::str::from_utf8(headers)
.map_err(|e| TransportError::Malformed(format!("non-UTF-8 headers: {e}")))?;
for line in s.split("\r\n") {
let Some((name, value)) = line.split_once(':') else {
continue;
};
if name.eq_ignore_ascii_case("content-length") {
return value
.trim()
.parse()
.map_err(|e| TransportError::Malformed(format!("invalid Content-Length: {e}")));
}
}
Err(TransportError::Malformed(
"missing Content-Length header".into(),
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decodes_a_single_message() {
let mut codec = ContentLengthCodec::default();
let mut buf = BytesMut::from(&b"Content-Length: 17\r\n\r\n{\"jsonrpc\":\"2.0\"}"[..]);
let body = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(&body[..], br#"{"jsonrpc":"2.0"}"#);
assert!(buf.is_empty());
}
#[test]
fn waits_for_full_body() {
let mut codec = ContentLengthCodec::default();
let mut buf = BytesMut::from(&b"Content-Length: 17\r\n\r\n{\"jsonrpc\""[..]);
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.extend_from_slice(b":\"2.0\"}");
let body = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(&body[..], br#"{"jsonrpc":"2.0"}"#);
}
#[test]
fn decodes_a_message_split_at_every_byte_boundary() {
let mut codec = ContentLengthCodec::default();
let mut buf = BytesMut::new();
let frame = b"Content-Length: 17\r\n\r\n{\"jsonrpc\":\"2.0\"}";
for (index, byte) in frame.iter().enumerate() {
buf.extend_from_slice(std::slice::from_ref(byte));
let decoded = codec.decode(&mut buf).unwrap();
if index + 1 == frame.len() {
assert_eq!(decoded.as_deref(), Some(&br#"{"jsonrpc":"2.0"}"#[..]));
} else {
assert!(decoded.is_none());
}
}
}
#[test]
fn decodes_consecutive_messages_without_losing_buffered_bytes() {
let mut codec = ContentLengthCodec::default();
let mut buf =
BytesMut::from(&b"Content-Length: 3\r\n\r\noneContent-Length: 3\r\n\r\ntwo"[..]);
assert_eq!(
codec.decode(&mut buf).unwrap().as_deref(),
Some(&b"one"[..])
);
assert_eq!(
codec.decode(&mut buf).unwrap().as_deref(),
Some(&b"two"[..])
);
assert!(buf.is_empty());
}
#[test]
fn ignores_extra_headers() {
let mut codec = ContentLengthCodec::default();
let mut buf = BytesMut::from(
&b"Content-Type: application/vscode-jsonrpc; charset=utf-8\r\nContent-Length: 17\r\n\r\n{\"jsonrpc\":\"2.0\"}"[..],
);
let body = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(&body[..], br#"{"jsonrpc":"2.0"}"#);
}
#[test]
fn rejects_a_message_without_content_length() {
let mut codec = ContentLengthCodec::default();
let mut buf = BytesMut::from(&b"Not-Content-Length: 17\r\n\r\n"[..]);
let error = codec.decode(&mut buf).unwrap_err();
assert!(
matches!(error, TransportError::Malformed(message) if message == "missing Content-Length header")
);
}
#[test]
fn rejects_a_non_numeric_content_length() {
let mut codec = ContentLengthCodec::default();
let mut buf = BytesMut::from(&b"Content-Length: NaN\r\n\r\n"[..]);
let error = codec.decode(&mut buf).unwrap_err();
assert!(
matches!(error, TransportError::Malformed(message) if message.starts_with("invalid Content-Length:"))
);
}
#[test]
fn rejects_oversized() {
let mut codec = ContentLengthCodec::new(8);
let mut buf = BytesMut::from(&b"Content-Length: 17\r\n\r\n"[..]);
let err = codec.decode(&mut buf).unwrap_err();
assert!(matches!(err, TransportError::OversizedMessage { .. }));
}
#[test]
fn rejects_oversized_on_encode() {
let mut codec = ContentLengthCodec::new(8);
let mut buf = BytesMut::new();
let err = codec
.encode(Bytes::from_static(b"123456789"), &mut buf)
.unwrap_err();
assert!(matches!(
err,
TransportError::OversizedMessage {
length: 9,
limit: 8
}
));
assert!(buf.is_empty(), "an oversized frame writes no partial bytes");
}
#[test]
fn default_limit_is_sixteen_mib_on_receive() {
let mut codec = ContentLengthCodec::default();
let mut buf =
BytesMut::from(format!("Content-Length: {}\r\n\r\n", DEFAULT_MAX_SIZE + 1).as_bytes());
let err = codec.decode(&mut buf).unwrap_err();
assert!(matches!(
err,
TransportError::OversizedMessage { length, limit }
if length == DEFAULT_MAX_SIZE + 1 && limit == DEFAULT_MAX_SIZE
));
}
#[test]
fn encodes_with_header() {
let mut codec = ContentLengthCodec::default();
let mut buf = BytesMut::new();
codec
.encode(Bytes::from_static(br#"{"jsonrpc":"2.0"}"#), &mut buf)
.unwrap();
assert_eq!(&buf[..], b"Content-Length: 17\r\n\r\n{\"jsonrpc\":\"2.0\"}");
}
}