1use std::{
4 io,
5 pin::Pin,
6 sync::Arc,
7 task::{Context, Poll},
8};
9
10use rustls::{
11 ClientConfig, DigitallySignedStruct, Error as TlsError, RootCertStore, ServerConfig,
12 SignatureScheme,
13 client::{
14 WebPkiServerVerifier,
15 danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier},
16 },
17 crypto::{CryptoProvider, WebPkiSupportedAlgorithms},
18 pki_types::{CertificateDer, ServerName, UnixTime},
19};
20use sha2::{Digest, Sha256, Sha384, Sha512};
21use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
22use tokio_rustls::{TlsAcceptor, TlsConnector};
23use x509_parser::{
24 prelude::{FromDer, X509Certificate},
25 signature_algorithm::SignatureAlgorithm,
26};
27
28use crate::{
29 auth::TlsServerEndPoint,
30 pre_startup::{CertificateVerification, SslMode},
31};
32
33#[derive(Debug)]
34struct CertificateVerifier {
35 verification: CertificateVerification,
36 roots: RootCertStore,
37 supported: WebPkiSupportedAlgorithms,
38}
39
40impl ServerCertVerifier for CertificateVerifier {
41 fn verify_server_cert(
42 &self,
43 end_entity: &CertificateDer<'_>,
44 intermediates: &[CertificateDer<'_>],
45 server_name: &ServerName<'_>,
46 _ocsp_response: &[u8],
47 now: UnixTime,
48 ) -> Result<ServerCertVerified, TlsError> {
49 match self.verification {
50 CertificateVerification::None => Ok(ServerCertVerified::assertion()),
51 CertificateVerification::CertificateAuthority => {
52 let parsed = rustls::server::ParsedCertificate::try_from(end_entity)?;
53 rustls::client::verify_server_cert_signed_by_trust_anchor(
54 &parsed,
55 &self.roots,
56 intermediates,
57 now,
58 self.supported.all,
59 )?;
60 Ok(ServerCertVerified::assertion())
61 }
62 CertificateVerification::CertificateAuthorityAndHost => {
63 let verifier = WebPkiServerVerifier::builder(Arc::new(self.roots.clone()))
64 .build()
65 .map_err(|error| TlsError::General(error.to_string()))?;
66 verifier.verify_server_cert(end_entity, intermediates, server_name, &[], now)
67 }
68 }
69 }
70
71 fn verify_tls12_signature(
72 &self,
73 message: &[u8],
74 cert: &CertificateDer<'_>,
75 dss: &DigitallySignedStruct,
76 ) -> Result<HandshakeSignatureValid, TlsError> {
77 rustls::crypto::verify_tls12_signature(message, cert, dss, &self.supported)
78 }
79
80 fn verify_tls13_signature(
81 &self,
82 message: &[u8],
83 cert: &CertificateDer<'_>,
84 dss: &DigitallySignedStruct,
85 ) -> Result<HandshakeSignatureValid, TlsError> {
86 rustls::crypto::verify_tls13_signature(message, cert, dss, &self.supported)
87 }
88
89 fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
90 self.supported.supported_schemes()
91 }
92}
93
94#[must_use]
99pub fn client_config(mode: SslMode, roots: RootCertStore) -> ClientConfig {
100 let verification = mode.strategy().verification;
101 if verification == CertificateVerification::CertificateAuthorityAndHost {
102 return ClientConfig::builder()
103 .with_root_certificates(roots)
104 .with_no_client_auth();
105 }
106
107 let provider = CryptoProvider::get_default()
108 .cloned()
109 .unwrap_or_else(|| Arc::new(rustls::crypto::aws_lc_rs::default_provider()));
110 let verifier = CertificateVerifier {
111 verification,
112 roots,
113 supported: provider.signature_verification_algorithms,
114 };
115 ClientConfig::builder()
116 .dangerous()
117 .with_custom_certificate_verifier(Arc::new(verifier))
118 .with_no_client_auth()
119}
120
121#[derive(Debug)]
123pub struct ClientTls<S> {
124 inner: tokio_rustls::client::TlsStream<S>,
125 tls_server_end_point: Vec<u8>,
126}
127
128#[derive(Debug)]
130pub struct ServerTls<S> {
131 inner: tokio_rustls::server::TlsStream<S>,
132 tls_server_end_point: Vec<u8>,
133}
134
135impl<S> TlsServerEndPoint for ClientTls<S> {
136 fn tls_server_end_point(&self) -> &[u8] {
137 &self.tls_server_end_point
138 }
139}
140
141impl<S> TlsServerEndPoint for ServerTls<S> {
142 fn tls_server_end_point(&self) -> &[u8] {
143 &self.tls_server_end_point
144 }
145}
146
147macro_rules! delegate_io {
148 ($wrapper:ident) => {
149 impl<S: AsyncRead + AsyncWrite + Unpin> AsyncRead for $wrapper<S> {
150 fn poll_read(
151 mut self: Pin<&mut Self>,
152 cx: &mut Context<'_>,
153 buffer: &mut ReadBuf<'_>,
154 ) -> Poll<io::Result<()>> {
155 Pin::new(&mut self.inner).poll_read(cx, buffer)
156 }
157 }
158
159 impl<S: AsyncRead + AsyncWrite + Unpin> AsyncWrite for $wrapper<S> {
160 fn poll_write(
161 mut self: Pin<&mut Self>,
162 cx: &mut Context<'_>,
163 buffer: &[u8],
164 ) -> Poll<Result<usize, io::Error>> {
165 Pin::new(&mut self.inner).poll_write(cx, buffer)
166 }
167
168 fn poll_flush(
169 mut self: Pin<&mut Self>,
170 cx: &mut Context<'_>,
171 ) -> Poll<Result<(), io::Error>> {
172 Pin::new(&mut self.inner).poll_flush(cx)
173 }
174
175 fn poll_shutdown(
176 mut self: Pin<&mut Self>,
177 cx: &mut Context<'_>,
178 ) -> Poll<Result<(), io::Error>> {
179 Pin::new(&mut self.inner).poll_shutdown(cx)
180 }
181 }
182 };
183}
184
185delegate_io!(ClientTls);
186delegate_io!(ServerTls);
187
188pub async fn connect<S>(
194 stream: S,
195 server_name: ServerName<'static>,
196 config: Arc<ClientConfig>,
197) -> io::Result<ClientTls<S>>
198where
199 S: AsyncRead + AsyncWrite + Unpin,
200{
201 let inner = TlsConnector::from(config)
202 .connect(server_name, stream)
203 .await?;
204 let certificate = inner
205 .get_ref()
206 .1
207 .peer_certificates()
208 .and_then(|certificates| certificates.first())
209 .ok_or_else(|| {
210 io::Error::new(io::ErrorKind::InvalidData, "TLS peer sent no certificate")
211 })?;
212
213 Ok(ClientTls {
214 tls_server_end_point: channel_binding(certificate)?,
215 inner,
216 })
217}
218
219pub async fn accept<S>(
225 stream: S,
226 config: Arc<ServerConfig>,
227 leaf_certificate: &CertificateDer<'_>,
228) -> io::Result<ServerTls<S>>
229where
230 S: AsyncRead + AsyncWrite + Unpin,
231{
232 let tls_server_end_point = channel_binding(leaf_certificate)?;
233 let inner = TlsAcceptor::from(config).accept(stream).await?;
234 Ok(ServerTls {
235 inner,
236 tls_server_end_point,
237 })
238}
239
240pub fn channel_binding(certificate: &CertificateDer<'_>) -> io::Result<Vec<u8>> {
246 let (_, parsed) = X509Certificate::from_der(certificate.as_ref()).map_err(|error| {
247 io::Error::new(
248 io::ErrorKind::InvalidData,
249 format!("invalid TLS certificate: {error}"),
250 )
251 })?;
252 let signature_oid = match SignatureAlgorithm::try_from(&parsed.signature_algorithm) {
253 Ok(SignatureAlgorithm::RSASSA_PSS(parameters)) => {
254 parameters.hash_algorithm_oid().to_id_string()
255 }
256 _ => parsed.signature_algorithm.algorithm.to_id_string(),
257 };
258
259 Ok(match signature_oid.as_str() {
260 "2.16.840.1.101.3.4.2.2" | "1.2.840.113549.1.1.12" | "1.2.840.10045.4.3.3" => {
262 Sha384::digest(certificate.as_ref()).to_vec()
263 }
264 "2.16.840.1.101.3.4.2.3" | "1.2.840.113549.1.1.13" | "1.2.840.10045.4.3.4" => {
266 Sha512::digest(certificate.as_ref()).to_vec()
267 }
268 _ => Sha256::digest(certificate.as_ref()).to_vec(),
272 })
273}
274
275#[cfg(test)]
276mod tests {
277 use super::*;
278 use crate::{
279 Conn,
280 pre_startup::{Negotiation, PreStartupOffer},
281 transport::Buffered,
282 };
283 use rcgen::{CertificateParams, KeyPair, PKCS_ECDSA_P384_SHA384, generate_simple_self_signed};
284 use rustls::{RootCertStore, pki_types::PrivateKeyDer};
285 use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
286
287 async fn handshake_with_mode(
288 mode: SslMode,
289 roots: RootCertStore,
290 certificate: CertificateDer<'static>,
291 key_der: &[u8],
292 ) -> io::Result<ClientTls<tokio::io::DuplexStream>> {
293 let key = PrivateKeyDer::try_from(key_der.to_vec()).unwrap();
294 let server_config = Arc::new(
295 ServerConfig::builder()
296 .with_no_client_auth()
297 .with_single_cert(vec![certificate.clone()], key)
298 .unwrap(),
299 );
300 let client_config = Arc::new(client_config(mode, roots));
301 let (client_io, server_io) = tokio::io::duplex(16 * 1024);
302 let (client, _) = tokio::join!(
303 connect(
304 client_io,
305 ServerName::try_from("localhost").unwrap(),
306 client_config,
307 ),
308 accept(server_io, server_config, &certificate),
309 );
310 client
311 }
312
313 #[tokio::test]
314 async fn sslmode_controls_chain_and_hostname_verification() {
315 let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
316 let generated = generate_simple_self_signed(["database.example".into()]).unwrap();
317 let certificate = CertificateDer::from(generated.cert.der().to_vec());
318 let key = generated.signing_key.serialize_der();
319
320 handshake_with_mode(
321 SslMode::Require,
322 RootCertStore::empty(),
323 certificate.clone(),
324 &key,
325 )
326 .await
327 .expect("require encrypts without validating the self-signed certificate");
328
329 let mut roots = RootCertStore::empty();
330 roots.add(certificate.clone()).unwrap();
331 handshake_with_mode(SslMode::VerifyCa, roots.clone(), certificate.clone(), &key)
332 .await
333 .expect("verify-ca validates the chain but deliberately ignores the host");
334
335 let error = handshake_with_mode(SslMode::VerifyFull, roots, certificate, &key)
336 .await
337 .expect_err("verify-full must reject a certificate for another host");
338 assert_eq!(error.kind(), io::ErrorKind::InvalidData);
339 }
340
341 #[tokio::test]
342 async fn negotiates_tls_and_exposes_equal_channel_bindings() {
343 let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
344 let generated = generate_simple_self_signed(["localhost".into()]).unwrap();
345 let certificate = CertificateDer::from(generated.cert.der().to_vec());
346 let key = PrivateKeyDer::try_from(generated.signing_key.serialize_der()).unwrap();
347
348 let server_config = ServerConfig::builder()
349 .with_no_client_auth()
350 .with_single_cert(vec![certificate.clone()], key)
351 .unwrap();
352 let mut roots = RootCertStore::empty();
353 roots.add(certificate.clone()).unwrap();
354 let client_config = ClientConfig::builder()
355 .with_root_certificates(roots)
356 .with_no_client_auth();
357 let expected = channel_binding(&certificate).unwrap();
358 let (client_io, server_io) = tokio::io::duplex(16 * 1024);
359
360 let (client, server) = tokio::join!(
361 connect(
362 client_io,
363 ServerName::try_from("localhost").unwrap(),
364 Arc::new(client_config),
365 ),
366 accept(server_io, Arc::new(server_config), &certificate),
367 );
368 let client = client.unwrap();
369 let server = server.unwrap();
370
371 assert_eq!(client.tls_server_end_point(), expected);
372 assert_eq!(server.tls_server_end_point(), expected);
373 }
374
375 #[tokio::test]
376 async fn typed_pre_startup_upgrade_changes_the_buffered_transport() {
377 let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
378 let generated = generate_simple_self_signed(["localhost".into()]).unwrap();
379 let certificate = CertificateDer::from(generated.cert.der().to_vec());
380 let key = PrivateKeyDer::try_from(generated.signing_key.serialize_der()).unwrap();
381 let server_config = Arc::new(
382 ServerConfig::builder()
383 .with_no_client_auth()
384 .with_single_cert(vec![certificate.clone()], key)
385 .unwrap(),
386 );
387 let mut roots = RootCertStore::empty();
388 roots.add(certificate.clone()).unwrap();
389 let client_config = Arc::new(
390 ClientConfig::builder()
391 .with_root_certificates(roots)
392 .with_no_client_auth(),
393 );
394 let expected = channel_binding(&certificate).unwrap();
395 let (client_io, mut server_io) = tokio::io::duplex(16 * 1024);
396
397 let server = tokio::spawn(async move {
398 let mut request = [0; 8];
399 server_io.read_exact(&mut request).await.unwrap();
400 assert_eq!(request, [0, 0, 0, 8, 4, 210, 22, 47]);
401 server_io.write_all(b"S").await.unwrap();
402 accept(server_io, server_config, &certificate)
403 .await
404 .unwrap()
405 });
406
407 let mut awaiting_reply = Conn::new(Buffered::new(client_io)).request_ssl();
408 awaiting_reply.flush().await.unwrap();
409 let Negotiation::Accepted(handshake) = awaiting_reply.receive_ssl_reply().await.unwrap()
410 else {
411 panic!("test server rejected TLS")
412 };
413 let pre_startup = handshake
414 .connect_tls(ServerName::try_from("localhost").unwrap(), client_config)
415 .await
416 .unwrap();
417 let client = pre_startup.into_transport().into_inner();
418 let server = server.await.unwrap();
419
420 assert_eq!(client.tls_server_end_point(), expected);
421 assert_eq!(server.tls_server_end_point(), expected);
422 }
423
424 #[tokio::test]
425 async fn server_role_terminates_tls_before_accepting_startup() {
426 let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
427 let generated = generate_simple_self_signed(["localhost".into()]).unwrap();
428 let certificate = CertificateDer::from(generated.cert.der().to_vec());
429 let key = PrivateKeyDer::try_from(generated.signing_key.serialize_der()).unwrap();
430 let server_config = Arc::new(
431 ServerConfig::builder()
432 .with_no_client_auth()
433 .with_single_cert(vec![certificate.clone()], key)
434 .unwrap(),
435 );
436 let mut roots = RootCertStore::empty();
437 roots.add(certificate.clone()).unwrap();
438 let client_config = Arc::new(
439 ClientConfig::builder()
440 .with_root_certificates(roots)
441 .with_no_client_auth(),
442 );
443 let expected = channel_binding(&certificate).unwrap();
444 let (mut client_io, server_io) = tokio::io::duplex(16 * 1024);
445
446 let client = tokio::spawn(async move {
447 client_io
448 .write_all(&[0, 0, 0, 8, 4, 210, 22, 47])
449 .await
450 .unwrap();
451 assert_eq!(client_io.read_u8().await.unwrap(), b'S');
452 connect(
453 client_io,
454 ServerName::try_from("localhost").unwrap(),
455 client_config,
456 )
457 .await
458 .unwrap()
459 });
460
461 let mut pre_startup = Conn::new(Buffered::new_frontend(server_io));
462 let message = pre_startup.receive_pre_startup_wire().await.unwrap();
463 let PreStartupOffer::Ssl(decision) = pre_startup.offer_pre_startup(message) else {
464 panic!("test client did not request TLS")
465 };
466 let mut handshake = decision.approve_ssl();
467 handshake.flush().await.unwrap();
468 let pre_startup = handshake
469 .accept_tls(server_config, certificate)
470 .await
471 .unwrap();
472 let server = pre_startup.into_transport().into_inner();
473 let client = client.await.unwrap();
474
475 assert_eq!(client.tls_server_end_point(), expected);
476 assert_eq!(server.tls_server_end_point(), expected);
477 }
478
479 #[test]
480 fn channel_binding_uses_the_certificate_signature_digest() {
481 let key = KeyPair::generate_for(&PKCS_ECDSA_P384_SHA384).unwrap();
482 let params = CertificateParams::new(["localhost".to_owned()]).unwrap();
483 let generated = params.self_signed(&key).unwrap();
484 let certificate = CertificateDer::from(generated.der().to_vec());
485
486 assert_eq!(
487 channel_binding(&certificate).unwrap(),
488 Sha384::digest(certificate.as_ref()).to_vec()
489 );
490 }
491}