refactor(tools): fold cloud STT request builders — single-use locals, hoisted ASR regex, merged with-blocks

This commit is contained in:
Teknium
2026-09-03 00:46:20 -07:00
parent f09b63a26e
commit 35c32f6498

View File

@@ -26,6 +26,10 @@ from tools.transcription_common import (
# Log-record parity with the origin module.
logger = logging.getLogger("tools.transcription_tools")
# Voxtral-style ``language xx <asr_text> ...`` prefix some SDK responses carry.
_ASR_TEXT_RE = re.compile(
r"\s*language\s+[\w.-]+(?:\s*<audio_language>[^<]*</audio_language>)?\s*<asr_text>\s*(?P<text>.*)",
flags=re.IGNORECASE | re.DOTALL)
def _has_xai_stt_credentials() -> bool:
@@ -92,9 +96,9 @@ def _transcribe_groq(
language = language or _resolve_stt_language("groq")
def _run(client):
create_kwargs = {"model": model_name, "response_format": "text", **_sdk_prompt_kwargs(language, prompt)}
with open(file_path, "rb") as audio_file:
transcription = client.audio.transcriptions.create(file=audio_file, **create_kwargs)
transcription = client.audio.transcriptions.create(
file=audio_file, model=model_name, response_format="text", **_sdk_prompt_kwargs(language, prompt))
transcript_text = str(transcription).strip()
logger.info("Transcribed %s via Groq API (%s, lang=%s, %d chars)",
Path(file_path).name, model_name, language or "auto", len(transcript_text))
@@ -137,8 +141,8 @@ def _transcribe_openai(
"model": model_name, "response_format": "text" if model_name == "whisper-1" else "json",
}
if language:
# gpt-transcribe takes a ``languages`` list and rejects the legacy field.
if model_name == "gpt-transcribe":
# gpt-transcribe takes a ``languages`` list and rejects the legacy field.
create_kwargs["extra_body"] = {"languages": [language]}
else:
create_kwargs["language"] = language
@@ -152,8 +156,7 @@ def _transcribe_openai(
try:
transcription = _create_transcription(file_path)
except BadRequestError as exc:
message = str(exc).lower()
if not any(k in message for k in ("unsupported", "corrupted", "invalid file")):
if not any(k in str(exc).lower() for k in ("unsupported", "corrupted", "invalid file")):
raise
# Newer models reject containers whisper-1 accepted (Ogg/Opus voice notes): transcode, retry once.
converted_path, transcode_error = _transcode_audio_for_stt(file_path, work_dir)
@@ -181,18 +184,17 @@ def _transcribe_mistral(
try:
_lazy_ensure_quietly("stt.mistral")
from mistralai.client import Mistral
with Mistral(api_key=api_key) as client:
with open(file_path, "rb") as audio_file:
# Language: hook override > stt.mistral.language > stt.language > env > auto.
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),
)
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))
return _ok_result(transcript_text, "mistral")
with Mistral(api_key=api_key) as client, open(file_path, "rb") as audio_file:
# Language: hook override > stt.mistral.language > stt.language > env > auto.
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),
)
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))
return _ok_result(transcript_text, "mistral")
except Exception as e:
return _cloud_failure(e, file_path, "Mistral transcription", type(e).__name__)
@@ -245,11 +247,9 @@ def _transcribe_xai(
# STT is API-billed: prefer the explicit XAI_API_KEY over the xAI OAuth/Grok-subscription
# credential, which may be valid for Grok yet hit spending-limit errors on /v1/stt.
direct_api_key = str(get_env_value("XAI_API_KEY") or "").strip()
if direct_api_key:
creds = {"provider": "xai", "api_key": direct_api_key,
"base_url": str(get_env_value("XAI_BASE_URL") or "https://api.x.ai/v1").strip().rstrip("/")}
else:
creds = resolve_xai_http_credentials()
creds = {"provider": "xai", "api_key": direct_api_key,
"base_url": str(get_env_value("XAI_BASE_URL") or "https://api.x.ai/v1").strip().rstrip("/")
} if direct_api_key else resolve_xai_http_credentials()
api_key = str(creds.get("api_key") or "").strip()
if not api_key:
return _error_result("No xAI credentials found. Configure xAI OAuth in `hermes model` or set XAI_API_KEY")
@@ -258,10 +258,9 @@ def _transcribe_xai(
def _resolve_base_url(resolved_creds: Dict[str, str]) -> str:
# OAuth bearers are pinned to the resolver-validated origin; overrides apply to API keys only.
if resolved_creds.get("provider") == "xai-oauth":
url = resolved_creds.get("base_url")
else:
url = xai_config.get("base_url") or get_env_value("XAI_STT_BASE_URL") or resolved_creds.get("base_url")
url = resolved_creds.get("base_url")
if resolved_creds.get("provider") != "xai-oauth":
url = xai_config.get("base_url") or get_env_value("XAI_STT_BASE_URL") or url
return str(url or XAI_STT_BASE_URL).strip().rstrip("/")
# Language: hook override > stt.xai.language > stt.language > env.
@@ -270,9 +269,8 @@ def _transcribe_xai(
def _post() -> Any:
from tools.xai_http import hermes_xai_user_agent
data: Dict[str, str] = {"language": language} if language else {}
for flag, default in (("format", True), ("diarize", False)):
if is_truthy_value(xai_config.get(flag, default)):
data[flag] = "true"
data.update({flag: "true" for flag, default in (("format", True), ("diarize", False))
if is_truthy_value(xai_config.get(flag, default))})
def _post_transcription(bearer: str, endpoint_base_url: str):
headers = {"Authorization": f"Bearer {bearer}", "User-Agent": hermes_xai_user_agent()}
@@ -329,9 +327,8 @@ 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 {}),
}
if language_code:
data["language_code"] = language_code
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:
@@ -353,14 +350,12 @@ def _transcribe_deepinfra(
from hermes_cli.models import deepinfra_base_url, deepinfra_model_ids
# ``stt.deepinfra: null`` in YAML yields None, not {} — coalesce.
base_url = deepinfra_base_url(_get_stt_section(_load_stt_config(), "deepinfra"))
model_name = model_name or next(iter(deepinfra_model_ids("stt")), None)
if not model_name:
candidates = deepinfra_model_ids("stt")
if not candidates:
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."
)
model_name = candidates[0]
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."
)
return _transcribe_openai(file_path, model_name, api_key=api_key, base_url=base_url,
provider_label="deepinfra", language=language, prompt=prompt)
@@ -374,11 +369,9 @@ def _is_local_or_private_url(url: str) -> bool:
from urllib.parse import urlparse
import ipaddress
host = (urlparse(url).hostname or "").lower()
if not host:
return False
if host == "localhost" or host.endswith((".local", ".lan", ".internal")):
return True
addr = ipaddress.ip_address(host)
addr = ipaddress.ip_address(host) # raises for "" and non-IP hostnames
return addr.is_private or addr.is_loopback
except Exception: # unparsable URL or non-IP hostname
return False
@@ -419,18 +412,14 @@ def _resolve_openai_audio_client_config() -> tuple[str, str]:
managed = _managed()
if managed is None:
raise ValueError(selection_error(
"stt", NOUS_MANAGED_PROVIDER,
"the Nous Tool Gateway is not available (not entitled or unreachable)",
))
"stt", NOUS_MANAGED_PROVIDER, "the Nous Tool Gateway is not available (not entitled or unreachable)"))
return managed
direct = _direct_openai_credentials(openai_cfg.get("api_key", ""), openai_cfg.get("base_url", ""))
if direct is not None:
return direct
if selected is not None:
raise ValueError(selection_error(
"stt", selected,
"neither stt.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set",
))
"stt", selected, "neither stt.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set"))
managed = _managed()
if managed is None:
message = "Neither stt.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set"
@@ -446,8 +435,5 @@ def _extract_transcript_text(transcription: Any) -> str:
if not isinstance(value, str) and isinstance(transcription, dict):
value = transcription.get("text")
text = (value if isinstance(value, str) else str(transcription)).strip()
match = re.match(
r"\s*language\s+[\w.-]+(?:\s*<audio_language>[^<]*</audio_language>)?\s*<asr_text>\s*(?P<text>.*)",
text, flags=re.IGNORECASE | re.DOTALL,
)
match = _ASR_TEXT_RE.match(text)
return match.group("text").strip() if match else text