Skip to main content

pg_proto/
tls.rs

1//! TLS transport upgrades and RFC 5929 channel binding.
2
3use std::{
4    io,
5    pin::Pin,
6    sync::Arc,
7    task::{Context, Poll},
8};
9
10use rustls::{
11    ClientConfig, DigitallySignedStruct, Error as TlsError, RootCertStore, ServerConfig,
12    SignatureScheme,
13    client::{
14        WebPkiServerVerifier,
15        danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier},
16    },
17    crypto::{CryptoProvider, WebPkiSupportedAlgorithms},
18    pki_types::{CertificateDer, ServerName, UnixTime},
19};
20use sha2::{Digest, Sha256, Sha384, Sha512};
21use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
22use tokio_rustls::{TlsAcceptor, TlsConnector};
23use x509_parser::{
24    prelude::{FromDer, X509Certificate},
25    signature_algorithm::SignatureAlgorithm,
26};
27
28use crate::{
29    auth::TlsServerEndPoint,
30    pre_startup::{CertificateVerification, SslMode},
31};
32
33#[derive(Debug)]
34struct CertificateVerifier {
35    verification: CertificateVerification,
36    roots: RootCertStore,
37    supported: WebPkiSupportedAlgorithms,
38}
39
40impl ServerCertVerifier for CertificateVerifier {
41    fn verify_server_cert(
42        &self,
43        end_entity: &CertificateDer<'_>,
44        intermediates: &[CertificateDer<'_>],
45        server_name: &ServerName<'_>,
46        _ocsp_response: &[u8],
47        now: UnixTime,
48    ) -> Result<ServerCertVerified, TlsError> {
49        match self.verification {
50            CertificateVerification::None => Ok(ServerCertVerified::assertion()),
51            CertificateVerification::CertificateAuthority => {
52                let parsed = rustls::server::ParsedCertificate::try_from(end_entity)?;
53                rustls::client::verify_server_cert_signed_by_trust_anchor(
54                    &parsed,
55                    &self.roots,
56                    intermediates,
57                    now,
58                    self.supported.all,
59                )?;
60                Ok(ServerCertVerified::assertion())
61            }
62            CertificateVerification::CertificateAuthorityAndHost => {
63                let verifier = WebPkiServerVerifier::builder(Arc::new(self.roots.clone()))
64                    .build()
65                    .map_err(|error| TlsError::General(error.to_string()))?;
66                verifier.verify_server_cert(end_entity, intermediates, server_name, &[], now)
67            }
68        }
69    }
70
71    fn verify_tls12_signature(
72        &self,
73        message: &[u8],
74        cert: &CertificateDer<'_>,
75        dss: &DigitallySignedStruct,
76    ) -> Result<HandshakeSignatureValid, TlsError> {
77        rustls::crypto::verify_tls12_signature(message, cert, dss, &self.supported)
78    }
79
80    fn verify_tls13_signature(
81        &self,
82        message: &[u8],
83        cert: &CertificateDer<'_>,
84        dss: &DigitallySignedStruct,
85    ) -> Result<HandshakeSignatureValid, TlsError> {
86        rustls::crypto::verify_tls13_signature(message, cert, dss, &self.supported)
87    }
88
89    fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
90        self.supported.supported_schemes()
91    }
92}
93
94/// Builds a client TLS configuration with libpq-compatible `sslmode` verification.
95///
96/// `require`, `prefer`, and `allow` encrypt without checking the certificate chain;
97/// `verify-ca` checks the chain without checking the host; `verify-full` checks both.
98#[must_use]
99pub fn client_config(mode: SslMode, roots: RootCertStore) -> ClientConfig {
100    let verification = mode.strategy().verification;
101    if verification == CertificateVerification::CertificateAuthorityAndHost {
102        return ClientConfig::builder()
103            .with_root_certificates(roots)
104            .with_no_client_auth();
105    }
106
107    let provider = CryptoProvider::get_default()
108        .cloned()
109        .unwrap_or_else(|| Arc::new(rustls::crypto::aws_lc_rs::default_provider()));
110    let verifier = CertificateVerifier {
111        verification,
112        roots,
113        supported: provider.signature_verification_algorithms,
114    };
115    ClientConfig::builder()
116        .dangerous()
117        .with_custom_certificate_verifier(Arc::new(verifier))
118        .with_no_client_auth()
119}
120
121/// A client-side TLS stream carrying its `PostgreSQL` channel-binding value.
122#[derive(Debug)]
123pub struct ClientTls<S> {
124    inner: tokio_rustls::client::TlsStream<S>,
125    tls_server_end_point: Vec<u8>,
126}
127
128/// A server-side TLS stream carrying its `PostgreSQL` channel-binding value.
129#[derive(Debug)]
130pub struct ServerTls<S> {
131    inner: tokio_rustls::server::TlsStream<S>,
132    tls_server_end_point: Vec<u8>,
133}
134
135impl<S> TlsServerEndPoint for ClientTls<S> {
136    fn tls_server_end_point(&self) -> &[u8] {
137        &self.tls_server_end_point
138    }
139}
140
141impl<S> TlsServerEndPoint for ServerTls<S> {
142    fn tls_server_end_point(&self) -> &[u8] {
143        &self.tls_server_end_point
144    }
145}
146
147macro_rules! delegate_io {
148    ($wrapper:ident) => {
149        impl<S: AsyncRead + AsyncWrite + Unpin> AsyncRead for $wrapper<S> {
150            fn poll_read(
151                mut self: Pin<&mut Self>,
152                cx: &mut Context<'_>,
153                buffer: &mut ReadBuf<'_>,
154            ) -> Poll<io::Result<()>> {
155                Pin::new(&mut self.inner).poll_read(cx, buffer)
156            }
157        }
158
159        impl<S: AsyncRead + AsyncWrite + Unpin> AsyncWrite for $wrapper<S> {
160            fn poll_write(
161                mut self: Pin<&mut Self>,
162                cx: &mut Context<'_>,
163                buffer: &[u8],
164            ) -> Poll<Result<usize, io::Error>> {
165                Pin::new(&mut self.inner).poll_write(cx, buffer)
166            }
167
168            fn poll_flush(
169                mut self: Pin<&mut Self>,
170                cx: &mut Context<'_>,
171            ) -> Poll<Result<(), io::Error>> {
172                Pin::new(&mut self.inner).poll_flush(cx)
173            }
174
175            fn poll_shutdown(
176                mut self: Pin<&mut Self>,
177                cx: &mut Context<'_>,
178            ) -> Poll<Result<(), io::Error>> {
179                Pin::new(&mut self.inner).poll_shutdown(cx)
180            }
181        }
182    };
183}
184
185delegate_io!(ClientTls);
186delegate_io!(ServerTls);
187
188/// Negotiates TLS as a `PostgreSQL` client and records the peer certificate binding.
189///
190/// # Errors
191///
192/// Returns a TLS handshake, certificate, or channel-binding error.
193pub async fn connect<S>(
194    stream: S,
195    server_name: ServerName<'static>,
196    config: Arc<ClientConfig>,
197) -> io::Result<ClientTls<S>>
198where
199    S: AsyncRead + AsyncWrite + Unpin,
200{
201    let inner = TlsConnector::from(config)
202        .connect(server_name, stream)
203        .await?;
204    let certificate = inner
205        .get_ref()
206        .1
207        .peer_certificates()
208        .and_then(|certificates| certificates.first())
209        .ok_or_else(|| {
210            io::Error::new(io::ErrorKind::InvalidData, "TLS peer sent no certificate")
211        })?;
212
213    Ok(ClientTls {
214        tls_server_end_point: channel_binding(certificate)?,
215        inner,
216    })
217}
218
219/// Negotiates TLS as a `PostgreSQL` server using the configured leaf certificate.
220///
221/// # Errors
222///
223/// Returns a TLS handshake or invalid-certificate error.
224pub async fn accept<S>(
225    stream: S,
226    config: Arc<ServerConfig>,
227    leaf_certificate: &CertificateDer<'_>,
228) -> io::Result<ServerTls<S>>
229where
230    S: AsyncRead + AsyncWrite + Unpin,
231{
232    let tls_server_end_point = channel_binding(leaf_certificate)?;
233    let inner = TlsAcceptor::from(config).accept(stream).await?;
234    Ok(ServerTls {
235        inner,
236        tls_server_end_point,
237    })
238}
239
240/// Computes the RFC 5929 `tls-server-end-point` value from a DER certificate.
241///
242/// # Errors
243///
244/// Returns an error if the certificate is not valid DER.
245pub fn channel_binding(certificate: &CertificateDer<'_>) -> io::Result<Vec<u8>> {
246    let (_, parsed) = X509Certificate::from_der(certificate.as_ref()).map_err(|error| {
247        io::Error::new(
248            io::ErrorKind::InvalidData,
249            format!("invalid TLS certificate: {error}"),
250        )
251    })?;
252    let signature_oid = match SignatureAlgorithm::try_from(&parsed.signature_algorithm) {
253        Ok(SignatureAlgorithm::RSASSA_PSS(parameters)) => {
254            parameters.hash_algorithm_oid().to_id_string()
255        }
256        _ => parsed.signature_algorithm.algorithm.to_id_string(),
257    };
258
259    Ok(match signature_oid.as_str() {
260        // SHA-384, sha384WithRSAEncryption, and ecdsa-with-SHA384.
261        "2.16.840.1.101.3.4.2.2" | "1.2.840.113549.1.1.12" | "1.2.840.10045.4.3.3" => {
262            Sha384::digest(certificate.as_ref()).to_vec()
263        }
264        // SHA-512, sha512WithRSAEncryption, and ecdsa-with-SHA512.
265        "2.16.840.1.101.3.4.2.3" | "1.2.840.113549.1.1.13" | "1.2.840.10045.4.3.4" => {
266            Sha512::digest(certificate.as_ref()).to_vec()
267        }
268        // RFC 5929 promotes MD5 and SHA-1 to SHA-256. SHA-256 is also the
269        // interoperable fallback for signature schemes whose digest is not
270        // expressed directly by their algorithm identifier.
271        _ => Sha256::digest(certificate.as_ref()).to_vec(),
272    })
273}
274
275#[cfg(test)]
276mod tests {
277    use super::*;
278    use crate::{
279        Conn,
280        pre_startup::{Negotiation, PreStartupOffer},
281        transport::Buffered,
282    };
283    use rcgen::{CertificateParams, KeyPair, PKCS_ECDSA_P384_SHA384, generate_simple_self_signed};
284    use rustls::{RootCertStore, pki_types::PrivateKeyDer};
285    use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
286
287    async fn handshake_with_mode(
288        mode: SslMode,
289        roots: RootCertStore,
290        certificate: CertificateDer<'static>,
291        key_der: &[u8],
292    ) -> io::Result<ClientTls<tokio::io::DuplexStream>> {
293        let key = PrivateKeyDer::try_from(key_der.to_vec()).unwrap();
294        let server_config = Arc::new(
295            ServerConfig::builder()
296                .with_no_client_auth()
297                .with_single_cert(vec![certificate.clone()], key)
298                .unwrap(),
299        );
300        let client_config = Arc::new(client_config(mode, roots));
301        let (client_io, server_io) = tokio::io::duplex(16 * 1024);
302        let (client, _) = tokio::join!(
303            connect(
304                client_io,
305                ServerName::try_from("localhost").unwrap(),
306                client_config,
307            ),
308            accept(server_io, server_config, &certificate),
309        );
310        client
311    }
312
313    #[tokio::test]
314    async fn sslmode_controls_chain_and_hostname_verification() {
315        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
316        let generated = generate_simple_self_signed(["database.example".into()]).unwrap();
317        let certificate = CertificateDer::from(generated.cert.der().to_vec());
318        let key = generated.signing_key.serialize_der();
319
320        handshake_with_mode(
321            SslMode::Require,
322            RootCertStore::empty(),
323            certificate.clone(),
324            &key,
325        )
326        .await
327        .expect("require encrypts without validating the self-signed certificate");
328
329        let mut roots = RootCertStore::empty();
330        roots.add(certificate.clone()).unwrap();
331        handshake_with_mode(SslMode::VerifyCa, roots.clone(), certificate.clone(), &key)
332            .await
333            .expect("verify-ca validates the chain but deliberately ignores the host");
334
335        let error = handshake_with_mode(SslMode::VerifyFull, roots, certificate, &key)
336            .await
337            .expect_err("verify-full must reject a certificate for another host");
338        assert_eq!(error.kind(), io::ErrorKind::InvalidData);
339    }
340
341    #[tokio::test]
342    async fn negotiates_tls_and_exposes_equal_channel_bindings() {
343        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
344        let generated = generate_simple_self_signed(["localhost".into()]).unwrap();
345        let certificate = CertificateDer::from(generated.cert.der().to_vec());
346        let key = PrivateKeyDer::try_from(generated.signing_key.serialize_der()).unwrap();
347
348        let server_config = ServerConfig::builder()
349            .with_no_client_auth()
350            .with_single_cert(vec![certificate.clone()], key)
351            .unwrap();
352        let mut roots = RootCertStore::empty();
353        roots.add(certificate.clone()).unwrap();
354        let client_config = ClientConfig::builder()
355            .with_root_certificates(roots)
356            .with_no_client_auth();
357        let expected = channel_binding(&certificate).unwrap();
358        let (client_io, server_io) = tokio::io::duplex(16 * 1024);
359
360        let (client, server) = tokio::join!(
361            connect(
362                client_io,
363                ServerName::try_from("localhost").unwrap(),
364                Arc::new(client_config),
365            ),
366            accept(server_io, Arc::new(server_config), &certificate),
367        );
368        let client = client.unwrap();
369        let server = server.unwrap();
370
371        assert_eq!(client.tls_server_end_point(), expected);
372        assert_eq!(server.tls_server_end_point(), expected);
373    }
374
375    #[tokio::test]
376    async fn typed_pre_startup_upgrade_changes_the_buffered_transport() {
377        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
378        let generated = generate_simple_self_signed(["localhost".into()]).unwrap();
379        let certificate = CertificateDer::from(generated.cert.der().to_vec());
380        let key = PrivateKeyDer::try_from(generated.signing_key.serialize_der()).unwrap();
381        let server_config = Arc::new(
382            ServerConfig::builder()
383                .with_no_client_auth()
384                .with_single_cert(vec![certificate.clone()], key)
385                .unwrap(),
386        );
387        let mut roots = RootCertStore::empty();
388        roots.add(certificate.clone()).unwrap();
389        let client_config = Arc::new(
390            ClientConfig::builder()
391                .with_root_certificates(roots)
392                .with_no_client_auth(),
393        );
394        let expected = channel_binding(&certificate).unwrap();
395        let (client_io, mut server_io) = tokio::io::duplex(16 * 1024);
396
397        let server = tokio::spawn(async move {
398            let mut request = [0; 8];
399            server_io.read_exact(&mut request).await.unwrap();
400            assert_eq!(request, [0, 0, 0, 8, 4, 210, 22, 47]);
401            server_io.write_all(b"S").await.unwrap();
402            accept(server_io, server_config, &certificate)
403                .await
404                .unwrap()
405        });
406
407        let mut awaiting_reply = Conn::new(Buffered::new(client_io)).request_ssl();
408        awaiting_reply.flush().await.unwrap();
409        let Negotiation::Accepted(handshake) = awaiting_reply.receive_ssl_reply().await.unwrap()
410        else {
411            panic!("test server rejected TLS")
412        };
413        let pre_startup = handshake
414            .connect_tls(ServerName::try_from("localhost").unwrap(), client_config)
415            .await
416            .unwrap();
417        let client = pre_startup.into_transport().into_inner();
418        let server = server.await.unwrap();
419
420        assert_eq!(client.tls_server_end_point(), expected);
421        assert_eq!(server.tls_server_end_point(), expected);
422    }
423
424    #[tokio::test]
425    async fn server_role_terminates_tls_before_accepting_startup() {
426        let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
427        let generated = generate_simple_self_signed(["localhost".into()]).unwrap();
428        let certificate = CertificateDer::from(generated.cert.der().to_vec());
429        let key = PrivateKeyDer::try_from(generated.signing_key.serialize_der()).unwrap();
430        let server_config = Arc::new(
431            ServerConfig::builder()
432                .with_no_client_auth()
433                .with_single_cert(vec![certificate.clone()], key)
434                .unwrap(),
435        );
436        let mut roots = RootCertStore::empty();
437        roots.add(certificate.clone()).unwrap();
438        let client_config = Arc::new(
439            ClientConfig::builder()
440                .with_root_certificates(roots)
441                .with_no_client_auth(),
442        );
443        let expected = channel_binding(&certificate).unwrap();
444        let (mut client_io, server_io) = tokio::io::duplex(16 * 1024);
445
446        let client = tokio::spawn(async move {
447            client_io
448                .write_all(&[0, 0, 0, 8, 4, 210, 22, 47])
449                .await
450                .unwrap();
451            assert_eq!(client_io.read_u8().await.unwrap(), b'S');
452            connect(
453                client_io,
454                ServerName::try_from("localhost").unwrap(),
455                client_config,
456            )
457            .await
458            .unwrap()
459        });
460
461        let mut pre_startup = Conn::new(Buffered::new_frontend(server_io));
462        let message = pre_startup.receive_pre_startup_wire().await.unwrap();
463        let PreStartupOffer::Ssl(decision) = pre_startup.offer_pre_startup(message) else {
464            panic!("test client did not request TLS")
465        };
466        let mut handshake = decision.approve_ssl();
467        handshake.flush().await.unwrap();
468        let pre_startup = handshake
469            .accept_tls(server_config, certificate)
470            .await
471            .unwrap();
472        let server = pre_startup.into_transport().into_inner();
473        let client = client.await.unwrap();
474
475        assert_eq!(client.tls_server_end_point(), expected);
476        assert_eq!(server.tls_server_end_point(), expected);
477    }
478
479    #[test]
480    fn channel_binding_uses_the_certificate_signature_digest() {
481        let key = KeyPair::generate_for(&PKCS_ECDSA_P384_SHA384).unwrap();
482        let params = CertificateParams::new(["localhost".to_owned()]).unwrap();
483        let generated = params.self_signed(&key).unwrap();
484        let certificate = CertificateDer::from(generated.der().to_vec());
485
486        assert_eq!(
487            channel_binding(&certificate).unwrap(),
488            Sha384::digest(certificate.as_ref()).to_vec()
489        );
490    }
491}