#[cfg(not(feature = "std"))]
extern crate alloc;
#[cfg(all(
not(feature = "std"),
any(feature = "transport-cms", feature = "transport-ecies")
))]
use alloc::boxed::Box;
#[cfg(not(feature = "std"))]
use alloc::sync::Arc;
#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
#[cfg(feature = "std")]
use std::sync::Arc;
use crate::asn1::Frame;
use crate::der::{Decode, Encode};
use crate::encode;
use crate::policy::TransitStatus;
use crate::transport::envelopes::{TransportEnvelope, WireEnvelope, WireMode};
use crate::transport::error::TransportError;
use crate::transport::messaging::ResponseHandler;
use crate::transport::TransportResult;
#[cfg(feature = "x509")]
mod x509 {
pub use crate::crypto::aead::Decryptor;
pub use crate::transport::builders::{EnvelopeBuilder, EnvelopeLimits};
pub use crate::transport::handshake::TcpHandshakeState;
pub use crate::transport::state::EncryptedProtocolState;
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
mod handshake {
pub use crate::crypto::aead::KeyInit;
pub use crate::crypto::profiles::{CryptoProvider, SecurityProfileDesc, TightbeamProfile};
pub use crate::crypto::sign::elliptic_curve::sec1::{FromEncodedPoint, ModulusSize, ToEncodedPoint};
pub use crate::crypto::sign::elliptic_curve::{AffinePoint, Curve, CurveArithmetic, PublicKey};
pub use crate::crypto::sign::Verifier;
pub use crate::spki::EncodePublicKey;
pub use crate::transport::handshake::{
ClientHandshakeProtocol, HandshakeError, HandshakeProtocolKind, ServerHandshakeProtocol,
};
}
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
pub use handshake::*;
#[cfg(feature = "transport-ecies")]
mod ecies {
pub use crate::cms::enveloped_data::EnvelopedData;
pub use crate::cms::signed_data::SignedData;
pub use crate::crypto::aead::RuntimeAead;
pub use crate::crypto::ecies::{EciesEphemeral, EciesMessageOps, EciesPublicKeyOps};
pub use crate::crypto::sign::SignatureEncoding;
pub use crate::der::oid::AssociatedOid;
pub use crate::transport::handshake::client::{EciesHandshakeClient, ExtractVerifyingKey};
pub use crate::transport::handshake::{ClientHello, ClientKeyExchange, HandshakeFinalization, ServerHandshake};
#[cfg(feature = "std")]
pub use crate::crypto::x509::policy::CertificateValidation;
}
#[cfg(feature = "transport-ecies")]
pub use ecies::*;
#[cfg(feature = "transport-cms")]
mod cms {
pub use crate::transport::handshake::negotiation::SecurityOffer;
pub use crate::transport::handshake::CmsClientConfig;
}
#[cfg(feature = "transport-cms")]
pub use cms::*;
}
#[cfg(feature = "x509")]
use x509::*;
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
pub(crate) const HANDSHAKE_MAX_WIRE: usize = 16 * 1024;
pub(crate) fn parse_der_length(first_byte: u8, length_octets: &[u8]) -> Option<usize> {
if first_byte & 0x80 == 0 {
return Some(first_byte as usize);
}
let octet_count = (first_byte & 0x7F) as usize;
if octet_count == 0 || octet_count != length_octets.len() || octet_count > core::mem::size_of::<usize>() {
return None;
}
if length_octets[0] == 0 {
return None;
}
let mut length = 0usize;
for &byte in length_octets.iter() {
length = (length << 8) | (byte as usize);
}
if length < 0x80 {
return None;
}
Some(length)
}
pub(crate) fn reconstruct_der_encoding(tag: u8, length_first: u8, length_octets: &[u8], content: &[u8]) -> Vec<u8> {
let length_bytes_count = if length_first & 0x80 == 0 {
1
} else {
1 + length_octets.len()
};
let mut buffer = Vec::with_capacity(1 + length_bytes_count + content.len());
buffer.push(tag);
buffer.push(length_first);
if length_first & 0x80 != 0 {
buffer.extend_from_slice(length_octets);
}
buffer.extend_from_slice(content);
buffer
}
fn envelope_versions_compatible(envelope: &TransportEnvelope) -> bool {
match envelope {
TransportEnvelope::Request(pkg) => pkg.message.validate_version_compatibility(),
TransportEnvelope::Response(pkg) => {
pkg.message.as_ref().is_none_or(|frame| frame.validate_version_compatibility())
}
#[cfg(feature = "x509")]
TransportEnvelope::EnvelopedData(_) | TransportEnvelope::SignedData(_) => true,
}
}
pub trait MessageIO: ResponseHandler {
#[allow(async_fn_in_trait)]
async fn read_envelope(&mut self) -> TransportResult<Vec<u8>>;
#[allow(async_fn_in_trait)]
async fn write_envelope(&mut self, buffer: &[u8]) -> TransportResult<()>;
fn decode_envelope(buffer: &[u8]) -> TransportResult<TransportEnvelope> {
let envelope = TransportEnvelope::from_der(buffer)?;
if !envelope_versions_compatible(&envelope) {
return Err(TransportError::InvalidMessage);
}
Ok(envelope)
}
fn encode_envelope(envelope: &TransportEnvelope) -> TransportResult<Vec<u8>> {
Ok(encode(envelope)?)
}
#[allow(async_fn_in_trait)]
async fn read_decoded_envelope(&mut self) -> TransportResult<TransportEnvelope> {
let bytes = self.read_envelope().await?;
Self::decode_envelope(&bytes)
}
#[allow(async_fn_in_trait)]
async fn try_read_decoded_envelope(&mut self) -> TransportResult<Option<TransportEnvelope>> {
match self.read_decoded_envelope().await {
Ok(envelope) => Ok(Some(envelope)),
Err(TransportError::ConnectionClosed) => Ok(None),
Err(e) => Err(e),
}
}
fn handle_message(&self, message: Arc<Frame>) -> Option<Frame> {
let frame = Arc::try_unwrap(message).unwrap_or_else(|arc| (*arc).clone());
self.handler().and_then(|handler| handler(frame))
}
fn parse_der_length(first_byte: u8, length_octets: &[u8]) -> Option<usize> {
parse_der_length(first_byte, length_octets)
}
fn reconstruct_der_encoding(tag: u8, length_first: u8, length_octets: &[u8], content: &[u8]) -> Vec<u8> {
reconstruct_der_encoding(tag, length_first, length_octets, content)
}
}
#[cfg(feature = "x509")]
pub trait EncryptedMessageIO: MessageIO {
#[allow(async_fn_in_trait)]
async fn relay_message(&mut self) -> TransportResult<TransportEnvelope>
where
Self: EncryptedProtocolState,
{
let wire_bytes = self.read_envelope().await?;
let wire_envelope = WireEnvelope::from_der(&wire_bytes)?;
match wire_envelope {
WireEnvelope::Cleartext(transport_envelope) => {
if self.to_decryptor_ref().is_ok() {
return Err(TransportError::MissingEncryption);
}
Ok(transport_envelope)
}
WireEnvelope::Encrypted(encrypted_info) => {
let decrypted_bytes = self.to_decryptor_ref()?.decrypt_content(&encrypted_info)?;
decrypted_bytes
.with(|bytes| Self::decode_envelope(bytes))
.map_err(crate::error::TightBeamError::from)?
}
}
}
#[allow(async_fn_in_trait)]
async fn send_envelope(&mut self, envelope: TransportEnvelope, encrypt: bool) -> TransportResult<()>
where
Self: EncryptedProtocolState,
{
let wire_envelope = if encrypt {
let envelope_bytes = Self::encode_envelope(&envelope)?;
let encrypted_info = self.to_encryptor_ref()?.encrypt_content(&envelope_bytes, [], None)?;
WireEnvelope::Encrypted(encrypted_info)
} else {
WireEnvelope::Cleartext(envelope)
};
let wire_bytes = wire_envelope.to_der()?;
self.write_envelope(&wire_bytes).await
}
fn wrap_message(message: Frame) -> TransportEnvelope {
TransportEnvelope::new_request(message)
}
#[allow(async_fn_in_trait)]
async fn wrap_and_encrypt_message(&mut self, message: Frame) -> TransportResult<WireEnvelope>
where
Self: EncryptedProtocolState,
{
let limits = EnvelopeLimits::from_pair(self.to_max_cleartext_envelope(), self.to_max_encrypted_envelope());
let mut builder = limits.apply(EnvelopeBuilder::request(message));
if self.to_handshake_state() == TcpHandshakeState::Complete {
let encryptor = self.to_encryptor_ref()?;
let wire_mode = WireMode::Encrypted;
builder = builder.with_wire_mode(wire_mode);
builder = builder.with_encryptor(encryptor);
} else {
let wire_mode = WireMode::Cleartext;
builder = builder.with_wire_mode(wire_mode);
}
builder.finish()
}
#[allow(async_fn_in_trait)]
async fn decrypt_response(&mut self, wire_bytes: Vec<u8>) -> TransportResult<TransportEnvelope>
where
Self: EncryptedProtocolState,
{
let wire_envelope = WireEnvelope::from_der(&wire_bytes)?;
match wire_envelope {
WireEnvelope::Cleartext(env) => Ok(env),
WireEnvelope::Encrypted(encrypted_info) => {
let decrypted_bytes = self.to_decryptor_ref()?.decrypt_content(&encrypted_info)?;
decrypted_bytes
.with(|bytes| Self::decode_envelope(bytes))
.map_err(crate::error::TightBeamError::from)?
}
}
}
#[cfg(feature = "transport-ecies")]
#[allow(async_fn_in_trait)]
async fn ensure_handshake_complete<P>(&mut self) -> TransportResult<()>
where
Self: Sized + EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Default + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic + AssociatedOid,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EciesPublicKeyOps + EncodePublicKey,
<PublicKey<P::Curve> as EciesPublicKeyOps>::SecretKey: EciesEphemeral<PublicKey = PublicKey<P::Curve>>,
P::Signature: SignatureEncoding,
for<'b> P::Signature: TryFrom<&'b [u8]>,
for<'b> <P::Signature as TryFrom<&'b [u8]>>::Error: Into<HandshakeError>,
P::VerifyingKey: Verifier<P::Signature> + ExtractVerifyingKey + From<PublicKey<P::Curve>> + EncodePublicKey,
P::AeadCipher: KeyInit,
{
let should_handshake = (self.to_server_certificate_ref().is_some()
|| self.to_trust_store_ref().is_some()
|| self.is_client_validators_present())
&& self.to_handshake_state() == TcpHandshakeState::None;
if should_handshake {
self.perform_client_handshake().await?;
}
Ok(())
}
#[cfg(all(not(feature = "transport-ecies"), feature = "transport-cms"))]
#[allow(async_fn_in_trait)]
async fn ensure_handshake_complete<P>(&mut self) -> TransportResult<()>
where
Self: Sized + EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Default + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EncodePublicKey,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + Verifier<P::Signature> + 'static,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit + Send + Sync,
{
let should_handshake = (self.to_server_certificate_ref().is_some()
|| self.to_trust_store_ref().is_some()
|| self.is_client_validators_present())
&& self.to_handshake_state() == TcpHandshakeState::None;
if should_handshake {
self.perform_client_handshake().await?;
}
Ok(())
}
#[cfg(feature = "transport-ecies")]
#[allow(async_fn_in_trait)]
async fn perform_client_handshake_no_mutual_auth<P>(&mut self) -> TransportResult<()>
where
Self: Sized + MessageIO + EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Default + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic + AssociatedOid,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EciesPublicKeyOps + EncodePublicKey,
<PublicKey<P::Curve> as EciesPublicKeyOps>::SecretKey: EciesEphemeral<PublicKey = PublicKey<P::Curve>>,
P::Signature: SignatureEncoding,
for<'b> P::Signature: TryFrom<&'b [u8]>,
for<'b> <P::Signature as TryFrom<&'b [u8]>>::Error: Into<HandshakeError>,
P::VerifyingKey: Verifier<P::Signature> + ExtractVerifyingKey + From<PublicKey<P::Curve>> + EncodePublicKey,
P::AeadCipher: KeyInit,
P::EciesMessage: EciesMessageOps,
{
let mut client = EciesHandshakeClient::<P, P::EciesMessage>::new(None);
#[cfg(all(feature = "x509", feature = "std"))]
if let Some(store) = self.to_trust_store_ref() {
let validator = Arc::clone(store) as Arc<dyn CertificateValidation>;
client = client.with_certificate_validator(validator);
}
let initial_message = client.build_client_hello()?;
if initial_message.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let client_hello = ClientHello::from_der(&initial_message)?;
let signed_data: SignedData = (&client_hello).try_into().map_err(|_| TransportError::InvalidMessage)?;
let signed_data = Box::new(signed_data);
let initial_envelope = TransportEnvelope::SignedData(signed_data);
let wire_envelope = WireEnvelope::Cleartext(initial_envelope);
self.write_envelope(&wire_envelope.to_der()?).await?;
#[cfg(all(feature = "std", not(target_arch = "wasm32")))]
{
self.set_handshake_state(TcpHandshakeState::AwaitingServerResponse {
initiated_at: std::time::Instant::now(),
});
}
#[cfg(not(all(feature = "std", not(target_arch = "wasm32"))))]
{
self.set_handshake_state(TcpHandshakeState::AwaitingServerResponse { initiated_at: 0 });
}
let response_wire_bytes = self.read_envelope().await?;
if response_wire_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let response_wire = WireEnvelope::from_der(&response_wire_bytes)?;
let response_envelope = match response_wire {
WireEnvelope::Cleartext(env) => env,
WireEnvelope::Encrypted(_) => return Err(TransportError::InvalidMessage),
};
let signed_data = match response_envelope {
TransportEnvelope::SignedData(sd) => sd,
_ => return Err(TransportError::InvalidMessage),
};
let server_handshake: ServerHandshake =
signed_data.as_ref().try_into().map_err(|_| TransportError::InvalidMessage)?;
let response_bytes = server_handshake.to_der()?;
if response_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let next_message_bytes = client.process_server_handshake(&response_bytes).await?;
if next_message_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let client_kex = ClientKeyExchange::from_der(&next_message_bytes)?;
let enveloped_data: EnvelopedData = (&client_kex).try_into().map_err(|_| TransportError::InvalidMessage)?;
let enveloped_data = Box::new(enveloped_data);
let msg_envelope = TransportEnvelope::EnvelopedData(enveloped_data);
let wire_envelope = WireEnvelope::Cleartext(msg_envelope);
self.write_envelope(&wire_envelope.to_der()?).await?;
let cipher = client.complete()?;
let profile = HandshakeFinalization::selected_profile(&client).ok_or(TransportError::InvalidMessage)?;
let aead_oid = profile.aead.ok_or(TransportError::InvalidMessage)?;
let session_key = RuntimeAead::new(cipher, aead_oid);
self.set_symmetric_key(session_key);
self.set_handshake_state(TcpHandshakeState::Complete);
Ok(())
}
#[cfg(feature = "transport-ecies")]
fn build_ecies_client_orchestrator<P>(
&self,
) -> TransportResult<Box<dyn ClientHandshakeProtocol<Error = HandshakeError> + Send + 'static>>
where
Self: EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Default + Send + Sync + 'static,
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 + 'static,
for<'b> P::Signature: TryFrom<&'b [u8]>,
for<'b> <P::Signature as TryFrom<&'b [u8]>>::Error: Into<HandshakeError>,
P::VerifyingKey: Verifier<P::Signature> + ExtractVerifyingKey + 'static,
P::AeadCipher: KeyInit + Send + Sync,
{
#[cfg(all(feature = "x509", feature = "std"))]
let validator = self
.to_trust_store_ref()
.map(|store| Arc::clone(store) as Arc<dyn CertificateValidation>);
#[cfg(not(all(feature = "x509", feature = "std")))]
let validator = None;
let key = self.to_key_manager_ref().ok_or(TransportError::MissingEncryption)?;
let client_cert = self.to_client_certificate_ref().map(Arc::clone);
Ok(key.create_ecies_client::<crate::crypto::ecies::Secp256k1EciesMessage>(
None,
client_cert,
None,
validator,
)?)
}
#[cfg(feature = "transport-cms")]
fn build_cms_client_orchestrator<P>(
&self,
) -> TransportResult<Box<dyn ClientHandshakeProtocol<Error = HandshakeError> + Send + 'static>>
where
Self: EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Default + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EncodePublicKey,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + Verifier<P::Signature> + 'static,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit + Send + Sync,
{
let key = self.to_key_manager_ref().ok_or(TransportError::MissingEncryption)?;
let store = self
.to_trust_store_ref()
.ok_or(TransportError::HandshakeError(HandshakeError::MissingTrustStore))?;
let chain = self
.to_server_certificate_chain_ref()
.ok_or(TransportError::MissingServerCertificateChain)?;
let trust_store = Arc::clone(store);
let server_identity = Arc::clone(chain).into();
let security_offer = Some(SecurityOffer::new(vec![SecurityProfileDesc::from(&TightbeamProfile)]));
let client_certificate = self.to_client_certificate_ref().map(Arc::clone);
Ok(key.create_cms_client(CmsClientConfig {
server_identity,
trust_store,
security_offer,
client_certificate,
})?)
}
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
#[allow(async_fn_in_trait)]
async fn drive_client_handshake(
&mut self,
kind: HandshakeProtocolKind,
mut orchestrator: Box<dyn ClientHandshakeProtocol<Error = HandshakeError> + Send>,
) -> TransportResult<()>
where
Self: Sized + MessageIO + EncryptedProtocolState,
{
let initial_message = orchestrator.start().await?;
if initial_message.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let initial_envelope = kind.wrap_client_start(&initial_message)?;
let wire_envelope = WireEnvelope::Cleartext(initial_envelope);
self.write_envelope(&wire_envelope.to_der()?).await?;
#[cfg(all(feature = "std", not(target_arch = "wasm32")))]
{
self.set_handshake_state(TcpHandshakeState::AwaitingServerResponse {
initiated_at: std::time::Instant::now(),
});
}
#[cfg(not(all(feature = "std", not(target_arch = "wasm32"))))]
{
self.set_handshake_state(TcpHandshakeState::AwaitingServerResponse { initiated_at: 0 });
}
let response_wire_bytes = self.read_envelope().await?;
if response_wire_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let response_wire = WireEnvelope::from_der(&response_wire_bytes)?;
let response_envelope = match response_wire {
WireEnvelope::Cleartext(env) => env,
WireEnvelope::Encrypted(_) => {
return Err(TransportError::InvalidMessage);
}
};
let response_bytes = kind.unwrap_server_response(response_envelope)?;
if response_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let next_message = orchestrator.handle_response(&response_bytes).await?;
if let Some(msg_bytes) = next_message {
if msg_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let msg_envelope = kind.wrap_client_followup(&msg_bytes)?;
let wire_envelope = WireEnvelope::Cleartext(msg_envelope);
self.write_envelope(&wire_envelope.to_der()?).await?;
}
let session_key = orchestrator.complete().await?;
self.set_symmetric_key(session_key);
self.set_handshake_state(TcpHandshakeState::Complete);
Ok(())
}
#[cfg(feature = "transport-ecies")]
#[allow(async_fn_in_trait)]
async fn perform_client_handshake<P>(&mut self) -> TransportResult<()>
where
Self: Sized + MessageIO + EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Default + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic + AssociatedOid,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EciesPublicKeyOps + EncodePublicKey,
<PublicKey<P::Curve> as EciesPublicKeyOps>::SecretKey: EciesEphemeral<PublicKey = PublicKey<P::Curve>>,
P::Signature: SignatureEncoding + 'static,
for<'b> P::Signature: TryFrom<&'b [u8]>,
for<'b> <P::Signature as TryFrom<&'b [u8]>>::Error: Into<HandshakeError>,
P::VerifyingKey:
Verifier<P::Signature> + ExtractVerifyingKey + From<PublicKey<P::Curve>> + EncodePublicKey + 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit + Send + Sync,
{
let kind = self.to_handshake_protocol_kind();
if matches!(kind, HandshakeProtocolKind::Ecies) && self.to_key_manager_ref().is_none() {
return self.perform_client_handshake_no_mutual_auth().await;
}
let orchestrator = match kind {
HandshakeProtocolKind::Ecies => self.build_ecies_client_orchestrator()?,
#[cfg(feature = "transport-cms")]
HandshakeProtocolKind::Cms => self.build_cms_client_orchestrator()?,
#[cfg(not(feature = "transport-cms"))]
HandshakeProtocolKind::Cms => {
return Err(TransportError::UnsupportedHandshakeProtocol(HandshakeProtocolKind::Cms));
}
};
self.drive_client_handshake(kind, orchestrator).await
}
#[cfg(all(not(feature = "transport-ecies"), feature = "transport-cms"))]
#[allow(async_fn_in_trait)]
async fn perform_client_handshake<P>(&mut self) -> TransportResult<()>
where
Self: Sized + MessageIO + EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Default + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EncodePublicKey,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + Verifier<P::Signature> + 'static,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit + Send + Sync,
{
let kind = self.to_handshake_protocol_kind();
let orchestrator = match kind {
HandshakeProtocolKind::Ecies => {
return Err(TransportError::UnsupportedHandshakeProtocol(HandshakeProtocolKind::Ecies));
}
HandshakeProtocolKind::Cms => self.build_cms_client_orchestrator()?,
};
self.drive_client_handshake(kind, orchestrator).await
}
#[cfg(feature = "transport-ecies")]
fn build_ecies_server_orchestrator<P>(
&self,
) -> TransportResult<Box<dyn ServerHandshakeProtocol<Error = HandshakeError> + Send + Sync + 'static>>
where
Self: EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
P::Signature: SignatureEncoding,
for<'b> P::Signature: TryFrom<&'b [u8]>,
P::VerifyingKey: Verifier<P::Signature> + for<'b> From<&'b PublicKey<P::Curve>>,
P::AeadCipher: KeyInit + Send + Sync + 'static,
{
let cert_arc = self.to_server_certificate_arc().ok_or(TransportError::MissingEncryption)?;
let key_manager = self.to_key_manager_ref().ok_or(TransportError::MissingEncryption)?;
let client_validators = self.to_client_validators_ref().map(Arc::clone);
let supported_profiles = vec![SecurityProfileDesc::from(&TightbeamProfile)];
Ok(key_manager.create_ecies_server(cert_arc, None, supported_profiles, client_validators)?)
}
#[cfg(feature = "transport-cms")]
fn build_cms_server_orchestrator<P>(
&self,
) -> TransportResult<Box<dyn ServerHandshakeProtocol<Error = HandshakeError> + Send + Sync + 'static>>
where
Self: EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + Verifier<P::Signature> + 'static,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit + Send + Sync + 'static,
{
let key_manager = self.to_key_manager_ref().ok_or(TransportError::MissingEncryption)?;
let client_validators = self.to_client_validators_ref().map(Arc::clone);
let supported_profiles = vec![SecurityProfileDesc::from(&TightbeamProfile)];
Ok(key_manager.create_cms_server(client_validators, supported_profiles)?)
}
#[cfg(any(feature = "transport-cms", feature = "transport-ecies"))]
#[allow(async_fn_in_trait)]
async fn drive_server_handshake(&mut self, kind: HandshakeProtocolKind, request: &[u8]) -> TransportResult<()>
where
Self: Sized + MessageIO + EncryptedProtocolState,
{
let orchestrator = self.to_server_handshake_mut().as_mut().ok_or(TransportError::InvalidState)?;
let response_bytes = orchestrator.handle_request(request).await?;
if let Some(response) = response_bytes {
if response.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let server_envelope = kind.wrap_server_response(&response)?;
let wire_envelope = WireEnvelope::Cleartext(server_envelope);
self.write_envelope(&wire_envelope.to_der()?).await?;
#[cfg(all(feature = "std", not(target_arch = "wasm32")))]
{
self.set_handshake_state(TcpHandshakeState::AwaitingClientFinish {
initiated_at: std::time::Instant::now(),
});
}
#[cfg(not(all(feature = "std", not(target_arch = "wasm32"))))]
{
self.set_handshake_state(TcpHandshakeState::AwaitingClientFinish { initiated_at: 0 });
}
} else {
let session_key = orchestrator.complete().await?;
if let Some(peer_cert) = orchestrator.peer_certificate().cloned() {
self.set_peer_certificate(peer_cert);
}
self.set_symmetric_key(session_key);
self.set_handshake_state(TcpHandshakeState::Complete);
*self.to_server_handshake_mut() = None;
}
Ok(())
}
#[cfg(feature = "transport-ecies")]
#[allow(async_fn_in_trait)]
async fn perform_server_handshake<P>(&mut self, handshake_bytes: &[u8]) -> TransportResult<()>
where
Self: Sized + MessageIO + EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
PublicKey<P::Curve>: EciesPublicKeyOps,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + Verifier<P::Signature> + 'static,
for<'b> P::VerifyingKey: From<&'b PublicKey<P::Curve>>,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit + Send + Sync + 'static,
{
if handshake_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let kind = self.to_handshake_protocol_kind();
let transport_envelope = TransportEnvelope::from_der(handshake_bytes)?;
let raw_message = kind.unwrap_client_request(&transport_envelope)?;
if self.to_server_handshake_mut().is_none() {
let orchestrator = match kind {
HandshakeProtocolKind::Ecies => self.build_ecies_server_orchestrator()?,
#[cfg(feature = "transport-cms")]
HandshakeProtocolKind::Cms => self.build_cms_server_orchestrator()?,
#[cfg(not(feature = "transport-cms"))]
HandshakeProtocolKind::Cms => {
return Err(TransportError::UnsupportedHandshakeProtocol(HandshakeProtocolKind::Cms));
}
};
*self.to_server_handshake_mut() = Some(orchestrator);
}
self.drive_server_handshake(kind, &raw_message).await
}
#[cfg(all(not(feature = "transport-ecies"), feature = "transport-cms"))]
#[allow(async_fn_in_trait)]
async fn perform_server_handshake<P>(&mut self, handshake_bytes: &[u8]) -> TransportResult<()>
where
Self: Sized + MessageIO + EncryptedProtocolState<CryptoProvider = P>,
P: CryptoProvider + Send + Sync + 'static,
P::Curve: Curve + CurveArithmetic,
<P::Curve as Curve>::FieldBytesSize: ModulusSize,
AffinePoint<P::Curve>: FromEncodedPoint<P::Curve> + ToEncodedPoint<P::Curve>,
P::VerifyingKey: From<PublicKey<P::Curve>> + EncodePublicKey + Verifier<P::Signature> + 'static,
P::Signature: 'static,
P::Digest: Send + 'static,
P::AeadCipher: KeyInit + Send + Sync + 'static,
{
if handshake_bytes.len() > HANDSHAKE_MAX_WIRE {
return Err(TransportError::InvalidMessage);
}
let kind = self.to_handshake_protocol_kind();
let transport_envelope = TransportEnvelope::from_der(handshake_bytes)?;
let raw_message = kind.unwrap_client_request(&transport_envelope)?;
if self.to_server_handshake_mut().is_none() {
let orchestrator = match kind {
HandshakeProtocolKind::Ecies => {
return Err(TransportError::UnsupportedHandshakeProtocol(HandshakeProtocolKind::Ecies));
}
HandshakeProtocolKind::Cms => self.build_cms_server_orchestrator()?,
};
*self.to_server_handshake_mut() = Some(orchestrator);
}
self.drive_server_handshake(kind, &raw_message).await
}
#[cfg(feature = "x509")]
#[allow(async_fn_in_trait)]
async fn perform_emit_cycle(
&mut self,
message: Frame,
) -> TransportResult<(TransitStatus, Option<Frame>, Option<Frame>)>
where
Self: Sized + MessageIO + EncryptedProtocolState,
{
let wire_envelope = self.wrap_and_encrypt_message(message).await?;
let wire_bytes = wire_envelope.to_der()?;
self.write_envelope(&wire_bytes).await?;
let response_bytes = self.read_envelope().await?;
let response_envelope = self.decrypt_response(response_bytes).await?;
let (status, response) = match response_envelope {
TransportEnvelope::Response(pkg) => (pkg.status, pkg.message),
TransportEnvelope::Request(_) => return Err(TransportError::InvalidMessage),
TransportEnvelope::EnvelopedData(_) | TransportEnvelope::SignedData(_) => {
return Err(TransportError::InvalidMessage)
}
};
let returned_message = if status != TransitStatus::Accepted {
match wire_envelope {
WireEnvelope::Cleartext(TransportEnvelope::Request(pkg)) => Some(pkg.message),
_ => None, }
} else {
None
};
let response_frame = response.map(|arc| Arc::try_unwrap(arc).unwrap_or_else(|a| (*a).clone()));
let returned_frame = returned_message.map(|arc| Arc::try_unwrap(arc).unwrap_or_else(|a| (*a).clone()));
Ok((status, response_frame, returned_frame))
}
}
pub trait Pingable {
fn ping(&mut self) -> TransportResult<()>;
}
#[cfg(test)]
mod tests {
use super::*;
use crate::asn1::{MessagePriority, Metadata};
use crate::transport::envelopes::{RequestPackage, ResponsePackage};
use crate::Version;
const PARSE_DER_LENGTH_CASES: &[(u8, &[u8], Option<usize>)] = &[
(0x00, &[], Some(0)),
(0x7F, &[], Some(127)),
(0x81, &[0x80], Some(128)),
(0x81, &[0xFF], Some(255)),
(0x82, &[0x01, 0x00], Some(256)),
(0x80, &[], None),
(0x81, &[0x05], None),
(0x82, &[0x00, 0x80], None),
(0x89, &[0x01; 9], None),
(0x82, &[0x01], None),
(0x81, &[], None),
];
#[test]
fn parse_der_length_cases() {
for &(first, rest, expected) in PARSE_DER_LENGTH_CASES {
assert_eq!(parse_der_length(first, rest), expected);
}
}
fn frame_with_priority(version: Version) -> Frame {
let mut metadata = Metadata::default();
metadata.priority = Some(MessagePriority::Standard);
Frame { version, metadata, message: Vec::new(), integrity: None, nonrepudiation: None }
}
struct DecodeProbe;
impl crate::transport::ResponseHandler for DecodeProbe {
fn with_handler<F>(self, _handler: F) -> Self
where
F: Fn(Frame) -> Option<Frame> + Send + Sync + 'static,
{
self
}
fn handler(&self) -> Option<&(dyn Fn(Frame) -> Option<Frame> + Send + Sync)> {
None
}
}
impl MessageIO for DecodeProbe {
async fn read_envelope(&mut self) -> TransportResult<Vec<u8>> {
Err(TransportError::ConnectionClosed)
}
async fn write_envelope(&mut self, _buffer: &[u8]) -> TransportResult<()> {
Ok(())
}
}
fn version_envelope_cases() -> Vec<(&'static str, TransportEnvelope, bool)> {
vec![
(
"request V0+priority",
TransportEnvelope::Request(RequestPackage::new(frame_with_priority(Version::V0))),
false,
),
(
"request V2+priority",
TransportEnvelope::Request(RequestPackage::new(frame_with_priority(Version::V2))),
true,
),
(
"response V0+priority",
TransportEnvelope::Response(ResponsePackage::new(
TransitStatus::Accepted,
Some(frame_with_priority(Version::V0)),
)),
false,
),
(
"response without frame",
TransportEnvelope::Response(ResponsePackage::new(TransitStatus::Accepted, None)),
true,
),
]
}
#[test]
fn envelope_version_compatibility_and_decode_ingress() {
for (_label, envelope, compatible) in version_envelope_cases() {
assert_eq!(envelope_versions_compatible(&envelope), compatible);
let bytes = crate::encode(&envelope).expect("encode probe envelope");
let decoded = <DecodeProbe as MessageIO>::decode_envelope(&bytes);
assert_eq!(decoded.is_ok(), compatible);
if !compatible {
assert!(matches!(decoded, Err(TransportError::InvalidMessage)));
}
}
}
}