Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 10 additions & 3 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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 }
Expand All @@ -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",
]
74 changes: 12 additions & 62 deletions src/connection/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -275,28 +274,19 @@ fn connect_stream(endpoint: &Endpoint, connect_timeout: Duration) -> Result<TcpS
})
}

/// Run the TLS handshake over an established TCP stream.
/// Wrap an established TCP stream in Thrift's TLS channel and run the handshake.
#[cfg(feature = "tls")]
fn tls_handshake(
fn tls_channel(
endpoint: &Endpoint,
mut stream: TcpStream,
stream: TcpStream,
tls: &TlsOptions,
) -> Result<StreamOwned<ClientConnection, TcpStream>> {
) -> Result<TTlsClientChannel> {
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")]
Expand Down Expand Up @@ -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<std::sync::Mutex<StreamOwned<ClientConnection, TcpStream>>>);

#[cfg(feature = "tls")]
impl SharedTlsStream {
fn new(stream: StreamOwned<ClientConnection, TcpStream>) -> Self {
Self(std::sync::Arc::new(std::sync::Mutex::new(stream)))
}

fn lock(&self) -> std::sync::MutexGuard<'_, StreamOwned<ClientConnection, TcpStream>> {
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<usize> {
self.lock().read(buf)
}
}

#[cfg(feature = "tls")]
impl Write for SharedTlsStream {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.lock().write(buf)
}

fn flush(&mut self) -> std::io::Result<()> {
self.lock().flush()
}
}

#[cfg(test)]
mod tests {
use super::*;
Expand Down
Loading