use libdictenstein::substring::{BidirectionalDictionaryNode, SubstringMatch};
use libdictenstein::DictionaryNode;
#[derive(Debug, Clone)]
pub struct ExtensionState<N: DictionaryNode> {
pub node: N,
pub query_pos: usize,
pub distance: usize,
pub labels: Vec<N::Unit>,
}
impl<N: DictionaryNode> ExtensionState<N> {
pub fn new(node: N, query_pos: usize, distance: usize) -> Self {
ExtensionState {
node,
query_pos,
distance,
labels: Vec::new(),
}
}
pub fn with_labels(node: N, query_pos: usize, distance: usize, labels: Vec<N::Unit>) -> Self {
ExtensionState {
node,
query_pos,
distance,
labels,
}
}
}
pub struct BidirectionalExtension<'a, N>
where
N: BidirectionalDictionaryNode,
N::Unit: Into<u32>,
{
match_info: &'a SubstringMatch<N>,
max_distance: usize,
query_chars: Vec<char>,
piece_start: usize,
piece_end: usize,
}
impl<'a, N> BidirectionalExtension<'a, N>
where
N: BidirectionalDictionaryNode,
N::Unit: Into<u32>,
{
pub fn new(
match_info: &'a SubstringMatch<N>,
query: &str,
piece_start: usize,
piece_end: usize,
max_distance: usize,
) -> Self {
BidirectionalExtension {
match_info,
max_distance,
query_chars: query.chars().collect(),
piece_start,
piece_end,
}
}
pub fn extend(&self) -> Vec<(String, usize)> {
let mut results = Vec::new();
let left_states = self.extend_left();
for left_state in left_states {
let right_results = self.extend_right(&left_state);
results.extend(right_results);
}
results
}
fn extend_left(&self) -> Vec<LeftExtensionState<N>> {
let mut states = Vec::new();
let query_prefix: Vec<char> = self.query_chars[..self.piece_start]
.iter()
.copied()
.collect();
let start_node = self.find_match_start_node();
if let Some(start) = start_node {
let initial = LeftExtensionState {
node: start,
query_remaining: query_prefix.len(),
distance: 0,
prefix_labels: Vec::new(),
};
self.extend_left_recursive(initial, &query_prefix, &mut states);
}
states
}
fn find_match_start_node(&self) -> Option<N> {
let _term_chars: Vec<char> = self.match_info.term.chars().collect();
if self.match_info.position == 0 {
let mut current = self.match_info.node.clone();
let mut depth = 0;
while let Some(_) = current.parent() {
depth += 1;
if depth >= self.match_info.position + self.match_info.length {
break;
}
if let Some(parent) = current.parent() {
current = parent;
} else {
break;
}
}
Some(current)
} else {
let mut current = self.match_info.node.clone();
for _ in 0..self.match_info.length {
if let Some(parent) = current.parent() {
current = parent;
} else {
break;
}
}
Some(current)
}
}
fn extend_left_recursive(
&self,
state: LeftExtensionState<N>,
query_prefix: &[char],
results: &mut Vec<LeftExtensionState<N>>,
) {
if state.query_remaining == 0 {
if state.node.is_root() || state.distance <= self.max_distance {
results.push(state);
}
return;
}
if state.distance > self.max_distance {
return;
}
let query_idx = query_prefix.len() - state.query_remaining;
let query_char = query_prefix[query_idx];
if let (Some(parent), Some(label)) = (state.node.parent(), state.node.parent_label()) {
let matches = label_matches_char(label, query_char);
if matches {
let mut new_labels = state.prefix_labels.clone();
new_labels.push(label);
let new_state = LeftExtensionState {
node: parent,
query_remaining: state.query_remaining - 1,
distance: state.distance,
prefix_labels: new_labels,
};
self.extend_left_recursive(new_state, query_prefix, results);
} else {
if state.distance + 1 <= self.max_distance {
let mut new_labels = state.prefix_labels.clone();
new_labels.push(label);
let new_state = LeftExtensionState {
node: parent,
query_remaining: state.query_remaining - 1,
distance: state.distance + 1,
prefix_labels: new_labels,
};
self.extend_left_recursive(new_state, query_prefix, results);
}
}
}
if state.distance + 1 <= self.max_distance {
let new_state = LeftExtensionState {
node: state.node.clone(),
query_remaining: state.query_remaining - 1,
distance: state.distance + 1,
prefix_labels: state.prefix_labels.clone(),
};
self.extend_left_recursive(new_state, query_prefix, results);
}
if let (Some(parent), Some(label)) = (state.node.parent(), state.node.parent_label()) {
if state.distance + 1 <= self.max_distance {
let mut new_labels = state.prefix_labels.clone();
new_labels.push(label);
let new_state = LeftExtensionState {
node: parent,
query_remaining: state.query_remaining,
distance: state.distance + 1,
prefix_labels: new_labels,
};
self.extend_left_recursive(new_state, query_prefix, results);
}
}
}
fn extend_right(&self, left_state: &LeftExtensionState<N>) -> Vec<(String, usize)> {
let mut results = Vec::new();
let query_suffix: Vec<char> = self.query_chars[self.piece_end..].iter().copied().collect();
let initial = RightExtensionState {
node: self.match_info.node.clone(),
query_remaining: query_suffix.len(),
distance: left_state.distance,
suffix_labels: Vec::new(),
};
let right_states = self.extend_right_recursive(initial, &query_suffix);
for right_state in right_states {
if right_state.distance <= self.max_distance && right_state.node.is_final() {
let mut term_chars: Vec<char> = Vec::new();
for label in left_state.prefix_labels.iter().rev() {
if let Some(c) = label_to_char(*label) {
term_chars.push(c);
}
}
for c in self.match_info.term.chars() {
term_chars.push(c);
}
let term: String = term_chars.into_iter().collect();
results.push((term, right_state.distance));
}
}
results
}
fn extend_right_recursive(
&self,
state: RightExtensionState<N>,
query_suffix: &[char],
) -> Vec<RightExtensionState<N>> {
let mut results = Vec::new();
if state.query_remaining == 0 {
if state.distance <= self.max_distance {
results.push(state);
}
return results;
}
if state.distance > self.max_distance {
return results;
}
let query_idx = query_suffix.len() - state.query_remaining;
let query_char = query_suffix[query_idx];
for (label, child) in state.node.edges() {
let matches = label_matches_char(label, query_char);
if matches {
let mut new_labels = state.suffix_labels.clone();
new_labels.push(label);
let new_state = RightExtensionState {
node: child,
query_remaining: state.query_remaining - 1,
distance: state.distance,
suffix_labels: new_labels,
};
results.extend(self.extend_right_recursive(new_state, query_suffix));
} else if state.distance + 1 <= self.max_distance {
let mut new_labels = state.suffix_labels.clone();
new_labels.push(label);
let new_state = RightExtensionState {
node: child,
query_remaining: state.query_remaining - 1,
distance: state.distance + 1,
suffix_labels: new_labels,
};
results.extend(self.extend_right_recursive(new_state, query_suffix));
}
}
if state.distance + 1 <= self.max_distance {
let new_state = RightExtensionState {
node: state.node.clone(),
query_remaining: state.query_remaining - 1,
distance: state.distance + 1,
suffix_labels: state.suffix_labels.clone(),
};
results.extend(self.extend_right_recursive(new_state, query_suffix));
}
for (label, child) in state.node.edges() {
if state.distance + 1 <= self.max_distance {
let mut new_labels = state.suffix_labels.clone();
new_labels.push(label);
let new_state = RightExtensionState {
node: child,
query_remaining: state.query_remaining,
distance: state.distance + 1,
suffix_labels: new_labels,
};
results.extend(self.extend_right_recursive(new_state, query_suffix));
}
}
results
}
}
#[derive(Clone)]
struct LeftExtensionState<N: DictionaryNode> {
node: N,
query_remaining: usize,
distance: usize,
prefix_labels: Vec<N::Unit>,
}
#[derive(Clone)]
struct RightExtensionState<N: DictionaryNode> {
node: N,
query_remaining: usize,
distance: usize,
suffix_labels: Vec<N::Unit>,
}
fn label_matches_char<U>(label: U, query_char: char) -> bool
where
U: Copy + PartialEq + Into<u32>,
{
let label_u32: u32 = label.into();
label_u32 == query_char as u32
}
fn label_to_char<U>(label: U) -> Option<char>
where
U: Copy + Into<u32>,
{
let label_u32: u32 = label.into();
char::from_u32(label_u32)
}
#[cfg(test)]
mod tests {
}