use log::{debug, error, info, warn};
use sunset::packets::ParseContext;
use sunset::sshwire::{SSHDecode, SSHEncode, SSHSource, WireError};
use crate::tlv;
pub type OtaTlvType = u8;
pub type OtaTlvLen = u8;
pub const OTA_TYPE_VALUE_SSH_STAMP: u32 = 0x7373_6873;
pub const CHECKSUM_LEN: u32 = 32;
pub const MAX_TLV_SIZE: u32 = 257;
fn enc_len_val<SE>(
value: &SE,
s: &mut dyn sunset::sshwire::SSHSink,
) -> sunset::sshwire::WireResult<()>
where
SE: Sized + SSHEncode,
{
OtaTlvLen::try_from(core::mem::size_of::<SE>())
.map_err(|_| sunset::sshwire::WireError::PacketWrong)?
.enc(s)?;
value.enc(s)
}
fn dec_check_val_len<'de, S, SE>(s: &mut S) -> sunset::sshwire::WireResult<()>
where
S: sunset::sshwire::SSHSource<'de>,
SE: Sized,
{
let val_len = OtaTlvLen::dec(s)?;
let expected_len = OtaTlvLen::try_from(core::mem::size_of::<SE>())
.map_err(|_| sunset::sshwire::WireError::PacketWrong)?;
if val_len != expected_len {
return Err(sunset::sshwire::WireError::PacketWrong);
}
Ok(())
}
pub const OTA_TYPE: OtaTlvType = 0;
pub const FIRMWARE_BLOB: OtaTlvType = 1;
pub const SHA256_CHECKSUM: OtaTlvType = 2;
#[derive(Debug)]
#[repr(u8)] pub enum Tlv {
OtaType { ota_type: u32 },
Sha256Checksum {
checksum: [u8; CHECKSUM_LEN as usize],
},
FirmwareBlob { size: u32 },
}
impl SSHEncode for Tlv {
fn enc(&self, s: &mut dyn sunset::sshwire::SSHSink) -> sunset::sshwire::WireResult<()> {
match self {
Tlv::OtaType { ota_type } => {
OTA_TYPE.enc(s)?;
enc_len_val(ota_type, s)
}
Tlv::FirmwareBlob { size } => {
FIRMWARE_BLOB.enc(s)?;
enc_len_val(size, s)
}
Tlv::Sha256Checksum { checksum } => {
SHA256_CHECKSUM.enc(s)?;
enc_len_val(checksum, s)
}
}
}
}
impl<'de> SSHDecode<'de> for Tlv {
fn dec<S>(s: &mut S) -> sunset::sshwire::WireResult<Self>
where
S: sunset::sshwire::SSHSource<'de>,
{
OtaTlvType::dec(s).and_then(|tlv_type| match tlv_type {
FIRMWARE_BLOB => {
dec_check_val_len::<S, u32>(s)?;
Ok(Tlv::FirmwareBlob { size: u32::dec(s)? })
}
SHA256_CHECKSUM => {
let expected_len = OtaTlvLen::try_from(tlv::CHECKSUM_LEN)
.map_err(|_| sunset::sshwire::WireError::PacketWrong)?;
if OtaTlvLen::dec(s)? != expected_len {
return Err(sunset::sshwire::WireError::PacketWrong);
}
let mut checksum = [0u8; tlv::CHECKSUM_LEN as usize];
for element in &mut checksum {
*element = u8::dec(s)?;
}
Ok(Tlv::Sha256Checksum { checksum })
}
OTA_TYPE => {
dec_check_val_len::<S, u32>(s)?;
let ota_type = u32::dec(s)?;
Ok(Tlv::OtaType { ota_type })
}
_ => {
error!("Unknown TLV type encountered: {tlv_type}");
let len = OtaTlvLen::dec(s)?;
s.take(len as usize)?; Err(sunset::sshwire::WireError::UnknownPacket { number: tlv_type })
}
})
}
}
pub struct TlvsSource<'a> {
remaining_buf: &'a [u8],
ctx: ParseContext,
used: usize,
}
impl<'a> TlvsSource<'a> {
#[must_use]
pub fn new(buf: &'a [u8]) -> Self {
Self {
remaining_buf: buf,
ctx: ParseContext::default(),
used: 0,
}
}
#[must_use]
pub fn used(&self) -> usize {
self.used
}
pub fn try_taking_bytes_for_tlv(
&mut self,
tlv_holder: &mut [u8],
current_len: &mut usize,
) -> Result<(), WireError> {
if *current_len
< core::mem::size_of::<tlv::OtaTlvType>() + core::mem::size_of::<tlv::OtaTlvLen>()
{
let needed = core::mem::size_of::<tlv::OtaTlvType>()
+ core::mem::size_of::<tlv::OtaTlvLen>()
- *current_len;
debug!("Adding {needed} bytes to have up to TLV type and length");
let to_read = core::cmp::min(needed, self.remaining());
let type_len_bytes = self.take(to_read)?;
tlv_holder[*current_len..*current_len + to_read].copy_from_slice(type_len_bytes);
*current_len += to_read;
if to_read < needed {
info!("Will get more data to complete TLV type/length");
return Err(WireError::RanOut);
}
}
let slice_len_start = core::mem::size_of::<tlv::OtaTlvType>();
let slice_value_start =
core::mem::size_of::<tlv::OtaTlvType>() + core::mem::size_of::<tlv::OtaTlvLen>();
if *current_len >= slice_value_start {
let val_len = tlv_holder[slice_len_start] as usize;
debug!(
"value length: {}, Source remaining bytes: {}",
val_len,
self.remaining()
);
let needed = val_len + slice_value_start - *current_len;
let to_read = needed.min(self.remaining());
let needed_type_len_bytes = self.take(to_read)?;
tlv_holder[*current_len..*current_len + to_read].copy_from_slice(needed_type_len_bytes);
*current_len += to_read;
if to_read < needed {
info!("Will get more data to complete TLV value");
return Err(WireError::RanOut);
}
}
Ok(())
}
}
impl<'de> SSHSource<'de> for TlvsSource<'de> {
fn take(&mut self, len: usize) -> sunset::sshwire::WireResult<&'de [u8]> {
if len > self.remaining_buf.len() {
return Err(sunset::sshwire::WireError::RanOut);
}
let t;
(t, self.remaining_buf) = self.remaining_buf.split_at(len);
self.used += len;
Ok(t)
}
fn remaining(&self) -> usize {
self.remaining_buf.len()
}
fn ctx(&mut self) -> &mut ParseContext {
&mut self.ctx
}
}
#[derive(Debug)]
pub struct OtaHeader {
pub(crate) ota_type: Option<u32>,
pub(crate) firmware_blob_size: Option<u32>,
pub sha256_checksum: Option<[u8; tlv::CHECKSUM_LEN as usize]>,
}
impl OtaHeader {
#[cfg(not(target_os = "none"))]
#[must_use]
pub fn new(ota_type: u32, sha256_checksum: &[u8], firmware_blob_size: u32) -> Self {
let mut checksum_array = [0u8; tlv::CHECKSUM_LEN as usize];
checksum_array.copy_from_slice(sha256_checksum);
Self {
ota_type: Some(ota_type),
firmware_blob_size: Some(firmware_blob_size),
sha256_checksum: Some(checksum_array),
}
}
pub fn serialize(&self, buf: &mut [u8]) -> usize {
let mut offset = 0;
if let Some(ota_type) = self.ota_type {
let tlv = tlv::Tlv::OtaType { ota_type };
let used = sunset::sshwire::write_ssh(&mut buf[offset..], &tlv)
.expect("Failed to serialize OTA Type TLV");
offset += used;
}
if let Some(checksum) = &self.sha256_checksum {
let tlv = tlv::Tlv::Sha256Checksum {
checksum: *checksum,
};
let used = sunset::sshwire::write_ssh(&mut buf[offset..], &tlv)
.expect("Failed to serialize SHA256 Checksum TLV");
offset += used;
}
if let Some(size) = self.firmware_blob_size {
let tlv = tlv::Tlv::FirmwareBlob { size };
let used = sunset::sshwire::write_ssh(&mut buf[offset..], &tlv)
.expect("Failed to serialize Firmware Blob TLV");
offset += used;
}
offset
}
pub fn deserialize(buf: &[u8]) -> Result<(Self, usize), sunset::sshwire::WireError> {
let mut source = tlv::TlvsSource::new(buf);
let mut ota_type = None;
let mut firmware_blob_size = None;
let mut sha256_checksum = None;
while source.remaining() > 0 {
match tlv::Tlv::dec(&mut source) {
Err(sunset::sshwire::WireError::UnknownPacket { number }) => {
warn!(
"Unknown packet type encountered: {number}. TLV skipping it and continuing"
);
}
Err(e) => {
return Err(e);
}
Ok(tlv) => {
match tlv {
tlv::Tlv::OtaType { ota_type: ot } => {
ota_type = Some(ot);
}
tlv::Tlv::Sha256Checksum { checksum } => {
Self::check_ota_is_first_tlv(ota_type)?;
sha256_checksum = Some(checksum);
}
tlv::Tlv::FirmwareBlob { size } => {
Self::check_ota_is_first_tlv(ota_type)?;
firmware_blob_size = Some(size);
break;
}
}
}
}
}
Ok((
Self {
ota_type,
firmware_blob_size,
sha256_checksum,
},
source.used(),
))
}
fn check_ota_is_first_tlv(ota_type: Option<u32>) -> Result<(), WireError> {
if ota_type.is_none() {
error!("SHA256 Checksum TLV encountered before OTA Type TLV. Ignoring it");
Err(sunset::sshwire::WireError::PacketWrong)
} else {
Ok(())
}
}
}