#![forbid(unsafe_code)]
mod base38;
mod manual_packer;
mod qr_packer;
mod verhoeff;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SetupPayload {
pub version: u8,
pub vendor_id: Option<u16>,
pub product_id: Option<u16>,
pub commissioning_flow: CommissioningFlow,
pub discovery_capabilities: DiscoveryCapabilities,
pub discriminator: Discriminator,
pub passcode: Passcode,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Discriminator(u16);
impl Discriminator {
pub const fn new(value: u16) -> Result<Self> {
if value > 0x0FFF {
Err(Error::DiscriminatorOutOfRange(value))
} else {
Ok(Self(value))
}
}
pub const fn as_u16(self) -> u16 {
self.0
}
pub const fn short(self) -> u8 {
((self.0 >> 8) & 0x0F) as u8
}
}
pub const DISALLOWED_PASSCODES: &[u32] = &[
0, 11_111_111, 22_222_222, 33_333_333, 44_444_444, 55_555_555, 66_666_666, 77_777_777,
88_888_888, 99_999_999, 12_345_678, 87_654_321,
];
pub const MAX_PASSCODE: u32 = 99_999_998;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Passcode(u32);
impl Passcode {
pub fn new(value: u32) -> Result<Self> {
if DISALLOWED_PASSCODES.contains(&value) {
return Err(Error::PasscodeDisallowedTrivial(value));
}
if value > MAX_PASSCODE {
return Err(Error::PasscodeOutOfRange(value));
}
Ok(Self(value))
}
pub const fn as_u32(self) -> u32 {
self.0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum CommissioningFlow {
Standard,
UserIntent,
Custom,
}
impl CommissioningFlow {
pub const fn from_u8(value: u8) -> Result<Self> {
match value {
0 => Ok(Self::Standard),
1 => Ok(Self::UserIntent),
2 => Ok(Self::Custom),
other => Err(Error::CommissioningFlowReserved(other)),
}
}
pub const fn as_u8(self) -> u8 {
match self {
Self::Standard => 0,
Self::UserIntent => 1,
Self::Custom => 2,
}
}
}
bitflags::bitflags! {
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct DiscoveryCapabilities: u8 {
const SOFT_AP = 0b0000_0001;
const BLE = 0b0000_0010;
const ON_NETWORK = 0b0000_0100;
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum Error {
#[error("QR string is missing the `MT:` prefix")]
MissingMtPrefix,
#[error("invalid Base38 character `{0}` at position {1}")]
InvalidBase38Char(char, usize),
#[error("Base38 chunk at position {position} decodes to an out-of-range value")]
Base38ChunkOutOfRange {
position: usize,
},
#[error("QR payload is the wrong length: {got} bytes, expected exactly {need}")]
QrPayloadWrongLength {
got: usize,
need: usize,
},
#[error("QR payload has {extra} byte(s) after the fixed 11-byte block; vendor TLV blobs are not supported in this release")]
QrTrailingBytes {
extra: usize,
},
#[error("manual code must be 11 or 21 digits; got {0}")]
ManualCodeWrongLength(usize),
#[error("manual code contains non-digit `{0}` at position {1}")]
ManualCodeNonDigit(char, usize),
#[error("manual code Verhoeff check digit failed")]
ManualCodeBadChecksum,
#[error("manual code {field} value {value} exceeds the 16-bit field width")]
FieldOutOfRange {
field: &'static str,
value: u32,
},
#[error("discriminator {0} exceeds the 12-bit field width")]
DiscriminatorOutOfRange(u16),
#[error("passcode {0} exceeds the 27-bit field width")]
PasscodeOutOfRange(u32),
#[error("passcode {0} is in the disallowed-trivial list (spec §5.1.7.1)")]
PasscodeDisallowedTrivial(u32),
#[error("commissioning flow value {0} is reserved")]
CommissioningFlowReserved(u8),
#[error("QR-form payload requires both vendor_id and product_id to be present")]
QrRequiresVidPid,
#[error("commissioning flow `Custom` requires vendor-specific QR fields not supported by matter-rust")]
CustomFlowUnsupported,
}
pub type Result<T> = core::result::Result<T, Error>;
const QR_PREFIX: &str = "MT:";
pub fn encode_qr(payload: &SetupPayload) -> Result<String> {
let bytes = qr_packer::pack(payload)?;
Ok(format!("{QR_PREFIX}{}", base38::encode(&bytes)))
}
pub fn parse_qr(s: &str) -> Result<SetupPayload> {
let payload = s.strip_prefix(QR_PREFIX).ok_or(Error::MissingMtPrefix)?;
let bytes = base38::decode(payload)?;
let need = qr_packer::FIXED_BYTE_LEN;
if bytes.len() < need {
return Err(Error::QrPayloadWrongLength {
got: bytes.len(),
need,
});
}
if bytes.len() > need {
return Err(Error::QrTrailingBytes {
extra: bytes.len() - need,
});
}
let mut fixed = [0u8; qr_packer::FIXED_BYTE_LEN];
fixed.copy_from_slice(&bytes[..need]);
qr_packer::unpack(&fixed)
}
pub fn encode_manual_code(payload: &SetupPayload) -> String {
manual_packer::pack(payload)
}
pub fn parse_manual_code(s: &str) -> Result<SetupPayload> {
manual_packer::unpack(s)
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod error_tests {
use super::Error;
#[test]
fn display_missing_mt_prefix() {
assert_eq!(
Error::MissingMtPrefix.to_string(),
"QR string is missing the `MT:` prefix"
);
}
#[test]
fn display_invalid_base38_char() {
assert_eq!(
Error::InvalidBase38Char('?', 7).to_string(),
"invalid Base38 character `?` at position 7"
);
}
#[test]
fn display_qr_trailing_bytes() {
assert_eq!(
Error::QrTrailingBytes { extra: 3 }.to_string(),
"QR payload has 3 byte(s) after the fixed 11-byte block; vendor TLV blobs are not supported in this release"
);
}
#[test]
fn display_manual_bad_checksum() {
assert_eq!(
Error::ManualCodeBadChecksum.to_string(),
"manual code Verhoeff check digit failed"
);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod discriminator_tests {
use super::{Discriminator, Error};
#[test]
fn new_accepts_zero() {
let d = Discriminator::new(0).unwrap();
assert_eq!(d.as_u16(), 0);
assert_eq!(d.short(), 0);
}
#[test]
fn new_accepts_max_12_bit() {
let d = Discriminator::new(0x0FFF).unwrap();
assert_eq!(d.as_u16(), 0x0FFF);
assert_eq!(d.short(), 0x0F);
}
#[test]
fn new_rejects_13_bit() {
let err = Discriminator::new(0x1000).unwrap_err();
assert!(matches!(err, Error::DiscriminatorOutOfRange(0x1000)));
}
#[test]
fn short_is_upper_4_bits() {
let d = Discriminator::new(0x0ABC).unwrap();
assert_eq!(d.short(), 0xA);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod passcode_tests {
use super::{Error, Passcode};
#[test]
fn new_accepts_normal_value() {
let p = Passcode::new(20_202_021).unwrap();
assert_eq!(p.as_u32(), 20_202_021);
}
#[test]
fn new_rejects_28_bit_value() {
let too_large = 1u32 << 27;
let err = Passcode::new(too_large).unwrap_err();
assert!(matches!(err, Error::PasscodeOutOfRange(v) if v == too_large));
}
#[test]
fn new_accepts_high_valid_value() {
let p = Passcode::new(99_000_001).unwrap();
assert_eq!(p.as_u32(), 99_000_001);
}
#[test]
fn new_accepts_max_passcode() {
let p = Passcode::new(super::MAX_PASSCODE).unwrap();
assert_eq!(p.as_u32(), 99_999_998);
}
#[test]
fn new_rejects_values_above_max_but_below_2_27() {
for &v in &[100_000_000_u32, 102_950_749, (1 << 27) - 1] {
let err = Passcode::new(v).unwrap_err();
assert!(
matches!(err, Error::PasscodeOutOfRange(x) if x == v),
"expected PasscodeOutOfRange for {v}, got {err:?}"
);
}
}
#[test]
fn new_rejects_all_zeros() {
let err = Passcode::new(0).unwrap_err();
assert!(matches!(err, Error::PasscodeDisallowedTrivial(0)));
}
#[test]
fn new_rejects_all_ones() {
let err = Passcode::new(11_111_111).unwrap_err();
assert!(matches!(err, Error::PasscodeDisallowedTrivial(11_111_111)));
}
#[test]
fn new_rejects_counting_up() {
let err = Passcode::new(12_345_678).unwrap_err();
assert!(matches!(err, Error::PasscodeDisallowedTrivial(12_345_678)));
}
#[test]
fn new_rejects_counting_down() {
let err = Passcode::new(87_654_321).unwrap_err();
assert!(matches!(err, Error::PasscodeDisallowedTrivial(87_654_321)));
}
#[test]
fn new_rejects_all_disallowed() {
for &v in super::DISALLOWED_PASSCODES {
let err = Passcode::new(v).unwrap_err();
assert!(
matches!(err, Error::PasscodeDisallowedTrivial(x) if x == v),
"expected DisallowedTrivial for {v}, got {err:?}"
);
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod commissioning_flow_tests {
use super::{CommissioningFlow, Error};
#[test]
fn from_u8_standard() {
assert_eq!(
CommissioningFlow::from_u8(0).unwrap(),
CommissioningFlow::Standard
);
}
#[test]
fn from_u8_user_intent() {
assert_eq!(
CommissioningFlow::from_u8(1).unwrap(),
CommissioningFlow::UserIntent
);
}
#[test]
fn from_u8_custom() {
assert_eq!(
CommissioningFlow::from_u8(2).unwrap(),
CommissioningFlow::Custom
);
}
#[test]
fn from_u8_reserved() {
let err = CommissioningFlow::from_u8(3).unwrap_err();
assert!(matches!(err, Error::CommissioningFlowReserved(3)));
}
#[test]
fn from_u8_out_of_range() {
let err = CommissioningFlow::from_u8(99).unwrap_err();
assert!(matches!(err, Error::CommissioningFlowReserved(99)));
}
#[test]
fn as_u8_roundtrip() {
assert_eq!(CommissioningFlow::Standard.as_u8(), 0);
assert_eq!(CommissioningFlow::UserIntent.as_u8(), 1);
assert_eq!(CommissioningFlow::Custom.as_u8(), 2);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod discovery_capabilities_tests {
use super::DiscoveryCapabilities;
#[test]
fn empty_set() {
let d = DiscoveryCapabilities::empty();
assert_eq!(d.bits(), 0);
assert!(!d.contains(DiscoveryCapabilities::BLE));
}
#[test]
fn ble_only() {
let d = DiscoveryCapabilities::BLE;
assert_eq!(d.bits(), 0b0000_0010);
assert!(d.contains(DiscoveryCapabilities::BLE));
assert!(!d.contains(DiscoveryCapabilities::ON_NETWORK));
}
#[test]
fn on_network_only() {
let d = DiscoveryCapabilities::ON_NETWORK;
assert_eq!(d.bits(), 0b0000_0100);
}
#[test]
fn combined() {
let d = DiscoveryCapabilities::BLE | DiscoveryCapabilities::ON_NETWORK;
assert_eq!(d.bits(), 0b0000_0110);
}
#[test]
fn from_bits_preserves_reserved() {
let d = DiscoveryCapabilities::from_bits_retain(0b1100_0001);
assert_eq!(d.bits(), 0b1100_0001);
assert!(d.contains(DiscoveryCapabilities::SOFT_AP));
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod setup_payload_tests {
use super::*;
pub(super) fn spec_example_payload() -> SetupPayload {
SetupPayload {
version: 0,
vendor_id: Some(0xFFF1),
product_id: Some(0x8000),
commissioning_flow: CommissioningFlow::Standard,
discovery_capabilities: DiscoveryCapabilities::ON_NETWORK,
discriminator: Discriminator::new(0xF00).unwrap(),
passcode: Passcode::new(20_202_021).unwrap(),
}
}
#[test]
fn spec_example_round_trips_through_struct() {
let p = spec_example_payload();
assert_eq!(p.vendor_id, Some(0xFFF1));
assert_eq!(p.product_id, Some(0x8000));
assert_eq!(p.discriminator.as_u16(), 0xF00);
assert_eq!(p.passcode.as_u32(), 20_202_021);
assert_eq!(p.commissioning_flow, CommissioningFlow::Standard);
assert!(p
.discovery_capabilities
.contains(DiscoveryCapabilities::ON_NETWORK));
}
#[test]
fn manual_only_payload_has_no_vid_pid() {
let p = SetupPayload {
version: 0,
vendor_id: None,
product_id: None,
commissioning_flow: CommissioningFlow::Standard,
discovery_capabilities: DiscoveryCapabilities::empty(),
discriminator: Discriminator::new(0xA00).unwrap(),
passcode: Passcode::new(20_202_021).unwrap(),
};
assert!(p.vendor_id.is_none());
assert!(p.product_id.is_none());
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod qr_api_tests {
use super::*;
use crate::setup::setup_payload_tests::spec_example_payload;
#[test]
fn spec_example_qr_encode_decode_roundtrip() {
let p = spec_example_payload();
let s = encode_qr(&p).unwrap();
assert!(s.starts_with("MT:"), "got {s:?}");
let back = parse_qr(&s).unwrap();
assert_eq!(back, p);
}
#[test]
fn parse_qr_rejects_missing_prefix() {
let err = parse_qr("Y.K9042C00KA0648G00").unwrap_err();
assert!(matches!(err, Error::MissingMtPrefix));
}
#[test]
fn parse_qr_rejects_trailing_bytes() {
let p = spec_example_payload();
let mut s = encode_qr(&p).unwrap();
s.push_str("000");
let err = parse_qr(&s).unwrap_err();
assert!(
matches!(err, Error::QrTrailingBytes { extra: 2 }),
"got {err:?}"
);
}
#[test]
fn parse_qr_rejects_short_payload() {
let err = parse_qr("MT:00000").unwrap_err();
assert!(
matches!(err, Error::QrPayloadWrongLength { .. }),
"got {err:?}"
);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod manual_api_tests {
use super::*;
fn payload_11() -> SetupPayload {
SetupPayload {
version: 0,
vendor_id: None,
product_id: None,
commissioning_flow: CommissioningFlow::Standard,
discovery_capabilities: DiscoveryCapabilities::empty(),
discriminator: Discriminator::new(0x0F00).unwrap(),
passcode: Passcode::new(20_202_021).unwrap(),
}
}
#[test]
fn encode_manual_11_then_parse() {
let p = payload_11();
let s = encode_manual_code(&p);
assert_eq!(s.len(), 11);
let back = parse_manual_code(&s).unwrap();
assert_eq!(back, p);
}
#[test]
fn encode_manual_21_then_parse() {
let mut p = payload_11();
p.vendor_id = Some(0xFFF1);
p.product_id = Some(0x8000);
let s = encode_manual_code(&p);
assert_eq!(s.len(), 21);
let back = parse_manual_code(&s).unwrap();
assert_eq!(back, p);
}
#[test]
fn parse_manual_rejects_wrong_length() {
let err = parse_manual_code("12345").unwrap_err();
assert!(matches!(err, Error::ManualCodeWrongLength(5)));
}
#[test]
fn parse_manual_rejects_non_digit() {
let err = parse_manual_code("1234567890A").unwrap_err();
assert!(matches!(err, Error::ManualCodeNonDigit('A', 10)));
}
}