feat(transcription): add AssemblyAI as transcription provider

Add AssemblyAI as a third transcription provider option alongside
OpenAI and Groq. AssemblyAI offers better accuracy for certain
audio types (distant voices, noisy environments) and serves as a
reliable fallback when other providers struggle.

Changes:
- Add AssemblyAITranscriptionProvider class in providers/transcription.py
- Add 'assemblyai' option in base channel's transcribe_audio()
- Per-channel configuration via transcriptionProvider in config

Usage:
  Set transcriptionProvider: 'assemblyai' and provide an AssemblyAI
  API key via transcriptionApiKey in the channel config.
This commit is contained in:
comadreja
2026-06-09 05:33:18 +08:00
committed by Xubin Ren
parent f183b37542
commit f3eb2aa08b
17 changed files with 780 additions and 113 deletions
+25 -56
View File
@@ -11,26 +11,20 @@ from __future__ import annotations
from contextlib import suppress
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Literal
from typing import Any
from loguru import logger
from nanobot.audio.transcription_registry import (
get_transcription_provider,
resolve_transcription_provider,
)
from nanobot.config.paths import get_media_dir
from nanobot.utils.media_decode import FileSizeExceeded, save_base64_data_url
TranscriptionProviderName = Literal["groq", "openai", "openrouter", "xiaomi_mimo"]
TranscriptionProviderName = str
_DEFAULT_PROVIDER: TranscriptionProviderName = "groq"
_DEFAULT_MODELS: dict[TranscriptionProviderName, str] = {
"groq": "whisper-large-v3",
"openai": "whisper-1",
"openrouter": "openai/whisper-1",
"xiaomi_mimo": "mimo-v2.5-asr",
}
_PROVIDER_ALIASES: dict[str, TranscriptionProviderName] = {
"mimo": "xiaomi_mimo",
"xiaomi": "xiaomi_mimo",
}
_MAX_AUDIO_BYTES_FALLBACK = 25 * 1024 * 1024
_AUDIO_MIME_ALLOWED: frozenset[str] = frozenset({
"audio/aac",
@@ -72,13 +66,8 @@ class TranscriptionIngressError(Exception):
def _as_provider(value: Any) -> TranscriptionProviderName | None:
if isinstance(value, str):
name = value.strip().lower()
if name in _PROVIDER_ALIASES:
return _PROVIDER_ALIASES[name]
if name in _DEFAULT_MODELS:
return name # type: ignore[return-value]
return None
spec = resolve_transcription_provider(value)
return spec.name if spec else None
def _provider_config(config: Any, provider: str) -> Any:
@@ -101,11 +90,17 @@ def resolve_transcription_config(config: Any) -> EffectiveTranscriptionConfig:
or _as_provider(getattr(channels, "transcription_provider", None))
or _DEFAULT_PROVIDER
)
spec = get_transcription_provider(provider)
if spec is None:
logger.warning("Unknown transcription provider {}; falling back to {}", provider, _DEFAULT_PROVIDER)
provider = _DEFAULT_PROVIDER
spec = get_transcription_provider(provider)
default_model = spec.default_model if spec else ""
provider_cfg = _provider_config(config, provider)
return EffectiveTranscriptionConfig(
enabled=bool(getattr(top, "enabled", True)),
provider=provider,
model=(getattr(top, "model", None) or _DEFAULT_MODELS[provider]).strip(),
model=(getattr(top, "model", None) or default_model).strip(),
language=getattr(top, "language", None) or getattr(channels, "transcription_language", None),
api_key=getattr(provider_cfg, "api_key", None) or "",
api_base=getattr(provider_cfg, "api_base", None) or "",
@@ -170,40 +165,14 @@ async def transcribe_audio_file(
"""Transcribe *file_path* using the already-resolved transcription config."""
if not config.enabled or not config.configured:
return ""
if config.provider == "openai":
from nanobot.providers.transcription import OpenAITranscriptionProvider
provider = OpenAITranscriptionProvider(
api_key=config.api_key,
api_base=config.api_base or None,
language=config.language,
model=config.model,
)
elif config.provider == "openrouter":
from nanobot.providers.transcription import OpenRouterTranscriptionProvider
provider = OpenRouterTranscriptionProvider(
api_key=config.api_key,
api_base=config.api_base or None,
language=config.language,
model=config.model,
)
elif config.provider == "xiaomi_mimo":
from nanobot.providers.transcription import XiaomiMiMoTranscriptionProvider
provider = XiaomiMiMoTranscriptionProvider(
api_key=config.api_key,
api_base=config.api_base or None,
language=config.language,
model=config.model,
)
else:
from nanobot.providers.transcription import GroqTranscriptionProvider
provider = GroqTranscriptionProvider(
api_key=config.api_key,
api_base=config.api_base or None,
language=config.language,
model=config.model,
)
spec = get_transcription_provider(config.provider)
if spec is None:
logger.warning("Unknown transcription provider: {}", config.provider)
return ""
provider = spec.load_adapter()(
api_key=config.api_key,
api_base=config.api_base or None,
language=config.language,
model=config.model,
)
return await provider.transcribe(file_path)
+90
View File
@@ -0,0 +1,90 @@
"""Registry for speech-to-text providers.
Provider-specific HTTP adapters live in ``nanobot.providers.transcription``.
This module is the app-level source of truth for provider names, aliases,
default models, and adapter class paths.
"""
from __future__ import annotations
from dataclasses import dataclass
from importlib import import_module
from pathlib import Path
from typing import Any, Protocol
class TranscriptionProviderAdapter(Protocol):
"""Runtime protocol implemented by provider-specific transcription adapters."""
def __init__(
self,
api_key: str | None = None,
api_base: str | None = None,
language: str | None = None,
model: str | None = None,
) -> None: ...
async def transcribe(self, file_path: str | Path) -> str: ...
@dataclass(frozen=True)
class TranscriptionProviderSpec:
name: str
default_model: str
adapter: str
aliases: tuple[str, ...] = ()
def load_adapter(self) -> type[TranscriptionProviderAdapter]:
module_name, _, class_name = self.adapter.partition(":")
if not module_name or not class_name:
raise RuntimeError(f"Invalid transcription adapter path: {self.adapter}")
adapter = getattr(import_module(module_name), class_name)
return adapter
TRANSCRIPTION_PROVIDERS: tuple[TranscriptionProviderSpec, ...] = (
TranscriptionProviderSpec(
name="groq",
default_model="whisper-large-v3",
adapter="nanobot.providers.transcription:GroqTranscriptionProvider",
),
TranscriptionProviderSpec(
name="openai",
default_model="whisper-1",
adapter="nanobot.providers.transcription:OpenAITranscriptionProvider",
),
TranscriptionProviderSpec(
name="openrouter",
default_model="openai/whisper-1",
adapter="nanobot.providers.transcription:OpenRouterTranscriptionProvider",
),
TranscriptionProviderSpec(
name="xiaomi_mimo",
default_model="mimo-v2.5-asr",
adapter="nanobot.providers.transcription:XiaomiMiMoTranscriptionProvider",
aliases=("mimo", "xiaomi"),
),
TranscriptionProviderSpec(
name="assemblyai",
default_model="universal-3-pro,universal-2",
adapter="nanobot.providers.transcription:AssemblyAITranscriptionProvider",
),
)
_BY_NAME = {spec.name: spec for spec in TRANSCRIPTION_PROVIDERS}
_BY_ALIAS = {alias: spec for spec in TRANSCRIPTION_PROVIDERS for alias in spec.aliases}
def transcription_provider_names() -> tuple[str, ...]:
return tuple(spec.name for spec in TRANSCRIPTION_PROVIDERS)
def get_transcription_provider(name: str) -> TranscriptionProviderSpec | None:
return _BY_NAME.get(name)
def resolve_transcription_provider(value: Any) -> TranscriptionProviderSpec | None:
if not isinstance(value, str):
return None
name = value.strip().lower()
return _BY_NAME.get(name) or _BY_ALIAS.get(name)
+7 -2
View File
@@ -47,7 +47,7 @@ class TranscriptionConfig(Base):
"""Cross-channel audio transcription configuration."""
enabled: bool = True
provider: Literal["groq", "openai", "openrouter", "xiaomi_mimo"] | None = None
provider: str | None = None # Validated by nanobot.audio.transcription_registry.
model: str | None = None
language: str | None = Field(default=None, pattern=r"^[a-z]{2,3}$")
max_duration_sec: int = Field(default=120, ge=1, le=600)
@@ -202,6 +202,7 @@ class ProvidersConfig(Base):
anthropic: ProviderConfig = Field(default_factory=ProviderConfig)
openai: ProviderConfig = Field(default_factory=ProviderConfig)
openrouter: ProviderConfig = Field(default_factory=ProviderConfig)
assemblyai: ProviderConfig = Field(default_factory=ProviderConfig) # AssemblyAI voice transcription
huggingface: ProviderConfig = Field(default_factory=ProviderConfig)
skywork: ProviderConfig = Field(default_factory=ProviderConfig) # Skywork / APIFree API gateway
deepseek: ProviderConfig = Field(default_factory=ProviderConfig)
@@ -402,6 +403,8 @@ class Config(BaseSettings):
# Explicit provider prefix wins — prevents `github-copilot/...codex` matching openai_codex.
for spec in PROVIDERS:
if spec.is_transcription_only:
continue
p = getattr(self.providers, spec.name, None)
if p and model_prefix and normalized_prefix == spec.name:
if spec.is_oauth or spec.is_local or spec.is_direct or p.api_key:
@@ -409,6 +412,8 @@ class Config(BaseSettings):
# Match by keyword (order follows PROVIDERS registry)
for spec in PROVIDERS:
if spec.is_transcription_only:
continue
p = getattr(self.providers, spec.name, None)
if p and any(_kw_matches(kw) for kw in spec.keywords):
if spec.is_oauth or spec.is_local or spec.is_direct or p.api_key:
@@ -435,7 +440,7 @@ class Config(BaseSettings):
# Fallback: gateways first, then others (follows registry order)
# OAuth providers are NOT valid fallbacks — they require explicit model selection
for spec in PROVIDERS:
if spec.is_oauth:
if spec.is_oauth or spec.is_transcription_only:
continue
p = getattr(self.providers, spec.name, None)
if p and p.api_key:
+2
View File
@@ -41,6 +41,8 @@ def _make_provider_core(
provider_name = config.get_provider_name(model, preset=resolved)
p = config.get_provider(model, preset=resolved)
spec = find_by_name(provider_name) if provider_name else None
if spec and spec.is_transcription_only:
raise ValueError(f"Provider '{provider_name}' only supports transcription.")
backend = spec.backend if spec else "openai_compat"
if backend == "azure_openai":
+14
View File
@@ -60,6 +60,9 @@ class ProviderSpec:
# Direct providers skip API-key validation (user supplies everything)
is_direct: bool = False
# Provider is listed for shared credentials but cannot serve chat completions.
is_transcription_only: bool = False
# Provider supports cache_control on content blocks (e.g. Anthropic prompt caching)
supports_prompt_caching: bool = False
@@ -507,6 +510,17 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
backend="openai_compat",
default_api_base="https://api.groq.com/openai/v1",
),
# AssemblyAI: voice transcription only. It appears in provider settings so
# users can manage credentials, but WebUI excludes it from chat model pickers.
ProviderSpec(
name="assemblyai",
keywords=("assemblyai",),
env_key="ASSEMBLYAI_API_KEY",
display_name="AssemblyAI",
backend="openai_compat",
default_api_base="https://api.assemblyai.com/v2",
is_transcription_only=True,
),
# Qianfan (百度千帆): OpenAI-compatible API
ProviderSpec(
name="qianfan",
+194 -1
View File
@@ -1,7 +1,7 @@
"""Provider-specific voice transcription adapters.
This module only knows how to call external transcription APIs such as Groq,
OpenAI Whisper, OpenRouter, and Xiaomi MiMo ASR. Product-level config fallback,
OpenAI Whisper, OpenRouter, Xiaomi MiMo ASR, and AssemblyAI. Product-level config fallback,
WebUI upload validation, and channel integration live in
``nanobot.audio.transcription``.
"""
@@ -19,6 +19,9 @@ from loguru import logger
_CHAT_COMPLETIONS_PATH = "chat/completions"
_TRANSCRIPTIONS_PATH = "audio/transcriptions"
_ASSEMBLYAI_DEFAULT_API_BASE = "https://api.assemblyai.com/v2"
_ASSEMBLYAI_POLL_ATTEMPTS = 60
_ASSEMBLYAI_POLL_INTERVAL_S = 2.0
_AUDIO_MIME_OVERRIDES = {
".m4a": "audio/mp4",
".mpga": "audio/mpeg",
@@ -63,6 +66,11 @@ def _resolve_chat_completions_url(api_base: str | None, default_url: str) -> str
return f"{base}/{_CHAT_COMPLETIONS_PATH}"
def _resolve_api_path(api_base: str | None, default_base: str, path: str) -> str:
base = (api_base or default_base).rstrip("/")
return f"{base}/{path.lstrip('/')}"
def _audio_mime_type(path: Path) -> str:
return (
_AUDIO_MIME_OVERRIDES.get(path.suffix.lower())
@@ -93,6 +101,90 @@ _RETRYABLE_EXCEPTIONS = (
)
async def _request_json_with_retry(
client: httpx.AsyncClient,
method: str,
url: str,
*,
provider_label: str,
**kwargs: object,
) -> dict[str, Any] | None:
for attempt in range(_MAX_RETRIES + 1):
try:
request = getattr(client, method.lower(), None)
if request is None:
response = await client.request(method, url, **kwargs)
else:
response = await request(url, **kwargs)
except _RETRYABLE_EXCEPTIONS as e:
if attempt < _MAX_RETRIES:
logger.warning(
"{} transcription transient error (attempt {}/{}): {}",
provider_label,
attempt + 1,
_MAX_RETRIES + 1,
e,
)
await asyncio.sleep(_BACKOFF_S[attempt])
continue
logger.exception(
"{} transcription error after {} attempts: {}",
provider_label,
_MAX_RETRIES + 1,
e,
)
return None
except Exception as e:
logger.exception("{} transcription error: {}", provider_label, e)
return None
if response.status_code in _RETRYABLE_STATUS and attempt < _MAX_RETRIES:
logger.warning(
"{} transcription transient HTTP {} (attempt {}/{})",
provider_label,
response.status_code,
attempt + 1,
_MAX_RETRIES + 1,
)
await asyncio.sleep(_BACKOFF_S[attempt])
continue
try:
response.raise_for_status()
except httpx.HTTPStatusError:
body = response.text.strip().replace("\n", " ")[:500]
logger.error(
"{} transcription HTTP {}{}{}",
provider_label,
response.status_code,
f" {response.reason_phrase}" if response.reason_phrase else "",
f": {body}" if body else "",
)
return None
except Exception as e:
logger.exception("{} transcription error: {}", provider_label, e)
return None
try:
payload = response.json()
except Exception as e:
logger.exception(
"{} transcription error: malformed response body: {}",
provider_label,
e,
)
return None
if not isinstance(payload, dict):
logger.error(
"{} transcription error: unexpected response shape: {!r}",
provider_label,
type(payload).__name__,
)
return None
return payload
return None
async def _post_transcription_with_retry(
url: str,
*,
@@ -305,6 +397,107 @@ def _text_from_chat_payload(payload: dict[str, Any]) -> str:
return text if isinstance(text, str) else ""
def _assemblyai_speech_models(model: str | None) -> list[str]:
return [part for part in (part.strip() for part in (model or "").split(",")) if part]
class AssemblyAITranscriptionProvider:
"""Voice transcription provider using AssemblyAI's asynchronous REST API."""
def __init__(
self,
api_key: str | None = None,
api_base: str | None = None,
language: str | None = None,
model: str | None = None,
):
base = api_base or os.environ.get("ASSEMBLYAI_BASE_URL")
self.api_key = api_key or os.environ.get("ASSEMBLYAI_API_KEY")
self.upload_url = _resolve_api_path(base, _ASSEMBLYAI_DEFAULT_API_BASE, "upload")
self.transcript_url = _resolve_api_path(base, _ASSEMBLYAI_DEFAULT_API_BASE, "transcript")
self.language = language or None
self.model = model or "universal-3-pro,universal-2"
logger.debug("AssemblyAI transcription endpoint: {}", self.transcript_url)
async def transcribe(self, file_path: str | Path) -> str:
if not self.api_key:
logger.warning("AssemblyAI API key not configured for transcription")
return ""
path = Path(file_path)
if not path.exists():
logger.error("Audio file not found: {}", file_path)
return ""
try:
data = path.read_bytes()
except OSError as e:
logger.exception("AssemblyAI transcription error: cannot read audio file: {}", e)
return ""
headers = {"Authorization": self.api_key}
async with httpx.AsyncClient() as client:
upload = await _request_json_with_retry(
client,
"POST",
self.upload_url,
provider_label="AssemblyAI",
headers={**headers, "Content-Type": "application/octet-stream"},
content=data,
timeout=60.0,
)
upload_url = upload.get("upload_url") if upload else None
if not isinstance(upload_url, str) or not upload_url:
logger.error("AssemblyAI transcription error: upload_url missing")
return ""
body: dict[str, object] = {"audio_url": upload_url}
speech_models = _assemblyai_speech_models(self.model)
if speech_models:
body["speech_models"] = speech_models
if self.language:
body["language_code"] = self.language
transcript = await _request_json_with_retry(
client,
"POST",
self.transcript_url,
provider_label="AssemblyAI",
headers=headers,
json=body,
timeout=30.0,
)
transcript_id = transcript.get("id") if transcript else None
if not isinstance(transcript_id, str) or not transcript_id:
logger.error("AssemblyAI transcription error: transcript id missing")
return ""
poll_url = f"{self.transcript_url.rstrip('/')}/{transcript_id}"
for attempt in range(_ASSEMBLYAI_POLL_ATTEMPTS):
payload = await _request_json_with_retry(
client,
"GET",
poll_url,
provider_label="AssemblyAI",
headers=headers,
timeout=30.0,
)
if not payload:
return ""
status = str(payload.get("status") or "").lower()
if status == "completed":
text = payload.get("text")
return text if isinstance(text, str) else ""
if status in {"error", "failed"}:
logger.error(
"AssemblyAI transcription failed: {}",
payload.get("error") or payload,
)
return ""
if attempt < _ASSEMBLYAI_POLL_ATTEMPTS - 1:
await asyncio.sleep(_ASSEMBLYAI_POLL_INTERVAL_S)
logger.error("AssemblyAI transcription timed out while polling transcript")
return ""
class OpenAITranscriptionProvider:
"""Voice transcription provider using OpenAI's Whisper API."""
+19 -7
View File
@@ -16,6 +16,10 @@ from zoneinfo import ZoneInfo
import httpx
from nanobot.audio.transcription import resolve_transcription_config
from nanobot.audio.transcription_registry import (
resolve_transcription_provider,
transcription_provider_names,
)
from nanobot.config.loader import get_config_path, load_config, save_config
from nanobot.config.schema import ModelPresetConfig
from nanobot.providers.image_generation import (
@@ -91,7 +95,6 @@ _IMAGE_GENERATION_ASPECT_RATIOS = {
"2:3",
"21:9",
}
_TRANSCRIPTION_PROVIDERS = ("groq", "openai", "openrouter", "xiaomi_mimo")
_CONTEXT_WINDOW_TOKEN_OPTIONS = {65_536, 262_144}
_MODEL_CONFIGURATION_SLUG_RE = re.compile(r"[^a-z0-9_-]+")
_ENV_REF_RE = re.compile(r"\$\{([A-Za-z_][A-Za-z0-9_]*)\}")
@@ -424,9 +427,13 @@ def provider_models_payload(query: QueryParams) -> dict[str, Any]:
"fetched_at": time.time(),
}
if (
spec.backend in _MODEL_LIST_UNSUPPORTED_BACKENDS
and spec.name != "minimax_anthropic"
) or spec.is_oauth:
spec.is_transcription_only
or (
spec.backend in _MODEL_LIST_UNSUPPORTED_BACKENDS
and spec.name != "minimax_anthropic"
)
or spec.is_oauth
):
return {
**base_payload,
"status": "unsupported",
@@ -542,6 +549,8 @@ def _validate_configured_provider(config: Any, provider: str) -> None:
spec = find_by_name(provider)
if spec is None:
raise WebUISettingsError("unknown provider")
if spec.is_transcription_only:
raise WebUISettingsError("provider does not support chat models")
provider_config = getattr(config.providers, provider, None)
if (
provider_config is None
@@ -580,7 +589,7 @@ def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]:
def _transcription_provider_rows(config: Any) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
for name in _TRANSCRIPTION_PROVIDERS:
for name in transcription_provider_names():
spec = find_by_name(name)
provider_config = getattr(config.providers, name, None)
rows.append({
@@ -640,6 +649,7 @@ def settings_payload(
"api_key_hint": _mask_secret_hint(provider_config.api_key),
"api_base": provider_config.api_base,
"default_api_base": spec.default_api_base or None,
"model_selectable": not spec.is_transcription_only,
}
if oauth_status is not None:
row["oauth_account"] = oauth_status["account"]
@@ -1357,10 +1367,12 @@ def update_transcription_settings(query: QueryParams) -> dict[str, Any]:
provider = _query_first(query, "provider")
if provider is not None:
provider = provider.strip().lower()
if provider not in _TRANSCRIPTION_PROVIDERS:
provider_spec = resolve_transcription_provider(provider)
if provider_spec is None:
raise WebUISettingsError("unknown transcription provider")
provider = provider_spec.name
if transcription.provider != provider:
transcription.provider = provider # type: ignore[assignment]
transcription.provider = provider
changed = True
model = _query_first(query, "model")