diff --git a/grpc/Cargo.toml b/grpc/Cargo.toml index ca9f3dd22..5553df011 100644 --- a/grpc/Cargo.toml +++ b/grpc/Cargo.toml @@ -56,6 +56,7 @@ bytes = "1.10.1" futures = { version = "0.3", default-features = false, optional = true } hickory-resolver = { version = "0.26.1", optional = true } http = "1.1.0" +httparse = "1.9" http-body = "1.0.1" hyper = { version = "1.6.0", features = ["client", "http2"] } hyper-util = { version = "0.1", features = ["client-proxy"] } diff --git a/grpc/src/client/channel.rs b/grpc/src/client/channel.rs index 87928b00d..bd9c6b58f 100644 --- a/grpc/src/client/channel.rs +++ b/grpc/src/client/channel.rs @@ -64,6 +64,7 @@ use crate::client::name_resolution::ResolverUpdate; use crate::client::name_resolution::Target; use crate::client::name_resolution::dns; use crate::client::name_resolution::global_registry; +use crate::client::name_resolution::proxy_resolver; use crate::client::name_resolution::{self}; use crate::client::service_config::LbPolicyType; use crate::client::service_config::ServiceConfig; @@ -207,6 +208,7 @@ impl ChannelBuilder { // TODO(nathanielford): Return errors here instead of panicking. let target = Target::from_str(self.target.as_str()).unwrap(); let resolver_builder = global_registry().get(target.scheme()).unwrap(); + let resolver_builder = proxy_resolver::Builder::new_arc(resolver_builder); let authority = self .authority diff --git a/grpc/src/client/name_resolution/proxy_resolver.rs b/grpc/src/client/name_resolution/proxy_resolver.rs index 6e6488794..74446e48e 100644 --- a/grpc/src/client/name_resolution/proxy_resolver.rs +++ b/grpc/src/client/name_resolution/proxy_resolver.rs @@ -95,8 +95,13 @@ impl ResolverBuilder for Builder { impl Builder { /// Creates a new `Builder` that wraps the given `child_builder`. - pub(crate) fn new(child_builder: Arc) -> Self { - Self { child_builder } + pub(crate) fn new_arc(child_builder: Arc) -> Arc { + // Skip proxy lookup for non-DNS targets. + if child_builder.scheme() != "dns" { + return child_builder; + } + + Arc::new(Self { child_builder }) } fn new_resolver( @@ -105,10 +110,6 @@ impl Builder { options: ResolverOptions, matcher: Option<&Matcher>, ) -> Box { - // Skip proxy lookup for non-DNS targets. - if target.scheme() != "dns" { - return self.child_builder.build(target, options); - } // If HTTPS_PROXY is unset, avoid parsing the target as a DNS hostname. let Some(matcher) = matcher else { return self.child_builder.build(target, options); @@ -353,7 +354,7 @@ mod tests { matcher: Option<&Matcher>, child_builder: Arc, ) -> Vec
{ - let builder = Builder::new(child_builder); + let builder = Builder { child_builder }; let target: Target = target_uri.parse().unwrap(); let (work_scheduler, mut work_rx) = TestWorkScheduler::new_pair(); @@ -475,7 +476,7 @@ mod tests { let target_uri = "dns:///var%20/run/grpc.sock"; let child_builder = Arc::new(MockResolverBuilder {}); - let builder = Builder::new(child_builder); + let builder = Builder { child_builder }; let target: Target = target_uri.parse().unwrap(); let (work_scheduler, mut work_rx) = TestWorkScheduler::new_pair(); @@ -503,35 +504,15 @@ mod tests { assert!(err.contains("invalid target host in URL")); } + #[cfg(unix)] #[tokio::test] async fn unix_path_bypass() { - let matcher = Matcher::builder() - .https("http://proxy.example.com:8080") - .build(); - - // Proxy lookup for unix targets should be skipped. - let addresses = run_resolver_and_get_addresses( - "unix:///var%20/run/grpc.sock", - vec!["127.0.0.1".parse().unwrap()], - Some(&matcher), - ) - .await; - - assert_eq!(addresses.len(), 1); - assert_eq!(&*addresses[0].address, DIRECT_ADDRESS); - assert!(ProxyOptions::from_addr(&addresses[0]).is_none()); - - // Check for abstract-unix scheme. - let addresses = run_resolver_and_get_addresses( - "unix-abstract:grpc.sock", - vec!["127.0.0.1".parse().unwrap()], - Some(&matcher), - ) - .await; - - assert_eq!(addresses.len(), 1); - assert_eq!(&*addresses[0].address, DIRECT_ADDRESS); - assert!(ProxyOptions::from_addr(&addresses[0]).is_none()); + crate::client::name_resolution::unix::reg(); + let unix_builder = crate::client::name_resolution::global_registry() + .get("unix") + .expect("unix resolver not registered"); + let proxy_builder = Builder::new_arc(unix_builder.clone()); + assert!(Arc::ptr_eq(&unix_builder, &proxy_builder)); } #[tokio::test] diff --git a/grpc/src/client/subchannel.rs b/grpc/src/client/subchannel.rs index 874371aff..c59ef6b9e 100644 --- a/grpc/src/client/subchannel.rs +++ b/grpc/src/client/subchannel.rs @@ -52,8 +52,10 @@ use crate::client::load_balancing::subchannel::private::Sealed; use crate::client::name_resolution::Address; use crate::client::stream_util::FailingRecvStream; use crate::client::transport::DynTransport; +use crate::client::transport::ProxyOptions; use crate::client::transport::SecurityOpts; use crate::client::transport::TransportOptions; +use crate::client::transport::http_connect::HttpConnectHandshaker; use crate::core::RequestHeaders; use crate::credentials::call::CallDetails; use crate::credentials::call::ClientConnectionSecurityInfo as CallClientConnectionSecurityInfo; @@ -306,11 +308,17 @@ impl InternalSubchannel { transport: Arc, backoff: Arc, runtime: GrpcRuntime, - security_opts: SecurityOpts, + mut security_opts: SecurityOpts, work_queue: WorkQueueTx, ) -> Arc { let on_drop = Arc::new(Notify::new()); let address_string = address.address.to_string(); + if let Some(proxy_opts) = ProxyOptions::from_addr(&address) { + security_opts.credentials = Arc::new(HttpConnectHandshaker::new( + security_opts.credentials, + proxy_opts, + )); + } let this = Arc::new_cyclic(|weak_self| Self { address, on_drop: on_drop.clone(), diff --git a/grpc/src/client/transport/http_connect/mod.rs b/grpc/src/client/transport/http_connect/mod.rs new file mode 100644 index 000000000..deac88919 --- /dev/null +++ b/grpc/src/client/transport/http_connect/mod.rs @@ -0,0 +1,188 @@ +/* + * + * Copyright 2026 gRPC authors. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to + * deal in the Software without restriction, including without limitation the + * rights to use, copy, modify, merge, publish, distribute, sublicense, and/or + * sell copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING + * FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS + * IN THE SOFTWARE. + * + */ + +use std::sync::Arc; + +use bytes::Bytes; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tonic::async_trait; + +use crate::client::transport::ProxyOptions; +use crate::client::transport::http_connect::rewind::Rewind; +use crate::credentials::ChannelCredentials; +use crate::credentials::ProtocolInfo; +use crate::credentials::call::CallCredentials; +use crate::credentials::client::ClientHandshakeInfo; +use crate::credentials::client::HandshakeOutput; +use crate::credentials::common::Authority; +use crate::private; +use crate::rt::BoxEndpoint; +use crate::rt::EndpointIoStream; +use crate::rt::GrpcEndpoint; +use crate::rt::GrpcRuntime; +use crate::rt::StreamEndpoint; + +mod rewind; + +/// Performs the HTTP CONNECT handshake on the given endpoint. +/// +/// This function sends an HTTP CONNECT request to the proxy specified in `opts`, +/// reads the response, and returns a new endpoint that yields any buffered data +/// read during the handshake before delegating to the original endpoint. +async fn do_connect_handshake( + input: I, + opts: &ProxyOptions, +) -> Result, String> { + let mut io = EndpointIoStream::new(input); + + let mut req = format!( + "CONNECT {} HTTP/1.1\r\nHost: {}\r\n", + opts.target_authority(), + opts.target_authority() + ) + .into_bytes(); + if let Some(creds) = opts.proxy_authorization_header() { + req.extend_from_slice(b"Proxy-Authorization: "); + req.extend_from_slice(creds.as_bytes()); + req.extend_from_slice(b"\r\n"); + } + req.extend_from_slice(b"\r\n"); // headers end + + io.write_all(&req).await.map_err(|e| e.to_string())?; + io.flush().await.map_err(|e| e.to_string())?; + + const READ_BUF_SIZE: usize = 8192; + let mut buf = vec![0u8; READ_BUF_SIZE]; + let mut read = 0; + + // Read the response. + loop { + let n = io.read(&mut buf[read..]).await.map_err(|e| e.to_string())?; + if n == 0 { + return Err("Connection closed by proxy".to_string()); + } + read += n; + + // Allocate space on the stack to read up to 16 headers from the proxy. + let mut headers = [httparse::EMPTY_HEADER; 16]; + let mut res = httparse::Response::new(&mut headers); + match res.parse(&buf[..read]) { + Ok(httparse::Status::Complete(len)) => { + if res.code != Some(200) { + return Err(format!("Proxy returned status {}", res.code.unwrap_or(0))); + } + // Success! + let remaining = read - len; + let buffered_data = if remaining > 0 { + Some(Bytes::copy_from_slice(&buf[len..read])) + } else { + None + }; + + let local_addr = io + .get_ref() + .get_local_address() + .to_string() + .into_boxed_str(); + let peer_addr = io.get_ref().get_peer_address().to_string().into_boxed_str(); + let network_type = io.get_ref().get_network_type(); + // Check for buffered data. In most cases, the buffer should be + // empty as the server waits for the client to send the first + // message, e.g. in TLS. + let endpoint = if let Some(data) = buffered_data { + Rewind::new_buffered(io, data) + } else { + Rewind::new_unbuffered(io) + }; + return Ok(StreamEndpoint::new( + endpoint, + local_addr, + peer_addr, + network_type, + )); + } + Ok(httparse::Status::Partial) => { + if read >= READ_BUF_SIZE { + return Err("Response too large".to_string()); + } + } + Err(e) => { + return Err(format!("Failed to parse HTTP response: {}", e)); + } + } + } +} + +/// A credential wrapper that performs an HTTP CONNECT handshake before +/// delegating to an inner security credential (like TLS). +pub(crate) struct HttpConnectHandshaker { + inner: Arc, + options: ProxyOptions, +} + +impl HttpConnectHandshaker { + /// Constructs a new `HttpConnectHandshaker` wrapping the inner credentials. + pub(crate) fn new(inner: Arc, options: &ProxyOptions) -> Self { + Self { + inner, + options: options.clone(), + } + } +} + +/// The I/O stream wrapper returned after the HTTP CONNECT handshake succeeds. +type ProxyStream = StreamEndpoint>>; + +#[async_trait] +impl ChannelCredentials for HttpConnectHandshaker { + fn info(&self) -> &ProtocolInfo { + self.inner.info() + } + + fn get_call_credentials(&self, token: private::Internal) -> Option<&Arc> { + self.inner.get_call_credentials(token) + } + + async fn connect( + &self, + authority: &Authority, + source: BoxEndpoint, + info: &ClientHandshakeInfo, + runtime: &GrpcRuntime, + token: private::Internal, + ) -> Result { + // Perform the HTTP CONNECT handshake. + let proxied_stream = do_connect_handshake(source, &self.options).await?; + + // Delegate the actual security handshake (e.g., TLS) to the wrapped + // credentials. + self.inner + .connect(authority, Box::new(proxied_stream), info, runtime, token) + .await + } +} + +#[cfg(test)] +mod test; diff --git a/grpc/src/client/transport/http_connect/rewind.rs b/grpc/src/client/transport/http_connect/rewind.rs new file mode 100644 index 000000000..dfce7b4fe --- /dev/null +++ b/grpc/src/client/transport/http_connect/rewind.rs @@ -0,0 +1,119 @@ +// Copyright (c) 2023-2025 Sean McArthur +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to deal +// in the Software without restriction, including without limitation the rights +// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +// copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in +// all copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +// THE SOFTWARE. +// +// Note: This file contains modifications by the gRPC authors. +// See git revision history for subsequent changes and details. +// Original source: https://github.com/hyperium/hyper-util/blob/911d0f256342fa6740721104418c6102dc7d1d0b/src/common/rewind.rs +// +// Modifications: +// - Add a constructor for creating objects without a buffer. +// - Format imports. + +use std::cmp; +use std::io; +use std::pin::Pin; +use std::task::Context; +use std::task::Poll; + +use bytes::Buf; +use bytes::Bytes; +use tokio::io::AsyncRead; +use tokio::io::AsyncWrite; +use tokio::io::ReadBuf; + +/// Combine a buffer with an IO, rewinding reads to use the buffer. +#[derive(Debug)] +pub(crate) struct Rewind { + pub(crate) pre: Option, + pub(crate) inner: T, +} + +impl Rewind { + pub(crate) fn new_buffered(io: T, buf: Bytes) -> Self { + Rewind { + pre: Some(buf), + inner: io, + } + } + pub(crate) fn new_unbuffered(io: T) -> Self { + Rewind { + pre: None, + inner: io, + } + } +} + +impl AsyncRead for Rewind +where + T: AsyncRead + Unpin, +{ + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + if let Some(mut prefix) = self.pre.take() + && !prefix.is_empty() + { + let copy_len = cmp::min(prefix.len(), buf.remaining()); + buf.put_slice(&prefix[..copy_len]); + prefix.advance(copy_len); + if !prefix.is_empty() { + self.pre = Some(prefix); + } + + return Poll::Ready(Ok(())); + } + Pin::new(&mut self.inner).poll_read(cx, buf) + } +} + +impl AsyncWrite for Rewind +where + T: AsyncWrite + Unpin, +{ + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.inner).poll_write(cx, buf) + } + + fn poll_write_vectored( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + bufs: &[io::IoSlice<'_>], + ) -> Poll> { + Pin::new(&mut self.inner).poll_write_vectored(cx, bufs) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_shutdown(cx) + } + + fn is_write_vectored(&self) -> bool { + self.inner.is_write_vectored() + } +} diff --git a/grpc/src/client/transport/http_connect/test.rs b/grpc/src/client/transport/http_connect/test.rs new file mode 100644 index 000000000..71a508e84 --- /dev/null +++ b/grpc/src/client/transport/http_connect/test.rs @@ -0,0 +1,481 @@ +/* + * + * Copyright 2026 gRPC authors. + * + * Permission is hereby granted, free of charge, to any person obtaining a copy + * of this software and associated documentation files (the "Software"), to + * deal in the Software without restriction, including without limitation the + * rights to use, copy, modify, merge, publish, distribute, sublicense, and/or + * sell copies of the Software, and to permit persons to whom the Software is + * furnished to do so, subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in + * all copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, + * FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE + * AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER + * LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING + * FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS + * IN THE SOFTWARE. + * + */ + +use std::fs; +use std::net::SocketAddr; +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::Once; + +use http::HeaderValue; +use rustls::crypto::ring; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; + +use crate::client::transport::ProxyOptions; +use crate::client::transport::http_connect::HttpConnectHandshaker; +use crate::credentials::ChannelCredentials; +use crate::credentials::LocalChannelCredentials; +use crate::credentials::ServerCredentials; +use crate::credentials::client::ClientHandshakeInfo; +use crate::credentials::client::HandshakeOutput; +use crate::credentials::common::Authority; +use crate::credentials::rustls::Identity; +use crate::credentials::rustls::RootCertificates; +use crate::credentials::rustls::StaticProvider; +use crate::credentials::rustls::client::ClientTlsConfig; +use crate::credentials::rustls::client::RustlsChannelCredentials; +use crate::credentials::rustls::server::RustlsServerCredentials; +use crate::credentials::rustls::server::ServerTlsConfig; +use crate::private; +use crate::rt::EndpointIoStream; +use crate::rt::GrpcRuntime; +use crate::rt::StreamEndpoint; +use crate::rt::tokio::TokioRuntime; + +static INIT: Once = Once::new(); + +fn init_provider() { + INIT.call_once(|| { + let _ = ring::default_provider().install_default(); + }); +} + +fn tls_credentials() -> (RustlsServerCredentials, RustlsChannelCredentials) { + init_provider(); + + let certs_path = PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .parent() + .unwrap() + .join("examples/data/tls"); + + let server_cert = fs::read(certs_path.join("server.pem")).expect("failed to read server.pem"); + let server_key = fs::read(certs_path.join("server.key")).expect("failed to read server.key"); + let ca_cert = fs::read(certs_path.join("ca.pem")).expect("failed to read ca.pem"); + + let identity = Identity::from_pem(server_cert, server_key); + let identity_provider = StaticProvider::new(vec![identity]); + let server_tls_config = ServerTlsConfig::new(identity_provider); + let server_creds = RustlsServerCredentials::new(server_tls_config).unwrap(); + + let root_certs = RootCertificates::from_pem(ca_cert); + let root_provider = StaticProvider::new(root_certs); + let tls_client_config = ClientTlsConfig::new().with_root_certificates_provider(root_provider); + let rustls_creds = RustlsChannelCredentials::new(tls_client_config).unwrap(); + + (server_creds, rustls_creds) +} + +async fn run_mock_tls_server(listener: TcpListener, creds: RustlsServerCredentials) { + let (stream, _) = listener.accept().await.unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); + let runtime = GrpcRuntime::new(TokioRuntime::default()); + let handshake_res = creds + .accept(stream, runtime, private::Internal) + .await + .unwrap(); + + let mut tls_stream = EndpointIoStream::new(handshake_res.endpoint); + let mut buf = vec![0u8; 5]; + tls_stream.read_exact(&mut buf).await.unwrap(); + assert_eq!(buf, b"hello"); + tls_stream.write_all(b"world").await.unwrap(); +} + +async fn perform_handshake( + handshaker: &HttpConnectHandshaker, + proxy_addr: SocketAddr, + target_port: u16, +) -> Result { + let source = TcpStream::connect(proxy_addr).await.unwrap(); + let endpoint = StreamEndpoint::new_from_tcp(source).unwrap(); + + let info = ClientHandshakeInfo::default(); + let runtime = GrpcRuntime::new(TokioRuntime::default()); + let authority = Authority::new("localhost".to_string(), Some(target_port)); + + handshaker + .connect( + &authority, + Box::new(endpoint), + &info, + &runtime, + private::Internal, + ) + .await +} + +async fn verify_client_handshake( + handshaker: &HttpConnectHandshaker, + proxy_addr: SocketAddr, + server_port: u16, +) { + let handshake_output = perform_handshake(handshaker, proxy_addr, server_port) + .await + .unwrap(); + + let mut client_stream = EndpointIoStream::new(handshake_output.endpoint); + client_stream.write_all(b"hello").await.unwrap(); + let mut buf = vec![0u8; 5]; + client_stream.read_exact(&mut buf).await.unwrap(); + assert_eq!(buf, b"world"); +} + +async fn run_client_handshake_fail( + handshaker: &HttpConnectHandshaker, + proxy_addr: SocketAddr, +) -> String { + match perform_handshake(handshaker, proxy_addr, 12345).await { + Ok(_) => panic!("Expected connection failure"), + Err(e) => e, + } +} + +#[tokio::test] +async fn test_proxy_success_no_auth() { + let (server_creds, rustls_creds) = tls_credentials(); + + let server_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let server_addr = server_listener.local_addr().unwrap(); + + let server_handle = tokio::spawn(run_mock_tls_server(server_listener, server_creds)); + + // Start Proxy + let target_host = format!("localhost:{}", server_addr.port()); + let proxy_addr = spawn_proxy(ProxyConfig { + expected_host: target_host.clone(), + expected_auth: None, + connect_response: connect_success_response(), + target_addr: Some(server_addr), + }) + .await; + + // Connect Client via Proxy using HttpConnectHandshaker directly. + let proxy_options = ProxyOptions::new(target_host, None); + let handshaker = HttpConnectHandshaker::new(Arc::new(rustls_creds), &proxy_options); + + verify_client_handshake(&handshaker, proxy_addr, server_addr.port()).await; + server_handle.await.unwrap(); +} + +#[tokio::test] +async fn test_proxy_success_with_auth() { + let (server_creds, rustls_creds) = tls_credentials(); + + let server_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let server_addr = server_listener.local_addr().unwrap(); + + let server_handle = tokio::spawn(run_mock_tls_server(server_listener, server_creds)); + + // Start Proxy (expects Basic auth). + let target_host = format!("localhost:{}", server_addr.port()); + let expected_auth = "Basic dXNlcjpwYXNzd29yZA=="; // user:password + let proxy_addr = spawn_proxy(ProxyConfig { + expected_host: target_host.clone(), + expected_auth: Some(expected_auth.to_string()), + connect_response: connect_success_response(), + target_addr: Some(server_addr), + }) + .await; + + // Connect Client via Proxy with Auth Header. + let auth_header = HeaderValue::from_str(expected_auth).unwrap(); + let proxy_options = ProxyOptions::new(target_host, Some(auth_header)); + let handshaker = HttpConnectHandshaker::new(Arc::new(rustls_creds), &proxy_options); + + verify_client_handshake(&handshaker, proxy_addr, server_addr.port()).await; + server_handle.await.unwrap(); +} + +#[tokio::test] +async fn test_proxy_failure_large_header() { + let (_, rustls_creds) = tls_credentials(); + + let target_host = "localhost:12345".to_string(); + let proxy_addr = spawn_proxy(ProxyConfig { + expected_host: target_host.clone(), + expected_auth: None, + connect_response: connect_large_header_response(), + target_addr: None, // Will close after response + }) + .await; + + let proxy_options = ProxyOptions::new(target_host, None); + let handshaker = HttpConnectHandshaker::new(Arc::new(rustls_creds), &proxy_options); + + let err_msg = run_client_handshake_fail(&handshaker, proxy_addr).await; + assert!( + err_msg.contains("Response too large"), + "Expected 'Response too large', got: {}", + err_msg + ); +} + +#[tokio::test] +async fn test_proxy_failure_invalid_response() { + let (_, rustls_creds) = tls_credentials(); + + let target_host = "localhost:12345".to_string(); + let proxy_addr = spawn_proxy(ProxyConfig { + expected_host: target_host.clone(), + expected_auth: None, + connect_response: connect_invalid_response(), + target_addr: None, + }) + .await; + + let proxy_options = ProxyOptions::new(target_host, None); + let handshaker = HttpConnectHandshaker::new(Arc::new(rustls_creds), &proxy_options); + + let err_msg = run_client_handshake_fail(&handshaker, proxy_addr).await; + assert!( + err_msg.contains("Failed to parse HTTP response"), + "Expected parse error, got: {}", + err_msg + ); +} + +#[tokio::test] +async fn test_proxy_failure_bad_status() { + let (_, rustls_creds) = tls_credentials(); + + let target_host = "localhost:12345".to_string(); + let proxy_addr = spawn_proxy(ProxyConfig { + expected_host: target_host.clone(), + expected_auth: None, + connect_response: connect_status_response(502, "Bad Gateway"), + target_addr: None, + }) + .await; + + let proxy_options = ProxyOptions::new(target_host, None); + let handshaker = HttpConnectHandshaker::new(Arc::new(rustls_creds), &proxy_options); + + let err_msg = run_client_handshake_fail(&handshaker, proxy_addr).await; + assert!( + err_msg.contains("Proxy returned status 502"), + "Expected 'Proxy returned status 502', got: {}", + err_msg + ); +} + +#[tokio::test] +async fn test_proxy_success_local_rewind_batched() { + let local_creds = LocalChannelCredentials::new_arc(); + + // Start a mock target TCP server that sends the "later" part of the + // response. + let target_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let target_addr = target_listener.local_addr().unwrap(); + + let initial_bytes = b"initial_batched_bytes_"; + let later_bytes = b"later_server_bytes"; + + let target_handle = tokio::spawn(async move { + if let Ok((mut stream, _)) = target_listener.accept().await { + stream.write_all(later_bytes).await.unwrap(); + stream.flush().await.unwrap(); + } + }); + + let target_host = format!("localhost:{}", target_addr.port()); + let proxy_options = ProxyOptions::new(target_host.clone(), None); + let handshaker = HttpConnectHandshaker::new(local_creds, &proxy_options); + + // Start Proxy configured to batch initial_bytes and tunnel to target + // server. + let proxy_addr = spawn_proxy(ProxyConfig { + expected_host: target_host.clone(), + expected_auth: None, + connect_response: connect_batched_response(initial_bytes), + target_addr: Some(target_addr), + }) + .await; + + // Connect Client via Proxy. + let handshake_output = perform_handshake(&handshaker, proxy_addr, target_addr.port()) + .await + .unwrap(); + + let mut proxied_stream = EndpointIoStream::new(handshake_output.endpoint); + + // Verify we can read BOTH the batched initial_bytes and the subsequent + // later_bytes. + let mut expected_total = initial_bytes.to_vec(); + expected_total.extend_from_slice(later_bytes); + + let mut buf = vec![0u8; expected_total.len()]; + proxied_stream.read_exact(&mut buf).await.unwrap(); + + assert_eq!(buf, expected_total); + target_handle.await.unwrap(); +} + +// --- Helper Functions for Payloads --- + +fn connect_success_response() -> Vec { + b"HTTP/1.1 200 Connection Established\r\n\r\n".to_vec() +} + +fn connect_status_response(code: u16, phrase: &str) -> Vec { + format!("HTTP/1.1 {} {}\r\n\r\n", code, phrase).into_bytes() +} + +fn connect_large_header_response() -> Vec { + let mut res = b"HTTP/1.1 200 OK\r\n".to_vec(); + // Exceed the 8KB buffer limit of HttpConnectHandshaker + let large_header = vec![b'A'; 9000]; + res.extend_from_slice(b"X-Large: "); + res.extend_from_slice(&large_header); + res.extend_from_slice(b"\r\n\r\n"); + res +} + +fn connect_invalid_response() -> Vec { + b"NOT HTTP RESPONSE\r\n\r\n".to_vec() +} + +fn connect_batched_response(server_bytes: &[u8]) -> Vec { + let mut res = connect_success_response(); + res.extend_from_slice(server_bytes); + res +} + +// --- Mock Proxy Server --- + +struct ProxyConfig { + expected_host: String, + expected_auth: Option, + connect_response: Vec, + target_addr: Option, +} + +async fn spawn_proxy(config: ProxyConfig) -> SocketAddr { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + if let Ok((mut client_stream, _)) = listener.accept().await { + let mut buf = vec![0u8; 16384]; + let mut read = 0; + + // Read request + loop { + let n = client_stream.read(&mut buf[read..]).await.unwrap(); + if n == 0 { + return; // Client disconnected + } + read += n; + + let mut headers = [httparse::EMPTY_HEADER; 16]; + let mut req = httparse::Request::new(&mut headers); + match req.parse(&buf[..read]) { + Ok(httparse::Status::Complete(len)) => { + // Validate method is CONNECT + if req.method != Some("CONNECT") { + let res = b"HTTP/1.1 405 Method Not Allowed\r\n\r\n"; + client_stream.write_all(res).await.unwrap(); + return; + } + + // Validate path (host). + let path = req.path.unwrap(); + if path != config.expected_host { + let res = b"HTTP/1.1 400 Bad Request\r\n\r\n"; + client_stream.write_all(res).await.unwrap(); + return; + } + + // Validate Host header. + let mut host_ok = false; + for header in req.headers.iter() { + if header.name.eq_ignore_ascii_case("host") + && header.value == config.expected_host.as_bytes() + { + host_ok = true; + } + } + if !host_ok { + let res = b"HTTP/1.1 400 Bad Request\r\n\r\n"; + client_stream.write_all(res).await.unwrap(); + return; + } + + // Validate Auth if expected. + if let Some(ref expected_auth) = config.expected_auth { + let mut auth_ok = false; + for header in req.headers.iter() { + if header.name.eq_ignore_ascii_case("proxy-authorization") + && header.value == expected_auth.as_bytes() + { + auth_ok = true; + } + } + if !auth_ok { + let res = b"HTTP/1.1 407 Proxy Authentication Required\nProxy-Authenticate: Basic realm=\"proxy\"\r\n\r\n"; + client_stream.write_all(res).await.unwrap(); + return; + } + } + + // Send the configured response + client_stream + .write_all(&config.connect_response) + .await + .unwrap(); + client_stream.flush().await.unwrap(); + + // Tunnel if target_addr is Some. + if let Some(target_addr) = config.target_addr { + let mut backend_stream = match TcpStream::connect(target_addr).await { + Ok(s) => s, + Err(_) => { + return; + } + }; + let _ = tokio::io::copy_bidirectional( + &mut client_stream, + &mut backend_stream, + ) + .await; + } + return; + } + Ok(httparse::Status::Partial) => { + if read >= buf.len() { + return; // Too large + } + } + Err(_) => { + let res = b"HTTP/1.1 400 Bad Request\r\n\r\n"; + client_stream.write_all(res).await.unwrap(); + return; + } + } + } + } + }); + addr +} diff --git a/grpc/src/client/transport/mod.rs b/grpc/src/client/transport/mod.rs index 8bce87acb..9ba07616d 100644 --- a/grpc/src/client/transport/mod.rs +++ b/grpc/src/client/transport/mod.rs @@ -36,6 +36,7 @@ use crate::credentials::client::ClientHandshakeInfo; use crate::credentials::common::Authority; use crate::rt::GrpcRuntime; +pub(crate) mod http_connect; mod registry; // Using tower/buffer enables tokio's rt feature even though it's possible to diff --git a/grpc/src/credentials/dyn_wrapper.rs b/grpc/src/credentials/dyn_wrapper.rs index f6067c32e..779d93785 100644 --- a/grpc/src/credentials/dyn_wrapper.rs +++ b/grpc/src/credentials/dyn_wrapper.rs @@ -79,8 +79,8 @@ mod tests { use crate::credentials::LocalServerCredentials; use crate::credentials::SecurityLevel; use crate::rt; - use crate::rt::AsyncIoAdapter; - use crate::rt::tokio::TokioIoStream; + use crate::rt::EndpointIoStream; + use crate::rt::StreamEndpoint; #[tokio::test] async fn test_dyn_server_credential_dispatch() { @@ -106,7 +106,7 @@ mod tests { }); let (stream, _) = listener.accept().await.unwrap(); - let server_stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let server_stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let result = dyn_creds .dyn_accept(Box::new(server_stream) as Box, runtime) @@ -121,7 +121,7 @@ mod tests { assert_eq!(security_info.security_level(), SecurityLevel::NoSecurity); let mut buf = vec![0u8; 25]; - AsyncIoAdapter::new(endpoint) + EndpointIoStream::new(endpoint) .read_exact(&mut buf) .await .unwrap(); diff --git a/grpc/src/credentials/local.rs b/grpc/src/credentials/local.rs index 0a8905b5c..bd8b5f594 100644 --- a/grpc/src/credentials/local.rs +++ b/grpc/src/credentials/local.rs @@ -203,10 +203,10 @@ mod test { use crate::credentials::client::ClientHandshakeInfo; use crate::credentials::common::Authority; use crate::rt; - use crate::rt::AsyncIoAdapter; + use crate::rt::EndpointIoStream; use crate::rt::GrpcEndpoint; + use crate::rt::StreamEndpoint; use crate::rt::TcpOptions; - use crate::rt::tokio::TokioIoStream; #[test] fn test_security_level_for_endpoint_success() { @@ -281,7 +281,7 @@ mod test { server_stream.write_all(test_data).await.unwrap(); let mut buf = vec![0u8; test_data.len()]; - AsyncIoAdapter::new(endpoint) + EndpointIoStream::new(endpoint) .read_exact(&mut buf) .await .unwrap(); @@ -318,7 +318,7 @@ mod test { }); let (stream, _) = listener.accept().await.unwrap(); - let server_stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let server_stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let output = creds .accept(server_stream, runtime, private::Internal) @@ -331,7 +331,7 @@ mod test { assert_eq!(security_info.security_level(), SecurityLevel::NoSecurity); let mut buf = vec![0u8; 10]; - AsyncIoAdapter::new(endpoint) + EndpointIoStream::new(endpoint) .read_exact(&mut buf) .await .unwrap(); diff --git a/grpc/src/credentials/rustls/client/mod.rs b/grpc/src/credentials/rustls/client/mod.rs index 03e32909b..4fa92ebfa 100644 --- a/grpc/src/credentials/rustls/client/mod.rs +++ b/grpc/src/credentials/rustls/client/mod.rs @@ -57,8 +57,8 @@ use crate::credentials::rustls::parse_key; use crate::credentials::rustls::sanitize_crypto_provider; use crate::credentials::rustls::tls_stream::TlsStream; use crate::private; -use crate::rt::AsyncIoAdapter; use crate::rt::BoxEndpoint; +use crate::rt::EndpointIoStream; use crate::rt::GrpcRuntime; #[cfg(test)] @@ -205,7 +205,7 @@ impl RustlsChannelCredentials { let server_name = ServerName::try_from(authority.host()) .map_err(|e| format!("invalid authority: {}", e))? .to_owned(); - let input_io = AsyncIoAdapter::new(source); + let input_io = EndpointIoStream::new(source); let tls_stream = self .connector diff --git a/grpc/src/credentials/rustls/client/test.rs b/grpc/src/credentials/rustls/client/test.rs index 2cbcd866d..dfb2524c1 100644 --- a/grpc/src/credentials/rustls/client/test.rs +++ b/grpc/src/credentials/rustls/client/test.rs @@ -50,7 +50,7 @@ use crate::credentials::rustls::client::ClientTlsConfig; use crate::credentials::rustls::client::RustlsChannelCredentials; use crate::private; use crate::rt; -use crate::rt::AsyncIoAdapter; +use crate::rt::EndpointIoStream; use crate::rt::TcpOptions; static INIT: Once = Once::new(); @@ -181,7 +181,7 @@ async fn test_tls_key_log() { .expect("Handshake failed"); let stream = result.endpoint; let mut buf = Vec::new(); - let _ = AsyncIoAdapter::new(stream).read_to_end(&mut buf).await; + let _ = EndpointIoStream::new(stream).read_to_end(&mut buf).await; assert_eq!(buf, b"Hello world"); server_task.await.unwrap(); @@ -339,7 +339,7 @@ async fn test_mtls_handshake_no_identity() { let stream = result.endpoint; let mut buf = Vec::new(); - let res = AsyncIoAdapter::new(stream).read_to_end(&mut buf).await; + let res = EndpointIoStream::new(stream).read_to_end(&mut buf).await; assert!( res.is_err(), "read from TLS stream should fail due to missing client identity" @@ -387,7 +387,7 @@ async fn test_mtls_handshake_with_identitiy() { let stream = result.endpoint; let mut buf = Vec::new(); - let _ = AsyncIoAdapter::new(stream).read_to_end(&mut buf).await; + let _ = EndpointIoStream::new(stream).read_to_end(&mut buf).await; assert_eq!(buf, b"Hello world"); server_task.await.unwrap(); @@ -449,7 +449,9 @@ async fn check_client_resumption_disabled( ); let mut buf = Vec::new(); - let _ = AsyncIoAdapter::new(tls_stream).read_to_end(&mut buf).await; + let _ = EndpointIoStream::new(tls_stream) + .read_to_end(&mut buf) + .await; assert_eq!(buf, b"Hello world"); } @@ -614,7 +616,7 @@ async fn run_handshake_test(server_alpn: Vec>, expect_success: bool) { let stream = result.endpoint; let mut buf = Vec::new(); // Ignore read errors if server closed connection abruptly (which happens in failure cases, but here we expect success) - let _ = AsyncIoAdapter::new(stream).read_to_end(&mut buf).await; + let _ = EndpointIoStream::new(stream).read_to_end(&mut buf).await; assert_eq!(buf, b"Hello world"); } else { assert!(result.is_err(), "Handshake succeeded but expected failure"); diff --git a/grpc/src/credentials/rustls/server/mod.rs b/grpc/src/credentials/rustls/server/mod.rs index b851eb89d..3e99c93f5 100644 --- a/grpc/src/credentials/rustls/server/mod.rs +++ b/grpc/src/credentials/rustls/server/mod.rs @@ -55,7 +55,7 @@ use crate::credentials::rustls::tls_stream::TlsStream; use crate::credentials::server::HandshakeOutput; use crate::credentials::server::ServerConnectionSecurityInfo; use crate::private; -use crate::rt::AsyncIoAdapter; +use crate::rt::EndpointIoStream; use crate::rt::GrpcEndpoint; use crate::rt::GrpcRuntime; @@ -350,7 +350,7 @@ impl ServerCredentials for RustlsServerCredentials { _runtime: GrpcRuntime, _token: private::Internal, ) -> Result>, String> { - let input_io = AsyncIoAdapter::new(source); + let input_io = EndpointIoStream::new(source); let tls_stream = self .acceptor .accept(input_io) diff --git a/grpc/src/credentials/rustls/server/test.rs b/grpc/src/credentials/rustls/server/test.rs index 9a079ceb2..04c710cfd 100644 --- a/grpc/src/credentials/rustls/server/test.rs +++ b/grpc/src/credentials/rustls/server/test.rs @@ -45,8 +45,8 @@ use crate::credentials::rustls::server::RustlsServerCredentials; use crate::credentials::rustls::server::ServerTlsConfig; use crate::credentials::rustls::server::TlsClientCertificateRequestType; use crate::private; -use crate::rt::AsyncIoAdapter; -use crate::rt::tokio::TokioIoStream; +use crate::rt::EndpointIoStream; +use crate::rt::StreamEndpoint; use crate::rt::{self}; static INIT: Once = Once::new(); @@ -73,14 +73,14 @@ async fn test_tls_server_handshake() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let result = creds.accept(stream, runtime, private::Internal).await; assert!( result.is_ok(), "Server handshake failed: {:?}", result.err() ); - let mut stream = AsyncIoAdapter::new(result.unwrap().endpoint); + let mut stream = EndpointIoStream::new(result.unwrap().endpoint); let mut buf = [0u8; 5]; stream.read_exact(&mut buf).await.unwrap(); assert_eq!(&buf, b"ping!"); @@ -129,7 +129,7 @@ async fn test_tls_server_handshake_no_alpn() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let result = creds.accept(stream, runtime, private::Internal).await; assert!(result.is_err(), "Server handshake should have failed"); }); @@ -172,7 +172,7 @@ async fn test_tls_server_handshake_bad_alpn() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let runtime = rt::default_runtime(); let result = creds.accept(stream, runtime, private::Internal).await; assert!(result.is_err(), "Server handshake should have failed"); @@ -209,7 +209,7 @@ async fn test_tls_handshake_alpn_h1_and_h2() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let runtime = rt::default_runtime(); creds .accept(stream, runtime, private::Internal) @@ -256,7 +256,7 @@ async fn test_tls_server_mtls_require_fail() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let result = creds.accept(stream, runtime, private::Internal).await; assert!(result.is_err(), "Handshake should fail without client cert"); }); @@ -307,12 +307,12 @@ async fn test_tls_server_mtls_success() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let result = creds .accept(stream, runtime, private::Internal) .await .expect("Server handshake failed"); - let mut stream = AsyncIoAdapter::new(result.endpoint); + let mut stream = EndpointIoStream::new(result.endpoint); let mut buf = [0u8; 5]; stream.read_exact(&mut buf).await.unwrap(); assert_eq!(&buf, b"ping!"); @@ -367,12 +367,12 @@ async fn test_tls_server_mtls_optional() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let result = creds .accept(stream, runtime, private::Internal) .await .expect("Server handshake failed"); - let mut stream = AsyncIoAdapter::new(result.endpoint); + let mut stream = EndpointIoStream::new(result.endpoint); let mut buf = [0u8; 5]; stream.read_exact(&mut buf).await.unwrap(); assert_eq!(&buf, b"ping!"); @@ -417,12 +417,12 @@ async fn test_tls_server_key_log() { let server_task = tokio::spawn(async move { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let result = creds .accept(stream, runtime, private::Internal) .await .expect("Server handshake failed"); - let mut stream = AsyncIoAdapter::new(result.endpoint); + let mut stream = EndpointIoStream::new(result.endpoint); let mut buf = [0u8; 5]; stream.read_exact(&mut buf).await.unwrap(); assert_eq!(&buf, b"ping!"); @@ -471,12 +471,12 @@ async fn check_resumption_disabled(versions: Vec<&'static rustls::SupportedProto let server_task = tokio::spawn(async move { for _ in 0..2 { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let runtime = rt::default_runtime(); let result = creds.accept(stream, runtime, private::Internal).await; assert!(result.is_ok()); let stream = result.unwrap().endpoint; - AsyncIoAdapter::new(stream) + EndpointIoStream::new(stream) .write_all(b"pong!") .await .unwrap(); @@ -547,7 +547,7 @@ async fn test_tls_server_sni() { let server_task = tokio::spawn(async move { for _ in 0..2 { let (stream, _) = listener.accept().await.unwrap(); - let stream = TokioIoStream::new_from_tcp(stream).unwrap(); + let stream = StreamEndpoint::new_from_tcp(stream).unwrap(); let runtime = rt::default_runtime(); let result = creds.accept(stream, runtime, private::Internal).await; assert!( @@ -555,7 +555,7 @@ async fn test_tls_server_sni() { "Server handshake failed: {:?}", result.err() ); - let mut stream = AsyncIoAdapter::new(result.unwrap().endpoint); + let mut stream = EndpointIoStream::new(result.unwrap().endpoint); let mut buf = [0u8; 5]; stream.read_exact(&mut buf).await.unwrap(); assert_eq!(&buf, b"ping!"); diff --git a/grpc/src/credentials/rustls/tls_stream.rs b/grpc/src/credentials/rustls/tls_stream.rs index c56b7dd4d..bf15cee99 100644 --- a/grpc/src/credentials/rustls/tls_stream.rs +++ b/grpc/src/credentials/rustls/tls_stream.rs @@ -33,11 +33,11 @@ use tokio::io::ReadBuf; use tokio_rustls::TlsStream as RustlsStream; use crate::private; -use crate::rt::AsyncIoAdapter; +use crate::rt::EndpointIoStream; use crate::rt::GrpcEndpoint; pub struct TlsStream { - inner: RustlsStream>, + inner: RustlsStream>, } impl GrpcEndpoint for TlsStream @@ -119,11 +119,11 @@ where } impl TlsStream { - pub(crate) fn new(inner: RustlsStream>) -> Self { + pub(crate) fn new(inner: RustlsStream>) -> Self { Self { inner } } - pub(crate) fn inner(&self) -> &RustlsStream> { + pub(crate) fn inner(&self) -> &RustlsStream> { &self.inner } } diff --git a/grpc/src/rt/mod.rs b/grpc/src/rt/mod.rs index fb1b3166f..0c73af8b7 100644 --- a/grpc/src/rt/mod.rs +++ b/grpc/src/rt/mod.rs @@ -195,11 +195,11 @@ pub trait GrpcEndpoint: Send + Unpin + 'static { /// An adapter that exposes `AsyncRead` and `AsyncWrite` functionality for /// interfacing with `hyper` and `rustls`. This type is kept private to avoid /// exposing its read and write methods to external crates. -pub(crate) struct AsyncIoAdapter { +pub(crate) struct EndpointIoStream { inner: T, } -impl AsyncIoAdapter { +impl EndpointIoStream { pub(crate) fn new(inner: T) -> Self { Self { inner } } @@ -209,7 +209,7 @@ impl AsyncIoAdapter { } } -impl AsyncRead for AsyncIoAdapter { +impl AsyncRead for EndpointIoStream { fn poll_read( mut self: Pin<&mut Self>, cx: &mut Context<'_>, @@ -219,7 +219,7 @@ impl AsyncRead for AsyncIoAdapter { } } -impl AsyncWrite for AsyncIoAdapter { +impl AsyncWrite for EndpointIoStream { fn poll_write( mut self: Pin<&mut Self>, cx: &mut Context<'_>, @@ -249,6 +249,91 @@ impl AsyncWrite for AsyncIoAdapter { } } +/// A wrapper that implements [GrpcEndpoint] for an asynchronous I/O stream. +pub(crate) struct StreamEndpoint { + inner: T, + peer_addr: Box, + local_addr: Box, + network_type: &'static str, +} + +impl StreamEndpoint { + pub(crate) fn new( + inner: T, + local_addr: Box, + peer_addr: Box, + network_type: &'static str, + ) -> Self { + Self { + inner, + peer_addr, + local_addr, + network_type, + } + } +} + +impl GrpcEndpoint for StreamEndpoint { + fn get_local_address(&self) -> &str { + &self.local_addr + } + + fn get_peer_address(&self) -> &str { + &self.peer_addr + } + + fn get_network_type(&self) -> &'static str { + self.network_type + } + + fn poll_read_private( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + _token: private::Internal, + ) -> Poll> { + Pin::new(&mut self.inner).poll_read(cx, buf) + } + + fn poll_write_private( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + _token: private::Internal, + ) -> Poll> { + Pin::new(&mut self.inner).poll_write(cx, buf) + } + + fn poll_write_vectored_private( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + bufs: &[io::IoSlice<'_>], + _token: private::Internal, + ) -> Poll> { + Pin::new(&mut self.inner).poll_write_vectored(cx, bufs) + } + + fn is_write_vectored_private(&self, _token: private::Internal) -> bool { + self.inner.is_write_vectored() + } + + fn poll_flush_private( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + _token: private::Internal, + ) -> Poll> { + Pin::new(&mut self.inner).poll_flush(cx) + } + + fn poll_shutdown_private( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + _token: private::Internal, + ) -> Poll> { + Pin::new(&mut self.inner).poll_shutdown(cx) + } +} + impl GrpcEndpoint for Box { fn get_local_address(&self) -> &str { (**self).get_local_address() diff --git a/grpc/src/rt/tokio/mod.rs b/grpc/src/rt/tokio/mod.rs index ebde593d3..40074ba7d 100644 --- a/grpc/src/rt/tokio/mod.rs +++ b/grpc/src/rt/tokio/mod.rs @@ -28,13 +28,10 @@ use std::net::SocketAddr; use std::pin::Pin; use std::time::Duration; -use tokio::io::AsyncRead; -use tokio::io::AsyncWrite; use tokio::net::TcpStream; use tokio::task::JoinHandle; use crate::client::name_resolution::TCP_IP_NETWORK_TYPE; -use crate::private; use crate::rt::BoxEndpoint; use crate::rt::BoxFuture; use crate::rt::BoxedTaskHandle; @@ -42,6 +39,7 @@ use crate::rt::DnsResolver; use crate::rt::ResolverOptions; use crate::rt::Runtime; use crate::rt::Sleep; +use crate::rt::StreamEndpoint; use crate::rt::TaskHandle; use crate::rt::TcpOptions; @@ -125,7 +123,7 @@ impl Runtime for TokioRuntime { .map_err(|err| err.to_string())?; } let stream: Box = - Box::new(TokioIoStream::new_from_tcp(stream)?); + Box::new(StreamEndpoint::new_from_tcp(stream)?); Ok(stream) }) } @@ -147,7 +145,7 @@ impl Runtime for TokioRuntime { let peer_addr = stream.peer_addr().map_err(|err| err.to_string())?; let local_addr = stream.local_addr().map_err(|err| err.to_string())?; - let stream: Box = Box::new(TokioIoStream { + let stream: Box = Box::new(StreamEndpoint { peer_addr: format!("{peer_addr:?}").into_boxed_str(), local_addr: format!("{local_addr:?}").into_boxed_str(), network_type: UNIX_NETWORK_TYPE, @@ -166,17 +164,9 @@ impl TokioDefaultDnsResolver { Ok(TokioDefaultDnsResolver { _priv: () }) } } - -pub(crate) struct TokioIoStream { - inner: T, - peer_addr: Box, - local_addr: Box, - network_type: &'static str, -} - -impl TokioIoStream { +impl StreamEndpoint { pub(crate) fn new_from_tcp(stream: TcpStream) -> Result { - Ok(TokioIoStream { + Ok(StreamEndpoint { local_addr: stream .local_addr() .map_err(|err| err.to_string())? @@ -193,67 +183,6 @@ impl TokioIoStream { } } -impl super::GrpcEndpoint for TokioIoStream { - fn get_local_address(&self) -> &str { - &self.local_addr - } - - fn get_peer_address(&self) -> &str { - &self.peer_addr - } - - fn get_network_type(&self) -> &'static str { - self.network_type - } - - fn poll_read_private( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - buf: &mut tokio::io::ReadBuf<'_>, - _token: private::Internal, - ) -> std::task::Poll> { - Pin::new(&mut self.inner).poll_read(cx, buf) - } - - fn poll_write_private( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - buf: &[u8], - _token: private::Internal, - ) -> std::task::Poll> { - Pin::new(&mut self.inner).poll_write(cx, buf) - } - - fn poll_write_vectored_private( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - bufs: &[std::io::IoSlice<'_>], - _token: private::Internal, - ) -> std::task::Poll> { - Pin::new(&mut self.inner).poll_write_vectored(cx, bufs) - } - - fn is_write_vectored_private(&self, _token: private::Internal) -> bool { - self.inner.is_write_vectored() - } - - fn poll_flush_private( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - _token: private::Internal, - ) -> std::task::Poll> { - Pin::new(&mut self.inner).poll_flush(cx) - } - - fn poll_shutdown_private( - mut self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - _token: private::Internal, - ) -> std::task::Poll> { - Pin::new(&mut self.inner).poll_shutdown(cx) - } -} - #[cfg(test)] mod tests { use super::DnsResolver;