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:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user