Skip to content

Commit 54bfa1f

Browse files
authored
Upgrade Thrift to 0.25.0 (#13)
* Upgrade Thrift to 0.25.0 * Use Thrift native TLS channel
1 parent 8046bd3 commit 54bfa1f

2 files changed

Lines changed: 22 additions & 65 deletions

File tree

‎Cargo.toml‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -32,11 +32,11 @@ categories = ["database"]
3232
name = "iotdb_client"
3333

3434
[dependencies]
35-
thrift = "0.23"
35+
thrift = "0.25"
3636
byteorder = "1.5"
3737
chrono = "0.4"
3838
log = "0.4"
39-
# thrift 0.23 allows any uuid 1.x; 1.21+ raises MSRV to Rust 1.85.
39+
# thrift 0.25 allows any uuid 1.x; 1.21+ raises MSRV to Rust 1.85.
4040
uuid = "=1.20.0"
4141
rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"], optional = true }
4242
rustls-pemfile = { version = "2.2", optional = true }
@@ -55,4 +55,11 @@ env_logger = "0.11"
5555
default = []
5656
# TLS 1.2/1.3 via rustls, using the ring crypto provider and platform
5757
# trust roots with consistent WebPKI verification. Adds `use_ssl` & friends to SessionConfig.
58-
tls = ["dep:rustls", "dep:rustls-pemfile", "dep:rustls-native-certs", "dep:zeroize", "dep:security-framework"]
58+
tls = [
59+
"thrift/rustls",
60+
"dep:rustls",
61+
"dep:rustls-pemfile",
62+
"dep:rustls-native-certs",
63+
"dep:zeroize",
64+
"dep:security-framework",
65+
]

‎src/connection/mod.rs‎

Lines changed: 12 additions & 62 deletions
Original file line numberDiff line numberDiff line change
@@ -35,15 +35,14 @@ use rustls::crypto::WebPkiSupportedAlgorithms;
3535
#[cfg(feature = "tls")]
3636
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
3737
#[cfg(feature = "tls")]
38-
use rustls::{
39-
ClientConfig, ClientConnection, DigitallySignedStruct, RootCertStore, SignatureScheme,
40-
StreamOwned,
41-
};
38+
use rustls::{ClientConfig, DigitallySignedStruct, RootCertStore, SignatureScheme};
4239

4340
use thrift::protocol::{
4441
TBinaryInputProtocol, TBinaryOutputProtocol, TCompactInputProtocol, TCompactOutputProtocol,
4542
TInputProtocol, TOutputProtocol,
4643
};
44+
#[cfg(feature = "tls")]
45+
use thrift::transport::TTlsClientChannel;
4746
use thrift::transport::{TFramedReadTransport, TFramedWriteTransport, TIoChannel, TTcpChannel};
4847

4948
use crate::error::{Error, Result};
@@ -197,9 +196,9 @@ impl Connection {
197196

198197
#[cfg(feature = "tls")]
199198
if let Some(tls) = &options.tls {
200-
let stream = tls_handshake(&endpoint, stream, tls)?;
201-
let shared = SharedTlsStream::new(stream);
202-
let (input, output) = build_protocols(shared.clone(), shared, options.protocol);
199+
let channel = tls_channel(&endpoint, stream, tls)?;
200+
let (read_half, write_half) = channel.split()?;
201+
let (input, output) = build_protocols(read_half, write_half, options.protocol);
203202
return Ok(Self {
204203
endpoint,
205204
protocol: options.protocol,
@@ -275,28 +274,19 @@ fn connect_stream(endpoint: &Endpoint, connect_timeout: Duration) -> Result<TcpS
275274
})
276275
}
277276

278-
/// Run the TLS handshake over an established TCP stream.
277+
/// Wrap an established TCP stream in Thrift's TLS channel and run the handshake.
279278
#[cfg(feature = "tls")]
280-
fn tls_handshake(
279+
fn tls_channel(
281280
endpoint: &Endpoint,
282-
mut stream: TcpStream,
281+
stream: TcpStream,
283282
tls: &TlsOptions,
284-
) -> Result<StreamOwned<ClientConnection, TcpStream>> {
283+
) -> Result<TTlsClientChannel> {
285284
let config = tls_client_config(tls)?;
286285
let domain = tls.domain_override.as_deref().unwrap_or(&endpoint.host);
287286
let server_name = ServerName::try_from(domain.to_owned())
288287
.map_err(|e| Error::Client(format!("invalid TLS server name '{domain}': {e}")))?;
289-
let mut connection =
290-
ClientConnection::new(config, server_name).map_err(|e| Error::Tls(e.to_string()))?;
291-
292-
connection
293-
.complete_io(&mut stream)
294-
.map_err(|e| Error::Tls(e.to_string()))?;
295-
if connection.is_handshaking() {
296-
return Err(Error::Tls("TLS handshake did not complete".into()));
297-
}
298-
299-
Ok(StreamOwned::new(connection, stream))
288+
TTlsClientChannel::with_stream(stream, server_name, config)
289+
.map_err(|e| Error::Tls(e.to_string()))
300290
}
301291

302292
#[cfg(feature = "tls")]
@@ -445,46 +435,6 @@ impl ServerCertVerifier for NoCertificateVerification {
445435
}
446436
}
447437

448-
/// A `TlsStream` shared between the read and write transports.
449-
///
450-
/// `TTcpChannel::split` clones the underlying OS socket, but a TLS stream
451-
/// cannot be split that way (record layer state is shared), so both framed
452-
/// transports hold the same stream behind a mutex. The generated sync
453-
/// client fully writes + flushes a request before reading the response, so
454-
/// read and write never contend.
455-
#[cfg(feature = "tls")]
456-
#[derive(Clone)]
457-
struct SharedTlsStream(std::sync::Arc<std::sync::Mutex<StreamOwned<ClientConnection, TcpStream>>>);
458-
459-
#[cfg(feature = "tls")]
460-
impl SharedTlsStream {
461-
fn new(stream: StreamOwned<ClientConnection, TcpStream>) -> Self {
462-
Self(std::sync::Arc::new(std::sync::Mutex::new(stream)))
463-
}
464-
465-
fn lock(&self) -> std::sync::MutexGuard<'_, StreamOwned<ClientConnection, TcpStream>> {
466-
self.0.lock().unwrap_or_else(|p| p.into_inner())
467-
}
468-
}
469-
470-
#[cfg(feature = "tls")]
471-
impl Read for SharedTlsStream {
472-
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
473-
self.lock().read(buf)
474-
}
475-
}
476-
477-
#[cfg(feature = "tls")]
478-
impl Write for SharedTlsStream {
479-
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
480-
self.lock().write(buf)
481-
}
482-
483-
fn flush(&mut self) -> std::io::Result<()> {
484-
self.lock().flush()
485-
}
486-
}
487-
488438
#[cfg(test)]
489439
mod tests {
490440
use super::*;

0 commit comments

Comments
 (0)