use std::io;
use bytes::{Buf, BufMut, Bytes, BytesMut};
use memchr::memchr;
use super::{Decoder, Encoder};
#[derive(Debug, Copy, Clone)]
#[non_exhaustive]
pub struct LinesCodec {
max_length: usize,
next_index: usize,
}
impl LinesCodec {
pub const fn new() -> Self {
Self {
max_length: usize::MAX,
next_index: 0,
}
}
pub const fn new_with_max_length(max_length: usize) -> Self {
Self {
max_length,
next_index: 0,
}
}
pub const fn max_length(&self) -> usize {
self.max_length
}
}
impl Default for LinesCodec {
fn default() -> Self {
Self::new()
}
}
impl<T: AsRef<str>> Encoder<T> for LinesCodec {
type Error = io::Error;
#[inline]
fn encode(&mut self, item: T, dst: &mut BytesMut) -> Result<(), Self::Error> {
let item = item.as_ref();
dst.reserve(item.len() + 1);
dst.put_slice(item.as_bytes());
dst.put_u8(b'\n');
Ok(())
}
}
impl Decoder for LinesCodec {
type Item = String;
type Error = io::Error;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
if src.is_empty() {
self.next_index = 0;
return Ok(None);
}
let start = if self.next_index < src.len() {
self.next_index
} else {
0
};
let len = match memchr(b'\n', &src[start..]) {
Some(n) => start + n,
None => {
let max = self.max_length;
if max != usize::MAX {
let max_cr = max.saturating_add(1);
if src.len() > max && !(src.len() == max_cr && src.last() == Some(&b'\r')) {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"max line length exceeded",
));
}
}
self.next_index = src.len();
return Ok(None);
}
};
let max = self.max_length;
if max != usize::MAX {
let max_cr = max.saturating_add(1);
if len > max && !(len == max_cr && src.get(len - 1) == Some(&b'\r')) {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"max line length exceeded",
));
}
}
self.next_index = 0;
let mut buf = src.split_to(len);
debug_assert_eq!(len, buf.len());
src.advance(1);
match buf.last() {
Some(b'\r') => buf.truncate(len - 1),
None => return Ok(Some(String::new())),
_ => {}
}
try_into_utf8(buf.freeze())
}
fn decode_eof(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
match self.decode(src)? {
Some(frame) => Ok(Some(frame)),
None if src.is_empty() => Ok(None),
None => {
self.next_index = 0;
let buf = match src.last() {
Some(b'\r') => src.split_to(src.len() - 1),
_ => src.split(),
};
if buf.is_empty() {
return Ok(None);
}
try_into_utf8(buf.freeze())
}
}
}
}
fn try_into_utf8(buf: Bytes) -> io::Result<Option<String>> {
String::from_utf8(buf.to_vec())
.map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err))
.map(Some)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn lines_decoder() {
let mut codec = LinesCodec::default();
let mut buf = BytesMut::from("\nline 1\nline 2\r\nline 3\n\r\n\r");
assert_eq!("", codec.decode(&mut buf).unwrap().unwrap());
assert_eq!("line 1", codec.decode(&mut buf).unwrap().unwrap());
assert_eq!("line 2", codec.decode(&mut buf).unwrap().unwrap());
assert_eq!("line 3", codec.decode(&mut buf).unwrap().unwrap());
assert_eq!("", codec.decode(&mut buf).unwrap().unwrap());
assert!(codec.decode(&mut buf).unwrap().is_none());
assert!(codec.decode_eof(&mut buf).unwrap().is_none());
buf.put_slice(b"k");
assert!(codec.decode(&mut buf).unwrap().is_none());
assert_eq!("\rk", codec.decode_eof(&mut buf).unwrap().unwrap());
assert!(codec.decode(&mut buf).unwrap().is_none());
assert!(codec.decode_eof(&mut buf).unwrap().is_none());
}
#[test]
fn lines_encoder() {
let mut codec = LinesCodec::default();
let mut buf = BytesMut::new();
codec.encode("", &mut buf).unwrap();
assert_eq!(&buf[..], b"\n");
codec.encode("test", &mut buf).unwrap();
assert_eq!(&buf[..], b"\ntest\n");
codec.encode("a\nb", &mut buf).unwrap();
assert_eq!(&buf[..], b"\ntest\na\nb\n");
}
#[test]
fn lines_encoder_no_overflow() {
let mut codec = LinesCodec::default();
let mut buf = BytesMut::new();
codec.encode("1234567", &mut buf).unwrap();
assert_eq!(&buf[..], b"1234567\n");
let mut buf = BytesMut::new();
codec.encode("12345678", &mut buf).unwrap();
assert_eq!(&buf[..], b"12345678\n");
let mut buf = BytesMut::new();
codec.encode("123456789111213", &mut buf).unwrap();
assert_eq!(&buf[..], b"123456789111213\n");
let mut buf = BytesMut::new();
codec.encode("1234567891112131", &mut buf).unwrap();
assert_eq!(&buf[..], b"1234567891112131\n");
}
#[test]
fn lines_decoder_errors_on_overlong_line_without_delimiter() {
let mut codec = LinesCodec::new_with_max_length(4);
let mut buf = BytesMut::from(&b"aaaaa"[..]);
let err = codec.decode(&mut buf).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn lines_decoder_resumes_from_previous_search() {
let mut codec = LinesCodec::default();
let mut buf = BytesMut::from(&b"partial"[..]);
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.put_slice(b" line\n");
assert_eq!("partial line", codec.decode(&mut buf).unwrap().unwrap());
assert!(codec.decode(&mut buf).unwrap().is_none());
}
#[test]
fn lines_decoder_resumes_across_multiple_partial_chunks() {
let mut codec = LinesCodec::default();
let mut buf = BytesMut::new();
buf.put_slice(b"partial");
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.put_slice(b" line");
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.put_slice(b" across chunks");
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.put_slice(b"\n");
assert_eq!(
"partial line across chunks",
codec.decode(&mut buf).unwrap().unwrap()
);
}
#[test]
fn lines_decoder_resumes_with_max_length() {
let mut codec = LinesCodec::new_with_max_length(18);
let mut buf = BytesMut::new();
buf.put_slice(b"partial");
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.put_slice(b" line");
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.put_slice(b" ok\n");
assert_eq!("partial line ok", codec.decode(&mut buf).unwrap().unwrap());
}
#[test]
fn lines_decoder_errors_on_overlong_partial_line() {
let mut codec = LinesCodec::new_with_max_length(4);
let mut buf = BytesMut::from(&b"aa"[..]);
assert!(codec.decode(&mut buf).unwrap().is_none());
buf.put_slice(b"aaa");
let err = codec.decode(&mut buf).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn lines_decoder_resets_search_after_decode_eof() {
let mut codec = LinesCodec::default();
let mut buf = BytesMut::from(&b"partial"[..]);
assert!(codec.decode(&mut buf).unwrap().is_none());
assert_eq!("partial", codec.decode_eof(&mut buf).unwrap().unwrap());
buf.put_slice(b"next\n");
assert_eq!("next", codec.decode(&mut buf).unwrap().unwrap());
}
}