use crate::connection::bulk_copy_state::ATTENTION_TIMEOUT_SECONDS;
use crate::connection::client_context::{IPAddressPreference, TransportContext};
use crate::connection::transport::buffers::TdsReadBuffer;
use crate::connection::transport::extractable_stream;
use crate::connection::transport::parallel_connect::{ParallelConnectConfig, parallel_connect};
use crate::connection::transport::ssl_handler::SslHandler;
use crate::connection_provider::tds_connection_provider::PARSER_REGISTRY;
use crate::core::{
CancelHandle, EncryptionOptions, EncryptionSetting, NegotiatedEncryptionSetting, TdsResult,
};
use crate::datatypes::column_values::ColumnValues;
use crate::datatypes::decoder::{GenericDecoder, PlpColumnStream};
use crate::datatypes::row_writer::RowWriter;
use crate::datatypes::sqldatatypes::TdsDataType;
use crate::error::Error::{OperationCancelledError, TimeoutError};
use crate::error::TimeoutErrorType;
use crate::handler::handler_factory::SessionSettings;
use crate::io::packet_reader::{LENGTH_NULL, TdsPacketReader};
use crate::io::packet_writer::PacketWriter;
use crate::io::reader_writer::{NetworkReader, NetworkReaderWriter, NetworkWriter};
use crate::io::token_stream::{
ColumnPolicy, ParserContext, PlpPauseState, RowHeader, RowPauseState, RowReadResult,
TdsTokenStreamReader, read_active_plp_bytes_internal, receive_row_header_internal,
receive_row_into_internal, receive_token_internal, resume_row_into_internal,
};
use crate::message::attention::AttentionRequest;
use crate::message::login_options::TdsVersion;
use crate::message::messages::{PacketStatusFlags, PacketType, Request, ResetConnectionMode};
use crate::token::tokens::{ColMetadataToken, DoneStatus, TokenType, Tokens};
use async_trait::async_trait;
use byteorder::{BigEndian, ByteOrder, LittleEndian};
use std::cmp::min;
use std::future::{Future, poll_fn};
use std::io::Error;
use std::io::ErrorKind;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard};
use std::task::{Context, Poll, Waker};
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
use tokio::net::{self, TcpStream};
use tokio::time::{Instant, timeout, timeout_at};
use tracing::{debug, error, event, info, trace, warn};
type CompleteBufferedPlp = Option<Option<(usize, Option<u64>, usize)>>;
const MAX_ATTENTION_SETTLEMENT_TOKENS: usize = 1024;
#[derive(Debug, Default)]
pub(crate) struct AttentionSettlement {
pub(crate) tokens: Vec<Tokens>,
pub(crate) overflowed: bool,
}
impl AttentionSettlement {
fn push(&mut self, token: Tokens) {
if self.tokens.len() < MAX_ATTENTION_SETTLEMENT_TOKENS {
self.tokens.push(token);
} else {
self.overflowed = true;
}
}
pub(crate) fn retained_token_count(&self) -> usize {
self.tokens.len()
}
}
enum ReadInterruption {
Cancelled,
TimedOut(tokio::time::error::Elapsed),
}
enum InterruptibleRead<T> {
Completed(TdsResult<T>),
Interrupted {
error: Box<crate::error::Error>,
boundary: Option<(Instant, TdsResult<T>)>,
},
}
impl ReadInterruption {
fn into_error(self) -> crate::error::Error {
match self {
Self::Cancelled => OperationCancelledError("Request was cancelled".to_string()),
Self::TimedOut(elapsed) => TimeoutError(TimeoutErrorType::Elapsed(elapsed)),
}
}
}
async fn await_read_or_interrupt<F, T>(
mut read: Pin<&mut F>,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
) -> Result<TdsResult<T>, ReadInterruption>
where
F: Future<Output = TdsResult<T>>,
{
if cancel_handle.is_some_and(|handle| handle.cancel_token.is_cancelled()) {
return Err(ReadInterruption::Cancelled);
}
let first = poll_fn(|cx| Poll::Ready(read.as_mut().poll(cx))).await;
if let Poll::Ready(result) = first {
return Ok(result);
}
let timed_read = async {
match remaining_request_timeout {
Some(remaining) => match timeout(remaining, read.as_mut()).await {
Ok(result) => Ok(result),
Err(elapsed) => Err(ReadInterruption::TimedOut(elapsed)),
},
None => Ok(read.await),
}
};
match cancel_handle {
Some(handle) => {
tokio::select! {
biased;
_ = handle.cancel_token.cancelled() => Err(ReadInterruption::Cancelled),
result = timed_read => result,
}
}
None => timed_read.await,
}
}
#[cfg(windows)]
use crate::connection::transport::localdb::resolve_localdb_instance;
#[cfg(windows)]
use crate::connection::transport::named_pipes::open_named_pipe_with_retry;
pub(crate) const PRE_NEGOTIATED_PACKET_SIZE: u32 = 4096;
async fn create_base_stream(
ipaddress_preference: IPAddressPreference,
transport_context: &TransportContext,
keep_alive_in_ms: u32,
keep_alive_interval_in_ms: u32,
multi_subnet_failover: bool,
connect_timeout_ms: u64,
) -> TdsResult<Box<dyn Stream>> {
match transport_context {
TransportContext::Tcp { host, port, .. } => {
if multi_subnet_failover {
create_base_stream_parallel(
host,
*port,
keep_alive_in_ms,
keep_alive_interval_in_ms,
connect_timeout_ms,
)
.await
} else {
create_base_stream_sequential(
ipaddress_preference,
host,
*port,
keep_alive_in_ms,
keep_alive_interval_in_ms,
connect_timeout_ms,
)
.await
}
}
#[cfg(windows)]
TransportContext::NamedPipe { pipe_name } => {
if multi_subnet_failover {
return Err(crate::error::Error::UsageError(
"MultiSubnetFailover is only supported with TCP connections. \
Named Pipes do not support MultiSubnetFailover."
.to_string(),
));
}
info!("Connecting to Named Pipe: {}", pipe_name);
let pipe_client = open_named_pipe_with_retry(pipe_name).await?;
info!("Connected to Named Pipe: {}", pipe_name);
Ok(Box::new(pipe_client))
}
#[cfg(not(windows))]
TransportContext::NamedPipe { .. } => Err(crate::error::Error::from(std::io::Error::new(
std::io::ErrorKind::Unsupported,
"Named Pipes are only supported on Windows",
))),
#[cfg(windows)]
TransportContext::SharedMemory { instance_name } => {
if multi_subnet_failover {
return Err(crate::error::Error::UsageError(
"MultiSubnetFailover is only supported with TCP connections. \
Shared Memory does not support MultiSubnetFailover."
.to_string(),
));
}
let actual_instance = if instance_name.is_empty() {
"MSSQLSERVER"
} else {
instance_name.as_str()
};
info!(
"Connecting via Shared Memory (LPC-over-Named-Pipes) to instance: {}",
actual_instance
);
let pipe_name = format!(r"\\.\pipe\SQLLocal\{actual_instance}");
info!("Connecting to Shared Memory pipe: {}", pipe_name);
let pipe_client = open_named_pipe_with_retry(&pipe_name).await?;
info!("Connected to Shared Memory (LPC-over-NP): {}", pipe_name);
Ok(Box::new(pipe_client))
}
#[cfg(not(windows))]
TransportContext::SharedMemory { .. } => {
Err(crate::error::Error::from(std::io::Error::new(
std::io::ErrorKind::Unsupported,
"Shared Memory is only supported on Windows",
)))
}
#[cfg(windows)]
TransportContext::LocalDB { instance_name } => {
if multi_subnet_failover {
return Err(crate::error::Error::UsageError(
"MultiSubnetFailover is only supported with TCP connections. \
LocalDB does not support MultiSubnetFailover."
.to_string(),
));
}
info!("Connecting to LocalDB instance: {}", instance_name);
let pipe_name = resolve_localdb_instance(instance_name).await?;
info!("LocalDB instance resolved to pipe: {}", pipe_name);
let pipe_client = open_named_pipe_with_retry(&pipe_name).await?;
info!("Connected to LocalDB instance: {}", instance_name);
Ok(Box::new(pipe_client))
}
}
}
fn sort_by_ip_preference(
socket_addresses: &mut [SocketAddr],
ipaddress_preference: IPAddressPreference,
) {
match ipaddress_preference {
IPAddressPreference::UsePlatformDefault => {
trace!("Using platform default IP address preference");
}
IPAddressPreference::IPv4First => {
socket_addresses.sort_by_key(|a| a.is_ipv6());
trace!("IPv4 addresses first");
}
IPAddressPreference::IPv6First => {
socket_addresses.sort_by_key(|b| std::cmp::Reverse(b.is_ipv6()));
trace!("IPv6 addresses first");
}
}
}
async fn create_base_stream_sequential(
ipaddress_preference: IPAddressPreference,
host: &str,
port: u16,
keep_alive_in_ms: u32,
keep_alive_interval_in_ms: u32,
connect_timeout_ms: u64,
) -> TdsResult<Box<dyn Stream>> {
info!(
"Connecting to TCP transport (sequential): {}:{}",
host, port
);
let mut socket_addresses: Vec<SocketAddr> =
tokio::net::lookup_host((host, port)).await?.collect();
let mut last_error = None;
let mut tcp_stream = None;
sort_by_ip_preference(&mut socket_addresses, ipaddress_preference);
info!("Socket addresses: {:?}", socket_addresses);
for socket_address in socket_addresses {
let socket = if socket_address.is_ipv6() {
net::TcpSocket::new_v6()?
} else {
net::TcpSocket::new_v4()?
};
let keep_alive_settings = socket2::TcpKeepalive::new()
.with_time(Duration::from_millis(keep_alive_in_ms as u64))
.with_interval(Duration::from_millis(keep_alive_interval_in_ms as u64));
let socket2_socket = socket2::SockRef::from(&socket);
socket2_socket.set_tcp_keepalive(&keep_alive_settings)?;
socket2_socket.set_nodelay(true)?;
let connect_future = socket.connect(socket_address);
tcp_stream = match timeout(Duration::from_millis(connect_timeout_ms), connect_future).await
{
Ok(Ok(stream)) => {
info!("Connected to TCP transport: {}:{}", host, port);
Some(stream)
}
Ok(Err(e)) => {
last_error = Some(e);
None
}
Err(_elapsed) => {
last_error = Some(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!(
"Connection to {} timed out after {}ms",
socket_address, connect_timeout_ms
),
));
None
}
};
if tcp_stream.is_some() {
break;
}
}
if let Some(stream) = tcp_stream {
Ok(Box::new(stream))
} else {
Err(crate::error::Error::from(last_error.unwrap_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::NotConnected,
format!("Failed to connect to {host}:{port}"),
)
})))
}
}
async fn create_base_stream_parallel(
host: &str,
port: u16,
keep_alive_in_ms: u32,
keep_alive_interval_in_ms: u32,
connect_timeout_ms: u64,
) -> TdsResult<Box<dyn Stream>> {
info!(
"Connecting to TCP transport (parallel/MultiSubnetFailover): {}:{}",
host, port
);
let config = ParallelConnectConfig {
timeout_ms: connect_timeout_ms,
keep_alive_in_ms,
keep_alive_interval_in_ms,
};
let result = parallel_connect(host, port, &config).await?;
info!(
"Parallel connection succeeded to {} (tried {} addresses, {} failed)",
result.connected_address, result.total_addresses, result.failed_attempts
);
Ok(Box::new(result.stream))
}
async fn create_transport_for_version(
stream: Box<dyn Stream>,
tds_version: TdsVersion,
transport_context: &TransportContext,
encryption_options: EncryptionOptions,
encryption_mode: EncryptionSetting,
) -> TdsResult<NetworkTransport> {
let ssl_handler = SslHandler {
server_host_name: transport_context.get_server_name().to_string(),
encryption_options,
};
match tds_version {
TdsVersion::V7_4 => {
info!("Creating NetworkTransport for TDS 7.4 with TLS wrapping");
Ok(NetworkTransport::new(
stream,
ssl_handler,
PRE_NEGOTIATED_PACKET_SIZE,
encryption_mode,
true, ))
}
TdsVersion::V8_0 => {
info!("Creating NetworkTransport for TDS 8.0 with immediate TLS");
let encrypted_stream = ssl_handler
.enable_ssl_async(stream, NegotiatedEncryptionSetting::Strict)
.await?;
Ok(NetworkTransport::new(
encrypted_stream,
ssl_handler,
PRE_NEGOTIATED_PACKET_SIZE,
encryption_mode,
false, ))
}
TdsVersion::Unknown(version_value) => Err(crate::error::Error::ProtocolError(format!(
"Unsupported TDS version: 0x{version_value:08X}. Only TDS 7.4 and TDS 8.0 are supported."
))),
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn create_transport(
ipaddress_preference: IPAddressPreference,
tds_version: TdsVersion,
transport_context: &TransportContext,
encryption_options: EncryptionOptions,
keep_alive_in_ms: u32,
keep_alive_interval_in_ms: u32,
multi_subnet_failover: bool,
connect_timeout_ms: u64,
) -> TdsResult<NetworkTransport> {
let encryption_mode = encryption_options.mode;
let stream = create_base_stream(
ipaddress_preference,
transport_context,
keep_alive_in_ms,
keep_alive_interval_in_ms,
multi_subnet_failover,
connect_timeout_ms,
)
.await?;
create_transport_for_version(
stream,
tds_version,
transport_context,
encryption_options,
encryption_mode,
)
.await
}
#[async_trait]
pub trait TransportSslHandler {
async fn enable_ssl(&mut self) -> TdsResult<()>;
async fn disable_ssl(&mut self) -> TdsResult<()>;
}
pub trait Stream: AsyncRead + AsyncWrite + Unpin + Send + Sync {
fn tls_handshake_starting(&mut self);
fn tls_handshake_completed(&mut self);
fn is_connection_dead(&self) -> bool {
false
}
fn channel_binding_token(&self) -> Option<Vec<u8>> {
None
}
}
impl Stream for TcpStream {
fn tls_handshake_starting(&mut self) {
}
fn tls_handshake_completed(&mut self) {
}
fn is_connection_dead(&self) -> bool {
match self.try_read(&mut [0u8; 1]) {
Err(ref e) if e.kind() == ErrorKind::WouldBlock => false,
Ok(0) => true,
Err(_) => true,
Ok(_) => false,
}
}
}
impl Stream for Box<dyn Stream> {
fn tls_handshake_starting(&mut self) {
(**self).tls_handshake_starting();
}
fn tls_handshake_completed(&mut self) {
(**self).tls_handshake_completed();
}
fn is_connection_dead(&self) -> bool {
(**self).is_connection_dead()
}
fn channel_binding_token(&self) -> Option<Vec<u8>> {
(**self).channel_binding_token()
}
}
#[derive(Clone)]
struct SharedStream {
inner: Arc<Mutex<Box<dyn Stream>>>,
}
impl SharedStream {
fn new(stream: Box<dyn Stream>) -> Self {
Self {
inner: Arc::new(Mutex::new(stream)),
}
}
fn lock(&self) -> MutexGuard<'_, Box<dyn Stream>> {
self.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn into_inner(self) -> TdsResult<Box<dyn Stream>> {
Arc::try_unwrap(self.inner)
.map_err(|_| {
crate::error::Error::ImplementationError(
"Cannot replace a network stream while an I/O operation still holds it"
.to_string(),
)
})?
.into_inner()
.map_err(|error| {
crate::error::Error::ImplementationError(format!(
"Cannot replace a poisoned network stream: {error}"
))
})
}
}
impl AsyncRead for SharedStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut **self.lock()).poll_read(cx, buf)
}
}
impl AsyncWrite for SharedStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut **self.lock()).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut **self.lock()).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut **self.lock()).poll_shutdown(cx)
}
}
impl Stream for SharedStream {
fn tls_handshake_starting(&mut self) {
self.lock().tls_handshake_starting();
}
fn tls_handshake_completed(&mut self) {
self.lock().tls_handshake_completed();
}
fn is_connection_dead(&self) -> bool {
self.lock().is_connection_dead()
}
fn channel_binding_token(&self) -> Option<Vec<u8>> {
self.lock().channel_binding_token()
}
}
async fn send_attention_packet(mut stream: SharedStream) -> TdsResult<()> {
let mut packet = Vec::with_capacity(PacketWriter::PACKET_HEADER_SIZE);
PacketWriter::build_header(
&mut packet,
PacketWriter::PACKET_HEADER_SIZE,
PacketType::Attention,
1,
true,
false,
ResetConnectionMode::None,
)?;
stream.write_all(&packet).await?;
Ok(())
}
async fn send_attention_and_complete_read<F, T>(
stream: SharedStream,
deadline: Instant,
read: Pin<&mut F>,
) -> TdsResult<T>
where
F: Future<Output = TdsResult<T>>,
{
timeout_at(deadline, send_attention_packet(stream))
.await
.map_err(|_| {
TimeoutError(TimeoutErrorType::String(
"Timed out sending the attention packet".to_string(),
))
})??;
timeout_at(deadline, read).await.map_err(|_| {
TimeoutError(TimeoutErrorType::String(
"Timed out finishing the in-flight token after attention".to_string(),
))
})?
}
async fn read_to_attention_boundary<F, T>(
mut read: Pin<&mut F>,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
attention_stream: Option<SharedStream>,
already_dead: bool,
) -> InterruptibleRead<T>
where
F: Future<Output = TdsResult<T>>,
{
let interruption = match await_read_or_interrupt(
read.as_mut(),
remaining_request_timeout,
cancel_handle,
)
.await
{
Ok(result) => return InterruptibleRead::Completed(result),
Err(interruption) => interruption,
};
let error = Box::new(interruption.into_error());
let boundary = if already_dead {
None
} else {
attention_stream.map(|stream| {
let deadline = Instant::now() + Duration::from_secs(ATTENTION_TIMEOUT_SECONDS);
(deadline, stream)
})
};
let boundary = match boundary {
Some((deadline, stream)) => Some((
deadline,
Box::pin(send_attention_and_complete_read(stream, deadline, read)).await,
)),
None => None,
};
InterruptibleRead::Interrupted { error, boundary }
}
struct AttentionDrainContext {
metadata: Option<Arc<ColMetadataToken>>,
column_encryption_supported: bool,
}
pub(crate) struct NetworkTransport {
encryption: Option<NegotiatedEncryptionSetting>,
packet_size: u32,
stream: Option<SharedStream>,
ssl_handler: SslHandler,
encryption_setting: EncryptionSetting,
tds_read_buffer: TdsReadBuffer,
use_tds74_tls_wrapping: bool,
extractable_stream_handle: Option<extractable_stream::ExtractableStreamHandle>,
pending_reset: ResetConnectionMode,
reset_dispatched: bool,
known_dead: bool,
nbc_bitmap_scratch: Option<Arc<[u8]>>,
column_encryption_supported: bool,
attention_settlement: Option<Box<AttentionSettlement>>,
}
impl std::fmt::Debug for NetworkTransport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NetworkTransport")
.field("encryption", &self.encryption)
.field("packet_size", &self.packet_size)
.field("stream", &"<stream>")
.field("ssl_handler", &self.ssl_handler)
.field("encryption_setting", &self.encryption_setting)
.finish()
}
}
impl NetworkReaderWriter for NetworkTransport {
fn notify_encryption_setting_change(&mut self, setting: NegotiatedEncryptionSetting) {
self.notify_encryption_negotiation(setting);
}
fn notify_session_setting_change(&mut self, setting: &SessionSettings) {
self.packet_size = setting.packet_size;
}
fn as_writer(&mut self) -> &mut dyn NetworkWriter {
self
}
}
#[async_trait]
impl NetworkReader for NetworkTransport {
fn packet_size(&self) -> u32 {
self.packet_size
}
}
#[async_trait]
impl NetworkWriter for NetworkTransport {
async fn send(&mut self, data: &[u8]) -> TdsResult<()> {
let stream = self.stream.as_mut().ok_or_else(|| {
crate::error::Error::ConnectionClosed(
"Cannot send: connection has been closed".to_string(),
)
})?;
if let Err(e) = stream.write_all(data).await {
self.known_dead = true;
return Err(e.into());
}
Ok(())
}
fn packet_size(&self) -> u32 {
self.packet_size
}
fn get_encryption_setting(&self) -> NegotiatedEncryptionSetting {
self.encryption
.unwrap_or(NegotiatedEncryptionSetting::NoEncryption)
}
fn set_reset_mode(&mut self, mode: ResetConnectionMode) {
self.pending_reset = mode;
self.reset_dispatched = false;
}
fn take_reset_mode(&mut self) -> ResetConnectionMode {
std::mem::replace(&mut self.pending_reset, ResetConnectionMode::None)
}
fn note_reset_dispatched(&mut self) {
self.reset_dispatched = true;
}
fn take_reset_dispatched(&mut self) -> bool {
std::mem::replace(&mut self.reset_dispatched, false)
}
fn channel_binding_token(&self) -> Option<Vec<u8>> {
self.stream.as_ref()?.channel_binding_token()
}
}
impl NetworkTransport {
pub fn new(
stream: Box<dyn Stream>,
ssl_handler: SslHandler,
packet_size: u32,
encryption_setting: EncryptionSetting,
use_tds74_tls_wrapping: bool,
) -> Self {
Self {
encryption: None,
stream: Some(SharedStream::new(stream)),
ssl_handler,
packet_size,
encryption_setting,
tds_read_buffer: TdsReadBuffer::new(packet_size as usize),
use_tds74_tls_wrapping,
extractable_stream_handle: None,
pending_reset: ResetConnectionMode::None,
reset_dispatched: false,
known_dead: false,
nbc_bitmap_scratch: None,
column_encryption_supported: false,
attention_settlement: None,
}
}
pub(crate) fn notify_encryption_negotiation(
&mut self,
encryption: NegotiatedEncryptionSetting,
) {
assert!(self.encryption.is_none());
self.encryption = Some(encryption);
}
pub(crate) fn take_attention_settlement(&mut self) -> Option<AttentionSettlement> {
self.attention_settlement
.take()
.map(|settlement| *settlement)
}
async fn enable_ssl_internal(&mut self) -> TdsResult<()> {
let base_stream = self
.stream
.take()
.expect("Stream already taken")
.into_inner()?;
let base_stream: Box<dyn Stream> = if self.use_tds74_tls_wrapping {
#[cfg(target_os = "macos")]
{
let tls_over_tds =
crate::connection::transport::ssl_handler::TlsOverTdsStream::new(base_stream);
Box::new(
crate::connection::transport::ssl_handler::BufferedTdsStream::new(tls_over_tds),
)
}
#[cfg(not(target_os = "macos"))]
{
Box::new(
crate::connection::transport::ssl_handler::TlsOverTdsStream::new(base_stream),
)
}
} else {
base_stream
};
let (handle, extractable_stream) =
extractable_stream::ExtractableStreamHandle::new(base_stream);
self.extractable_stream_handle = Some(handle);
let negotiated = self
.encryption
.unwrap_or(NegotiatedEncryptionSetting::Mandatory);
let encrypted_stream = self
.ssl_handler
.enable_ssl_async(Box::new(extractable_stream), negotiated)
.await?;
self.stream = Some(SharedStream::new(encrypted_stream));
Ok(())
}
async fn disable_ssl_internal(&mut self) -> TdsResult<()> {
let encrypted_stream = self
.stream
.take()
.ok_or_else(|| {
crate::error::Error::ImplementationError(
"disable_ssl called but stream is not available".to_string(),
)
})?
.into_inner()?;
std::mem::forget(encrypted_stream);
let handle = self.extractable_stream_handle.take().ok_or_else(|| {
error!("disable_ssl called but enable_ssl was never called");
crate::error::Error::ImplementationError(
"Cannot disable TLS: TLS was never enabled (no extractable stream handle)"
.to_string(),
)
})?;
let base_stream = handle
.extract()
.map_err(|e| {
error!("Failed to lock extractable stream: {e}");
crate::error::Error::ImplementationError(format!("Cannot disable TLS: {e}"))
})?
.ok_or_else(|| {
error!("Failed to extract underlying stream - was disable_ssl called twice?");
crate::error::Error::ImplementationError(
"Cannot disable TLS: underlying stream was already extracted".to_string(),
)
})?;
info!("Successfully disabled TLS, reverting to unencrypted stream");
self.stream = Some(SharedStream::new(base_stream));
Ok(())
}
pub(crate) async fn close_transport(&mut self) -> TdsResult<()> {
if let Some(stream) = self.stream.as_mut() {
stream.shutdown().await?;
}
self.stream = None;
self.known_dead = true;
Ok(())
}
const MAX_CONSECUTIVE_EMPTY_MESSAGES: u32 = 1;
#[cold]
fn stalled_on_empty_messages_error(outstanding: usize) -> crate::error::Error {
crate::error::Error::ProtocolError(format!(
"TDS read stalled: {outstanding} more byte(s) required but the peer sent {} \
consecutive payload-free end-of-message packets without advancing.",
Self::MAX_CONSECUTIVE_EMPTY_MESSAGES + 1
))
}
#[cold]
fn no_progress_error(outstanding: usize) -> crate::error::Error {
crate::error::Error::ProtocolError(format!(
"TDS read made no progress: {outstanding} more byte(s) required but the packet \
carried no payload. The value extends past the end of the message."
))
}
fn message_cannot_satisfy(&self, needed: usize) -> bool {
self.tds_read_buffer.end_of_message
&& needed > self.tds_read_buffer.get_remaining_byte_count()
}
fn value_truncated_by_message_end(&self, consumed_any: bool, needed: usize) -> bool {
(consumed_any || self.tds_read_buffer.get_remaining_byte_count() > 0)
&& self.message_cannot_satisfy(needed)
}
#[cold]
fn past_end_of_message_error(outstanding: usize, available: usize) -> crate::error::Error {
crate::error::Error::ProtocolError(format!(
"TDS read extends past the end of the message: {outstanding} more byte(s) required \
but only {available} remain in the final packet of the message."
))
}
async fn refill_for(&mut self, needed: usize) -> TdsResult<()> {
if self.value_truncated_by_message_end(false, needed) {
return Err(Self::past_end_of_message_error(
needed,
self.tds_read_buffer.get_remaining_byte_count(),
));
}
self.read_tds_packet().await
}
async fn read_tds_packet(&mut self) -> TdsResult<()> {
let remaining_bytes = self.tds_read_buffer.get_remaining_byte_count();
if remaining_bytes > 0 {
self.tds_read_buffer.shift_data_to_front();
let new_packet_size = self.get_new_tds_packet().await?;
self.tds_read_buffer
.remove_header_from_packet(new_packet_size);
} else {
self.tds_read_buffer.reset_to_length(0);
let new_packet_size = self.get_new_tds_packet().await?;
self.tds_read_buffer
.remove_header_from_packet(new_packet_size);
}
Ok(())
}
fn try_read_tds_packet(&mut self) -> TdsResult<bool> {
if self.tds_read_buffer.end_of_message {
return Ok(false);
}
let remaining_bytes = self.tds_read_buffer.get_remaining_byte_count();
if remaining_bytes.saturating_add(self.tds_read_buffer.max_packet_size)
> self.tds_read_buffer.working_buffer.len()
{
return Ok(false);
}
if remaining_bytes > 0 {
self.tds_read_buffer.shift_data_to_front();
} else {
self.tds_read_buffer.reset_to_length(0);
}
let Some(new_packet_size) = self.try_get_new_tds_packet()? else {
return Ok(false);
};
self.tds_read_buffer
.remove_header_from_packet(new_packet_size);
Ok(true)
}
fn move_pending_packet_bytes(&mut self, base_offset: usize) -> TdsResult<usize> {
let bytes_available = self.tds_read_buffer.pending_bytes;
let pending_offset = self.tds_read_buffer.pending_bytes_offset;
if bytes_available > 0 {
let src_end = pending_offset.saturating_add(bytes_available);
let dest_end = base_offset.saturating_add(bytes_available);
let buffer_len = self.tds_read_buffer.working_buffer.len();
if src_end > buffer_len || dest_end > buffer_len {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid pending bytes range: src {}..{}, dest {}, buffer_len {}",
pending_offset, src_end, base_offset, buffer_len
)));
}
self.tds_read_buffer
.working_buffer
.copy_within(pending_offset..src_end, base_offset);
self.tds_read_buffer.pending_bytes = 0;
self.tds_read_buffer.pending_bytes_offset = 0;
}
Ok(bytes_available)
}
fn complete_tds_packet_if_available(
&mut self,
base_offset: usize,
bytes_available: usize,
) -> TdsResult<Option<usize>> {
if bytes_available < PacketWriter::PACKET_HEADER_SIZE {
return Ok(None);
}
let length_from_packet_header = BigEndian::read_u16(
&self.tds_read_buffer.working_buffer[base_offset + 2..base_offset + 4],
);
let packet_size = usize::from(length_from_packet_header);
if packet_size < PacketWriter::PACKET_HEADER_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid TDS packet length {}: must be at least {} bytes (header size)",
packet_size,
PacketWriter::PACKET_HEADER_SIZE
)));
}
if packet_size > self.tds_read_buffer.max_packet_size {
return Err(crate::error::Error::ProtocolError(format!(
"TDS packet length {} exceeds negotiated max packet size {}",
packet_size, self.tds_read_buffer.max_packet_size
)));
}
let buffer_len = self.tds_read_buffer.working_buffer.len();
if base_offset.saturating_add(packet_size) > buffer_len {
return Err(crate::error::Error::ProtocolError(format!(
"TDS packet length {} at offset {} exceeds buffer capacity {}",
packet_size, base_offset, buffer_len
)));
}
if bytes_available < packet_size {
return Ok(None);
}
let is_end_of_message = self.tds_read_buffer.working_buffer[base_offset + 1]
& PacketStatusFlags::Eom as u8
!= 0;
if packet_size == PacketWriter::PACKET_HEADER_SIZE && !is_end_of_message {
return Err(crate::error::Error::ProtocolError(
"Received a payload-free TDS packet that is not end-of-message".to_string(),
));
}
self.tds_read_buffer.end_of_message = is_end_of_message;
let extra_bytes = bytes_available - packet_size;
if extra_bytes > 0 {
self.tds_read_buffer.pending_bytes = extra_bytes;
self.tds_read_buffer.pending_bytes_offset = base_offset + packet_size;
} else {
self.tds_read_buffer.pending_bytes = 0;
self.tds_read_buffer.pending_bytes_offset = 0;
}
event!(
tracing::Level::DEBUG,
"Received packet of size: {:?}",
packet_size
);
use pretty_hex::PrettyHex;
event!(
tracing::Level::DEBUG,
"Packet content: {:?}",
&mut self.tds_read_buffer.working_buffer[base_offset..base_offset + packet_size]
.hex_dump()
);
Ok(Some(packet_size))
}
fn park_partial_packet(&mut self, base_offset: usize, bytes_available: usize) {
debug_assert_eq!(self.tds_read_buffer.pending_bytes, 0);
self.tds_read_buffer.pending_bytes = bytes_available;
self.tds_read_buffer.pending_bytes_offset = base_offset;
}
fn try_get_new_tds_packet(&mut self) -> TdsResult<Option<usize>> {
let base_offset = self.tds_read_buffer.buffer_length;
let mut bytes_available = self.move_pending_packet_bytes(base_offset)?;
let waker = Waker::noop();
let mut context = Context::from_waker(waker);
loop {
if let Some(packet_size) =
self.complete_tds_packet_if_available(base_offset, bytes_available)?
{
return Ok(Some(packet_size));
}
let stream = self.stream.as_mut().ok_or_else(|| {
crate::error::Error::ConnectionClosed(
"Cannot read TDS packet: connection has been closed".to_string(),
)
})?;
let mut read_buffer = ReadBuf::new(
&mut self.tds_read_buffer.working_buffer[base_offset + bytes_available..],
);
match Pin::new(stream).poll_read(&mut context, &mut read_buffer) {
Poll::Pending => {
self.park_partial_packet(base_offset, bytes_available);
return Ok(None);
}
Poll::Ready(Err(error)) => {
self.known_dead = true;
return Err(error.into());
}
Poll::Ready(Ok(())) => {
let bytes_read = read_buffer.filled().len();
if bytes_read == 0 {
self.known_dead = true;
let section = if bytes_available < PacketWriter::PACKET_HEADER_SIZE {
"header"
} else {
"payload"
};
return Err(crate::error::Error::ConnectionClosed(format!(
"Connection closed by server while reading TDS packet {section}"
)));
}
bytes_available += bytes_read;
}
}
}
}
async fn get_new_tds_packet(&mut self) -> TdsResult<usize> {
let base_offset = self.tds_read_buffer.buffer_length;
let mut bytes_available = self.move_pending_packet_bytes(base_offset)?;
loop {
if let Some(packet_size) =
self.complete_tds_packet_if_available(base_offset, bytes_available)?
{
return Ok(packet_size);
}
let stream = self.stream.as_mut().ok_or_else(|| {
crate::error::Error::ConnectionClosed(
"Cannot read TDS packet: connection has been closed".to_string(),
)
})?;
let bytes_read = match stream
.read(&mut self.tds_read_buffer.working_buffer[base_offset + bytes_available..])
.await
{
Ok(bytes_read) => bytes_read,
Err(error) => {
self.known_dead = true;
return Err(error.into());
}
};
if bytes_read == 0 {
self.known_dead = true;
let section = if bytes_available < PacketWriter::PACKET_HEADER_SIZE {
"header"
} else {
"payload"
};
return Err(crate::error::Error::ConnectionClosed(format!(
"Connection closed by server while reading TDS packet {section}"
)));
}
bytes_available += bytes_read;
}
}
async fn send_attention_and_wait(
&mut self,
parser_context: &ParserContext,
attention_timeout: Duration,
) -> TdsResult<bool> {
self.attention_settlement = None;
let deadline = Instant::now() + attention_timeout;
match timeout_at(deadline, self.cancel_read_stream()).await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
self.known_dead = true;
return Err(e);
}
Err(_elapsed) => {
warn!(
timeout = ?attention_timeout,
"Timed out sending the attention packet; marking the connection dead"
);
self.known_dead = true;
return Ok(false);
}
}
match self
.wait_for_attention_ack(parser_context, None, deadline)
.await
{
Ok(true) => Ok(true),
Ok(false) => {
warn!(
timeout = ?attention_timeout,
"Attention went unacknowledged within the bound; marking the connection dead"
);
self.known_dead = true;
Ok(false)
}
Err(e) => {
self.known_dead = true;
Err(e)
}
}
}
async fn wait_for_attention_ack(
&mut self,
parser_context: &ParserContext,
first_token: Option<Tokens>,
deadline: Instant,
) -> TdsResult<bool> {
let start = Instant::now();
match timeout_at(
deadline,
self.drain_to_attention_ack(parser_context, first_token),
)
.await
{
Ok(Ok(settlement)) => {
self.attention_settlement = Some(Box::new(settlement));
debug!("Attention ACK received after {:?}", start.elapsed());
Ok(true)
}
Ok(Err(e)) => Err(e),
Err(_elapsed) => {
debug!("Attention ACK timeout after {:?}", start.elapsed());
Ok(false)
}
}
}
fn attention_drain_context(&self, parser_context: &ParserContext) -> AttentionDrainContext {
match parser_context {
ParserContext::ColumnMetadata(metadata, _) => AttentionDrainContext {
metadata: Some(Arc::clone(metadata)),
column_encryption_supported: self.column_encryption_supported,
},
ParserContext::ColumnEncryption(enabled) => AttentionDrainContext {
metadata: None,
column_encryption_supported: *enabled,
},
ParserContext::None(()) => AttentionDrainContext {
metadata: None,
column_encryption_supported: self.column_encryption_supported,
},
}
}
fn apply_attention_token(
context: &mut AttentionDrainContext,
settlement: &mut AttentionSettlement,
token: Tokens,
) -> bool {
match token {
Tokens::ColMetadata(metadata) => {
context.metadata = Some(Arc::new(metadata));
false
}
Tokens::Row(_) => false,
Tokens::Done(done) => {
let acknowledged = done.status.contains(DoneStatus::ATTN) && !done.has_more();
context.metadata = None;
settlement.push(Tokens::Done(done));
acknowledged
}
Tokens::DoneProc(done) => {
let acknowledged = done.status.contains(DoneStatus::ATTN) && !done.has_more();
context.metadata = None;
settlement.push(Tokens::DoneProc(done));
acknowledged
}
Tokens::DoneInProc(done) => {
let acknowledged = done.status.contains(DoneStatus::ATTN) && !done.has_more();
context.metadata = None;
settlement.push(Tokens::DoneInProc(done));
acknowledged
}
token => {
settlement.push(token);
false
}
}
}
async fn discard_paused_row(&mut self, pause_state: RowPauseState) -> TdsResult<()> {
let mut writer = crate::datatypes::row_writer::DiscardRowWriter;
match resume_row_into_internal(self, pause_state, ColumnPolicy::SkipAll, &mut writer)
.await?
{
RowReadResult::RowWritten => Ok(()),
RowReadResult::RowPaused(_) | RowReadResult::PlpPaused(_) => {
Err(crate::error::Error::ProtocolError(
"Attention drain paused while discarding a row".to_string(),
))
}
RowReadResult::Token(_) => Err(crate::error::Error::ProtocolError(
"Attention drain reached a control token inside a row".to_string(),
)),
}
}
async fn discard_paused_plp(&mut self, mut plp_state: PlpPauseState) -> TdsResult<()> {
let mut buffer = vec![0u8; 8192];
while !plp_state.reached_end() {
let read = read_active_plp_bytes_internal(self, &mut plp_state, &mut buffer).await?;
if read == 0 && !plp_state.reached_end() {
return Err(crate::error::Error::ProtocolError(
"Attention drain made no progress while discarding a PLP value".to_string(),
));
}
}
self.discard_paused_row(plp_state.row_pause_state).await
}
async fn discard_active_plp(&mut self, plp_state: &mut PlpPauseState) -> TdsResult<()> {
let mut buffer = vec![0u8; 8192];
while !plp_state.reached_end() {
let read = read_active_plp_bytes_internal(self, plp_state, &mut buffer).await?;
if read == 0 && !plp_state.reached_end() {
return Err(crate::error::Error::ProtocolError(
"Attention drain made no progress while discarding an active PLP value"
.to_string(),
));
}
}
self.discard_paused_row(plp_state.row_pause_state.clone())
.await
}
async fn discard_interrupted_row_result(
&mut self,
result: RowReadResult,
) -> TdsResult<Option<Tokens>> {
match result {
RowReadResult::RowWritten => Ok(None),
RowReadResult::Token(token) => Ok(Some(token)),
RowReadResult::RowPaused(pause_state) => {
self.discard_paused_row(pause_state).await?;
Ok(None)
}
RowReadResult::PlpPaused(plp_state) => {
self.discard_paused_plp(plp_state).await?;
Ok(None)
}
}
}
async fn wait_for_attention_after_row_result(
&mut self,
parser_context: &ParserContext,
result: RowReadResult,
deadline: Instant,
) -> TdsResult<bool> {
let first_token =
match timeout_at(deadline, self.discard_interrupted_row_result(result)).await {
Ok(result) => result?,
Err(_) => return Ok(false),
};
self.wait_for_attention_ack(parser_context, first_token, deadline)
.await
}
async fn wait_for_attention_after_row_header(
&mut self,
parser_context: &ParserContext,
header: RowHeader,
deadline: Instant,
) -> TdsResult<bool> {
let first_token = match header {
RowHeader::Positioned(pause_state) => {
match timeout_at(deadline, self.discard_paused_row(pause_state)).await {
Ok(result) => result?,
Err(_) => return Ok(false),
}
None
}
RowHeader::Token(token) => Some(token),
};
self.wait_for_attention_ack(parser_context, first_token, deadline)
.await
}
async fn wait_for_attention_after_plp(
&mut self,
parser_context: &ParserContext,
plp_state: &mut PlpPauseState,
deadline: Instant,
) -> TdsResult<bool> {
match timeout_at(deadline, self.discard_active_plp(plp_state)).await {
Ok(result) => result?,
Err(_) => return Ok(false),
}
self.wait_for_attention_ack(parser_context, None, deadline)
.await
}
async fn drain_to_attention_ack(
&mut self,
parser_context: &ParserContext,
first_token: Option<Tokens>,
) -> TdsResult<AttentionSettlement> {
let mut context = self.attention_drain_context(parser_context);
let mut settlement = AttentionSettlement::default();
if let Some(token) = first_token
&& Self::apply_attention_token(&mut context, &mut settlement, token)
{
return Ok(settlement);
}
loop {
if let Some(metadata) = context.metadata.as_ref().cloned() {
let parser_context = ParserContext::ColumnMetadata(metadata, None);
let mut writer = crate::datatypes::row_writer::DiscardRowWriter;
let mut nbc_bitmap_scratch = self.nbc_bitmap_scratch.take();
let result = receive_row_into_internal(
self,
&*PARSER_REGISTRY,
&parser_context,
ColumnPolicy::SkipAll,
&mut writer,
&mut nbc_bitmap_scratch,
)
.await;
self.nbc_bitmap_scratch = nbc_bitmap_scratch;
if let Some(token) = self.discard_interrupted_row_result(result?).await?
&& Self::apply_attention_token(&mut context, &mut settlement, token)
{
return Ok(settlement);
}
continue;
}
let parser_context =
ParserContext::ColumnEncryption(context.column_encryption_supported);
let token = receive_token_internal(self, &*PARSER_REGISTRY, &parser_context).await?;
if Self::apply_attention_token(&mut context, &mut settlement, token) {
return Ok(settlement);
}
}
}
}
#[async_trait]
impl TransportSslHandler for NetworkTransport {
async fn enable_ssl(&mut self) -> TdsResult<()> {
self.enable_ssl_internal().await
}
async fn disable_ssl(&mut self) -> TdsResult<()> {
if self.encryption_setting == EncryptionSetting::Strict {
return Err(crate::error::Error::from(Error::new(
std::io::ErrorKind::InvalidInput,
"Under strict mode the client must communicate over TLS",
)));
}
self.disable_ssl_internal().await
}
}
impl TdsPacketReader for NetworkTransport {
fn reset_reader(&mut self) {
let unread = self
.tds_read_buffer
.buffer_length
.saturating_sub(self.tds_read_buffer.buffer_position);
if unread > 0 {
warn!(
unread_bytes = unread,
"Discarding unread bytes while resetting the packet reader"
);
}
self.tds_read_buffer
.change_packet_size(NetworkReader::packet_size(self));
self.tds_read_buffer.reset_to_length(0);
}
#[inline(always)]
fn try_read_byte(&mut self) -> Option<u8> {
self.tds_read_buffer.try_read_byte()
}
#[inline(always)]
fn try_read_int16(&mut self) -> Option<i16> {
self.tds_read_buffer.try_read_int16()
}
#[inline(always)]
fn try_read_uint16(&mut self) -> Option<u16> {
self.tds_read_buffer.try_read_uint16()
}
#[inline(always)]
fn try_read_slice(&mut self, length: usize) -> Option<&[u8]> {
self.tds_read_buffer.try_read_slice(length)
}
#[inline(always)]
fn try_read_uint24(&mut self) -> Option<u32> {
self.tds_read_buffer.try_read_uint24()
}
#[inline(always)]
fn try_read_int32(&mut self) -> Option<i32> {
self.tds_read_buffer.try_read_int32()
}
#[inline(always)]
fn try_read_uint32(&mut self) -> Option<u32> {
self.tds_read_buffer.try_read_uint32()
}
#[inline(always)]
fn try_read_uint40(&mut self) -> Option<u64> {
self.tds_read_buffer.try_read_uint40()
}
#[inline(always)]
fn try_read_int64(&mut self) -> Option<i64> {
self.tds_read_buffer.try_read_int64()
}
#[inline(always)]
fn try_read_float32(&mut self) -> Option<f32> {
self.tds_read_buffer.try_read_float32()
}
#[inline(always)]
fn try_read_float64(&mut self) -> Option<f64> {
self.tds_read_buffer.try_read_float64()
}
async fn read_byte(&mut self) -> TdsResult<u8> {
loop {
if let Some(value) = self.try_read_byte() {
return Ok(value);
}
self.refill_for(1).await?;
}
}
async fn read_int16_big_endian(&mut self) -> TdsResult<i16> {
while !self.tds_read_buffer.do_we_have_enough_data(2) {
self.refill_for(2).await?;
}
let result = BigEndian::read_i16(self.tds_read_buffer.get_slice());
self.tds_read_buffer.consume_bytes(2)?;
Ok(result)
}
async fn read_int32_big_endian(&mut self) -> TdsResult<i32> {
while !self.tds_read_buffer.do_we_have_enough_data(4) {
self.refill_for(4).await?;
}
let result = BigEndian::read_i32(self.tds_read_buffer.get_slice());
self.tds_read_buffer.consume_bytes(4)?;
Ok(result)
}
async fn read_uint40(&mut self) -> TdsResult<u64> {
loop {
if let Some(value) = self.try_read_uint40() {
return Ok(value);
}
self.refill_for(5).await?;
}
}
async fn read_float32(&mut self) -> TdsResult<f32> {
loop {
if let Some(value) = self.try_read_float32() {
return Ok(value);
}
self.refill_for(4).await?;
}
}
async fn read_float64(&mut self) -> TdsResult<f64> {
loop {
if let Some(value) = self.try_read_float64() {
return Ok(value);
}
self.refill_for(8).await?;
}
}
async fn read_int16(&mut self) -> TdsResult<i16> {
loop {
if let Some(value) = self.try_read_int16() {
return Ok(value);
}
self.refill_for(2).await?;
}
}
async fn read_uint16(&mut self) -> TdsResult<u16> {
loop {
if let Some(value) = self.try_read_uint16() {
return Ok(value);
}
self.refill_for(2).await?;
}
}
async fn read_uint24(&mut self) -> TdsResult<u32> {
loop {
if let Some(value) = self.try_read_uint24() {
return Ok(value);
}
self.refill_for(3).await?;
}
}
async fn read_int32(&mut self) -> TdsResult<i32> {
loop {
if let Some(value) = self.try_read_int32() {
return Ok(value);
}
self.refill_for(4).await?;
}
}
async fn read_uint32(&mut self) -> TdsResult<u32> {
loop {
if let Some(value) = self.try_read_uint32() {
return Ok(value);
}
self.refill_for(4).await?;
}
}
async fn read_int64(&mut self) -> TdsResult<i64> {
loop {
if let Some(value) = self.try_read_int64() {
return Ok(value);
}
self.refill_for(8).await?;
}
}
async fn read_uint64(&mut self) -> TdsResult<u64> {
while !self.tds_read_buffer.do_we_have_enough_data(8) {
self.refill_for(8).await?;
}
let result = LittleEndian::read_u64(self.tds_read_buffer.get_slice());
self.tds_read_buffer.consume_bytes(8)?;
Ok(result)
}
async fn read_bytes(&mut self, buffer: &mut [u8]) -> TdsResult<usize> {
let mut total_read = 0;
let mut length_to_read = buffer.len();
let mut offset = 0;
let mut empty_messages = 0u32;
while length_to_read > 0 {
if self.value_truncated_by_message_end(total_read > 0, length_to_read) {
return Err(Self::past_end_of_message_error(
length_to_read,
self.tds_read_buffer.get_remaining_byte_count(),
));
}
if !self
.tds_read_buffer
.do_we_have_enough_data(min(self.tds_read_buffer.max_packet_size, length_to_read))
{
self.read_tds_packet().await?;
}
let available = self.tds_read_buffer.get_remaining_byte_count();
let to_read = min(
available,
min(length_to_read, self.tds_read_buffer.max_packet_size - 8),
);
if to_read == 0 {
if !self.tds_read_buffer.end_of_message {
return Err(Self::no_progress_error(length_to_read));
}
if total_read > 0 {
return Err(Self::past_end_of_message_error(length_to_read, available));
}
empty_messages += 1;
if empty_messages > Self::MAX_CONSECUTIVE_EMPTY_MESSAGES {
return Err(Self::stalled_on_empty_messages_error(length_to_read));
}
continue;
}
buffer[offset..offset + to_read].copy_from_slice(
&self.tds_read_buffer.working_buffer[self.tds_read_buffer.buffer_position
..self.tds_read_buffer.buffer_position + to_read],
);
offset += to_read;
length_to_read -= to_read;
total_read += to_read;
self.tds_read_buffer.consume_bytes(to_read)?;
}
Ok(total_read)
}
async fn read_bytes_uninit(
&mut self,
buffer: &mut [std::mem::MaybeUninit<u8>],
) -> TdsResult<usize> {
let mut total_read = 0;
let mut length_to_read = buffer.len();
let mut offset = 0;
let mut empty_messages = 0u32;
while length_to_read > 0 {
if self.value_truncated_by_message_end(total_read > 0, length_to_read) {
return Err(Self::past_end_of_message_error(
length_to_read,
self.tds_read_buffer.get_remaining_byte_count(),
));
}
if !self
.tds_read_buffer
.do_we_have_enough_data(min(self.tds_read_buffer.max_packet_size, length_to_read))
{
self.read_tds_packet().await?;
}
let available = self.tds_read_buffer.get_remaining_byte_count();
let to_read = min(
available,
min(length_to_read, self.tds_read_buffer.max_packet_size - 8),
);
if to_read == 0 {
if !self.tds_read_buffer.end_of_message {
return Err(Self::no_progress_error(length_to_read));
}
if total_read > 0 {
return Err(Self::past_end_of_message_error(length_to_read, available));
}
empty_messages += 1;
if empty_messages > Self::MAX_CONSECUTIVE_EMPTY_MESSAGES {
return Err(Self::stalled_on_empty_messages_error(length_to_read));
}
continue;
}
let source = &self.tds_read_buffer.working_buffer[self.tds_read_buffer.buffer_position
..self.tds_read_buffer.buffer_position + to_read];
unsafe {
std::ptr::copy_nonoverlapping(
source.as_ptr(),
buffer.as_mut_ptr().cast::<u8>().add(offset),
to_read,
);
}
offset += to_read;
length_to_read -= to_read;
total_read += to_read;
self.tds_read_buffer.consume_bytes(to_read)?;
}
Ok(total_read)
}
async fn read_u8_varbyte(&mut self) -> TdsResult<Vec<u8>> {
let length: u8 = self.read_byte().await?;
let mut result: Vec<u8> = vec![0; length as usize];
self.read_bytes(&mut result[0..]).await?;
Ok(result)
}
async fn read_u16_varbyte(&mut self) -> TdsResult<Vec<u8>> {
let length: u16 = self.read_uint16().await?;
let mut result: Vec<u8> = vec![0; length as usize];
self.read_bytes(&mut result[0..]).await?;
Ok(result)
}
async fn read_varchar_u16_length(&mut self) -> TdsResult<Option<String>> {
let length: u16 = self.read_uint16().await?;
if length == LENGTH_NULL {
return Ok(None);
}
let string = self
.read_unicode_with_byte_length((length as usize) << 1)
.await?;
Ok(Some(string))
}
async fn read_varchar_u8_length(&mut self) -> TdsResult<String> {
let length: u8 = self.read_byte().await?;
let string = self
.read_unicode_with_byte_length((length as usize) << 1)
.await?;
Ok(string)
}
async fn read_unicode(&mut self, string_length: usize) -> TdsResult<String> {
let result = self
.read_unicode_with_byte_length(string_length * 2)
.await?;
Ok(result)
}
async fn read_unicode_with_byte_length(&mut self, byte_length: usize) -> TdsResult<String> {
let mut byte_buffer: Vec<u8> = vec![0; byte_length];
let _ = self.read_bytes(&mut byte_buffer[0..]).await?;
let mut u16_buffer = Vec::with_capacity(byte_buffer.len() / 2);
for chunk in byte_buffer.chunks(2) {
let value = u16::from_le_bytes([chunk[0], chunk[1]]);
u16_buffer.push(value);
}
let string =
String::from_utf16(&u16_buffer).map_err(|e| Error::new(ErrorKind::InvalidData, e))?;
Ok(string)
}
async fn skip_bytes(&mut self, skip_count: usize) -> TdsResult<()> {
let mut length_to_read = skip_count;
let mut empty_messages = 0u32;
while length_to_read > 0 {
let skipped_any = length_to_read < skip_count;
if self.value_truncated_by_message_end(skipped_any, length_to_read) {
return Err(Self::past_end_of_message_error(
length_to_read,
self.tds_read_buffer.get_remaining_byte_count(),
));
}
if !self.tds_read_buffer.do_we_have_enough_data(min(
self.tds_read_buffer.max_packet_size - 8,
length_to_read,
)) {
self.read_tds_packet().await?;
}
let available = self.tds_read_buffer.get_remaining_byte_count();
let to_read = min(
available,
min(length_to_read, self.tds_read_buffer.max_packet_size - 8),
);
if to_read == 0 {
if !self.tds_read_buffer.end_of_message {
return Err(Self::no_progress_error(length_to_read));
}
if skipped_any {
return Err(Self::past_end_of_message_error(length_to_read, available));
}
empty_messages += 1;
if empty_messages > Self::MAX_CONSECUTIVE_EMPTY_MESSAGES {
return Err(Self::stalled_on_empty_messages_error(length_to_read));
}
continue;
}
length_to_read -= to_read;
self.tds_read_buffer.consume_bytes(to_read)?;
}
Ok(())
}
async fn cancel_read_stream(&mut self) -> TdsResult<()> {
let attention = AttentionRequest::new();
let mut packet_writer = attention.create_packet_writer(self.as_writer(), None, None);
attention.serialize(&mut packet_writer).await?;
Ok(())
}
}
impl NetworkTransport {
pub(crate) fn try_receive_row_header(
&mut self,
context: &ParserContext,
) -> TdsResult<Option<RowPauseState>> {
let ParserContext::ColumnMetadata(metadata, decryptor) = context else {
return Err(crate::error::Error::ProtocolError(
"Expected ColumnMetadata in context for row decoding".to_string(),
));
};
let buffered = self.tds_read_buffer.get_buffered_slice();
let Some(&token) = buffered.first() else {
return Ok(None);
};
if token == TokenType::Row as u8 {
self.tds_read_buffer.consume_bytes(1)?;
return Ok(Some(RowPauseState {
next_column_index: 0,
metadata: Arc::clone(metadata),
nbc_null_bitmap: None,
decryptor: decryptor.clone(),
}));
}
if token != TokenType::NbcRow as u8 {
return Ok(None);
}
let bitmap_len = metadata.columns.len().div_ceil(8);
let Some(bitmap_bytes) = buffered.get(1..1 + bitmap_len) else {
return Ok(None);
};
let bitmap = if let Some(mut cached) = self.nbc_bitmap_scratch.take()
&& cached.len() == bitmap_len
&& let Some(buffer) = Arc::get_mut(&mut cached)
{
buffer.copy_from_slice(bitmap_bytes);
self.nbc_bitmap_scratch = Some(Arc::clone(&cached));
cached
} else {
let bitmap: Arc<[u8]> = Arc::from(bitmap_bytes);
self.nbc_bitmap_scratch = Some(Arc::clone(&bitmap));
bitmap
};
self.tds_read_buffer.consume_bytes(1 + bitmap_len)?;
Ok(Some(RowPauseState {
next_column_index: 0,
metadata: Arc::clone(metadata),
nbc_null_bitmap: Some(bitmap),
decryptor: decryptor.clone(),
}))
}
pub(crate) fn try_read_buffered_column(
&mut self,
pause_state: &RowPauseState,
target: usize,
) -> TdsResult<Option<ColumnValues>> {
if target != pause_state.next_column_index {
return Ok(None);
}
let Some(metadata) = pause_state.metadata.columns.get(target) else {
return Ok(None);
};
if pause_state
.nbc_null_bitmap
.as_ref()
.is_some_and(|bitmap| bitmap[target / 8] & (1 << (target % 8)) != 0)
{
return Ok(Some(ColumnValues::Null));
}
if pause_state.decryptor.is_some() {
return Ok(None);
}
let decoder = GenericDecoder::default();
let Some((value, used)) =
decoder.try_decode_buffered(self.tds_read_buffer.get_buffered_slice(), metadata)?
else {
return Ok(None);
};
self.tds_read_buffer.consume_bytes(used)?;
Ok(Some(value))
}
pub(crate) fn try_read_buffered_column_with_base(
&mut self,
pause_state: &RowPauseState,
target: usize,
) -> TdsResult<Option<(ColumnValues, Option<TdsDataType>)>> {
if target != pause_state.next_column_index {
return Ok(None);
}
let Some(metadata) = pause_state.metadata.columns.get(target) else {
return Ok(None);
};
if pause_state
.nbc_null_bitmap
.as_ref()
.is_some_and(|bitmap| bitmap[target / 8] & (1 << (target % 8)) != 0)
{
return Ok(Some((ColumnValues::Null, None)));
}
if metadata.data_type != TdsDataType::SsVariant {
return self
.try_read_buffered_column(pause_state, target)
.map(|value| value.map(|value| (value, None)));
}
if pause_state.decryptor.is_some() {
return Ok(None);
}
let decoder = GenericDecoder::default();
let Some((base, value, used)) =
decoder.try_decode_buffered_variant(self.tds_read_buffer.get_buffered_slice())?
else {
return Ok(None);
};
self.tds_read_buffer.consume_bytes(used)?;
Ok(Some((value, base)))
}
pub(crate) fn try_begin_buffered_plp(
&mut self,
pause_state: &RowPauseState,
target: usize,
) -> TdsResult<Option<Option<PlpColumnStream>>> {
let Some(metadata) = pause_state.metadata.columns.get(target) else {
return Ok(None);
};
let Some((stream, used)) = PlpColumnStream::try_begin_buffered(
metadata,
self.tds_read_buffer.get_buffered_slice(),
)?
else {
return Ok(None);
};
self.tds_read_buffer.consume_bytes(used)?;
Ok(Some(stream))
}
pub(crate) fn try_read_complete_buffered_plp_column(
&mut self,
pause_state: &RowPauseState,
target: usize,
out: &mut [u8],
) -> TdsResult<CompleteBufferedPlp> {
let Some(metadata) = pause_state.metadata.columns.get(target) else {
return Ok(None);
};
let buffered = self.tds_read_buffer.get_buffered_slice();
let Some((stream, header_used)) = PlpColumnStream::try_begin_buffered(metadata, buffered)?
else {
return Ok(None);
};
let Some(mut stream) = stream else {
self.tds_read_buffer.consume_bytes(header_used)?;
return Ok(Some(None));
};
let Some(remaining) = buffered.get(header_used..) else {
return Ok(None);
};
let Some((payload_used, written)) = stream.try_read_complete_buffered(remaining, out)?
else {
return Ok(None);
};
let total_used = header_used.checked_add(payload_used).ok_or_else(|| {
crate::error::Error::ProtocolError("Buffered PLP byte count overflowed".to_string())
})?;
let known_total = stream.known_len();
let total_read = stream.total_read();
self.tds_read_buffer.consume_bytes(total_used)?;
Ok(Some(Some((written, known_total, total_read))))
}
pub(crate) fn try_read_buffered_plp(
&mut self,
plp_state: &mut PlpPauseState,
out: &mut [u8],
) -> TdsResult<Option<usize>> {
loop {
if let Some((used, written)) = plp_state
.plp_stream
.try_read_buffered(self.tds_read_buffer.get_buffered_slice(), out)?
{
self.tds_read_buffer.consume_bytes(used)?;
return Ok(Some(written));
}
if !self.try_read_tds_packet()? {
return Ok(None);
}
}
}
#[cfg(test)]
pub(crate) fn try_read_buffered_row_into<W: RowWriter + ?Sized>(
&mut self,
pause_state: &mut RowPauseState,
writer: &mut W,
) -> TdsResult<bool> {
self.try_read_buffered_row_prefix_into(pause_state, usize::MAX, writer)
}
pub(crate) fn try_read_buffered_row_prefix_into<W: RowWriter + ?Sized>(
&mut self,
pause_state: &mut RowPauseState,
end_column: usize,
writer: &mut W,
) -> TdsResult<bool> {
if pause_state.decryptor.is_some() {
return Ok(false);
}
let decoder = GenericDecoder::default();
let mut consumed = 0usize;
let outcome = {
let buffered = self.tds_read_buffer.get_buffered_slice();
let mut outcome = Ok(true);
while pause_state.next_column_index < end_column
&& let Some(metadata) = pause_state
.metadata
.columns
.get(pause_state.next_column_index)
{
let col = pause_state.next_column_index;
if pause_state
.nbc_null_bitmap
.as_ref()
.is_some_and(|bitmap| bitmap[col / 8] & (1 << (col % 8)) != 0)
{
writer.write_null(col);
pause_state.next_column_index += 1;
continue;
}
let Some(remaining) = buffered.get(consumed..) else {
outcome = Err(crate::error::Error::ProtocolError(
"Buffered row decoder consumed past the available data".to_string(),
));
break;
};
match decoder.try_decode_buffered_into(remaining, metadata, col, writer) {
Ok(Some(used)) => {
let Some(next) = consumed.checked_add(used) else {
outcome = Err(crate::error::Error::ProtocolError(
"Buffered row decoder byte count overflowed".to_string(),
));
break;
};
consumed = next;
pause_state.next_column_index += 1;
}
Ok(None) => {
outcome = Ok(false);
break;
}
Err(error) => {
outcome = Err(error);
break;
}
}
}
outcome
};
self.tds_read_buffer.consume_bytes(consumed)?;
outcome.map(|complete| {
complete
&& (pause_state.next_column_index >= end_column
|| pause_state.next_column_index >= pause_state.metadata.columns.len())
})
}
pub(crate) async fn receive_token(
&mut self,
context: &ParserContext,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
) -> TdsResult<Tokens> {
if let ParserContext::ColumnEncryption(enabled) = context {
self.column_encryption_supported = *enabled;
}
self.attention_settlement = None;
let attention_stream = self.stream.as_ref().cloned();
let already_dead = self.known_dead;
let outcome = {
let mut read = std::pin::pin!(receive_token_internal(self, &*PARSER_REGISTRY, context));
read_to_attention_boundary(
read.as_mut(),
remaining_request_timeout,
cancel_handle,
attention_stream,
already_dead,
)
.await
};
let (error, boundary) = match outcome {
InterruptibleRead::Completed(result) => return result,
InterruptibleRead::Interrupted { error, boundary } => (error, boundary),
};
let Some((deadline, completed)) = boundary else {
self.known_dead = true;
return Err(*error);
};
match completed {
Ok(token) => match self
.wait_for_attention_ack(context, Some(token), deadline)
.await
{
Ok(true) => {}
Ok(false) => self.known_dead = true,
Err(attention_error) => {
debug!(
?attention_error,
"Failed to settle an interrupted token read"
);
self.known_dead = true;
}
},
Err(attention_error) => {
debug!(
?attention_error,
"Failed to finish an interrupted token read"
);
self.known_dead = true;
}
}
Err(*error)
}
pub(crate) async fn receive_row_into<W>(
&mut self,
context: &ParserContext,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
plan: ColumnPolicy,
writer: &mut W,
) -> TdsResult<RowReadResult>
where
W: RowWriter + Send + ?Sized,
{
self.attention_settlement = None;
let attention_stream = self.stream.as_ref().cloned();
let already_dead = self.known_dead;
let mut nbc_bitmap_scratch = self.nbc_bitmap_scratch.take();
let outcome = {
let mut read = std::pin::pin!(receive_row_into_internal(
self,
&*PARSER_REGISTRY,
context,
plan,
writer,
&mut nbc_bitmap_scratch,
));
read_to_attention_boundary(
read.as_mut(),
remaining_request_timeout,
cancel_handle,
attention_stream,
already_dead,
)
.await
};
self.nbc_bitmap_scratch = nbc_bitmap_scratch;
let (error, boundary) = match outcome {
InterruptibleRead::Completed(result) => return result,
InterruptibleRead::Interrupted { error, boundary } => (error, boundary),
};
let Some((deadline, completed)) = boundary else {
self.known_dead = true;
return Err(*error);
};
match completed {
Ok(result) => match self
.wait_for_attention_after_row_result(context, result, deadline)
.await
{
Ok(true) => {}
Ok(false) => self.known_dead = true,
Err(attention_error) => {
debug!(?attention_error, "Failed to settle an interrupted row read");
self.known_dead = true;
}
},
Err(attention_error) => {
debug!(?attention_error, "Failed to finish an interrupted row read");
self.known_dead = true;
}
}
Err(*error)
}
pub(crate) async fn receive_row_header(
&mut self,
context: &ParserContext,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
) -> TdsResult<RowHeader> {
self.attention_settlement = None;
let attention_stream = self.stream.as_ref().cloned();
let already_dead = self.known_dead;
let mut nbc_bitmap_scratch = self.nbc_bitmap_scratch.take();
let outcome = {
let mut read = std::pin::pin!(receive_row_header_internal(
self,
&*PARSER_REGISTRY,
context,
&mut nbc_bitmap_scratch,
));
read_to_attention_boundary(
read.as_mut(),
remaining_request_timeout,
cancel_handle,
attention_stream,
already_dead,
)
.await
};
self.nbc_bitmap_scratch = nbc_bitmap_scratch;
let (error, boundary) = match outcome {
InterruptibleRead::Completed(result) => return result,
InterruptibleRead::Interrupted { error, boundary } => (error, boundary),
};
let Some((deadline, completed)) = boundary else {
self.known_dead = true;
return Err(*error);
};
match completed {
Ok(header) => match self
.wait_for_attention_after_row_header(context, header, deadline)
.await
{
Ok(true) => {}
Ok(false) => self.known_dead = true,
Err(attention_error) => {
debug!(
?attention_error,
"Failed to settle an interrupted row header"
);
self.known_dead = true;
}
},
Err(attention_error) => {
debug!(
?attention_error,
"Failed to finish an interrupted row header"
);
self.known_dead = true;
}
}
Err(*error)
}
pub(crate) async fn resume_row_into<W>(
&mut self,
pause_state: RowPauseState,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
plan: ColumnPolicy,
writer: &mut W,
) -> TdsResult<RowReadResult>
where
W: RowWriter + Send + ?Sized,
{
self.attention_settlement = None;
let attention_stream = self.stream.as_ref().cloned();
let already_dead = self.known_dead;
let drain_context = ParserContext::ColumnMetadata(Arc::clone(&pause_state.metadata), None);
let outcome = {
let mut read =
std::pin::pin!(resume_row_into_internal(self, pause_state, plan, writer));
read_to_attention_boundary(
read.as_mut(),
remaining_request_timeout,
cancel_handle,
attention_stream,
already_dead,
)
.await
};
let (error, boundary) = match outcome {
InterruptibleRead::Completed(result) => return result,
InterruptibleRead::Interrupted { error, boundary } => (error, boundary),
};
let Some((deadline, completed)) = boundary else {
self.known_dead = true;
return Err(*error);
};
match completed {
Ok(result) => match self
.wait_for_attention_after_row_result(&drain_context, result, deadline)
.await
{
Ok(true) => {}
Ok(false) => self.known_dead = true,
Err(attention_error) => {
debug!(
?attention_error,
"Failed to settle an interrupted row continuation"
);
self.known_dead = true;
}
},
Err(attention_error) => {
debug!(
?attention_error,
"Failed to finish an interrupted row continuation"
);
self.known_dead = true;
}
}
Err(*error)
}
pub(crate) async fn read_active_plp_bytes(
&mut self,
plp_state: &mut PlpPauseState,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
out: &mut [u8],
) -> TdsResult<usize> {
self.attention_settlement = None;
let attention_stream = self.stream.as_ref().cloned();
let already_dead = self.known_dead;
let drain_context =
ParserContext::ColumnMetadata(Arc::clone(&plp_state.row_pause_state.metadata), None);
let outcome = {
let mut read = std::pin::pin!(read_active_plp_bytes_internal(self, plp_state, out));
read_to_attention_boundary(
read.as_mut(),
remaining_request_timeout,
cancel_handle,
attention_stream,
already_dead,
)
.await
};
let (error, boundary) = match outcome {
InterruptibleRead::Completed(result) => return result,
InterruptibleRead::Interrupted { error, boundary } => (error, boundary),
};
let Some((deadline, completed)) = boundary else {
self.known_dead = true;
return Err(*error);
};
match completed {
Ok(_) => match self
.wait_for_attention_after_plp(&drain_context, plp_state, deadline)
.await
{
Ok(true) => {}
Ok(false) => self.known_dead = true,
Err(attention_error) => {
debug!(?attention_error, "Failed to settle an interrupted PLP read");
self.known_dead = true;
}
},
Err(attention_error) => {
debug!(?attention_error, "Failed to finish an interrupted PLP read");
self.known_dead = true;
}
}
Err(*error)
}
}
#[async_trait]
impl TdsTokenStreamReader for NetworkTransport {
fn try_receive_row_header(
&mut self,
context: &ParserContext,
) -> TdsResult<Option<RowPauseState>> {
NetworkTransport::try_receive_row_header(self, context)
}
fn try_read_buffered_column(
&mut self,
pause_state: &RowPauseState,
target: usize,
) -> TdsResult<Option<ColumnValues>> {
NetworkTransport::try_read_buffered_column(self, pause_state, target)
}
async fn receive_token(
&mut self,
context: &ParserContext,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
) -> TdsResult<Tokens> {
NetworkTransport::receive_token(self, context, remaining_request_timeout, cancel_handle)
.await
}
async fn receive_row_into(
&mut self,
context: &ParserContext,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
plan: ColumnPolicy,
writer: &mut (dyn RowWriter + Send),
) -> TdsResult<RowReadResult> {
NetworkTransport::receive_row_into(
self,
context,
remaining_request_timeout,
cancel_handle,
plan,
writer,
)
.await
}
async fn receive_row_header(
&mut self,
context: &ParserContext,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
) -> TdsResult<RowHeader> {
NetworkTransport::receive_row_header(
self,
context,
remaining_request_timeout,
cancel_handle,
)
.await
}
async fn resume_row_into(
&mut self,
pause_state: RowPauseState,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
plan: ColumnPolicy,
writer: &mut (dyn RowWriter + Send),
) -> TdsResult<RowReadResult> {
NetworkTransport::resume_row_into(
self,
pause_state,
remaining_request_timeout,
cancel_handle,
plan,
writer,
)
.await
}
async fn read_active_plp_bytes(
&mut self,
plp_state: &mut PlpPauseState,
remaining_request_timeout: Option<Duration>,
cancel_handle: Option<&CancelHandle>,
out: &mut [u8],
) -> TdsResult<usize> {
NetworkTransport::read_active_plp_bytes(
self,
plp_state,
remaining_request_timeout,
cancel_handle,
out,
)
.await
}
}
#[async_trait]
impl crate::connection::transport::tds_transport::TdsTransport for NetworkTransport {
fn as_writer_ref(&self) -> &dyn NetworkWriter {
self
}
fn as_writer(&mut self) -> &mut dyn NetworkWriter {
self
}
fn reset_reader(&mut self) {
self.tds_read_buffer.change_packet_size(self.packet_size);
self.tds_read_buffer.reset_to_length(0);
}
fn packet_size(&self) -> u32 {
self.packet_size
}
async fn close_transport(&mut self) -> TdsResult<()> {
if let Some(stream) = self.stream.as_mut() {
stream.shutdown().await?;
}
self.stream = None;
self.known_dead = true;
Ok(())
}
async fn send_attention_with_timeout(
&mut self,
context: &ParserContext,
attention_timeout: Duration,
) -> TdsResult<bool> {
self.send_attention_and_wait(context, attention_timeout)
.await
}
fn is_connection_dead(&self) -> bool {
match &self.stream {
Some(stream) => stream.is_connection_dead(),
None => true,
}
}
fn connection_known_dead(&self) -> bool {
self.known_dead
}
fn mark_known_dead(&mut self) {
self.known_dead = true;
}
}
#[cfg(test)]
pub(crate) mod tests {
use super::*; use crate::connection::client_context::ClientContext;
use crate::connection::transport::network_transport::Stream;
use crate::connection::transport::ssl_handler::SslHandler;
use crate::core::EncryptionOptions;
use crate::datatypes::row_writer::DefaultRowWriter;
use crate::datatypes::sqldatatypes::{TdsDataType, TypeInfo};
use crate::message::messages::PacketType;
use crate::query::metadata::ColumnMetadata;
use crate::test_packet_support::{
TestPacketBuilder, build_duplex_transport, create_network_transport_with_chunked_data,
create_network_transport_with_data, create_network_transport_with_live_peer,
create_network_transport_with_live_peer_capturing_writes, encode_utf16_le,
};
use crate::token::tokens::ColMetadataToken;
use bytes::Bytes;
use futures::SinkExt;
use futures::StreamExt;
use rand::Rng;
use tokio::io::{DuplexStream, duplex};
use tokio_util::codec::{BytesCodec, FramedRead, FramedWrite};
pub(crate) const MAX_BUFFER_SIZE: usize = 8192;
fn int4_row_context(column_count: usize) -> ParserContext {
ParserContext::ColumnMetadata(
Arc::new(ColMetadataToken {
column_count: u16::try_from(column_count).unwrap(),
columns: (0..column_count)
.map(|index| ColumnMetadata {
user_type: 0,
flags: 0,
type_info: TypeInfo::fixed_len(TdsDataType::Int4).unwrap(),
data_type: TdsDataType::Int4,
column_name: format!("value{index}"),
multi_part_name: None,
crypto_metadata: None,
})
.collect(),
cek_table: vec![],
}),
None,
)
}
fn plp_varbinary_metadata() -> Arc<ColMetadataToken> {
Arc::new(ColMetadataToken {
column_count: 1,
columns: vec![ColumnMetadata {
user_type: 0,
flags: 0,
type_info: TypeInfo::partial_len(
TdsDataType::BigVarBinary,
usize::from(u16::MAX),
None,
)
.unwrap(),
data_type: TdsDataType::BigVarBinary,
column_name: "payload".to_string(),
multi_part_name: None,
crypto_metadata: None,
}],
cek_table: vec![],
})
}
fn plp_varbinary_row_context() -> ParserContext {
ParserContext::ColumnMetadata(plp_varbinary_metadata(), None)
}
impl Stream for DuplexStream {
fn tls_handshake_starting(&mut self) {
}
fn tls_handshake_completed(&mut self) {
}
}
pub(crate) fn create_readable_network_transport(
context: &ClientContext,
) -> (NetworkTransport, DuplexStream) {
let (client_side, server_side) = duplex(MAX_BUFFER_SIZE);
let ssl_handler = SslHandler {
server_host_name: context.transport_context.get_server_name().clone(),
encryption_options: context.encryption_options.clone(),
};
(
NetworkTransport::new(
Box::new(client_side),
ssl_handler,
context.packet_size as u32,
context.encryption_options.mode,
false,
),
server_side,
)
}
struct ErroringStream;
impl AsyncRead for ErroringStream {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Poll::Ready(Err(Error::new(
ErrorKind::ConnectionReset,
"synthetic read failure",
)))
}
}
impl AsyncWrite for ErroringStream {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
}
impl Stream for ErroringStream {
fn tls_handshake_starting(&mut self) {}
fn tls_handshake_completed(&mut self) {}
}
struct HookTrackingStream {
inner: DuplexStream,
events: Arc<Mutex<[bool; 4]>>,
}
impl AsyncRead for HookTrackingStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_read(cx, buf)
}
}
impl AsyncWrite for HookTrackingStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.get_mut().inner).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
let this = self.get_mut();
this.events.lock().unwrap()[0] = true;
Pin::new(&mut this.inner).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
let this = self.get_mut();
this.events.lock().unwrap()[1] = true;
Pin::new(&mut this.inner).poll_shutdown(cx)
}
}
impl Stream for HookTrackingStream {
fn tls_handshake_starting(&mut self) {
self.events.lock().unwrap()[2] = true;
}
fn tls_handshake_completed(&mut self) {
self.events.lock().unwrap()[3] = true;
}
fn channel_binding_token(&self) -> Option<Vec<u8>> {
Some(vec![1, 2, 3])
}
}
#[tokio::test]
async fn shared_stream_forwards_hooks_and_requires_exclusive_extraction() {
let (inner, mut peer) = duplex(MAX_BUFFER_SIZE);
let events = Arc::new(Mutex::new([false; 4]));
let mut stream = SharedStream::new(Box::new(HookTrackingStream {
inner,
events: Arc::clone(&events),
}));
stream.write_all(b"x").await.unwrap();
let mut byte = [0_u8; 1];
peer.read_exact(&mut byte).await.unwrap();
assert_eq!(byte, *b"x");
peer.write_all(b"y").await.unwrap();
stream.read_exact(&mut byte).await.unwrap();
assert_eq!(byte, *b"y");
stream.flush().await.unwrap();
assert_eq!(stream.channel_binding_token(), Some(vec![1, 2, 3]));
stream.tls_handshake_starting();
stream.tls_handshake_completed();
stream.shutdown().await.unwrap();
assert_eq!(*events.lock().unwrap(), [true; 4]);
let outstanding = stream.clone();
let error = stream
.into_inner()
.err()
.expect("a shared stream must not be extracted");
assert!(matches!(
error,
crate::error::Error::ImplementationError(message)
if message.contains("still holds it")
));
drop(outstanding.into_inner().unwrap());
}
#[tokio::test]
async fn get_new_tds_packet_surfaces_a_read_error_and_marks_the_connection_dead() {
let context = ClientContext::default();
let ssl_handler = SslHandler {
server_host_name: context.transport_context.get_server_name().clone(),
encryption_options: context.encryption_options.clone(),
};
let mut transport = NetworkTransport::new(
Box::new(ErroringStream),
ssl_handler,
context.packet_size as u32,
context.encryption_options.mode,
false,
);
let err = transport
.read_tds_packet()
.await
.expect_err("a failing stream must surface its read error");
assert!(matches!(err, crate::error::Error::Io(_)));
assert!(transport.known_dead);
}
#[tokio::test]
async fn try_get_new_tds_packet_surfaces_a_read_error_and_marks_the_connection_dead() {
let context = ClientContext::default();
let ssl_handler = SslHandler {
server_host_name: context.transport_context.get_server_name().clone(),
encryption_options: context.encryption_options.clone(),
};
let mut transport = NetworkTransport::new(
Box::new(ErroringStream),
ssl_handler,
context.packet_size as u32,
context.encryption_options.mode,
false,
);
let err = transport
.try_read_tds_packet()
.expect_err("a failing stream must surface its read error");
assert!(matches!(err, crate::error::Error::Io(_)));
assert!(transport.known_dead);
}
#[tokio::test]
async fn test_network_transport_send() {
let context = ClientContext {
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
let (mut transport, server_side) = create_readable_network_transport(&context);
let mut rng = rand::rng();
let data_vector: Vec<u8> = (0..MAX_BUFFER_SIZE).map(|_| rng.random()).collect();
let mut framed_reader = FramedRead::new(server_side, BytesCodec::new());
let result = transport.send(&data_vector[..]).await;
match result {
Ok(_) => {}
Err(e) => panic!("Error sending data: {e}"),
}
let received = framed_reader
.next()
.await
.expect("No data")
.expect("Decode error");
assert_eq!(received.as_ref(), &data_vector[..]);
}
#[test]
fn test_tds_transport_reset_reader_resizes_buffer_after_packet_size_change() {
use crate::connection::transport::tds_transport::TdsTransport;
let initial_packet_size: u32 = 4096;
let negotiated_packet_size: u32 = 8000;
let context = ClientContext {
packet_size: initial_packet_size as u16,
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
let (mut transport, _server_side) = create_readable_network_transport(&context);
assert_eq!(transport.packet_size, initial_packet_size);
assert_eq!(transport.tds_read_buffer.working_buffer.len(), 8192); assert_eq!(transport.tds_read_buffer.max_packet_size, 4096);
transport.packet_size = negotiated_packet_size;
TdsTransport::reset_reader(&mut transport);
assert_eq!(
transport.tds_read_buffer.working_buffer.len(),
16000,
"Buffer should be resized to 8000 * 2 = 16000 bytes after reset_reader()"
);
assert_eq!(
transport.tds_read_buffer.max_packet_size, 8000,
"max_packet_size should be updated to 8000"
);
assert_eq!(
transport.tds_read_buffer.buffer_position, 0,
"buffer_position should be reset to 0"
);
assert_eq!(
transport.tds_read_buffer.buffer_length, 0,
"buffer_length should be reset to 0"
);
}
#[test]
fn test_tds_transport_reset_reader_same_size_preserves_buffer() {
use crate::connection::transport::tds_transport::TdsTransport;
let packet_size: u32 = 4096;
let context = ClientContext {
packet_size: packet_size as u16,
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
let (mut transport, _server_side) = create_readable_network_transport(&context);
let initial_buffer_len = transport.tds_read_buffer.working_buffer.len();
assert_eq!(initial_buffer_len, 8192);
TdsTransport::reset_reader(&mut transport);
assert_eq!(
transport.tds_read_buffer.working_buffer.len(),
initial_buffer_len
);
assert_eq!(transport.tds_read_buffer.max_packet_size, 4096);
}
#[tokio::test]
async fn test_get_new_tds_packet_handles_multiple_packets_in_single_read() {
use byteorder::{BigEndian, ByteOrder};
let packet_size: u32 = 512;
let context = ClientContext {
packet_size: packet_size as u16,
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
let (mut transport, server_side) = create_readable_network_transport(&context);
let packet1_payload = vec![0xAA; 24]; let packet2_payload = vec![0xBB; 24];
let packet1_total_len: u16 = 8 + packet1_payload.len() as u16; let packet2_total_len: u16 = 8 + packet2_payload.len() as u16;
let mut packet1 = vec![0u8; packet1_total_len as usize];
packet1[0] = 0x04; packet1[1] = 0x00; BigEndian::write_u16(&mut packet1[2..4], packet1_total_len);
packet1[4] = 0x00; packet1[5] = 0x00; packet1[6] = 0x01; packet1[7] = 0x00; packet1[8..].copy_from_slice(&packet1_payload);
let mut packet2 = vec![0u8; packet2_total_len as usize];
packet2[0] = 0x04; packet2[1] = 0x01; BigEndian::write_u16(&mut packet2[2..4], packet2_total_len);
packet2[4] = 0x00;
packet2[5] = 0x00;
packet2[6] = 0x02; packet2[7] = 0x00;
packet2[8..].copy_from_slice(&packet2_payload);
let mut combined_data = packet1.clone();
combined_data.extend_from_slice(&packet2);
let mut framed_writer = FramedWrite::new(server_side, BytesCodec::new());
framed_writer
.send(Bytes::copy_from_slice(&combined_data))
.await
.expect("Failed to send test data");
let size1 = transport
.get_new_tds_packet()
.await
.expect("Failed to read first packet");
assert_eq!(
size1, packet1_total_len as usize,
"First packet size mismatch"
);
let packet1_in_buffer = &transport.tds_read_buffer.working_buffer[0..size1];
assert_eq!(
packet1_in_buffer,
&packet1[..],
"First packet content mismatch"
);
assert_eq!(
transport.tds_read_buffer.pending_bytes, packet2_total_len as usize,
"pending_bytes should track the second packet"
);
transport.tds_read_buffer.buffer_length = 0;
transport.tds_read_buffer.buffer_position = 0;
let size2 = transport
.get_new_tds_packet()
.await
.expect("Failed to read second packet");
assert_eq!(
size2, packet2_total_len as usize,
"Second packet size mismatch"
);
let packet2_in_buffer = &transport.tds_read_buffer.working_buffer[0..size2];
assert_eq!(
packet2_in_buffer,
&packet2[..],
"Second packet content mismatch"
);
assert_eq!(
transport.tds_read_buffer.pending_bytes, 0,
"pending_bytes should be 0 after consuming all data"
);
}
#[tokio::test]
async fn test_get_new_tds_packet_single_packet_per_read() {
use byteorder::{BigEndian, ByteOrder};
let packet_size: u32 = 512;
let context = ClientContext {
packet_size: packet_size as u16,
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
let (mut transport, server_side) = create_readable_network_transport(&context);
let payload = vec![0xCC; 100];
let total_len: u16 = 8 + payload.len() as u16;
let mut packet = vec![0u8; total_len as usize];
packet[0] = 0x04; packet[1] = 0x01; BigEndian::write_u16(&mut packet[2..4], total_len);
packet[4] = 0x00;
packet[5] = 0x00;
packet[6] = 0x01;
packet[7] = 0x00;
packet[8..].copy_from_slice(&payload);
let mut framed_writer = FramedWrite::new(server_side, BytesCodec::new());
framed_writer
.send(Bytes::copy_from_slice(&packet))
.await
.expect("Failed to send test data");
let size = transport
.get_new_tds_packet()
.await
.expect("Failed to read packet");
assert_eq!(size, total_len as usize);
let packet_in_buffer = &transport.tds_read_buffer.working_buffer[0..size];
assert_eq!(packet_in_buffer, &packet[..]);
assert_eq!(transport.tds_read_buffer.pending_bytes, 0);
}
fn tabular_packet(payload: &[u8], end_of_message: bool) -> Vec<u8> {
let packet_len = PacketWriter::PACKET_HEADER_SIZE + payload.len();
let mut packet = vec![0; packet_len];
packet[0] = PacketType::TabularResult as u8;
packet[1] = u8::from(end_of_message);
BigEndian::write_u16(&mut packet[2..4], u16::try_from(packet_len).unwrap());
packet[PacketWriter::PACKET_HEADER_SIZE..].copy_from_slice(payload);
packet
}
#[tokio::test]
async fn nonblocking_packet_probe_appends_an_available_packet() {
let context = ClientContext {
packet_size: 512,
..Default::default()
};
let (mut transport, mut server) = create_readable_network_transport(&context);
server
.write_all(&tabular_packet(b"first", false))
.await
.unwrap();
transport.read_tds_packet().await.unwrap();
server
.write_all(&tabular_packet(b"second", true))
.await
.unwrap();
assert!(transport.try_read_tds_packet().unwrap());
assert_eq!(
transport.tds_read_buffer.get_buffered_slice(),
b"firstsecond"
);
}
#[tokio::test]
async fn nonblocking_packet_probe_returns_pending_without_losing_payload() {
let context = ClientContext {
packet_size: 512,
..Default::default()
};
let (mut transport, mut server) = create_readable_network_transport(&context);
server
.write_all(&tabular_packet(b"first", false))
.await
.unwrap();
transport.read_tds_packet().await.unwrap();
assert!(!transport.try_read_tds_packet().unwrap());
assert_eq!(transport.tds_read_buffer.get_buffered_slice(), b"first");
}
#[tokio::test]
async fn buffered_plp_probe_defers_before_exhausting_packet_buffer() {
const PACKET_SIZE: usize = 512;
const PACKET_PAYLOAD: usize = PACKET_SIZE - PacketWriter::PACKET_HEADER_SIZE;
const VALUE_LEN: usize = 1600;
let mut plp_wire = Vec::with_capacity(VALUE_LEN + 8);
plp_wire.extend_from_slice(&u32::try_from(VALUE_LEN).unwrap().to_le_bytes());
plp_wire.extend(std::iter::repeat_n(0xAB, VALUE_LEN));
plp_wire.extend_from_slice(&0_u32.to_le_bytes());
let chunks = plp_wire.chunks(PACKET_PAYLOAD);
let packet_count = chunks.len();
let mut packets = Vec::new();
for (index, chunk) in chunks.enumerate() {
packets.extend_from_slice(&tabular_packet(chunk, index + 1 == packet_count));
}
let context = ClientContext {
packet_size: u16::try_from(PACKET_SIZE).unwrap(),
..Default::default()
};
let (mut transport, mut server) = create_readable_network_transport(&context);
server.write_all(&packets).await.unwrap();
let metadata = ColumnMetadata {
user_type: 0,
flags: 0,
type_info: TypeInfo::partial_len(
TdsDataType::BigVarBinary,
usize::from(u16::MAX),
None,
)
.unwrap(),
data_type: TdsDataType::BigVarBinary,
column_name: "payload".to_string(),
multi_part_name: None,
crypto_metadata: None,
};
let (stream, _) =
PlpColumnStream::try_begin_buffered(&metadata, &(VALUE_LEN as u64).to_le_bytes())
.unwrap()
.unwrap();
let mut plp_state = PlpPauseState {
row_pause_state: RowPauseState {
next_column_index: 1,
metadata: Arc::new(ColMetadataToken {
column_count: 1,
columns: vec![metadata],
cek_table: Vec::new(),
}),
nbc_null_bitmap: None,
decryptor: None,
},
plp_stream: stream.unwrap(),
};
let mut out = vec![0; 1300];
assert!(matches!(
transport.try_read_buffered_plp(&mut plp_state, &mut out),
Ok(None)
));
assert_eq!(
transport
.read_active_plp_bytes(&mut plp_state, None, None, &mut out)
.await
.unwrap(),
out.len()
);
assert!(out.iter().all(|byte| *byte == 0xAB));
}
#[tokio::test]
async fn nonblocking_packet_probe_preserves_a_fragmented_header() {
let context = ClientContext {
packet_size: 512,
..Default::default()
};
let (mut transport, mut server) = create_readable_network_transport(&context);
server
.write_all(&tabular_packet(b"first", false))
.await
.unwrap();
transport.read_tds_packet().await.unwrap();
let second = tabular_packet(b"second", true);
server.write_all(&second[..4]).await.unwrap();
assert!(!transport.try_read_tds_packet().unwrap());
server.write_all(&second[4..]).await.unwrap();
assert!(transport.try_read_tds_packet().unwrap());
assert_eq!(
transport.tds_read_buffer.get_buffered_slice(),
b"firstsecond"
);
}
#[tokio::test]
async fn nonblocking_packet_probe_preserves_a_fragmented_payload() {
let context = ClientContext {
packet_size: 512,
..Default::default()
};
let (mut transport, mut server) = create_readable_network_transport(&context);
server
.write_all(&tabular_packet(b"first", false))
.await
.unwrap();
transport.read_tds_packet().await.unwrap();
let second = tabular_packet(b"second", true);
server.write_all(&second[..10]).await.unwrap();
assert!(!transport.try_read_tds_packet().unwrap());
server.write_all(&second[10..]).await.unwrap();
assert!(transport.try_read_tds_packet().unwrap());
assert_eq!(
transport.tds_read_buffer.get_buffered_slice(),
b"firstsecond"
);
}
#[tokio::test]
async fn nonblocking_packet_probe_stops_at_end_of_message() {
let context = ClientContext {
packet_size: 512,
..Default::default()
};
let (mut transport, mut server) = create_readable_network_transport(&context);
server
.write_all(&tabular_packet(b"only", true))
.await
.unwrap();
transport.read_tds_packet().await.unwrap();
assert!(!transport.try_read_tds_packet().unwrap());
assert_eq!(transport.tds_read_buffer.get_buffered_slice(), b"only");
}
#[tokio::test]
async fn test_multi_packet_coalescing_behavior_only() {
use byteorder::{BigEndian, ByteOrder};
use tokio::time::{Duration, timeout};
let packet_size: u32 = 512;
let context = ClientContext {
packet_size: packet_size as u16,
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
let (mut transport, server_side) = create_readable_network_transport(&context);
let packet1_payload: Vec<u8> = (0..24).map(|i| i as u8).collect(); let packet2_payload: Vec<u8> = (0..24).map(|i| (100 + i) as u8).collect();
let packet1_total_len: u16 = 8 + packet1_payload.len() as u16;
let packet2_total_len: u16 = 8 + packet2_payload.len() as u16;
let mut packet1 = vec![0u8; packet1_total_len as usize];
packet1[0] = 0x04; packet1[1] = 0x00; BigEndian::write_u16(&mut packet1[2..4], packet1_total_len);
packet1[6] = 0x01; packet1[8..].copy_from_slice(&packet1_payload);
let mut packet2 = vec![0u8; packet2_total_len as usize];
packet2[0] = 0x04; packet2[1] = 0x01; BigEndian::write_u16(&mut packet2[2..4], packet2_total_len);
packet2[6] = 0x02; packet2[8..].copy_from_slice(&packet2_payload);
let mut combined_data = packet1.clone();
combined_data.extend_from_slice(&packet2);
let mut framed_writer = FramedWrite::new(server_side, BytesCodec::new());
framed_writer
.send(Bytes::copy_from_slice(&combined_data))
.await
.expect("Failed to send test data");
let size1 = transport
.get_new_tds_packet()
.await
.expect("Failed to read first packet");
assert_eq!(size1, packet1_total_len as usize, "First packet size wrong");
let read_packet1_id = transport.tds_read_buffer.working_buffer[6];
let read_packet1_payload: Vec<u8> =
transport.tds_read_buffer.working_buffer[8..size1].to_vec();
assert_eq!(
&read_packet1_payload[..],
&packet1_payload[..],
"First packet payload corrupted"
);
assert_eq!(read_packet1_id, 0x01, "First packet should have ID=1");
transport.tds_read_buffer.reset_to_length(0);
let read_result = timeout(
Duration::from_millis(500), transport.get_new_tds_packet(),
)
.await;
let size2 = match read_result {
Ok(Ok(size)) => size,
Ok(Err(e)) => panic!("Error reading second packet: {:?}", e),
Err(_elapsed) => {
panic!(
"BUG DETECTED: Timed out waiting for second packet!\n\
The second packet's bytes were discarded after the first read.\n\
This is the multi-packet coalescing bug."
);
}
};
assert_eq!(
size2, packet2_total_len as usize,
"Second packet size wrong"
);
let read_packet2_id = transport.tds_read_buffer.working_buffer[6];
let read_packet2_payload: Vec<u8> =
transport.tds_read_buffer.working_buffer[8..size2].to_vec();
assert_eq!(
&read_packet2_payload[..],
&packet2_payload[..],
"Second packet payload corrupted - got wrong data!"
);
assert_eq!(read_packet2_id, 0x02, "Second packet should have ID=2");
}
#[tokio::test]
async fn test_get_new_tds_packet_bounds_check_on_pending_bytes() {
let packet_size: u32 = 512;
let context = ClientContext {
packet_size: packet_size as u16,
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
let (mut transport, _server_side) = create_readable_network_transport(&context);
let buffer_len = transport.tds_read_buffer.working_buffer.len();
transport.tds_read_buffer.pending_bytes = 100;
transport.tds_read_buffer.pending_bytes_offset = buffer_len + 1000;
let result = transport.get_new_tds_packet().await;
assert!(
result.is_err(),
"Expected error for out-of-bounds pending_bytes_offset"
);
let err = result.unwrap_err();
assert!(
matches!(err, crate::error::Error::ProtocolError(_)),
"Expected ProtocolError, got {:?}",
err
);
transport.tds_read_buffer.reset_to_length(0);
transport.tds_read_buffer.pending_bytes = buffer_len + 500; transport.tds_read_buffer.pending_bytes_offset = 0;
let result = transport.get_new_tds_packet().await;
assert!(
result.is_err(),
"Expected error for oversized pending_bytes"
);
transport.tds_read_buffer.reset_to_length(buffer_len - 10);
transport.tds_read_buffer.pending_bytes = 100; transport.tds_read_buffer.pending_bytes_offset = 0;
let result = transport.get_new_tds_packet().await;
assert!(
result.is_err(),
"Expected error when dest range exceeds buffer"
);
}
#[tokio::test]
async fn test_get_new_tds_packet_validates_packet_length_from_header() {
use byteorder::{BigEndian, ByteOrder};
let packet_size: u32 = 512;
let context = ClientContext {
packet_size: packet_size as u16,
encryption_options: EncryptionOptions {
mode: EncryptionSetting::On,
trust_server_certificate: true,
..EncryptionOptions::default()
},
..Default::default()
};
{
let (mut transport, server_side) = create_readable_network_transport(&context);
let mut malformed_packet = vec![0u8; 16];
malformed_packet[0] = 0x04; malformed_packet[1] = 0x01; BigEndian::write_u16(&mut malformed_packet[2..4], 4); malformed_packet[6] = 0x01;
let mut framed_writer = FramedWrite::new(server_side, BytesCodec::new());
framed_writer
.send(Bytes::copy_from_slice(&malformed_packet))
.await
.expect("Failed to send test data");
let result = transport.get_new_tds_packet().await;
assert!(
result.is_err(),
"Expected error for packet length < header size"
);
let err = result.unwrap_err();
assert!(
matches!(err, crate::error::Error::ProtocolError(_)),
"Expected ProtocolError, got {:?}",
err
);
}
{
let (mut transport, server_side) = create_readable_network_transport(&context);
let mut oversized_packet = vec![0u8; 16];
oversized_packet[0] = 0x04;
oversized_packet[1] = 0x01;
BigEndian::write_u16(&mut oversized_packet[2..4], 60000); oversized_packet[6] = 0x01;
let mut framed_writer = FramedWrite::new(server_side, BytesCodec::new());
framed_writer
.send(Bytes::copy_from_slice(&oversized_packet))
.await
.expect("Failed to send test data");
let result = transport.get_new_tds_packet().await;
assert!(
result.is_err(),
"Expected error for packet length > max_packet_size"
);
let err = result.unwrap_err();
assert!(
matches!(err, crate::error::Error::ProtocolError(_)),
"Expected ProtocolError, got {:?}",
err
);
}
{
let (mut transport, server_side) = create_readable_network_transport(&context);
let payload = vec![0xAA; 24];
let total_len: u16 = 8 + payload.len() as u16;
let mut valid_packet = vec![0u8; total_len as usize];
valid_packet[0] = 0x04;
valid_packet[1] = 0x01;
BigEndian::write_u16(&mut valid_packet[2..4], total_len);
valid_packet[6] = 0x01;
valid_packet[8..].copy_from_slice(&payload);
let mut framed_writer = FramedWrite::new(server_side, BytesCodec::new());
framed_writer
.send(Bytes::copy_from_slice(&valid_packet))
.await
.expect("Failed to send test data");
let result = transport.get_new_tds_packet().await;
assert!(result.is_ok(), "Valid packet should succeed");
assert_eq!(result.unwrap(), total_len as usize);
}
}
mod is_connection_dead_tests {
use super::*;
use crate::connection::transport::tds_transport::TdsTransport;
use tokio::net::TcpListener;
#[tokio::test]
async fn tcp_stream_alive_returns_false() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).await.unwrap();
let _server = listener.accept().await.unwrap();
assert!(!client.is_connection_dead());
}
#[tokio::test]
async fn tcp_stream_server_closed_returns_true() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).await.unwrap();
let (server, _) = listener.accept().await.unwrap();
drop(server);
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(client.is_connection_dead());
}
#[tokio::test]
async fn network_transport_no_stream_returns_true() {
let ssl_handler = SslHandler {
server_host_name: "test".to_string(),
encryption_options: EncryptionOptions::new(),
};
let mut transport = NetworkTransport::new(
Box::new(tokio::io::duplex(64).0),
ssl_handler,
4096,
EncryptionSetting::Strict,
false,
);
transport.stream = None;
assert!(transport.is_connection_dead());
}
#[tokio::test]
async fn network_transport_alive_tcp_returns_false() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).await.unwrap();
let _server = listener.accept().await.unwrap();
let ssl_handler = SslHandler {
server_host_name: "test".to_string(),
encryption_options: EncryptionOptions::new(),
};
let transport = NetworkTransport::new(
Box::new(client),
ssl_handler,
4096,
EncryptionSetting::Strict,
false,
);
assert!(!transport.is_connection_dead());
}
#[tokio::test]
async fn network_transport_dead_tcp_returns_true() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).await.unwrap();
let (server, _) = listener.accept().await.unwrap();
let ssl_handler = SslHandler {
server_host_name: "test".to_string(),
encryption_options: EncryptionOptions::new(),
};
let transport = NetworkTransport::new(
Box::new(client),
ssl_handler,
4096,
EncryptionSetting::Strict,
false,
);
drop(server);
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(transport.is_connection_dead());
}
#[tokio::test]
async fn duplex_stream_uses_default_false() {
let (client_side, _server_side) = duplex(64);
assert!(!client_side.is_connection_dead());
}
#[tokio::test]
async fn box_dyn_stream_delegates_to_inner() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).await.unwrap();
let (server, _) = listener.accept().await.unwrap();
let boxed: Box<dyn Stream> = Box::new(client);
assert!(!boxed.is_connection_dead());
drop(server);
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(boxed.is_connection_dead());
}
}
mod payload_eof_marks_dead_tests {
use super::*;
use crate::connection::transport::tds_transport::TdsTransport;
use tokio::io::AsyncWriteExt;
#[tokio::test]
async fn partial_packet_then_eof_sets_known_dead() {
let (client_side, mut server_side) = duplex(MAX_BUFFER_SIZE);
let ssl_handler = SslHandler {
server_host_name: "test".to_string(),
encryption_options: EncryptionOptions::new(),
};
let mut transport = NetworkTransport::new(
Box::new(client_side),
ssl_handler,
4096,
EncryptionSetting::On,
false,
);
let header = [0x04u8, 0x01, 0x00, 0x10, 0x00, 0x00, 0x01, 0x00];
server_side.write_all(&header).await.unwrap();
server_side.flush().await.unwrap();
drop(server_side);
let result = transport.read_tds_packet().await;
assert!(
result.is_err(),
"reading a truncated packet must return an error"
);
assert!(
transport.connection_known_dead(),
"payload EOF must mark the connection known-dead"
);
}
}
fn generate_random_bytes(length: usize) -> Vec<u8> {
let mut rng = rand::rng();
let mut bytes = vec![0u8; length];
rng.fill(&mut bytes[..]);
bytes
}
#[tokio::test]
async fn test_read_byte() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let byte_value = rand::rng().random::<u8>();
let builder = binding.append_byte(byte_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_byte().await.unwrap(), byte_value);
}
#[tokio::test]
async fn test_read_int16() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let int16_value = rand::rng().random::<i16>();
let builder = binding.append_i16(int16_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_int16().await.unwrap(), int16_value);
}
#[tokio::test]
async fn test_read_uint16() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let uint16_value = rand::rng().random::<u16>();
let builder = binding.append_u16(uint16_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_uint16().await.unwrap(), uint16_value);
}
#[tokio::test]
async fn test_read_int32() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let int32_value = rand::rng().random::<i32>();
let builder = binding.append_i32(int32_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_int32().await.unwrap(), int32_value);
}
#[tokio::test]
async fn test_read_uint32() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let uint32_value = rand::rng().random::<u32>();
let builder = binding.append_u32(uint32_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_uint32().await.unwrap(), uint32_value);
}
#[tokio::test]
async fn test_read_int64() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let int64_value = rand::rng().random::<i64>();
let builder = binding.append_i64(int64_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_int64().await.unwrap(), int64_value);
}
#[tokio::test]
async fn test_read_uint64() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let uint64_value = rand::rng().random::<u64>();
let builder = binding.append_u64(uint64_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_uint64().await.unwrap(), uint64_value);
}
#[tokio::test]
async fn test_read_float32() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let float32_value = rand::rng().random::<f32>();
let builder = binding.append_f32(float32_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_float32().await.unwrap(), float32_value);
}
#[tokio::test]
async fn test_read_float64() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let float64_value = rand::rng().random::<f64>();
let builder = binding.append_f64(float64_value);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_float64().await.unwrap(), float64_value);
}
#[tokio::test]
async fn test_read_unicode() {
let unicode_string = "Hello, world";
let char_count = unicode_string.encode_utf16().count();
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let builder = binding.append_bytes(&encode_utf16_le(unicode_string));
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(
reader.read_unicode(char_count).await.unwrap(),
unicode_string
);
}
#[tokio::test]
async fn test_read_bytes() {
let bytes_len = 2000;
let bytes = generate_random_bytes(bytes_len);
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let builder = binding.append_bytes(&bytes);
let mut reader = create_network_transport_with_data(&builder.build());
let mut buffer = vec![0; bytes_len];
assert_eq!(reader.read_bytes(&mut buffer).await.unwrap(), bytes_len);
assert_eq!(buffer, bytes);
}
#[tokio::test]
async fn test_read_u8_varbyte() {
let bytes_len: u8 = 200;
let data_bytes = generate_random_bytes(bytes_len as usize);
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
binding.append_byte(bytes_len);
let builder = binding.append_bytes(&data_bytes);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_u8_varbyte().await.unwrap(), data_bytes);
}
#[tokio::test]
async fn test_read_u16_varbyte() {
let bytes_len: u16 = 1000;
let data_bytes = generate_random_bytes(bytes_len as usize);
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
binding.append_u16(bytes_len);
let builder = binding.append_bytes(&data_bytes);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_u16_varbyte().await.unwrap(), data_bytes);
}
#[tokio::test]
async fn test_read_varchar_u16_length() {
let unicode_string = "Hello, world";
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
binding.append_u16(unicode_string.encode_utf16().count() as u16);
let builder = binding.append_bytes(&encode_utf16_le(unicode_string));
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(
reader.read_varchar_u16_length().await.unwrap(),
Some(unicode_string.to_string())
);
}
#[tokio::test]
async fn test_read_varchar_u16_length_null() {
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
let builder = binding.append_u16(LENGTH_NULL);
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(reader.read_varchar_u16_length().await.unwrap(), None);
}
#[tokio::test]
async fn test_read_varchar_u8_length() {
let unicode_string = "Hello, world";
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
binding.append_byte(unicode_string.encode_utf16().count() as u8);
let builder = binding.append_bytes(&encode_utf16_le(unicode_string));
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(
reader.read_varchar_u8_length().await.unwrap(),
unicode_string
);
}
#[tokio::test]
async fn test_read_varchar_u8_length_long_string() {
let unicode_string = "a".repeat(200);
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
binding.append_byte(unicode_string.encode_utf16().count() as u8);
let builder = binding.append_bytes(&encode_utf16_le(&unicode_string));
let mut reader = create_network_transport_with_data(&builder.build());
assert_eq!(
reader.read_varchar_u8_length().await.unwrap(),
unicode_string
);
}
#[tokio::test]
async fn test_packet_header_split_across_reads() {
let unicode_string = "Hello, world";
let mut binding = TestPacketBuilder::new(PacketType::PreLogin);
binding.append_byte(unicode_string.encode_utf16().count() as u8);
let builder = binding.append_bytes(&encode_utf16_le(unicode_string));
let mut reader = create_network_transport_with_chunked_data(&builder.build(), 3);
assert_eq!(
reader.read_varchar_u8_length().await.unwrap(),
unicode_string
);
}
#[tokio::test]
async fn test_truncated_packet_reports_error() {
let mut binding = TestPacketBuilder::new(PacketType::TabularResult);
let mut packet = binding
.append_bytes(&[0xAB; 100 - PacketWriter::PACKET_HEADER_SIZE])
.build();
assert_eq!(packet.len(), 100);
BigEndian::write_u16(&mut packet[2..4], 200);
let mut reader = create_network_transport_with_data(&packet);
assert!(reader.read_byte().await.is_err());
}
#[tokio::test]
async fn test_read_value_spanning_packet_boundary() {
let mut first = TestPacketBuilder::new(PacketType::TabularResult);
let mut second = TestPacketBuilder::new(PacketType::TabularResult);
let mut stream = first.continuation().append_bytes(&[0x11, 0x22]).build();
stream.extend_from_slice(&second.append_bytes(&[0x33, 0x44]).build());
let mut reader = create_network_transport_with_data(&stream);
assert_eq!(reader.read_uint32().await.unwrap(), 0x4433_2211);
}
#[tokio::test]
async fn test_read_value_spanning_packet_boundary_fragmented() {
let mut first = TestPacketBuilder::new(PacketType::TabularResult);
let mut second = TestPacketBuilder::new(PacketType::TabularResult);
let mut stream = first.continuation().append_bytes(&[0x11, 0x22]).build();
stream.extend_from_slice(&second.append_bytes(&[0x33, 0x44]).build());
let mut reader = create_network_transport_with_chunked_data(&stream, 3);
assert_eq!(reader.read_uint32().await.unwrap(), 0x4433_2211);
}
#[tokio::test]
async fn buffered_cursor_reads_complete_row_header_and_column() {
let expected = 0x1234_5678_i32;
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let mut payload = vec![TokenType::Row as u8];
payload.extend_from_slice(&expected.to_le_bytes());
let mut reader = create_network_transport_with_data(&packet.append_bytes(&payload).build());
reader.read_tds_packet().await.unwrap();
let pause_state = reader
.try_receive_row_header(&int4_row_context(1))
.unwrap()
.expect("complete buffered row header");
assert_eq!(
reader.try_read_buffered_column(&pause_state, 0).unwrap(),
Some(ColumnValues::Int(expected))
);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 0);
}
#[tokio::test]
async fn buffered_row_writer_finishes_a_complete_row_without_continuation() {
let expected = [0x1234_5678_i32, -42_i32];
let mut payload = vec![TokenType::Row as u8];
payload.extend(expected.iter().flat_map(|value| value.to_le_bytes()));
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let mut reader = create_network_transport_with_data(&packet.append_bytes(&payload).build());
reader.read_tds_packet().await.unwrap();
let mut pause_state = reader
.try_receive_row_header(&int4_row_context(expected.len()))
.unwrap()
.expect("complete buffered row header");
let mut writer = DefaultRowWriter::new(expected.len());
assert!(
reader
.try_read_buffered_row_into(&mut pause_state, &mut writer)
.unwrap()
);
assert_eq!(
writer.take_row(),
expected
.into_iter()
.map(ColumnValues::Int)
.collect::<Vec<_>>()
);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 0);
}
#[tokio::test]
async fn buffered_row_writer_keeps_partial_column_for_async_continuation() {
let expected = [0x1234_5678_i32, -42_i32];
let second = expected[1].to_le_bytes();
let mut first_payload = vec![TokenType::Row as u8];
first_payload.extend_from_slice(&expected[0].to_le_bytes());
first_payload.extend_from_slice(&second[..2]);
let mut first = TestPacketBuilder::new(PacketType::TabularResult);
let mut second_packet = TestPacketBuilder::new(PacketType::TabularResult);
let mut stream = first.continuation().append_bytes(&first_payload).build();
stream.extend_from_slice(&second_packet.append_bytes(&second[2..]).build());
let mut reader = create_network_transport_with_data(&stream);
reader.read_tds_packet().await.unwrap();
let mut pause_state = reader
.try_receive_row_header(&int4_row_context(expected.len()))
.unwrap()
.expect("complete buffered row header");
let mut writer = DefaultRowWriter::new(expected.len());
assert!(
!reader
.try_read_buffered_row_into(&mut pause_state, &mut writer)
.unwrap()
);
assert_eq!(pause_state.next_column_index, 1);
assert_eq!(
reader.tds_read_buffer.get_remaining_byte_count(),
2,
"the partial second value must remain buffered"
);
let result = reader
.resume_row_into(
pause_state,
None,
None,
ColumnPolicy::DecodeAll,
&mut writer,
)
.await
.unwrap();
assert!(matches!(result, RowReadResult::RowWritten));
assert_eq!(
writer.take_row(),
expected
.into_iter()
.map(ColumnValues::Int)
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn buffered_cursor_miss_preserves_bytes_for_async_continuation() {
let expected = 0x1234_5678_i32;
let value = expected.to_le_bytes();
let mut first = TestPacketBuilder::new(PacketType::TabularResult);
let mut second = TestPacketBuilder::new(PacketType::TabularResult);
let mut first_payload = vec![TokenType::Row as u8];
first_payload.extend_from_slice(&value[..2]);
let mut stream = first.continuation().append_bytes(&first_payload).build();
stream.extend_from_slice(&second.append_bytes(&value[2..]).build());
let mut reader = create_network_transport_with_data(&stream);
reader.read_tds_packet().await.unwrap();
let pause_state = reader
.try_receive_row_header(&int4_row_context(1))
.unwrap()
.expect("row header is wholly buffered");
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 2);
assert_eq!(
reader.try_read_buffered_column(&pause_state, 0).unwrap(),
None
);
assert_eq!(
reader.tds_read_buffer.get_remaining_byte_count(),
2,
"a miss must not consume the partial scalar"
);
let mut writer = DefaultRowWriter::new(1);
let result = reader
.resume_row_into(
pause_state,
None,
None,
ColumnPolicy::DecodeOne(0),
&mut writer,
)
.await
.unwrap();
assert!(matches!(result, RowReadResult::RowWritten));
assert_eq!(writer.take_row(), vec![ColumnValues::Int(expected)]);
}
#[tokio::test]
async fn buffered_nbcrow_null_column_needs_no_payload_bytes() {
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let payload = [TokenType::NbcRow as u8, 0b0000_0001];
let mut reader = create_network_transport_with_data(&packet.append_bytes(&payload).build());
reader.read_tds_packet().await.unwrap();
let pause_state = reader
.try_receive_row_header(&int4_row_context(1))
.unwrap()
.expect("complete NBCROW header");
assert_eq!(
reader.try_read_buffered_column(&pause_state, 0).unwrap(),
Some(ColumnValues::Null)
);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 0);
}
#[tokio::test]
async fn buffered_cursor_rejects_invalid_context_and_preserves_non_rows() {
let mut empty_packet = TestPacketBuilder::new(PacketType::TabularResult);
let mut empty = create_network_transport_with_data(&empty_packet.build());
empty.read_tds_packet().await.unwrap();
assert!(
empty
.try_receive_row_header(&int4_row_context(1))
.unwrap()
.is_none()
);
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let payload = [TokenType::Done as u8];
let mut reader = create_network_transport_with_data(&packet.append_bytes(&payload).build());
reader.read_tds_packet().await.unwrap();
assert!(
reader
.try_receive_row_header(&ParserContext::None(()))
.is_err()
);
assert!(
reader
.try_receive_row_header(&int4_row_context(1))
.unwrap()
.is_none()
);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 1);
let pause_state = RowPauseState {
next_column_index: 1,
metadata: match int4_row_context(1) {
ParserContext::ColumnMetadata(metadata, _) => metadata,
_ => unreachable!(),
},
nbc_null_bitmap: None,
decryptor: None,
};
assert_eq!(
reader.try_read_buffered_column(&pause_state, 1).unwrap(),
None
);
}
#[tokio::test]
async fn buffered_row_writer_propagates_decoder_errors() {
let metadata = Arc::new(ColMetadataToken {
column_count: 1,
columns: vec![ColumnMetadata {
user_type: 0,
flags: 0,
type_info: TypeInfo::var_len(TdsDataType::IntN, 8).unwrap(),
data_type: TdsDataType::IntN,
column_name: "value".to_string(),
multi_part_name: None,
crypto_metadata: None,
}],
cek_table: Vec::new(),
});
let mut pause_state = RowPauseState {
next_column_index: 0,
metadata,
nbc_null_bitmap: None,
decryptor: None,
};
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let mut reader =
create_network_transport_with_data(&packet.append_bytes(&[3, 0, 0, 0]).build());
reader.read_tds_packet().await.unwrap();
let mut writer = DefaultRowWriter::new(1);
assert!(
reader
.try_read_buffered_row_into(&mut pause_state, &mut writer)
.is_err()
);
}
#[tokio::test]
async fn buffered_row_writer_writes_nbcrow_nulls() {
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let payload = [TokenType::NbcRow as u8, 0b0000_0001];
let mut reader = create_network_transport_with_data(&packet.append_bytes(&payload).build());
reader.read_tds_packet().await.unwrap();
let mut pause_state = reader
.try_receive_row_header(&int4_row_context(1))
.unwrap()
.unwrap();
let mut writer = DefaultRowWriter::new(1);
assert!(
reader
.try_read_buffered_row_into(&mut pause_state, &mut writer)
.unwrap()
);
assert_eq!(writer.take_row(), vec![ColumnValues::Null]);
}
#[tokio::test]
async fn buffered_variant_column_honors_nbcrow_null_bitmap() {
let metadata = Arc::new(ColMetadataToken {
column_count: 1,
columns: vec![ColumnMetadata {
user_type: 0,
flags: 0,
type_info: TypeInfo::var_len(TdsDataType::SsVariant, 8009).unwrap(),
data_type: TdsDataType::SsVariant,
column_name: "variant".to_string(),
multi_part_name: None,
crypto_metadata: None,
}],
cek_table: Vec::new(),
});
let pause_state = RowPauseState {
next_column_index: 0,
metadata,
nbc_null_bitmap: Some(Arc::from([1_u8])),
decryptor: None,
};
let mut reader = create_network_transport_with_data(&[]);
assert_eq!(
reader
.try_read_buffered_column_with_base(&pause_state, 0)
.unwrap(),
Some((ColumnValues::Null, None))
);
}
#[tokio::test]
async fn buffered_nbcrow_reuses_unaliased_bitmap_allocation() {
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let payload = [
TokenType::NbcRow as u8,
0b0000_0001,
0b0000_0010,
TokenType::NbcRow as u8,
0b0000_0100,
0b0000_1000,
];
let mut reader = create_network_transport_with_data(&packet.append_bytes(&payload).build());
reader.read_tds_packet().await.unwrap();
let context = int4_row_context(9);
let first = reader
.try_receive_row_header(&context)
.unwrap()
.expect("first NBCROW header");
let first_bitmap = first.nbc_null_bitmap.as_ref().expect("first bitmap");
assert_eq!(first_bitmap.as_ref(), &[0b0000_0001, 0b0000_0010]);
let first_allocation = first_bitmap.as_ptr();
drop(first);
let second = reader
.try_receive_row_header(&context)
.unwrap()
.expect("second NBCROW header");
let second_bitmap = second.nbc_null_bitmap.as_ref().expect("second bitmap");
assert_eq!(second_bitmap.as_ref(), &[0b0000_0100, 0b0000_1000]);
assert_eq!(
second_bitmap.as_ptr(),
first_allocation,
"the uniquely owned scratch bitmap should be refilled in place"
);
}
#[tokio::test]
async fn buffered_nbcrow_bitmap_miss_preserves_header_for_async_continuation() {
let mut first = TestPacketBuilder::new(PacketType::TabularResult);
let mut second = TestPacketBuilder::new(PacketType::TabularResult);
let mut stream = first
.continuation()
.append_bytes(&[TokenType::NbcRow as u8, 0])
.build();
stream.extend_from_slice(&second.append_bytes(&[0]).build());
let mut reader = create_network_transport_with_data(&stream);
reader.read_tds_packet().await.unwrap();
let context = int4_row_context(9);
assert!(reader.try_receive_row_header(&context).unwrap().is_none());
assert_eq!(
reader.tds_read_buffer.get_remaining_byte_count(),
2,
"the token and partial bitmap must remain buffered"
);
let header = reader
.receive_row_header(&context, None, None)
.await
.unwrap();
let RowHeader::Positioned(pause_state) = header else {
panic!("expected an NBCROW position");
};
assert_eq!(
pause_state
.nbc_null_bitmap
.as_ref()
.expect("NBCROW bitmap")
.as_ref(),
&[0, 0]
);
}
#[tokio::test]
async fn test_sync_scalar_probe_fallback_across_packet_boundaries() {
let expected_uint16 = 0x1234u16;
let expected_int16 = -0x1234i16;
let expected_uint24 = 0x00A1_B2C3u32;
let expected_int32 = -0x0123_4567i32;
let expected_uint32 = 0x89AB_CDEFu32;
let expected_uint40 = 0xAB_CDEF_0123u64;
let expected_int64 = -0x0102_0304_0506_0708i64;
let expected_float32 = 1.5f32;
let expected_float64 = -2.25f64;
let uint16 = expected_uint16.to_le_bytes();
let int16 = expected_int16.to_le_bytes();
let uint24 = expected_uint24.to_le_bytes();
let int32 = expected_int32.to_le_bytes();
let uint32 = expected_uint32.to_le_bytes();
let uint40 = expected_uint40.to_le_bytes();
let int64 = expected_int64.to_le_bytes();
let float32 = expected_float32.to_le_bytes();
let float64 = expected_float64.to_le_bytes();
let payloads = [
vec![0xAB, uint16[0]],
vec![uint16[1], int16[0]],
vec![int16[1], uint24[0], uint24[1]],
vec![uint24[2], int32[0], int32[1], int32[2]],
vec![int32[3], uint32[0], uint32[1], uint32[2]],
vec![uint32[3], uint40[0], uint40[1], uint40[2], uint40[3]],
vec![
uint40[4], int64[0], int64[1], int64[2], int64[3], int64[4], int64[5], int64[6],
],
vec![int64[7], float32[0], float32[1], float32[2]],
vec![
float32[3], float64[0], float64[1], float64[2], float64[3], float64[4], float64[5],
float64[6],
],
vec![float64[7]],
];
let mut stream = Vec::new();
let last_index = payloads.len() - 1;
for (index, payload) in payloads.iter().enumerate() {
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
if index != last_index {
packet.continuation();
}
stream.extend_from_slice(&packet.append_bytes(payload).build());
}
let mut reader = create_network_transport_with_data(&stream);
assert_eq!(reader.try_read_byte(), None);
assert_eq!(reader.read_byte().await.unwrap(), 0xAB);
assert_eq!(reader.try_read_uint16(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 1);
assert_eq!(reader.read_uint16().await.unwrap(), expected_uint16);
assert_eq!(reader.try_read_int16(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 1);
assert_eq!(reader.read_int16().await.unwrap(), expected_int16);
assert_eq!(reader.try_read_uint24(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 2);
assert_eq!(reader.read_uint24().await.unwrap(), expected_uint24);
assert_eq!(reader.try_read_int32(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 3);
assert_eq!(reader.read_int32().await.unwrap(), expected_int32);
assert_eq!(reader.try_read_uint32(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 3);
assert_eq!(reader.read_uint32().await.unwrap(), expected_uint32);
assert_eq!(reader.try_read_uint40(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 4);
assert_eq!(reader.read_uint40().await.unwrap(), expected_uint40);
assert_eq!(reader.try_read_int64(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 7);
assert_eq!(reader.read_int64().await.unwrap(), expected_int64);
assert_eq!(reader.try_read_float32(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 3);
assert_eq!(reader.read_float32().await.unwrap(), expected_float32);
assert_eq!(reader.try_read_float64(), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 7);
assert_eq!(reader.read_float64().await.unwrap(), expected_float64);
}
#[tokio::test]
async fn test_slice_probe_hits_within_a_packet() {
let payload: Vec<u8> = (0..32u8).collect();
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
let stream = packet.append_bytes(&payload).build();
let mut reader = create_network_transport_with_data(&stream);
assert_eq!(reader.try_read_slice(1), None);
assert_eq!(reader.read_byte().await.unwrap(), 0);
assert_eq!(reader.try_read_slice(31), Some(&payload[1..]));
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 0);
}
#[tokio::test]
async fn test_slice_probe_falls_back_across_a_packet_boundary() {
let mut first = TestPacketBuilder::new(PacketType::TabularResult);
let mut second = TestPacketBuilder::new(PacketType::TabularResult);
let mut stream = first.append_bytes(&[0, 1, 2, 3, 4]).continuation().build();
stream.extend_from_slice(&second.append_bytes(&[5, 6, 7, 8, 9]).build());
let mut reader = create_network_transport_with_data(&stream);
assert_eq!(reader.read_byte().await.unwrap(), 0);
assert_eq!(reader.try_read_slice(9), None);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 4);
let mut owned = vec![0u8; 9];
reader.read_bytes(&mut owned).await.unwrap();
assert_eq!(owned, vec![1, 2, 3, 4, 5, 6, 7, 8, 9]);
}
#[tokio::test]
async fn test_payload_free_non_eom_packet_is_rejected() {
let mut packet = TestPacketBuilder::new(PacketType::TabularResult).build();
assert_eq!(packet.len(), PacketWriter::PACKET_HEADER_SIZE);
packet[1] = 0x00;
let mut reader = create_network_transport_with_data(&packet);
assert!(matches!(
reader.read_byte().await,
Err(crate::error::Error::ProtocolError(_))
));
}
#[tokio::test]
async fn test_payload_free_eom_packet_is_accepted() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult).build();
let mut next = TestPacketBuilder::new(PacketType::TabularResult);
stream.extend_from_slice(&next.append_byte(0x7F).build());
let mut reader = create_network_transport_with_data(&stream);
assert_eq!(reader.read_byte().await.unwrap(), 0x7F);
let mut reader = create_network_transport_with_data(&stream);
let mut one = [0u8; 1];
assert_eq!(reader.read_bytes(&mut one).await.unwrap(), 1);
assert_eq!(one[0], 0x7F, "read_bytes must agree with read_byte");
let mut reader = create_network_transport_with_data(&stream);
reader
.skip_bytes(1)
.await
.expect("skip_bytes must agree with read_byte");
}
#[tokio::test]
async fn read_bytes_past_end_of_message_errors_instead_of_hanging() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0x11)
.build();
stream.extend_from_slice(&TestPacketBuilder::new(PacketType::TabularResult).build());
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [0u8; 2];
let result = timeout(Duration::from_secs(5), reader.read_bytes(&mut destination))
.await
.expect("read_bytes hung: the refill loop consumed an empty EOM packet forever");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn read_bytes_uninit_past_end_of_message_errors_instead_of_hanging() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0x22)
.build();
stream.extend_from_slice(&TestPacketBuilder::new(PacketType::TabularResult).build());
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [std::mem::MaybeUninit::<u8>::uninit(); 2];
let result = timeout(
Duration::from_secs(5),
reader.read_bytes_uninit(&mut destination),
)
.await
.expect("read_bytes_uninit hung: the refill loop consumed an empty EOM packet forever");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn read_uint32_spanning_end_of_message_errors_instead_of_hanging() {
let stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0x01)
.append_byte(0x02)
.build();
let mut reader = create_network_transport_with_live_peer(&stream);
let result = timeout(Duration::from_secs(5), reader.read_uint32())
.await
.expect("read_uint32 hung on a value truncated by the end of the message");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn skip_bytes_past_end_of_message_errors_instead_of_hanging() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0x33)
.build();
stream.extend_from_slice(&TestPacketBuilder::new(PacketType::TabularResult).build());
let mut reader = create_network_transport_with_live_peer(&stream);
let result = timeout(Duration::from_secs(5), reader.skip_bytes(2))
.await
.expect("skip_bytes hung: the refill loop consumed an empty EOM packet forever");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn bulk_reads_still_span_packets_and_tolerate_a_trailing_empty_eom() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult)
.continuation()
.append_bytes(&[0xA1, 0xA2])
.build();
stream.extend_from_slice(
&TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&[0xA3, 0xA4])
.build(),
);
stream.extend_from_slice(&TestPacketBuilder::new(PacketType::TabularResult).build());
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [0u8; 4];
let read = timeout(Duration::from_secs(5), reader.read_bytes(&mut destination))
.await
.expect("a value spanning two packets must not hang")
.expect("a value spanning two packets must read successfully");
assert_eq!(read, 4);
assert_eq!(destination, [0xA1, 0xA2, 0xA3, 0xA4]);
}
#[tokio::test]
async fn read_bytes_errors_on_consecutive_empty_messages_instead_of_spinning() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult).build();
stream.extend_from_slice(&TestPacketBuilder::new(PacketType::TabularResult).build());
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [0u8; 2];
let result = timeout(Duration::from_secs(5), reader.read_bytes(&mut destination))
.await
.expect("read_bytes spun: consecutive empty messages advanced nothing");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn read_bytes_uninit_errors_on_consecutive_empty_messages_instead_of_spinning() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult).build();
stream.extend_from_slice(&TestPacketBuilder::new(PacketType::TabularResult).build());
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [std::mem::MaybeUninit::<u8>::uninit(); 2];
let result = timeout(
Duration::from_secs(5),
reader.read_bytes_uninit(&mut destination),
)
.await
.expect("read_bytes_uninit spun: consecutive empty messages advanced nothing");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
fn addr(ip: &str, port: u16) -> SocketAddr {
SocketAddr::new(ip.parse().unwrap(), port)
}
#[test]
fn sort_by_ip_preference_platform_default_leaves_order_untouched() {
let mut addrs = vec![
addr("2001:db8::1", 1433),
addr("192.0.2.1", 1433),
addr("2001:db8::2", 1433),
];
let original = addrs.clone();
sort_by_ip_preference(&mut addrs, IPAddressPreference::UsePlatformDefault);
assert_eq!(addrs, original);
}
#[test]
fn sort_by_ip_preference_ipv4_first_orders_v4_before_v6() {
let mut addrs = vec![
addr("2001:db8::1", 1433),
addr("192.0.2.1", 1433),
addr("2001:db8::2", 1433),
addr("192.0.2.2", 1433),
];
sort_by_ip_preference(&mut addrs, IPAddressPreference::IPv4First);
assert_eq!(
addrs,
vec![
addr("192.0.2.1", 1433),
addr("192.0.2.2", 1433),
addr("2001:db8::1", 1433),
addr("2001:db8::2", 1433),
],
"IPv4 addresses must sort before IPv6, preserving relative order within each family"
);
}
#[test]
fn sort_by_ip_preference_ipv6_first_orders_v6_before_v4() {
let mut addrs = vec![
addr("192.0.2.1", 1433),
addr("2001:db8::1", 1433),
addr("192.0.2.2", 1433),
addr("2001:db8::2", 1433),
];
sort_by_ip_preference(&mut addrs, IPAddressPreference::IPv6First);
assert_eq!(
addrs,
vec![
addr("2001:db8::1", 1433),
addr("2001:db8::2", 1433),
addr("192.0.2.1", 1433),
addr("192.0.2.2", 1433),
],
"IPv6 addresses must sort before IPv4, preserving relative order within each family"
);
}
#[test]
fn sort_by_ip_preference_handles_single_family_lists() {
let mut v4_only = vec![addr("192.0.2.1", 1433), addr("192.0.2.2", 1433)];
let expected = v4_only.clone();
sort_by_ip_preference(&mut v4_only, IPAddressPreference::IPv6First);
assert_eq!(v4_only, expected, "no IPv6 entries to reorder against");
}
#[tokio::test(flavor = "current_thread")]
async fn create_base_stream_sequential_resolution_yields_to_the_executor() {
use std::sync::atomic::{AtomicUsize, Ordering};
let heartbeats = Arc::new(AtomicUsize::new(0));
let heartbeats_task = heartbeats.clone();
let heartbeat = tokio::spawn(async move {
loop {
heartbeats_task.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(1)).await;
}
});
let _ = create_base_stream_sequential(
IPAddressPreference::UsePlatformDefault,
"localhost",
0,
30_000,
1_000,
200,
)
.await;
heartbeat.abort();
assert!(
heartbeats.load(Ordering::SeqCst) > 0,
"the heartbeat task never ran while resolving 'localhost' — \
resolution is blocking the executor instead of awaiting it"
);
}
#[tokio::test]
async fn skip_bytes_errors_on_consecutive_empty_messages_instead_of_spinning() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult).build();
stream.extend_from_slice(&TestPacketBuilder::new(PacketType::TabularResult).build());
let mut reader = create_network_transport_with_live_peer(&stream);
let result = timeout(Duration::from_secs(5), reader.skip_bytes(2))
.await
.expect("skip_bytes spun: consecutive empty messages advanced nothing");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn bulk_reads_tolerate_a_single_leading_empty_message() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult).build();
stream.extend_from_slice(
&TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&[0xB1, 0xB2])
.build(),
);
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [0u8; 2];
let read = timeout(Duration::from_secs(5), reader.read_bytes(&mut destination))
.await
.expect("a read starting at a message boundary must not hang")
.expect("a leading empty message must be consumed, not rejected");
assert_eq!(read, 2);
assert_eq!(destination, [0xB1, 0xB2]);
}
#[tokio::test]
async fn read_bytes_past_end_of_message_without_a_trailing_packet_errors_instead_of_hanging() {
let stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0x11)
.build();
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [0u8; 2];
let result = timeout(Duration::from_secs(5), reader.read_bytes(&mut destination))
.await
.expect("read_bytes hung: the reader re-entered the socket after end-of-message");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn read_bytes_uninit_past_end_of_message_without_a_trailing_packet_errors() {
let stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0x22)
.build();
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [std::mem::MaybeUninit::<u8>::uninit(); 2];
let result = timeout(
Duration::from_secs(5),
reader.read_bytes_uninit(&mut destination),
)
.await
.expect("read_bytes_uninit hung: the reader re-entered the socket after end-of-message");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn skip_bytes_past_end_of_message_without_a_trailing_packet_errors() {
let stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0x33)
.build();
let mut reader = create_network_transport_with_live_peer(&stream);
let result = timeout(Duration::from_secs(5), reader.skip_bytes(2))
.await
.expect("skip_bytes hung: the reader re-entered the socket after end-of-message");
assert!(
matches!(result, Err(crate::error::Error::ProtocolError(_))),
"expected a protocol error, got {result:?}"
);
}
#[tokio::test]
async fn a_value_ending_exactly_at_end_of_message_reads_successfully() {
let mut stream = TestPacketBuilder::new(PacketType::TabularResult)
.continuation()
.append_bytes(&[0xB1, 0xB2])
.build();
stream.extend_from_slice(
&TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&[0xB3])
.build(),
);
let mut reader = create_network_transport_with_live_peer(&stream);
let mut destination = [0u8; 3];
let read = timeout(Duration::from_secs(5), reader.read_bytes(&mut destination))
.await
.expect("a value ending at the message boundary must not hang")
.expect("a value ending at the message boundary must read successfully");
assert_eq!(read, 3);
assert_eq!(destination, [0xB1, 0xB2, 0xB3]);
}
#[tokio::test]
async fn reset_reader_discards_unread_bytes_instead_of_aborting() {
let stream = TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&[0xC1, 0xC2, 0xC3, 0xC4])
.build();
let mut reader = create_network_transport_with_live_peer(&stream);
let mut first = [0u8; 1];
reader
.read_bytes(&mut first)
.await
.expect("the first byte must read");
assert_eq!(first, [0xC1]);
assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 3);
TdsPacketReader::reset_reader(&mut reader);
assert_eq!(
reader.tds_read_buffer.get_remaining_byte_count(),
0,
"reset_reader must leave an empty buffer"
);
}
fn done_token_message_with_type(token_type: TokenType, status: u16) -> Vec<u8> {
TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(token_type as u8)
.append_u16(status)
.append_u16(0) .append_u64(0) .build()
}
fn done_token_message(status: u16) -> Vec<u8> {
done_token_message_with_type(TokenType::Done, status)
}
async fn transport_responding_after_attention(
first_packet: Vec<u8>,
completion_packet: Vec<u8>,
) -> (NetworkTransport, tokio::task::JoinHandle<()>) {
let acknowledgement = done_token_message(DoneStatus::ATTN.bits());
let (client_side, mut peer) = duplex(MAX_BUFFER_SIZE);
peer.write_all(&first_packet).await.unwrap();
let transport = build_duplex_transport(client_side);
let peer_task = tokio::spawn(async move {
let mut attention = [0_u8; PacketWriter::PACKET_HEADER_SIZE];
peer.read_exact(&mut attention).await.unwrap();
assert_eq!(attention[0], PacketType::Attention as u8);
peer.write_all(&completion_packet).await.unwrap();
peer.write_all(&acknowledgement).await.unwrap();
});
(transport, peer_task)
}
async fn cancel_row_read_after_first_poll(
first_packet: Vec<u8>,
completion_packet: Vec<u8>,
context: &ParserContext,
plan: ColumnPolicy,
) -> (TdsResult<RowReadResult>, NetworkTransport) {
let column_count = match context {
ParserContext::ColumnMetadata(metadata, _) => metadata.columns.len(),
_ => panic!("row cancellation requires column metadata"),
};
let (mut transport, peer_task) =
transport_responding_after_attention(first_packet, completion_packet).await;
let parent = CancelHandle::new();
let child = parent.child_handle();
let mut writer = DefaultRowWriter::new(column_count);
let result = {
let mut read = std::pin::pin!(transport.receive_row_into(
context,
None,
Some(&child),
plan,
&mut writer,
));
let first_poll = poll_fn(|cx| Poll::Ready(read.as_mut().poll(cx))).await;
assert!(
first_poll.is_pending(),
"the row read unexpectedly completed before cancellation"
);
parent.cancel();
timeout(Duration::from_secs(5), read)
.await
.expect("the interrupted row did not settle")
};
peer_task.await.unwrap();
(result, transport)
}
fn int4_colmetadata_bytes(name: &str) -> Vec<u8> {
let mut bytes = vec![TokenType::ColMetadata as u8];
bytes.extend_from_slice(&1_u16.to_le_bytes());
bytes.extend_from_slice(&0_u32.to_le_bytes());
bytes.extend_from_slice(&0_u16.to_le_bytes());
bytes.push(TdsDataType::Int4 as u8);
bytes.push(u8::try_from(name.chars().count()).unwrap());
bytes.extend_from_slice(&encode_utf16_le(name));
bytes
}
fn int4_row_message(token_type: u8, is_nbc: bool, value: i32) -> Vec<u8> {
let mut packet = TestPacketBuilder::new(PacketType::TabularResult);
packet.append_byte(token_type);
if is_nbc {
packet.append_byte(0);
}
packet.append_i32(value).build()
}
#[test]
fn ready_read_does_not_require_a_time_driver() {
let runtime = tokio::runtime::Builder::new_current_thread()
.build()
.expect("current-thread runtime without the time driver");
runtime.block_on(async {
let mut read = std::pin::pin!(async { Ok::<_, crate::error::Error>(7) });
let result =
await_read_or_interrupt(read.as_mut(), Some(Duration::from_secs(30)), None).await;
assert!(matches!(result, Ok(Ok(7))));
});
}
#[tokio::test(start_paused = true)]
async fn suspended_read_observes_an_exhausted_timeout_budget() {
let mut read = std::pin::pin!(std::future::pending::<TdsResult<()>>());
let result = await_read_or_interrupt(read.as_mut(), Some(Duration::ZERO), None).await;
assert!(matches!(result, Err(ReadInterruption::TimedOut(_))));
}
fn cancelled_handle() -> CancelHandle {
let parent = CancelHandle::new();
let child = parent.child_handle();
parent.cancel();
child
}
fn is_known_dead(transport: &NetworkTransport) -> bool {
use crate::connection::transport::tds_transport::TdsTransport;
TdsTransport::connection_known_dead(transport)
}
async fn time_one_cancelled_read(transport: &mut NetworkTransport) -> Duration {
let started = Instant::now();
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("cancellation hung waiting for a DONE_ATTN the server never sent");
assert!(
matches!(result, Err(OperationCancelledError(_))),
"the caller must still see its cancellation, got {result:?}"
);
started.elapsed()
}
#[tokio::test(start_paused = true)]
async fn cancellation_does_not_wait_forever_for_an_attention_acknowledgement() {
let (mut transport, mut written) =
create_network_transport_with_live_peer_capturing_writes(&[]);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("cancellation hung waiting for a DONE_ATTN the server never sent");
assert!(
matches!(result, Err(OperationCancelledError(_))),
"the caller must still see its cancellation, got {result:?}"
);
let sent = timeout(Duration::from_secs(600), written.recv())
.await
.expect("cancellation returned without ever putting ATTENTION on the wire")
.expect("the peer must observe a write");
assert_eq!(
sent[0],
PacketType::Attention as u8,
"cancellation must put an ATTENTION packet on the wire"
);
}
#[tokio::test(start_paused = true)]
async fn a_stalled_attention_write_does_not_park_cancellation() {
const ATTENTION_PACKET_LEN: usize = 8;
let (client_side, mut peer) = duplex(1);
let mut transport = build_duplex_transport(client_side);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("cancellation hung writing ATTENTION to a peer that had stopped reading");
assert!(
matches!(result, Err(OperationCancelledError(_))),
"the caller must still see its cancellation, got {result:?}"
);
assert!(
is_known_dead(&transport),
"a half-written attention leaves the stream unresynchronisable"
);
let mut buffer = [0u8; ATTENTION_PACKET_LEN];
let delivered = timeout(Duration::from_secs(600), peer.read(&mut buffer))
.await
.expect("the write left nothing with the peer, so nothing stalled")
.expect("reading the peer end must not fail");
assert!(
delivered < ATTENTION_PACKET_LEN,
"the write must have stalled part-way; a completed one would have \
delivered all {ATTENTION_PACKET_LEN} bytes, so this test would be \
measuring the drain instead"
);
}
#[tokio::test(start_paused = true)]
async fn a_second_cancellation_does_not_spend_the_bound_again() {
let bound = Duration::from_secs(ATTENTION_TIMEOUT_SECONDS);
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&[]);
let first = time_one_cancelled_read(&mut transport).await;
assert!(
first >= bound,
"the first cancellation waits out the whole bound for an \
acknowledgement that never comes, took {first:?}"
);
assert!(
is_known_dead(&transport),
"an unacknowledged attention must leave the connection dead"
);
let second = time_one_cancelled_read(&mut transport).await;
assert_eq!(
second,
Duration::ZERO,
"a connection already given up on has nothing left to acknowledge, \
so a later cancellation must return without waiting at all"
);
}
#[tokio::test(start_paused = true)]
async fn request_timeout_does_not_wait_forever_for_an_attention_acknowledgement() {
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&[]);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(
&ParserContext::None(()),
Some(Duration::from_millis(50)),
None,
),
)
.await
.expect("the timeout path hung waiting for a DONE_ATTN the server never sent");
assert!(
matches!(result, Err(TimeoutError(_))),
"the caller must still see its timeout, got {result:?}"
);
}
#[tokio::test(start_paused = true)]
async fn an_unacknowledged_attention_marks_the_connection_dead() {
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&[]);
let _ = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("cancellation hung waiting for a DONE_ATTN the server never sent");
assert!(
is_known_dead(&transport),
"a connection whose attention went unacknowledged must not be reused"
);
}
#[tokio::test(start_paused = true)]
async fn an_acknowledged_attention_leaves_the_connection_usable() {
let (mut transport, _written) = create_network_transport_with_live_peer_capturing_writes(
&done_token_message(DoneStatus::ATTN.bits()),
);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("an acknowledged cancellation must not hang");
assert!(
matches!(result, Err(OperationCancelledError(_))),
"the caller must still see its cancellation, got {result:?}"
);
assert!(
!is_known_dead(&transport),
"an acknowledged attention leaves the connection reusable"
);
}
#[tokio::test(start_paused = true)]
async fn all_final_done_family_tokens_can_acknowledge_attention() {
for (name, token_type) in [
("DONE", TokenType::Done),
("DONEPROC", TokenType::DoneProc),
("DONEINPROC", TokenType::DoneInProc),
] {
let response = done_token_message_with_type(token_type, DoneStatus::ATTN.bits());
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&response);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("a final DONE-family acknowledgement timed out");
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(
!is_known_dead(&transport),
"{name}_ATTN must leave the connection reusable"
);
}
}
#[tokio::test(start_paused = true)]
async fn done_more_cannot_acknowledge_attention() {
for (name, token_type) in [
("DONE", TokenType::Done),
("DONEPROC", TokenType::DoneProc),
("DONEINPROC", TokenType::DoneInProc),
] {
let mut response = done_token_message_with_type(
token_type,
(DoneStatus::ATTN | DoneStatus::MORE).bits(),
);
response.extend_from_slice(&done_token_message(DoneStatus::ATTN.bits()));
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&response);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("the drain stopped at a DONE_MORE token");
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(!is_known_dead(&transport));
let settlement = transport
.take_attention_settlement()
.expect("the final acknowledgement must produce settlement state");
assert_eq!(
settlement.tokens.len(),
2,
"{name}_ATTN_MORE terminated the drain early"
);
}
}
#[test]
fn attention_settlement_stops_retaining_tokens_at_its_limit() {
let mut settlement = AttentionSettlement::default();
for _ in 0..=MAX_ATTENTION_SETTLEMENT_TOKENS {
settlement.push(Tokens::TabName);
}
assert_eq!(settlement.tokens.len(), MAX_ATTENTION_SETTLEMENT_TOKENS);
assert!(settlement.overflowed);
}
#[tokio::test(start_paused = true)]
async fn the_drain_discards_trailing_tokens_before_the_acknowledgement() {
let mut stream = done_token_message(DoneStatus::FINAL.bits());
stream.extend_from_slice(&done_token_message(DoneStatus::ATTN.bits()));
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&stream);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&ParserContext::None(()), None, Some(&cancelled_handle())),
)
.await
.expect("the drain hung instead of skipping past the trailing DONE");
assert!(
matches!(result, Err(OperationCancelledError(_))),
"the caller must still see its cancellation, got {result:?}"
);
assert!(
!is_known_dead(&transport),
"the acknowledgement arrived, so the connection stays usable"
);
}
#[tokio::test(start_paused = true)]
async fn queued_rows_are_drained_with_current_metadata() {
for (name, token_type, is_nbc) in [
("ROW", TokenType::Row as u8, false),
("NBCROW", TokenType::NbcRow as u8, true),
] {
let mut stream = int4_row_message(token_type, is_nbc, 42);
stream.extend_from_slice(&int4_row_message(token_type, is_nbc, 43));
stream.extend_from_slice(&done_token_message(DoneStatus::ATTN.bits()));
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&stream);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(&int4_row_context(1), None, Some(&cancelled_handle())),
)
.await
.expect("the attention drain hung on a queued row");
assert!(
matches!(result, Err(OperationCancelledError(_))),
"the caller must see its cancellation, got {result:?}"
);
assert!(
!is_known_dead(&transport),
"{name} was drained through DONE_ATTN, so the connection is reusable"
);
}
}
#[tokio::test(start_paused = true)]
async fn attention_drain_adopts_new_colmetadata() {
let mut response = TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&int4_colmetadata_bytes("value"))
.append_byte(TokenType::Row as u8)
.append_i32(42)
.build();
response.extend_from_slice(&done_token_message(DoneStatus::ATTN.bits()));
let (mut transport, _written) =
create_network_transport_with_live_peer_capturing_writes(&response);
let result = timeout(
Duration::from_secs(600),
transport.receive_token(
&ParserContext::ColumnEncryption(false),
None,
Some(&cancelled_handle()),
),
)
.await
.expect("the attention drain hung after new column metadata");
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(
!is_known_dead(&transport),
"the new metadata and its row were drained through DONE_ATTN"
);
}
#[tokio::test]
async fn cancellation_preserves_a_parser_paused_mid_row() {
let first_packet = TestPacketBuilder::new(PacketType::TabularResult)
.continuation()
.append_byte(TokenType::Row as u8)
.append_bytes(&42_i32.to_le_bytes()[..2])
.build();
let second_packet = TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&42_i32.to_le_bytes()[2..])
.build();
let context = int4_row_context(1);
let (result, transport) = cancel_row_read_after_first_poll(
first_packet,
second_packet,
&context,
ColumnPolicy::DecodeAll,
)
.await;
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(
!is_known_dead(&transport),
"finishing the in-flight row reached DONE_ATTN and preserved the connection"
);
}
#[tokio::test(start_paused = true)]
async fn attention_timeout_while_draining_after_a_partial_row_retires_the_connection() {
use std::sync::atomic::{AtomicBool, Ordering};
let first_packet = TestPacketBuilder::new(PacketType::TabularResult)
.continuation()
.append_byte(TokenType::Row as u8)
.append_bytes(&42_i32.to_le_bytes()[..2])
.build();
let completion_packet = TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&42_i32.to_le_bytes()[2..])
.build();
let (client_side, mut peer) = duplex(MAX_BUFFER_SIZE);
peer.write_all(&first_packet).await.unwrap();
let mut transport = build_duplex_transport(client_side);
let row_completed = Arc::new(AtomicBool::new(false));
let peer_row_completed = Arc::clone(&row_completed);
let peer_task = tokio::spawn(async move {
let mut attention = [0_u8; PacketWriter::PACKET_HEADER_SIZE];
peer.read_exact(&mut attention).await.unwrap();
assert_eq!(attention[0], PacketType::Attention as u8);
peer.write_all(&completion_packet).await.unwrap();
peer_row_completed.store(true, Ordering::Release);
std::future::pending::<()>().await;
});
let parent = CancelHandle::new();
let child = parent.child_handle();
let context = int4_row_context(1);
let mut writer = DefaultRowWriter::new(1);
let result = {
let mut read = std::pin::pin!(transport.receive_row_into(
&context,
None,
Some(&child),
ColumnPolicy::DecodeAll,
&mut writer,
));
let first_poll = poll_fn(|cx| Poll::Ready(read.as_mut().poll(cx))).await;
assert!(first_poll.is_pending());
parent.cancel();
timeout(Duration::from_secs(600), read)
.await
.expect("the shared ATTENTION deadline did not bound the drain")
};
peer_task.abort();
assert!(
row_completed.load(Ordering::Acquire),
"the peer must finish the interrupted row before withholding DONE_ATTN"
);
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(
is_known_dead(&transport),
"a mid-drain timeout must retire the unsynchronized connection"
);
}
#[tokio::test]
async fn cancellation_discards_columns_after_a_row_pause() {
let first_packet = TestPacketBuilder::new(PacketType::TabularResult)
.continuation()
.append_byte(TokenType::Row as u8)
.append_bytes(&42_i32.to_le_bytes()[..2])
.build();
let second_packet = TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&42_i32.to_le_bytes()[2..])
.append_i32(43)
.build();
let context = int4_row_context(2);
let (result, transport) = cancel_row_read_after_first_poll(
first_packet,
second_packet,
&context,
ColumnPolicy::DecodeOne(0),
)
.await;
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(
!is_known_dead(&transport),
"the remaining column was discarded before DONE_ATTN"
);
}
#[tokio::test]
async fn cancellation_discards_a_paused_plp_before_the_acknowledgement() {
let payload = b"cancelled PLP payload";
let total_length = u64::try_from(payload.len()).unwrap().to_le_bytes();
let first_packet = TestPacketBuilder::new(PacketType::TabularResult)
.continuation()
.append_byte(TokenType::Row as u8)
.append_bytes(&total_length[..4])
.build();
let second_packet = TestPacketBuilder::new(PacketType::TabularResult)
.append_bytes(&total_length[4..])
.append_u32(u32::try_from(payload.len()).unwrap())
.append_bytes(payload)
.append_u32(0)
.build();
let context = plp_varbinary_row_context();
let (result, transport) = cancel_row_read_after_first_poll(
first_packet,
second_packet,
&context,
ColumnPolicy::DecodeOne(0),
)
.await;
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(
!is_known_dead(&transport),
"the PLP payload was discarded before DONE_ATTN"
);
}
#[tokio::test]
async fn cancellation_completes_an_interrupted_nbcrow_header() {
let first_packet = TestPacketBuilder::new(PacketType::TabularResult)
.continuation()
.append_byte(TokenType::NbcRow as u8)
.build();
let completion_packet = TestPacketBuilder::new(PacketType::TabularResult)
.append_byte(0)
.append_i32(42)
.build();
let (mut transport, peer_task) =
transport_responding_after_attention(first_packet, completion_packet).await;
let context = int4_row_context(1);
let parent = CancelHandle::new();
let child = parent.child_handle();
let result = {
let mut read =
std::pin::pin!(transport.receive_row_header(&context, None, Some(&child),));
let first_poll = poll_fn(|cx| Poll::Ready(read.as_mut().poll(cx))).await;
assert!(
first_poll.is_pending(),
"the NBCROW header unexpectedly completed before cancellation"
);
parent.cancel();
timeout(Duration::from_secs(5), read)
.await
.expect("the interrupted row header did not settle")
};
peer_task.await.unwrap();
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(
!is_known_dead(&transport),
"finishing the header and row reached DONE_ATTN"
);
}
#[tokio::test]
async fn cancellation_discards_an_active_plp_before_the_acknowledgement() {
let payload = b"active PLP payload";
let metadata = plp_varbinary_metadata();
let (plp_stream, _) = PlpColumnStream::try_begin_buffered(
&metadata.columns[0],
&u64::try_from(payload.len()).unwrap().to_le_bytes(),
)
.unwrap()
.unwrap();
let mut plp_state = PlpPauseState {
row_pause_state: RowPauseState {
next_column_index: 1,
metadata,
nbc_null_bitmap: None,
decryptor: None,
},
plp_stream: plp_stream.unwrap(),
};
let completion_packet = TestPacketBuilder::new(PacketType::TabularResult)
.append_u32(u32::try_from(payload.len()).unwrap())
.append_bytes(payload)
.append_u32(0)
.build();
let (mut transport, peer_task) =
transport_responding_after_attention(Vec::new(), completion_packet).await;
let parent = CancelHandle::new();
let child = parent.child_handle();
let mut out = [0_u8; 1];
let result = {
let mut read = std::pin::pin!(transport.read_active_plp_bytes(
&mut plp_state,
None,
Some(&child),
&mut out,
));
let first_poll = poll_fn(|cx| Poll::Ready(read.as_mut().poll(cx))).await;
assert!(
first_poll.is_pending(),
"the PLP read unexpectedly completed before cancellation"
);
parent.cancel();
timeout(Duration::from_secs(5), read)
.await
.expect("the interrupted PLP read did not settle")
};
peer_task.await.unwrap();
assert!(matches!(result, Err(OperationCancelledError(_))));
assert!(plp_state.reached_end());
assert!(
!is_known_dead(&transport),
"draining the active PLP reached DONE_ATTN"
);
}
}