1use 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
20const TAIL_TRIM: usize = 2_000;
22const PEAK_LIMIT: f32 = 0.95;
24
25pub struct LocalTtsProvider {
27 cache_dir: PathBuf,
28 show_progress: bool,
29 local_only: bool,
30 max_chars: usize,
31 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 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 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 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 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 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 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 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
361pub 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(); p.clear_sessions();
416 }
417}