use std::net::SocketAddr;
use matter_crypto::pase::PaseProver;
use matter_transport::{PeerHint, SessionId, SessionManager, SessionRole};
use crate::driver::datagram::AsyncDatagram;
use crate::driver::error::DriverError;
use crate::driver::unsecured::{parse_status_report, require_handshake_opcode, UnsecuredExchange};
use crate::driver::TransportReliability;
const OP_PBKDF_PARAM_REQUEST: u8 = 0x20;
const OP_PBKDF_PARAM_RESPONSE: u8 = 0x21;
const OP_PASE_PAKE1: u8 = 0x22;
const OP_PASE_PAKE2: u8 = 0x23;
const OP_PASE_PAKE3: u8 = 0x24;
const OP_STATUS_REPORT: u8 = 0x40;
const PASE_EXCHANGE_ID: u16 = 1;
pub async fn run_pase<T: AsyncDatagram>(
transport: &T,
sessions: &mut SessionManager,
peer: SocketAddr,
passcode: u32,
) -> Result<SessionId, DriverError> {
run_pase_with(
transport,
sessions,
peer,
passcode,
TransportReliability::Mrp,
)
.await
}
pub async fn run_pase_with<T: AsyncDatagram>(
transport: &T,
sessions: &mut SessionManager,
peer: SocketAddr,
passcode: u32,
reliability: TransportReliability,
) -> Result<SessionId, DriverError> {
let local = sessions.allocate_session_id();
let mut prover = PaseProver::new_with_negotiation(passcode, local.0)?;
let mut exch = UnsecuredExchange::new_ephemeral_with(PASE_EXCHANGE_ID, reliability)?;
let request = prover.start()?;
let resp = exch
.send_and_recv(
transport,
peer,
OP_PBKDF_PARAM_REQUEST,
OP_PBKDF_PARAM_RESPONSE,
&request,
None,
)
.await?;
if let Err(e) = require_handshake_opcode(&resp, OP_PBKDF_PARAM_RESPONSE) {
let _ = exch
.send_standalone_ack(transport, peer, resp.message_counter)
.await;
return Err(e);
}
prover.handle_pbkdf_response(&resp.payload)?;
let pake1 = prover.next_message()?;
let pake2 = exch
.send_and_recv(
transport,
peer,
OP_PASE_PAKE1,
OP_PASE_PAKE2,
&pake1,
Some(resp.message_counter),
)
.await?;
if let Err(e) = require_handshake_opcode(&pake2, OP_PASE_PAKE2) {
let _ = exch
.send_standalone_ack(transport, peer, pake2.message_counter)
.await;
return Err(e);
}
prover.handle_pake2(&pake2.payload)?;
let pake3 = prover.next_message()?;
let report = exch
.send_and_recv(
transport,
peer,
OP_PASE_PAKE3,
OP_STATUS_REPORT,
&pake3,
Some(pake2.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 peer_session_id = prover.responder_session_id().ok_or(DriverError::Handshake(
"PASE negotiation produced no responder session id",
))?;
let keys = prover.finish()?;
sessions.register_pase_with_local_id(
local,
keys,
SessionRole::Initiator,
peer_session_id,
PeerHint::default(),
);
if reliability == TransportReliability::TransportProvides {
sessions.set_transport_reliable(local, true)?;
}
Ok(local)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)] mod tests {
use matter_crypto::pase::{PasePbkdfParams, PaseVerifier};
use matter_transport::{SessionKeys, SessionManager};
use super::*;
use crate::driver::datagram::{AsyncDatagram, InMemoryDatagram};
use crate::driver::unsecured::{decode_unsecured, encode_unsecured};
const OP_PBKDF_RESP: u8 = 0x21;
const OP_PAKE2: u8 = 0x23;
const OP_STANDALONE_ACK: u8 = 0x10;
fn status_report_body(general: u16, protocol_code: u16) -> Vec<u8> {
let mut b = Vec::with_capacity(8);
b.extend_from_slice(&general.to_le_bytes());
b.extend_from_slice(&0u32.to_le_bytes()); b.extend_from_slice(&protocol_code.to_le_bytes());
b
}
#[tokio::test]
async fn run_pase_establishes_matching_session() {
let pin = 20_202_021;
let params = PasePbkdfParams {
iterations: 1000,
salt: vec![0x55; 16],
};
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 verifier = PaseVerifier::new_from_pin(pin, params, 0x00BB).unwrap();
let mut ctr: u32 = 100;
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
verifier.handle_pbkdf_request(&m.payload).unwrap();
let resp = verifier.next_message().unwrap();
let wire = encode_unsecured(
ctr,
m.exchange_id,
OP_PBKDF_RESP,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&resp,
);
ctr += 1;
dev_io.send_to(&wire, ctrl_addr).await.unwrap();
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
verifier.handle_pake1(&m.payload).unwrap();
let pake2 = verifier.next_message().unwrap();
let wire = encode_unsecured(
ctr,
m.exchange_id,
OP_PAKE2,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&pake2,
);
dev_io.send_to(&wire, ctrl_addr).await.unwrap();
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
verifier.handle_pake3(&m.payload).unwrap();
ctr += 1;
let report = encode_unsecured(
ctr,
m.exchange_id,
OP_STATUS_REPORT,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&status_report_body(0, 0),
);
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, OP_STANDALONE_ACK);
assert_eq!(ack.ack_counter, Some(ctr));
verifier.finish().unwrap()
};
let controller = run_pase(&ctrl_io, &mut sessions, dev_addr, pin);
let (ctrl_result, dev_keys) = tokio::join!(controller, device);
let sid = ctrl_result.unwrap();
let registered = sessions.get(sid).unwrap();
assert_eq!(registered.keys, SessionKeys::from(dev_keys));
assert_eq!(registered.peer_id, matter_transport::SessionId(0x00BB));
assert_eq!(sessions.is_transport_reliable(sid), Some(false));
}
#[tokio::test]
async fn run_pase_with_marks_session_reliable() {
let pin = 20_202_021;
let params = PasePbkdfParams {
iterations: 1000,
salt: vec![0x55; 16],
};
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 verifier = PaseVerifier::new_from_pin(pin, params, 0x00BB).unwrap();
let mut ctr: u32 = 100;
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
verifier.handle_pbkdf_request(&m.payload).unwrap();
let resp = verifier.next_message().unwrap();
dev_io
.send_to(
&encode_unsecured(
ctr,
m.exchange_id,
OP_PBKDF_RESP,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&resp,
),
ctrl_addr,
)
.await
.unwrap();
ctr += 1;
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
verifier.handle_pake1(&m.payload).unwrap();
let pake2 = verifier.next_message().unwrap();
dev_io
.send_to(
&encode_unsecured(
ctr,
m.exchange_id,
OP_PAKE2,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&pake2,
),
ctrl_addr,
)
.await
.unwrap();
ctr += 1;
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
verifier.handle_pake3(&m.payload).unwrap();
dev_io
.send_to(
&encode_unsecured(
ctr,
m.exchange_id,
OP_STATUS_REPORT,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&status_report_body(0, 0),
),
ctrl_addr,
)
.await
.unwrap();
verifier.finish().unwrap()
};
let controller = run_pase_with(
&ctrl_io,
&mut sessions,
dev_addr,
pin,
TransportReliability::TransportProvides,
);
let (ctrl_result, _dev_keys) = tokio::join!(controller, device);
let sid = ctrl_result.unwrap();
assert_eq!(
sessions.is_transport_reliable(sid),
Some(true),
"a TransportProvides PASE must flag the session transport_reliable"
);
}
#[tokio::test]
async fn run_pase_surfaces_status_report_rejection() {
let pin = 20_202_021;
let params = PasePbkdfParams {
iterations: 1000,
salt: vec![0x55; 16],
};
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 verifier = PaseVerifier::new_from_pin(pin, params, 0x00BB).unwrap();
let mut ctr: u32 = 100;
for op in [OP_PBKDF_RESP, OP_PAKE2] {
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
match op {
OP_PBKDF_RESP => verifier.handle_pbkdf_request(&m.payload).unwrap(),
_ => verifier.handle_pake1(&m.payload).unwrap(),
}
let reply = verifier.next_message().unwrap();
let wire = encode_unsecured(
ctr,
m.exchange_id,
op,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&reply,
);
ctr += 1;
dev_io.send_to(&wire, ctrl_addr).await.unwrap();
}
let (p, _) = dev_io.recv_from().await.unwrap();
let m = decode_unsecured(&p).unwrap();
let report = encode_unsecured(
ctr,
m.exchange_id,
OP_STATUS_REPORT,
matter_transport::ProtocolId::SECURE_CHANNEL,
false,
true,
Some(m.message_counter),
None,
&status_report_body(1, 0x0002),
);
dev_io.send_to(&report, ctrl_addr).await.unwrap();
};
let controller = run_pase(&ctrl_io, &mut sessions, dev_addr, pin);
let (ctrl_result, ()) = tokio::join!(controller, device);
let err = ctrl_result.unwrap_err();
assert!(
matches!(
err,
DriverError::SessionEstablishmentFailed {
general_code: 1,
protocol_code: 0x0002,
}
),
"expected SessionEstablishmentFailed, got: {err:?}"
);
}
}