#[cfg(feature = "alloc")]
use alloc::boxed::Box;
use core::ops::DerefMut;
use core::time::Duration;
use core::ffi::{self, c_void};
use core::net::SocketAddr;
use embedded_mbedtls_sys::{
mbedtls_ssl_config, mbedtls_ssl_config_init, mbedtls_ssl_context, mbedtls_ssl_init,
MBEDTLS_ERR_SSL_WANT_READ, MBEDTLS_ERR_SSL_WANT_WRITE,
};
use embedded_nal::UdpClientStack;
use embedded_timers::clock::Clock;
use rand_core::{CryptoRng, RngCore};
use crate::{error::Error, rng::rng_try_fill_bytes_callback_fn, timing, udp};
pub struct SslContext<'a, Net, C: Clock, R: RngCore + CryptoRng> {
config: mbedtls_ssl_config,
net_context: Net,
timer_context: timing::MbedtlsTimer<'a, C>,
csrng: R,
}
impl<'a, U: UdpClientStack, C: Clock, R: RngCore + CryptoRng>
SslContext<'a, udp::UdpContext<U>, C, R>
{
pub fn new_udp_client_side(
net_stack: U,
clock: &'a C,
csrng: R,
server_addr: SocketAddr,
) -> Self {
let mut config = mbedtls_ssl_config::default();
unsafe { mbedtls_ssl_config_init(&mut config) };
let net_context = udp::UdpContext::new(net_stack, server_addr);
let timer_context = timing::MbedtlsTimer::new(clock);
SslContext {
config,
net_context,
timer_context,
csrng,
}
}
}
impl<'a, Net, C: Clock, R: RngCore + CryptoRng> Drop for SslContext<'a, Net, C, R> {
fn drop(&mut self) {
unsafe {
embedded_mbedtls_sys::mbedtls_ssl_config_free(&mut self.config);
}
}
}
pub struct SslConnection<
'a,
Net,
C: Clock + 'a,
R: RngCore + CryptoRng,
CTX: DerefMut<Target = SslContext<'a, Net, C, R>>,
> {
mbedtls_ctx: mbedtls_ssl_context,
ssl_ctx: CTX,
}
impl<'a, 'b: 'a, U: UdpClientStack, C: Clock, R: RngCore + CryptoRng>
SslConnection<'b, udp::UdpContext<U>, C, R, &'a mut SslContext<'b, udp::UdpContext<U>, C, R>>
{
pub fn new_dtls_client(
ssl_context: &'a mut SslContext<'b, udp::UdpContext<U>, C, R>,
preset: Preset,
) -> Result<Self, Error> {
Self::new_generic_dtls_client(ssl_context, preset)
}
}
#[cfg(feature = "alloc")]
impl<'a, U: UdpClientStack, C: Clock, R: RngCore + CryptoRng>
SslConnection<'a, udp::UdpContext<U>, C, R, Box<SslContext<'a, udp::UdpContext<U>, C, R>>>
{
pub fn new_dtls_client_heap_context(
ssl_context: SslContext<'a, udp::UdpContext<U>, C, R>,
preset: Preset,
) -> Result<Self, Error> {
Self::new_generic_dtls_client(Box::new(ssl_context), preset)
}
}
impl<
'a,
U: UdpClientStack,
C: Clock,
R: RngCore + CryptoRng,
CTX: DerefMut<Target = SslContext<'a, udp::UdpContext<U>, C, R>>,
> SslConnection<'a, udp::UdpContext<U>, C, R, CTX>
{
fn new_generic_dtls_client(ssl_context: CTX, preset: Preset) -> Result<Self, Error> {
let mut context = mbedtls_ssl_context::default();
unsafe { mbedtls_ssl_init(&mut context) };
let mut this = SslConnection {
mbedtls_ctx: context,
ssl_ctx: ssl_context,
};
use embedded_mbedtls_sys::MBEDTLS_SSL_IS_CLIENT;
use embedded_mbedtls_sys::MBEDTLS_SSL_TRANSPORT_DATAGRAM;
let ret = unsafe {
embedded_mbedtls_sys::mbedtls_ssl_config_defaults(
&mut this.ssl_ctx.config as *mut mbedtls_ssl_config,
MBEDTLS_SSL_IS_CLIENT as i32,
MBEDTLS_SSL_TRANSPORT_DATAGRAM as i32,
preset.into(),
)
};
if ret < 0 {
return Err(ret.into());
}
unsafe {
embedded_mbedtls_sys::mbedtls_ssl_conf_rng(
&mut this.ssl_ctx.config,
Some(rng_try_fill_bytes_callback_fn::<R>),
&mut this.ssl_ctx.csrng as *mut R as *mut c_void,
);
}
let ret = unsafe {
embedded_mbedtls_sys::mbedtls_ssl_setup(&mut this.mbedtls_ctx, &this.ssl_ctx.config)
};
if ret < 0 {
return Err(ret.into());
}
unsafe {
embedded_mbedtls_sys::mbedtls_ssl_set_bio(
&mut this.mbedtls_ctx,
&mut this.ssl_ctx.net_context as *mut udp::UdpContext<U> as *mut c_void,
Some(crate::udp::udp_send::<U>),
Some(crate::udp::udp_recv::<U>),
None, );
embedded_mbedtls_sys::mbedtls_ssl_set_timer_cb(
&mut this.mbedtls_ctx,
&mut this.ssl_ctx.timer_context as *mut timing::MbedtlsTimer<C> as *mut c_void,
Some(timing::set_timer::<C>),
Some(timing::get_timer::<C>),
);
}
Ok(this)
}
pub fn conf_handshake_timeout(&mut self, min: Duration, max: Duration) {
unsafe {
embedded_mbedtls_sys::mbedtls_ssl_conf_handshake_timeout(
&mut self.ssl_ctx.config,
min.as_millis().try_into().unwrap_or(u32::MAX / 4),
max.as_millis().try_into().unwrap_or(u32::MAX),
);
}
}
}
impl<'a, Net, C, R, CTX> SslConnection<'a, Net, C, R, CTX>
where
C: Clock,
R: RngCore + CryptoRng,
CTX: DerefMut<Target = SslContext<'a, Net, C, R>>,
{
pub fn handshake(&mut self) -> nb::Result<(), Error> {
unsafe {
use embedded_mbedtls_sys::mbedtls_ssl_handshake;
let ret = mbedtls_ssl_handshake(&mut self.mbedtls_ctx);
if matches!(ret, MBEDTLS_ERR_SSL_WANT_READ | MBEDTLS_ERR_SSL_WANT_WRITE) {
return Err(nb::Error::WouldBlock);
}
if ret < 0 {
return Err(nb::Error::Other(ret.into()));
}
}
Ok(())
}
pub fn configure_psk(&mut self, psk: &[u8], psk_identity: &[u8]) -> Result<(), Error> {
unsafe {
let ret = embedded_mbedtls_sys::mbedtls_ssl_conf_psk(
&mut self.ssl_ctx.config,
psk.as_ptr(),
psk.len(),
psk_identity.as_ptr(),
psk_identity.len(),
);
if ret < 0 {
Err(ret.into())
} else {
Ok(())
}
}
}
pub fn write(&mut self, data: &[u8]) -> nb::Result<usize, Error> {
use embedded_mbedtls_sys::{
MBEDTLS_ERR_SSL_ASYNC_IN_PROGRESS, MBEDTLS_ERR_SSL_CRYPTO_IN_PROGRESS,
};
unsafe {
use embedded_mbedtls_sys::mbedtls_ssl_write;
let ret = mbedtls_ssl_write(&mut self.mbedtls_ctx, data.as_ptr(), data.len());
if matches!(
ret,
MBEDTLS_ERR_SSL_WANT_READ
| MBEDTLS_ERR_SSL_WANT_WRITE
| MBEDTLS_ERR_SSL_CRYPTO_IN_PROGRESS
| MBEDTLS_ERR_SSL_ASYNC_IN_PROGRESS
) {
return Err(nb::Error::WouldBlock);
}
if ret < 0 {
return Err(nb::Error::Other(ret.into()));
}
Ok(ret as usize)
}
}
pub fn read(&mut self, buf: &mut [u8]) -> nb::Result<usize, Error> {
unsafe {
use embedded_mbedtls_sys::mbedtls_ssl_read;
use embedded_mbedtls_sys::{
MBEDTLS_ERR_SSL_ASYNC_IN_PROGRESS, MBEDTLS_ERR_SSL_CRYPTO_IN_PROGRESS,
};
let ret = mbedtls_ssl_read(&mut self.mbedtls_ctx, buf.as_mut_ptr(), buf.len());
if matches!(
ret,
MBEDTLS_ERR_SSL_WANT_READ
| MBEDTLS_ERR_SSL_WANT_WRITE
| MBEDTLS_ERR_SSL_ASYNC_IN_PROGRESS
| MBEDTLS_ERR_SSL_CRYPTO_IN_PROGRESS
) {
return Err(nb::Error::WouldBlock);
}
let len = match ret {
x if x >= 0 => x,
e => {
return Err(nb::Error::Other(e.into()));
}
};
Ok(len as usize)
}
}
pub fn close_notify(&mut self) -> nb::Result<(), Error> {
let ret = unsafe { embedded_mbedtls_sys::mbedtls_ssl_close_notify(&mut self.mbedtls_ctx) };
if ret == 0 {
return Ok(());
}
if matches!(ret, MBEDTLS_ERR_SSL_WANT_READ | MBEDTLS_ERR_SSL_WANT_WRITE) {
return Err(nb::Error::WouldBlock);
}
Err(nb::Error::Other(ret.into()))
}
pub fn session_reset(&mut self) -> Result<(), Error> {
let ret = unsafe { embedded_mbedtls_sys::mbedtls_ssl_session_reset(&mut self.mbedtls_ctx) };
if ret < 0 {
Err(ret.into())
} else {
Ok(())
}
}
}
impl<'a, Net, C, R, CTX> Drop for SslConnection<'a, Net, C, R, CTX>
where
C: Clock,
R: RngCore + CryptoRng,
CTX: DerefMut<Target = SslContext<'a, Net, C, R>>,
{
fn drop(&mut self) {
unsafe {
embedded_mbedtls_sys::mbedtls_ssl_free(&mut self.mbedtls_ctx);
}
}
}
#[derive(Debug, Clone, Copy)]
pub enum Preset {
Default,
SuiteB,
}
impl From<Preset> for core::ffi::c_int {
fn from(value: Preset) -> Self {
match value {
Preset::Default => embedded_mbedtls_sys::MBEDTLS_SSL_PRESET_DEFAULT as ffi::c_int,
Preset::SuiteB => embedded_mbedtls_sys::MBEDTLS_SSL_PRESET_SUITEB as ffi::c_int,
}
}
}
#[cfg(test)]
mod test {
use core::net::{IpAddr, Ipv4Addr, SocketAddr};
use embedded_nal::UdpClientStack;
use embedded_timers::clock::Clock;
use rand_core::{CryptoRng, RngCore};
use crate::udp;
use super::{SslConnection, SslContext};
fn _setup_ssl_stack<U: UdpClientStack, R: RngCore + CryptoRng>(
net_stack: U,
clock: &impl Clock,
rng: R,
) {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 22);
let mut ctx = SslContext::new_udp_client_side(net_stack, clock, rng, server_addr);
let _connection = SslConnection::new_dtls_client(&mut ctx, super::Preset::Default).unwrap();
}
type _BoxedDtls<'a, U, C, R> =
SslConnection<'a, udp::UdpContext<U>, C, R, Box<SslContext<'a, udp::UdpContext<U>, C, R>>>;
#[cfg(feature = "alloc")]
fn _setup_ssl_heap<'a, U: UdpClientStack, C: Clock, R: RngCore + CryptoRng>(
net_stack: U,
clock: &'a C,
rng: R,
) -> _BoxedDtls<'a, U, C, R> {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 22);
let ctx = SslContext::new_udp_client_side(net_stack, clock, rng, server_addr);
let connection =
SslConnection::new_dtls_client_heap_context(ctx, super::Preset::Default).unwrap();
connection
}
}