use alloc::string::String;
use alloc::vec::Vec;
use super::Bundle;
use crate::bytes::Bytes;
use crate::persistence::NamespaceSummary;
pub const MAGIC: &[u8; 8] = b"CUBECLB\x01";
pub const FORMAT_VERSION: u32 = 1;
pub(crate) const ENTRY_SIZE: usize = 20;
pub fn flat_bundle_version(bytes: &[u8]) -> Option<u32> {
if bytes.len() < MAGIC.len() + 4 || &bytes[..MAGIC.len()] != MAGIC {
return None;
}
read_u32(bytes, MAGIC.len())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EmbeddedBundleError {
NotABundle,
UnsupportedFormat(u32),
Corrupted(&'static str),
}
impl core::fmt::Display for EmbeddedBundleError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::NotABundle => write!(f, "not a cubecl flat bundle"),
Self::UnsupportedFormat(found) => write!(
f,
"flat bundle format {found} is not {FORMAT_VERSION}, its entries are unreadable"
),
Self::Corrupted(what) => write!(f, "corrupted flat bundle: {what}"),
}
}
}
impl core::error::Error for EmbeddedBundleError {}
#[derive(Debug)]
pub struct EmbeddedBundle {
bytes: Bytes,
metadata: (usize, usize),
namespaces: Vec<(usize, usize)>,
entries: (usize, usize),
data: usize,
}
impl EmbeddedBundle {
pub fn open(bytes: Bytes) -> Result<Self, EmbeddedBundleError> {
Self::parse(bytes.shared())
}
pub fn from_static(blob: &'static [u8]) -> Result<Self, EmbeddedBundleError> {
#[cfg(feature = "shared-bytes")]
let bytes = Bytes::from_shared(
bytes::Bytes::from_static(blob),
crate::bytes::AllocationProperty::Other,
);
#[cfg(not(feature = "shared-bytes"))]
let bytes = Bytes::from_bytes_vec(blob.to_vec());
Self::open(bytes)
}
fn parse(bytes: Bytes) -> Result<Self, EmbeddedBundleError> {
use EmbeddedBundleError::{Corrupted, NotABundle, UnsupportedFormat};
if bytes.len() < MAGIC.len() + 8 || &bytes[..MAGIC.len()] != MAGIC {
return Err(NotABundle);
}
let mut cursor = MAGIC.len();
let take_u32 = |cursor: &mut usize| -> Option<u32> {
let value = read_u32(&bytes, *cursor)?;
*cursor += 4;
Some(value)
};
let format = take_u32(&mut cursor).ok_or(Corrupted("truncated header"))?;
if format != FORMAT_VERSION {
return Err(UnsupportedFormat(format));
}
let metadata_len = take_u32(&mut cursor).ok_or(Corrupted("truncated header"))? as usize;
let metadata_start = cursor;
let metadata_end = metadata_start
.checked_add(metadata_len)
.ok_or(Corrupted("metadata length overflows"))?;
if metadata_end > bytes.len() {
return Err(Corrupted("metadata runs past the end"));
}
cursor = metadata_end;
let namespace_count = take_u32(&mut cursor).ok_or(Corrupted("truncated header"))? as usize;
let mut namespaces = Vec::with_capacity(namespace_count.min(bytes.len() / 4));
for _ in 0..namespace_count {
let len =
read_u32(&bytes, cursor).ok_or(Corrupted("truncated namespace table"))? as usize;
let start = cursor + 4;
cursor = start
.checked_add(len)
.ok_or(Corrupted("namespace length overflows"))?;
if cursor > bytes.len() {
return Err(Corrupted("namespace table runs past the end"));
}
namespaces.push((start, len));
}
let entry_count = read_u32(&bytes, cursor).ok_or(Corrupted("truncated header"))? as usize;
cursor += 4;
let entries_start = cursor;
let index_len = entry_count
.checked_mul(ENTRY_SIZE)
.ok_or(Corrupted("entry count overflows"))?;
let data_start = entries_start
.checked_add(index_len)
.ok_or(Corrupted("entry index overflows"))?;
if data_start > bytes.len() {
return Err(Corrupted("entry index runs past the end"));
}
let this = Self {
bytes,
metadata: (metadata_start, metadata_end),
namespaces,
entries: (entries_start, entry_count),
data: data_start,
};
this.validate_entries()?;
Ok(this)
}
fn validate_entries(&self) -> Result<(), EmbeddedBundleError> {
use EmbeddedBundleError::Corrupted;
let available = self.bytes.len() - self.data;
let namespace_count = self.namespaces.len();
let mut previous: Option<Entry> = None;
for index in 0..self.entries.1 {
let entry = self
.entry(index)
.ok_or(Corrupted("truncated entry index"))?;
if entry.namespace as usize >= namespace_count {
return Err(Corrupted("entry points at an unknown namespace"));
}
for (offset, len) in [
(entry.key_offset, entry.key_len),
(entry.value_offset, entry.value_len),
] {
let end = (offset as usize)
.checked_add(len as usize)
.ok_or(Corrupted("entry span overflows"))?;
if end > available {
return Err(Corrupted("entry span runs past the end"));
}
}
if let Some(previous) = &previous {
let order = (previous.namespace, self.key_of(previous))
.cmp(&(entry.namespace, self.key_of(&entry)));
if order != core::cmp::Ordering::Less {
return Err(Corrupted("entry index is not sorted by (namespace, key)"));
}
}
previous = Some(entry);
}
for index in 0..namespace_count {
self.namespace(index)
.ok_or(Corrupted("namespace is not valid UTF-8"))?;
}
Ok(())
}
pub fn metadata(&self) -> &[u8] {
&self.bytes[self.metadata.0..self.metadata.1]
}
pub fn manifest(&self) -> Result<super::BundleManifest, super::BundleError> {
let manifest = super::BundleManifest::parse(self.metadata())?;
manifest.warn_on_version_mismatch();
Ok(manifest)
}
pub fn len(&self) -> usize {
self.entries.1
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn summary(&self) -> Vec<NamespaceSummary> {
let mut summary: Vec<NamespaceSummary> = (0..self.namespaces.len())
.filter_map(|index| {
Some(NamespaceSummary {
namespace: alloc::string::ToString::to_string(self.namespace(index)?),
entries: 0,
bytes: 0,
})
})
.collect();
for index in 0..self.entries.1 {
let Some(entry) = self.entry(index) else {
break;
};
if let Some(namespace) = summary.get_mut(entry.namespace as usize) {
namespace.entries += 1;
namespace.bytes += u64::from(entry.key_len) + u64::from(entry.value_len);
}
}
summary
}
fn namespace(&self, index: usize) -> Option<&str> {
let &(start, len) = self.namespaces.get(index)?;
core::str::from_utf8(self.bytes.get(start..start + len)?).ok()
}
fn namespace_id(&self, namespace: &str) -> Option<u32> {
let at = lower_bound(self.namespaces.len(), |index| {
Some(self.namespace(index)?.cmp(namespace))
})?;
(self.namespace(at)? == namespace).then_some(at as u32)
}
fn entry(&self, index: usize) -> Option<Entry> {
let at = self.entries.0 + index * ENTRY_SIZE;
Some(Entry {
namespace: read_u32(&self.bytes, at)?,
key_offset: read_u32(&self.bytes, at + 4)?,
key_len: read_u32(&self.bytes, at + 8)?,
value_offset: read_u32(&self.bytes, at + 12)?,
value_len: read_u32(&self.bytes, at + 16)?,
})
}
fn span(&self, offset: u32, len: u32) -> core::ops::Range<usize> {
let start = self.data + offset as usize;
start..start + len as usize
}
fn key_of(&self, entry: &Entry) -> &[u8] {
&self.bytes[self.span(entry.key_offset, entry.key_len)]
}
fn value_of(&self, entry: &Entry) -> &[u8] {
&self.bytes[self.span(entry.value_offset, entry.value_len)]
}
fn value_window(&self, entry: &Entry) -> Option<Bytes> {
let span = self.span(entry.value_offset, entry.value_len);
self.bytes
.view(span.start, span.end)
.inspect_err(|err| log::warn!("Embedded bundle: can't view an entry: {err:?}"))
.ok()
}
fn first_of(&self, namespace: u32) -> Option<usize> {
let at = lower_bound(self.entries.1, |index| {
Some(self.entry(index)?.namespace.cmp(&namespace))
})?;
(self.entry(at)?.namespace == namespace).then_some(at)
}
}
impl Bundle for EmbeddedBundle {
fn get(&self, namespace: &str, key: &[u8]) -> Option<Bytes> {
let namespace = self.namespace_id(namespace)?;
let at = lower_bound(self.entries.1, |index| {
let entry = self.entry(index)?;
Some((entry.namespace, self.key_of(&entry)).cmp(&(namespace, key)))
})?;
let entry = self.entry(at)?;
((entry.namespace, self.key_of(&entry)) == (namespace, key))
.then(|| self.value_window(&entry))
.flatten()
}
fn scan(&self, namespace: &str, visit: &mut dyn FnMut(&[u8], &[u8])) {
let Some(id) = self.namespace_id(namespace) else {
return;
};
let Some(first) = self.first_of(id) else {
return;
};
for index in first..self.entries.1 {
let Some(entry) = self.entry(index) else {
return;
};
if entry.namespace != id {
return;
}
visit(self.key_of(&entry), self.value_of(&entry));
}
}
fn namespaces(&self) -> Vec<String> {
(0..self.namespaces.len())
.filter_map(|index| Some(alloc::string::ToString::to_string(self.namespace(index)?)))
.collect()
}
fn describe(&self) -> String {
alloc::format!("embedded bundle ({} entries)", self.entries.1)
}
}
struct Entry {
namespace: u32,
key_offset: u32,
key_len: u32,
value_offset: u32,
value_len: u32,
}
fn lower_bound(
len: usize,
compare: impl Fn(usize) -> Option<core::cmp::Ordering>,
) -> Option<usize> {
let (mut low, mut high) = (0usize, len);
while low < high {
let mid = low + (high - low) / 2;
if compare(mid)? == core::cmp::Ordering::Less {
low = mid + 1;
} else {
high = mid;
}
}
Some(low)
}
fn read_u32(bytes: &[u8], at: usize) -> Option<u32> {
let raw: [u8; 4] = bytes.get(at..at + 4)?.try_into().ok()?;
Some(u32::from_le_bytes(raw))
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
fn empty_bundle() -> Bytes {
let mut bytes = Vec::new();
bytes.extend_from_slice(MAGIC);
bytes.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&0u32.to_le_bytes()); Bytes::from_bytes_vec(bytes)
}
#[test]
fn an_empty_bundle_reads_as_empty() {
let bundle = EmbeddedBundle::open(empty_bundle()).unwrap();
assert!(bundle.is_empty());
assert_eq!(bundle.get("anything", b"key"), None);
bundle.scan("anything", &mut |_, _| panic!("no entries to visit"));
}
#[test]
fn foreign_bytes_are_rejected() {
assert_eq!(
EmbeddedBundle::open(Bytes::from_bytes_vec(b"definitely not a bundle".to_vec()))
.unwrap_err(),
EmbeddedBundleError::NotABundle
);
assert_eq!(
EmbeddedBundle::open(Bytes::from_bytes_vec(vec![])).unwrap_err(),
EmbeddedBundleError::NotABundle
);
}
#[test]
fn another_format_version_is_rejected() {
let mut bytes = empty_bundle().to_vec();
bytes[MAGIC.len()..MAGIC.len() + 4].copy_from_slice(&99u32.to_le_bytes());
let bytes = Bytes::from_bytes_vec(bytes);
assert_eq!(
EmbeddedBundle::open(bytes).unwrap_err(),
EmbeddedBundleError::UnsupportedFormat(99)
);
}
#[test]
fn truncated_bytes_are_rejected() {
let full = empty_bundle();
for len in MAGIC.len()..full.len() {
let result = EmbeddedBundle::open(Bytes::from_bytes_vec(full[..len].to_vec()));
assert!(result.is_err(), "a {len}-byte prefix must not open");
}
}
#[test]
fn out_of_range_spans_are_rejected() {
let mut bytes = Vec::new();
bytes.extend_from_slice(MAGIC);
bytes.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&1u32.to_le_bytes()); bytes.extend_from_slice(&2u32.to_le_bytes());
bytes.extend_from_slice(b"ns");
bytes.extend_from_slice(&1u32.to_le_bytes()); bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&99u32.to_le_bytes()); bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&0u32.to_le_bytes());
assert!(matches!(
EmbeddedBundle::open(Bytes::from_bytes_vec(bytes)).unwrap_err(),
EmbeddedBundleError::Corrupted(_)
));
}
#[test]
fn an_unsorted_entry_index_is_rejected() {
let bundle = |keys: [&[u8]; 2]| {
let mut data = Vec::new();
let mut index = Vec::new();
for key in keys {
index.extend_from_slice(&0u32.to_le_bytes()); index.extend_from_slice(&(data.len() as u32).to_le_bytes());
index.extend_from_slice(&(key.len() as u32).to_le_bytes());
data.extend_from_slice(key);
index.extend_from_slice(&(data.len() as u32).to_le_bytes());
index.extend_from_slice(&0u32.to_le_bytes()); }
let mut bytes = Vec::new();
bytes.extend_from_slice(MAGIC);
bytes.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&1u32.to_le_bytes()); bytes.extend_from_slice(&2u32.to_le_bytes());
bytes.extend_from_slice(b"ns");
bytes.extend_from_slice(&2u32.to_le_bytes()); bytes.extend_from_slice(&index);
bytes.extend_from_slice(&data);
EmbeddedBundle::open(Bytes::from_bytes_vec(bytes))
};
assert!(bundle([b"a", b"b"]).is_ok());
assert!(matches!(
bundle([b"b", b"a"]).unwrap_err(),
EmbeddedBundleError::Corrupted(_)
));
assert!(matches!(
bundle([b"a", b"a"]).unwrap_err(),
EmbeddedBundleError::Corrupted(_)
));
}
#[test]
fn an_unknown_namespace_id_is_rejected() {
let mut bytes = Vec::new();
bytes.extend_from_slice(MAGIC);
bytes.extend_from_slice(&FORMAT_VERSION.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&1u32.to_le_bytes()); bytes.extend_from_slice(&7u32.to_le_bytes()); bytes.extend_from_slice(&0u32.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes());
bytes.extend_from_slice(&0u32.to_le_bytes());
assert!(matches!(
EmbeddedBundle::open(Bytes::from_bytes_vec(bytes)).unwrap_err(),
EmbeddedBundleError::Corrupted(_)
));
}
}