use core::future::Future;
use core::pin::Pin;
use crate::cms::content_info::CmsVersion;
use crate::cms::enveloped_data::{KeyAgreeRecipientIdentifier, UserKeyingMaterial};
use crate::cms::signed_data::{EncapsulatedContentInfo, SignedData, SignerInfo};
use crate::cms::{cert::IssuerAndSerialNumber, signed_data::SignerIdentifier};
use crate::crypto::aead::KeyInit;
use crate::crypto::hash::Digest;
use crate::crypto::key::SigningKeyProvider;
use crate::crypto::profiles::{CryptoProvider, SecurityProfile, SecurityProfileDesc};
use crate::crypto::secret::Secret;
use crate::crypto::sign::elliptic_curve::sec1::{FromEncodedPoint, ModulusSize, ToEncodedPoint};
use crate::crypto::sign::elliptic_curve::{AffinePoint, PublicKey, SecretKey};
use crate::crypto::sign::{EcdsaSignatureVerifier, SignatureAlgorithmIdentifier};
use crate::crypto::x509::store::CertificateTrust;
use crate::crypto::x509::utils::validate_certificate_expiry;
use crate::crypto::x509::Certificate;
use crate::der::asn1::OctetString;
use crate::der::oid::AssociatedOid;
use crate::der::{Decode, Encode};
use crate::random::{generate_nonce, OsRng};
use crate::spki::{AlgorithmIdentifierOwned, EncodePublicKey, SubjectPublicKeyInfoOwned};
use crate::transport::handshake::builders::{TightBeamEnvelopedDataBuilder, TightBeamKariBuilder};
use crate::transport::handshake::error::HandshakeError;
use crate::transport::handshake::negotiation::SecurityOffer;
use crate::transport::handshake::processors::TightBeamSignedDataProcessor;
use crate::transport::handshake::state::HandshakeInvariant;
use crate::transport::handshake::state::{ClientHandshakeState, ClientStateMachine};
use crate::transport::handshake::utils::{compute_transcript_digest, extract_verifying_key_from_cert, validate_state};
use crate::transport::handshake::{Arc, ClientHandshakeProtocol, HandshakeAlertHandler, HandshakeFinalization};
pub struct CmsHandshakeClient<P>
where
P: CryptoProvider,
{
state: ClientStateMachine,
client_key_provider: Arc<dyn SigningKeyProvider>,
client_certificate: Option<Arc<Certificate>>,
server_cert: Arc<Certificate>,
transcript_hash: Option<[u8; 32]>,
transcript_buffer: Vec<u8>,
session_key: Option<Secret<Vec<u8>>>,
security_offer: Option<SecurityOffer>,
selected_profile: Option<SecurityProfileDesc>,
provider: P,
trust_store: Option<Arc<dyn CertificateTrust>>,
invariants: HandshakeInvariant,
}
impl<P> CmsHandshakeClient<P>
where
P: CryptoProvider,
P::Curve: elliptic_curve::Curve + elliptic_curve::CurveArithmetic,
<P::Curve as elliptic_curve::Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EncodePublicKey,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + signature::Verifier<P::Signature> + 'static,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit,
{
pub fn new(provider: P, client_key_provider: Arc<dyn SigningKeyProvider>, server_cert: Arc<Certificate>) -> Self {
Self {
state: ClientStateMachine::default(),
client_key_provider,
client_certificate: None,
server_cert,
transcript_hash: None,
transcript_buffer: Vec::new(),
session_key: None,
security_offer: None,
selected_profile: None,
provider,
trust_store: None,
invariants: HandshakeInvariant::default(),
}
}
#[must_use]
pub fn with_transcript_hash(mut self, hash: [u8; 32]) -> Self {
self.transcript_hash = Some(hash);
self
}
#[must_use]
pub fn with_trust_store(mut self, store: Arc<dyn CertificateTrust>) -> Self {
self.trust_store = Some(store);
self
}
pub fn with_client_certificate(mut self, certificate: impl Into<Certificate>) -> Self {
self.client_certificate = Some(Arc::new(certificate.into()));
self
}
#[must_use]
pub fn with_security_offer(mut self, offer: SecurityOffer) -> Self {
self.security_offer = Some(offer);
self
}
pub fn selected_profile(&self) -> Option<SecurityProfileDesc> {
self.selected_profile
}
fn validate_expected_state(&self, expected: ClientHandshakeState) -> Result<(), HandshakeError> {
validate_state(self.state.state(), expected)
}
fn validate_state_and_certificate(&self) -> Result<(), HandshakeError> {
self.validate_expected_state(ClientHandshakeState::Init)?;
if let Some(store) = &self.trust_store {
store.evaluate(&self.server_cert)?;
} else {
validate_certificate_expiry(&self.server_cert)?;
}
Ok(())
}
fn extract_server_public_key(&self) -> Result<PublicKey<P::Curve>, HandshakeError> {
Ok(PublicKey::<P::Curve>::from_sec1_bytes(
self.server_cert
.tbs_certificate
.subject_public_key_info
.subject_public_key
.raw_bytes(),
)?)
}
fn create_ephemeral_keypair(&self) -> Result<(SecretKey<P::Curve>, SubjectPublicKeyInfoOwned), HandshakeError> {
let sender_ephemeral = SecretKey::<P::Curve>::random(&mut OsRng);
let sender_public = sender_ephemeral.public_key();
let sender_pub_spki = sender_public.to_public_key_der()?;
let sender_pub_spki = SubjectPublicKeyInfoOwned::from_der(sender_pub_spki.as_bytes())?;
Ok((sender_ephemeral, sender_pub_spki))
}
fn build_recipient_identifier(&self) -> KeyAgreeRecipientIdentifier {
KeyAgreeRecipientIdentifier::IssuerAndSerialNumber(IssuerAndSerialNumber {
issuer: self.server_cert.tbs_certificate.issuer.clone(),
serial_number: self.server_cert.tbs_certificate.serial_number.clone(),
})
}
fn extract_server_verifying_key(&self, server_cert: &Certificate) -> Result<P::VerifyingKey, HandshakeError> {
let server_public_key = extract_verifying_key_from_cert::<P::Curve>(server_cert)?;
Ok(P::VerifyingKey::from(server_public_key))
}
fn compute_signer_identifier(&self, verifying_key: &P::VerifyingKey) -> Result<SignerIdentifier, HandshakeError> {
Ok(crate::crypto::x509::utils::compute_signer_identifier::<P::Digest, _>(
verifying_key,
)?)
}
fn compute_transcript_hash(&self) -> [u8; 32] {
compute_transcript_digest::<P::Digest>(&self.transcript_buffer)
}
fn verify_signature(
&self,
signed_data_der: &[u8],
server_verifying_key: P::VerifyingKey,
expected_sid: SignerIdentifier,
) -> Result<Vec<u8>, HandshakeError> {
let verifier = EcdsaSignatureVerifier::<P::VerifyingKey, P::Signature, P::Digest>::from_verifying_key_with_sid(
server_verifying_key,
expected_sid,
);
let processor = TightBeamSignedDataProcessor::new(verifier);
let digest_oid = P::Digest::OID;
let verified_content = processor.process_der(signed_data_der, &digest_oid)?;
let expected_hash = self.transcript_hash.ok_or(HandshakeError::InvalidState)?;
if verified_content.len() != 32 || verified_content.as_slice() != expected_hash {
Err(HandshakeError::SignatureVerificationFailed)
} else {
Ok(verified_content)
}
}
pub fn build_key_exchange(&mut self, session_key: Vec<u8>) -> Result<Vec<u8>, HandshakeError> {
self.validate_key_exchange_prerequisites()?;
let (server_public_key, sender_ephemeral, sender_pub_spki) = self.extract_key_exchange_crypto_material()?;
let ukm = self.create_user_keying_material()?;
let rid = self.build_recipient_identifier();
let kari_builder = self.build_kari_structure(sender_ephemeral, sender_pub_spki, server_public_key, rid, ukm)?;
let enveloped_data_der = self.build_enveloped_data(kari_builder, &session_key)?;
self.finalize_key_exchange(&enveloped_data_der, session_key)?;
Ok(enveloped_data_der)
}
pub fn process_server_finished(&mut self, signed_data_der: &[u8]) -> Result<Vec<u8>, HandshakeError> {
self.validate_expected_state(ClientHandshakeState::KeyExchangeSent)?;
if self.transcript_hash.is_none() {
self.transcript_hash = Some(self.compute_transcript_hash());
}
let server_verifying_key = self.extract_server_verifying_key(&self.server_cert)?;
let expected_signer_identifier = self.compute_signer_identifier(&server_verifying_key)?;
let verified_content =
self.verify_signature(signed_data_der, server_verifying_key, expected_signer_identifier)?;
self.transcript_buffer.extend_from_slice(signed_data_der);
self.state.transition(ClientHandshakeState::ServerFinishedReceived)?;
self.invariants.lock_transcript()?;
Ok(verified_content)
}
pub async fn build_client_finished(&mut self) -> Result<Vec<u8>, HandshakeError> {
self.validate_client_finished_prerequisites()?;
let (transcript_hash, digest) = self.prepare_finished_digest()?;
let signature_bytes = self.sign_finished_digest(&digest).await?;
let (signer_id, digest_alg, signature_alg) = self.build_finished_crypto_components().await?;
let signed_data_der =
self.build_signed_data(transcript_hash, &signature_bytes, signer_id, digest_alg, signature_alg)?;
self.finalize_client_finished()?;
Ok(signed_data_der)
}
pub fn complete(&mut self) -> Result<(), HandshakeError> {
self.validate_expected_state(ClientHandshakeState::ClientFinishedSent)?;
self.state.transition(ClientHandshakeState::Completed)?;
Ok(())
}
pub fn state(&self) -> ClientHandshakeState {
self.state.state()
}
pub fn is_complete(&self) -> bool {
self.state.state().is_completed()
}
pub fn session_key(&self) -> Option<&Secret<Vec<u8>>> {
self.session_key.as_ref()
}
fn validate_key_exchange_prerequisites(&self) -> Result<(), HandshakeError> {
if self.state.state() == ClientHandshakeState::Init {
self.validate_state_and_certificate()?;
} else if self.state.state() != ClientHandshakeState::HelloSent {
return Err(HandshakeError::InvalidState);
}
Ok(())
}
#[allow(clippy::type_complexity)]
fn extract_key_exchange_crypto_material(
&self,
) -> Result<(PublicKey<P::Curve>, SecretKey<P::Curve>, SubjectPublicKeyInfoOwned), HandshakeError> {
let server_public_key = self.extract_server_public_key()?;
let (sender_ephemeral, sender_pub_spki) = self.create_ephemeral_keypair()?;
Ok((server_public_key, sender_ephemeral, sender_pub_spki))
}
fn create_user_keying_material(&self) -> Result<UserKeyingMaterial, HandshakeError> {
let ukm_bytes = generate_nonce::<64>(None)?;
UserKeyingMaterial::new(ukm_bytes.to_vec()).map_err(Into::into)
}
fn build_kari_structure(
&self,
sender_ephemeral: SecretKey<P::Curve>,
sender_pub_spki: SubjectPublicKeyInfoOwned,
server_public_key: PublicKey<P::Curve>,
rid: KeyAgreeRecipientIdentifier,
ukm: UserKeyingMaterial,
) -> Result<TightBeamKariBuilder<P>, HandshakeError> {
let key_wrap_oid =
<P::Profile as SecurityProfile>::KEY_WRAP_OID.ok_or(HandshakeError::MissingKeyWrapAlgorithm)?;
let key_enc_alg = AlgorithmIdentifierOwned { oid: key_wrap_oid, parameters: None };
let kari_builder = TightBeamKariBuilder::new(self.provider)
.with_sender_priv(sender_ephemeral)
.with_sender_pub_spki(sender_pub_spki)
.with_recipient_pub(server_public_key)
.with_recipient_rid(rid)
.with_ukm(ukm)
.with_key_enc_alg(key_enc_alg);
Ok(kari_builder)
}
fn build_enveloped_data(
&self,
kari_builder: TightBeamKariBuilder<P>,
session_key: &[u8],
) -> Result<Vec<u8>, HandshakeError> {
let mut enveloped_builder = TightBeamEnvelopedDataBuilder::new(kari_builder);
if let Some(ref offer) = self.security_offer {
let offer_attr = crate::transport::handshake::attributes::encode_security_offer(offer)?;
enveloped_builder = enveloped_builder.with_unprotected_attr(offer_attr);
}
let enveloped_data = enveloped_builder.build(session_key, None)?;
enveloped_data.to_der().map_err(Into::into)
}
fn finalize_key_exchange(&mut self, enveloped_data_der: &[u8], session_key: Vec<u8>) -> Result<(), HandshakeError> {
if self.transcript_hash.is_none() {
self.transcript_buffer.extend_from_slice(enveloped_data_der);
}
self.session_key = Some(Secret::from(session_key));
self.state.transition(ClientHandshakeState::KeyExchangeSent)?;
Ok(())
}
fn validate_client_finished_prerequisites(&self) -> Result<(), HandshakeError> {
self.validate_expected_state(ClientHandshakeState::ServerFinishedReceived)
}
fn prepare_finished_digest(&self) -> Result<([u8; 32], Vec<u8>), HandshakeError> {
let transcript_hash = self.transcript_hash.ok_or(HandshakeError::InvalidState)?;
let mut hasher = P::Digest::new();
hasher.update(transcript_hash);
let digest = hasher.finalize();
let digest_bytes = digest.to_vec();
Ok((transcript_hash, digest_bytes))
}
async fn sign_finished_digest(&self, digest: &[u8]) -> Result<Vec<u8>, HandshakeError> {
let signature_bytes = self.client_key_provider.sign(digest).await?;
Ok(signature_bytes)
}
async fn build_finished_crypto_components(
&self,
) -> Result<(SignerIdentifier, AlgorithmIdentifierOwned, AlgorithmIdentifierOwned), HandshakeError> {
use crate::crypto::x509::utils::compute_signer_identifier_from_der;
let public_key_bytes = self.client_key_provider.to_public_key_bytes().await?;
let signer_id = compute_signer_identifier_from_der::<P::Digest>(&public_key_bytes)?;
let digest_alg = AlgorithmIdentifierOwned { oid: P::Digest::OID, parameters: None };
let signature_alg = AlgorithmIdentifierOwned { oid: P::Signature::ALGORITHM_OID, parameters: None };
Ok((signer_id, digest_alg, signature_alg))
}
fn build_signed_data(
&self,
transcript_hash: [u8; 32],
signature_bytes: &[u8],
signer_id: SignerIdentifier,
digest_alg: AlgorithmIdentifierOwned,
signature_alg: AlgorithmIdentifierOwned,
) -> Result<Vec<u8>, HandshakeError> {
let signer_info = SignerInfo {
version: CmsVersion::V1,
sid: signer_id,
digest_alg: digest_alg.clone(),
signed_attrs: None,
signature_algorithm: signature_alg,
signature: OctetString::new(signature_bytes)?,
unsigned_attrs: None,
};
let octet_string = OctetString::new(transcript_hash)?;
let econtent_der = octet_string.to_der()?;
let econtent_any = crate::der::Any::from_der(&econtent_der)?;
let encap_content_info =
EncapsulatedContentInfo { econtent_type: crate::oids::DATA, econtent: Some(econtent_any) };
let signed_data = SignedData {
version: CmsVersion::V1,
digest_algorithms: vec![digest_alg].try_into()?,
encap_content_info,
certificates: None,
crls: None,
signer_infos: vec![signer_info].try_into()?,
};
signed_data.to_der().map_err(Into::into)
}
fn finalize_client_finished(&mut self) -> Result<(), HandshakeError> {
self.state.transition(ClientHandshakeState::ClientFinishedSent)?;
self.invariants.mark_finished_sent()?;
Ok(())
}
}
impl<P> HandshakeFinalization<P> for CmsHandshakeClient<P>
where
P: CryptoProvider,
{
fn selected_profile(&self) -> Option<SecurityProfileDesc> {
self.selected_profile
}
}
impl<P> HandshakeAlertHandler for CmsHandshakeClient<P> where P: CryptoProvider {}
impl<P> ClientHandshakeProtocol for CmsHandshakeClient<P>
where
P: CryptoProvider + Send + Sync + 'static,
P::Curve: elliptic_curve::Curve + elliptic_curve::CurveArithmetic,
<P::Curve as elliptic_curve::Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EncodePublicKey,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + signature::Verifier<P::Signature> + 'static,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: Send + Sync + KeyInit,
{
type Error = HandshakeError;
fn start<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<Vec<u8>, Self::Error>> + Send + 'a>> {
Box::pin(async move { self.build_key_exchange(vec![0u8; 32]) })
}
fn handle_response<'a, 'b>(
&'a mut self,
msg: &'b [u8],
) -> Pin<Box<dyn Future<Output = Result<Option<Vec<u8>>, Self::Error>> + Send + 'a>>
where
'b: 'a,
{
Box::pin(async move {
self.process_server_finished(msg)?;
let client_finished = self.build_client_finished().await?;
Ok(Some(client_finished))
})
}
#[cfg(feature = "aead")]
fn complete<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<crate::crypto::aead::RuntimeAead, Self::Error>> + Send + 'a>> {
Box::pin(async move {
if self.state.state() != ClientHandshakeState::ClientFinishedSent {
return Err(HandshakeError::InvalidState);
}
let cek = self.session_key.as_ref().ok_or(HandshakeError::InvalidState)?;
let profile = self.selected_profile.ok_or(HandshakeError::InvalidState)?;
let aead_oid = profile.aead.ok_or(HandshakeError::InvalidState)?;
let transcript_hash = self.transcript_hash.ok_or(HandshakeError::InvalidState)?;
let cipher = cek.with(|key_bytes| self.derive_session_aead(key_bytes, &transcript_hash))??;
self.state.transition(ClientHandshakeState::Completed)?;
Ok(crate::crypto::aead::RuntimeAead::new(cipher, aead_oid))
})
}
fn is_complete(&self) -> bool {
self.is_complete()
}
fn selected_profile(&self) -> Option<SecurityProfileDesc> {
self.selected_profile
}
}
#[cfg(test)]
mod tests {
use crate::cms::enveloped_data::EnvelopedData;
use crate::crypto::profiles::DefaultCryptoProvider;
use crate::crypto::sign::elliptic_curve::SecretKey;
use crate::der::{Decode, Encode};
use crate::oids::{HASH_SHA3_256, SIGNER_ECDSA_WITH_SHA3_256};
use crate::spki::AlgorithmIdentifierOwned;
use crate::transport::handshake::builders::TightBeamSignedDataBuilder;
use crate::transport::handshake::processors::{TightBeamEnvelopedDataProcessor, TightBeamKariRecipient};
use crate::transport::handshake::state::ClientHandshakeState;
use crate::transport::handshake::tests::*;
#[tokio::test]
async fn test_client_state_flow() -> Result<(), Box<dyn std::error::Error>> {
let transcript_hash = [1u8; 32];
let server_test_cert = create_test_certificate();
let mut client = TestCmsClientBuilder::new()
.with_server_cert(server_test_cert.certificate.clone())
.with_transcript_hash(transcript_hash)
.build()?;
assert_eq!(client.state(), ClientHandshakeState::Init);
let session_key = vec![2u8; 32];
let key_exchange = client.build_key_exchange(session_key.clone())?;
assert_eq!(client.state(), ClientHandshakeState::KeyExchangeSent);
assert!(client.session_key().is_some());
let enveloped_data = EnvelopedData::from_der(&key_exchange)?;
let server_secret = SecretKey::from(server_test_cert.signing_key.clone());
let provider = DefaultCryptoProvider::default();
let kari_processor = TightBeamKariRecipient::new(provider, server_secret);
let processor = TightBeamEnvelopedDataProcessor::<DefaultCryptoProvider>::new(kari_processor);
let decrypted = processor.process(&enveloped_data)?;
assert_eq!(decrypted, session_key);
let digest_alg = AlgorithmIdentifierOwned { oid: HASH_SHA3_256, parameters: None };
let signature_alg = AlgorithmIdentifierOwned { oid: SIGNER_ECDSA_WITH_SHA3_256, parameters: None };
let server_finished_builder = TightBeamSignedDataBuilder::<DefaultCryptoProvider, _>::new(
&server_test_cert.signing_key,
digest_alg,
signature_alg,
)?;
let server_finished = server_finished_builder.build(&transcript_hash)?;
let server_finished = server_finished.to_der()?;
let verified = client.process_server_finished(&server_finished)?;
assert_eq!(verified, transcript_hash);
assert_eq!(client.state(), ClientHandshakeState::ServerFinishedReceived);
let _client_finished = client.build_client_finished().await?;
assert_eq!(client.state(), ClientHandshakeState::ClientFinishedSent);
client.complete()?;
assert!(client.is_complete());
assert_eq!(client.state(), ClientHandshakeState::Completed);
Ok(())
}
#[tokio::test]
async fn test_invalid_state_transitions() -> Result<(), Box<dyn std::error::Error>> {
let mut client = TestCmsClientBuilder::new().build()?;
let result = client.process_server_finished(&[]);
assert!(result.is_err());
let result = client.build_client_finished().await;
assert!(result.is_err());
Ok(())
}
}