use alloc::vec::Vec;
use dvb_common::traits::{Parse, Serialize};
use crate::error::{Error, Result};
use crate::registry::{Interface, MessageType, ParameterType};
pub const HEADER_LEN: usize = 5;
pub const PARAMETER_HEADER_LEN: usize = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct Parameter<'a> {
pub ptype: ParameterType,
#[cfg_attr(feature = "serde", serde(borrow))]
pub value: &'a [u8],
}
impl<'a> Parameter<'a> {
#[must_use]
pub const fn new(ptype: ParameterType, value: &'a [u8]) -> Self {
Self { ptype, value }
}
#[must_use]
pub const fn wire_len(&self) -> usize {
PARAMETER_HEADER_LEN + self.value.len()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct SimulcryptMessage<'a> {
pub protocol_version: u8,
pub message_type: MessageType,
#[cfg_attr(feature = "serde", serde(borrow))]
pub parameters: Vec<Parameter<'a>>,
}
impl<'a> SimulcryptMessage<'a> {
#[must_use]
pub fn new(
protocol_version: u8,
message_type: MessageType,
parameters: Vec<Parameter<'a>>,
) -> Self {
Self {
protocol_version,
message_type,
parameters,
}
}
#[must_use]
pub const fn interface(&self) -> Interface {
self.message_type.interface()
}
#[must_use]
pub fn body_len(&self) -> usize {
self.parameters.iter().map(Parameter::wire_len).sum()
}
#[must_use]
pub fn find(&self, ptype: ParameterType) -> Option<&Parameter<'a>> {
self.parameters.iter().find(|p| p.ptype == ptype)
}
pub fn parse_on(iface: Interface, bytes: &'a [u8]) -> Result<Self> {
if bytes.len() < HEADER_LEN {
return Err(Error::BufferTooShort {
need: HEADER_LEN,
have: bytes.len(),
what: "generic_message header",
});
}
let protocol_version = bytes[0];
let raw_message_type = u16::from_be_bytes([bytes[1], bytes[2]]);
let message_length = u16::from_be_bytes([bytes[3], bytes[4]]) as usize;
let body = &bytes[HEADER_LEN..];
if body.len() < message_length {
return Err(Error::InvalidMessageLength {
length: message_length as u16,
reason: "message_length exceeds available bytes",
});
}
let body = &body[..message_length];
let message_type = MessageType::from_u16(iface, raw_message_type);
let mut parameters = Vec::new();
let mut off = 0usize;
while off < body.len() {
if body.len() - off < PARAMETER_HEADER_LEN {
return Err(Error::BufferTooShort {
need: PARAMETER_HEADER_LEN,
have: body.len() - off,
what: "parameter TLV header",
});
}
let raw_ptype = u16::from_be_bytes([body[off], body[off + 1]]);
let plen = u16::from_be_bytes([body[off + 2], body[off + 3]]) as usize;
let vstart = off + PARAMETER_HEADER_LEN;
let remaining = body.len() - vstart;
if remaining < plen {
return Err(Error::TruncatedParameter {
ptype: raw_ptype,
need: plen,
have: remaining,
});
}
let value = &body[vstart..vstart + plen];
parameters.push(Parameter::new(
ParameterType::from_u16(iface, raw_ptype),
value,
));
off = vstart + plen;
}
Ok(Self {
protocol_version,
message_type,
parameters,
})
}
}
impl<'a> Parse<'a> for SimulcryptMessage<'a> {
type Error = Error;
fn parse(bytes: &'a [u8]) -> Result<Self> {
Self::parse_on(Interface::EcmgScs, bytes)
}
}
impl<'a> Serialize for SimulcryptMessage<'a> {
type Error = Error;
fn serialized_len(&self) -> usize {
HEADER_LEN + self.body_len()
}
fn serialize_into(&self, buf: &mut [u8]) -> Result<usize> {
let total = self.serialized_len();
if buf.len() < total {
return Err(Error::OutputBufferTooSmall {
need: total,
have: buf.len(),
});
}
let body_len = self.body_len();
if body_len > u16::MAX as usize {
return Err(Error::FieldTooWide {
what: "message_length",
value: body_len,
bits: 16,
});
}
buf[0] = self.protocol_version;
buf[1..3].copy_from_slice(&self.message_type.to_u16().to_be_bytes());
buf[3..5].copy_from_slice(&(body_len as u16).to_be_bytes());
let mut off = HEADER_LEN;
for p in &self.parameters {
let plen = p.value.len();
if plen > u16::MAX as usize {
return Err(Error::FieldTooWide {
what: "parameter_length",
value: plen,
bits: 16,
});
}
buf[off..off + 2].copy_from_slice(&p.ptype.to_u16().to_be_bytes());
buf[off + 2..off + 4].copy_from_slice(&(plen as u16).to_be_bytes());
off += PARAMETER_HEADER_LEN;
buf[off..off + plen].copy_from_slice(p.value);
off += plen;
}
debug_assert_eq!(off, total);
Ok(total)
}
}