refactor(tools): fold cloud STT request builders — single-use locals, hoisted ASR regex, merged with-blocks
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user