Skip to main content

mz_ccsr/
tls.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10//! TLS certificates and identities.
11
12use std::fmt;
13use std::sync::Arc;
14
15use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
16use rustls::client::verify_server_name;
17use rustls::pki_types::pem::PemObject;
18use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime};
19use rustls::server::ParsedCertificate;
20use rustls::{CertificateError, DigitallySignedStruct, SignatureScheme};
21use serde::{Deserialize, Serialize};
22use x509_cert::der::asn1::{Ia5StringRef, PrintableStringRef, Utf8StringRef};
23use x509_cert::der::oid::db::rfc4519;
24use x509_cert::der::{Decode, Tag, Tagged};
25use x509_cert::ext::pkix::SubjectAltName;
26use x509_cert::ext::pkix::name::GeneralName;
27use zeroize::{Zeroize, Zeroizing};
28
29/// An error constructing a [`Certificate`] or [`Identity`].
30#[derive(Debug, thiserror::Error)]
31pub enum TlsError {
32    #[error("invalid PEM: {0}")]
33    Pem(#[from] rustls::pki_types::pem::Error),
34    #[error("no certificate found in PEM input")]
35    NoCertificate,
36    #[error("no private key found in PEM input")]
37    NoPrivateKey,
38    #[error("invalid certificate: {0}")]
39    Certificate(rustls::CertificateError),
40    #[error("invalid TLS identity: {0}")]
41    Identity(rustls::Error),
42    #[error("invalid TLS configuration: {0}")]
43    Config(rustls::Error),
44    #[error(transparent)]
45    Reqwest(#[from] reqwest::Error),
46}
47
48/// A [Serde][serde]-enabled wrapper around [`reqwest::Identity`].
49///
50/// Holds the PEM-encoded private key and certificate chain. The buffer is
51/// zeroized on drop.
52///
53/// [Serde]: serde
54#[derive(Clone, Eq, PartialEq, Hash, Serialize, Deserialize)]
55pub struct Identity {
56    pem: Vec<u8>,
57}
58
59impl fmt::Debug for Identity {
60    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
61        f.debug_struct("Identity").finish_non_exhaustive()
62    }
63}
64
65impl Zeroize for Identity {
66    fn zeroize(&mut self) {
67        self.pem.zeroize();
68    }
69}
70
71impl Drop for Identity {
72    fn drop(&mut self) {
73        self.zeroize();
74    }
75}
76
77impl Identity {
78    /// Constructs an identity from a PEM-formatted private key and certificate
79    /// chain, leaf certificate first.
80    ///
81    /// The key may be PKCS #8, PKCS #1 (RSA) or SEC1 (EC). Returns an error if
82    /// the key does not match the leaf certificate.
83    pub fn from_pem(key: &[u8], cert: &[u8]) -> Result<Self, TlsError> {
84        let mut pem = Zeroizing::new(Vec::with_capacity(key.len() + cert.len() + 1));
85        pem.extend_from_slice(key);
86        pem.push(b'\n');
87        pem.extend_from_slice(cert);
88
89        let (certs, key) = parse_identity_pem(&pem)?;
90        // reqwest only checks that the key matches the certificate when the
91        // client is built, so check here to report the error up front.
92        let provider = rustls::crypto::aws_lc_rs::default_provider();
93        rustls::sign::CertifiedKey::from_der(certs, key, &provider).map_err(TlsError::Identity)?;
94        let _ = reqwest::Identity::from_pem(&pem)?;
95
96        Ok(Identity {
97            pem: std::mem::take(&mut *pem),
98        })
99    }
100}
101
102/// Splits an identity PEM buffer into its certificate chain and private key.
103fn parse_identity_pem(
104    pem: &[u8],
105) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), TlsError> {
106    // Mirror `reqwest::Identity::from_pem`, which uses the last private key in
107    // the buffer.
108    let mut keys = PrivateKeyDer::pem_slice_iter(pem).collect::<Result<Vec<_>, _>>()?;
109    let key = keys.pop().ok_or(TlsError::NoPrivateKey)?;
110    keys.iter_mut().for_each(Zeroize::zeroize);
111    let certs = CertificateDer::pem_slice_iter(pem).collect::<Result<Vec<_>, _>>()?;
112    if certs.is_empty() {
113        return Err(TlsError::NoCertificate);
114    }
115    Ok((certs, key))
116}
117
118/// Builds the rustls configuration for a client that trusts `roots` in
119/// addition to the platform's trust store, and that presents `identity`, if
120/// any, for client authentication.
121///
122/// Server certificates are verified by `rustls-platform-verifier`, as reqwest
123/// does by default, with the [`ExactRootMatch`] fallback.
124pub(crate) fn rustls_config(
125    roots: &[Certificate],
126    identity: Option<&Identity>,
127) -> Result<rustls::ClientConfig, TlsError> {
128    let provider = Arc::new(rustls::crypto::aws_lc_rs::default_provider());
129    let roots: Vec<_> = roots
130        .iter()
131        .map(|cert| CertificateDer::from(cert.der.clone()))
132        .collect();
133    let inner = rustls_platform_verifier::Verifier::new_with_extra_roots(
134        roots.clone(),
135        Arc::clone(&provider),
136    )
137    .map_err(TlsError::Config)?;
138    let builder = rustls::ClientConfig::builder_with_provider(provider)
139        .with_safe_default_protocol_versions()
140        .map_err(TlsError::Config)?
141        .dangerous()
142        .with_custom_certificate_verifier(Arc::new(ExactRootMatch { inner, roots }));
143    let mut config = match identity {
144        Some(identity) => {
145            let (certs, key) = parse_identity_pem(&identity.pem)?;
146            builder
147                .with_client_auth_cert(certs, key)
148                .map_err(TlsError::Identity)?
149        }
150        None => builder.with_no_client_auth(),
151    };
152    // reqwest only sets ALPN on TLS configurations it builds itself. This
153    // mirrors its choice while the workspace enables reqwest's `http2` feature.
154    config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
155    Ok(config)
156}
157
158/// A server certificate verifier that accepts a server certificate that is
159/// byte-for-byte identical to one of `roots`, and otherwise defers to `inner`.
160///
161/// An exact match must still be within its validity period and valid for the
162/// server name, but skips the chain, basic constraints and key usage checks.
163/// This keeps a self-signed `CA:TRUE` certificate working when it is supplied
164/// as its own certificate authority, which webpki rejects as
165/// `CaUsedAsEndEntity`. The name check falls back to the subject common names
166/// when an exact match has no DNS or IP subjectAltName, see
167/// [`common_name_matches`]. Every other certificate is checked against
168/// subjectAltNames only. An exact match is decided without consulting
169/// `inner`. Handshake signatures are always verified by `inner`.
170#[derive(Debug)]
171struct ExactRootMatch {
172    inner: rustls_platform_verifier::Verifier,
173    roots: Vec<CertificateDer<'static>>,
174}
175
176/// Returns whether `server_name` is a DNS name equal, ignoring ASCII case, to
177/// any subject common name of `cert`, and `cert` has no DNS or IP
178/// subjectAltName.
179///
180/// This is stricter than OpenSSL's fallback, which also applies when only IP
181/// subjectAltNames are present, matches wildcard common names, and decodes
182/// string types other than UTF8String, PrintableString and IA5String.
183fn common_name_matches(cert: &x509_cert::Certificate, server_name: &ServerName<'_>) -> bool {
184    let ServerName::DnsName(name) = server_name else {
185        return false;
186    };
187    let tbs = &cert.tbs_certificate;
188    // An undecodable subjectAltName extension counts as present.
189    let has_san = tbs.filter::<SubjectAltName>().any(|san| match san {
190        Ok((_, SubjectAltName(names))) => names
191            .iter()
192            .any(|name| matches!(name, GeneralName::DnsName(_) | GeneralName::IpAddress(_))),
193        Err(_) => true,
194    });
195    if has_san {
196        return false;
197    }
198    tbs.subject
199        .0
200        .iter()
201        .flat_map(|rdn| rdn.0.iter())
202        .filter(|atv| atv.oid == rfc4519::CN)
203        .any(|cn| {
204            let cn = match cn.value.tag() {
205                Tag::Utf8String => cn
206                    .value
207                    .decode_as::<Utf8StringRef<'_>>()
208                    .map(|s| s.as_str().to_owned()),
209                Tag::PrintableString => cn
210                    .value
211                    .decode_as::<PrintableStringRef<'_>>()
212                    .map(|s| s.as_str().to_owned()),
213                Tag::Ia5String => cn
214                    .value
215                    .decode_as::<Ia5StringRef<'_>>()
216                    .map(|s| s.as_str().to_owned()),
217                _ => return false,
218            };
219            cn.is_ok_and(|cn| cn.eq_ignore_ascii_case(name.as_ref()))
220        })
221}
222
223impl ServerCertVerifier for ExactRootMatch {
224    fn verify_server_cert(
225        &self,
226        end_entity: &CertificateDer<'_>,
227        intermediates: &[CertificateDer<'_>],
228        server_name: &ServerName<'_>,
229        ocsp_response: &[u8],
230        now: UnixTime,
231    ) -> Result<ServerCertVerified, rustls::Error> {
232        // Checked before `inner`, which would reject a `CA:TRUE` certificate and
233        // logs every rejection at error level.
234        if self
235            .roots
236            .iter()
237            .any(|root| root.as_ref() == end_entity.as_ref())
238        {
239            let cert = x509_cert::Certificate::from_der(end_entity)
240                .map_err(|_| rustls::Error::InvalidCertificate(CertificateError::BadEncoding))?;
241            let validity = &cert.tbs_certificate.validity;
242            let now = now.as_secs();
243            if now < validity.not_before.to_unix_duration().as_secs() {
244                return Err(rustls::Error::InvalidCertificate(
245                    CertificateError::NotValidYet,
246                ));
247            }
248            if now > validity.not_after.to_unix_duration().as_secs() {
249                return Err(rustls::Error::InvalidCertificate(CertificateError::Expired));
250            }
251            // Checks subjectAltNames only, not basic constraints, so a
252            // `CA:TRUE` certificate passes. Errors match `inner`'s on Linux.
253            let parsed = ParsedCertificate::try_from(end_entity)?;
254            match verify_server_name(&parsed, server_name) {
255                Ok(()) => {}
256                Err(_) if common_name_matches(&cert, server_name) => {}
257                Err(e) => return Err(e),
258            }
259            return Ok(ServerCertVerified::assertion());
260        }
261        self.inner
262            .verify_server_cert(end_entity, intermediates, server_name, ocsp_response, now)
263    }
264
265    fn verify_tls12_signature(
266        &self,
267        message: &[u8],
268        cert: &CertificateDer<'_>,
269        dss: &DigitallySignedStruct,
270    ) -> Result<HandshakeSignatureValid, rustls::Error> {
271        self.inner.verify_tls12_signature(message, cert, dss)
272    }
273
274    fn verify_tls13_signature(
275        &self,
276        message: &[u8],
277        cert: &CertificateDer<'_>,
278        dss: &DigitallySignedStruct,
279    ) -> Result<HandshakeSignatureValid, rustls::Error> {
280        self.inner.verify_tls13_signature(message, cert, dss)
281    }
282
283    fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
284        self.inner.supported_verify_schemes()
285    }
286}
287
288impl From<Identity> for reqwest::Identity {
289    fn from(id: Identity) -> Self {
290        reqwest::Identity::from_pem(&id.pem).expect("known to be a valid identity")
291    }
292}
293
294/// A [Serde][serde]-enabled wrapper around [`reqwest::Certificate`].
295///
296/// [Serde]: serde
297#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
298pub struct Certificate {
299    der: Vec<u8>,
300}
301
302impl Certificate {
303    /// Constructs a certificate from the first certificate in a PEM-formatted
304    /// buffer.
305    pub fn from_pem(pem: &[u8]) -> Result<Certificate, TlsError> {
306        let der = CertificateDer::pem_slice_iter(pem)
307            .next()
308            .ok_or(TlsError::NoCertificate)??;
309        Self::from_der(&der)
310    }
311
312    /// Constructs a certificate from a DER-formatted buffer.
313    pub fn from_der(der: &[u8]) -> Result<Certificate, TlsError> {
314        // Parse the certificate as a trust anchor, as the verifier does when
315        // the client is built.
316        rustls::RootCertStore::empty()
317            .add(CertificateDer::from_slice(der).into_owned())
318            .map_err(|e| match e {
319                rustls::Error::InvalidCertificate(e) => TlsError::Certificate(e),
320                e => TlsError::Certificate(rustls::CertificateError::Other(rustls::OtherError(
321                    Arc::new(e),
322                ))),
323            })?;
324        Ok(Certificate { der: der.into() })
325    }
326}
327
328impl From<Certificate> for reqwest::Certificate {
329    fn from(cert: Certificate) -> Self {
330        reqwest::Certificate::from_der(&cert.der).expect("known to be a valid cert")
331    }
332}