use crate::error::TalkError;
use fs2::FileExt;
use futures::StreamExt;
use std::fs::OpenOptions;
use std::io::{IsTerminal, Write};
use std::path::{Path, PathBuf};
const PROGRESS_LOG_EVERY_BYTES: u64 = 32 * 1024 * 1024;
#[derive(Debug, Clone)]
pub struct ModelSpec {
pub display_name: &'static str,
pub tarball_url: &'static str,
pub inner_dir: &'static str,
pub required_files: &'static [&'static str],
pub approx_size: &'static str,
pub manual_files_hint: &'static str,
}
impl ModelSpec {
fn manual_fallback_message(&self, cause: &str, dir: &Path) -> String {
format!(
"Failed to download {} model ({}). To install manually: \
download {}, extract it, and place {} into {}.",
self.display_name,
cause,
self.tarball_url,
self.manual_files_hint,
dir.display()
)
}
}
pub fn is_present(dir: &Path, spec: &ModelSpec) -> bool {
spec.required_files.iter().all(|name| {
let p = dir.join(name);
match std::fs::metadata(&p) {
Ok(m) => m.is_file() && m.len() > 0,
Err(_) => false,
}
})
}
pub fn ensure_present(dir: &Path, spec: &ModelSpec) -> Result<(), TalkError> {
if is_present(dir, spec) {
return Ok(());
}
Err(TalkError::Config(format!(
"{} model not found in {}. The model is not downloaded \
automatically; it must be fetched with explicit consent. \
To install manually: download {}, extract it, and place \
{} into {}.",
spec.display_name,
dir.display(),
spec.tarball_url,
spec.manual_files_hint,
dir.display()
)))
}
struct LockGuard<'a>(&'a std::fs::File);
impl Drop for LockGuard<'_> {
fn drop(&mut self) {
let _ = fs2::FileExt::unlock(self.0);
}
}
pub async fn download_and_install(dir: &Path, spec: &ModelSpec) -> Result<(), TalkError> {
if is_present(dir, spec) {
return Ok(());
}
let parent = dir.parent().ok_or_else(|| {
TalkError::Config(format!(
"{} model_dir has no parent directory: {}",
spec.display_name,
dir.display()
))
})?;
std::fs::create_dir_all(parent).map_err(|e| {
TalkError::Config(format!(
"Failed to create {} model parent dir {}: {}",
spec.display_name,
parent.display(),
e
))
})?;
let lock_path = sibling_lock_path(dir);
let lock_file = OpenOptions::new()
.create(true)
.read(true)
.write(true)
.truncate(false)
.open(&lock_path)
.map_err(|e| {
TalkError::Config(format!(
"Failed to open {} model lock file {}: {}",
spec.display_name,
lock_path.display(),
e
))
})?;
lock_file.lock_exclusive().map_err(|e| {
TalkError::Config(format!(
"Failed to acquire exclusive flock on {}: {}",
lock_path.display(),
e
))
})?;
let _lock_guard = LockGuard(&lock_file);
if is_present(dir, spec) {
return Ok(());
}
let tarball_tmp = parent.join(format!(
"{}.download.tmp",
dir.file_name()
.and_then(|n| n.to_str())
.unwrap_or("talk-rs-model")
));
let _ = std::fs::remove_file(&tarball_tmp);
if let Err(e) = download_to(spec.tarball_url, &tarball_tmp).await {
let _ = std::fs::remove_file(&tarball_tmp);
return Err(TalkError::Config(
spec.manual_fallback_message(&e.to_string(), dir),
));
}
let res = install_from_tarball(&tarball_tmp, dir, spec);
let _ = std::fs::remove_file(&tarball_tmp);
if let Err(e) = res {
return Err(TalkError::Config(
spec.manual_fallback_message(&e.to_string(), dir),
));
}
Ok(())
}
fn sibling_lock_path(dir: &Path) -> PathBuf {
let mut p = dir.to_path_buf();
let name = p
.file_name()
.and_then(|n| n.to_str())
.unwrap_or("talk-rs-model");
let lock_name = format!("{}.lock", name);
p.set_file_name(lock_name);
p
}
async fn download_to(url: &str, dest: &Path) -> Result<(), TalkError> {
log::info!("model fetch: downloading {} -> {}", url, dest.display());
let response = reqwest::get(url)
.await
.map_err(|e| TalkError::Transcription(format!("HTTP GET failed: {}", e)))?;
let status = response.status();
if !status.is_success() {
return Err(TalkError::Transcription(format!(
"HTTP {} from {}",
status, url
)));
}
let total_hint = response.content_length();
let mut out = std::fs::File::create(dest).map_err(|e| {
TalkError::Transcription(format!("create temp tarball {}: {}", dest.display(), e))
})?;
let mut stream = response.bytes_stream();
let mut written: u64 = 0;
let mut next_log_at: u64 = PROGRESS_LOG_EVERY_BYTES;
while let Some(chunk) = stream.next().await {
let chunk =
chunk.map_err(|e| TalkError::Transcription(format!("HTTP stream error: {}", e)))?;
out.write_all(&chunk).map_err(|e| {
TalkError::Transcription(format!("write to temp tarball {}: {}", dest.display(), e))
})?;
written += chunk.len() as u64;
if written >= next_log_at {
match total_hint {
Some(total) => log::info!("model fetch: downloaded {} / {} bytes", written, total),
None => log::info!("model fetch: downloaded {} bytes", written),
}
next_log_at = next_log_at.saturating_add(PROGRESS_LOG_EVERY_BYTES);
}
}
out.flush().map_err(|e| {
TalkError::Transcription(format!("flush temp tarball {}: {}", dest.display(), e))
})?;
drop(out);
log::info!("model fetch: download complete ({} bytes)", written);
Ok(())
}
pub fn install_from_tarball(tarball: &Path, dir: &Path, spec: &ModelSpec) -> Result<(), TalkError> {
let parent = dir.parent().ok_or_else(|| {
TalkError::Transcription(format!("install target {} has no parent", dir.display()))
})?;
std::fs::create_dir_all(parent).map_err(|e| {
TalkError::Transcription(format!("create parent {}: {}", parent.display(), e))
})?;
let staging = parent.join(format!(
"{}.staging.tmp",
dir.file_name()
.and_then(|n| n.to_str())
.unwrap_or("talk-rs-model")
));
let _ = std::fs::remove_dir_all(&staging);
std::fs::create_dir_all(&staging).map_err(|e| {
TalkError::Transcription(format!("create staging dir {}: {}", staging.display(), e))
})?;
let cleanup_staging = |s: &Path| {
let _ = std::fs::remove_dir_all(s);
};
let tarball_file = std::fs::File::open(tarball).map_err(|e| {
cleanup_staging(&staging);
TalkError::Transcription(format!("open tarball {}: {}", tarball.display(), e))
})?;
let decoder = bzip2::read::BzDecoder::new(tarball_file);
let mut archive = tar::Archive::new(decoder);
if let Err(e) = archive.unpack(&staging) {
cleanup_staging(&staging);
return Err(TalkError::Transcription(format!(
"extract tarball into {}: {}",
staging.display(),
e
)));
}
let inner_dir = locate_inner_dir(&staging, spec)?;
for name in spec.required_files {
let p = inner_dir.join(name);
let md = std::fs::metadata(&p).map_err(|e| {
cleanup_staging(&staging);
TalkError::Transcription(format!("extracted archive missing {}: {}", p.display(), e))
})?;
if !md.is_file() || md.len() == 0 {
cleanup_staging(&staging);
return Err(TalkError::Transcription(format!(
"extracted archive has empty or non-file entry: {}",
p.display()
)));
}
}
if let Err(e) = std::fs::remove_dir_all(dir) {
if e.kind() != std::io::ErrorKind::NotFound {
cleanup_staging(&staging);
return Err(TalkError::Transcription(format!(
"remove pre-existing model dir {}: {}",
dir.display(),
e
)));
}
}
if let Err(e) = std::fs::rename(&inner_dir, dir) {
cleanup_staging(&staging);
return Err(TalkError::Transcription(format!(
"promote {} -> {}: {}",
inner_dir.display(),
dir.display(),
e
)));
}
cleanup_staging(&staging);
Ok(())
}
fn locate_inner_dir(staging: &Path, spec: &ModelSpec) -> Result<PathBuf, TalkError> {
let preferred = staging.join(spec.inner_dir);
if preferred.is_dir() {
return Ok(preferred);
}
let mut candidates: Vec<PathBuf> = Vec::new();
let entries = std::fs::read_dir(staging).map_err(|e| {
TalkError::Transcription(format!("read staging dir {}: {}", staging.display(), e))
})?;
for entry in entries.flatten() {
let path = entry.path();
let name = match path.file_name().and_then(|n| n.to_str()) {
Some(n) => n,
None => continue,
};
if name.starts_with('.') {
continue;
}
if path.is_dir() {
candidates.push(path);
}
}
if candidates.len() == 1 {
return Ok(candidates.remove(0));
}
Err(TalkError::Transcription(format!(
"could not locate {} model dir inside extracted archive at {} \
(expected `{}/` or a single top-level dir, found {} dirs)",
spec.display_name,
staging.display(),
spec.inner_dir,
candidates.len()
)))
}
pub async fn ensure_with_cli_consent(dir: &Path, spec: &ModelSpec) -> Result<(), TalkError> {
if is_present(dir, spec) {
return Ok(());
}
if std::io::stdin().is_terminal() {
eprint!(
"The {} model ({}) is not installed at {}.\n\
Download it now? [y/N] ",
spec.display_name,
spec.approx_size,
dir.display()
);
let _ = std::io::stderr().flush();
let mut answer = String::new();
std::io::stdin()
.read_line(&mut answer)
.map_err(TalkError::Io)?;
let answer = answer.trim().to_ascii_lowercase();
if answer != "y" && answer != "yes" {
return Err(TalkError::Config(format!(
"{} model download declined; nothing was downloaded. \
Install it later by re-running and accepting the prompt, \
or place the files manually in {}.",
spec.display_name,
dir.display()
)));
}
eprintln!(
"Downloading {} model to {} …",
spec.display_name,
dir.display()
);
} else {
eprintln!(
"{} model not found at {}; downloading ({}) — selecting this \
local backend implies consent. Choose a different provider to \
avoid this.",
spec.display_name,
dir.display(),
spec.approx_size,
);
}
download_and_install(dir, spec).await
}
#[cfg(test)]
mod tests {
use super::*;
use bzip2::Compression;
use tempfile::TempDir;
const TEST_SPEC: ModelSpec = ModelSpec {
display_name: "Test",
tarball_url: "https://example.com/test-model.tar.bz2",
inner_dir: "test-model-inner",
required_files: &["a.onnx", "b.onnx", "tokens.txt"],
approx_size: "~1 MB",
manual_files_hint: "a.onnx/b.onnx/tokens.txt",
};
fn make_synthetic_tarball(
path: &Path,
entries: &[(&str, &[u8])],
inner: &str,
) -> Result<(), TalkError> {
let file = std::fs::File::create(path).map_err(|e| {
TalkError::Transcription(format!("create fixture {}: {}", path.display(), e))
})?;
let encoder = bzip2::write::BzEncoder::new(file, Compression::fast());
let mut builder = tar::Builder::new(encoder);
for (name, bytes) in entries {
let mut header = tar::Header::new_gnu();
header.set_size(bytes.len() as u64);
header.set_mode(0o644);
header.set_cksum();
let entry_path = format!("{}/{}", inner, name);
builder
.append_data(&mut header, &entry_path, *bytes)
.map_err(|e| TalkError::Transcription(format!("append {}: {}", entry_path, e)))?;
}
let encoder = builder
.into_inner()
.map_err(|e| TalkError::Transcription(format!("close tar: {}", e)))?;
encoder
.finish()
.map_err(|e| TalkError::Transcription(format!("close bz2: {}", e)))?;
Ok(())
}
fn populate_dir(dir: &Path, spec: &ModelSpec) {
std::fs::create_dir_all(dir).unwrap();
for name in spec.required_files {
let mut f = std::fs::File::create(dir.join(name)).unwrap();
f.write_all(b"dummy-bytes").unwrap();
}
}
#[test]
fn is_present_empty_dir_is_false() {
let tmp = TempDir::new().unwrap();
assert!(!is_present(tmp.path(), &TEST_SPEC));
}
#[test]
fn is_present_complete_dir_is_true() {
let tmp = TempDir::new().unwrap();
populate_dir(tmp.path(), &TEST_SPEC);
assert!(is_present(tmp.path(), &TEST_SPEC));
}
#[test]
fn is_present_zero_byte_file_is_false() {
let tmp = TempDir::new().unwrap();
populate_dir(tmp.path(), &TEST_SPEC);
std::fs::File::create(tmp.path().join("tokens.txt")).unwrap();
assert!(!is_present(tmp.path(), &TEST_SPEC));
}
#[test]
fn ensure_present_errors_without_download_when_absent() {
let tmp = TempDir::new().unwrap();
let dir = tmp.path().join("m");
let err = ensure_present(&dir, &TEST_SPEC).expect_err("absent model must error");
let msg = err.to_string();
assert!(msg.contains("not downloaded automatically") || msg.contains("explicit consent"));
assert!(msg.contains(TEST_SPEC.tarball_url));
assert!(!dir.exists());
}
#[test]
fn install_from_tarball_promotes_and_flattens() {
let tmp = TempDir::new().unwrap();
let tarball = tmp.path().join("model.tar.bz2");
make_synthetic_tarball(
&tarball,
&[
("a.onnx", b"A"),
("b.onnx", b"B"),
("tokens.txt", b"T"),
("README.md", b"R"),
],
TEST_SPEC.inner_dir,
)
.unwrap();
let dir = tmp.path().join("model");
install_from_tarball(&tarball, &dir, &TEST_SPEC).expect("install must succeed");
assert!(is_present(&dir, &TEST_SPEC));
assert!(!dir.join(TEST_SPEC.inner_dir).exists(), "nesting flattened");
assert!(!tmp.path().join("model.staging.tmp").exists());
}
#[test]
fn install_from_tarball_missing_file_leaves_dir_absent() {
let tmp = TempDir::new().unwrap();
let tarball = tmp.path().join("model.tar.bz2");
make_synthetic_tarball(
&tarball,
&[("a.onnx", b"A"), ("tokens.txt", b"T")],
TEST_SPEC.inner_dir,
)
.unwrap();
let dir = tmp.path().join("model");
let err =
install_from_tarball(&tarball, &dir, &TEST_SPEC).expect_err("missing file must fail");
assert!(err.to_string().contains("b.onnx"));
assert!(!dir.exists(), "failed install must not leave a dir");
}
#[test]
fn install_from_tarball_single_subdir_fallback() {
let tmp = TempDir::new().unwrap();
let tarball = tmp.path().join("model.tar.bz2");
make_synthetic_tarball(
&tarball,
&[("a.onnx", b"A"), ("b.onnx", b"B"), ("tokens.txt", b"T")],
"some-other-name",
)
.unwrap();
let dir = tmp.path().join("model");
install_from_tarball(&tarball, &dir, &TEST_SPEC).expect("fallback install must succeed");
assert!(is_present(&dir, &TEST_SPEC));
}
#[test]
fn sibling_lock_path_is_sibling_not_child() {
let dir = Path::new("/tmp/some/kokoro-model");
let lock = sibling_lock_path(dir);
assert_eq!(lock, Path::new("/tmp/some/kokoro-model.lock"));
}
#[tokio::test]
async fn download_and_install_fast_path_when_present() {
let tmp = TempDir::new().unwrap();
let dir = tmp.path().join("m");
populate_dir(&dir, &TEST_SPEC);
let start = std::time::Instant::now();
download_and_install(&dir, &TEST_SPEC)
.await
.expect("fast path must succeed");
assert!(start.elapsed() < std::time::Duration::from_millis(100));
}
#[test]
fn manual_fallback_message_contains_url_and_dir() {
let msg = TEST_SPEC.manual_fallback_message("boom", Path::new("/some/where/m"));
assert!(msg.contains(TEST_SPEC.tarball_url));
assert!(msg.contains("/some/where/m"));
assert!(msg.contains("a.onnx"));
assert!(msg.contains("boom"));
}
}