@@ -35,15 +35,14 @@ use rustls::crypto::WebPkiSupportedAlgorithms;
3535#[ cfg( feature = "tls" ) ]
3636use 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
4340use thrift:: protocol:: {
4441 TBinaryInputProtocol , TBinaryOutputProtocol , TCompactInputProtocol , TCompactOutputProtocol ,
4542 TInputProtocol , TOutputProtocol ,
4643} ;
44+ #[ cfg( feature = "tls" ) ]
45+ use thrift:: transport:: TTlsClientChannel ;
4746use thrift:: transport:: { TFramedReadTransport , TFramedWriteTransport , TIoChannel , TTcpChannel } ;
4847
4948use 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) ]
489439mod tests {
490440 use super :: * ;
0 commit comments