use std::net::SocketAddr;
use std::time::Duration;
use matter_transport::{
decode_header, decode_protocol_header, encode_header, encode_protocol_header, DestNodeId,
ExchangeFlags, MessageCounter, NodeId, ProtocolHeader, ProtocolId, SecuredMessageFlags,
SecuredMessageHeader, SecurityFlags, SessionId,
};
use crate::driver::datagram::AsyncDatagram;
use crate::driver::error::DriverError;
use crate::driver::TransportReliability;
const UNSECURED_RESPONSE_TIMEOUT: Duration = Duration::from_secs(30);
const OPCODE_MRP_STANDALONE_ACK: u8 = 0x10;
const OPCODE_STATUS_REPORT: u8 = 0x40;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SecureChannelStatus {
pub general_code: u16,
pub protocol_id: u32,
pub protocol_code: u16,
}
impl SecureChannelStatus {
#[must_use]
pub fn is_session_establishment_success(self) -> bool {
self.general_code == 0 && self.protocol_id == 0 && self.protocol_code == 0
}
}
pub fn parse_status_report(msg: &UnsecuredMessage) -> Result<SecureChannelStatus, DriverError> {
if msg.opcode != OPCODE_STATUS_REPORT {
return Err(DriverError::Handshake(
"expected a SecureChannel StatusReport to close the handshake",
));
}
let b = &msg.payload;
if b.len() < 8 {
return Err(DriverError::Handshake("StatusReport body truncated"));
}
Ok(SecureChannelStatus {
general_code: u16::from_le_bytes([b[0], b[1]]),
protocol_id: u32::from_le_bytes([b[2], b[3], b[4], b[5]]),
protocol_code: u16::from_le_bytes([b[6], b[7]]),
})
}
pub fn require_handshake_opcode(msg: &UnsecuredMessage, opcode: u8) -> Result<(), DriverError> {
if msg.protocol_id == ProtocolId::SECURE_CHANNEL && msg.opcode == opcode {
return Ok(());
}
if msg.protocol_id == ProtocolId::SECURE_CHANNEL && msg.opcode == OPCODE_STATUS_REPORT {
let status = parse_status_report(msg)?;
return Err(DriverError::SessionEstablishmentFailed {
general_code: status.general_code,
protocol_code: status.protocol_code,
});
}
Err(DriverError::Handshake(
"unexpected opcode in session-establishment exchange",
))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UnsecuredMessage {
pub message_counter: u32,
pub exchange_id: u16,
pub opcode: u8,
pub protocol_id: ProtocolId,
pub is_initiator: bool,
pub ack_counter: Option<u32>,
pub source_node_id: Option<u64>,
pub payload: Vec<u8>,
}
#[allow(clippy::too_many_arguments)] #[must_use]
pub fn encode_unsecured(
message_counter: u32,
exchange_id: u16,
opcode: u8,
protocol_id: ProtocolId,
initiator: bool,
reliable: bool,
ack: Option<u32>,
source_node_id: Option<u64>,
app_payload: &[u8],
) -> Vec<u8> {
let mut flags = SecuredMessageFlags::empty();
if source_node_id.is_some() {
flags |= SecuredMessageFlags::SOURCE_PRESENT;
}
let header = SecuredMessageHeader {
flags,
session_id: SessionId(0),
security_flags: SecurityFlags::empty(),
message_counter: MessageCounter(message_counter),
source_node_id: source_node_id.map(NodeId),
destination_node_id: None,
};
let mut buf = encode_header(&header);
let mut exchange_flags = ExchangeFlags::empty();
if initiator {
exchange_flags |= ExchangeFlags::INITIATOR;
}
if reliable {
exchange_flags |= ExchangeFlags::RELIABLE;
}
let protocol_header = ProtocolHeader {
exchange_flags,
opcode,
exchange_id,
protocol_id,
ack_counter: ack.map(MessageCounter),
};
encode_protocol_header(&protocol_header, &mut buf);
buf.extend_from_slice(app_payload);
buf
}
#[allow(clippy::too_many_arguments)] #[must_use]
pub fn encode_unsecured_reply(
message_counter: u32,
exchange_id: u16,
opcode: u8,
protocol_id: ProtocolId,
reliable: bool,
ack: Option<u32>,
destination_node_id: Option<u64>,
app_payload: &[u8],
) -> Vec<u8> {
let mut flags = SecuredMessageFlags::empty();
if destination_node_id.is_some() {
flags |= SecuredMessageFlags::DEST_UNICAST;
}
let header = SecuredMessageHeader {
flags,
session_id: SessionId(0),
security_flags: SecurityFlags::empty(),
message_counter: MessageCounter(message_counter),
source_node_id: None,
destination_node_id: destination_node_id.map(|n| DestNodeId::Node(NodeId(n))),
};
let mut buf = encode_header(&header);
let mut exchange_flags = ExchangeFlags::empty();
if reliable {
exchange_flags |= ExchangeFlags::RELIABLE;
}
let protocol_header = ProtocolHeader {
exchange_flags,
opcode,
exchange_id,
protocol_id,
ack_counter: ack.map(MessageCounter),
};
encode_protocol_header(&protocol_header, &mut buf);
buf.extend_from_slice(app_payload);
buf
}
pub fn decode_unsecured(bytes: &[u8]) -> Result<UnsecuredMessage, DriverError> {
let (msg_header, rest) = decode_header(bytes)?;
if msg_header.session_id.0 != 0 {
return Err(DriverError::UnexpectedSecuredMessage(
msg_header.session_id.0,
));
}
let (protocol_header, app) = decode_protocol_header(rest)?;
Ok(UnsecuredMessage {
message_counter: msg_header.message_counter.0,
exchange_id: protocol_header.exchange_id,
opcode: protocol_header.opcode,
protocol_id: protocol_header.protocol_id,
is_initiator: protocol_header
.exchange_flags
.contains(ExchangeFlags::INITIATOR),
ack_counter: protocol_header.ack_counter.map(|c| c.0),
source_node_id: msg_header.source_node_id.map(|n| n.0),
payload: app.to_vec(),
})
}
pub struct UnsecuredExchange {
counter: u32,
exchange_id: u16,
source_node_id: u64,
retransmit: Duration,
response_timeout: Duration,
max_attempts: u8,
last_consumed_peer_counter: Option<u32>,
reliability: TransportReliability,
}
impl UnsecuredExchange {
#[must_use]
pub fn new(initial_counter: u32, exchange_id: u16, source_node_id: u64) -> Self {
Self {
counter: initial_counter,
exchange_id,
source_node_id,
retransmit: Duration::from_millis(300),
response_timeout: UNSECURED_RESPONSE_TIMEOUT,
max_attempts: 5,
last_consumed_peer_counter: None,
reliability: TransportReliability::Mrp,
}
}
pub fn new_ephemeral(exchange_id: u16) -> Result<Self, DriverError> {
Self::new_ephemeral_with(exchange_id, TransportReliability::Mrp)
}
pub fn new_ephemeral_with(
exchange_id: u16,
reliability: TransportReliability,
) -> Result<Self, DriverError> {
let rng = ring::rand::SystemRandom::new();
let mut bytes = [0u8; 12];
ring::rand::SecureRandom::fill(&rng, &mut bytes).map_err(|_| {
DriverError::Handshake("system CSPRNG failure seeding unsecured session")
})?;
let counter_seed: [u8; 4] = [bytes[0], bytes[1], bytes[2], bytes[3]];
let node_seed: [u8; 8] = [
bytes[4], bytes[5], bytes[6], bytes[7], bytes[8], bytes[9], bytes[10], bytes[11],
];
let counter = (u32::from_le_bytes(counter_seed) & 0x0FFF_FFFF) + 1;
let source_node_id = (u64::from_le_bytes(node_seed) & 0x0FFF_FFFF_FFFF_FFFF).max(1);
let mut exch = Self::new(counter, exchange_id, source_node_id);
exch.reliability = reliability;
Ok(exch)
}
pub async fn send_standalone_ack<T: AsyncDatagram>(
&mut self,
transport: &T,
peer: SocketAddr,
peer_counter: u32,
) -> Result<(), DriverError> {
if self.reliability == TransportReliability::TransportProvides {
return Ok(());
}
let counter = self.counter;
self.counter = self.counter.wrapping_add(1);
let wire = encode_unsecured(
counter,
self.exchange_id,
OPCODE_MRP_STANDALONE_ACK,
ProtocolId::SECURE_CHANNEL,
true,
false,
Some(peer_counter),
Some(self.source_node_id),
&[],
);
transport.send_to(&wire, peer).await?;
#[cfg(feature = "tracing")]
tracing::debug!(
target: "matter_wire",
dir = "tx",
session_id = 0_u64,
exchange_id = u64::from(self.exchange_id),
protocol = u64::from(ProtocolId::SECURE_CHANNEL.protocol),
opcode = u64::from(OPCODE_MRP_STANDALONE_ACK),
payload = "",
"wire"
);
Ok(())
}
#[allow(clippy::too_many_lines)]
pub async fn send_and_recv<T: AsyncDatagram>(
&mut self,
transport: &T,
peer: SocketAddr,
opcode: u8,
expected_opcode: u8,
app_payload: &[u8],
ack: Option<u32>,
) -> Result<UnsecuredMessage, DriverError> {
let counter = self.counter;
self.counter = self.counter.wrapping_add(1);
let transport_provides = self.reliability == TransportReliability::TransportProvides;
let reliable_bit = !transport_provides;
let ack_field = if transport_provides { None } else { ack };
let wire = encode_unsecured(
counter,
self.exchange_id,
opcode,
ProtocolId::SECURE_CHANNEL,
true,
reliable_bit,
ack_field,
Some(self.source_node_id),
app_payload,
);
#[cfg(feature = "tracing")]
tracing::debug!(
opcode = format_args!("{opcode:#04x}"),
exchange_id = self.exchange_id,
wire = %crate::hexdump::hex(&wire),
"unsecured send"
);
#[cfg(feature = "tracing")]
tracing::debug!(
target: "matter_wire",
dir = "tx",
session_id = 0_u64,
exchange_id = u64::from(self.exchange_id),
protocol = u64::from(ProtocolId::SECURE_CHANNEL.protocol),
opcode = u64::from(opcode),
payload = %crate::hexdump::hex(app_payload),
"wire"
);
let mut attempts: u8 = 0;
let mut acked = false;
let mut sent = false;
loop {
let should_send = if transport_provides { !sent } else { !acked };
if should_send {
transport.send_to(&wire, peer).await?;
sent = true;
}
let wait = if transport_provides || acked {
self.response_timeout
} else {
self.retransmit
};
match tokio::time::timeout(wait, transport.recv_from()).await {
Ok(recv) => {
let (packet, _from) = recv?;
#[cfg(feature = "tracing")]
tracing::debug!(
len = packet.len(),
head = %crate::hexdump::hex(&packet[..packet.len().min(24)]),
"unsecured recv"
);
if packet.len() >= 3 && (packet[1] != 0 || packet[2] != 0) {
continue;
}
let msg = decode_unsecured(&packet)?;
if msg.exchange_id != self.exchange_id {
continue;
}
if msg.protocol_id == ProtocolId::SECURE_CHANNEL
&& msg.opcode == OPCODE_MRP_STANDALONE_ACK
{
acked = true;
continue;
}
if let Some(last) = self.last_consumed_peer_counter {
if msg.message_counter <= last {
continue;
}
}
if msg.protocol_id == ProtocolId::SECURE_CHANNEL
&& msg.opcode != expected_opcode
&& msg.opcode != OPCODE_STATUS_REPORT
{
continue;
}
self.last_consumed_peer_counter = Some(msg.message_counter);
#[cfg(feature = "tracing")]
tracing::debug!(
target: "matter_wire",
dir = "rx",
session_id = 0_u64,
exchange_id = u64::from(msg.exchange_id),
protocol = u64::from(msg.protocol_id.protocol),
opcode = u64::from(msg.opcode),
payload = %crate::hexdump::hex(&msg.payload),
"wire"
);
return Ok(msg);
}
Err(_elapsed) => {
if transport_provides || acked {
return Err(DriverError::Timeout {
exchange_id: self.exchange_id,
});
}
attempts += 1;
if attempts >= self.max_attempts {
return Err(DriverError::Timeout {
exchange_id: self.exchange_id,
});
}
}
}
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)] mod tests {
use matter_transport::ProtocolId;
use crate::driver::datagram::{AsyncDatagram, InMemoryDatagram};
use super::*;
#[test]
fn unsecured_roundtrip_preserves_fields() {
let wire = encode_unsecured(
0x8000_0001,
42,
0x22,
ProtocolId::SECURE_CHANNEL,
true,
true,
Some(7),
None,
b"pake1-bytes",
);
let msg = decode_unsecured(&wire).unwrap();
assert_eq!(msg.message_counter, 0x8000_0001);
assert_eq!(msg.exchange_id, 42);
assert_eq!(msg.opcode, 0x22);
assert_eq!(msg.protocol_id, ProtocolId::SECURE_CHANNEL);
assert!(msg.is_initiator);
assert_eq!(msg.ack_counter, Some(7));
assert_eq!(msg.payload, b"pake1-bytes");
}
#[test]
fn unsecured_rejects_secured_session_id() {
use matter_transport::{
encode_header, MessageCounter, SecuredMessageFlags, SecuredMessageHeader,
SecurityFlags, SessionId,
};
let hdr = SecuredMessageHeader {
flags: SecuredMessageFlags::empty(),
session_id: SessionId(5), security_flags: SecurityFlags::empty(),
message_counter: MessageCounter(1),
source_node_id: None,
destination_node_id: None,
};
let mut wire = encode_header(&hdr);
wire.extend_from_slice(&[0u8; 6]); let err = decode_unsecured(&wire).unwrap_err();
assert!(matches!(err, DriverError::UnexpectedSecuredMessage(5)));
}
#[tokio::test]
async fn unsecured_send_and_recv_roundtrips() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch = UnsecuredExchange::new(1, 7, 0xE0E0);
let controller = exch.send_and_recv(
&ctrl_io, dev_addr, 0x20,
0x21,
b"req", None,
);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let msg = decode_unsecured(&pkt).unwrap();
assert_eq!(msg.opcode, 0x20);
assert_eq!(msg.payload, b"req");
let reply = encode_unsecured(
100,
msg.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
true,
Some(msg.message_counter),
None,
b"resp",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
let got = got.unwrap();
assert_eq!(got.opcode, 0x21);
assert_eq!(got.payload, b"resp");
}
#[tokio::test]
async fn unsecured_standalone_ack_has_expected_shape() {
let (a, b) = InMemoryDatagram::pair();
let b_addr = b.local_addr();
let mut exch = UnsecuredExchange::new(5, 9, 0xE0E0);
exch.send_standalone_ack(&a, b_addr, 7).await.unwrap();
let (pkt, _) = b.recv_from().await.unwrap();
let msg = decode_unsecured(&pkt).unwrap();
assert_eq!(msg.opcode, 0x10);
assert!(msg.payload.is_empty());
assert_eq!(msg.message_counter, 5);
assert_eq!(msg.ack_counter, Some(7));
assert!(msg.is_initiator);
assert_eq!(msg.source_node_id, Some(0xE0E0));
}
#[test]
fn status_report_parses_success_and_failure() {
let mk = |general: u16, code: u16| UnsecuredMessage {
message_counter: 1,
exchange_id: 1,
opcode: 0x40,
protocol_id: ProtocolId::SECURE_CHANNEL,
is_initiator: false,
ack_counter: None,
source_node_id: None,
payload: {
let mut b = Vec::new();
b.extend_from_slice(&general.to_le_bytes());
b.extend_from_slice(&0u32.to_le_bytes());
b.extend_from_slice(&code.to_le_bytes());
b
},
};
let ok = parse_status_report(&mk(0, 0)).unwrap();
assert!(ok.is_session_establishment_success());
let no = parse_status_report(&mk(1, 0x0002)).unwrap();
assert!(!no.is_session_establishment_success());
assert_eq!(no.general_code, 1);
assert_eq!(no.protocol_code, 0x0002);
let mut wrong = mk(0, 0);
wrong.opcode = 0x21;
assert!(parse_status_report(&wrong).is_err());
}
#[test]
fn unsecured_encode_carries_source_node_id() {
let wire = encode_unsecured(
1,
7,
0x20,
ProtocolId::SECURE_CHANNEL,
true,
true,
None,
Some(0x1122_3344_5566_7788),
b"req",
);
assert_eq!(wire[0] & 0b0000_0100, 0b0000_0100, "S flag must be set");
assert_eq!(wire[8..16], 0x1122_3344_5566_7788u64.to_le_bytes());
let msg = decode_unsecured(&wire).unwrap();
assert_eq!(msg.source_node_id, Some(0x1122_3344_5566_7788));
}
#[tokio::test]
async fn unsecured_exchange_frames_carry_source_node_id() {
let (a, b) = InMemoryDatagram::pair();
let b_addr = b.local_addr();
let mut exch = UnsecuredExchange::new(5, 9, 0xABCD);
exch.send_standalone_ack(&a, b_addr, 3).await.unwrap();
let (pkt, _) = b.recv_from().await.unwrap();
let msg = decode_unsecured(&pkt).unwrap();
assert_eq!(msg.source_node_id, Some(0xABCD));
}
#[tokio::test]
async fn unsecured_new_ephemeral_seeds_counter_and_node_id() {
let (a, b) = InMemoryDatagram::pair();
let b_addr = b.local_addr();
let mut exch = UnsecuredExchange::new_ephemeral(9).unwrap();
exch.send_standalone_ack(&a, b_addr, 3).await.unwrap();
let (pkt, _) = b.recv_from().await.unwrap();
let msg = decode_unsecured(&pkt).unwrap();
let node_id = msg.source_node_id.expect("ephemeral node id present");
assert_ne!(node_id, 0, "node id must be nonzero (operational range)");
assert_ne!(msg.message_counter, 0, "counter must be nonzero");
assert!(
msg.message_counter <= 0x0FFF_FFFF + 1,
"counter seeded as random 28-bit + 1"
);
}
#[tokio::test]
async fn unsecured_send_and_recv_skips_standalone_ack() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch = UnsecuredExchange::new(1, 7, 0xE0E0);
let controller = exch.send_and_recv(&ctrl_io, dev_addr, 0x20, 0x21, b"req", None);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let msg = decode_unsecured(&pkt).unwrap();
let ack = encode_unsecured(
100,
msg.exchange_id,
0x10,
ProtocolId::SECURE_CHANNEL,
false,
false,
Some(msg.message_counter),
None,
b"",
);
dev_io.send_to(&ack, ctrl_addr).await.unwrap();
let reply = encode_unsecured(
101,
msg.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
true,
Some(msg.message_counter),
None,
b"resp",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
let got = got.unwrap();
assert_eq!(got.opcode, 0x21, "standalone ack must be skipped");
assert_eq!(got.payload, b"resp");
}
#[tokio::test]
async fn unsecured_send_and_recv_skips_secured_frames() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch = UnsecuredExchange::new(1, 7, 0xE0E0);
let controller = exch.send_and_recv(&ctrl_io, dev_addr, 0x30, 0x31, b"sigma1", None);
let device = async {
use matter_transport::{
encode_header, MessageCounter, SecuredMessageFlags, SecuredMessageHeader,
SecurityFlags, SessionId,
};
let (pkt, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&pkt).unwrap();
let hdr = SecuredMessageHeader {
flags: SecuredMessageFlags::empty(),
session_id: SessionId(1),
security_flags: SecurityFlags::empty(),
message_counter: MessageCounter(99),
source_node_id: None,
destination_node_id: None,
};
let mut stray = encode_header(&hdr);
stray.extend_from_slice(&[0xAA; 24]); dev_io.send_to(&stray, ctrl_addr).await.unwrap();
let reply = encode_unsecured(
100,
m.exchange_id,
0x31,
ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
b"sigma2",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
let got = got.unwrap();
assert_eq!(got.opcode, 0x31, "secured straggler must be skipped");
assert_eq!(got.payload, b"sigma2");
}
#[tokio::test]
async fn unsecured_send_and_recv_skips_stale_prior_step_response() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch = UnsecuredExchange::new(1, 7, 0xE0E0);
let controller = exch.send_and_recv(
&ctrl_io, dev_addr, 0x22,
0x23,
b"pake1", None,
);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&pkt).unwrap();
let stale = encode_unsecured(
100,
m.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
b"stale-pbkdf-response",
);
dev_io.send_to(&stale, ctrl_addr).await.unwrap();
let reply = encode_unsecured(
101,
m.exchange_id,
0x23,
ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
b"pake2",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
let got = got.unwrap();
assert_eq!(
got.opcode, 0x23,
"stale prior-step response must be skipped"
);
assert_eq!(got.payload, b"pake2");
}
#[tokio::test]
async fn unsecured_send_and_recv_unexpected_frame_times_out() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch = UnsecuredExchange::new(1, 7, 0xE0E0);
exch.retransmit = Duration::from_millis(30);
exch.response_timeout = Duration::from_millis(60);
exch.max_attempts = 2;
let controller = exch.send_and_recv(
&ctrl_io, dev_addr, 0x22,
0x23,
b"pake1", None,
);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&pkt).unwrap();
let ack = encode_unsecured(
100,
m.exchange_id,
OPCODE_MRP_STANDALONE_ACK,
ProtocolId::SECURE_CHANNEL,
false,
false,
Some(m.message_counter),
None,
b"",
);
dev_io.send_to(&ack, ctrl_addr).await.unwrap();
let stale = encode_unsecured(
101,
m.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
b"stale",
);
dev_io.send_to(&stale, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
assert!(
matches!(got, Err(DriverError::Timeout { exchange_id: 7 })),
"unexpected unresolving frame must time out, got: {got:?}"
);
}
#[tokio::test]
async fn unsecured_send_and_recv_retransmits_dropped_send() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch = UnsecuredExchange::new(1, 7, 0xE0E0);
ctrl_io.set_drops(1);
let controller = exch.send_and_recv(&ctrl_io, dev_addr, 0x20, 0x21, b"req", None);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap(); let msg = decode_unsecured(&pkt).unwrap();
let reply = encode_unsecured(
100,
msg.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
true,
Some(msg.message_counter),
None,
b"resp",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
assert_eq!(got.unwrap().opcode, 0x21);
}
#[tokio::test]
async fn transport_provides_sets_no_r_bit() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch =
UnsecuredExchange::new_ephemeral_with(7, TransportReliability::TransportProvides)
.unwrap();
let controller = exch.send_and_recv(&ctrl_io, dev_addr, 0x20, 0x21, b"req", Some(5));
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let (_hdr, rest) = decode_header(&pkt).unwrap();
let (ph, _) = decode_protocol_header(rest).unwrap();
assert!(
!ph.exchange_flags.contains(ExchangeFlags::RELIABLE),
"R-flag must be clear under TransportProvides"
);
assert!(
ph.ack_counter.is_none(),
"ack must not be attached under TransportProvides"
);
let m = decode_unsecured(&pkt).unwrap();
let reply = encode_unsecured(
100,
m.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
false,
None,
None,
b"resp",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
assert_eq!(got.unwrap().opcode, 0x21);
}
#[tokio::test]
async fn mrp_sets_r_bit_and_ack() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch = UnsecuredExchange::new_ephemeral_with(7, TransportReliability::Mrp).unwrap();
let controller = exch.send_and_recv(&ctrl_io, dev_addr, 0x20, 0x21, b"req", Some(5));
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let (_hdr, rest) = decode_header(&pkt).unwrap();
let (ph, _) = decode_protocol_header(rest).unwrap();
assert!(
ph.exchange_flags.contains(ExchangeFlags::RELIABLE),
"R-flag must be set under Mrp"
);
assert_eq!(
ph.ack_counter.map(|c| c.0),
Some(5),
"ack must be attached under Mrp"
);
let m = decode_unsecured(&pkt).unwrap();
let reply = encode_unsecured(
100,
m.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
false,
None,
None,
b"resp",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
assert_eq!(got.unwrap().opcode, 0x21);
}
#[tokio::test(start_paused = true)]
async fn transport_provides_sends_exactly_once() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut exch =
UnsecuredExchange::new_ephemeral_with(7, TransportReliability::TransportProvides)
.unwrap();
let controller = exch.send_and_recv(&ctrl_io, dev_addr, 0x20, 0x21, b"req", None);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&pkt).unwrap();
tokio::time::sleep(Duration::from_millis(450)).await;
let reply = encode_unsecured(
100,
m.exchange_id,
0x21,
ProtocolId::SECURE_CHANNEL,
false,
false,
None,
None,
b"resp",
);
dev_io.send_to(&reply, ctrl_addr).await.unwrap();
let extra = tokio::time::timeout(Duration::from_millis(500), dev_io.recv_from()).await;
assert!(
extra.is_err(),
"TransportProvides must send exactly one packet (no retransmit)"
);
};
let (got, ()) = tokio::join!(controller, device);
assert_eq!(got.unwrap().opcode, 0x21);
}
#[tokio::test(start_paused = true)]
async fn transport_provides_times_out_at_30s() {
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = ctrl_io.local_addr(); let _keep_peer_alive = dev_io; let mut exch =
UnsecuredExchange::new_ephemeral_with(7, TransportReliability::TransportProvides)
.unwrap();
let start = tokio::time::Instant::now();
let res = exch
.send_and_recv(&ctrl_io, dev_addr, 0x20, 0x21, b"req", None)
.await;
let elapsed = start.elapsed();
assert!(
matches!(res, Err(DriverError::Timeout { exchange_id: 7 })),
"expected a Timeout, got: {res:?}"
);
assert!(
elapsed >= Duration::from_secs(30),
"TransportProvides must wait the full 30 s response deadline, waited {elapsed:?}"
);
}
#[tokio::test(start_paused = true)]
async fn standalone_ack_noop_under_transport_provides() {
let (a, b) = InMemoryDatagram::pair();
let b_addr = b.local_addr();
let mut exch =
UnsecuredExchange::new_ephemeral_with(9, TransportReliability::TransportProvides)
.unwrap();
let counter_before = exch.counter;
exch.send_standalone_ack(&a, b_addr, 7).await.unwrap();
let got = tokio::time::timeout(Duration::from_millis(50), b.recv_from()).await;
assert!(
got.is_err(),
"send_standalone_ack must be a no-op under TransportProvides"
);
assert_eq!(
exch.counter, counter_before,
"the no-op must not advance the message counter"
);
}
}