use super::{Decode, Encode};
use crate::Error;
use bytes::{Buf, BufMut, BytesMut};
use std::convert::TryFrom;
uint_enum! {
#[repr(u8)]
pub enum PacketType {
SQLBatch = 1,
PreTDSv7Login = 2,
Rpc = 3,
TabularResult = 4,
AttentionSignal = 6,
BulkLoad = 7,
Fat = 8,
TransactionManagerReq = 14,
TDSv7Login = 16,
Sspi = 17,
PreLogin = 18,
}
}
uint_enum! {
#[repr(u8)]
pub enum PacketStatus {
NormalMessage = 0,
EndOfMessage = 1,
IgnoreEvent = 3,
ResetConnection = 0x08,
ResetConnectionSkipTran = 0x10,
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct PacketHeader {
ty: PacketType,
status: PacketStatus,
length: u16,
spid: u16,
id: u8,
window: u8,
}
impl PacketHeader {
pub fn new(length: usize, id: u8) -> PacketHeader {
assert!(length <= u16::MAX as usize);
PacketHeader {
ty: PacketType::TDSv7Login,
status: PacketStatus::ResetConnection,
length: length as u16,
spid: 0,
id,
window: 0,
}
}
pub fn rpc(id: u8) -> Self {
Self {
ty: PacketType::Rpc,
status: PacketStatus::NormalMessage,
..Self::new(0, id)
}
}
pub fn pre_login(id: u8) -> Self {
Self {
ty: PacketType::PreLogin,
status: PacketStatus::EndOfMessage,
..Self::new(0, id)
}
}
pub fn login(id: u8) -> Self {
Self {
ty: PacketType::TDSv7Login,
status: PacketStatus::EndOfMessage,
..Self::new(0, id)
}
}
#[allow(dead_code)]
pub fn sspi(id: u8) -> Self {
Self {
ty: PacketType::Sspi,
status: PacketStatus::EndOfMessage,
..Self::new(0, id)
}
}
pub fn batch(id: u8) -> Self {
Self {
ty: PacketType::SQLBatch,
status: PacketStatus::NormalMessage,
..Self::new(0, id)
}
}
pub fn bulk_load(id: u8) -> Self {
Self {
ty: PacketType::BulkLoad,
status: PacketStatus::NormalMessage,
..Self::new(0, id)
}
}
pub fn attention(id: u8) -> Self {
Self {
ty: PacketType::AttentionSignal,
status: PacketStatus::EndOfMessage,
..Self::new(0, id)
}
}
pub fn transaction_manager(id: u8) -> Self {
Self {
ty: PacketType::TransactionManagerReq,
status: PacketStatus::EndOfMessage,
..Self::new(0, id)
}
}
pub fn set_status(&mut self, status: PacketStatus) {
self.status = status;
}
#[cfg(any(
feature = "rustls",
feature = "native-tls",
feature = "vendored-openssl"
))]
pub fn set_type(&mut self, ty: PacketType) {
self.ty = ty;
}
pub fn status(&self) -> PacketStatus {
self.status
}
pub fn r#type(&self) -> PacketType {
self.ty
}
pub fn length(&self) -> u16 {
self.length
}
}
impl<B> Encode<B> for PacketHeader
where
B: BufMut,
{
fn encode(self, dst: &mut B) -> crate::Result<()> {
dst.put_u8(self.ty as u8);
dst.put_u8(self.status as u8);
dst.put_u16(self.length);
dst.put_u16(self.spid);
dst.put_u8(self.id);
dst.put_u8(self.window);
Ok(())
}
}
impl Decode<BytesMut> for PacketHeader {
fn decode(src: &mut BytesMut) -> crate::Result<Self>
where
Self: Sized,
{
let raw_ty = src.get_u8();
let ty = PacketType::try_from(raw_ty).map_err(|_| {
Error::Protocol(format!("header: invalid packet type: {}", raw_ty).into())
})?;
let status = PacketStatus::try_from(src.get_u8())
.map_err(|_| Error::Protocol("header: invalid packet status".into()))?;
let header = PacketHeader {
ty,
status,
length: src.get_u16(),
spid: src.get_u16(),
id: src.get_u8(),
window: src.get_u8(),
};
Ok(header)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tds::codec::Packet;
#[test]
fn attention_packet_header_fields() {
let header = PacketHeader::attention(42);
assert_eq!(header.r#type() as u8, PacketType::AttentionSignal as u8);
assert_eq!(header.r#type() as u8, 0x06);
assert_eq!(header.status(), PacketStatus::EndOfMessage);
}
#[test]
fn attention_packet_header_encodes_to_eight_bytes() {
let header = PacketHeader::attention(42);
let mut buf = BytesMut::new();
header.encode(&mut buf).unwrap();
assert_eq!(buf.len(), 8);
assert_eq!(
&buf[..],
&[
0x06, 0x01, 0x00, 0x00, 0x00, 0x00, 42, 0x00, ]
);
}
#[test]
fn attention_packet_encodes_with_length_of_header() {
let packet = Packet::new(PacketHeader::attention(1), BytesMut::new());
let mut buf = BytesMut::new();
packet.encode(&mut buf).unwrap();
assert_eq!(&buf[..], &[0x06, 0x01, 0x00, 0x08, 0x00, 0x00, 0x01, 0x00]);
}
#[test]
fn new_sets_login_type_and_reset_connection_status() {
let header = PacketHeader::new(123, 5);
assert_eq!(header.r#type() as u8, PacketType::TDSv7Login as u8);
assert_eq!(header.status(), PacketStatus::ResetConnection);
assert_eq!(header.length(), 123);
}
#[test]
#[should_panic]
fn new_panics_on_length_overflow() {
PacketHeader::new(usize::from(u16::MAX) + 1, 0);
}
#[test]
fn rpc_header_type_and_status() {
let header = PacketHeader::rpc(7);
assert_eq!(header.r#type() as u8, PacketType::Rpc as u8);
assert_eq!(header.status(), PacketStatus::NormalMessage);
}
#[test]
fn pre_login_header_type_and_status() {
let header = PacketHeader::pre_login(7);
assert_eq!(header.r#type() as u8, PacketType::PreLogin as u8);
assert_eq!(header.status(), PacketStatus::EndOfMessage);
}
#[test]
fn login_header_type_and_status() {
let header = PacketHeader::login(7);
assert_eq!(header.r#type() as u8, PacketType::TDSv7Login as u8);
assert_eq!(header.status(), PacketStatus::EndOfMessage);
}
#[test]
fn batch_header_type_and_status() {
let header = PacketHeader::batch(7);
assert_eq!(header.r#type() as u8, PacketType::SQLBatch as u8);
assert_eq!(header.status(), PacketStatus::NormalMessage);
}
#[test]
fn bulk_load_header_type_and_status() {
let header = PacketHeader::bulk_load(7);
assert_eq!(header.r#type() as u8, PacketType::BulkLoad as u8);
assert_eq!(header.status(), PacketStatus::NormalMessage);
}
#[test]
fn transaction_manager_header_type_and_status() {
let header = PacketHeader::transaction_manager(7);
assert_eq!(
header.r#type() as u8,
PacketType::TransactionManagerReq as u8
);
assert_eq!(header.status(), PacketStatus::EndOfMessage);
}
#[test]
fn set_status_mutates_header() {
let mut header = PacketHeader::batch(1);
assert_eq!(header.status(), PacketStatus::NormalMessage);
header.set_status(PacketStatus::IgnoreEvent);
assert_eq!(header.status(), PacketStatus::IgnoreEvent);
}
#[test]
fn decode_round_trips_header_fields() {
let header = PacketHeader::rpc(9);
let mut buf = BytesMut::new();
header.encode(&mut buf).unwrap();
let decoded = PacketHeader::decode(&mut buf).unwrap();
assert_eq!(decoded.r#type() as u8, PacketType::Rpc as u8);
assert_eq!(decoded.status(), PacketStatus::NormalMessage);
assert_eq!(decoded.length(), 0);
}
#[test]
fn decode_invalid_packet_type_errors() {
let mut buf = BytesMut::from(&[0xffu8, 0x01, 0x00, 0x08, 0x00, 0x00, 0x01, 0x00][..]);
let err = PacketHeader::decode(&mut buf).unwrap_err();
assert!(format!("{}", err).contains("invalid packet type"));
}
#[test]
fn decode_invalid_packet_status_errors() {
let mut buf = BytesMut::from(&[0x01u8, 0xff, 0x00, 0x08, 0x00, 0x00, 0x01, 0x00][..]);
let err = PacketHeader::decode(&mut buf).unwrap_err();
assert!(format!("{}", err).contains("invalid packet status"));
}
}