refactor(tools): AST-neutral closer/bracket hug pass over STT/TTS modules

This commit is contained in:
Teknium
2026-09-03 00:53:14 -07:00
parent 404a324ea5
commit d85b86e97a
9 changed files with 50 additions and 101 deletions
+3 -6
View File
@@ -21,8 +21,7 @@ from hermes_cli._subprocess_compat import windows_hide_flags
from utils import is_truthy_value
from tools.transcription_common import (
COMMON_LOCAL_BIN_DIRS, LOCAL_NATIVE_AUDIO_FORMATS, MAX_FILE_SIZE, SUPPORTED_FORMATS,
_config_number, _error_result, _lazy_ensure_quietly, _process_error_detail,
)
_config_number, _error_result, _lazy_ensure_quietly, _process_error_detail)
# Log-record parity with the origin module.
logger = logging.getLogger("tools.transcription_tools")
@@ -54,8 +53,7 @@ def _run_quiet(command: list, *, timeout: float, env: Optional[dict] = None) ->
return subprocess.run(
command, check=True, capture_output=True, text=True,
encoding="utf-8", errors="replace", timeout=timeout,
stdin=subprocess.DEVNULL, env=env, creationflags=windows_hide_flags(),
)
stdin=subprocess.DEVNULL, env=env, creationflags=windows_hide_flags())
# Shared encode profile for every STT-bound m4a (transcode and silence-trim):
@@ -268,8 +266,7 @@ def _trim_silence_for_cloud_stt(file_path: str, stt_config: Dict[str, Any]) -> O
filter_expr = (
f"silenceremove="
f"start_periods=1:start_threshold={threshold_db}dB:start_silence={keep_seconds}:"
f"stop_periods=-1:stop_threshold={threshold_db}dB:stop_silence={keep_seconds}"
)
f"stop_periods=-1:stop_threshold={threshold_db}dB:stop_silence={keep_seconds}")
work_dir = tempfile.mkdtemp(prefix="hermes-stt-trim-")
trimmed_path = os.path.join(work_dir, f"{Path(file_path).stem or 'audio'}-trimmed.m4a")
# Scale the all-silence guard with keep_ms: output that is solely kept pause must never upload as "speech".
+7 -14
View File
@@ -21,8 +21,7 @@ from tools.transcription_audio import _transcode_audio_for_stt
from tools.transcription_common import (
DEFAULT_GROQ_STT_MODEL, DEFAULT_STT_MODEL, ELEVENLABS_STT_BASE_URL, GROQ_BASE_URL, GROQ_MODELS,
OPENAI_BASE_URL, OPENAI_MODELS, XAI_STT_BASE_URL, _error_result, _get_stt_section,
_lazy_ensure_quietly, _log_prompt_unsupported, _ok_result,
)
_lazy_ensure_quietly, _log_prompt_unsupported, _ok_result)
# Log-record parity with the origin module.
logger = logging.getLogger("tools.transcription_tools")
@@ -110,8 +109,7 @@ def _transcribe_groq(
def _transcribe_openai(
file_path: str, model_name: str, *, api_key: Optional[str] = None,
base_url: Optional[str] = None, provider_label: str = "openai", language: Optional[str] = None,
prompt: Optional[str] = None,
) -> Dict[str, Any]:
prompt: Optional[str] = None) -> Dict[str, Any]:
"""Transcribe via the OpenAI ``audio.transcriptions.create`` SDK shape.
Shared by every OpenAI-compatible endpoint (DeepInfra etc.): explicit ``api_key``/
@@ -189,8 +187,7 @@ def _transcribe_mistral(
language = language or _resolve_stt_language("mistral")
result = client.audio.transcriptions.complete(
model=model_name, file={"content": audio_file, "file_name": Path(file_path).name},
**_sdk_prompt_kwargs(language, prompt),
)
**_sdk_prompt_kwargs(language, prompt))
transcript_text = _extract_transcript_text(result)
logger.info("Transcribed %s via Mistral API (%s, %d chars)",
Path(file_path).name, model_name, len(transcript_text))
@@ -210,8 +207,7 @@ def _post_audio_multipart(url: str, headers: Dict[str, str], file_path: str, dat
def _rest_provider(
file_path: str, provider: str, label: str, post: Callable[[], Any], extract_detail,
extract_text, log: Callable[[str, Dict[str, Any]], None],
) -> Dict[str, Any]:
extract_text, log: Callable[[str, Dict[str, Any]], None]) -> Dict[str, Any]:
"""Shared multipart REST flow: ``post()`` -> ``log(text, body)`` -> ok envelope.
Non-200 -> ``"<label> API error (HTTP n): detail"`` (JSON detail via *extract_detail*, else
@@ -327,8 +323,7 @@ def _transcribe_elevenlabs(
"model_id": model_name,
"tag_audio_events": str(is_truthy_value(elevenlabs_config.get("tag_audio_events", False))).lower(),
"diarize": str(is_truthy_value(elevenlabs_config.get("diarize", False))).lower(),
**({"language_code": language_code} if language_code else {}),
}
**({"language_code": language_code} if language_code else {})}
return _post_audio_multipart(f"{base_url}/speech-to-text", {"xi-api-key": api_key}, file_path, data)
def _log(transcript_text: str, _body: Dict[str, Any]) -> None:
@@ -354,8 +349,7 @@ def _transcribe_deepinfra(
if not model_name:
return _error_result(
"No DeepInfra STT model available. Pin one in config.yaml under stt.deepinfra.model, "
"or check connectivity to api.deepinfra.com so the live catalog can be fetched."
)
"or check connectivity to api.deepinfra.com so the live catalog can be fetched.")
return _transcribe_openai(file_path, model_name, api_key=api_key, base_url=base_url,
provider_label="deepinfra", language=language, prompt=prompt)
@@ -396,8 +390,7 @@ def _resolve_openai_audio_client_config() -> tuple[str, str]:
legacy ladder: direct credentials, then the managed gateway. Failures raise ValueError."""
from tools.transcription_tools import (
_load_stt_config, managed_nous_tools_enabled, nous_tool_gateway_unavailable_message,
resolve_managed_tool_gateway,
)
resolve_managed_tool_gateway)
from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER, read_selection, selection_error
openai_cfg = _load_stt_config().get("openai") or {}
selected = read_selection("stt")
+8 -16
View File
@@ -19,11 +19,9 @@ from tools.tts_command_provider import (
_command_output_format, _command_timeout, _is_command_provider_config as _is_command_stt_provider_config,
_named_provider_config, _resolve_command_config, command_env_passthrough as _command_stt_env_passthrough,
command_failure_detail, render_command_template as _render_command_stt_template,
run_command_provider as _run_command_stt,
)
run_command_provider as _run_command_stt)
from tools.transcription_common import (
BUILTIN_STT_PROVIDERS, _error_result, _log_prompt_unsupported, _ok_result,
)
BUILTIN_STT_PROVIDERS, _error_result, _log_prompt_unsupported, _ok_result)
# Log-record parity with the origin module.
logger = logging.getLogger("tools.transcription_tools")
@@ -72,8 +70,7 @@ def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str:
def _transcribe_command_stt(
file_path: str, provider_name: str, config: Dict[str, Any], stt_config: Dict[str, Any],
model_override: Optional[str] = None, language_override: Optional[str] = None,
prompt: Optional[str] = None,
) -> Dict[str, Any]:
prompt: Optional[str] = None) -> Dict[str, Any]:
"""Transcribe via a user-declared ``stt.providers.<name>: type: command``.
Placeholders (shell-quote-aware; ``{{``/``}}`` stay literal): ``{input_path}``,
@@ -130,8 +127,7 @@ def _unregistered_stt_provider_error(provider: str) -> Dict[str, Any]:
"installed STT plugins, or configure a command provider under "
f"`stt.providers.{key}.command`.",
provider=key,
error_type="provider_not_registered",
)
error_type="provider_not_registered")
def _dispatch_to_plugin_provider(
@@ -178,8 +174,7 @@ def _dispatch_to_plugin_provider(
logger.info("STT plugin provider '%s' reports not available; returning unavailability envelope.", key)
return _error_result(
f"STT plugin '{key}' is not available — check that its required credentials / dependencies are configured.",
provider=key,
)
provider=key)
logger.info("Transcribing with plugin STT provider '%s'...", key)
# The prompt travels via the ABC's ``**extra`` kwargs and is only sent when
# set, so pre-prompt providers see byte-identical calls on the no-prompt path.
@@ -215,8 +210,7 @@ def _enforce_prompt_length_limit(prompt: Optional[str], provider: str) -> Option
logger.warning(
"Transcription prompt is ~%d tokens; whisper-family provider '%s' "
"only uses the final ~%d — truncating to the last %d characters.",
len(prompt) // _PROMPT_CHARS_PER_TOKEN, provider, _WHISPER_PROMPT_TOKEN_CAP, max_chars,
)
len(prompt) // _PROMPT_CHARS_PER_TOKEN, provider, _WHISPER_PROMPT_TOKEN_CAP, max_chars)
return prompt[-max_chars:]
@@ -237,16 +231,14 @@ def _apply_pre_transcription_hook(
return model, None, prompt
hook_results = invoke_hook(
"pre_transcription", file_path=file_path, provider=provider,
model=model, language=language, prompt=prompt, source=source,
)
model=model, language=language, prompt=prompt, source=source)
overrides: Dict[str, Any] = {}
for hook_result in hook_results:
for key, value in (hook_result.items() if isinstance(hook_result, dict) else ()):
if key == "file_path":
logger.warning(
"pre_transcription hook attempted to change "
"file_path (read-only) — ignoring the attempt."
)
"file_path (read-only) — ignoring the attempt.")
elif key not in _PRE_TRANSCRIPTION_MUTABLE_FIELDS:
logger.debug("pre_transcription hook returned unsupported field %r — ignoring.", key)
elif not isinstance(value, str):
+1 -2
View File
@@ -45,8 +45,7 @@ GROQ_MODELS = {"whisper-large-v3", "whisper-large-v3-turbo", "distil-whisper-lar
# (a regression test fails on drift); plugins may not register under these names and the
# dispatcher short-circuits them before command/plugin lookup.
BUILTIN_STT_PROVIDERS = frozenset({
"local", "local_command", "groq", "openai", "mistral", "xai", "elevenlabs", "deepinfra",
})
"local", "local_command", "groq", "openai", "mistral", "xai", "elevenlabs", "deepinfra"})
# Built-in providers that upload audio to a remote API.
CLOUD_STT_PROVIDERS = frozenset(BUILTIN_STT_PROVIDERS - {"local", "local_command"})
+6 -12
View File
@@ -23,8 +23,7 @@ from tools.transcription_audio import _run_quiet
from tools.transcription_common import (
DEFAULT_LOCAL_MODEL, DEFAULT_LOCAL_STT_LANGUAGE, GROQ_MODELS, LOCAL_STT_COMMAND_ENV,
OPENAI_MODELS, _config_number, _error_result, _log_prompt_unsupported, _ok_result,
_process_error_detail,
)
_process_error_detail)
# Log-record parity with the origin module.
logger = logging.getLogger("tools.transcription_tools")
@@ -53,8 +52,7 @@ def _normalize_local_model(model_name: Optional[str]) -> str:
"STT model '%s' is a cloud-only name and cannot be used with the local "
"provider. Falling back to '%s'. Set stt.local.model to a valid "
"faster-whisper size (tiny, base, small, medium, large-v3).",
model_name, DEFAULT_LOCAL_MODEL,
)
model_name, DEFAULT_LOCAL_MODEL)
return DEFAULT_LOCAL_MODEL
return model_name
@@ -77,8 +75,7 @@ def _try_lazy_install_stt() -> bool:
"venv owner: `stat -c '%%u' '$(dirname $(dirname $(which python3)))'` "
"then `su - <owner> -c 'VIRTUAL_ENV=/opt/hermes/.venv "
"uv pip install faster-whisper==1.2.1'`",
exc,
)
exc)
return False
@@ -89,8 +86,7 @@ def _try_lazy_install_stt() -> bool:
_CUDA_LIB_ERROR_MARKERS = (
"libcublas", "libcudnn", "libcudart", "cannot be loaded", "cannot open shared object",
"no kernel image is available", "CUBLAS_STATUS_NOT_SUPPORTED", "no CUDA-capable device",
"CUDA driver version is insufficient",
)
"CUDA driver version is insufficient")
def _looks_like_cuda_lib_error(exc: BaseException) -> bool:
@@ -172,8 +168,7 @@ def build_local_transcribe_kwargs(stt_config: Optional[Dict[str, Any]] = None) -
kwargs: Dict[str, Any] = {
"beam_size": 5,
"condition_on_previous_text": False,
"vad_filter": vad_enabled is None or bool(vad_enabled),
}
"vad_filter": vad_enabled is None or bool(vad_enabled)}
if kwargs["vad_filter"]:
kwargs["vad_parameters"] = {
"min_silence_duration_ms": _config_number(local_cfg, "vad_min_silence_ms", _VAD_MIN_SILENCE_MS_DEFAULT, int)
@@ -246,8 +241,7 @@ def _transcribe_local_command(
return _error_result(prep_error)
command = command_template.format(
input_path=shlex.quote(prepared_input), output_dir=shlex.quote(output_dir),
language=shlex.quote(language), model=shlex.quote(normalized_model),
)
language=shlex.quote(language), model=shlex.quote(normalized_model))
# Scrub Hermes secrets from the child env (same policy as _run_command_stt).
from tools.environments.local import hermes_subprocess_env
_run_quiet(shlex.split(command), timeout=300, env=hermes_subprocess_env(inherit_credentials=False))
+14 -28
View File
@@ -25,42 +25,36 @@ from hermes_cli._subprocess_compat import windows_hide_flags # noqa: F401 (imp
from utils import is_truthy_value
from tools.managed_tool_gateway import resolve_managed_tool_gateway # noqa: F401 (patched by tests)
from tools.tool_backend_helpers import ( # noqa: F401 (patched by tests; read lazily by transcription_cloud)
managed_nous_tools_enabled, nous_tool_gateway_unavailable_message, resolve_openai_audio_api_key,
)
managed_nous_tools_enabled, nous_tool_gateway_unavailable_message, resolve_openai_audio_api_key)
from tools.transcription_common import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
BUILTIN_STT_PROVIDERS, CLOUD_STT_PROVIDERS, DEFAULT_ELEVENLABS_STT_MODEL,
DEFAULT_GROQ_STT_MODEL, DEFAULT_LOCAL_MODEL, DEFAULT_MISTRAL_STT_MODEL, DEFAULT_PROVIDER,
DEFAULT_STT_MODEL, ELEVENLABS_STT_BASE_URL, GROQ_MODELS, LOCAL_STT_COMMAND_ENV,
LOCAL_STT_LANGUAGE_ENV, MAX_FILE_SIZE, OPENAI_MODELS, SUPPORTED_FORMATS, XAI_STT_BASE_URL,
_error_result, _get_stt_section, _ok_result,
)
_error_result, _get_stt_section, _ok_result)
from tools.transcription_audio import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
_CLOUD_TRIM_KEEP_MS_DEFAULT, _CLOUD_TRIM_MIN_INPUT_SECONDS, _CLOUD_TRIM_THRESHOLD_DB_DEFAULT,
_cloud_trim_settings, _convert_caf_to_wav, _find_ffmpeg_binary, _find_ffprobe_binary,
_find_whisper_binary, _prepare_audio_for_transcription, _prepare_local_audio,
_probe_audio_duration, _run_ffmpeg_stt_encode, _trim_silence_for_cloud_stt,
_validate_audio_file, _validate_audio_file_size, _validate_audio_source_file,
)
_validate_audio_file, _validate_audio_file_size, _validate_audio_source_file)
from tools.transcription_local import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
_LOGPROB_THRESHOLD_DEFAULT, _NO_SPEECH_PROB_THRESHOLD_DEFAULT, _get_idle_unload_seconds,
_get_local_command_template, _has_local_command, _is_hallucinated_segment,
_join_confident_segments, _load_local_whisper_model, _looks_like_cuda_lib_error,
_normalize_local_model, _transcribe_local_command, _try_lazy_install_stt,
build_local_transcribe_kwargs,
)
build_local_transcribe_kwargs)
from tools.transcription_cloud import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
_extract_transcript_text, _has_xai_stt_credentials, _is_local_or_private_url,
_resolve_openai_audio_client_config, _transcribe_deepinfra, _transcribe_elevenlabs,
_transcribe_groq, _transcribe_mistral, _transcribe_openai, _transcribe_xai,
)
_transcribe_groq, _transcribe_mistral, _transcribe_openai, _transcribe_xai)
from tools.transcription_command import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
COMMAND_STT_OUTPUT_FORMATS, DEFAULT_COMMAND_STT_LANGUAGE, DEFAULT_COMMAND_STT_OUTPUT_FORMAT,
DEFAULT_COMMAND_STT_TIMEOUT_SECONDS, _PROMPT_CHARS_PER_TOKEN, _WHISPER_PROMPT_TOKEN_CAP,
_apply_pre_transcription_hook, _dispatch_to_plugin_provider, _enforce_prompt_length_limit,
_get_command_stt_output_format, _get_command_stt_timeout, _get_named_stt_provider_config,
_render_command_stt_template, _resolve_command_stt_provider_config, _run_command_stt,
_transcribe_command_stt, _unregistered_stt_provider_error,
)
_transcribe_command_stt, _unregistered_stt_provider_error)
logger = logging.getLogger(__name__)
@@ -245,15 +239,13 @@ _CLOUD_PROVIDER_SPECS = {
"No local STT available, using ElevenLabs Scribe STT API"),
"deepinfra": (_has_deepinfra_key, _has_deepinfra_key,
"STT provider 'deepinfra' configured but DEEPINFRA_API_KEY not set (or openai package missing)",
"No local STT available, using DeepInfra Whisper API"),
}
"No local STT available, using DeepInfra Whisper API")}
# Explicit selections whose resolution is more than a probe + warning.
_EXPLICIT_RESOLVERS = {
"local": _resolve_explicit_local,
"local_command": _resolve_explicit_local_command,
"openai": _resolve_explicit_openai,
}
"openai": _resolve_explicit_openai}
def _resolve_explicit_provider(provider: str) -> str:
@@ -429,8 +421,7 @@ def _read_block_error(file_path: str) -> Optional[Dict[str, Any]]:
def _transcribe_prepared_audio(
file_path: str, model: Optional[str] = None, source: Optional[str] = None
) -> Dict[str, Any]:
file_path: str, model: Optional[str] = None, source: Optional[str] = None) -> Dict[str, Any]:
"""Transcribe a validated audio file with the configured STT provider.
``model`` overrides the config default; ``source`` is a caller-surface label
@@ -477,8 +468,7 @@ _BUILTIN_MODEL_KEYS = {
"openai": ("openai", "model", DEFAULT_STT_MODEL, False),
"mistral": ("mistral", "model", DEFAULT_MISTRAL_STT_MODEL, False),
"elevenlabs": ("elevenlabs", "model_id", DEFAULT_ELEVENLABS_STT_MODEL, False),
"deepinfra": ("deepinfra", "model", "", True),
}
"deepinfra": ("deepinfra", "model", "", True)}
def _builtin_model_name(provider: str, stt_config: Dict[str, Any], model: Optional[str]) -> str:
@@ -494,8 +484,7 @@ def _builtin_model_name(provider: str, stt_config: Dict[str, Any], model: Option
def _dispatch_stt_provider(
file_path: str, provider: str, stt_config: Dict[str, Any], model: Optional[str] = None,
source: Optional[str] = None,
) -> Dict[str, Any]:
source: Optional[str] = None) -> Dict[str, Any]:
"""Route *file_path* to the handler for *provider* (built-in > command > plugin)."""
# Static ``stt.prompt`` is the base; hook results mutate on top (last hook to set a field wins).
prompt = stt_config.get("prompt")
@@ -522,8 +511,7 @@ def _dispatch_stt_provider(
# Plugin backend: reads ``stt.<provider>`` like built-ins; the ``model`` argument overrides it.
plugin_result = _dispatch_to_plugin_provider(
file_path, provider, stt_config, model=model or _get_stt_section(stt_config, provider).get("model"),
language=language or _resolve_stt_language(provider, stt_config), prompt=prompt,
)
language=language or _resolve_stt_language(provider, stt_config), prompt=prompt)
return plugin_result if plugin_result is not None else _no_provider_error(provider, stt_config)
@@ -543,13 +531,11 @@ def _no_provider_error(provider: str, stt_config: Dict[str, Any]) -> Dict[str, A
"set GROQ_API_KEY for free Groq Whisper, set MISTRAL_API_KEY for Mistral "
"Voxtral Transcribe, configure xAI OAuth or set XAI_API_KEY for xAI Grok STT, "
"set ELEVENLABS_API_KEY for ElevenLabs Scribe, or set VOICE_TOOLS_OPENAI_KEY "
"or OPENAI_API_KEY for the OpenAI Whisper API."
)
"or OPENAI_API_KEY for the OpenAI Whisper API.")
def transcribe_audio(
file_path: str, model: Optional[str] = None, source: Optional[str] = None
) -> Dict[str, Any]:
file_path: str, model: Optional[str] = None, source: Optional[str] = None) -> Dict[str, Any]:
"""Validate, preprocess supported inputs, and dispatch transcription.
``source`` is a caller-surface label (``"gateway"``, ``"voice_mode"``) forwarded to
+2 -4
View File
@@ -255,8 +255,7 @@ def _is_command_provider_config(config: Dict[str, Any]) -> bool:
def _resolve_command_config(
provider: str, config: Dict[str, Any], reserved: FrozenSet[str],
) -> Optional[Dict[str, Any]]:
provider: str, config: Dict[str, Any], reserved: FrozenSet[str]) -> Optional[Dict[str, Any]]:
"""Config of a user-declared command provider; None for *reserved* names, unknown or non-command."""
key = (provider or "").lower().strip()
if not key or key in reserved:
@@ -287,8 +286,7 @@ def _command_output_format(config: Dict[str, Any], formats: FrozenSet[str], defa
# Any ``tts.provider`` value NOT in this set refers to ``tts.providers.<name>``.
BUILTIN_TTS_PROVIDERS = frozenset({
"edge", "elevenlabs", "openai", "minimax", "xai", "mistral", "gemini",
"neutts", "kittentts", "piper", "deepinfra",
})
"neutts", "kittentts", "piper", "deepinfra"})
DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS = 120
DEFAULT_COMMAND_TTS_OUTPUT_FORMAT = "mp3"
+5 -11
View File
@@ -150,8 +150,7 @@ _PROVIDER_PRIORITY: List[str] = ["elevenlabs", "gemini", "openai", "xai"]
def resolve_streaming_provider(
tts_config: Dict, preferred: Optional[str] = None
) -> Optional[StreamingTTSProvider]:
tts_config: Dict, preferred: Optional[str] = None) -> Optional[StreamingTTSProvider]:
"""Return a ready streamer for the *configured* provider, else ``None``.
``tts.streaming.provider`` when set: a name pins that exact streamer (``None``
@@ -199,8 +198,7 @@ class ElevenLabsStreamer(StreamingTTSProvider):
text=text, voice_id=self.section.get("voice_id", DEFAULT_ELEVENLABS_VOICE_ID),
model_id=self.section.get("streaming_model_id",
self.section.get("model_id", DEFAULT_ELEVENLABS_STREAMING_MODEL_ID)),
output_format="pcm_24000",
)
output_format="pcm_24000")
def _openai_config_api_key() -> str:
@@ -225,8 +223,7 @@ class OpenAIStreamer(StreamingTTSProvider):
from openai import OpenAI
client = OpenAI(
api_key=(self.section.get("api_key") or resolve_openai_audio_api_key()),
base_url=(self.section.get("base_url") or get_env_value("OPENAI_BASE_URL") or None),
)
base_url=(self.section.get("base_url") or get_env_value("OPENAI_BASE_URL") or None))
with client.audio.speech.with_streaming_response.create(
model=self.section.get("model", "gpt-4o-mini-tts"), voice=self.section.get("voice", "alloy"),
input=text, response_format="pcm",
@@ -249,8 +246,7 @@ class GeminiStreamer(StreamingTTSProvider):
import json as _json
import requests
from tools.tts_tool_providers import (
DEFAULT_GEMINI_TTS_BASE_URL, DEFAULT_GEMINI_TTS_MODEL, DEFAULT_GEMINI_TTS_VOICE,
)
DEFAULT_GEMINI_TTS_BASE_URL, DEFAULT_GEMINI_TTS_MODEL, DEFAULT_GEMINI_TTS_VOICE)
api_key = _gemini_key()
model = str(self.section.get("model", DEFAULT_GEMINI_TTS_MODEL)).strip() or DEFAULT_GEMINI_TTS_MODEL
voice = str(self.section.get("voice", DEFAULT_GEMINI_TTS_VOICE)).strip() or DEFAULT_GEMINI_TTS_VOICE
@@ -261,9 +257,7 @@ class GeminiStreamer(StreamingTTSProvider):
"contents": [{"parts": [{"text": text}]}],
"generationConfig": {
"responseModalities": ["AUDIO"],
"speechConfig": {"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": voice}}},
},
}
"speechConfig": {"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": voice}}}}}
url = f"{base_url}/models/{model}:streamGenerateContent"
def _sse_chunks() -> Iterator[bytes]:
+4 -8
View File
@@ -35,8 +35,7 @@ _DEGREE_UNITS = (("C", "Celsius"), ("F", "Fahrenheit"))
# Unit suffix (regex, after a digit) -> spoken word; km/h variants before the bare "m".
_UNIT_WORDS = (
(r"km\s*/\s*h", "kilometres per hour"), (r"km/h", "kilometres per hour"),
(r"mm", "millimetres"), (r"cm", "centimetres"), (r"m", "metres"),
)
(r"mm", "millimetres"), (r"cm", "centimetres"), (r"m", "metres"))
# Currency prefix (regex) -> spoken word; order matters (NZ$/A$/US$ before bare $).
_CURRENCY_WORDS = (
(r"NZ\$", "New Zealand dollars", re.IGNORECASE), (r"A\$", "Australian dollars", re.IGNORECASE),
@@ -48,8 +47,7 @@ _EMOJI_RE = re.compile(
"[\U0001F1E6-\U0001F1FF\U0001F300-\U0001F5FF\U0001F600-\U0001F64F\U0001F680-\U0001F6FF"
"\U0001F700-\U0001F77F\U0001F780-\U0001F7FF\U0001F800-\U0001F8FF\U0001F900-\U0001F9FF"
"\U0001FA00-\U0001FAFF☀-➿]+",
flags=re.UNICODE,
)
flags=re.UNICODE)
_VARIATION_SELECTOR_RE = re.compile("[︎️]")
@@ -86,8 +84,7 @@ def _normalize_temperature_ranges(text: str) -> str:
lambda m, w=word: (
f"{m.group(1).replace(chr(0x2212), '-')} to {m.group(2).replace(chr(0x2212), '-')} degrees {w}"
),
text, flags=re.IGNORECASE,
)
text, flags=re.IGNORECASE)
return text
@@ -233,8 +230,7 @@ _LEGACY_TTS_STRIP_STEPS = (
(re.compile(r'---+'), ''),
# Emoji + variation selectors/ZWJ: providers speak them as awkward labels.
(re.compile('[\U0001F000-\U0001FAFF\u2600-\u27BF\uFE0F\u200D\U000E0020-\U000E007F]+'), ' '),
(re.compile(r'\n{3,}'), '\n\n'),
)
(re.compile(r'\n{3,}'), '\n\n'))
def _strip_markdown_for_tts(text: str) -> str: