use std::net::SocketAddr;
use std::time::{Duration, Instant};
use matter_transport::{
DecodeInboundOutput, MrpEvent, MrpFlags, ProtocolId, SessionId, SessionManager,
};
use crate::driver::datagram::AsyncDatagram;
use crate::driver::error::DriverError;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SecuredResponse {
pub exchange_id: u16,
pub payload: Vec<u8>,
}
const IDLE_SLEEP: Duration = Duration::from_secs(3600);
pub async fn secured_round_trip<T: AsyncDatagram>(
transport: &T,
sessions: &mut SessionManager,
session_id: SessionId,
peer: SocketAddr,
opcode: u8,
protocol_id: ProtocolId,
app_payload: &[u8],
) -> Result<SecuredResponse, DriverError> {
let out = sessions.encode_outbound(
session_id,
None, opcode,
protocol_id,
app_payload,
MrpFlags { reliable: true },
Instant::now(),
)?;
let our_exchange = out.exchange_id;
transport.send_to(&out.wire_bytes, peer).await?;
#[cfg(feature = "tracing")]
tracing::debug!(
target: "matter_wire",
dir = "tx",
session_id = u64::from(session_id.0),
exchange_id = u64::from(our_exchange),
protocol = u64::from(protocol_id.protocol),
opcode = u64::from(opcode),
payload = %crate::hexdump::hex(app_payload),
"wire"
);
loop {
let now = Instant::now();
let sleep_for = sessions.poll_timeout().map_or(IDLE_SLEEP, |deadline| {
deadline.saturating_duration_since(now)
});
tokio::select! {
biased;
recv = transport.recv_from() => {
let (packet, _from) = recv?;
if packet.len() >= 3 && packet[1] == 0 && packet[2] == 0 {
continue;
}
let decoded = match sessions.decode_inbound(&packet, Instant::now()) {
Ok(d) => d,
Err(
matter_transport::Error::UnknownSession(_)
| matter_transport::Error::DecryptionFailed,
) => continue,
Err(e) => return Err(e.into()),
};
match decoded {
DecodeInboundOutput::AppMessage {
exchange_id,
payload,
protocol_id: msg_protocol_id,
opcode: msg_opcode,
..
} if exchange_id == our_exchange => {
#[cfg(feature = "tracing")]
tracing::debug!(
target: "matter_wire",
dir = "rx",
session_id = u64::from(session_id.0),
exchange_id = u64::from(exchange_id),
protocol = u64::from(msg_protocol_id.protocol),
opcode = u64::from(msg_opcode),
payload = %crate::hexdump::hex(&payload),
"wire"
);
#[cfg(not(feature = "tracing"))]
let _ = (&msg_protocol_id, &msg_opcode);
return Ok(SecuredResponse { exchange_id, payload });
}
DecodeInboundOutput::DuplicateReliableAckResent { ack_packet, .. } => {
transport.send_to(&ack_packet, peer).await?;
}
_ => {}
}
}
() = tokio::time::sleep(sleep_for) => {
for event in sessions.handle_timeout(Instant::now()) {
match event {
MrpEvent::Retransmit { packet, .. }
| MrpEvent::SendStandaloneAck { packet, .. } => {
transport.send_to(&packet, peer).await?;
}
MrpEvent::Expired { exchange_id, .. } => {
return Err(DriverError::Timeout { exchange_id });
}
_ => {}
}
}
}
}
}
}
pub const MAX_READ_CHUNKS: usize = 64;
pub const MAX_READ_BYTES: usize = 256 * 1024;
pub const MAX_READ_CHUNK_BYTES: usize = 64 * 1024;
const OP_STATUS_RESPONSE: u8 = 0x01;
fn enforce_read_caps(
existing_chunks: usize,
chunk_len: usize,
total_bytes: usize,
) -> Result<usize, DriverError> {
if chunk_len > MAX_READ_CHUNK_BYTES {
return Err(DriverError::ReadTooLarge {
limit: "MAX_READ_CHUNK_BYTES",
});
}
if existing_chunks + 1 > MAX_READ_CHUNKS {
return Err(DriverError::ReadTooLarge {
limit: "MAX_READ_CHUNKS",
});
}
let total = total_bytes.saturating_add(chunk_len);
if total > MAX_READ_BYTES {
return Err(DriverError::ReadTooLarge {
limit: "MAX_READ_BYTES",
});
}
Ok(total)
}
pub async fn secured_read<T: AsyncDatagram>(
transport: &T,
sessions: &mut SessionManager,
session_id: SessionId,
peer: SocketAddr,
opcode: u8,
protocol_id: ProtocolId,
app_payload: &[u8],
) -> Result<Vec<Vec<u8>>, DriverError> {
let out = sessions.encode_outbound(
session_id,
None,
opcode,
protocol_id,
app_payload,
MrpFlags { reliable: true },
Instant::now(),
)?;
let our_exchange = out.exchange_id;
transport.send_to(&out.wire_bytes, peer).await?;
let mut chunks: Vec<Vec<u8>> = Vec::new();
let mut total_bytes = 0usize;
loop {
let now = Instant::now();
let sleep_for = sessions.poll_timeout().map_or(IDLE_SLEEP, |deadline| {
deadline.saturating_duration_since(now)
});
tokio::select! {
biased;
recv = transport.recv_from() => {
let (packet, _from) = recv?;
if packet.len() >= 3 && packet[1] == 0 && packet[2] == 0 {
continue;
}
let decoded = match sessions.decode_inbound(&packet, Instant::now()) {
Ok(d) => d,
Err(
matter_transport::Error::UnknownSession(_)
| matter_transport::Error::DecryptionFailed,
) => continue,
Err(e) => return Err(e.into()),
};
match decoded {
DecodeInboundOutput::AppMessage { exchange_id, payload, .. }
if exchange_id == our_exchange =>
{
total_bytes = enforce_read_caps(chunks.len(), payload.len(), total_bytes)?;
let more = crate::im::parse_report_data(&payload)?.more_chunked_messages;
chunks.push(payload);
if !more {
return Ok(chunks);
}
let status = crate::im::build_status_response(0);
let ack = sessions.encode_outbound(
session_id,
Some(our_exchange),
OP_STATUS_RESPONSE,
ProtocolId::INTERACTION_MODEL,
&status,
MrpFlags { reliable: true },
Instant::now(),
)?;
transport.send_to(&ack.wire_bytes, peer).await?;
}
DecodeInboundOutput::DuplicateReliableAckResent { ack_packet, .. } => {
transport.send_to(&ack_packet, peer).await?;
}
_ => {}
}
}
() = tokio::time::sleep(sleep_for) => {
for event in sessions.handle_timeout(Instant::now()) {
match event {
MrpEvent::Retransmit { packet, .. }
| MrpEvent::SendStandaloneAck { packet, .. } => {
transport.send_to(&packet, peer).await?;
}
MrpEvent::Expired { exchange_id, .. } => {
return Err(DriverError::Timeout { exchange_id });
}
_ => {}
}
}
}
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)] mod tests {
use std::time::Instant;
use matter_crypto::pase::PaseSessionKeys;
use matter_transport::{
DecodeInboundOutput, MrpFlags, PeerHint, ProtocolId, SessionManager, SessionRole,
};
use super::*;
use crate::driver::datagram::{AsyncDatagram, InMemoryDatagram};
fn paired_pase_sessions() -> (SessionManager, SessionManager) {
let keys = PaseSessionKeys {
ke: [0u8; 16],
i2r_key: [1u8; 16],
r2i_key: [2u8; 16],
attestation_key: [3u8; 16],
};
let mut ctrl = SessionManager::new();
let mut dev = SessionManager::new();
let _ctrl_sid =
ctrl.register_pase(keys.clone(), SessionRole::Initiator, 1, PeerHint::default());
let _dev_sid = dev.register_pase(keys, SessionRole::Responder, 1, PeerHint::default());
(ctrl, dev)
}
#[tokio::test]
async fn secured_round_trip_retransmits_dropped_request() {
let (mut ctrl, mut dev) = paired_pase_sessions();
let session = matter_transport::SessionId(1);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
ctrl_io.set_drops(1);
let request = b"req".as_slice();
let response = b"resp".as_slice();
let controller = secured_round_trip(
&ctrl_io,
&mut ctrl,
session,
dev_addr,
0x08,
ProtocolId::INTERACTION_MODEL,
request,
);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let DecodeInboundOutput::AppMessage {
exchange_id,
payload,
..
} = dev.decode_inbound(&pkt, Instant::now()).unwrap()
else {
panic!("expected an application message");
};
assert_eq!(payload, request);
let out = dev
.encode_outbound(
session,
Some(exchange_id),
0x09,
ProtocolId::INTERACTION_MODEL,
response,
MrpFlags { reliable: true },
Instant::now(),
)
.unwrap();
dev_io.send_to(&out.wire_bytes, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
assert_eq!(got.unwrap().payload, response);
}
#[tokio::test]
async fn secured_round_trip_skips_unsecured_frames() {
let (mut ctrl, mut dev) = paired_pase_sessions();
let session = matter_transport::SessionId(1);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let request = b"req".as_slice();
let response = b"resp".as_slice();
let controller = secured_round_trip(
&ctrl_io,
&mut ctrl,
session,
dev_addr,
0x08,
ProtocolId::INTERACTION_MODEL,
request,
);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let DecodeInboundOutput::AppMessage { exchange_id, .. } =
dev.decode_inbound(&pkt, Instant::now()).unwrap()
else {
panic!("expected an application message");
};
let stray = crate::driver::unsecured::encode_unsecured(
7,
1,
0x40,
ProtocolId::SECURE_CHANNEL,
false,
true,
None,
None,
&[0u8; 8],
);
dev_io.send_to(&stray, ctrl_addr).await.unwrap();
let out = dev
.encode_outbound(
session,
Some(exchange_id),
0x09,
ProtocolId::INTERACTION_MODEL,
response,
MrpFlags { reliable: true },
Instant::now(),
)
.unwrap();
dev_io.send_to(&out.wire_bytes, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
assert_eq!(got.unwrap().payload, response);
}
#[tokio::test]
async fn secured_round_trip_skips_unknown_session_frames() {
let (mut ctrl, mut dev) = paired_pase_sessions();
let session = matter_transport::SessionId(1);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let stale_keys = PaseSessionKeys {
ke: [9u8; 16],
i2r_key: [9u8; 16],
r2i_key: [9u8; 16],
attestation_key: [9u8; 16],
};
let mut stale = SessionManager::new();
let stale_sid =
stale.register_pase(stale_keys, SessionRole::Initiator, 99, PeerHint::default());
let request = b"req".as_slice();
let response = b"resp".as_slice();
let controller = secured_round_trip(
&ctrl_io,
&mut ctrl,
session,
dev_addr,
0x08,
ProtocolId::INTERACTION_MODEL,
request,
);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let DecodeInboundOutput::AppMessage { exchange_id, .. } =
dev.decode_inbound(&pkt, Instant::now()).unwrap()
else {
panic!("expected an application message");
};
let stray = stale
.encode_outbound(
stale_sid,
None,
0x09,
ProtocolId::INTERACTION_MODEL,
b"stale",
MrpFlags { reliable: false },
Instant::now(),
)
.unwrap();
dev_io.send_to(&stray.wire_bytes, ctrl_addr).await.unwrap();
let out = dev
.encode_outbound(
session,
Some(exchange_id),
0x09,
ProtocolId::INTERACTION_MODEL,
response,
MrpFlags { reliable: true },
Instant::now(),
)
.unwrap();
dev_io.send_to(&out.wire_bytes, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
assert_eq!(got.unwrap().payload, response);
}
#[tokio::test]
async fn secured_round_trip_returns_response_payload() {
let (mut ctrl, mut dev) = paired_pase_sessions();
let session = matter_transport::SessionId(1);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let request = b"invoke-request-tlv".as_slice();
let response = b"invoke-response-tlv".as_slice();
let controller = secured_round_trip(
&ctrl_io,
&mut ctrl,
session,
dev_addr,
0x08, ProtocolId::INTERACTION_MODEL,
request,
);
let device = async {
loop {
let (pkt, _) = dev_io.recv_from().await.unwrap();
if let DecodeInboundOutput::AppMessage {
exchange_id,
payload,
..
} = dev.decode_inbound(&pkt, Instant::now()).unwrap()
{
assert_eq!(payload, request);
let out = dev
.encode_outbound(
session,
Some(exchange_id),
0x09, ProtocolId::INTERACTION_MODEL,
response,
MrpFlags { reliable: true },
Instant::now(),
)
.unwrap();
dev_io.send_to(&out.wire_bytes, ctrl_addr).await.unwrap();
break;
}
}
};
let (got, ()) = tokio::join!(controller, device);
let got = got.unwrap();
assert_eq!(got.payload, response);
}
fn report_data(ep: u16, cl: u32, at: u32, val: u64, more: bool) -> Vec<u8> {
use matter_codec::{Tag, TlvWriter};
let mut buf = Vec::new();
let mut w = TlvWriter::new(&mut buf);
w.start_structure(Tag::Anonymous).unwrap();
w.start_array(Tag::Context(1)).unwrap(); w.start_structure(Tag::Anonymous).unwrap(); w.start_structure(Tag::Context(1)).unwrap(); w.start_list(Tag::Context(1)).unwrap(); w.put_uint(Tag::Context(2), u64::from(ep)).unwrap();
w.put_uint(Tag::Context(3), u64::from(cl)).unwrap();
w.put_uint(Tag::Context(4), u64::from(at)).unwrap();
w.end_container().unwrap();
w.put_uint(Tag::Context(2), val).unwrap(); w.end_container().unwrap(); w.end_container().unwrap(); w.end_container().unwrap(); if more {
w.put_bool(Tag::Context(3), true).unwrap(); }
w.put_uint(Tag::Context(0xFF), 11).unwrap();
w.end_container().unwrap();
buf
}
#[test]
fn enforce_read_caps_rejects_oversized_single_chunk() {
let err = enforce_read_caps(0, MAX_READ_CHUNK_BYTES + 1, 0).unwrap_err();
match err {
DriverError::ReadTooLarge { limit } => assert_eq!(limit, "MAX_READ_CHUNK_BYTES"),
other => panic!("expected ReadTooLarge(MAX_READ_CHUNK_BYTES), got {other:?}"),
}
}
#[test]
fn enforce_read_caps_accepts_normal_chunk() {
let total = enforce_read_caps(0, 1024, 0).unwrap();
assert_eq!(total, 1024);
let total = enforce_read_caps(1, 512, total).unwrap();
assert_eq!(total, 1536);
}
#[test]
fn enforce_read_caps_rejects_too_many_chunks() {
let err = enforce_read_caps(MAX_READ_CHUNKS, 1, 0).unwrap_err();
match err {
DriverError::ReadTooLarge { limit } => assert_eq!(limit, "MAX_READ_CHUNKS"),
other => panic!("expected MAX_READ_CHUNKS, got {other:?}"),
}
}
#[test]
fn enforce_read_caps_rejects_cumulative_overflow() {
let err = enforce_read_caps(1, 1, MAX_READ_BYTES).unwrap_err();
match err {
DriverError::ReadTooLarge { limit } => assert_eq!(limit, "MAX_READ_BYTES"),
other => panic!("expected MAX_READ_BYTES, got {other:?}"),
}
}
#[tokio::test]
async fn secured_read_reassembles_two_chunks() {
let (mut ctrl, mut dev) = paired_pase_sessions();
let session = matter_transport::SessionId(1);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let controller = secured_read(
&ctrl_io,
&mut ctrl,
session,
dev_addr,
0x02, ProtocolId::INTERACTION_MODEL,
b"readreq",
);
let device = async {
let (pkt, _) = dev_io.recv_from().await.unwrap();
let DecodeInboundOutput::AppMessage { exchange_id, .. } =
dev.decode_inbound(&pkt, Instant::now()).unwrap()
else {
panic!("expected ReadRequest");
};
let c0 = report_data(0, 0x28, 0x0002, 5010, true);
let out = dev
.encode_outbound(
session,
Some(exchange_id),
0x05, ProtocolId::INTERACTION_MODEL,
&c0,
MrpFlags { reliable: true },
Instant::now(),
)
.unwrap();
dev_io.send_to(&out.wire_bytes, ctrl_addr).await.unwrap();
let (ack, _) = dev_io.recv_from().await.unwrap();
let DecodeInboundOutput::AppMessage {
opcode,
exchange_id: ack_exchange,
..
} = dev.decode_inbound(&ack, Instant::now()).unwrap()
else {
panic!("expected StatusResponse");
};
assert_eq!(
opcode, 0x01,
"controller must ack the chunk with StatusResponse"
);
assert_eq!(
ack_exchange, exchange_id,
"StatusResponse must ride the read exchange (enables the chunk-ack piggyback)"
);
let c1 = report_data(1, 0x06, 0x0000, 1, false);
let out = dev
.encode_outbound(
session,
Some(exchange_id),
0x05,
ProtocolId::INTERACTION_MODEL,
&c1,
MrpFlags { reliable: true },
Instant::now(),
)
.unwrap();
dev_io.send_to(&out.wire_bytes, ctrl_addr).await.unwrap();
};
let (got, ()) = tokio::join!(controller, device);
let chunks = got.unwrap();
assert_eq!(chunks.len(), 2, "both chunks returned");
let mut acc = crate::im::ReportAccumulator::new();
for c in &chunks {
acc.push(crate::im::parse_report_data(c).unwrap()).unwrap();
}
let attrs = acc.finish();
assert_eq!(attrs.len(), 2);
assert_eq!(attrs[0].0.endpoint, 0);
assert_eq!(attrs[1].0.endpoint, 1);
}
}