diff --git a/tbf-parser/src/types.rs b/tbf-parser/src/types.rs index 5f4dca45..8491924c 100644 --- a/tbf-parser/src/types.rs +++ b/tbf-parser/src/types.rs @@ -1183,4 +1183,85 @@ impl TbfHeader { _ => None, } } + + /// Returns the checksum for TBF header according to the `parse_tbf_header` function + /// `new_flags` is the new value we want to set + pub fn compute_checksum(header: &[u8], new_flags: u32) -> Result { + let mut checksum: u32 = 0; + + let header_iter = header.chunks_exact(4); + + // Iterate all chunks and XOR the chunks to compute the checksum. + for (i, chunk) in header_iter.enumerate() { + let word = if i == 2 { + new_flags + } else if i == 3 { + continue; + } else { + u32::from_le_bytes(chunk.try_into()?) + }; + checksum ^= word; + } + Ok(checksum) + } + + /// Sets the flag field and updates the checksum, it modifies the state accordingly + pub fn set_flags(&mut self, flags: u32, header: &[u8]) -> Result<(), TbfParseError> { + let new_checksum = Self::compute_checksum(header, flags)?; + match self { + TbfHeader::TbfHeaderV2(hd) => { + hd.base.flags = flags; + hd.base.checksum = new_checksum; + } + TbfHeader::Padding(base) => { + base.flags = flags; + base.checksum = new_checksum; + } + } + + Ok(()) + } + + /// Returns the TBF Header's flags (avoid duplication code for the setting functions below). + fn get_flags(&self) -> u32 { + match self { + TbfHeader::TbfHeaderV2(hd) => hd.base.flags, + TbfHeader::Padding(base) => base.flags, + } + } + + /// Enables or disables the application by setting the enabled flag. + pub fn set_enabled(&mut self, enabled: bool, header: &[u8]) -> Result<(), TbfParseError> { + let flags: u32 = if enabled { + self.get_flags() | 0x00000001 + } else { + self.get_flags() & !0x00000001 + }; + self.set_flags(flags, header) + } + + /// Enable or disables erase confirmation by setting the sticky flag. + pub fn set_sticky(&mut self, enabled: bool, header: &[u8]) -> Result<(), TbfParseError> { + let flags: u32 = if enabled { + self.get_flags() | 0x00000002 + } else { + self.get_flags() & !0x00000002 + }; + self.set_flags(flags, header) + } + + /// Returns a 16 byte array with the serialized base header (little-endian format). + pub fn serialize(&self) -> Result<[u8; 16], TbfParseError> { + let base = match self { + TbfHeader::TbfHeaderV2(hd) => &hd.base, + TbfHeader::Padding(base) => base, + }; + let mut bytes = [0u8; 16]; + bytes[0..2].copy_from_slice(&base.version.to_le_bytes()); + bytes[2..4].copy_from_slice(&base.header_size.to_le_bytes()); + bytes[4..8].copy_from_slice(&base.total_size.to_le_bytes()); + bytes[8..12].copy_from_slice(&base.flags.to_le_bytes()); + bytes[12..16].copy_from_slice(&base.checksum.to_le_bytes()); + Ok(bytes) + } } diff --git a/tbf-parser/tests/serialization.rs b/tbf-parser/tests/serialization.rs new file mode 100644 index 00000000..42173b02 --- /dev/null +++ b/tbf-parser/tests/serialization.rs @@ -0,0 +1,224 @@ +use tbf_parser::parse::*; +use tbf_parser::types::TbfHeader; + +// Serialization + +#[test] +fn serialize_identical_with_original() { + let buffer: Vec = include_bytes!("./flashes/simple.dat").to_vec(); + + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + + let header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + let serialized = header.serialize().unwrap(); + + // Check if serialize matches original buffer + assert_eq!(&buffer[0..16], &serialized[..]); +} + +// Flag modifications +#[test] +fn flags_modifications() { + let mut buffer = include_bytes!("./flashes/footerSHA256.dat").to_vec(); + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + + let mut header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + assert!(header.enabled()); + header + .set_sticky(true, &buffer[0..header_len as usize]) + .unwrap(); + // Set sticky without parsing + assert!(header.sticky()); + // Unset + header + .set_sticky(false, &buffer[0..header_len as usize]) + .unwrap(); + + // Disable + header + .set_enabled(false, &buffer[0..header_len as usize]) + .unwrap(); + let serialized = header.serialize().unwrap(); + buffer[0..16].copy_from_slice(&serialized); + + let reparsed = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + assert!(!reparsed.enabled()); + + // Enable + let mut header = reparsed; + header + .set_enabled(true, &buffer[0..header_len as usize]) + .unwrap(); + let serialized = header.serialize().unwrap(); + buffer[0..16].copy_from_slice(&serialized); + + let reparsed = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + assert!(reparsed.enabled()); +} + +#[test] +fn padding_header_set_flags() { + let buffer = vec![ + 0x02, 0x00, 0x10, 0x00, 0x10, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x12, 0x00, 0x10, + 0x00, + ]; + + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + let mut header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + + assert!(!header.is_app()); + + let serialized = header.serialize(); + assert!(serialized.is_ok()); + + assert!(header.set_flags(0x06000001, &buffer).is_ok()); + + let serialized = header.serialize().unwrap(); + let flags = u32::from_le_bytes(serialized[8..12].try_into().unwrap()); + assert_eq!(flags, 0x06000001); +} + +#[test] +fn fields_preserved() { + let buffer = include_bytes!("./flashes/simple.dat").to_vec(); + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + + let mut header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + + let header_size = header.header_size(); + let total_size = header.total_size(); + + header + .set_flags(0x00000003, &buffer[0..header_len as usize]) + .unwrap(); + let serialized = header.serialize().unwrap(); + + let _version = u16::from_le_bytes(serialized[0..2].try_into().unwrap()); + let _header_size = u16::from_le_bytes(serialized[2..4].try_into().unwrap()); + let _total_size = u32::from_le_bytes(serialized[4..8].try_into().unwrap()); + + // Check other fields are unchanged + assert_eq!(_version, 2); + assert_eq!(header_size, _header_size); + assert_eq!(total_size, _total_size); +} + +#[test] +fn multiple_flags_set() { + let buffer = include_bytes!("./flashes/simple.dat").to_vec(); + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + + let mut header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + + // Trying to set multiple times to check consistency + for i in 1..21 { + header + .set_flags(i, &buffer[0..header_len as usize]) + .unwrap(); + assert_eq!(header.enabled(), i % 2 == 1); + } +} + +// Checksum // +#[test] +fn checksum() { + // Try with empty_buffer + let empty_buffer: Vec = vec![]; + let result = TbfHeader::compute_checksum(&empty_buffer, 0x00000006D); + assert_eq!(result.unwrap(), 0x00000000); + + let mut buffer = include_bytes!("./flashes/simple.dat").to_vec(); + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + + let mut header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + + // Test with array of multiple flags for checksum validation + for flags in [0x000FABCD, 0x00000001, 0x00000002, 0x00000003] { + header + .set_flags(flags, &buffer[0..header_len as usize]) + .unwrap(); + let serialized = header.serialize().unwrap(); + buffer[0..16].copy_from_slice(&serialized); + + let result = parse_tbf_header(&buffer[0..header_len as usize], 2); + assert!(result.is_ok(), "Checksum validation failed for {flags}"); + } +} + +// Complete use // +#[test] +fn serialization_multiple_checks() { + let mut buffer = include_bytes!("./flashes/footerRSA4096.dat").to_vec(); + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + + let mut header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + assert!(header.enabled()); + + // Disable + header + .set_enabled(false, &buffer[0..header_len as usize]) + .unwrap(); + let serialized = header.serialize().unwrap(); + buffer[0..16].copy_from_slice(&serialized); + + let header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + assert!(!header.enabled()); + + // Enable and set sticky + let mut header = header; + header + .set_enabled(true, &buffer[0..header_len as usize]) + .unwrap(); + header + .set_sticky(true, &buffer[0..header_len as usize]) + .unwrap(); + let serialized = header.serialize().unwrap(); + buffer[0..16].copy_from_slice(&serialized); + + let header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + assert!(header.enabled()); + assert!(header.sticky()); + + // Disable sticky with high bits + let flags = 0xD6D6FFF1; + let mut header = header; + header + .set_flags(flags, &buffer[0..header_len as usize]) + .unwrap(); + let serialized = header.serialize().unwrap(); + buffer[0..16].copy_from_slice(&serialized); + + let header = parse_tbf_header(&buffer[0..header_len as usize], 2).unwrap(); + assert!(header.enabled()); + assert!(!header.sticky()); + let flags_buffer = u32::from_le_bytes(buffer[8..12].try_into().unwrap()); + assert_eq!(flags_buffer, flags); +} + +#[test] +fn corrupt() { + let mut buffer = include_bytes!("./flashes/simple.dat").to_vec(); + let (_, header_len, _) = parse_tbf_header_lengths(&buffer[0..8].try_into().unwrap()) + .ok() + .unwrap(); + + // Corrupt the checksum manually and check for parsing error + buffer[12] ^= 0x6D; + + let result = parse_tbf_header(&buffer[0..header_len as usize], 2); + assert!(result.is_err()); +}