use std::time::Duration;
use miette::Diagnostic;
use thiserror::Error;
use polonius_the_crab::prelude::*;
pub mod serial;
#[cfg(feature = "ble")]
pub mod ble;
pub mod udp;
#[derive(Debug, PartialEq, Clone, Copy)]
struct SmpHeader {
ver: u8,
op: u8,
flags: u8,
data_length: u16,
group_id: u16,
sequence_num: u8,
command_id: u8,
}
impl SmpHeader {
fn from_bytes(data: [u8; SMP_HEADER_SIZE]) -> Self {
Self {
ver: (data[0] >> 3) & 0b11,
op: data[0] & 0b111,
flags: data[1],
data_length: u16::from_be_bytes([data[2], data[3]]),
group_id: u16::from_be_bytes([data[4], data[5]]),
sequence_num: data[6],
command_id: data[7],
}
}
fn to_bytes(self) -> [u8; SMP_HEADER_SIZE] {
let [length_0, length_1] = self.data_length.to_be_bytes();
let [group_id_0, group_id_1] = self.group_id.to_be_bytes();
[
((self.ver & 0b11) << 3) | (self.op & 0b111),
self.flags,
length_0,
length_1,
group_id_0,
group_id_1,
self.sequence_num,
self.command_id,
]
}
}
pub const SMP_HEADER_SIZE: usize = 8;
pub const SMP_BODY_MAX_SIZE: usize = u16::MAX as usize;
pub const SMP_TRANSFER_BUFFER_SIZE: usize = SMP_HEADER_SIZE + SMP_BODY_MAX_SIZE;
mod smp_op {
pub(super) const READ: u8 = 0;
pub(super) const READ_RSP: u8 = 1;
pub(super) const WRITE: u8 = 2;
pub(super) const WRITE_RSP: u8 = 3;
}
#[derive(Error, Debug, Diagnostic)]
pub enum SendError {
#[error("A timeout occurred")]
#[diagnostic(code(mcumgr_toolkit::transport::send::timeout))]
Timeout,
#[error("Transport error")]
#[diagnostic(code(mcumgr_toolkit::transport::send::transport))]
TransportError(#[source] Box<dyn std::error::Error + Sync + Send + 'static>),
#[error("Given data slice was too big")]
#[diagnostic(code(mcumgr_toolkit::transport::send::too_big))]
DataTooBig,
}
#[derive(Error, Debug, Diagnostic)]
pub enum ReceiveError {
#[error("A timeout occurred")]
#[diagnostic(code(mcumgr_toolkit::transport::recv::timeout))]
Timeout,
#[error("Transport error")]
#[diagnostic(code(mcumgr_toolkit::transport::recv::transport))]
TransportError(#[source] Box<dyn std::error::Error + Sync + Send + 'static>),
#[error("Received unexpected response")]
#[diagnostic(code(mcumgr_toolkit::transport::recv::unexpected))]
UnexpectedResponse,
#[error("Received frame that exceeds configured MTU")]
#[diagnostic(code(mcumgr_toolkit::transport::recv::too_big))]
FrameTooBig,
#[error("Failed to decode base64 data")]
#[diagnostic(code(mcumgr_toolkit::transport::recv::base64_decode))]
Base64DecodeError(#[from] base64::DecodeSliceError),
}
impl From<std::io::Error> for SendError {
fn from(e: std::io::Error) -> Self {
if std::io::ErrorKind::TimedOut == e.kind() {
Self::Timeout
} else {
Self::TransportError(e.into())
}
}
}
impl From<std::io::Error> for ReceiveError {
fn from(e: std::io::Error) -> Self {
if std::io::ErrorKind::TimedOut == e.kind() {
Self::Timeout
} else {
Self::TransportError(e.into())
}
}
}
#[cfg(feature = "ble")]
impl From<btleplug::Error> for SendError {
fn from(e: btleplug::Error) -> Self {
if let btleplug::Error::TimedOut(_) = e {
Self::Timeout
} else {
Self::TransportError(e.into())
}
}
}
#[cfg(feature = "ble")]
impl From<btleplug::Error> for ReceiveError {
fn from(e: btleplug::Error) -> Self {
if let btleplug::Error::TimedOut(_) = e {
Self::Timeout
} else {
Self::TransportError(e.into())
}
}
}
pub trait Transport {
fn send_raw_frame(
&mut self,
header: [u8; SMP_HEADER_SIZE],
data: &[u8],
) -> Result<(), SendError>;
fn recv_raw_frame<'a>(
&mut self,
buffer: &'a mut [u8; SMP_TRANSFER_BUFFER_SIZE],
) -> Result<&'a [u8], ReceiveError>;
fn send_frame(
&mut self,
write_operation: bool,
sequence_num: u8,
group_id: u16,
command_id: u8,
data: &[u8],
) -> Result<(), SendError> {
let header = SmpHeader {
ver: 0b01,
op: if write_operation {
smp_op::WRITE
} else {
smp_op::READ
},
flags: 0,
data_length: data.len().try_into().map_err(|_| SendError::DataTooBig)?,
group_id,
sequence_num,
command_id,
};
let header_data = header.to_bytes();
self.send_raw_frame(header_data, data)
}
fn receive_frame<'a>(
&mut self,
mut buffer: &'a mut [u8; SMP_TRANSFER_BUFFER_SIZE],
write_operation: bool,
sequence_num: u8,
group_id: u16,
command_id: u8,
) -> Result<&'a [u8], ReceiveError> {
polonius_loop!(|buffer| -> Result<&'polonius [u8], ReceiveError> {
let frame = polonius_try!(self.recv_raw_frame(buffer));
let (header_data, data) = match frame.split_first_chunk::<SMP_HEADER_SIZE>() {
Some(parts) => parts,
None => polonius_break!(Err(ReceiveError::UnexpectedResponse)),
};
let header = SmpHeader::from_bytes(*header_data);
if header.sequence_num != sequence_num {
polonius_continue!();
}
let expected_op = if write_operation {
smp_op::WRITE_RSP
} else {
smp_op::READ_RSP
};
if (header.group_id != group_id)
|| (header.command_id != command_id)
|| (header.op != expected_op)
|| (usize::from(header.data_length) != data.len())
{
polonius_break!(Err(ReceiveError::UnexpectedResponse));
}
polonius_return!(Ok(data));
})
}
fn set_timeout(
&mut self,
timeout: Duration,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>>;
fn max_smp_frame_size(&self) -> usize {
usize::MAX
}
}
pub trait IntoTransport {
fn into_transport(self) -> Box<dyn Transport + Send>;
}
impl<T> IntoTransport for T
where
T: Transport + Send + 'static,
{
fn into_transport(self) -> Box<dyn Transport + Send> {
Box::new(self)
}
}
impl IntoTransport for Box<dyn Transport + Send> {
fn into_transport(self) -> Box<dyn Transport + Send> {
self
}
}