use crate::connection::client_context::{ClientContext, TransportContext};
use crate::connection::session_recovery::SessionRecoveryData;
use crate::core::{EncryptionSetting, NegotiatedEncryptionSetting, TdsResult, Version};
use crate::error::{Error, SqlInfoMessage, SqlServerDiagnostics};
use crate::handler::sspi_handler::SspiAuthHandler;
use crate::io::packet_reader::TdsPacketReader;
use crate::io::reader_writer::NetworkReaderWriter;
use crate::io::token_stream::TdsTokenStreamReader;
use crate::message::login::{
EnvChangeProperties, Feature, FeatureExtension, FeaturesRequest, FedAuthTokenRequest,
LoginRequest, LoginRequestModel, LoginResponse, LoginResponseModel, LoginResponseStatus,
SspiRequest,
};
use crate::message::login_options::TdsVersion;
use crate::message::messages::Request;
use crate::message::prelogin::{
EncryptionType, PreloginRequest, PreloginRequestModel, PreloginResponse,
};
use crate::token::login_ack::LoginAckToken;
use crate::token::tokens::{SessionStateToken, SqlCollation};
use tracing::{debug, warn};
use uuid::Uuid;
pub(crate) struct HandlerFactory {
pub(crate) context: ClientContext,
pub(crate) recovery_data: Option<Box<SessionRecoveryData>>,
}
impl HandlerFactory {
pub(crate) fn prelogin_handler(&self) -> PreloginHandler<'_> {
PreloginHandler { factory: self }
}
pub(crate) fn login_handler(&self, prelogin_fedauth_supported: bool) -> LoginHandler<'_> {
LoginHandler {
factory: self,
prelogin_fedauth_supported,
}
}
pub(crate) fn session_handler<'a>(
&'a self,
transport_context: &'a TransportContext,
) -> SessionHandler<'a, 'a> {
SessionHandler {
factory: self,
transport_context,
}
}
fn create_login_request<'a, 'b>(
&'a self,
prelogin_fedauth_supported: bool,
transport_context: &'b TransportContext,
) -> LoginRequest<'a>
where
'b: 'a,
{
let model = self.create_login_model(prelogin_fedauth_supported, transport_context);
LoginRequest { model }
}
fn create_login_model<'a, 'b>(
&'a self,
prelogin_fedauth_supported: bool,
transport_context: &'b TransportContext,
) -> LoginRequestModel<'a>
where
'b: 'a,
{
LoginRequestModel::from_context(
&self.context,
prelogin_fedauth_supported,
transport_context,
self.recovery_data.as_deref(),
)
}
fn create_login_response(&self) -> LoginResponse {
LoginResponse::new()
}
}
#[derive(Debug)]
pub(crate) struct NegotiatedSettings {
pub session_settings: SessionSettings,
pub database_collation: SqlCollation,
pub language: String,
pub database: String,
login_database: String,
login_language: String,
login_database_collation: SqlCollation,
#[allow(dead_code)] pub char_set: Option<String>,
pub login_ack_tds_version: Option<TdsVersion>,
pub login_ack_server_version: Option<Version>,
pub server_reported_name: Option<String>,
}
impl NegotiatedSettings {
fn new(
session_settings: SessionSettings,
database_collation: SqlCollation,
language: String,
database: String,
char_set: Option<String>,
login_ack_tds_version: Option<TdsVersion>,
login_ack_server_version: Option<Version>,
) -> Self {
NegotiatedSettings {
session_settings,
database_collation,
login_database: database.clone(),
login_language: language.clone(),
login_database_collation: database_collation,
language,
database,
char_set,
login_ack_tds_version,
login_ack_server_version,
server_reported_name: None,
}
}
pub(crate) fn restore_login_defaults(&mut self) {
self.database = self.login_database.clone();
self.language = self.login_language.clone();
self.database_collation = self.login_database_collation;
}
pub(crate) fn is_session_recovery_acknowledged(&self) -> bool {
self.session_settings
.supported_features
.iter()
.any(|f| f.feature_identifier() == FeatureExtension::SRecovery && f.is_acknowledged())
}
pub(crate) fn session_recovery_initial_state(&self) -> Option<&[u8]> {
self.session_settings
.supported_features
.iter()
.find(|f| f.feature_identifier() == FeatureExtension::SRecovery && f.is_acknowledged())
.and_then(|f| f.session_recovery_initial_state())
}
pub(crate) fn is_column_encryption_supported(&self) -> bool {
self.session_settings.supported_features.iter().any(|f| {
f.feature_identifier() == FeatureExtension::AlwaysEncrypted && f.is_acknowledged()
})
}
}
#[derive(Debug)]
pub(crate) struct SessionSettings {
pub packet_size: u32,
#[allow(dead_code)] pub user_name: String,
pub(crate) supported_features: Vec<Box<dyn Feature>>,
#[allow(dead_code)] pub(crate) mars_enabled: bool,
#[allow(dead_code)] pub pre_login_has_fedauth_supported: bool,
#[allow(dead_code)] pub negotiated_encryption_settings: NegotiatedEncryptionSetting,
}
impl SessionSettings {
fn new(
context: &ClientContext,
pre_login_has_fedauth_supported: bool,
packet_size: u32,
negotiated_encryption_settings: NegotiatedEncryptionSetting,
features: &mut Vec<Box<dyn Feature>>,
) -> Self {
let mut result = SessionSettings {
packet_size,
user_name: context.user_name.clone(),
supported_features: vec![],
mars_enabled: context.mars_enabled,
pre_login_has_fedauth_supported,
negotiated_encryption_settings,
};
result.supported_features.append(features);
result
}
}
#[cfg(any(fuzzing, test, feature = "test-util"))]
pub(crate) fn create_test_negotiated_settings_internal() -> NegotiatedSettings {
let session_settings = SessionSettings {
packet_size: 4096,
user_name: "test".to_string(),
supported_features: Vec::new(),
mars_enabled: false,
pre_login_has_fedauth_supported: false,
negotiated_encryption_settings: NegotiatedEncryptionSetting::NoEncryption,
};
let database_collation = SqlCollation {
info: 0,
lcid_language_id: 0x0409,
col_flags: 0,
sort_id: 0,
};
NegotiatedSettings {
session_settings,
database_collation,
language: "us_english".to_string(),
database: "master".to_string(),
login_database: "master".to_string(),
login_language: "us_english".to_string(),
login_database_collation: database_collation,
char_set: None,
login_ack_tds_version: None,
login_ack_server_version: None,
server_reported_name: None,
}
}
pub(crate) struct SessionHandler<'a, 'b> {
pub(crate) factory: &'a HandlerFactory,
pub(crate) transport_context: &'b TransportContext,
}
fn server_name_from_login_messages(messages: &[SqlInfoMessage]) -> Option<String> {
messages
.iter()
.filter_map(|m| m.server_name.as_deref())
.find(|name| !name.is_empty())
.map(str::to_string)
}
impl<'a, 'b> SessionHandler<'a, 'b> {
pub(crate) async fn execute<T: NetworkReaderWriter + TdsTokenStreamReader + TdsPacketReader>(
&mut self,
reader_writer: &mut T,
) -> TdsResult<(
NegotiatedSettings,
Vec<SqlInfoMessage>,
Vec<SessionStateToken>,
)> {
let pre_login_result = self.get_pre_login_result(reader_writer).await?;
self.validate_prelogin_result(&pre_login_result)?;
reader_writer.notify_encryption_setting_change(pre_login_result.encryption_setting);
let mut login_result = self
.get_login_result(reader_writer, pre_login_result.is_fed_auth_supported)
.await?;
self.validate_login_result(&mut login_result)?;
let negotiated_settings =
self.infer_negotiated_settings(&pre_login_result, &mut login_result)?;
let info_messages = std::mem::take(&mut login_result.diagnostics.info_messages);
let session_state_tokens = std::mem::take(&mut login_result.session_state_tokens);
reader_writer.notify_session_setting_change(&negotiated_settings.session_settings);
Ok((negotiated_settings, info_messages, session_state_tokens))
}
async fn get_pre_login_result<T: NetworkReaderWriter + TdsPacketReader>(
&self,
reader_writer: &mut T,
) -> TdsResult<PreloginResult> {
let result = self
.factory
.prelogin_handler()
.execute(reader_writer)
.await?;
Ok(result)
}
fn validate_prelogin_result(&self, _result: &PreloginResult) -> TdsResult<()> {
Ok(())
}
fn validate_login_result(&self, result: &mut LoginResult) -> TdsResult<()> {
if result.status == LoginResponseStatus::Error {
let diagnostics = std::mem::take(&mut result.diagnostics);
if diagnostics.errors.is_empty() {
return Err(Error::ProtocolError(
"Login failed: Server did not send login acknowledgement".to_string(),
));
}
return Err(Error::from_sql_diagnostics(diagnostics));
}
Ok(())
}
fn infer_negotiated_settings(
&self,
prelogin_result: &PreloginResult,
login_result: &mut LoginResult,
) -> TdsResult<NegotiatedSettings> {
let change_props = &login_result.change_properties;
let packet_size = change_props.packet_size as u32;
let session_settings = SessionSettings::new(
&self.factory.context,
prelogin_result.is_fed_auth_supported,
packet_size,
prelogin_result.encryption_setting,
&mut login_result.supported_features,
);
let database_collation = match change_props.database_collation {
Some(ref collation) => *collation,
None => {
return Err(Error::ProtocolError(
"Database collation missing after login. Server must send database collation in EnvChange token.".to_string()
));
}
};
let database = match change_props.database {
Some(ref db) => db.clone(),
None => {
return Err(Error::ProtocolError(
"Database name missing after login. Server must send database name in EnvChange token.".to_string()
));
}
};
let (login_ack_tds_version, login_ack_server_version) = match &login_result.login_ack {
Some(ack) => (Some(ack.tds_version), Some(ack.prog_version)),
None => (None, None),
};
let language = change_props.language.clone().unwrap_or_default();
let mut settings = NegotiatedSettings::new(
session_settings,
database_collation,
language,
database,
change_props.char_set.clone(),
login_ack_tds_version,
login_ack_server_version,
);
settings.server_reported_name =
server_name_from_login_messages(&login_result.diagnostics.info_messages);
Ok(settings)
}
async fn get_login_result<T: TdsTokenStreamReader + NetworkReaderWriter>(
&mut self,
reader_writer: &mut T,
prelogin_fedauth_supported: bool,
) -> TdsResult<LoginResult> {
self.factory
.login_handler(prelogin_fedauth_supported)
.execute(reader_writer, self.transport_context)
.await
}
}
struct PreloginResult {
encryption_setting: NegotiatedEncryptionSetting,
is_fed_auth_supported: bool,
}
impl PreloginResult {}
pub(crate) struct PreloginHandler<'a> {
factory: &'a HandlerFactory,
}
impl PreloginHandler<'_> {
async fn execute<T: NetworkReaderWriter + TdsPacketReader>(
&self,
reader_writer: &mut T,
) -> TdsResult<PreloginResult> {
let request_model = PreloginRequestModel::new(
Uuid::new_v4(),
Option::from(self.factory.context.mars_enabled),
Option::from(self.factory.context.encryption_options.mode),
Option::from(self.factory.context.database_instance.as_str()),
);
let prelogin_request = PreloginRequest {
model: &request_model,
};
let mut packet_writer = prelogin_request.create_packet_writer(reader_writer, None, None);
prelogin_request.serialize(&mut packet_writer).await?;
let response = PreloginResponse {};
let response_model = &response.deserialize(reader_writer).await?;
if request_model.mars_enabled && !response_model.mars_enabled {
return Err(Error::ProtocolError(
"Server does not support MARS (Multiple Active Result Sets)".to_string(),
));
}
if !response_model.dbinstance_valid {
warn!("Database instance validation failed");
}
if request_model.encryption_setting == EncryptionSetting::Strict {
return Ok(PreloginResult {
encryption_setting: NegotiatedEncryptionSetting::Strict,
is_fed_auth_supported: response_model.federated_auth_supported,
});
}
match &response_model.encryption {
EncryptionType::Off => {
if request_model.encryption_setting == EncryptionSetting::PreferOff {
Ok(PreloginResult {
encryption_setting: NegotiatedEncryptionSetting::LoginOnly,
is_fed_auth_supported: response_model.federated_auth_supported,
})
} else {
Err(Error::ProtocolError(format!(
"Server disallowed encryption but client requires it (client: {:?}, server: Off)",
request_model.encryption_setting
)))
}
}
EncryptionType::NotSupported => {
if request_model.encryption_setting == EncryptionSetting::PreferOff {
Ok(PreloginResult {
encryption_setting: NegotiatedEncryptionSetting::NoEncryption,
is_fed_auth_supported: response_model.federated_auth_supported,
})
} else {
Err(Error::ProtocolError(format!(
"Server does not support encryption but client requires it (client: {:?}, server: NotSupported)",
request_model.encryption_setting
)))
}
}
_ => Ok(PreloginResult {
encryption_setting: NegotiatedEncryptionSetting::Mandatory,
is_fed_auth_supported: response_model.federated_auth_supported,
}),
}
}
}
struct LoginResult {
supported_features: Vec<Box<dyn Feature>>,
change_properties: EnvChangeProperties,
status: LoginResponseStatus,
login_ack: Option<LoginAckToken>,
diagnostics: SqlServerDiagnostics,
session_state_tokens: Vec<SessionStateToken>,
}
pub struct LoginHandler<'a> {
factory: &'a HandlerFactory,
prelogin_fedauth_supported: bool,
}
impl LoginHandler<'_> {
async fn execute<T: TdsTokenStreamReader + NetworkReaderWriter>(
&self,
reader_writer: &mut T,
transport_context: &TransportContext,
) -> TdsResult<LoginResult> {
let encryption = reader_writer.get_encryption_setting();
if encryption != NegotiatedEncryptionSetting::Strict
&& encryption != NegotiatedEncryptionSetting::NoEncryption
{
reader_writer.enable_ssl().await?;
}
let (request_model, mut sspi_handler) = self
.send_login7_request(reader_writer, transport_context)
.await?;
let requested_features = request_model.features_request;
let mut login_response = self
.get_login_response(reader_writer, requested_features.clone())
.await?;
let mut info_messages = std::mem::take(&mut login_response.info_messages);
while login_response.get_status() == LoginResponseStatus::WaitingForSspi {
let sspi_challenge = login_response.sspi_token.as_ref().ok_or_else(|| {
Error::ProtocolError(
"Login response status is WaitingForSspi but no sspi_token present".to_string(),
)
})?;
let handler = sspi_handler.as_mut().ok_or_else(|| {
Error::ProtocolError(
"Received SSPI challenge but no SSPI handler available".to_string(),
)
})?;
debug!(
"Processing SSPI challenge ({} bytes)",
sspi_challenge.data.len()
);
let response_opt = handler.process_challenge(&sspi_challenge.data)?;
match response_opt {
Some(response_token) => {
debug!("Sending SSPI response ({} bytes)", response_token.len());
let sspi_request = SspiRequest {
token_data: response_token,
};
let mut packet_writer =
sspi_request.create_packet_writer(reader_writer, None, None);
sspi_request.serialize(&mut packet_writer).await?;
let mut next_login_response = self
.get_login_response(reader_writer, requested_features.clone())
.await?;
info_messages.append(&mut next_login_response.info_messages);
login_response = next_login_response;
}
None => {
debug!(
"SSPI authentication complete, no response needed - reading final server response"
);
let mut next_login_response = self
.get_login_response(reader_writer, requested_features.clone())
.await?;
info_messages.append(&mut next_login_response.info_messages);
login_response = next_login_response;
}
}
}
login_response = if login_response.get_status() == LoginResponseStatus::WaitingForFedAuth {
let fed_auth_info = match &login_response.fed_auth_info {
Some(fed_auth_info) => fed_auth_info,
None => {
return Err(Error::ProtocolError(
"Login response status is WaitingForFedAuth but no fed_auth_info present. Protocol error.".to_string()
));
}
};
let context = &self.factory.context;
let entra_id_token_factory = context.entra_id_token_factory()?;
let token = entra_id_token_factory
.create_token(
fed_auth_info.spn.clone(),
fed_auth_info.sts_url.clone(),
context.tds_authentication_method.clone(),
)
.await?;
let fed_auth_request = FedAuthTokenRequest {
access_token_bytes: token,
};
let mut packet_writer =
fed_auth_request.create_packet_writer(reader_writer, None, None);
fed_auth_request.serialize(&mut packet_writer).await?;
let mut next_login_response = self
.get_login_response(reader_writer, requested_features)
.await?;
info_messages.append(&mut next_login_response.info_messages);
next_login_response
} else {
login_response
};
let response_status = login_response.get_status();
let errors = std::mem::take(&mut login_response.errors);
let supported_features = login_response
.features
.get_acknowledged_features()
.iter()
.map(|f| f.clone_box())
.collect();
let session_state_tokens = std::mem::take(&mut login_response.session_state_tokens);
Ok(LoginResult {
supported_features,
change_properties: login_response.change_properties,
status: response_status,
login_ack: login_response.success_token,
diagnostics: SqlServerDiagnostics::new(errors, info_messages),
session_state_tokens,
})
}
async fn send_login7_request<'a, 'b>(
&'a self,
reader_writer: &mut impl NetworkReaderWriter,
transport_context: &'b TransportContext,
) -> TdsResult<(LoginRequestModel<'a>, Option<SspiAuthHandler>)>
where
'b: 'a,
{
let context = &self.factory.context;
let is_integrated_security = context.integrated_security();
let (sspi_token, sspi_handler) = if is_integrated_security {
let mut config = context.integrated_auth_config();
if let Some(token) = reader_writer.channel_binding_token() {
debug!(
"Applying TLS channel binding token ({} bytes) for Extended Protection",
token.len()
);
config = config.with_channel_bindings(token);
}
let server = transport_context.get_server_name();
let port = transport_context.get_port();
debug!(
"Creating SSPI handler for integrated authentication to {}:{}",
server, port
);
let mut handler = SspiAuthHandler::new(&config, &server, port)?;
let initial_token = handler.get_initial_token()?;
debug!(
"Generated initial SSPI token ({} bytes) using {}",
initial_token.len(),
handler.package_name()
);
(Some(initial_token), Some(handler))
} else {
(None, None)
};
let request = if let Some(token) = sspi_token {
LoginRequest {
model: LoginRequestModel::from_context_with_sspi(
context,
self.prelogin_fedauth_supported,
transport_context,
token,
self.factory.recovery_data.as_deref(),
),
}
} else {
self.factory
.create_login_request(self.prelogin_fedauth_supported, transport_context)
};
let mut packet_writer = request.create_packet_writer(reader_writer, None, None);
request.serialize(&mut packet_writer).await?;
Ok((request.model, sspi_handler))
}
async fn get_login_response<T: TdsTokenStreamReader>(
&self,
reader_writer: &mut T,
requested_features: FeaturesRequest,
) -> TdsResult<LoginResponseModel> {
let response = self.factory.create_login_response();
response
.deserialize(reader_writer, requested_features)
.await
}
}
#[cfg(test)]
mod tests {
use super::{
NegotiatedSettings, create_test_negotiated_settings_internal,
server_name_from_login_messages,
};
use crate::error::SqlInfoMessage;
use crate::message::features::always_encrypted::AlwaysEncryptedFeature;
use crate::message::features::session_recovery::SessionRecoveryFeature;
use crate::message::login::Feature;
fn info_message(server_name: Option<&str>) -> SqlInfoMessage {
SqlInfoMessage {
message: "Changed database context to 'master'.".to_string(),
state: 1,
class: 0,
number: 5701,
server_name: server_name.map(str::to_string),
proc_name: None,
line_number: None,
}
}
#[test]
fn server_name_comes_from_the_first_login_message_that_carries_one() {
let messages = vec![info_message(Some("SQLPROD01")), info_message(Some("OTHER"))];
assert_eq!(
server_name_from_login_messages(&messages).as_deref(),
Some("SQLPROD01")
);
}
#[test]
fn blank_and_absent_server_names_are_skipped_not_accepted() {
let messages = vec![
info_message(None),
info_message(Some("")),
info_message(Some("SQLPROD01")),
];
assert_eq!(
server_name_from_login_messages(&messages).as_deref(),
Some("SQLPROD01")
);
}
#[test]
fn no_usable_server_name_yields_none() {
assert_eq!(server_name_from_login_messages(&[]), None);
let blank = vec![info_message(None), info_message(Some(""))];
assert_eq!(server_name_from_login_messages(&blank), None);
}
fn settings_with_ae(acknowledged: bool) -> NegotiatedSettings {
let mut settings = create_test_negotiated_settings_internal();
let mut feature = AlwaysEncryptedFeature::default();
feature.set_acknowledged(acknowledged);
settings
.session_settings
.supported_features
.push(Box::new(feature));
settings
}
#[test]
fn is_column_encryption_supported_true_when_acknowledged() {
let settings = settings_with_ae(true);
assert!(settings.is_column_encryption_supported());
}
#[test]
fn is_column_encryption_supported_false_when_feature_present_but_unacknowledged() {
let settings = settings_with_ae(false);
assert!(!settings.is_column_encryption_supported());
}
#[test]
fn is_column_encryption_supported_false_when_feature_absent() {
let settings = create_test_negotiated_settings_internal();
assert!(!settings.is_column_encryption_supported());
}
fn settings_with_session_recovery(acknowledged: bool) -> NegotiatedSettings {
let mut settings = create_test_negotiated_settings_internal();
let mut feature = SessionRecoveryFeature::new(1);
feature.set_acknowledged(acknowledged);
settings
.session_settings
.supported_features
.push(Box::new(feature));
settings
}
#[test]
fn is_session_recovery_acknowledged_true_when_acknowledged() {
let settings = settings_with_session_recovery(true);
assert!(settings.is_session_recovery_acknowledged());
}
#[test]
fn is_session_recovery_acknowledged_false_when_unacknowledged() {
let settings = settings_with_session_recovery(false);
assert!(!settings.is_session_recovery_acknowledged());
}
#[test]
fn is_session_recovery_acknowledged_false_when_feature_absent() {
let settings = create_test_negotiated_settings_internal();
assert!(!settings.is_session_recovery_acknowledged());
}
#[test]
fn session_recovery_initial_state_returns_ack_payload() {
let mut settings = create_test_negotiated_settings_internal();
let mut feature = SessionRecoveryFeature::new(1);
feature.set_acknowledged(true);
feature.deserialize(&[0x07, 0x02, 0x01, 0x02]).unwrap();
settings
.session_settings
.supported_features
.push(Box::new(feature));
assert_eq!(
settings.session_recovery_initial_state(),
Some([0x07, 0x02, 0x01, 0x02].as_slice())
);
}
#[test]
fn session_recovery_initial_state_none_without_feature() {
let settings = create_test_negotiated_settings_internal();
assert!(settings.session_recovery_initial_state().is_none());
}
}