use fsst::Compressor;
use integer_encoding::VarIntWriter as _;
use usize_cast::IntoUsize as _;
use super::model::StagedStrings;
use crate::MltResult;
use crate::codecs::fsst::{FsstRawData, compress_fsst, compress_fsst_with};
use crate::decoder::stream::header01;
use crate::decoder::strings::{checked_string_end, encode_null_end};
use crate::decoder::{DictionaryType, LengthType, OffsetType, StreamMeta, StreamType};
use crate::encoder::model::{StrEncoding, StreamCtx};
use crate::encoder::stream::{dedup_strings, write_stream_payload};
use crate::encoder::{Codecs, Encoder};
use crate::utils::strings_to_lengths;
const FSST_OVERHEAD_THRESHOLD: usize = 2_048;
const FSST_SAMPLE_STRINGS: usize = 256;
#[hotpath::measure]
pub(crate) fn fsst_try_train(strings: &[&str]) -> Option<Compressor> {
if strings.is_empty() {
return None;
}
let total_plain_size: usize = strings.iter().map(|s| s.len()).sum();
if total_plain_size < FSST_OVERHEAD_THRESHOLD {
return None;
}
let byte_slices: Vec<&[u8]> = strings.iter().map(|s| s.as_bytes()).collect();
let compressor = Compressor::train(&byte_slices);
let symbols = compressor.symbol_table();
let symbol_lengths = compressor.symbol_lengths();
let symbol_overhead: usize = symbol_lengths
.iter()
.take(symbols.len())
.map(|&l| usize::from(l))
.sum();
let sample = if strings.len() <= FSST_SAMPLE_STRINGS {
strings
} else {
&strings[..FSST_SAMPLE_STRINGS]
};
let plain_size: usize = sample.iter().map(|s| s.len()).sum();
let compressed_size: usize = sample
.iter()
.map(|s| compressor.compress(s.as_bytes()).len())
.sum();
if symbol_overhead + compressed_size < plain_size {
Some(compressor)
} else {
None
}
}
impl Encoder {
pub(crate) fn fsst_compressor(&mut self, key: &str, corpus: &[&str]) -> Option<&Compressor> {
if !self.config().allow_fsst() {
return None;
}
self.fsst_cache
.entry(key.to_owned())
.or_insert_with(|| fsst_try_train(corpus))
.as_ref()
}
}
impl Codecs {
#[hotpath::measure]
pub(crate) fn write_str_col(
&mut self,
v: &StagedStrings,
presence: Option<&StagedStrings>,
enc: &mut Encoder,
) -> MltResult<()> {
let non_null = v.dense_values();
let name = &v.name;
if let Some(str_enc) = enc.override_str_enc(name) {
match str_enc {
StrEncoding::Plain => write_str_plain(&non_null, presence, name, enc, self)?,
StrEncoding::Dict => write_str_dict(&non_null, presence, name, enc, self)?,
StrEncoding::Fsst => write_str_fsst(&non_null, presence, name, enc, self)?,
StrEncoding::FsstDict => write_str_fsst_dict(&non_null, presence, name, enc, self)?,
}
} else {
let (unique, offset_indices) = dedup_strings(&non_null)?;
let compressor = enc.fsst_compressor(name, &unique);
let count = non_null.len();
let plain_fsst = compressor.map(|c| compress_fsst_with(&non_null, c));
let dict_fsst = compressor.map(|c| compress_fsst_with(&unique, c));
let mut alt = enc.try_alternatives();
alt.with(|enc| write_str_plain(&non_null, presence, name, enc, self))?;
alt.with(|enc| {
write_str_dict_raw(&unique, &offset_indices, presence, name, enc, self)
})?;
if let Some(ref raw) = plain_fsst {
alt.with(|enc| write_str_fsst_raw(raw, count, presence, name, enc, self))?;
}
if let Some(ref raw) = dict_fsst {
alt.with(|enc| {
write_str_fsst_dict_raw(raw, &offset_indices, presence, name, enc, self)
})?;
}
}
Ok(())
}
}
#[hotpath::measure]
fn write_str_plain(
non_null: &[&str],
presence: Option<&StagedStrings>,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
let lengths = strings_to_lengths(non_null)?;
enc.write_varint(2u32 + u32::from(presence.is_some()))?;
write_presence_stream(presence, enc, codecs)?;
let ctx = StreamCtx::prop(StreamType::Length(LengthType::VarBinary), name);
codecs.write_int_stream(&lengths, &ctx, enc)?;
write_raw_str_data(non_null, DictionaryType::None, enc)
}
#[hotpath::measure]
fn write_str_dict(
non_null: &[&str],
presence: Option<&StagedStrings>,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
let (unique, offset_indices) = dedup_strings(non_null)?;
write_str_dict_raw(&unique, &offset_indices, presence, name, enc, codecs)
}
fn write_str_dict_raw(
unique: &[&str],
offset_indices: &[u32],
presence: Option<&StagedStrings>,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
let lengths = strings_to_lengths(unique)?;
enc.write_varint(3u32 + u32::from(presence.is_some()))?;
write_presence_stream(presence, enc, codecs)?;
let ctx = StreamCtx::prop(StreamType::Length(LengthType::Dictionary), name);
codecs.write_int_stream(&lengths, &ctx, enc)?;
let ctx = StreamCtx::prop(StreamType::Offset(OffsetType::String), name);
codecs.write_int_stream(offset_indices, &ctx, enc)?;
write_raw_str_data(unique, DictionaryType::Single, enc)
}
#[hotpath::measure]
fn write_str_fsst(
non_null: &[&str],
presence: Option<&StagedStrings>,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
let raw = compress_fsst(non_null);
write_str_fsst_raw(&raw, non_null.len(), presence, name, enc, codecs)
}
fn write_str_fsst_raw(
raw: &FsstRawData,
count: usize,
presence: Option<&StagedStrings>,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
let offsets: Vec<u32> = (0..u32::try_from(count)?).collect();
enc.write_varint(5u32 + u32::from(presence.is_some()))?;
write_presence_stream(presence, enc, codecs)?;
write_fsst_data(raw, DictionaryType::Single, name, enc, codecs)?;
let ctx = StreamCtx::prop(StreamType::Offset(OffsetType::String), name);
codecs.write_int_stream(&offsets, &ctx, enc)
}
#[hotpath::measure]
fn write_str_fsst_dict(
non_null: &[&str],
presence: Option<&StagedStrings>,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
let (unique, offset_indices) = dedup_strings(non_null)?;
let raw = compress_fsst(&unique);
write_str_fsst_dict_raw(&raw, &offset_indices, presence, name, enc, codecs)
}
fn write_str_fsst_dict_raw(
raw: &FsstRawData,
offset_indices: &[u32],
presence: Option<&StagedStrings>,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
enc.write_varint(5u32 + u32::from(presence.is_some()))?;
write_presence_stream(presence, enc, codecs)?;
write_fsst_data(raw, DictionaryType::Single, name, enc, codecs)?;
let ctx = StreamCtx::prop(StreamType::Offset(OffsetType::String), name);
codecs.write_int_stream(offset_indices, &ctx, enc)
}
fn write_presence_stream(
presence: Option<&StagedStrings>,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
if let Some(strings) = presence {
codecs.write_presence_stream(strings.presence_bools(), enc)?;
}
Ok(())
}
#[hotpath::measure]
pub fn write_fsst_data(
raw: &FsstRawData,
dict_type: DictionaryType,
name: &str,
enc: &mut Encoder,
codecs: &mut Codecs,
) -> MltResult<()> {
let ctx = StreamCtx::prop(StreamType::Length(LengthType::Symbol), name);
codecs.write_int_stream(&raw.symbol_lengths, &ctx, enc)?;
let typ = StreamType::Data(DictionaryType::Fsst);
let meta = StreamMeta::new_none(typ, raw.symbol_lengths.len())?;
write_stream_payload(enc, meta, false, &raw.symbol_bytes)?;
let ctx = StreamCtx::prop(StreamType::Length(LengthType::Dictionary), name);
codecs.write_int_stream(&raw.value_lengths, &ctx, enc)?;
let meta = StreamMeta::new_none(StreamType::Data(dict_type), raw.value_lengths.len())?;
write_stream_payload(enc, meta, false, &raw.corpus)?;
Ok(())
}
#[hotpath::measure]
pub fn write_raw_str_data(
strings: &[&str],
dict_type: DictionaryType,
enc: &mut Encoder,
) -> MltResult<()> {
let total_len: usize = strings.iter().map(|s| s.len()).sum();
let typ = StreamType::Data(dict_type);
let meta = StreamMeta::new_none(typ, strings.len())?;
header01::write_stream_meta(&meta, enc, false, u32::try_from(total_len)?)?;
enc.data_mut().reserve(total_len);
for s in strings {
enc.data_mut().extend_from_slice(s.as_bytes());
}
Ok(())
}
impl StagedStrings {
#[must_use]
pub fn from_strings(
name: impl Into<String>,
values: impl IntoIterator<Item = impl AsRef<str>>,
) -> Self {
let name = name.into();
let iter = values.into_iter();
let (lower, _) = iter.size_hint();
let mut lengths = Vec::with_capacity(lower);
let mut data = String::new();
let mut end = 0_i32;
for value in iter {
let value = value.as_ref();
end = checked_string_end(end, value.len())
.expect("staged string corpus exceeds supported i32 range");
lengths.push(end);
data.push_str(value);
}
Self {
name,
lengths,
data,
}
}
#[must_use]
pub fn from_optional(
name: impl Into<String>,
values: impl IntoIterator<Item = Option<impl AsRef<str>>>,
) -> Self {
let name = name.into();
let iter = values.into_iter();
let (lower, _) = iter.size_hint();
let mut lengths = Vec::with_capacity(lower);
let mut data = String::new();
let mut end = 0_i32;
for value in iter {
match value {
Some(value) => {
let value = value.as_ref();
end = checked_string_end(end, value.len())
.expect("staged string corpus exceeds supported i32 range");
lengths.push(end);
data.push_str(value);
}
None => lengths.push(encode_null_end(end)),
}
}
Self {
name,
lengths,
data,
}
}
#[must_use]
pub fn feature_count(&self) -> usize {
self.lengths.len()
}
pub fn presence_bools(&self) -> impl ExactSizeIterator<Item = bool> + '_ {
self.lengths.iter().map(|&end| end >= 0)
}
#[must_use]
pub fn dense_values(&self) -> Vec<&str> {
let mut values = Vec::new();
let mut start = 0_u32;
for &end in &self.lengths {
if end >= 0 {
let end = end.cast_unsigned();
values.push(&self.data[start.into_usize()..end.into_usize()]);
start = end;
} else {
start = (!end).cast_unsigned();
}
}
values
}
}