use crate::asn1::OctetString;
use crate::constants::TIGHTBEAM_AAD_DOMAIN_TAG;
use crate::crypto::aead::{KeyInit, RuntimeAead};
use crate::crypto::ecies::EciesEphemeral;
use crate::crypto::ecies::{encrypt, EciesMessageOps, EciesPublicKeyOps};
use crate::crypto::key::SigningKeyProvider;
use crate::crypto::profiles::{CryptoProvider, SecurityProfileDesc};
use crate::crypto::sign::elliptic_curve::sec1::{FromEncodedPoint, ModulusSize, ToEncodedPoint};
use crate::crypto::sign::elliptic_curve::{AffinePoint, Curve, CurveArithmetic, PublicKey};
use crate::crypto::sign::SignatureEncoding;
use crate::crypto::sign::Verifier;
use crate::crypto::x509::policy::CertificateValidation;
use crate::crypto::x509::utils::validate_certificate_expiry;
use crate::der::{Decode, Encode};
use crate::random::generate_nonce;
use crate::transport::handshake::error::HandshakeError;
use crate::transport::handshake::negotiation::SecurityOffer;
use crate::transport::handshake::state::HandshakeInvariant;
use crate::transport::handshake::state::{ClientHandshakeState, ClientStateMachine};
use crate::transport::handshake::utils::{compute_transcript_digest, octet_string_to_32_byte_array, validate_state};
use crate::transport::handshake::{Arc, ClientHandshakeProtocol, ClientHello, ClientKeyExchange, ServerHandshake};
use crate::transport::handshake::{HandshakeAlertHandler, HandshakeFinalization}; use crate::x509::Certificate;
pub struct EciesHandshakeClient<P, M>
where
P: CryptoProvider,
{
state: ClientStateMachine,
client_random: Option<[u8; 32]>,
base_session_key: Option<[u8; 32]>,
server_random: Option<[u8; 32]>,
transcript_hash: Option<[u8; 32]>,
aad_domain_tag: Option<&'static [u8]>,
security_offer: Option<SecurityOffer>,
selected_profile: Option<SecurityProfileDesc>,
certificate_validator: Option<Arc<dyn CertificateValidation>>,
client_certificate: Option<Arc<Certificate>>,
client_key_provider: Option<Arc<dyn crate::crypto::key::SigningKeyProvider>>,
_phantom_provider: ::core::marker::PhantomData<P>,
_phantom_message: ::core::marker::PhantomData<M>,
invariants: HandshakeInvariant,
}
pub trait ExtractVerifyingKey: Sized {
fn extract_from_certificate(cert: &Certificate) -> Result<Self, HandshakeError>;
}
impl<P, M> EciesHandshakeClient<P, M>
where
P: CryptoProvider,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EciesPublicKeyOps,
<PublicKey<P::Curve> as EciesPublicKeyOps>::SecretKey: EciesEphemeral<PublicKey = PublicKey<P::Curve>>,
P::Signature: SignatureEncoding,
for<'a> P::Signature: TryFrom<&'a [u8]>,
for<'a> <P::Signature as TryFrom<&'a [u8]>>::Error: Into<HandshakeError>,
P::VerifyingKey: Verifier<P::Signature> + ExtractVerifyingKey,
P::AeadCipher: KeyInit,
M: EciesMessageOps,
{
pub fn new(aad_domain_tag: Option<&'static [u8]>) -> Self {
Self {
state: ClientStateMachine::default(),
client_random: None,
base_session_key: None,
server_random: None,
transcript_hash: None,
aad_domain_tag: aad_domain_tag.or(Some(TIGHTBEAM_AAD_DOMAIN_TAG)),
security_offer: None, selected_profile: None,
certificate_validator: None,
client_certificate: None,
client_key_provider: None,
invariants: HandshakeInvariant::default(),
_phantom_provider: ::core::marker::PhantomData,
_phantom_message: ::core::marker::PhantomData,
}
}
pub fn new_with_identity(
aad_domain_tag: Option<&'static [u8]>,
client_certificate: Option<Arc<Certificate>>,
client_key_provider: Option<Arc<dyn crate::crypto::key::SigningKeyProvider>>,
) -> Self {
Self {
state: ClientStateMachine::default(),
client_random: None,
base_session_key: None,
server_random: None,
transcript_hash: None,
aad_domain_tag: aad_domain_tag.or(Some(TIGHTBEAM_AAD_DOMAIN_TAG)),
security_offer: None, selected_profile: None,
certificate_validator: None,
client_certificate,
client_key_provider,
invariants: HandshakeInvariant::default(),
_phantom_provider: ::core::marker::PhantomData,
_phantom_message: ::core::marker::PhantomData,
}
}
pub fn with_certificate_validator(mut self, validator: Arc<dyn CertificateValidation>) -> Self {
self.certificate_validator = Some(validator);
self
}
pub fn with_client_identity(
mut self,
certificate: Arc<Certificate>,
key_provider: Arc<dyn SigningKeyProvider>,
) -> Self {
self.client_certificate = Some(certificate);
self.client_key_provider = Some(key_provider);
self
}
pub fn with_security_offer(mut self, offer: SecurityOffer) -> Self {
self.security_offer = Some(offer);
self
}
fn validate_expected_state(&self, expected: ClientHandshakeState) -> Result<(), HandshakeError> {
validate_state(self.state.state(), expected)
}
fn validate_and_extract_server_handshake(
&self,
server_handshake_der: &[u8],
) -> Result<ServerHandshake, HandshakeError> {
let server_handshake = ServerHandshake::from_der(server_handshake_der)?;
if let Some(validator) = &self.certificate_validator {
validator.evaluate(&server_handshake.certificate)?;
} else {
validate_certificate_expiry(&server_handshake.certificate)?;
}
Ok(server_handshake)
}
fn extract_server_random(&mut self, server_handshake: &ServerHandshake) -> Result<(), HandshakeError> {
let server_random = octet_string_to_32_byte_array(&server_handshake.server_random)?;
self.server_random = Some(server_random);
Ok(())
}
fn compute_and_store_transcript_hash(&mut self, server_handshake: &ServerHandshake) -> Result<(), HandshakeError> {
let client_random = self.client_random.ok_or(HandshakeError::InvalidState)?;
let server_random = self.server_random.ok_or(HandshakeError::InvalidState)?;
let spki_bytes = server_handshake
.certificate
.tbs_certificate
.subject_public_key_info
.subject_public_key
.raw_bytes();
let transcript_digest = self.compute_transcript_hash(&client_random, &server_random, spki_bytes);
self.transcript_hash = Some(transcript_digest);
self.invariants.lock_transcript()?;
Ok(())
}
fn generate_base_session_key(&mut self) -> Result<(), HandshakeError> {
let base_key = generate_nonce::<32>(None)?;
self.base_session_key = Some(base_key);
Ok(())
}
pub fn build_client_hello(&mut self) -> Result<Vec<u8>, HandshakeError> {
self.validate_expected_state(ClientHandshakeState::Init)?;
let client_random = generate_nonce::<32>(None)?;
self.client_random = Some(client_random);
let client_hello = ClientHello {
client_random: OctetString::new(client_random)?,
security_offer: self.security_offer.clone(),
};
self.state.transition(ClientHandshakeState::HelloSent)?;
Ok(client_hello.to_der()?)
}
pub async fn process_server_handshake(&mut self, server_handshake_der: &[u8]) -> Result<Vec<u8>, HandshakeError> {
self.validate_expected_state(ClientHandshakeState::HelloSent)?;
let _client_random_check = self.client_random.ok_or(HandshakeError::InvalidState)?;
self.state.transition(ClientHandshakeState::ServerHelloReceived)?;
let server_handshake = self.validate_and_extract_server_handshake(server_handshake_der)?;
self.validate_profile_selection(&server_handshake)?;
self.extract_server_random(&server_handshake)?;
self.verify_server_handshake_signature(&server_handshake)?;
let encrypted_bytes = self.generate_and_encrypt_session_key(&server_handshake)?;
let (client_certificate, client_signature) = self.prepare_client_auth(&server_handshake).await?;
let client_kex = ClientKeyExchange {
encrypted_data: OctetString::new(encrypted_bytes)?,
#[cfg(feature = "x509")]
client_certificate,
#[cfg(feature = "x509")]
client_signature,
};
self.state.transition(ClientHandshakeState::KeyExchangeSent)?;
Ok(client_kex.to_der()?)
}
fn validate_profile_selection(&mut self, server_handshake: &ServerHandshake) -> Result<(), HandshakeError> {
let accept = server_handshake.security_accept.as_ref().ok_or(HandshakeError::InvalidState)?;
match &self.security_offer {
Some(offer) => {
if !offer.profiles.contains(&accept.profile) {
return Err(HandshakeError::InvalidProfileSelection);
}
self.selected_profile = Some(accept.profile);
}
None => {
self.selected_profile = Some(accept.profile);
}
}
Ok(())
}
fn verify_server_handshake_signature(&mut self, server_handshake: &ServerHandshake) -> Result<(), HandshakeError> {
let verifying_key = self.extract_verifying_key(&server_handshake.certificate)?;
self.compute_and_store_transcript_hash(server_handshake)?;
let transcript_digest = self.transcript_hash.ok_or(HandshakeError::InvalidState)?;
self.verify_server_signature(&verifying_key, &transcript_digest, server_handshake.signature.as_bytes())
}
fn generate_and_encrypt_session_key(
&mut self,
server_handshake: &ServerHandshake,
) -> Result<Vec<u8>, HandshakeError> {
self.generate_base_session_key()?;
let base_key = self.base_session_key.ok_or(HandshakeError::InvalidState)?;
let client_random = self.client_random.ok_or(HandshakeError::InvalidState)?;
self.perform_ecies_encryption(&base_key, &client_random, &server_handshake.certificate, self.aad_domain_tag)
}
async fn prepare_client_auth(
&self,
server_handshake: &ServerHandshake,
) -> Result<(Option<Certificate>, Option<OctetString>), HandshakeError> {
let transcript_digest = self.transcript_hash.ok_or(HandshakeError::InvalidState)?;
if server_handshake.client_cert_required {
let cert = self.client_certificate.as_ref().ok_or(HandshakeError::MutualAuthRequired)?;
let key_provider = self.client_key_provider.as_ref().ok_or(HandshakeError::MutualAuthRequired)?;
let signature_bytes = key_provider.sign(&transcript_digest).await?;
Ok((Some(Certificate::clone(cert)), Some(OctetString::new(signature_bytes)?)))
} else if let Some(cert) = &self.client_certificate {
let key_provider = self.client_key_provider.as_ref().ok_or(HandshakeError::InvalidState)?;
let signature_bytes = key_provider.sign(&transcript_digest).await?;
Ok((Some(Certificate::clone(cert)), Some(OctetString::new(signature_bytes)?)))
} else {
Ok((None, None))
}
}
pub fn complete(&mut self) -> Result<P::AeadCipher, HandshakeError> {
self.validate_expected_state(ClientHandshakeState::KeyExchangeSent)?;
let base_key = self.base_session_key.as_ref().ok_or(HandshakeError::InvalidState)?;
let client_random = self.client_random.as_ref().ok_or(HandshakeError::InvalidState)?;
let server_random = self.server_random.as_ref().ok_or(HandshakeError::InvalidState)?;
let mut salt = [0u8; 64];
salt[..32].copy_from_slice(client_random);
salt[32..].copy_from_slice(server_random);
let session_key = self.derive_session_aead(base_key, &salt)?;
self.invariants.derive_aead_once()?;
self.state.transition(ClientHandshakeState::Completed)?;
if let Some(mut bk) = self.base_session_key.take() {
bk.fill(0);
}
if let Some(mut cr) = self.client_random.take() {
cr.fill(0);
}
if let Some(mut sr) = self.server_random.take() {
sr.fill(0);
}
Ok(session_key)
}
pub fn state(&self) -> ClientHandshakeState {
self.state.state()
}
pub fn is_complete(&self) -> bool {
self.state.state().is_completed()
}
pub fn transcript_hash(&self) -> Option<[u8; 32]> {
self.transcript_hash
}
fn extract_verifying_key(&self, cert: &Certificate) -> Result<P::VerifyingKey, HandshakeError> {
P::VerifyingKey::extract_from_certificate(cert)
}
fn compute_transcript_hash(
&self,
client_random: &[u8; 32],
server_random: &[u8; 32],
spki_bytes: &[u8],
) -> [u8; 32] {
let mut data = Vec::with_capacity(32 + 32 + spki_bytes.len());
data.extend_from_slice(client_random);
data.extend_from_slice(server_random);
data.extend_from_slice(spki_bytes);
compute_transcript_digest::<P::Digest>(&data)
}
fn verify_server_signature(
&self,
verifying_key: &P::VerifyingKey,
digest: &[u8; 32],
signature_bytes: &[u8],
) -> Result<(), HandshakeError> {
let signature = P::Signature::try_from(signature_bytes).map_err(|e| e.into())?;
verifying_key.verify(digest, &signature)?;
Ok(())
}
fn perform_ecies_encryption(
&self,
base_key: &[u8; 32],
client_random: &[u8; 32],
server_certificate: &Certificate,
associated_data: Option<&[u8]>,
) -> Result<Vec<u8>, HandshakeError> {
let mut plaintext = [0u8; 64];
plaintext[..32].copy_from_slice(base_key);
plaintext[32..].copy_from_slice(client_random);
let recipient_pubkey = PublicKey::<P::Curve>::from_sec1_bytes(
server_certificate
.tbs_certificate
.subject_public_key_info
.subject_public_key
.raw_bytes(),
)?;
let encrypted_message =
encrypt::<_, _, _, M>(&recipient_pubkey, &plaintext, associated_data, Some(&mut rand_core::OsRng))?;
Ok(encrypted_message.to_bytes())
}
}
impl<P, M> HandshakeFinalization<P> for EciesHandshakeClient<P, M>
where
P: CryptoProvider,
{
fn selected_profile(&self) -> Option<SecurityProfileDesc> {
self.selected_profile
}
}
impl<P, M> HandshakeAlertHandler for EciesHandshakeClient<P, M> where P: CryptoProvider {}
impl<P, M> ClientHandshakeProtocol for EciesHandshakeClient<P, M>
where
P: CryptoProvider + Send + Sync,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EciesPublicKeyOps,
<PublicKey<P::Curve> as EciesPublicKeyOps>::SecretKey: EciesEphemeral<PublicKey = PublicKey<P::Curve>>,
P::Signature: SignatureEncoding + Send + Sync,
for<'a> P::Signature: TryFrom<&'a [u8]>,
for<'a> <P::Signature as TryFrom<&'a [u8]>>::Error: Into<HandshakeError>,
P::VerifyingKey: Verifier<P::Signature> + ExtractVerifyingKey + Send + Sync,
P::AeadCipher: KeyInit + Send + Sync + 'static,
M: EciesMessageOps + Send + Sync,
{
type Error = HandshakeError;
fn start<'a>(
&'a mut self,
) -> core::pin::Pin<Box<dyn core::future::Future<Output = Result<Vec<u8>, Self::Error>> + Send + 'a>> {
Box::pin(async move { self.build_client_hello() })
}
fn handle_response<'a, 'b>(
&'a mut self,
msg: &'b [u8],
) -> core::pin::Pin<Box<dyn core::future::Future<Output = Result<Option<Vec<u8>>, Self::Error>> + Send + 'a>>
where
'b: 'a,
{
Box::pin(async move {
let client_kex = self.process_server_handshake(msg).await?;
Ok(Some(client_kex))
})
}
#[cfg(feature = "aead")]
fn complete<'a>(
&'a mut self,
) -> core::pin::Pin<Box<dyn core::future::Future<Output = Result<RuntimeAead, Self::Error>> + Send + 'a>> {
Box::pin(async move {
if self.state.state() != ClientHandshakeState::KeyExchangeSent {
return Err(HandshakeError::InvalidState);
}
let profile = self.selected_profile.ok_or(HandshakeError::InvalidState)?;
let aead_oid = profile.aead.ok_or(HandshakeError::InvalidState)?;
let base_key = self.base_session_key.as_ref().ok_or(HandshakeError::InvalidState)?;
let client_random = self.client_random.as_ref().ok_or(HandshakeError::InvalidState)?;
let server_random = self.server_random.as_ref().ok_or(HandshakeError::InvalidState)?;
let mut salt = [0u8; 64];
salt[..32].copy_from_slice(client_random);
salt[32..].copy_from_slice(server_random);
let cipher = self.derive_session_aead(base_key, &salt)?;
self.state.transition(ClientHandshakeState::Completed)?;
if let Some(mut bk) = self.base_session_key.take() {
bk.fill(0);
}
if let Some(mut cr) = self.client_random.take() {
cr.fill(0);
}
if let Some(mut sr) = self.server_random.take() {
sr.fill(0);
}
Ok(RuntimeAead::new(cipher, aead_oid))
})
}
fn is_complete(&self) -> bool {
self.state.state().is_completed()
}
fn selected_profile(&self) -> Option<SecurityProfileDesc> {
self.selected_profile
}
}
#[cfg(feature = "secp256k1")]
impl ExtractVerifyingKey for crate::crypto::sign::ecdsa::Secp256k1VerifyingKey {
fn extract_from_certificate(cert: &Certificate) -> Result<Self, HandshakeError> {
let public_key_bytes = crate::crypto::x509::utils::extract_verifying_key_bytes(cert);
let public_key = k256::PublicKey::from_sec1_bytes(public_key_bytes)?;
Ok(Self::from(public_key))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::crypto::ecies::Secp256k1EciesMessage;
use crate::crypto::profiles::{DefaultCryptoProvider, SecurityProfileDesc};
use crate::crypto::sign::ecdsa::Secp256k1Signature;
use crate::crypto::sign::Signer;
use crate::der::Encode;
use crate::transport::handshake::negotiation::{SecurityAccept, SecurityOffer};
use crate::transport::handshake::tests::*;
use crate::transport::handshake::ServerHandshake;
use crate::oids::{
AES_256_GCM, AES_256_WRAP, CURVE_SECP256K1, HASH_SHA3_256, HASH_SHA3_384, HASH_SHA3_512,
SIGNER_ECDSA_WITH_SHA3_512,
};
#[tokio::test]
async fn test_client_state_flow() -> Result<(), Box<dyn core::error::Error>> {
let mut client = TestEciesClientBuilder::new().build();
assert_eq!(client.state(), ClientHandshakeState::Init);
let _client_hello_der = client.build_client_hello()?;
assert_eq!(client.state(), ClientHandshakeState::HelloSent); assert!(client.client_random.is_some());
let test_cert = create_test_certificate();
let client_random = client.client_random.unwrap();
let server_random = crate::random::generate_nonce::<32>(None)?;
let transcript_hash = compute_test_transcript_hash(
&client_random,
&server_random,
test_cert
.certificate
.tbs_certificate
.subject_public_key_info
.subject_public_key
.raw_bytes(),
);
let signature_bytes: Secp256k1Signature = test_cert.signing_key.try_sign(&transcript_hash)?;
let server_handshake_der =
create_test_server_handshake(&test_cert.certificate, &server_random, &signature_bytes.to_bytes())?;
let client_kex_der = client.process_server_handshake(&server_handshake_der).await?;
assert_eq!(client.state(), ClientHandshakeState::KeyExchangeSent);
assert!(client.base_session_key.is_some());
assert!(client.transcript_hash.is_some());
let _client_kex = ClientKeyExchange::from_der(&client_kex_der)?;
let _session_key = client.complete()?;
assert!(client.is_complete());
assert_eq!(client.state(), ClientHandshakeState::Completed);
Ok(())
}
#[tokio::test]
async fn test_invalid_state_transitions() -> Result<(), Box<dyn core::error::Error>> {
let mut client = TestEciesClientBuilder::new().build();
let result = client.process_server_handshake(&[]).await;
assert!(result.is_err());
let _client_hello = client.build_client_hello()?;
assert_eq!(client.state(), ClientHandshakeState::HelloSent);
let result = client.complete();
assert!(result.is_err());
Ok(())
}
#[tokio::test]
async fn test_client_profile_validation() -> Result<(), Box<dyn core::error::Error>> {
let mk_profile = |id: u8| SecurityProfileDesc {
digest: match id {
1 => HASH_SHA3_256,
2 => HASH_SHA3_384,
_ => HASH_SHA3_512,
},
aead: Some(AES_256_GCM),
aead_key_size: Some(32),
signature: Some(SIGNER_ECDSA_WITH_SHA3_512),
kdf: Some(HASH_SHA3_256), curve: Some(CURVE_SECP256K1),
key_wrap: Some(AES_256_WRAP),
kem: None,
};
let (p_a, p_b, p_c) = (mk_profile(1), mk_profile(2), mk_profile(3));
let test_cert = create_test_certificate();
#[allow(clippy::type_complexity)]
let setup_client = |offer: Option<SecurityOffer>| -> Result<
(EciesHandshakeClient<DefaultCryptoProvider, Secp256k1EciesMessage>, [u8; 32]),
Box<dyn std::error::Error>,
> {
let mut client = TestEciesClientBuilder::new().build();
if let Some(offer) = offer {
client = client.with_security_offer(offer);
}
let _hello = client.build_client_hello()?;
let client_random = client.client_random.ok_or("No client random")?;
Ok((client, client_random))
};
let create_server_response = |client_random: &[u8; 32],
server_random: [u8; 32],
accepted_profile: &SecurityProfileDesc|
-> Result<Vec<u8>, Box<dyn core::error::Error>> {
let transcript_hash = compute_test_transcript_hash(
client_random,
&server_random,
test_cert
.certificate
.tbs_certificate
.subject_public_key_info
.subject_public_key
.raw_bytes(),
);
let signature: Secp256k1Signature = test_cert.signing_key.try_sign(&transcript_hash)?;
let signature_bytes = signature.to_bytes().to_vec();
let response = ServerHandshake {
certificate: test_cert.certificate.clone(),
server_random: OctetString::new(server_random)?,
signature: OctetString::new(signature_bytes)?,
security_accept: Some(SecurityAccept::new(*accepted_profile)),
client_cert_required: false,
};
Ok(response.to_der()?)
};
{
let (mut client, client_random) = setup_client(Some(SecurityOffer::new(vec![p_a, p_b])))?;
let server_response = create_server_response(&client_random, [2u8; 32], &p_b)?;
let _kex = client.process_server_handshake(&server_response).await?;
assert_eq!(client.selected_profile, Some(p_b));
}
{
let (mut client, client_random) = setup_client(Some(SecurityOffer::new(vec![p_a, p_b])))?;
let server_response = create_server_response(&client_random, [3u8; 32], &p_c)?;
let result = client.process_server_handshake(&server_response).await;
assert!(matches!(result, Err(HandshakeError::InvalidProfileSelection)));
}
{
let (mut client, client_random) = setup_client(None)?;
let server_response = create_server_response(&client_random, [4u8; 32], &p_a)?;
let _kex = client.process_server_handshake(&server_response).await?;
assert_eq!(client.selected_profile, Some(p_a));
}
Ok(())
}
}