diff --git a/Cargo.toml b/Cargo.toml index 5b58379..50579e6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -32,11 +32,11 @@ categories = ["database"] name = "iotdb_client" [dependencies] -thrift = "0.23" +thrift = "0.25" byteorder = "1.5" chrono = "0.4" log = "0.4" -# thrift 0.23 allows any uuid 1.x; 1.21+ raises MSRV to Rust 1.85. +# thrift 0.25 allows any uuid 1.x; 1.21+ raises MSRV to Rust 1.85. uuid = "=1.20.0" rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"], optional = true } rustls-pemfile = { version = "2.2", optional = true } @@ -55,4 +55,11 @@ env_logger = "0.11" default = [] # TLS 1.2/1.3 via rustls, using the ring crypto provider and platform # trust roots with consistent WebPKI verification. Adds `use_ssl` & friends to SessionConfig. -tls = ["dep:rustls", "dep:rustls-pemfile", "dep:rustls-native-certs", "dep:zeroize", "dep:security-framework"] +tls = [ + "thrift/rustls", + "dep:rustls", + "dep:rustls-pemfile", + "dep:rustls-native-certs", + "dep:zeroize", + "dep:security-framework", +] diff --git a/src/connection/mod.rs b/src/connection/mod.rs index 7bcdf81..33f9a7c 100644 --- a/src/connection/mod.rs +++ b/src/connection/mod.rs @@ -35,15 +35,14 @@ use rustls::crypto::WebPkiSupportedAlgorithms; #[cfg(feature = "tls")] use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; #[cfg(feature = "tls")] -use rustls::{ - ClientConfig, ClientConnection, DigitallySignedStruct, RootCertStore, SignatureScheme, - StreamOwned, -}; +use rustls::{ClientConfig, DigitallySignedStruct, RootCertStore, SignatureScheme}; use thrift::protocol::{ TBinaryInputProtocol, TBinaryOutputProtocol, TCompactInputProtocol, TCompactOutputProtocol, TInputProtocol, TOutputProtocol, }; +#[cfg(feature = "tls")] +use thrift::transport::TTlsClientChannel; use thrift::transport::{TFramedReadTransport, TFramedWriteTransport, TIoChannel, TTcpChannel}; use crate::error::{Error, Result}; @@ -197,9 +196,9 @@ impl Connection { #[cfg(feature = "tls")] if let Some(tls) = &options.tls { - let stream = tls_handshake(&endpoint, stream, tls)?; - let shared = SharedTlsStream::new(stream); - let (input, output) = build_protocols(shared.clone(), shared, options.protocol); + let channel = tls_channel(&endpoint, stream, tls)?; + let (read_half, write_half) = channel.split()?; + let (input, output) = build_protocols(read_half, write_half, options.protocol); return Ok(Self { endpoint, protocol: options.protocol, @@ -275,28 +274,19 @@ fn connect_stream(endpoint: &Endpoint, connect_timeout: Duration) -> Result Result> { +) -> Result { let config = tls_client_config(tls)?; let domain = tls.domain_override.as_deref().unwrap_or(&endpoint.host); let server_name = ServerName::try_from(domain.to_owned()) .map_err(|e| Error::Client(format!("invalid TLS server name '{domain}': {e}")))?; - let mut connection = - ClientConnection::new(config, server_name).map_err(|e| Error::Tls(e.to_string()))?; - - connection - .complete_io(&mut stream) - .map_err(|e| Error::Tls(e.to_string()))?; - if connection.is_handshaking() { - return Err(Error::Tls("TLS handshake did not complete".into())); - } - - Ok(StreamOwned::new(connection, stream)) + TTlsClientChannel::with_stream(stream, server_name, config) + .map_err(|e| Error::Tls(e.to_string())) } #[cfg(feature = "tls")] @@ -445,46 +435,6 @@ impl ServerCertVerifier for NoCertificateVerification { } } -/// A `TlsStream` shared between the read and write transports. -/// -/// `TTcpChannel::split` clones the underlying OS socket, but a TLS stream -/// cannot be split that way (record layer state is shared), so both framed -/// transports hold the same stream behind a mutex. The generated sync -/// client fully writes + flushes a request before reading the response, so -/// read and write never contend. -#[cfg(feature = "tls")] -#[derive(Clone)] -struct SharedTlsStream(std::sync::Arc>>); - -#[cfg(feature = "tls")] -impl SharedTlsStream { - fn new(stream: StreamOwned) -> Self { - Self(std::sync::Arc::new(std::sync::Mutex::new(stream))) - } - - fn lock(&self) -> std::sync::MutexGuard<'_, StreamOwned> { - self.0.lock().unwrap_or_else(|p| p.into_inner()) - } -} - -#[cfg(feature = "tls")] -impl Read for SharedTlsStream { - fn read(&mut self, buf: &mut [u8]) -> std::io::Result { - self.lock().read(buf) - } -} - -#[cfg(feature = "tls")] -impl Write for SharedTlsStream { - fn write(&mut self, buf: &[u8]) -> std::io::Result { - self.lock().write(buf) - } - - fn flush(&mut self) -> std::io::Result<()> { - self.lock().flush() - } -} - #[cfg(test)] mod tests { use super::*;