use revision::Error as RevisionError;
use revision::optimised::tag::{SizeClass, read_tag};
#[cfg(test)]
mod variant_ids {
pub const NONE: u8 = 0;
pub const NULL: u8 = 1;
pub const BOOL: u8 = 2;
pub const NUMBER: u8 = 3;
pub const STRING: u8 = 4;
pub const REGEX: u8 = 16;
}
const VALUE_FIXED_SIZES: [u8; 32] = {
let mut t = [0u8; 32];
t[2] = 1; t
};
#[inline]
pub(crate) fn skip_value_wire(reader: &mut &[u8]) -> Result<(), RevisionError> {
let rev = <u16 as revision::DeserializeRevisioned>::deserialize_revisioned(reader)?;
if rev != 2 {
return Err(RevisionError::Deserialize(format!(
"skip_value_wire: expected Value revision 2, got {rev}"
)));
}
let tag = read_tag(reader)?;
let sc = tag.size_class()?;
match sc {
SizeClass::Inline => Ok(()),
SizeClass::Fixed => {
let n = VALUE_FIXED_SIZES[tag.variant_id() as usize] as usize;
if reader.len() < n {
return Err(RevisionError::Deserialize(
"skip_value_wire: fixed payload exceeds remaining bytes".into(),
));
}
*reader = &reader[n..];
Ok(())
}
SizeClass::Varlen => {
if reader.len() < 4 {
return Err(RevisionError::Deserialize(
"skip_value_wire: truncated varlen length prefix".into(),
));
}
let mut len_buf = [0u8; 4];
len_buf.copy_from_slice(&reader[..4]);
*reader = &reader[4..];
let len = u32::from_le_bytes(len_buf) as usize;
if reader.len() < len {
return Err(RevisionError::Deserialize(
"skip_value_wire: varlen payload exceeds remaining bytes".into(),
));
}
*reader = &reader[len..];
Ok(())
}
}
}
#[inline]
pub(crate) fn rev2_optimised_payload_unchecked(wire: &[u8]) -> Result<&[u8], RevisionError> {
if wire.is_empty() {
return Err(RevisionError::Deserialize(
"rev2_optimised_payload_unchecked: empty wire".into(),
));
}
debug_assert_eq!(
wire[0], 2u8,
"rev2_optimised_payload_unchecked: caller-asserted rev 2 invariant violated",
);
let r = &wire[1..];
if r.len() < 4 {
return Err(RevisionError::Deserialize(
"rev2_optimised_payload_unchecked: truncated u32_le envelope length".into(),
));
}
let payload_len = u32::from_le_bytes([r[0], r[1], r[2], r[3]]) as usize;
let r = &r[4..];
if r.len() < payload_len {
return Err(RevisionError::Deserialize(
"rev2_optimised_payload_unchecked: payload exceeds remaining bytes".into(),
));
}
Ok(&r[..payload_len])
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use revision::SerializeRevisioned;
use surrealdb_strand::Strand;
use super::*;
use crate::val::{Array, Number, Object, Value};
fn encoded(v: &Value) -> Vec<u8> {
let mut out = Vec::new();
v.serialize_revisioned(&mut out).expect("serialize");
out
}
#[test]
fn skips_inline_none() {
let bytes = encoded(&Value::None);
assert_eq!(bytes.len(), 2);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty(), "skip_value_wire must consume entire wire");
}
#[test]
fn skips_inline_null() {
let bytes = encoded(&Value::Null);
assert_eq!(bytes.len(), 2);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn skips_fixed_bool() {
for v in [Value::Bool(true), Value::Bool(false)] {
let bytes = encoded(&v);
assert_eq!(bytes.len(), 3);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
}
#[test]
fn skips_varlen_string() {
let bytes = encoded(&Value::String(Strand::from("hello world")));
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn skips_varlen_number() {
let bytes = encoded(&Value::Number(Number::Int(42)));
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn skips_varlen_array() {
let arr = Value::Array(Array::from(vec![
Value::Number(Number::Int(1)),
Value::String("x".into()),
]));
let bytes = encoded(&arr);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn skips_varlen_object() {
let obj = Value::Object(Object::from(BTreeMap::from([
(Strand::from("a"), Value::Number(Number::Int(1))),
(Strand::from("b"), Value::String("y".into())),
])));
let bytes = encoded(&obj);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn skips_through_concatenated_values() {
let mut buf = Vec::new();
buf.extend(encoded(&Value::None));
buf.extend(encoded(&Value::Bool(true)));
buf.extend(encoded(&Value::Number(Number::Int(7))));
buf.extend(encoded(&Value::String("payload".into())));
let mut r: &[u8] = &buf;
skip_value_wire(&mut r).unwrap();
skip_value_wire(&mut r).unwrap();
skip_value_wire(&mut r).unwrap();
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn errors_on_wrong_revision() {
let bytes = [0x01u8, 0x00];
let mut r: &[u8] = &bytes;
let err = skip_value_wire(&mut r).expect_err("must bail on rev != 2");
match err {
RevisionError::Deserialize(s) => assert!(s.contains("revision 2")),
other => panic!("expected Deserialize error, got {other:?}"),
}
}
#[test]
fn errors_on_truncated_varlen() {
let bytes = [0x02u8, (4 | (0b10 << 5))];
let mut r: &[u8] = &bytes;
assert!(skip_value_wire(&mut r).is_err());
}
#[test]
fn errors_on_empty_buffer() {
let mut r: &[u8] = &[];
assert!(skip_value_wire(&mut r).is_err());
}
#[test]
fn errors_on_truncated_fixed_payload() {
let bytes = [0x02u8, (2u8 | (0b01u8 << 5))];
let mut r: &[u8] = &bytes;
let err = skip_value_wire(&mut r).expect_err("truncated Fixed payload must error");
match err {
RevisionError::Deserialize(s) => {
assert!(s.contains("fixed payload"), "unexpected error message: {s}")
}
other => panic!("expected Deserialize error, got {other:?}"),
}
}
#[test]
fn errors_on_reserved_size_class() {
let bytes = [0x02u8, 0b0110_0000];
let mut r: &[u8] = &bytes;
assert!(skip_value_wire(&mut r).is_err());
}
#[test]
fn skips_all_varlen_variants() {
use std::str::FromStr;
use std::time::Duration as StdDuration;
use chrono::{TimeZone, Utc};
use crate::val::{
Bytes, Datetime, Duration as SurrealDuration, File, RecordId, Regex, Set, TableName,
Uuid as SurrealUuid,
};
let cases: Vec<(Value, u8)> = vec![
(Value::Duration(SurrealDuration(StdDuration::from_secs(1))), 5),
(Value::Datetime(Datetime(Utc.with_ymd_and_hms(2024, 5, 21, 0, 0, 0).unwrap())), 6),
(Value::Uuid(SurrealUuid(uuid::Uuid::from_u128(0xa1))), 7),
(Value::Set(Set::new()), 9),
(Value::Bytes(Bytes::from(vec![1u8, 2, 3])), 12),
(Value::Table(TableName::from("foo")), 13),
(
Value::RecordId(RecordId {
table: TableName::from("foo"),
key: crate::val::RecordIdKey::from(Strand::from("bar")),
}),
14,
),
(Value::File(File::new("bucket".to_string(), "key".to_string())), 15),
(Value::Regex(Regex::from_str("/x/").unwrap()), 16),
];
for (v, expected_id) in cases {
let bytes = encoded(&v);
let tag = revision::optimised::tag::Tag(bytes[1]);
assert_eq!(
tag.variant_id(),
expected_id,
"variant_id drift for {v:?} (got {}, expected {expected_id})",
tag.variant_id(),
);
assert_eq!(
tag.size_class().unwrap(),
revision::optimised::tag::SizeClass::Varlen,
"size_class drift for {v:?}",
);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).expect("skip_value_wire must succeed");
assert!(r.is_empty(), "skip_value_wire must consume entire wire for {v:?}");
}
}
#[test]
fn skips_geometry_variant() {
use crate::val::Geometry;
let v = Value::Geometry(Geometry::Point(geo::Point::new(0.0, 0.0)));
let bytes = encoded(&v);
let tag = revision::optimised::tag::Tag(bytes[1]);
assert_eq!(tag.variant_id(), 11, "Geometry variant_id drift");
assert_eq!(tag.size_class().unwrap(), revision::optimised::tag::SizeClass::Varlen);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn skips_range_variant() {
use std::ops::Bound;
use crate::val::Range;
let v = Value::Range(Box::new(Range {
start: Bound::Included(Value::Number(Number::Int(1))),
end: Bound::Excluded(Value::Number(Number::Int(10))),
}));
let bytes = encoded(&v);
let tag = revision::optimised::tag::Tag(bytes[1]);
assert_eq!(tag.variant_id(), 17, "Range variant_id drift");
assert_eq!(tag.size_class().unwrap(), revision::optimised::tag::SizeClass::Varlen);
let mut r: &[u8] = &bytes;
skip_value_wire(&mut r).unwrap();
assert!(r.is_empty());
}
#[test]
fn variant_layout_matches_value_declaration() {
use std::str::FromStr;
use crate::val::Regex;
fn header(v: &Value) -> (u8, revision::optimised::tag::SizeClass) {
let bytes = encoded(v);
let tag = revision::optimised::tag::Tag(bytes[1]);
(tag.variant_id(), tag.size_class().unwrap())
}
assert_eq!(header(&Value::None).0, variant_ids::NONE);
assert_eq!(header(&Value::Null).0, variant_ids::NULL);
assert_eq!(header(&Value::Bool(false)).0, variant_ids::BOOL);
assert_eq!(header(&Value::Number(Number::Int(0))).0, variant_ids::NUMBER);
assert_eq!(header(&Value::String(Strand::from("x"))).0, variant_ids::STRING);
assert_eq!(header(&Value::Regex(Regex::from_str("/x/").unwrap())).0, variant_ids::REGEX);
let fixed_samples: &[(Value, u8)] = &[(Value::Bool(false), 1), (Value::Bool(true), 1)];
for (v, expected_size) in fixed_samples {
let (vid, sc) = header(v);
assert_eq!(sc, revision::optimised::tag::SizeClass::Fixed);
assert_eq!(VALUE_FIXED_SIZES[vid as usize], *expected_size);
}
}
}