wfloat 2.0.0__py3-none-macosx_12_0_arm64.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
wfloat/__init__.py ADDED
@@ -0,0 +1,57 @@
1
+ from ._constants import SPEAKER_IDS, VALID_EMOTIONS, VALID_SIDS
2
+ from ._llm import LlmModel
3
+ from ._llm_load import load_llm_model
4
+ from ._model import Model, TtsModel, load, load_tts_model
5
+ from ._stt import SttModel, SttSession
6
+ from ._stt_load import load_moonshine_tiny_en, load_stt_model, load_whisper_tiny_en
7
+ from ._vad import VadModel
8
+ from ._vad_load import load_silero_vad, load_vad_model
9
+ from ._results import (
10
+ Audio,
11
+ AudioResult,
12
+ GenerationResult,
13
+ LlmGenerationResult,
14
+ StreamingTranscriptionResult,
15
+ TranscriptionResult,
16
+ TranscriptionSegment,
17
+ TranscriptionToken,
18
+ Timeline,
19
+ TimelineChunk,
20
+ TtsSynthesisResult,
21
+ VadDetectionResult,
22
+ VadSegment,
23
+ )
24
+ from ._version import __version__
25
+
26
+ __all__ = [
27
+ "Audio",
28
+ "AudioResult",
29
+ "GenerationResult",
30
+ "LlmGenerationResult",
31
+ "LlmModel",
32
+ "Model",
33
+ "SPEAKER_IDS",
34
+ "SttModel",
35
+ "SttSession",
36
+ "StreamingTranscriptionResult",
37
+ "TtsModel",
38
+ "Timeline",
39
+ "TimelineChunk",
40
+ "TranscriptionResult",
41
+ "TranscriptionSegment",
42
+ "TranscriptionToken",
43
+ "TtsSynthesisResult",
44
+ "VALID_EMOTIONS",
45
+ "VALID_SIDS",
46
+ "VadDetectionResult",
47
+ "VadModel",
48
+ "VadSegment",
49
+ "load",
50
+ "load_llm_model",
51
+ "load_moonshine_tiny_en",
52
+ "load_silero_vad",
53
+ "load_stt_model",
54
+ "load_whisper_tiny_en",
55
+ "load_tts_model",
56
+ "load_vad_model",
57
+ ]
wfloat/__main__.py ADDED
@@ -0,0 +1,5 @@
1
+ from ._cli import main
2
+
3
+
4
+ if __name__ == "__main__": # pragma: no cover
5
+ raise SystemExit(main())
wfloat/_assets.py ADDED
@@ -0,0 +1,421 @@
1
+ from dataclasses import dataclass
2
+ from pathlib import Path
3
+ from typing import Dict, Mapping, Optional
4
+ from urllib.parse import urlparse
5
+
6
+ from ._generated_model_urls import MODEL_ASSETS, REGISTRY_ORIGIN, SHARED_ASSETS
7
+
8
+ REGISTRY_BASE_URL = REGISTRY_ORIGIN
9
+ WFLOAT_TTS_MODEL_ID = "wfloat/wfloat-tts"
10
+ SILERO_VAD_MODEL_ID = "snakers4/silero-vad"
11
+ SMOLLM2_360M_INSTRUCT_MODEL_ID = "HuggingFaceTB/SmolLM2-360M-Instruct"
12
+
13
+
14
+ @dataclass(frozen=True)
15
+ class ModelAssets:
16
+ model_onnx: str
17
+ model_onnx_checksum: str
18
+ model_tokens: str
19
+ model_tokens_checksum: str
20
+ espeak_data: str
21
+ espeak_checksum: str
22
+
23
+ @classmethod
24
+ def from_dict(cls, data: Dict[str, object]) -> "ModelAssets":
25
+ required_fields = (
26
+ "model_onnx",
27
+ "model_onnx_checksum",
28
+ "model_tokens",
29
+ "model_tokens_checksum",
30
+ "espeak_data",
31
+ "espeak_checksum",
32
+ )
33
+
34
+ missing = [
35
+ field_name
36
+ for field_name in required_fields
37
+ if not isinstance(data.get(field_name), str) or not str(data.get(field_name)).strip()
38
+ ]
39
+ if missing:
40
+ raise ValueError(
41
+ "Model asset response is missing required fields: %s"
42
+ % ", ".join(missing)
43
+ )
44
+
45
+ return cls(
46
+ model_onnx=str(data["model_onnx"]),
47
+ model_onnx_checksum=str(data["model_onnx_checksum"]),
48
+ model_tokens=str(data["model_tokens"]),
49
+ model_tokens_checksum=str(data["model_tokens_checksum"]),
50
+ espeak_data=str(data["espeak_data"]),
51
+ espeak_checksum=str(data["espeak_checksum"]),
52
+ )
53
+
54
+ def to_dict(self) -> Dict[str, str]:
55
+ return {
56
+ "model_onnx": self.model_onnx,
57
+ "model_onnx_checksum": self.model_onnx_checksum,
58
+ "model_tokens": self.model_tokens,
59
+ "model_tokens_checksum": self.model_tokens_checksum,
60
+ "espeak_data": self.espeak_data,
61
+ "espeak_checksum": self.espeak_checksum,
62
+ }
63
+
64
+
65
+ @dataclass(frozen=True)
66
+ class SttModelAssets:
67
+ family: str
68
+ tokens: str
69
+ tokens_checksum: Optional[str] = None
70
+ model: Optional[str] = None
71
+ model_checksum: Optional[str] = None
72
+ preprocessor: Optional[str] = None
73
+ preprocessor_checksum: Optional[str] = None
74
+ encoder: Optional[str] = None
75
+ encoder_checksum: Optional[str] = None
76
+ decoder: Optional[str] = None
77
+ decoder_checksum: Optional[str] = None
78
+ joiner: Optional[str] = None
79
+ joiner_checksum: Optional[str] = None
80
+ uncached_decoder: Optional[str] = None
81
+ uncached_decoder_checksum: Optional[str] = None
82
+ cached_decoder: Optional[str] = None
83
+ cached_decoder_checksum: Optional[str] = None
84
+
85
+ @classmethod
86
+ def from_dict(cls, data: Dict[str, object]) -> "SttModelAssets":
87
+ family = str(data.get("family") or "").strip()
88
+ if not family:
89
+ raise ValueError("STT asset response is missing required field: family")
90
+
91
+ files = data.get("files")
92
+ if isinstance(files, Mapping):
93
+ normalized: Dict[str, str] = {}
94
+ for key, value in files.items():
95
+ if not isinstance(key, str) or not isinstance(value, Mapping):
96
+ continue
97
+ url = value.get("url")
98
+ checksum = value.get("checksum")
99
+ if isinstance(url, str) and url.strip():
100
+ normalized[key] = url
101
+ if isinstance(checksum, str) and checksum.strip():
102
+ normalized[f"{key}_checksum"] = checksum
103
+ merged = dict(data)
104
+ merged.update(normalized)
105
+ data = merged
106
+
107
+ required_fields = ("tokens",)
108
+ missing = [
109
+ field_name
110
+ for field_name in required_fields
111
+ if not isinstance(data.get(field_name), str)
112
+ or not str(data.get(field_name)).strip()
113
+ ]
114
+ if missing:
115
+ raise ValueError(
116
+ "STT asset response is missing required fields: %s"
117
+ % ", ".join(missing)
118
+ )
119
+
120
+ def optional_string(key: str) -> Optional[str]:
121
+ value = data.get(key)
122
+ if isinstance(value, str) and value.strip():
123
+ return value.strip()
124
+ return None
125
+
126
+ return cls(
127
+ family=family,
128
+ tokens=str(data["tokens"]).strip(),
129
+ tokens_checksum=optional_string("tokens_checksum"),
130
+ model=optional_string("model"),
131
+ model_checksum=optional_string("model_checksum"),
132
+ preprocessor=optional_string("preprocessor"),
133
+ preprocessor_checksum=optional_string("preprocessor_checksum"),
134
+ encoder=optional_string("encoder"),
135
+ encoder_checksum=optional_string("encoder_checksum"),
136
+ decoder=optional_string("decoder"),
137
+ decoder_checksum=optional_string("decoder_checksum"),
138
+ joiner=optional_string("joiner"),
139
+ joiner_checksum=optional_string("joiner_checksum"),
140
+ uncached_decoder=optional_string("uncached_decoder"),
141
+ uncached_decoder_checksum=optional_string("uncached_decoder_checksum"),
142
+ cached_decoder=optional_string("cached_decoder"),
143
+ cached_decoder_checksum=optional_string("cached_decoder_checksum"),
144
+ )
145
+
146
+ def to_dict(self) -> Dict[str, str]:
147
+ data = {
148
+ "family": self.family,
149
+ "tokens": self.tokens,
150
+ }
151
+ optional_fields = {
152
+ "tokens_checksum": self.tokens_checksum,
153
+ "model": self.model,
154
+ "model_checksum": self.model_checksum,
155
+ "preprocessor": self.preprocessor,
156
+ "preprocessor_checksum": self.preprocessor_checksum,
157
+ "encoder": self.encoder,
158
+ "encoder_checksum": self.encoder_checksum,
159
+ "decoder": self.decoder,
160
+ "decoder_checksum": self.decoder_checksum,
161
+ "joiner": self.joiner,
162
+ "joiner_checksum": self.joiner_checksum,
163
+ "uncached_decoder": self.uncached_decoder,
164
+ "uncached_decoder_checksum": self.uncached_decoder_checksum,
165
+ "cached_decoder": self.cached_decoder,
166
+ "cached_decoder_checksum": self.cached_decoder_checksum,
167
+ }
168
+ for key, value in optional_fields.items():
169
+ if value:
170
+ data[key] = value
171
+ return data
172
+
173
+
174
+ @dataclass(frozen=True)
175
+ class VadModelAssets:
176
+ family: str
177
+ model: str
178
+ model_checksum: Optional[str] = None
179
+
180
+ @classmethod
181
+ def from_dict(cls, data: Dict[str, object]) -> "VadModelAssets":
182
+ family = str(data.get("family") or "").strip()
183
+ if not family:
184
+ raise ValueError("VAD asset response is missing required field: family")
185
+
186
+ files = data.get("files")
187
+ if isinstance(files, Mapping):
188
+ model_file = files.get("model")
189
+ if isinstance(model_file, Mapping):
190
+ merged = dict(data)
191
+ url = model_file.get("url")
192
+ checksum = model_file.get("checksum")
193
+ if isinstance(url, str) and url.strip():
194
+ merged["model"] = url
195
+ if isinstance(checksum, str) and checksum.strip():
196
+ merged["model_checksum"] = checksum
197
+ data = merged
198
+
199
+ model = str(data.get("model") or "").strip()
200
+ if not model:
201
+ raise ValueError("VAD asset response is missing required field: model")
202
+
203
+ def optional_string(key: str) -> Optional[str]:
204
+ value = data.get(key)
205
+ if isinstance(value, str) and value.strip():
206
+ return value.strip()
207
+ return None
208
+
209
+ return cls(
210
+ family=family,
211
+ model=model,
212
+ model_checksum=optional_string("model_checksum"),
213
+ )
214
+
215
+ def to_dict(self) -> Dict[str, str]:
216
+ data = {
217
+ "family": self.family,
218
+ "model": self.model,
219
+ }
220
+ if self.model_checksum:
221
+ data["model_checksum"] = self.model_checksum
222
+ return data
223
+
224
+
225
+ @dataclass(frozen=True)
226
+ class LlmModelAssets:
227
+ family: str
228
+ model: str
229
+ model_checksum: Optional[str] = None
230
+ context_size: Optional[int] = None
231
+ chat_template: Optional[str] = None
232
+ chat_template_format: Optional[str] = None
233
+
234
+ @classmethod
235
+ def from_dict(cls, data: Dict[str, object]) -> "LlmModelAssets":
236
+ family = str(data.get("family") or "").strip()
237
+ if not family:
238
+ raise ValueError("LLM asset response is missing required field: family")
239
+
240
+ files = data.get("files")
241
+ if isinstance(files, Mapping):
242
+ model_file = files.get("model")
243
+ if isinstance(model_file, Mapping):
244
+ merged = dict(data)
245
+ url = model_file.get("url")
246
+ checksum = model_file.get("checksum")
247
+ if isinstance(url, str) and url.strip():
248
+ merged["model"] = url
249
+ if isinstance(checksum, str) and checksum.strip():
250
+ merged["model_checksum"] = checksum
251
+ data = merged
252
+
253
+ model = str(data.get("model") or "").strip()
254
+ if not model:
255
+ raise ValueError("LLM asset response is missing required field: model")
256
+
257
+ def optional_string(key: str) -> Optional[str]:
258
+ value = data.get(key)
259
+ if isinstance(value, str) and value.strip():
260
+ return value.strip()
261
+ return None
262
+
263
+ context_size_value = data.get("context_size")
264
+ context_size = None
265
+ if context_size_value is not None:
266
+ context_size = int(context_size_value)
267
+
268
+ return cls(
269
+ family=family,
270
+ model=model,
271
+ model_checksum=optional_string("model_checksum"),
272
+ context_size=context_size,
273
+ chat_template=optional_string("chat_template"),
274
+ chat_template_format=optional_string("chat_template_format"),
275
+ )
276
+
277
+ def to_dict(self) -> Dict[str, object]:
278
+ data: Dict[str, object] = {
279
+ "family": self.family,
280
+ "model": self.model,
281
+ }
282
+ if self.model_checksum:
283
+ data["model_checksum"] = self.model_checksum
284
+ if self.context_size is not None:
285
+ data["context_size"] = self.context_size
286
+ if self.chat_template:
287
+ data["chat_template"] = self.chat_template
288
+ if self.chat_template_format:
289
+ data["chat_template_format"] = self.chat_template_format
290
+ return data
291
+
292
+
293
+ def filename_from_url(url: str, fallback: str) -> str:
294
+ parsed = urlparse(url)
295
+ filename = Path(parsed.path).name
296
+ return filename or fallback
297
+
298
+
299
+ def _registry_url(asset: Mapping[str, object]) -> str:
300
+ path = asset.get("path")
301
+ if not isinstance(path, str) or not path.startswith("/"):
302
+ raise RuntimeError("Registry asset is missing a valid path.")
303
+ return REGISTRY_ORIGIN + path
304
+
305
+
306
+ def _registry_checksum(asset: Mapping[str, object]) -> Optional[str]:
307
+ checksum = asset.get("sha256")
308
+ if isinstance(checksum, str) and checksum.strip():
309
+ return checksum
310
+ return None
311
+
312
+
313
+ def _required_checksum(asset: Mapping[str, object], name: str) -> str:
314
+ checksum = _registry_checksum(asset)
315
+ if checksum is None:
316
+ raise RuntimeError(f"Registry asset is missing required checksum: {name}")
317
+ return checksum
318
+
319
+
320
+ def _model_assets(model_name: str) -> Mapping[str, object]:
321
+ data = MODEL_ASSETS.get(model_name)
322
+ if not isinstance(data, Mapping):
323
+ raise ValueError("Unsupported model: %s" % model_name)
324
+ return data
325
+
326
+
327
+ def _file_asset(data: Mapping[str, object], name: str) -> Mapping[str, object]:
328
+ asset = data.get(name)
329
+ if not isinstance(asset, Mapping):
330
+ raise RuntimeError(f"Registry model entry is missing asset: {name}")
331
+ return asset
332
+
333
+
334
+ def _optional_file_url(data: Mapping[str, object], name: str) -> Optional[str]:
335
+ asset = data.get(name)
336
+ if not isinstance(asset, Mapping):
337
+ return None
338
+ return _registry_url(asset)
339
+
340
+
341
+ def _optional_file_checksum(data: Mapping[str, object], name: str) -> Optional[str]:
342
+ asset = data.get(name)
343
+ if not isinstance(asset, Mapping):
344
+ return None
345
+ return _registry_checksum(asset)
346
+
347
+
348
+ def fetch_model_assets(model_name: str) -> ModelAssets:
349
+ if model_name != WFLOAT_TTS_MODEL_ID:
350
+ raise ValueError("Unsupported TTS model: %s" % model_name)
351
+
352
+ model_assets = _model_assets(model_name)
353
+ model_onnx = _file_asset(model_assets, "model_onnx")
354
+ model_tokens = _file_asset(model_assets, "model_tokens")
355
+ espeak_data = SHARED_ASSETS["espeak_ng_data_zip"]
356
+
357
+ return ModelAssets(
358
+ model_onnx=_registry_url(model_onnx),
359
+ model_onnx_checksum=_required_checksum(model_onnx, "model_onnx"),
360
+ model_tokens=_registry_url(model_tokens),
361
+ model_tokens_checksum=_required_checksum(model_tokens, "model_tokens"),
362
+ espeak_data=_registry_url(espeak_data),
363
+ espeak_checksum=_required_checksum(espeak_data, "espeak_ng_data_zip"),
364
+ )
365
+
366
+
367
+ def fetch_stt_assets(model_name: str) -> SttModelAssets:
368
+ model_assets = _model_assets(model_name)
369
+ family = model_assets.get("family")
370
+ if family not in {"whisper", "zipformer-transducer", "moonshine"}:
371
+ raise ValueError("Unsupported STT model: %s" % model_name)
372
+
373
+ return SttModelAssets(
374
+ family=str(family),
375
+ model=_optional_file_url(model_assets, "model"),
376
+ model_checksum=_optional_file_checksum(model_assets, "model"),
377
+ preprocessor=_optional_file_url(model_assets, "preprocessor"),
378
+ preprocessor_checksum=_optional_file_checksum(model_assets, "preprocessor"),
379
+ encoder=_optional_file_url(model_assets, "encoder"),
380
+ encoder_checksum=_optional_file_checksum(model_assets, "encoder"),
381
+ decoder=_optional_file_url(model_assets, "decoder"),
382
+ decoder_checksum=_optional_file_checksum(model_assets, "decoder"),
383
+ joiner=_optional_file_url(model_assets, "joiner"),
384
+ joiner_checksum=_optional_file_checksum(model_assets, "joiner"),
385
+ uncached_decoder=_optional_file_url(model_assets, "uncached_decoder"),
386
+ uncached_decoder_checksum=_optional_file_checksum(model_assets, "uncached_decoder"),
387
+ cached_decoder=_optional_file_url(model_assets, "cached_decoder"),
388
+ cached_decoder_checksum=_optional_file_checksum(model_assets, "cached_decoder"),
389
+ tokens=_registry_url(_file_asset(model_assets, "tokens")),
390
+ tokens_checksum=_optional_file_checksum(model_assets, "tokens"),
391
+ )
392
+
393
+
394
+ def fetch_vad_assets(model_name: str) -> VadModelAssets:
395
+ if model_name != SILERO_VAD_MODEL_ID:
396
+ raise ValueError("Unsupported VAD model: %s" % model_name)
397
+
398
+ model_assets = _model_assets(model_name)
399
+ model = _file_asset(model_assets, "model")
400
+
401
+ return VadModelAssets(
402
+ family=str(model_assets.get("family") or "silero-vad"),
403
+ model=_registry_url(model),
404
+ model_checksum=_registry_checksum(model),
405
+ )
406
+
407
+
408
+ def fetch_llm_assets(model_name: str) -> LlmModelAssets:
409
+ if model_name != SMOLLM2_360M_INSTRUCT_MODEL_ID:
410
+ raise ValueError("Unsupported LLM model: %s" % model_name)
411
+
412
+ model_assets = _model_assets(model_name)
413
+ model = _file_asset(model_assets, "model")
414
+
415
+ return LlmModelAssets(
416
+ family=str(model_assets.get("family") or "smollm"),
417
+ model=_registry_url(model),
418
+ model_checksum=_registry_checksum(model),
419
+ context_size=8192,
420
+ chat_template_format="chatml",
421
+ )
wfloat/_cache.py ADDED
@@ -0,0 +1,208 @@
1
+ import json
2
+ import os
3
+ import shutil
4
+ import tempfile
5
+ import uuid
6
+ from dataclasses import dataclass
7
+ from pathlib import Path
8
+ from typing import Optional
9
+
10
+ from ._assets import ModelAssets, filename_from_url
11
+ from ._download import (
12
+ download_file,
13
+ extract_archive,
14
+ normalize_checksum,
15
+ resolve_extracted_data_directory,
16
+ verify_checksum,
17
+ )
18
+
19
+
20
+ @dataclass(frozen=True)
21
+ class CachedModelAssets:
22
+ model_name: str
23
+ cache_dir: Path
24
+ model_path: Path
25
+ tokens_path: Path
26
+ espeak_data_dir: Path
27
+ manifest_path: Path
28
+
29
+
30
+ def get_default_cache_dir() -> Path:
31
+ if os.name == "nt":
32
+ local_appdata = os.environ.get("LOCALAPPDATA")
33
+ if local_appdata:
34
+ return Path(local_appdata) / "wfloat" / "Cache"
35
+ return Path.home() / "AppData" / "Local" / "wfloat" / "Cache"
36
+
37
+ if sys_platform_startswith("darwin"):
38
+ return Path.home() / "Library" / "Caches" / "wfloat"
39
+
40
+ xdg_cache_home = os.environ.get("XDG_CACHE_HOME")
41
+ if xdg_cache_home:
42
+ return Path(xdg_cache_home) / "wfloat"
43
+
44
+ return Path.home() / ".cache" / "wfloat"
45
+
46
+
47
+ def sys_platform_startswith(prefix: str) -> bool:
48
+ import sys
49
+
50
+ return sys.platform.startswith(prefix)
51
+
52
+
53
+ def normalize_model_name(model_name: str) -> str:
54
+ normalized = model_name.strip().replace("\\", "/")
55
+ normalized = normalized.replace("/", "--")
56
+ normalized = normalized.replace(" ", "-")
57
+ return normalized
58
+
59
+
60
+ def _ensure_directory(path: Path) -> None:
61
+ path.mkdir(parents=True, exist_ok=True)
62
+
63
+
64
+ def _cleanup_stale_model_files(model_dir: Path, active_names) -> None:
65
+ if not model_dir.exists():
66
+ return
67
+
68
+ for child in model_dir.iterdir():
69
+ if child.name in active_names:
70
+ continue
71
+ if child.is_file():
72
+ child.unlink()
73
+
74
+
75
+ def _write_manifest(manifest_path: Path, model_name: str, assets: ModelAssets) -> None:
76
+ manifest_payload = {
77
+ "model_name": model_name,
78
+ "assets": assets.to_dict(),
79
+ }
80
+ manifest_path.write_text(json.dumps(manifest_payload, indent=2, sort_keys=True) + "\n")
81
+
82
+
83
+ def _ensure_cached_file(
84
+ source_url: str,
85
+ checksum: str,
86
+ destination: Path,
87
+ downloads_dir: Path,
88
+ *,
89
+ force_download: bool,
90
+ ) -> Path:
91
+ if not force_download and verify_checksum(destination, checksum):
92
+ return destination
93
+
94
+ suffix = destination.suffix or ".bin"
95
+ temp_download = downloads_dir / (uuid.uuid4().hex + suffix)
96
+ download_file(source_url, temp_download, expected_checksum=checksum)
97
+ _ensure_directory(destination.parent)
98
+ os.replace(str(temp_download), str(destination))
99
+ return destination
100
+
101
+
102
+ def _install_espeak_data(
103
+ assets: ModelAssets,
104
+ cache_root: Path,
105
+ downloads_dir: Path,
106
+ *,
107
+ force_download: bool,
108
+ ) -> Path:
109
+ checksum = normalize_checksum(assets.espeak_checksum)
110
+ espeak_root = cache_root / "espeak" / checksum
111
+ data_dir = espeak_root / "espeak-ng-data"
112
+ ready_marker = espeak_root / ".ready"
113
+
114
+ if (
115
+ not force_download
116
+ and ready_marker.is_file()
117
+ and data_dir.is_dir()
118
+ ):
119
+ return data_dir
120
+
121
+ if espeak_root.exists():
122
+ shutil.rmtree(espeak_root)
123
+
124
+ _ensure_directory(espeak_root)
125
+ archive_name = filename_from_url(assets.espeak_data, checksum + ".zip")
126
+ temp_archive = downloads_dir / (uuid.uuid4().hex + "-" + archive_name)
127
+ download_file(
128
+ assets.espeak_data,
129
+ temp_archive,
130
+ expected_checksum=assets.espeak_checksum,
131
+ )
132
+
133
+ extraction_root = Path(
134
+ tempfile.mkdtemp(prefix="wfloat-espeak-", dir=str(downloads_dir))
135
+ )
136
+ try:
137
+ extract_archive(temp_archive, extraction_root)
138
+ resolved_data_dir = resolve_extracted_data_directory(extraction_root)
139
+ if data_dir.exists():
140
+ shutil.rmtree(data_dir)
141
+ shutil.copytree(str(resolved_data_dir), str(data_dir))
142
+ ready_marker.write_text("ready\n")
143
+ finally:
144
+ if temp_archive.exists():
145
+ temp_archive.unlink()
146
+ if extraction_root.exists():
147
+ shutil.rmtree(extraction_root)
148
+
149
+ return data_dir
150
+
151
+
152
+ def cache_model_assets(
153
+ model_name: str,
154
+ assets: ModelAssets,
155
+ *,
156
+ cache_dir: Optional[Path] = None,
157
+ force_download: bool = False,
158
+ ) -> CachedModelAssets:
159
+ cache_root = Path(cache_dir) if cache_dir is not None else get_default_cache_dir()
160
+ models_dir = cache_root / "models"
161
+ downloads_dir = cache_root / "downloads"
162
+
163
+ _ensure_directory(models_dir)
164
+ _ensure_directory(downloads_dir)
165
+
166
+ model_dir = models_dir / normalize_model_name(model_name)
167
+ _ensure_directory(model_dir)
168
+
169
+ model_filename = filename_from_url(assets.model_onnx, "model.onnx")
170
+ tokens_filename = filename_from_url(assets.model_tokens, "tokens.txt")
171
+ manifest_path = model_dir / "manifest.json"
172
+
173
+ _cleanup_stale_model_files(
174
+ model_dir,
175
+ active_names={model_filename, tokens_filename, "manifest.json"},
176
+ )
177
+
178
+ model_path = _ensure_cached_file(
179
+ assets.model_onnx,
180
+ assets.model_onnx_checksum,
181
+ model_dir / model_filename,
182
+ downloads_dir,
183
+ force_download=force_download,
184
+ )
185
+ tokens_path = _ensure_cached_file(
186
+ assets.model_tokens,
187
+ assets.model_tokens_checksum,
188
+ model_dir / tokens_filename,
189
+ downloads_dir,
190
+ force_download=force_download,
191
+ )
192
+ espeak_data_dir = _install_espeak_data(
193
+ assets,
194
+ cache_root,
195
+ downloads_dir,
196
+ force_download=force_download,
197
+ )
198
+
199
+ _write_manifest(manifest_path, model_name, assets)
200
+
201
+ return CachedModelAssets(
202
+ model_name=model_name,
203
+ cache_dir=cache_root,
204
+ model_path=model_path,
205
+ tokens_path=tokens_path,
206
+ espeak_data_dir=espeak_data_dir,
207
+ manifest_path=manifest_path,
208
+ )