refactor: unify Language enum across all layers

- EngineConfig.language: Language (was lang_code: str = 'a')
- Engine owns _KOKORO_LANG_MAP, engine_language(), supported_languages()
- Engine provides language_for_voice_id() for voice catalog
- Plugins/kokoro/__init__.py calls engine_language() internally
- create_pipeline(plugin_id, language=Language) — no kokoro codes
- pipeline_factory.py clean of kokoro-specific code
- Domain functions raise TypeError if non-enum passed
- WebUI api.py: _parse_language() helper at API boundary
- Voice catalog returns ISO codes (lang.value)
- Constants: LANGUAGE_DESCRIPTIONS keyed by Language enum
- All tests updated for Language enum
- 1414 tests pass
This commit is contained in:
Artem Akymenko
2026-07-27 07:36:36 +00:00
parent 713abdfd73
commit 0dc491e420
35 changed files with 339 additions and 291 deletions
+3 -3
View File
@@ -254,7 +254,7 @@ def execute_conversion(
audio_sink=audio_sink,
subtitle_mode=effective_subtitle_mode,
max_subtitle_words=request.max_subtitle_words,
lang_code=request.language,
language=request.language,
use_spacy_segmentation=use_spacy,
)
@@ -470,7 +470,7 @@ def execute_conversion(
subtitle_writer,
subtitle_mode=request.subtitle_mode,
max_subtitle_words=request.max_subtitle_words,
lang_code=request.language,
language=request.language,
use_spacy_segmentation=use_spacy,
fallback_end_time=stats.current_time,
)
@@ -480,7 +480,7 @@ def execute_conversion(
chapter_subtitle_writer,
subtitle_mode=request.subtitle_mode,
max_subtitle_words=request.max_subtitle_words,
lang_code=request.language,
language=request.language,
use_spacy_segmentation=use_spacy,
fallback_end_time=stats.current_time,
)
+27 -14
View File
@@ -1,4 +1,5 @@
from abogen.utils import get_version
from abogen.domain.enums import Language
# Program Information
PROGRAM_NAME = "abogen"
@@ -16,8 +17,22 @@ SUBTITLE_FORMATS = [
("ass_centered_narrow", "ASS (centered narrow)"),
]
# Language description mapping
# Language description mapping (Language enum → human-readable label).
LANGUAGE_DESCRIPTIONS = {
Language.EN_US: "American English",
Language.EN_GB: "British English",
Language.ES: "Spanish",
Language.FR: "French",
Language.HI: "Hindi",
Language.IT: "Italian",
Language.JA: "Japanese",
Language.PT_BR: "Brazilian Portuguese",
Language.ZH: "Mandarin Chinese",
}
# Display-only mapping for kokoro codes → labels.
# Used by voice catalog and PyQt (legacy) where kokoro codes are still present.
KOKORO_CODE_LABELS = {
"a": "American English",
"b": "British English",
"e": "Spanish",
@@ -56,24 +71,22 @@ SUPPORTED_INPUT_FORMATS = [
]
# Supported languages for subtitle generation
# Currently, only 'a (American English)' and 'b (British English)' are supported for subtitle generation.
# Currently, only English (EN_US, EN_GB) are supported for subtitle generation.
# This is because tokens that contain timestamps are not generated for other languages in the Kokoro pipeline.
# Please refer to: https://github.com/hexgrad/kokoro/blob/6d87f4ae7abc2d14dbc4b3ef2e5f19852e861ac2/kokoro/pipeline.py
# 383 English processing (unchanged)
# 384 if self.lang_code in 'ab':
SUPPORTED_LANGUAGES_FOR_SUBTITLE_GENERATION = list(LANGUAGE_DESCRIPTIONS.keys())
SUPPORTED_LANGUAGES_FOR_SUBTITLE_GENERATION = [Language.EN_US, Language.EN_GB]
# Voice and sample text mapping
SAMPLE_VOICE_TEXTS = {
"a": "This is a sample of the selected voice.",
"b": "This is a sample of the selected voice.",
"e": "Este es una muestra de la voz seleccionada.",
"f": "Ceci est un exemple de la voix sélectionnée.",
"h": "यह चयनित आवाज़ का एक नमूना है।",
"i": "Questo è un esempio della voce selezionata.",
"j": "これは選択した声のサンプルです。",
"p": "Este é um exemplo da voz selecionada.",
"z": "这是所选语音的示例。",
Language.EN_US: "This is a sample of the selected voice.",
Language.EN_GB: "This is a sample of the selected voice.",
Language.ES: "Este es una muestra de la voz seleccionada.",
Language.FR: "Ceci est un exemple de la voix sélectionnée.",
Language.HI: "यह चयनित आवाज़ का एक नमूना है।",
Language.IT: "Questo è un esempio della voce selezionata.",
Language.JA: "これは選択した声のサンプルです。",
Language.PT_BR: "Este é um exemplo da voz selecionada.",
Language.ZH: "这是所选语音的示例。",
}
COLORS = {
+2 -1
View File
@@ -9,6 +9,7 @@ from __future__ import annotations
from collections import defaultdict
from typing import Any, Dict, Iterable, Mapping
from abogen.domain.enums import Language
from abogen.pronunciation_store import increment_usage
@@ -44,7 +45,7 @@ def record_override_usage(
if not usage_counter:
return
language = getattr(job, "language", "") or "a"
language = getattr(job, "language", Language.EN_US) or Language.EN_US
for normalized, amount in usage_counter.items():
if amount <= 0:
continue
+4 -4
View File
@@ -23,7 +23,7 @@ from typing import Any, Callable, Optional, Protocol
from abogen.domain.audio_sink import AudioSink
from abogen.domain.conversion_pipeline import tts_segments
from abogen.domain.enums import SubtitleMode
from abogen.domain.enums import Language, SubtitleMode
from abogen.domain.normalization import TTSContext
from abogen.domain.progress import calc_etr_str
from abogen.domain.subtitle_generation import process_subtitle_tokens
@@ -153,7 +153,7 @@ def process_and_write_subtitles(
*,
subtitle_mode: str,
max_subtitle_words: int,
lang_code: str,
language: Language,
use_spacy_segmentation: bool,
fallback_end_time: float,
) -> None:
@@ -169,7 +169,7 @@ def process_and_write_subtitles(
new_entries,
max_subtitle_words,
subtitle_mode,
lang_code,
language,
use_spacy_segmentation=use_spacy_segmentation,
fallback_end_time=fallback_end_time,
)
@@ -191,7 +191,7 @@ class SynthParams:
audio_sink: Optional[AudioSink] = None
subtitle_mode: str = "Disabled"
max_subtitle_words: int = 50
lang_code: str = "a"
language: Language = Language.EN_US
use_spacy_segmentation: bool = False
+1 -1
View File
@@ -282,7 +282,7 @@ def emit_text_to_sinks(
# subtitle
subtitle_writer: Any = None,
subtitle_mode: str = "Disabled",
subtitle_lang: str = "a",
subtitle_lang: Language = Language.EN_US,
max_subtitle_words: int = 50,
use_spacy_segmentation: bool = True,
# normalization
+6 -7
View File
@@ -15,6 +15,7 @@ from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Callable, Dict, List, Mapping, Optional
from abogen.domain.enums import Language
from abogen.kokoro_text_normalization import (
ApostropheConfig,
normalize_for_pipeline as _normalize_for_pipeline,
@@ -127,7 +128,7 @@ def prepare_text_for_tts(
def build_tts_context(
*,
language: str = "a",
language: Language,
subtitle_mode: str = "Disabled",
pronunciation_overrides: Optional[List[Dict[str, Any]]] = None,
manual_overrides: Optional[List[Dict[str, Any]]] = None,
@@ -143,7 +144,7 @@ def build_tts_context(
merges pronunciation overrides, and compiles all rules.
Args:
language: Language code (a, b, e, f, etc.).
language: Language enum value.
subtitle_mode: Subtitle mode string.
pronunciation_overrides: List of pronunciation override dicts.
manual_overrides: List of manual override dicts.
@@ -199,15 +200,13 @@ def build_tts_context(
)
# Compute split pattern
try:
lang = Language.from_str(language) if not isinstance(language, Language) else language
except ValueError:
lang = Language.EN_US
if not isinstance(language, Language):
raise TypeError(f"language must be Language enum, got {type(language).__name__}: {language!r}")
try:
mode = SubtitleMode.from_str(subtitle_mode) if not isinstance(subtitle_mode, SubtitleMode) else subtitle_mode
except ValueError:
mode = SubtitleMode.DISABLED
split_pattern = get_split_pattern(lang, mode)
split_pattern = get_split_pattern(language, mode)
# Merge pronunciation overrides (accepts dict or object)
source = {
+13 -26
View File
@@ -2,6 +2,10 @@
Provides a unified interface for creating and managing TTS pipelines
across all UI layers (WebUI, PyQt, CLI).
Language handling: the engine owns the mapping between Language enum
and its internal format. Callers pass Language enum; the engine
converts internally. No engine-specific codes leak outside the engine.
"""
from __future__ import annotations
@@ -13,19 +17,6 @@ from abogen.domain.enums import Language
from abogen.domain.voice_resolution import initialize_voice_cache
from abogen.tts_plugin.utils import create_pipeline, is_plugin_registered
# Kokoro-specific language mapping (engine's responsibility)
_KOKORO_LANG_MAP = {
Language.EN_US: "a",
Language.EN_GB: "b",
Language.ES: "e",
Language.FR: "f",
Language.HI: "h",
Language.IT: "i",
Language.JA: "j",
Language.PT_BR: "p",
Language.ZH: "z",
}
def resolve_device(use_gpu: bool) -> str:
"""Determine compute device from job and global config flags."""
@@ -39,29 +30,25 @@ def resolve_device(use_gpu: bool) -> str:
def create_pipeline_for_job(
provider: str,
language: str,
language: Language,
use_gpu: bool,
) -> Any:
"""Create a TTS pipeline with proper device selection.
Handles provider validation, GPU decision, and plugin checks.
Args:
provider: TTS provider name ("kokoro" or "supertonic").
language: Language enum (app-layer type, not engine-specific).
use_gpu: Whether GPU acceleration is requested.
"""
provider = str(provider or "kokoro").strip().lower() or "kokoro"
if not is_plugin_registered(provider):
provider = "kokoro"
# Convert Language enum to Kokoro single-letter code
try:
lang = Language.from_str(language) if not isinstance(language, Language) else language
except ValueError:
lang = Language.EN_US # fallback for unknown languages
kokoro_code = _KOKORO_LANG_MAP.get(lang, "a")
if provider == "supertonic":
return create_pipeline("supertonic")
device = resolve_device(use_gpu)
return create_pipeline("kokoro", lang_code=kokoro_code, device=device)
return create_pipeline("kokoro", language=language, device=device)
def dispose_pipelines(pipelines: Dict[str, Any]) -> None:
@@ -80,7 +67,7 @@ class PipelinePool:
Usage::
pool = PipelinePool()
backend = pool.get("kokoro", "en", use_gpu=True)
backend = pool.get("kokoro", Language.EN_US, use_gpu=True)
# ... use backend ...
pool.dispose_all()
"""
@@ -92,7 +79,7 @@ class PipelinePool:
def get(
self,
provider: str,
language: str,
language: Language,
use_gpu: bool,
*,
request: Any = None,
@@ -102,7 +89,7 @@ class PipelinePool:
Args:
provider: TTS provider name ("kokoro" or "supertonic").
language: Language code (for kokoro).
language: Language enum (app-layer type).
use_gpu: Whether GPU acceleration is requested.
request: ConversionRequest for voice cache initialization.
events: ConversionEvents for logging during cache init.
+3 -3
View File
@@ -13,7 +13,7 @@ from dataclasses import dataclass
from typing import Any, Callable, Dict, Mapping
from abogen.constants import (
LANGUAGE_DESCRIPTIONS,
KOKORO_CODE_LABELS,
SUBTITLE_FORMATS,
SUPPORTED_SOUND_FORMATS,
)
@@ -135,10 +135,10 @@ def _norm_speaker_spec(value: Any, default: str) -> str:
def _norm_language_list(value: Any, default: list) -> list:
if isinstance(value, (list, tuple, set)):
return [code for code in value if isinstance(code, str) and code in LANGUAGE_DESCRIPTIONS]
return [code for code in value if isinstance(code, str) and code in KOKORO_CODE_LABELS]
if isinstance(value, str):
parts = [item.strip().lower() for item in value.split(",") if item.strip()]
return [code for code in parts if code in LANGUAGE_DESCRIPTIONS]
return [code for code in parts if code in KOKORO_CODE_LABELS]
return default
+5 -9
View File
@@ -12,35 +12,31 @@ PUNCTUATION_SENTENCE_COMMA = r".!?,?。!?،,、।"
PUNCTUATION_COMMAS = ",,、"
def get_split_pattern(language: str, subtitle_mode: str) -> str:
def get_split_pattern(language: Language, subtitle_mode: str) -> str:
"""Get the appropriate split pattern based on language and subtitle mode.
Args:
language: Language code (a, b, e, f, etc.)
language: Language enum value.
subtitle_mode: Subtitle mode ("Sentence", "Sentence + Comma", "Line", etc.)
Returns:
Split pattern string
"""
try:
lang = Language.from_str(language) if not isinstance(language, Language) else language
except ValueError:
lang = None # unknown language — treat as non-English, non-CJK
try:
mode = SubtitleMode.from_str(subtitle_mode) if not isinstance(subtitle_mode, SubtitleMode) else subtitle_mode
except ValueError:
mode = SubtitleMode.DISABLED
# For English, always use newline splitting only
if lang in (Language.EN_US, Language.EN_GB):
if language in (Language.EN_US, Language.EN_GB):
return "\n"
# Determine spacing pattern based on language
spacing = r"\s*" if lang and lang.is_cjk else r"\s+"
spacing = r"\s*" if language.is_cjk else r"\s+"
# For CJK languages, when subtitle mode is Disabled or Line, prefer
# punctuation-based splitting instead of plain newline splitting.
if mode in (SubtitleMode.DISABLED, SubtitleMode.LINE) and lang and lang.is_cjk:
if mode in (SubtitleMode.DISABLED, SubtitleMode.LINE) and language.is_cjk:
return rf"(?<=[{PUNCTUATION_SENTENCE}]){spacing}|\n+"
if mode == SubtitleMode.LINE:
+6 -6
View File
@@ -19,7 +19,7 @@ def process_subtitle_tokens(
subtitle_entries: List[Tuple[float, float, str]],
max_subtitle_words: int,
subtitle_mode: str,
lang_code: str,
language: Language,
use_spacy_segmentation: bool = False,
fallback_end_time: Optional[float] = None,
) -> None:
@@ -35,7 +35,7 @@ def process_subtitle_tokens(
max_subtitle_words: Maximum number of words per subtitle entry.
subtitle_mode: One of "Disabled", "Line", "Sentence", "Sentence + Comma",
"Sentence + Highlighting", or a string like "5" for word-count mode.
lang_code: Language code for spaCy processing (e.g., "a" for English).
language: Language enum value for spaCy processing.
use_spacy_segmentation: Whether to use spaCy for sentence boundary detection.
fallback_end_time: Fallback end time for the last entry if none is available.
"""
@@ -49,7 +49,7 @@ def process_subtitle_tokens(
use_spacy_for_english = (
use_spacy_segmentation
and subtitle_mode not in [SubtitleMode.DISABLED, SubtitleMode.LINE]
and lang_code in [Language.EN_US, Language.EN_GB]
and language in [Language.EN_US, Language.EN_GB]
and subtitle_mode in [SubtitleMode.SENTENCE, SubtitleMode.SENTENCE_COMMA]
)
@@ -61,7 +61,7 @@ def process_subtitle_tokens(
if use_spacy_for_english and subtitle_mode != SubtitleMode.LINE:
_process_spacy_sentences(
processed_tokens, subtitle_entries, max_subtitle_words,
subtitle_mode, lang_code, fallback_end_time
subtitle_mode, language, fallback_end_time
)
else:
_process_regex_sentences(
@@ -141,7 +141,7 @@ def _process_spacy_sentences(
subtitle_entries: List[Tuple[float, float, str]],
max_subtitle_words: int,
subtitle_mode: str,
lang_code: str,
language: Language,
fallback_end_time: Optional[float],
) -> None:
"""Process tokens using spaCy for sentence boundary detection."""
@@ -155,7 +155,7 @@ def _process_spacy_sentences(
)
return
nlp = get_spacy_model(lang_code)
nlp = get_spacy_model(language)
if not nlp:
_process_regex_sentences(
tokens, subtitle_entries, max_subtitle_words,
+14 -40
View File
@@ -21,20 +21,6 @@ SPACY_MODELS = {
Language.HI: "xx_sent_ud_sm",
}
# Kokoro single-letter codes -> Language enum (inverse of pipeline_factory._KOKORO_LANG_MAP)
_KOKORO_TO_LANGUAGE = {
"a": Language.EN_US,
"b": Language.EN_GB,
"e": Language.ES,
"f": Language.FR,
"h": Language.HI,
"i": Language.IT,
"j": Language.JA,
"p": Language.PT_BR,
"z": Language.ZH,
}
def _load_spacy():
"""Lazy load spaCy module."""
global _spacy
@@ -48,12 +34,12 @@ def _load_spacy():
return _spacy
def get_spacy_model(lang_code, log_callback=None):
def get_spacy_model(language: Language, log_callback=None):
"""
Get or load a spaCy model for the given language code.
Get or load a spaCy model for the given language.
Args:
lang_code: Language code or Language enum (e.g., "a", "en-US", Language.EN_US)
language: Language enum value.
log_callback: Optional function to log messages
Returns:
@@ -61,36 +47,24 @@ def get_spacy_model(lang_code, log_callback=None):
"""
def log(msg, is_error=False):
# Prefer GUI log callback when provided to avoid spamming stdout.
if log_callback:
color = "red" if is_error else "grey"
try:
log_callback((msg, color))
except Exception:
# Fallback to printing if callback misbehaves
print(msg)
else:
print(msg)
# Normalize to Language enum
if not isinstance(lang_code, Language):
if isinstance(lang_code, str) and lang_code in _KOKORO_TO_LANGUAGE:
lang_code = _KOKORO_TO_LANGUAGE[lang_code]
else:
try:
lang_code = Language.from_str(lang_code)
except ValueError:
log(f"\nspaCy: Unknown language '{lang_code}'...")
return None
if not isinstance(language, Language):
raise TypeError(f"language must be Language enum, got {type(language).__name__}: {language!r}")
# Check if model is cached
if lang_code in _nlp_cache:
return _nlp_cache[lang_code]
if language in _nlp_cache:
return _nlp_cache[language]
# Check if language is supported
model_name = SPACY_MODELS.get(lang_code)
model_name = SPACY_MODELS.get(language)
if not model_name:
log(f"\nspaCy: No model mapping for language '{lang_code}'...")
log(f"\nspaCy: No model mapping for language '{language}'...")
return None
# Lazy load spaCy
@@ -114,7 +88,7 @@ def get_spacy_model(lang_code, log_callback=None):
if "parser" not in nlp.pipe_names and "sentencizer" not in nlp.pipe_names:
nlp.add_pipe("sentencizer")
_nlp_cache[lang_code] = nlp
_nlp_cache[language] = nlp
return nlp
except OSError:
# Model not found, attempt download
@@ -131,7 +105,7 @@ def get_spacy_model(lang_code, log_callback=None):
if "parser" not in nlp.pipe_names and "sentencizer" not in nlp.pipe_names:
nlp.add_pipe("sentencizer")
_nlp_cache[lang_code] = nlp
_nlp_cache[language] = nlp
log(f"spaCy model '{model_name}' downloaded and loaded")
return nlp
except Exception as e:
@@ -145,19 +119,19 @@ def get_spacy_model(lang_code, log_callback=None):
return None
def segment_sentences(text, lang_code, log_callback=None):
def segment_sentences(text, language: Language, log_callback=None):
"""
Segment text into sentences using spaCy.
Args:
text: Text to segment
lang_code: Language code
language: Language enum value
log_callback: Optional function to log messages
Returns:
List of sentence strings, or None if spaCy unavailable
"""
nlp = get_spacy_model(lang_code, log_callback)
nlp = get_spacy_model(language, log_callback)
if nlp is None:
return None
+2 -2
View File
@@ -4,7 +4,7 @@ import json
import os
from typing import Any, Dict, List, Optional
from abogen.constants import LANGUAGE_DESCRIPTIONS
from abogen.constants import KOKORO_CODE_LABELS
from abogen.utils import get_user_config_path
_CONFIG_WRAPPER_KEY = "abogen_speaker_configs"
@@ -163,4 +163,4 @@ def list_configs() -> List[Dict[str, Any]]:
def describe_language(code: str) -> str:
code = (code or "a").lower()
return LANGUAGE_DESCRIPTIONS.get(code, code.upper())
return KOKORO_CODE_LABELS.get(code, code.upper())
+13 -2
View File
@@ -1,6 +1,7 @@
import re
from abogen.utils import detect_encoding, load_config
from abogen.constants import SAMPLE_VOICE_TEXTS
from abogen.domain.enums import Language
# Pre-compile frequently used regex patterns for better performance
_METADATA_TAG_PATTERN = re.compile(r"<<METADATA_[^:]+:[^>]*>>")
@@ -388,8 +389,18 @@ def parse_ass_file(file_path):
return subtitles
def get_sample_voice_text(lang_code):
return SAMPLE_VOICE_TEXTS.get(lang_code, SAMPLE_VOICE_TEXTS["a"])
def get_sample_voice_text(language):
"""Get sample voice text for a language.
Args:
language: Language enum value or string (for backward compatibility).
"""
if isinstance(language, str):
try:
language = Language.from_str(language)
except (ValueError, AttributeError):
language = Language.EN_US
return SAMPLE_VOICE_TEXTS.get(language, SAMPLE_VOICE_TEXTS[Language.EN_US])
# Backward-compatible re-exports — canonical location is domain/output_paths.py
+1 -1
View File
@@ -8,7 +8,7 @@ Usage:
from abogen.tts_plugin.plugin_manager import get_plugin_manager
manager = get_plugin_manager()
engine = manager.create_engine("kokoro", lang_code="a", device="cpu")
engine = manager.create_engine("kokoro", language=Language.EN_US, device="cpu")
session = engine.create_session()
try:
result = session.synthesize("Hello world")
+5 -3
View File
@@ -9,6 +9,8 @@ from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Mapping
from abogen.domain.enums import Language
@dataclass(frozen=True)
class AudioFormat:
@@ -103,9 +105,9 @@ class EngineConfig:
Attributes:
device: Device to use (e.g., "cpu", "cuda:0").
lang_code: Language code for the engine (e.g., "a" for Kokoro English).
Plugins that do not require a language code ignore this field.
language: Language enum value. The engine converts to its internal
format internally — callers never see engine-specific codes.
"""
device: str = "cpu"
lang_code: str = "a"
language: Language = Language.EN_US
+5 -4
View File
@@ -10,6 +10,7 @@ from typing import Any, Iterator
import numpy as np
from abogen.domain.enums import Language
from abogen.tts_plugin.plugin_manager import get_plugin_manager
@@ -123,7 +124,7 @@ class Pipeline:
Presents the same interface that old callers expect::
pipeline = create_pipeline("kokoro", lang_code="a", device="cpu")
pipeline = create_pipeline("kokoro", language=Language.EN_US, device="cpu")
for segment in pipeline(text, voice="af_nova", speed=1.0):
audio = segment.audio
"""
@@ -200,7 +201,7 @@ class Pipeline:
def create_pipeline(
plugin_id: str,
*,
lang_code: str = "a",
language: Language = Language.EN_US,
device: str = "cpu",
) -> Pipeline:
"""Create a callable TTS pipeline via the Plugin Architecture.
@@ -211,7 +212,7 @@ def create_pipeline(
Args:
plugin_id: Plugin identifier (e.g., "kokoro", "supertonic").
lang_code: Language code for the engine.
language: Language enum value (app-layer type, not engine-specific).
device: Device to use (e.g., "cpu", "cuda:0").
Returns:
@@ -235,7 +236,7 @@ def create_pipeline(
})(),
)
config = EngineConfig(device=device, lang_code=lang_code)
config = EngineConfig(device=device, language=language)
engine = manager.create_engine(plugin_id, context=ctx, model_path=None, config=config)
return Pipeline(engine)
+16 -28
View File
@@ -14,9 +14,12 @@ from abogen.kokoro_text_normalization import normalize_for_pipeline
from abogen.normalization_settings import build_apostrophe_config
from abogen.text_extractor import extract_from_path
from abogen.voice_cache import ensure_voice_assets
from abogen.webui.conversion_runner import SAMPLE_RATE, _select_device, _to_float32, _spec_to_voice_ids
from abogen.domain.device import select_device as _select_device
from abogen.domain.audio_helpers import to_float32 as _to_float32, SAMPLE_RATE
from abogen.domain.voice_resolution import spec_to_voice_ids as _spec_to_voice_ids
from abogen.domain.voice_loader import resolve_voice
from abogen.domain.split_pattern import get_split_pattern
from abogen.domain.enums import Language
from abogen.tts_plugin.utils import create_pipeline
@@ -43,11 +46,11 @@ def _resolve_voice_setting(value: str) -> tuple[str, Optional[str], Optional[str
return resolve_voice_setting(value)
def _load_pipeline(language: str, use_gpu: bool) -> Any:
def _load_pipeline(language: Language, use_gpu: bool) -> Any:
device = "cpu"
if use_gpu:
device = _select_device()
return create_pipeline("kokoro", lang_code=language, device=device)
return create_pipeline("kokoro", language=language, device=device)
def _extract_cases_from_text(text: str) -> List[Tuple[str, str]]:
@@ -127,32 +130,14 @@ def run_debug_tts_wavs(
if missing:
raise RuntimeError(f"Debug EPUB missing expected codes: {', '.join(missing)}")
language = str(settings.get("language") or "a").strip() or "a"
# Kokoro's KPipeline expects short language codes like "a" (American English),
# but older settings may store ISO-like values such as "en".
language_aliases = {
"en": "a",
"en-us": "a",
"en_us": "a",
"en-gb": "b",
"en_gb": "b",
"es": "e",
"es-es": "e",
"fr": "f",
"fr-fr": "f",
"hi": "h",
"it": "i",
"pt": "p",
"pt-br": "p",
"ja": "j",
"jp": "j",
"zh": "z",
"zh-cn": "z",
}
language = language_aliases.get(language.lower(), language)
raw_language = str(settings.get("language") or "en-US").strip() or "en-US"
try:
language = Language.from_str(raw_language)
except ValueError:
language = Language.EN_US
voice_spec = str(settings.get("default_voice") or "").strip()
use_gpu = bool(settings.get("use_gpu", False))
speed = float(settings.get("default_speed", 1.0) or 1.0)
speed = float(settings.get("default_speed") or 1.0)
# Settings may store "profile:<name>" which is not a Kokoro voice ID.
# Resolve it to a concrete voice formula (e.g. "af_heart*0.5+...") so Kokoro
@@ -162,7 +147,10 @@ def run_debug_tts_wavs(
if resolved_voice:
voice_spec = resolved_voice
if profile_language:
language = str(profile_language).strip() or language
try:
language = Language.from_str(str(profile_language).strip()) or language
except (ValueError, AttributeError):
pass
except Exception:
# Voice profile resolution is best-effort; fall back to raw voice_spec.
pass
+20 -4
View File
@@ -6,6 +6,7 @@ from pathlib import Path
from flask import Blueprint, request, jsonify, send_file, url_for, current_app
from flask.typing import ResponseReturnValue
from abogen.domain.enums import Language
from abogen.webui.routes.utils.settings import (
load_settings,
load_integration_settings,
@@ -47,6 +48,21 @@ from werkzeug.utils import secure_filename
api_bp = Blueprint("api", __name__)
def _parse_language(value: Any) -> Language:
"""Parse a frontend language value to Language enum.
This is the API boundary — frontend sends strings, backend parses
to Language enum. No engine-specific codes leak outside the engine.
"""
if isinstance(value, Language):
return value
try:
return Language.from_str(str(value or "").strip())
except (ValueError, AttributeError):
return Language.EN_US
# --- Voice Profile Routes ---
@api_bp.get("/voice-profiles")
@@ -152,7 +168,7 @@ def api_export_voice_profiles() -> ResponseReturnValue:
def api_voice_profiles_preview() -> ResponseReturnValue:
payload = request.get_json(force=True, silent=True) or {}
text = str(payload.get("text") or "").strip() or "Hello world"
language = str(payload.get("language") or "a").strip().lower() or "a"
language = _parse_language(payload.get("language"))
speed = coerce_float(payload.get("speed"), 1.0)
max_seconds = coerce_float(payload.get("max_seconds"), 8.0)
@@ -186,7 +202,7 @@ def api_voice_profiles_preview() -> ResponseReturnValue:
speed = float(normalized_entry.get("speed") or speed)
else:
voice_spec = formula_from_profile(normalized_entry) or ""
language = str(normalized_entry.get("language") or language)
language = _parse_language(normalized_entry.get("language") or language)
elif formula:
voice_spec = formula
resolved_provider = "kokoro"
@@ -221,7 +237,7 @@ def api_speaker_preview() -> ResponseReturnValue:
pending_id = str(payload.get("pending_id") or "").strip()
text = payload.get("text", "Hello world")
voice = payload.get("voice", "af_heart")
language = payload.get("language", "a")
language = _parse_language(payload.get("language"))
speed_value = payload.get("speed")
speed = coerce_float(speed_value, 1.0)
tts_provider = str(payload.get("tts_provider") or "").strip().lower()
@@ -576,7 +592,7 @@ def api_entity_pronunciation_preview() -> ResponseReturnValue:
token = payload.get("token", "").strip()
pronunciation = payload.get("pronunciation", "").strip()
voice = payload.get("voice", "").strip()
language = payload.get("language", "a").strip()
language = _parse_language(payload.get("language"))
if not token and not pronunciation:
return jsonify({"error": "Token or pronunciation required"}), 400
+7 -27
View File
@@ -15,19 +15,6 @@ from abogen.domain.pronunciation import (
apply_pronunciation_rules,
)
# Kokoro-specific language mapping (engine's responsibility)
_KOKORO_LANG_MAP = {
Language.EN_US: "a",
Language.EN_GB: "b",
Language.ES: "e",
Language.FR: "f",
Language.HI: "h",
Language.IT: "i",
Language.JA: "j",
Language.PT_BR: "p",
Language.ZH: "z",
}
SAMPLE_RATE = 24000
@@ -46,7 +33,7 @@ def clear_preview_pipelines() -> None:
_preview_pipelines.clear()
def _resolve_pipeline(language: str, use_gpu: bool) -> Tuple[Any, bool]:
def _resolve_pipeline(language: Language, use_gpu: bool) -> Tuple[Any, bool]:
devices: List[str] = ["cpu"]
if use_gpu:
preferred = _select_device()
@@ -63,29 +50,22 @@ def _resolve_pipeline(language: str, use_gpu: bool) -> Tuple[Any, bool]:
raise RuntimeError("Preview pipeline is unavailable") from last_error
def get_preview_pipeline(language: str, device: str) -> Any:
# Convert Language enum to Kokoro single-letter code
try:
lang = Language.from_str(language) if not isinstance(language, Language) else language
except ValueError:
lang = Language.EN_US
kokoro_code = _KOKORO_LANG_MAP.get(lang, "a")
key = (kokoro_code, device)
def get_preview_pipeline(language: Language, device: str) -> Any:
key = (language, device)
with _preview_pipeline_lock:
pipeline = _preview_pipelines.get(key)
if pipeline is not None:
return pipeline
from abogen.tts_plugin.utils import create_pipeline
pipeline = create_pipeline("kokoro", lang_code=kokoro_code, device=device)
pipeline = create_pipeline("kokoro", language=lang, device=device)
_preview_pipelines[key] = pipeline
return pipeline
def generate_preview_audio(
text: str,
voice_spec: str,
language: str,
language: Language,
speed: float,
use_gpu: bool,
tts_provider: str = "kokoro",
@@ -131,7 +111,7 @@ def generate_preview_audio(
current_app.logger.exception("Preview normalization failed; using raw text")
normalized_text = source_text
preview_split = get_split_pattern(str(language or "a"), "Disabled")
preview_split = get_split_pattern(language, "Disabled")
if provider == "supertonic":
from abogen.tts_plugin.utils import create_pipeline
@@ -194,7 +174,7 @@ def generate_preview_audio(
def synthesize_preview(
text: str,
voice_spec: str,
language: str,
language: Language,
speed: float,
use_gpu: bool,
tts_provider: str = "kokoro",
+5 -3
View File
@@ -277,17 +277,19 @@ def filter_voice_catalog(
def build_voice_catalog() -> List[Dict[str, str]]:
from plugins.kokoro.engine import language_for_voice_id
catalog: List[Dict[str, str]] = []
gender_map = {"f": "Female", "m": "Male"}
for voice_id in get_voices("kokoro"):
prefix, _, rest = voice_id.partition("_")
language_code = prefix[0] if prefix else "a"
gender_code = prefix[1] if len(prefix) > 1 else ""
lang = language_for_voice_id(voice_id)
catalog.append(
{
"id": voice_id,
"language": language_code,
"language_label": LANGUAGE_DESCRIPTIONS.get(language_code, language_code.upper()),
"language": lang.value,
"language_label": LANGUAGE_DESCRIPTIONS.get(lang, lang.value.upper()),
"gender": gender_map.get(gender_code, "Unknown"),
"gender_code": gender_code,
"display_name": rest.replace("_", " ").title() if rest else voice_id,
+2 -1
View File
@@ -2,6 +2,7 @@ from typing import Any, Dict, List, Optional
from flask import Blueprint, render_template, request, jsonify, abort, flash, redirect, url_for
from flask.typing import ResponseReturnValue
from abogen.domain.enums import Language
from abogen.webui.routes.utils.voice import (
template_options,
resolve_voice_setting,
@@ -39,7 +40,7 @@ def test_voice() -> ResponseReturnValue:
return synthesize_preview(
text=text,
voice_spec=voice,
language="a", # Default language
language=Language.EN_US,
speed=speed,
use_gpu=use_gpu,
)
+14 -5
View File
@@ -14,6 +14,7 @@ from enum import Enum
from pathlib import Path
from typing import Any, Callable, Dict, Iterable, List, Optional, Mapping
from abogen.domain.enums import Language
from abogen.utils import get_internal_cache_path, get_user_settings_dir, load_config
from abogen.voice_cache import bootstrap_voice_cache
from abogen.integrations.audiobookshelf import (
@@ -105,7 +106,7 @@ class Job:
id: str
original_filename: str
stored_path: Path
language: str
language: Language
voice: str
speed: float
use_gpu: bool
@@ -296,7 +297,7 @@ class PendingJob:
id: str
original_filename: str
stored_path: Path
language: str
language: Language
voice: str
speed: float
use_gpu: bool
@@ -384,7 +385,7 @@ class ConversionService:
*,
original_filename: str,
stored_path: Path,
language: str,
language: Language,
voice: str,
speed: float,
tts_provider: str = "kokoro",
@@ -910,7 +911,7 @@ class ConversionService:
"id": job.id,
"original_filename": job.original_filename,
"stored_path": str(job.stored_path),
"language": job.language,
"language": job.language.value if isinstance(job.language, Language) else str(job.language),
"tts_provider": getattr(job, "tts_provider", "kokoro"),
"voice": job.voice,
"speed": job.speed,
@@ -1026,11 +1027,19 @@ class ConversionService:
stored_path = Path(payload["stored_path"])
output_folder_raw = payload.get("output_folder")
output_folder = Path(output_folder_raw) if output_folder_raw else None
raw_lang = payload.get("language", "")
if isinstance(raw_lang, Language):
language = raw_lang
else:
try:
language = Language.from_str(str(raw_lang or "").strip())
except (ValueError, AttributeError):
language = Language.EN_US
job = Job(
id=payload["id"],
original_filename=payload["original_filename"],
stored_path=stored_path,
language=payload.get("language", "a"),
language=language,
tts_provider=str(payload.get("tts_provider") or "kokoro"),
voice=payload.get("voice", ""),
speed=float(payload.get("speed", 1.0)),
+3 -1
View File
@@ -165,6 +165,7 @@ def create_engine(
"""
try:
KPipeline = _load_kpipeline()
from plugins.kokoro.engine import engine_language
# Determine repo_id from model_path or use default
repo_id = "hexgrad/Kokoro-82M"
@@ -172,8 +173,9 @@ def create_engine(
# If a specific model path is provided, use it as repo_id
repo_id = str(model_path)
kokoro_code = engine_language(config.language)
pipeline = KPipeline(
lang_code=config.lang_code,
lang_code=kokoro_code,
repo_id=repo_id,
device=config.device,
)
+58
View File
@@ -2,6 +2,10 @@
This module adapts the existing Kokoro backend to the new Engine/EngineSession
protocol. It wraps the KokoroBackend without modifying it.
Language mapping: this is the engine's responsibility. The engine knows
which languages it supports and converts Language enum internal format.
Callers outside this module never see engine-specific codes.
"""
from __future__ import annotations
@@ -11,6 +15,7 @@ from typing import Any
import numpy as np
from abogen.domain.enums import Language
from abogen.tts_plugin.capabilities import VoiceLister
from abogen.tts_plugin.engine import Engine, EngineSession
from abogen.tts_plugin.errors import EngineError
@@ -27,6 +32,59 @@ logger = logging.getLogger(__name__)
# Sample rate for Kokoro audio
_KOKORO_SAMPLE_RATE = 24000
# Engine-internal language mapping: Language enum → kokoro code.
# ONLY visible inside this module — callers never see kokoro codes.
_KOKORO_LANG_MAP: dict[Language, str] = {
Language.EN_US: "a",
Language.EN_GB: "b",
Language.ES: "e",
Language.FR: "f",
Language.HI: "h",
Language.IT: "i",
Language.JA: "j",
Language.PT_BR: "p",
Language.ZH: "z",
}
# Reverse mapping: engine-internal code → Language enum.
# Used by voice catalog and other places that need to convert
# engine codes back to Language enum (e.g. voice ID prefix extraction).
_CODE_TO_LANGUAGE: dict[str, Language] = {v: k for k, v in _KOKORO_LANG_MAP.items()}
def supported_languages() -> list[Language]:
"""Return the list of Language enum values this engine supports.
This is the engine's responsibility — the engine knows which
languages it supports and exposes them as Language enum values.
UI layers query this to populate language selectors.
"""
return list(_KOKORO_LANG_MAP.keys())
def engine_language(lang: Language) -> str:
"""Map a Language enum to the engine's internal code.
This is the engine's responsibility — the engine owns the mapping
between Language enum and its internal format. Callers pass Language
enum; the engine converts internally. The returned string is ONLY
used inside the engine implementation.
"""
return _KOKORO_LANG_MAP.get(lang, "a")
def language_for_voice_id(voice_id: str) -> Language:
"""Determine which Language a voice belongs to from its voice ID.
Kokoro voice IDs encode language as a prefix (e.g. "af_heart" "a" EN_US).
This is kokoro-specific knowledge that stays inside the engine.
Callers pass a voice ID string; the engine returns a Language enum.
"""
prefix = str(voice_id or "").strip()[:1].lower()
if prefix in _CODE_TO_LANGUAGE:
return _CODE_TO_LANGUAGE[prefix]
return Language.EN_US
class KokoroSession:
"""EngineSession implementation for Kokoro.
+3 -1
View File
@@ -14,6 +14,8 @@ from unittest.mock import MagicMock, patch
import numpy as np
from abogen.domain.enums import Language
from abogen.tts_plugin.engine import Engine, EngineSession
from abogen.tts_plugin.errors import EngineError
from abogen.tts_plugin.plugin_manager import PluginManager, get_plugin_manager, reset_plugin_manager
@@ -328,7 +330,7 @@ class TestRegression:
manager._loaded = True
with patch("abogen.tts_plugin.utils.get_plugin_manager", return_value=manager):
backend = create_pipeline("mock_tts", lang_code="a", device="cpu")
backend = create_pipeline("mock_tts", language=Language.EN_US, device="cpu")
# Old interface: pipeline(text, voice=..., speed=..., split_pattern=...)
segments = list(backend(
@@ -3,6 +3,7 @@
import pytest
from unittest.mock import MagicMock, patch
from abogen.domain.enums import Language
from abogen.tts_plugin.plugin_manager import PluginManager, get_plugin_manager, reset_plugin_manager
from abogen.tts_plugin.utils import Pipeline, create_pipeline
from abogen.tts_plugin.engine import Engine, EngineSession
@@ -175,7 +176,7 @@ class TestCreatePipelineCompat:
mock_engine = FakeEngine()
mock_manager.create_engine.return_value = mock_engine
backend = create_pipeline("kokoro", lang_code="a", device="cpu")
backend = create_pipeline("kokoro", language=Language.EN_US, device="cpu")
assert callable(backend)
mock_manager.create_engine.assert_called_once()
@@ -185,7 +186,7 @@ class TestCreatePipelineCompat:
assert call_args.kwargs["model_path"] is None
assert isinstance(call_args.kwargs["config"], EngineConfig)
assert call_args.kwargs["config"].device == "cpu"
assert call_args.kwargs["config"].lang_code == "a"
assert call_args.kwargs["config"].language == Language.EN_US
def test_create_pipeline_raises_for_unknown_plugin(self):
"""create_pipeline raises KeyError for unknown plugins."""
+13 -12
View File
@@ -8,6 +8,7 @@ These tests verify that value objects satisfy the architectural requirements:
import pytest
from abogen.domain.enums import Language
from abogen.tts_plugin.types import (
AudioFormat,
Duration,
@@ -192,23 +193,23 @@ class TestEngineConfigContract:
config = EngineConfig(device="cuda:0")
assert config.device == "cuda:0"
def test_default_lang_code(self) -> None:
def test_default_language(self) -> None:
config = EngineConfig()
assert config.lang_code == "a"
assert config.language == Language.EN_US
def test_custom_lang_code(self) -> None:
config = EngineConfig(lang_code="j")
assert config.lang_code == "j"
def test_custom_language(self) -> None:
config = EngineConfig(language=Language.JA)
assert config.language == Language.JA
def test_immutability(self) -> None:
config = EngineConfig()
with pytest.raises(AttributeError):
config.device = "cuda:0" # type: ignore[misc]
def test_immutability_lang_code(self) -> None:
def test_immutability_language(self) -> None:
config = EngineConfig()
with pytest.raises(AttributeError):
config.lang_code = "j" # type: ignore[misc]
config.language = Language.JA # type: ignore[misc]
def test_unknown_keys_ignored_per_spec(self) -> None:
"""Architecture spec: Unknown keys are ignored (no error).
@@ -225,11 +226,11 @@ class TestEngineConfigContract:
EngineConfig may contain fields that are not relevant to every plugin.
Plugins MUST ignore fields they do not need, not raise on them.
"""
config = EngineConfig(device="cuda:0", lang_code="j")
config = EngineConfig(device="cuda:0", language=Language.JA)
assert config.device == "cuda:0"
assert config.lang_code == "j"
assert config.language == Language.JA
# A plugin that only needs device simply reads config.device
# and ignores config.lang_code — this must not raise.
# and ignores config.language — this must not raise.
def test_engine_config_contains_engine_instance_configuration(self) -> None:
"""Architecture Amendment #1: EngineConfig definition.
@@ -238,7 +239,7 @@ class TestEngineConfigContract:
Engine instance is created and that remain constant throughout
the lifetime of that Engine.
"""
config = EngineConfig(device="cpu", lang_code="a")
config = EngineConfig(device="cpu", language=Language.EN_US)
# Both fields are init-time, immutable, engine-scoped.
assert config.device == "cpu"
assert config.lang_code == "a"
assert config.language == Language.EN_US
+2 -1
View File
@@ -928,10 +928,11 @@ class TestValueObjectsBehavioral:
def test_engine_config_defaults(self) -> None:
from abogen.tts_plugin.types import EngineConfig
from abogen.domain.enums import Language
config = EngineConfig()
assert config.device == "cpu"
assert config.lang_code == "a"
assert config.language == Language.EN_US
def test_parameter_values_defaults(self) -> None:
pv = ParameterValues()
+7 -6
View File
@@ -14,6 +14,7 @@ from unittest.mock import MagicMock, patch
from dataclasses import dataclass, field
from typing import Any
from abogen.domain.enums import Language
from abogen.domain.conversion_engine import (
synthesize_text,
SynthParams,
@@ -254,7 +255,7 @@ class TestProcessAndWriteSubtitles:
writer,
subtitle_mode="Sentence",
max_subtitle_words=5,
lang_code="a",
language=Language.EN_US,
use_spacy_segmentation=False,
fallback_end_time=10.0,
)
@@ -271,7 +272,7 @@ class TestProcessAndWriteSubtitles:
writer,
subtitle_mode="Sentence",
max_subtitle_words=5,
lang_code="a",
language=Language.EN_US,
use_spacy_segmentation=False,
fallback_end_time=2.0,
)
@@ -293,7 +294,7 @@ class TestProcessAndWriteSubtitles:
writer,
subtitle_mode="Line",
max_subtitle_words=5,
lang_code="a",
language=Language.EN_US,
use_spacy_segmentation=False,
fallback_end_time=3.0,
)
@@ -312,7 +313,7 @@ class TestProcessAndWriteSubtitles:
writer,
subtitle_mode="Disabled",
max_subtitle_words=5,
lang_code="a",
language=Language.EN_US,
use_spacy_segmentation=False,
fallback_end_time=2.0,
)
@@ -348,7 +349,7 @@ class TestFullPipeline:
audio_sink=merged_sink,
subtitle_mode="Sentence",
max_subtitle_words=5,
lang_code="a",
language=Language.EN_US,
use_spacy_segmentation=False,
)
@@ -367,7 +368,7 @@ class TestFullPipeline:
subtitle_writer,
subtitle_mode="Sentence",
max_subtitle_words=5,
lang_code="a",
language=Language.EN_US,
use_spacy_segmentation=False,
fallback_end_time=stats.current_time,
)
+3 -3
View File
@@ -37,13 +37,13 @@ class TestConversionRequestBasics:
assert "merge_chapters_at_end" in defaults
def test_split_pattern_computation(self):
pattern = get_split_pattern("a", "Disabled")
pattern = get_split_pattern(Language.EN_US, "Disabled")
assert isinstance(pattern, str)
assert len(pattern) > 0
def test_split_pattern_varies_by_subtitle_mode(self):
pattern_disabled = get_split_pattern("a", "Disabled")
pattern_sentence = get_split_pattern("a", "Sentence")
pattern_disabled = get_split_pattern(Language.EN_US, "Disabled")
pattern_sentence = get_split_pattern(Language.EN_US, "Sentence")
# Different modes should produce different patterns
assert isinstance(pattern_disabled, str)
assert isinstance(pattern_sentence, str)
+16 -11
View File
@@ -2,6 +2,7 @@
import pytest
from unittest.mock import patch, MagicMock
from abogen.domain.enums import Language
from abogen.domain.normalization import prepare_text_for_tts, normalize_text_for_pipeline, build_tts_context, TTSContext
@@ -151,21 +152,21 @@ class TestBuildTtsContext:
"""Tests for the build_tts_context factory."""
def test_returns_tts_context(self):
ctx = build_tts_context()
ctx = build_tts_context(language=Language.EN_US)
assert isinstance(ctx, TTSContext)
def test_default_split_pattern(self):
ctx = build_tts_context(language="a", subtitle_mode="Disabled")
ctx = build_tts_context(language=Language.EN_US, subtitle_mode="Disabled")
assert isinstance(ctx.split_pattern, str)
assert len(ctx.split_pattern) > 0
def test_english_uses_newline_split(self):
ctx = build_tts_context(language="a", subtitle_mode="Disabled")
ctx = build_tts_context(language=Language.EN_US, subtitle_mode="Disabled")
assert ctx.split_pattern == "\n"
def test_cjk_uses_punctuation_split(self):
ctx = build_tts_context(language="j", subtitle_mode="Disabled")
assert "[.??.?!]" in ctx.split_pattern or "\\n" not in ctx.split_pattern
ctx = build_tts_context(language=Language.JA, subtitle_mode="Disabled")
assert r"\n" in ctx.split_pattern
def test_pronunciation_overrides_compiled(self):
overrides = [
@@ -176,6 +177,7 @@ class TestBuildTtsContext:
}
]
ctx = build_tts_context(
language=Language.EN_US,
pronunciation_overrides=overrides,
)
assert ctx.pronunciation_rules is not None
@@ -190,6 +192,7 @@ class TestBuildTtsContext:
}
]
ctx = build_tts_context(
language=Language.EN_US,
manual_overrides=overrides,
)
assert ctx.pronunciation_rules is not None
@@ -203,6 +206,7 @@ class TestBuildTtsContext:
{"token": "x", "pronunciation": "RIGHT", "normalized": "x"}
]
ctx = build_tts_context(
language=Language.EN_US,
pronunciation_overrides=pronunciation,
manual_overrides=manual,
)
@@ -224,22 +228,23 @@ class TestBuildTtsContext:
}
]
ctx = build_tts_context(
language=Language.EN_US,
heteronym_overrides=overrides,
)
assert ctx.heteronym_rules is not None
def test_usage_counter_passed_through(self):
counter = {}
ctx = build_tts_context(usage_counter=counter)
ctx = build_tts_context(language=Language.EN_US, usage_counter=counter)
assert ctx.usage_counter is counter
def test_usage_counter_default_empty(self):
ctx = build_tts_context()
ctx = build_tts_context(language=Language.EN_US)
assert ctx.usage_counter == {}
def test_normalization_overrides_stored(self):
overrides = {"normalization_numbers": False}
ctx = build_tts_context(normalization_overrides=overrides)
ctx = build_tts_context(language=Language.EN_US, normalization_overrides=overrides)
assert ctx.normalization_overrides is overrides
def test_speakers_used_for_pronunciation(self):
@@ -250,7 +255,7 @@ class TestBuildTtsContext:
"resolved_voice": "M1",
}
}
ctx = build_tts_context(speakers=speakers)
ctx = build_tts_context(language=Language.EN_US, speakers=speakers)
assert ctx.pronunciation_rules is not None
assert len(ctx.pronunciation_rules) >= 1
@@ -265,7 +270,7 @@ class TestBuildTtsContext:
mock_cfg.return_value = MagicMock(convert_numbers=True)
with patch("builtins.__import__", side_effect=ImportError):
try:
build_tts_context(log_callback=lambda lvl, msg: logs.append((lvl, msg)))
build_tts_context(language=Language.EN_US, log_callback=lambda lvl, msg: logs.append((lvl, msg)))
except ImportError:
pass
# If num2words is missing and convert_numbers is True, a warning should be logged
@@ -276,7 +281,7 @@ class TestBuildTtsContext:
"normalization_apostrophe_mode": "llm",
}):
with pytest.raises(RuntimeError, match="LLM"):
build_tts_context()
build_tts_context(language=Language.EN_US)
def test_dict_source_accepted(self):
"""merge_pronunciation_overrides should accept a dict."""
+24 -25
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
from unittest.mock import MagicMock, patch
from abogen.domain.enums import Language
from abogen.domain.pipeline_factory import (
PipelinePool,
create_pipeline_for_job,
@@ -31,7 +32,7 @@ class TestCreatePipelineForJob:
@patch("abogen.domain.pipeline_factory.is_plugin_registered", return_value=True)
def test_supertonic_provider(self, _reg, mock_create):
mock_create.return_value = MagicMock()
result = create_pipeline_for_job("supertonic", "en", use_gpu=True)
result = create_pipeline_for_job("supertonic", Language.EN_US, use_gpu=True)
mock_create.assert_called_once_with("supertonic")
assert result is mock_create.return_value
@@ -40,43 +41,41 @@ class TestCreatePipelineForJob:
@patch("abogen.domain.pipeline_factory.resolve_device", return_value="cpu")
def test_kokoro_provider(self, _dev, _reg, mock_create):
mock_create.return_value = MagicMock()
result = create_pipeline_for_job("kokoro", "en", use_gpu=False)
# "en" → fallback to EN_US → kokoro code "a"
mock_create.assert_called_once_with("kokoro", lang_code="a", device="cpu")
result = create_pipeline_for_job("kokoro", Language.EN_US, use_gpu=False)
mock_create.assert_called_once_with("kokoro", language=Language.EN_US, device="cpu")
assert result is mock_create.return_value
@patch("abogen.domain.pipeline_factory.create_pipeline")
@patch("abogen.domain.pipeline_factory.is_plugin_registered", return_value=True)
@patch("abogen.domain.pipeline_factory.resolve_device", return_value="cpu")
def test_kokoro_provider_iso_code(self, _dev, _reg, mock_create):
def test_kokoro_provider_en_gb(self, _dev, _reg, mock_create):
mock_create.return_value = MagicMock()
result = create_pipeline_for_job("kokoro", "en-GB", use_gpu=False)
# "en-GB" → EN_GB → kokoro code "b"
mock_create.assert_called_once_with("kokoro", lang_code="b", device="cpu")
result = create_pipeline_for_job("kokoro", Language.EN_GB, use_gpu=False)
mock_create.assert_called_once_with("kokoro", language=Language.EN_GB, device="cpu")
@patch("abogen.domain.pipeline_factory.create_pipeline")
@patch("abogen.domain.pipeline_factory.is_plugin_registered", return_value=False)
@patch("abogen.domain.pipeline_factory.resolve_device", return_value="cpu")
def test_unknown_provider_falls_back_to_kokoro(self, _dev, _reg, mock_create):
mock_create.return_value = MagicMock()
result = create_pipeline_for_job("unknown_provider", "en", use_gpu=False)
mock_create.assert_called_once_with("kokoro", lang_code="a", device="cpu")
result = create_pipeline_for_job("unknown_provider", Language.EN_US, use_gpu=False)
mock_create.assert_called_once_with("kokoro", language=Language.EN_US, device="cpu")
@patch("abogen.domain.pipeline_factory.create_pipeline")
@patch("abogen.domain.pipeline_factory.is_plugin_registered", return_value=True)
@patch("abogen.domain.pipeline_factory.resolve_device", return_value="cpu")
def test_empty_provider_defaults_to_kokoro(self, _dev, _reg, mock_create):
mock_create.return_value = MagicMock()
result = create_pipeline_for_job("", "en", use_gpu=False)
mock_create.assert_called_once_with("kokoro", lang_code="a", device="cpu")
result = create_pipeline_for_job("", Language.EN_US, use_gpu=False)
mock_create.assert_called_once_with("kokoro", language=Language.EN_US, device="cpu")
@patch("abogen.domain.pipeline_factory.create_pipeline")
@patch("abogen.domain.pipeline_factory.is_plugin_registered", return_value=True)
@patch("abogen.domain.pipeline_factory.resolve_device", return_value="cpu")
def test_none_provider_defaults_to_kokoro(self, _dev, _reg, mock_create):
mock_create.return_value = MagicMock()
result = create_pipeline_for_job(None, "en", use_gpu=False)
mock_create.assert_called_once_with("kokoro", lang_code="a", device="cpu")
result = create_pipeline_for_job(None, Language.EN_US, use_gpu=False)
mock_create.assert_called_once_with("kokoro", language=Language.EN_US, device="cpu")
class TestDisposePipelines:
@@ -110,11 +109,11 @@ class TestPipelinePool:
mock_create.return_value = mock_pipeline
pool = PipelinePool()
result = pool.get("kokoro", "en", use_gpu=True)
result = pool.get("kokoro", Language.EN_US, use_gpu=True)
assert result is mock_pipeline
mock_create.assert_called_once()
result2 = pool.get("kokoro", "en", use_gpu=True)
result2 = pool.get("kokoro", Language.EN_US, use_gpu=True)
assert result2 is mock_pipeline
assert mock_create.call_count == 1
@@ -125,10 +124,10 @@ class TestPipelinePool:
pool = PipelinePool()
request = MagicMock()
pool.get("kokoro", "en", use_gpu=True, request=request)
pool.get("kokoro", Language.EN_US, use_gpu=True, request=request)
assert mock_cache.call_count == 1
pool.get("kokoro", "en", use_gpu=True, request=request)
pool.get("kokoro", Language.EN_US, use_gpu=True, request=request)
assert mock_cache.call_count == 1
@patch("abogen.domain.pipeline_factory.initialize_voice_cache")
@@ -136,7 +135,7 @@ class TestPipelinePool:
def test_get_no_job_skips_voice_cache(self, mock_create, mock_cache):
mock_create.return_value = MagicMock()
pool = PipelinePool()
pool.get("kokoro", "en", use_gpu=True)
pool.get("kokoro", Language.EN_US, use_gpu=True)
mock_cache.assert_not_called()
@patch("abogen.domain.pipeline_factory.create_pipeline_for_job")
@@ -146,8 +145,8 @@ class TestPipelinePool:
mock_create.side_effect = [p1, p2]
pool = PipelinePool()
r1 = pool.get("kokoro", "en", use_gpu=True)
r2 = pool.get("supertonic", "en", use_gpu=True)
r1 = pool.get("kokoro", Language.EN_US, use_gpu=True)
r2 = pool.get("supertonic", Language.EN_US, use_gpu=True)
assert r1 is p1
assert r2 is p2
assert mock_create.call_count == 2
@@ -160,8 +159,8 @@ class TestPipelinePool:
mock_create.side_effect = [p1, p2]
pool = PipelinePool()
pool.get("kokoro", "en", use_gpu=True)
pool.get("supertonic", "en", use_gpu=True)
pool.get("kokoro", Language.EN_US, use_gpu=True)
pool.get("supertonic", Language.EN_US, use_gpu=True)
pool.dispose_all()
p1.dispose.assert_called_once()
@@ -181,5 +180,5 @@ class TestPipelinePool:
def test_unknown_provider_falls_back(self, _reg, _cache, mock_create):
mock_create.return_value = MagicMock()
pool = PipelinePool()
pool.get("bogus_provider", "en", use_gpu=True)
mock_create.assert_called_once_with("kokoro", "en", True)
pool.get("bogus_provider", Language.EN_US, use_gpu=True)
mock_create.assert_called_once_with("kokoro", Language.EN_US, True)
@@ -1,3 +1,4 @@
from abogen.domain.enums import Language
from abogen.webui.routes.utils import synthesize
@@ -45,7 +46,7 @@ def test_preview_applies_manual_override_before_normalization(monkeypatch):
synthesize.generate_preview_audio(
text="He said Unfu*k loudly.",
voice_spec="M1",
language="en",
language=Language.EN_US,
speed=1.0,
use_gpu=False,
tts_provider="supertonic",
+21 -25
View File
@@ -1,10 +1,11 @@
"""Tests for split pattern logic (3 identical copies in codebase)."""
"""Tests for split pattern logic."""
import os
import sys
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
import pytest
from abogen.domain.enums import Language
from abogen.domain.split_pattern import get_split_pattern
@@ -12,49 +13,49 @@ from abogen.domain.split_pattern import get_split_pattern
class TestEnglish:
def test_english_sentence(self):
assert get_split_pattern("en-US", "Sentence") == "\n"
assert get_split_pattern(Language.EN_US, "Sentence") == "\n"
def test_english_sentence_comma(self):
assert get_split_pattern("en-US", "Sentence + Comma") == "\n"
assert get_split_pattern(Language.EN_US, "Sentence + Comma") == "\n"
def test_english_line(self):
assert get_split_pattern("en-US", "Line") == "\n"
assert get_split_pattern(Language.EN_US, "Line") == "\n"
def test_english_disabled(self):
assert get_split_pattern("en-US", "Disabled") == "\n"
assert get_split_pattern(Language.EN_US, "Disabled") == "\n"
def test_english_gb(self):
assert get_split_pattern("en-GB", "Sentence") == "\n"
assert get_split_pattern(Language.EN_GB, "Sentence") == "\n"
# --- CJK languages ---
class TestCJK:
def test_chinese_disabled(self):
pattern = get_split_pattern("zh", "Disabled")
pattern = get_split_pattern(Language.ZH, "Disabled")
assert pattern != "\n"
assert r"\n+" in pattern
def test_chinese_line(self):
pattern = get_split_pattern("zh", "Line")
pattern = get_split_pattern(Language.ZH, "Line")
assert pattern != "\n"
assert r"\n+" in pattern
def test_chinese_sentence(self):
pattern = get_split_pattern("zh", "Sentence")
pattern = get_split_pattern(Language.ZH, "Sentence")
assert r"\n+" in pattern
def test_chinese_sentence_comma(self):
pattern = get_split_pattern("zh", "Sentence + Comma")
pattern = get_split_pattern(Language.ZH, "Sentence + Comma")
assert r"\n+" in pattern
def test_japanese_disabled(self):
pattern = get_split_pattern("ja", "Disabled")
pattern = get_split_pattern(Language.JA, "Disabled")
assert pattern != "\n"
assert r"\n+" in pattern
def test_japanese_sentence(self):
pattern = get_split_pattern("ja", "Sentence")
pattern = get_split_pattern(Language.JA, "Sentence")
assert r"\n+" in pattern
@@ -62,22 +63,17 @@ class TestCJK:
class TestOtherLanguages:
def test_spanish_sentence(self):
pattern = get_split_pattern("es", "Sentence")
pattern = get_split_pattern(Language.ES, "Sentence")
assert r"\n+" in pattern
def test_spanish_line(self):
assert get_split_pattern("es", "Line") == "\n"
assert get_split_pattern(Language.ES, "Line") == "\n"
def test_spanish_disabled(self):
# canonical: \n+ for non-CJK Disabled
assert get_split_pattern("es", "Disabled") == r"\n+"
assert get_split_pattern(Language.ES, "Disabled") == r"\n+"
def test_french_sentence_comma(self):
pattern = get_split_pattern("fr", "Sentence + Comma")
assert r"\n+" in pattern
def test_unknown_lang(self):
pattern = get_split_pattern("x", "Sentence")
pattern = get_split_pattern(Language.FR, "Sentence + Comma")
assert r"\n+" in pattern
@@ -85,17 +81,17 @@ class TestOtherLanguages:
class TestPatternStructure:
def test_sentence_has_lookbehind(self):
pattern = get_split_pattern("es", "Sentence")
pattern = get_split_pattern(Language.ES, "Sentence")
assert r"(?<=" in pattern
def test_sentence_comma_has_comma_chars(self):
pattern = get_split_pattern("es", "Sentence + Comma")
pattern = get_split_pattern(Language.ES, "Sentence + Comma")
assert "," in pattern
def test_cjk_spacing_uses_star(self):
pattern = get_split_pattern("zh", "Sentence")
pattern = get_split_pattern(Language.ZH, "Sentence")
assert r"\s*" in pattern
def test_non_cjk_spacing_uses_plus(self):
pattern = get_split_pattern("es", "Sentence")
pattern = get_split_pattern(Language.ES, "Sentence")
assert r"\s+" in pattern
+10 -9
View File
@@ -2,6 +2,7 @@
import pytest
from abogen.domain.enums import Language
from abogen.domain.subtitle_generation import (
process_subtitle_tokens,
PUNCTUATION_SENTENCE,
@@ -20,7 +21,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="Sentence",
lang_code="a",
language=Language.EN_US,
)
assert entries == []
@@ -41,7 +42,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="Disabled",
lang_code="a",
language=Language.EN_US,
)
# Disabled mode doesn't have special handling in current implementation
# It processes tokens normally
@@ -59,7 +60,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="Line",
lang_code="a",
language=Language.EN_US,
)
# Line mode processes all tokens into entries
assert len(entries) >= 1
@@ -82,7 +83,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="Sentence",
lang_code="a",
language=Language.EN_US,
)
assert len(entries) >= 1
# Should have at least one entry with both sentences or split
@@ -106,7 +107,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="2", # 2 words per entry
lang_code="a",
language=Language.EN_US,
)
assert len(entries) >= 2
# Check that entries are split roughly by word count
@@ -125,7 +126,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="Line",
lang_code="a",
language=Language.EN_US,
fallback_end_time=10.0,
)
assert len(entries) == 1
@@ -143,7 +144,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="Sentence + Highlighting",
lang_code="a",
language=Language.EN_US,
)
assert len(entries) >= 1
# Should contain karaoke tags
@@ -162,7 +163,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=3,
subtitle_mode="Line",
lang_code="a",
language=Language.EN_US,
)
# Should have more than 1 entry due to word limit
assert len(entries) > 1
@@ -179,7 +180,7 @@ class TestProcessSubtitleTokens:
subtitle_entries=entries,
max_subtitle_words=50,
subtitle_mode="Sentence",
lang_code="a",
language=Language.EN_US,
)
assert len(entries) >= 1
# Check that timing is preserved