Skip to main content

aurum_core/tts/
local.rs

1//! Local ONNX KittenTTS provider (MIT binary path — no GPL phonemizer).
2
3use super::catalogue::{ensure_voice_pack, lookup_model, lookup_voice, onnx_path, voices_path};
4use super::npz::load_voices_npz;
5use super::provider::{BackendKind, SynthesisOptions, SynthesisProvider, SynthesisResult};
6use super::tokenize::ipa_to_ids;
7use super::validate::{
8    clamp_speaking_rate, normalize_tts_language, prepare_text, DEFAULT_MAX_CHARS,
9};
10use super::wav::peak_guard_f32_to_i16;
11use crate::error::{ProviderError, Result, UserError};
12use async_trait::async_trait;
13use ort::session::Session;
14use ort::value::Tensor;
15use std::collections::HashMap;
16use std::path::{Path, PathBuf};
17use std::sync::{Arc, Mutex};
18use std::time::Duration;
19
20/// Samples trimmed from the tail of every chunk (trailing silence artifact).
21const TAIL_TRIM: usize = 2_000;
22/// Peak limit before i16 quantization.
23const PEAK_LIMIT: f32 = 0.95;
24
25/// On-device KittenTTS via ONNX Runtime + misaki-rs G2P (no espeak / GPL).
26pub struct LocalTtsProvider {
27    cache_dir: PathBuf,
28    show_progress: bool,
29    local_only: bool,
30    max_chars: usize,
31    /// Lazily loaded sessions keyed by model id.
32    sessions: Mutex<HashMap<String, Arc<LoadedPack>>>,
33}
34
35struct LoadedPack {
36    session: Mutex<Session>,
37    voices: HashMap<String, super::npz::VoiceMatrix>,
38    sample_rate_hz: u32,
39    /// Optional speed priors from config.json (internal key → multiplier).
40    speed_priors: HashMap<String, f32>,
41}
42
43impl LocalTtsProvider {
44    pub fn new(cache_dir: PathBuf) -> Self {
45        Self {
46            cache_dir,
47            show_progress: false,
48            local_only: false,
49            max_chars: DEFAULT_MAX_CHARS,
50            sessions: Mutex::new(HashMap::new()),
51        }
52    }
53
54    pub fn with_progress(mut self, v: bool) -> Self {
55        self.show_progress = v;
56        self
57    }
58
59    pub fn with_local_only(mut self, v: bool) -> Self {
60        self.local_only = v;
61        self
62    }
63
64    pub fn with_max_chars(mut self, n: usize) -> Self {
65        self.max_chars = n.max(1);
66        self
67    }
68
69    /// Drop loaded ONNX sessions for this provider (frees ORT graphs held in RAM).
70    ///
71    /// Safe to call anytime; the next synthesize/preload reloads from the on-disk pack.
72    /// Does not delete cached files under the TTS cache directory.
73    pub fn clear_sessions(&self) {
74        if let Ok(mut guard) = self.sessions.lock() {
75            guard.clear();
76        }
77    }
78
79    async fn ensure_loaded(&self, model: &str, local_only: bool) -> Result<Arc<LoadedPack>> {
80        {
81            let guard = self.sessions.lock().map_err(|_| {
82                crate::error::TranscriptionError::internal("TTS session map poisoned")
83            })?;
84            if let Some(pack) = guard.get(model) {
85                return Ok(Arc::clone(pack));
86            }
87        }
88
89        let info = lookup_model(model)?;
90        let _pack_dir = ensure_voice_pack(
91            &self.cache_dir,
92            model,
93            self.show_progress,
94            local_only || self.local_only,
95        )
96        .await?;
97
98        let onnx = onnx_path(&self.cache_dir, info);
99        let voices_file = voices_path(&self.cache_dir, info);
100        let speed_priors = load_speed_priors(&self.cache_dir, info);
101        let sample_rate = info.sample_rate_hz;
102
103        let loaded = tokio::task::spawn_blocking(move || {
104            load_pack(&onnx, &voices_file, sample_rate, speed_priors)
105        })
106        .await
107        .map_err(|e| crate::error::TranscriptionError::internal(format!("TTS load join: {e}")))??;
108
109        let arc = Arc::new(loaded);
110        let mut guard = self
111            .sessions
112            .lock()
113            .map_err(|_| crate::error::TranscriptionError::internal("TTS session map poisoned"))?;
114        let entry = guard
115            .entry(model.to_string())
116            .or_insert_with(|| Arc::clone(&arc));
117        Ok(Arc::clone(entry))
118    }
119}
120
121fn synthesize_with_pack(
122    pack: &LoadedPack,
123    text: &str,
124    opts: &SynthesisOptions,
125    text_chars: usize,
126    text_truncated: bool,
127) -> Result<SynthesisResult> {
128    if let Some(flag) = &opts.cancel {
129        if flag.is_cancelled() {
130            return Err(ProviderError::Cancelled.into());
131        }
132    }
133
134    let voice = lookup_voice(&opts.voice)?;
135    let internal = voice.internal_key;
136    let rate = clamp_speaking_rate(opts.speaking_rate);
137
138    // G2P (MIT misaki-rs, no espeak).
139    let g2p = misaki_rs::G2P::new(misaki_rs::Language::EnglishUS);
140    let (ipa, _) = g2p.g2p(text).map_err(|e| ProviderError::Other {
141        message: format!("G2P failed: {e}"),
142    })?;
143    // Strip unknown markers that misaki may emit without espeak fallback.
144    let ipa = ipa.replace('❓', "");
145    if ipa.trim().is_empty() {
146        return Err(ProviderError::Other {
147            message: "G2P produced empty phonemes for input text".into(),
148        }
149        .into());
150    }
151
152    let voice_mat = pack
153        .voices
154        .get(internal)
155        .or_else(|| pack.voices.get(&opts.voice))
156        .ok_or_else(|| UserError::Other {
157            message: format!(
158                "voice embedding '{internal}' missing from pack; available: {:?}",
159                pack.voices.keys().collect::<Vec<_>>()
160            ),
161        })?;
162
163    let effective_speed = rate * pack.speed_priors.get(internal).copied().unwrap_or(1.0);
164    let ids = ipa_to_ids(&ipa);
165    if ids.len() <= 2 {
166        return Err(ProviderError::Other {
167            message: format!("no tokenizable phonemes from IPA {ipa:?}"),
168        }
169        .into());
170    }
171    let seq_len = ids.len();
172    let style = voice_mat.style_row(text.chars().count()).to_vec();
173    let style_dim = style.len();
174
175    let t_ids =
176        Tensor::<i64>::from_array(([1usize, seq_len], ids)).map_err(|e| ProviderError::Other {
177            message: format!("input_ids tensor: {e}"),
178        })?;
179    let t_style = Tensor::<f32>::from_array(([1usize, style_dim], style)).map_err(|e| {
180        ProviderError::Other {
181            message: format!("style tensor: {e}"),
182        }
183    })?;
184    let t_speed = Tensor::<f32>::from_array(([1usize], vec![effective_speed])).map_err(|e| {
185        ProviderError::Other {
186            message: format!("speed tensor: {e}"),
187        }
188    })?;
189
190    let mut session = pack
191        .session
192        .lock()
193        .map_err(|_| crate::error::TranscriptionError::internal("ORT session mutex poisoned"))?;
194    let outputs = session
195        .run(ort::inputs![t_ids, t_style, t_speed])
196        .map_err(|e| ProviderError::Other {
197            message: format!("ONNX inference failed: {e}"),
198        })?;
199    let (_shape, audio_data) =
200        outputs[0]
201            .try_extract_tensor::<f32>()
202            .map_err(|e| ProviderError::Other {
203                message: format!("extract audio tensor: {e}"),
204            })?;
205    let audio_flat: Vec<f32> = audio_data.to_vec();
206    let trimmed_len = audio_flat.len().saturating_sub(TAIL_TRIM);
207    let audio = &audio_flat[..trimmed_len];
208    if audio.is_empty() {
209        return Err(ProviderError::Other {
210            message: "synthesis produced empty audio".into(),
211        }
212        .into());
213    }
214
215    let pcm = peak_guard_f32_to_i16(audio, PEAK_LIMIT);
216    let sample_rate = opts.sample_rate_hz.unwrap_or(pack.sample_rate_hz);
217    let duration_ms = (pcm.len() as u64)
218        .saturating_mul(1000)
219        .checked_div(sample_rate as u64)
220        .unwrap_or(0);
221
222    // Report the engine language actually used (English-only G2P), not a raw echo.
223    let language = normalize_tts_language(&opts.language).unwrap_or_else(|_| "en".into());
224
225    Ok(SynthesisResult {
226        pcm_i16_mono: pcm,
227        sample_rate_hz: sample_rate,
228        channels: 1,
229        backend_kind: BackendKind::Local,
230        provider: "local".into(),
231        model: opts.model.clone(),
232        voice: voice.id.to_string(),
233        language,
234        duration_ms,
235        text_chars,
236        text_truncated,
237    })
238}
239
240fn load_pack(
241    onnx: &Path,
242    voices_file: &Path,
243    sample_rate_hz: u32,
244    speed_priors: HashMap<String, f32>,
245) -> Result<LoadedPack> {
246    let session = Session::builder()
247        .map_err(|e| ProviderError::ModelLoad {
248            model: onnx.display().to_string(),
249            reason: format!("ORT session builder: {e}"),
250        })?
251        .commit_from_file(onnx)
252        .map_err(|e| ProviderError::ModelLoad {
253            model: onnx.display().to_string(),
254            reason: format!("load ONNX: {e}"),
255        })?;
256    let voices = load_voices_npz(voices_file)?;
257    Ok(LoadedPack {
258        session: Mutex::new(session),
259        voices,
260        sample_rate_hz,
261        speed_priors,
262    })
263}
264
265fn load_speed_priors(
266    cache_dir: &Path,
267    info: &super::catalogue::TtsModelInfo,
268) -> HashMap<String, f32> {
269    let path = super::catalogue::config_path(cache_dir, info);
270    let Ok(bytes) = std::fs::read(&path) else {
271        return HashMap::new();
272    };
273    let Ok(v) = serde_json::from_slice::<serde_json::Value>(&bytes) else {
274        return HashMap::new();
275    };
276    let mut out = HashMap::new();
277    if let Some(map) = v.get("speed_priors").and_then(|x| x.as_object()) {
278        for (k, val) in map {
279            if let Some(f) = val.as_f64() {
280                out.insert(k.clone(), f as f32);
281            }
282        }
283    }
284    out
285}
286
287#[async_trait]
288impl SynthesisProvider for LocalTtsProvider {
289    fn name(&self) -> &'static str {
290        "local"
291    }
292
293    async fn synthesize(&self, text: &str, opts: &SynthesisOptions) -> Result<SynthesisResult> {
294        let prepared = prepare_text(text, self.max_chars)?;
295        // Reject unsupported languages up front (G2P is English-only).
296        let mut opts = opts.clone();
297        opts.language = normalize_tts_language(&opts.language)?;
298        lookup_model(&opts.model)?;
299        lookup_voice(&opts.voice)?;
300
301        let local_only = opts.local_only || self.local_only;
302        let pack = self.ensure_loaded(&opts.model, local_only).await?;
303
304        if let Some(flag) = &opts.cancel {
305            if flag.is_cancelled() {
306                return Err(ProviderError::Cancelled.into());
307            }
308        }
309
310        let timeout = Duration::from_millis(if opts.timeout_ms == 0 {
311            super::validate::DEFAULT_TIMEOUT_MS
312        } else {
313            opts.timeout_ms
314        });
315
316        let text_owned = prepared.text.clone();
317        let text_chars = prepared.text_chars;
318        let text_truncated = prepared.text_truncated;
319        let opts_owned = opts.clone();
320        let pack_clone = Arc::clone(&pack);
321        // Best-effort cancel if a wall-clock timeout fires mid-inference.
322        let cancel_on_timeout = opts.cancel.clone();
323
324        let join = tokio::task::spawn_blocking(move || {
325            synthesize_with_pack(
326                pack_clone.as_ref(),
327                &text_owned,
328                &opts_owned,
329                text_chars,
330                text_truncated,
331            )
332        });
333
334        match tokio::time::timeout(timeout, join).await {
335            Ok(Ok(result)) => result,
336            Ok(Err(e)) => Err(crate::error::TranscriptionError::internal(format!(
337                "TTS synth join: {e}"
338            ))),
339            Err(_elapsed) => {
340                if let Some(flag) = cancel_on_timeout {
341                    flag.cancel();
342                }
343                Err(ProviderError::Other {
344                    message: format!(
345                        "TTS synthesis exceeded timeout ({} ms)",
346                        timeout.as_millis()
347                    ),
348                }
349                .into())
350            }
351        }
352    }
353
354    async fn preload(&self, model: &str, voice: &str) -> Result<()> {
355        lookup_voice(voice)?;
356        let _ = self.ensure_loaded(model, self.local_only).await?;
357        Ok(())
358    }
359}
360
361/// Convenience: synthesize with a one-shot provider.
362pub async fn synthesize_local(
363    cache_dir: impl Into<PathBuf>,
364    text: &str,
365    opts: &SynthesisOptions,
366) -> Result<SynthesisResult> {
367    let provider = LocalTtsProvider::new(cache_dir.into()).with_progress(false);
368    provider.synthesize(text, opts).await
369}
370
371#[cfg(test)]
372mod tests {
373    use super::*;
374
375    #[tokio::test]
376    async fn empty_text_user_error() {
377        let dir = tempfile::tempdir().unwrap();
378        let p = LocalTtsProvider::new(dir.path().to_path_buf()).with_local_only(true);
379        let err = p
380            .synthesize("  ", &SynthesisOptions::default())
381            .await
382            .unwrap_err();
383        assert_eq!(err.exit_code(), 2);
384    }
385
386    #[tokio::test]
387    async fn missing_pack_local_only() {
388        let dir = tempfile::tempdir().unwrap();
389        let p = LocalTtsProvider::new(dir.path().to_path_buf()).with_local_only(true);
390        let err = p
391            .synthesize("Hello", &SynthesisOptions::default())
392            .await
393            .unwrap_err();
394        assert!(matches!(err.exit_code(), 2 | 4));
395    }
396
397    #[tokio::test]
398    async fn unsupported_language_user_error() {
399        let dir = tempfile::tempdir().unwrap();
400        let p = LocalTtsProvider::new(dir.path().to_path_buf()).with_local_only(true);
401        let opts = SynthesisOptions {
402            language: "fr".into(),
403            ..Default::default()
404        };
405        let err = p.synthesize("Hello", &opts).await.unwrap_err();
406        assert_eq!(err.exit_code(), 2);
407        assert!(err.to_string().contains("unsupported TTS language"));
408    }
409
410    #[test]
411    fn clear_sessions_is_safe_when_empty() {
412        let dir = tempfile::tempdir().unwrap();
413        let p = LocalTtsProvider::new(dir.path().to_path_buf()).with_local_only(true);
414        p.clear_sessions(); // must not panic
415        p.clear_sessions();
416    }
417}