diff --git a/mssql-tds/src/connection/tds_client.rs b/mssql-tds/src/connection/tds_client.rs index 9743dc8c..beed722c 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_NULL, PLP_TERMINATOR, PLP_UNKNOWN_LEN}; 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::{ @@ -184,6 +186,54 @@ pub enum CursorColumn { RowEnded, } +/// Result of beginning or continuing a streamed PLP parameter write. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum StreamedParamStatus { + /// The server still expects a streamed parameter value. Stream its chunks + /// next via [`TdsClient::write_streamed_chunk`], then call + /// [`TdsClient::end_streamed_param`]. + NeedData { + /// Name of the streamed parameter now awaiting its value chunks. + param_name: String, + }, + /// All streamed parameters have been written, and the server response has + /// been positioned exactly as after a normal execute call. + Complete(StatementResult), +} + +/// 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 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. /// /// Created by [`TdsConnectionProvider::create_client()`](crate::connection_provider::tds_connection_provider::TdsConnectionProvider::create_client). @@ -295,6 +345,10 @@ pub struct TdsClient { /// budget-exhaustion paths are exercised without live reconnect timing. #[cfg(test)] reconnect_elapsed_for_test: Option, + + // 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 { @@ -345,6 +399,7 @@ impl TdsClient { active_row_read_state: ActiveRowReadState::Idle, #[cfg(test)] reconnect_elapsed_for_test: None, + streamed_write_state: StreamedWriteState::Idle, } } @@ -883,36 +938,24 @@ impl TdsClient { if self.execution_context.has_open_batch() { return Err(UsageError(ALREADY_EXECUTING_ERROR.to_string())); }; + if named_params.iter().any(RpcParameter::is_data_at_exec) { + return Err(UsageError( + "Data-at-execution parameters require begin_sp_executesql.".to_string(), + )); + } - self.begin_command(); - let reconnect_elapsed = self.check_and_reconnect(timeout_sec, cancel_handle).await?; - let budget = Self::deduct_timeout(timeout_sec, reconnect_elapsed); - let resolved = budget.into_timeout()?; - let timeout_sec = resolved.seconds(); - let request_timeout = resolved.duration(); - - // Store timeout and cancel handle for this operation - self.remaining_request_timeout = request_timeout; - self.cancel_handle = cancel_handle.map(|handle| handle.child_handle()); - - self.transport.reset_reader(); - let database_collation = self.negotiated_settings.database_collation; - - let sql_statement_value = - SqlType::NVarcharMax(Some(SqlString::from_utf8_string(sql.clone()))); - - // Create the parameter list for sp_execute_sql - let statement_parameter = RpcParameter::new(None, StatusFlags::NONE, sql_statement_value); - - // Build the comma separated list of parameters + let declaration_params = named_params.clone(); let mut params_list_as_string = String::new(); - - build_parameter_list_string(&named_params, &mut params_list_as_string)?; + build_parameter_list_string(&declaration_params, &mut params_list_as_string)?; // Always Encrypted: when the connection enabled column encryption and the // server acknowledged the feature, ask the server which parameters need // encryption and encrypt them in place before sending the real RPC. self.ensure_force_column_encryption_supported(named_params.iter())?; + let (timeout_sec, request_timeout, database_collation) = self + .prepare_sp_executesql_command(timeout_sec, cancel_handle) + .await?; + if self.should_encrypt_parameters() && !named_params.is_empty() { self.encrypt_parameters( &sql, @@ -928,32 +971,537 @@ impl TdsClient { self.cancel_handle = cancel_handle.map(|handle| handle.child_handle()); } - let params_as_sql_string = SqlType::NVarcharMax(Some(SqlString::from_utf8_string( - params_list_as_string.clone(), - ))); + let rpc = self.build_sp_executesql_rpc( + sql, + named_params, + params_list_as_string, + &database_collation, + ); - let params_parameter = RpcParameter::new(None, StatusFlags::NONE, params_as_sql_string); + let mut packet_writer = + rpc.create_packet_writer(self.transport.as_writer(), timeout_sec, cancel_handle); + rpc.serialize(&mut packet_writer).await?; - // Create the parameter list for positional parameters of sp_execute_sql. - // These could be named parameters as well, but we want to avoid sending the name - // to send less data over the wire. - let positional_parameters_vec = vec![statement_parameter, params_parameter]; - let positional_parameters = Some(positional_parameters_vec); + self.position_on_first_result().await + } - // Build the RPC request. - let rpc = SqlRpc::new( + /// Starts a parameterized `sp_executesql` whose MAX parameter values are + /// supplied later, in chunks (data-at-execution). + /// + /// This is the additive streaming counterpart to + /// [`TdsClient::execute_sp_executesql`]. Existing materialized-parameter + /// callers should continue using that method; this method is for + /// [`RpcParameter::data_at_exec`] parameters. + /// + /// Supported streamed types are `nvarchar(max)`, `varchar(max)`, and + /// `varbinary(max)`. Chunks are raw wire bytes: callers must encode + /// `nvarchar(max)` chunks as UTF-16LE, while `varchar(max)` and + /// `varbinary(max)` chunks use their corresponding single-byte/binary + /// representation. A zero-length stream is a present empty value; use + /// [`TdsClient::write_streamed_null`] for SQL `NULL`. + /// + /// The method returns [`StreamedParamStatus::NeedData`] after the RPC + /// prefix is parked. Supply zero or more chunks and call + /// [`TdsClient::end_streamed_param`]. When all parameters are complete, + /// the final server result is returned as + /// [`StreamedParamStatus::Complete`]. + /// + /// Streamed parameters must be named and must use one of the supported MAX + /// types. Streaming is not supported when Always Encrypted is active. + /// + /// # Errors + /// Returns a usage error for invalid streamed parameters or an active + /// command, and returns transport/serialization errors from the request. + 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())); + } + + self.current_command_ce_setting = ExecutionColumnEncryptionSetting::UseConnectionSetting; + self.ensure_force_column_encryption_supported(named_params.iter())?; + if self.should_encrypt_parameters() { + return Err(UsageError( + "Streamed PLP parameter writes are not supported with Always Encrypted." + .to_string(), + )); + } + + let (streamed_params, materialized_params) = + Self::split_and_validate_streamed_params(named_params)?; + + if streamed_params.is_empty() { + let result = self + .execute_sp_executesql( + sql, + materialized_params, + ExecuteOptions { + timeout: timeout_sec, + cancel: cancel_handle, + column_encryption: ExecutionColumnEncryptionSetting::UseConnectionSetting, + }, + ) + .await?; + return Ok(StreamedParamStatus::Complete(result)); + } + + let mut declaration_params = materialized_params.clone(); + declaration_params.extend(streamed_params.iter().cloned()); + let mut params_list_as_string = String::new(); + build_parameter_list_string(&declaration_params, &mut params_list_as_string)?; + + let (timeout_sec, _request_timeout, database_collation) = self + .prepare_sp_executesql_command(timeout_sec, cancel_handle) + .await?; + let rpc = self.build_sp_executesql_rpc( + sql, + materialized_params, + params_list_as_string, + &database_collation, + ); + + self.start_sp_executesql_streamed( + rpc, + streamed_params, + timeout_sec, + cancel_handle, + database_collation, + ) + .await + } + + async fn prepare_sp_executesql_command( + &mut self, + timeout_sec: Option, + cancel_handle: Option<&CancelHandle>, + ) -> TdsResult<(Option, Option, SqlCollation)> { + self.begin_command(); + let reconnect_elapsed = self.check_and_reconnect(timeout_sec, cancel_handle).await?; + let budget = Self::deduct_timeout(timeout_sec, reconnect_elapsed); + let resolved = budget.into_timeout()?; + let timeout_sec = resolved.seconds(); + let request_timeout = resolved.duration(); + + self.remaining_request_timeout = request_timeout; + self.cancel_handle = cancel_handle.map(|handle| handle.child_handle()); + self.transport.reset_reader(); + + Ok(( + timeout_sec, + request_timeout, + self.negotiated_settings.database_collation, + )) + } + + fn build_sp_executesql_rpc<'a>( + &self, + sql: String, + named_params: Vec, + params_list_as_string: String, + database_collation: &'a SqlCollation, + ) -> SqlRpc<'a> { + let statement_parameter = RpcParameter::new( + None, + StatusFlags::NONE, + SqlType::NVarcharMax(Some(SqlString::from_utf8_string(sql))), + ); + + 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]); + + SqlRpc::new( RpcType::ProcId(RpcProcs::ExecuteSql), positional_parameters, Some(named_params), - &database_collation, + database_collation, &self.execution_context, - ); + ) + } + + fn split_and_validate_streamed_params( + named_params: Vec, + ) -> TdsResult<(Vec, Vec)> { + let (streamed_params, materialized_params): (Vec<_>, Vec<_>) = named_params + .into_iter() + .partition(RpcParameter::is_data_at_exec); + + 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())); + } + } + + Ok((streamed_params, materialized_params)) + } + + async fn start_sp_executesql_streamed( + &mut self, + rpc: SqlRpc<'_>, + streamed_params: Vec, + timeout_sec: Option, + cancel_handle: Option<&CancelHandle>, + database_collation: SqlCollation, + ) -> TdsResult { + 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); - rpc.serialize(&mut packet_writer).await?; + // 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. + let serialization_result = async { + rpc.serialize_prefix(&mut packet_writer).await?; + first + .serialize( + &mut packet_writer, + &database_collation, + false, + &GenericEncoder::new(), + ) + .await + } + .await; + if let Err(error) = serialization_result { + drop(packet_writer); + self.abort_streamed_write(); + return Err(error); + } + let message = packet_writer.suspend(); - self.position_on_first_result().await + self.streamed_write_state = StreamedWriteState::Active(Box::new(StreamedWriteContext { + message, + pending, + db_collation: database_collation, + value_opened: false, + null_signaled: false, + })); + + 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 + /// a [`StreamedParamStatus::NeedData`] from + /// [`begin_sp_executesql`](Self::begin_sp_executesql) (or + /// [`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, 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(()); + } + 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, + 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 + } + .await; + let message = packet_writer.suspend(); + + 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 + // 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(()) + } + 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); + } + + /// 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). + /// When the last streamed parameter closes, the RPC message is finalized and + /// sent, then the real server result is returned in + /// [`StreamedParamStatus::Complete`]. + /// + /// # 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, + value_opened, + null_signaled, + } = *ctx; + + let mut packet_writer = PacketWriter::resume(message, self.transport.as_writer()); + + // 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 { + 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 + .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; + + 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, + // 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, + }) + } + // 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(result) => Ok(StreamedParamStatus::Complete(result)), + 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. @@ -4756,6 +5304,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 { @@ -4768,6 +5319,7 @@ mod tests { packet_data: Vec::new(), packet_pos: 0, resume_results: VecDeque::new(), + send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4780,6 +5332,7 @@ mod tests { packet_data: Vec::new(), packet_pos: 0, resume_results: VecDeque::new(), + send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4792,6 +5345,7 @@ mod tests { packet_data, packet_pos: 0, resume_results: VecDeque::new(), + send_should_fail: Arc::new(std::sync::atomic::AtomicBool::new(false)), } } @@ -4893,6 +5447,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(()) } @@ -5126,6 +5688,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, @@ -7761,4 +8344,765 @@ mod tests { "the RETURNVALUE must be surfaced as an output parameter" ); } + + // ── 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]; + /// 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 { + 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::Complete( + StatementResult::NoRows { .. } | StatementResult::Rows | StatementResult::End + ) + )); + + 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 + )); + } + + /// 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::Complete( + StatementResult::NoRows { .. } | StatementResult::Rows | StatementResult::End + ) + )); + + 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::Complete( + StatementResult::NoRows { .. } | StatementResult::Rows | StatementResult::End + ) + )); + } + + /// 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::Complete( + StatementResult::NoRows { .. } | StatementResult::Rows | StatementResult::End + ) + )); + + 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] + 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::Complete( + StatementResult::NoRows { .. } | StatementResult::Rows | StatementResult::End + ) + )); + + // @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` delegates to + /// the atomic path: it sends the RPC, consumes the response, and returns + /// Complete without parking any streamed state. + #[tokio::test] + async fn begin_without_data_at_exec_delegates_and_returns_complete() { + 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::Complete( + StatementResult::NoRows { .. } | StatementResult::Rows | StatementResult::End + ) + )); + 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() { + 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::Complete( + StatementResult::NoRows { .. } | StatementResult::Rows | StatementResult::End + ) + )); + + 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/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..6f2ecee0 100644 --- a/mssql-tds/src/message/parameters/rpc_parameters.rs +++ b/mssql-tds/src/message/parameters/rpc_parameters.rs @@ -113,6 +113,16 @@ 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 (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, } 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,35 @@ impl RpcParameter { } } + // Data-at-execution: the value is streamed later in chunks. Reuse the + // 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!( + "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?; + 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 +373,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 +919,110 @@ 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 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( + 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 + + 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 + + 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 header (no length field). + #[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))); + + 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..73db6a56 --- /dev/null +++ b/mssql-tds/tests/test_client_write_apis.rs @@ -0,0 +1,798 @@ +// 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, StatementResult, 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(), + (), + ) + .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"), + other => panic!("expected NeedData for the first streamed param, got {other:?}"), + } + + // 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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_nvm WHERE id = 1".to_string(), ()) + .await?; + { + let row = client.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:?}"), + } + } + 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(), + (), + ) + .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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_mix WHERE id = 7".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(()) + } + + /// 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(), + (), + ) + .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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + client.close_query().await?; + + client + .execute("SELECT val FROM #plp_vbm WHERE id = 1".to_string(), ()) + .await?; + { + 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:?}"), + } + } + 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(), + (), + ) + .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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + client.close_query().await?; + + client + .execute( + "SELECT LEN(a), LEN(b) FROM #plp_two WHERE id = 1".to_string(), + (), + ) + .await?; + { + let row = client.next_row().await?.expect("expected a row"); + assert_eq!(row.len(), 2); + } + client.close_query().await?; + Ok(()) + } + + /// Streams two PLP parameters with a materialized integer parameter between + /// them, verifying that streamed parameters resume in the original RPC order. + #[tokio::test] + async fn stream_plp_int_plp_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_int_plp (a NVARCHAR(MAX), id INT, b NVARCHAR(MAX))".to_string(), + (), + ) + .await?; + client.close_query().await?; + + let a = "A".repeat(8_000); + let b = "B".repeat(6_000); + let params = vec![ + RpcParameter::new( + Some("@a".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(), + RpcParameter::new( + Some("@id".to_string()), + StatusFlags::NONE, + SqlType::Int(Some(42)), + ), + RpcParameter::new( + Some("@b".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ) + .data_at_exec(), + ]; + + let status = client + .begin_sp_executesql( + "INSERT INTO #plp_int_plp (a, id, b) VALUES (@a, @id, @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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + client.close_query().await?; + + client + .execute("SELECT a, id, b FROM #plp_int_plp".to_string(), ()) + .await?; + { + let row = client.next_row().await?.expect("expected a row"); + assert_eq!(row.len(), 3); + match &row[0] { + ColumnValues::String(value) => assert_eq!(value.to_utf8_string(), a), + other => panic!("Expected String for column a, got {other:?}"), + } + match &row[1] { + ColumnValues::Int(value) => assert_eq!(*value, 42), + other => panic!("Expected I32 for column id, got {other:?}"), + } + match &row[2] { + ColumnValues::String(value) => assert_eq!(value.to_utf8_string(), b), + other => panic!("Expected String for column b, got {other:?}"), + } + } + 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(), + (), + ) + .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::Complete( + StatementResult::NoRows { .. } | StatementResult::End + ) + )); + client.close_query().await?; + } + + // Every streamed row must be present. + client + .execute("SELECT COUNT(*) FROM #plp_rows".to_string(), ()) + .await?; + { + 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:?}"), + } + } + 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}"), + (), + ) + .await?; + { + 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) + ); + } + other => panic!("Expected String for varchar(max), got {other:?}"), + } + } + 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(), + (), + ) + .await?; + client.close_query().await?; + + // A NULL max parameter is materialized (value None -> PLP_NULL), so + // execute_sp_executesql completes atomically with no NeedData. + let null_param = RpcParameter::new( + Some("@v".to_string()), + StatusFlags::NONE, + SqlType::NVarcharMax(None), + ); + + let status = client + .execute_sp_executesql( + "INSERT INTO #plp_null (id, val) VALUES (1, @v)".to_string(), + vec![null_param], + (), + ) + .await?; + assert!( + matches!( + status, + StatementResult::NoRows { .. } | StatementResult::End + ), + "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(), ()) + .await?; + { + let row = client.next_row().await?.expect("expected a row"); + assert!( + matches!(&row[0], ColumnValues::Null), + "expected SQL NULL, got {:?}", + &row[0] + ); + } + client.close_query().await?; + Ok(()) + } + + /// 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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + 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(); + 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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + 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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + 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::Complete(StatementResult::NoRows { .. } | StatementResult::End) + )); + 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(()) + } +}