Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 66 additions & 8 deletions mssql-mock-tds/src/protocol.rs
Original file line number Diff line number Diff line change
Expand Up @@ -853,20 +853,78 @@ pub fn build_query_result(response: &crate::query_response::QueryResponse) -> By
result.extend_from_slice(&build_info_token(info));
}

// Serialize each row
for row in &response.rows {
result.put_u8(TokenType::Row as u8);
for value in &row.values {
value.write_to_buffer(&mut result);
match &response.error_after {
Some(err) => {
// Stream the first `after_rows` rows, then an ERROR token and a
// terminal DONE (row count 0) — no trailing rows.
for row in response.rows.iter().take(err.after_rows) {
result.put_u8(TokenType::Row as u8);
for value in &row.values {
value.write_to_buffer(&mut result);
}
}
result.extend_from_slice(&build_error_token(
err.number,
err.state,
err.severity,
&err.message,
));
// INFO tokens emitted during the drain (after ERROR, before DONE).
for info in &err.drain_info {
result.extend_from_slice(&build_info_token(info));
}
result.extend_from_slice(&build_done_token(0));
}
}
None => {
// Serialize each row
for row in &response.rows {
result.put_u8(TokenType::Row as u8);
for value in &row.values {
value.write_to_buffer(&mut result);
}
}

// DONE token
result.extend_from_slice(&build_done_token(response.rows.len() as u64));
// DONE token
result.extend_from_slice(&build_done_token(response.rows.len() as u64));
}
}

wrap_in_packet(PacketType::TabularResult, result)
}

/// Build a bare ERROR token (0xAA) with no surrounding DONE or packet framing,
/// for injecting mid-stream into a result set.
fn build_error_token(number: u32, state: u8, severity: u8, message: &str) -> BytesMut {
let mut token = BytesMut::new();
token.put_u8(TokenType::Error as u8);

let length_pos = token.len();
token.put_u16_le(0); // Placeholder for length (little-endian on the wire)

token.put_u32_le(number);
token.put_u8(state);
token.put_u8(severity);

// Message (US_VARCHAR: u16 code-unit count + UTF-16LE)
token.put_u16_le(message.chars().count() as u16);
for ch in message.encode_utf16() {
token.put_u16_le(ch);
}

// Server name / procedure name (empty B_VARCHARs)
token.put_u8(0);
token.put_u8(0);

// Line number
token.put_u32_le(1);

let token_length = (token.len() - length_pos - 2) as u16;
let mut length_bytes = &mut token[length_pos..length_pos + 2];
length_bytes.put_u16_le(token_length);

token
}

/// Build an error response
pub fn build_error_response(message: &str) -> BytesMut {
let mut response = BytesMut::new();
Expand Down
28 changes: 28 additions & 0 deletions mssql-mock-tds/src/query_response.rs
Original file line number Diff line number Diff line change
Expand Up @@ -162,12 +162,31 @@ impl InfoMessage {
}
}

/// A server ERROR token injected partway through a result set, after
/// `after_rows` rows have been streamed, followed by a terminal DONE. Used to
/// exercise the fetch-time error/drain path.
#[derive(Debug, Clone)]
pub struct MidStreamError {
pub after_rows: usize,
pub number: u32,
pub state: u8,
pub severity: u8,
pub message: String,
/// INFO tokens emitted after the ERROR and before the terminal DONE, so the
/// fetch-time drain path (async `drain_stream` / sync blocking drain) is
/// exercised on Info capture, not just the happy pre-error stream.
pub drain_info: Vec<InfoMessage>,
}

/// A complete query response definition
#[derive(Debug, Clone)]
pub struct QueryResponse {
pub columns: Vec<ColumnDefinition>,
pub rows: Vec<Row>,
pub info_tokens: Vec<InfoMessage>,
/// When set, only the first `after_rows` rows are streamed, then an ERROR
/// token and a terminal DONE are emitted (no trailing rows).
pub error_after: Option<MidStreamError>,
}

impl QueryResponse {
Expand All @@ -177,6 +196,7 @@ impl QueryResponse {
columns,
rows,
info_tokens: Vec::new(),
error_after: None,
}
}

Expand All @@ -185,12 +205,19 @@ impl QueryResponse {
self
}

/// Inject a mid-stream ERROR token after `after_rows` rows.
pub fn with_error_after(mut self, error_after: MidStreamError) -> Self {
self.error_after = Some(error_after);
self
}

/// Helper to create a response for SELECT 1
pub fn select_one() -> Self {
Self {
columns: vec![ColumnDefinition::new("", SqlDataType::Int)],
rows: vec![Row::new(vec![ColumnValue::Int(1)])],
info_tokens: Vec::new(),
error_after: None,
}
}

Expand All @@ -208,6 +235,7 @@ impl QueryResponse {
ColumnValue::Int(3),
])],
info_tokens: Vec::new(),
error_after: None,
}
}
}
Expand Down
2 changes: 2 additions & 0 deletions mssql-tds/src/connection.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,5 +23,7 @@ pub(crate) mod metadata_retriever;
pub(crate) mod session_recovery;
/// Primary client type and result set traits.
pub mod tds_client;
/// Synchronous, reactor-free row-fetch client over the blocking TDS edge.
pub mod tds_sync_client;
/// Transport layer (TCP, Named Pipes, Shared Memory).
pub mod transport;
Loading