use crate::format::bytes::read_le_uint as read_size;
use crate::format::{FormatContext, FormatError, FormatResult};
const GCOL_SIGNATURE: [u8; 4] = *b"GCOL";
const GCOL_VERSION: u8 = 1;
const GCOL_MIN_SIZE: usize = 4096;
pub const GCOL_MAX_SIZE: usize = 65536;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GlobalHeapObject {
pub index: u16,
pub ref_count: u16,
pub data: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GlobalHeapCollection {
pub objects: Vec<GlobalHeapObject>,
}
impl GlobalHeapCollection {
pub fn new() -> Self {
Self {
objects: Vec::new(),
}
}
pub fn add_object(&mut self, data: Vec<u8>) -> FormatResult<u16> {
let max_index = self.objects.iter().map(|o| o.index).max().unwrap_or(0);
if max_index == u16::MAX {
return Err(FormatError::InvalidData(
"global heap collection is full (65535 objects)".into(),
));
}
let index = max_index + 1;
self.objects.push(GlobalHeapObject {
index,
ref_count: 0,
data,
});
Ok(index)
}
pub fn get_object(&self, index: u16) -> Option<&[u8]> {
self.objects
.iter()
.find(|o| o.index == index)
.map(|o| o.data.as_slice())
}
pub fn remove_object(&mut self, index: u16) -> bool {
match self.objects.iter().position(|o| o.index == index) {
Some(at) => {
self.objects.remove(at);
true
}
None => false,
}
}
pub fn is_empty(&self) -> bool {
self.objects.is_empty()
}
pub fn free_space_at(&self, ctx: &FormatContext, collection_size: usize) -> Option<usize> {
let (header_size, objhdr_size, objects_size) = self.layout(ctx);
let content_size = header_size + objects_size + objhdr_size;
if collection_size < content_size {
return None;
}
Some(collection_size - header_size - objects_size)
}
pub fn object_disk_size(ctx: &FormatContext, data_len: usize) -> usize {
let ss = ctx.sizeof_size as usize;
pad_to_8(2 + 2 + 4 + ss) + pad_to_8(data_len)
}
pub fn max_index(&self) -> u16 {
self.objects.iter().map(|o| o.index).max().unwrap_or(0)
}
pub fn encode(&self, ctx: &FormatContext) -> Vec<u8> {
self.encode_at_size(ctx, self.encoded_size(ctx)).unwrap()
}
pub fn encoded_size(&self, ctx: &FormatContext) -> usize {
let (header_size, objhdr_size, objects_size) = self.layout(ctx);
(header_size + objects_size + objhdr_size).max(GCOL_MIN_SIZE)
}
pub fn encode_at_size(
&self,
ctx: &FormatContext,
collection_size: usize,
) -> FormatResult<Vec<u8>> {
let ss = ctx.sizeof_size as usize;
let (header_size, objhdr_size, objects_size) = self.layout(ctx);
let content_size = header_size + objects_size + objhdr_size;
if collection_size < content_size {
return Err(FormatError::InvalidData(format!(
"global heap collection needs {content_size} bytes but was given {collection_size}"
)));
}
let free_space = collection_size - header_size - objects_size;
let mut buf = Vec::with_capacity(collection_size);
buf.extend_from_slice(&GCOL_SIGNATURE);
buf.push(GCOL_VERSION);
buf.extend_from_slice(&[0u8; 3]); buf.extend_from_slice(&(collection_size as u64).to_le_bytes()[..ss]);
buf.resize(header_size, 0);
for obj in &self.objects {
let obj_start = buf.len();
buf.extend_from_slice(&obj.index.to_le_bytes());
buf.extend_from_slice(&obj.ref_count.to_le_bytes());
buf.extend_from_slice(&0u32.to_le_bytes()); buf.extend_from_slice(&(obj.data.len() as u64).to_le_bytes()[..ss]);
buf.resize(obj_start + objhdr_size, 0); buf.extend_from_slice(&obj.data);
buf.resize(buf.len() + (pad_to_8(obj.data.len()) - obj.data.len()), 0);
}
buf.extend_from_slice(&0u16.to_le_bytes()); buf.extend_from_slice(&0u16.to_le_bytes()); buf.extend_from_slice(&0u32.to_le_bytes()); buf.extend_from_slice(&(free_space as u64).to_le_bytes()[..ss]);
buf.resize(collection_size, 0);
debug_assert_eq!(buf.len(), collection_size);
Ok(buf)
}
fn layout(&self, ctx: &FormatContext) -> (usize, usize, usize) {
let ss = ctx.sizeof_size as usize;
let header_size = pad_to_8(4 + 1 + 3 + ss); let objhdr_size = pad_to_8(2 + 2 + 4 + ss); let objects_size = self
.objects
.iter()
.map(|obj| objhdr_size + pad_to_8(obj.data.len()))
.sum();
(header_size, objhdr_size, objects_size)
}
pub fn decode_size(buf: &[u8], ctx: &FormatContext) -> FormatResult<usize> {
let ss = ctx.sizeof_size as usize;
let header_size = pad_to_8(4 + 1 + 3 + ss);
if buf.len() < header_size {
return Err(FormatError::BufferTooShort {
needed: header_size,
available: buf.len(),
});
}
if buf[0..4] != GCOL_SIGNATURE {
return Err(FormatError::InvalidSignature);
}
if buf[4] != GCOL_VERSION {
return Err(FormatError::InvalidVersion(buf[4]));
}
Ok(read_size(&buf[8..], ss) as usize)
}
pub fn decode(buf: &[u8], ctx: &FormatContext) -> FormatResult<(Self, usize)> {
let ss = ctx.sizeof_size as usize;
let header_size = pad_to_8(4 + 1 + 3 + ss);
let objhdr_size = pad_to_8(2 + 2 + 4 + ss);
if buf.len() < header_size {
return Err(FormatError::BufferTooShort {
needed: header_size,
available: buf.len(),
});
}
if buf[0..4] != GCOL_SIGNATURE {
return Err(FormatError::InvalidSignature);
}
let version = buf[4];
if version != GCOL_VERSION {
return Err(FormatError::InvalidVersion(version));
}
let collection_size = read_size(&buf[8..], ss) as usize;
if buf.len() < collection_size {
return Err(FormatError::BufferTooShort {
needed: collection_size,
available: buf.len(),
});
}
let mut pos = header_size;
let mut objects = Vec::new();
while pos + objhdr_size <= collection_size {
let obj_start = pos;
let index = u16::from_le_bytes([buf[pos], buf[pos + 1]]);
pos += 2;
let ref_count = u16::from_le_bytes([buf[pos], buf[pos + 1]]);
pos += 2;
let _reserved =
u32::from_le_bytes([buf[pos], buf[pos + 1], buf[pos + 2], buf[pos + 3]]);
pos += 4;
let size = read_size(&buf[pos..], ss) as usize;
pos = obj_start + objhdr_size;
if index == 0 {
break;
}
let obj_end = pos
.checked_add(size)
.filter(|&end| end <= collection_size)
.ok_or_else(|| {
FormatError::InvalidData(format!(
"global heap object {} extends past collection boundary",
index,
))
})?;
let data = buf[pos..obj_end].to_vec();
let padded = pad_to_8(size);
pos += padded;
objects.push(GlobalHeapObject {
index,
ref_count,
data,
});
}
Ok((Self { objects }, collection_size))
}
}
impl Default for GlobalHeapCollection {
fn default() -> Self {
Self::new()
}
}
pub fn encode_vlen_reference(
sequence_length: u32,
collection_addr: u64,
object_index: u32,
ctx: &FormatContext,
) -> Vec<u8> {
let sa = ctx.sizeof_addr as usize;
let mut buf = Vec::with_capacity(4 + sa + 4);
buf.extend_from_slice(&sequence_length.to_le_bytes());
buf.extend_from_slice(&collection_addr.to_le_bytes()[..sa]);
buf.extend_from_slice(&object_index.to_le_bytes());
buf
}
pub fn decode_vlen_reference(buf: &[u8], ctx: &FormatContext) -> FormatResult<(u32, u64, u32)> {
let sa = ctx.sizeof_addr as usize;
let total = 4 + sa + 4;
if buf.len() < total {
return Err(FormatError::BufferTooShort {
needed: total,
available: buf.len(),
});
}
let seq_len = u32::from_le_bytes([buf[0], buf[1], buf[2], buf[3]]);
let addr = read_size(&buf[4..], sa);
let index = u32::from_le_bytes([
buf[4 + sa],
buf[4 + sa + 1],
buf[4 + sa + 2],
buf[4 + sa + 3],
]);
Ok((seq_len, addr, index))
}
pub fn vlen_reference_size(ctx: &FormatContext) -> usize {
4 + ctx.sizeof_addr as usize + 4
}
pub fn vlen_seq_len(data_len: usize) -> FormatResult<u32> {
u32::try_from(data_len).map_err(|_| {
FormatError::InvalidData(format!(
"a {data_len}-byte vlen sequence does not fit the 32-bit length field"
))
})
}
fn pad_to_8(n: usize) -> usize {
(n + 7) & !7
}
#[cfg(test)]
mod tests {
use super::*;
fn ctx() -> FormatContext {
FormatContext {
sizeof_addr: 8,
sizeof_size: 8,
}
}
fn ctx4() -> FormatContext {
FormatContext {
sizeof_addr: 4,
sizeof_size: 4,
}
}
#[test]
fn empty_collection_roundtrip() {
let coll = GlobalHeapCollection::new();
let encoded = coll.encode(&ctx());
let (decoded, consumed) = GlobalHeapCollection::decode(&encoded, &ctx()).unwrap();
assert_eq!(consumed, encoded.len());
assert_eq!(decoded, coll);
assert!(decoded.objects.is_empty());
}
#[test]
fn single_object_roundtrip() {
let mut coll = GlobalHeapCollection::new();
let idx = coll.add_object(b"hello".to_vec()).unwrap();
assert_eq!(idx, 1);
let encoded = coll.encode(&ctx());
let (decoded, consumed) = GlobalHeapCollection::decode(&encoded, &ctx()).unwrap();
assert_eq!(consumed, encoded.len());
assert_eq!(decoded.objects.len(), 1);
assert_eq!(decoded.objects[0].index, 1);
assert_eq!(decoded.objects[0].data, b"hello");
}
#[test]
fn multiple_objects_roundtrip() {
let mut coll = GlobalHeapCollection::new();
let i1 = coll.add_object(b"alpha".to_vec()).unwrap();
let i2 = coll.add_object(b"beta".to_vec()).unwrap();
let i3 = coll.add_object(b"gamma delta".to_vec()).unwrap();
assert_eq!(i1, 1);
assert_eq!(i2, 2);
assert_eq!(i3, 3);
let encoded = coll.encode(&ctx());
let (decoded, _) = GlobalHeapCollection::decode(&encoded, &ctx()).unwrap();
assert_eq!(decoded.objects.len(), 3);
assert_eq!(decoded.get_object(1), Some(b"alpha".as_slice()));
assert_eq!(decoded.get_object(2), Some(b"beta".as_slice()));
assert_eq!(decoded.get_object(3), Some(b"gamma delta".as_slice()));
}
#[test]
fn get_object_not_found() {
let coll = GlobalHeapCollection::new();
assert_eq!(coll.get_object(1), None);
}
#[test]
fn padding_to_8() {
assert_eq!(pad_to_8(0), 0);
assert_eq!(pad_to_8(1), 8);
assert_eq!(pad_to_8(7), 8);
assert_eq!(pad_to_8(8), 8);
assert_eq!(pad_to_8(9), 16);
assert_eq!(pad_to_8(16), 16);
}
#[test]
fn vlen_reference_roundtrip() {
let c = ctx();
let encoded = encode_vlen_reference(5, 0x1234_5678_9ABC_DEF0, 42, &c);
assert_eq!(encoded.len(), vlen_reference_size(&c));
let (seq_len, addr, idx) = decode_vlen_reference(&encoded, &c).unwrap();
assert_eq!(seq_len, 5);
assert_eq!(addr, 0x1234_5678_9ABC_DEF0);
assert_eq!(idx, 42);
}
#[test]
fn vlen_reference_4byte_roundtrip() {
let c = ctx4();
let encoded = encode_vlen_reference(10, 0x1234_5678, 7, &c);
assert_eq!(encoded.len(), 12); let (seq_len, addr, idx) = decode_vlen_reference(&encoded, &c).unwrap();
assert_eq!(seq_len, 10);
assert_eq!(addr, 0x1234_5678);
assert_eq!(idx, 7);
}
#[test]
fn vlen_seq_len_bounds() {
assert_eq!(vlen_seq_len(0).unwrap(), 0);
assert_eq!(vlen_seq_len(u32::MAX as usize).unwrap(), u32::MAX);
#[cfg(target_pointer_width = "64")]
assert!(vlen_seq_len(u32::MAX as usize + 1).is_err());
}
#[test]
fn vlen_reference_size_check() {
assert_eq!(vlen_reference_size(&ctx()), 16);
assert_eq!(vlen_reference_size(&ctx4()), 12);
}
#[test]
fn decode_bad_signature() {
let mut buf = vec![0u8; 32];
buf[0..4].copy_from_slice(b"XYZW");
let err = GlobalHeapCollection::decode(&buf, &ctx()).unwrap_err();
assert!(matches!(err, FormatError::InvalidSignature));
}
#[test]
fn decode_bad_version() {
let coll = GlobalHeapCollection::new();
let mut encoded = coll.encode(&ctx());
encoded[4] = 99;
let err = GlobalHeapCollection::decode(&encoded, &ctx()).unwrap_err();
assert!(matches!(err, FormatError::InvalidVersion(99)));
}
#[test]
fn decode_buffer_too_short() {
let buf = [0u8; 4];
let err = GlobalHeapCollection::decode(&buf, &ctx()).unwrap_err();
assert!(matches!(err, FormatError::BufferTooShort { .. }));
}
#[test]
fn ctx4_roundtrip() {
let c = ctx4();
let mut coll = GlobalHeapCollection::new();
coll.add_object(b"test data".to_vec()).unwrap();
let encoded = coll.encode(&c);
let (decoded, consumed) = GlobalHeapCollection::decode(&encoded, &c).unwrap();
assert_eq!(consumed, encoded.len());
assert_eq!(decoded.get_object(1), Some(b"test data".as_slice()));
}
#[test]
fn object_data_alignment() {
let mut coll = GlobalHeapCollection::new();
coll.add_object(vec![1]).unwrap(); coll.add_object(vec![2, 3, 4, 5, 6, 7, 8, 9, 10]).unwrap(); coll.add_object(vec![11, 12, 13, 14, 15, 16, 17, 18])
.unwrap();
let encoded = coll.encode(&ctx());
let (decoded, _) = GlobalHeapCollection::decode(&encoded, &ctx()).unwrap();
assert_eq!(decoded.get_object(1), Some([1u8].as_slice()));
assert_eq!(
decoded.get_object(2),
Some([2, 3, 4, 5, 6, 7, 8, 9, 10].as_slice())
);
assert_eq!(
decoded.get_object(3),
Some([11, 12, 13, 14, 15, 16, 17, 18].as_slice())
);
}
#[test]
fn empty_data_object() {
let mut coll = GlobalHeapCollection::new();
coll.add_object(vec![]).unwrap();
let encoded = coll.encode(&ctx());
let (decoded, _) = GlobalHeapCollection::decode(&encoded, &ctx()).unwrap();
assert_eq!(decoded.get_object(1), Some([].as_slice()));
}
#[test]
fn remove_object_keeps_the_survivors_indices() {
let mut coll = GlobalHeapCollection::new();
assert_eq!(coll.add_object(b"first".to_vec()).unwrap(), 1);
assert_eq!(coll.add_object(b"second".to_vec()).unwrap(), 2);
assert_eq!(coll.add_object(b"third".to_vec()).unwrap(), 3);
assert!(coll.remove_object(2));
assert!(!coll.remove_object(2));
let encoded = coll.encode(&ctx());
let (decoded, _) = GlobalHeapCollection::decode(&encoded, &ctx()).unwrap();
assert_eq!(decoded.get_object(1), Some(b"first".as_slice()));
assert_eq!(decoded.get_object(2), None);
assert_eq!(decoded.get_object(3), Some(b"third".as_slice()));
assert!(coll.remove_object(1) && coll.remove_object(3));
assert!(coll.is_empty());
}
#[test]
fn encode_at_size_holds_an_oversized_collection_open() {
let mut coll = GlobalHeapCollection::new();
for _ in 0..4 {
coll.add_object(vec![0xab; 2000]).unwrap();
}
let block = coll.encode(&ctx()).len();
assert!(block > GCOL_MIN_SIZE, "{block} is not above the minimum");
coll.remove_object(2);
assert!(
coll.encoded_size(&ctx()) < block,
"the test needs the natural size to have shrunk"
);
let held = coll.encode_at_size(&ctx(), block).unwrap();
assert_eq!(held.len(), block);
assert_eq!(
GlobalHeapCollection::decode(&held, &ctx()).unwrap().1,
block,
"the collection must still declare the whole block"
);
let (decoded, _) = GlobalHeapCollection::decode(&held, &ctx()).unwrap();
assert_eq!(decoded.get_object(3), Some(vec![0xab; 2000].as_slice()));
assert!(coll.encode_at_size(&ctx(), 64).is_err());
}
#[test]
fn new_objects_encode_a_zero_ref_count() {
let mut coll = GlobalHeapCollection::new();
coll.add_object(b"x".to_vec()).unwrap();
let img = coll.encode(&ctx());
let header_size = pad_to_8(4 + 1 + 3 + ctx().sizeof_size as usize);
assert_eq!(&img[header_size + 2..header_size + 4], &[0, 0]);
}
#[test]
fn rewrite_preserves_a_foreign_objects_ref_count() {
let mut coll = GlobalHeapCollection::new();
coll.add_object(b"doomed".to_vec()).unwrap();
coll.add_object(b"vds linked".to_vec()).unwrap();
coll.objects[1].ref_count = 3;
let img = coll.encode(&ctx());
let (mut back, _) = GlobalHeapCollection::decode(&img, &ctx()).unwrap();
assert_eq!(back.objects[1].ref_count, 3, "decode must keep the count");
assert!(back.remove_object(1));
let rewritten = back.encode_at_size(&ctx(), img.len()).unwrap();
let (again, _) = GlobalHeapCollection::decode(&rewritten, &ctx()).unwrap();
assert_eq!(again.get_object(2), Some(b"vds linked".as_slice()));
assert_eq!(again.objects[0].ref_count, 3, "rewrite must keep the count");
}
}