use std::net::SocketAddr;
use matter_cert::{MatterTime, TrustedRoots};
use matter_crypto::{CaseCredentials, CaseInitiator};
use matter_transport::{Discovery, ServiceKind, SessionId, SessionManager, SessionRole};
use crate::driver::datagram::AsyncDatagram;
use crate::driver::error::DriverError;
use crate::driver::unsecured::{parse_status_report, require_handshake_opcode, UnsecuredExchange};
#[must_use]
pub fn operational_instance_name(compressed_fabric_id: [u8; 8], node_id: u64) -> String {
let cfid = u64::from_be_bytes(compressed_fabric_id);
format!("{cfid:016X}-{node_id:016X}")
}
pub(crate) const RESOLVE_POLL_ATTEMPTS: u32 = 300;
const RESOLVE_POLL_INTERVAL: std::time::Duration = std::time::Duration::from_millis(100);
pub const BLE_RESOLVE_POLL_ATTEMPTS: u32 = 600;
pub(crate) fn preferred_address(addresses: &[std::net::IpAddr]) -> Option<std::net::IpAddr> {
let is_v6_link_local = |a: &std::net::IpAddr| match a {
std::net::IpAddr::V6(v6) => (v6.segments()[0] & 0xffc0) == 0xfe80,
std::net::IpAddr::V4(_) => false,
};
addresses
.iter()
.find(|a| a.is_ipv4())
.or_else(|| addresses.iter().find(|a| !is_v6_link_local(a)))
.or_else(|| addresses.first())
.copied()
}
pub async fn resolve_operational<D: Discovery>(
discovery: &mut D,
compressed_fabric_id: [u8; 8],
node_id: u64,
) -> Result<SocketAddr, DriverError> {
resolve_operational_with_attempts(
discovery,
compressed_fabric_id,
node_id,
RESOLVE_POLL_ATTEMPTS,
)
.await
}
pub async fn resolve_operational_with_attempts<D: Discovery>(
discovery: &mut D,
compressed_fabric_id: [u8; 8],
node_id: u64,
attempts: u32,
) -> Result<SocketAddr, DriverError> {
let target = operational_instance_name(compressed_fabric_id, node_id);
let handle = discovery
.query(ServiceKind::Operational)
.map_err(DriverError::Transport)?;
for _ in 0..attempts {
for svc in discovery.poll_results(handle) {
if svc.instance_name.eq_ignore_ascii_case(&target) {
if let Some(addr) = preferred_address(&svc.addresses) {
discovery.stop_query(handle);
return Ok(SocketAddr::new(addr, svc.port));
}
}
}
tokio::time::sleep(RESOLVE_POLL_INTERVAL).await;
}
discovery.stop_query(handle);
Err(DriverError::Discovery(format!(
"operational node {target} not found via mDNS"
)))
}
const OP_SIGMA1: u8 = 0x30;
const OP_SIGMA2: u8 = 0x31;
const OP_SIGMA3: u8 = 0x32;
const OP_STATUS_REPORT: u8 = 0x40;
const CASE_EXCHANGE_ID: u16 = 1;
#[allow(clippy::too_many_arguments)]
pub async fn run_case<T: AsyncDatagram>(
transport: &T,
sessions: &mut SessionManager,
peer: SocketAddr,
credentials: CaseCredentials,
trusted_roots: TrustedRoots,
peer_node_id: u64,
peer_fabric_id: u64,
now: MatterTime,
) -> Result<SessionId, DriverError> {
let local = sessions.allocate_session_id();
let output = Box::pin(run_case_establish(
transport,
peer,
local.0,
credentials,
trusted_roots,
peer_node_id,
peer_fabric_id,
now,
))
.await?;
let sid = sessions.register_case(&output, SessionRole::Initiator);
Ok(sid)
}
#[allow(clippy::too_many_arguments)]
pub async fn run_case_establish<T: AsyncDatagram>(
transport: &T,
peer: SocketAddr,
local_session_id: u16,
credentials: CaseCredentials,
trusted_roots: TrustedRoots,
peer_node_id: u64,
peer_fabric_id: u64,
now: MatterTime,
) -> Result<matter_crypto::CaseSessionOutput, DriverError> {
let mut initiator = CaseInitiator::new(
credentials,
trusted_roots,
peer_node_id,
peer_fabric_id,
local_session_id,
now,
)?;
let mut exch = UnsecuredExchange::new_ephemeral(CASE_EXCHANGE_ID)?;
let sigma1 = initiator.start()?;
let sigma2 = exch
.send_and_recv(transport, peer, OP_SIGMA1, OP_SIGMA2, &sigma1, None)
.await?;
if let Err(e) = require_handshake_opcode(&sigma2, OP_SIGMA2) {
let _ = exch
.send_standalone_ack(transport, peer, sigma2.message_counter)
.await;
return Err(e);
}
#[cfg(feature = "tracing")]
tracing::debug!(
sigma2 = %crate::hexdump::hex(&sigma2.payload),
"received Sigma2"
);
initiator.handle_sigma2(&sigma2.payload)?;
let sigma3 = initiator.next_message()?;
let report = exch
.send_and_recv(
transport,
peer,
OP_SIGMA3,
OP_STATUS_REPORT,
&sigma3,
Some(sigma2.message_counter),
)
.await?;
let status = parse_status_report(&report)?;
exch.send_standalone_ack(transport, peer, report.message_counter)
.await?;
if !status.is_session_establishment_success() {
return Err(DriverError::SessionEstablishmentFailed {
general_code: status.general_code,
protocol_code: status.protocol_code,
});
}
let output = initiator.finish()?;
Ok(output)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)] mod tests {
use super::*;
#[test]
fn operational_instance_name_formats_16_16_uppercase_hex() {
let cfid = [0x87, 0xe1, 0xb0, 0x04, 0xe2, 0x35, 0xa1, 0x30];
let node_id: u64 = 0x0000_0000_0000_0001;
assert_eq!(
operational_instance_name(cfid, node_id),
"87E1B004E235A130-0000000000000001"
);
}
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr};
use matter_transport::{MatterService, QueryHandle};
struct FakeDiscovery {
service: MatterService,
}
impl Discovery for FakeDiscovery {
fn publish(&mut self, _s: &MatterService) -> matter_transport::Result<()> {
Ok(())
}
fn unpublish(&mut self, _n: &str, _k: ServiceKind) -> matter_transport::Result<()> {
Ok(())
}
fn query(&mut self, _k: ServiceKind) -> matter_transport::Result<QueryHandle> {
Ok(QueryHandle(1))
}
fn stop_query(&mut self, _h: QueryHandle) {}
fn poll_results(&mut self, _h: QueryHandle) -> Vec<MatterService> {
vec![self.service.clone()]
}
}
#[tokio::test]
async fn resolve_operational_returns_matching_addr() {
let cfid = [0x87, 0xe1, 0xb0, 0x04, 0xe2, 0x35, 0xa1, 0x30];
let node_id: u64 = 1;
let name = operational_instance_name(cfid, node_id);
let mut disc = FakeDiscovery {
service: MatterService::new(
name,
ServiceKind::Operational,
vec![IpAddr::V4(Ipv4Addr::new(192, 0, 2, 7))],
5540,
HashMap::new(),
),
};
let addr = resolve_operational(&mut disc, cfid, node_id).await.unwrap();
assert_eq!(
addr,
std::net::SocketAddr::new(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 7)), 5540)
);
}
#[tokio::test]
async fn resolve_operational_prefers_routable_addresses() {
use std::net::Ipv6Addr;
let cfid = [0x87, 0xe1, 0xb0, 0x04, 0xe2, 0x35, 0xa1, 0x30];
let node_id: u64 = 1;
let name = operational_instance_name(cfid, node_id);
let link_local = IpAddr::V6(Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 0x1d42));
let ula = IpAddr::V6(Ipv6Addr::new(0xfdfc, 0x20da, 0x4273, 0x126f, 0, 0, 0, 1));
let v4 = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 248));
let mut disc = FakeDiscovery {
service: MatterService::new(
name.clone(),
ServiceKind::Operational,
vec![link_local, ula, v4],
5540,
HashMap::new(),
),
};
let addr = resolve_operational(&mut disc, cfid, node_id).await.unwrap();
assert_eq!(addr, std::net::SocketAddr::new(v4, 5540));
let mut disc = FakeDiscovery {
service: MatterService::new(
name,
ServiceKind::Operational,
vec![link_local, ula],
5540,
HashMap::new(),
),
};
let addr = resolve_operational(&mut disc, cfid, node_id).await.unwrap();
assert_eq!(addr, std::net::SocketAddr::new(ula, 5540));
}
struct CountingDiscovery {
service: MatterService,
succeed_on: u32,
polls: u32,
}
impl Discovery for CountingDiscovery {
fn publish(&mut self, _s: &MatterService) -> matter_transport::Result<()> {
Ok(())
}
fn unpublish(&mut self, _n: &str, _k: ServiceKind) -> matter_transport::Result<()> {
Ok(())
}
fn query(&mut self, _k: ServiceKind) -> matter_transport::Result<QueryHandle> {
Ok(QueryHandle(1))
}
fn stop_query(&mut self, _h: QueryHandle) {}
fn poll_results(&mut self, _h: QueryHandle) -> Vec<MatterService> {
self.polls += 1;
if self.polls >= self.succeed_on {
vec![self.service.clone()]
} else {
vec![]
}
}
}
fn counting_discovery(cfid: [u8; 8], node_id: u64, succeed_on: u32) -> CountingDiscovery {
CountingDiscovery {
service: MatterService::new(
operational_instance_name(cfid, node_id),
ServiceKind::Operational,
vec![IpAddr::V4(Ipv4Addr::new(192, 0, 2, 9))],
5540,
HashMap::new(),
),
succeed_on,
polls: 0,
}
}
#[tokio::test(start_paused = true)]
async fn resolve_operational_with_attempts_succeeds_within_budget() {
let cfid = [0x87, 0xe1, 0xb0, 0x04, 0xe2, 0x35, 0xa1, 0x30];
let node_id: u64 = 1;
let mut disc = counting_discovery(cfid, node_id, 3);
let addr = resolve_operational_with_attempts(&mut disc, cfid, node_id, 3)
.await
.unwrap();
assert_eq!(
addr,
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 9)), 5540)
);
assert_eq!(disc.polls, 3, "must poll exactly until the record appears");
}
#[tokio::test(start_paused = true)]
async fn resolve_operational_with_attempts_respects_budget() {
let cfid = [0x87, 0xe1, 0xb0, 0x04, 0xe2, 0x35, 0xa1, 0x30];
let node_id: u64 = 1;
let mut disc = counting_discovery(cfid, node_id, 5);
let err = resolve_operational_with_attempts(&mut disc, cfid, node_id, 4)
.await
.expect_err("a 4-poll budget must give up before the 5th-poll record");
assert!(matches!(err, DriverError::Discovery(_)), "got {err:?}");
assert_eq!(disc.polls, 4, "must poll exactly the budget then stop");
}
#[tokio::test]
async fn resolve_operational_still_resolves_via_delegation() {
let cfid = [0x87, 0xe1, 0xb0, 0x04, 0xe2, 0x35, 0xa1, 0x30];
let node_id: u64 = 1;
let mut disc = counting_discovery(cfid, node_id, 1);
let addr = resolve_operational(&mut disc, cfid, node_id).await.unwrap();
assert_eq!(
addr,
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 9)), 5540)
);
}
use matter_cert::test_support::{build_unsigned, with_signature, TestCertFields};
use matter_cert::{
BasicConstraints, DistinguishedName, DnAttribute, Extensions, KeyIdentifier, KeyUsage,
MatterCertificate, MatterTime, PublicKey, Signature, TrustAnchor, TrustedRoots,
};
use matter_crypto::{CaseCredentials, CaseResponder, CaseSigner, RingSigner, Sigma1Outcome};
use matter_transport::{SessionKeys, SessionManager};
use crate::driver::datagram::InMemoryDatagram;
use crate::driver::unsecured::{decode_unsecured, encode_unsecured};
const T_FABRIC_ID: u64 = 0x4242_4242_4242_4242;
const T_INITIATOR_NODE: u64 = 0xDEAD_BEEF_CAFE_F00D;
const T_RESPONDER_NODE: u64 = 0xBABE_FEED_1234_5678;
const T_IPK: [u8; 16] = [0x77; 16];
const T_RCAC_SKI: [u8; 20] = [0x01; 20];
const T_NOC_SKI: [u8; 20] = [0x02; 20];
fn build_test_rcac() -> (MatterCertificate, RingSigner, [u8; 65]) {
let (rcac_signer, _pkcs8) = RingSigner::generate().unwrap();
let rcac_pub = *rcac_signer.public_key().as_bytes();
let rcac_dn = DistinguishedName::new(vec![DnAttribute::RcacId(1)]);
let extensions = Extensions::builder()
.basic_constraints(Some(BasicConstraints::new(true, Some(1))))
.key_usage(Some(KeyUsage::KEY_CERT_SIGN))
.subject_key_identifier(Some(KeyIdentifier(T_RCAC_SKI)))
.authority_key_identifier(Some(KeyIdentifier(T_RCAC_SKI)))
.build();
let fields = TestCertFields {
serial: vec![0x01],
issuer: rcac_dn.clone(),
not_before: MatterTime::from_unix_secs(1_700_000_000),
not_after: MatterTime::from_unix_secs(2_500_000_000),
subject: rcac_dn,
public_key: PublicKey::new(rcac_pub).unwrap(),
extensions,
signature: Signature::new([0u8; 64]),
};
let unsigned = build_unsigned(fields);
let tbs = unsigned.to_x509_tbs_der().unwrap();
let sig = rcac_signer.sign_p256_sha256(&tbs).unwrap();
let rcac = with_signature(&unsigned, Signature::new(sig));
(rcac, rcac_signer, rcac_pub)
}
fn roots_for(rcac: &MatterCertificate) -> TrustedRoots {
let mut roots = TrustedRoots::new();
roots.add(TrustAnchor::from_root_cert(rcac));
roots
}
fn build_test_noc(rcac_signer: &RingSigner, node_id: u64) -> (MatterCertificate, RingSigner) {
let (noc_signer, _) = RingSigner::generate().unwrap();
let noc_pub = *noc_signer.public_key().as_bytes();
let subject_dn = DistinguishedName::new(vec![
DnAttribute::FabricId(T_FABRIC_ID),
DnAttribute::NodeId(node_id),
]);
let issuer_dn = DistinguishedName::new(vec![DnAttribute::RcacId(1)]);
let extensions = Extensions::builder()
.basic_constraints(Some(BasicConstraints::new(false, None)))
.key_usage(Some(KeyUsage::DIGITAL_SIGNATURE))
.subject_key_identifier(Some(KeyIdentifier(T_NOC_SKI)))
.authority_key_identifier(Some(KeyIdentifier(T_RCAC_SKI)))
.build();
let fields = TestCertFields {
serial: vec![0x02],
issuer: issuer_dn,
not_before: MatterTime::from_unix_secs(1_700_000_000),
not_after: MatterTime::from_unix_secs(2_500_000_000),
subject: subject_dn,
public_key: PublicKey::new(noc_pub).unwrap(),
extensions,
signature: Signature::new([0u8; 64]),
};
let unsigned = build_unsigned(fields);
let tbs = unsigned.to_x509_tbs_der().unwrap();
let sig = rcac_signer.sign_p256_sha256(&tbs).unwrap();
let noc = with_signature(&unsigned, Signature::new(sig));
(noc, noc_signer)
}
fn creds(
noc: MatterCertificate,
signer: RingSigner,
node_id: u64,
rcac_pub: [u8; 65],
) -> CaseCredentials {
CaseCredentials {
noc,
icac: None,
signer: Box::new(signer),
fabric_id: T_FABRIC_ID,
node_id,
ipk: T_IPK,
rcac_public_key: rcac_pub,
}
}
#[tokio::test]
async fn run_case_surfaces_sigma1_status_report_rejection() {
let (rcac, rcac_signer, rcac_pub) = build_test_rcac();
let (init_noc, init_signer) = build_test_noc(&rcac_signer, T_INITIATOR_NODE);
let init_creds = creds(init_noc, init_signer, T_INITIATOR_NODE, rcac_pub);
let ctrl_roots = roots_for(&rcac);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut sessions = SessionManager::new();
let device = async {
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
let mut body = Vec::new();
body.extend_from_slice(&1u16.to_le_bytes());
body.extend_from_slice(&0u32.to_le_bytes());
body.extend_from_slice(&0x0001u16.to_le_bytes());
let report = encode_unsecured(
200,
m.exchange_id,
0x40,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&body,
);
dev_io.send_to(&report, ctrl_addr).await.unwrap();
};
let controller = run_case(
&ctrl_io,
&mut sessions,
dev_addr,
init_creds,
ctrl_roots,
T_RESPONDER_NODE,
T_FABRIC_ID,
MatterTime::from_unix_secs(2_000_000_000),
);
let (ctrl_result, ()) = tokio::join!(controller, device);
let err = ctrl_result.unwrap_err();
assert!(
matches!(
err,
DriverError::SessionEstablishmentFailed {
general_code: 1,
protocol_code: 0x0001,
}
),
"expected SessionEstablishmentFailed, got: {err:?}"
);
}
#[tokio::test]
async fn run_case_establishes_matching_session() {
let (rcac, rcac_signer, rcac_pub) = build_test_rcac();
let (init_noc, init_signer) = build_test_noc(&rcac_signer, T_INITIATOR_NODE);
let (resp_noc, resp_signer) = build_test_noc(&rcac_signer, T_RESPONDER_NODE);
let init_creds = creds(init_noc, init_signer, T_INITIATOR_NODE, rcac_pub);
let resp_creds = creds(resp_noc, resp_signer, T_RESPONDER_NODE, rcac_pub);
let resp_roots = roots_for(&rcac);
let ctrl_roots = roots_for(&rcac);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let mut sessions = SessionManager::new();
let device = async {
let mut responder = CaseResponder::new(
resp_creds,
resp_roots,
0x00D2,
MatterTime::from_unix_secs(2_000_000_000),
)
.unwrap();
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
assert!(matches!(
responder.handle_sigma1(&m.payload).unwrap(),
Sigma1Outcome::NewSession
));
let sigma2 = responder.next_message().unwrap();
let wire = encode_unsecured(
200,
m.exchange_id,
0x31,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&sigma2,
);
dev_io.send_to(&wire, ctrl_addr).await.unwrap();
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
responder.handle_sigma3(&m.payload).unwrap();
let mut body = Vec::new();
body.extend_from_slice(&0u16.to_le_bytes());
body.extend_from_slice(&0u32.to_le_bytes());
body.extend_from_slice(&0u16.to_le_bytes());
let report = encode_unsecured(
201,
m.exchange_id,
0x40,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&body,
);
dev_io.send_to(&report, ctrl_addr).await.unwrap();
let ack = tokio::time::timeout(std::time::Duration::from_secs(2), dev_io.recv_from())
.await
.expect("controller must ack the StatusReport")
.unwrap();
let ack = decode_unsecured(&ack.0).unwrap();
assert_eq!(ack.opcode, 0x10);
assert_eq!(ack.ack_counter, Some(201));
responder.finish().unwrap()
};
let controller = run_case(
&ctrl_io,
&mut sessions,
dev_addr,
init_creds,
ctrl_roots,
T_RESPONDER_NODE,
T_FABRIC_ID,
MatterTime::from_unix_secs(2_000_000_000),
);
let (ctrl_result, dev_out) = tokio::join!(controller, device);
let sid = ctrl_result.unwrap();
let registered = sessions.get(sid).unwrap();
assert_eq!(registered.keys, SessionKeys::from_case_output(&dev_out));
assert_eq!(registered.peer_id, matter_transport::SessionId(0x00D2));
}
#[tokio::test]
async fn run_case_establish_returns_registerable_output() {
let (rcac, rcac_signer, rcac_pub) = build_test_rcac();
let (init_noc, init_signer) = build_test_noc(&rcac_signer, T_INITIATOR_NODE);
let (resp_noc, resp_signer) = build_test_noc(&rcac_signer, T_RESPONDER_NODE);
let init_creds = creds(init_noc, init_signer, T_INITIATOR_NODE, rcac_pub);
let resp_creds = creds(resp_noc, resp_signer, T_RESPONDER_NODE, rcac_pub);
let resp_roots = roots_for(&rcac);
let ctrl_roots = roots_for(&rcac);
let (ctrl_io, dev_io) = InMemoryDatagram::pair();
let dev_addr = dev_io.local_addr();
let ctrl_addr = ctrl_io.local_addr();
let local_session_id: u16 = 0x0777;
let device = async {
let mut responder = CaseResponder::new(
resp_creds,
resp_roots,
0x00D2,
MatterTime::from_unix_secs(2_000_000_000),
)
.unwrap();
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
assert!(matches!(
responder.handle_sigma1(&m.payload).unwrap(),
Sigma1Outcome::NewSession
));
let sigma2 = responder.next_message().unwrap();
let wire = encode_unsecured(
200,
m.exchange_id,
0x31,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&sigma2,
);
dev_io.send_to(&wire, ctrl_addr).await.unwrap();
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
responder.handle_sigma3(&m.payload).unwrap();
let mut body = Vec::new();
body.extend_from_slice(&0u16.to_le_bytes());
body.extend_from_slice(&0u32.to_le_bytes());
body.extend_from_slice(&0u16.to_le_bytes());
let report = encode_unsecured(
201,
m.exchange_id,
0x40,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&body,
);
dev_io.send_to(&report, ctrl_addr).await.unwrap();
let ack = tokio::time::timeout(std::time::Duration::from_secs(2), dev_io.recv_from())
.await
.expect("controller must ack the StatusReport")
.unwrap();
let ack = decode_unsecured(&ack.0).unwrap();
assert_eq!(ack.opcode, 0x10);
assert_eq!(ack.ack_counter, Some(201));
responder.finish().unwrap()
};
let controller = run_case_establish(
&ctrl_io,
dev_addr,
local_session_id,
init_creds,
ctrl_roots,
T_RESPONDER_NODE,
T_FABRIC_ID,
MatterTime::from_unix_secs(2_000_000_000),
);
let (ctrl_result, dev_out) = tokio::join!(controller, device);
let output = ctrl_result.unwrap();
assert_eq!(output.local.session_id, local_session_id);
let mut sessions = SessionManager::new();
let sid = sessions.register_case(&output, SessionRole::Initiator);
let registered = sessions.get(sid).unwrap();
assert_eq!(registered.keys, SessionKeys::from_case_output(&dev_out));
assert_eq!(registered.peer_id, matter_transport::SessionId(0x00D2));
}
}