#[cfg(feature = "tls")]
use std::io::{Read, Write};
#[cfg(feature = "tls")]
use std::sync::Arc;
#[cfg(feature = "tls")]
use sha2::{Digest, Sha224, Sha256, Sha384, Sha512};
use super::{AsyncTransport, BufferedTransport, TransportError};
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct TlsConfig {
pub mode: SslMode,
pub server_name: String,
pub ca_cert: Option<Vec<u8>>,
pub client_cert: Option<Vec<u8>>,
pub client_key: Option<Vec<u8>>,
pub accept_invalid_certs: bool,
#[cfg(feature = "tls")]
pub crypto_provider: Option<Arc<rustls::crypto::CryptoProvider>>,
}
impl TlsConfig {
pub fn new(mode: SslMode, server_name: impl Into<String>) -> Self {
Self {
mode,
server_name: server_name.into(),
..Default::default()
}
}
pub fn mode(mut self, mode: SslMode) -> Self {
self.mode = mode;
self
}
pub fn server_name(mut self, name: impl Into<String>) -> Self {
self.server_name = name.into();
self
}
pub fn ca_cert(mut self, cert: Vec<u8>) -> Self {
self.ca_cert = Some(cert);
self
}
pub fn client_cert(mut self, cert: Vec<u8>) -> Self {
self.client_cert = Some(cert);
self
}
pub fn client_key(mut self, key: Vec<u8>) -> Self {
self.client_key = Some(key);
self
}
pub fn accept_invalid_certs(mut self, accept: bool) -> Self {
self.accept_invalid_certs = accept;
self
}
}
impl Default for TlsConfig {
fn default() -> Self {
Self {
mode: SslMode::VerifyFull,
server_name: String::new(),
ca_cert: None,
client_cert: None,
client_key: None,
accept_invalid_certs: false,
#[cfg(feature = "tls")]
crypto_provider: None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
#[non_exhaustive]
pub enum SslMode {
Disable,
Prefer,
Require,
VerifyCa,
VerifyFull,
}
impl SslMode {
#[allow(clippy::should_implement_trait)]
pub fn from_str(s: &str) -> Result<Self, TransportError> {
match s.to_lowercase().as_str() {
"disable" => Ok(SslMode::Disable),
"prefer" => Ok(SslMode::Prefer),
"require" => Ok(SslMode::Require),
"verify-ca" => Ok(SslMode::VerifyCa),
"verify-full" => Ok(SslMode::VerifyFull),
_ => Err(TransportError::InvalidConfig(format!(
"invalid sslmode: {}",
s
))),
}
}
}
impl std::fmt::Display for SslMode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SslMode::Disable => write!(f, "disable"),
SslMode::Prefer => write!(f, "prefer"),
SslMode::Require => write!(f, "require"),
SslMode::VerifyCa => write!(f, "verify-ca"),
SslMode::VerifyFull => write!(f, "verify-full"),
}
}
}
#[cfg(feature = "tls")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum TlsVerificationPolicy {
InsecureNoVerify,
VerifyCaOnly,
VerifyFull,
}
#[cfg(feature = "tls")]
fn verification_policy(config: &TlsConfig) -> TlsVerificationPolicy {
if config.accept_invalid_certs || matches!(config.mode, SslMode::Require) {
TlsVerificationPolicy::InsecureNoVerify
} else if matches!(config.mode, SslMode::VerifyCa) {
TlsVerificationPolicy::VerifyCaOnly
} else {
TlsVerificationPolicy::VerifyFull
}
}
#[cfg(feature = "tls")]
fn build_rustls_config(config: &TlsConfig) -> Result<Arc<rustls::ClientConfig>, TransportError> {
use rustls::client::ClientConfig as RustlsClientConfig;
let crypto_provider = config
.crypto_provider
.clone()
.unwrap_or_else(default_crypto_provider);
let config_builder = RustlsClientConfig::builder_with_provider(crypto_provider.clone())
.with_protocol_versions(&[&rustls::version::TLS13, &rustls::version::TLS12])
.map_err(|e| {
TransportError::TlsHandshake(format!("unsupported protocol versions: {}", e))
})?;
let client_config = match verification_policy(config) {
TlsVerificationPolicy::InsecureNoVerify => config_builder
.dangerous()
.with_custom_certificate_verifier(Arc::new(NoVerifier))
.with_no_client_auth(),
TlsVerificationPolicy::VerifyCaOnly => {
let root_store = build_root_store(config)?;
let verifier = CertificateVerifier::new(
Arc::new(root_store),
crypto_provider.signature_verification_algorithms,
false,
);
config_builder
.dangerous()
.with_custom_certificate_verifier(Arc::new(verifier))
.with_no_client_auth()
}
TlsVerificationPolicy::VerifyFull => {
let root_store = build_root_store(config)?;
config_builder
.with_root_certificates(root_store)
.with_no_client_auth()
}
};
let mut client_config = client_config;
if let (Some(cert_bytes), Some(key_bytes)) = (&config.client_cert, &config.client_key) {
let certs = parse_certs(cert_bytes)?;
let key = parse_private_key(key_bytes)?;
let certified_key = rustls::sign::CertifiedKey::from_der(certs, key, &crypto_provider)
.map_err(|e| TransportError::TlsHandshake(format!("invalid client cert/key: {}", e)))?;
client_config.client_auth_cert_resolver =
Arc::new(rustls::sign::SingleCertAndKey::from(certified_key));
}
client_config.alpn_protocols.clear();
Ok(Arc::new(client_config))
}
#[cfg(feature = "tls")]
fn default_crypto_provider() -> Arc<rustls::crypto::CryptoProvider> {
Arc::new(rustls_rustcrypto::provider())
}
#[cfg(feature = "tls")]
fn build_root_store(config: &TlsConfig) -> Result<rustls::RootCertStore, TransportError> {
let mut root_store = rustls::RootCertStore::empty();
if let Some(ref ca_bytes) = config.ca_cert {
let certs = parse_certs(ca_bytes)?;
for cert in certs {
root_store.add(cert).map_err(|e| {
TransportError::TlsHandshake(format!("failed to add CA certificate: {}", e))
})?;
}
} else {
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
}
Ok(root_store)
}
#[cfg(feature = "tls")]
fn parse_certs(
bytes: &[u8],
) -> Result<Vec<rustls::pki_types::CertificateDer<'static>>, TransportError> {
let mut cursor = std::io::Cursor::new(bytes);
let pem_result: Result<Vec<_>, _> = rustls_pemfile::certs(&mut cursor).collect();
if let Ok(certs) = pem_result {
if !certs.is_empty() {
return Ok(certs);
}
}
Ok(vec![rustls::pki_types::CertificateDer::from(
bytes.to_vec(),
)])
}
#[cfg(feature = "tls")]
fn parse_private_key(
bytes: &[u8],
) -> Result<rustls::pki_types::PrivateKeyDer<'static>, TransportError> {
let mut cursor = std::io::Cursor::new(bytes);
if let Ok(Some(key)) = rustls_pemfile::private_key(&mut cursor) {
return Ok(key);
}
Ok(rustls::pki_types::PrivateKeyDer::Pkcs8(
rustls::pki_types::PrivatePkcs8KeyDer::from(bytes.to_vec()),
))
}
#[cfg(feature = "tls")]
fn compute_tls_server_end_point(
cert: &rustls::pki_types::CertificateDer<'_>,
) -> Result<Option<Vec<u8>>, TransportError> {
let (_, cert) = x509_parser::parse_x509_certificate(cert.as_ref()).map_err(|e| {
TransportError::TlsHandshake(format!("failed to parse peer certificate: {e}"))
})?;
let signature_oid = cert.signature_algorithm.algorithm.to_id_string();
let der = cert.as_ref();
let digest = match signature_oid.as_str() {
"1.2.840.113549.1.1.4" | "1.2.840.113549.1.1.5" | "1.2.840.10045.4.1" => {
Sha256::digest(der).to_vec()
}
"1.2.840.113549.1.1.14" | "1.2.840.10045.4.3.1" => Sha224::digest(der).to_vec(),
"1.2.840.113549.1.1.11" | "1.2.840.10045.4.3.2" => Sha256::digest(der).to_vec(),
"1.2.840.113549.1.1.12" | "1.2.840.10045.4.3.3" => Sha384::digest(der).to_vec(),
"1.2.840.113549.1.1.13" | "1.2.840.10045.4.3.4" => Sha512::digest(der).to_vec(),
_ => return Ok(None),
};
Ok(Some(digest))
}
#[cfg(feature = "tls")]
#[derive(Debug)]
struct CertificateVerifier {
roots: Arc<rustls::RootCertStore>,
supported_algs: rustls::crypto::WebPkiSupportedAlgorithms,
verify_hostname: bool,
}
#[cfg(feature = "tls")]
impl CertificateVerifier {
fn new(
roots: Arc<rustls::RootCertStore>,
supported_algs: rustls::crypto::WebPkiSupportedAlgorithms,
verify_hostname: bool,
) -> Self {
Self {
roots,
supported_algs,
verify_hostname,
}
}
}
#[cfg(feature = "tls")]
impl rustls::client::danger::ServerCertVerifier for CertificateVerifier {
fn verify_server_cert(
&self,
end_entity: &rustls::pki_types::CertificateDer<'_>,
intermediates: &[rustls::pki_types::CertificateDer<'_>],
server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
let cert = rustls::server::ParsedCertificate::try_from(end_entity)?;
rustls::client::verify_server_cert_signed_by_trust_anchor(
&cert,
&self.roots,
intermediates,
now,
self.supported_algs.all,
)?;
if self.verify_hostname {
rustls::client::verify_server_name(&cert, server_name)?;
}
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(message, cert, dss, &self.supported_algs)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(message, cert, dss, &self.supported_algs)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
self.supported_algs.supported_schemes()
}
}
#[cfg(feature = "tls")]
#[derive(Debug)]
struct NoVerifier;
#[cfg(feature = "tls")]
impl rustls::client::danger::ServerCertVerifier for NoVerifier {
fn verify_server_cert(
&self,
_end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
vec![
rustls::SignatureScheme::ECDSA_NISTP256_SHA256,
rustls::SignatureScheme::ECDSA_NISTP384_SHA384,
rustls::SignatureScheme::ED25519,
rustls::SignatureScheme::RSA_PKCS1_SHA256,
rustls::SignatureScheme::RSA_PKCS1_SHA384,
rustls::SignatureScheme::RSA_PKCS1_SHA512,
rustls::SignatureScheme::RSA_PSS_SHA256,
rustls::SignatureScheme::RSA_PSS_SHA384,
]
}
}
#[cfg(feature = "tls")]
pub struct TlsTransport<T: AsyncTransport> {
tls_conn: rustls::ClientConnection,
inner: T,
tls_server_end_point: Option<Vec<u8>>,
}
#[cfg(feature = "tls")]
impl<T: AsyncTransport> TlsTransport<T> {
pub async fn handshake(
inner: T,
config: Arc<rustls::ClientConfig>,
server_name: &str,
) -> Result<Self, TransportError> {
let server_name = rustls::pki_types::ServerName::try_from(server_name.to_string())
.map_err(|e| {
TransportError::TlsHandshake(format!(
"invalid server name '{}': {}",
server_name, e
))
})?;
let mut tls_conn = rustls::ClientConnection::new(config, server_name).map_err(|e| {
TransportError::TlsHandshake(format!("TLS connection creation failed: {}", e))
})?;
let mut inner = inner;
let mut handshake_buf = [0u8; 8192];
let mut iterations = 0;
const MAX_HANDSHAKE_ITERATIONS: u32 = 100;
loop {
iterations += 1;
if iterations > MAX_HANDSHAKE_ITERATIONS {
return Err(TransportError::TlsHandshake(
"TLS handshake did not complete within iteration limit".into(),
));
}
let mut outgoing = Vec::new();
tls_conn
.write_tls(&mut outgoing)
.map_err(|e| TransportError::TlsHandshake(format!("write_tls: {}", e)))?;
if !outgoing.is_empty() {
inner.write_all(&outgoing).await?;
inner.flush().await?;
}
if !tls_conn.is_handshaking() {
break;
}
let n = inner.read(&mut handshake_buf).await?;
if n == 0 {
return Err(TransportError::UnexpectedEof);
}
let bytes_read = tls_conn
.read_tls(&mut &handshake_buf[..n])
.map_err(|e| TransportError::TlsHandshake(format!("read_tls: {}", e)))?;
if bytes_read == 0 {
return Err(TransportError::TlsHandshake(
"TLS handshake stalled: no data consumed".into(),
));
}
tls_conn
.process_new_packets()
.map_err(|e| TransportError::TlsHandshake(format!("process_new_packets: {}", e)))?;
}
if tls_conn.is_handshaking() {
return Err(TransportError::TlsHandshake(
"TLS handshake incomplete after loop exit".into(),
));
}
let tls_server_end_point = tls_conn
.peer_certificates()
.and_then(|certs| certs.first())
.map(compute_tls_server_end_point)
.transpose()?
.flatten();
Ok(TlsTransport {
tls_conn,
inner,
tls_server_end_point,
})
}
pub fn protocol_version(&self) -> Option<rustls::ProtocolVersion> {
self.tls_conn.protocol_version()
}
pub fn negotiated_cipher_suite(&self) -> Option<rustls::SupportedCipherSuite> {
self.tls_conn.negotiated_cipher_suite()
}
pub fn peer_certificate(&self) -> Option<rustls::pki_types::CertificateDer<'static>> {
self.tls_conn
.peer_certificates()
.and_then(|certs| certs.first())
.cloned()
}
async fn flush_tls_outgoing(&mut self) -> Result<(), TransportError> {
let mut outgoing = Vec::new();
self.tls_conn
.write_tls(&mut outgoing)
.map_err(|e| TransportError::TlsHandshake(format!("write_tls: {}", e)))?;
if !outgoing.is_empty() {
self.inner.write_all(&outgoing).await?;
}
Ok(())
}
}
#[cfg(feature = "tls")]
impl<T: AsyncTransport> AsyncTransport for TlsTransport<T> {
fn is_secure(&self) -> bool {
true
}
fn tls_server_end_point(&self) -> Option<Vec<u8>> {
self.tls_server_end_point.clone()
}
async fn read(&mut self, buf: &mut [u8]) -> Result<usize, TransportError> {
loop {
match self.tls_conn.reader().read(buf) {
Ok(n) => {
if n == 0 {
return Ok(0);
}
return Ok(n);
}
Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {}
Err(ref e) if e.kind() == std::io::ErrorKind::ConnectionAborted => {
return Ok(0);
}
Err(e) => {
return Err(TransportError::TlsHandshake(format!("TLS read: {}", e)));
}
}
let mut cipher_buf = [0u8; 8192];
let n = self.inner.read(&mut cipher_buf).await?;
if n == 0 {
return Err(TransportError::UnexpectedEof);
}
self.tls_conn
.read_tls(&mut &cipher_buf[..n])
.map_err(|e| TransportError::TlsHandshake(format!("read_tls: {}", e)))?;
self.tls_conn
.process_new_packets()
.map_err(|e| TransportError::TlsHandshake(format!("process_new_packets: {}", e)))?;
}
}
async fn write(&mut self, buf: &[u8]) -> Result<usize, TransportError> {
let n = self
.tls_conn
.writer()
.write(buf)
.map_err(|e| TransportError::TlsHandshake(format!("TLS write: {}", e)))?;
self.flush_tls_outgoing().await?;
Ok(n)
}
async fn write_all(&mut self, buf: &[u8]) -> Result<(), TransportError> {
let mut written = 0;
while written < buf.len() {
let n = self
.tls_conn
.writer()
.write(&buf[written..])
.map_err(|e| TransportError::TlsHandshake(format!("TLS write: {}", e)))?;
written += n;
}
self.flush_tls_outgoing().await?;
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.tls_conn
.writer()
.flush()
.map_err(|e| TransportError::TlsHandshake(format!("TLS flush: {}", e)))?;
self.flush_tls_outgoing().await?;
self.inner.flush().await
}
async fn shutdown(&mut self) -> Result<(), TransportError> {
self.tls_conn.send_close_notify();
self.flush_tls_outgoing().await?;
self.inner.shutdown().await
}
}
#[allow(clippy::large_enum_variant)]
pub enum PgTransport<T: AsyncTransport> {
Plain(BufferedTransport<T>),
#[cfg(feature = "tls")]
Tls(BufferedTransport<TlsTransport<T>>),
}
impl<T: AsyncTransport> AsyncTransport for PgTransport<T> {
fn is_secure(&self) -> bool {
#[cfg(feature = "tls")]
{
self.is_tls()
}
#[cfg(not(feature = "tls"))]
{
false
}
}
fn tls_server_end_point(&self) -> Option<Vec<u8>> {
match self {
Self::Plain(t) => t.tls_server_end_point(),
#[cfg(feature = "tls")]
Self::Tls(t) => t.tls_server_end_point(),
}
}
async fn read(&mut self, buf: &mut [u8]) -> Result<usize, TransportError> {
match self {
Self::Plain(t) => t.read(buf).await,
#[cfg(feature = "tls")]
Self::Tls(t) => t.read(buf).await,
}
}
async fn write(&mut self, buf: &[u8]) -> Result<usize, TransportError> {
match self {
Self::Plain(t) => t.write(buf).await,
#[cfg(feature = "tls")]
Self::Tls(t) => t.write(buf).await,
}
}
async fn write_all(&mut self, buf: &[u8]) -> Result<(), TransportError> {
match self {
Self::Plain(t) => t.write_all(buf).await,
#[cfg(feature = "tls")]
Self::Tls(t) => t.write_all(buf).await,
}
}
async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), TransportError> {
match self {
Self::Plain(t) => t.read_exact(buf).await,
#[cfg(feature = "tls")]
Self::Tls(t) => t.read_exact(buf).await,
}
}
async fn flush(&mut self) -> Result<(), TransportError> {
match self {
Self::Plain(t) => t.flush().await,
#[cfg(feature = "tls")]
Self::Tls(t) => t.flush().await,
}
}
async fn shutdown(&mut self) -> Result<(), TransportError> {
match self {
Self::Plain(t) => t.shutdown().await,
#[cfg(feature = "tls")]
Self::Tls(t) => t.shutdown().await,
}
}
}
impl<T: AsyncTransport> PgTransport<T> {
#[cfg(feature = "tls")]
pub fn is_tls(&self) -> bool {
matches!(self, Self::Tls(_))
}
#[cfg(not(feature = "tls"))]
pub fn is_tls(&self) -> bool {
false
}
#[cfg(feature = "tls")]
pub fn tls_info(&self) -> Option<TlsInfo> {
match self {
Self::Tls(t) => {
let inner = t.inner();
Some(TlsInfo {
protocol_version: inner.protocol_version().map(|v| format!("{:?}", v)),
cipher_suite: inner.negotiated_cipher_suite().map(|v| format!("{:?}", v)),
peer_certificate: inner.peer_certificate().map(|c| c.to_vec()),
})
}
Self::Plain(_) => None,
}
}
#[cfg(not(feature = "tls"))]
pub fn tls_info(&self) -> Option<TlsInfo> {
None
}
}
#[derive(Debug)]
#[non_exhaustive]
pub struct TlsInfo {
pub protocol_version: Option<String>,
pub cipher_suite: Option<String>,
pub peer_certificate: Option<Vec<u8>>,
}
fn check_time_available() -> Result<(), TransportError> {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_err(|e| {
TransportError::TlsHandshake(format!(
"SystemTime::now() is not available on this platform. \
TLS certificate validation requires the current time. \
Error: {}",
e
))
})?;
Ok(())
}
#[cfg(feature = "tls")]
pub async fn negotiate_tls<T: AsyncTransport>(
tcp: T,
config: &TlsConfig,
) -> Result<PgTransport<T>, TransportError> {
const MAX_SSL_NEGOTIATION_ERROR_LEN: usize = 8 * 1024;
check_time_available()?;
let mut tcp = tcp;
let ssl_request: [u8; 8] = [
0x00, 0x00, 0x00, 0x08, 0x04, 0xD2, 0x16, 0x2F, ];
tcp.write_all(&ssl_request).await?;
tcp.flush().await?;
let mut response = [0u8; 1];
tcp.read_exact(&mut response).await?;
match response[0] {
b'S' => {
let tls_config = build_rustls_config(config)?;
let tls = TlsTransport::handshake(tcp, tls_config, &config.server_name).await?;
Ok(PgTransport::Tls(BufferedTransport::new(tls)))
}
b'N' => match config.mode {
SslMode::Disable => Ok(PgTransport::Plain(BufferedTransport::new(tcp))),
SslMode::Prefer => {
#[cfg(feature = "tracing")]
tracing::warn!(
server_name = %config.server_name,
"Server does not support TLS; falling back to plaintext"
);
Ok(PgTransport::Plain(BufferedTransport::new(tcp)))
}
SslMode::Require | SslMode::VerifyCa | SslMode::VerifyFull => {
Err(TransportError::TlsNotSupported)
}
},
b'E' => {
let mut len_buf = [0u8; 4];
tcp.read_exact(&mut len_buf).await?;
let len = i32::from_be_bytes(len_buf);
if len < 4 {
return Err(TransportError::TlsHandshake(
"server sent malformed error response during SSL negotiation".into(),
));
}
let len = len as usize;
if len > MAX_SSL_NEGOTIATION_ERROR_LEN {
return Err(TransportError::TlsHandshake(format!(
"server sent oversized SSL negotiation error response: {} bytes",
len
)));
}
let mut error_buf = vec![0u8; len - 4];
tcp.read_exact(&mut error_buf).await?;
Err(TransportError::TlsHandshake(format!(
"server rejected SSL request: {:?}",
String::from_utf8_lossy(&error_buf)
)))
}
other => Err(TransportError::TlsHandshake(format!(
"unexpected response byte during SSL negotiation: 0x{:02x} ('{}')",
other,
char::from_u32(other as u32).unwrap_or('?')
))),
}
}
#[cfg(not(feature = "tls"))]
pub async fn negotiate_tls<T: AsyncTransport>(
tcp: T,
config: &TlsConfig,
) -> Result<PgTransport<T>, TransportError> {
match config.mode {
SslMode::Disable => Ok(PgTransport::Plain(BufferedTransport::new(tcp))),
_ => Err(TransportError::TlsHandshake(
"TLS support is not compiled in. Enable the 'tls' feature flag.".into(),
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[allow(unused_imports)]
use crate::transport::MockTransport;
#[test]
fn test_ssl_mode_from_str() {
assert_eq!(SslMode::from_str("disable").unwrap(), SslMode::Disable);
assert_eq!(SslMode::from_str("prefer").unwrap(), SslMode::Prefer);
assert_eq!(SslMode::from_str("require").unwrap(), SslMode::Require);
assert_eq!(SslMode::from_str("verify-ca").unwrap(), SslMode::VerifyCa);
assert_eq!(
SslMode::from_str("verify-full").unwrap(),
SslMode::VerifyFull
);
assert!(matches!(
SslMode::from_str("invalid"),
Err(TransportError::InvalidConfig(_))
));
}
#[test]
fn test_ssl_mode_display() {
assert_eq!(SslMode::Disable.to_string(), "disable");
assert_eq!(SslMode::Prefer.to_string(), "prefer");
assert_eq!(SslMode::Require.to_string(), "require");
assert_eq!(SslMode::VerifyCa.to_string(), "verify-ca");
assert_eq!(SslMode::VerifyFull.to_string(), "verify-full");
}
#[cfg(feature = "tls")]
#[test]
fn test_verification_policy_matrix() {
let base = TlsConfig::new(SslMode::VerifyFull, "localhost");
assert_eq!(
verification_policy(&TlsConfig::new(SslMode::Require, "localhost")),
TlsVerificationPolicy::InsecureNoVerify
);
assert_eq!(
verification_policy(&TlsConfig::new(SslMode::VerifyCa, "localhost")),
TlsVerificationPolicy::VerifyCaOnly
);
assert_eq!(
verification_policy(&TlsConfig::new(SslMode::VerifyFull, "localhost")),
TlsVerificationPolicy::VerifyFull
);
assert_eq!(
verification_policy(&base.accept_invalid_certs(true)),
TlsVerificationPolicy::InsecureNoVerify
);
}
#[test]
fn test_check_time_available() {
assert!(check_time_available().is_ok());
}
#[test]
#[cfg(feature = "tls")]
fn test_parse_certs_der() {
let der = vec![0x30, 0x03, 0x01, 0x01, 0xFF]; let certs = parse_certs(&der).unwrap();
assert_eq!(certs.len(), 1);
}
#[test]
#[cfg(feature = "tls")]
fn test_parse_private_key_der() {
let der = vec![0x30, 0x03, 0x01, 0x01, 0xFF];
let key = parse_private_key(&der).unwrap();
assert!(matches!(key, rustls::pki_types::PrivateKeyDer::Pkcs8(_)));
}
#[tokio::test]
#[cfg(feature = "tls")]
async fn test_negotiate_tls_server_supports_ssl() {
use crate::transport::MockTransport;
let mock = MockTransport::new(vec![b'S']);
let config = TlsConfig {
mode: SslMode::Require,
server_name: "localhost".into(),
accept_invalid_certs: true,
..Default::default()
};
let result = negotiate_tls(mock, &config).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_negotiate_tls_server_rejects_ssl() {
use crate::transport::MockTransport;
let mock = MockTransport::new(vec![b'N']);
let config = TlsConfig {
mode: SslMode::Disable,
server_name: "localhost".into(),
..Default::default()
};
let result = negotiate_tls(mock, &config).await;
assert!(result.is_ok());
assert!(!result.unwrap().is_tls());
}
#[tokio::test]
#[cfg(feature = "tls")]
async fn test_negotiate_tls_server_rejects_ssl_require_mode() {
use crate::transport::MockTransport;
let mock = MockTransport::new(vec![b'N']);
let config = TlsConfig {
mode: SslMode::Require,
server_name: "localhost".into(),
..Default::default()
};
let result = negotiate_tls(mock, &config).await;
assert!(matches!(result, Err(TransportError::TlsNotSupported)));
}
#[tokio::test]
#[cfg(feature = "tls")]
async fn test_negotiate_tls_server_sends_error() {
use crate::transport::MockTransport;
let mut response = vec![b'E'];
response.extend_from_slice(&i32::to_be_bytes(8)); response.extend_from_slice(b"M\0test");
let mock = MockTransport::new(response);
let config = TlsConfig {
mode: SslMode::Require,
server_name: "localhost".into(),
..Default::default()
};
let result = negotiate_tls(mock, &config).await;
assert!(matches!(result, Err(TransportError::TlsHandshake(_))));
}
#[tokio::test]
#[cfg(feature = "tls")]
async fn test_negotiate_tls_unexpected_byte() {
use crate::transport::MockTransport;
let mock = MockTransport::new(vec![b'X']);
let config = TlsConfig {
mode: SslMode::Require,
server_name: "localhost".into(),
..Default::default()
};
let result = negotiate_tls(mock, &config).await;
assert!(matches!(result, Err(TransportError::TlsHandshake(_))));
}
#[test]
fn test_pg_transport_is_tls_without_feature() {
let mock = MockTransport::new(vec![]);
let buf = BufferedTransport::new(mock);
let pg = PgTransport::Plain(buf);
assert!(!pg.is_tls());
assert!(pg.tls_info().is_none());
}
}