use byteorder::{BigEndian, ByteOrder, LittleEndian};
use tokio::io::{AsyncReadExt, AsyncWriteExt, DuplexStream, duplex};
use tokio::sync::mpsc::{UnboundedReceiver, unbounded_channel};
use crate::connection::client_context::ClientContext;
use crate::connection::transport::network_transport::NetworkTransport;
use crate::connection::transport::ssl_handler::SslHandler;
use crate::message::messages::{PacketStatusFlags, PacketType};
macro_rules! append_method {
($name:ident, $type:ty, $size:expr_2021, $write_fn:ident) => {
pub(crate) fn $name(&mut self, number: $type) -> &mut TestPacketBuilder {
let mut buffer = [0u8; $size];
LittleEndian::$write_fn(&mut buffer, number);
self.data.extend_from_slice(&buffer);
self
}
};
}
pub(crate) struct TestPacketBuilder {
data: Vec<u8>,
}
impl TestPacketBuilder {
pub(crate) fn new(packet_type: PacketType) -> TestPacketBuilder {
let mut data: Vec<u8> = vec![0; 8];
data[1] = 0x1;
data[0] = packet_type as u8;
TestPacketBuilder { data }
}
pub(crate) fn append_byte(&mut self, byte: u8) -> &mut TestPacketBuilder {
self.data.push(byte);
self
}
pub(crate) fn append_bytes(&mut self, bytes: &[u8]) -> &mut TestPacketBuilder {
self.data.extend_from_slice(bytes);
self
}
pub(crate) fn continuation(&mut self) -> &mut TestPacketBuilder {
self.data[1] &= !(PacketStatusFlags::Eom as u8);
self
}
append_method!(append_u16, u16, 2, write_u16);
append_method!(append_i16, i16, 2, write_i16);
append_method!(append_f32, f32, 4, write_f32);
append_method!(append_f64, f64, 8, write_f64);
append_method!(append_i64, i64, 8, write_i64);
append_method!(append_u32, u32, 4, write_u32);
append_method!(append_i32, i32, 4, write_i32);
append_method!(append_u64, u64, 8, write_u64);
pub(crate) fn build(&mut self) -> Vec<u8> {
let total = u16::try_from(self.data.len()).expect("test packet exceeds u16 length");
BigEndian::write_u16(&mut self.data[2..4], total);
self.data.clone()
}
}
pub(crate) fn encode_utf16_le(value: &str) -> Vec<u8> {
value
.encode_utf16()
.flat_map(|unit| unit.to_le_bytes())
.collect()
}
pub(crate) fn build_duplex_transport(client_side: DuplexStream) -> NetworkTransport {
let context = ClientContext::default();
NetworkTransport::new(
Box::new(client_side),
SslHandler {
server_host_name: context.transport_context.get_server_name().clone(),
encryption_options: context.encryption_options.clone(),
},
context.packet_size as u32,
context.encryption_options.mode,
false,
)
}
pub(crate) fn create_network_transport_with_data(data: &[u8]) -> NetworkTransport {
create_network_transport_with_chunked_data(data, data.len().max(1))
}
pub(crate) fn create_network_transport_with_chunked_data(
data: &[u8],
chunk_size: usize,
) -> NetworkTransport {
let chunk_size = chunk_size.max(1);
let (client_side, mut server_side) = duplex(chunk_size);
let owned = data.to_vec();
tokio::spawn(async move {
for chunk in owned.chunks(chunk_size) {
if server_side.write_all(chunk).await.is_err() {
return;
}
}
});
build_duplex_transport(client_side)
}
pub(crate) fn create_network_transport_with_live_peer(data: &[u8]) -> NetworkTransport {
let (client_side, mut server_side) = duplex(data.len().max(1));
let owned = data.to_vec();
tokio::spawn(async move {
let _ = server_side.write_all(&owned).await;
std::future::pending::<()>().await;
});
build_duplex_transport(client_side)
}
pub(crate) fn create_network_transport_with_live_peer_capturing_writes(
data: &[u8],
) -> (NetworkTransport, UnboundedReceiver<Vec<u8>>) {
let (client_side, mut server_side) = duplex(data.len().max(4096));
let (sender, receiver) = unbounded_channel();
let owned = data.to_vec();
tokio::spawn(async move {
let _ = server_side.write_all(&owned).await;
let mut buffer = [0u8; 512];
while let Ok(read) = server_side.read(&mut buffer).await {
if read == 0 || sender.send(buffer[..read].to_vec()).is_err() {
break;
}
}
std::future::pending::<()>().await;
});
(build_duplex_transport(client_side), receiver)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_packet_builder_writes_total_length_in_header() {
let mut builder = TestPacketBuilder::new(PacketType::PreLogin);
builder.append_bytes(&[0u8; 12]);
let packet = builder.build();
assert_eq!(packet.len(), 20);
assert_eq!(
BigEndian::read_u16(&packet[2..4]),
20,
"header must carry total length, not payload length"
);
}
#[test]
fn test_packet_builder_empty_payload_length_is_header_only() {
let packet = TestPacketBuilder::new(PacketType::TabularResult).build();
assert_eq!(packet.len(), 8);
assert_eq!(BigEndian::read_u16(&packet[2..4]), 8);
}
}