use super::types::{self, Column};
use crate::value::Value;
use rustlavel_core::{Error, Result};
use std::sync::Arc;
pub const HEADER_LEN: usize = 8;
pub const DEFAULT_PACKET_SIZE: usize = 4096;
pub const TDS_VERSION_7_4: u32 = 0x7400_0004;
pub mod packet {
pub const SQL_BATCH: u8 = 0x01;
pub const RPC: u8 = 0x03;
pub const TABULAR_RESULT: u8 = 0x04;
pub const ATTENTION: u8 = 0x06;
pub const LOGIN7: u8 = 0x10;
pub const SSPI: u8 = 0x11;
pub const PRE_LOGIN: u8 = 0x12;
}
pub mod status {
pub const NORMAL: u8 = 0x00;
pub const END_OF_MESSAGE: u8 = 0x01;
pub const IGNORE: u8 = 0x02;
pub const RESET_CONNECTION: u8 = 0x08;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PacketHeader {
pub kind: u8,
pub status: u8,
pub length: u16,
pub spid: u16,
pub id: u8,
pub window: u8,
}
impl PacketHeader {
pub fn parse(bytes: &[u8]) -> Result<PacketHeader> {
if bytes.len() < HEADER_LEN {
return Err(Error::Protocol("truncated packet header from the server".into()));
}
Ok(PacketHeader {
kind: bytes[0],
status: bytes[1],
length: u16::from_be_bytes([bytes[2], bytes[3]]),
spid: u16::from_be_bytes([bytes[4], bytes[5]]),
id: bytes[6],
window: bytes[7],
})
}
pub fn write_into(&self, out: &mut Vec<u8>) {
out.push(self.kind);
out.push(self.status);
out.extend_from_slice(&self.length.to_be_bytes());
out.extend_from_slice(&self.spid.to_be_bytes());
out.push(self.id);
out.push(self.window);
}
pub fn is_end_of_message(&self) -> bool {
self.status & status::END_OF_MESSAGE != 0
}
}
pub fn split_message(kind: u8, payload: &[u8], packet_size: usize) -> Vec<Vec<u8>> {
let capacity = packet_size.max(HEADER_LEN + 1) - HEADER_LEN;
let mut packets = Vec::with_capacity(payload.len() / capacity + 1);
let mut id: u8 = 1;
let mut offset = 0;
loop {
let end = (offset + capacity).min(payload.len());
let chunk = &payload[offset..end];
let last = end == payload.len();
let mut packet = Vec::with_capacity(HEADER_LEN + chunk.len());
PacketHeader {
kind,
status: if last { status::END_OF_MESSAGE } else { status::NORMAL },
length: (HEADER_LEN + chunk.len()) as u16,
spid: 0,
id,
window: 0,
}
.write_into(&mut packet);
packet.extend_from_slice(chunk);
packets.push(packet);
if last {
return packets;
}
offset = end;
id = id.wrapping_add(1);
}
}
pub mod prelogin_option {
pub const VERSION: u8 = 0x00;
pub const ENCRYPTION: u8 = 0x01;
pub const INSTOPT: u8 = 0x02;
pub const THREADID: u8 = 0x03;
pub const MARS: u8 = 0x04;
pub const TERMINATOR: u8 = 0xFF;
}
pub mod encryption {
pub const OFF: u8 = 0x00;
pub const ON: u8 = 0x01;
pub const NOT_SUPPORTED: u8 = 0x02;
pub const REQUIRED: u8 = 0x03;
}
pub fn prelogin(encryption: u8) -> Vec<u8> {
let options: [(u8, Vec<u8>); 5] = [
(prelogin_option::VERSION, vec![9, 0, 0, 0, 0, 0]),
(prelogin_option::ENCRYPTION, vec![encryption]),
(prelogin_option::INSTOPT, vec![0]),
(prelogin_option::THREADID, 0u32.to_le_bytes().to_vec()),
(prelogin_option::MARS, vec![0]),
];
let header_len = options.len() * 5 + 1;
let mut head = Vec::with_capacity(header_len);
let mut data = Vec::new();
for (token, value) in &options {
head.push(*token);
head.extend_from_slice(&((header_len + data.len()) as u16).to_be_bytes());
head.extend_from_slice(&(value.len() as u16).to_be_bytes());
data.extend_from_slice(value);
}
head.push(prelogin_option::TERMINATOR);
head.extend_from_slice(&data);
head
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PreloginResponse {
pub encryption: u8,
}
pub fn parse_prelogin(payload: &[u8]) -> Result<PreloginResponse> {
let mut encryption = encryption::NOT_SUPPORTED;
let mut at = 0;
while at < payload.len() {
let token = payload[at];
if token == prelogin_option::TERMINATOR {
break;
}
if at + 5 > payload.len() {
return Err(Error::Protocol("truncated PRELOGIN option header".into()));
}
let offset = u16::from_be_bytes([payload[at + 1], payload[at + 2]]) as usize;
let length = u16::from_be_bytes([payload[at + 3], payload[at + 4]]) as usize;
if offset + length > payload.len() {
return Err(Error::Protocol("PRELOGIN option points past the packet".into()));
}
if token == prelogin_option::ENCRYPTION && length >= 1 {
encryption = payload[offset];
}
at += 5;
}
Ok(PreloginResponse { encryption })
}
#[derive(Debug, Clone)]
pub struct Login7<'a> {
pub hostname: &'a str,
pub username: &'a str,
pub password: &'a [u8],
pub application: &'a str,
pub server: &'a str,
pub library: &'a str,
pub language: &'a str,
pub database: &'a str,
pub packet_size: usize,
}
const LOGIN7_FIXED_LEN: usize = 94;
const OPTION_FLAGS_1: u8 = 0xE0;
const OPTION_FLAGS_2: u8 = 0x02;
pub fn login7(login: &Login7<'_>) -> Vec<u8> {
let mut out = Vec::with_capacity(256);
out.extend_from_slice(&0u32.to_le_bytes());
out.extend_from_slice(&TDS_VERSION_7_4.to_le_bytes());
out.extend_from_slice(&(login.packet_size as u32).to_le_bytes());
out.extend_from_slice(&0x0100_0000u32.to_le_bytes()); out.extend_from_slice(&std::process::id().to_le_bytes());
out.extend_from_slice(&0u32.to_le_bytes()); out.push(OPTION_FLAGS_1);
out.push(OPTION_FLAGS_2);
out.push(0); out.push(0); out.extend_from_slice(&0i32.to_le_bytes()); out.extend_from_slice(&0u32.to_le_bytes());
let mut data: Vec<u8> = Vec::new();
let place = |data: &mut Vec<u8>, text: &str| -> [u8; 4] {
let offset = (LOGIN7_FIXED_LEN + data.len()) as u16;
let mut characters = 0u16;
for unit in text.encode_utf16() {
data.extend_from_slice(&unit.to_le_bytes());
characters += 1;
}
let mut entry = [0u8; 4];
entry[..2].copy_from_slice(&offset.to_le_bytes());
entry[2..].copy_from_slice(&characters.to_le_bytes());
entry
};
let hostname = place(&mut data, login.hostname);
let username = place(&mut data, login.username);
let password_offset = (LOGIN7_FIXED_LEN + data.len()) as u16;
data.extend_from_slice(login.password);
let password_characters = (login.password.len() / 2) as u16;
let application = place(&mut data, login.application);
let server = place(&mut data, login.server);
let library = place(&mut data, login.library);
let language = place(&mut data, login.language);
let database = place(&mut data, login.database);
let tail = (LOGIN7_FIXED_LEN + data.len()) as u16;
out.extend_from_slice(&hostname);
out.extend_from_slice(&username);
out.extend_from_slice(&password_offset.to_le_bytes());
out.extend_from_slice(&password_characters.to_le_bytes());
out.extend_from_slice(&application);
out.extend_from_slice(&server);
out.extend_from_slice(&[0u8; 4]); out.extend_from_slice(&library);
out.extend_from_slice(&language);
out.extend_from_slice(&database);
out.extend_from_slice(&[0u8; 6]); out.extend_from_slice(&tail.to_le_bytes()); out.extend_from_slice(&0u16.to_le_bytes()); out.extend_from_slice(&tail.to_le_bytes()); out.extend_from_slice(&0u16.to_le_bytes());
out.extend_from_slice(&tail.to_le_bytes()); out.extend_from_slice(&0u16.to_le_bytes());
out.extend_from_slice(&0u32.to_le_bytes());
debug_assert_eq!(out.len(), LOGIN7_FIXED_LEN);
out.extend_from_slice(&data);
let length = out.len() as u32;
out[..4].copy_from_slice(&length.to_le_bytes());
out
}
pub fn all_headers(transaction: u64) -> Vec<u8> {
let mut out = Vec::with_capacity(22);
out.extend_from_slice(&22u32.to_le_bytes()); out.extend_from_slice(&18u32.to_le_bytes()); out.extend_from_slice(&2u16.to_le_bytes()); out.extend_from_slice(&transaction.to_le_bytes());
out.extend_from_slice(&1u32.to_le_bytes()); out
}
pub fn sql_batch(sql: &str, transaction: u64) -> Vec<u8> {
let mut out = all_headers(transaction);
for unit in sql.encode_utf16() {
out.extend_from_slice(&unit.to_le_bytes());
}
out
}
pub const SP_EXECUTESQL: u16 = 10;
#[derive(Debug, Clone)]
pub struct RpcParameter {
pub name: String,
pub bytes: Vec<u8>,
}
pub fn rpc(proc_id: u16, parameters: &[RpcParameter], transaction: u64) -> Vec<u8> {
let mut out = all_headers(transaction);
out.extend_from_slice(&0xFFFFu16.to_le_bytes());
out.extend_from_slice(&proc_id.to_le_bytes());
out.extend_from_slice(&0u16.to_le_bytes());
for parameter in parameters {
let name: Vec<u16> = parameter.name.encode_utf16().collect();
out.push(name.len() as u8);
for unit in &name {
out.extend_from_slice(&unit.to_le_bytes());
}
out.push(0); out.extend_from_slice(¶meter.bytes);
}
out
}
pub fn execute_sql(sql: &str, params: &[Value], transaction: u64) -> Vec<u8> {
let declaration = types::declare(params);
let mut parameters = Vec::with_capacity(params.len() + 2);
parameters.push(RpcParameter {
name: String::new(),
bytes: types::encode(&Value::Text(sql.to_string())),
});
parameters.push(RpcParameter {
name: String::new(),
bytes: types::encode(&Value::Text(declaration)),
});
for (index, value) in params.iter().enumerate() {
parameters.push(RpcParameter {
name: format!("@P{}", index + 1),
bytes: types::encode(value),
});
}
rpc(SP_EXECUTESQL, ¶meters, transaction)
}
pub mod token {
pub const RETURN_STATUS: u8 = 0x79;
pub const COLMETADATA: u8 = 0x81;
pub const ALTMETADATA: u8 = 0x88;
pub const TABNAME: u8 = 0xA4;
pub const COLINFO: u8 = 0xA5;
pub const ORDER: u8 = 0xA9;
pub const ERROR: u8 = 0xAA;
pub const INFO: u8 = 0xAB;
pub const RETURN_VALUE: u8 = 0xAC;
pub const LOGINACK: u8 = 0xAD;
pub const FEATUREEXTACK: u8 = 0xAE;
pub const ROW: u8 = 0xD1;
pub const NBCROW: u8 = 0xD2;
pub const ENVCHANGE: u8 = 0xE3;
pub const SSPI: u8 = 0xED;
pub const DONE: u8 = 0xFD;
pub const DONEPROC: u8 = 0xFE;
pub const DONEINPROC: u8 = 0xFF;
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ServerError {
pub number: i32,
pub state: u8,
pub severity: u8,
pub message: String,
pub server: String,
pub procedure: String,
pub line: u32,
}
impl ServerError {
pub fn into_error(self, sql: Option<&str>) -> Error {
let mut text = format!(
"SQL Server error {} (severity {}, state {}): {}",
self.number, self.severity, self.state, self.message
);
if !self.procedure.is_empty() {
text.push_str(&format!(" — in {}, line {}", self.procedure, self.line));
}
if let Some(sql) = sql {
text.push_str(&format!("\n SQL: {sql}"));
}
Error::msg(text)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DoneKind {
Batch,
Procedure,
InProcedure,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Done {
pub kind: DoneKind,
pub status: u16,
pub current_command: u16,
pub rows: u64,
}
impl Done {
pub fn has_count(&self) -> bool {
self.status & 0x0010 != 0
}
pub fn has_error(&self) -> bool {
self.status & 0x0002 != 0
}
pub fn has_more(&self) -> bool {
self.status & 0x0001 != 0
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LoginAck {
pub interface: u8,
pub tds_version: u32,
pub program: String,
pub version: (u8, u8, u16),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EnvChange {
Database(String),
PacketSize(usize),
BeginTransaction(u64),
CommitTransaction,
RollbackTransaction,
Other(u8),
}
#[derive(Debug, Clone)]
pub enum Token {
ColumnMetadata(Arc<Vec<Column>>),
Row(Vec<Value>),
Done(Done),
Error(ServerError),
Info(ServerError),
LoginAck(LoginAck),
EnvChange(EnvChange),
ReturnStatus(i32),
Ignored(u8),
}
pub struct TokenStream<'a> {
reader: Reader<'a>,
columns: Arc<Vec<Column>>,
}
impl<'a> TokenStream<'a> {
pub fn new(bytes: &'a [u8]) -> Self {
TokenStream { reader: Reader::new(bytes), columns: Arc::new(Vec::new()) }
}
pub fn columns(&self) -> &Arc<Vec<Column>> {
&self.columns
}
pub fn next_token(&mut self) -> Result<Option<Token>> {
if self.reader.is_empty() {
return Ok(None);
}
let tag = self.reader.u8()?;
let parsed = match tag {
token::COLMETADATA => {
self.columns = Arc::new(types::parse_column_metadata(&mut self.reader)?);
Token::ColumnMetadata(Arc::clone(&self.columns))
}
token::ROW => Token::Row(types::read_row(&mut self.reader, &self.columns)?),
token::NBCROW => Token::Row(types::read_nbc_row(&mut self.reader, &self.columns)?),
token::DONE => Token::Done(parse_done(DoneKind::Batch, &mut self.reader)?),
token::DONEPROC => Token::Done(parse_done(DoneKind::Procedure, &mut self.reader)?),
token::DONEINPROC => Token::Done(parse_done(DoneKind::InProcedure, &mut self.reader)?),
token::ERROR => Token::Error(parse_server_error(&mut self.reader)?),
token::INFO => Token::Info(parse_server_error(&mut self.reader)?),
token::LOGINACK => Token::LoginAck(parse_login_ack(&mut self.reader)?),
token::ENVCHANGE => Token::EnvChange(parse_env_change(&mut self.reader)?),
token::RETURN_STATUS => Token::ReturnStatus(self.reader.i32()?),
token::RETURN_VALUE => {
skip_return_value(&mut self.reader)?;
Token::Ignored(tag)
}
token::FEATUREEXTACK => {
skip_feature_ext_ack(&mut self.reader)?;
Token::Ignored(tag)
}
token::ORDER | token::TABNAME | token::COLINFO | token::SSPI | token::ALTMETADATA => {
let length = self.reader.u16()? as usize;
self.reader.skip(length)?;
Token::Ignored(tag)
}
other => {
return Err(Error::Protocol(format!(
"unknown TDS token 0x{other:02X} in the response stream"
)));
}
};
Ok(Some(parsed))
}
}
fn parse_done(kind: DoneKind, reader: &mut Reader<'_>) -> Result<Done> {
Ok(Done {
kind,
status: reader.u16()?,
current_command: reader.u16()?,
rows: reader.u64()?,
})
}
fn parse_server_error(reader: &mut Reader<'_>) -> Result<ServerError> {
let _length = reader.u16()?;
Ok(ServerError {
number: reader.i32()?,
state: reader.u8()?,
severity: reader.u8()?,
message: reader.us_varchar()?,
server: reader.b_varchar()?,
procedure: reader.b_varchar()?,
line: reader.u32()?,
})
}
fn parse_login_ack(reader: &mut Reader<'_>) -> Result<LoginAck> {
let _length = reader.u16()?;
let interface = reader.u8()?;
let tds_version = reader.u32()?;
let program = reader.b_varchar()?;
let major = reader.u8()?;
let minor = reader.u8()?;
let build_high = reader.u8()?;
let build_low = reader.u8()?;
Ok(LoginAck {
interface,
tds_version,
program,
version: (major, minor, u16::from_be_bytes([build_high, build_low])),
})
}
fn parse_env_change(reader: &mut Reader<'_>) -> Result<EnvChange> {
let length = reader.u16()? as usize;
let body = reader.take(length)?;
let mut inner = Reader::new(body);
Ok(match inner.u8()? {
1 => EnvChange::Database(inner.b_varchar()?),
4 => EnvChange::PacketSize(
inner.b_varchar()?.parse().unwrap_or(DEFAULT_PACKET_SIZE),
),
8 => {
let descriptor = inner.b_varbyte()?;
let mut bytes = [0u8; 8];
let taken = descriptor.len().min(8);
bytes[..taken].copy_from_slice(&descriptor[..taken]);
EnvChange::BeginTransaction(u64::from_le_bytes(bytes))
}
9 => EnvChange::CommitTransaction,
10 => EnvChange::RollbackTransaction,
other => EnvChange::Other(other),
})
}
fn skip_return_value(reader: &mut Reader<'_>) -> Result<()> {
reader.u16()?; reader.b_varchar()?; reader.u8()?; reader.u32()?; reader.u16()?; let type_info = types::parse_type_info(reader)?;
types::read_value(reader, &type_info)?;
Ok(())
}
fn skip_feature_ext_ack(reader: &mut Reader<'_>) -> Result<()> {
loop {
if reader.u8()? == 0xFF {
return Ok(());
}
let length = reader.u32()? as usize;
reader.skip(length)?;
}
}
pub struct Reader<'a> {
bytes: &'a [u8],
position: usize,
}
impl<'a> Reader<'a> {
pub fn new(bytes: &'a [u8]) -> Self {
Reader { bytes, position: 0 }
}
pub fn is_empty(&self) -> bool {
self.position >= self.bytes.len()
}
pub fn remaining(&self) -> usize {
self.bytes.len().saturating_sub(self.position)
}
pub fn take(&mut self, count: usize) -> Result<&'a [u8]> {
let end = self.position.checked_add(count).ok_or_else(too_short)?;
if end > self.bytes.len() {
return Err(too_short());
}
let slice = &self.bytes[self.position..end];
self.position = end;
Ok(slice)
}
pub fn skip(&mut self, count: usize) -> Result<()> {
self.take(count).map(|_| ())
}
pub fn u8(&mut self) -> Result<u8> {
Ok(self.take(1)?[0])
}
pub fn u16(&mut self) -> Result<u16> {
Ok(u16::from_le_bytes(self.take(2)?.try_into().expect("2 bytes")))
}
pub fn i32(&mut self) -> Result<i32> {
Ok(i32::from_le_bytes(self.take(4)?.try_into().expect("4 bytes")))
}
pub fn u32(&mut self) -> Result<u32> {
Ok(u32::from_le_bytes(self.take(4)?.try_into().expect("4 bytes")))
}
pub fn u64(&mut self) -> Result<u64> {
Ok(u64::from_le_bytes(self.take(8)?.try_into().expect("8 bytes")))
}
pub fn b_varchar(&mut self) -> Result<String> {
let characters = self.u8()? as usize;
self.ucs2(characters)
}
pub fn us_varchar(&mut self) -> Result<String> {
let characters = self.u16()? as usize;
self.ucs2(characters)
}
pub fn b_varbyte(&mut self) -> Result<&'a [u8]> {
let length = self.u8()? as usize;
self.take(length)
}
fn ucs2(&mut self, characters: usize) -> Result<String> {
let bytes = self.take(characters * 2)?;
Ok(decode_ucs2(bytes))
}
}
pub fn decode_ucs2(bytes: &[u8]) -> String {
let (pairs, _odd_trailing_byte) = bytes.as_chunks::<2>();
let units: Vec<u16> = pairs.iter().copied().map(u16::from_le_bytes).collect();
String::from_utf16_lossy(&units)
}
fn too_short() -> Error {
Error::Protocol("truncated message from the server".into())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_packet_header_survives_a_round_trip() {
let header = PacketHeader {
kind: packet::SQL_BATCH,
status: status::END_OF_MESSAGE,
length: 4096,
spid: 53,
id: 7,
window: 0,
};
let mut bytes = Vec::new();
header.write_into(&mut bytes);
assert_eq!(bytes.len(), HEADER_LEN);
assert_eq!(&bytes[2..4], &4096u16.to_be_bytes());
assert_eq!(PacketHeader::parse(&bytes).unwrap(), header);
}
#[test]
fn a_short_header_is_a_protocol_error() {
assert!(PacketHeader::parse(&[0x04, 0x01, 0x00]).is_err());
}
#[test]
fn a_payload_that_fits_becomes_one_packet_marked_final() {
let packets = split_message(packet::SQL_BATCH, b"hello", DEFAULT_PACKET_SIZE);
assert_eq!(packets.len(), 1);
let header = PacketHeader::parse(&packets[0]).unwrap();
assert_eq!(header.kind, packet::SQL_BATCH);
assert_eq!(header.length as usize, packets[0].len());
assert_eq!(header.id, 1);
assert!(header.is_end_of_message());
assert_eq!(&packets[0][HEADER_LEN..], b"hello");
}
#[test]
fn a_payload_larger_than_the_packet_size_is_split_and_only_the_last_ends_it() {
let payload: Vec<u8> = (0..60u8).collect();
let packets = split_message(packet::SQL_BATCH, &payload, 32);
assert_eq!(packets.len(), 3);
let headers: Vec<PacketHeader> =
packets.iter().map(|p| PacketHeader::parse(p).unwrap()).collect();
assert_eq!(headers.iter().map(|h| h.id).collect::<Vec<_>>(), vec![1, 2, 3]);
assert!(!headers[0].is_end_of_message());
assert!(!headers[1].is_end_of_message());
assert!(headers[2].is_end_of_message());
assert!(packets.iter().all(|p| p.len() <= 32));
let rebuilt: Vec<u8> =
packets.iter().flat_map(|p| p[HEADER_LEN..].iter().copied()).collect();
assert_eq!(rebuilt, payload);
}
#[test]
fn an_empty_payload_is_still_one_end_of_message_packet() {
let packets = split_message(packet::PRE_LOGIN, &[], DEFAULT_PACKET_SIZE);
assert_eq!(packets.len(), 1);
assert_eq!(packets[0].len(), HEADER_LEN);
assert!(PacketHeader::parse(&packets[0]).unwrap().is_end_of_message());
}
#[test]
fn a_prelogin_requests_encryption_and_its_offsets_point_at_its_data() {
let payload = prelogin(encryption::ON);
assert_eq!(payload[25], prelogin_option::TERMINATOR);
assert_eq!(payload[5], prelogin_option::ENCRYPTION);
let offset = u16::from_be_bytes([payload[6], payload[7]]) as usize;
let length = u16::from_be_bytes([payload[8], payload[9]]) as usize;
assert_eq!(length, 1);
assert_eq!(payload[offset], encryption::ON);
}
#[test]
fn reads_the_encryption_level_the_server_chose() {
let mut answer = prelogin(encryption::REQUIRED);
assert_eq!(
parse_prelogin(&answer).unwrap(),
PreloginResponse { encryption: encryption::REQUIRED }
);
answer.truncate(1);
answer[0] = prelogin_option::TERMINATOR;
assert_eq!(
parse_prelogin(&answer).unwrap().encryption,
encryption::NOT_SUPPORTED
);
}
#[test]
fn a_prelogin_option_pointing_past_the_packet_is_rejected() {
let payload = vec![prelogin_option::ENCRYPTION, 0xFF, 0xFF, 0x00, 0x01, 0xFF];
assert!(parse_prelogin(&payload).is_err());
}
#[test]
fn login7_declares_its_own_length_and_counts_strings_in_characters() {
let payload = login7(&Login7 {
hostname: "laptop",
username: "sa",
password: &[0xB3, 0xA5, 0x83, 0xA5],
application: "rustlavel",
server: "db",
library: "rustlavel-db",
language: "",
database: "blog",
packet_size: DEFAULT_PACKET_SIZE,
});
assert_eq!(
u32::from_le_bytes(payload[..4].try_into().unwrap()) as usize,
payload.len()
);
assert_eq!(u32::from_le_bytes(payload[4..8].try_into().unwrap()), TDS_VERSION_7_4);
let username_offset = u16::from_le_bytes(payload[40..42].try_into().unwrap()) as usize;
let username_length = u16::from_le_bytes(payload[42..44].try_into().unwrap()) as usize;
assert_eq!(username_length, 2);
assert_eq!(
decode_ucs2(&payload[username_offset..username_offset + username_length * 2]),
"sa"
);
let password_length = u16::from_le_bytes(payload[46..48].try_into().unwrap());
assert_eq!(password_length, 2);
}
#[test]
fn a_batch_carries_the_transaction_it_belongs_to() {
let payload = sql_batch("select 1", 0xDEAD_BEEF);
assert_eq!(u32::from_le_bytes(payload[..4].try_into().unwrap()), 22);
assert_eq!(u64::from_le_bytes(payload[10..18].try_into().unwrap()), 0xDEAD_BEEF);
assert_eq!(decode_ucs2(&payload[22..]), "select 1");
}
#[test]
fn an_rpc_names_its_procedure_by_id() {
let payload = rpc(SP_EXECUTESQL, &[], 0);
assert_eq!(u16::from_le_bytes(payload[22..24].try_into().unwrap()), 0xFFFF);
assert_eq!(u16::from_le_bytes(payload[24..26].try_into().unwrap()), SP_EXECUTESQL);
}
#[test]
fn a_parameterised_call_sends_the_statement_as_data_not_as_sql() {
let hostile = "'; drop table users; --";
let payload = execute_sql("select @P1", &[Value::Text(hostile.into())], 0);
let statement: Vec<u8> = "select @P1".encode_utf16().flat_map(u16::to_le_bytes).collect();
let value: Vec<u8> = hostile.encode_utf16().flat_map(u16::to_le_bytes).collect();
assert!(payload.windows(statement.len()).any(|w| w == statement));
assert!(payload.windows(value.len()).any(|w| w == value));
let declaration: Vec<u8> =
"@P1 nvarchar(max)".encode_utf16().flat_map(u16::to_le_bytes).collect();
assert!(payload.windows(declaration.len()).any(|w| w == declaration));
}
#[test]
fn an_error_token_names_its_number_and_severity() {
let mut body = vec![token::ERROR];
let mut fields = Vec::new();
fields.extend_from_slice(&18456i32.to_le_bytes());
fields.push(1); fields.push(14); let message = "Login failed for user 'sa'.";
fields.extend_from_slice(&(message.encode_utf16().count() as u16).to_le_bytes());
fields.extend(message.encode_utf16().flat_map(u16::to_le_bytes));
fields.push(2); fields.extend("db".encode_utf16().flat_map(u16::to_le_bytes));
fields.push(0); fields.extend_from_slice(&1u32.to_le_bytes());
body.extend_from_slice(&(fields.len() as u16).to_le_bytes());
body.extend_from_slice(&fields);
let mut stream = TokenStream::new(&body);
let error = match stream.next_token().unwrap().unwrap() {
Token::Error(error) => error,
other => panic!("expected an error token, got {other:?}"),
};
assert_eq!(error.number, 18456);
assert_eq!(error.severity, 14);
assert_eq!(error.message, message);
assert_eq!(error.server, "db");
let rendered = error.into_error(Some("select 1")).to_string();
assert!(rendered.contains("18456"), "{rendered}");
assert!(rendered.contains("severity 14"), "{rendered}");
assert!(rendered.contains("SQL: select 1"), "{rendered}");
}
#[test]
fn a_done_token_reports_the_rows_a_statement_touched() {
let mut body = vec![token::DONEINPROC];
body.extend_from_slice(&0x0011u16.to_le_bytes()); body.extend_from_slice(&0xC1u16.to_le_bytes()); body.extend_from_slice(&3u64.to_le_bytes());
let mut stream = TokenStream::new(&body);
let done = match stream.next_token().unwrap().unwrap() {
Token::Done(done) => done,
other => panic!("expected a done token, got {other:?}"),
};
assert_eq!(done.kind, DoneKind::InProcedure);
assert_eq!(done.rows, 3);
assert!(done.has_count());
assert!(done.has_more());
assert!(!done.has_error());
}
#[test]
fn a_done_token_without_a_count_bit_reports_no_rows() {
let mut body = vec![token::DONE];
body.extend_from_slice(&0u16.to_le_bytes());
body.extend_from_slice(&0u16.to_le_bytes());
body.extend_from_slice(&99u64.to_le_bytes());
let mut stream = TokenStream::new(&body);
match stream.next_token().unwrap().unwrap() {
Token::Done(done) => assert!(!done.has_count()),
other => panic!("expected a done token, got {other:?}"),
}
}
#[test]
fn an_env_change_announces_a_transaction_and_then_ends_it() {
let mut body = vec![token::ENVCHANGE];
let mut change = vec![8u8]; change.push(8); change.extend_from_slice(&0x0102_0304_0506_0708u64.to_le_bytes());
change.push(0); body.extend_from_slice(&(change.len() as u16).to_le_bytes());
body.extend_from_slice(&change);
body.push(token::ENVCHANGE);
let ended = vec![9u8, 0, 0];
body.extend_from_slice(&(ended.len() as u16).to_le_bytes());
body.extend_from_slice(&ended);
let mut stream = TokenStream::new(&body);
match stream.next_token().unwrap().unwrap() {
Token::EnvChange(EnvChange::BeginTransaction(descriptor)) => {
assert_eq!(descriptor, 0x0102_0304_0506_0708)
}
other => panic!("expected a transaction to begin, got {other:?}"),
}
match stream.next_token().unwrap().unwrap() {
Token::EnvChange(EnvChange::CommitTransaction) => {}
other => panic!("expected a commit, got {other:?}"),
}
}
#[test]
fn the_negotiated_packet_size_arrives_as_a_decimal_string() {
let mut body = vec![token::ENVCHANGE];
let mut change = vec![4u8];
change.push(4); change.extend("8192".encode_utf16().flat_map(u16::to_le_bytes));
change.push(0);
body.extend_from_slice(&(change.len() as u16).to_le_bytes());
body.extend_from_slice(&change);
let mut stream = TokenStream::new(&body);
match stream.next_token().unwrap().unwrap() {
Token::EnvChange(EnvChange::PacketSize(size)) => assert_eq!(size, 8192),
other => panic!("expected a packet size change, got {other:?}"),
}
}
#[test]
fn a_login_ack_reports_the_version_that_was_negotiated() {
let mut fields = vec![1u8]; fields.extend_from_slice(&TDS_VERSION_7_4.to_le_bytes());
fields.push(4);
fields.extend("mssq".encode_utf16().flat_map(u16::to_le_bytes));
fields.extend_from_slice(&[16, 0, 0x0F, 0xA0]);
let mut body = vec![token::LOGINACK];
body.extend_from_slice(&(fields.len() as u16).to_le_bytes());
body.extend_from_slice(&fields);
let mut stream = TokenStream::new(&body);
match stream.next_token().unwrap().unwrap() {
Token::LoginAck(ack) => {
assert_eq!(ack.program, "mssq");
assert_eq!(ack.version, (16, 0, 0x0FA0));
assert_eq!(ack.tds_version, TDS_VERSION_7_4);
}
other => panic!("expected a login ack, got {other:?}"),
}
}
#[test]
fn an_unknown_token_stops_the_stream_rather_than_guessing() {
let error = TokenStream::new(&[0x42]).next_token().unwrap_err().to_string();
assert!(error.contains("0x42"), "{error}");
}
#[test]
fn an_empty_message_yields_no_tokens() {
assert!(TokenStream::new(&[]).next_token().unwrap().is_none());
}
#[test]
fn a_truncated_token_is_a_protocol_error_not_a_panic() {
let mut body = vec![token::DONE];
body.extend_from_slice(&[0, 0, 0, 0, 1, 2]);
assert!(TokenStream::new(&body).next_token().is_err());
}
#[test]
fn reads_little_endian_scalars_and_counted_strings() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&0x1234u16.to_le_bytes());
bytes.extend_from_slice(&(-5i32).to_le_bytes());
bytes.push(3);
bytes.extend("ada".encode_utf16().flat_map(u16::to_le_bytes));
let mut reader = Reader::new(&bytes);
assert_eq!(reader.u16().unwrap(), 0x1234);
assert_eq!(reader.i32().unwrap(), -5);
assert_eq!(reader.b_varchar().unwrap(), "ada");
assert!(reader.is_empty());
}
}