use super::{Packet, PacketCodec, PacketHeader, HEADER_BYTES};
use crate::Error;
use asynchronous_codec::Decoder;
use bytes::{Buf, BytesMut};
use tracing::{event, Level};
pub trait Decode<B: Buf> {
fn decode(src: &mut B) -> crate::Result<Self>
where
Self: Sized;
}
impl Decoder for PacketCodec {
type Item = Packet;
type Error = Error;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
if src.len() < HEADER_BYTES {
src.reserve(HEADER_BYTES);
return Ok(None);
}
let header = PacketHeader::decode(&mut BytesMut::from(&src[0..HEADER_BYTES]))?;
let length = header.length() as usize;
if src.len() < length {
src.reserve(length);
return Ok(None);
}
event!(
Level::TRACE,
"Reading a {:?} ({} bytes)",
header.r#type(),
length,
);
let header = PacketHeader::decode(src)?;
if length < HEADER_BYTES {
return Err(Error::Protocol("Invalid packet length".into()));
}
let payload = src.split_to(length - HEADER_BYTES);
Ok(Some(Packet::new(header, payload)))
}
fn decode_eof(&mut self, buf: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
match self.decode(buf)? {
Some(frame) => Ok(Some(frame)),
None => {
if buf.is_empty() {
Ok(None)
} else {
Err(std::io::Error::other("bytes remaining on stream").into())
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tds::codec::{Encode, PacketHeader, PacketType};
fn full_packet_bytes() -> BytesMut {
let payload = BytesMut::from(&b"hello world"[..]);
let packet = Packet::new(PacketHeader::batch(1), payload);
let mut buf = BytesMut::new();
packet.encode(&mut buf).unwrap();
assert_eq!(buf.len(), 19);
buf
}
#[test]
fn decode_partial_header_returns_none() {
let mut src = BytesMut::from(&full_packet_bytes()[0..4]);
let mut codec = PacketCodec;
let out = codec.decode(&mut src).unwrap();
assert!(out.is_none());
}
#[test]
fn decode_complete_packet_returns_some() {
let mut src = full_packet_bytes();
let mut codec = PacketCodec;
let packet = codec
.decode(&mut src)
.unwrap()
.expect("a complete packet must decode to Some");
assert_eq!(packet.header.r#type() as u8, PacketType::SQLBatch as u8);
assert_eq!(&packet.payload[..], b"hello world");
assert_eq!(packet.payload.len(), 11);
assert!(src.is_empty());
}
#[test]
fn decode_incomplete_body_returns_none() {
let mut src = BytesMut::from(&full_packet_bytes()[0..11]);
let mut codec = PacketCodec;
let out = codec.decode(&mut src).unwrap();
assert!(out.is_none());
}
#[test]
fn decode_minimal_packet_exactly_header_bytes() {
let packet = Packet::new(PacketHeader::attention(1), BytesMut::new());
let mut src = BytesMut::new();
packet.encode(&mut src).unwrap();
assert_eq!(src.len(), 8);
let mut codec = PacketCodec;
let packet = codec
.decode(&mut src)
.unwrap()
.expect("an 8-byte packet must decode to Some");
assert_eq!(
packet.header.r#type() as u8,
PacketType::AttentionSignal as u8
);
assert!(packet.payload.is_empty());
}
#[test]
fn decode_rejects_length_below_header() {
let packet = Packet::new(PacketHeader::attention(1), BytesMut::new());
let mut src = BytesMut::new();
packet.encode(&mut src).unwrap();
src[2] = 0;
src[3] = 5;
let mut codec = PacketCodec;
let err = codec
.decode(&mut src)
.expect_err("length below header must error");
assert!(matches!(err, Error::Protocol(_)));
}
#[test]
fn decode_eof_returns_some_for_complete_packet() {
let mut src = full_packet_bytes();
let mut codec = PacketCodec;
let packet = codec
.decode_eof(&mut src)
.unwrap()
.expect("decode_eof must yield a complete packet");
assert_eq!(packet.header.r#type() as u8, PacketType::SQLBatch as u8);
assert_eq!(&packet.payload[..], b"hello world");
}
#[test]
fn decode_eof_errors_on_trailing_partial_bytes() {
let mut src = BytesMut::from(&full_packet_bytes()[0..4]);
let mut codec = PacketCodec;
assert!(codec.decode_eof(&mut src).is_err());
}
}