use std::collections::VecDeque;
use std::time::Duration;
use async_trait::async_trait;
use crate::connection::client_context::ClientContext;
use crate::connection::execution_context::ExecutionContext;
use crate::connection::tds_client::TdsClient;
use crate::connection::transport::any_transport::AnyTransport;
use crate::connection::transport::network_transport::TransportSslHandler;
use crate::connection::transport::tds_transport::TdsTransport;
use crate::core::{CancelHandle, NegotiatedEncryptionSetting, TdsResult};
use crate::datatypes::row_writer::RowWriter;
use crate::datatypes::sqldatatypes::{
PartialLengthType, TdsDataType, TypeInfo, TypeInfoVariant, UdtInfo, UdtInfoInColMetadata,
};
use crate::handler::handler_factory::create_test_negotiated_settings_internal;
use crate::io::reader_writer::{NetworkReader, NetworkWriter};
use crate::io::token_stream::{
ColumnPolicy, ParserContext, PlpPauseState, RowHeader, RowPauseState, RowReadResult,
TdsTokenStreamReader,
};
use crate::message::messages::ResetConnectionMode;
use crate::query::metadata::ColumnMetadata;
use crate::token::tokens::{
ColMetadataToken, CurrentCommand, DoneStatus, DoneToken, EnvChangeContainer, EnvChangeToken,
EnvChangeTokenSubType, ErrorToken, InfoToken, Tokens,
};
pub struct ScriptedToken(Tokens);
#[derive(Debug)]
struct TokenReplayTransport {
pending_tokens: VecDeque<Tokens>,
pending_rows: VecDeque<VecDeque<i32>>,
active_row: Option<VecDeque<i32>>,
buffered_prefix_columns: Option<usize>,
reset_mode: ResetConnectionMode,
reset_dispatched: bool,
known_dead: bool,
}
impl TokenReplayTransport {
fn new(tokens: Vec<Tokens>) -> Self {
Self {
pending_tokens: VecDeque::from(tokens),
pending_rows: VecDeque::new(),
active_row: None,
buffered_prefix_columns: None,
reset_mode: ResetConnectionMode::None,
reset_dispatched: false,
known_dead: false,
}
}
fn with_int_rows(
metadata: Tokens,
rows: Vec<Vec<i32>>,
buffered_prefix_columns: Option<usize>,
) -> Self {
let done = Tokens::Done(DoneToken {
status: DoneStatus::FINAL,
cur_cmd: CurrentCommand::Select,
row_count: 0,
});
let mut transport = Self::new(vec![metadata, done]);
transport.pending_rows = rows.into_iter().map(VecDeque::from).collect();
transport.buffered_prefix_columns = buffered_prefix_columns;
transport
}
fn position_int_row(&mut self, context: &ParserContext) -> TdsResult<Option<RowPauseState>> {
let Some(row) = self.pending_rows.pop_front() else {
return Ok(None);
};
let ParserContext::ColumnMetadata(metadata, decryptor) = context else {
return Err(crate::error::Error::ProtocolError(
"Expected column metadata while positioning a scripted row".to_string(),
));
};
self.active_row = Some(row);
Ok(Some(RowPauseState {
next_column_index: 0,
metadata: std::sync::Arc::clone(metadata),
nbc_null_bitmap: None,
decryptor: decryptor.clone(),
}))
}
}
#[async_trait]
impl TdsTokenStreamReader for TokenReplayTransport {
fn try_receive_row_header(
&mut self,
context: &ParserContext,
) -> TdsResult<Option<RowPauseState>> {
self.position_int_row(context)
}
fn try_read_buffered_column(
&mut self,
_pause_state: &RowPauseState,
_target: usize,
) -> TdsResult<Option<crate::datatypes::column_values::ColumnValues>> {
Ok(self
.active_row
.as_mut()
.and_then(VecDeque::pop_front)
.map(crate::datatypes::column_values::ColumnValues::Int))
}
fn try_read_buffered_test_row(
&mut self,
_pause_state: &mut RowPauseState,
) -> TdsResult<Option<(Vec<i32>, bool)>> {
let Some(row) = self.active_row.as_mut() else {
return Ok(None);
};
let take = self
.buffered_prefix_columns
.unwrap_or(row.len())
.min(row.len());
let prefix = row.drain(..take).collect();
let complete = row.is_empty();
if complete {
self.active_row = None;
}
Ok(Some((prefix, complete)))
}
async fn receive_token(
&mut self,
_context: &ParserContext,
_remaining_request_timeout: Option<Duration>,
_cancel_handle: Option<&CancelHandle>,
) -> TdsResult<Tokens> {
if let Some(tok) = self.pending_tokens.pop_front() {
return Ok(tok);
}
Err(crate::error::Error::ConnectionClosed("test".to_string()))
}
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> {
if let Some(tok) = self.pending_tokens.pop_front() {
return Ok(RowReadResult::Token(tok));
}
Err(crate::error::Error::ConnectionClosed("test".to_string()))
}
async fn receive_row_header(
&mut self,
context: &ParserContext,
_remaining_request_timeout: Option<Duration>,
_cancel_handle: Option<&CancelHandle>,
) -> TdsResult<RowHeader> {
if let Some(row) = self.position_int_row(context)? {
return Ok(RowHeader::Positioned(row));
}
if let Some(tok) = self.pending_tokens.pop_front() {
return Ok(RowHeader::Token(tok));
}
Err(crate::error::Error::ConnectionClosed("test".to_string()))
}
async fn resume_row_into(
&mut self,
mut pause_state: RowPauseState,
_remaining_request_timeout: Option<Duration>,
_cancel_handle: Option<&CancelHandle>,
_plan: ColumnPolicy,
writer: &mut (dyn RowWriter + Send),
) -> TdsResult<RowReadResult> {
if let Some(mut row) = self.active_row.take() {
while let Some(value) = row.pop_front() {
writer.write_i32(pause_state.next_column_index, value);
pause_state.next_column_index += 1;
}
return Ok(RowReadResult::RowWritten);
}
Err(crate::error::Error::ConnectionClosed("test".to_string()))
}
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> {
Err(crate::error::Error::ConnectionClosed("test".to_string()))
}
}
#[async_trait]
impl TransportSslHandler for TokenReplayTransport {
async fn enable_ssl(&mut self) -> TdsResult<()> {
Ok(())
}
async fn disable_ssl(&mut self) -> TdsResult<()> {
Ok(())
}
}
#[async_trait]
impl NetworkWriter for TokenReplayTransport {
async fn send(&mut self, _data: &[u8]) -> TdsResult<()> {
Ok(())
}
fn packet_size(&self) -> u32 {
4096
}
fn get_encryption_setting(&self) -> NegotiatedEncryptionSetting {
NegotiatedEncryptionSetting::NoEncryption
}
fn set_reset_mode(&mut self, mode: ResetConnectionMode) {
self.reset_mode = mode;
self.reset_dispatched = false;
}
fn take_reset_mode(&mut self) -> ResetConnectionMode {
std::mem::replace(&mut self.reset_mode, 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)
}
}
#[async_trait]
impl NetworkReader for TokenReplayTransport {
fn packet_size(&self) -> u32 {
4096
}
}
#[async_trait]
impl TdsTransport for TokenReplayTransport {
fn as_writer_ref(&self) -> &dyn NetworkWriter {
self
}
fn as_writer(&mut self) -> &mut dyn NetworkWriter {
self
}
fn reset_reader(&mut self) {}
fn packet_size(&self) -> u32 {
4096
}
async fn close_transport(&mut self) -> TdsResult<()> {
Ok(())
}
async fn send_attention_with_timeout(
&mut self,
_context: &ParserContext,
_timeout: Duration,
) -> TdsResult<bool> {
Ok(false)
}
fn is_connection_dead(&self) -> bool {
true
}
fn connection_known_dead(&self) -> bool {
self.known_dead
}
fn mark_known_dead(&mut self) {
self.known_dead = true;
}
}
pub fn tds_client_from_tokens(tokens: Vec<ScriptedToken>) -> TdsClient {
let tokens: Vec<Tokens> = tokens.into_iter().map(|t| t.0).collect();
let transport = AnyTransport::dynamic(TokenReplayTransport::new(tokens));
let negotiated_settings = create_test_negotiated_settings_internal();
let execution_context = ExecutionContext::new();
let client_context = ClientContext::with_data_source("tcp:localhost,1433");
TdsClient::new(
transport,
negotiated_settings,
execution_context,
client_context,
Vec::new(),
)
}
pub fn tds_client_from_int_rows(rows: Vec<Vec<i32>>) -> TdsClient {
let width = rows.first().map_or(0, Vec::len);
let metadata = Tokens::ColMetadata(ColMetadataToken {
column_count: u16::try_from(width).unwrap_or(u16::MAX),
columns: int_columns(width),
cek_table: Vec::new(),
});
let transport =
AnyTransport::dynamic(TokenReplayTransport::with_int_rows(metadata, rows, None));
let negotiated_settings = create_test_negotiated_settings_internal();
let execution_context = ExecutionContext::new();
let client_context = ClientContext::with_data_source("tcp:localhost,1433");
TdsClient::new(
transport,
negotiated_settings,
execution_context,
client_context,
Vec::new(),
)
}
pub fn tds_client_from_partial_int_rows(
rows: Vec<Vec<i32>>,
buffered_prefix_columns: usize,
) -> TdsClient {
let width = rows.first().map_or(0, Vec::len);
let metadata = Tokens::ColMetadata(ColMetadataToken {
column_count: u16::try_from(width).unwrap_or(u16::MAX),
columns: int_columns(width),
cek_table: Vec::new(),
});
let transport = AnyTransport::dynamic(TokenReplayTransport::with_int_rows(
metadata,
rows,
Some(buffered_prefix_columns),
));
let negotiated_settings = create_test_negotiated_settings_internal();
let execution_context = ExecutionContext::new();
let client_context = ClientContext::with_data_source("tcp:localhost,1433");
TdsClient::new(
transport,
negotiated_settings,
execution_context,
client_context,
Vec::new(),
)
}
pub fn tds_client_from_tokens_in_transaction(
tokens: Vec<ScriptedToken>,
descriptor: u64,
) -> TdsClient {
let tokens: Vec<Tokens> = tokens.into_iter().map(|t| t.0).collect();
let transport = AnyTransport::dynamic(TokenReplayTransport::new(tokens));
let negotiated_settings = create_test_negotiated_settings_internal();
let mut execution_context = ExecutionContext::new();
execution_context.set_transaction_descriptor(descriptor);
let client_context = ClientContext::with_data_source("tcp:localhost,1433");
TdsClient::new(
transport,
negotiated_settings,
execution_context,
client_context,
Vec::new(),
)
}
pub fn col_metadata_empty() -> ScriptedToken {
ScriptedToken(Tokens::ColMetadata(ColMetadataToken::default()))
}
pub fn col_metadata(columns: Vec<ColumnMetadata>) -> ScriptedToken {
ScriptedToken(Tokens::ColMetadata(ColMetadataToken {
column_count: u16::try_from(columns.len()).unwrap_or(u16::MAX),
columns,
cek_table: Vec::new(),
}))
}
pub fn int_columns(n: usize) -> Vec<ColumnMetadata> {
(1..=n)
.map(|i| ColumnMetadata {
user_type: 0,
flags: 0x01, type_info: TypeInfo::fixed_len(TdsDataType::Int4).expect("Int4 is a fixed-length type"),
data_type: TdsDataType::Int4,
column_name: format!("c{i}"),
multi_part_name: None,
crypto_metadata: None,
})
.collect()
}
pub fn udt_column(max_byte_size: u16) -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0x01,
type_info: TypeInfo::partial_len(TdsDataType::Udt, usize::from(max_byte_size), None)
.expect("UDT is a PLP type"),
data_type: TdsDataType::Udt,
column_name: "udt".to_string(),
multi_part_name: None,
crypto_metadata: None,
}
}
pub fn udt_column_with_metadata(
max_byte_size: u16,
db_name: &str,
schema_name: &str,
type_name: &str,
assembly_qualified_name: &str,
) -> ColumnMetadata {
let mut column = udt_column(max_byte_size);
column.type_info.type_info_variant = TypeInfoVariant::PartialLen(
PartialLengthType::Udt,
Some(usize::from(max_byte_size)),
None,
None,
Some(UdtInfo::InColMetadata(UdtInfoInColMetadata::new(
max_byte_size,
db_name.to_string(),
schema_name.to_string(),
type_name.to_string(),
assembly_qualified_name.to_string(),
))),
);
column
}
pub fn mixed_lob_columns(prefix_columns: usize) -> Vec<ColumnMetadata> {
let mut columns = int_columns(prefix_columns);
columns.push(ColumnMetadata {
user_type: 0,
flags: 0x01,
type_info: TypeInfo::partial_len(TdsDataType::NVarChar, usize::from(u16::MAX), None)
.expect("nvarchar(max) is a PLP type"),
data_type: TdsDataType::NVarChar,
column_name: "lob".to_string(),
multi_part_name: None,
crypto_metadata: None,
});
columns
}
pub fn tds_client_from_mixed_lob_prefix_rows(rows: Vec<Vec<i32>>) -> TdsClient {
let prefix_columns = rows.first().map_or(0, Vec::len);
let columns = mixed_lob_columns(prefix_columns);
let metadata = Tokens::ColMetadata(ColMetadataToken {
column_count: u16::try_from(columns.len()).unwrap_or(u16::MAX),
columns,
cek_table: Vec::new(),
});
let transport =
AnyTransport::dynamic(TokenReplayTransport::with_int_rows(metadata, rows, None));
let negotiated_settings = create_test_negotiated_settings_internal();
let execution_context = ExecutionContext::new();
let client_context = ClientContext::with_data_source("tcp:localhost,1433");
TdsClient::new(
transport,
negotiated_settings,
execution_context,
client_context,
Vec::new(),
)
}
pub fn done_more() -> ScriptedToken {
ScriptedToken(Tokens::Done(DoneToken {
status: DoneStatus::MORE,
cur_cmd: CurrentCommand::Insert,
row_count: 0,
}))
}
pub fn done_in_proc_more() -> ScriptedToken {
ScriptedToken(Tokens::DoneInProc(DoneToken {
status: DoneStatus::MORE,
cur_cmd: CurrentCommand::Select,
row_count: 0,
}))
}
pub fn done_proc_no_more() -> ScriptedToken {
ScriptedToken(Tokens::DoneProc(DoneToken {
status: DoneStatus::FINAL,
cur_cmd: CurrentCommand::Select,
row_count: 0,
}))
}
pub fn done_no_more() -> ScriptedToken {
ScriptedToken(Tokens::Done(DoneToken {
status: DoneStatus::FINAL,
cur_cmd: CurrentCommand::Insert,
row_count: 0,
}))
}
pub fn env_change_rollback_transaction() -> ScriptedToken {
ScriptedToken(Tokens::EnvChange(EnvChangeToken {
sub_type: EnvChangeTokenSubType::RollbackTransaction,
change_type: EnvChangeContainer::from((0u64, 0u64)),
}))
}
pub fn env_change_reset_connection() -> ScriptedToken {
ScriptedToken(Tokens::EnvChange(EnvChangeToken {
sub_type: EnvChangeTokenSubType::ResetConnection,
change_type: EnvChangeContainer::from((0u32, 0u32)),
}))
}
pub fn done_more_with_count(row_count: u64) -> ScriptedToken {
ScriptedToken(Tokens::Done(DoneToken {
status: DoneStatus::MORE | DoneStatus::COUNT,
cur_cmd: CurrentCommand::Insert,
row_count,
}))
}
pub fn done_more_select_with_count(row_count: u64) -> ScriptedToken {
ScriptedToken(Tokens::Done(DoneToken {
status: DoneStatus::MORE | DoneStatus::COUNT,
cur_cmd: CurrentCommand::Select,
row_count,
}))
}
pub fn info(number: u32, severity: u8, message: &str) -> ScriptedToken {
ScriptedToken(Tokens::Info(InfoToken {
number,
state: 1,
severity,
message: message.to_string(),
server_name: "test-server".to_string(),
proc_name: String::new(),
line_number: 1,
}))
}
pub fn sql_error(number: u32, severity: u8, message: &str) -> ScriptedToken {
ScriptedToken(Tokens::Error(ErrorToken {
number,
state: 1,
severity,
message: message.to_string(),
server_name: "test-server".to_string(),
proc_name: String::new(),
line_number: 1,
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn buffered_rows_require_column_metadata_context() {
let metadata = Tokens::ColMetadata(ColMetadataToken::default());
let mut transport = TokenReplayTransport::with_int_rows(metadata, vec![vec![1]], None);
assert!(
transport
.position_int_row(&ParserContext::None(()))
.is_err()
);
}
}
#[cfg(test)]
pub(crate) mod byte_stream {
use super::*;
use crate::token::parsers::common::test_utils::MockReader;
struct ByteStreamTransport {
reader: MockReader,
registry: crate::io::token_stream::GenericTokenParserRegistry,
nbc_bitmap_scratch: Option<std::sync::Arc<[u8]>>,
known_dead: bool,
}
impl ByteStreamTransport {
fn new(bytes: Vec<u8>) -> Self {
Self {
reader: MockReader::new(bytes),
registry: crate::io::token_stream::GenericTokenParserRegistry::default(),
nbc_bitmap_scratch: None,
known_dead: false,
}
}
}
impl std::fmt::Debug for ByteStreamTransport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ByteStreamTransport").finish()
}
}
#[async_trait]
impl TdsTokenStreamReader for ByteStreamTransport {
async fn receive_token(
&mut self,
context: &ParserContext,
_remaining_request_timeout: Option<Duration>,
_cancel_handle: Option<&CancelHandle>,
) -> TdsResult<Tokens> {
crate::io::token_stream::receive_token_internal(
&mut self.reader,
&self.registry,
context,
)
.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> {
crate::io::token_stream::receive_row_into_internal(
&mut self.reader,
&self.registry,
context,
plan,
writer,
&mut self.nbc_bitmap_scratch,
)
.await
}
async fn receive_row_header(
&mut self,
context: &ParserContext,
_remaining_request_timeout: Option<Duration>,
_cancel_handle: Option<&CancelHandle>,
) -> TdsResult<RowHeader> {
let token = crate::io::token_stream::receive_token_internal(
&mut self.reader,
&self.registry,
context,
)
.await?;
Ok(RowHeader::Token(token))
}
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> {
Err(crate::error::Error::ConnectionClosed("test".to_string()))
}
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> {
Err(crate::error::Error::ConnectionClosed("test".to_string()))
}
}
#[async_trait]
impl TransportSslHandler for ByteStreamTransport {
async fn enable_ssl(&mut self) -> TdsResult<()> {
Ok(())
}
async fn disable_ssl(&mut self) -> TdsResult<()> {
Ok(())
}
}
#[async_trait]
impl NetworkWriter for ByteStreamTransport {
async fn send(&mut self, _data: &[u8]) -> TdsResult<()> {
Ok(())
}
fn packet_size(&self) -> u32 {
4096
}
fn get_encryption_setting(&self) -> NegotiatedEncryptionSetting {
NegotiatedEncryptionSetting::NoEncryption
}
fn set_reset_mode(&mut self, _mode: ResetConnectionMode) {}
fn take_reset_mode(&mut self) -> ResetConnectionMode {
ResetConnectionMode::None
}
fn note_reset_dispatched(&mut self) {}
fn take_reset_dispatched(&mut self) -> bool {
false
}
}
#[async_trait]
impl NetworkReader for ByteStreamTransport {
fn packet_size(&self) -> u32 {
4096
}
}
#[async_trait]
impl TdsTransport for ByteStreamTransport {
fn as_writer_ref(&self) -> &dyn NetworkWriter {
self
}
fn as_writer(&mut self) -> &mut dyn NetworkWriter {
self
}
fn reset_reader(&mut self) {}
fn packet_size(&self) -> u32 {
4096
}
async fn close_transport(&mut self) -> TdsResult<()> {
Ok(())
}
async fn send_attention_with_timeout(
&mut self,
_context: &ParserContext,
_timeout: Duration,
) -> TdsResult<bool> {
Ok(false)
}
fn is_connection_dead(&self) -> bool {
false
}
fn connection_known_dead(&self) -> bool {
self.known_dead
}
fn mark_known_dead(&mut self) {
self.known_dead = true;
}
}
pub(crate) fn tds_client_over_raw_bytes(bytes: Vec<u8>) -> TdsClient {
let transport = AnyTransport::dynamic(ByteStreamTransport::new(bytes));
let negotiated_settings = create_test_negotiated_settings_internal();
let execution_context = ExecutionContext::new();
let client_context = ClientContext::with_data_source("tcp:localhost,1433");
TdsClient::new(
transport,
negotiated_settings,
execution_context,
client_context,
Vec::new(),
)
}
pub(crate) fn tds_client_over_raw_bytes_with_column_encryption(bytes: Vec<u8>) -> TdsClient {
use crate::message::features::always_encrypted::AlwaysEncryptedFeature;
use crate::message::login::Feature;
let transport = AnyTransport::dynamic(ByteStreamTransport::new(bytes));
let mut negotiated_settings = create_test_negotiated_settings_internal();
let mut feature = AlwaysEncryptedFeature::default();
feature.set_acknowledged(true);
negotiated_settings
.session_settings
.supported_features
.push(Box::new(feature));
let execution_context = ExecutionContext::new();
let client_context = ClientContext::with_data_source("tcp:localhost,1433");
TdsClient::new(
transport,
negotiated_settings,
execution_context,
client_context,
Vec::new(),
)
}
}