1use crate::egress::wire::varint;
46use crate::error::{Result, fmt};
47
48pub(crate) const MAX_CONN_DICT_HEAP_BYTES: usize = 256 * 1024 * 1024;
54
55pub(crate) const MAX_CONN_DICT_SIZE: usize = 8_388_608;
58
59#[derive(Debug, Clone, Copy)]
61#[repr(C)]
62pub struct SymbolEntry {
63 pub offset: u32,
64 pub len: u32,
65}
66
67#[derive(Debug, Default, Clone)]
69pub struct SymbolDict {
70 arena: Vec<u8>,
71 entries: Vec<SymbolEntry>,
72}
73
74impl SymbolDict {
75 pub fn new() -> Self {
76 Self::default()
77 }
78
79 pub fn len(&self) -> usize {
81 self.entries.len()
82 }
83
84 pub fn is_empty(&self) -> bool {
85 self.entries.is_empty()
86 }
87
88 pub fn heap_bytes(&self) -> usize {
90 self.arena.len()
91 }
92
93 pub fn get(&self, id: u32) -> Option<&str> {
95 let entry = self.entries.get(id as usize)?;
96 let start = entry.offset as usize;
97 let end = start + entry.len as usize;
98 debug_assert!(
99 end <= self.arena.len(),
100 "entry {id} offset+len={end} exceeds arena len {}",
101 self.arena.len()
102 );
103 Some(unsafe { std::str::from_utf8_unchecked(&self.arena[start..end]) })
106 }
107
108 pub fn arena(&self) -> &[u8] {
110 &self.arena
111 }
112
113 pub fn entries(&self) -> &[SymbolEntry] {
116 &self.entries
117 }
118
119 pub fn reset(&mut self) {
123 self.entries.clear();
124 self.arena.clear();
125 self.entries.shrink_to(1024);
126 self.arena.shrink_to(64 * 1024);
127 }
128
129 pub fn apply_delta<'a, I>(&mut self, delta_start: u64, entries: I) -> Result<()>
132 where
133 I: IntoIterator<Item = &'a [u8]>,
134 {
135 let expected = self.entries.len() as u64;
136 if delta_start != expected {
137 return Err(fmt!(
138 ProtocolError,
139 "symbol dict delta_start={} but registry len={}",
140 delta_start,
141 expected
142 ));
143 }
144 for bytes in entries {
145 self.push_one(bytes)?;
146 }
147 Ok(())
148 }
149
150 pub fn apply_delta_from_bytes(&mut self, bytes: &[u8]) -> Result<usize> {
159 let mut cursor = 0usize;
160 let (delta_start, n) = varint::decode_u64(&bytes[cursor..])?;
161 cursor += n;
162 let (delta_count, n) = varint::decode_u64(&bytes[cursor..])?;
163 cursor += n;
164
165 let expected = self.entries.len() as u64;
166 if delta_start != expected {
167 return Err(fmt!(
168 ProtocolError,
169 "symbol dict delta_start={} but registry len={}",
170 delta_start,
171 expected
172 ));
173 }
174
175 let headroom = MAX_CONN_DICT_SIZE.saturating_sub(self.entries.len()) as u64;
191 if delta_count > headroom {
192 return Err(fmt!(
193 ProtocolError,
194 "symbol dict delta_count={} exceeds remaining capacity {} \
195 (current entries={}, max={})",
196 delta_count,
197 headroom,
198 self.entries.len(),
199 MAX_CONN_DICT_SIZE
200 ));
201 }
202
203 let snapshot_entries = self.entries.len();
204 let snapshot_arena = self.arena.len();
205 let result: Result<usize> = (|| {
206 for i in 0..delta_count {
207 let (entry_len, n) = varint::decode_usize(&bytes[cursor..])?;
208 cursor += n;
209 let end = cursor.checked_add(entry_len).ok_or_else(|| {
210 fmt!(
211 ProtocolError,
212 "symbol dict entry length overflow at i={}",
213 i
214 )
215 })?;
216 if end > bytes.len() {
217 return Err(fmt!(
218 ProtocolError,
219 "symbol dict truncated at entry {}: need {} bytes, have {}",
220 i,
221 entry_len,
222 bytes.len() - cursor
223 ));
224 }
225 self.push_one(&bytes[cursor..end])?;
226 cursor = end;
227 }
228 Ok(cursor)
229 })();
230 if result.is_err() {
231 self.entries.truncate(snapshot_entries);
232 self.arena.truncate(snapshot_arena);
233 }
234 result
235 }
236
237 fn push_one(&mut self, bytes: &[u8]) -> Result<()> {
238 let s = std::str::from_utf8(bytes).map_err(|e| {
239 fmt!(
240 InvalidUtf8,
241 "symbol dict entry {} is not valid UTF-8: {}",
242 self.entries.len(),
243 e
244 )
245 })?;
246 if self.entries.len() >= MAX_CONN_DICT_SIZE {
247 return Err(fmt!(
248 ProtocolError,
249 "symbol dict full: {} entries (max {}); server must emit \
250 CACHE_RESET(dict) before adding more",
251 self.entries.len(),
252 MAX_CONN_DICT_SIZE
253 ));
254 }
255 let new_heap = self
256 .arena
257 .len()
258 .checked_add(s.len())
259 .ok_or_else(|| fmt!(ProtocolError, "symbol dict heap overflow"))?;
260 if new_heap > MAX_CONN_DICT_HEAP_BYTES {
261 return Err(fmt!(
262 ProtocolError,
263 "symbol dict heap would reach {} bytes (max {}); server \
264 must emit CACHE_RESET(dict) before adding more",
265 new_heap,
266 MAX_CONN_DICT_HEAP_BYTES
267 ));
268 }
269 let offset = u32::try_from(self.arena.len())
270 .map_err(|_| fmt!(ProtocolError, "symbol dict arena exceeds u32"))?;
271 let len = u32::try_from(s.len())
272 .map_err(|_| fmt!(ProtocolError, "symbol dict entry exceeds u32 length"))?;
273 self.arena.extend_from_slice(s.as_bytes());
274 self.entries.push(SymbolEntry { offset, len });
275 Ok(())
276 }
277}
278
279#[cfg(test)]
280mod tests {
281 use super::*;
282 use crate::egress::wire::varint::encode_u64;
283 use crate::error::ErrorCode;
284
285 fn build_delta(start: u64, entries: &[&str]) -> Vec<u8> {
286 let mut out = Vec::new();
287 encode_u64(start, &mut out);
288 encode_u64(entries.len() as u64, &mut out);
289 for e in entries {
290 encode_u64(e.len() as u64, &mut out);
291 out.extend_from_slice(e.as_bytes());
292 }
293 out
294 }
295
296 #[test]
297 fn empty_dict() {
298 let d = SymbolDict::new();
299 assert_eq!(d.len(), 0);
300 assert!(d.is_empty());
301 assert_eq!(d.heap_bytes(), 0);
302 assert!(d.get(0).is_none());
303 }
304
305 #[test]
306 fn apply_first_delta_via_iter() {
307 let mut d = SymbolDict::new();
308 let entries: Vec<&[u8]> = vec![b"AAPL", b"MSFT", b"GOOG"];
309 d.apply_delta(0, entries).unwrap();
310 assert_eq!(d.len(), 3);
311 assert_eq!(d.get(0), Some("AAPL"));
312 assert_eq!(d.get(1), Some("MSFT"));
313 assert_eq!(d.get(2), Some("GOOG"));
314 assert_eq!(d.get(3), None);
315 assert_eq!(d.heap_bytes(), 4 + 4 + 4);
316 }
317
318 #[test]
319 fn second_delta_appends() {
320 let mut d = SymbolDict::new();
321 d.apply_delta(0, [b"a".as_slice()]).unwrap();
322 d.apply_delta(1, [b"bb".as_slice(), b"ccc".as_slice()])
323 .unwrap();
324 assert_eq!(d.len(), 3);
325 assert_eq!(d.get(2), Some("ccc"));
326 }
327
328 #[test]
329 fn delta_start_mismatch_rejected() {
330 let mut d = SymbolDict::new();
331 d.apply_delta(0, [b"x".as_slice()]).unwrap();
332 let err = d.apply_delta(5, [b"y".as_slice()]).unwrap_err();
334 assert_eq!(err.code(), ErrorCode::ProtocolError);
335 }
336
337 #[test]
338 fn from_bytes_roundtrip() {
339 let mut d = SymbolDict::new();
340 let bytes = build_delta(0, &["AAPL", "MSFT"]);
341 let consumed = d.apply_delta_from_bytes(&bytes).unwrap();
342 assert_eq!(consumed, bytes.len());
343 assert_eq!(d.get(0), Some("AAPL"));
344 assert_eq!(d.get(1), Some("MSFT"));
345
346 let bytes2 = build_delta(2, &["GOOG"]);
347 d.apply_delta_from_bytes(&bytes2).unwrap();
348 assert_eq!(d.get(2), Some("GOOG"));
349 }
350
351 #[test]
352 fn from_bytes_partial_failure_rolls_back() {
353 let mut d = SymbolDict::new();
358 d.apply_delta(0, [b"first".as_slice()]).unwrap();
359 let snapshot_len = d.len();
360 let snapshot_heap = d.heap_bytes();
361
362 let mut bytes = Vec::new();
363 encode_u64(snapshot_len as u64, &mut bytes); encode_u64(2, &mut bytes); encode_u64(2, &mut bytes); bytes.extend_from_slice(b"ok");
367 encode_u64(10, &mut bytes); bytes.extend_from_slice(b"abc"); let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
371 assert_eq!(err.code(), ErrorCode::ProtocolError);
372 assert_eq!(d.len(), snapshot_len);
375 assert_eq!(d.heap_bytes(), snapshot_heap);
376 let next = build_delta(snapshot_len as u64, &["recovered"]);
377 d.apply_delta_from_bytes(&next).unwrap();
378 assert_eq!(d.get(snapshot_len as u32), Some("recovered"));
379 }
380
381 #[test]
382 fn from_bytes_truncated_entry_rejected() {
383 let mut d = SymbolDict::new();
384 let mut bytes = build_delta(0, &["hello"]);
385 bytes.truncate(bytes.len() - 1); let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
387 assert_eq!(err.code(), ErrorCode::ProtocolError);
388 }
389
390 #[test]
391 fn from_bytes_invalid_utf8_rejected() {
392 let mut bytes = Vec::new();
393 encode_u64(0, &mut bytes);
394 encode_u64(1, &mut bytes);
395 encode_u64(2, &mut bytes);
396 bytes.extend_from_slice(&[0xFF, 0xFE]); let mut d = SymbolDict::new();
398 let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
399 assert_eq!(err.code(), ErrorCode::InvalidUtf8);
400 }
401
402 #[test]
403 fn reset_clears_state() {
404 let mut d = SymbolDict::new();
405 d.apply_delta(0, [b"x".as_slice(), b"yy".as_slice()])
406 .unwrap();
407 assert_eq!(d.len(), 2);
408 d.reset();
409 assert_eq!(d.len(), 0);
410 assert_eq!(d.heap_bytes(), 0);
411 d.apply_delta(0, [b"new".as_slice()]).unwrap();
413 assert_eq!(d.get(0), Some("new"));
414 }
415
416 #[test]
417 fn delta_with_zero_entries_is_noop() {
418 let mut d = SymbolDict::new();
419 d.apply_delta(0, std::iter::empty::<&[u8]>()).unwrap();
420 let bytes = build_delta(0, &[]);
421 let consumed = d.apply_delta_from_bytes(&bytes).unwrap();
422 assert_eq!(consumed, bytes.len());
423 assert_eq!(d.len(), 0);
424 }
425
426 #[test]
427 fn delta_count_exceeding_capacity_rejected_upfront() {
428 let mut d = SymbolDict::new();
432 let mut bytes = Vec::new();
433 encode_u64(0, &mut bytes); encode_u64(u64::MAX, &mut bytes); let err = d.apply_delta_from_bytes(&bytes).unwrap_err();
443 assert_eq!(err.code(), ErrorCode::ProtocolError);
444 assert!(
445 err.msg().contains("exceeds remaining capacity"),
446 "expected upfront-cap rejection, got: {}",
447 err.msg()
448 );
449 assert_eq!(d.len(), 0);
450 }
451
452 #[test]
453 fn unicode_entries_preserved() {
454 let mut d = SymbolDict::new();
455 let bytes = build_delta(0, &["café", "日本語"]);
456 d.apply_delta_from_bytes(&bytes).unwrap();
457 assert_eq!(d.get(0), Some("café"));
458 assert_eq!(d.get(1), Some("日本語"));
459 }
460}