diff --git a/abogen/application/conversion_request.py b/abogen/application/conversion_request.py index ca21192..0193ac2 100644 --- a/abogen/application/conversion_request.py +++ b/abogen/application/conversion_request.py @@ -13,7 +13,7 @@ from dataclasses import dataclass, field from pathlib import Path from typing import Any, Dict, List, Optional -from abogen.domain.enums import OutputFormat, SaveMode, SubtitleFormat, SubtitleMode +from abogen.domain.enums import Language, OutputFormat, SaveMode, SubtitleFormat, SubtitleMode @dataclass @@ -30,7 +30,7 @@ class ConversionRequest: original_filename: str = "" # --- TTS Settings --- - language: str = "a" + language: Language = Language.EN_US tts_provider: str = "kokoro" voice: str = "M1" voice_profile: Optional[str] = None diff --git a/abogen/application/conversion_service.py b/abogen/application/conversion_service.py index 90e5e15..c4d8b89 100644 --- a/abogen/application/conversion_service.py +++ b/abogen/application/conversion_service.py @@ -135,7 +135,7 @@ def _prepare_tts_context( # Compute split pattern split_pattern = get_split_pattern( - str(request.language or "a"), + request.language or Language.EN_US, request.subtitle_mode or SubtitleMode.DISABLED, ) diff --git a/abogen/domain/enums.py b/abogen/domain/enums.py index 5345f5b..8bc74c6 100644 --- a/abogen/domain/enums.py +++ b/abogen/domain/enums.py @@ -123,3 +123,57 @@ class InputFormat(str, Enum): return cls(suffix) except ValueError: raise ValueError(f"Unsupported input format: {path.suffix!r}. Supported: {[m.value for m in cls]}") + + +class Language(str, Enum): + """TTS language code (ISO 639-1 with region where needed). + + Each engine (Kokoro, Supertonic) maps these to its own + internal language identifiers. + """ + EN_US = "en-US" + EN_GB = "en-GB" + ES = "es" + FR = "fr" + HI = "hi" + IT = "it" + JA = "ja" + PT_BR = "pt-BR" + ZH = "zh" + + @property + def display_name(self) -> str: + """Human-readable language name.""" + _names = { + "en-US": "American English", + "en-GB": "British English", + "es": "Spanish", + "fr": "French", + "hi": "Hindi", + "it": "Italian", + "ja": "Japanese", + "pt-BR": "Brazilian Portuguese", + "zh": "Mandarin Chinese", + } + return _names[self.value] + + @property + def is_cjk(self) -> bool: + """True for CJK languages (Chinese, Japanese).""" + return self in (self.ZH, self.JA) + + @property + def supports_subtitle_tokens(self) -> bool: + """True if this language generates timestamped tokens for subtitles.""" + return self in (self.EN_US, self.EN_GB) + + @classmethod + def from_str(cls, value: str) -> Language: + """Parse from user input: ISO code, case-insensitive.""" + if isinstance(value, Language): + return value + normalized = value.strip() + for member in cls: + if member.value.lower() == normalized.lower(): + return member + raise ValueError(f"Invalid Language: {value!r}. Valid: {[m.value for m in cls]}") diff --git a/abogen/domain/pipeline_factory.py b/abogen/domain/pipeline_factory.py index 6ee8ea7..829489e 100644 --- a/abogen/domain/pipeline_factory.py +++ b/abogen/domain/pipeline_factory.py @@ -9,9 +9,23 @@ from __future__ import annotations from typing import Any, Dict, Optional from abogen.domain.device import select_device +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.""" @@ -36,11 +50,18 @@ def create_pipeline_for_job( 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=language, device=device) + return create_pipeline("kokoro", lang_code=kokoro_code, device=device) def dispose_pipelines(pipelines: Dict[str, Any]) -> None: diff --git a/abogen/domain/split_pattern.py b/abogen/domain/split_pattern.py index de87224..27b9691 100644 --- a/abogen/domain/split_pattern.py +++ b/abogen/domain/split_pattern.py @@ -3,7 +3,7 @@ from __future__ import annotations """Unified split pattern logic extracted from 3 copies.""" import re -from abogen.domain.enums import SubtitleMode +from abogen.domain.enums import Language, SubtitleMode PUNCTUATION_SENTENCE = r".!?。!?" PUNCTUATION_SENTENCE_COMMA = r".!?,。!?、," @@ -19,23 +19,32 @@ def get_split_pattern(language: str, subtitle_mode: str) -> str: 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 language in ("a", "b"): + if lang in (Language.EN_US, Language.EN_GB): return "\n" # Determine spacing pattern based on language - spacing = r"\s*" if language in ("z", "j") else r"\s+" + spacing = r"\s*" if lang and lang.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 subtitle_mode in (SubtitleMode.DISABLED, SubtitleMode.LINE) and language in ("z", "j"): + if mode in (SubtitleMode.DISABLED, SubtitleMode.LINE) and lang and lang.is_cjk: return rf"(?<=[{PUNCTUATION_SENTENCE}]){spacing}|\n+" - if subtitle_mode == SubtitleMode.LINE: + if mode == SubtitleMode.LINE: return "\n" - elif subtitle_mode == SubtitleMode.SENTENCE: + elif mode == SubtitleMode.SENTENCE: return rf"(?<=[{PUNCTUATION_SENTENCE}]){spacing}|\n+" - elif subtitle_mode == SubtitleMode.SENTENCE_COMMA: + elif mode == SubtitleMode.SENTENCE_COMMA: return rf"(?<=[{PUNCTUATION_SENTENCE_COMMA}]){spacing}|\n+" else: return r"\n+" diff --git a/abogen/domain/subtitle_generation.py b/abogen/domain/subtitle_generation.py index e74b9b6..1428f71 100644 --- a/abogen/domain/subtitle_generation.py +++ b/abogen/domain/subtitle_generation.py @@ -10,7 +10,7 @@ from __future__ import annotations import re from typing import List, Optional, Tuple -from abogen.domain.enums import SubtitleMode +from abogen.domain.enums import Language, SubtitleMode # Punctuation constants for sentence splitting @@ -53,7 +53,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 ["a", "b"] + and lang_code in [Language.EN_US, Language.EN_GB] and subtitle_mode in [SubtitleMode.SENTENCE, SubtitleMode.SENTENCE_COMMA] ) diff --git a/abogen/spacy_utils.py b/abogen/spacy_utils.py index 46d6510..d03e859 100644 --- a/abogen/spacy_utils.py +++ b/abogen/spacy_utils.py @@ -2,21 +2,23 @@ Lazy-loaded spaCy utilities for sentence segmentation. """ +from abogen.domain.enums import Language + # Cached spaCy module and models (lazy loaded) _spacy = None _nlp_cache = {} # Language code to spaCy model mapping SPACY_MODELS = { - "a": "en_core_web_sm", # American English - "b": "en_core_web_sm", # British English - "e": "es_core_news_sm", # Spanish - "f": "fr_core_news_sm", # French - "i": "it_core_news_sm", # Italian - "p": "pt_core_news_sm", # Brazilian Portuguese - "z": "zh_core_web_sm", # Mandarin Chinese - "j": "ja_core_news_sm", # Japanese - "h": "xx_sent_ud_sm", # Hindi (multi-language model) + Language.EN_US: "en_core_web_sm", + Language.EN_GB: "en_core_web_sm", + Language.ES: "es_core_news_sm", + Language.FR: "fr_core_news_sm", + Language.IT: "it_core_news_sm", + Language.PT_BR: "pt_core_news_sm", + Language.ZH: "zh_core_web_sm", + Language.JA: "ja_core_news_sm", + Language.HI: "xx_sent_ud_sm", } @@ -36,10 +38,9 @@ def _load_spacy(): def get_spacy_model(lang_code, log_callback=None): """ Get or load a spaCy model for the given language code. - Downloads the model automatically if not available. Args: - lang_code: Language code (a, b, e, f, etc.) + lang_code: Language code or Language enum (e.g., "a", "en-US", Language.EN_US) log_callback: Optional function to log messages Returns: @@ -58,6 +59,14 @@ def get_spacy_model(lang_code, log_callback=None): else: print(msg) + # Normalize to Language enum + if not isinstance(lang_code, Language): + try: + lang_code = Language.from_str(lang_code) + except ValueError: + log(f"\nspaCy: Unknown language '{lang_code}'...") + return None + # Check if model is cached if lang_code in _nlp_cache: return _nlp_cache[lang_code] diff --git a/abogen/webui/routes/utils/synthesize.py b/abogen/webui/routes/utils/synthesize.py index 5f1feb7..bff19f5 100644 --- a/abogen/webui/routes/utils/synthesize.py +++ b/abogen/webui/routes/utils/synthesize.py @@ -7,8 +7,22 @@ from flask import current_app, send_file from flask.typing import ResponseReturnValue from abogen.domain.device import select_device as _select_device +from abogen.domain.enums import Language from abogen.domain.split_pattern import get_split_pattern +# 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 @@ -45,14 +59,21 @@ def _resolve_pipeline(language: str, use_gpu: bool) -> Tuple[Any, bool]: def get_preview_pipeline(language: str, device: str) -> Any: - key = (language, device) + # 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) 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=language, device=device) + pipeline = create_pipeline("kokoro", lang_code=kokoro_code, device=device) _preview_pipelines[key] = pipeline return pipeline diff --git a/tests/test_domain_pipeline_factory.py b/tests/test_domain_pipeline_factory.py index 74cf67d..255cc1e 100644 --- a/tests/test_domain_pipeline_factory.py +++ b/tests/test_domain_pipeline_factory.py @@ -41,16 +41,26 @@ class TestCreatePipelineForJob: def test_kokoro_provider(self, _dev, _reg, mock_create): mock_create.return_value = MagicMock() result = create_pipeline_for_job("kokoro", "en", use_gpu=False) - mock_create.assert_called_once_with("kokoro", lang_code="en", device="cpu") + # "en" → fallback to EN_US → kokoro code "a" + mock_create.assert_called_once_with("kokoro", lang_code="a", 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): + 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") + @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="en", device="cpu") + mock_create.assert_called_once_with("kokoro", lang_code="a", device="cpu") @patch("abogen.domain.pipeline_factory.create_pipeline") @patch("abogen.domain.pipeline_factory.is_plugin_registered", return_value=True) @@ -58,7 +68,7 @@ class TestCreatePipelineForJob: 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="en", device="cpu") + mock_create.assert_called_once_with("kokoro", lang_code="a", device="cpu") @patch("abogen.domain.pipeline_factory.create_pipeline") @patch("abogen.domain.pipeline_factory.is_plugin_registered", return_value=True) @@ -66,7 +76,7 @@ class TestCreatePipelineForJob: 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="en", device="cpu") + mock_create.assert_called_once_with("kokoro", lang_code="a", device="cpu") class TestDisposePipelines: diff --git a/tests/test_enums.py b/tests/test_enums.py index 2a929f3..c5cc422 100644 --- a/tests/test_enums.py +++ b/tests/test_enums.py @@ -5,6 +5,7 @@ from pathlib import Path from abogen.domain.enums import ( InputFormat, + Language, OutputFormat, SaveMode, SubtitleFormat, @@ -97,3 +98,34 @@ class TestInputFormat: def test_dot_ext(self): assert InputFormat.EPUB.dot_ext == ".epub" assert InputFormat.SRT.dot_ext == ".srt" + + +class TestLanguage: + def test_iso_codes(self): + assert Language.EN_US == "en-US" + assert Language.EN_GB == "en-GB" + assert Language.ZH == "zh" + assert Language.JA == "ja" + + def test_display_name(self): + assert Language.EN_US.display_name == "American English" + assert Language.JA.display_name == "Japanese" + + def test_is_cjk(self): + assert Language.ZH.is_cjk is True + assert Language.JA.is_cjk is True + assert Language.EN_US.is_cjk is False + + def test_supports_subtitle_tokens(self): + assert Language.EN_US.supports_subtitle_tokens is True + assert Language.EN_GB.supports_subtitle_tokens is True + assert Language.ZH.supports_subtitle_tokens is False + + def test_from_str_case_insensitive(self): + assert Language.from_str("EN-US") == Language.EN_US + assert Language.from_str("en-gb") == Language.EN_GB + assert Language.from_str("ZH") == Language.ZH + + def test_from_str_invalid(self): + with pytest.raises(ValueError, match="Invalid Language"): + Language.from_str("en") diff --git a/tests/test_split_pattern.py b/tests/test_split_pattern.py index 152179c..a2ef7dc 100644 --- a/tests/test_split_pattern.py +++ b/tests/test_split_pattern.py @@ -12,49 +12,49 @@ from abogen.domain.split_pattern import get_split_pattern class TestEnglish: def test_english_sentence(self): - assert get_split_pattern("a", "Sentence") == "\n" + assert get_split_pattern("en-US", "Sentence") == "\n" def test_english_sentence_comma(self): - assert get_split_pattern("a", "Sentence + Comma") == "\n" + assert get_split_pattern("en-US", "Sentence + Comma") == "\n" def test_english_line(self): - assert get_split_pattern("a", "Line") == "\n" + assert get_split_pattern("en-US", "Line") == "\n" def test_english_disabled(self): - assert get_split_pattern("a", "Disabled") == "\n" + assert get_split_pattern("en-US", "Disabled") == "\n" - def test_english_b(self): - assert get_split_pattern("b", "Sentence") == "\n" + def test_english_gb(self): + assert get_split_pattern("en-GB", "Sentence") == "\n" # --- CJK languages --- class TestCJK: def test_chinese_disabled(self): - pattern = get_split_pattern("z", "Disabled") + pattern = get_split_pattern("zh", "Disabled") assert pattern != "\n" assert r"\n+" in pattern def test_chinese_line(self): - pattern = get_split_pattern("z", "Line") + pattern = get_split_pattern("zh", "Line") assert pattern != "\n" assert r"\n+" in pattern def test_chinese_sentence(self): - pattern = get_split_pattern("z", "Sentence") + pattern = get_split_pattern("zh", "Sentence") assert r"\n+" in pattern def test_chinese_sentence_comma(self): - pattern = get_split_pattern("z", "Sentence + Comma") + pattern = get_split_pattern("zh", "Sentence + Comma") assert r"\n+" in pattern def test_japanese_disabled(self): - pattern = get_split_pattern("j", "Disabled") + pattern = get_split_pattern("ja", "Disabled") assert pattern != "\n" assert r"\n+" in pattern def test_japanese_sentence(self): - pattern = get_split_pattern("j", "Sentence") + pattern = get_split_pattern("ja", "Sentence") assert r"\n+" in pattern @@ -62,18 +62,18 @@ class TestCJK: class TestOtherLanguages: def test_spanish_sentence(self): - pattern = get_split_pattern("e", "Sentence") + pattern = get_split_pattern("es", "Sentence") assert r"\n+" in pattern def test_spanish_line(self): - assert get_split_pattern("e", "Line") == "\n" + assert get_split_pattern("es", "Line") == "\n" def test_spanish_disabled(self): # canonical: \n+ for non-CJK Disabled - assert get_split_pattern("e", "Disabled") == r"\n+" + assert get_split_pattern("es", "Disabled") == r"\n+" def test_french_sentence_comma(self): - pattern = get_split_pattern("f", "Sentence + Comma") + pattern = get_split_pattern("fr", "Sentence + Comma") assert r"\n+" in pattern def test_unknown_lang(self): @@ -85,17 +85,17 @@ class TestOtherLanguages: class TestPatternStructure: def test_sentence_has_lookbehind(self): - pattern = get_split_pattern("e", "Sentence") + pattern = get_split_pattern("es", "Sentence") assert r"(?<=" in pattern def test_sentence_comma_has_comma_chars(self): - pattern = get_split_pattern("e", "Sentence + Comma") + pattern = get_split_pattern("es", "Sentence + Comma") assert "," in pattern def test_cjk_spacing_uses_star(self): - pattern = get_split_pattern("z", "Sentence") + pattern = get_split_pattern("zh", "Sentence") assert r"\s*" in pattern def test_non_cjk_spacing_uses_plus(self): - pattern = get_split_pattern("e", "Sentence") + pattern = get_split_pattern("es", "Sentence") assert r"\s+" in pattern