From bd07290ba999f7cc36d17d3dffbb746c60d18a2a Mon Sep 17 00:00:00 2001 From: Shiwani Gupta Date: Sun, 2 Aug 2026 12:47:52 +0530 Subject: [PATCH 1/4] Add msodbcsql-style incremental PLP parameter write to mssql-tds Extend the existing sp_executesql serialize flow with a data-at-execution pause point rather than adding a parallel send path. RpcParameter gains a data_at_exec marker; its serialize writes the parameter header and opens an unknown-length PLP value, then stops, reusing the same write_type_info the atomic path uses. begin_sp_executesql takes a single named_params list (some marked data_at_exec, mirroring ODBC SQL_DATA_AT_EXEC), partitions materialized vs streamed, sends materialized params through the normal path, and streams the rest via write_streamed_chunk/end_streamed_param. PacketWriter suspend/resume parks the in-progress message as owned client state, the write analogue of the incremental read pause. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql-tds/src/connection/tds_client.rs | 335 +++++++++++++++++- mssql-tds/src/io/packet_writer.rs | 137 +++++++ mssql-tds/src/message/messages.rs | 2 +- .../src/message/parameters/rpc_parameters.rs | 155 ++++++++ mssql-tds/src/message/rpc.rs | 27 +- mssql-tds/tests/test_client_write_apis.rs | 315 ++++++++++++++++ 6 files changed, 965 insertions(+), 6 deletions(-) create mode 100644 mssql-tds/tests/test_client_write_apis.rs diff --git a/mssql-tds/src/connection/tds_client.rs b/mssql-tds/src/connection/tds_client.rs index efd71882..d76bac4f 100644 --- a/mssql-tds/src/connection/tds_client.rs +++ b/mssql-tds/src/connection/tds_client.rs @@ -6,12 +6,14 @@ use crate::connection::bulk_copy_state::ATTENTION_TIMEOUT_SECONDS; use crate::connection::client_context::{ClientContext, ExecutionColumnEncryptionSetting}; use crate::connection::session_recovery::RecoveryContext; use crate::datatypes::bulk_copy_metadata::BulkCopyColumnMetadata; +use crate::datatypes::encoder::GenericEncoder; use crate::datatypes::row_writer::{DefaultRowWriter, DiscardRowWriter, RowWriter}; use crate::datatypes::sql_string::SqlString; use crate::datatypes::sqltypes::SqlType; +use crate::datatypes::tds_value_serializer::PLP_TERMINATOR; use crate::error::Error::UsageError; use crate::error::{SqlErrorInfo, SqlInfoMessage}; -use crate::io::packet_writer::PacketWriter; +use crate::io::packet_writer::{PacketWriter, SuspendedMessage, TdsPacketWriter}; use crate::message::bulk_load::{StreamingBulkLoadWriter, build_insert_bulk_command}; use crate::message::messages::{PacketType, ResetConnectionMode}; use crate::message::parameters::rpc_parameters::{ @@ -130,6 +132,47 @@ pub enum CursorColumn { RowEnded, } +/// Result of closing the current streamed PLP parameter via +/// [`TdsClient::end_streamed_param`]. +#[derive(Debug)] +pub enum StreamedParamStatus { + /// The server still expects another streamed parameter. Stream its value + /// chunks next via [`TdsClient::write_streamed_chunk`], then call + /// [`TdsClient::end_streamed_param`] again. Carries the name of the + /// parameter now open for data. + NeedData { + /// Name of the streamed parameter now awaiting its value chunks. + param_name: String, + }, + /// All streamed parameters have been written; the RPC message has been sent + /// and its response consumed. Any result set is now available exactly as + /// after a normal `execute_*` call. + Done, +} + +/// State machine for an in-progress incremental (streamed) PLP parameter write, +/// parked as owned state on the client between calls (mirrors the read side's +/// [`ActiveRowReadState`]). `Active` holds the suspended RPC message plus the +/// streamed parameters whose headers have not yet been emitted. +#[derive(Debug)] +enum StreamedWriteState { + Idle, + Active(Box), +} + +/// Owned context for a suspended, partially-written streamed RPC. +#[derive(Debug)] +struct StreamedWriteContext { + /// The parked outgoing RPC message. The parameter currently open for data + /// has already had its header + `PLP_UNKNOWN_LEN` written; the next bytes + /// appended are its value chunks. + message: SuspendedMessage, + /// Streamed parameters whose headers have not yet been written, in order. + pending: std::collections::VecDeque, + /// Collation used to write `TYPE_INFO` for subsequent streamed parameters. + db_collation: SqlCollation, +} + /// Active TDS connection to a SQL Server instance. /// /// Created by [`TdsConnectionProvider::create_client()`](crate::connection_provider::tds_connection_provider::TdsConnectionProvider::create_client). @@ -218,6 +261,10 @@ pub struct TdsClient { // Active PLP stream state when row decoding paused at a PLP target column. active_row_read_state: ActiveRowReadState, + + // Active incremental (streamed) PLP parameter-write state, if a streamed + // RPC is mid-flight. Idle for the normal atomic send path. + streamed_write_state: StreamedWriteState, } impl TdsClient { @@ -264,6 +311,7 @@ impl TdsClient { cancel_handle: None, empty_metadata: Vec::new(), active_row_read_state: ActiveRowReadState::Idle, + streamed_write_state: StreamedWriteState::Idle, } } @@ -829,6 +877,291 @@ impl TdsClient { self.position_on_first_result().await } + /// Executes a parameterized `sp_executesql` where one or more MAX-type + /// parameter values are supplied later, in chunks (data-at-execution). + /// + /// This is the streamed counterpart of + /// [`execute_sp_executesql`](Self::execute_sp_executesql): it takes the same + /// single `named_params` list, but any entry marked with + /// [`RpcParameter::data_at_exec`] has its value streamed afterwards instead + /// of materialized up front. Such a parameter's `SqlType` acts purely as a + /// type template (any inner value is ignored, so `SqlType::NVarcharMax(None)` + /// is the intended form) and it must be named and be a MAX type + /// (`nvarchar(max)`, `varchar(max)` or `varbinary(max)`). + /// + /// If no parameter is marked data-at-execution this simply delegates to the + /// atomic path and returns [`StreamedParamStatus::Done`]. Otherwise the RPC + /// header, positional `@statement`/`@params` arguments (whose declaration + /// covers every parameter) and all materialized parameters are written + /// eagerly through the normal serialize path; the header + + /// `PLP_UNKNOWN_LEN` of the first data-at-execution parameter is written via + /// the same [`RpcParameter::serialize`], which stops at the value; and the + /// in-progress message is parked as owned state on the client. Drive each + /// value with [`write_streamed_chunk`](Self::write_streamed_chunk) and close + /// it with [`end_streamed_param`](Self::end_streamed_param). + /// + /// Value chunk bytes are raw wire bytes and their lengths are byte counts: + /// UTF-16LE for `nvarchar(max)`, the server's single-byte encoding for + /// `varchar(max)`, and raw bytes for `varbinary(max)`. Encoding is the + /// caller's responsibility, mirroring the read side which yields raw PLP + /// bytes. + /// + /// Returns [`StreamedParamStatus::NeedData`] naming the first data-at-execution + /// parameter now awaiting its value. + /// + /// # Errors + /// Returns a usage error if a batch is already open, a streamed write is + /// already in progress, a data-at-execution parameter is a non-MAX type or + /// unnamed, or Always Encrypted is active (streamed values cannot be + /// encrypted). + pub async fn begin_sp_executesql( + &mut self, + sql: String, + named_params: Vec, + timeout_sec: Option, + cancel_handle: Option<&CancelHandle>, + ) -> TdsResult { + if self.execution_context.has_open_batch() + || !matches!(self.streamed_write_state, StreamedWriteState::Idle) + { + return Err(UsageError(ALREADY_EXECUTING_ERROR.to_string())); + } + + // Split the bound parameters into materialized values (sent now, via the + // existing serialize path) and data-at-execution values (streamed later), + // preserving order. Mirrors ODBC, where every parameter is bound and some + // are flagged SQL_DATA_AT_EXEC. + let (streamed_params, materialized_params): (Vec<_>, Vec<_>) = + named_params.into_iter().partition(RpcParameter::is_data_at_exec); + + // No data-at-execution parameters: this is just an ordinary execution. + if streamed_params.is_empty() { + self.execute_sp_executesql( + sql, + materialized_params, + ExecuteOptions { + timeout: timeout_sec, + cancel: cancel_handle, + column_encryption: ExecutionColumnEncryptionSetting::UseConnectionSetting, + }, + ) + .await?; + return Ok(StreamedParamStatus::Done); + } + + for param in &streamed_params { + if !param.is_streamable_plp() { + return Err(UsageError( + "Streamed parameters must be nvarchar(max), varchar(max) or varbinary(max)." + .to_string(), + )); + } + if param.name.is_none() { + return Err(UsageError( + "Streamed parameters must be named.".to_string(), + )); + } + } + + self.current_command_ce_setting = + ExecutionColumnEncryptionSetting::UseConnectionSetting; + + self.begin_command(); + let reconnect_elapsed = self.check_and_reconnect(timeout_sec, cancel_handle).await?; + let timeout_sec = Self::deduct_timeout(timeout_sec, reconnect_elapsed); + + self.remaining_request_timeout = Self::timeout_to_duration(timeout_sec); + self.cancel_handle = cancel_handle.map(|handle| handle.child_handle()); + + // Always Encrypted is not supported when streaming parameter values. + if self.should_encrypt_parameters() { + return Err(UsageError( + "Streamed PLP parameter writes are not supported with Always Encrypted.".to_string(), + )); + } + + self.transport.reset_reader(); + let database_collation = self.negotiated_settings.database_collation; + + let statement_parameter = RpcParameter::new( + None, + StatusFlags::NONE, + SqlType::NVarcharMax(Some(SqlString::from_utf8_string(sql))), + ); + + // The @params declaration must cover BOTH the materialized and the + // streamed parameters so the server knows every parameter's type. + let mut params_list_as_string = String::new(); + let mut all_declarations = materialized_params.clone(); + all_declarations.extend(streamed_params.iter().cloned()); + build_parameter_list_string(&all_declarations, &mut params_list_as_string)?; + + let params_parameter = RpcParameter::new( + None, + StatusFlags::NONE, + SqlType::NVarcharMax(Some(SqlString::from_utf8_string(params_list_as_string))), + ); + + let positional_parameters = Some(vec![statement_parameter, params_parameter]); + + let rpc = SqlRpc::new( + RpcType::ProcId(RpcProcs::ExecuteSql), + positional_parameters, + Some(materialized_params), + &database_collation, + &self.execution_context, + ); + + let mut pending: std::collections::VecDeque = + streamed_params.into_iter().collect(); + let first = pending + .pop_front() + .expect("streamed_params checked non-empty above"); + let first_name = first + .name + .clone() + .expect("streamed parameter names validated above"); + + let mut packet_writer = + rpc.create_packet_writer(self.transport.as_writer(), timeout_sec, cancel_handle); + // Write the RPC prefix (headers, proc, positional + materialized named + // params) then the first streamed parameter's header. The data-at-exec + // branch of `serialize` writes name + status + TYPE_INFO + PLP_UNKNOWN_LEN + // and stops at the value — the same method that serialized the + // materialized params, parked partway through. + rpc.serialize_prefix(&mut packet_writer).await?; + first + .serialize(&mut packet_writer, &database_collation, false, &GenericEncoder::new()) + .await?; + let message = packet_writer.suspend(); + + self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation: database_collation, + })); + + Ok(StreamedParamStatus::NeedData { + param_name: first_name, + }) + } + + /// Appends one value chunk to the streamed parameter currently open for + /// data. Call zero or more times between + /// [`begin_sp_executesql`](Self::begin_sp_executesql) (or + /// a [`StreamedParamStatus::NeedData`] from + /// [`end_streamed_param`](Self::end_streamed_param)) and the matching + /// [`end_streamed_param`](Self::end_streamed_param). + /// + /// Empty chunks are ignored: a zero-length PLP chunk header is the value + /// terminator, so it must never be emitted mid-value. + /// + /// # Errors + /// Returns a usage error if no streamed parameter is currently open, or if + /// `chunk` is longer than [`u32::MAX`] bytes (the PLP chunk-length field is + /// 32-bit). + pub async fn write_streamed_chunk(&mut self, chunk: &[u8]) -> TdsResult<()> { + if chunk.is_empty() { + return Ok(()); + } + if chunk.len() > u32::MAX as usize { + return Err(UsageError(format!( + "Streamed PLP chunk length {} exceeds the maximum chunk size of {} bytes.", + chunk.len(), + u32::MAX + ))); + } + + let ctx = match std::mem::replace(&mut self.streamed_write_state, StreamedWriteState::Idle) { + StreamedWriteState::Active(ctx) => ctx, + StreamedWriteState::Idle => { + return Err(UsageError( + "write_streamed_chunk called with no active streamed parameter.".to_string(), + )); + } + }; + let StreamedWriteContext { + message, + pending, + db_collation, + } = *ctx; + + let mut packet_writer = PacketWriter::resume(message, self.transport.as_writer()); + let result = async { + packet_writer.write_u32_async(chunk.len() as u32).await?; + packet_writer.write_async(chunk).await + } + .await; + let message = packet_writer.suspend(); + + // Re-park the message regardless of outcome so the state machine stays + // consistent; a write error fails the whole streamed operation. + self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation, + })); + result + } + + /// Closes the streamed parameter currently open for data by writing its PLP + /// terminator. + /// + /// If more streamed parameters remain, the next one's header is written and + /// [`StreamedParamStatus::NeedData`] is returned (stream its chunks next). + /// When the last streamed parameter closes, the RPC message is finalized and + /// sent, the response is advanced to the first column metadata, and + /// [`StreamedParamStatus::Done`] is returned — the result set is then + /// available exactly as after a normal `execute_*` call. + /// + /// # Errors + /// Returns a usage error if no streamed parameter is currently open, or an + /// I/O error if sending fails. + pub async fn end_streamed_param(&mut self) -> TdsResult { + let ctx = match std::mem::replace(&mut self.streamed_write_state, StreamedWriteState::Idle) { + StreamedWriteState::Active(ctx) => ctx, + StreamedWriteState::Idle => { + return Err(UsageError( + "end_streamed_param called with no active streamed parameter.".to_string(), + )); + } + }; + let StreamedWriteContext { + message, + mut pending, + db_collation, + } = *ctx; + + let mut packet_writer = PacketWriter::resume(message, self.transport.as_writer()); + packet_writer.write_u32_async(PLP_TERMINATOR).await?; + + if let Some(next) = pending.pop_front() { + let next_name = next + .name + .clone() + .expect("streamed parameter names validated at begin"); + next.serialize(&mut packet_writer, &db_collation, false, &GenericEncoder::new()) + .await?; + let message = packet_writer.suspend(); + self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation, + })); + return Ok(StreamedParamStatus::NeedData { + param_name: next_name, + }); + } + + // Last streamed parameter closed: send the message and consume the + // response exactly like execute_sp_executesql does. + packet_writer.finalize().await?; + drop(packet_writer); + + self.position_on_first_result().await?; + Ok(StreamedParamStatus::Done) + } + /// Executes a bulk load operation using zero-copy streaming. /// /// This method provides superior performance by eliminating per-row Vec allocations. diff --git a/mssql-tds/src/io/packet_writer.rs b/mssql-tds/src/io/packet_writer.rs index 693ae2da..3b5dbe41 100644 --- a/mssql-tds/src/io/packet_writer.rs +++ b/mssql-tds/src/io/packet_writer.rs @@ -111,6 +111,25 @@ pub struct PacketWriter<'a> { eom_pending: bool, } +/// Owned, detached state of an in-progress outgoing message, produced by +/// [`PacketWriter::suspend`] and consumed by [`PacketWriter::resume`]. Holds +/// every field of [`PacketWriter`] except the borrowed network writer, allowing +/// a partially-written message to be parked as owned state between calls. +#[derive(Debug)] +pub(crate) struct SuspendedMessage { + packet_type: PacketType, + max_payload_size: usize, + packet_id: u8, + payload_cursor: Cursor>, + packet_size: usize, + is_first_packet: bool, + start_time: Instant, + max_timeout_sec: Option, + cancel_handle: Option, + reset_mode: ResetConnectionMode, + eom_pending: bool, +} + impl<'a> PacketWriter<'a> { pub(crate) const PACKET_HEADER_SIZE: usize = 8; @@ -158,6 +177,59 @@ impl<'a> PacketWriter<'a> { } } + /// Detaches this writer's in-progress message state from the borrowed + /// network writer so it can be parked as owned state (e.g. on the TDS + /// client) across `await` points and multiple public calls, then later + /// reattached with [`resume`](Self::resume). + /// + /// This is the enabling primitive for incremental (streamed) PLP parameter + /// writes: the RPC header and any fully-materialized parameters are written + /// eagerly, then the message is suspended while the caller streams parameter + /// chunks one call at a time, resuming for each chunk and the final + /// terminator + `finalize`. + /// + /// The timeout budget (`start_time`) and packet accounting (`packet_id`, + /// `is_first_packet`, `eom_pending`, buffered payload) are preserved so the + /// resumed message behaves as one continuous send. + pub(crate) fn suspend(self) -> SuspendedMessage { + SuspendedMessage { + packet_type: self.packet_type, + max_payload_size: self.max_payload_size, + packet_id: self.packet_id, + payload_cursor: self.payload_cursor, + packet_size: self.packet_size, + is_first_packet: self.is_first_packet, + start_time: self.start_time, + max_timeout_sec: self.max_timeout_sec, + cancel_handle: self.cancel_handle, + reset_mode: self.reset_mode, + eom_pending: self.eom_pending, + } + } + + /// Reattaches a previously [`suspend`](Self::suspend)ed message to a network + /// writer, restoring all packet/timeout accounting so writing can continue + /// exactly where it left off. + pub(crate) fn resume( + state: SuspendedMessage, + network_writer: &'a mut dyn NetworkWriter, + ) -> PacketWriter<'a> { + PacketWriter { + packet_type: state.packet_type, + network_writer, + max_payload_size: state.max_payload_size, + packet_id: state.packet_id, + payload_cursor: state.payload_cursor, + packet_size: state.packet_size, + is_first_packet: state.is_first_packet, + start_time: state.start_time, + max_timeout_sec: state.max_timeout_sec, + cancel_handle: state.cancel_handle, + reset_mode: state.reset_mode, + eom_pending: state.eom_pending, + } + } + #[cfg(test)] pub(crate) async fn cancel_current_message(&mut self) -> TdsResult<()> { self.populate_header_and_send(true, true).await @@ -1301,4 +1373,69 @@ pub(crate) mod tests { let writer = PacketWriter::new(PacketType::TabularResult, &mut mock, None, None); assert!(writer.max_timeout_sec.is_none()); } + + /// Reassembles the TDS packet stream captured by the mock into the original + /// contiguous payload, stripping each 8-byte packet header. + fn reassemble_payload(sent: &[u8]) -> Vec { + let mut reconstructed: Vec = Vec::new(); + let mut offset = 0; + while offset < sent.len() { + let packet_len = u16::from_be_bytes([sent[offset + 2], sent[offset + 3]]) as usize; + reconstructed.extend_from_slice(&sent[offset + 8..offset + packet_len]); + offset += packet_len; + } + reconstructed + } + + /// A message written across a suspend/resume boundary produces the same + /// bytes as if written in one go: payload is preserved and the final packet + /// still terminates the message. + #[test] + fn suspend_resume_preserves_payload_within_single_packet() { + let packet_size = 4096u32; + let mut mock = MockNetworkWriter::new(packet_size); + + let mut writer = PacketWriter::new(PacketType::RpcRequest, &mut mock, None, None); + block_on(writer.write_async(&[0x01, 0x02, 0x03, 0x04])).unwrap(); + let suspended = writer.suspend(); + + // Nothing should have been sent yet (payload fits one packet, unflushed). + assert!(mock.data.is_empty()); + + let mut writer = PacketWriter::resume(suspended, &mut mock); + block_on(writer.write_async(&[0x05, 0x06, 0x07, 0x08])).unwrap(); + block_on(writer.finalize()).unwrap(); + + assert_eq!( + reassemble_payload(&mock.data), + vec![0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08] + ); + // Single packet, EOM status bit set on the last (only) packet. + assert_eq!(mock.data[0], PacketType::RpcRequest as u8); + assert_eq!(mock.data[1] & 0x01, 0x01); + } + + /// Suspending mid-message preserves packet accounting so a resumed write that + /// overflows the packet boundary frames continuous, correctly ordered + /// packets. + #[test] + fn suspend_resume_spans_packet_boundary() { + let packet_size = 16u32; // 8-byte payload per packet + let mut mock = MockNetworkWriter::new(packet_size); + + let first: Vec = (0..8).collect(); + let second: Vec = (8..16).collect(); + + let mut writer = PacketWriter::new(PacketType::RpcRequest, &mut mock, None, None); + block_on(writer.write_async(&first)).unwrap(); + let suspended = writer.suspend(); + + let mut writer = PacketWriter::resume(suspended, &mut mock); + block_on(writer.write_async(&second)).unwrap(); + block_on(writer.finalize()).unwrap(); + + // Two packets were needed; payload reassembles to the full 16 bytes. + assert!(mock.data.len() > packet_size as usize); + assert_eq!(reassemble_payload(&mock.data), (0..16).collect::>()); + } } diff --git a/mssql-tds/src/message/messages.rs b/mssql-tds/src/message/messages.rs index 212594da..a49f953e 100644 --- a/mssql-tds/src/message/messages.rs +++ b/mssql-tds/src/message/messages.rs @@ -5,7 +5,7 @@ use crate::core::{CancelHandle, NegotiatedEncryptionSetting, TdsResult}; use crate::io::{packet_writer::PacketWriter, reader_writer::NetworkWriter}; use async_trait::async_trait; -#[derive(Copy, Clone)] +#[derive(Copy, Clone, Debug)] #[allow(dead_code, clippy::upper_case_acronyms)] pub enum PacketType { Unknown = 0x00, diff --git a/mssql-tds/src/message/parameters/rpc_parameters.rs b/mssql-tds/src/message/parameters/rpc_parameters.rs index 7f8600e0..f287d22e 100644 --- a/mssql-tds/src/message/parameters/rpc_parameters.rs +++ b/mssql-tds/src/message/parameters/rpc_parameters.rs @@ -7,6 +7,7 @@ use crate::datatypes::column_values::DEFAULT_VARTIME_SCALE; use crate::datatypes::encoder::SqlValueEncoder; use crate::datatypes::sql_tvp::TvpTypeName; use crate::datatypes::sqltypes::SqlType; +use crate::datatypes::tds_value_serializer::PLP_UNKNOWN_LEN; use crate::{ core::TdsResult, datatypes::sqldatatypes::TdsDataType, @@ -113,6 +114,15 @@ pub struct RpcParameter { /// `SqlParameter.ForceColumnEncryption`; a client-side directive that is /// never sent on the wire. force_column_encryption: bool, + + /// When `true`, this parameter's value is supplied later, in chunks, via the + /// data-at-execution path (ODBC `SQL_DATA_AT_EXEC`). During serialization the + /// `value` is treated purely as a type template: [`serialize`](Self::serialize) + /// writes the parameter header and opens an unknown-length PLP value, then + /// stops. The value chunks and PLP terminator are written afterwards by the + /// streaming driver. Only the MAX (PLP) types are eligible. Never sent on the + /// wire as a flag. + data_at_exec: bool, } impl RpcParameter { @@ -124,9 +134,28 @@ impl RpcParameter { value, encrypted: None, force_column_encryption: false, + data_at_exec: false, } } + /// Marks this parameter as data-at-execution: its value is streamed later in + /// chunks rather than materialized up front. The `value` supplied to + /// [`new`](Self::new) is used only as a type template (a `None`-valued MAX + /// type such as `SqlType::NVarcharMax(None)` is the intended form). + /// + /// Only `nvarchar(max)`, `varchar(max)` and `varbinary(max)` may be streamed; + /// see [`is_streamable_plp`](Self::is_streamable_plp). + pub fn data_at_exec(mut self) -> Self { + self.data_at_exec = true; + self + } + + /// Returns `true` if this parameter's value is supplied via the + /// data-at-execution (streamed) path. + pub(crate) fn is_data_at_exec(&self) -> bool { + self.data_at_exec + } + /// Requires this parameter to be encrypted under Always Encrypted. /// /// When set, the driver fails with a usage error if the server reports the @@ -297,6 +326,33 @@ impl RpcParameter { } } + // Data-at-execution: the value is streamed later in chunks. Reuse the + // exact opening the atomic PLP path emits — status byte, TYPE_INFO, then + // the unknown-total-length sentinel — and stop. The value chunks and PLP + // terminator are written afterwards by the streaming driver. This is the + // write analogue of the incremental read's pause point: the same + // serialize method, parked partway through the value. + if self.data_at_exec { + if !self.is_streamable_plp() { + return Err(Error::UsageError(format!( + "Parameter '{}' is not a streamable PLP (MAX) type; only \ + nvarchar(max), varchar(max) and varbinary(max) may be streamed.", + self.name.as_deref().unwrap_or("") + ))); + } + if self.encrypted.is_some() { + return Err(Error::UsageError( + "Encrypted parameters cannot be streamed incrementally.".to_string(), + )); + } + packet_writer.write_byte_async(self.options.bits()).await?; + self.value + .write_type_info(packet_writer, db_collation, None, None) + .await?; + packet_writer.write_u64_async(PLP_UNKNOWN_LEN).await?; + return Ok(()); + } + // Encrypted parameters bypass the normal value encoder: the ciphertext // is sent as a BIGVARBINARY with the ENCRYPTED status flag and a // trailing CryptoMetaData block (Always Encrypted). @@ -315,6 +371,16 @@ impl RpcParameter { Ok(()) } + /// Returns `true` when this parameter's declared type is a MAX (PLP) type + /// eligible for incremental streaming: `nvarchar(max)`, `varchar(max)`, or + /// `varbinary(max)`. + pub(crate) fn is_streamable_plp(&self) -> bool { + matches!( + self.value, + SqlType::NVarcharMax(_) | SqlType::VarcharMax(_) | SqlType::VarBinaryMax(_) + ) + } + /// Marks this parameter as encrypted, supplying the ciphertext (or `None` /// for an encrypted NULL) and the cipher metadata. When set, [`serialize`] /// writes the value as a BIGVARBINARY with the ENCRYPTED status flag and a @@ -851,4 +917,93 @@ mod tests { ); assert_eq!(param.value(), &SqlType::Int(Some(42))); } + + /// Serializes a data-at-execution PLP parameter via the normal `serialize` + /// path (which, for a `data_at_exec` param, writes only the header and opens + /// the value), returning the payload bytes. + fn streamed_header_bytes(param: &RpcParameter, is_positional: bool) -> Vec { + let mut mock = MockNetworkWriter::new(16384); + let mut w = PacketWriter::new(PacketType::RpcRequest, &mut mock, None, None); + let collation = SqlCollation::default(); + let encoder = GenericEncoder {}; + block_on(param.serialize(&mut w, &collation, is_positional, &encoder)).unwrap(); + payload(&w) + } + + /// A named data-at-execution `nvarchar(max)` param serializes to: name + /// prefix, status-flags byte, the value's TYPE_INFO, then the 8-byte + /// `PLP_UNKNOWN_LEN` sentinel that opens the value. No value bytes or + /// terminator are written — those are streamed later. + #[test] + fn serialize_data_at_exec_named() { + let param = RpcParameter::new( + Some("@p".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(); + + let mut expected = vec![0x02, 0x40, 0x00, 0x70, 0x00]; // name: len 2, "@p" UTF-16LE + expected.push(StatusFlags::NONE.bits()); // status flags + expected.extend_from_slice(&type_info_bytes(&SqlType::NVarcharMax(None))); // TYPE_INFO + expected.extend_from_slice(&0xFFFF_FFFF_FFFF_FFFEu64.to_le_bytes()); // PLP_UNKNOWN_LEN + + assert_eq!(streamed_header_bytes(¶m, false), expected); + } + + /// A positional data-at-execution param writes a zero-length name byte in + /// place of the name, then the same status/TYPE_INFO/PLP_UNKNOWN_LEN sequence. + #[test] + fn serialize_data_at_exec_positional() { + let param = RpcParameter::new(None, StatusFlags::NONE, SqlType::VarBinaryMax(None)) + .data_at_exec(); + + let mut expected = vec![0x00]; // zero-length name (positional) + expected.push(StatusFlags::NONE.bits()); + expected.extend_from_slice(&type_info_bytes(&SqlType::VarBinaryMax(None))); + expected.extend_from_slice(&0xFFFF_FFFF_FFFF_FFFEu64.to_le_bytes()); + + assert_eq!(streamed_header_bytes(¶m, true), expected); + } + + /// Non-MAX types are rejected on the data-at-execution path: only + /// nvarchar(max)/varchar(max)/varbinary(max) may be streamed. + #[test] + fn serialize_data_at_exec_rejects_non_max_type() { + let param = RpcParameter::new( + Some("@p".to_string()), + StatusFlags::NONE, + SqlType::Int(Some(1)), + ) + .data_at_exec(); + assert!(!param.is_streamable_plp()); + + let mut mock = MockNetworkWriter::new(16384); + let mut w = PacketWriter::new(PacketType::RpcRequest, &mut mock, None, None); + let collation = SqlCollation::default(); + let encoder = GenericEncoder {}; + let err = block_on(param.serialize(&mut w, &collation, false, &encoder)) + .expect_err("non-max type must be rejected"); + assert!(matches!(err, Error::UsageError(_))); + } + + /// Encrypted parameters cannot be streamed incrementally. + #[test] + fn serialize_data_at_exec_rejects_encrypted() { + let mut param = RpcParameter::new( + Some("@p".to_string()), + StatusFlags::NONE, + SqlType::VarBinaryMax(None), + ) + .data_at_exec(); + param.set_encrypted(Some(vec![0x01, 0x02]), sample_metadata()); + + let mut mock = MockNetworkWriter::new(16384); + let mut w = PacketWriter::new(PacketType::RpcRequest, &mut mock, None, None); + let collation = SqlCollation::default(); + let encoder = GenericEncoder {}; + let err = block_on(param.serialize(&mut w, &collation, false, &encoder)) + .expect_err("encrypted parameter must be rejected"); + assert!(matches!(err, Error::UsageError(_))); + } } diff --git a/mssql-tds/src/message/rpc.rs b/mssql-tds/src/message/rpc.rs index 374c0cb8..175bf8ce 100644 --- a/mssql-tds/src/message/rpc.rs +++ b/mssql-tds/src/message/rpc.rs @@ -97,6 +97,28 @@ impl<'a> SqlRpc<'a> { Ok(()) } + /// Serializes the RPC up to and including all fully-materialized + /// parameters, but does **not** send the terminating `finalize` packet. + /// + /// Used by the incremental (streamed) PLP write path: the caller writes the + /// header, proc, positional and materialized named parameters here, then + /// appends one or more streamed parameters chunk-by-chunk before calling + /// `finalize` itself. For the normal atomic send, use [`serialize`], which + /// wraps this and then finalizes. + pub(crate) async fn serialize_prefix<'s, 'b>( + &'s self, + packet_writer: &'s mut PacketWriter<'b>, + ) -> TdsResult<()> + where + 'b: 's, + { + write_headers(&self.headers, packet_writer).await?; + self.write_proc(packet_writer).await?; + self.write_positional_parameters(packet_writer).await?; + self.write_named_parameters(packet_writer).await?; + Ok(()) + } + async fn write_proc(&self, packet_writer: &mut PacketWriter<'_>) -> TdsResult<()> { match &self.rpc_type { RpcType::Named(stored_proc_name) => { @@ -175,10 +197,7 @@ impl Request for SqlRpc<'_> { where 'b: 'a, { - write_headers(&self.headers, packet_writer).await?; - self.write_proc(packet_writer).await?; - self.write_positional_parameters(packet_writer).await?; - self.write_named_parameters(packet_writer).await?; + self.serialize_prefix(packet_writer).await?; packet_writer.finalize().await?; Ok(()) } diff --git a/mssql-tds/tests/test_client_write_apis.rs b/mssql-tds/tests/test_client_write_apis.rs new file mode 100644 index 00000000..00263043 --- /dev/null +++ b/mssql-tds/tests/test_client_write_apis.rs @@ -0,0 +1,315 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! End-to-end tests for the incremental (streamed) PLP parameter-write path +//! (`begin_sp_executesql` / `write_streamed_chunk` / `end_streamed_param`). +//! +//! These mirror the read-side PLP tests: they require a live SQL Server and are +//! driven by the `DB_HOST` / `DB_USERNAME` / `SQL_PASSWORD` environment +//! variables (see `common`), so they only run in CI. + +#[cfg(test)] +mod common; + +mod streamed_plp_write { + use crate::common::{build_tcp_datasource, create_context, init_tracing}; + use mssql_tds::connection::tds_client::{ResultSet, ResultSetClient, StreamedParamStatus}; + use mssql_tds::connection_provider::tds_connection_provider::TdsConnectionProvider; + use mssql_tds::datatypes::column_values::ColumnValues; + use mssql_tds::datatypes::sqltypes::SqlType; + use mssql_tds::message::parameters::rpc_parameters::{RpcParameter, StatusFlags}; + + /// Encodes a string to the UTF-16LE wire bytes an `nvarchar(max)` value uses. + fn utf16le(text: &str) -> Vec { + text.encode_utf16().flat_map(|u| u.to_le_bytes()).collect() + } + + /// Streams a large `nvarchar(max)` value into a temp table in multiple + /// chunks, then reads it back and verifies the round-trip. + #[tokio::test] + async fn stream_nvarchar_max_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_nvm (id INT, val NVARCHAR(MAX))".to_string(), + None, + None, + ) + .await?; + client.close_query().await?; + + let value = "Z".repeat(20_000); + let wire = utf16le(&value); + + let streamed = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(); + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_nvm (id, val) VALUES (1, @v)".to_string(), + vec![streamed], + None, + None, + ) + .await?; + match status { + StreamedParamStatus::NeedData { param_name } => assert_eq!(param_name, "@v"), + StreamedParamStatus::Done => panic!("expected NeedData for the first streamed param"), + } + + // Stream the value in two chunks split on an even (code-unit) boundary. + let split = (wire.len() / 2) & !1; + client.write_streamed_chunk(&wire[..split]).await?; + client.write_streamed_chunk(&wire[split..]).await?; + + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_nvm WHERE id = 1".to_string(), None, None) + .await?; + if let Some(resultset) = client.get_current_resultset() { + let row = resultset.next_row().await?.expect("expected a row"); + match &row[0] { + ColumnValues::String(s) => { + let round_tripped = s.to_utf8_string(); + assert_eq!(round_tripped.len(), value.len()); + assert_eq!(round_tripped, value); + } + other => panic!("Expected String for nvarchar(max), got {other:?}"), + } + } else { + panic!("expected a result set"); + } + client.close_query().await?; + Ok(()) + } + + /// Mixes a fully-materialized parameter with a data-at-execution one in the + /// same `begin_sp_executesql` call: the materialized `@id` is sent up front + /// via the normal serialize path, while `@v` is streamed. Verifies the + /// integrated single-`named_params` list (not a separate streamed argument). + #[tokio::test] + async fn stream_mixed_materialized_and_data_at_exec() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_mix (id INT, val NVARCHAR(MAX))".to_string(), + None, + None, + ) + .await?; + client.close_query().await?; + + let value = "M".repeat(15_000); + let wire = utf16le(&value); + + let params = vec![ + RpcParameter::new( + Some("@id".to_string()), + StatusFlags::NONE, + SqlType::Int(Some(7)), + ), + RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(), + ]; + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_mix (id, val) VALUES (@id, @v)".to_string(), + params, + None, + None, + ) + .await?; + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@v") + ); + + client.write_streamed_chunk(&wire).await?; + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + client + .execute( + "SELECT val FROM #plp_mix WHERE id = 7".to_string(), + None, + None, + ) + .await?; + if let Some(resultset) = client.get_current_resultset() { + let row = resultset.next_row().await?.expect("expected a row"); + match &row[0] { + ColumnValues::String(s) => assert_eq!(s.to_utf8_string(), value), + other => panic!("Expected String for nvarchar(max), got {other:?}"), + } + } else { + panic!("expected a result set"); + } + client.close_query().await?; + Ok(()) + } + + /// Streams a large `varbinary(max)` value in multiple chunks and verifies the + /// round-trip. + #[tokio::test] + async fn stream_varbinary_max_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_vbm (id INT, val VARBINARY(MAX))".to_string(), + None, + None, + ) + .await?; + client.close_query().await?; + + let value: Vec = (0..30_000u32).map(|i| (i % 256) as u8).collect(); + + let streamed = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::VarBinaryMax(None), + ) + .data_at_exec(); + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_vbm (id, val) VALUES (1, @v)".to_string(), + vec![streamed], + None, + None, + ) + .await?; + assert!(matches!(status, StreamedParamStatus::NeedData { .. })); + + for chunk in value.chunks(7_000) { + client.write_streamed_chunk(chunk).await?; + } + + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_vbm WHERE id = 1".to_string(), None, None) + .await?; + if let Some(resultset) = client.get_current_resultset() { + let row = resultset.next_row().await?.expect("expected a row"); + match &row[0] { + ColumnValues::Bytes(b) => assert_eq!(b.as_slice(), value.as_slice()), + other => panic!("Expected Bytes for varbinary(max), got {other:?}"), + } + } else { + panic!("expected a result set"); + } + client.close_query().await?; + Ok(()) + } + + /// Streams two `nvarchar(max)` parameters in one RPC, advancing through the + /// `NeedData` -> `NeedData` -> `Done` lifecycle. + #[tokio::test] + async fn stream_two_params_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_two (id INT, a NVARCHAR(MAX), b NVARCHAR(MAX))".to_string(), + None, + None, + ) + .await?; + client.close_query().await?; + + let a = "A".repeat(12_000); + let b = "B".repeat(9_000); + + let params = vec![ + RpcParameter::new( + Some("@a".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(), + RpcParameter::new( + Some("@b".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(), + ]; + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_two (id, a, b) VALUES (1, @a, @b)".to_string(), + params, + None, + None, + ) + .await?; + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@a") + ); + + client.write_streamed_chunk(&utf16le(&a)).await?; + let status = client.end_streamed_param().await?; + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@b") + ); + + client.write_streamed_chunk(&utf16le(&b)).await?; + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + client + .execute( + "SELECT LEN(a), LEN(b) FROM #plp_two WHERE id = 1".to_string(), + None, + None, + ) + .await?; + if let Some(resultset) = client.get_current_resultset() { + let row = resultset.next_row().await?.expect("expected a row"); + assert_eq!(row.len(), 2); + } else { + panic!("expected a result set"); + } + client.close_query().await?; + Ok(()) + } +} From 7e8c9eb04c2c62e6d1721506ac79d143e359f96a Mon Sep 17 00:00:00 2001 From: Shiwani Gupta Date: Tue, 4 Aug 2026 14:10:43 +0530 Subject: [PATCH 2/4] Add streamed PLP write test coverage Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql-tds/src/connection/tds_client.rs | 462 ++++++++++++++++++ .../src/message/parameters/rpc_parameters.rs | 20 + mssql-tds/tests/test_client_write_apis.rs | 164 ++++++- 3 files changed, 645 insertions(+), 1 deletion(-) diff --git a/mssql-tds/src/connection/tds_client.rs b/mssql-tds/src/connection/tds_client.rs index d76bac4f..2b91a6ee 100644 --- a/mssql-tds/src/connection/tds_client.rs +++ b/mssql-tds/src/connection/tds_client.rs @@ -5992,4 +5992,466 @@ mod tests { "expected a ForceColumnEncryption column-encryption error, got: {err}" ); } + + // ── Streamed (data-at-execution) PLP parameter write ── + + /// Reassembles the TDS packet stream captured by the mock transport into the + /// original contiguous request payload, stripping each 8-byte packet header. + fn reassemble_sent(sent: &[u8]) -> Vec { + let mut out = Vec::new(); + let mut off = 0; + while off < sent.len() { + let packet_len = u16::from_be_bytes([sent[off + 2], sent[off + 3]]) as usize; + out.extend_from_slice(&sent[off + 8..off + packet_len]); + off += packet_len; + } + out + } + + /// The 8-byte little-endian `PLP_UNKNOWN_LEN` sentinel that opens every + /// unknown-length PLP value on the wire. + const PLP_UNKNOWN_LEN_BYTES: [u8; 8] = [0xFE, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF]; + /// The 4-byte PLP terminator (a zero-length chunk header) that closes a value. + const PLP_TERMINATOR_BYTES: [u8; 4] = [0x00, 0x00, 0x00, 0x00]; + + /// Index of the last occurrence of `needle` in `haystack`, if any. + fn find_last(haystack: &[u8], needle: &[u8]) -> Option { + if needle.is_empty() || haystack.len() < needle.len() { + return None; + } + (0..=haystack.len() - needle.len()) + .rev() + .find(|&i| &haystack[i..i + needle.len()] == needle) + } + + /// A data-at-execution `varbinary(max)` parameter template named `name`. + /// varbinary keeps chunk bytes raw (no encoding), so tests can assert the + /// exact wire bytes they streamed. + fn streamed_varbinary(name: &str) -> RpcParameter { + RpcParameter::new( + Some(name.to_string()), + StatusFlags::NONE, + SqlType::VarBinaryMax(None), + ) + .data_at_exec() + } + + /// A single streamed chunk is framed as `[u32 len][bytes]` and the value is + /// then closed with the PLP terminator. + #[tokio::test] + async fn streamed_write_single_chunk_frames_value_and_terminator() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + let status = client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@v") + ); + + let chunk = [0xAAu8; 6]; + client.write_streamed_chunk(&chunk).await.unwrap(); + let status = client.end_streamed_param().await.unwrap(); + assert!(matches!(status, StreamedParamStatus::Done)); + + let payload = reassemble_sent(&sent.lock().unwrap()); + let pos = find_last(&payload, &PLP_UNKNOWN_LEN_BYTES).expect("value opener present"); + let after = &payload[pos + PLP_UNKNOWN_LEN_BYTES.len()..]; + + let mut expected = Vec::new(); + expected.extend_from_slice(&(chunk.len() as u32).to_le_bytes()); + expected.extend_from_slice(&chunk); + expected.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert_eq!(after, expected.as_slice()); + + // The lifecycle is complete: no streamed write remains parked. + assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + } + + /// Multiple chunks are each length-prefixed independently and the value is + /// closed by exactly one terminator — the incremental multi-chunk case. + #[tokio::test] + async fn streamed_write_multiple_chunks_each_length_prefixed_single_terminator() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + let chunks: [&[u8]; 3] = [&[0x01, 0x02], &[0x03, 0x04, 0x05], &[0x06]]; + for c in chunks { + client.write_streamed_chunk(c).await.unwrap(); + } + client.end_streamed_param().await.unwrap(); + + let payload = reassemble_sent(&sent.lock().unwrap()); + let pos = find_last(&payload, &PLP_UNKNOWN_LEN_BYTES).unwrap(); + let after = &payload[pos + PLP_UNKNOWN_LEN_BYTES.len()..]; + + let mut expected = Vec::new(); + for c in chunks { + expected.extend_from_slice(&(c.len() as u32).to_le_bytes()); + expected.extend_from_slice(c); + } + expected.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert_eq!(after, expected.as_slice()); + } + + /// An empty chunk must be ignored, not written: a zero-length chunk header + /// IS the terminator, so emitting one mid-value would truncate the value on + /// the server. Protocol-safety analogue of the read-side early-terminator + /// test. + #[tokio::test] + async fn streamed_write_empty_chunk_is_skipped_not_terminator() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + client.write_streamed_chunk(&[0x11, 0x22]).await.unwrap(); + client.write_streamed_chunk(&[]).await.unwrap(); // must be a no-op + client.end_streamed_param().await.unwrap(); + + let payload = reassemble_sent(&sent.lock().unwrap()); + let pos = find_last(&payload, &PLP_UNKNOWN_LEN_BYTES).unwrap(); + let after = &payload[pos + PLP_UNKNOWN_LEN_BYTES.len()..]; + + // Exactly one framed chunk followed by exactly one terminator; the empty + // chunk contributed nothing. + let mut expected = Vec::new(); + expected.extend_from_slice(&2u32.to_le_bytes()); + expected.extend_from_slice(&[0x11, 0x22]); + expected.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert_eq!(after, expected.as_slice()); + } + + /// Two data-at-execution parameters advance NeedData(@a) -> NeedData(@b) -> + /// Done, and both values are framed in order on the wire. + #[tokio::test] + async fn streamed_write_two_params_advances_need_data_then_done() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + let status = client + .begin_sp_executesql( + "INSERT INTO t(a, b) VALUES (@a, @b)".to_string(), + vec![streamed_varbinary("@a"), streamed_varbinary("@b")], + None, + None, + ) + .await + .unwrap(); + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@a") + ); + + let a_value = [0xA1u8; 4]; + client.write_streamed_chunk(&a_value).await.unwrap(); + let status = client.end_streamed_param().await.unwrap(); + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@b") + ); + + let b_value = [0xB2u8; 5]; + client.write_streamed_chunk(&b_value).await.unwrap(); + let status = client.end_streamed_param().await.unwrap(); + assert!(matches!(status, StreamedParamStatus::Done)); + + // @b is the last streamed param: its opener is the last sentinel, and the + // bytes after it are exactly @b's framed value + terminator. + let payload = reassemble_sent(&sent.lock().unwrap()); + let b_pos = find_last(&payload, &PLP_UNKNOWN_LEN_BYTES).unwrap(); + let mut expected_b = Vec::new(); + expected_b.extend_from_slice(&(b_value.len() as u32).to_le_bytes()); + expected_b.extend_from_slice(&b_value); + expected_b.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert_eq!( + &payload[b_pos + PLP_UNKNOWN_LEN_BYTES.len()..], + expected_b.as_slice() + ); + + // @a's opener precedes @b's, and its framed value + terminator sits right + // after it (followed by @b's parameter header). + let a_pos = find_last(&payload[..b_pos], &PLP_UNKNOWN_LEN_BYTES).unwrap(); + let mut expected_a = Vec::new(); + expected_a.extend_from_slice(&(a_value.len() as u32).to_le_bytes()); + expected_a.extend_from_slice(&a_value); + expected_a.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert!( + payload[a_pos + PLP_UNKNOWN_LEN_BYTES.len()..].starts_with(&expected_a), + "@a's framed value + terminator must immediately follow its opener" + ); + assert!(a_pos < b_pos, "@a must be serialized before @b"); + } + + /// With no data-at-execution parameter, `begin_sp_executesql` behaves like + /// the atomic path: it sends the RPC, consumes the response, and returns + /// Done without parking any streamed state. + #[tokio::test] + async fn begin_without_data_at_exec_delegates_and_returns_done() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + let status = client + .begin_sp_executesql( + "INSERT INTO t(id) VALUES (@id)".to_string(), + vec![RpcParameter::new( + Some("@id".to_string()), + StatusFlags::NONE, + SqlType::Int(Some(7)), + )], + None, + None, + ) + .await + .unwrap(); + + assert!(matches!(status, StreamedParamStatus::Done)); + assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + assert!( + !sent.lock().unwrap().is_empty(), + "the atomic RPC should have been sent" + ); + } + + /// Beginning a streamed execution while one is already active is rejected. + #[tokio::test] + async fn begin_while_stream_active_errors() { + let (mut client, _sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + let err = client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .expect_err("cannot begin a second streamed execution while one is active"); + assert!(matches!(err, UsageError(_))); + } + + /// A non-MAX data-at-execution parameter is rejected before any wire I/O, and + /// the client is left in the Idle state. + #[tokio::test] + async fn begin_rejects_non_max_data_at_exec_param() { + let mut client = create_test_client_with_tokens(vec![]); + + let bad = RpcParameter::new( + Some("@n".to_string()), + StatusFlags::NONE, + SqlType::Int(Some(1)), + ) + .data_at_exec(); + + let err = client + .begin_sp_executesql("SELECT @n".to_string(), vec![bad], None, None) + .await + .expect_err("non-max data-at-exec parameter must be rejected"); + assert!(matches!(err, UsageError(_))); + assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + } + + /// An unnamed data-at-execution parameter is rejected: streamed values must + /// be named so `sp_executesql` can match them by name. + #[tokio::test] + async fn begin_rejects_unnamed_data_at_exec_param() { + let mut client = create_test_client_with_tokens(vec![]); + + let bad = RpcParameter::new(None, StatusFlags::NONE, SqlType::VarBinaryMax(None)) + .data_at_exec(); + + let err = client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![bad], + None, + None, + ) + .await + .expect_err("unnamed data-at-exec parameter must be rejected"); + assert!(matches!(err, UsageError(_))); + assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + } + + /// `write_streamed_chunk` with no active streamed parameter is a usage error. + #[tokio::test] + async fn write_streamed_chunk_without_active_stream_errors() { + let mut client = create_test_client_with_tokens(vec![]); + let err = client + .write_streamed_chunk(&[0x01]) + .await + .expect_err("no active streamed parameter"); + assert!(matches!(err, UsageError(_))); + } + + /// `end_streamed_param` with no active streamed parameter is a usage error. + #[tokio::test] + async fn end_streamed_param_without_active_stream_errors() { + let mut client = create_test_client_with_tokens(vec![]); + let err = client + .end_streamed_param() + .await + .expect_err("no active streamed parameter"); + assert!(matches!(err, UsageError(_))); + } + + /// Little-endian UTF-16 encoding of `s`, matching how parameter names are + /// written on the wire (length-prefixed unicode). + fn utf16le(s: &str) -> Vec { + s.encode_utf16().flat_map(u16::to_le_bytes).collect() + } + + /// A materialized (send-now) named parameter and a data-at-execution one may + /// be mixed: the materialized value is sent atomically in the RPC prefix and + /// the streamed value follows. Offline analogue of the sparse non-PLP + PLP + /// e2e test. The materialized name precedes the streamed value opener, and + /// the streamed value frames correctly. + #[tokio::test] + async fn streamed_write_mixes_materialized_and_data_at_exec() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + let materialized = RpcParameter::new( + Some("@id".to_string()), + StatusFlags::NONE, + SqlType::Int(Some(7)), + ); + + let status = client + .begin_sp_executesql( + "INSERT INTO t(id, v) VALUES (@id, @v)".to_string(), + vec![materialized, streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@v") + ); + + let value = [0xCDu8; 8]; + client.write_streamed_chunk(&value).await.unwrap(); + assert!(matches!( + client.end_streamed_param().await.unwrap(), + StreamedParamStatus::Done + )); + + let payload = reassemble_sent(&sent.lock().unwrap()); + + // The streamed value opener is the last unknown-length sentinel, and the + // materialized parameter's name is serialized before it. + let v_pos = find_last(&payload, &PLP_UNKNOWN_LEN_BYTES).unwrap(); + let id_name = utf16le("@id"); + let id_pos = find_last(&payload[..v_pos], &id_name) + .expect("materialized parameter name must be serialized before the streamed value"); + assert!(id_pos < v_pos, "@id must precede the streamed @v value"); + + let mut expected = Vec::new(); + expected.extend_from_slice(&(value.len() as u32).to_le_bytes()); + expected.extend_from_slice(&value); + expected.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert_eq!(&payload[v_pos + PLP_UNKNOWN_LEN_BYTES.len()..], expected.as_slice()); + } + + /// After `begin_sp_executesql` parks the message, the streamed state must + /// retain the un-emitted parameters and the collation used to write their + /// TYPE_INFO. Analogue of the read side's pause-state-preserves-collation + /// test. + #[tokio::test] + async fn streamed_write_pause_state_preserves_pending_and_collation() { + let (mut client, _sent) = create_capturing_client(vec![done_no_more()]); + let expected_collation = client.negotiated_settings.database_collation; + + client + .begin_sp_executesql( + "INSERT INTO t(a, b, c) VALUES (@a, @b, @c)".to_string(), + vec![ + streamed_varbinary("@a"), + streamed_varbinary("@b"), + streamed_varbinary("@c"), + ], + None, + None, + ) + .await + .unwrap(); + + match &client.streamed_write_state { + StreamedWriteState::Active(ctx) => { + // @a is open for data; @b and @c remain pending, in order. + assert_eq!(ctx.pending.len(), 2); + assert_eq!(ctx.pending[0].name.as_deref(), Some("@b")); + assert_eq!(ctx.pending[1].name.as_deref(), Some("@c")); + assert_eq!(ctx.db_collation, expected_collation); + } + StreamedWriteState::Idle => panic!("streamed write should be Active after begin"), + } + } + + /// A single chunk larger than the TDS packet payload is split across multiple + /// packets on the wire yet reassembles into one contiguous length-prefixed + /// value. Analogue of multi-packet incremental reads. + #[tokio::test] + async fn streamed_write_large_chunk_spans_multiple_packets() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + // Larger than the mock transport's 4096-byte packet so the value must + // span several TDS packets. + let big: Vec = (0..10_000u32).map(|i| (i % 251) as u8).collect(); + client.write_streamed_chunk(&big).await.unwrap(); + client.end_streamed_param().await.unwrap(); + + // The captured stream is more than one packet. + assert!( + sent.lock().unwrap().len() > 4096, + "a 10 KB value must produce multiple TDS packets" + ); + + let payload = reassemble_sent(&sent.lock().unwrap()); + let pos = find_last(&payload, &PLP_UNKNOWN_LEN_BYTES).unwrap(); + + let mut expected = Vec::new(); + expected.extend_from_slice(&(big.len() as u32).to_le_bytes()); + expected.extend_from_slice(&big); + expected.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert_eq!(&payload[pos + PLP_UNKNOWN_LEN_BYTES.len()..], expected.as_slice()); + } } diff --git a/mssql-tds/src/message/parameters/rpc_parameters.rs b/mssql-tds/src/message/parameters/rpc_parameters.rs index f287d22e..5d90ddd1 100644 --- a/mssql-tds/src/message/parameters/rpc_parameters.rs +++ b/mssql-tds/src/message/parameters/rpc_parameters.rs @@ -951,6 +951,26 @@ mod tests { assert_eq!(streamed_header_bytes(¶m, false), expected); } + /// A named data-at-execution `varchar(max)` param serializes with the same + /// header shape as the other MAX types, using the varchar TYPE_INFO. Covers + /// the third streamable type (nvarchar/varchar/varbinary all supported). + #[test] + fn serialize_data_at_exec_varchar_max_named() { + let param = RpcParameter::new( + Some("@p".to_string()), + StatusFlags::NONE, + SqlType::VarcharMax(None), + ) + .data_at_exec(); + + let mut expected = vec![0x02, 0x40, 0x00, 0x70, 0x00]; // name: len 2, "@p" UTF-16LE + expected.push(StatusFlags::NONE.bits()); // status flags + expected.extend_from_slice(&type_info_bytes(&SqlType::VarcharMax(None))); // TYPE_INFO + expected.extend_from_slice(&0xFFFF_FFFF_FFFF_FFFEu64.to_le_bytes()); // PLP_UNKNOWN_LEN + + assert_eq!(streamed_header_bytes(¶m, false), expected); + } + /// A positional data-at-execution param writes a zero-length name byte in /// place of the name, then the same status/TYPE_INFO/PLP_UNKNOWN_LEN sequence. #[test] diff --git a/mssql-tds/tests/test_client_write_apis.rs b/mssql-tds/tests/test_client_write_apis.rs index 00263043..2018439f 100644 --- a/mssql-tds/tests/test_client_write_apis.rs +++ b/mssql-tds/tests/test_client_write_apis.rs @@ -312,4 +312,166 @@ mod streamed_plp_write { client.close_query().await?; Ok(()) } -} + + /// Streams a `varchar(max)` value into several rows in sequence, each via its + /// own `begin`/chunks/`end` cycle on the same connection, then verifies the + /// row count with `SELECT COUNT(*)`. Proves the streamed-write state machine + /// resets cleanly between rows so many rows can be written back-to-back. + #[tokio::test] + async fn stream_varchar_max_multiple_rows_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_rows (id INT, val VARCHAR(MAX))".to_string(), + None, + None, + ) + .await?; + client.close_query().await?; + + const ROW_COUNT: i32 = 5; + for id in 1..=ROW_COUNT { + // varchar(max) wire bytes are single-byte encoded; ASCII payload + // bytes equal the value's UTF-8 bytes, so stream them directly. + let value = format!("row-{id}-").repeat(3_000); + + let streamed = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::VarcharMax(None), + ) + .data_at_exec(); + + let status = client + .begin_sp_executesql( + format!("INSERT INTO #plp_rows (id, val) VALUES ({id}, @v)"), + vec![streamed], + None, + None, + ) + .await?; + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@v") + ); + + for chunk in value.as_bytes().chunks(4_096) { + client.write_streamed_chunk(chunk).await?; + } + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + } + + // Every streamed row must be present. + client + .execute( + "SELECT COUNT(*) FROM #plp_rows".to_string(), + None, + None, + ) + .await?; + if let Some(resultset) = client.get_current_resultset() { + let row = resultset.next_row().await?.expect("expected a count row"); + match &row[0] { + ColumnValues::Int(count) => assert_eq!(*count, ROW_COUNT), + other => panic!("Expected Int for COUNT(*), got {other:?}"), + } + } else { + panic!("expected a result set"); + } + client.close_query().await?; + + // Spot-check the last row's value survived the multi-row stream intact. + client + .execute( + format!("SELECT val FROM #plp_rows WHERE id = {ROW_COUNT}"), + None, + None, + ) + .await?; + if let Some(resultset) = client.get_current_resultset() { + let row = resultset.next_row().await?.expect("expected the last row"); + match &row[0] { + ColumnValues::String(s) => { + assert_eq!(s.to_utf8_string(), format!("row-{ROW_COUNT}-").repeat(3_000)); + } + other => panic!("Expected String for varchar(max), got {other:?}"), + } + } else { + panic!("expected a result set"); + } + client.close_query().await?; + Ok(()) + } + + /// A NULL value for a `nvarchar(max)` column round-trips as SQL NULL. NULL is + /// never streamed: it is bound directly as a materialized `NVarcharMax(None)` + /// value (which serializes to `PLP_NULL`), mirroring how a NULL data-at-exec + /// indicator is sent inline without ever requesting streamed data. + #[tokio::test] + async fn write_null_max_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_null (id INT, val NVARCHAR(MAX))".to_string(), + None, + None, + ) + .await?; + client.close_query().await?; + + // A NULL max parameter is materialized (value None -> PLP_NULL), so + // begin_sp_executesql completes atomically with no NeedData. + let null_param = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ); + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_null (id, val) VALUES (1, @v)".to_string(), + vec![null_param], + None, + None, + ) + .await?; + assert!( + matches!(status, StreamedParamStatus::Done), + "a materialized NULL parameter must not request streamed data" + ); + client.close_query().await?; + + client + .execute( + "SELECT val FROM #plp_null WHERE id = 1".to_string(), + None, + None, + ) + .await?; + if let Some(resultset) = client.get_current_resultset() { + let row = resultset.next_row().await?.expect("expected a row"); + assert!( + matches!(&row[0], ColumnValues::Null), + "expected SQL NULL, got {:?}", + &row[0] + ); + } else { + panic!("expected a result set"); + } + client.close_query().await?; + Ok(()) + } +} \ No newline at end of file From 792b4f1b1aac5d11fd8fea4470b888db745250fa Mon Sep 17 00:00:00 2001 From: Shiwani Gupta Date: Sun, 9 Aug 2026 21:42:29 +0530 Subject: [PATCH 3/4] Abort streamed PLP write on mid-stream failure A failed write_streamed_chunk/end_streamed_param leaves a partial message on the wire, so re-parking it as Active let the caller keep appending to a corrupt message. Add abort_streamed_write to drop the message (state -> Idle, not resumable) and flag the connection for reset, matching msodbcsql's DAE teardown on a failed send. Add fault-injecting unit tests for both abort paths and e2e variations (empty value, many small chunks, connection reuse). Migrate the e2e test file to the current execute/ResultSet API. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql-tds/src/connection/tds_client.rs | 335 +++++++++++++++--- .../src/message/parameters/rpc_parameters.rs | 4 +- mssql-tds/tests/test_client_write_apis.rs | 273 ++++++++++---- 3 files changed, 491 insertions(+), 121 deletions(-) diff --git a/mssql-tds/src/connection/tds_client.rs b/mssql-tds/src/connection/tds_client.rs index 2b91a6ee..cc7fc508 100644 --- a/mssql-tds/src/connection/tds_client.rs +++ b/mssql-tds/src/connection/tds_client.rs @@ -931,8 +931,9 @@ impl TdsClient { // existing serialize path) and data-at-execution values (streamed later), // preserving order. Mirrors ODBC, where every parameter is bound and some // are flagged SQL_DATA_AT_EXEC. - let (streamed_params, materialized_params): (Vec<_>, Vec<_>) = - named_params.into_iter().partition(RpcParameter::is_data_at_exec); + let (streamed_params, materialized_params): (Vec<_>, Vec<_>) = named_params + .into_iter() + .partition(RpcParameter::is_data_at_exec); // No data-at-execution parameters: this is just an ordinary execution. if streamed_params.is_empty() { @@ -957,14 +958,11 @@ impl TdsClient { )); } if param.name.is_none() { - return Err(UsageError( - "Streamed parameters must be named.".to_string(), - )); + return Err(UsageError("Streamed parameters must be named.".to_string())); } } - self.current_command_ce_setting = - ExecutionColumnEncryptionSetting::UseConnectionSetting; + self.current_command_ce_setting = ExecutionColumnEncryptionSetting::UseConnectionSetting; self.begin_command(); let reconnect_elapsed = self.check_and_reconnect(timeout_sec, cancel_handle).await?; @@ -976,7 +974,8 @@ impl TdsClient { // Always Encrypted is not supported when streaming parameter values. if self.should_encrypt_parameters() { return Err(UsageError( - "Streamed PLP parameter writes are not supported with Always Encrypted.".to_string(), + "Streamed PLP parameter writes are not supported with Always Encrypted." + .to_string(), )); } @@ -1031,7 +1030,12 @@ impl TdsClient { // materialized params, parked partway through. rpc.serialize_prefix(&mut packet_writer).await?; first - .serialize(&mut packet_writer, &database_collation, false, &GenericEncoder::new()) + .serialize( + &mut packet_writer, + &database_collation, + false, + &GenericEncoder::new(), + ) .await?; let message = packet_writer.suspend(); @@ -1072,7 +1076,8 @@ impl TdsClient { ))); } - let ctx = match std::mem::replace(&mut self.streamed_write_state, StreamedWriteState::Idle) { + let ctx = match std::mem::replace(&mut self.streamed_write_state, StreamedWriteState::Idle) + { StreamedWriteState::Active(ctx) => ctx, StreamedWriteState::Idle => { return Err(UsageError( @@ -1094,14 +1099,43 @@ impl TdsClient { .await; let message = packet_writer.suspend(); - // Re-park the message regardless of outcome so the state machine stays - // consistent; a write error fails the whole streamed operation. - self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { - message, - pending, - db_collation, - })); - result + match result { + Ok(()) => { + // Chunk framed cleanly: re-park the message so the next chunk or + // the terminator continues exactly where this one left off. + self.streamed_write_state = + StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation, + })); + Ok(()) + } + Err(e) => { + // A partial value chunk is now on the wire, so this message can no + // longer be continued safely. Drop it and abort the streamed + // write rather than re-parking it as resumable. + drop(message); + self.abort_streamed_write(); + Err(e) + } + } + } + + /// Aborts an in-progress streamed PLP write after a mid-stream failure. + /// + /// The half-written RPC message is dropped rather than re-parked, so a + /// subsequent [`write_streamed_chunk`](Self::write_streamed_chunk) or + /// [`end_streamed_param`](Self::end_streamed_param) fails cleanly (no active + /// stream) instead of appending to a corrupt message, and the connection is + /// marked for reset so the desynced wire stream cannot leak into the next + /// command. Mirrors msodbcsql, which tears down data-at-execution state on a + /// failed DAE send (`FlushStmt` / `ClearNonBLOBDAEParam`) rather than leaving + /// the value resumable. + fn abort_streamed_write(&mut self) { + self.streamed_write_state = StreamedWriteState::Idle; + self.execution_context.set_has_open_batch(false); + self.prepare_reset_connection(false); } /// Closes the streamed parameter currently open for data by writing its PLP @@ -1118,7 +1152,8 @@ impl TdsClient { /// Returns a usage error if no streamed parameter is currently open, or an /// I/O error if sending fails. pub async fn end_streamed_param(&mut self) -> TdsResult { - let ctx = match std::mem::replace(&mut self.streamed_write_state, StreamedWriteState::Idle) { + let ctx = match std::mem::replace(&mut self.streamed_write_state, StreamedWriteState::Idle) + { StreamedWriteState::Active(ctx) => ctx, StreamedWriteState::Idle => { return Err(UsageError( @@ -1133,33 +1168,71 @@ impl TdsClient { } = *ctx; let mut packet_writer = PacketWriter::resume(message, self.transport.as_writer()); - packet_writer.write_u32_async(PLP_TERMINATOR).await?; - - if let Some(next) = pending.pop_front() { - let next_name = next - .name - .clone() - .expect("streamed parameter names validated at begin"); - next.serialize(&mut packet_writer, &db_collation, false, &GenericEncoder::new()) - .await?; - let message = packet_writer.suspend(); - self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { - message, - pending, - db_collation, - })); - return Ok(StreamedParamStatus::NeedData { - param_name: next_name, - }); - } - // Last streamed parameter closed: send the message and consume the - // response exactly like execute_sp_executesql does. - packet_writer.finalize().await?; - drop(packet_writer); + // Close the current value with its terminator, then either open the next + // streamed parameter's header or (for the last one) finalize the send. + // Anything that fails mid-message aborts the whole streamed write. + let write_outcome = async { + packet_writer.write_u32_async(PLP_TERMINATOR).await?; + match pending.pop_front() { + Some(next) => { + let next_name = next + .name + .clone() + .expect("streamed parameter names validated at begin"); + next.serialize( + &mut packet_writer, + &db_collation, + false, + &GenericEncoder::new(), + ) + .await?; + Ok(Some(next_name)) + } + None => { + packet_writer.finalize().await?; + Ok(None) + } + } + } + .await; - self.position_on_first_result().await?; - Ok(StreamedParamStatus::Done) + match write_outcome { + // Another streamed parameter is now open for data. + Ok(Some(next_name)) => { + let message = packet_writer.suspend(); + self.streamed_write_state = + StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation, + })); + Ok(StreamedParamStatus::NeedData { + param_name: next_name, + }) + } + // Last streamed parameter closed and the message was sent: consume + // the response exactly like execute_sp_executesql does. A failure + // reading the response also aborts (the request is already on the + // wire, so the connection must be reset before reuse). + Ok(None) => { + drop(packet_writer); + match self.position_on_first_result().await { + Ok(_) => Ok(StreamedParamStatus::Done), + Err(e) => { + self.abort_streamed_write(); + Err(e) + } + } + } + // Terminator or next-parameter header write failed mid-message: drop + // the half-written message and abort rather than leave it resumable. + Err(e) => { + drop(packet_writer); + self.abort_streamed_write(); + Err(e) + } + } } /// Executes a bulk load operation using zero-copy streaming. @@ -4317,6 +4390,9 @@ mod tests { /// `read_row_column` down a specific arm (e.g. a `PlpPaused` result that /// makes the cursor emit `CursorColumn::PlpStreaming`). resume_results: VecDeque, + /// When set, the next (and every subsequent) `send` fails, simulating a + /// mid-message wire failure. Shared so a test can flip it after setup. + send_should_fail: Arc, } impl TestTransport { @@ -4328,7 +4404,8 @@ mod tests { sent: Arc::new(std::sync::Mutex::new(Vec::new())), packet_data: Vec::new(), packet_pos: 0, - resume_results: VecDeque::new(), +resume_results: VecDeque::new(), + send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4340,7 +4417,8 @@ mod tests { sent: Arc::new(std::sync::Mutex::new(Vec::new())), packet_data: Vec::new(), packet_pos: 0, - resume_results: VecDeque::new(), +resume_results: VecDeque::new(), + send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4352,7 +4430,8 @@ mod tests { sent: Arc::new(std::sync::Mutex::new(Vec::new())), packet_data, packet_pos: 0, - resume_results: VecDeque::new(), +resume_results: VecDeque::new(), + send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4454,6 +4533,14 @@ mod tests { #[async_trait] impl NetworkWriter for TestTransport { async fn send(&mut self, data: &[u8]) -> TdsResult<()> { + if self + .send_should_fail + .load(std::sync::atomic::Ordering::SeqCst) + { + return Err(crate::error::Error::ConnectionClosed( + "injected send failure".to_string(), + )); + } self.sent.lock().unwrap().extend_from_slice(data); Ok(()) } @@ -4658,6 +4745,27 @@ mod tests { (client, sent) } + /// Like [`create_capturing_client`], but also returns a shared flag that, + /// when set, makes the transport's next `send` fail — used to exercise the + /// mid-stream abort path of the streamed PLP write. + fn create_failing_capturing_client( + tokens: Vec, + ) -> (TdsClient, std::sync::Arc) { + let transport = Box::new(TestTransport::with_tokens(tokens)); + let fail = std::sync::Arc::clone(&transport.send_should_fail); + let negotiated_settings = + crate::handler::handler_factory::create_test_negotiated_settings_internal(); + let execution_context = crate::connection::execution_context::ExecutionContext::new(); + let client_context = ClientContext::with_data_source("tcp:localhost,1433"); + let client = TdsClient::new( + transport, + negotiated_settings, + execution_context, + client_context, + ); + (client, fail) + } + fn done_no_more() -> Tokens { Tokens::Done(DoneToken { status: DoneStatus::FINAL, @@ -6071,7 +6179,10 @@ mod tests { assert_eq!(after, expected.as_slice()); // The lifecycle is complete: no streamed write remains parked. - assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Idle + )); } /// Multiple chunks are each length-prefixed independently and the value is @@ -6224,13 +6335,115 @@ mod tests { .unwrap(); assert!(matches!(status, StreamedParamStatus::Done)); - assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Idle + )); assert!( !sent.lock().unwrap().is_empty(), "the atomic RPC should have been sent" ); } + /// A mid-value send failure aborts the streamed write: the error is + /// surfaced, the parked message is dropped (state returns to `Idle`, not left + /// `Active`/resumable), the connection is flagged for reset so the desynced + /// wire stream cannot leak into the next command, and further streamed calls + /// fail cleanly. Mirrors msodbcsql tearing down data-at-execution state on a + /// failed DAE send rather than leaving the value resumable. + #[tokio::test] + async fn streamed_write_chunk_send_failure_aborts_stream() { + let (mut client, fail) = create_failing_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Active(_) + )); + + // Fail the next wire send; a chunk large enough to overflow the packet + // payload buffer forces a flush (and thus a `send`) inside + // write_streamed_chunk. + fail.store(true, std::sync::atomic::Ordering::SeqCst); + let big = vec![0xABu8; 10_000]; + let _err = client + .write_streamed_chunk(&big) + .await + .expect_err("a send failure must surface as an error"); + + // The streamed write is aborted, not left resumable. + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Idle + )); + // The connection is flagged for reset so the next command re-syncs. + assert!(matches!( + client.transport.as_writer().take_reset_mode(), + ResetConnectionMode::Reset + )); + + // Further streamed calls now fail cleanly (no active stream) rather than + // appending to the corrupt message. + let followup = client + .write_streamed_chunk(&[0x00]) + .await + .expect_err("no active streamed parameter after abort"); + assert!(matches!(followup, UsageError(_))); + let end = client + .end_streamed_param() + .await + .expect_err("no active streamed parameter after abort"); + assert!(matches!(end, UsageError(_))); + } + + /// A send failure while finalizing the last streamed parameter (flushing the + /// terminator + message) aborts the same way: error surfaced, state `Idle`, + /// connection flagged for reset. + #[tokio::test] + async fn end_streamed_param_send_failure_aborts_stream() { + let (mut client, fail) = create_failing_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + // A small chunk stays buffered (fits one packet, no send yet). + client + .write_streamed_chunk(&[0x01, 0x02, 0x03]) + .await + .unwrap(); + + // finalize() inside end_streamed_param flushes the message -> send fails. + fail.store(true, std::sync::atomic::Ordering::SeqCst); + let _err = client + .end_streamed_param() + .await + .expect_err("finalize send failure must surface"); + + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Idle + )); + assert!(matches!( + client.transport.as_writer().take_reset_mode(), + ResetConnectionMode::Reset + )); + } + /// Beginning a streamed execution while one is already active is rejected. #[tokio::test] async fn begin_while_stream_active_errors() { @@ -6276,7 +6489,10 @@ mod tests { .await .expect_err("non-max data-at-exec parameter must be rejected"); assert!(matches!(err, UsageError(_))); - assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Idle + )); } /// An unnamed data-at-execution parameter is rejected: streamed values must @@ -6285,8 +6501,8 @@ mod tests { async fn begin_rejects_unnamed_data_at_exec_param() { let mut client = create_test_client_with_tokens(vec![]); - let bad = RpcParameter::new(None, StatusFlags::NONE, SqlType::VarBinaryMax(None)) - .data_at_exec(); + let bad = + RpcParameter::new(None, StatusFlags::NONE, SqlType::VarBinaryMax(None)).data_at_exec(); let err = client .begin_sp_executesql( @@ -6298,7 +6514,10 @@ mod tests { .await .expect_err("unnamed data-at-exec parameter must be rejected"); assert!(matches!(err, UsageError(_))); - assert!(matches!(client.streamed_write_state, StreamedWriteState::Idle)); + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Idle + )); } /// `write_streamed_chunk` with no active streamed parameter is a usage error. @@ -6378,7 +6597,10 @@ mod tests { expected.extend_from_slice(&(value.len() as u32).to_le_bytes()); expected.extend_from_slice(&value); expected.extend_from_slice(&PLP_TERMINATOR_BYTES); - assert_eq!(&payload[v_pos + PLP_UNKNOWN_LEN_BYTES.len()..], expected.as_slice()); + assert_eq!( + &payload[v_pos + PLP_UNKNOWN_LEN_BYTES.len()..], + expected.as_slice() + ); } /// After `begin_sp_executesql` parks the message, the streamed state must @@ -6452,6 +6674,9 @@ mod tests { expected.extend_from_slice(&(big.len() as u32).to_le_bytes()); expected.extend_from_slice(&big); expected.extend_from_slice(&PLP_TERMINATOR_BYTES); - assert_eq!(&payload[pos + PLP_UNKNOWN_LEN_BYTES.len()..], expected.as_slice()); + assert_eq!( + &payload[pos + PLP_UNKNOWN_LEN_BYTES.len()..], + expected.as_slice() + ); } } diff --git a/mssql-tds/src/message/parameters/rpc_parameters.rs b/mssql-tds/src/message/parameters/rpc_parameters.rs index 5d90ddd1..fe317d74 100644 --- a/mssql-tds/src/message/parameters/rpc_parameters.rs +++ b/mssql-tds/src/message/parameters/rpc_parameters.rs @@ -975,8 +975,8 @@ mod tests { /// place of the name, then the same status/TYPE_INFO/PLP_UNKNOWN_LEN sequence. #[test] fn serialize_data_at_exec_positional() { - let param = RpcParameter::new(None, StatusFlags::NONE, SqlType::VarBinaryMax(None)) - .data_at_exec(); + let param = + RpcParameter::new(None, StatusFlags::NONE, SqlType::VarBinaryMax(None)).data_at_exec(); let mut expected = vec![0x00]; // zero-length name (positional) expected.push(StatusFlags::NONE.bits()); diff --git a/mssql-tds/tests/test_client_write_apis.rs b/mssql-tds/tests/test_client_write_apis.rs index 2018439f..92785f8f 100644 --- a/mssql-tds/tests/test_client_write_apis.rs +++ b/mssql-tds/tests/test_client_write_apis.rs @@ -13,7 +13,7 @@ mod common; mod streamed_plp_write { use crate::common::{build_tcp_datasource, create_context, init_tracing}; - use mssql_tds::connection::tds_client::{ResultSet, ResultSetClient, StreamedParamStatus}; + use mssql_tds::connection::tds_client::{ResultSet, StreamedParamStatus}; use mssql_tds::connection_provider::tds_connection_provider::TdsConnectionProvider; use mssql_tds::datatypes::column_values::ColumnValues; use mssql_tds::datatypes::sqltypes::SqlType; @@ -38,8 +38,7 @@ mod streamed_plp_write { client .execute( "CREATE TABLE #plp_nvm (id INT, val NVARCHAR(MAX))".to_string(), - None, - None, + (), ) .await?; client.close_query().await?; @@ -77,10 +76,10 @@ mod streamed_plp_write { client.close_query().await?; client - .execute("SELECT val FROM #plp_nvm WHERE id = 1".to_string(), None, None) + .execute("SELECT val FROM #plp_nvm WHERE id = 1".to_string(), ()) .await?; - if let Some(resultset) = client.get_current_resultset() { - let row = resultset.next_row().await?.expect("expected a row"); + { + let row = client.next_row().await?.expect("expected a row"); match &row[0] { ColumnValues::String(s) => { let round_tripped = s.to_utf8_string(); @@ -89,8 +88,6 @@ mod streamed_plp_write { } other => panic!("Expected String for nvarchar(max), got {other:?}"), } - } else { - panic!("expected a result set"); } client.close_query().await?; Ok(()) @@ -112,8 +109,7 @@ mod streamed_plp_write { client .execute( "CREATE TABLE #plp_mix (id INT, val NVARCHAR(MAX))".to_string(), - None, - None, + (), ) .await?; client.close_query().await?; @@ -153,20 +149,14 @@ mod streamed_plp_write { client.close_query().await?; client - .execute( - "SELECT val FROM #plp_mix WHERE id = 7".to_string(), - None, - None, - ) + .execute("SELECT val FROM #plp_mix WHERE id = 7".to_string(), ()) .await?; - if let Some(resultset) = client.get_current_resultset() { - let row = resultset.next_row().await?.expect("expected a row"); + { + let row = client.next_row().await?.expect("expected a row"); match &row[0] { ColumnValues::String(s) => assert_eq!(s.to_utf8_string(), value), other => panic!("Expected String for nvarchar(max), got {other:?}"), } - } else { - panic!("expected a result set"); } client.close_query().await?; Ok(()) @@ -186,8 +176,7 @@ mod streamed_plp_write { client .execute( "CREATE TABLE #plp_vbm (id INT, val VARBINARY(MAX))".to_string(), - None, - None, + (), ) .await?; client.close_query().await?; @@ -220,16 +209,14 @@ mod streamed_plp_write { client.close_query().await?; client - .execute("SELECT val FROM #plp_vbm WHERE id = 1".to_string(), None, None) + .execute("SELECT val FROM #plp_vbm WHERE id = 1".to_string(), ()) .await?; - if let Some(resultset) = client.get_current_resultset() { - let row = resultset.next_row().await?.expect("expected a row"); + { + let row = client.next_row().await?.expect("expected a row"); match &row[0] { ColumnValues::Bytes(b) => assert_eq!(b.as_slice(), value.as_slice()), other => panic!("Expected Bytes for varbinary(max), got {other:?}"), } - } else { - panic!("expected a result set"); } client.close_query().await?; Ok(()) @@ -249,8 +236,7 @@ mod streamed_plp_write { client .execute( "CREATE TABLE #plp_two (id INT, a NVARCHAR(MAX), b NVARCHAR(MAX))".to_string(), - None, - None, + (), ) .await?; client.close_query().await?; @@ -299,15 +285,12 @@ mod streamed_plp_write { client .execute( "SELECT LEN(a), LEN(b) FROM #plp_two WHERE id = 1".to_string(), - None, - None, + (), ) .await?; - if let Some(resultset) = client.get_current_resultset() { - let row = resultset.next_row().await?.expect("expected a row"); + { + let row = client.next_row().await?.expect("expected a row"); assert_eq!(row.len(), 2); - } else { - panic!("expected a result set"); } client.close_query().await?; Ok(()) @@ -329,8 +312,7 @@ mod streamed_plp_write { client .execute( "CREATE TABLE #plp_rows (id INT, val VARCHAR(MAX))".to_string(), - None, - None, + (), ) .await?; client.close_query().await?; @@ -370,20 +352,14 @@ mod streamed_plp_write { // Every streamed row must be present. client - .execute( - "SELECT COUNT(*) FROM #plp_rows".to_string(), - None, - None, - ) + .execute("SELECT COUNT(*) FROM #plp_rows".to_string(), ()) .await?; - if let Some(resultset) = client.get_current_resultset() { - let row = resultset.next_row().await?.expect("expected a count row"); + { + let row = client.next_row().await?.expect("expected a count row"); match &row[0] { ColumnValues::Int(count) => assert_eq!(*count, ROW_COUNT), other => panic!("Expected Int for COUNT(*), got {other:?}"), } - } else { - panic!("expected a result set"); } client.close_query().await?; @@ -391,20 +367,20 @@ mod streamed_plp_write { client .execute( format!("SELECT val FROM #plp_rows WHERE id = {ROW_COUNT}"), - None, - None, + (), ) .await?; - if let Some(resultset) = client.get_current_resultset() { - let row = resultset.next_row().await?.expect("expected the last row"); + { + let row = client.next_row().await?.expect("expected the last row"); match &row[0] { ColumnValues::String(s) => { - assert_eq!(s.to_utf8_string(), format!("row-{ROW_COUNT}-").repeat(3_000)); + assert_eq!( + s.to_utf8_string(), + format!("row-{ROW_COUNT}-").repeat(3_000) + ); } other => panic!("Expected String for varchar(max), got {other:?}"), } - } else { - panic!("expected a result set"); } client.close_query().await?; Ok(()) @@ -426,8 +402,7 @@ mod streamed_plp_write { client .execute( "CREATE TABLE #plp_null (id INT, val NVARCHAR(MAX))".to_string(), - None, - None, + (), ) .await?; client.close_query().await?; @@ -455,23 +430,193 @@ mod streamed_plp_write { client.close_query().await?; client - .execute( - "SELECT val FROM #plp_null WHERE id = 1".to_string(), - None, - None, - ) + .execute("SELECT val FROM #plp_null WHERE id = 1".to_string(), ()) .await?; - if let Some(resultset) = client.get_current_resultset() { - let row = resultset.next_row().await?.expect("expected a row"); + { + let row = client.next_row().await?.expect("expected a row"); assert!( matches!(&row[0], ColumnValues::Null), "expected SQL NULL, got {:?}", &row[0] ); - } else { - panic!("expected a result set"); } client.close_query().await?; Ok(()) } -} \ No newline at end of file + + /// Streaming zero chunks then ending yields an empty (non-NULL) value: the + /// value opener followed immediately by the terminator encodes a present, + /// zero-length value. Distinct from the NULL path + /// (`write_null_max_round_trips`). + #[tokio::test] + async fn stream_empty_value_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_empty (id INT, val NVARCHAR(MAX))".to_string(), + (), + ) + .await?; + client.close_query().await?; + + let streamed = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(); + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_empty (id, val) VALUES (1, @v)".to_string(), + vec![streamed], + None, + None, + ) + .await?; + assert!(matches!(status, StreamedParamStatus::NeedData { .. })); + + // No chunks: close the value immediately. + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_empty WHERE id = 1".to_string(), ()) + .await?; + { + let row = client.next_row().await?.expect("expected a row"); + match &row[0] { + ColumnValues::String(s) => assert_eq!(s.to_utf8_string(), ""), + other => panic!("Expected empty String for nvarchar(max), got {other:?}"), + } + } + client.close_query().await?; + Ok(()) + } + + /// The same value split into many tiny (2-byte) chunks reassembles intact, + /// stressing per-chunk length-prefixing across a large call count. Each chunk + /// is one UTF-16 code unit, so the boundaries never split a code unit. + #[tokio::test] + async fn stream_nvarchar_max_many_small_chunks_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_small (id INT, val NVARCHAR(MAX))".to_string(), + (), + ) + .await?; + client.close_query().await?; + + let value = "abcd".repeat(2_000); // 8000 chars + let wire = utf16le(&value); + + let streamed = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(); + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_small (id, val) VALUES (1, @v)".to_string(), + vec![streamed], + None, + None, + ) + .await?; + assert!(matches!(status, StreamedParamStatus::NeedData { .. })); + + // One UTF-16 code unit (2 bytes) per chunk: 8000 chunks. + for chunk in wire.chunks(2) { + client.write_streamed_chunk(chunk).await?; + } + + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_small WHERE id = 1".to_string(), ()) + .await?; + { + let row = client.next_row().await?.expect("expected a row"); + match &row[0] { + ColumnValues::String(s) => assert_eq!(s.to_utf8_string(), value), + other => panic!("Expected String for nvarchar(max), got {other:?}"), + } + } + client.close_query().await?; + Ok(()) + } + + /// After a streamed write completes, the same connection is reusable for an + /// ordinary query: the happy path leaves no desynced wire state behind. (The + /// failure path, which flags the connection for reset, is covered by the + /// offline abort unit tests.) + #[tokio::test] + async fn stream_then_normal_execute_reuses_connection() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_reuse (id INT, val VARBINARY(MAX))".to_string(), + (), + ) + .await?; + client.close_query().await?; + + let value: Vec = (0..5_000u32).map(|i| (i % 256) as u8).collect(); + let streamed = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::VarBinaryMax(None), + ) + .data_at_exec(); + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_reuse (id, val) VALUES (1, @v)".to_string(), + vec![streamed], + None, + None, + ) + .await?; + assert!(matches!(status, StreamedParamStatus::NeedData { .. })); + client.write_streamed_chunk(&value).await?; + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + // Reuse the same client for a plain query. + client.execute("SELECT 42".to_string(), ()).await?; + { + let row = client.next_row().await?.expect("expected a row"); + match &row[0] { + ColumnValues::Int(v) => assert_eq!(*v, 42), + other => panic!("Expected Int, got {other:?}"), + } + } + client.close_query().await?; + Ok(()) + } +} From 907fd5460984cb2e0a068a12fbb259caf23ce3ec Mon Sep 17 00:00:00 2001 From: Shiwani Gupta Date: Mon, 10 Aug 2026 11:08:31 +0530 Subject: [PATCH 4/4] Support NULL on the streamed PLP parameter write path Defer the PLP length field: the data-at-exec serialize now writes only the parameter header (status + TYPE_INFO), and the length field (PLP_UNKNOWN_LEN opener or PLP_NULL) is written lazily by the streaming driver. This lets a streamed parameter resolve to NULL before any data is sent, matching msodbcsql path 2 (SQLPutData(SQL_NULL_DATA)). Add write_streamed_null(); end_streamed_param emits PLP_NULL for a NULL-signalled param, the terminator for a value, or opener+terminator for an untouched (empty) param. Guard both orderings (chunk-after-null, null-after-chunk). Add unit and e2e tests. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- mssql-tds/src/connection/tds_client.rs | 311 +++++++++++++++++- .../src/message/parameters/rpc_parameters.rs | 37 +-- mssql-tds/tests/test_client_write_apis.rs | 63 +++- 3 files changed, 372 insertions(+), 39 deletions(-) diff --git a/mssql-tds/src/connection/tds_client.rs b/mssql-tds/src/connection/tds_client.rs index cc7fc508..142811f7 100644 --- a/mssql-tds/src/connection/tds_client.rs +++ b/mssql-tds/src/connection/tds_client.rs @@ -10,7 +10,7 @@ use crate::datatypes::encoder::GenericEncoder; use crate::datatypes::row_writer::{DefaultRowWriter, DiscardRowWriter, RowWriter}; use crate::datatypes::sql_string::SqlString; use crate::datatypes::sqltypes::SqlType; -use crate::datatypes::tds_value_serializer::PLP_TERMINATOR; +use crate::datatypes::tds_value_serializer::{PLP_NULL, PLP_TERMINATOR, PLP_UNKNOWN_LEN}; use crate::error::Error::UsageError; use crate::error::{SqlErrorInfo, SqlInfoMessage}; use crate::io::packet_writer::{PacketWriter, SuspendedMessage, TdsPacketWriter}; @@ -164,13 +164,23 @@ enum StreamedWriteState { #[derive(Debug)] struct StreamedWriteContext { /// The parked outgoing RPC message. The parameter currently open for data - /// has already had its header + `PLP_UNKNOWN_LEN` written; the next bytes - /// appended are its value chunks. + /// has had its header (status byte + `TYPE_INFO`) written; the PLP length + /// field (the unknown-length opener, or `PLP_NULL`) and the value chunks are + /// appended afterwards by the streaming driver. message: SuspendedMessage, /// Streamed parameters whose headers have not yet been written, in order. pending: std::collections::VecDeque, /// Collation used to write `TYPE_INFO` for subsequent streamed parameters. db_collation: SqlCollation, + /// `true` once the unknown-length PLP opener (`PLP_UNKNOWN_LEN`) has been + /// written for the currently-open parameter — i.e. at least one value chunk + /// has been emitted. The opener is written lazily (on the first chunk) so a + /// parameter can still resolve to NULL before any data is sent. + value_opened: bool, + /// `true` when the caller has signalled the currently-open parameter is NULL + /// via [`TdsClient::write_streamed_null`]. Closing the parameter then writes + /// `PLP_NULL` instead of an opener + terminator, and no chunks may follow. + null_signaled: bool, } /// Active TDS connection to a SQL Server instance. @@ -1043,6 +1053,8 @@ impl TdsClient { message, pending, db_collation: database_collation, + value_opened: false, + null_signaled: false, })); Ok(StreamedParamStatus::NeedData { @@ -1061,9 +1073,10 @@ impl TdsClient { /// terminator, so it must never be emitted mid-value. /// /// # Errors - /// Returns a usage error if no streamed parameter is currently open, or if - /// `chunk` is longer than [`u32::MAX`] bytes (the PLP chunk-length field is - /// 32-bit). + /// Returns a usage error if no streamed parameter is currently open, if the + /// parameter was already marked NULL via + /// [`write_streamed_null`](Self::write_streamed_null), or if `chunk` is longer + /// than [`u32::MAX`] bytes (the PLP chunk-length field is 32-bit). pub async fn write_streamed_chunk(&mut self, chunk: &[u8]) -> TdsResult<()> { if chunk.is_empty() { return Ok(()); @@ -1089,10 +1102,35 @@ impl TdsClient { message, pending, db_collation, + value_opened, + null_signaled, } = *ctx; + // A parameter already signalled NULL cannot carry data. Re-park the + // (still-clean) message and reject; this is a caller sequencing error, + // not a wire failure, so the stream stays usable. + if null_signaled { + self.streamed_write_state = + StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation, + value_opened, + null_signaled, + })); + return Err(UsageError( + "write_streamed_chunk called after the parameter was marked NULL.".to_string(), + )); + } + let mut packet_writer = PacketWriter::resume(message, self.transport.as_writer()); let result = async { + // Lazily write the unknown-length opener on the first chunk. Deferring + // it (rather than emitting it at header time) is what lets a streamed + // parameter still resolve to NULL before any data is sent. + if !value_opened { + packet_writer.write_u64_async(PLP_UNKNOWN_LEN).await?; + } packet_writer.write_u32_async(chunk.len() as u32).await?; packet_writer.write_async(chunk).await } @@ -1102,12 +1140,15 @@ impl TdsClient { match result { Ok(()) => { // Chunk framed cleanly: re-park the message so the next chunk or - // the terminator continues exactly where this one left off. + // the terminator continues exactly where this one left off. The + // opener is now on the wire, so the value is committed as present. self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { message, pending, db_collation, + value_opened: true, + null_signaled: false, })); Ok(()) } @@ -1138,8 +1179,73 @@ impl TdsClient { self.prepare_reset_connection(false); } - /// Closes the streamed parameter currently open for data by writing its PLP - /// terminator. + /// Marks the streamed parameter currently open for data as SQL NULL. + /// + /// Call this instead of [`write_streamed_chunk`](Self::write_streamed_chunk) + /// when a data-at-execution parameter resolves to NULL (the ODBC + /// `SQLPutData(SQL_NULL_DATA)` case). No bytes are written now; when the + /// parameter is closed with [`end_streamed_param`](Self::end_streamed_param) + /// the driver emits `PLP_NULL` instead of an unknown-length opener + + /// terminator. Mirrors msodbcsql, which writes `VARMAX_LENGTH_NULL` with no + /// chunks for a DAE parameter that resolves to NULL. + /// + /// # Errors + /// Returns a usage error if no streamed parameter is currently open, or if + /// value chunks have already been written for this parameter (a value that + /// has begun streaming cannot become NULL). + pub fn write_streamed_null(&mut self) -> TdsResult<()> { + let ctx = match std::mem::replace(&mut self.streamed_write_state, StreamedWriteState::Idle) + { + StreamedWriteState::Active(ctx) => ctx, + StreamedWriteState::Idle => { + return Err(UsageError( + "write_streamed_null called with no active streamed parameter.".to_string(), + )); + } + }; + let StreamedWriteContext { + message, + pending, + db_collation, + value_opened, + .. + } = *ctx; + + if value_opened { + // Chunks are already on the wire; the value cannot become NULL. Re-park + // the (still-clean) message — this is a caller sequencing error. + self.streamed_write_state = + StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation, + value_opened, + null_signaled: false, + })); + return Err(UsageError( + "write_streamed_null called after value chunks were already written.".to_string(), + )); + } + + self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation, + value_opened: false, + null_signaled: true, + })); + Ok(()) + } + + /// Closes the streamed parameter currently open for data. + /// Closes the streamed parameter currently open for data. + /// + /// The current parameter's value is closed on the wire: + /// - if it was marked NULL via [`write_streamed_null`](Self::write_streamed_null), + /// `PLP_NULL` is written (no opener, no terminator); + /// - if one or more chunks were written, the `PLP_TERMINATOR` is written; and + /// - if neither (an untouched parameter), the unknown-length opener + + /// terminator are written, encoding a present, zero-length value. /// /// If more streamed parameters remain, the next one's header is written and /// [`StreamedParamStatus::NeedData`] is returned (stream its chunks next). @@ -1165,15 +1271,28 @@ impl TdsClient { message, mut pending, db_collation, + value_opened, + null_signaled, } = *ctx; let mut packet_writer = PacketWriter::resume(message, self.transport.as_writer()); - // Close the current value with its terminator, then either open the next - // streamed parameter's header or (for the last one) finalize the send. - // Anything that fails mid-message aborts the whole streamed write. + // Close the current value, then either open the next streamed parameter's + // header or (for the last one) finalize the send. Anything that fails + // mid-message aborts the whole streamed write. let write_outcome = async { - packet_writer.write_u32_async(PLP_TERMINATOR).await?; + if null_signaled { + // NULL: the length field is PLP_NULL and no chunks/terminator + // follow. + packet_writer.write_u64_async(PLP_NULL).await?; + } else { + // An untouched parameter never wrote its opener; write it now so + // the value is a present, zero-length value rather than absent. + if !value_opened { + packet_writer.write_u64_async(PLP_UNKNOWN_LEN).await?; + } + packet_writer.write_u32_async(PLP_TERMINATOR).await?; + } match pending.pop_front() { Some(next) => { let next_name = next @@ -1206,6 +1325,10 @@ impl TdsClient { message, pending, db_collation, + // Fresh parameter: its opener has not been written and it + // has not been marked NULL. + value_opened: false, + null_signaled: false, })); Ok(StreamedParamStatus::NeedData { param_name: next_name, @@ -4404,7 +4527,7 @@ mod tests { sent: Arc::new(std::sync::Mutex::new(Vec::new())), packet_data: Vec::new(), packet_pos: 0, -resume_results: VecDeque::new(), + resume_results: VecDeque::new(), send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4417,7 +4540,7 @@ resume_results: VecDeque::new(), sent: Arc::new(std::sync::Mutex::new(Vec::new())), packet_data: Vec::new(), packet_pos: 0, -resume_results: VecDeque::new(), + resume_results: VecDeque::new(), send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4430,7 +4553,7 @@ resume_results: VecDeque::new(), sent: Arc::new(std::sync::Mutex::new(Vec::new())), packet_data, packet_pos: 0, -resume_results: VecDeque::new(), + resume_results: VecDeque::new(), send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -6121,6 +6244,9 @@ resume_results: VecDeque::new(), const PLP_UNKNOWN_LEN_BYTES: [u8; 8] = [0xFE, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF]; /// The 4-byte PLP terminator (a zero-length chunk header) that closes a value. const PLP_TERMINATOR_BYTES: [u8; 4] = [0x00, 0x00, 0x00, 0x00]; + /// The 8-byte little-endian `PLP_NULL` sentinel: a MAX-type value that is SQL + /// NULL. No chunks or terminator follow it. + const PLP_NULL_BYTES: [u8; 8] = [0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF]; /// Index of the last occurrence of `needle` in `haystack`, if any. fn find_last(haystack: &[u8], needle: &[u8]) -> Option { @@ -6185,6 +6311,159 @@ resume_results: VecDeque::new(), )); } + /// A streamed parameter marked NULL (before any chunk) is closed with + /// `PLP_NULL` — no unknown-length opener, no chunks, no terminator. Mirrors + /// msodbcsql's `VARMAX_LENGTH_NULL` for a data-at-execution NULL. + #[tokio::test] + async fn streamed_write_null_frames_plp_null() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + client.write_streamed_null().unwrap(); + let status = client.end_streamed_param().await.unwrap(); + assert!(matches!(status, StreamedParamStatus::Done)); + + let payload = reassemble_sent(&sent.lock().unwrap()); + // The streamed @v value's length field is PLP_NULL and nothing follows it + // (no chunks, no terminator), so it is the final value on the wire. Note + // the positional @statement/@params args are themselves nvarchar(max) and + // legitimately use PLP_UNKNOWN_LEN openers, so we assert on @v's PLP_NULL + // being last rather than the absence of any opener in the whole payload. + assert!( + payload.ends_with(&PLP_NULL_BYTES), + "a NULL streamed value must end the message with PLP_NULL and no terminator" + ); + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Idle + )); + } + + /// Writing a chunk after marking the parameter NULL is a usage error, and the + /// stream stays usable (the message is not corrupted — nothing was written). + #[tokio::test] + async fn streamed_write_chunk_after_null_errors() { + let (mut client, _sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + client.write_streamed_null().unwrap(); + let err = client + .write_streamed_chunk(&[0x01]) + .await + .expect_err("cannot write data after NULL"); + assert!(matches!(err, UsageError(_))); + // Still active: end can still close it as NULL. + assert!(matches!( + client.streamed_write_state, + StreamedWriteState::Active(_) + )); + assert!(matches!( + client.end_streamed_param().await.unwrap(), + StreamedParamStatus::Done + )); + } + + /// Marking a parameter NULL after chunks were already written is a usage + /// error: a value that has begun streaming cannot become NULL. + #[tokio::test] + async fn streamed_null_after_chunk_errors() { + let (mut client, _sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(v) VALUES (@v)".to_string(), + vec![streamed_varbinary("@v")], + None, + None, + ) + .await + .unwrap(); + + client.write_streamed_chunk(&[0x01, 0x02]).await.unwrap(); + let err = client + .write_streamed_null() + .expect_err("cannot mark NULL after chunks were written"); + assert!(matches!(err, UsageError(_))); + } + + /// `write_streamed_null` with no active streamed parameter is a usage error. + #[tokio::test] + async fn streamed_null_without_active_stream_errors() { + let (mut client, _sent) = create_capturing_client(vec![done_no_more()]); + let err = client + .write_streamed_null() + .expect_err("no active streamed parameter"); + assert!(matches!(err, UsageError(_))); + } + + /// A NULL streamed parameter followed by a normal streamed parameter: the + /// first closes with `PLP_NULL`, the second frames its value normally, and the + /// lifecycle advances NeedData -> NeedData -> Done. + #[tokio::test] + async fn streamed_write_null_then_value_param() { + let (mut client, sent) = create_capturing_client(vec![done_no_more()]); + + client + .begin_sp_executesql( + "INSERT INTO t(a, b) VALUES (@a, @b)".to_string(), + vec![streamed_varbinary("@a"), streamed_varbinary("@b")], + None, + None, + ) + .await + .unwrap(); + + // @a is NULL. + client.write_streamed_null().unwrap(); + let status = client.end_streamed_param().await.unwrap(); + assert!( + matches!(&status, StreamedParamStatus::NeedData { param_name } if param_name == "@b") + ); + + // @b carries a value. + let value = [0x7Au8; 5]; + client.write_streamed_chunk(&value).await.unwrap(); + assert!(matches!( + client.end_streamed_param().await.unwrap(), + StreamedParamStatus::Done + )); + + let payload = reassemble_sent(&sent.lock().unwrap()); + // @b's value is framed after an unknown-length opener; @a contributed a + // PLP_NULL and no opener. + let b_pos = find_last(&payload, &PLP_UNKNOWN_LEN_BYTES).unwrap(); + let mut expected_b = Vec::new(); + expected_b.extend_from_slice(&(value.len() as u32).to_le_bytes()); + expected_b.extend_from_slice(&value); + expected_b.extend_from_slice(&PLP_TERMINATOR_BYTES); + assert_eq!( + &payload[b_pos + PLP_UNKNOWN_LEN_BYTES.len()..], + expected_b.as_slice() + ); + assert!( + find_last(&payload[..b_pos], &PLP_NULL_BYTES).is_some(), + "@a must have emitted PLP_NULL before @b's value" + ); + } + /// Multiple chunks are each length-prefixed independently and the value is /// closed by exactly one terminator — the incremental multi-chunk case. #[tokio::test] diff --git a/mssql-tds/src/message/parameters/rpc_parameters.rs b/mssql-tds/src/message/parameters/rpc_parameters.rs index fe317d74..6f2ecee0 100644 --- a/mssql-tds/src/message/parameters/rpc_parameters.rs +++ b/mssql-tds/src/message/parameters/rpc_parameters.rs @@ -7,7 +7,6 @@ use crate::datatypes::column_values::DEFAULT_VARTIME_SCALE; use crate::datatypes::encoder::SqlValueEncoder; use crate::datatypes::sql_tvp::TvpTypeName; use crate::datatypes::sqltypes::SqlType; -use crate::datatypes::tds_value_serializer::PLP_UNKNOWN_LEN; use crate::{ core::TdsResult, datatypes::sqldatatypes::TdsDataType, @@ -118,10 +117,11 @@ pub struct RpcParameter { /// When `true`, this parameter's value is supplied later, in chunks, via the /// data-at-execution path (ODBC `SQL_DATA_AT_EXEC`). During serialization the /// `value` is treated purely as a type template: [`serialize`](Self::serialize) - /// writes the parameter header and opens an unknown-length PLP value, then - /// stops. The value chunks and PLP terminator are written afterwards by the - /// streaming driver. Only the MAX (PLP) types are eligible. Never sent on the - /// wire as a flag. + /// writes the parameter header (status byte + `TYPE_INFO`) and stops *before* + /// the PLP length field. The length field (the unknown-length opener, or + /// `PLP_NULL`), the value chunks and the terminator are written afterwards by + /// the streaming driver. Only the MAX (PLP) types are eligible. Never sent on + /// the wire as a flag. data_at_exec: bool, } @@ -327,11 +327,14 @@ impl RpcParameter { } // Data-at-execution: the value is streamed later in chunks. Reuse the - // exact opening the atomic PLP path emits — status byte, TYPE_INFO, then - // the unknown-total-length sentinel — and stop. The value chunks and PLP - // terminator are written afterwards by the streaming driver. This is the - // write analogue of the incremental read's pause point: the same - // serialize method, parked partway through the value. + // exact opening the atomic PLP path emits — status byte and TYPE_INFO — + // and stop *before* the PLP length field. The length field (the + // unknown-length opener `PLP_UNKNOWN_LEN`, or `PLP_NULL`), the value + // chunks and the terminator are written afterwards by the streaming + // driver. Deferring the length field is what lets a streamed parameter + // still resolve to NULL before any data is sent. This is the write + // analogue of the incremental read's pause point: the same serialize + // method, parked partway through the value. if self.data_at_exec { if !self.is_streamable_plp() { return Err(Error::UsageError(format!( @@ -349,7 +352,6 @@ impl RpcParameter { self.value .write_type_info(packet_writer, db_collation, None, None) .await?; - packet_writer.write_u64_async(PLP_UNKNOWN_LEN).await?; return Ok(()); } @@ -930,10 +932,10 @@ mod tests { payload(&w) } - /// A named data-at-execution `nvarchar(max)` param serializes to: name - /// prefix, status-flags byte, the value's TYPE_INFO, then the 8-byte - /// `PLP_UNKNOWN_LEN` sentinel that opens the value. No value bytes or - /// terminator are written — those are streamed later. + /// A named data-at-execution `nvarchar(max)` param serializes to just the + /// header: name prefix, status-flags byte, then the value's TYPE_INFO. The + /// PLP length field (opener or NULL), value bytes and terminator are all + /// written later by the streaming driver, not here. #[test] fn serialize_data_at_exec_named() { let param = RpcParameter::new( @@ -946,7 +948,6 @@ mod tests { let mut expected = vec![0x02, 0x40, 0x00, 0x70, 0x00]; // name: len 2, "@p" UTF-16LE expected.push(StatusFlags::NONE.bits()); // status flags expected.extend_from_slice(&type_info_bytes(&SqlType::NVarcharMax(None))); // TYPE_INFO - expected.extend_from_slice(&0xFFFF_FFFF_FFFF_FFFEu64.to_le_bytes()); // PLP_UNKNOWN_LEN assert_eq!(streamed_header_bytes(¶m, false), expected); } @@ -966,13 +967,12 @@ mod tests { let mut expected = vec![0x02, 0x40, 0x00, 0x70, 0x00]; // name: len 2, "@p" UTF-16LE expected.push(StatusFlags::NONE.bits()); // status flags expected.extend_from_slice(&type_info_bytes(&SqlType::VarcharMax(None))); // TYPE_INFO - expected.extend_from_slice(&0xFFFF_FFFF_FFFF_FFFEu64.to_le_bytes()); // PLP_UNKNOWN_LEN assert_eq!(streamed_header_bytes(¶m, false), expected); } /// A positional data-at-execution param writes a zero-length name byte in - /// place of the name, then the same status/TYPE_INFO/PLP_UNKNOWN_LEN sequence. + /// place of the name, then the same status/TYPE_INFO header (no length field). #[test] fn serialize_data_at_exec_positional() { let param = @@ -981,7 +981,6 @@ mod tests { let mut expected = vec![0x00]; // zero-length name (positional) expected.push(StatusFlags::NONE.bits()); expected.extend_from_slice(&type_info_bytes(&SqlType::VarBinaryMax(None))); - expected.extend_from_slice(&0xFFFF_FFFF_FFFF_FFFEu64.to_le_bytes()); assert_eq!(streamed_header_bytes(¶m, true), expected); } diff --git a/mssql-tds/tests/test_client_write_apis.rs b/mssql-tds/tests/test_client_write_apis.rs index 92785f8f..ff8e5565 100644 --- a/mssql-tds/tests/test_client_write_apis.rs +++ b/mssql-tds/tests/test_client_write_apis.rs @@ -444,10 +444,65 @@ mod streamed_plp_write { Ok(()) } - /// Streaming zero chunks then ending yields an empty (non-NULL) value: the - /// value opener followed immediately by the terminator encodes a present, - /// zero-length value. Distinct from the NULL path - /// (`write_null_max_round_trips`). + /// A data-at-execution parameter that resolves to NULL: `begin` returns + /// `NeedData`, the caller signals NULL via `write_streamed_null` instead of + /// streaming chunks, and `end` closes it with `PLP_NULL`. Round-trips as SQL + /// NULL, and is distinct from the empty-value path. Mirrors msodbcsql's + /// `SQLPutData(SQL_NULL_DATA)` on a DAE-bound parameter. + #[tokio::test] + async fn stream_null_value_round_trips() -> mssql_tds::core::TdsResult<()> { + init_tracing(); + let context = create_context(); + let provider = TdsConnectionProvider {}; + let mut client = provider + .create_client(context, &build_tcp_datasource(), None) + .await?; + + client + .execute( + "CREATE TABLE #plp_snull (id INT, val NVARCHAR(MAX))".to_string(), + (), + ) + .await?; + client.close_query().await?; + + let streamed = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(); + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_snull (id, val) VALUES (1, @v)".to_string(), + vec![streamed], + None, + None, + ) + .await?; + assert!(matches!(status, StreamedParamStatus::NeedData { .. })); + + // Signal NULL instead of streaming chunks. + client.write_streamed_null()?; + let status = client.end_streamed_param().await?; + assert!(matches!(status, StreamedParamStatus::Done)); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_snull WHERE id = 1".to_string(), ()) + .await?; + { + let row = client.next_row().await?.expect("expected a row"); + assert!( + matches!(&row[0], ColumnValues::Null), + "expected SQL NULL from a streamed-NULL parameter, got {:?}", + &row[0] + ); + } + client.close_query().await?; + Ok(()) + } #[tokio::test] async fn stream_empty_value_round_trips() -> mssql_tds::core::TdsResult<()> { init_tracing();