use super::protocol::{HEADER_LEN, PacketHeader, packet, split_message};
use rustlavel_core::{Error, Result};
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll, ready};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, ReadBuf};
use tokio::net::TcpStream;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Encryption {
Disabled,
LoginOnly,
#[default]
Required,
}
impl Encryption {
pub fn as_byte(self) -> u8 {
match self {
Encryption::Disabled => super::protocol::encryption::NOT_SUPPORTED,
Encryption::LoginOnly => super::protocol::encryption::OFF,
Encryption::Required => super::protocol::encryption::ON,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Negotiated {
None,
LoginOnly,
Session,
}
pub fn negotiate(requested: Encryption, server: u8) -> Result<Negotiated> {
use super::protocol::encryption as level;
Ok(match (requested, server) {
(Encryption::Disabled, level::NOT_SUPPORTED) => Negotiated::None,
(Encryption::Disabled, level::REQUIRED) | (Encryption::Disabled, level::ON) => {
return Err(Error::msg(
"the server requires an encrypted connection, but this connection asked for none. \
Use the default encryption setting.",
));
}
(Encryption::Disabled, _) => Negotiated::None,
(Encryption::LoginOnly, level::NOT_SUPPORTED) => Negotiated::None,
(Encryption::Required, level::NOT_SUPPORTED) => {
return Err(Error::msg(
"this connection requires encryption, but the server reports it cannot encrypt. \
Give SQL Server a certificate, or connect with encryption set to login-only.",
));
}
(_, level::ON) | (_, level::REQUIRED) => Negotiated::Session,
(Encryption::Required, _) => Negotiated::Session,
(Encryption::LoginOnly, _) => Negotiated::LoginOnly,
})
}
pub fn obfuscate_password(password: &str) -> Vec<u8> {
password
.encode_utf16()
.flat_map(u16::to_le_bytes)
.map(|byte| byte.rotate_left(4) ^ 0xA5)
.collect()
}
pub fn deobfuscate_password(bytes: &[u8]) -> String {
let plain: Vec<u8> = bytes
.iter()
.map(|byte| {
let plain = byte ^ 0xA5;
plain.rotate_left(4)
})
.collect();
let (pairs, _odd_trailing_byte) = plain.as_chunks::<2>();
let units: Vec<u16> = pairs.iter().copied().map(u16::from_le_bytes).collect();
String::from_utf16_lossy(&units)
}
pub struct TdsHandshakeStream {
socket: TcpStream,
wrapping: bool,
packet_size: usize,
outgoing: Vec<u8>,
pending: Vec<u8>,
pending_at: usize,
incoming: Vec<u8>,
ready: Vec<u8>,
ready_at: usize,
}
impl TdsHandshakeStream {
pub fn new(socket: TcpStream, packet_size: usize) -> Self {
TdsHandshakeStream {
socket,
wrapping: true,
packet_size,
outgoing: Vec::new(),
pending: Vec::new(),
pending_at: 0,
incoming: Vec::new(),
ready: Vec::new(),
ready_at: 0,
}
}
pub fn stop_wrapping(&mut self) {
self.wrapping = false;
}
pub fn into_socket(self) -> Result<TcpStream> {
if self.ready_at < self.ready.len() || !self.incoming.is_empty() {
return Err(Error::Protocol(
"the server sent data before encryption was torn down".into(),
));
}
Ok(self.socket)
}
fn take_packet(&mut self) -> Result<bool> {
if self.incoming.len() < HEADER_LEN {
return Ok(false);
}
let header = PacketHeader::parse(&self.incoming)?;
let total = header.length as usize;
if total < HEADER_LEN {
return Err(Error::Protocol("packet length is impossibly small".into()));
}
if self.incoming.len() < total {
return Ok(false);
}
self.ready = self.incoming[HEADER_LEN..total].to_vec();
self.ready_at = 0;
self.incoming.drain(..total);
Ok(true)
}
}
fn protocol_io(error: Error) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, error.to_string())
}
impl AsyncRead for TdsHandshakeStream {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
loop {
if self.ready_at < self.ready.len() {
let take = buf.remaining().min(self.ready.len() - self.ready_at);
let at = self.ready_at;
buf.put_slice(&self.ready[at..at + take]);
self.ready_at += take;
return Poll::Ready(Ok(()));
}
if !self.wrapping {
if !self.incoming.is_empty() {
self.ready = std::mem::take(&mut self.incoming);
self.ready_at = 0;
continue;
}
return Pin::new(&mut self.socket).poll_read(cx, buf);
}
if self.take_packet().map_err(protocol_io)? {
continue;
}
let mut chunk = [0u8; 8192];
let mut incoming = ReadBuf::new(&mut chunk);
ready!(Pin::new(&mut self.socket).poll_read(cx, &mut incoming))?;
let filled = incoming.filled().len();
if filled == 0 {
return Poll::Ready(Ok(()));
}
let bytes = incoming.filled().to_vec();
self.incoming.extend_from_slice(&bytes);
}
}
}
impl AsyncWrite for TdsHandshakeStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
if !self.wrapping {
return Pin::new(&mut self.socket).poll_write(cx, buf);
}
self.outgoing.extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
if self.wrapping && !self.outgoing.is_empty() {
let flight = std::mem::take(&mut self.outgoing);
let packet_size = self.packet_size;
for packet in split_message(packet::PRE_LOGIN, &flight, packet_size) {
self.pending.extend_from_slice(&packet);
}
}
while self.pending_at < self.pending.len() {
let this = &mut *self;
let written =
ready!(Pin::new(&mut this.socket).poll_write(cx, &this.pending[this.pending_at..]))?;
if written == 0 {
return Poll::Ready(Err(io::ErrorKind::WriteZero.into()));
}
self.pending_at += written;
}
self.pending.clear();
self.pending_at = 0;
Pin::new(&mut self.socket).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.socket).poll_shutdown(cx)
}
}
pub enum TdsStream {
Plain(TcpStream),
Tls(Box<tokio_rustls::client::TlsStream<TdsHandshakeStream>>),
Closed,
}
impl TdsStream {
pub async fn write_all(&mut self, bytes: &[u8]) -> io::Result<()> {
match self {
TdsStream::Plain(stream) => stream.write_all(bytes).await,
TdsStream::Tls(stream) => stream.write_all(bytes).await,
TdsStream::Closed => Err(closed()),
}
}
pub async fn flush(&mut self) -> io::Result<()> {
match self {
TdsStream::Plain(stream) => stream.flush().await,
TdsStream::Tls(stream) => stream.flush().await,
TdsStream::Closed => Err(closed()),
}
}
pub async fn read(&mut self, buffer: &mut [u8]) -> io::Result<usize> {
use tokio::io::AsyncReadExt;
match self {
TdsStream::Plain(stream) => stream.read(buffer).await,
TdsStream::Tls(stream) => stream.read(buffer).await,
TdsStream::Closed => Err(closed()),
}
}
pub async fn shutdown(&mut self) -> io::Result<()> {
match self {
TdsStream::Plain(stream) => stream.shutdown().await,
TdsStream::Tls(stream) => stream.shutdown().await,
TdsStream::Closed => Ok(()),
}
}
pub fn take(&mut self) -> TdsStream {
std::mem::replace(self, TdsStream::Closed)
}
pub fn into_plain(self) -> Result<TdsStream> {
match self {
TdsStream::Tls(stream) => {
let (wrapper, _session) = stream.into_inner();
Ok(TdsStream::Plain(wrapper.into_socket()?))
}
other => Ok(other),
}
}
}
fn closed() -> io::Error {
io::Error::new(io::ErrorKind::NotConnected, "the TDS connection has no transport")
}
#[derive(Debug, Clone, Copy)]
pub struct TlsOptions {
pub trust_server_certificate: bool,
}
impl Default for TlsOptions {
fn default() -> Self {
TlsOptions { trust_server_certificate: true }
}
}
pub async fn start_tls(
socket: TcpStream,
host: &str,
options: TlsOptions,
packet_size: usize,
) -> Result<tokio_rustls::client::TlsStream<TdsHandshakeStream>> {
let connector = tokio_rustls::TlsConnector::from(client_config(options));
let name = rustls::pki_types::ServerName::try_from(host.to_string())
.map_err(|_| Error::msg(format!("`{host}` is not a valid TLS server name")))?;
let wrapper = TdsHandshakeStream::new(socket, packet_size);
let mut stream = connector.connect(name, wrapper).await.map_err(|e| {
Error::msg(format!(
"the TLS handshake inside SQL Server's pre-login exchange failed: {e}"
))
})?;
stream.get_mut().0.stop_wrapping();
Ok(stream)
}
fn client_config(options: TlsOptions) -> Arc<rustls::ClientConfig> {
use std::sync::OnceLock;
static TRUSTING: OnceLock<Arc<rustls::ClientConfig>> = OnceLock::new();
static VERIFYING: OnceLock<Arc<rustls::ClientConfig>> = OnceLock::new();
if options.trust_server_certificate {
Arc::clone(TRUSTING.get_or_init(|| {
let provider = Arc::new(rustls::crypto::aws_lc_rs::default_provider());
let verifier = Arc::new(TrustAnyCertificate(Arc::clone(&provider)));
Arc::new(
builder(provider)
.dangerous()
.with_custom_certificate_verifier(verifier)
.with_no_client_auth(),
)
}))
} else {
Arc::clone(VERIFYING.get_or_init(|| {
let roots = rustls::RootCertStore { roots: webpki_roots::TLS_SERVER_ROOTS.to_vec() };
Arc::new(
builder(Arc::new(rustls::crypto::aws_lc_rs::default_provider()))
.with_root_certificates(roots)
.with_no_client_auth(),
)
}))
}
}
fn builder(
provider: Arc<rustls::crypto::CryptoProvider>,
) -> rustls::ConfigBuilder<rustls::ClientConfig, rustls::WantsVerifier> {
rustls::ClientConfig::builder_with_provider(provider)
.with_protocol_versions(&[&rustls::version::TLS12])
.expect("TLS 1.2 is enabled by this crate's rustls features")
}
#[derive(Debug)]
struct TrustAnyCertificate(Arc<rustls::crypto::CryptoProvider>);
impl rustls::client::danger::ServerCertVerifier for TrustAnyCertificate {
fn verify_server_cert(
&self,
_end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> std::result::Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
self.0.signature_verification_algorithms.supported_schemes()
}
}
#[cfg(test)]
mod tests {
use super::*;
use super::super::protocol::encryption as level;
#[test]
fn the_password_scheme_swaps_nibbles_then_xors_with_a5() {
assert_eq!(obfuscate_password("a"), vec![0xB3, 0xA5]);
assert_eq!(obfuscate_password("abc"), vec![0xB3, 0xA5, 0x83, 0xA5, 0x93, 0xA5]);
assert_eq!(obfuscate_password("é").len(), 2);
assert_eq!(obfuscate_password(""), Vec::<u8>::new());
}
#[test]
fn obfuscation_is_reversible_which_is_the_whole_point_of_calling_it_that() {
for password in ["", "a", "Rustlavel!2026", "pässwörd", "日本語"] {
assert_eq!(deobfuscate_password(&obfuscate_password(password)), password);
}
}
#[test]
fn encryption_choices_map_onto_the_prelogin_option_bytes() {
assert_eq!(Encryption::Disabled.as_byte(), level::NOT_SUPPORTED);
assert_eq!(Encryption::LoginOnly.as_byte(), level::OFF);
assert_eq!(Encryption::Required.as_byte(), level::ON);
assert_eq!(Encryption::default(), Encryption::Required);
}
#[test]
fn a_server_that_offers_only_login_encryption_gets_login_encryption() {
assert_eq!(negotiate(Encryption::LoginOnly, level::OFF).unwrap(), Negotiated::LoginOnly);
}
#[test]
fn a_server_that_wants_full_encryption_gets_it_whatever_the_client_preferred() {
for server in [level::ON, level::REQUIRED] {
assert_eq!(negotiate(Encryption::LoginOnly, server).unwrap(), Negotiated::Session);
assert_eq!(negotiate(Encryption::Required, server).unwrap(), Negotiated::Session);
}
}
#[test]
fn a_server_that_cannot_encrypt_fails_a_connection_that_requires_it() {
let error = negotiate(Encryption::Required, level::NOT_SUPPORTED).unwrap_err().to_string();
assert!(error.contains("cannot encrypt"), "{error}");
assert_eq!(negotiate(Encryption::LoginOnly, level::NOT_SUPPORTED).unwrap(), Negotiated::None);
}
#[test]
fn refusing_encryption_a_server_requires_is_an_error_not_a_downgrade() {
let error = negotiate(Encryption::Disabled, level::REQUIRED).unwrap_err().to_string();
assert!(error.contains("requires an encrypted connection"), "{error}");
assert_eq!(negotiate(Encryption::Disabled, level::NOT_SUPPORTED).unwrap(), Negotiated::None);
}
#[tokio::test]
async fn the_handshake_wrapper_frames_a_flight_and_then_gets_out_of_the_way() {
use tokio::io::AsyncReadExt;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut received = Vec::new();
socket.read_buf(&mut received).await.unwrap();
let mut more = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_millis(200),
socket.read_buf(&mut more),
)
.await;
received.extend_from_slice(&more);
received
});
let mut wrapper =
TdsHandshakeStream::new(TcpStream::connect(address).await.unwrap(), 4096);
wrapper.write_all(b"handshake").await.unwrap();
wrapper.flush().await.unwrap();
wrapper.stop_wrapping();
wrapper.write_all(b"raw").await.unwrap();
wrapper.flush().await.unwrap();
let received = server.await.unwrap();
let header = PacketHeader::parse(&received).unwrap();
assert_eq!(header.kind, packet::PRE_LOGIN);
assert!(header.is_end_of_message());
assert_eq!(&received[HEADER_LEN..header.length as usize], b"handshake");
assert_eq!(&received[header.length as usize..], b"raw");
}
}