use crate::connection::client_context::ClientContext;
use crate::connection::execution_context::ExecutionContext;
use crate::core::{NegotiatedEncryptionSetting, TdsResult, Version};
use crate::message::login_options::TdsVersion;
use crate::token::tokens::{SessionStateToken, SqlCollation};
pub(crate) struct RecoveryContext {
pub session_recovery_negotiated: bool,
pub session_state_table: Box<SessionStateTable>,
pub client_context: Option<Box<ClientContext>>,
pub original_tds_version: Option<TdsVersion>,
pub original_server_version: Option<Version>,
pub original_encryption_level: Option<NegotiatedEncryptionSetting>,
pub original_mars_enabled: bool,
pub recovery_count: u32,
}
impl std::fmt::Debug for RecoveryContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RecoveryContext")
.field(
"session_recovery_negotiated",
&self.session_recovery_negotiated,
)
.field("session_state_table", &self.session_state_table)
.field(
"client_context",
&self.client_context.as_ref().map(|_| "<ClientContext>"),
)
.field("original_tds_version", &self.original_tds_version)
.field("original_server_version", &self.original_server_version)
.field("original_encryption_level", &self.original_encryption_level)
.field("original_mars_enabled", &self.original_mars_enabled)
.field("recovery_count", &self.recovery_count)
.finish()
}
}
impl RecoveryContext {
pub fn new() -> Self {
Self {
session_recovery_negotiated: false,
session_state_table: Box::new(SessionStateTable::new()),
client_context: None,
original_tds_version: None,
original_server_version: None,
original_encryption_level: None,
original_mars_enabled: false,
recovery_count: 0,
}
}
pub fn initialize(
&mut self,
client_context: ClientContext,
negotiated_settings: &crate::handler::handler_factory::NegotiatedSettings,
login_session_state_tokens: &[SessionStateToken],
) {
self.client_context = Some(Box::new(client_context));
self.original_tds_version = negotiated_settings.login_ack_tds_version;
self.original_server_version = negotiated_settings.login_ack_server_version;
self.original_encryption_level = Some(
negotiated_settings
.session_settings
.negotiated_encryption_settings,
);
self.original_mars_enabled = negotiated_settings.session_settings.mars_enabled;
self.session_recovery_negotiated = negotiated_settings.is_session_recovery_acknowledged();
self.session_state_table.seed_initial_state(
negotiated_settings.database.clone(),
negotiated_settings.language.clone(),
negotiated_settings.database_collation,
login_session_state_tokens,
negotiated_settings.session_recovery_initial_state(),
);
}
pub fn is_recovery_possible(&self, execution_context: &ExecutionContext) -> bool {
self.session_recovery_negotiated
&& self.session_state_table.is_session_recoverable()
&& !execution_context.has_open_batch()
&& !execution_context.has_active_transaction()
}
pub fn validate_reconnection(
&self,
new_settings: &crate::handler::handler_factory::NegotiatedSettings,
) -> TdsResult<()> {
use crate::error::Error;
if !new_settings.is_session_recovery_acknowledged() {
return Err(Error::ReconnectionValidationFailed(
"Server did not acknowledge session recovery on reconnection".to_string(),
));
}
if self.original_tds_version != new_settings.login_ack_tds_version {
return Err(Error::ReconnectionValidationFailed(format!(
"TDS version mismatch: original {:?}, reconnected {:?}",
self.original_tds_version, new_settings.login_ack_tds_version
)));
}
let major_versions_match = match (
self.original_server_version,
new_settings.login_ack_server_version,
) {
(Some(orig), Some(new_ver)) => orig.major == new_ver.major,
(None, None) => true,
_ => false,
};
if !major_versions_match {
return Err(Error::ReconnectionValidationFailed(format!(
"Server major version mismatch: original {:?}, reconnected {:?}",
self.original_server_version.map(|v| v.major),
new_settings.login_ack_server_version.map(|v| v.major)
)));
}
if self.original_encryption_level
!= Some(new_settings.session_settings.negotiated_encryption_settings)
{
return Err(Error::ReconnectionValidationFailed(format!(
"Encryption level mismatch: original {:?}, reconnected {:?}",
self.original_encryption_level,
new_settings.session_settings.negotiated_encryption_settings
)));
}
if self.original_mars_enabled != new_settings.session_settings.mars_enabled {
return Err(Error::ReconnectionValidationFailed(format!(
"MARS setting mismatch: original {}, reconnected {}",
self.original_mars_enabled, new_settings.session_settings.mars_enabled
)));
}
Ok(())
}
pub fn process_session_state(&mut self, token: &SessionStateToken) -> TdsResult<()> {
if token.sequence_number == u32::MAX {
self.session_state_table.master_recovery_disabled = true;
return Ok(());
}
for entry in &token.states {
self.session_state_table.update_state(
entry.state_id,
token.sequence_number,
entry.recoverable,
entry.data.clone(),
);
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub(crate) struct SessionStateRecord {
pub recoverable: bool,
#[allow(dead_code)] pub sequence: u32,
pub data: Vec<u8>,
}
#[derive(Debug)]
pub(crate) struct SessionStateTable {
pub initial_state: [Option<SessionStateRecord>; 256],
pub delta: [Option<SessionStateRecord>; 256],
pub initial_database: String,
pub initial_language: String,
pub initial_collation: SqlCollation,
pub master_recovery_disabled: bool,
}
impl Default for SessionStateTable {
fn default() -> Self {
Self {
initial_state: std::array::from_fn(|_| None),
delta: std::array::from_fn(|_| None),
initial_database: String::new(),
initial_language: String::new(),
initial_collation: SqlCollation::default(),
master_recovery_disabled: false,
}
}
}
type FeatureAckEntry = (u8, Vec<u8>);
type FeatureAckParseError = (usize, u8);
impl SessionStateTable {
pub fn new() -> Self {
Self::default()
}
pub fn update_state(&mut self, state_id: u8, sequence: u32, recoverable: bool, data: Vec<u8>) {
self.delta[state_id as usize] = Some(SessionStateRecord {
recoverable,
sequence,
data,
});
}
pub fn seed_initial_state(
&mut self,
database: String,
language: String,
collation: SqlCollation,
tokens: &[SessionStateToken],
feature_ack_initial_state: Option<&[u8]>,
) {
self.initial_database = database;
self.initial_language = language;
self.initial_collation = collation;
if let Some(data) = feature_ack_initial_state {
self.seed_initial_state_from_feature_ack(data);
}
for token in tokens {
if token.sequence_number == u32::MAX {
self.master_recovery_disabled = true;
continue;
}
for entry in &token.states {
self.update_state(
entry.state_id,
token.sequence_number,
entry.recoverable,
entry.data.clone(),
);
}
}
}
fn parse_feature_ack(data: &[u8]) -> Result<Vec<FeatureAckEntry>, FeatureAckParseError> {
let mut entries = Vec::new();
let mut i = 0;
while i < data.len() {
let entry_start = i;
let state_id = data[i];
i += 1;
let Some(&len_byte) = data.get(i) else {
return Err((entry_start, state_id));
};
i += 1;
let len = if len_byte == 0xFF {
let Some(bytes) = i.checked_add(4).and_then(|end| data.get(i..end)) else {
return Err((entry_start, state_id));
};
i += 4;
u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]) as usize
} else {
len_byte as usize
};
let Some(value) = i.checked_add(len).and_then(|end| data.get(i..end)) else {
return Err((entry_start, state_id));
};
i += len;
entries.push((state_id, value.to_vec()));
}
Ok(entries)
}
fn seed_initial_state_from_feature_ack(&mut self, data: &[u8]) {
match Self::parse_feature_ack(data) {
Ok(entries) => {
for (state_id, value) in entries {
self.initial_state[state_id as usize] = Some(SessionStateRecord {
recoverable: true,
sequence: 0,
data: value,
});
}
}
Err((offset, state_id)) => self.abort_partial_baseline(offset, state_id),
}
}
pub(crate) fn apply_feature_ack_to_delta(&mut self, data: &[u8]) {
self.delta = std::array::from_fn(|_| None);
match Self::parse_feature_ack(data) {
Ok(entries) => {
for (state_id, value) in entries {
self.update_state(state_id, 0, true, value);
}
}
Err((offset, state_id)) => self.abort_partial_baseline(offset, state_id),
}
}
fn abort_partial_baseline(&mut self, offset: usize, state_id: u8) {
tracing::warn!(
offset,
state_id,
"Malformed session-recovery FEATUREEXTACK payload; disabling recovery"
);
self.master_recovery_disabled = true;
}
pub fn is_session_recoverable(&self) -> bool {
!self.master_recovery_disabled
&& !self
.initial_state
.iter()
.chain(self.delta.iter())
.flatten()
.any(|record| !record.recoverable)
}
pub fn reset(&mut self) {
self.delta = std::array::from_fn(|_| None);
}
pub fn snapshot(
&self,
current_database: Option<&str>,
current_language: Option<&str>,
current_collation: Option<SqlCollation>,
) -> Box<SessionRecoveryData> {
Box::new(SessionRecoveryData {
initial_state: self.initial_state.clone(),
delta: self.delta.clone(),
initial_database: self.initial_database.clone(),
initial_language: self.initial_language.clone(),
initial_collation: self.initial_collation,
database: current_database
.unwrap_or(&self.initial_database)
.to_string(),
language: current_language
.unwrap_or(&self.initial_language)
.to_string(),
collation: current_collation.unwrap_or(self.initial_collation),
})
}
}
#[derive(Debug, Clone)]
pub(crate) struct SessionRecoveryData {
pub initial_state: [Option<SessionStateRecord>; 256],
pub delta: [Option<SessionStateRecord>; 256],
pub initial_database: String,
pub initial_language: String,
pub initial_collation: SqlCollation,
pub database: String,
pub language: String,
pub collation: SqlCollation,
}
impl SessionRecoveryData {
pub fn initial_block_length(&self) -> u32 {
let mut len: u32 = 0;
len += 1 + 2 * self.initial_database.encode_utf16().count() as u32;
len += if self.initial_collation == SqlCollation::default() {
1 } else {
6 };
len += 1 + 2 * self.initial_language.encode_utf16().count() as u32;
for entry in &self.initial_state {
if let Some(record) = entry.as_ref() {
len += 1; len += state_value_wire_length(record.data.len());
}
}
len
}
pub fn delta_block_length(&self) -> u32 {
let mut len: u32 = 0;
if self.database == self.initial_database {
len += 1;
} else {
len += 1 + 2 * self.database.encode_utf16().count() as u32;
}
if self.collation == self.initial_collation {
len += 1;
} else {
len += 6;
}
if self.language == self.initial_language {
len += 1;
} else {
len += 1 + 2 * self.language.encode_utf16().count() as u32;
}
for entry in &self.delta {
if let Some(record) = entry.as_ref() {
len += 1; len += state_value_wire_length(record.data.len());
}
}
len
}
pub fn total_data_length(&self) -> u32 {
8 + self.initial_block_length() + self.delta_block_length()
}
}
fn state_value_wire_length(data_len: usize) -> u32 {
if data_len < 0xFF {
1 + data_len as u32
} else {
5 + data_len as u32
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_table_is_recoverable() {
let table = SessionStateTable::new();
assert!(table.is_session_recoverable());
assert!(!table.master_recovery_disabled);
}
#[test]
fn update_recoverable_state_keeps_session_recoverable() {
let mut table = SessionStateTable::new();
table.update_state(0, 1, true, vec![0x01, 0x02]);
assert!(table.is_session_recoverable());
assert!(table.delta[0].is_some());
let record = table.delta[0].as_ref().unwrap();
assert!(record.recoverable);
assert_eq!(record.sequence, 1);
assert_eq!(record.data, vec![0x01, 0x02]);
}
#[test]
fn update_unrecoverable_state_blocks_recovery() {
let mut table = SessionStateTable::new();
table.update_state(5, 1, false, vec![0xAA]);
assert!(!table.is_session_recoverable());
}
#[test]
fn transition_unrecoverable_to_recoverable_restores_recovery() {
let mut table = SessionStateTable::new();
table.update_state(5, 1, false, vec![0xAA]);
assert!(!table.is_session_recoverable());
table.update_state(5, 2, true, vec![0xBB]);
assert!(table.is_session_recoverable());
let record = table.delta[5].as_ref().unwrap();
assert!(record.recoverable);
assert_eq!(record.sequence, 2);
assert_eq!(record.data, vec![0xBB]);
}
#[test]
fn transition_recoverable_to_unrecoverable_blocks_recovery() {
let mut table = SessionStateTable::new();
table.update_state(10, 1, true, vec![0x01]);
assert!(table.is_session_recoverable());
table.update_state(10, 2, false, vec![0x02]);
assert!(!table.is_session_recoverable());
}
#[test]
fn multiple_unrecoverable_states_require_all_cleared() {
let mut table = SessionStateTable::new();
table.update_state(1, 1, false, vec![0x01]);
table.update_state(2, 1, false, vec![0x02]);
assert!(!table.is_session_recoverable());
table.update_state(1, 2, true, vec![0x03]);
assert!(!table.is_session_recoverable());
table.update_state(2, 2, true, vec![0x04]);
assert!(table.is_session_recoverable());
}
#[test]
fn repeated_update_same_recoverability_no_count_change() {
let mut table = SessionStateTable::new();
table.update_state(0, 1, false, vec![0x01]);
assert!(!table.is_session_recoverable());
table.update_state(0, 2, false, vec![0x02]);
assert!(!table.is_session_recoverable());
table.update_state(0, 3, true, vec![0x03]);
assert!(table.is_session_recoverable());
}
#[test]
fn master_recovery_disabled_blocks_even_with_no_unrecoverable_states() {
let mut table = SessionStateTable::new();
table.master_recovery_disabled = true;
assert!(!table.is_session_recoverable());
}
#[test]
fn reset_clears_delta_and_unrecoverable_count() {
let mut table = SessionStateTable::new();
table.update_state(0, 1, false, vec![0x01]);
table.update_state(5, 1, true, vec![0x02]);
assert!(!table.is_session_recoverable());
table.reset();
assert!(table.is_session_recoverable());
assert!(table.delta[0].is_none());
assert!(table.delta[5].is_none());
}
#[test]
fn seed_from_feature_ack_parses_server_baseline() {
let ack = vec![
0x00, 0x09, 0x00, 0x60, 0x81, 0x14, 0xFF, 0xE7, 0xFF, 0xFF, 0x00, 0x02, 0x02, 0x07, 0x01, 0x04, 0x01, 0x00, 0x05, 0x04, 0xFF, 0xFF, 0xFF, 0xFF, 0x06, 0x01, 0x00, 0x07, 0x01, 0x02, 0x08, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x09, 0x04, 0xFF, 0xFF, 0xFF, 0xFF, ];
let mut table = SessionStateTable::new();
table.seed_initial_state_from_feature_ack(&ack);
for (id, expected) in [
(
0usize,
vec![0x00, 0x60, 0x81, 0x14, 0xFF, 0xE7, 0xFF, 0xFF, 0x00],
),
(2, vec![0x07, 0x01]),
(4, vec![0x00]),
(5, vec![0xFF, 0xFF, 0xFF, 0xFF]),
(6, vec![0x00]),
(7, vec![0x02]),
(8, vec![0x00; 8]),
(9, vec![0xFF, 0xFF, 0xFF, 0xFF]),
] {
let record = table.initial_state[id]
.as_ref()
.unwrap_or_else(|| panic!("state {id} missing"));
assert_eq!(record.data, expected);
assert!(record.recoverable);
assert_eq!(record.sequence, 0);
}
assert!(table.initial_state[1].is_none());
assert!(table.initial_state[3].is_none());
assert!(table.is_session_recoverable());
}
#[test]
fn seed_from_feature_ack_extended_length() {
let mut ack = vec![10u8, 0xFF];
ack.extend_from_slice(&300u32.to_le_bytes());
ack.extend_from_slice(&[0xAB; 300]);
let mut table = SessionStateTable::new();
table.seed_initial_state_from_feature_ack(&ack);
let record = table.initial_state[10].as_ref().unwrap();
assert_eq!(record.data.len(), 300);
assert!(record.data.iter().all(|&b| b == 0xAB));
}
#[test]
fn seed_from_feature_ack_truncated_disables_recovery() {
let ack = vec![3, 1, 0xAA, 5, 10, 0x01];
let mut table = SessionStateTable::new();
table.seed_initial_state_from_feature_ack(&ack);
assert!(table.initial_state[3].is_none());
assert!(table.initial_state[5].is_none());
assert!(table.master_recovery_disabled);
assert!(!table.is_session_recoverable());
}
#[test]
fn apply_feature_ack_to_delta_resets_then_reloads() {
let ack = vec![7, 2, 0x01, 0x02];
let mut table = SessionStateTable::new();
table.update_state(3, 5, true, vec![0xAA]);
table.apply_feature_ack_to_delta(&ack);
assert!(table.delta[3].is_none());
assert!(table.initial_state[7].is_none());
let record = table.delta[7].as_ref().unwrap();
assert_eq!(record.data, vec![0x01, 0x02]);
assert!(record.recoverable);
assert_eq!(record.sequence, 0);
}
#[test]
fn seed_initial_state_forwards_feature_ack() {
let ack = vec![7, 2, 0x01, 0x02];
let mut table = SessionStateTable::new();
table.seed_initial_state(
"master".to_string(),
"us_english".to_string(),
SqlCollation::default(),
&[],
Some(&ack),
);
assert_eq!(table.initial_database, "master");
assert_eq!(
table.initial_state[7].as_ref().unwrap().data,
vec![0x01, 0x02]
);
}
#[test]
fn seed_initial_state_from_tokens_tracks_recoverability() {
use crate::token::tokens::SessionStateEntry;
let mut table = SessionStateTable::new();
table.seed_initial_state(
"master".to_string(),
"us_english".to_string(),
SqlCollation::default(),
&[SessionStateToken {
sequence_number: 1,
status: 0,
states: vec![SessionStateEntry {
state_id: 2,
recoverable: false,
data: vec![0x01],
}],
}],
None,
);
assert!(!table.is_session_recoverable());
let mut table = SessionStateTable::new();
table.seed_initial_state(
String::new(),
String::new(),
SqlCollation::default(),
&[SessionStateToken {
sequence_number: u32::MAX,
status: 0,
states: vec![],
}],
None,
);
assert!(table.master_recovery_disabled);
assert!(!table.is_session_recoverable());
}
#[test]
fn reset_preserves_master_recovery_disabled() {
let mut table = SessionStateTable::new();
table.master_recovery_disabled = true;
table.reset();
assert!(table.master_recovery_disabled);
assert!(!table.is_session_recoverable());
}
#[test]
fn reset_preserves_initial_state() {
let mut table = SessionStateTable::new();
table.initial_state[0] = Some(SessionStateRecord {
recoverable: true,
sequence: 1,
data: vec![0x01],
});
table.initial_database = "testdb".to_string();
table.reset();
assert!(table.initial_state[0].is_some());
assert_eq!(table.initial_database, "testdb");
}
#[test]
fn snapshot_creates_independent_copy() {
let mut table = SessionStateTable::new();
table.initial_database = "mydb".to_string();
table.initial_language = "us_english".to_string();
table.update_state(0, 1, true, vec![0x01]);
let snapshot = table.snapshot(None, None, None);
assert_eq!(snapshot.initial_database, "mydb");
assert_eq!(snapshot.initial_language, "us_english");
assert_eq!(snapshot.database, "mydb");
assert_eq!(snapshot.language, "us_english");
assert!(snapshot.delta[0].is_some());
table.initial_database = "other".to_string();
table.update_state(0, 2, false, vec![0x02]);
assert_eq!(snapshot.initial_database, "mydb");
assert!(snapshot.delta[0].as_ref().unwrap().recoverable);
}
#[test]
fn all_256_state_ids_can_be_used() {
let mut table = SessionStateTable::new();
for i in 0..=255u8 {
table.update_state(i, 1, true, vec![i]);
}
assert!(table.is_session_recoverable());
assert!(table.delta[0].is_some());
assert!(table.delta[255].is_some());
}
use crate::token::tokens::{SessionStateEntry, SessionStateToken};
fn make_session_state_token(
sequence_number: u32,
states: Vec<SessionStateEntry>,
) -> SessionStateToken {
SessionStateToken {
sequence_number,
status: 0,
states,
}
}
fn make_entry(state_id: u8, recoverable: bool, data: Vec<u8>) -> SessionStateEntry {
SessionStateEntry {
state_id,
recoverable,
data,
}
}
#[test]
fn recovery_context_new_defaults() {
let ctx = RecoveryContext::new();
assert!(!ctx.session_recovery_negotiated);
assert!(ctx.session_state_table.is_session_recoverable());
}
#[test]
fn process_session_state_updates_table() {
let mut ctx = RecoveryContext::new();
let token = make_session_state_token(
1,
vec![
make_entry(0, true, vec![0x01]),
make_entry(5, false, vec![0x02]),
],
);
ctx.process_session_state(&token).unwrap();
assert!(ctx.session_state_table.delta[0].is_some());
assert!(ctx.session_state_table.delta[5].is_some());
assert!(!ctx.session_state_table.is_session_recoverable()); }
#[test]
fn process_session_state_master_disable() {
let mut ctx = RecoveryContext::new();
let token = make_session_state_token(u32::MAX, vec![make_entry(0, true, vec![0x01])]);
ctx.process_session_state(&token).unwrap();
assert!(ctx.session_state_table.master_recovery_disabled);
assert!(ctx.session_state_table.delta[0].is_none());
}
#[test]
fn process_session_state_multiple_tokens_accumulate() {
let mut ctx = RecoveryContext::new();
let token1 = make_session_state_token(1, vec![make_entry(0, true, vec![0x01])]);
ctx.process_session_state(&token1).unwrap();
let token2 = make_session_state_token(2, vec![make_entry(1, true, vec![0x02])]);
ctx.process_session_state(&token2).unwrap();
assert!(ctx.session_state_table.delta[0].is_some());
assert!(ctx.session_state_table.delta[1].is_some());
assert_eq!(
ctx.session_state_table.delta[0].as_ref().unwrap().sequence,
1
);
assert_eq!(
ctx.session_state_table.delta[1].as_ref().unwrap().sequence,
2
);
}
#[test]
fn process_session_state_overwrites_same_state_id() {
let mut ctx = RecoveryContext::new();
let token1 = make_session_state_token(1, vec![make_entry(0, true, vec![0x01])]);
ctx.process_session_state(&token1).unwrap();
let token2 = make_session_state_token(2, vec![make_entry(0, true, vec![0xFF])]);
ctx.process_session_state(&token2).unwrap();
let record = ctx.session_state_table.delta[0].as_ref().unwrap();
assert_eq!(record.sequence, 2);
assert_eq!(record.data, vec![0xFF]);
}
fn make_initialized_context() -> RecoveryContext {
use crate::message::features::session_recovery::SessionRecoveryFeature;
use crate::message::login::Feature;
let mut ctx = RecoveryContext::new();
let client_ctx = crate::connection::client_context::ClientContext::with_data_source(
"tcp:localhost,1433",
);
let mut settings =
crate::handler::handler_factory::create_test_negotiated_settings_internal();
settings.login_ack_tds_version = Some(TdsVersion::V7_4);
settings.login_ack_server_version = Some(Version::new(16, 0, 1000, 0));
settings.session_settings.negotiated_encryption_settings =
NegotiatedEncryptionSetting::Mandatory;
let mut feature = SessionRecoveryFeature::new(1);
feature.set_acknowledged(true);
settings
.session_settings
.supported_features
.push(Box::new(feature));
ctx.initialize(client_ctx, &settings, &[]);
ctx
}
#[test]
fn initialize_captures_all_fields() {
let ctx = make_initialized_context();
assert!(ctx.client_context.is_some());
assert_eq!(ctx.original_tds_version, Some(TdsVersion::V7_4));
assert_eq!(
ctx.original_server_version,
Some(Version::new(16, 0, 1000, 0))
);
assert_eq!(
ctx.original_encryption_level,
Some(NegotiatedEncryptionSetting::Mandatory)
);
assert!(!ctx.original_mars_enabled);
assert_eq!(ctx.recovery_count, 0);
}
#[test]
fn is_recovery_possible_all_conditions_met() {
let ctx = make_initialized_context();
let exec = crate::connection::execution_context::ExecutionContext::new();
assert!(ctx.is_recovery_possible(&exec));
}
#[test]
fn is_recovery_possible_false_when_not_negotiated() {
let mut ctx = make_initialized_context();
ctx.session_recovery_negotiated = false;
let exec = crate::connection::execution_context::ExecutionContext::new();
assert!(!ctx.is_recovery_possible(&exec));
}
#[test]
fn is_recovery_possible_false_when_master_disabled() {
let mut ctx = make_initialized_context();
ctx.session_state_table.master_recovery_disabled = true;
let exec = crate::connection::execution_context::ExecutionContext::new();
assert!(!ctx.is_recovery_possible(&exec));
}
#[test]
fn is_recovery_possible_false_when_unrecoverable_state() {
let mut ctx = make_initialized_context();
ctx.session_state_table
.update_state(5, 1, false, vec![0x01]);
let exec = crate::connection::execution_context::ExecutionContext::new();
assert!(!ctx.is_recovery_possible(&exec));
}
#[test]
fn is_recovery_possible_false_when_batch_open() {
let ctx = make_initialized_context();
let mut exec = crate::connection::execution_context::ExecutionContext::new();
exec.set_has_open_batch(true);
assert!(!ctx.is_recovery_possible(&exec));
}
#[test]
fn is_recovery_possible_false_when_transaction_active() {
let ctx = make_initialized_context();
let mut exec = crate::connection::execution_context::ExecutionContext::new();
exec.set_transaction_descriptor(12345);
assert!(!ctx.is_recovery_possible(&exec));
}
#[test]
fn is_recovery_possible_false_multiple_blockers() {
let mut ctx = make_initialized_context();
ctx.session_state_table.master_recovery_disabled = true;
let mut exec = crate::connection::execution_context::ExecutionContext::new();
exec.set_has_open_batch(true);
exec.set_transaction_descriptor(1);
assert!(!ctx.is_recovery_possible(&exec));
}
#[test]
fn is_recovery_possible_recovers_after_clearing_state() {
let mut ctx = make_initialized_context();
ctx.session_state_table
.update_state(5, 1, false, vec![0x01]);
let exec = crate::connection::execution_context::ExecutionContext::new();
assert!(!ctx.is_recovery_possible(&exec));
ctx.session_state_table.update_state(5, 2, true, vec![0x02]);
assert!(ctx.is_recovery_possible(&exec));
}
#[test]
fn debug_format_hides_client_context() {
let ctx = make_initialized_context();
let debug_str = format!("{:?}", ctx);
assert!(debug_str.contains("<ClientContext>"));
assert!(!debug_str.contains("password"));
}
use crate::handler::handler_factory::{
NegotiatedSettings, create_test_negotiated_settings_internal,
};
use crate::message::features::session_recovery::SessionRecoveryFeature;
use crate::message::login::Feature;
fn make_matching_negotiated_settings() -> NegotiatedSettings {
let mut settings = create_test_negotiated_settings_internal();
settings.login_ack_tds_version = Some(TdsVersion::V7_4);
settings.login_ack_server_version = Some(Version::new(16, 0, 1000, 0));
settings.session_settings.negotiated_encryption_settings =
NegotiatedEncryptionSetting::Mandatory;
settings.session_settings.mars_enabled = false;
let mut feature = SessionRecoveryFeature::new(1);
feature.set_acknowledged(true);
settings
.session_settings
.supported_features
.push(Box::new(feature));
settings
}
#[test]
fn validate_reconnection_all_matching() {
let ctx = make_initialized_context();
let settings = make_matching_negotiated_settings();
assert!(ctx.validate_reconnection(&settings).is_ok());
}
#[test]
fn validate_reconnection_fails_when_session_recovery_not_acknowledged() {
let ctx = make_initialized_context();
let mut settings = make_matching_negotiated_settings();
settings.session_settings.supported_features.clear();
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("did not acknowledge session recovery"));
}
#[test]
fn validate_reconnection_fails_when_feature_present_but_not_acknowledged() {
let ctx = make_initialized_context();
let mut settings = make_matching_negotiated_settings();
settings.session_settings.supported_features.clear();
let feature = SessionRecoveryFeature::new(1); settings
.session_settings
.supported_features
.push(Box::new(feature));
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("did not acknowledge session recovery"));
}
#[test]
fn validate_reconnection_fails_on_tds_version_mismatch() {
let ctx = make_initialized_context();
let mut settings = make_matching_negotiated_settings();
settings.login_ack_tds_version = Some(TdsVersion::V8_0);
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("TDS version mismatch"));
}
#[test]
fn validate_reconnection_fails_on_server_major_version_mismatch() {
let ctx = make_initialized_context();
let mut settings = make_matching_negotiated_settings();
settings.login_ack_server_version = Some(Version::new(15, 0, 2000, 0));
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("Server major version mismatch"));
}
#[test]
fn validate_reconnection_ok_when_server_minor_version_differs() {
let ctx = make_initialized_context();
let mut settings = make_matching_negotiated_settings();
settings.login_ack_server_version = Some(Version::new(16, 5, 9999, 0));
assert!(ctx.validate_reconnection(&settings).is_ok());
}
#[test]
fn validate_reconnection_fails_on_encryption_mismatch() {
let ctx = make_initialized_context();
let mut settings = make_matching_negotiated_settings();
settings.session_settings.negotiated_encryption_settings =
NegotiatedEncryptionSetting::NoEncryption;
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("Encryption level mismatch"));
}
#[test]
fn validate_reconnection_fails_on_mars_mismatch() {
let ctx = make_initialized_context();
let mut settings = make_matching_negotiated_settings();
settings.session_settings.mars_enabled = true;
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("MARS setting mismatch"));
}
#[test]
fn validate_reconnection_fails_when_original_tds_version_is_none() {
let mut ctx = make_initialized_context();
ctx.original_tds_version = None;
let settings = make_matching_negotiated_settings();
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("TDS version mismatch"));
}
#[test]
fn validate_reconnection_fails_when_original_server_version_is_none() {
let mut ctx = make_initialized_context();
ctx.original_server_version = None;
let settings = make_matching_negotiated_settings();
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("Server major version mismatch"));
}
#[test]
fn validate_reconnection_fails_when_original_encryption_is_none() {
let mut ctx = make_initialized_context();
ctx.original_encryption_level = None;
let settings = make_matching_negotiated_settings();
let err = ctx.validate_reconnection(&settings).unwrap_err();
let msg = format!("{}", err);
assert!(msg.contains("Encryption level mismatch"));
}
}