mod buffered;
mod error;
mod params;
#[cfg(feature = "tls")]
pub mod tls;
#[cfg(not(feature = "tls"))]
mod tls;
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
mod native;
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
mod tokio_tcp;
#[cfg(target_arch = "wasm32")]
mod tcp;
#[cfg(target_arch = "wasm32")]
pub use tcp::connect_with_timeout;
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
pub use tokio_tcp::connect_with_timeout;
#[allow(unused_imports)]
pub use buffered::BufferedTransport;
pub use error::TransportError;
#[allow(unused_imports)]
pub use params::ConnectionParams;
pub use tls::{negotiate_tls, PgTransport, SslMode, TlsConfig, TlsInfo};
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
#[allow(unused_imports)]
pub use native::NativeTcpTransport;
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
#[allow(unused_imports)]
pub use tokio_tcp::TokioTcpTransport;
#[cfg(target_arch = "wasm32")]
#[allow(unused_imports)]
pub use tcp::WasiTcpTransport;
#[non_exhaustive]
#[derive(Debug)]
pub enum ClientTransport {
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
Native(NativeTcpTransport),
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
Tokio(TokioTcpTransport),
#[cfg(target_arch = "wasm32")]
Wasi(WasiTcpTransport),
#[cfg(test)]
Mock(MockTransport),
}
impl AsyncTransport for ClientTransport {
#[inline]
async fn read(&mut self, _buf: &mut [u8]) -> Result<usize, TransportError> {
match self {
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
ClientTransport::Native(t) => t.read(_buf).await,
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
ClientTransport::Tokio(t) => t.read(_buf).await,
#[cfg(target_arch = "wasm32")]
ClientTransport::Wasi(t) => t.read(_buf).await,
#[cfg(test)]
ClientTransport::Mock(t) => t.read(_buf).await,
#[cfg(not(any(
all(not(target_arch = "wasm32"), feature = "test-native"),
all(not(target_arch = "wasm32"), feature = "tokio-transport"),
target_arch = "wasm32",
test
)))]
_ => unreachable!("no transport enabled: enable 'tokio-transport' or 'test-native' feature, or compile for wasm32-wasip2"),
}
}
#[inline]
async fn write(&mut self, _buf: &[u8]) -> Result<usize, TransportError> {
match self {
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
ClientTransport::Native(t) => t.write(_buf).await,
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
ClientTransport::Tokio(t) => t.write(_buf).await,
#[cfg(target_arch = "wasm32")]
ClientTransport::Wasi(t) => t.write(_buf).await,
#[cfg(test)]
ClientTransport::Mock(t) => t.write(_buf).await,
#[cfg(not(any(
all(not(target_arch = "wasm32"), feature = "test-native"),
all(not(target_arch = "wasm32"), feature = "tokio-transport"),
target_arch = "wasm32",
test
)))]
_ => unreachable!("no transport enabled: enable 'tokio-transport' or 'test-native' feature, or compile for wasm32-wasip2"),
}
}
#[inline]
async fn write_all(&mut self, _buf: &[u8]) -> Result<(), TransportError> {
match self {
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
ClientTransport::Native(t) => t.write_all(_buf).await,
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
ClientTransport::Tokio(t) => t.write_all(_buf).await,
#[cfg(target_arch = "wasm32")]
ClientTransport::Wasi(t) => t.write_all(_buf).await,
#[cfg(test)]
ClientTransport::Mock(t) => t.write_all(_buf).await,
#[cfg(not(any(
all(not(target_arch = "wasm32"), feature = "test-native"),
all(not(target_arch = "wasm32"), feature = "tokio-transport"),
target_arch = "wasm32",
test
)))]
_ => unreachable!("no transport enabled: enable 'tokio-transport' or 'test-native' feature, or compile for wasm32-wasip2"),
}
}
#[inline]
async fn read_exact(&mut self, _buf: &mut [u8]) -> Result<(), TransportError> {
match self {
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
ClientTransport::Native(t) => t.read_exact(_buf).await,
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
ClientTransport::Tokio(t) => t.read_exact(_buf).await,
#[cfg(target_arch = "wasm32")]
ClientTransport::Wasi(t) => t.read_exact(_buf).await,
#[cfg(test)]
ClientTransport::Mock(t) => t.read_exact(_buf).await,
#[cfg(not(any(
all(not(target_arch = "wasm32"), feature = "test-native"),
all(not(target_arch = "wasm32"), feature = "tokio-transport"),
target_arch = "wasm32",
test
)))]
_ => unreachable!("no transport enabled: enable 'tokio-transport' or 'test-native' feature, or compile for wasm32-wasip2"),
}
}
#[inline]
async fn flush(&mut self) -> Result<(), TransportError> {
match self {
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
ClientTransport::Native(t) => t.flush().await,
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
ClientTransport::Tokio(t) => t.flush().await,
#[cfg(target_arch = "wasm32")]
ClientTransport::Wasi(t) => t.flush().await,
#[cfg(test)]
ClientTransport::Mock(t) => t.flush().await,
#[cfg(not(any(
all(not(target_arch = "wasm32"), feature = "test-native"),
all(not(target_arch = "wasm32"), feature = "tokio-transport"),
target_arch = "wasm32",
test
)))]
_ => unreachable!("no transport enabled: enable 'tokio-transport' or 'test-native' feature, or compile for wasm32-wasip2"),
}
}
#[inline]
async fn shutdown(&mut self) -> Result<(), TransportError> {
match self {
#[cfg(all(not(target_arch = "wasm32"), feature = "test-native"))]
ClientTransport::Native(t) => t.shutdown().await,
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
ClientTransport::Tokio(t) => t.shutdown().await,
#[cfg(target_arch = "wasm32")]
ClientTransport::Wasi(t) => t.shutdown().await,
#[cfg(test)]
ClientTransport::Mock(t) => t.shutdown().await,
#[cfg(not(any(
all(not(target_arch = "wasm32"), feature = "test-native"),
all(not(target_arch = "wasm32"), feature = "tokio-transport"),
target_arch = "wasm32",
test
)))]
_ => unreachable!("no transport enabled: enable 'tokio-transport' or 'test-native' feature, or compile for wasm32-wasip2"),
}
}
}
#[allow(async_fn_in_trait)]
pub trait AsyncTransport {
fn is_secure(&self) -> bool {
false
}
fn tls_server_end_point(&self) -> Option<Vec<u8>> {
None
}
async fn read(&mut self, buf: &mut [u8]) -> Result<usize, TransportError>;
async fn write(&mut self, buf: &[u8]) -> Result<usize, TransportError>;
async fn write_all(&mut self, buf: &[u8]) -> Result<(), TransportError>;
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), TransportError>;
async fn flush(&mut self) -> Result<(), TransportError>;
async fn shutdown(&mut self) -> Result<(), TransportError>;
}
#[cfg(test)]
#[derive(Debug)]
pub struct MockTransport {
read_data: Vec<u8>,
read_pos: usize,
max_read_chunk: usize,
pub written: Vec<u8>,
closed: bool,
pub flushed: bool,
pub shutdown_called: bool,
}
#[cfg(test)]
impl MockTransport {
pub fn new(read_data: Vec<u8>) -> Self {
Self {
read_data,
read_pos: 0,
max_read_chunk: 0,
written: Vec::new(),
closed: false,
flushed: false,
shutdown_called: false,
}
}
pub fn with_max_read_chunk(mut self, chunk: usize) -> Self {
self.max_read_chunk = chunk;
self
}
pub fn written(&self) -> &[u8] {
&self.written
}
pub fn is_read_exhausted(&self) -> bool {
self.read_pos >= self.read_data.len()
}
}
#[cfg(test)]
impl AsyncTransport for MockTransport {
fn tls_server_end_point(&self) -> Option<Vec<u8>> {
None
}
async fn read(&mut self, buf: &mut [u8]) -> Result<usize, TransportError> {
if self.closed {
return Ok(0);
}
if self.read_pos >= self.read_data.len() {
return Ok(0);
}
let remaining = &self.read_data[self.read_pos..];
let mut to_read = remaining.len().min(buf.len());
if self.max_read_chunk > 0 {
to_read = to_read.min(self.max_read_chunk);
}
buf[..to_read].copy_from_slice(&remaining[..to_read]);
self.read_pos += to_read;
Ok(to_read)
}
async fn write(&mut self, buf: &[u8]) -> Result<usize, TransportError> {
if self.closed {
return Err(TransportError::ConnectionReset);
}
self.written.extend_from_slice(buf);
Ok(buf.len())
}
async fn write_all(&mut self, buf: &[u8]) -> Result<(), TransportError> {
if self.closed {
return Err(TransportError::ConnectionReset);
}
self.written.extend_from_slice(buf);
Ok(())
}
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), TransportError> {
let mut filled = 0;
while filled < buf.len() {
let n = self.read(&mut buf[filled..]).await?;
if n == 0 {
return Err(TransportError::UnexpectedEof);
}
filled += n;
}
Ok(())
}
async fn flush(&mut self) -> Result<(), TransportError> {
self.flushed = true;
Ok(())
}
async fn shutdown(&mut self) -> Result<(), TransportError> {
self.shutdown_called = true;
self.closed = true;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_transport_error_classification() {
assert!(TransportError::ConnectionReset.is_connection_broken());
assert!(TransportError::UnexpectedEof.is_connection_broken());
assert!(TransportError::ConnectionRefused.is_connection_broken());
assert!(!TransportError::Timeout.is_connection_broken());
assert!(TransportError::Timeout.is_transient());
assert!(TransportError::DnsResolutionFailed {
host: "example.com".into()
}
.is_transient());
assert!(!TransportError::ConnectionReset.is_transient());
}
#[tokio::test]
async fn test_mock_transport_basic_read_write() {
let mut mock = MockTransport::new(vec![1, 2, 3, 4, 5]);
let mut buf = [0u8; 3];
assert_eq!(mock.read(&mut buf).await.unwrap(), 3);
assert_eq!(&buf, &[1, 2, 3]);
assert_eq!(mock.write(&[10, 11]).await.unwrap(), 2);
assert_eq!(mock.written(), &[10, 11]);
}
#[tokio::test]
async fn test_mock_transport_read_exact() {
let mut mock = MockTransport::new(vec![1, 2, 3, 4, 5]).with_max_read_chunk(2);
let mut buf = [0u8; 4];
mock.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, &[1, 2, 3, 4]);
}
#[tokio::test]
async fn test_mock_transport_partial_reads() {
let mut mock = MockTransport::new(vec![1, 2, 3, 4, 5]).with_max_read_chunk(2);
let mut buf = [0u8; 5];
assert_eq!(mock.read(&mut buf).await.unwrap(), 2);
assert_eq!(&buf[..2], &[1, 2]);
assert_eq!(mock.read(&mut buf[2..]).await.unwrap(), 2);
assert_eq!(&buf[..4], &[1, 2, 3, 4]);
assert_eq!(mock.read(&mut buf[4..]).await.unwrap(), 1);
assert_eq!(&buf, &[1, 2, 3, 4, 5]);
assert_eq!(mock.read(&mut buf).await.unwrap(), 0); }
#[tokio::test]
async fn test_mock_transport_eof_on_read_exact() {
let mut mock = MockTransport::new(vec![1, 2]);
let mut buf = [0u8; 5];
assert!(matches!(
mock.read_exact(&mut buf).await,
Err(TransportError::UnexpectedEof)
));
}
}