Skip to main content

flare_core/common/cert/
pinning.rs

1//! TLS certificate pinning helpers.
2
3use std::fmt;
4use std::sync::Arc;
5
6use base64::Engine;
7use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
8use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
9use rustls::{CertificateError, DigitallySignedStruct, Error as RustlsError, SignatureScheme};
10use sha2::{Digest, Sha256};
11use x509_parser::prelude::{FromDer, X509Certificate};
12
13use crate::common::config_types::TlsConfig;
14use crate::common::error::{FlareError, Result};
15
16#[derive(Clone, Debug, Default)]
17pub struct TlsPinningPolicy {
18    spki_sha256_pins: Vec<Vec<u8>>,
19    certificate_sha256_pins: Vec<Vec<u8>>,
20}
21
22impl TlsPinningPolicy {
23    pub fn from_tls_config(tls: &TlsConfig) -> Result<Self> {
24        let spki_sha256_pins = normalize_pin_list(&tls.spki_sha256_pins, "SPKI SHA-256")?;
25        let certificate_sha256_pins =
26            normalize_pin_list(&tls.certificate_sha256_pins, "certificate SHA-256")?;
27        Ok(Self {
28            spki_sha256_pins,
29            certificate_sha256_pins,
30        })
31    }
32
33    pub fn is_enabled(&self) -> bool {
34        !self.spki_sha256_pins.is_empty() || !self.certificate_sha256_pins.is_empty()
35    }
36
37    pub fn verify_chain(
38        &self,
39        end_entity: &CertificateDer<'_>,
40        intermediates: &[CertificateDer<'_>],
41    ) -> std::result::Result<(), String> {
42        if !self.spki_sha256_pins.is_empty()
43            && certificate_chain_der(end_entity, intermediates).any(|certificate| {
44                spki_sha256(certificate.as_ref())
45                    .map(|actual| self.spki_sha256_pins.iter().any(|pin| pin == &actual))
46                    .unwrap_or(false)
47            })
48        {
49            return Ok(());
50        }
51
52        if !self.certificate_sha256_pins.is_empty()
53            && certificate_chain_der(end_entity, intermediates).any(|certificate| {
54                let actual = Sha256::digest(certificate.as_ref()).to_vec();
55                self.certificate_sha256_pins
56                    .iter()
57                    .any(|pin| pin == &actual)
58            })
59        {
60            return Ok(());
61        }
62
63        let leaf_spki = spki_sha256(end_entity.as_ref())
64            .map(|hash| format!("spki-sha256/{}", base64_pin(&hash)))
65            .unwrap_or_else(|err| format!("spki-unavailable({err})"));
66        Err(format!("TLS certificate pin mismatch; leaf {leaf_spki}"))
67    }
68}
69
70#[derive(Debug)]
71pub struct PinnedServerCertVerifier {
72    delegate: Arc<dyn ServerCertVerifier>,
73    policy: TlsPinningPolicy,
74}
75
76impl PinnedServerCertVerifier {
77    pub fn new(delegate: Arc<dyn ServerCertVerifier>, policy: TlsPinningPolicy) -> Self {
78        Self { delegate, policy }
79    }
80}
81
82impl ServerCertVerifier for PinnedServerCertVerifier {
83    fn verify_server_cert(
84        &self,
85        end_entity: &CertificateDer<'_>,
86        intermediates: &[CertificateDer<'_>],
87        server_name: &ServerName<'_>,
88        ocsp_response: &[u8],
89        now: UnixTime,
90    ) -> std::result::Result<ServerCertVerified, RustlsError> {
91        self.delegate.verify_server_cert(
92            end_entity,
93            intermediates,
94            server_name,
95            ocsp_response,
96            now,
97        )?;
98
99        self.policy
100            .verify_chain(end_entity, intermediates)
101            .map_err(|_| {
102                RustlsError::InvalidCertificate(CertificateError::ApplicationVerificationFailure)
103            })?;
104
105        Ok(ServerCertVerified::assertion())
106    }
107
108    fn verify_tls12_signature(
109        &self,
110        message: &[u8],
111        cert: &CertificateDer<'_>,
112        dss: &DigitallySignedStruct,
113    ) -> std::result::Result<HandshakeSignatureValid, RustlsError> {
114        self.delegate.verify_tls12_signature(message, cert, dss)
115    }
116
117    fn verify_tls13_signature(
118        &self,
119        message: &[u8],
120        cert: &CertificateDer<'_>,
121        dss: &DigitallySignedStruct,
122    ) -> std::result::Result<HandshakeSignatureValid, RustlsError> {
123        self.delegate.verify_tls13_signature(message, cert, dss)
124    }
125
126    fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
127        self.delegate.supported_verify_schemes()
128    }
129}
130
131pub fn normalize_sha256_pin(pin: &str) -> Option<Vec<u8>> {
132    let pin = pin.trim();
133    let pin = pin
134        .strip_prefix("spki-sha256/")
135        .or_else(|| pin.strip_prefix("sha256/"))
136        .unwrap_or(pin);
137    let hex_candidate: String = pin.chars().filter(|ch| *ch != ':').collect();
138
139    if hex_candidate.len() == 64 && hex_candidate.chars().all(|ch| ch.is_ascii_hexdigit()) {
140        return hex::decode(hex_candidate).ok();
141    }
142
143    base64::engine::general_purpose::STANDARD.decode(pin).ok()
144}
145
146pub fn spki_sha256(certificate_der: &[u8]) -> std::result::Result<Vec<u8>, String> {
147    let (_, certificate) = X509Certificate::from_der(certificate_der)
148        .map_err(|error| format!("parse certificate DER failed: {error}"))?;
149    Ok(Sha256::digest(certificate.tbs_certificate.subject_pki.raw).to_vec())
150}
151
152pub fn spki_sha256_pin(certificate_der: &[u8]) -> std::result::Result<String, String> {
153    spki_sha256(certificate_der).map(|hash| format!("spki-sha256/{}", base64_pin(&hash)))
154}
155
156fn normalize_pin_list(raw_pins: &[String], label: &str) -> Result<Vec<Vec<u8>>> {
157    raw_pins
158        .iter()
159        .map(|pin| {
160            normalize_sha256_pin(pin)
161                .ok_or_else(|| FlareError::protocol_error(format!("invalid {label} pin: {pin}")))
162        })
163        .collect()
164}
165
166fn base64_pin(hash: &[u8]) -> String {
167    base64::engine::general_purpose::STANDARD.encode(hash)
168}
169
170fn certificate_chain_der<'a>(
171    end_entity: &'a CertificateDer<'a>,
172    intermediates: &'a [CertificateDer<'a>],
173) -> impl Iterator<Item = &'a CertificateDer<'a>> {
174    std::iter::once(end_entity).chain(intermediates.iter())
175}
176
177impl fmt::Display for TlsPinningPolicy {
178    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
179        f.debug_struct("TlsPinningPolicy")
180            .field("spki_sha256_pins", &self.spki_sha256_pins.len())
181            .field(
182                "certificate_sha256_pins",
183                &self.certificate_sha256_pins.len(),
184            )
185            .finish()
186    }
187}
188
189#[cfg(test)]
190mod tests {
191    use super::*;
192
193    #[test]
194    fn pin_normalization_accepts_base64_and_hex() {
195        let bytes = [7u8; 32];
196        let base64_pin = format!(
197            "spki-sha256/{}",
198            base64::engine::general_purpose::STANDARD.encode(bytes)
199        );
200        let legacy_base64_pin = format!(
201            "sha256/{}",
202            base64::engine::general_purpose::STANDARD.encode(bytes)
203        );
204        let hex_pin = hex::encode(bytes);
205        let colon_hex_pin = hex_pin
206            .as_bytes()
207            .chunks(2)
208            .map(|chunk| std::str::from_utf8(chunk).unwrap())
209            .collect::<Vec<_>>()
210            .join(":");
211
212        assert_eq!(normalize_sha256_pin(&base64_pin), Some(bytes.to_vec()));
213        assert_eq!(
214            normalize_sha256_pin(&legacy_base64_pin),
215            Some(bytes.to_vec())
216        );
217        assert_eq!(normalize_sha256_pin(&hex_pin), Some(bytes.to_vec()));
218        assert_eq!(normalize_sha256_pin(&colon_hex_pin), Some(bytes.to_vec()));
219    }
220
221    #[test]
222    fn spki_pin_is_stable_for_generated_certificate() {
223        let certified = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
224            .expect("generate certificate");
225        let cert_der = certified.cert.der().to_vec();
226
227        let pin = spki_sha256_pin(&cert_der).expect("spki pin");
228
229        assert!(pin.starts_with("spki-sha256/"));
230        assert_eq!(normalize_sha256_pin(&pin).expect("normalize").len(), 32);
231    }
232
233    #[test]
234    fn policy_matches_leaf_spki_pin() {
235        let certified = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
236            .expect("generate certificate");
237        let cert_der = certified.cert.der().to_vec();
238        let pin = spki_sha256_pin(&cert_der).expect("spki pin");
239        let policy =
240            TlsPinningPolicy::from_tls_config(&TlsConfig::none().with_spki_sha256_pin(pin))
241                .expect("policy");
242        let certificate = CertificateDer::from(cert_der);
243
244        policy
245            .verify_chain(&certificate, &[])
246            .expect("matching pin");
247    }
248
249    #[test]
250    fn policy_rejects_non_matching_spki_pin() {
251        let certified = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
252            .expect("generate certificate");
253        let cert_der = certified.cert.der().to_vec();
254        let policy =
255            TlsPinningPolicy::from_tls_config(&TlsConfig::none().with_spki_sha256_pin(format!(
256                "spki-sha256/{}",
257                base64::engine::general_purpose::STANDARD.encode([9u8; 32])
258            )))
259            .expect("policy");
260        let certificate = CertificateDer::from(cert_der);
261
262        policy
263            .verify_chain(&certificate, &[])
264            .expect_err("non-matching pin");
265    }
266}