Skip to main content

corium_protocol/
codec.rs

1//! Composite wire encoding for values, EDN forms, datoms, and schema.
2//!
3//! Single values reuse the sortable tag space from `corium-core`; composite
4//! payloads extend it with container tags and a per-message interning table
5//! for keywords and repeated strings (see `docs/design/protocol.md`). The
6//! composite variant is length-prefixed rather than escaped, and keywords
7//! travel by name so no shared interner state is required across processes.
8
9use std::collections::HashMap;
10use std::sync::Arc;
11
12use corium_core::{
13    Attribute, Cardinality, Datom, EntityId, Keyword, KeywordInterner, Schema, TotalF64, Unique,
14    Value, ValueType,
15};
16use corium_db::Idents;
17use corium_query::edn::Edn;
18use thiserror::Error;
19
20// Scalar tags shared with `corium_core::encoding`.
21const BOOL: u8 = 0x10;
22const LONG: u8 = 0x20;
23const DOUBLE: u8 = 0x30;
24const INSTANT: u8 = 0x40;
25const UUID: u8 = 0x50;
26const REF: u8 = 0x90;
27// Composite-variant tags.
28const NIL: u8 = 0x00;
29const KEYWORD_NAME: u8 = 0x61;
30const STR_INTERNED: u8 = 0x71;
31const BYTES_PREFIXED: u8 = 0x81;
32const LIST: u8 = 0xA0;
33const VECTOR: u8 = 0xA1;
34const MAP: u8 = 0xA2;
35const SET: u8 = 0xA3;
36const TAGGED: u8 = 0xA4;
37const SYMBOL: u8 = 0xA5;
38
39/// Codec failure.
40#[derive(Debug, Error, Eq, PartialEq)]
41pub enum CodecError {
42    /// Input ended before a complete item was read.
43    #[error("truncated wire payload")]
44    Truncated,
45    /// Unknown tag byte.
46    #[error("unknown wire tag {0:#x}")]
47    UnknownTag(u8),
48    /// Interning table reference out of range.
49    #[error("invalid intern reference {0}")]
50    InvalidIntern(u64),
51    /// String payload is not UTF-8.
52    #[error("invalid UTF-8 string")]
53    InvalidUtf8,
54    /// A keyword id has no entry in the supplied interner.
55    #[error("keyword id {0} is not interned")]
56    UnknownKeyword(u64),
57    /// A count or length does not fit the platform.
58    #[error("wire length out of range")]
59    Length,
60    /// Payload decoded but trailing bytes remain.
61    #[error("trailing bytes after wire payload")]
62    Trailing,
63    /// Field value outside its legal range.
64    #[error("invalid wire field: {0}")]
65    InvalidField(&'static str),
66}
67
68/// Streaming writer with a per-message string interning table.
69#[derive(Default)]
70pub struct Writer {
71    buf: Vec<u8>,
72    table: HashMap<String, u64>,
73}
74
75impl Writer {
76    /// Creates an empty writer.
77    #[must_use]
78    pub fn new() -> Self {
79        Self::default()
80    }
81
82    /// Consumes the writer, returning the message bytes.
83    #[must_use]
84    pub fn finish(self) -> Vec<u8> {
85        self.buf
86    }
87
88    fn varint(&mut self, mut n: u64) {
89        loop {
90            let byte = (n & 0x7f) as u8;
91            n >>= 7;
92            if n == 0 {
93                self.buf.push(byte);
94                return;
95            }
96            self.buf.push(byte | 0x80);
97        }
98    }
99
100    /// Writes an interned string: `0 len bytes` defines the next table
101    /// index on first use; later uses write `index` (1-based).
102    fn intern(&mut self, text: &str) {
103        if let Some(&index) = self.table.get(text) {
104            self.varint(index);
105            return;
106        }
107        let index = self.table.len() as u64 + 1;
108        self.table.insert(text.to_owned(), index);
109        self.varint(0);
110        self.varint(text.len() as u64);
111        self.buf.extend_from_slice(text.as_bytes());
112    }
113
114    fn keyword(&mut self, keyword: &Keyword) {
115        self.buf.push(KEYWORD_NAME);
116        match &keyword.namespace {
117            Some(namespace) => self.intern(&format!("{namespace}/{}", keyword.name)),
118            None => self.intern(&keyword.name),
119        }
120    }
121
122    /// Writes a `u64` as a varint (for counts and ids).
123    pub fn u64(&mut self, n: u64) {
124        self.varint(n);
125    }
126
127    /// Writes an `i64` (zigzag varint).
128    pub fn i64(&mut self, n: i64) {
129        self.varint(zigzag(n));
130    }
131
132    /// Writes a raw byte.
133    pub fn byte(&mut self, b: u8) {
134        self.buf.push(b);
135    }
136
137    /// Writes one EDN form.
138    pub fn edn(&mut self, form: &Edn) {
139        match form {
140            Edn::Nil => self.buf.push(NIL),
141            Edn::Bool(v) => {
142                self.buf.push(BOOL);
143                self.buf.push(u8::from(*v));
144            }
145            Edn::Long(v) => {
146                self.buf.push(LONG);
147                self.i64(*v);
148            }
149            Edn::Double(v) => {
150                self.buf.push(DOUBLE);
151                self.buf.extend_from_slice(&v.sortable_bits().to_be_bytes());
152            }
153            Edn::Str(v) => {
154                self.buf.push(STR_INTERNED);
155                self.intern(v);
156            }
157            Edn::Keyword(k) => self.keyword(k),
158            Edn::Symbol(s) => {
159                self.buf.push(SYMBOL);
160                self.intern(s);
161            }
162            Edn::List(items) => self.seq(LIST, items),
163            Edn::Vector(items) => self.seq(VECTOR, items),
164            Edn::Set(items) => self.seq(SET, items),
165            Edn::Map(pairs) => {
166                self.buf.push(MAP);
167                self.varint(pairs.len() as u64);
168                for (key, value) in pairs {
169                    self.edn(key);
170                    self.edn(value);
171                }
172            }
173            Edn::Tagged(tag, value) => {
174                self.buf.push(TAGGED);
175                self.intern(tag);
176                self.edn(value);
177            }
178        }
179    }
180
181    fn seq(&mut self, tag: u8, items: &[Edn]) {
182        self.buf.push(tag);
183        self.varint(items.len() as u64);
184        for item in items {
185            self.edn(item);
186        }
187    }
188
189    /// Writes one engine value. Keywords travel by name via `interner`.
190    ///
191    /// # Errors
192    /// Returns [`CodecError::UnknownKeyword`] for an unresolvable keyword id.
193    pub fn value(&mut self, value: &Value, interner: &KeywordInterner) -> Result<(), CodecError> {
194        match value {
195            Value::Bool(v) => {
196                self.buf.push(BOOL);
197                self.buf.push(u8::from(*v));
198            }
199            Value::Long(v) => {
200                self.buf.push(LONG);
201                self.i64(*v);
202            }
203            Value::Double(v) => {
204                self.buf.push(DOUBLE);
205                self.buf.extend_from_slice(&v.sortable_bits().to_be_bytes());
206            }
207            Value::Instant(v) => {
208                self.buf.push(INSTANT);
209                self.i64(*v);
210            }
211            Value::Uuid(v) => {
212                self.buf.push(UUID);
213                self.buf.extend_from_slice(&v.to_be_bytes());
214            }
215            Value::Keyword(id) => {
216                let keyword = interner
217                    .resolve(*id)
218                    .ok_or(CodecError::UnknownKeyword(*id))?
219                    .clone();
220                self.keyword(&keyword);
221            }
222            Value::Str(v) => {
223                self.buf.push(STR_INTERNED);
224                self.intern(v);
225            }
226            Value::Bytes(v) => {
227                self.buf.push(BYTES_PREFIXED);
228                self.varint(v.len() as u64);
229                self.buf.extend_from_slice(v);
230            }
231            Value::Ref(e) => {
232                self.buf.push(REF);
233                self.varint(e.raw());
234            }
235        }
236        Ok(())
237    }
238}
239
240/// Streaming reader over one wire message.
241pub struct Reader<'a> {
242    input: &'a [u8],
243    table: Vec<String>,
244}
245
246impl<'a> Reader<'a> {
247    /// Creates a reader over message bytes.
248    #[must_use]
249    pub fn new(input: &'a [u8]) -> Self {
250        Self {
251            input,
252            table: Vec::new(),
253        }
254    }
255
256    /// Fails unless every input byte was consumed.
257    ///
258    /// # Errors
259    /// Returns [`CodecError::Trailing`] when bytes remain.
260    pub fn expect_end(&self) -> Result<(), CodecError> {
261        if self.input.is_empty() {
262            Ok(())
263        } else {
264            Err(CodecError::Trailing)
265        }
266    }
267
268    fn take(&mut self, n: usize) -> Result<&'a [u8], CodecError> {
269        let bytes = self.input.get(..n).ok_or(CodecError::Truncated)?;
270        self.input = &self.input[n..];
271        Ok(bytes)
272    }
273
274    fn tag(&mut self) -> Result<u8, CodecError> {
275        Ok(self.take(1)?[0])
276    }
277
278    /// Reads a varint `u64`.
279    ///
280    /// # Errors
281    /// Returns an error for truncated input.
282    pub fn u64(&mut self) -> Result<u64, CodecError> {
283        let mut out = 0_u64;
284        let mut shift = 0_u32;
285        loop {
286            let byte = self.take(1)?[0];
287            out |= u64::from(byte & 0x7f)
288                .checked_shl(shift)
289                .ok_or(CodecError::Length)?;
290            if byte & 0x80 == 0 {
291                return Ok(out);
292            }
293            shift += 7;
294            if shift > 63 {
295                return Err(CodecError::Length);
296            }
297        }
298    }
299
300    /// Reads a zigzag varint `i64`.
301    ///
302    /// # Errors
303    /// Returns an error for truncated input.
304    pub fn i64(&mut self) -> Result<i64, CodecError> {
305        Ok(unzigzag(self.u64()?))
306    }
307
308    /// Reads a raw byte.
309    ///
310    /// # Errors
311    /// Returns an error for truncated input.
312    pub fn byte(&mut self) -> Result<u8, CodecError> {
313        self.tag()
314    }
315
316    fn count(&mut self) -> Result<usize, CodecError> {
317        usize::try_from(self.u64()?).map_err(|_| CodecError::Length)
318    }
319
320    fn intern(&mut self) -> Result<String, CodecError> {
321        let index = self.u64()?;
322        if index == 0 {
323            let len = self.count()?;
324            let text = std::str::from_utf8(self.take(len)?)
325                .map_err(|_| CodecError::InvalidUtf8)?
326                .to_owned();
327            self.table.push(text.clone());
328            return Ok(text);
329        }
330        let position = usize::try_from(index - 1).map_err(|_| CodecError::Length)?;
331        self.table
332            .get(position)
333            .cloned()
334            .ok_or(CodecError::InvalidIntern(index))
335    }
336
337    fn double(&mut self) -> Result<TotalF64, CodecError> {
338        let sortable = u64::from_be_bytes(
339            self.take(8)?
340                .try_into()
341                .map_err(|_| CodecError::Truncated)?,
342        );
343        let bits = if sortable & (1_u64 << 63) == 0 {
344            !sortable
345        } else {
346            sortable ^ (1_u64 << 63)
347        };
348        Ok(TotalF64(f64::from_bits(bits)))
349    }
350
351    /// Reads one EDN form.
352    ///
353    /// # Errors
354    /// Returns [`CodecError`] for malformed input.
355    pub fn edn(&mut self) -> Result<Edn, CodecError> {
356        Ok(match self.tag()? {
357            NIL => Edn::Nil,
358            BOOL => Edn::Bool(self.take(1)?[0] != 0),
359            LONG => Edn::Long(self.i64()?),
360            DOUBLE => Edn::Double(self.double()?),
361            STR_INTERNED => Edn::Str(self.intern()?),
362            KEYWORD_NAME => Edn::Keyword(Keyword::parse(&self.intern()?)),
363            SYMBOL => Edn::Symbol(self.intern()?),
364            LIST => Edn::List(self.items()?),
365            VECTOR => Edn::Vector(self.items()?),
366            SET => {
367                let mut items = self.items()?;
368                items.sort();
369                items.dedup();
370                Edn::Set(items)
371            }
372            MAP => {
373                let count = self.count()?;
374                let mut pairs = Vec::with_capacity(count.min(4096));
375                for _ in 0..count {
376                    let key = self.edn()?;
377                    let value = self.edn()?;
378                    pairs.push((key, value));
379                }
380                pairs.sort_by(|left, right| left.0.cmp(&right.0));
381                Edn::Map(pairs)
382            }
383            TAGGED => {
384                let tag = self.intern()?;
385                Edn::Tagged(tag, Box::new(self.edn()?))
386            }
387            other => return Err(CodecError::UnknownTag(other)),
388        })
389    }
390
391    fn items(&mut self) -> Result<Vec<Edn>, CodecError> {
392        let count = self.count()?;
393        let mut items = Vec::with_capacity(count.min(4096));
394        for _ in 0..count {
395            items.push(self.edn()?);
396        }
397        Ok(items)
398    }
399
400    /// Reads one engine value, interning keyword names into `interner`.
401    ///
402    /// # Errors
403    /// Returns [`CodecError`] for malformed input.
404    pub fn value(&mut self, interner: &mut KeywordInterner) -> Result<Value, CodecError> {
405        Ok(match self.tag()? {
406            BOOL => Value::Bool(self.take(1)?[0] != 0),
407            LONG => Value::Long(self.i64()?),
408            DOUBLE => Value::Double(self.double()?),
409            INSTANT => Value::Instant(self.i64()?),
410            UUID => Value::Uuid(u128::from_be_bytes(
411                self.take(16)?
412                    .try_into()
413                    .map_err(|_| CodecError::Truncated)?,
414            )),
415            KEYWORD_NAME => {
416                let keyword = Keyword::parse(&self.intern()?);
417                Value::Keyword(interner.intern(keyword))
418            }
419            STR_INTERNED => Value::Str(Arc::from(self.intern()?.as_str())),
420            BYTES_PREFIXED => {
421                let len = self.count()?;
422                Value::Bytes(Arc::from(self.take(len)?))
423            }
424            REF => Value::Ref(EntityId::from_raw(self.u64()?)),
425            other => return Err(CodecError::UnknownTag(other)),
426        })
427    }
428}
429
430#[allow(clippy::cast_sign_loss)]
431const fn zigzag(n: i64) -> u64 {
432    ((n << 1) ^ (n >> 63)) as u64
433}
434
435#[allow(clippy::cast_possible_wrap)]
436const fn unzigzag(n: u64) -> i64 {
437    ((n >> 1) as i64) ^ -((n & 1) as i64)
438}
439
440/// Encodes one EDN form as a standalone message.
441#[must_use]
442pub fn encode_edn(form: &Edn) -> Vec<u8> {
443    let mut writer = Writer::new();
444    writer.edn(form);
445    writer.finish()
446}
447
448/// Decodes one EDN form from a standalone message.
449///
450/// # Errors
451/// Returns [`CodecError`] for malformed or trailing input.
452pub fn decode_edn(bytes: &[u8]) -> Result<Edn, CodecError> {
453    let mut reader = Reader::new(bytes);
454    let form = reader.edn()?;
455    reader.expect_end()?;
456    Ok(form)
457}
458
459/// Encodes a datom list; keyword values travel by name via `interner`.
460///
461/// # Errors
462/// Returns [`CodecError::UnknownKeyword`] for unresolvable keyword ids.
463pub fn encode_datoms(datoms: &[Datom], interner: &KeywordInterner) -> Result<Vec<u8>, CodecError> {
464    let mut writer = Writer::new();
465    writer.u64(datoms.len() as u64);
466    for datom in datoms {
467        writer.u64(datom.e.raw());
468        writer.u64(datom.a.raw());
469        writer.u64(datom.tx.raw());
470        writer.byte(u8::from(datom.added));
471        writer.value(&datom.v, interner)?;
472    }
473    Ok(writer.finish())
474}
475
476/// Decodes a datom list, interning keyword names into `interner`.
477///
478/// # Errors
479/// Returns [`CodecError`] for malformed input.
480pub fn decode_datoms(
481    bytes: &[u8],
482    interner: &mut KeywordInterner,
483) -> Result<Vec<Datom>, CodecError> {
484    let mut reader = Reader::new(bytes);
485    let count = usize::try_from(reader.u64()?).map_err(|_| CodecError::Length)?;
486    let mut datoms = Vec::with_capacity(count.min(65_536));
487    for _ in 0..count {
488        let e = EntityId::from_raw(reader.u64()?);
489        let a = EntityId::from_raw(reader.u64()?);
490        let tx = EntityId::from_raw(reader.u64()?);
491        let added = reader.byte()? != 0;
492        let v = reader.value(interner)?;
493        datoms.push(Datom { e, a, v, tx, added });
494    }
495    reader.expect_end()?;
496    Ok(datoms)
497}
498
499/// Encodes schema attributes plus the ident registry (handshake payload).
500#[must_use]
501pub fn encode_schema(schema: &Schema, idents: &Idents) -> Vec<u8> {
502    let mut writer = Writer::new();
503    let attrs: Vec<_> = schema.iter().collect();
504    writer.u64(attrs.len() as u64);
505    for (_, attr) in attrs {
506        writer.u64(attr.id.raw());
507        writer.byte(value_type_tag(attr.value_type));
508        writer.byte(match attr.cardinality {
509            Cardinality::One => 0,
510            Cardinality::Many => 1,
511        });
512        writer.byte(match attr.unique {
513            None => 0,
514            Some(Unique::Identity) => 1,
515            Some(Unique::Value) => 2,
516        });
517        writer.byte(
518            u8::from(attr.is_component)
519                | (u8::from(attr.indexed) << 1)
520                | (u8::from(attr.no_history) << 2),
521        );
522    }
523    let idents: Vec<_> = idents.iter().collect();
524    writer.u64(idents.len() as u64);
525    for (keyword, id) in idents {
526        writer.edn(&Edn::Keyword(keyword.clone()));
527        writer.u64(id.raw());
528    }
529    writer.finish()
530}
531
532/// Decodes a schema/ident handshake payload.
533///
534/// # Errors
535/// Returns [`CodecError`] for malformed input.
536pub fn decode_schema(bytes: &[u8]) -> Result<(Schema, Idents), CodecError> {
537    let mut reader = Reader::new(bytes);
538    let mut schema = Schema::default();
539    let attr_count = usize::try_from(reader.u64()?).map_err(|_| CodecError::Length)?;
540    for _ in 0..attr_count {
541        let id = EntityId::from_raw(reader.u64()?);
542        let value_type = value_type_from(reader.byte()?)?;
543        let cardinality = match reader.byte()? {
544            0 => Cardinality::One,
545            1 => Cardinality::Many,
546            _ => return Err(CodecError::InvalidField("cardinality")),
547        };
548        let unique = match reader.byte()? {
549            0 => None,
550            1 => Some(Unique::Identity),
551            2 => Some(Unique::Value),
552            _ => return Err(CodecError::InvalidField("unique")),
553        };
554        let flags = reader.byte()?;
555        schema.insert(Attribute {
556            id,
557            value_type,
558            cardinality,
559            unique,
560            is_component: flags & 1 != 0,
561            indexed: flags & 2 != 0,
562            no_history: flags & 4 != 0,
563        });
564    }
565    let mut idents = Idents::default();
566    let ident_count = usize::try_from(reader.u64()?).map_err(|_| CodecError::Length)?;
567    for _ in 0..ident_count {
568        let Edn::Keyword(keyword) = reader.edn()? else {
569            return Err(CodecError::InvalidField("ident keyword"));
570        };
571        let id = EntityId::from_raw(reader.u64()?);
572        idents.insert(keyword, id);
573    }
574    reader.expect_end()?;
575    Ok((schema, idents))
576}
577
578/// Encodes an interner snapshot (keywords in dense id order).
579#[must_use]
580pub fn encode_naming(interner: &KeywordInterner) -> Vec<u8> {
581    let mut writer = Writer::new();
582    let entries: Vec<_> = interner.iter().collect();
583    writer.u64(entries.len() as u64);
584    for (_, keyword) in entries {
585        writer.edn(&Edn::Keyword(keyword.clone()));
586    }
587    writer.finish()
588}
589
590/// Decodes an interner snapshot; re-interning in order reproduces dense ids.
591///
592/// # Errors
593/// Returns [`CodecError`] for malformed input.
594pub fn decode_naming(bytes: &[u8]) -> Result<KeywordInterner, CodecError> {
595    let mut reader = Reader::new(bytes);
596    let mut interner = KeywordInterner::default();
597    let count = usize::try_from(reader.u64()?).map_err(|_| CodecError::Length)?;
598    for _ in 0..count {
599        let Edn::Keyword(keyword) = reader.edn()? else {
600            return Err(CodecError::InvalidField("interner keyword"));
601        };
602        interner.intern(keyword);
603    }
604    reader.expect_end()?;
605    Ok(interner)
606}
607
608const fn value_type_tag(value_type: ValueType) -> u8 {
609    match value_type {
610        ValueType::Bool => 0,
611        ValueType::Long => 1,
612        ValueType::Double => 2,
613        ValueType::Instant => 3,
614        ValueType::Uuid => 4,
615        ValueType::Keyword => 5,
616        ValueType::Str => 6,
617        ValueType::Bytes => 7,
618        ValueType::Ref => 8,
619    }
620}
621
622const fn value_type_from(tag: u8) -> Result<ValueType, CodecError> {
623    Ok(match tag {
624        0 => ValueType::Bool,
625        1 => ValueType::Long,
626        2 => ValueType::Double,
627        3 => ValueType::Instant,
628        4 => ValueType::Uuid,
629        5 => ValueType::Keyword,
630        6 => ValueType::Str,
631        7 => ValueType::Bytes,
632        8 => ValueType::Ref,
633        _ => return Err(CodecError::InvalidField("value type")),
634    })
635}