1pub use crate::verifier::nitro::AttestationDocument;
14use crate::verifier::{self, Policy, TrustStore, VerifiedEvidence};
15use axum::http::{HeaderMap, Method, Request, StatusCode, Uri};
16use bytes::Buf;
17use log::{debug, info};
18use quinn::Endpoint;
19use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
20use rustls::crypto::WebPkiSupportedAlgorithms;
21use rustls::Error as RustlsError;
22use rustls_pki_types::{CertificateDer, ServerName, UnixTime};
23use serde::Serialize;
24use sha2::{Digest, Sha256};
25use std::collections::BTreeMap;
26use std::net::SocketAddr;
27use std::sync::{Arc, Mutex};
28use std::time::Duration;
29use x509_parser::prelude::*;
30
31#[derive(Debug, Clone)]
50pub struct EnclaveCertVerifier {
51 received_cert: Arc<Mutex<Option<CertificateDer<'static>>>>,
52 verified_evidence: Arc<Mutex<Option<VerifiedEvidence>>>,
53 expected_measurements: BTreeMap<String, Vec<u8>>,
54 policy: Policy,
55 trust: Arc<TrustStore>,
56 algorithms: WebPkiSupportedAlgorithms,
57}
58
59impl EnclaveCertVerifier {
61 pub fn new() -> Self {
64 Self {
65 received_cert: Arc::new(Mutex::new(None)),
66 verified_evidence: Arc::new(Mutex::new(None)),
67 expected_measurements: BTreeMap::new(),
68 policy: Policy::default(),
69 trust: Arc::new(TrustStore::builtin()),
70 algorithms: rustls::crypto::ring::default_provider().signature_verification_algorithms,
71 }
72 }
73
74 pub fn with_expected_measurement(
79 mut self,
80 name: impl Into<String>,
81 value: impl Into<Vec<u8>>,
82 ) -> Self {
83 self.expected_measurements
84 .insert(name.into().to_lowercase(), value.into());
85 self
86 }
87
88 pub fn with_expected_pcr(self, index: usize, value: impl Into<Vec<u8>>) -> Self {
90 self.with_expected_measurement(format!("pcr{index}"), value)
91 }
92
93 pub fn with_trust_store(mut self, trust: TrustStore) -> Self {
95 self.trust = Arc::new(trust);
96 self
97 }
98
99 pub fn allow_mock(mut self) -> Self {
104 self.policy.allow_mock = true;
105 self
106 }
107
108 pub fn allow_debug(mut self) -> Self {
110 self.policy.allow_debug = true;
111 self
112 }
113
114 pub fn received_certificate(&self) -> Option<CertificateDer<'static>> {
116 self.received_cert
117 .lock()
118 .ok()
119 .and_then(|guard| guard.clone())
120 }
121
122 pub fn verified_evidence(&self) -> Option<VerifiedEvidence> {
124 self.verified_evidence
125 .lock()
126 .ok()
127 .and_then(|guard| guard.clone())
128 }
129
130 pub fn verified_attestation(&self) -> Option<AttestationDocument> {
133 self.verified_evidence().and_then(|evidence| evidence.nitro)
134 }
135
136 fn verify(
138 &self,
139 end_entity: &CertificateDer<'_>,
140 now: UnixTime,
141 ) -> Result<VerifiedEvidence, String> {
142 let (_, cert) = X509Certificate::from_der(end_entity.as_ref())
144 .map_err(|e| format!("malformed certificate: {e}"))?;
145 let now_secs = now.as_secs() as i64;
146 if now_secs < cert.validity().not_before.timestamp() {
147 return Err("certificate is not valid yet".into());
148 }
149 if now_secs > cert.validity().not_after.timestamp() {
150 return Err("certificate has expired".into());
151 }
152 cert.verify_signature(None)
153 .map_err(|e| format!("certificate is not correctly self-signed: {e}"))?;
154
155 let eat_bytes = extract_attestation_doc(end_entity.as_ref())
157 .map_err(|e| format!("missing attestation extension: {e}"))?;
158 let binding = Sha256::digest(cert.public_key().raw);
159 let evidence =
160 verifier::verify_evidence(&eat_bytes, &binding, now, &self.trust, self.policy)?;
161
162 for (name, expected) in &self.expected_measurements {
164 match evidence.measurements.get(name) {
165 Some(actual) if actual == expected => {}
166 Some(_) => {
167 return Err(format!(
168 "{} does not match the expected value",
169 name.to_uppercase()
170 ))
171 }
172 None => {
173 return Err(format!(
174 "{} evidence has no measurement '{name}'",
175 evidence.tee
176 ))
177 }
178 }
179 }
180
181 Ok(evidence)
182 }
183}
184
185impl Default for EnclaveCertVerifier {
187 fn default() -> Self {
189 Self::new()
190 }
191}
192
193const ATTESTATION_OID: &[u64] = &[1, 3, 6, 1, 4, 1, 99999, 1];
195
196impl ServerCertVerifier for EnclaveCertVerifier {
198 fn verify_server_cert(
202 &self,
203 end_entity: &CertificateDer<'_>,
204 _intermediates: &[CertificateDer<'_>],
205 _server_name: &ServerName<'_>,
206 _ocsp_response: &[u8],
207 now: UnixTime,
208 ) -> Result<ServerCertVerified, rustls::Error> {
209 let evidence = self
210 .verify(end_entity, now)
211 .map_err(|e| RustlsError::General(format!("RA-TLS verification failed: {e}")))?;
212 debug!("{} attestation verified", evidence.tee);
213
214 if let Ok(mut guard) = self.received_cert.lock() {
215 *guard = Some(end_entity.clone().into_owned());
216 }
217 if let Ok(mut guard) = self.verified_evidence.lock() {
218 *guard = Some(evidence);
219 }
220 Ok(ServerCertVerified::assertion())
221 }
222
223 fn verify_tls12_signature(
225 &self,
226 message: &[u8],
227 cert: &CertificateDer<'_>,
228 dss: &rustls::DigitallySignedStruct,
229 ) -> Result<HandshakeSignatureValid, rustls::Error> {
230 rustls::crypto::verify_tls12_signature(message, cert, dss, &self.algorithms)
231 }
232
233 fn verify_tls13_signature(
235 &self,
236 message: &[u8],
237 cert: &CertificateDer<'_>,
238 dss: &rustls::DigitallySignedStruct,
239 ) -> Result<HandshakeSignatureValid, rustls::Error> {
240 rustls::crypto::verify_tls13_signature(message, cert, dss, &self.algorithms)
241 }
242
243 fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
245 self.algorithms.supported_schemes()
246 }
247}
248
249pub fn extract_attestation_doc(
251 cert_der: &[u8],
252) -> Result<Vec<u8>, Box<dyn std::error::Error + Send + Sync>> {
253 let (_, cert) = X509Certificate::from_der(cert_der)?;
255
256 for ext in cert.extensions() {
258 if ext
259 .oid
260 .iter()
261 .into_iter()
262 .flatten()
263 .eq(ATTESTATION_OID.iter().copied())
264 {
265 let raw_value = ext.value;
266
267 if !raw_value.is_empty() && raw_value[0] == 0x04 {
269 let (_, octet_string) = der_parser::der::parse_der_octetstring(raw_value)?;
270 return Ok(octet_string.as_slice()?.to_vec());
271 }
272
273 return Ok(raw_value.to_vec());
274 }
275 }
276
277 Err("Attestation extension OID not found in certificate".into())
278}
279
280#[derive(Debug, Clone)]
282pub struct ClientResponse {
283 pub status: StatusCode,
284 pub headers: HeaderMap,
285 pub body: Vec<u8>,
286}
287
288impl ClientResponse {
290 pub fn text(&self) -> Result<String, std::string::FromUtf8Error> {
292 String::from_utf8(self.body.clone())
293 }
294}
295
296pub struct TtkClient {
298 endpoint: Endpoint,
299 send_request: h3::client::SendRequest<h3_quinn::OpenStreams, axum::body::Bytes>,
300 driver_handle: tokio::task::JoinHandle<Result<(), h3::Error>>,
301 server_addr: SocketAddr,
302 server_name: String,
303 peer_cert: Option<CertificateDer<'static>>,
304}
305
306impl TtkClient {
308 pub async fn connect(
311 server_addr: SocketAddr,
312 server_name: &str,
313 ) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
314 Self::connect_with_verifier(server_addr, server_name, EnclaveCertVerifier::new()).await
315 }
316
317 pub async fn connect_with_verifier(
319 server_addr: SocketAddr,
320 server_name: &str,
321 verifier: EnclaveCertVerifier,
322 ) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
323 let _ = rustls::crypto::ring::default_provider().install_default();
325
326 let cert_verifier = Arc::new(verifier);
327
328 let mut client_crypto = rustls::ClientConfig::builder()
329 .dangerous()
330 .with_custom_certificate_verifier(cert_verifier.clone())
331 .with_no_client_auth();
332
333 client_crypto.alpn_protocols = vec![b"h3".to_vec()];
335
336 let quic_client_config = quinn::ClientConfig::new(Arc::new(
337 quinn::crypto::rustls::QuicClientConfig::try_from(client_crypto)?,
338 ));
339
340 let bind_addr: SocketAddr = if server_addr.is_ipv6() {
342 "[::]:0".parse()?
343 } else {
344 "0.0.0.0:0".parse()?
345 };
346 let mut endpoint = Endpoint::client(bind_addr)?;
347 endpoint.set_default_client_config(quic_client_config);
348
349 info!(
350 "Initiating QUIC connection to {} ({})",
351 server_addr, server_name
352 );
353 let connecting = endpoint.connect(server_addr, server_name)?;
354 let connection = connecting.await?;
355 info!("QUIC connection established with {}", server_addr);
356
357 let peer_cert = cert_verifier.received_certificate();
358 if let Some(ref cert) = peer_cert {
359 let hash = Sha256::digest(cert.as_ref());
360 info!(
361 "Server Certificate SHA-256 fingerprint: {}",
362 hex_encode(&hash)
363 );
364 }
365
366 let h3_quic_conn = h3_quinn::Connection::new(connection);
368 let (mut driver, send_request) = h3::client::new(h3_quic_conn).await?;
369
370 let driver_handle =
372 tokio::spawn(async move { std::future::poll_fn(|cx| driver.poll_close(cx)).await });
373
374 Ok(Self {
375 endpoint,
376 send_request,
377 driver_handle,
378 server_addr,
379 server_name: server_name.to_string(),
380 peer_cert,
381 })
382 }
383
384 pub fn server_addr(&self) -> SocketAddr {
386 self.server_addr
387 }
388
389 pub fn server_name(&self) -> &str {
391 &self.server_name
392 }
393
394 pub fn peer_cert(&self) -> Option<&CertificateDer<'static>> {
396 self.peer_cert.as_ref()
397 }
398
399 pub fn peer_cert_sha256(&self) -> Option<[u8; 32]> {
401 self.peer_cert.as_ref().map(|c| {
402 let mut arr = [0u8; 32];
403 arr.copy_from_slice(&Sha256::digest(c.as_ref()));
404 arr
405 })
406 }
407
408 pub fn peer_cert_sha256_hex(&self) -> Option<String> {
410 self.peer_cert_sha256().map(|h| hex_encode(&h))
411 }
412
413 pub async fn get(
415 &mut self,
416 path: &str,
417 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
418 let uri: Uri = if path.starts_with('/') {
419 format!("https://{}{}", self.server_name, path).parse()?
420 } else {
421 format!("https://{}/{}", self.server_name, path).parse()?
422 };
423
424 let req = Request::builder()
425 .method(Method::GET)
426 .uri(uri)
427 .header("Host", &self.server_name)
428 .header("User-Agent", "TTKClient/0.6.0")
429 .header("Accept", "*/*")
430 .body(())?;
431
432 self.send(req, None).await
433 }
434
435 pub async fn post(
437 &mut self,
438 path: &str,
439 body: &[u8],
440 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
441 self.post_with_content_type(path, "application/octet-stream", body)
442 .await
443 }
444
445 pub async fn post_json<T: Serialize + ?Sized>(
447 &mut self,
448 path: &str,
449 value: &T,
450 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
451 let body = serde_json::to_vec(value)?;
452 self.post_with_content_type(path, "application/json", &body)
453 .await
454 }
455
456 async fn post_with_content_type(
458 &mut self,
459 path: &str,
460 content_type: &str,
461 body: &[u8],
462 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
463 let uri: Uri = if path.starts_with('/') {
464 format!("https://{}{}", self.server_name, path).parse()?
465 } else {
466 format!("https://{}/{}", self.server_name, path).parse()?
467 };
468
469 let req = Request::builder()
470 .method(Method::POST)
471 .uri(uri)
472 .header("Host", &self.server_name)
473 .header("User-Agent", "TTKClient/0.6.0")
474 .header("Content-Type", content_type)
475 .header("Content-Length", body.len().to_string())
476 .body(())?;
477
478 self.send(req, Some(body)).await
479 }
480
481 pub async fn send(
483 &mut self,
484 req: Request<()>,
485 payload: Option<&[u8]>,
486 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
487 debug!("Sending HTTP/3 request: {} {}", req.method(), req.uri());
488 let mut stream = self.send_request.send_request(req).await?;
489
490 if let Some(data) = payload {
491 if !data.is_empty() {
492 stream
493 .send_data(axum::body::Bytes::copy_from_slice(data))
494 .await?;
495 }
496 }
497 stream.finish().await?;
498
499 let response = stream.recv_response().await?;
500 let status = response.status();
501 let headers = response.headers().clone();
502
503 let mut body = Vec::new();
504 while let Some(mut chunk) = stream.recv_data().await? {
505 while chunk.has_remaining() {
506 let slice = chunk.chunk();
507 body.extend_from_slice(slice);
508 let len = slice.len();
509 chunk.advance(len);
510 }
511 }
512
513 Ok(ClientResponse {
514 status,
515 headers,
516 body,
517 })
518 }
519
520 pub async fn close(self) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
522 drop(self.send_request);
523 let _ = tokio::time::timeout(Duration::from_secs(1), self.driver_handle).await;
525 self.endpoint.wait_idle().await;
526 Ok(())
527 }
528}
529
530pub fn hex_encode(bytes: &[u8]) -> String {
532 bytes.iter().map(|b| format!("{:02x}", b)).collect()
533}