use std::{str::FromStr, sync::LazyLock};
use asn1_rs::FromDer;
use lexe_byte_array::ByteArray;
use lexe_crypto::ed25519;
use lexe_sha256::sha256;
use rcgen::{DistinguishedName, DnType, string::Ia5String};
use x509_parser::{
certificate::X509Certificate, extensions::GeneralName, time::ASN1Time,
};
pub mod attest_client;
pub mod ed25519_ext;
pub mod lexe_ca;
pub mod p256;
pub mod shared_seed;
pub mod types;
pub use lexe_tls_core::*;
#[must_use]
pub fn cert_contains_dns(cert_der: &[u8], expected_dns: &[&str]) -> bool {
fn contains_dns(cert_der: &[u8], expected_dns: &[&str]) -> Option<()> {
if expected_dns.is_empty() {
return Some(());
}
let (_unparsed, cert) = X509Certificate::from_der(cert_der).ok()?;
let sans = &cert.subject_alternative_name().ok()??.value.general_names;
expected_dns
.iter()
.all(|dns_name| sans.contains(&GeneralName::DNSName(dns_name)))
.then_some(())
}
contains_dns(cert_der, expected_dns).is_some()
}
#[must_use]
pub fn cert_is_valid_for_at_least(cert_der: &[u8], buffer_days: u16) -> bool {
fn is_valid_for_at_least(cert_der: &[u8], buffer_days: i64) -> Option<()> {
use std::ops::Add;
let (_unparsed, cert) = X509Certificate::from_der(cert_der).ok()?;
let now = ASN1Time::now();
let validity = cert.validity();
if now < validity.not_before {
return None;
}
if now > validity.not_after {
return None;
}
let buffer_days_dur = time::Duration::days(buffer_days);
let now_plus_buffer = now.add(buffer_days_dur)?;
if now_plus_buffer < validity.not_before {
return None;
}
if now_plus_buffer > validity.not_after {
return None;
}
Some(())
}
is_valid_for_at_least(cert_der, i64::from(buffer_days)).is_some()
}
pub static DEFAULT_SUBJECT_ALT_NAMES: LazyLock<Vec<rcgen::SanType>> =
LazyLock::new(|| {
vec![rcgen::SanType::DnsName(
Ia5String::from_str("lexe.app").unwrap(),
)]
});
pub fn build_rcgen_cert_params(
common_name: &str,
not_before: time::OffsetDateTime,
not_after: time::OffsetDateTime,
subject_alt_names: Vec<rcgen::SanType>,
public_key: &ed25519::PublicKey,
overrides: impl FnOnce(&mut rcgen::CertificateParams),
) -> rcgen::CertificateParams {
let mut params = rcgen::CertificateParams::default();
params.not_before = not_before;
params.not_after = not_after;
params.subject_alt_names = subject_alt_names;
params.distinguished_name = lexe_distinguished_name(common_name);
overrides(&mut params);
let pubkey_hash = {
let hash = sha256::digest(public_key.as_slice());
hash.as_slice()[0..20].to_vec()
};
if matches!(params.is_ca, rcgen::IsCa::Ca(_) | rcgen::IsCa::ExplicitNoCa) {
params.key_identifier_method =
rcgen::KeyIdMethod::PreSpecified(pubkey_hash.clone());
}
let mut serial = pubkey_hash;
serial[0] &= 0x7f; params.serial_number = Some(rcgen::SerialNumber::from(serial));
params
}
pub fn lexe_distinguished_name(common_name: &str) -> DistinguishedName {
let mut name = DistinguishedName::new();
name.push(DnType::CountryName, "US");
name.push(DnType::StateOrProvinceName, "CA");
name.push(DnType::OrganizationName, "lexe-app");
name.push(DnType::CommonName, common_name);
name
}
#[cfg(any(test, feature = "test-utils"))]
pub mod test_utils {
use std::sync::Arc;
use anyhow::Context;
use rustls::{ClientConfig, ServerConfig, pki_types::ServerName};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
pub async fn do_tls_handshake(
client_config: Arc<ClientConfig>,
server_config: Arc<ServerConfig>,
expected_dns: &str,
) -> [Result<(), String>; 2] {
let (client_stream, server_stream) = tokio::io::duplex(4096);
let client = async move {
let connector = tokio_rustls::TlsConnector::from(client_config);
let sni = ServerName::try_from(expected_dns.to_owned()).unwrap();
let mut stream = connector
.connect(sni, client_stream)
.await
.context("Client didn't connect")?;
stream
.write_all(b"hello")
.await
.context("Could not write hello")?;
stream.flush().await.context("Toilet clogged")?;
stream.shutdown().await.context("Could not shutdown")?;
let mut resp = Vec::new();
stream.read_to_end(&mut resp).await.context("Read failed")?;
assert_eq!(&resp, b"goodbye");
Ok::<_, anyhow::Error>(())
};
let server = async move {
let acceptor = tokio_rustls::TlsAcceptor::from(server_config);
let mut stream = acceptor
.accept(server_stream)
.await
.context("Server didn't accept")?;
let mut req = Vec::new();
stream.read_to_end(&mut req).await.context("Read failed")?;
assert_eq!(&req, b"hello");
stream
.write_all(b"goodbye")
.await
.context("Could not write goodbye")?;
stream.shutdown().await.context("Could not shutdown")?;
Ok::<_, anyhow::Error>(())
};
let (client_result, server_result) = tokio::join!(client, server);
let (client_result, server_result) = (
client_result.map_err(|e| format!("{e:#}")),
server_result.map_err(|e| format!("{e:#}")),
);
println!("Client result: {client_result:?}");
println!("Server result: {server_result:?}");
println!("---");
[client_result, server_result]
}
}