use std::io;
use std::io::{Cursor, Write};
use super::bytes32::Bytes32;
use super::object_cache::{serialized_length, treehash, ObjectCache};
use super::read_cache_lookup::ReadCacheLookup;
use super::write_atom::write_atom;
use crate::allocator::{Allocator, NodePtr, SExp};
const BACK_REFERENCE: u8 = 0xfe;
const CONS_BOX_MARKER: u8 = 0xff;
#[derive(PartialEq, Eq, Clone)]
enum ReadOp {
Parse,
Cons,
}
pub struct Serializer {
read_op_stack: Vec<ReadOp>,
write_stack: Vec<NodePtr>,
read_cache_lookup: ReadCacheLookup,
thc: ObjectCache<Bytes32>,
slc: ObjectCache<u64>,
output: Cursor<Vec<u8>>,
}
impl Default for Serializer {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone)]
pub struct UndoState {
read_op_stack: Vec<ReadOp>,
write_stack: Vec<NodePtr>,
output_position: u64,
}
impl Serializer {
pub fn new() -> Self {
Self {
read_op_stack: vec![ReadOp::Parse],
write_stack: vec![],
read_cache_lookup: ReadCacheLookup::new(),
thc: ObjectCache::new(treehash),
slc: ObjectCache::new(serialized_length),
output: Cursor::new(vec![]),
}
}
fn serialize_pair(&mut self, left: NodePtr, right: NodePtr) -> io::Result<()> {
self.output.write_all(&[CONS_BOX_MARKER])?;
self.write_stack.push(right);
self.write_stack.push(left);
self.read_op_stack.push(ReadOp::Cons);
self.read_op_stack.push(ReadOp::Parse);
self.read_op_stack.push(ReadOp::Parse);
Ok(())
}
pub fn add(
&mut self,
a: &Allocator,
node: NodePtr,
sentinel: Option<NodePtr>,
) -> io::Result<(bool, UndoState)> {
assert!(!self.read_op_stack.is_empty());
let undo_state = UndoState {
read_op_stack: self.read_op_stack.clone(),
write_stack: self.write_stack.clone(),
output_position: self.output.position(),
};
self.write_stack.push(node);
while let Some(node_to_write) = self.write_stack.pop() {
if Some(node_to_write) == sentinel {
return Ok((false, undo_state));
}
let op = self.read_op_stack.pop();
assert!(op == Some(ReadOp::Parse));
let node_serialized_length = self.slc.get_or_calculate(a, &node_to_write, sentinel);
let node_tree_hash = self.thc.get_or_calculate(a, &node_to_write, sentinel);
if let (Some(node_tree_hash), Some(node_serialized_length)) =
(node_tree_hash, node_serialized_length)
{
match self
.read_cache_lookup
.find_path(node_tree_hash, *node_serialized_length)
{
Some(path) => {
self.output.write_all(&[BACK_REFERENCE])?;
write_atom(&mut self.output, &path)?;
self.read_cache_lookup.push(*node_tree_hash);
}
None => match a.sexp(node_to_write) {
SExp::Pair(left, right) => {
self.serialize_pair(left, right)?;
}
SExp::Atom => {
let atom = a.atom(node_to_write);
write_atom(&mut self.output, atom.as_ref())?;
self.read_cache_lookup.push(*node_tree_hash);
}
},
}
} else {
match a.sexp(node_to_write) {
SExp::Pair(left, right) => {
self.serialize_pair(left, right)?;
}
SExp::Atom => {
let atom = a.atom(node_to_write);
write_atom(&mut self.output, atom.as_ref())?;
}
}
}
while !self.read_op_stack.is_empty()
&& self.read_op_stack[self.read_op_stack.len() - 1] == ReadOp::Cons
{
self.read_op_stack.pop();
self.read_cache_lookup.pop2_and_cons();
}
}
Ok((true, undo_state))
}
pub fn restore(&mut self, state: UndoState) {
self.read_op_stack = state.read_op_stack;
self.write_stack = state.write_stack;
self.output.set_position(state.output_position);
self.output
.get_mut()
.truncate(state.output_position as usize);
}
pub fn size(&self) -> u64 {
self.output.position()
}
pub fn get_ref(&self) -> &Vec<u8> {
self.output.get_ref()
}
pub fn into_inner(self) -> Vec<u8> {
assert!(self.read_op_stack.is_empty());
self.output.into_inner()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::serde::{
node_from_bytes, node_from_bytes_backrefs, node_to_bytes, node_to_bytes_backrefs,
};
use hex_literal::hex;
#[test]
fn test_simple_incremental() {
let mut a = Allocator::new();
let sentinel = a.new_pair(NodePtr::NIL, NodePtr::NIL).unwrap();
let item = node_from_bytes(&mut a, &hex!("ffff0102ff0304")).unwrap();
let list = a.new_pair(item, sentinel).unwrap();
let mut ser = Serializer::new();
let mut size = ser.size();
for _ in 0..10 {
let (done, _) = ser.add(&a, list, Some(sentinel)).unwrap();
assert!(!done);
assert!(ser.size() > size);
size = ser.size();
}
let (done, _) = ser.add(&a, NodePtr::NIL, None).unwrap();
assert!(done);
let output = ser.into_inner();
assert_eq!(
hex::encode(&output),
"ffffff0102ff0304fffe02fffe02fffe02fffe02fffe02fffe02fffe02fffe02fffe0280"
);
let parsed = node_from_bytes_backrefs(&mut a, &output).unwrap();
let round_trip = node_to_bytes_backrefs(&a, parsed).unwrap();
assert_eq!(
hex::encode(&round_trip),
"ffffff0102ff0304fffe02fffe02fffe02fffe02fe01"
);
let round_trip = node_to_bytes(&a, parsed).unwrap();
assert_eq!(hex::encode(&round_trip), "ffffff0102ff0304ffffff0102ff0304ffffff0102ff0304ffffff0102ff0304ffffff0102ff0304ffffff0102ff0304ffffff0102ff0304ffffff0102ff0304ffffff0102ff0304ffffff0102ff030480");
}
#[test]
fn test_incremental() {
let mut a = Allocator::new();
let sentinel = a.new_pair(NodePtr::NIL, NodePtr::NIL).unwrap();
let node1 = a.new_small_number(1).unwrap();
let node2 = a.new_pair(node1, sentinel).unwrap();
let node3 = a.new_small_number(3).unwrap();
let node4 = a.new_atom(b"foobar").unwrap();
let node5 = a.new_pair(node3, node4).unwrap();
let item = a.new_pair(node2, node5).unwrap();
let mut ser = Serializer::new();
let mut size = ser.size();
let (done, _) = ser.add(&a, item, Some(sentinel)).unwrap();
assert!(!done);
assert!(ser.size() > size);
size = ser.size();
let node1 = a.new_small_number(1).unwrap();
let node2 = a.new_pair(node1, sentinel).unwrap();
let node3 = a.new_small_number(3).unwrap();
let node4 = a.new_atom(b"barfoo").unwrap();
let node5 = a.new_pair(node3, node4).unwrap();
let item = a.new_pair(node2, node5).unwrap();
for _ in 0..10 {
let (done, _) = ser.add(&a, item, Some(sentinel)).unwrap();
assert!(!done);
assert!(ser.size() > size);
size = ser.size();
}
let (done, _) = ser.add(&a, NodePtr::NIL, None).unwrap();
assert!(done);
let output = ser.into_inner();
assert_eq!(hex::encode(&output), "ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff0180ff0386626172666f6ffe0efe0efe0efe0efe0efe0efe0efe0efe0eff0386666f6f626172");
let parsed = node_from_bytes_backrefs(&mut a, &output).unwrap();
let round_trip = node_to_bytes_backrefs(&a, parsed).unwrap();
assert_eq!(hex::encode(&round_trip), "ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff0180ff0386626172666f6ffe0efe0efe0efe0efe0efe0efe0efe0efe0eff0386666f6f626172");
let round_trip = node_to_bytes(&a, parsed).unwrap();
assert_eq!(hex::encode(&round_trip), "ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff01ffff0180ff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386626172666f6fff0386666f6f626172");
}
#[test]
fn test_restore() {
let mut a = Allocator::new();
let sentinel = a.new_pair(NodePtr::NIL, NodePtr::NIL).unwrap();
let item = node_from_bytes(&mut a, &hex!("ffff0102ff0304")).unwrap();
let list = a.new_pair(item, sentinel).unwrap();
let mut ser = Serializer::new();
let (done, _) = ser.add(&a, list, Some(sentinel)).unwrap();
assert!(!done);
assert_eq!(ser.size(), 8);
assert_eq!(hex::encode(ser.get_ref()), "ffffff0102ff0304");
let (done, state) = ser.add(&a, NodePtr::NIL, None).unwrap();
assert!(done);
assert_eq!(ser.size(), 9);
assert_eq!(hex::encode(ser.get_ref()), "ffffff0102ff030480");
ser.restore(state.clone());
assert_eq!(ser.size(), 8);
assert_eq!(hex::encode(ser.get_ref()), "ffffff0102ff0304");
let (done, _) = ser.add(&a, item, None).unwrap();
assert!(done);
assert_eq!(ser.size(), 10);
assert_eq!(hex::encode(ser.get_ref()), "ffffff0102ff0304fe04");
ser.restore(state);
let item = a.new_small_number(1337).unwrap();
let (done, _) = ser.add(&a, item, None).unwrap();
assert!(done);
assert_eq!(ser.size(), 11);
assert_eq!(hex::encode(ser.get_ref()), "ffffff0102ff0304820539");
let output = ser.into_inner();
assert_eq!(hex::encode(&output), "ffffff0102ff0304820539");
}
}