Skip to main content

mcumgr_toolkit/
connection.rs

1use std::{io::Cursor, sync::Mutex, time::Duration};
2
3use crate::{
4    DEFAULT_RETRIES,
5    commands::{ErrResponse, ErrResponseV2, McuMgrCommand},
6    smp_errors::{DeviceError, MCUmgrErr},
7    transport::{
8        IntoTransport, ReceiveError, SMP_BODY_MAX_SIZE, SMP_TRANSFER_BUFFER_SIZE, SendError,
9        Transport,
10    },
11};
12
13use miette::{Diagnostic, IntoDiagnostic};
14use polonius_the_crab::prelude::*;
15use thiserror::Error;
16
17struct Transceiver {
18    transport: Box<dyn Transport + Send>,
19    next_seqnum: u8,
20    receive_buffer: Box<[u8; SMP_TRANSFER_BUFFER_SIZE]>,
21}
22
23struct Inner {
24    transceiver: Transceiver,
25    send_buffer: Box<[u8; SMP_BODY_MAX_SIZE]>,
26    retries: u8,
27}
28
29/// An SMP protocol layer connection to a device.
30///
31/// In most cases this struct will not be used directly by the user,
32/// but instead it is used indirectly through [`MCUmgrClient`](crate::MCUmgrClient).
33pub struct Connection {
34    inner: Mutex<Inner>,
35}
36
37/// Errors that can happen on SMP protocol level
38#[derive(Error, Debug, Diagnostic)]
39pub enum ExecuteError {
40    /// An error happened on SMP transport level while sending a request
41    #[error("Sending failed")]
42    #[diagnostic(code(mcumgr_toolkit::connection::execute::send))]
43    SendFailed(#[from] SendError),
44    /// An error happened on SMP transport level while receiving a response
45    #[error("Receiving failed")]
46    #[diagnostic(code(mcumgr_toolkit::connection::execute::receive))]
47    ReceiveFailed(#[from] ReceiveError),
48    /// An error happened while CBOR encoding the request payload
49    #[error("CBOR encoding failed")]
50    #[diagnostic(code(mcumgr_toolkit::connection::execute::encode))]
51    EncodeFailed(#[source] Box<dyn miette::Diagnostic + Send + Sync>),
52    /// An error happened while CBOR decoding the response payload
53    #[error("CBOR decoding failed")]
54    #[diagnostic(code(mcumgr_toolkit::connection::execute::decode))]
55    DecodeFailed(#[source] Box<dyn miette::Diagnostic + Send + Sync>),
56    /// The device returned an SMP error
57    #[error("Device returned error code: {0}")]
58    #[diagnostic(code(mcumgr_toolkit::connection::execute::device_error))]
59    ErrorResponse(DeviceError),
60}
61
62impl ExecuteError {
63    /// Checks if the device reported the command as unsupported
64    pub fn command_not_supported(&self) -> bool {
65        if let Self::ErrorResponse(DeviceError::V1 { rc, .. }) = self {
66            *rc == MCUmgrErr::MGMT_ERR_ENOTSUP as i32
67        } else {
68            false
69        }
70    }
71}
72
73impl Transceiver {
74    fn transceive_command(
75        &mut self,
76        write_operation: bool,
77        group_id: u16,
78        command_id: u8,
79        data: &[u8],
80    ) -> Result<&'_ [u8], ExecuteError> {
81        let sequence_num = self.next_seqnum;
82        self.next_seqnum = self.next_seqnum.wrapping_add(1);
83
84        self.transport
85            .send_frame(write_operation, sequence_num, group_id, command_id, data)?;
86
87        self.transport
88            .receive_frame(
89                &mut self.receive_buffer,
90                write_operation,
91                sequence_num,
92                group_id,
93                command_id,
94            )
95            .map_err(Into::into)
96    }
97
98    fn transceive_command_with_retries(
99        &mut self,
100        write_operation: bool,
101        group_id: u16,
102        command_id: u8,
103        data: &[u8],
104        num_retries: u8,
105    ) -> Result<&'_ [u8], ExecuteError> {
106        let mut this = self;
107
108        let mut counter = 0;
109
110        polonius_loop!(|this| -> Result<&'polonius [u8], ExecuteError> {
111            let result = this.transceive_command(write_operation, group_id, command_id, data);
112
113            if counter >= num_retries {
114                polonius_return!(result)
115            }
116            counter += 1;
117
118            match result {
119                Ok(_) => polonius_return!(result),
120                Err(e) => {
121                    let mut lowest_err: &dyn std::error::Error = &e;
122                    while let Some(lower_err) = lowest_err.source() {
123                        lowest_err = lower_err;
124                    }
125                    log::warn!("Retry transmission, error occurred: {lowest_err}");
126                }
127            }
128        })
129    }
130}
131
132impl Connection {
133    /// Creates a new SMP connection
134    pub fn new<T: IntoTransport>(transport: T) -> Self {
135        Self {
136            inner: Mutex::new(Inner {
137                transceiver: Transceiver {
138                    transport: transport.into_transport(),
139                    next_seqnum: rand::random(),
140                    receive_buffer: Box::new([0; SMP_TRANSFER_BUFFER_SIZE]),
141                },
142                send_buffer: Box::new([0; SMP_BODY_MAX_SIZE]),
143                retries: DEFAULT_RETRIES,
144            }),
145        }
146    }
147
148    /// Returns the maximum SMP frame size the underlying transport can
149    /// deliver reliably.
150    pub fn max_transport_frame_size(&self) -> usize {
151        self.inner
152            .lock()
153            .unwrap()
154            .transceiver
155            .transport
156            .max_smp_frame_size()
157    }
158
159    /// Returns how many bytes the device needs in its receive buffer
160    /// in addition to the SMP frame itself.
161    pub fn device_rx_buffer_overhead(&self) -> usize {
162        self.inner
163            .lock()
164            .unwrap()
165            .transceiver
166            .transport
167            .device_rx_buffer_overhead()
168    }
169
170    /// Changes the communication timeout.
171    ///
172    /// When the device does not respond to packets within the set
173    /// duration, an error will be raised.
174    pub fn set_timeout(
175        &self,
176        timeout: Duration,
177    ) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
178        self.inner
179            .lock()
180            .unwrap()
181            .transceiver
182            .transport
183            .set_timeout(timeout)
184    }
185
186    /// Changes the retry amount.
187    ///
188    /// When the device encounters a transport error, it will retry
189    /// this many times until giving up.
190    pub fn set_retries(&self, retries: u8) {
191        self.inner.lock().unwrap().retries = retries;
192    }
193
194    /// Executes a given CBOR based SMP command.
195    pub fn execute_command<R: McuMgrCommand>(
196        &self,
197        request: &R,
198    ) -> Result<R::Response, ExecuteError> {
199        self.execute_command_impl(request, true)
200    }
201
202    /// Executes a given CBOR based SMP command.
203    ///
204    /// Does not use retries.
205    pub fn execute_command_without_retries<R: McuMgrCommand>(
206        &self,
207        request: &R,
208    ) -> Result<R::Response, ExecuteError> {
209        self.execute_command_impl(request, false)
210    }
211
212    fn execute_command_impl<R: McuMgrCommand>(
213        &self,
214        request: &R,
215        use_retries: bool,
216    ) -> Result<R::Response, ExecuteError> {
217        let mut lock_guard = self.inner.lock().unwrap();
218        let locked_self: &mut Inner = &mut lock_guard;
219
220        let mut cursor = Cursor::new(locked_self.send_buffer.as_mut_slice());
221        ciborium::into_writer(request.data(), &mut cursor)
222            .into_diagnostic()
223            .map_err(Into::into)
224            .map_err(ExecuteError::EncodeFailed)?;
225        let data_size = cursor.position() as usize;
226        let data = &locked_self.send_buffer[..data_size];
227
228        log::debug!("TX data: {}", hex::encode(data));
229
230        let write_operation = request.is_write_operation();
231        let group_id = request.group_id();
232        let command_id = request.command_id();
233
234        let response = locked_self.transceiver.transceive_command_with_retries(
235            write_operation,
236            group_id,
237            command_id,
238            data,
239            if use_retries { locked_self.retries } else { 0 },
240        )?;
241
242        log::debug!("RX data: {}", hex::encode(response));
243
244        let err: ErrResponse = ciborium::from_reader(Cursor::new(response))
245            .into_diagnostic()
246            .map_err(Into::into)
247            .map_err(ExecuteError::DecodeFailed)?;
248
249        if let Some(ErrResponseV2 { rc, group }) = err.err {
250            return Err(ExecuteError::ErrorResponse(DeviceError::V2 { group, rc }));
251        }
252
253        if let Some(rc) = err.rc {
254            if rc != MCUmgrErr::MGMT_ERR_EOK as i32 {
255                return Err(ExecuteError::ErrorResponse(DeviceError::V1 {
256                    rc,
257                    rsn: err.rsn,
258                }));
259            }
260        }
261
262        let decoded_response: R::Response = ciborium::from_reader(Cursor::new(response))
263            .into_diagnostic()
264            .map_err(Into::into)
265            .map_err(ExecuteError::DecodeFailed)?;
266
267        Ok(decoded_response)
268    }
269
270    /// Executes a raw SMP command.
271    ///
272    /// Same as [`Connection::execute_command`], but the payload can be anything and must not
273    /// necessarily be CBOR encoded.
274    ///
275    /// Errors are also not decoded but instead will be returned as raw CBOR data.
276    ///
277    /// Read Zephyr's [SMP Protocol Specification](https://docs.zephyrproject.org/latest/services/device_mgmt/smp_protocol.html)
278    /// for more information.
279    pub fn execute_raw_command(
280        &self,
281        write_operation: bool,
282        group_id: u16,
283        command_id: u8,
284        data: &[u8],
285        use_retries: bool,
286    ) -> Result<Box<[u8]>, ExecuteError> {
287        let mut lock_guard = self.inner.lock().unwrap();
288        let locked_self: &mut Inner = &mut lock_guard;
289
290        locked_self
291            .transceiver
292            .transceive_command_with_retries(
293                write_operation,
294                group_id,
295                command_id,
296                data,
297                if use_retries { locked_self.retries } else { 0 },
298            )
299            .map(|val| val.into())
300    }
301}