use std::io;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::net::TcpStream;
use tokio::time::timeout_at;
use tracing::{debug, info};
use super::connection::{ConnectionConfig, ProxyConfig};
use crate::error::{KrafkaError, Result};
#[cfg(feature = "test-broker")]
pub(crate) type Connector =
std::sync::Arc<dyn Fn(&str) -> io::Result<tokio::io::DuplexStream> + Send + Sync>;
pub(crate) enum BrokerStream {
Tcp(TcpStream),
#[cfg(feature = "test-broker")]
Memory(tokio::io::DuplexStream),
}
impl BrokerStream {
pub(crate) fn into_tcp(self) -> Result<TcpStream> {
match self {
Self::Tcp(stream) => Ok(stream),
#[cfg(feature = "test-broker")]
Self::Memory(_) => Err(KrafkaError::config(
"an in-memory connection carries no TLS",
)),
}
}
}
impl AsyncRead for BrokerStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_read(cx, buf),
#[cfg(feature = "test-broker")]
Self::Memory(s) => Pin::new(s).poll_read(cx, buf),
}
}
}
impl AsyncWrite for BrokerStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_write(cx, buf),
#[cfg(feature = "test-broker")]
Self::Memory(s) => Pin::new(s).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_flush(cx),
#[cfg(feature = "test-broker")]
Self::Memory(s) => Pin::new(s).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_shutdown(cx),
#[cfg(feature = "test-broker")]
Self::Memory(s) => Pin::new(s).poll_shutdown(cx),
}
}
}
pub(crate) async fn dial(address: &str, config: &ConnectionConfig) -> Result<BrokerStream> {
#[cfg(feature = "test-broker")]
if let Some(connector) = &config.connector {
return connector(address)
.map(BrokerStream::Memory)
.map_err(KrafkaError::network);
}
let stream = match &config.proxy {
Some(proxy) => connect_via_proxy(address, proxy, config).await?,
None => super::happy_eyeballs::connect_happy_eyeballs(address, config).await?,
};
stream.set_nodelay(config.nodelay)?;
Ok(BrokerStream::Tcp(stream))
}
#[cfg(feature = "oauth-oidc")]
pub(crate) async fn dial_http(host: &str, port: u16) -> io::Result<TcpStream> {
TcpStream::connect((host, port)).await
}
pub(super) async fn connect_via_proxy(
address: &str,
proxy: &ProxyConfig,
config: &ConnectionConfig,
) -> Result<TcpStream> {
use tokio_socks::tcp::Socks5Stream;
debug!("Connecting to {address} via SOCKS5 proxy {}", proxy.address);
let deadline = tokio::time::Instant::now() + config.connect_timeout;
let addrs: Vec<std::net::SocketAddr> =
timeout_at(deadline, tokio::net::lookup_host(&proxy.address))
.await
.map_err(|_| KrafkaError::timeout("SOCKS5 proxy DNS resolution"))?
.map_err(KrafkaError::network)?
.collect();
let Some(&proxy_addr) = addrs.first() else {
return Err(KrafkaError::unavailable(format!(
"no addresses resolved for SOCKS5 proxy '{}'",
proxy.address
)));
};
let socket = super::happy_eyeballs::create_socket(proxy_addr, config)?;
let proxy_stream = timeout_at(deadline, async {
let tcp = socket
.connect(proxy_addr)
.await
.map_err(KrafkaError::network)?;
let socks = if let Some(ref creds) = proxy.credentials {
Socks5Stream::connect_with_password_and_socket(
tcp,
address,
creds.username(),
creds.password(),
)
.await
} else {
Socks5Stream::connect_with_socket(tcp, address).await
}
.map_err(|e| {
KrafkaError::network(std::io::Error::other(format!("SOCKS5 proxy error: {e}")))
})?;
Ok::<_, KrafkaError>(socks.into_inner())
})
.await
.map_err(|_| KrafkaError::timeout("SOCKS5 proxy connection"))??;
info!(
"SOCKS5 tunnel established to {address} via {}",
proxy.address
);
Ok(proxy_stream)
}