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 sha2::{Digest, Sha256};
24use std::collections::BTreeMap;
25use std::net::SocketAddr;
26use std::sync::{Arc, Mutex};
27use std::time::Duration;
28use x509_parser::prelude::*;
29
30#[derive(Debug, Clone)]
49pub struct EnclaveCertVerifier {
50 received_cert: Arc<Mutex<Option<CertificateDer<'static>>>>,
51 verified_evidence: Arc<Mutex<Option<VerifiedEvidence>>>,
52 expected_measurements: BTreeMap<String, Vec<u8>>,
53 policy: Policy,
54 trust: Arc<TrustStore>,
55 algorithms: WebPkiSupportedAlgorithms,
56}
57
58impl EnclaveCertVerifier {
60 pub fn new() -> Self {
63 Self {
64 received_cert: Arc::new(Mutex::new(None)),
65 verified_evidence: Arc::new(Mutex::new(None)),
66 expected_measurements: BTreeMap::new(),
67 policy: Policy::default(),
68 trust: Arc::new(TrustStore::builtin()),
69 algorithms: rustls::crypto::ring::default_provider().signature_verification_algorithms,
70 }
71 }
72
73 pub fn with_expected_measurement(
78 mut self,
79 name: impl Into<String>,
80 value: impl Into<Vec<u8>>,
81 ) -> Self {
82 self.expected_measurements
83 .insert(name.into().to_lowercase(), value.into());
84 self
85 }
86
87 pub fn with_expected_pcr(self, index: usize, value: impl Into<Vec<u8>>) -> Self {
89 self.with_expected_measurement(format!("pcr{index}"), value)
90 }
91
92 pub fn with_trust_store(mut self, trust: TrustStore) -> Self {
94 self.trust = Arc::new(trust);
95 self
96 }
97
98 pub fn allow_mock(mut self) -> Self {
103 self.policy.allow_mock = true;
104 self
105 }
106
107 pub fn allow_debug(mut self) -> Self {
109 self.policy.allow_debug = true;
110 self
111 }
112
113 pub fn received_certificate(&self) -> Option<CertificateDer<'static>> {
115 self.received_cert
116 .lock()
117 .ok()
118 .and_then(|guard| guard.clone())
119 }
120
121 pub fn verified_evidence(&self) -> Option<VerifiedEvidence> {
123 self.verified_evidence
124 .lock()
125 .ok()
126 .and_then(|guard| guard.clone())
127 }
128
129 pub fn verified_attestation(&self) -> Option<AttestationDocument> {
132 self.verified_evidence().and_then(|evidence| evidence.nitro)
133 }
134
135 fn verify(
137 &self,
138 end_entity: &CertificateDer<'_>,
139 now: UnixTime,
140 ) -> Result<VerifiedEvidence, String> {
141 let (_, cert) = X509Certificate::from_der(end_entity.as_ref())
143 .map_err(|e| format!("malformed certificate: {e}"))?;
144 let now_secs = now.as_secs() as i64;
145 if now_secs < cert.validity().not_before.timestamp() {
146 return Err("certificate is not valid yet".into());
147 }
148 if now_secs > cert.validity().not_after.timestamp() {
149 return Err("certificate has expired".into());
150 }
151 cert.verify_signature(None)
152 .map_err(|e| format!("certificate is not correctly self-signed: {e}"))?;
153
154 let eat_bytes = extract_attestation_doc(end_entity.as_ref())
156 .map_err(|e| format!("missing attestation extension: {e}"))?;
157 let binding = Sha256::digest(cert.public_key().raw);
158 let evidence =
159 verifier::verify_evidence(&eat_bytes, &binding, now, &self.trust, self.policy)?;
160
161 for (name, expected) in &self.expected_measurements {
163 match evidence.measurements.get(name) {
164 Some(actual) if actual == expected => {}
165 Some(_) => {
166 return Err(format!(
167 "{} does not match the expected value",
168 name.to_uppercase()
169 ))
170 }
171 None => {
172 return Err(format!(
173 "{} evidence has no measurement '{name}'",
174 evidence.tee
175 ))
176 }
177 }
178 }
179
180 Ok(evidence)
181 }
182}
183
184impl Default for EnclaveCertVerifier {
186 fn default() -> Self {
188 Self::new()
189 }
190}
191
192const ATTESTATION_OID: &[u64] = &[1, 3, 6, 1, 4, 1, 99999, 1];
194
195impl ServerCertVerifier for EnclaveCertVerifier {
197 fn verify_server_cert(
201 &self,
202 end_entity: &CertificateDer<'_>,
203 _intermediates: &[CertificateDer<'_>],
204 _server_name: &ServerName<'_>,
205 _ocsp_response: &[u8],
206 now: UnixTime,
207 ) -> Result<ServerCertVerified, rustls::Error> {
208 let evidence = self
209 .verify(end_entity, now)
210 .map_err(|e| RustlsError::General(format!("RA-TLS verification failed: {e}")))?;
211 debug!("{} attestation verified", evidence.tee);
212
213 if let Ok(mut guard) = self.received_cert.lock() {
214 *guard = Some(end_entity.clone().into_owned());
215 }
216 if let Ok(mut guard) = self.verified_evidence.lock() {
217 *guard = Some(evidence);
218 }
219 Ok(ServerCertVerified::assertion())
220 }
221
222 fn verify_tls12_signature(
224 &self,
225 message: &[u8],
226 cert: &CertificateDer<'_>,
227 dss: &rustls::DigitallySignedStruct,
228 ) -> Result<HandshakeSignatureValid, rustls::Error> {
229 rustls::crypto::verify_tls12_signature(message, cert, dss, &self.algorithms)
230 }
231
232 fn verify_tls13_signature(
234 &self,
235 message: &[u8],
236 cert: &CertificateDer<'_>,
237 dss: &rustls::DigitallySignedStruct,
238 ) -> Result<HandshakeSignatureValid, rustls::Error> {
239 rustls::crypto::verify_tls13_signature(message, cert, dss, &self.algorithms)
240 }
241
242 fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
244 self.algorithms.supported_schemes()
245 }
246}
247
248pub fn extract_attestation_doc(
250 cert_der: &[u8],
251) -> Result<Vec<u8>, Box<dyn std::error::Error + Send + Sync>> {
252 let (_, cert) = X509Certificate::from_der(cert_der)?;
254
255 for ext in cert.extensions() {
257 if ext
258 .oid
259 .iter()
260 .into_iter()
261 .flatten()
262 .eq(ATTESTATION_OID.iter().copied())
263 {
264 let raw_value = ext.value;
265
266 if !raw_value.is_empty() && raw_value[0] == 0x04 {
268 let (_, octet_string) = der_parser::der::parse_der_octetstring(raw_value)?;
269 return Ok(octet_string.as_slice()?.to_vec());
270 }
271
272 return Ok(raw_value.to_vec());
273 }
274 }
275
276 Err("Attestation extension OID not found in certificate".into())
277}
278
279#[derive(Debug, Clone)]
281pub struct ClientResponse {
282 pub status: StatusCode,
283 pub headers: HeaderMap,
284 pub body: Vec<u8>,
285}
286
287impl ClientResponse {
289 pub fn text(&self) -> Result<String, std::string::FromUtf8Error> {
291 String::from_utf8(self.body.clone())
292 }
293}
294
295pub struct TtkClient {
297 endpoint: Endpoint,
298 send_request: h3::client::SendRequest<h3_quinn::OpenStreams, axum::body::Bytes>,
299 driver_handle: tokio::task::JoinHandle<Result<(), h3::Error>>,
300 server_addr: SocketAddr,
301 server_name: String,
302 peer_cert: Option<CertificateDer<'static>>,
303}
304
305impl TtkClient {
307 pub async fn connect(
310 server_addr: SocketAddr,
311 server_name: &str,
312 ) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
313 Self::connect_with_verifier(server_addr, server_name, EnclaveCertVerifier::new()).await
314 }
315
316 pub async fn connect_with_verifier(
318 server_addr: SocketAddr,
319 server_name: &str,
320 verifier: EnclaveCertVerifier,
321 ) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
322 let _ = rustls::crypto::ring::default_provider().install_default();
324
325 let cert_verifier = Arc::new(verifier);
326
327 let mut client_crypto = rustls::ClientConfig::builder()
328 .dangerous()
329 .with_custom_certificate_verifier(cert_verifier.clone())
330 .with_no_client_auth();
331
332 client_crypto.alpn_protocols = vec![b"h3".to_vec()];
334
335 let quic_client_config = quinn::ClientConfig::new(Arc::new(
336 quinn::crypto::rustls::QuicClientConfig::try_from(client_crypto)?,
337 ));
338
339 let bind_addr: SocketAddr = if server_addr.is_ipv6() {
341 "[::]:0".parse()?
342 } else {
343 "0.0.0.0:0".parse()?
344 };
345 let mut endpoint = Endpoint::client(bind_addr)?;
346 endpoint.set_default_client_config(quic_client_config);
347
348 info!(
349 "Initiating QUIC connection to {} ({})",
350 server_addr, server_name
351 );
352 let connecting = endpoint.connect(server_addr, server_name)?;
353 let connection = connecting.await?;
354 info!("QUIC connection established with {}", server_addr);
355
356 let peer_cert = cert_verifier.received_certificate();
357 if let Some(ref cert) = peer_cert {
358 let hash = Sha256::digest(cert.as_ref());
359 info!(
360 "Server Certificate SHA-256 fingerprint: {}",
361 hex_encode(&hash)
362 );
363 }
364
365 let h3_quic_conn = h3_quinn::Connection::new(connection);
367 let (mut driver, send_request) = h3::client::new(h3_quic_conn).await?;
368
369 let driver_handle =
371 tokio::spawn(async move { std::future::poll_fn(|cx| driver.poll_close(cx)).await });
372
373 Ok(Self {
374 endpoint,
375 send_request,
376 driver_handle,
377 server_addr,
378 server_name: server_name.to_string(),
379 peer_cert,
380 })
381 }
382
383 pub fn server_addr(&self) -> SocketAddr {
385 self.server_addr
386 }
387
388 pub fn server_name(&self) -> &str {
390 &self.server_name
391 }
392
393 pub fn peer_cert(&self) -> Option<&CertificateDer<'static>> {
395 self.peer_cert.as_ref()
396 }
397
398 pub fn peer_cert_sha256(&self) -> Option<[u8; 32]> {
400 self.peer_cert.as_ref().map(|c| {
401 let mut arr = [0u8; 32];
402 arr.copy_from_slice(&Sha256::digest(c.as_ref()));
403 arr
404 })
405 }
406
407 pub fn peer_cert_sha256_hex(&self) -> Option<String> {
409 self.peer_cert_sha256().map(|h| hex_encode(&h))
410 }
411
412 pub async fn get(
414 &mut self,
415 path: &str,
416 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
417 let uri: Uri = if path.starts_with('/') {
418 format!("https://{}{}", self.server_name, path).parse()?
419 } else {
420 format!("https://{}/{}", self.server_name, path).parse()?
421 };
422
423 let req = Request::builder()
424 .method(Method::GET)
425 .uri(uri)
426 .header("Host", &self.server_name)
427 .header("User-Agent", "TTKClient/0.6.0")
428 .header("Accept", "*/*")
429 .body(())?;
430
431 self.send(req, None).await
432 }
433
434 pub async fn post(
436 &mut self,
437 path: &str,
438 body: &[u8],
439 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
440 let uri: Uri = if path.starts_with('/') {
441 format!("https://{}{}", self.server_name, path).parse()?
442 } else {
443 format!("https://{}/{}", self.server_name, path).parse()?
444 };
445
446 let req = Request::builder()
447 .method(Method::POST)
448 .uri(uri)
449 .header("Host", &self.server_name)
450 .header("User-Agent", "TTKClient/0.6.0")
451 .header("Content-Type", "application/octet-stream")
452 .header("Content-Length", body.len().to_string())
453 .body(())?;
454
455 self.send(req, Some(body)).await
456 }
457
458 pub async fn send(
460 &mut self,
461 req: Request<()>,
462 payload: Option<&[u8]>,
463 ) -> Result<ClientResponse, Box<dyn std::error::Error + Send + Sync>> {
464 debug!("Sending HTTP/3 request: {} {}", req.method(), req.uri());
465 let mut stream = self.send_request.send_request(req).await?;
466
467 if let Some(data) = payload {
468 if !data.is_empty() {
469 stream
470 .send_data(axum::body::Bytes::copy_from_slice(data))
471 .await?;
472 }
473 }
474 stream.finish().await?;
475
476 let response = stream.recv_response().await?;
477 let status = response.status();
478 let headers = response.headers().clone();
479
480 let mut body = Vec::new();
481 while let Some(mut chunk) = stream.recv_data().await? {
482 while chunk.has_remaining() {
483 let slice = chunk.chunk();
484 body.extend_from_slice(slice);
485 let len = slice.len();
486 chunk.advance(len);
487 }
488 }
489
490 Ok(ClientResponse {
491 status,
492 headers,
493 body,
494 })
495 }
496
497 pub async fn close(self) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
499 drop(self.send_request);
500 let _ = tokio::time::timeout(Duration::from_secs(1), self.driver_handle).await;
502 self.endpoint.wait_idle().await;
503 Ok(())
504 }
505}
506
507pub fn hex_encode(bytes: &[u8]) -> String {
509 bytes.iter().map(|b| format!("{:02x}", b)).collect()
510}
511
512pub const CLIENT_USAGE: &str = "\
514Usage: client [OPTIONS] [URL]
515
516Options:
517 -s, --server-name <NAME> SNI server name (default: localhost)
518 -p, --path <PATH> Request path (default: /)
519 -a, --addr <ADDR> Server socket address (default: 127.0.0.1:4433)
520 -h, --help Print help information
521
522Examples:
523 client
524 client https://127.0.0.1:4433/hello
525 client --addr 127.0.0.1:4433 --server-name enclave.local --path /evidence
526";
527
528#[derive(Debug, Clone, PartialEq, Eq)]
530pub struct ClientTarget {
531 pub server_addr: SocketAddr,
533 pub server_name: String,
535 pub path: String,
537}
538
539pub fn parse_client_args(
545 args: &[String],
546 default_addr: Option<String>,
547 default_name: Option<String>,
548) -> Option<ClientTarget> {
549 let mut server_addr_str = default_addr.unwrap_or_else(|| "127.0.0.1:4433".to_string());
550 let mut server_name = default_name.unwrap_or_else(|| "localhost".to_string());
551 let mut path = "/".to_string();
552
553 let mut i = 0;
554 while i < args.len() {
555 let arg = &args[i];
556 if arg == "--help" || arg == "-h" {
557 return None;
558 } else if (arg == "--server-name" || arg == "-s") && i + 1 < args.len() {
559 i += 1;
560 server_name = args[i].clone();
561 } else if (arg == "--path" || arg == "-p") && i + 1 < args.len() {
562 i += 1;
563 path = args[i].clone();
564 } else if (arg == "--addr" || arg == "-a") && i + 1 < args.len() {
565 i += 1;
566 server_addr_str = args[i].clone();
567 } else if !arg.starts_with('-') {
568 if let Ok(uri) = arg.parse::<Uri>() {
570 if let Some(host) = uri.host() {
571 let port = uri.port_u16().unwrap_or(4433);
572 server_addr_str = format!("{}:{}", host, port);
573 if host != "127.0.0.1" && host != "0.0.0.0" {
574 server_name = host.to_string();
575 }
576 }
577 if !uri.path().is_empty() {
578 path = uri.path().to_string();
579 if let Some(query) = uri.query() {
580 path.push('?');
581 path.push_str(query);
582 }
583 }
584 } else {
585 server_addr_str = arg.clone();
586 }
587 }
588 i += 1;
589 }
590
591 let server_addr: SocketAddr = server_addr_str
592 .parse()
593 .unwrap_or_else(|_| "127.0.0.1:4433".parse().unwrap());
594
595 Some(ClientTarget {
596 server_addr,
597 server_name,
598 path,
599 })
600}