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::pki_types::pem::PemObject;
16use rustls::pki_types::{CertificateDer, PrivateKeyDer};
17use serde::{Deserialize, Serialize};
18use zeroize::{Zeroize, Zeroizing};
19
20/// An error constructing a [`Certificate`] or [`Identity`].
21#[derive(Debug, thiserror::Error)]
22pub enum TlsError {
23    #[error("invalid PEM: {0}")]
24    Pem(#[from] rustls::pki_types::pem::Error),
25    #[error("no certificate found in PEM input")]
26    NoCertificate,
27    #[error("no private key found in PEM input")]
28    NoPrivateKey,
29    #[error("invalid certificate: {0}")]
30    Certificate(rustls::CertificateError),
31    #[error("invalid TLS identity: {0}")]
32    Identity(rustls::Error),
33    #[error(transparent)]
34    Reqwest(#[from] reqwest::Error),
35}
36
37/// A [Serde][serde]-enabled wrapper around [`reqwest::Identity`].
38///
39/// Holds the PEM-encoded private key and certificate chain. The buffer is
40/// zeroized on drop.
41///
42/// [Serde]: serde
43#[derive(Clone, Eq, PartialEq, Hash, Serialize, Deserialize)]
44pub struct Identity {
45    pem: Vec<u8>,
46}
47
48impl fmt::Debug for Identity {
49    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
50        f.debug_struct("Identity").finish_non_exhaustive()
51    }
52}
53
54impl Zeroize for Identity {
55    fn zeroize(&mut self) {
56        self.pem.zeroize();
57    }
58}
59
60impl Drop for Identity {
61    fn drop(&mut self) {
62        self.zeroize();
63    }
64}
65
66impl Identity {
67    /// Constructs an identity from a PEM-formatted private key and certificate
68    /// chain, leaf certificate first.
69    ///
70    /// The key may be PKCS #8, PKCS #1 (RSA) or SEC1 (EC). Returns an error if
71    /// the key does not match the leaf certificate.
72    pub fn from_pem(key: &[u8], cert: &[u8]) -> Result<Self, TlsError> {
73        let mut pem = Zeroizing::new(Vec::with_capacity(key.len() + cert.len() + 1));
74        pem.extend_from_slice(key);
75        pem.push(b'\n');
76        pem.extend_from_slice(cert);
77
78        // Mirror `reqwest::Identity::from_pem`, which uses the last private
79        // key in the buffer.
80        let mut keys = PrivateKeyDer::pem_slice_iter(&pem).collect::<Result<Vec<_>, _>>()?;
81        let key = keys.pop().ok_or(TlsError::NoPrivateKey)?;
82        keys.iter_mut().for_each(Zeroize::zeroize);
83        let certs = CertificateDer::pem_slice_iter(&pem).collect::<Result<Vec<_>, _>>()?;
84        if certs.is_empty() {
85            return Err(TlsError::NoCertificate);
86        }
87
88        // reqwest only checks that the key matches the certificate when the
89        // client is built, so check here to report the error up front.
90        let provider = rustls::crypto::aws_lc_rs::default_provider();
91        rustls::sign::CertifiedKey::from_der(certs, key, &provider).map_err(TlsError::Identity)?;
92        let _ = reqwest::Identity::from_pem(&pem)?;
93
94        Ok(Identity {
95            pem: std::mem::take(&mut *pem),
96        })
97    }
98}
99
100impl From<Identity> for reqwest::Identity {
101    fn from(id: Identity) -> Self {
102        reqwest::Identity::from_pem(&id.pem).expect("known to be a valid identity")
103    }
104}
105
106/// A [Serde][serde]-enabled wrapper around [`reqwest::Certificate`].
107///
108/// [Serde]: serde
109#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
110pub struct Certificate {
111    der: Vec<u8>,
112}
113
114impl Certificate {
115    /// Constructs a certificate from the first certificate in a PEM-formatted
116    /// buffer.
117    pub fn from_pem(pem: &[u8]) -> Result<Certificate, TlsError> {
118        let der = CertificateDer::pem_slice_iter(pem)
119            .next()
120            .ok_or(TlsError::NoCertificate)??;
121        Self::from_der(&der)
122    }
123
124    /// Constructs a certificate from a DER-formatted buffer.
125    pub fn from_der(der: &[u8]) -> Result<Certificate, TlsError> {
126        // Parse the certificate as a trust anchor, as the verifier does when
127        // the client is built.
128        rustls::RootCertStore::empty()
129            .add(CertificateDer::from_slice(der).into_owned())
130            .map_err(|e| match e {
131                rustls::Error::InvalidCertificate(e) => TlsError::Certificate(e),
132                e => TlsError::Certificate(rustls::CertificateError::Other(rustls::OtherError(
133                    Arc::new(e),
134                ))),
135            })?;
136        Ok(Certificate { der: der.into() })
137    }
138}
139
140impl From<Certificate> for reqwest::Certificate {
141    fn from(cert: Certificate) -> Self {
142        reqwest::Certificate::from_der(&cert.der).expect("known to be a valid cert")
143    }
144}