use std::collections::HashMap;
use std::{io, mem};
use fsst::Compressor;
use integer_encoding::VarIntWriter as _;
use crate::decoder::{ColumnType, Morton};
use crate::encoder::model::{CurveParams, ExplicitEncoder, StrEncoding, StreamCtx};
use crate::encoder::{EncoderConfig, IntEncoder, VertexBufferType};
use crate::utils::BinarySerializer as _;
use crate::{MltError, MltResult};
#[derive(Default)]
pub struct Encoder {
cfg: EncoderConfig,
pub(crate) explicit: Option<ExplicitEncoder>,
hdr: Vec<u8>,
meta: Vec<u8>,
data: Vec<u8>,
pub(crate) morton_cache: Option<Morton>,
pub(crate) hilbert_cache: Option<CurveParams>,
pub(crate) fsst_cache: HashMap<String, Option<Compressor>>,
alt_stack: Vec<AltLevel>,
}
impl Encoder {
#[inline]
#[must_use]
pub fn new(cfg: EncoderConfig) -> Self {
Self {
cfg,
..Self::default()
}
}
#[inline]
#[must_use]
pub fn with_explicit(cfg: EncoderConfig, explicit: ExplicitEncoder) -> Self {
Self {
cfg,
explicit: Some(explicit),
..Self::default()
}
}
#[must_use]
pub(crate) fn preserve_results(&mut self) -> Self {
assert_eq!(self.alt_stack.len(), 0, "Alternatives stack is not empty");
Self {
cfg: EncoderConfig::default(),
explicit: None,
hdr: mem::take(&mut self.hdr),
meta: mem::take(&mut self.meta),
data: mem::take(&mut self.data),
morton_cache: None,
hilbert_cache: None,
fsst_cache: HashMap::new(),
alt_stack: vec![],
}
}
#[inline]
pub(crate) fn write_column_type(&mut self, column_type: ColumnType) -> MltResult<()> {
column_type.write_to(&mut self.meta).map_err(MltError::from)
}
#[inline]
pub(crate) fn write_column_name(&mut self, name: &str) -> MltResult<()> {
self.meta.write_string(name).map_err(MltError::from)
}
#[inline]
#[must_use]
pub fn config(&self) -> EncoderConfig {
self.cfg
}
#[inline]
#[must_use]
pub fn data(&self) -> &[u8] {
&self.data
}
#[inline]
pub(crate) fn data_mut(&mut self) -> &mut Vec<u8> {
&mut self.data
}
#[inline]
#[must_use]
pub fn meta(&self) -> &[u8] {
&self.meta
}
#[inline]
pub(crate) fn meta_mut(&mut self) -> &mut Vec<u8> {
&mut self.meta
}
#[inline]
#[must_use]
pub fn section_lens(&self) -> (usize, usize, usize) {
(self.hdr.len(), self.meta.len(), self.data.len())
}
#[inline]
pub(crate) fn write_column_header(
&mut self,
column_type: ColumnType,
name: &str,
) -> MltResult<()> {
self.write_column_type(column_type)?;
self.write_column_name(name)
}
#[hotpath::measure]
pub fn write_header(&mut self, name: &str, extent: u32, column_count: usize) -> MltResult<()> {
if name.is_empty() {
return Err(MltError::MissingLayerName);
}
debug_assert!(
self.alt_stack.is_empty(),
"write_header called with an open alternatives session"
);
let name_len = u32::try_from(name.len())?;
let column_count = u32::try_from(column_count)?;
self.hdr.write_varint(name_len).map_err(MltError::from)?;
self.hdr.extend_from_slice(name.as_bytes());
self.hdr.write_varint(extent).map_err(MltError::from)?;
self.hdr
.write_varint(column_count)
.map_err(MltError::from)?;
Ok(())
}
#[inline]
pub(crate) fn override_int_enc(&self, ctx: &StreamCtx<'_>) -> Option<IntEncoder> {
self.explicit.as_ref().map(|e| (e.get_int_encoder)(ctx))
}
#[inline]
pub(crate) fn override_str_enc(&self, name: &str) -> Option<StrEncoding> {
self.explicit.as_ref().map(|e| (e.get_str_encoding)(name))
}
#[inline]
#[allow(clippy::unused_self)]
pub(crate) fn override_vertex_buffer_type(&self) -> Option<VertexBufferType> {
self.explicit.as_ref().map(|e| e.vertex_buffer_type)
}
#[inline]
pub(crate) fn force_stream(&self, ctx: &StreamCtx<'_>) -> bool {
self.explicit
.as_ref()
.is_some_and(|e| (e.force_stream)(ctx))
}
#[inline]
#[must_use]
pub fn total_len(&self) -> usize {
self.hdr.len() + self.meta.len() + self.data.len()
}
pub(crate) fn clear_results(&mut self) {
debug_assert!(self.alt_stack.is_empty(), "Alternatives stack is not empty");
self.hdr.clear();
self.meta.clear();
self.data.clear();
}
#[must_use]
pub fn into_raw_bytes(mut self) -> Vec<u8> {
if self.hdr.is_empty() && self.meta.is_empty() {
return self.data;
}
let mut out = Vec::with_capacity(self.hdr.len() + self.meta.len() + self.data.len());
out.append(&mut self.hdr);
out.append(&mut self.meta);
out.append(&mut self.data);
out
}
pub fn into_layer_bytes(self) -> MltResult<Vec<u8>> {
self.into_layer_bytes_with_tag(1)
}
fn into_layer_bytes_with_tag(mut self, tag: u8) -> MltResult<Vec<u8>> {
debug_assert!(
self.alt_stack.is_empty(),
"into_layer_bytes_with_tag called with an open alternatives session"
);
let body_len = self.hdr.len() + self.meta.len() + self.data.len();
let size = u32::try_from(body_len + 1)?; let mut out = Vec::with_capacity(5 + 1 + body_len);
out.write_varint(size).map_err(MltError::from)?;
out.push(tag);
out.append(&mut self.hdr);
out.append(&mut self.meta);
out.append(&mut self.data);
Ok(out)
}
pub fn try_alternatives(&mut self) -> AltSession<'_> {
self.alt_stack.push(AltLevel {
data_start: self.data.len(),
meta_start: self.meta.len(),
best_data: None,
best_meta: None,
});
AltSession { enc: self }
}
fn alt_commit(&mut self) {
debug_assert!(
!self.alt_stack.is_empty(),
"alt_commit called outside an active AltSession"
);
let (data, meta, stack) = (&mut self.data, &mut self.meta, &mut self.alt_stack);
let level = stack.last_mut().unwrap();
Self::close_candidate(data, meta, level);
}
fn alt_pop(&mut self) {
debug_assert!(
!self.alt_stack.is_empty(),
"alt_pop called outside an active AltSession"
);
{
let (data, meta, stack) = (&mut self.data, &mut self.meta, &mut self.alt_stack);
let level = stack.last_mut().unwrap();
let data_pending = data.len() - (level.data_start + level.best_data.unwrap_or(0));
let meta_pending = meta.len() - (level.meta_start + level.best_meta.unwrap_or(0));
if data_pending > 0 || meta_pending > 0 || level.best_data.is_none() {
Self::close_candidate(data, meta, level);
}
}
self.alt_stack.pop();
}
fn close_candidate(data: &mut Vec<u8>, meta: &mut Vec<u8>, level: &mut AltLevel) {
let best_data_end = level.data_start + level.best_data.unwrap_or(0);
let best_meta_end = level.meta_start + level.best_meta.unwrap_or(0);
let cand_data = data.len() - best_data_end;
let cand_meta = meta.len() - best_meta_end;
let cand_total = cand_data + cand_meta;
let best_total = level.best_data.unwrap_or(0) + level.best_meta.unwrap_or(0);
if level.best_data.is_none_or(|_| cand_total < best_total) {
if level.best_data.is_some() {
data.copy_within(best_data_end..best_data_end + cand_data, level.data_start);
meta.copy_within(best_meta_end..best_meta_end + cand_meta, level.meta_start);
}
data.truncate(level.data_start + cand_data);
meta.truncate(level.meta_start + cand_meta);
level.best_data = Some(cand_data);
level.best_meta = Some(cand_meta);
} else {
data.truncate(best_data_end);
meta.truncate(best_meta_end);
}
}
}
#[derive(Debug, Default, Clone)]
struct AltLevel {
data_start: usize,
meta_start: usize,
best_data: Option<usize>,
best_meta: Option<usize>,
}
#[must_use = "AltSession must be used; drop it to finalise the competition"]
pub struct AltSession<'a> {
enc: &'a mut Encoder,
}
impl AltSession<'_> {
#[hotpath::measure]
pub fn with<F>(&mut self, f: F) -> MltResult<()>
where
F: FnOnce(&mut Encoder) -> MltResult<()>,
{
let data_cp = self.enc.data.len();
let meta_cp = self.enc.meta.len();
match f(self.enc) {
Ok(()) => {
self.enc.alt_commit();
Ok(())
}
Err(e) => {
self.enc.data.truncate(data_cp);
self.enc.meta.truncate(meta_cp);
Err(e)
}
}
}
}
impl Drop for AltSession<'_> {
fn drop(&mut self) {
self.enc.alt_pop();
}
}
impl io::Write for Encoder {
#[inline]
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.data.write(buf)
}
#[inline]
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
#[inline]
fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
self.data.write_all(buf)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn push(enc: &mut Encoder, bytes: &[u8]) {
enc.data.extend_from_slice(bytes);
}
#[test]
fn alternatives_keeps_shortest() {
let mut enc = Encoder::default();
push(&mut enc, b"prefix");
let mut alt = enc.try_alternatives();
alt.with(|enc| {
push(enc, b"longer");
Ok(())
})
.unwrap(); alt.with(|enc| {
push(enc, b"ab");
Ok(())
})
.unwrap(); alt.with(|enc| {
push(enc, b"xyz");
Ok(())
})
.unwrap(); drop(alt);
assert_eq!(enc.data, b"prefixab");
}
#[test]
fn alternatives_tie_keeps_first() {
let mut enc = Encoder::default();
let mut alt = enc.try_alternatives();
alt.with(|enc| {
push(enc, b"aaa");
Ok(())
})
.unwrap(); alt.with(|enc| {
push(enc, b"bbb");
Ok(())
})
.unwrap(); drop(alt);
assert_eq!(enc.data, b"aaa");
}
#[test]
fn alternatives_single_candidate() {
let mut enc = Encoder::default();
let mut alt = enc.try_alternatives();
alt.with(|enc| {
push(enc, b"only");
Ok(())
})
.unwrap();
drop(alt);
assert_eq!(enc.data, b"only");
}
#[test]
fn prefix_bytes_are_preserved() {
let mut enc = Encoder::default();
push(&mut enc, b"HDR");
let mut alt = enc.try_alternatives();
alt.with(|enc| {
push(enc, b"long_encoding");
Ok(())
})
.unwrap(); alt.with(|enc| {
push(enc, b"short");
Ok(())
})
.unwrap(); drop(alt);
assert_eq!(&enc.data[..3], b"HDR");
assert_eq!(&enc.data[3..], b"short");
}
#[test]
fn drop_after_all_committed_is_noop() {
let mut enc = Encoder::default();
let mut alt = enc.try_alternatives();
alt.with(|enc| {
push(enc, b"best");
Ok(())
})
.unwrap();
drop(alt);
assert!(enc.alt_stack.is_empty(), "stack empty after drop");
assert_eq!(enc.data, b"best");
}
#[test]
fn nested_alternatives() {
let mut enc = Encoder::default();
let mut outer = enc.try_alternatives();
outer
.with(|enc| {
push(enc, b"A:");
let mut inner = enc.try_alternatives(); inner.with(|enc| {
push(enc, b"long_inner");
Ok(())
})?; inner.with(|enc| {
push(enc, b"in");
Ok(())
})?; drop(inner); push(enc, b"!");
Ok(())
})
.unwrap();
outer
.with(|enc| {
push(enc, b"B");
Ok(())
})
.unwrap(); drop(outer);
assert_eq!(enc.data, b"B");
}
#[test]
fn nesting_depth_reflected_in_stack() {
let mut enc = Encoder::default();
assert_eq!(enc.alt_stack.len(), 0);
let mut outer = enc.try_alternatives();
outer
.with(|enc| {
assert_eq!(enc.alt_stack.len(), 1); let mut inner = enc.try_alternatives();
inner.with(|enc| {
assert_eq!(enc.alt_stack.len(), 2); push(enc, b"x");
Ok(())
})?;
drop(inner); assert_eq!(enc.alt_stack.len(), 1);
push(enc, b"y");
Ok(())
})
.unwrap();
drop(outer); assert_eq!(enc.alt_stack.len(), 0);
}
#[test]
fn alternatives_tracks_meta_and_data() {
let mut enc = Encoder::default();
enc.data.extend_from_slice(b"D");
enc.meta.extend_from_slice(b"M");
let mut alt = enc.try_alternatives();
alt.with(|enc| {
push(enc, b"DDDD");
enc.meta.extend_from_slice(b"mm");
Ok(())
})
.unwrap();
alt.with(|enc| {
push(enc, b"d");
enc.meta.extend_from_slice(b"n");
Ok(())
})
.unwrap();
drop(alt);
assert_eq!(enc.data, b"Dd");
assert_eq!(enc.meta, b"Mn");
}
#[test]
fn error_candidate_is_rolled_back() {
let mut enc = Encoder::default();
let mut alt = enc.try_alternatives();
alt.with(|enc| {
push(enc, b"ok");
Ok(())
})
.unwrap();
let _ = alt.with(|enc| {
push(enc, b"partial");
Err(MltError::IntegerOverflow) });
drop(alt);
assert_eq!(enc.data, b"ok"); }
}