use async_trait::async_trait;
use std::collections::HashMap;
use std::path::Path;
use crate::processor::{Correction, CorrectionKind, ProcessError, ProcessResult, TextProcessor};
use crate::types::ContextSnapshot;
fn phoneme_distance(a: &str, b: &str) -> usize {
let a: Vec<char> = a.chars().collect();
let b: Vec<char> = b.chars().collect();
let (m, n) = (a.len(), b.len());
let mut prev = (0..=n).collect::<Vec<_>>();
let mut curr = vec![0; n + 1];
for i in 1..=m {
curr[0] = i;
for j in 1..=n {
let cost = if a[i - 1] == b[j - 1] { 0 } else { 1 };
curr[j] = (prev[j] + 1).min(curr[j - 1] + 1).min(prev[j - 1] + cost);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[n]
}
fn phoneme_similarity(a: &str, b: &str) -> f32 {
if a.is_empty() && b.is_empty() {
return 1.0;
}
let max_len = a.chars().count().max(b.chars().count());
let dist = phoneme_distance(a, b);
1.0 - (dist as f32 / max_len as f32)
}
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
a.iter()
.zip(b.iter())
.map(|(x, y)| x * y)
.sum::<f32>()
.max(0.0)
}
pub struct IpaDictionary {
entries: HashMap<String, String>,
}
impl IpaDictionary {
pub fn load(path: impl AsRef<Path>) -> Result<Self, ProcessError> {
let data = std::fs::read_to_string(path.as_ref()).map_err(|e| ProcessError::Unavailable(format!("load IPA dict: {e}")))?;
let entries: HashMap<String, String> =
serde_json::from_str(&data).map_err(|e| ProcessError::Unavailable(format!("parse IPA dict: {e}")))?;
tracing::info!(entries = entries.len(), "IPA dictionary loaded");
Ok(Self { entries })
}
pub fn empty() -> Self {
Self {
entries: HashMap::new(),
}
}
pub fn lookup(&self, word: &str) -> Option<&str> {
self.entries.get(&word.to_lowercase()).map(|s| s.as_str())
}
}
pub trait TextEmbedder: Send + Sync {
fn embed(&self, text: &str) -> Result<Vec<f32>, ProcessError>;
fn similarity_floor(&self) -> f32 {
0.0
}
}
#[derive(Debug, Clone)]
pub struct CustomEntry {
pub word: String,
pub phonemes: String,
pub embedding: Option<Vec<f32>>,
}
pub trait G2pBackend: Send + Sync {
fn phonemize(&self, word: &str) -> Result<String, ProcessError>;
}
pub struct PhonemeCorrector {
ipa_dict: IpaDictionary,
custom_entries: Vec<CustomEntry>,
g2p: Option<Box<dyn G2pBackend>>,
embedder: Option<Box<dyn TextEmbedder>>,
pub alpha: f32,
pub threshold: f32,
pub composite_threshold: f32,
pub max_merge: usize,
}
impl PhonemeCorrector {
pub fn new(ipa_dict: IpaDictionary, custom_entries: Vec<CustomEntry>) -> Self {
Self {
ipa_dict,
custom_entries,
g2p: None,
embedder: None,
alpha: 1.0,
threshold: 0.85,
composite_threshold: 0.65,
max_merge: 3,
}
}
pub fn with_threshold(mut self, threshold: f32) -> Self {
self.threshold = threshold;
self
}
pub fn with_composite_threshold(mut self, threshold: f32) -> Self {
self.composite_threshold = threshold;
self
}
pub fn with_g2p(mut self, g2p: impl G2pBackend + 'static) -> Self {
self.g2p = Some(Box::new(g2p));
self
}
pub fn with_embedder(mut self, embedder: impl TextEmbedder + 'static, alpha: f32) -> Self {
for entry in &mut self.custom_entries {
match embedder.embed(&entry.word) {
Ok(emb) => entry.embedding = Some(emb),
Err(e) => {
tracing::warn!(word = %entry.word, error = %e, "failed to embed custom entry")
}
}
}
self.embedder = Some(Box::new(embedder));
self.alpha = alpha.clamp(0.0, 1.0);
self
}
fn word_to_phonemes(&self, word: &str) -> Option<String> {
let clean: String = word
.chars()
.filter(|c| c.is_alphanumeric() || *c == '\'')
.collect();
if let Some(ipa) = self.ipa_dict.lookup(&clean) {
return Some(ipa.to_string());
}
if let Some(g2p) = &self.g2p {
match g2p.phonemize(&clean) {
Ok(ipa) if !ipa.is_empty() => {
tracing::debug!(word = %clean, phonemes = %ipa, "G2P fallback");
return Some(ipa);
}
Ok(_) => {}
Err(e) => {
tracing::warn!(word = %clean, error = %e, "G2P failed");
}
}
}
None
}
fn best_match(&self, phonemes: &str, text_span: &str) -> Option<(usize, f32)> {
let (text_emb, floor) = match (self.alpha < 1.0, self.embedder.as_ref()) {
(true, Some(e)) => (e.embed(text_span).ok(), e.similarity_floor()),
_ => (None, 0.0),
};
let mut best: Option<(usize, f32)> = None;
for (i, entry) in self.custom_entries.iter().enumerate() {
let phon_sim = phoneme_similarity(phonemes, &entry.phonemes);
let (score, accept_at) = match (&text_emb, &entry.embedding) {
(Some(span_emb), Some(entry_emb)) => {
let text_sim = cosine_similarity(span_emb, entry_emb);
let text_sim = crate::similarity::rescale(text_sim, floor);
(
self.alpha * phon_sim + (1.0 - self.alpha) * text_sim,
self.composite_threshold,
)
}
_ => (phon_sim, self.threshold),
};
if score >= accept_at && (best.is_none() || score > best.unwrap().1) {
best = Some((i, score));
}
}
best
}
}
#[async_trait]
impl TextProcessor for PhonemeCorrector {
async fn process(
&self,
text: &str,
_ctx: &ContextSnapshot,
) -> Result<ProcessResult, ProcessError> {
if self.custom_entries.is_empty() {
return Ok(ProcessResult {
text: text.to_string(),
corrections: vec![],
});
}
let words: Vec<&str> = text.split_whitespace().collect();
if words.is_empty() {
return Ok(ProcessResult {
text: String::new(),
corrections: vec![],
});
}
let word_phonemes: Vec<Option<String>> =
words.iter().map(|w| self.word_to_phonemes(w)).collect();
let mut result_words: Vec<String> = words.iter().map(|w| w.to_string()).collect();
let mut consumed = vec![false; words.len()]; let mut corrections = Vec::new();
struct Candidate {
start: usize,
len: usize,
entry_idx: usize,
similarity: f32,
}
let mut candidates: Vec<Candidate> = Vec::new();
for i in 0..words.len() {
if let Some(phonemes) = &word_phonemes[i] {
if let Some((idx, sim)) = self.best_match(phonemes, words[i]) {
if words[i].to_lowercase() != self.custom_entries[idx].word.to_lowercase() {
candidates.push(Candidate {
start: i,
len: 1,
entry_idx: idx,
similarity: sim,
});
}
}
}
for merge_len in 2..=self.max_merge.min(words.len() - i) {
let window_phonemes: Option<String> = (i..i + merge_len)
.map(|j| word_phonemes[j].as_deref())
.collect::<Option<Vec<_>>>()
.map(|parts| parts.concat());
if let Some(merged) = &window_phonemes {
let text_span: String = (i..i + merge_len)
.map(|j| words[j])
.collect::<Vec<_>>()
.join(" ");
if let Some((idx, sim)) = self.best_match(merged, &text_span) {
candidates.push(Candidate {
start: i,
len: merge_len,
entry_idx: idx,
similarity: sim,
});
}
}
}
}
candidates.sort_by(|a, b| {
b.similarity
.partial_cmp(&a.similarity)
.unwrap()
.then_with(|| a.len.cmp(&b.len))
});
for cand in &candidates {
let end = cand.start + cand.len;
if (cand.start..end).any(|j| consumed[j]) {
continue;
}
let original: Vec<&str> = (cand.start..end).map(|j| words[j]).collect();
let original_str = original.join(" ");
tracing::debug!(
original = %original_str,
replacement = %self.custom_entries[cand.entry_idx].word,
similarity = cand.similarity,
merge_len = cand.len,
"phoneme match"
);
corrections.push(Correction {
kind: CorrectionKind::DictionaryMatch,
original: original_str,
replacement: self.custom_entries[cand.entry_idx].word.clone(),
});
result_words[cand.start] = self.custom_entries[cand.entry_idx].word.clone();
consumed[cand.start..end].fill(true);
}
let final_words: Vec<&str> = result_words
.iter()
.enumerate()
.filter(|(i, _)| !consumed[*i] || *i < words.len() && result_words[*i] != words[*i])
.map(|(_, w)| w.as_str())
.collect();
Ok(ProcessResult {
text: final_words.join(" "),
corrections,
})
}
}
#[cfg(feature = "onnx")]
pub struct OnnxG2p {
session: std::sync::Mutex<ort::session::Session>,
text_to_idx: HashMap<String, i64>,
idx_to_phoneme: HashMap<i64, String>,
lang_token: i64,
char_repeats: usize,
}
#[cfg(feature = "onnx")]
impl OnnxG2p {
pub fn load(model_dir: impl AsRef<Path>) -> Result<Self, ProcessError> {
let dir = model_dir.as_ref();
let session = ort::session::Session::builder()
.and_then(|mut b| b.commit_from_file(dir.join("g2p.onnx")))
.map_err(|e| ProcessError::Unavailable(format!("load G2P model: {e}")))?;
let tok_data =
std::fs::read_to_string(dir.join("tokenizer.json")).map_err(|e| ProcessError::Unavailable(format!("load tokenizer: {e}")))?;
let tok: serde_json::Value = serde_json::from_str(&tok_data).map_err(|e| ProcessError::Failed(format!("parse tokenizer: {e}")))?;
let text_to_idx: HashMap<String, i64> = tok["text_to_idx"]
.as_object()
.ok_or_else(|| ProcessError::Unavailable("missing text_to_idx".into()))?
.iter()
.map(|(k, v)| (k.clone(), v.as_i64().unwrap_or(0)))
.collect();
let idx_to_phoneme: HashMap<i64, String> = tok["idx_to_phoneme"]
.as_object()
.ok_or_else(|| ProcessError::Unavailable("missing idx_to_phoneme".into()))?
.iter()
.map(|(k, v)| {
(
k.parse::<i64>().unwrap_or(0),
v.as_str().unwrap_or("").to_string(),
)
})
.collect();
let lang_token = *text_to_idx.get("<en_us>").unwrap_or(&2);
tracing::info!(
text_symbols = text_to_idx.len(),
phoneme_symbols = idx_to_phoneme.len(),
"ONNX G2P loaded"
);
Ok(Self {
session: std::sync::Mutex::new(session),
text_to_idx,
idx_to_phoneme,
lang_token,
char_repeats: 3,
})
}
fn tokenize(&self, word: &str) -> Vec<i64> {
let mut tokens = vec![self.lang_token];
for c in word.to_lowercase().chars() {
if let Some(&idx) = self.text_to_idx.get(&c.to_string()) {
tokens.push(idx);
}
}
let mut repeated = Vec::with_capacity(tokens.len() * self.char_repeats);
for t in &tokens {
for _ in 0..self.char_repeats {
repeated.push(*t);
}
}
repeated
}
fn ctc_decode(&self, logits: &[f32], seq_len: usize, n_classes: usize) -> String {
let mut decoded = Vec::new();
let mut prev: Option<i64> = None;
for t in 0..seq_len {
let offset = t * n_classes;
let best = (0..n_classes)
.max_by(|&a, &b| logits[offset + a].partial_cmp(&logits[offset + b]).unwrap())
.unwrap_or(0) as i64;
if best != 0 && Some(best) != prev {
decoded.push(best);
}
prev = Some(best);
}
decoded
.iter()
.filter_map(|&idx| self.idx_to_phoneme.get(&idx))
.filter(|s| !s.starts_with('<'))
.cloned()
.collect::<String>()
}
}
#[cfg(feature = "onnx")]
impl G2pBackend for OnnxG2p {
fn phonemize(&self, word: &str) -> Result<String, ProcessError> {
use ndarray::{Array1, Array2};
use ort::value::Value;
if word.is_empty() {
return Ok(String::new());
}
let tokens = self.tokenize(word);
let seq_len = tokens.len();
let text = Array2::from_shape_vec((1, seq_len), tokens).map_err(|e| ProcessError::Failed(format!("shape: {e}")))?;
let start_index =
Array2::from_shape_vec((1, 1), vec![0_i64]).map_err(|e| ProcessError::Failed(format!("shape: {e}")))?;
let text_len = Array1::from_vec(vec![seq_len as i64]);
let mut session = self.session.lock().unwrap();
let outputs = session
.run(vec![
(
"text",
Value::from_array(text)
.map_err(|e| ProcessError::Failed(format!("{e}")))?
.into_dyn(),
),
(
"start_index",
Value::from_array(start_index)
.map_err(|e| ProcessError::Failed(format!("{e}")))?
.into_dyn(),
),
(
"text_len",
Value::from_array(text_len)
.map_err(|e| ProcessError::Failed(format!("{e}")))?
.into_dyn(),
),
])
.map_err(|e| ProcessError::Inference(format!("G2P inference: {e}")))?;
let logits = outputs[0]
.try_extract_array::<f32>()
.map_err(|e| ProcessError::Failed(format!("extract: {e}")))?;
let view = logits.view();
let out_seq = view.shape()[1];
let n_classes = view.shape()[2];
let logits_flat: Vec<f32> = view.iter().copied().collect();
drop(outputs);
drop(session);
Ok(self.ctc_decode(&logits_flat, out_seq, n_classes))
}
}
#[cfg(feature = "onnx")]
pub struct OnnxTextEmbedder {
backend: std::sync::Mutex<crate::embedding::EmbeddingBackend>,
}
#[cfg(feature = "onnx")]
impl OnnxTextEmbedder {
pub fn load(model_dir: impl AsRef<Path>) -> Result<Self, ProcessError> {
let backend = crate::embedding::EmbeddingBackend::load(model_dir)
.map_err(ProcessError::Failed)?;
tracing::info!("ONNX text embedder loaded");
Ok(Self {
backend: std::sync::Mutex::new(backend),
})
}
}
#[cfg(feature = "onnx")]
impl TextEmbedder for OnnxTextEmbedder {
fn embed(&self, text: &str) -> Result<Vec<f32>, ProcessError> {
let mut backend = self
.backend
.lock()
.map_err(|e| ProcessError::Failed(format!("embedder mutex poisoned: {e}")))?;
backend.embed(text).map_err(ProcessError::Failed)
}
fn similarity_floor(&self) -> f32 {
match self.backend.lock() {
Ok(mut b) => b.similarity_floor(),
Err(e) => {
tracing::warn!(error = %e, "embedder mutex poisoned; floor defaults to 0");
0.0
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
struct FlooredEmbedder {
floor: f32,
vector: Vec<f32>,
}
impl TextEmbedder for FlooredEmbedder {
fn embed(&self, _text: &str) -> Result<Vec<f32>, ProcessError> {
Ok(self.vector.clone())
}
fn similarity_floor(&self) -> f32 {
self.floor
}
}
#[test]
fn default_embedder_floor_is_zero_so_behaviour_is_unchanged() {
struct Plain;
impl TextEmbedder for Plain {
fn embed(&self, _t: &str) -> Result<Vec<f32>, ProcessError> {
Ok(vec![1.0, 0.0])
}
}
assert_eq!(Plain.similarity_floor(), 0.0);
}
#[test]
fn two_backends_with_different_floors_score_alike() {
let entry = CustomEntry {
word: "TensorFlow".into(),
phonemes: "tɛnsɝfloʊ".into(),
embedding: Some(vec![1.0, 0.0]),
};
let low_floor = FlooredEmbedder {
floor: 0.45,
vector: vec![0.725, (1.0f32 - 0.725 * 0.725).sqrt()],
};
let high_floor = FlooredEmbedder {
floor: 0.70,
vector: vec![0.850, (1.0f32 - 0.850 * 0.850).sqrt()],
};
let score = |e: FlooredEmbedder| {
let c = PhonemeCorrector::new(IpaDictionary::empty(), vec![entry.clone()])
.with_embedder(e, 0.7);
c.best_match("tɛnsɝfloʊ", "tensor flow").map(|(_, s)| s)
};
let a = score(low_floor).expect("low-floor backend produced no match");
let b = score(high_floor).expect("high-floor backend produced no match");
assert!((a - b).abs() < 1e-4, "{a} vs {b}");
}
#[test]
fn without_rescaling_the_same_alpha_would_have_diverged() {
const ALPHA: f32 = 0.7;
let phon = 1.0f32;
let raw_low = ALPHA * phon + (1.0 - ALPHA) * 0.725;
let raw_high = ALPHA * phon + (1.0 - ALPHA) * 0.850;
assert!((raw_high - raw_low).abs() > 1e-3);
let scaled_low = ALPHA * phon + (1.0 - ALPHA) * crate::similarity::rescale(0.725, 0.45);
let scaled_high = ALPHA * phon + (1.0 - ALPHA) * crate::similarity::rescale(0.850, 0.70);
assert!((scaled_high - scaled_low).abs() < 1e-5);
}
#[test]
fn composite_threshold_is_lower_than_the_phoneme_only_one() {
let c = PhonemeCorrector::new(IpaDictionary::empty(), vec![]);
assert!(c.composite_threshold < c.threshold);
}
#[test]
fn a_failed_entry_embedding_is_judged_by_the_phoneme_only_bar() {
let entry = CustomEntry {
word: "Kubernetes".into(),
phonemes: "kubɝnɛtiz".into(),
embedding: None,
};
let corrector = PhonemeCorrector::new(IpaDictionary::empty(), vec![entry])
.with_composite_threshold(0.0);
assert!(corrector.best_match("kupɚnɛt", "cooper net").is_none());
}
#[test]
fn default_corrector_uses_phoneme_only_scoring() {
let c = PhonemeCorrector::new(IpaDictionary::empty(), vec![]);
assert_eq!(c.alpha, 1.0);
}
#[test]
fn test_phoneme_distance_identical() {
assert_eq!(phoneme_distance("həloʊ", "həloʊ"), 0);
}
#[test]
fn test_phoneme_distance_one_edit() {
assert_eq!(phoneme_distance("ɪfɛkt", "ɛfɛkt"), 1);
}
#[test]
fn test_phoneme_similarity() {
let sim = phoneme_similarity("juːsɪfɛkt", "juːsɪfɛkt");
assert!((sim - 1.0).abs() < 1e-6);
let sim2 = phoneme_similarity("juːs", "juːsɪfɛkt");
assert!(sim2 < 0.7); }
#[test]
fn test_phoneme_distance_empty() {
assert_eq!(phoneme_distance("", "abc"), 3);
assert_eq!(phoneme_distance("abc", ""), 3);
assert_eq!(phoneme_distance("", ""), 0);
}
#[tokio::test]
async fn test_corrector_single_word() {
let dict = IpaDictionary::empty();
let custom = vec![CustomEntry {
word: "Kubernetes".into(),
phonemes: "kuːbɝniːts".into(),
embedding: None,
}];
let corrector = PhonemeCorrector::new(dict, custom);
let ctx = ContextSnapshot::default();
let result = corrector.process("kuber nets", &ctx).await.unwrap();
assert_eq!(result.text, "kuber nets");
}
#[tokio::test]
async fn test_corrector_merge_with_dict() {
let mut entries = HashMap::new();
entries.insert("use".into(), "juːs".into());
entries.insert("effect".into(), "ɪfɛkt".into());
entries.insert("java".into(), "dʒɑːvə".into());
entries.insert("script".into(), "skrɪpt".into());
let dict = IpaDictionary { entries };
let custom = vec![
CustomEntry {
word: "useEffect".into(),
phonemes: "juːsɪfɛkt".into(),
embedding: None,
},
CustomEntry {
word: "JavaScript".into(),
phonemes: "dʒɑːvəskrɪpt".into(),
embedding: None,
},
];
let corrector = PhonemeCorrector::new(dict, custom);
let ctx = ContextSnapshot::default();
let r = corrector.process("use effect", &ctx).await.unwrap();
assert_eq!(r.text, "useEffect");
assert_eq!(r.corrections.len(), 1);
let r2 = corrector.process("java script", &ctx).await.unwrap();
assert_eq!(r2.text, "JavaScript");
let r3 = corrector
.process("I called use effect in java script", &ctx)
.await
.unwrap();
assert_eq!(r3.text, "I called useEffect in JavaScript");
}
#[tokio::test]
async fn test_corrector_no_false_positive() {
let mut entries = HashMap::new();
entries.insert("use".into(), "juːs".into());
entries.insert("the".into(), "ðə".into());
entries.insert("computer".into(), "kəmpjuːtɝ".into());
let dict = IpaDictionary { entries };
let custom = vec![CustomEntry {
word: "useEffect".into(),
phonemes: "juːsɪfɛkt".into(),
embedding: None,
}];
let corrector = PhonemeCorrector::new(dict, custom);
let ctx = ContextSnapshot::default();
let r = corrector.process("use the computer", &ctx).await.unwrap();
assert_eq!(r.text, "use the computer");
assert_eq!(r.corrections.len(), 0);
}
#[tokio::test]
async fn test_corrector_empty_custom_dict() {
let dict = IpaDictionary::empty();
let corrector = PhonemeCorrector::new(dict, vec![]);
let ctx = ContextSnapshot::default();
let r = corrector.process("hello world", &ctx).await.unwrap();
assert_eq!(r.text, "hello world");
assert_eq!(r.corrections.len(), 0);
}
}