1use 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#[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#[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 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 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
102fn parse_identity_pem(
104 pem: &[u8],
105) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), TlsError> {
106 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
118pub(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 config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
155 Ok(config)
156}
157
158#[derive(Debug)]
171struct ExactRootMatch {
172 inner: rustls_platform_verifier::Verifier,
173 roots: Vec<CertificateDer<'static>>,
174}
175
176fn 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 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 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 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#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
298pub struct Certificate {
299 der: Vec<u8>,
300}
301
302impl Certificate {
303 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 pub fn from_der(der: &[u8]) -> Result<Certificate, TlsError> {
314 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}