use crate::egress::wire::varint;
use crate::error::{Result, fmt};
pub(crate) const MAX_CONN_DICT_HEAP_BYTES: usize = 256 * 1024 * 1024;
pub(crate) const MAX_CONN_DICT_SIZE: usize = 8_388_608;
#[derive(Debug, Clone, Copy)]
#[repr(C)]
pub struct SymbolEntry {
pub offset: u32,
pub len: u32,
}
#[derive(Debug, Default, Clone)]
pub struct SymbolDict {
arena: Vec<u8>,
entries: Vec<SymbolEntry>,
}
impl SymbolDict {
pub fn new() -> Self {
Self::default()
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn heap_bytes(&self) -> usize {
self.arena.len()
}
pub fn get(&self, id: u32) -> Option<&str> {
let entry = self.entries.get(id as usize)?;
let start = entry.offset as usize;
let end = start + entry.len as usize;
debug_assert!(
end <= self.arena.len(),
"entry {id} offset+len={end} exceeds arena len {}",
self.arena.len()
);
Some(unsafe { std::str::from_utf8_unchecked(&self.arena[start..end]) })
}
pub fn arena(&self) -> &[u8] {
&self.arena
}
pub fn entries(&self) -> &[SymbolEntry] {
&self.entries
}
pub fn reset(&mut self) {
self.entries.clear();
self.arena.clear();
self.entries.shrink_to(1024);
self.arena.shrink_to(64 * 1024);
}
pub fn apply_delta<'a, I>(&mut self, delta_start: u64, entries: I) -> Result<()>
where
I: IntoIterator<Item = &'a [u8]>,
{
let expected = self.entries.len() as u64;
if delta_start != expected {
return Err(fmt!(
ProtocolError,
"symbol dict delta_start={} but registry len={}",
delta_start,
expected
));
}
for bytes in entries {
self.push_one(bytes)?;
}
Ok(())
}
pub fn apply_delta_from_bytes(&mut self, bytes: &[u8]) -> Result<usize> {
let mut cursor = 0usize;
let (delta_start, n) = varint::decode_u64(&bytes[cursor..])?;
cursor += n;
let (delta_count, n) = varint::decode_u64(&bytes[cursor..])?;
cursor += n;
let expected = self.entries.len() as u64;
if delta_start != expected {
return Err(fmt!(
ProtocolError,
"symbol dict delta_start={} but registry len={}",
delta_start,
expected
));
}
let headroom = MAX_CONN_DICT_SIZE.saturating_sub(self.entries.len()) as u64;
if delta_count > headroom {
return Err(fmt!(
ProtocolError,
"symbol dict delta_count={} exceeds remaining capacity {} \
(current entries={}, max={})",
delta_count,
headroom,
self.entries.len(),
MAX_CONN_DICT_SIZE
));
}
let snapshot_entries = self.entries.len();
let snapshot_arena = self.arena.len();
let result: Result<usize> = (|| {
for i in 0..delta_count {
let (entry_len, n) = varint::decode_usize(&bytes[cursor..])?;
cursor += n;
let end = cursor.checked_add(entry_len).ok_or_else(|| {
fmt!(
ProtocolError,
"symbol dict entry length overflow at i={}",
i
)
})?;
if end > bytes.len() {
return Err(fmt!(
ProtocolError,
"symbol dict truncated at entry {}: need {} bytes, have {}",
i,
entry_len,
bytes.len() - cursor
));
}
self.push_one(&bytes[cursor..end])?;
cursor = end;
}
Ok(cursor)
})();
if result.is_err() {
self.entries.truncate(snapshot_entries);
self.arena.truncate(snapshot_arena);
}
result
}
fn push_one(&mut self, bytes: &[u8]) -> Result<()> {
let s = std::str::from_utf8(bytes).map_err(|e| {
fmt!(
InvalidUtf8,
"symbol dict entry {} is not valid UTF-8: {}",
self.entries.len(),
e
)
})?;
if self.entries.len() >= MAX_CONN_DICT_SIZE {
return Err(fmt!(
ProtocolError,
"symbol dict full: {} entries (max {}); server must emit \
CACHE_RESET(dict) before adding more",
self.entries.len(),
MAX_CONN_DICT_SIZE
));
}
let new_heap = self
.arena
.len()
.checked_add(s.len())
.ok_or_else(|| fmt!(ProtocolError, "symbol dict heap overflow"))?;
if new_heap > MAX_CONN_DICT_HEAP_BYTES {
return Err(fmt!(
ProtocolError,
"symbol dict heap would reach {} bytes (max {}); server \
must emit CACHE_RESET(dict) before adding more",
new_heap,
MAX_CONN_DICT_HEAP_BYTES
));
}
let offset = u32::try_from(self.arena.len())
.map_err(|_| fmt!(ProtocolError, "symbol dict arena exceeds u32"))?;
let len = u32::try_from(s.len())
.map_err(|_| fmt!(ProtocolError, "symbol dict entry exceeds u32 length"))?;
self.arena.extend_from_slice(s.as_bytes());
self.entries.push(SymbolEntry { offset, len });
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::egress::wire::varint::encode_u64;
use crate::error::ErrorCode;
fn build_delta(start: u64, entries: &[&str]) -> Vec<u8> {
let mut out = Vec::new();
encode_u64(start, &mut out);
encode_u64(entries.len() as u64, &mut out);
for e in entries {
encode_u64(e.len() as u64, &mut out);
out.extend_from_slice(e.as_bytes());
}
out
}
#[test]
fn empty_dict() {
let d = SymbolDict::new();
assert_eq!(d.len(), 0);
assert!(d.is_empty());
assert_eq!(d.heap_bytes(), 0);
assert!(d.get(0).is_none());
}
#[test]
fn apply_first_delta_via_iter() {
let mut d = SymbolDict::new();
let entries: Vec<&[u8]> = vec![b"AAPL", b"MSFT", b"GOOG"];
d.apply_delta(0, entries).unwrap();
assert_eq!(d.len(), 3);
assert_eq!(d.get(0), Some("AAPL"));
assert_eq!(d.get(1), Some("MSFT"));
assert_eq!(d.get(2), Some("GOOG"));
assert_eq!(d.get(3), None);
assert_eq!(d.heap_bytes(), 4 + 4 + 4);
}
#[test]
fn second_delta_appends() {
let mut d = SymbolDict::new();
d.apply_delta(0, [b"a".as_slice()]).unwrap();
d.apply_delta(1, [b"bb".as_slice(), b"ccc".as_slice()])
.unwrap();
assert_eq!(d.len(), 3);
assert_eq!(d.get(2), Some("ccc"));
}
#[test]
fn delta_start_mismatch_rejected() {
let mut d = SymbolDict::new();
d.apply_delta(0, [b"x".as_slice()]).unwrap();
let err = d.apply_delta(5, [b"y".as_slice()]).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn from_bytes_roundtrip() {
let mut d = SymbolDict::new();
let bytes = build_delta(0, &["AAPL", "MSFT"]);
let consumed = d.apply_delta_from_bytes(&bytes).unwrap();
assert_eq!(consumed, bytes.len());
assert_eq!(d.get(0), Some("AAPL"));
assert_eq!(d.get(1), Some("MSFT"));
let bytes2 = build_delta(2, &["GOOG"]);
d.apply_delta_from_bytes(&bytes2).unwrap();
assert_eq!(d.get(2), Some("GOOG"));
}
#[test]
fn from_bytes_partial_failure_rolls_back() {
let mut d = SymbolDict::new();
d.apply_delta(0, [b"first".as_slice()]).unwrap();
let snapshot_len = d.len();
let snapshot_heap = d.heap_bytes();
let mut bytes = Vec::new();
encode_u64(snapshot_len as u64, &mut bytes); encode_u64(2, &mut bytes); encode_u64(2, &mut bytes); bytes.extend_from_slice(b"ok");
encode_u64(10, &mut bytes); bytes.extend_from_slice(b"abc");
let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert_eq!(d.len(), snapshot_len);
assert_eq!(d.heap_bytes(), snapshot_heap);
let next = build_delta(snapshot_len as u64, &["recovered"]);
d.apply_delta_from_bytes(&next).unwrap();
assert_eq!(d.get(snapshot_len as u32), Some("recovered"));
}
#[test]
fn from_bytes_truncated_entry_rejected() {
let mut d = SymbolDict::new();
let mut bytes = build_delta(0, &["hello"]);
bytes.truncate(bytes.len() - 1); let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn from_bytes_invalid_utf8_rejected() {
let mut bytes = Vec::new();
encode_u64(0, &mut bytes);
encode_u64(1, &mut bytes);
encode_u64(2, &mut bytes);
bytes.extend_from_slice(&[0xFF, 0xFE]); let mut d = SymbolDict::new();
let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
assert_eq!(err.code(), ErrorCode::InvalidUtf8);
}
#[test]
fn reset_clears_state() {
let mut d = SymbolDict::new();
d.apply_delta(0, [b"x".as_slice(), b"yy".as_slice()])
.unwrap();
assert_eq!(d.len(), 2);
d.reset();
assert_eq!(d.len(), 0);
assert_eq!(d.heap_bytes(), 0);
d.apply_delta(0, [b"new".as_slice()]).unwrap();
assert_eq!(d.get(0), Some("new"));
}
#[test]
fn delta_with_zero_entries_is_noop() {
let mut d = SymbolDict::new();
d.apply_delta(0, std::iter::empty::<&[u8]>()).unwrap();
let bytes = build_delta(0, &[]);
let consumed = d.apply_delta_from_bytes(&bytes).unwrap();
assert_eq!(consumed, bytes.len());
assert_eq!(d.len(), 0);
}
#[test]
fn delta_count_exceeding_capacity_rejected_upfront() {
let mut d = SymbolDict::new();
let mut bytes = Vec::new();
encode_u64(0, &mut bytes); encode_u64(u64::MAX, &mut bytes); let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(
err.msg().contains("exceeds remaining capacity"),
"expected upfront-cap rejection, got: {}",
err.msg()
);
assert_eq!(d.len(), 0);
}
#[test]
fn unicode_entries_preserved() {
let mut d = SymbolDict::new();
let bytes = build_delta(0, &["café", "日本語"]);
d.apply_delta_from_bytes(&bytes).unwrap();
assert_eq!(d.get(0), Some("café"));
assert_eq!(d.get(1), Some("日本語"));
}
}