use core::fmt;
use zerocopy::{FromBytes, Immutable, IntoBytes, KnownLayout};
pub const BLOCK_SIZE: usize = 512;
pub const MAX_PAYLOAD_SIZE: usize = 476;
const CHECKSUM_SIZE: usize = 24;
pub const MAX_PAYLOAD_SIZE_WITH_CHECKSUM: usize =
MAX_PAYLOAD_SIZE - CHECKSUM_SIZE;
pub const PADDING_BYTE: u8 = 0xFF;
pub const ALIGN: usize = 4;
pub const MAGIC_NUMBER: [u32; 3] = [0x0A324655, 0x9E5D5157, 0x0AB16F30];
#[derive(Debug, PartialEq, Eq)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub enum BlockError {
InputBuffer,
MagicNumber,
PayloadSize,
BlockNumberInvalid,
}
impl fmt::Display for BlockError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InputBuffer => write!(f, "Input buffer"),
Self::MagicNumber => write!(f, "Magic number incorrect"),
Self::PayloadSize => write!(f, "Payload size too large"),
Self::BlockNumberInvalid => write!(f, "Block number invalid"),
}
}
}
impl core::error::Error for BlockError {}
#[derive(Debug, Copy, Clone, Immutable, KnownLayout, FromBytes, IntoBytes)]
#[repr(C)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub struct Block {
pub magic_start_0: u32,
pub magic_start_1: u32,
pub flags: Flags,
pub target_addr: u32,
pub data_len: u32,
pub block: u32,
pub total_blocks: u32,
pub board_family_id_or_file_size: u32,
pub data: [u8; MAX_PAYLOAD_SIZE],
pub magic_end: u32,
}
const _: () = {
assert!(core::mem::size_of::<Block>() == BLOCK_SIZE);
};
impl Default for Block {
fn default() -> Self {
Self {
magic_start_0: MAGIC_NUMBER[0],
magic_start_1: MAGIC_NUMBER[1],
flags: Flags::default(),
target_addr: 0,
data_len: 0,
block: 0,
total_blocks: 0,
board_family_id_or_file_size: 0,
data: [0; MAX_PAYLOAD_SIZE],
magic_end: MAGIC_NUMBER[2],
}
}
}
impl Block {
pub fn new(
block: usize,
total_blocks: usize,
data: &[u8],
target_addr: usize,
) -> Self {
let mut this = Self::default();
assert!(block <= total_blocks);
assert!(block <= u32::MAX as usize);
this.block = block as u32;
assert!(total_blocks <= u32::MAX as usize);
this.total_blocks = total_blocks as u32;
assert!(target_addr <= u32::MAX as usize);
this.target_addr = target_addr as u32;
assert!(data.len() <= this.data.len());
this.data_len = data.len() as u32;
this.data[0..data.len()].copy_from_slice(data);
this
}
pub fn from_bytes(buf: &[u8]) -> Result<Block, BlockError> {
let block = match Block::ref_from_bytes(buf) {
Ok(b) => b,
Err(_e) => return Err(BlockError::InputBuffer),
};
if [block.magic_start_0, block.magic_start_1, block.magic_end]
!= MAGIC_NUMBER
{
return Err(BlockError::MagicNumber);
}
if block.data_len > MAX_PAYLOAD_SIZE as u32 {
return Err(BlockError::PayloadSize);
}
if block.block >= block.total_blocks {
return Err(BlockError::BlockNumberInvalid);
}
Ok(*block)
}
pub fn has_checksum(&self) -> bool {
self.flags.contains(Flags::Checksum)
}
pub fn checksum(&self) -> Option<&Checksum> {
if self.has_checksum() {
let len = self.data.len();
Checksum::ref_from_bytes(&self.data[len - CHECKSUM_SIZE..len]).ok()
} else {
None
}
}
pub fn set_checksum(&mut self, checksum: Checksum) {
let begin = self.data.len() - size_of::<Checksum>();
let end = self.data.len();
self.data[begin..end].copy_from_slice(checksum.as_bytes());
self.flags |= Flags::Checksum;
}
pub fn has_extensions(&self) -> bool {
self.flags.contains(Flags::ExtensionTags)
}
pub fn extensions(&self) -> Option<Extensions<'_>> {
if self.has_extensions() {
let start = self.data_len as usize;
let start = start.next_multiple_of(Extensions::ALIGN);
let end = self.data.len();
Some(Extensions::from_bytes(&self.data[start..end]))
} else {
None
}
}
pub fn board_family_id(&self) -> Option<u32> {
match self.flags.contains(Flags::FamilyId) {
true => Some(self.board_family_id_or_file_size),
false => None,
}
}
pub fn data(&self) -> &[u8] {
&self.data[0..self.data_len as usize]
}
pub fn file_size(&self) -> Option<u32> {
match self.flags.contains(Flags::FamilyId) {
false => Some(self.board_family_id_or_file_size),
true => None,
}
}
pub fn add_extension(
&mut self,
tag: ExtensionTag,
data: &[u8],
) -> Result<(), BlockError> {
let mut ext_start = self.data_len as usize;
ext_start = ext_start.next_multiple_of(Extensions::ALIGN);
if self.has_extensions() {
let existing_extensions = self.extensions().unwrap();
for ext in existing_extensions {
let ext_total_len = Extensions::HEADER_SIZE + ext.data.len();
ext_start += ext_total_len.next_multiple_of(Extensions::ALIGN);
}
}
let ext_len = Extensions::HEADER_SIZE + data.len();
let ext_end = ext_start + ext_len;
if ext_end > self.data.len() {
return Err(BlockError::PayloadSize);
}
self.data[ext_start] = ext_len as u8;
let tag_bytes = tag.to_bytes();
self.data[ext_start + 1..ext_start + 4].copy_from_slice(&tag_bytes);
self.data[ext_start + Extensions::HEADER_SIZE..ext_end]
.copy_from_slice(data);
self.flags |= Flags::ExtensionTags;
Ok(())
}
}
#[derive(
Debug, PartialEq, Eq, Immutable, KnownLayout, FromBytes, IntoBytes,
)]
#[repr(C)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub struct Checksum {
pub start: u32,
pub length: u32,
pub checksum: [u8; 16],
}
const _: () = {
assert!(core::mem::size_of::<Checksum>() == CHECKSUM_SIZE);
};
#[derive(
Debug, Default, Clone, Copy, PartialEq, Eq, Immutable, FromBytes, IntoBytes,
)]
#[repr(C)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub struct Flags(u32);
bitflags::bitflags! {
impl Flags: u32 {
const NotMainFlash = 0x00000001;
const FileContainer = 0x00001000;
const FamilyId = 0x00002000;
const Checksum = 0x00004000;
const ExtensionTags = 0x00008000;
const _ = !0; }
}
#[derive(Debug)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub struct Extensions<'a> {
start: usize,
data: &'a [u8],
}
impl<'a> Extensions<'a> {
const HEADER_SIZE: usize = 4;
const ALIGN: usize = 4;
pub fn from_bytes(data: &'a [u8]) -> Self {
Self { start: 0, data }
}
fn current_tag(&self) -> ExtensionTag {
let tag = u32::from_le_bytes([
self.data[self.start + 1],
self.data[self.start + 2],
self.data[self.start + 3],
0,
]);
ExtensionTag::from(tag)
}
}
impl<'a> Iterator for Extensions<'a> {
type Item = Extension<'a>;
fn next(&mut self) -> Option<Self::Item> {
if self.start > self.data.len() {
return None;
}
let len = self.data[self.start] as usize;
if self.start + Self::HEADER_SIZE > self.start + len {
return None;
}
let extension = Extension {
tag: self.current_tag(),
data: &self.data[self.start + Self::HEADER_SIZE..self.start + len],
};
self.start += len;
self.start = self.start.next_multiple_of(Self::ALIGN);
Some(extension)
}
}
#[derive(Debug)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub struct Extension<'a> {
pub tag: ExtensionTag,
pub data: &'a [u8],
}
impl<'a> Extension<'a> {
pub const HEADER_SIZE: usize = 4;
}
#[derive(Debug, PartialEq, Eq)]
#[repr(u32)]
#[cfg_attr(feature = "defmt", derive(defmt::Format))]
pub enum ExtensionTag {
SemverString = 0x9fc7bc,
DescriptionString = 0x650d9d,
TargetPageSize = 0x0be9f7,
Sha2Checksum = 0xb46db0,
DeviceTypeId = 0xc8a729,
Other(u32),
}
impl From<u32> for ExtensionTag {
fn from(value: u32) -> Self {
match value {
0x9fc7bc => Self::SemverString,
0x650d9d => Self::DescriptionString,
0x0be9f7 => Self::TargetPageSize,
0xb46db0 => Self::Sha2Checksum,
0xc8a729 => Self::DeviceTypeId,
_ => Self::Other(value), }
}
}
impl ExtensionTag {
pub fn to_bytes(&self) -> [u8; 3] {
match self {
ExtensionTag::SemverString => {
0x9fc7bc_u32.to_le_bytes()[0..3].try_into().unwrap()
}
ExtensionTag::DescriptionString => {
0x650d9d_u32.to_le_bytes()[0..3].try_into().unwrap()
}
ExtensionTag::TargetPageSize => {
0x0be9f7_u32.to_le_bytes()[0..3].try_into().unwrap()
}
ExtensionTag::Sha2Checksum => {
0xb46db0_u32.to_le_bytes()[0..3].try_into().unwrap()
}
ExtensionTag::DeviceTypeId => {
0xc8a729_u32.to_le_bytes()[0..3].try_into().unwrap()
}
ExtensionTag::Other(value) => {
value.to_le_bytes()[0..3].try_into().unwrap()
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn magic_number() {
assert_eq!(MAGIC_NUMBER[0].as_bytes(), b"UF2\n");
}
#[test]
fn block_checksum() {
let mut block = Block::default();
assert_eq!(block.has_checksum(), false);
block.flags |= Flags::Checksum;
assert_eq!(block.has_checksum(), true);
let cksm = block.checksum();
assert!(cksm.is_some());
}
#[test]
fn block_extension() {
let mut block = Block {
flags: Flags::ExtensionTags,
data_len: 0,
..Default::default()
};
block.data[0..12].copy_from_slice(&[
0x09, 0xbc, 0xc7, 0x9f, 0x30, 0x2e, 0x31, 0x2e, 0x32, 0x00, 0x00,
0x00,
]);
block.data[12..24].copy_from_slice(&[
0x09, 0xbc, 0xc7, 0x9f, 0x30, 0x2e, 0x31, 0x2e, 0x32, 0x00, 0x00,
0x00,
]);
block.data[24..44].copy_from_slice(&[
0x14, 0x9d, 0x0d, 0x65, 0x41, 0x43, 0x4d, 0x45, 0x20, 0x54, 0x6f,
0x61, 0x73, 0x74, 0x65, 0x72, 0x20, 0x6d, 0x6b, 0x33,
]);
assert!(block.extensions().is_some());
let mut extensions = block.extensions().unwrap();
let first = extensions.next().unwrap();
assert_eq!(first.tag, ExtensionTag::SemverString);
assert_eq!(first.data, b"0.1.2");
let second = extensions.next().unwrap();
assert_eq!(second.tag, ExtensionTag::SemverString);
assert_eq!(second.data, b"0.1.2");
let third = extensions.next().unwrap();
assert_eq!(third.tag, ExtensionTag::DescriptionString);
assert_eq!(third.data, b"ACME Toaster mk3");
}
#[test]
fn example_file() {
use std::io::prelude::*;
let mut f = std::fs::File::open("example.uf2").unwrap();
let mut buffer = [0; 512];
f.read(&mut buffer).unwrap();
let block = Block::from_bytes(&buffer).unwrap();
assert_eq!(block.magic_start_0, MAGIC_NUMBER[0]);
assert_eq!(block.magic_start_1, MAGIC_NUMBER[1]);
assert_eq!(block.magic_end, MAGIC_NUMBER[2]);
assert_eq!(block.target_addr, 0x2000);
assert_eq!(block.data_len, 256);
assert_eq!(block.block, 0);
assert_eq!(block.total_blocks, 1438);
assert_eq!(block.board_family_id_or_file_size, 0);
}
#[test]
fn test_block_new() {
let data = [0xAA; 256];
let block = Block::new(0, 1, &data, 0x08000000);
assert_eq!(block.magic_start_0, MAGIC_NUMBER[0]);
assert_eq!(block.magic_start_1, MAGIC_NUMBER[1]);
assert_eq!(block.magic_end, MAGIC_NUMBER[2]);
assert_eq!(block.block, 0);
assert_eq!(block.total_blocks, 1);
assert_eq!(block.target_addr, 0x08000000);
assert_eq!(block.data_len, 256);
assert_eq!(block.data(), &data);
}
#[test]
fn test_block_new_multiple_blocks() {
let data = [0xBB; 100];
let block = Block::new(2, 5, &data, 0x08000100);
assert_eq!(block.block, 2);
assert_eq!(block.total_blocks, 5);
assert_eq!(block.target_addr, 0x08000100);
assert_eq!(block.data_len, 100);
assert_eq!(block.data(), &data);
}
#[test]
fn test_block_new_empty_data() {
let data: &[u8] = &[];
let block = Block::new(0, 1, data, 0);
assert_eq!(block.data_len, 0);
assert_eq!(block.data(), &[]);
}
#[test]
fn test_block_new_max_payload() {
let data = [0xCC; MAX_PAYLOAD_SIZE];
let block = Block::new(0, 1, &data, 0);
assert_eq!(block.data_len, MAX_PAYLOAD_SIZE as u32);
assert_eq!(block.data(), &data);
}
#[test]
#[should_panic(expected = "block <= total_blocks")]
fn test_block_new_panics_on_invalid_index() {
let data = [0xDD; 100];
Block::new(5, 3, &data, 0); }
#[test]
fn test_from_bytes_block_number_exceeds_total() {
let mut block = Block::default();
block.block = 5; block.total_blocks = 3; let bytes = block.as_bytes();
let result = Block::from_bytes(&bytes);
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
BlockError::BlockNumberInvalid
));
}
#[test]
fn test_set_checksum() {
let mut block = Block::default();
assert_eq!(block.has_checksum(), false);
assert!(block.checksum().is_none());
let checksum = Checksum {
start: 0x08000000,
length: 256,
checksum: [0xAB; 16],
};
block.set_checksum(checksum);
assert_eq!(block.has_checksum(), true);
let retrieved = block.checksum().unwrap();
assert_eq!(retrieved.start, 0x08000000);
assert_eq!(retrieved.length, 256);
assert_eq!(retrieved.checksum, [0xAB; 16]);
}
#[test]
fn test_board_family_id() {
let mut block = Block::default();
assert_eq!(block.board_family_id(), None);
block.flags |= Flags::FamilyId;
block.board_family_id_or_file_size = 0x12345678;
assert_eq!(block.board_family_id(), Some(0x12345678));
block.board_family_id_or_file_size = 0x87654321;
assert_eq!(block.board_family_id(), Some(0x87654321));
}
#[test]
fn test_file_size() {
let mut block = Block::default();
assert_eq!(block.file_size(), Some(0));
block.board_family_id_or_file_size = 1024;
assert_eq!(block.file_size(), Some(1024));
block.flags |= Flags::FamilyId;
assert_eq!(block.file_size(), None);
block.flags &= !Flags::FamilyId;
assert_eq!(block.file_size(), Some(1024));
}
#[test]
fn test_board_family_id_vs_file_size() {
let mut block = Block::default();
block.board_family_id_or_file_size = 0xCAFEBABE;
assert_eq!(block.file_size(), Some(0xCAFEBABE));
assert_eq!(block.board_family_id(), None);
block.flags |= Flags::FamilyId;
assert_eq!(block.board_family_id(), Some(0xCAFEBABE));
assert_eq!(block.file_size(), None);
assert_eq!(block.board_family_id_or_file_size, 0xCAFEBABE);
}
#[test]
fn test_add_extension() {
let mut block = Block::new(0, 1, &[0xAA; 100], 0x08000000);
assert_eq!(block.has_extensions(), false);
let result = block.add_extension(ExtensionTag::SemverString, b"1.0.0");
assert!(result.is_ok());
assert_eq!(block.has_extensions(), true);
let mut extensions = block.extensions().unwrap();
let ext = extensions.next().unwrap();
assert_eq!(ext.tag, ExtensionTag::SemverString);
assert_eq!(ext.data, b"1.0.0");
}
#[test]
fn test_add_extension_multiple() {
let mut block = Block::new(0, 1, &[0xBB; 50], 0x08000000);
block
.add_extension(ExtensionTag::SemverString, b"1.0.0")
.unwrap();
let mut extensions = block.extensions().unwrap();
let ext1 = extensions.next().unwrap();
assert_eq!(ext1.tag, ExtensionTag::SemverString);
assert_eq!(ext1.data, b"1.0.0");
block
.add_extension(ExtensionTag::DescriptionString, b"Test")
.unwrap();
let mut extensions = block.extensions().unwrap();
let ext1 = extensions.next().unwrap();
let ext2 = extensions.next().unwrap();
let tag1 = ext1.tag;
let tag2 = ext2.tag;
assert_eq!(tag1, ExtensionTag::SemverString);
assert_eq!(tag2, ExtensionTag::DescriptionString);
}
#[test]
fn test_add_extension_no_space() {
let mut block = Block::new(0, 1, &[0xCC; MAX_PAYLOAD_SIZE], 0x08000000);
let result = block.add_extension(ExtensionTag::SemverString, b"1.0.0");
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), BlockError::PayloadSize));
}
}