flare_core/common/cert/
pinning.rs1use 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}