Skip to content

Commit 6b92a59

Browse files
committed
Support client auth
1 parent 6ceced8 commit 6b92a59

1 file changed

Lines changed: 206 additions & 32 deletions

File tree

crates/attested-tls/src/lib.rs

Lines changed: 206 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@ use rustls::{
2121
RootCertStore,
2222
SignatureScheme,
2323
client::{
24+
ResolvesClientCert,
2425
VerifierBuilderError,
2526
WebPkiServerVerifier,
2627
danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier},
@@ -34,7 +35,11 @@ use rustls::{
3435
UnixTime,
3536
pem::PemObject,
3637
},
37-
server::ResolvesServerCert,
38+
server::{
39+
ResolvesServerCert,
40+
WebPkiClientVerifier,
41+
danger::{ClientCertVerified, ClientCertVerifier},
42+
},
3843
sign::{CertifiedKey, SigningKey},
3944
};
4045
use sha2::{Digest as _, Sha512};
@@ -131,6 +136,8 @@ impl AttestedCertificateResolver {
131136
Ok(Self { state })
132137
}
133138

139+
/// Create an attested certificate chain - either self-signed or with
140+
/// the provided CA
134141
async fn issue_ra_cert_chain(
135142
key: &KeyPair,
136143
ca: Option<&CaCert>,
@@ -151,7 +158,7 @@ impl AttestedCertificateResolver {
151158
.not_before(now)
152159
.not_after(not_after)
153160
.usage_server_auth(true)
154-
.usage_client_auth(false)
161+
.usage_client_auth(true)
155162
.attestation(&attestation)
156163
.build();
157164

@@ -248,8 +255,23 @@ impl AttestedCertificateResolver {
248255

249256
impl ResolvesServerCert for AttestedCertificateResolver {
250257
fn resolve(&self, _: rustls::server::ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
251-
let certificate = self.state.certificate.read().expect("certificate lock poisoned").clone();
258+
self.current_certified_key()
259+
}
260+
}
261+
262+
impl ResolvesClientCert for AttestedCertificateResolver {
263+
fn resolve(&self, _: &[&[u8]], _: &[SignatureScheme]) -> Option<Arc<CertifiedKey>> {
264+
self.current_certified_key()
265+
}
266+
267+
fn has_certs(&self) -> bool {
268+
!self.state.certificate.read().expect("certificate lock poisoned").is_empty()
269+
}
270+
}
252271

272+
impl AttestedCertificateResolver {
273+
fn current_certified_key(&self) -> Option<Arc<CertifiedKey>> {
274+
let certificate = self.state.certificate.read().expect("certificate lock poisoned").clone();
253275
Some(Arc::new(CertifiedKey::new(certificate, self.state.key.clone())))
254276
}
255277
}
@@ -286,7 +308,8 @@ fn create_report_data(
286308

287309
#[derive(Debug)]
288310
pub struct AttestedCertificateVerifier {
289-
inner: Arc<WebPkiServerVerifier>,
311+
server_inner: Arc<WebPkiServerVerifier>,
312+
client_inner: Arc<dyn ClientCertVerifier>,
290313
attestation_verifier: AttestationVerifier,
291314
}
292315

@@ -303,24 +326,29 @@ impl AttestedCertificateVerifier {
303326
attestation_verifier: AttestationVerifier,
304327
provider: Arc<CryptoProvider>,
305328
) -> Result<Self, AttestedTlsError> {
306-
let inner = WebPkiServerVerifier::builder_with_provider(root_store.into(), provider)
329+
let root_store = Arc::new(root_store);
330+
let server_inner =
331+
WebPkiServerVerifier::builder_with_provider(root_store.clone(), provider.clone())
332+
.build()
333+
.map_err(AttestedTlsError::VerifierBuilder)?;
334+
let client_inner = WebPkiClientVerifier::builder_with_provider(root_store, provider)
307335
.build()
308336
.map_err(AttestedTlsError::VerifierBuilder)?;
309337

310-
Ok(Self { inner, attestation_verifier })
338+
Ok(Self { server_inner, client_inner, attestation_verifier })
311339
}
312340

313341
fn extract_custom_attestation_from_cert(
314342
cert: &CertificateDer<'_>,
315343
) -> Result<AttestationExchangeMessage, rustls::Error> {
316344
// First try to parse using ra_tls which assumes DCAP
317-
if let Ok(Some(attestation)) = ra_tls::attestation::from_der(cert.as_ref()) {
318-
if let AttestationQuote::DstackTdx(tdx_quote) = attestation.quote {
319-
return Ok(AttestationExchangeMessage {
320-
attestation_type: AttestationType::DcapTdx,
321-
attestation: tdx_quote.quote,
322-
});
323-
}
345+
if let Ok(Some(attestation)) = ra_tls::attestation::from_der(cert.as_ref()) &&
346+
let AttestationQuote::DstackTdx(tdx_quote) = attestation.quote
347+
{
348+
return Ok(AttestationExchangeMessage {
349+
attestation_type: AttestationType::DcapTdx,
350+
attestation: tdx_quote.quote,
351+
});
324352
}
325353

326354
// If that fails, extract and parse the extension
@@ -396,6 +424,25 @@ impl AttestedCertificateVerifier {
396424

397425
Ok(())
398426
}
427+
428+
fn verify_attestation_binding(
429+
&self,
430+
end_entity: &CertificateDer<'_>,
431+
) -> Result<(), rustls::Error> {
432+
let expected_input_data = Self::expected_input_data_from_cert(end_entity)?;
433+
let attestation = Self::extract_custom_attestation_from_cert(end_entity)?;
434+
435+
tokio::task::block_in_place(|| {
436+
tokio::runtime::Handle::current().block_on(async {
437+
self.attestation_verifier
438+
.verify_attestation(attestation, expected_input_data)
439+
.await
440+
.unwrap();
441+
})
442+
});
443+
444+
Ok(())
445+
}
399446
}
400447

401448
impl ServerCertVerifier for AttestedCertificateVerifier {
@@ -407,7 +454,7 @@ impl ServerCertVerifier for AttestedCertificateVerifier {
407454
ocsp_response: &[u8],
408455
now: UnixTime,
409456
) -> Result<ServerCertVerified, rustls::Error> {
410-
match self.inner.verify_server_cert(
457+
match self.server_inner.verify_server_cert(
411458
end_entity,
412459
intermediates,
413460
server_name,
@@ -421,20 +468,7 @@ impl ServerCertVerifier for AttestedCertificateVerifier {
421468
Err(err) => return Err(err),
422469
Ok(_) => {}
423470
};
424-
let expected_input_data =
425-
AttestedCertificateVerifier::expected_input_data_from_cert(end_entity)?;
426-
let attestation =
427-
AttestedCertificateVerifier::extract_custom_attestation_from_cert(end_entity)?;
428-
429-
// Block when calling the verify function as it is async
430-
tokio::task::block_in_place(|| {
431-
tokio::runtime::Handle::current().block_on(async {
432-
self.attestation_verifier
433-
.verify_attestation(attestation, expected_input_data)
434-
.await
435-
.unwrap();
436-
})
437-
});
471+
self.verify_attestation_binding(end_entity)?;
438472
Ok(ServerCertVerified::assertion())
439473
}
440474

@@ -444,7 +478,7 @@ impl ServerCertVerifier for AttestedCertificateVerifier {
444478
cert: &CertificateDer<'_>,
445479
dss: &DigitallySignedStruct,
446480
) -> Result<HandshakeSignatureValid, rustls::Error> {
447-
self.inner.verify_tls12_signature(message, cert, dss)
481+
self.server_inner.verify_tls12_signature(message, cert, dss)
448482
}
449483

450484
fn verify_tls13_signature(
@@ -453,15 +487,68 @@ impl ServerCertVerifier for AttestedCertificateVerifier {
453487
cert: &CertificateDer<'_>,
454488
dss: &DigitallySignedStruct,
455489
) -> Result<HandshakeSignatureValid, rustls::Error> {
456-
self.inner.verify_tls13_signature(message, cert, dss)
490+
self.server_inner.verify_tls13_signature(message, cert, dss)
457491
}
458492

459493
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
460-
self.inner.supported_verify_schemes()
494+
self.server_inner.supported_verify_schemes()
461495
}
462496

463497
fn root_hint_subjects(&self) -> Option<&[DistinguishedName]> {
464-
self.inner.root_hint_subjects()
498+
self.server_inner.root_hint_subjects()
499+
}
500+
}
501+
502+
impl ClientCertVerifier for AttestedCertificateVerifier {
503+
fn offer_client_auth(&self) -> bool {
504+
self.client_inner.offer_client_auth()
505+
}
506+
507+
fn client_auth_mandatory(&self) -> bool {
508+
self.client_inner.client_auth_mandatory()
509+
}
510+
511+
fn root_hint_subjects(&self) -> &[DistinguishedName] {
512+
self.client_inner.root_hint_subjects()
513+
}
514+
515+
fn verify_client_cert(
516+
&self,
517+
end_entity: &CertificateDer<'_>,
518+
intermediates: &[CertificateDer<'_>],
519+
now: UnixTime,
520+
) -> Result<ClientCertVerified, rustls::Error> {
521+
match self.client_inner.verify_client_cert(end_entity, intermediates, now) {
522+
Err(rustls::Error::InvalidCertificate(rustls::CertificateError::UnknownIssuer)) => {
523+
Self::verify_cert_time_validity(end_entity, now)?;
524+
}
525+
Err(err) => return Err(err),
526+
Ok(_) => {}
527+
};
528+
self.verify_attestation_binding(end_entity)?;
529+
Ok(ClientCertVerified::assertion())
530+
}
531+
532+
fn verify_tls12_signature(
533+
&self,
534+
message: &[u8],
535+
cert: &CertificateDer<'_>,
536+
dss: &DigitallySignedStruct,
537+
) -> Result<HandshakeSignatureValid, rustls::Error> {
538+
self.client_inner.verify_tls12_signature(message, cert, dss)
539+
}
540+
541+
fn verify_tls13_signature(
542+
&self,
543+
message: &[u8],
544+
cert: &CertificateDer<'_>,
545+
dss: &DigitallySignedStruct,
546+
) -> Result<HandshakeSignatureValid, rustls::Error> {
547+
self.client_inner.verify_tls13_signature(message, cert, dss)
548+
}
549+
550+
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
551+
self.client_inner.supported_verify_schemes()
465552
}
466553
}
467554

@@ -677,6 +764,93 @@ mod tests {
677764
assert_ne!(initial_certificate.as_ref(), renewed_certificate.as_ref());
678765
}
679766

767+
#[tokio::test(flavor = "multi_thread")]
768+
async fn server_and_client_configs_complete_a_mutual_auth_handshake() {
769+
let provider: Arc<CryptoProvider> = aws_lc_rs::default_provider().into();
770+
771+
let server_resolver = AttestedCertificateResolver::new_with_provider(
772+
AttestationGenerator::new(AttestationType::DcapTdx, None)
773+
.expect("mock generator construction should succeed"),
774+
None,
775+
provider.clone(),
776+
)
777+
.await
778+
.expect("server resolver construction should succeed");
779+
let client_resolver = AttestedCertificateResolver::new_with_provider(
780+
AttestationGenerator::new(AttestationType::DcapTdx, None)
781+
.expect("mock generator construction should succeed"),
782+
None,
783+
provider.clone(),
784+
)
785+
.await
786+
.expect("client resolver construction should succeed");
787+
788+
let server_certificate = server_resolver
789+
.state
790+
.certificate
791+
.read()
792+
.expect("certificate lock poisoned")
793+
.first()
794+
.expect("resolver should hold a certificate")
795+
.clone();
796+
let client_certificate = client_resolver
797+
.state
798+
.certificate
799+
.read()
800+
.expect("certificate lock poisoned")
801+
.first()
802+
.expect("resolver should hold a certificate")
803+
.clone();
804+
805+
let mut client_roots = RootCertStore::empty();
806+
client_roots.add(server_certificate).expect("server certificate should be trusted");
807+
let mut server_roots = RootCertStore::empty();
808+
server_roots.add(client_certificate).expect("client certificate should be trusted");
809+
810+
let server_verifier = AttestedCertificateVerifier::new_with_provider(
811+
server_roots,
812+
AttestationVerifier::mock(),
813+
provider.clone(),
814+
)
815+
.expect("server verifier construction should succeed");
816+
let client_verifier = AttestedCertificateVerifier::new_with_provider(
817+
client_roots,
818+
AttestationVerifier::mock(),
819+
provider.clone(),
820+
)
821+
.expect("client verifier construction should succeed");
822+
823+
let server_config = ServerConfig::builder_with_provider(provider.clone())
824+
.with_safe_default_protocol_versions()
825+
.expect("server config should support default protocol versions")
826+
.with_client_cert_verifier(Arc::new(server_verifier))
827+
.with_cert_resolver(Arc::new(server_resolver));
828+
let client_config = ClientConfig::builder_with_provider(provider)
829+
.with_safe_default_protocol_versions()
830+
.expect("client config should support default protocol versions")
831+
.dangerous()
832+
.with_custom_certificate_verifier(Arc::new(client_verifier))
833+
.with_client_cert_resolver(Arc::new(client_resolver));
834+
835+
let mut client = ClientConnection::new(
836+
Arc::new(client_config),
837+
ServerName::try_from("foo").expect("server name should be valid"),
838+
)
839+
.expect("client connection should be created");
840+
let mut server =
841+
ServerConnection::new(Arc::new(server_config)).expect("server connection should exist");
842+
843+
while client.is_handshaking() || server.is_handshaking() {
844+
transfer_tls_client_to_server(&mut client, &mut server);
845+
transfer_tls_server_to_client(&mut server, &mut client);
846+
}
847+
848+
assert!(!client.is_handshaking());
849+
assert!(!server.is_handshaking());
850+
assert!(client.peer_certificates().is_some());
851+
assert!(server.peer_certificates().is_some());
852+
}
853+
680854
fn test_ca() -> CaCert {
681855
let key = KeyPair::generate_for(&PKCS_ECDSA_P256_SHA256)
682856
.expect("test CA key generation should succeed");

0 commit comments

Comments
 (0)