use std::sync::Arc;
use std::time::Duration;
#[cfg(feature = "tokio")]
mod tokio_rt {
pub use std::time::Instant;
pub use crate::transport::protocols::PersistentConnection;
pub use crate::transport::{AsyncListenerTrait, Protocol};
pub use tokio::io::{AsyncReadExt, AsyncWriteExt};
pub use tokio::net::{TcpListener, TcpStream};
}
#[cfg(feature = "tokio")]
use tokio_rt::*;
use crate::builder::TypeBuilder;
use crate::der::Encode;
use crate::transport::error::TransportFailure;
use crate::transport::ResponsePackage;
use crate::transport::{
EnvelopeBuilder, EnvelopeLimits, MessageIO, Pingable, TransportError, TransportResult, WireMode,
};
use crate::Frame;
#[cfg(feature = "x509")]
mod x509 {
pub use crate::crypto::aead::RuntimeAead;
pub use crate::crypto::profiles::CryptoProvider;
pub use crate::crypto::x509::policy::CertificateValidation;
pub use crate::crypto::x509::store::CertificateTrust;
pub use crate::transport::handshake::{
HandshakeError, HandshakeKeyManager, HandshakeProtocolKind, ServerHandshakeProtocol, TcpHandshakeState,
};
pub use crate::transport::state::EncryptedProtocolState;
pub use crate::transport::{EncryptedMessageIO, TransportEncryptionConfig};
pub use crate::x509::Certificate;
#[cfg(feature = "tokio")]
pub use crate::crypto::profiles::DefaultCryptoProvider;
#[cfg(feature = "tokio")]
pub use crate::transport::EncryptedProtocol;
}
#[cfg(feature = "x509")]
use x509::*;
#[cfg(feature = "transport-policy")]
mod policy {
pub use crate::policy::GatePolicy;
pub use crate::transport::policy::RestartPolicy;
}
#[cfg(feature = "transport-policy")]
use policy::*;
pub use crate::utils::marker::MaybeSend;
pub trait AsyncProtocolStream: MaybeSend + Unpin {
type Error: Into<TransportError>;
fn read_frame(
&mut self,
max_len: Option<usize>,
) -> impl core::future::Future<Output = Result<Vec<u8>, Self::Error>> + MaybeSend;
fn write_frame(&mut self, buffer: &[u8])
-> impl core::future::Future<Output = Result<(), Self::Error>> + MaybeSend;
fn is_alive(&self) -> bool;
}
#[cfg(feature = "tokio")]
pub struct TokioStream {
stream: TcpStream,
}
#[cfg(feature = "tokio")]
impl AsyncProtocolStream for TokioStream {
type Error = std::io::Error;
async fn read_frame(&mut self, max_len: Option<usize>) -> Result<Vec<u8>, Self::Error> {
let stream = &mut self.stream;
let mut tag = [0u8; 1];
stream.read_exact(&mut tag).await?;
let mut length_first = [0u8; 1];
stream.read_exact(&mut length_first).await?;
let (length_octets, content_length) = if length_first[0] & 0x80 == 0 {
(Vec::new(), length_first[0] as usize)
} else {
let octet_count = (length_first[0] & 0x7F) as usize;
let mut length_octets = vec![0u8; octet_count];
stream.read_exact(&mut length_octets).await?;
let length = crate::transport::io::parse_der_length(length_first[0], &length_octets)
.ok_or_else(|| std::io::Error::from(std::io::ErrorKind::InvalidData))?;
(length_octets, length)
};
if let Some(max) = max_len {
if content_length > max {
return Err(std::io::Error::from(std::io::ErrorKind::InvalidData));
}
}
let mut content = vec![0u8; content_length];
stream.read_exact(&mut content).await?;
Ok(crate::transport::io::reconstruct_der_encoding(
tag[0],
length_first[0],
&length_octets,
&content,
))
}
async fn write_frame(&mut self, buffer: &[u8]) -> Result<(), Self::Error> {
self.stream.write_all(buffer).await
}
fn is_alive(&self) -> bool {
self.stream.peer_addr().is_ok()
}
}
#[cfg(feature = "tokio")]
impl From<TcpStream> for TokioStream {
fn from(stream: TcpStream) -> Self {
Self { stream }
}
}
#[cfg(feature = "tokio")]
pub struct TokioListener<P: CryptoProvider = DefaultCryptoProvider> {
listener: TcpListener,
#[cfg(feature = "x509")]
certificate: Option<Arc<Certificate>>,
#[cfg(feature = "x509")]
client_validators: Option<Arc<Vec<Arc<dyn CertificateValidation>>>>,
#[cfg(feature = "x509")]
aad_domain_tag: Option<&'static [u8]>,
#[cfg(feature = "x509")]
max_cleartext_envelope: Option<usize>,
#[cfg(feature = "x509")]
max_encrypted_envelope: Option<usize>,
#[cfg(feature = "x509")]
handshake_timeout: Option<Duration>,
#[cfg(feature = "x509")]
key_manager: Option<Arc<HandshakeKeyManager<P>>>,
}
#[cfg(feature = "tokio")]
impl<P: CryptoProvider> TokioListener<P> {
pub fn local_addr(&self) -> std::io::Result<std::net::SocketAddr> {
self.listener.local_addr()
}
pub async fn bind(addr: &str) -> std::io::Result<Self> {
let listener = TcpListener::bind(addr).await?;
Ok(Self {
listener,
#[cfg(feature = "x509")]
certificate: None,
#[cfg(feature = "x509")]
client_validators: None,
#[cfg(feature = "x509")]
aad_domain_tag: None,
#[cfg(feature = "x509")]
max_cleartext_envelope: None,
#[cfg(feature = "x509")]
max_encrypted_envelope: None,
#[cfg(feature = "x509")]
handshake_timeout: None,
#[cfg(feature = "x509")]
key_manager: None,
})
}
#[cfg(not(feature = "x509"))]
pub async fn accept(&self) -> std::io::Result<(TokioStream, std::net::SocketAddr)> {
let (stream, addr) = self.listener.accept().await?;
Ok((TokioStream::from(stream), addr))
}
#[cfg(feature = "x509")]
pub async fn accept(&self) -> std::io::Result<(TcpTransport<TokioStream, P>, std::net::SocketAddr)> {
let (stream, addr) = self.listener.accept().await?;
let mut transport = TcpTransport::from(TokioStream::from(stream));
if let Some(cert) = &self.certificate {
transport.server_identity = Some(Arc::clone(cert));
}
if let Some(ref validators) = self.client_validators {
transport.client_validators = Some(Arc::clone(validators));
}
if let Some(aad) = self.aad_domain_tag {
transport.aad_domain_tag = Some(aad);
}
if let Some(max) = self.max_cleartext_envelope {
transport.max_cleartext_envelope = Some(max);
}
if let Some(max) = self.max_encrypted_envelope {
transport.max_encrypted_envelope = Some(max);
}
#[cfg(feature = "x509")]
if let Some(timeout) = self.handshake_timeout {
transport.handshake_timeout = timeout;
}
#[cfg(feature = "x509")]
if let Some(signatory) = &self.key_manager {
transport.key_manager = Some(Arc::clone(signatory));
}
Ok((transport, addr))
}
}
#[cfg(feature = "tokio")]
impl<P: CryptoProvider + Send + Sync> Protocol for TokioListener<P> {
type Listener = TokioListener<P>;
type Stream = TokioStream;
type Error = std::io::Error;
type Transport = TcpTransport<TokioStream, P>;
type Address = crate::transport::tcp::TightBeamSocketAddr;
fn default_bind_address() -> Result<Self::Address, Self::Error> {
"127.0.0.1:0"
.parse()
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidInput, e))
}
async fn bind(addr: Self::Address) -> Result<(Self::Listener, Self::Address), Self::Error> {
let listener = TcpListener::bind(addr.0).await?;
let bound_addr = listener.local_addr()?;
Ok((
Self {
listener,
#[cfg(feature = "x509")]
certificate: None,
#[cfg(feature = "x509")]
client_validators: None,
#[cfg(feature = "x509")]
aad_domain_tag: None,
#[cfg(feature = "x509")]
max_cleartext_envelope: None,
#[cfg(feature = "x509")]
max_encrypted_envelope: None,
#[cfg(feature = "x509")]
handshake_timeout: None,
#[cfg(feature = "x509")]
key_manager: None,
},
crate::transport::tcp::TightBeamSocketAddr(bound_addr),
))
}
async fn connect(addr: Self::Address) -> Result<Self::Stream, Self::Error> {
let stream = TcpStream::connect(addr.0).await?;
Ok(TokioStream::from(stream))
}
fn create_transport(stream: Self::Stream) -> Self::Transport {
TcpTransport::from(stream)
}
fn to_tightbeam_addr(&self) -> Result<Self::Address, Self::Error> {
Ok(crate::transport::tcp::TightBeamSocketAddr(self.local_addr()?))
}
}
#[cfg(all(feature = "tokio", feature = "x509"))]
impl<P: CryptoProvider + Send + Sync> EncryptedProtocol for TokioListener<P> {
type Encryptor = RuntimeAead;
type Decryptor = RuntimeAead;
type CryptoProvider = P;
async fn bind_with(
addr: Self::Address,
config: TransportEncryptionConfig<P>,
) -> Result<(Self::Listener, Self::Address), Self::Error> {
let listener = TcpListener::bind(addr.0).await?;
let bound_addr = listener.local_addr()?;
let certificate = Arc::new(config.certificate);
let client_validators = config.client_validators.as_ref().map(Arc::clone);
let key_manager = Arc::clone(&config.key_manager);
Ok((
Self {
listener,
certificate: Some(certificate),
client_validators,
aad_domain_tag: Some(config.aad_domain_tag),
max_cleartext_envelope: Some(config.max_cleartext_envelope),
max_encrypted_envelope: Some(config.max_encrypted_envelope),
handshake_timeout: Some(config.handshake_timeout),
key_manager: Some(key_manager),
},
crate::transport::tcp::TightBeamSocketAddr(bound_addr),
))
}
}
#[cfg(feature = "x509")]
impl<S: AsyncProtocolStream, P: CryptoProvider + Send + Sync + 'static> EncryptedProtocolState for TcpTransport<S, P>
where
TransportError: From<S::Error>,
{
type CryptoProvider = P;
fn to_encryptor_ref(&self) -> TransportResult<&RuntimeAead> {
self.symmetric_key
.as_ref()
.ok_or(TransportError::OperationFailed(TransportFailure::EncryptorUnavailable))
}
fn to_decryptor_ref(&self) -> TransportResult<&RuntimeAead> {
self.symmetric_key
.as_ref()
.ok_or(TransportError::OperationFailed(TransportFailure::EncryptorUnavailable))
}
fn to_handshake_state(&self) -> TcpHandshakeState {
self.handshake_state
}
fn set_handshake_state(&mut self, state: TcpHandshakeState) {
self.handshake_state = state;
}
fn to_server_certificate_ref(&self) -> Option<&Certificate> {
self.server_identity.as_ref().map(|arc| arc.as_ref())
}
fn to_server_certificate_arc(&self) -> Option<Arc<Certificate>> {
self.server_identity.as_ref().map(Arc::clone)
}
fn set_symmetric_key(&mut self, key: RuntimeAead) {
let _ = self.symmetric_key.take();
self.symmetric_key = Some(key);
}
fn to_max_cleartext_envelope(&self) -> Option<usize> {
self.max_cleartext_envelope
}
fn to_max_encrypted_envelope(&self) -> Option<usize> {
self.max_encrypted_envelope
}
fn is_client_validators_present(&self) -> bool {
self.client_validators.is_some()
}
fn to_handshake_protocol_kind(&self) -> HandshakeProtocolKind {
self.handshake_protocol_kind
}
fn to_key_manager_ref(&self) -> Option<&Arc<HandshakeKeyManager<P>>> {
self.key_manager.as_ref()
}
fn to_client_certificate_ref(&self) -> Option<&Arc<Certificate>> {
self.client_certificate.as_ref()
}
fn to_trust_store_ref(&self) -> Option<&Arc<dyn CertificateTrust>> {
self.trust_store.as_ref()
}
fn to_server_certificate_chain_ref(&self) -> Option<&Arc<[Certificate]>> {
self.server_certificate_chain.as_ref()
}
fn to_server_handshake_mut(
&mut self,
) -> &mut Option<Box<dyn ServerHandshakeProtocol<Error = HandshakeError> + Send + Sync>> {
&mut self.server_handshake
}
fn set_peer_certificate(&mut self, cert: Certificate) {
self.peer_certificate = Some(cert);
}
fn to_handshake_timeout(&self) -> Duration {
self.handshake_timeout
}
fn to_client_validators_ref(&self) -> Option<&Arc<Vec<Arc<dyn CertificateValidation>>>> {
self.client_validators.as_ref()
}
fn unset_symmetric_key(&mut self) {
self.symmetric_key = None;
}
}
#[cfg(feature = "x509")]
impl<S: AsyncProtocolStream> EncryptedMessageIO for TcpTransport<S> where TransportError: From<S::Error> {}
#[cfg(feature = "x509")]
impl<S: AsyncProtocolStream, P: CryptoProvider + Send + Sync> TcpTransport<S, P>
where
TransportError: From<S::Error>,
{
pub fn with_server_encryption(mut self, config: TransportEncryptionConfig<P>) -> Self {
let certificate = Arc::new(config.certificate);
self.server_identity = Some(certificate);
self.client_validators = config.client_validators;
self.aad_domain_tag = Some(config.aad_domain_tag);
self.max_cleartext_envelope = Some(config.max_cleartext_envelope);
self.max_encrypted_envelope = Some(config.max_encrypted_envelope);
self.handshake_timeout = config.handshake_timeout;
self.key_manager = Some(config.key_manager);
self
}
}
#[cfg(feature = "tokio")]
impl<P: CryptoProvider + Send + Sync> AsyncListenerTrait for TokioListener<P> {
async fn accept(&self) -> Result<(Self::Transport, Self::Address), Self::Error> {
let (stream, addr) = self.listener.accept().await?;
let mut transport = Self::create_transport(TokioStream::from(stream));
#[cfg(feature = "x509")]
if let Some(ref cert) = self.certificate {
transport.server_identity = Some(Arc::clone(cert));
}
#[cfg(feature = "x509")]
if let Some(ref signatory) = self.key_manager {
transport.key_manager = Some(Arc::clone(signatory));
}
#[cfg(feature = "x509")]
if let Some(timeout) = self.handshake_timeout {
transport.handshake_timeout = timeout;
}
Ok((transport, crate::transport::tcp::TightBeamSocketAddr(addr)))
}
}
#[cfg(feature = "tokio")]
impl<P: CryptoProvider + Send + Sync> crate::transport::Mycelial for TokioListener<P> {
async fn try_available_connect(&self) -> Result<(Self::Listener, Self::Address), Self::Error> {
let addr = "0.0.0.0:0"
.parse::<crate::transport::tcp::TightBeamSocketAddr>()
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidInput, e))?;
<TokioListener<P> as Protocol>::bind(addr).await
}
}
impl<S: AsyncProtocolStream> Pingable for TcpTransport<S>
where
TransportError: From<S::Error>,
TransportError: From<std::io::Error>,
{
fn ping(&mut self) -> TransportResult<()> {
if self.stream.is_alive() {
Ok(())
} else {
Err(TransportError::ConnectionClosed)
}
}
}
crate::impl_tcp_common!(TcpTransport, AsyncProtocolStream);
impl<S: AsyncProtocolStream> MessageIO for TcpTransport<S>
where
TransportError: From<S::Error>,
{
async fn read_envelope(&mut self) -> TransportResult<Vec<u8>> {
#[cfg(feature = "x509")]
let max_len = if self.is_handshake_pending() {
Some(crate::transport::tcp::HANDSHAKE_MAX_WIRE)
} else {
Some(
self.max_encrypted_envelope
.or(self.max_cleartext_envelope)
.unwrap_or(512 * 1024),
)
};
#[cfg(not(feature = "x509"))]
let max_len = None;
#[cfg(feature = "tokio")]
{
use tokio::time::timeout;
#[cfg(feature = "x509")]
let timeout_duration: Option<Duration> = {
match self.to_handshake_state() {
TcpHandshakeState::AwaitingServerResponse { initiated_at }
| TcpHandshakeState::AwaitingClientFinish { initiated_at } => {
let now = Instant::now();
let deadline = initiated_at + self.handshake_timeout;
if now >= deadline {
return Err(TransportError::OperationFailed(TransportFailure::Timeout));
}
Some(deadline.saturating_duration_since(now))
}
_ if self.is_handshake_pending() => Some(self.handshake_timeout),
_ => {
#[cfg(feature = "transport-policy")]
{
self.operation_timeout
}
#[cfg(not(feature = "transport-policy"))]
{
None
}
}
}
};
#[cfg(not(feature = "x509"))]
let timeout_duration: Option<Duration> = {
#[cfg(feature = "transport-policy")]
{
self.operation_timeout
}
#[cfg(not(feature = "transport-policy"))]
{
None
}
};
let buffer = if let Some(dur) = timeout_duration {
timeout(dur, self.stream.read_frame(max_len)).await??
} else {
self.stream.read_frame(max_len).await?
};
Ok(buffer)
}
#[cfg(not(feature = "tokio"))]
{
let buffer = self.stream.read_frame(max_len).await?;
Ok(buffer)
}
}
async fn write_envelope(&mut self, buffer: &[u8]) -> TransportResult<()> {
#[cfg(all(feature = "tokio", feature = "transport-policy"))]
if let Some(dur) = self.operation_timeout {
tokio::time::timeout(dur, self.stream.write_frame(buffer)).await??;
} else {
self.stream.write_frame(buffer).await?;
}
#[cfg(not(all(feature = "tokio", feature = "transport-policy")))]
self.stream.write_frame(buffer).await?;
Ok(())
}
}
#[cfg(all(feature = "x509", feature = "transport-policy"))]
impl<S: AsyncProtocolStream> crate::transport::MessageCollector for TcpTransport<S>
where
TransportError: From<S::Error>,
{
type CollectorGate = dyn crate::policy::GatePolicy;
fn collector_gate(&self) -> &Self::CollectorGate {
self.collector_gate.as_ref()
}
async fn collect_message(&mut self) -> TransportResult<(Arc<Frame>, crate::policy::TransitStatus)> {
self.collect_message_with_encryption().await
}
async fn send_response(
&mut self,
status: crate::policy::TransitStatus,
message: Option<Frame>,
) -> TransportResult<()> {
let response_pkg = ResponsePackage { status, message: message.map(Arc::new) };
let limits = EnvelopeLimits::from_pair(self.max_cleartext_envelope, self.max_encrypted_envelope);
let mut builder = limits.apply(EnvelopeBuilder::response(response_pkg));
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);
}
let wire_envelope = builder.build()?;
let wire_bytes = wire_envelope.to_der()?;
self.write_envelope(&wire_bytes).await?;
Ok(())
}
}
#[cfg(all(feature = "x509", feature = "transport-policy"))]
impl<S: AsyncProtocolStream> crate::transport::MessageEmitter for TcpTransport<S>
where
TransportError: From<S::Error>,
{
type EmitterGate = dyn crate::policy::GatePolicy;
type RestartPolicy = dyn crate::transport::policy::RestartPolicy;
fn to_restart_policy_ref(&self) -> &Self::RestartPolicy {
self.restart_policy.as_ref()
}
fn to_emitter_gate_policy_ref(&self) -> &Self::EmitterGate {
self.emitter_gate.as_ref()
}
async fn perform_send_receive(
&mut self,
message: Frame,
) -> TransportResult<(crate::policy::TransitStatus, Option<Frame>, Option<Frame>)> {
self.ensure_handshake_complete().await?;
#[cfg(feature = "tokio")]
{
let timeout_duration = self.operation_timeout;
if let Some(duration) = timeout_duration {
use tokio::time::timeout;
match timeout(duration, async { self.perform_emit_cycle(message).await }).await {
Ok(result) => result,
Err(_) => Err(TransportError::OperationFailed(TransportFailure::Timeout)),
}
} else {
self.perform_emit_cycle(message).await
}
}
#[cfg(not(feature = "tokio"))]
{
self.perform_emit_cycle(message).await
}
}
}
#[cfg(feature = "tokio")]
impl<P: CryptoProvider + Send + Sync> PersistentConnection for TokioListener<P> {
fn is_connected(transport: &Self::Transport) -> bool {
transport.stream.is_alive()
}
fn try_close(_transport: &mut Self::Transport) {
}
}
#[cfg(all(test, feature = "tokio"))]
mod tests {
use super::*;
use crate::crypto::sign::ecdsa::Secp256k1VerifyingKey;
use crate::crypto::sign::Sha3Signer;
use crate::testing::*;
use crate::transport::{MessageCollector, MessageEmitter, ResponseHandler};
#[cfg(feature = "x509")]
#[tokio::test]
async fn async_round_trip() -> TransportResult<()> {
let listener = TokioListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let test_message = create_v0_tightbeam(None, None);
let expected_response = create_v0_tightbeam(None, None);
let (tx, mut rx) = tokio::sync::mpsc::channel(1);
let response_msg = expected_response.clone();
let server = listener;
let server_handle = tokio::spawn(async move {
let (transport, _) = server.accept().await?;
let handler = Box::new(move |msg: Frame| {
let _ = tx.try_send(msg);
Some(response_msg.clone())
});
let mut transport = transport.with_handler(handler);
transport.handle_request().await
});
let stream = TcpStream::connect(addr).await?;
let mut transport = TcpTransport::from(TokioStream::from(stream));
let response = transport.emit(test_message.clone(), None).await?;
let received = rx.recv().await;
assert_eq!(Some(test_message), received);
assert_eq!(response.clone(), Some(expected_response));
server_handle.await??;
Ok(())
}
#[cfg(all(feature = "x509", feature = "transport-cms"))]
fn cms_test_client(stream: TcpStream) -> TcpTransport<TokioStream> {
use std::sync::Arc;
use crate::crypto::key::{Secp256k1KeyProvider, SigningKeyProvider};
use crate::crypto::sign::ecdsa::Secp256k1SigningKey;
use crate::transport::handshake::{HandshakeKeyManager, HandshakeProtocolKind};
let signing_key = Secp256k1SigningKey::from(create_test_signing_key());
let provider: Arc<dyn SigningKeyProvider> = Arc::new(Secp256k1KeyProvider::from(signing_key));
let mut transport = TcpTransport::from(TokioStream::from(stream));
transport.handshake_protocol_kind = HandshakeProtocolKind::Cms;
transport.key_manager = Some(Arc::new(HandshakeKeyManager::new(provider)));
transport
}
#[cfg(all(feature = "x509", feature = "transport-cms"))]
#[tokio::test]
async fn cms_client_without_trust_store_fails_closed() -> TransportResult<()> {
use crate::transport::handshake::HandshakeError;
use crate::transport::io::EncryptedMessageIO;
let listener: TokioListener = TokioListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let stream = TcpStream::connect(addr).await?;
let mut transport = cms_test_client(stream);
let result = transport.perform_client_handshake().await;
assert!(matches!(
result,
Err(TransportError::HandshakeError(HandshakeError::MissingTrustStore))
));
Ok(())
}
#[cfg(all(feature = "x509", feature = "transport-cms"))]
#[tokio::test]
async fn cms_client_without_server_chain_fails_closed() -> TransportResult<()> {
use std::sync::Arc;
use crate::crypto::hash::Sha3_256;
use crate::crypto::policy::Secp256k1Policy;
use crate::crypto::x509::store::{CertificateTrust, CertificateTrustBuilder, TrustBuilder};
use crate::transport::io::EncryptedMessageIO;
use crate::transport::X509ClientConfig;
let listener: TokioListener = TokioListener::bind("127.0.0.1:0").await?;
let addr = listener.local_addr()?;
let stream = TcpStream::connect(addr).await?;
let trust_store: Arc<dyn CertificateTrust> =
Arc::new(CertificateTrustBuilder::<Sha3_256>::from(Secp256k1Policy).build());
let mut transport = cms_test_client(stream).with_trust_store(trust_store);
let result = transport.perform_client_handshake().await;
assert!(matches!(result, Err(TransportError::MissingServerCertificateChain)));
Ok(())
}
#[cfg(all(feature = "transport-cms", feature = "transport-policy"))]
#[tokio::test]
async fn async_cms_round_trip() -> TransportResult<()> {
use core::str::FromStr;
use std::sync::Arc;
use crate::crypto::hash::Sha3_256;
use crate::crypto::key::{Secp256k1KeyProvider, SigningKeyProvider};
use crate::crypto::policy::Secp256k1Policy;
use crate::crypto::sign::ecdsa::{Secp256k1SigningKey, SigningKey};
use crate::crypto::x509::store::{CertificateTrust, CertificateTrustBuilder, TrustBuilder};
use crate::prelude::TightBeamSocketAddr;
use crate::spki::SubjectPublicKeyInfoOwned;
use crate::transport::handshake::{HandshakeKeyManager, HandshakeProtocolKind};
use crate::transport::{TransportEncryptionConfig, X509ClientConfig};
let signing_key = create_test_signing_key();
let verifying_key = Secp256k1VerifyingKey::from(&signing_key);
let sha3_signer = Sha3Signer::from(&signing_key);
let spki = SubjectPublicKeyInfoOwned::from_key(verifying_key)?;
let server_cert = crate::cert!(
profile: Root,
subject: "CN=Test Root CA,O=Test Org,C=US",
serial: 1u32,
duration: Duration::from_secs(365 * 24 * 60 * 60),
signer: &sha3_signer,
subject_public_key: spki
)?;
let addr = TightBeamSocketAddr::from_str("127.0.0.1:0")?;
let config = TransportEncryptionConfig::new(server_cert.clone(), signing_key.into());
let (listener, socket_addr) = TokioListener::bind_with(addr, config).await?;
let test_message = create_v0_tightbeam(None, None);
let expected_response = create_v0_tightbeam(None, None);
let (tx, mut rx) = tokio::sync::mpsc::channel(1);
let response_msg = expected_response.clone();
let server_handle = tokio::spawn(async move {
let (mut transport, _) = listener.accept().await?;
transport.handshake_protocol_kind = HandshakeProtocolKind::Cms;
let handler = Box::new(move |msg: Frame| {
let _ = tx.try_send(msg);
Some(response_msg.clone())
});
let mut transport = transport.with_handler(handler);
transport.handle_request().await
});
let client_key = SigningKey::from_bytes(&[2u8; 32].into()).map_err(|_| TransportError::InvalidState)?;
let client_cert = create_test_certificate(&client_key);
let signing_key = Secp256k1SigningKey::from(client_key);
let key_provider = Secp256k1KeyProvider::from(signing_key);
let client_provider: Arc<dyn SigningKeyProvider> = Arc::new(key_provider);
let trust_store: Arc<dyn CertificateTrust> = {
let certificate = server_cert.clone();
Arc::new(
CertificateTrustBuilder::<Sha3_256>::from(Secp256k1Policy)
.with_certificate(certificate)?
.build(),
)
};
let cert = Arc::new(client_cert);
let key = Arc::new(HandshakeKeyManager::new(client_provider));
let chain = Arc::from(vec![server_cert]);
let stream = TcpStream::connect(*socket_addr).await?;
let mut transport = TcpTransport::from(TokioStream::from(stream));
transport = transport.with_trust_store(trust_store);
transport = transport.with_client_identity(cert, key);
transport = transport.with_server_certificate_chain(chain);
transport = transport.with_handshake_protocol(HandshakeProtocolKind::Cms);
let response = transport.emit(test_message.clone(), None).await?;
let received = rx.recv().await;
assert_eq!(Some(test_message), received);
assert_eq!(response, Some(expected_response));
server_handle.await??;
Ok(())
}
#[cfg(all(feature = "x509", feature = "transport-policy"))]
fn encrypted_test_config() -> TransportResult<TransportEncryptionConfig<DefaultCryptoProvider>> {
use crate::spki::SubjectPublicKeyInfoOwned;
let signing_key = create_test_signing_key();
let verifying_key = Secp256k1VerifyingKey::from(&signing_key);
let sha3_signer = Sha3Signer::from(&signing_key);
let spki = SubjectPublicKeyInfoOwned::from_key(verifying_key)?;
let cert = crate::cert!(
profile: Root,
subject: "CN=Test Root CA,O=Test Org,C=US",
serial: 1u32,
duration: Duration::from_secs(365 * 24 * 60 * 60),
signer: &sha3_signer,
subject_public_key: spki
)?;
Ok(TransportEncryptionConfig::new(cert, signing_key.into()))
}
#[cfg(all(feature = "x509", feature = "transport-policy"))]
#[tokio::test]
async fn handshake_read_deadline_bounds_silent_client() -> TransportResult<()> {
use core::str::FromStr;
use crate::prelude::TightBeamSocketAddr;
let mut config = encrypted_test_config()?;
config.handshake_timeout = Duration::from_millis(500);
let addr = TightBeamSocketAddr::from_str("127.0.0.1:0")?;
let (listener, socket_addr) = TokioListener::bind_with(addr, config).await?;
let server_handle = tokio::spawn(async move {
let (mut transport, _) = listener.accept().await?;
transport.handle_request().await
});
let _silent_client = TcpStream::connect(*socket_addr).await?;
let joined = tokio::time::timeout(Duration::from_secs(5), server_handle).await;
assert!(matches!(joined, Ok(Ok(Err(_)))));
Ok(())
}
#[cfg(all(feature = "x509", feature = "transport-policy"))]
#[tokio::test]
async fn handshake_read_rejects_oversize_frame_before_body() -> TransportResult<()> {
use core::str::FromStr;
use crate::prelude::TightBeamSocketAddr;
let mut config = encrypted_test_config()?;
config.handshake_timeout = Duration::from_secs(5);
let addr = TightBeamSocketAddr::from_str("127.0.0.1:0")?;
let (listener, socket_addr) = TokioListener::bind_with(addr, config).await?;
let server_handle = tokio::spawn(async move {
let (mut transport, _) = listener.accept().await?;
transport.handle_request().await
});
let mut stream = TcpStream::connect(*socket_addr).await?;
stream.write_all(&[0x30, 0x83, 0x01, 0x00, 0x00]).await?;
let started = std::time::Instant::now();
let joined = tokio::time::timeout(Duration::from_secs(4), server_handle).await;
assert!(matches!(joined, Ok(Ok(Err(_)))));
assert!(started.elapsed() < Duration::from_secs(2));
Ok(())
}
#[cfg(all(feature = "transport-policy", feature = "transport-ecies"))]
#[tokio::test]
async fn async_with_encrypted_and_gate_policy() -> TransportResult<()> {
use core::str::FromStr;
use core::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use crate::crypto::hash::Sha3_256;
use crate::crypto::policy::Secp256k1Policy;
use crate::crypto::x509::store::{CertificateTrust, CertificateTrustBuilder, TrustBuilder};
use crate::policy::TransitStatus;
use crate::spki::SubjectPublicKeyInfoOwned;
use crate::transport::TransportEncryptionConfig;
use crate::transport::X509ClientConfig;
use crate::{prelude::TightBeamSocketAddr, transport::policy::PolicyConf};
struct BusyFirstGate {
first: AtomicBool,
}
impl BusyFirstGate {
fn new() -> Self {
Self { first: AtomicBool::new(true) }
}
}
impl GatePolicy for BusyFirstGate {
fn evaluate(&self, _msg: &Frame) -> TransitStatus {
if self.first.swap(false, Ordering::SeqCst) {
TransitStatus::Busy
} else {
TransitStatus::Accepted
}
}
}
let signing_key = create_test_signing_key();
let verifying_key = Secp256k1VerifyingKey::from(&signing_key);
let sha3_signer = Sha3Signer::from(&signing_key);
let spki = SubjectPublicKeyInfoOwned::from_key(verifying_key)?;
let cert = crate::cert!(
profile: Root,
subject: "CN=Test Root CA,O=Test Org,C=US",
serial: 1u32,
duration: Duration::from_secs(365 * 24 * 60 * 60),
signer: &sha3_signer,
subject_public_key: spki
)?;
let addr = TightBeamSocketAddr::from_str("127.0.0.1:0")?;
let config = TransportEncryptionConfig::new(cert.clone(), signing_key.clone().into());
let (listener, socket_addr) = TokioListener::bind_with(addr, config).await?;
let server = listener;
let test_message = create_v0_tightbeam(None, None);
let (tx, mut rx) = tokio::sync::mpsc::channel(2);
let server_handle = tokio::spawn(async move {
let (transport, _) = server.accept().await?;
let handler = Box::new(move |msg: Frame| {
let _ = tx.try_send(msg.clone());
Some(msg)
});
let mut transport = transport.with_collector_gate(BusyFirstGate::new()).with_handler(handler);
let result = transport.handle_request().await;
result?;
transport.handle_request().await
});
let stream = TcpStream::connect(*socket_addr).await?;
let trust_store: Arc<dyn CertificateTrust> = Arc::new(
CertificateTrustBuilder::<Sha3_256>::from(Secp256k1Policy)
.with_certificate(cert)?
.build(),
);
let mut transport = TcpTransport::from(TokioStream::from(stream)).with_trust_store(trust_store);
let first = transport.emit(test_message.clone(), None).await;
assert!(matches!(first, Err(TransportError::OperationFailed(TransportFailure::Busy))));
transport.emit(test_message.clone(), None).await?;
let received = rx.recv().await;
assert_eq!(Some(test_message), received);
assert!(rx.try_recv().is_err());
server_handle.await??;
Ok(())
}
}