use crate::error::TalkError;
use crate::model_fetch::ModelSpec;
use std::path::Path;
const KOKORO_TARBALL_URL: &str = "https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/kokoro-multi-lang-v1_0.tar.bz2";
const KOKORO_TARBALL_INNER_DIR: &str = "kokoro-multi-lang-v1_0";
const MODEL_FILES: &[&str] = &["model.onnx", "voices.bin", "tokens.txt"];
pub(crate) const MODEL_EN: &str = "model.onnx";
pub(crate) const EMPTY_LEXICON: &str = "empty_lexicon.txt";
pub(super) const KOKORO_SPEC: ModelSpec = ModelSpec {
display_name: "Kokoro",
tarball_url: KOKORO_TARBALL_URL,
inner_dir: KOKORO_TARBALL_INNER_DIR,
required_files: MODEL_FILES,
approx_size: "~350 MB",
manual_files_hint: "model.onnx/voices.bin/tokens.txt (plus espeak-ng-data/ and dict/)",
};
pub(crate) fn model_filename_for_lang(lang: &str) -> String {
format!("model-{}.onnx", lang)
}
pub(crate) fn is_present(dir: &Path) -> bool {
crate::model_fetch::is_present(dir, &KOKORO_SPEC)
}
pub(crate) fn ensure_present(dir: &Path) -> Result<(), TalkError> {
crate::model_fetch::ensure_present(dir, &KOKORO_SPEC)
}
pub(crate) async fn ensure_with_cli_consent(dir: &Path) -> Result<(), TalkError> {
crate::model_fetch::ensure_with_cli_consent(dir, &KOKORO_SPEC).await
}
pub(crate) fn baked_language(dir: &Path) -> Result<String, TalkError> {
let en_model = dir.join(MODEL_EN);
let bytes = std::fs::read(&en_model).map_err(|e| {
TalkError::Config(format!(
"failed to read Kokoro model {}: {}",
en_model.display(),
e
))
})?;
read_voice_metadata(&bytes)
}
pub(crate) fn ensure_lang_model(dir: &Path, lang: &str) -> Result<std::path::PathBuf, TalkError> {
let baked = baked_language(dir)?;
if lang == baked || baked.split('-').next() == Some(lang) {
return Ok(dir.join(MODEL_EN));
}
let empty_lexicon = dir.join(EMPTY_LEXICON);
if !empty_lexicon.exists() {
std::fs::write(&empty_lexicon, b"").map_err(|e| {
TalkError::Config(format!(
"failed to create {}: {}",
empty_lexicon.display(),
e
))
})?;
}
let derived = dir.join(model_filename_for_lang(lang));
if derived.exists() {
return Ok(derived);
}
let en_model = dir.join(MODEL_EN);
let bytes = std::fs::read(&en_model).map_err(|e| {
TalkError::Config(format!(
"failed to read stock Kokoro model {} for language patch: {}",
en_model.display(),
e
))
})?;
let patched = patch_voice_metadata(&bytes, lang)?;
let tmp = dir.join(format!("model-{}.onnx.tmp", lang));
std::fs::write(&tmp, &patched)
.map_err(|e| TalkError::Config(format!("failed to write {}: {}", tmp.display(), e)))?;
std::fs::rename(&tmp, &derived).map_err(|e| {
let _ = std::fs::remove_file(&tmp);
TalkError::Config(format!(
"failed to promote {} -> {}: {}",
tmp.display(),
derived.display(),
e
))
})?;
log::info!(
"kokoro: generated {} for language '{}'",
derived.display(),
lang
);
Ok(derived)
}
const VOICE_KEY_MARKER: &[u8] = b"\x0a\x05voice\x12";
const METADATA_ENTRY_TAG: u8 = 0x72;
fn locate_voice_entry(bytes: &[u8]) -> Result<(usize, usize, usize), TalkError> {
let mut hits = bytes
.windows(VOICE_KEY_MARKER.len())
.enumerate()
.filter(|(_, w)| *w == VOICE_KEY_MARKER)
.map(|(i, _)| i);
let marker = hits.next().ok_or_else(|| {
TalkError::Config(
"kokoro language patch: could not locate the `voice` ONNX metadata \
entry (unexpected model layout). Provide a hand-patched \
model-<lang>.onnx (voice metadata = the language code) next to \
model.onnx."
.to_string(),
)
})?;
if hits.next().is_some() {
return Err(TalkError::Config(
"kokoro language patch: the `voice` ONNX metadata entry appears \
more than once; refusing to patch ambiguously."
.to_string(),
));
}
let value_len_idx = marker + VOICE_KEY_MARKER.len();
let value_len_byte = *bytes.get(value_len_idx).ok_or_else(|| {
TalkError::Config("kokoro language patch: truncated voice metadata".to_string())
})?;
if value_len_byte >= 0x80 {
return Err(TalkError::Config(
"kokoro language patch: multi-byte value length not supported".to_string(),
));
}
let value_len = value_len_byte as usize;
let value_start = value_len_idx + 1;
if marker < 2 {
return Err(TalkError::Config(
"kokoro language patch: voice entry too close to start of file".to_string(),
));
}
let entry_tag_idx = marker - 2;
if bytes[entry_tag_idx] != METADATA_ENTRY_TAG {
return Err(TalkError::Config(
"kokoro language patch: unexpected metadata entry framing".to_string(),
));
}
if value_start + value_len > bytes.len() {
return Err(TalkError::Config(
"kokoro language patch: voice value runs past end of file".to_string(),
));
}
Ok((entry_tag_idx, value_start, value_len))
}
pub(crate) fn read_voice_metadata(bytes: &[u8]) -> Result<String, TalkError> {
let (_, value_start, value_len) = locate_voice_entry(bytes)?;
let raw = &bytes[value_start..value_start + value_len];
String::from_utf8(raw.to_vec()).map_err(|_| {
TalkError::Config("kokoro language patch: voice metadata is not UTF-8".to_string())
})
}
pub(crate) fn patch_voice_metadata(bytes: &[u8], lang: &str) -> Result<Vec<u8>, TalkError> {
if lang.is_empty() {
return Err(TalkError::Config(
"kokoro language patch: empty language code".to_string(),
));
}
if lang.len() >= 0x80 {
return Err(TalkError::Config(
"kokoro language patch: language code too long".to_string(),
));
}
let (entry_tag_idx, value_start, old_value_len) = locate_voice_entry(bytes)?;
let entry_len_idx = entry_tag_idx + 1;
let old_entry_len = bytes[entry_len_idx];
if old_entry_len >= 0x80 {
return Err(TalkError::Config(
"kokoro language patch: multi-byte entry length not supported".to_string(),
));
}
let new_value_len = lang.len();
let new_entry_len_i = old_entry_len as isize - old_value_len as isize + new_value_len as isize;
if !(0..0x80).contains(&new_entry_len_i) {
return Err(TalkError::Config(
"kokoro language patch: patched entry length out of single-byte range".to_string(),
));
}
let mut out = Vec::with_capacity(bytes.len() + new_value_len);
out.extend_from_slice(&bytes[..=entry_tag_idx]);
out.push(new_entry_len_i as u8);
out.extend_from_slice(VOICE_KEY_MARKER);
out.push(new_value_len as u8);
out.extend_from_slice(lang.as_bytes());
out.extend_from_slice(&bytes[value_start + old_value_len..]);
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
const VOICE_EN_ENTRY: &[u8] = b"\x72\x0e\x0a\x05voice\x12\x05en-us";
fn synthetic_model() -> Vec<u8> {
let mut v = b"onnx-header-noise".to_vec();
v.extend_from_slice(VOICE_EN_ENTRY);
v.extend_from_slice(b"onnx-trailing-noise");
v
}
#[test]
fn read_voice_metadata_returns_baked_language() {
let m = synthetic_model();
assert_eq!(read_voice_metadata(&m).unwrap(), "en-us");
}
#[test]
fn patch_voice_metadata_to_fr() {
let m = synthetic_model();
let patched = patch_voice_metadata(&m, "fr").expect("patch");
assert_eq!(patched.len(), m.len() - 3);
assert_eq!(read_voice_metadata(&patched).unwrap(), "fr");
assert!(patched.starts_with(b"onnx-header-noise"));
assert!(patched.ends_with(b"onnx-trailing-noise"));
}
#[test]
fn patch_voice_metadata_to_longer_code() {
let m = synthetic_model();
let patched = patch_voice_metadata(&m, "de-de").expect("patch");
assert_eq!(patched.len(), m.len()); assert_eq!(read_voice_metadata(&patched).unwrap(), "de-de");
}
#[test]
fn patch_voice_metadata_to_short_and_back_roundtrips_framing() {
let m = synthetic_model();
let de = patch_voice_metadata(&m, "de").expect("patch de");
let it = patch_voice_metadata(&de, "it").expect("patch it");
assert_eq!(read_voice_metadata(&it).unwrap(), "it");
}
#[test]
fn patch_voice_metadata_errors_when_absent() {
let buf = b"no voice metadata entry here".to_vec();
let err = patch_voice_metadata(&buf, "fr").expect_err("must error");
assert!(err.to_string().contains("could not locate"));
}
#[test]
fn patch_voice_metadata_errors_when_ambiguous() {
let mut buf = Vec::new();
buf.extend_from_slice(VOICE_EN_ENTRY);
buf.extend_from_slice(b"----");
buf.extend_from_slice(VOICE_EN_ENTRY);
let err = patch_voice_metadata(&buf, "fr").expect_err("must error");
assert!(err.to_string().contains("more than once"));
}
#[test]
fn patch_voice_metadata_rejects_empty_lang() {
let m = synthetic_model();
let err = patch_voice_metadata(&m, "").expect_err("must error");
assert!(err.to_string().contains("empty language code"));
}
#[test]
fn ensure_lang_model_stock_for_baked_language() {
let tmp = tempfile::TempDir::new().unwrap();
let dir = tmp.path();
std::fs::write(dir.join(MODEL_EN), synthetic_model()).unwrap();
std::fs::write(dir.join("voices.bin"), b"x").unwrap();
std::fs::write(dir.join("tokens.txt"), b"x").unwrap();
let p = ensure_lang_model(dir, "en").expect("resolve en");
assert_eq!(p, dir.join(MODEL_EN));
let p2 = ensure_lang_model(dir, "en-us").expect("resolve en-us");
assert_eq!(p2, dir.join(MODEL_EN));
}
#[test]
fn ensure_lang_model_derives_and_caches() {
let tmp = tempfile::TempDir::new().unwrap();
let dir = tmp.path();
std::fs::write(dir.join(MODEL_EN), synthetic_model()).unwrap();
let p = ensure_lang_model(dir, "fr").expect("derive fr");
assert_eq!(p, dir.join("model-fr.onnx"));
assert!(dir.join(EMPTY_LEXICON).exists());
assert_eq!(
read_voice_metadata(&std::fs::read(&p).unwrap()).unwrap(),
"fr"
);
let p2 = ensure_lang_model(dir, "fr").expect("cached fr");
assert_eq!(p, p2);
}
#[test]
fn kokoro_spec_has_expected_url_and_inner_dir() {
assert!(KOKORO_SPEC
.tarball_url
.ends_with("kokoro-multi-lang-v1_0.tar.bz2"));
assert_eq!(KOKORO_SPEC.inner_dir, "kokoro-multi-lang-v1_0");
}
}