use crate::crypto::Crypto;
use crate::error::{Error, ErrorCode};
use crate::fmt::Bytes;
use crate::sc::{check_opcode, OpCode};
use crate::transport::exchange::Exchange;
use crate::transport::session::SessionMode;
use crate::utils::storage::WriteBuf;
pub const MCSP_CHALLENGE_LEN: usize = 8;
pub const MCSP_SYNC_REQ_LEN: usize = MCSP_CHALLENGE_LEN;
pub const MCSP_SYNC_RSP_LEN: usize = 4 + MCSP_CHALLENGE_LEN;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub struct MsgCounterSyncReq {
pub challenge: [u8; MCSP_CHALLENGE_LEN],
}
impl MsgCounterSyncReq {
pub fn read(payload: &[u8]) -> Result<Self, Error> {
if payload.len() != MCSP_SYNC_REQ_LEN {
return Err(ErrorCode::InvalidData.into());
}
let mut challenge = [0u8; MCSP_CHALLENGE_LEN];
challenge.copy_from_slice(&payload[..MCSP_CHALLENGE_LEN]);
Ok(Self { challenge })
}
pub fn write(&self, wb: &mut WriteBuf<'_>) -> Result<(), Error> {
wb.copy_from_slice(&self.challenge)?;
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub struct MsgCounterSyncRsp {
pub synchronized_counter: u32,
pub response: [u8; MCSP_CHALLENGE_LEN],
}
impl MsgCounterSyncRsp {
pub fn read(payload: &[u8]) -> Result<Self, Error> {
if payload.len() != MCSP_SYNC_RSP_LEN {
return Err(ErrorCode::InvalidData.into());
}
let synchronized_counter =
u32::from_le_bytes([payload[0], payload[1], payload[2], payload[3]]);
let mut response = [0u8; MCSP_CHALLENGE_LEN];
response.copy_from_slice(&payload[4..4 + MCSP_CHALLENGE_LEN]);
Ok(Self {
synchronized_counter,
response,
})
}
pub fn write(&self, wb: &mut WriteBuf<'_>) -> Result<(), Error> {
wb.le_u32(self.synchronized_counter)?;
wb.copy_from_slice(&self.response)?;
Ok(())
}
}
pub async fn respond<C: Crypto>(crypto: C, mut exchange: Exchange<'_>) -> Result<(), Error> {
check_opcode(&exchange, OpCode::MsgCounterSyncReq)?;
let session_mode = exchange.with_state(|state| {
Ok(exchange
.id()
.session(&mut state.sessions)
.get_session_mode()
.clone())
})?;
if !matches!(session_mode, SessionMode::Group { .. }) {
error!("MCSP: MsgCounterSyncReq received on non-group session; dropping");
return Err(ErrorCode::Invalid.into());
}
let req = {
let rx = exchange.recv().await?;
MsgCounterSyncReq::read(rx.payload())?
};
debug!(
"MCSP: Received MsgCounterSyncReq (challenge={})",
Bytes(&req.challenge)
);
let synchronized_counter =
exchange.with_state(|state| state.sessions.get_or_init_global_group_data_ctr(crypto))?;
let rsp = MsgCounterSyncRsp {
synchronized_counter,
response: req.challenge,
};
debug!(
"MCSP: Sending MsgCounterSyncRsp (sync_ctr={}, response={})",
synchronized_counter,
Bytes(&rsp.response)
);
exchange
.send_with(|_, wb| {
rsp.write(wb)?;
Ok(Some(OpCode::MsgCounterSyncResp.meta()))
})
.await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn req_roundtrip() {
let req = MsgCounterSyncReq {
challenge: [1, 2, 3, 4, 5, 6, 7, 8],
};
let mut buf = [0u8; 32];
let mut wb = WriteBuf::new(&mut buf);
req.write(&mut wb).unwrap();
let slice = wb.as_slice();
assert_eq!(slice.len(), MCSP_SYNC_REQ_LEN);
assert_eq!(slice, &[1, 2, 3, 4, 5, 6, 7, 8]);
let decoded = MsgCounterSyncReq::read(slice).unwrap();
assert_eq!(decoded, req);
}
#[test]
fn req_wrong_len() {
assert!(MsgCounterSyncReq::read(&[0u8; 7]).is_err());
assert!(MsgCounterSyncReq::read(&[0u8; 9]).is_err());
assert!(MsgCounterSyncReq::read(&[]).is_err());
}
#[test]
fn rsp_roundtrip() {
let rsp = MsgCounterSyncRsp {
synchronized_counter: 0x1234_5678,
response: [0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff, 0x11, 0x22],
};
let mut buf = [0u8; 32];
let mut wb = WriteBuf::new(&mut buf);
rsp.write(&mut wb).unwrap();
let slice = wb.as_slice();
assert_eq!(slice.len(), MCSP_SYNC_RSP_LEN);
assert_eq!(&slice[0..4], &[0x78, 0x56, 0x34, 0x12]);
assert_eq!(
&slice[4..12],
&[0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff, 0x11, 0x22]
);
let decoded = MsgCounterSyncRsp::read(slice).unwrap();
assert_eq!(decoded, rsp);
}
#[test]
fn rsp_wrong_len() {
assert!(MsgCounterSyncRsp::read(&[0u8; 11]).is_err());
assert!(MsgCounterSyncRsp::read(&[0u8; 13]).is_err());
assert!(MsgCounterSyncRsp::read(&[]).is_err());
}
}