#![cfg(feature = "x509")]
#[cfg(not(feature = "std"))]
extern crate alloc;
#[cfg(not(feature = "std"))]
use alloc::{boxed::Box, sync::Arc, vec::Vec};
#[cfg(not(feature = "std"))]
use core::marker::PhantomData;
#[cfg(feature = "std")]
use std::marker::PhantomData;
#[cfg(feature = "std")]
use std::sync::Arc;
use crate::asn1::OctetString;
use crate::constants::{TIGHTBEAM_AAD_DOMAIN_TAG, TIGHTBEAM_ECIES_KDF_INFO};
use crate::crypto::aead::{Aead, AeadCore, KeyInit, Nonce, Payload};
use crate::crypto::common::{typenum::Unsigned, KeySizeUser};
use crate::crypto::ecies::EciesError;
use crate::crypto::ecies::EciesMessageOps;
use crate::crypto::kdf::ecies_kdf;
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::subtle::ConstantTimeEq;
use crate::crypto::sign::elliptic_curve::{AffinePoint, Curve, CurveArithmetic, PublicKey};
use crate::crypto::sign::{SignatureEncoding, Verifier};
use crate::crypto::x509::policy::CertificateValidation;
use crate::der::{Decode, Encode};
use crate::random::generate_nonce;
use crate::transport::handshake::error::HandshakeError;
use crate::transport::handshake::negotiation::SecurityAccept;
use crate::transport::handshake::state::HandshakeInvariant;
use crate::transport::handshake::state::{ServerHandshakeState, ServerStateMachine};
use crate::transport::handshake::utils::{
clear_session_randoms, compute_transcript_digest, extract_verifying_key_from_cert, octet_string_to_32_byte_array,
validate_state,
};
use crate::transport::handshake::{ClientHello, ClientKeyExchange, ServerHandshake, ServerHandshakeProtocol};
use crate::transport::handshake::{HandshakeAlertHandler, HandshakeFinalization, HandshakeNegotiation};
use crate::x509::Certificate;
pub struct EciesHandshakeServer<P>
where
P: CryptoProvider,
{
state: ServerStateMachine,
server_key_provider: Arc<dyn SigningKeyProvider>,
server_cert: Arc<Certificate>,
client_random: Option<[u8; 32]>,
server_random: Option<[u8; 32]>,
base_session_key: Option<[u8; 32]>,
transcript_hash: Option<[u8; 32]>,
aad_domain_tag: Option<&'static [u8]>,
supported_profiles: Vec<SecurityProfileDesc>,
selected_profile: Option<SecurityProfileDesc>,
client_validators: Option<Arc<Vec<Arc<dyn CertificateValidation>>>>,
validated_client_cert: Option<Arc<Certificate>>,
_phantom: PhantomData<P>,
invariants: HandshakeInvariant,
}
impl<P> EciesHandshakeServer<P>
where
P: CryptoProvider,
P::AeadCipher: KeyInit,
P::Signature: SignatureEncoding,
{
pub fn new(
server_key_provider: Arc<dyn SigningKeyProvider>,
server_cert: Arc<Certificate>,
aad_domain_tag: Option<&'static [u8]>,
client_validators: Option<Arc<Vec<Arc<dyn CertificateValidation>>>>,
) -> Self {
Self {
state: ServerStateMachine::default(),
server_key_provider,
server_cert,
client_random: None,
server_random: None,
base_session_key: None,
transcript_hash: None,
aad_domain_tag: aad_domain_tag.or(Some(TIGHTBEAM_AAD_DOMAIN_TAG)),
supported_profiles: Vec::new(), selected_profile: None,
client_validators,
validated_client_cert: None,
_phantom: PhantomData,
invariants: HandshakeInvariant::default(),
}
}
pub fn with_supported_profiles(mut self, profiles: Vec<SecurityProfileDesc>) -> Self {
self.supported_profiles = profiles;
self
}
pub async fn process_client_hello(&mut self, client_hello_der: &[u8]) -> Result<Vec<u8>, HandshakeError> {
self.validate_expected_state(ServerHandshakeState::Init)?;
let client_hello = self.decode_client_hello(client_hello_der)?;
let selected = self.negotiate_profile(client_hello.security_offer.as_ref())?;
self.selected_profile = Some(selected);
let security_accept = SecurityAccept::new(selected);
let client_random = octet_string_to_32_byte_array(&client_hello.client_random)?;
self.client_random = Some(client_random);
let server_random = self.generate_server_random()?;
let spki_bytes = self
.server_cert
.tbs_certificate
.subject_public_key_info
.subject_public_key
.raw_bytes();
let accept_der = security_accept.to_der()?;
let transcript_digest = self.compute_transcript_hash(&client_random, &server_random, spki_bytes, &accept_der);
self.transcript_hash = Some(transcript_digest);
self.invariants.lock_transcript()?;
let signature_bytes = self.sign_transcript_hash(&transcript_digest).await?;
let server_handshake_der =
self.build_server_handshake(server_random, signature_bytes, Some(security_accept))?;
self.state.transition(ServerHandshakeState::ClientHelloReceived)?;
self.state.transition(ServerHandshakeState::ServerHelloSent)?;
Ok(server_handshake_der)
}
pub async fn process_client_key_exchange(&mut self, client_kex_der: &[u8]) -> Result<(), HandshakeError>
where
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
for<'a> P::Signature: TryFrom<&'a [u8]>,
P::VerifyingKey: Verifier<P::Signature> + for<'a> From<&'a PublicKey<P::Curve>>,
{
self.validate_expected_state(ServerHandshakeState::ServerHelloSent)?;
let mut client_kex = self.decode_client_key_exchange(client_kex_der)?;
self.validate_client_certificate(&mut client_kex)?;
let encrypted_bytes = client_kex.encrypted_data.as_bytes();
let decrypted_payload = self.decrypt_ecies_payload(encrypted_bytes).await?;
let (base_session_key, client_random_from_payload) =
self.extract_session_data_from_payload(&decrypted_payload)?;
self.verify_client_random(&client_random_from_payload)?;
self.base_session_key = Some(base_session_key);
self.state.transition(ServerHandshakeState::KeyExchangeReceived)?;
Ok(())
}
pub fn complete(&mut self) -> Result<P::AeadCipher, HandshakeError> {
self.validate_expected_state(ServerHandshakeState::KeyExchangeReceived)?;
let base_session_key = self.base_session_key.as_ref().ok_or(HandshakeError::MissingBaseSessionKey)?;
let client_random = self.client_random.as_ref().ok_or(HandshakeError::MissingClientRandomState)?;
let server_random = self.server_random.as_ref().ok_or(HandshakeError::MissingServerRandom)?;
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_session_key, &salt)?;
self.invariants.derive_aead_once()?;
self.state.transition(ServerHandshakeState::Completed)?;
self.clear_sensitive_data();
Ok(session_key)
}
pub fn state(&self) -> ServerHandshakeState {
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 compute_transcript_hash(
&self,
client_random: &[u8; 32],
server_random: &[u8; 32],
spki_bytes: &[u8],
accept_der: &[u8],
) -> [u8; 32] {
let mut data = Vec::with_capacity(32 + 32 + spki_bytes.len() + accept_der.len());
data.extend_from_slice(client_random);
data.extend_from_slice(server_random);
data.extend_from_slice(spki_bytes);
data.extend_from_slice(accept_der);
compute_transcript_digest::<P::Digest>(&data)
}
fn validate_expected_state(&self, expected: ServerHandshakeState) -> Result<(), HandshakeError> {
validate_state(self.state.state(), expected)
}
fn decode_client_hello(&self, client_hello_der: &[u8]) -> Result<ClientHello, HandshakeError> {
Ok(ClientHello::from_der(client_hello_der)?)
}
fn generate_server_random(&mut self) -> Result<[u8; 32], HandshakeError> {
let server_random = generate_nonce::<32>(None)?;
self.server_random = Some(server_random);
Ok(server_random)
}
async fn sign_transcript_hash(&self, transcript_digest: &[u8; 32]) -> Result<Vec<u8>, HandshakeError> {
let sig = self.server_key_provider.sign(transcript_digest).await?;
Ok(sig.to_vec())
}
fn build_server_handshake(
&self,
server_random: [u8; 32],
signature_bytes: Vec<u8>,
security_accept: Option<SecurityAccept>,
) -> Result<Vec<u8>, HandshakeError> {
let server_handshake = ServerHandshake {
certificate: Certificate::clone(&self.server_cert),
server_random: OctetString::new(server_random)?,
signature: OctetString::new(signature_bytes)?,
security_accept,
client_cert_required: self.client_validators.is_some(),
};
Ok(server_handshake.to_der()?)
}
pub fn decode_client_key_exchange(&self, der_bytes: &[u8]) -> Result<ClientKeyExchange, HandshakeError> {
ClientKeyExchange::from_der(der_bytes).map_err(Into::into)
}
async fn decrypt_ecies_payload(&self, encrypted_bytes: &[u8]) -> Result<Vec<u8>, HandshakeError> {
let (ephemeral_pubkey, ciphertext_bytes) = {
let encrypted_message = <P::EciesMessage as EciesMessageOps>::from_bytes(encrypted_bytes)?;
(
encrypted_message.ephemeral_pubkey().to_vec(),
encrypted_message.ciphertext().to_vec(),
)
};
let shared_secret_bytes = self.server_key_provider.key_agreement(&ephemeral_pubkey).await?;
let k_enc = ecies_kdf::<P::Kdf>(&ephemeral_pubkey, shared_secret_bytes.into(), TIGHTBEAM_ECIES_KDF_INFO, None)?;
let nonce_size = <P::AeadCipher as AeadCore>::NonceSize::USIZE;
let tag_size = <P::AeadCipher as AeadCore>::TagSize::USIZE;
let key_size = <P::AeadCipher as KeySizeUser>::KeySize::USIZE;
let ciphertext_bytes = ciphertext_bytes.as_slice();
if ciphertext_bytes.len() < nonce_size + tag_size {
return Err(HandshakeError::EciesError(EciesError::InvalidCiphertext));
}
let nonce = Nonce::<P::AeadCipher>::from_slice(&ciphertext_bytes[..nonce_size]);
let ciphertext_with_tag = &ciphertext_bytes[nonce_size..];
let cipher = <P::AeadCipher as KeyInit>::new_from_slice(&k_enc[..key_size])
.map_err(|_| HandshakeError::InvalidKeySize { expected: key_size, received: k_enc.len() })?;
let payload = match self.aad_domain_tag {
Some(aad) => Payload { msg: ciphertext_with_tag, aad },
None => Payload { msg: ciphertext_with_tag, aad: b"" },
};
let plaintext = cipher.decrypt(nonce, payload)?;
Ok(plaintext)
}
fn extract_session_data_from_payload(
&self,
decrypted_payload: &[u8],
) -> Result<([u8; 32], [u8; 32]), HandshakeError> {
if decrypted_payload.len() != 64 {
return Err(HandshakeError::InvalidDecryptedPayloadSize);
}
let mut base_session_key = [0u8; 32];
let mut client_random_from_payload = [0u8; 32];
base_session_key.copy_from_slice(&decrypted_payload[..32]);
client_random_from_payload.copy_from_slice(&decrypted_payload[32..]);
Ok((base_session_key, client_random_from_payload))
}
fn verify_client_random(&self, client_random_from_payload: &[u8; 32]) -> Result<(), HandshakeError> {
let expected_client_random = self.client_random.ok_or(HandshakeError::MissingClientRandom)?;
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
let is_equal: bool = client_random_from_payload.ct_eq(&expected_client_random).into();
core::sync::atomic::compiler_fence(core::sync::atomic::Ordering::SeqCst);
if !is_equal {
Err(HandshakeError::ClientRandomMismatchReplay)
} else {
Ok(())
}
}
#[cfg(feature = "x509")]
fn validate_client_certificate(&mut self, client_kex: &mut ClientKeyExchange) -> Result<(), HandshakeError>
where
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
for<'a> P::Signature: TryFrom<&'a [u8]>,
P::VerifyingKey: Verifier<P::Signature> + for<'a> From<&'a PublicKey<P::Curve>>,
{
if let Some(validators) = &self.client_validators {
let client_cert = client_kex
.client_certificate
.take()
.ok_or(HandshakeError::MissingClientCertificate)?;
for validator in validators.iter() {
validator.evaluate(&client_cert)?;
}
let client_signature = client_kex
.client_signature
.as_ref()
.ok_or(HandshakeError::SignatureVerificationFailed)?;
let transcript_hash = self.transcript_hash.ok_or(HandshakeError::InvalidState)?;
let public_key = extract_verifying_key_from_cert::<P::Curve>(&client_cert)?;
let signature = P::Signature::try_from(client_signature.as_bytes())
.map_err(|_| HandshakeError::SignatureVerificationFailed)?;
let verifying_key = P::VerifyingKey::from(&public_key);
verifying_key.verify(&transcript_hash, &signature)?;
self.validated_client_cert = Some(Arc::new(client_cert));
}
Ok(())
}
fn clear_sensitive_data(&mut self) {
clear_session_randoms(&mut self.base_session_key, &mut self.client_random, &mut self.server_random);
}
}
impl<P> HandshakeNegotiation for EciesHandshakeServer<P>
where
P: CryptoProvider,
{
fn supported_profiles(&self) -> &[SecurityProfileDesc] {
&self.supported_profiles
}
}
impl<P> HandshakeFinalization<P> for EciesHandshakeServer<P>
where
P: CryptoProvider,
{
fn selected_profile(&self) -> Option<SecurityProfileDesc> {
self.selected_profile
}
}
impl<P> HandshakeAlertHandler for EciesHandshakeServer<P> where P: CryptoProvider {}
impl<P> ServerHandshakeProtocol for EciesHandshakeServer<P>
where
P: CryptoProvider + Send + Sync,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
for<'a> P::Signature: TryFrom<&'a [u8]>,
P::VerifyingKey: Verifier<P::Signature> + for<'a> From<&'a PublicKey<P::Curve>>,
P::AeadCipher: KeyInit + Send + Sync + 'static,
P::Signature: SignatureEncoding,
{
type Error = HandshakeError;
fn handle_request<'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 {
match self.state() {
ServerHandshakeState::Init => {
let server_handshake = self.process_client_hello(msg).await?;
Ok(Some(server_handshake))
}
ServerHandshakeState::ServerHelloSent => {
self.process_client_key_exchange(msg).await?;
Ok(None)
}
_ => Err(HandshakeError::InvalidState),
}
})
}
#[cfg(feature = "aead")]
fn complete<'a>(
&'a mut self,
) -> core::pin::Pin<
Box<dyn core::future::Future<Output = Result<crate::crypto::aead::RuntimeAead, Self::Error>> + Send + 'a>,
> {
Box::pin(async move {
self.validate_expected_state(ServerHandshakeState::KeyExchangeReceived)?;
let base_session_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 profile = self.selected_profile.ok_or(HandshakeError::InvalidState)?;
let aead_oid = profile.aead.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_session_key, &salt)?;
self.state.transition(ServerHandshakeState::Completed)?;
self.clear_sensitive_data();
Ok(crate::crypto::aead::RuntimeAead::new(cipher, aead_oid))
})
}
fn is_complete(&self) -> bool {
self.is_complete()
}
#[cfg(feature = "x509")]
fn peer_certificate(&self) -> Option<&Certificate> {
self.validated_client_cert.as_ref().map(|arc| arc.as_ref())
}
fn selected_profile(&self) -> Option<SecurityProfileDesc> {
self.selected_profile
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::crypto::profiles::SecurityProfileDesc;
use crate::random::OsRng;
use crate::transport::handshake::negotiation::{select_profile, SecurityOffer};
use crate::transport::handshake::tests::*;
fn create_test_client_hello_with_offer(
client_random: &[u8; 32],
offer: Option<SecurityOffer>,
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let client_hello = ClientHello {
client_random: crate::asn1::OctetString::new(*client_random)?,
security_offer: offer,
};
Ok(client_hello.to_der()?)
}
#[tokio::test]
async fn test_server_state_flow() -> Result<(), Box<dyn std::error::Error>> {
let mut server = TestEciesServerBuilder::new().build()?;
assert_eq!(server.state(), ServerHandshakeState::Init);
let client_random = crate::random::generate_nonce::<32>(None)?;
let client_hello_der = create_test_client_hello(&client_random)?;
let server_handshake_der = server.process_client_hello(&client_hello_der).await?;
assert_eq!(server.state(), ServerHandshakeState::ServerHelloSent);
assert!(server.client_random.is_some());
assert!(server.server_random.is_some());
assert!(server.transcript_hash.is_some());
let _server_handshake = ServerHandshake::from_der(&server_handshake_der)?;
let client_kex_der = build_test_client_key_exchange(&server)?;
server.process_client_key_exchange(&client_kex_der).await?;
assert_eq!(server.state(), ServerHandshakeState::KeyExchangeReceived);
assert!(server.base_session_key.is_some());
let _session_key = server.complete()?;
assert!(server.is_complete());
assert_eq!(server.state(), ServerHandshakeState::Completed);
Ok(())
}
#[tokio::test]
async fn test_invalid_state_transitions() -> Result<(), Box<dyn std::error::Error>> {
let mut server = TestEciesServerBuilder::new().build()?;
assert!(server.process_client_key_exchange(&[]).await.is_err());
assert!(server.complete().is_err());
let client_random = crate::random::generate_nonce::<32>(None)?;
let client_hello_der = create_test_client_hello(&client_random)?;
server.process_client_hello(&client_hello_der).await?;
assert!(server.process_client_hello(&client_hello_der).await.is_err());
assert!(server.complete().is_err());
let client_kex_der = build_test_client_key_exchange(&server)?;
server.process_client_key_exchange(&client_kex_der).await?;
assert!(server.process_client_key_exchange(&client_kex_der).await.is_err());
assert!(server.process_client_hello(&client_hello_der).await.is_err());
Ok(())
}
#[tokio::test]
async fn test_profile_negotiation() -> Result<(), Box<dyn std::error::Error>> {
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,
};
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 offer = SecurityOffer::new(vec![p_a, p_b]);
let selected = select_profile(&offer, &[p_b, p_c])?;
assert_eq!(selected, p_b);
let mut server = TestEciesServerBuilder::new().build()?.with_supported_profiles(vec![p_b, p_c]);
let client_random = [0u8; 32];
let client_hello_der = create_test_client_hello_with_offer(&client_random, Some(offer.clone()))?;
let _response = server.process_client_hello(&client_hello_der).await?;
assert_eq!(server.selected_profile, Some(p_b));
}
{
let mut server = TestEciesServerBuilder::new().build()?.with_supported_profiles(vec![p_a, p_b]);
let client_random = [1u8; 32];
let client_hello_der = create_test_client_hello(&client_random)?;
let _response = server.process_client_hello(&client_hello_der).await?;
assert_eq!(server.selected_profile, Some(p_a)); }
{
let offer = SecurityOffer::new(vec![p_a, p_b]);
let result = select_profile(&offer, &[p_c]);
assert!(result.is_err());
}
Ok(())
}
fn build_test_client_key_exchange<P>(
server: &EciesHandshakeServer<P>,
) -> Result<Vec<u8>, Box<dyn std::error::Error>>
where
P: CryptoProvider,
P::AeadCipher: KeyInit,
{
use crate::crypto::ecies::encrypt;
let server_pubkey = k256::PublicKey::from_sec1_bytes(
server
.server_cert
.tbs_certificate
.subject_public_key_info
.subject_public_key
.raw_bytes(),
)?;
let stored_client_random = server.client_random.ok_or("Missing client random")?;
let base_session_key = crate::random::generate_nonce::<32>(None)?;
let mut plaintext = [0u8; 64];
plaintext[..32].copy_from_slice(&base_session_key);
plaintext[32..].copy_from_slice(&stored_client_random);
let aad = server.aad_domain_tag.or(Some(crate::constants::TIGHTBEAM_AAD_DOMAIN_TAG));
let encrypted_message = encrypt::<_, _, _, crate::crypto::ecies::Secp256k1EciesMessage, P::Kdf, P::AeadCipher>(
&server_pubkey,
&plaintext,
aad,
Some(&mut OsRng),
)?;
create_test_client_key_exchange(&encrypted_message.to_bytes())
}
}