use std::convert::TryInto;
use std::error::Error as StdError;
use std::fmt;
use std::io::{self, Cursor};
use bytes::{Buf, BufMut, BytesMut};
use tokio::io::{AsyncRead, AsyncWrite};
use tokio_util::codec::{Decoder, Encoder, Framed, FramedRead, FramedWrite};
use flate2::Compress;
use flate2::Compression;
use flate2::Decompress;
use flate2::FlushCompress;
use flate2::FlushDecompress;
#[derive(Debug, Clone, Copy)]
pub struct Builder {
compression: bool,
compression_level: Compression,
max_frame_len: usize,
}
pub struct QuasselCodecError {
_priv: (),
}
#[derive(Debug)]
pub struct QuasselCodec {
builder: Builder,
state: DecodeState,
comp: Compress,
decomp: Decompress,
}
#[derive(Debug, Clone, Copy)]
enum DecodeState {
Head,
Data(usize),
}
impl QuasselCodec {
pub fn new() -> Self {
Self {
builder: Builder::new(),
state: DecodeState::Head,
comp: Compress::new(Compression::default(), true),
decomp: Decompress::new(true),
}
}
pub fn builder() -> Builder {
Builder::new()
}
pub fn max_frame_length(&self) -> usize {
self.builder.max_frame_len
}
pub fn compression(&self) -> bool {
self.builder.compression
}
pub fn compression_level(&self) -> Compression {
self.builder.compression_level
}
pub fn set_max_frame_length(&mut self, val: usize) {
self.builder.max_frame_length(val);
}
pub fn set_compression(&mut self, val: bool) {
self.builder.compression(val);
}
pub fn set_compression_level(&mut self, val: Compression) {
self.builder.compression_level(val);
}
fn decode_head(&mut self, src: &mut BytesMut) -> io::Result<Option<usize>> {
let head_len = 4;
if src.len() < head_len {
return Ok(None);
}
let field_len = {
let mut src = Cursor::new(&mut *src);
let field_len = src.get_uint(head_len);
if field_len > self.builder.max_frame_len as u64 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
QuasselCodecError { _priv: () },
));
}
field_len as usize
};
let _ = src.split_to(head_len);
src.reserve(field_len);
Ok(Some(field_len))
}
fn decode_data(&self, n: usize, src: &mut BytesMut) -> io::Result<Option<BytesMut>> {
if src.len() < n {
return Ok(None);
}
Ok(Some(src.split_to(n)))
}
}
impl Decoder for QuasselCodec {
type Item = BytesMut;
type Error = io::Error;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<BytesMut>, io::Error> {
let mut buf: &mut BytesMut = &mut BytesMut::new();
if self.builder.compression == true {
let mut msg = Vec::with_capacity(self.builder.max_frame_len);
let before_in = self.decomp.total_in();
let before_out = self.decomp.total_out();
self.decomp
.decompress_vec(&src, &mut msg, FlushDecompress::None)?;
src.clear();
let after_in = self.decomp.total_in();
let after_out = self.decomp.total_out();
let len = (after_out - before_out).try_into().unwrap();
buf.reserve(len);
buf.put(&msg[..]);
} else {
buf = src;
}
let n = match self.state {
DecodeState::Head => match self.decode_head(buf)? {
Some(n) => {
self.state = DecodeState::Data(n);
n
}
None => return Ok(None),
},
DecodeState::Data(n) => n,
};
match self.decode_data(n, buf)? {
Some(data) => {
self.state = DecodeState::Head;
buf.reserve(4);
Ok(Some(data))
}
None => Ok(None),
}
}
}
impl Encoder for QuasselCodec {
type Item = Vec<u8>;
type Error = io::Error;
fn encode(&mut self, data: Vec<u8>, dst: &mut BytesMut) -> Result<(), io::Error> {
let buf = &mut BytesMut::new();
let n = (&data).len();
if n > self.builder.max_frame_len {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
QuasselCodecError { _priv: () },
));
}
buf.reserve(4 + n);
buf.put_uint(n as u64, 4);
buf.extend_from_slice(&data[..]);
if self.builder.compression {
let mut cbuf: Vec<u8> = vec![0; 4 + n];
let before_in = self.comp.total_in();
let before_out = self.comp.total_out();
self.comp.compress(buf, &mut cbuf, FlushCompress::Full)?;
let after_in = self.comp.total_in();
let after_out = self.comp.total_out();
cbuf.truncate((after_out - before_out).try_into().unwrap());
*dst = BytesMut::from(&cbuf[..]);
} else {
*dst = buf.clone();
}
Ok(())
}
}
impl Default for QuasselCodec {
fn default() -> Self {
Self::new()
}
}
impl Builder {
pub fn new() -> Builder {
Builder {
compression: false,
compression_level: Compression::default(),
max_frame_len: 64 * 1024 * 1024,
}
}
pub fn compression(&mut self, val: bool) -> &mut Self {
self.compression = val;
self
}
pub fn compression_level(&mut self, val: Compression) -> &mut Self {
self.compression_level = val;
self
}
pub fn max_frame_length(&mut self, val: usize) -> &mut Self {
self.max_frame_len = val;
self
}
pub fn new_codec(&self) -> QuasselCodec {
QuasselCodec {
builder: *self,
state: DecodeState::Head,
comp: Compress::new(self.compression_level, true),
decomp: Decompress::new(true),
}
}
pub fn new_read<T>(&self, upstream: T) -> FramedRead<T, QuasselCodec>
where
T: AsyncRead,
{
FramedRead::new(upstream, self.new_codec())
}
pub fn new_write<T>(&self, inner: T) -> FramedWrite<T, QuasselCodec>
where
T: AsyncWrite,
{
FramedWrite::new(inner, self.new_codec())
}
pub fn new_framed<T>(&self, inner: T) -> Framed<T, QuasselCodec>
where
T: AsyncRead + AsyncWrite,
{
Framed::new(inner, self.new_codec())
}
}
impl fmt::Debug for QuasselCodecError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("QuasselCodecError").finish()
}
}
impl fmt::Display for QuasselCodecError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("frame size too big")
}
}
impl StdError for QuasselCodecError {}