use std::collections::HashMap;
use crate::error::{Error, Result};
use crate::index::lucene::codec::data_output::CodecOutput;
const FILE_FORMAT_NAME: &str = "FST";
const VERSION_CURRENT: i32 = 4;
const BIT_FINAL_ARC: u8 = 1 << 0;
const BIT_LAST_ARC: u8 = 1 << 1;
const BIT_TARGET_NEXT: u8 = 1 << 2;
const BIT_STOP_NODE: u8 = 1 << 3;
const BIT_ARC_HAS_OUTPUT: u8 = 1 << 4;
const BIT_ARC_HAS_FINAL_OUTPUT: u8 = 1 << 5;
const FINAL_END_NODE: i64 = -1;
const NON_FINAL_END_NODE: i64 = 0;
#[derive(Clone)]
struct BuilderArc {
label: u8,
target: i64,
output: Vec<u8>,
next_final_output: Vec<u8>,
is_final: bool,
}
#[derive(Clone, Default)]
struct UnCompiledNode {
arcs: Vec<BuilderArc>,
is_final: bool,
output: Vec<u8>,
}
pub struct FstBuilder {
frontier: Vec<UnCompiledNode>,
last_input: Vec<u8>,
bytes: Vec<u8>,
dedup: HashMap<Vec<u8>, i64>,
last_frozen_node: i64,
node_count: u64,
arc_count: u64,
arc_with_output_count: u64,
empty_output: Option<Vec<u8>>,
finished: bool,
}
impl Default for FstBuilder {
fn default() -> Self {
Self::new()
}
}
impl FstBuilder {
#[must_use]
pub fn new() -> Self {
Self {
frontier: vec![UnCompiledNode::default()],
last_input: Vec::new(),
bytes: vec![0],
dedup: HashMap::new(),
last_frozen_node: 0,
node_count: 0,
arc_count: 0,
arc_with_output_count: 0,
empty_output: None,
finished: false,
}
}
pub fn add(&mut self, key: &[u8], output: &[u8]) -> Result<()> {
if self.finished {
return Err(Error::InvalidFormat {
details: "this transducer is already finished".to_owned(),
});
}
if !self.last_input.is_empty() && key < self.last_input.as_slice() {
return Err(Error::InvalidFormat {
details: format!(
"transducer keys arrive in ascending order; {key:?} follows {:?}",
self.last_input
),
});
}
if key.is_empty() {
self.frontier[0].is_final = true;
self.empty_output = Some(output.to_vec());
return Ok(());
}
let mut shared = 0;
let stop = self.last_input.len().min(key.len());
while shared < stop && self.last_input[shared] == key[shared] {
shared += 1;
}
let prefix_len_plus_1 = shared + 1;
while self.frontier.len() < key.len() + 1 {
self.frontier.push(UnCompiledNode::default());
}
self.freeze_tail(prefix_len_plus_1)?;
for index in prefix_len_plus_1..=key.len() {
self.frontier[index - 1].arcs.push(BuilderArc {
label: key[index - 1],
target: 0,
output: Vec::new(),
next_final_output: Vec::new(),
is_final: false,
});
}
let last_index = key.len();
if self.last_input.len() != key.len() || prefix_len_plus_1 != key.len() + 1 {
self.frontier[last_index].is_final = true;
self.frontier[last_index].output.clear();
}
let mut remaining = output.to_vec();
for index in 1..prefix_len_plus_1 {
let label = key[index - 1];
let last_output = self.frontier[index - 1]
.arcs
.last()
.filter(|arc| arc.label == label)
.map(|arc| arc.output.clone())
.unwrap_or_default();
if last_output.is_empty() {
continue;
}
let common = common_prefix(&remaining, &last_output);
let word_suffix = last_output[common.len()..].to_vec();
if let Some(arc) = self.frontier[index - 1].arcs.last_mut() {
arc.output.clone_from(&common);
}
prepend_output(&mut self.frontier[index], &word_suffix);
remaining = remaining[common.len()..].to_vec();
}
if self.last_input.len() == key.len() && prefix_len_plus_1 == key.len() + 1 {
return Err(Error::InvalidFormat {
details: format!("{key:?} is added twice, and a byte-string output cannot merge"),
});
}
if let Some(arc) = self.frontier[prefix_len_plus_1 - 1].arcs.last_mut() {
arc.output = remaining;
}
self.last_input = key.to_vec();
Ok(())
}
fn freeze_tail(&mut self, prefix_len_plus_1: usize) -> Result<()> {
let down_to = prefix_len_plus_1.max(1);
let mut index = self.last_input.len();
while index >= down_to {
let node = std::mem::take(&mut self.frontier[index]);
let is_final = node.is_final || node.arcs.is_empty();
let next_final_output = node.output.clone();
let compiled = self.compile_node(&node)?;
let label = self.last_input[index - 1];
let parent = &mut self.frontier[index - 1];
if let Some(arc) = parent.arcs.last_mut() {
debug_assert_eq!(arc.label, label, "replaceLast targets the last arc");
arc.target = compiled;
arc.next_final_output = next_final_output;
arc.is_final = is_final;
}
if index == 0 {
break;
}
index -= 1;
}
Ok(())
}
fn compile_node(&mut self, node: &UnCompiledNode) -> Result<i64> {
if node.arcs.is_empty() {
return Ok(if node.is_final {
FINAL_END_NODE
} else {
NON_FINAL_END_NODE
});
}
if node.arcs.len() > 1 {
return self.add_node(node);
}
let key = dedup_key(node);
if let Some(address) = self.dedup.get(&key) {
return Ok(*address);
}
let address = self.add_node(node)?;
self.dedup.insert(key, address);
Ok(address)
}
fn add_node(&mut self, node: &UnCompiledNode) -> Result<i64> {
let start = self.bytes.len();
self.arc_count += node.arcs.len() as u64;
let last = node.arcs.len() - 1;
for (index, arc) in node.arcs.iter().enumerate() {
let mut flags = 0u8;
if index == last {
flags |= BIT_LAST_ARC;
}
if self.last_frozen_node == arc.target {
flags |= BIT_TARGET_NEXT;
}
if arc.is_final {
flags |= BIT_FINAL_ARC;
if !arc.next_final_output.is_empty() {
flags |= BIT_ARC_HAS_FINAL_OUTPUT;
}
}
let target_has_arcs = arc.target > 0;
if !target_has_arcs {
flags |= BIT_STOP_NODE;
}
if !arc.output.is_empty() {
flags |= BIT_ARC_HAS_OUTPUT;
}
let mut sink = CodecOutput::new(&mut self.bytes);
sink.write_byte(flags)?;
sink.write_byte(arc.label)?;
if !arc.output.is_empty() {
write_output(&mut sink, &arc.output)?;
self.arc_with_output_count += 1;
}
if !arc.next_final_output.is_empty() {
write_output(&mut sink, &arc.next_final_output)?;
}
if target_has_arcs && flags & BIT_TARGET_NEXT == 0 {
sink.write_vlong(arc.target)?;
}
}
let address = self.bytes.len() as i64 - 1;
self.bytes[start..].reverse();
self.node_count += 1;
self.last_frozen_node = address;
Ok(address)
}
pub fn finish(mut self) -> Result<Vec<u8>> {
if self.finished {
return Err(Error::InvalidFormat {
details: "this transducer is already finished".to_owned(),
});
}
self.freeze_tail(0)?;
self.finished = true;
let root = std::mem::take(&mut self.frontier[0]);
if root.arcs.is_empty() && self.empty_output.is_none() {
return Err(Error::InvalidFormat {
details: "a transducer with no key at all has no serialized form".to_owned(),
});
}
let mut start_node = self.compile_node(&root)?;
if start_node == FINAL_END_NODE && self.empty_output.is_some() {
start_node = 0;
}
let mut buffer = Vec::new();
let mut output = CodecOutput::new(&mut buffer);
output.write_header(FILE_FORMAT_NAME, VERSION_CURRENT)?;
output.write_byte(0)?;
match &self.empty_output {
Some(value) => {
output.write_byte(1)?;
let mut inner = Vec::new();
write_output(&mut CodecOutput::new(&mut inner), value)?;
inner.reverse();
output.write_vint(i32::try_from(inner.len()).unwrap_or(i32::MAX))?;
output.write_bytes(&inner)?;
}
None => output.write_byte(0)?,
}
output.write_byte(0)?;
output.write_vlong(start_node)?;
output.write_vlong(self.node_count as i64)?;
output.write_vlong(self.arc_count as i64)?;
output.write_vlong(self.arc_with_output_count as i64)?;
output.write_vlong(self.bytes.len() as i64)?;
output.write_bytes(&self.bytes)?;
Ok(buffer)
}
}
fn write_output<Sink: std::io::Write>(output: &mut CodecOutput<Sink>, value: &[u8]) -> Result<()> {
output.write_vint(i32::try_from(value.len()).unwrap_or(i32::MAX))?;
output.write_bytes(value)
}
fn common_prefix(left: &[u8], right: &[u8]) -> Vec<u8> {
let shared = left
.iter()
.zip(right.iter())
.take_while(|(a, b)| a == b)
.count();
left[..shared].to_vec()
}
fn prepend_output(node: &mut UnCompiledNode, prefix: &[u8]) {
if prefix.is_empty() {
return;
}
for arc in &mut node.arcs {
let mut combined = prefix.to_vec();
combined.extend_from_slice(&arc.output);
arc.output = combined;
}
if node.is_final {
let mut combined = prefix.to_vec();
combined.extend_from_slice(&node.output);
node.output = combined;
}
}
fn dedup_key(node: &UnCompiledNode) -> Vec<u8> {
let mut key = Vec::new();
for arc in &node.arcs {
key.push(arc.label);
key.extend_from_slice(&arc.target.to_be_bytes());
key.push(u8::from(arc.is_final));
key.extend_from_slice(&(arc.output.len() as u32).to_be_bytes());
key.extend_from_slice(&arc.output);
key.extend_from_slice(&(arc.next_final_output.len() as u32).to_be_bytes());
key.extend_from_slice(&arc.next_final_output);
}
key
}