1use super::protocol::{HEADER_LEN, PacketHeader, packet, split_message};
41use rustlavel_core::{Error, Result};
42use std::io;
43use std::pin::Pin;
44use std::sync::Arc;
45use std::task::{Context, Poll, ready};
46use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, ReadBuf};
47use tokio::net::TcpStream;
48
49#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
51pub enum Encryption {
52 Disabled,
55 LoginOnly,
58 #[default]
61 Required,
62}
63
64impl Encryption {
65 pub fn as_byte(self) -> u8 {
67 match self {
68 Encryption::Disabled => super::protocol::encryption::NOT_SUPPORTED,
69 Encryption::LoginOnly => super::protocol::encryption::OFF,
70 Encryption::Required => super::protocol::encryption::ON,
71 }
72 }
73}
74
75#[derive(Debug, Clone, Copy, PartialEq, Eq)]
77pub enum Negotiated {
78 None,
80 LoginOnly,
82 Session,
84}
85
86pub fn negotiate(requested: Encryption, server: u8) -> Result<Negotiated> {
92 use super::protocol::encryption as level;
93
94 Ok(match (requested, server) {
95 (Encryption::Disabled, level::NOT_SUPPORTED) => Negotiated::None,
96 (Encryption::Disabled, level::REQUIRED) | (Encryption::Disabled, level::ON) => {
97 return Err(Error::msg(
98 "the server requires an encrypted connection, but this connection asked for none. \
99 Use the default encryption setting.",
100 ));
101 }
102 (Encryption::Disabled, _) => Negotiated::None,
103
104 (Encryption::LoginOnly, level::NOT_SUPPORTED) => Negotiated::None,
107 (Encryption::Required, level::NOT_SUPPORTED) => {
108 return Err(Error::msg(
109 "this connection requires encryption, but the server reports it cannot encrypt. \
110 Give SQL Server a certificate, or connect with encryption set to login-only.",
111 ));
112 }
113
114 (_, level::ON) | (_, level::REQUIRED) => Negotiated::Session,
117 (Encryption::Required, _) => Negotiated::Session,
118 (Encryption::LoginOnly, _) => Negotiated::LoginOnly,
119 })
120}
121
122pub fn obfuscate_password(password: &str) -> Vec<u8> {
129 password
130 .encode_utf16()
131 .flat_map(u16::to_le_bytes)
132 .map(|byte| byte.rotate_left(4) ^ 0xA5)
134 .collect()
135}
136
137pub fn deobfuscate_password(bytes: &[u8]) -> String {
140 let plain: Vec<u8> = bytes
141 .iter()
142 .map(|byte| {
143 let plain = byte ^ 0xA5;
144 plain.rotate_left(4)
145 })
146 .collect();
147
148 let (pairs, _odd_trailing_byte) = plain.as_chunks::<2>();
152 let units: Vec<u16> = pairs.iter().copied().map(u16::from_le_bytes).collect();
153
154 String::from_utf16_lossy(&units)
155}
156
157pub struct TdsHandshakeStream {
165 socket: TcpStream,
166 wrapping: bool,
168 packet_size: usize,
169 outgoing: Vec<u8>,
171 pending: Vec<u8>,
173 pending_at: usize,
174 incoming: Vec<u8>,
176 ready: Vec<u8>,
178 ready_at: usize,
179}
180
181impl TdsHandshakeStream {
182 pub fn new(socket: TcpStream, packet_size: usize) -> Self {
183 TdsHandshakeStream {
184 socket,
185 wrapping: true,
186 packet_size,
187 outgoing: Vec::new(),
188 pending: Vec::new(),
189 pending_at: 0,
190 incoming: Vec::new(),
191 ready: Vec::new(),
192 ready_at: 0,
193 }
194 }
195
196 pub fn stop_wrapping(&mut self) {
199 self.wrapping = false;
200 }
201
202 pub fn into_socket(self) -> Result<TcpStream> {
207 if self.ready_at < self.ready.len() || !self.incoming.is_empty() {
208 return Err(Error::Protocol(
209 "the server sent data before encryption was torn down".into(),
210 ));
211 }
212 Ok(self.socket)
213 }
214
215 fn take_packet(&mut self) -> Result<bool> {
217 if self.incoming.len() < HEADER_LEN {
218 return Ok(false);
219 }
220 let header = PacketHeader::parse(&self.incoming)?;
221 let total = header.length as usize;
222 if total < HEADER_LEN {
223 return Err(Error::Protocol("packet length is impossibly small".into()));
224 }
225 if self.incoming.len() < total {
226 return Ok(false);
227 }
228
229 self.ready = self.incoming[HEADER_LEN..total].to_vec();
230 self.ready_at = 0;
231 self.incoming.drain(..total);
232 Ok(true)
233 }
234}
235
236fn protocol_io(error: Error) -> io::Error {
237 io::Error::new(io::ErrorKind::InvalidData, error.to_string())
238}
239
240impl AsyncRead for TdsHandshakeStream {
241 fn poll_read(
242 mut self: Pin<&mut Self>,
243 cx: &mut Context<'_>,
244 buf: &mut ReadBuf<'_>,
245 ) -> Poll<io::Result<()>> {
246 loop {
247 if self.ready_at < self.ready.len() {
248 let take = buf.remaining().min(self.ready.len() - self.ready_at);
249 let at = self.ready_at;
250 buf.put_slice(&self.ready[at..at + take]);
251 self.ready_at += take;
252 return Poll::Ready(Ok(()));
253 }
254
255 if !self.wrapping {
256 if !self.incoming.is_empty() {
259 self.ready = std::mem::take(&mut self.incoming);
260 self.ready_at = 0;
261 continue;
262 }
263 return Pin::new(&mut self.socket).poll_read(cx, buf);
264 }
265
266 if self.take_packet().map_err(protocol_io)? {
267 continue;
268 }
269
270 let mut chunk = [0u8; 8192];
271 let mut incoming = ReadBuf::new(&mut chunk);
272 ready!(Pin::new(&mut self.socket).poll_read(cx, &mut incoming))?;
273
274 let filled = incoming.filled().len();
275 if filled == 0 {
276 return Poll::Ready(Ok(()));
279 }
280 let bytes = incoming.filled().to_vec();
281 self.incoming.extend_from_slice(&bytes);
282 }
283 }
284}
285
286impl AsyncWrite for TdsHandshakeStream {
287 fn poll_write(
288 mut self: Pin<&mut Self>,
289 cx: &mut Context<'_>,
290 buf: &[u8],
291 ) -> Poll<io::Result<usize>> {
292 if !self.wrapping {
293 return Pin::new(&mut self.socket).poll_write(cx, buf);
294 }
295 self.outgoing.extend_from_slice(buf);
298 Poll::Ready(Ok(buf.len()))
299 }
300
301 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
302 if self.wrapping && !self.outgoing.is_empty() {
303 let flight = std::mem::take(&mut self.outgoing);
304 let packet_size = self.packet_size;
305 for packet in split_message(packet::PRE_LOGIN, &flight, packet_size) {
306 self.pending.extend_from_slice(&packet);
307 }
308 }
309
310 while self.pending_at < self.pending.len() {
311 let this = &mut *self;
312 let written =
313 ready!(Pin::new(&mut this.socket).poll_write(cx, &this.pending[this.pending_at..]))?;
314 if written == 0 {
315 return Poll::Ready(Err(io::ErrorKind::WriteZero.into()));
316 }
317 self.pending_at += written;
318 }
319 self.pending.clear();
320 self.pending_at = 0;
321
322 Pin::new(&mut self.socket).poll_flush(cx)
323 }
324
325 fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
326 Pin::new(&mut self.socket).poll_shutdown(cx)
327 }
328}
329
330pub enum TdsStream {
336 Plain(TcpStream),
337 Tls(Box<tokio_rustls::client::TlsStream<TdsHandshakeStream>>),
338 Closed,
343}
344
345impl TdsStream {
346 pub async fn write_all(&mut self, bytes: &[u8]) -> io::Result<()> {
347 match self {
348 TdsStream::Plain(stream) => stream.write_all(bytes).await,
349 TdsStream::Tls(stream) => stream.write_all(bytes).await,
350 TdsStream::Closed => Err(closed()),
351 }
352 }
353
354 pub async fn flush(&mut self) -> io::Result<()> {
355 match self {
356 TdsStream::Plain(stream) => stream.flush().await,
357 TdsStream::Tls(stream) => stream.flush().await,
358 TdsStream::Closed => Err(closed()),
359 }
360 }
361
362 pub async fn read(&mut self, buffer: &mut [u8]) -> io::Result<usize> {
363 use tokio::io::AsyncReadExt;
364 match self {
365 TdsStream::Plain(stream) => stream.read(buffer).await,
366 TdsStream::Tls(stream) => stream.read(buffer).await,
367 TdsStream::Closed => Err(closed()),
368 }
369 }
370
371 pub async fn shutdown(&mut self) -> io::Result<()> {
372 match self {
373 TdsStream::Plain(stream) => stream.shutdown().await,
374 TdsStream::Tls(stream) => stream.shutdown().await,
375 TdsStream::Closed => Ok(()),
376 }
377 }
378
379 pub fn take(&mut self) -> TdsStream {
381 std::mem::replace(self, TdsStream::Closed)
382 }
383
384 pub fn into_plain(self) -> Result<TdsStream> {
387 match self {
388 TdsStream::Tls(stream) => {
389 let (wrapper, _session) = stream.into_inner();
390 Ok(TdsStream::Plain(wrapper.into_socket()?))
391 }
392 other => Ok(other),
393 }
394 }
395}
396
397fn closed() -> io::Error {
398 io::Error::new(io::ErrorKind::NotConnected, "the TDS connection has no transport")
399}
400
401#[derive(Debug, Clone, Copy)]
403pub struct TlsOptions {
404 pub trust_server_certificate: bool,
414}
415
416impl Default for TlsOptions {
417 fn default() -> Self {
418 TlsOptions { trust_server_certificate: true }
419 }
420}
421
422pub async fn start_tls(
424 socket: TcpStream,
425 host: &str,
426 options: TlsOptions,
427 packet_size: usize,
428) -> Result<tokio_rustls::client::TlsStream<TdsHandshakeStream>> {
429 let connector = tokio_rustls::TlsConnector::from(client_config(options));
430
431 let name = rustls::pki_types::ServerName::try_from(host.to_string())
435 .map_err(|_| Error::msg(format!("`{host}` is not a valid TLS server name")))?;
436
437 let wrapper = TdsHandshakeStream::new(socket, packet_size);
438 let mut stream = connector.connect(name, wrapper).await.map_err(|e| {
439 Error::msg(format!(
440 "the TLS handshake inside SQL Server's pre-login exchange failed: {e}"
441 ))
442 })?;
443
444 stream.get_mut().0.stop_wrapping();
446 Ok(stream)
447}
448
449fn client_config(options: TlsOptions) -> Arc<rustls::ClientConfig> {
463 use std::sync::OnceLock;
464 static TRUSTING: OnceLock<Arc<rustls::ClientConfig>> = OnceLock::new();
465 static VERIFYING: OnceLock<Arc<rustls::ClientConfig>> = OnceLock::new();
466
467 if options.trust_server_certificate {
468 Arc::clone(TRUSTING.get_or_init(|| {
469 let provider = Arc::new(rustls::crypto::aws_lc_rs::default_provider());
470 let verifier = Arc::new(TrustAnyCertificate(Arc::clone(&provider)));
471 Arc::new(
472 builder(provider)
473 .dangerous()
474 .with_custom_certificate_verifier(verifier)
475 .with_no_client_auth(),
476 )
477 }))
478 } else {
479 Arc::clone(VERIFYING.get_or_init(|| {
480 let roots = rustls::RootCertStore { roots: webpki_roots::TLS_SERVER_ROOTS.to_vec() };
483 Arc::new(
484 builder(Arc::new(rustls::crypto::aws_lc_rs::default_provider()))
485 .with_root_certificates(roots)
486 .with_no_client_auth(),
487 )
488 }))
489 }
490}
491
492fn builder(
495 provider: Arc<rustls::crypto::CryptoProvider>,
496) -> rustls::ConfigBuilder<rustls::ClientConfig, rustls::WantsVerifier> {
497 rustls::ClientConfig::builder_with_provider(provider)
498 .with_protocol_versions(&[&rustls::version::TLS12])
499 .expect("TLS 1.2 is enabled by this crate's rustls features")
500}
501
502#[derive(Debug)]
517struct TrustAnyCertificate(Arc<rustls::crypto::CryptoProvider>);
518
519impl rustls::client::danger::ServerCertVerifier for TrustAnyCertificate {
520 fn verify_server_cert(
521 &self,
522 _end_entity: &rustls::pki_types::CertificateDer<'_>,
523 _intermediates: &[rustls::pki_types::CertificateDer<'_>],
524 _server_name: &rustls::pki_types::ServerName<'_>,
525 _ocsp_response: &[u8],
526 _now: rustls::pki_types::UnixTime,
527 ) -> std::result::Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
528 Ok(rustls::client::danger::ServerCertVerified::assertion())
529 }
530
531 fn verify_tls12_signature(
532 &self,
533 _message: &[u8],
534 _cert: &rustls::pki_types::CertificateDer<'_>,
535 _dss: &rustls::DigitallySignedStruct,
536 ) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
537 Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
538 }
539
540 fn verify_tls13_signature(
541 &self,
542 _message: &[u8],
543 _cert: &rustls::pki_types::CertificateDer<'_>,
544 _dss: &rustls::DigitallySignedStruct,
545 ) -> std::result::Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
546 Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
547 }
548
549 fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
550 self.0.signature_verification_algorithms.supported_schemes()
551 }
552}
553
554#[cfg(test)]
555mod tests {
556 use super::*;
557 use super::super::protocol::encryption as level;
558
559 #[test]
560 fn the_password_scheme_swaps_nibbles_then_xors_with_a5() {
561 assert_eq!(obfuscate_password("a"), vec![0xB3, 0xA5]);
565 assert_eq!(obfuscate_password("abc"), vec![0xB3, 0xA5, 0x83, 0xA5, 0x93, 0xA5]);
566
567 assert_eq!(obfuscate_password("é").len(), 2);
569 assert_eq!(obfuscate_password(""), Vec::<u8>::new());
570 }
571
572 #[test]
573 fn obfuscation_is_reversible_which_is_the_whole_point_of_calling_it_that() {
574 for password in ["", "a", "Rustlavel!2026", "pässwörd", "日本語"] {
575 assert_eq!(deobfuscate_password(&obfuscate_password(password)), password);
576 }
577 }
578
579 #[test]
580 fn encryption_choices_map_onto_the_prelogin_option_bytes() {
581 assert_eq!(Encryption::Disabled.as_byte(), level::NOT_SUPPORTED);
582 assert_eq!(Encryption::LoginOnly.as_byte(), level::OFF);
583 assert_eq!(Encryption::Required.as_byte(), level::ON);
584 assert_eq!(Encryption::default(), Encryption::Required);
585 }
586
587 #[test]
588 fn a_server_that_offers_only_login_encryption_gets_login_encryption() {
589 assert_eq!(negotiate(Encryption::LoginOnly, level::OFF).unwrap(), Negotiated::LoginOnly);
590 }
591
592 #[test]
593 fn a_server_that_wants_full_encryption_gets_it_whatever_the_client_preferred() {
594 for server in [level::ON, level::REQUIRED] {
595 assert_eq!(negotiate(Encryption::LoginOnly, server).unwrap(), Negotiated::Session);
596 assert_eq!(negotiate(Encryption::Required, server).unwrap(), Negotiated::Session);
597 }
598 }
599
600 #[test]
601 fn a_server_that_cannot_encrypt_fails_a_connection_that_requires_it() {
602 let error = negotiate(Encryption::Required, level::NOT_SUPPORTED).unwrap_err().to_string();
603 assert!(error.contains("cannot encrypt"), "{error}");
604
605 assert_eq!(negotiate(Encryption::LoginOnly, level::NOT_SUPPORTED).unwrap(), Negotiated::None);
607 }
608
609 #[test]
610 fn refusing_encryption_a_server_requires_is_an_error_not_a_downgrade() {
611 let error = negotiate(Encryption::Disabled, level::REQUIRED).unwrap_err().to_string();
614 assert!(error.contains("requires an encrypted connection"), "{error}");
615
616 assert_eq!(negotiate(Encryption::Disabled, level::NOT_SUPPORTED).unwrap(), Negotiated::None);
617 }
618
619 #[tokio::test]
620 async fn the_handshake_wrapper_frames_a_flight_and_then_gets_out_of_the_way() {
621 use tokio::io::AsyncReadExt;
622
623 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
624 let address = listener.local_addr().unwrap();
625
626 let server = tokio::spawn(async move {
627 let (mut socket, _) = listener.accept().await.unwrap();
628 let mut received = Vec::new();
629 socket.read_buf(&mut received).await.unwrap();
630 let mut more = Vec::new();
632 let _ = tokio::time::timeout(
633 std::time::Duration::from_millis(200),
634 socket.read_buf(&mut more),
635 )
636 .await;
637 received.extend_from_slice(&more);
638 received
639 });
640
641 let mut wrapper =
642 TdsHandshakeStream::new(TcpStream::connect(address).await.unwrap(), 4096);
643 wrapper.write_all(b"handshake").await.unwrap();
644 wrapper.flush().await.unwrap();
645 wrapper.stop_wrapping();
646 wrapper.write_all(b"raw").await.unwrap();
647 wrapper.flush().await.unwrap();
648
649 let received = server.await.unwrap();
650
651 let header = PacketHeader::parse(&received).unwrap();
653 assert_eq!(header.kind, packet::PRE_LOGIN);
654 assert!(header.is_end_of_message());
655 assert_eq!(&received[HEADER_LEN..header.length as usize], b"handshake");
656
657 assert_eq!(&received[header.length as usize..], b"raw");
659 }
660}