From cf0af1f50cb7e1e3ffda6a6130b76b9e7d90a8f8 Mon Sep 17 00:00:00 2001 From: seymourtang Date: Tue, 11 Aug 2026 17:55:25 +0800 Subject: [PATCH] refactor(agentkit): expose vendor params to static type checkers --- src/agora_agent/agentkit/vendors/avatar.py | 25 ++++-- src/agora_agent/agentkit/vendors/cn.py | 95 ++++++++++++++++------ src/agora_agent/agentkit/vendors/llm.py | 45 ++++++---- src/agora_agent/agentkit/vendors/mllm.py | 27 ++++-- src/agora_agent/agentkit/vendors/stt.py | 45 +++++++--- src/agora_agent/agentkit/vendors/tts.py | 83 +++++++++++++------ 6 files changed, 228 insertions(+), 92 deletions(-) diff --git a/src/agora_agent/agentkit/vendors/avatar.py b/src/agora_agent/agentkit/vendors/avatar.py index 23a709e..21618c3 100644 --- a/src/agora_agent/agentkit/vendors/avatar.py +++ b/src/agora_agent/agentkit/vendors/avatar.py @@ -1,16 +1,15 @@ import warnings from typing import Any, Dict, Optional -from pydantic import ConfigDict, Field, field_validator - from .base import BaseAvatar +from pydantic import BaseModel, ConfigDict, Field, field_validator LIVEAVATAR_SAMPLE_RATE = 24000 HEYGEN_SAMPLE_RATE = LIVEAVATAR_SAMPLE_RATE AKOOL_SAMPLE_RATE = 16000 -class LiveAvatarAvatar(BaseAvatar): +class LiveAvatarAvatarOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="LiveAvatar API key") @@ -31,6 +30,8 @@ def validate_quality(cls, v: str) -> str: raise ValueError(f"Invalid quality '{v}'. Must be one of: {', '.join(valid)}") return v + +class LiveAvatarAvatar(LiveAvatarAvatarOptions, BaseAvatar): @property def required_sample_rate(self) -> int: return LIVEAVATAR_SAMPLE_RATE @@ -57,7 +58,7 @@ def to_config(self) -> Dict[str, Any]: return {"enable": enable, "vendor": "liveavatar", "params": params} -class HeyGenAvatar(BaseAvatar): +class HeyGenAvatarOptions(BaseModel): """Deprecated: HeyGen has been renamed to LiveAvatar. Use LiveAvatarAvatar instead.""" model_config = ConfigDict(extra="forbid") @@ -89,6 +90,10 @@ def model_post_init(self, __context: Any) -> None: stacklevel=3, ) + +class HeyGenAvatar(HeyGenAvatarOptions, BaseAvatar): + """Deprecated: HeyGen has been renamed to LiveAvatar. Use LiveAvatarAvatar instead.""" + @property def required_sample_rate(self) -> int: return HEYGEN_SAMPLE_RATE @@ -115,7 +120,7 @@ def to_config(self) -> Dict[str, Any]: return {"enable": enable, "vendor": "heygen", "params": params} -class AkoolAvatar(BaseAvatar): +class AkoolAvatarOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Akool API key") @@ -123,6 +128,8 @@ class AkoolAvatar(BaseAvatar): enable: Optional[bool] = Field(default=None, description="Enable avatar (default: true)") additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional vendor-specific parameters") + +class AkoolAvatar(AkoolAvatarOptions, BaseAvatar): @property def required_sample_rate(self) -> int: return AKOOL_SAMPLE_RATE @@ -141,7 +148,7 @@ def to_config(self) -> Dict[str, Any]: return {"enable": enable, "vendor": "akool", "params": params} -class GenericAvatar(BaseAvatar): +class GenericAvatarOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Generic avatar provider API key") @@ -154,6 +161,8 @@ class GenericAvatar(BaseAvatar): enable: Optional[bool] = Field(default=None, description="Enable avatar (default: true)") additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional vendor-specific parameters") + +class GenericAvatar(GenericAvatarOptions, BaseAvatar): @property def required_sample_rate(self) -> int: return 0 @@ -179,7 +188,7 @@ def to_config(self) -> Dict[str, Any]: return {"enable": enable, "vendor": "generic", "params": params} -class AnamAvatar(BaseAvatar): +class AnamAvatarOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Anam API key") @@ -187,6 +196,8 @@ class AnamAvatar(BaseAvatar): enable: Optional[bool] = Field(default=None, description="Enable avatar (default: true)") additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional vendor-specific parameters") + +class AnamAvatar(AnamAvatarOptions, BaseAvatar): @property def required_sample_rate(self) -> int: return 0 diff --git a/src/agora_agent/agentkit/vendors/cn.py b/src/agora_agent/agentkit/vendors/cn.py index e336136..0d1c9c7 100644 --- a/src/agora_agent/agentkit/vendors/cn.py +++ b/src/agora_agent/agentkit/vendors/cn.py @@ -2,8 +2,6 @@ from typing import Any, Dict, List, Optional -from pydantic import ConfigDict, Field, model_validator - from ...types.mllm_turn_detection import MllmTurnDetection from .avatar import BaseAvatar from .base import BaseLLM, BaseMLLM @@ -15,9 +13,10 @@ ) from .stt import BaseSTT as _BaseSTTCompat from .tts import BaseTTS as _BaseTTSCompat +from pydantic import BaseModel, ConfigDict, Field, model_validator -class TencentSTT(_BaseSTTCompat): +class TencentSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Tencent ASR secret key") @@ -27,6 +26,8 @@ class TencentSTT(_BaseSTTCompat): voice_id: str = Field(..., description="Tencent ASR voice id") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class TencentSTT(TencentSTTOptions, _BaseSTTCompat): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update( @@ -41,11 +42,13 @@ def to_config(self) -> Dict[str, Any]: return {"vendor": "tencent", "params": params} -class FengmingSTT(_BaseSTTCompat): +class FengmingSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") keywords: Optional[List[str]] = Field(default=None, description="Hotwords that improve ASR accuracy") + +class FengmingSTT(FengmingSTTOptions, _BaseSTTCompat): def to_config(self) -> Dict[str, Any]: config: Dict[str, Any] = {"vendor": "fengming"} if self.keywords is not None: @@ -53,7 +56,7 @@ def to_config(self) -> Dict[str, Any]: return config -class XfyunSTT(_BaseSTTCompat): +class XfyunSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="Xfyun ASR API key") @@ -62,6 +65,8 @@ class XfyunSTT(_BaseSTTCompat): language: Optional[str] = Field(default=None, description="Xfyun ASR language") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class XfyunSTT(XfyunSTTOptions, _BaseSTTCompat): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) if self.api_key is not None: @@ -78,7 +83,7 @@ def to_config(self) -> Dict[str, Any]: } -class XfyunBigModelSTT(_BaseSTTCompat): +class XfyunBigModelSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="Xfyun BigModel ASR API key") @@ -88,6 +93,8 @@ class XfyunBigModelSTT(_BaseSTTCompat): language: Optional[str] = Field(default=None, description="Xfyun BigModel ASR language") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class XfyunBigModelSTT(XfyunBigModelSTTOptions, _BaseSTTCompat): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) if self.api_key is not None: @@ -106,7 +113,7 @@ def to_config(self) -> Dict[str, Any]: } -class XfyunDialectSTT(_BaseSTTCompat): +class XfyunDialectSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") app_id: Optional[str] = Field(default=None, description="Xfyun Dialect ASR app id") @@ -115,6 +122,8 @@ class XfyunDialectSTT(_BaseSTTCompat): language: Optional[str] = Field(default=None, description="Xfyun Dialect ASR language") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class XfyunDialectSTT(XfyunDialectSTTOptions, _BaseSTTCompat): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) if self.app_id is not None: @@ -131,7 +140,7 @@ def to_config(self) -> Dict[str, Any]: } -class MicrosoftSTT(_BaseSTTCompat): +class MicrosoftSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Azure subscription key") @@ -140,6 +149,8 @@ class MicrosoftSTT(_BaseSTTCompat): phrase_list: Optional[List[str]] = Field(default=None, description="Microsoft ASR phrase list") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class MicrosoftSTT(MicrosoftSTTOptions, _BaseSTTCompat): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -155,7 +166,7 @@ def to_config(self) -> Dict[str, Any]: } -class TencentTTS(_BaseTTSCompat): +class TencentTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") app_id: str = Field(..., description="Tencent TTS app id") @@ -169,6 +180,8 @@ class TencentTTS(_BaseTTSCompat): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Tencent TTS params") skip_patterns: Optional[List[int]] = Field(default=None) + +class TencentTTS(TencentTTSOptions, _BaseTTSCompat): @property def sample_rate(self) -> Optional[int]: audio_setting = (self.additional_params or {}).get("audio_setting") @@ -206,7 +219,7 @@ def to_config(self) -> Dict[str, Any]: return result -class BytedanceTTS(_BaseTTSCompat): +class BytedanceTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") token: str = Field(..., description="Bytedance TTS auth token") @@ -220,6 +233,8 @@ class BytedanceTTS(_BaseTTSCompat): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Bytedance TTS params") skip_patterns: Optional[List[int]] = Field(default=None) + +class BytedanceTTS(BytedanceTTSOptions, _BaseTTSCompat): @property def sample_rate(self) -> Optional[int]: audio_setting = (self.additional_params or {}).get("audio_setting") @@ -257,7 +272,7 @@ def to_config(self) -> Dict[str, Any]: return result -class BytedanceDuplexTTS(_BaseTTSCompat): +class BytedanceDuplexTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") token: str = Field(..., description="Bytedance Duplex TTS auth token") @@ -266,6 +281,8 @@ class BytedanceDuplexTTS(_BaseTTSCompat): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Bytedance Duplex TTS params") skip_patterns: Optional[List[int]] = Field(default=None) + +class BytedanceDuplexTTS(BytedanceDuplexTTSOptions, _BaseTTSCompat): @property def sample_rate(self) -> Optional[int]: audio_setting = (self.additional_params or {}).get("audio_setting") @@ -294,7 +311,7 @@ def to_config(self) -> Dict[str, Any]: return result -class CosyVoiceTTS(_BaseTTSCompat): +class CosyVoiceTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="CosyVoice API key") @@ -304,6 +321,8 @@ class CosyVoiceTTS(_BaseTTSCompat): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="CosyVoice TTS params from REST doc") skip_patterns: Optional[List[int]] = Field(default=None) + +class CosyVoiceTTS(CosyVoiceTTSOptions, _BaseTTSCompat): @property def resolved_sample_rate(self) -> Optional[int]: if self.sample_rate is not None: @@ -334,7 +353,7 @@ def to_config(self) -> Dict[str, Any]: return result -class StepFunTTS(_BaseTTSCompat): +class StepFunTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="StepFun TTS API key") @@ -343,6 +362,8 @@ class StepFunTTS(_BaseTTSCompat): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="StepFun TTS params from REST doc") skip_patterns: Optional[List[int]] = Field(default=None) + +class StepFunTTS(StepFunTTSOptions, _BaseTTSCompat): @property def sample_rate(self) -> Optional[int]: audio_setting = (self.additional_params or {}).get("audio_setting") @@ -369,7 +390,7 @@ def to_config(self) -> Dict[str, Any]: return result -class MicrosoftTTS(_BaseTTSCompat): +class MicrosoftTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Azure subscription key") @@ -381,6 +402,8 @@ class MicrosoftTTS(_BaseTTSCompat): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Microsoft TTS params") skip_patterns: Optional[List[int]] = Field(default=None) + +class MicrosoftTTS(MicrosoftTTSOptions, _BaseTTSCompat): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -401,7 +424,7 @@ def to_config(self) -> Dict[str, Any]: return result -class MiniMaxTTS(_BaseTTSCompat): +class MiniMaxTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: Optional[str] = Field(default=None, description="MiniMax API key") @@ -421,13 +444,15 @@ class MiniMaxTTS(_BaseTTSCompat): skip_patterns: Optional[List[int]] = Field(default=None) @model_validator(mode="after") - def _validate_params(self) -> "MiniMaxTTS": + def _validate_params(self) -> "MiniMaxTTSOptions": if self.voice_id is not None and self.timber_weights is not None: raise ValueError("MiniMaxTTS requires exactly one of voice_id or timber_weights") if self.voice_id is None and self.timber_weights is None: raise ValueError("MiniMaxTTS requires exactly one of voice_id or timber_weights") return self + +class MiniMaxTTS(MiniMaxTTSOptions, _BaseTTSCompat): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) if self.key is not None: @@ -466,7 +491,7 @@ def to_config(self) -> Dict[str, Any]: return result -class AliyunLLM(BaseLLM): +class AliyunLLMOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="OpenAI API key") @@ -490,7 +515,7 @@ class AliyunLLM(BaseLLM): max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") @model_validator(mode="after") - def _validate_byok_params(self) -> "AliyunLLM": + def _validate_byok_params(self) -> "AliyunLLMOptions": if not self.model: raise ValueError("AliyunLLM requires model") if self.api_key is not None and self.base_url is None: @@ -503,6 +528,8 @@ def _validate_byok_params(self) -> "AliyunLLM": raise ValueError("AliyunLLM Agora-managed mode does not allow vendor") return self + +class AliyunLLM(AliyunLLMOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -548,7 +575,7 @@ def to_config(self) -> Dict[str, Any]: return config -class BytedanceLLM(BaseLLM): +class BytedanceLLMOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="OpenAI API key") @@ -572,7 +599,7 @@ class BytedanceLLM(BaseLLM): max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") @model_validator(mode="after") - def _validate_byok_params(self) -> "BytedanceLLM": + def _validate_byok_params(self) -> "BytedanceLLMOptions": if not self.model: raise ValueError("BytedanceLLM requires model") if self.api_key is not None and self.base_url is None: @@ -585,6 +612,8 @@ def _validate_byok_params(self) -> "BytedanceLLM": raise ValueError("BytedanceLLM Agora-managed mode does not allow vendor") return self + +class BytedanceLLM(BytedanceLLMOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -630,7 +659,7 @@ def to_config(self) -> Dict[str, Any]: return config -class DeepSeekLLM(BaseLLM): +class DeepSeekLLMOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="OpenAI API key") @@ -654,7 +683,7 @@ class DeepSeekLLM(BaseLLM): max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") @model_validator(mode="after") - def _validate_byok_params(self) -> "DeepSeekLLM": + def _validate_byok_params(self) -> "DeepSeekLLMOptions": if not self.model: raise ValueError("DeepSeekLLM requires model") if self.api_key is not None and self.base_url is None: @@ -667,6 +696,8 @@ def _validate_byok_params(self) -> "DeepSeekLLM": raise ValueError("DeepSeekLLM Agora-managed mode does not allow vendor") return self + +class DeepSeekLLM(DeepSeekLLMOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -712,7 +743,7 @@ def to_config(self) -> Dict[str, Any]: return config -class TencentLLM(BaseLLM): +class TencentLLMOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="OpenAI API key") @@ -736,7 +767,7 @@ class TencentLLM(BaseLLM): max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") @model_validator(mode="after") - def _validate_byok_params(self) -> "TencentLLM": + def _validate_byok_params(self) -> "TencentLLMOptions": if not self.model: raise ValueError("TencentLLM requires model") if self.api_key is not None and self.base_url is None: @@ -749,6 +780,8 @@ def _validate_byok_params(self) -> "TencentLLM": raise ValueError("TencentLLM Agora-managed mode does not allow vendor") return self + +class TencentLLM(TencentLLMOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -794,7 +827,7 @@ def to_config(self) -> Dict[str, Any]: return config -class QwenOmni(BaseMLLM): +class QwenOmniOptions(BaseModel): """Alibaba Cloud Qwen Omni Realtime MLLM vendor (`mllm.vendor`: ``qwen_omni``).""" model_config = ConfigDict(extra="forbid") @@ -815,6 +848,10 @@ class QwenOmni(BaseMLLM): turn_detection: Optional[MllmTurnDetection] = Field(default=None, description="MLLM turn detection configuration") failure_message: Optional[str] = Field(default=None, description="Message played on failure") + +class QwenOmni(QwenOmniOptions, BaseMLLM): + """Alibaba Cloud Qwen Omni Realtime MLLM vendor (`mllm.vendor`: ``qwen_omni``).""" + def to_config(self) -> Dict[str, Any]: inner_params: Dict[str, Any] = dict(self.params or {}) if self.model is not None: @@ -848,7 +885,7 @@ def to_config(self) -> Dict[str, Any]: return config -class SenseTimeAvatar(BaseAvatar): +class SenseTimeAvatarOptions(BaseModel): model_config = ConfigDict(extra="forbid", populate_by_name=True) agora_token: Optional[str] = Field(default=None, description="RTC token for avatar publisher; generated by AgentSession when omitted") @@ -859,6 +896,8 @@ class SenseTimeAvatar(BaseAvatar): enable: Optional[bool] = Field(default=None) additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class SenseTimeAvatar(SenseTimeAvatarOptions, BaseAvatar): @property def required_sample_rate(self) -> int: return 0 @@ -881,7 +920,7 @@ def to_config(self) -> Dict[str, Any]: return {"enable": enable, "vendor": "sensetime", "params": params} -class SpatiusAvatar(BaseAvatar): +class SpatiusAvatarOptions(BaseModel): model_config = ConfigDict(extra="forbid") spatius_api_key: str = Field(..., description="Spatius API key") @@ -895,6 +934,8 @@ class SpatiusAvatar(BaseAvatar): enable: Optional[bool] = Field(default=None) additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class SpatiusAvatar(SpatiusAvatarOptions, BaseAvatar): @property def required_sample_rate(self) -> int: return self.sample_rate or 0 diff --git a/src/agora_agent/agentkit/vendors/llm.py b/src/agora_agent/agentkit/vendors/llm.py index 4e35491..1ba29f6 100644 --- a/src/agora_agent/agentkit/vendors/llm.py +++ b/src/agora_agent/agentkit/vendors/llm.py @@ -1,8 +1,7 @@ from typing import Any, Dict, List, Optional -from pydantic import ConfigDict, Field, model_validator - from .base import BaseLLM +from pydantic import BaseModel, ConfigDict, Field, model_validator LlmGreetingConfigs = Dict[str, Any] _OPENAI_MANAGED_MODELS = {"gpt-4o-mini", "gpt-4.1-mini", "gpt-5-nano", "gpt-5-mini"} @@ -27,7 +26,7 @@ def _dump_optional_model(value: Any) -> Any: return value -class OpenAI(BaseLLM): +class OpenAIOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="OpenAI API key") @@ -51,7 +50,7 @@ class OpenAI(BaseLLM): max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") @model_validator(mode="after") - def _validate_byok_params(self) -> "OpenAI": + def _validate_byok_params(self) -> "OpenAIOptions": if not self.model: raise ValueError("OpenAI requires model") if self.api_key is not None and self.base_url is None: @@ -64,6 +63,8 @@ def _validate_byok_params(self) -> "OpenAI": raise ValueError("OpenAI Agora-managed mode does not allow vendor") return self + +class OpenAI(OpenAIOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: # model is the default; explicit params entries extend/override it. # This matches the TS SDK behaviour: { model, ...params }. @@ -112,7 +113,7 @@ def to_config(self) -> Dict[str, Any]: return config -class AzureOpenAI(BaseLLM): +class AzureOpenAIOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Azure OpenAI API key") @@ -137,6 +138,8 @@ class AzureOpenAI(BaseLLM): mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None) max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") + +class AzureOpenAI(AzureOpenAIOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: url = ( f"{self.endpoint}/openai/deployments/" @@ -186,7 +189,7 @@ def to_config(self) -> Dict[str, Any]: return config -class Anthropic(BaseLLM): +class AnthropicOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Anthropic API key") @@ -209,6 +212,8 @@ class Anthropic(BaseLLM): mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None) max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") + +class Anthropic(AnthropicOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: # Named fields take precedence over anything in the generic params dict. params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -252,7 +257,7 @@ def to_config(self) -> Dict[str, Any]: return config -class Gemini(BaseLLM): +class GeminiOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Google AI API key") @@ -276,6 +281,8 @@ class Gemini(BaseLLM): mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None) max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") + +class Gemini(GeminiOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: # Named fields take precedence over anything in the generic params dict. params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -324,7 +331,7 @@ def to_config(self) -> Dict[str, Any]: return config -class Groq(BaseLLM): +class GroqOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Groq API key") @@ -348,11 +355,13 @@ class Groq(BaseLLM): max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") @model_validator(mode="after") - def _validate_byok_params(self) -> "Groq": + def _validate_byok_params(self) -> "GroqOptions": if not self.model: raise ValueError("Groq requires model") return self + +class Groq(GroqOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -398,7 +407,7 @@ def to_config(self) -> Dict[str, Any]: return config -class CustomLLM(BaseLLM): +class CustomLLMOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Custom LLM API key") @@ -422,11 +431,13 @@ class CustomLLM(BaseLLM): max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") @model_validator(mode="after") - def _validate_byok_params(self) -> "CustomLLM": + def _validate_byok_params(self) -> "CustomLLMOptions": if not self.model: raise ValueError("CustomLLM requires model") return self + +class CustomLLM(CustomLLMOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -473,7 +484,7 @@ def to_config(self) -> Dict[str, Any]: return config -class VertexAILLM(BaseLLM): +class VertexAILLMOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Vertex AI access token or API key") @@ -499,6 +510,8 @@ class VertexAILLM(BaseLLM): mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None) max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") + +class VertexAILLM(VertexAILLMOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: # Named fields take precedence over anything in the generic params dict. params: Dict[str, Any] = {"model": self.model, **(self.params or {})} @@ -551,7 +564,7 @@ def to_config(self) -> Dict[str, Any]: return config -class AmazonBedrock(BaseLLM): +class AmazonBedrockOptions(BaseModel): model_config = ConfigDict(extra="forbid") access_key: str = Field(..., description="AWS access key ID") @@ -576,6 +589,8 @@ class AmazonBedrock(BaseLLM): mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None) max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache") + +class AmazonBedrock(AmazonBedrockOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.params or {}) if self.max_tokens is not None: @@ -620,7 +635,7 @@ def to_config(self) -> Dict[str, Any]: return config -class Dify(BaseLLM): +class DifyOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Dify API key") @@ -642,6 +657,8 @@ class Dify(BaseLLM): mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None) max_history: Optional[int] = Field(default=None, gt=0) + +class Dify(DifyOptions, BaseLLM): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = {"model": self.model, **(self.params or {})} if self.user is not None: diff --git a/src/agora_agent/agentkit/vendors/mllm.py b/src/agora_agent/agentkit/vendors/mllm.py index 01563c4..bb207b7 100644 --- a/src/agora_agent/agentkit/vendors/mllm.py +++ b/src/agora_agent/agentkit/vendors/mllm.py @@ -1,14 +1,13 @@ from typing import Any, Dict, List, Optional -from pydantic import ConfigDict, Field - from ...types.mllm_turn_detection import MllmTurnDetection from .base import BaseMLLM +from pydantic import BaseModel, ConfigDict, Field MllmTurnDetectionConfig = MllmTurnDetection -class OpenAIRealtime(BaseMLLM): +class OpenAIRealtimeOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="OpenAI API key") @@ -28,6 +27,8 @@ class OpenAIRealtime(BaseMLLM): turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration") failure_message: Optional[str] = Field(default=None, description="Message played on failure") + +class OpenAIRealtime(OpenAIRealtimeOptions, BaseMLLM): def to_config(self) -> Dict[str, Any]: config: Dict[str, Any] = { "vendor": "openai", @@ -70,7 +71,7 @@ def to_config(self) -> Dict[str, Any]: return config -class AzureOpenAIRealtime(BaseMLLM): +class AzureOpenAIRealtimeOptions(BaseModel): """Azure OpenAI Realtime MLLM vendor (`mllm.vendor`: ``azure``).""" model_config = ConfigDict(extra="forbid") @@ -93,6 +94,10 @@ class AzureOpenAIRealtime(BaseMLLM): turn_detection: MllmTurnDetectionConfig = Field(..., description="MLLM turn detection configuration") failure_message: Optional[str] = Field(default=None, description="Message played on failure") + +class AzureOpenAIRealtime(AzureOpenAIRealtimeOptions, BaseMLLM): + """Azure OpenAI Realtime MLLM vendor (`mllm.vendor`: ``azure``).""" + def to_config(self) -> Dict[str, Any]: inner_params: Dict[str, Any] = dict(self.params or {}) if self.model is not None: @@ -129,7 +134,7 @@ def to_config(self) -> Dict[str, Any]: # is deprecated and reserved naming for future XaiSTT / XaiTTS cascading vendors. -class XaiGrok(BaseMLLM): +class XaiGrokOptions(BaseModel): """xAI Grok MLLM vendor (`mllm.vendor`: ``xai``).""" model_config = ConfigDict(extra="forbid") @@ -147,6 +152,10 @@ class XaiGrok(BaseMLLM): turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration") failure_message: Optional[str] = Field(default=None, description="Message played on failure") + +class XaiGrok(XaiGrokOptions, BaseMLLM): + """xAI Grok MLLM vendor (`mllm.vendor`: ``xai``).""" + def to_config(self) -> Dict[str, Any]: inner_params: Dict[str, Any] = dict(self.params or {}) if self.voice is not None: @@ -179,7 +188,7 @@ def to_config(self) -> Dict[str, Any]: return config -class VertexAI(BaseMLLM): +class VertexAIOptions(BaseModel): model_config = ConfigDict(extra="forbid") model: str = Field(..., description="Model name") @@ -202,6 +211,8 @@ class VertexAI(BaseMLLM): turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration") failure_message: Optional[str] = Field(default=None, description="Message played on failure") + +class VertexAI(VertexAIOptions, BaseMLLM): def to_config(self) -> Dict[str, Any]: # additional_params spread first so that explicit fields always win, # matching the TypeScript SDK. @@ -246,7 +257,7 @@ def to_config(self) -> Dict[str, Any]: return config -class GeminiLive(BaseMLLM): +class GeminiLiveOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Google API key") @@ -267,6 +278,8 @@ class GeminiLive(BaseMLLM): turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration") failure_message: Optional[str] = Field(default=None, description="Message played on failure") + +class GeminiLive(GeminiLiveOptions, BaseMLLM): def to_config(self) -> Dict[str, Any]: inner_params: Dict[str, Any] = {} if self.additional_params is not None: diff --git a/src/agora_agent/agentkit/vendors/stt.py b/src/agora_agent/agentkit/vendors/stt.py index c276fca..2dfdd56 100644 --- a/src/agora_agent/agentkit/vendors/stt.py +++ b/src/agora_agent/agentkit/vendors/stt.py @@ -1,13 +1,12 @@ from typing import Any, Dict, List, Optional -from pydantic import ConfigDict, Field, model_validator - from .base import BaseSTT +from pydantic import BaseModel, ConfigDict, Field, model_validator _DEEPGRAM_MANAGED_MODELS = {"nova-2", "nova-3"} -class SpeechmaticsSTT(BaseSTT): +class SpeechmaticsSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Speechmatics API key") @@ -16,6 +15,8 @@ class SpeechmaticsSTT(BaseSTT): uri: Optional[str] = Field(default=None, description="Speechmatics streaming WebSocket URL") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class SpeechmaticsSTT(SpeechmaticsSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -34,7 +35,7 @@ def to_config(self) -> Dict[str, Any]: return config -class DeepgramSTT(BaseSTT): +class DeepgramSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="Deepgram API key") @@ -46,11 +47,13 @@ class DeepgramSTT(BaseSTT): additional_params: Optional[Dict[str, Any]] = Field(default=None) @model_validator(mode="after") - def _validate_managed_model(self) -> "DeepgramSTT": + def _validate_managed_model(self) -> "DeepgramSTTOptions": if self.api_key is None and (self.model is None or self.model.strip().lower() not in _DEEPGRAM_MANAGED_MODELS): raise ValueError("DeepgramSTT requires api_key unless using a supported Agora-managed model") return self + +class DeepgramSTT(DeepgramSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) @@ -73,7 +76,7 @@ def to_config(self) -> Dict[str, Any]: return config -class MicrosoftSTT(BaseSTT): +class MicrosoftSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Azure subscription key") @@ -81,6 +84,8 @@ class MicrosoftSTT(BaseSTT): language: str = Field(..., description="Language code (e.g., en-US)") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class MicrosoftSTT(MicrosoftSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -97,7 +102,7 @@ def to_config(self) -> Dict[str, Any]: return config -class OpenAISTT(BaseSTT): +class OpenAISTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="OpenAI API key") @@ -107,6 +112,8 @@ class OpenAISTT(BaseSTT): input_audio_transcription: Optional[Dict[str, Any]] = Field(default=None, description="OpenAI transcription settings") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class OpenAISTT(OpenAISTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params["api_key"] = self.api_key @@ -134,7 +141,7 @@ def to_config(self) -> Dict[str, Any]: return config -class GoogleSTT(BaseSTT): +class GoogleSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") project_id: str = Field(..., description="Google Cloud project ID") @@ -144,6 +151,8 @@ class GoogleSTT(BaseSTT): model: Optional[str] = Field(default=None, description="Recognition model") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class GoogleSTT(GoogleSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -164,7 +173,7 @@ def to_config(self) -> Dict[str, Any]: return config -class AmazonSTT(BaseSTT): +class AmazonSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") access_key: str = Field(..., description="AWS Access Key ID") @@ -173,6 +182,8 @@ class AmazonSTT(BaseSTT): language: str = Field(..., description="Language code") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class AmazonSTT(AmazonSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -190,7 +201,7 @@ def to_config(self) -> Dict[str, Any]: return config -class AssemblyAISTT(BaseSTT): +class AssemblyAISTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="AssemblyAI API key") @@ -198,6 +209,8 @@ class AssemblyAISTT(BaseSTT): ws_url: Optional[str] = Field(default=None, description="AssemblyAI streaming WebSocket URL") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class AssemblyAISTT(AssemblyAISTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params["api_key"] = self.api_key @@ -213,12 +226,14 @@ def to_config(self) -> Dict[str, Any]: return config -class AresSTT(BaseSTT): +class AresSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") keywords: Optional[List[str]] = Field(default=None, description="Hotwords that improve ASR accuracy") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class AresSTT(AresSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) if self.keywords is not None: @@ -230,7 +245,7 @@ def to_config(self) -> Dict[str, Any]: return config -class SarvamSTT(BaseSTT): +class SarvamSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Sarvam API key") @@ -238,6 +253,8 @@ class SarvamSTT(BaseSTT): model: Optional[str] = Field(default=None, description="Model name") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class SarvamSTT(SarvamSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -254,7 +271,7 @@ def to_config(self) -> Dict[str, Any]: return config -class XaiSTT(BaseSTT): +class XaiSTTOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="xAI API key") @@ -263,6 +280,8 @@ class XaiSTT(BaseSTT): language: Optional[str] = Field(default=None, description="Language code for speech recognition") additional_params: Optional[Dict[str, Any]] = Field(default=None) + +class XaiSTT(XaiSTTOptions, BaseSTT): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params["api_key"] = self.api_key diff --git a/src/agora_agent/agentkit/vendors/tts.py b/src/agora_agent/agentkit/vendors/tts.py index f357ace..3289901 100644 --- a/src/agora_agent/agentkit/vendors/tts.py +++ b/src/agora_agent/agentkit/vendors/tts.py @@ -1,14 +1,13 @@ from typing import Any, Dict, List, Literal, Optional from urllib.parse import urlsplit -from pydantic import ConfigDict, Field, field_validator, model_validator - -from .base import BaseTTS, CartesiaSampleRate, ElevenLabsSampleRate, GoogleTTSSampleRate, MicrosoftSampleRate from ..constants import CredentialMode from ..presets import MiniMaxPresetModels, OpenAITtsPresetModels +from .base import BaseTTS, CartesiaSampleRate, ElevenLabsSampleRate, GoogleTTSSampleRate, MicrosoftSampleRate +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator -class ElevenLabsTTS(BaseTTS): +class ElevenLabsTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="ElevenLabs API key") @@ -23,6 +22,8 @@ class ElevenLabsTTS(BaseTTS): style: Optional[float] = Field(default=None, ge=0.0, le=1.0) use_speaker_boost: Optional[bool] = Field(default=None) + +class ElevenLabsTTS(ElevenLabsTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = { "key": self.key, @@ -50,7 +51,7 @@ def to_config(self) -> Dict[str, Any]: return result -class MicrosoftTTS(BaseTTS): +class MicrosoftTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Azure subscription key") @@ -62,6 +63,8 @@ class MicrosoftTTS(BaseTTS): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Microsoft TTS params") skip_patterns: Optional[List[int]] = Field(default=None) + +class MicrosoftTTS(MicrosoftTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -83,7 +86,7 @@ def to_config(self) -> Dict[str, Any]: return result -class OpenAITTS(BaseTTS): +class OpenAITTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: Optional[str] = Field(default=None, description="OpenAI API key") @@ -95,7 +98,7 @@ class OpenAITTS(BaseTTS): skip_patterns: Optional[List[int]] = Field(default=None) @model_validator(mode="after") - def _validate_byok_params(self) -> "OpenAITTS": + def _validate_byok_params(self) -> "OpenAITTSOptions": if self.api_key is not None: missing = [ name @@ -114,6 +117,8 @@ def _validate_byok_params(self) -> "OpenAITTS": raise ValueError("OpenAITTS base_url is only valid when api_key is set") return self + +class OpenAITTS(OpenAITTSOptions, BaseTTS): @property def sample_rate(self) -> Optional[int]: return 24000 @@ -140,7 +145,7 @@ def to_config(self) -> Dict[str, Any]: return result -class CartesiaTTS(BaseTTS): +class CartesiaTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Cartesia API key") @@ -151,6 +156,8 @@ class CartesiaTTS(BaseTTS): sample_rate: Optional[CartesiaSampleRate] = Field(default=None, description="Sample rate in Hz") skip_patterns: Optional[List[int]] = Field(default=None) + +class CartesiaTTS(CartesiaTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = { "api_key": self.api_key, @@ -171,7 +178,7 @@ def to_config(self) -> Dict[str, Any]: return result -class GoogleTTS(BaseTTS): +class GoogleTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Google Cloud service account credentials JSON string") @@ -180,6 +187,8 @@ class GoogleTTS(BaseTTS): sample_rate_hertz: Optional[GoogleTTSSampleRate] = Field(default=None, description="Sample rate in Hz") skip_patterns: Optional[List[int]] = Field(default=None) + +class GoogleTTS(GoogleTTSOptions, BaseTTS): @property def sample_rate(self) -> Optional[int]: return self.sample_rate_hertz @@ -201,7 +210,7 @@ def to_config(self) -> Dict[str, Any]: return result -class AmazonTTS(BaseTTS): +class AmazonTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") access_key: str = Field(..., description="AWS access key") @@ -211,6 +220,8 @@ class AmazonTTS(BaseTTS): engine: str = Field(..., description="Amazon Polly engine type") skip_patterns: Optional[List[int]] = Field(default=None) + +class AmazonTTS(AmazonTTSOptions, BaseTTS): @property def sample_rate(self) -> Optional[int]: return None @@ -230,7 +241,7 @@ def to_config(self) -> Dict[str, Any]: return result -class DeepgramTTS(BaseTTS): +class DeepgramTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Deepgram API key") @@ -240,6 +251,8 @@ class DeepgramTTS(BaseTTS): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Deepgram TTS parameters") skip_patterns: Optional[List[int]] = Field(default=None) + +class DeepgramTTS(DeepgramTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update({ @@ -257,7 +270,7 @@ def to_config(self) -> Dict[str, Any]: return result -class GradiumTTS(BaseTTS): +class GradiumTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Gradium API key") @@ -268,6 +281,8 @@ class GradiumTTS(BaseTTS): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Gradium TTS parameters") skip_patterns: Optional[List[int]] = Field(default=None) + +class GradiumTTS(GradiumTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params["api_key"] = self.api_key @@ -286,7 +301,7 @@ def to_config(self) -> Dict[str, Any]: return result -class HumeAITTS(BaseTTS): +class HumeAITTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Hume AI API key") @@ -298,6 +313,8 @@ class HumeAITTS(BaseTTS): trailing_silence: Optional[float] = Field(default=None, description="Trailing silence in seconds") skip_patterns: Optional[List[int]] = Field(default=None) + +class HumeAITTS(HumeAITTSOptions, BaseTTS): @property def sample_rate(self) -> Optional[int]: return None @@ -324,7 +341,7 @@ def to_config(self) -> Dict[str, Any]: return result -class RimeTTS(BaseTTS): +class RimeTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: Optional[str] = Field(default=None, description="Rime API key") @@ -335,7 +352,7 @@ class RimeTTS(BaseTTS): skip_patterns: Optional[List[int]] = Field(default=None) @model_validator(mode="after") - def _validate_credential_mode(self) -> "RimeTTS": + def _validate_credential_mode(self) -> "RimeTTSOptions": required: Dict[str, Optional[str]] if self.credential_mode == CredentialMode.MANAGED: required = {"base_url": self.base_url, "model_id": self.model_id} @@ -349,6 +366,8 @@ def _validate_credential_mode(self) -> "RimeTTS": raise ValueError(f"RimeTTS requires {', '.join(missing)} for {mode}") return self + +class RimeTTS(RimeTTSOptions, BaseTTS): @property def sample_rate(self) -> Optional[int]: return None @@ -370,7 +389,7 @@ def to_config(self) -> Dict[str, Any]: return result -class FishAudioTTS(BaseTTS): +class FishAudioTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Fish Audio API key") @@ -378,6 +397,8 @@ class FishAudioTTS(BaseTTS): backend: str = Field(..., description="Backend") skip_patterns: Optional[List[int]] = Field(default=None) + +class FishAudioTTS(FishAudioTTSOptions, BaseTTS): @property def sample_rate(self) -> Optional[int]: return None @@ -395,7 +416,7 @@ def to_config(self) -> Dict[str, Any]: return result -class MiniMaxTTS(BaseTTS): +class MiniMaxTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: Optional[str] = Field(default=None, description="MiniMax API key") @@ -417,7 +438,7 @@ class MiniMaxTTS(BaseTTS): skip_patterns: Optional[List[int]] = Field(default=None) @model_validator(mode="after") - def _validate_byok_params(self) -> "MiniMaxTTS": + def _validate_byok_params(self) -> "MiniMaxTTSOptions": if self.voice_id is not None and self.timber_weights is not None: raise ValueError("MiniMaxTTS requires exactly one of voice_id or timber_weights") if self.key is not None: @@ -440,6 +461,8 @@ def _validate_byok_params(self) -> "MiniMaxTTS": raise ValueError("MiniMaxTTS requires key unless using a supported Agora-managed model") return self + +class MiniMaxTTS(MiniMaxTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) if self.key is not None: @@ -484,7 +507,7 @@ def to_config(self) -> Dict[str, Any]: return result -class MistralTTS(BaseTTS): +class MistralTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Mistral API key") @@ -493,6 +516,8 @@ class MistralTTS(BaseTTS): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Mistral TTS parameters") skip_patterns: Optional[List[int]] = Field(default=None) + +class MistralTTS(MistralTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params["api_key"] = self.api_key @@ -507,7 +532,7 @@ def to_config(self) -> Dict[str, Any]: return result -class TypecastTTS(BaseTTS): +class TypecastTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="Typecast API key") @@ -516,6 +541,8 @@ class TypecastTTS(BaseTTS): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional Typecast TTS parameters") skip_patterns: Optional[List[int]] = Field(default=None) + +class TypecastTTS(TypecastTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params.update( @@ -532,7 +559,7 @@ def to_config(self) -> Dict[str, Any]: return result -class SarvamTTS(BaseTTS): +class SarvamTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Sarvam API subscription key") @@ -544,6 +571,8 @@ class SarvamTTS(BaseTTS): sample_rate: Optional[int] = Field(default=None, description="Audio sample rate in Hz") skip_patterns: Optional[List[int]] = Field(default=None) + +class SarvamTTS(SarvamTTSOptions, BaseTTS): @property def resolved_sample_rate(self) -> Optional[int]: return None @@ -569,7 +598,7 @@ def to_config(self) -> Dict[str, Any]: return result -class MurfTTS(BaseTTS): +class MurfTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") key: str = Field(..., description="Murf API key") @@ -582,6 +611,8 @@ class MurfTTS(BaseTTS): sample_rate: Optional[int] = Field(default=None, description="Audio sample rate") skip_patterns: Optional[List[int]] = Field(default=None) + +class MurfTTS(MurfTTSOptions, BaseTTS): @property def resolved_sample_rate(self) -> Optional[int]: return None @@ -610,7 +641,7 @@ def to_config(self) -> Dict[str, Any]: return result -class GenericTTS(BaseTTS): +class GenericTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") url: str = Field(..., description="HTTP(S) endpoint of the generic TTS service") @@ -633,6 +664,8 @@ def validate_url(cls, value: str) -> str: raise ValueError("GenericTTS url must be a valid HTTP(S) endpoint") return value + +class GenericTTS(GenericTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) if self.api_key is not None: @@ -662,7 +695,7 @@ def to_config(self) -> Dict[str, Any]: return result -class XaiTTS(BaseTTS): +class XaiTTSOptions(BaseModel): model_config = ConfigDict(extra="forbid") api_key: str = Field(..., description="xAI API key") @@ -672,6 +705,8 @@ class XaiTTS(BaseTTS): additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional xAI TTS params") skip_patterns: Optional[List[int]] = Field(default=None) + +class XaiTTS(XaiTTSOptions, BaseTTS): def to_config(self) -> Dict[str, Any]: params: Dict[str, Any] = dict(self.additional_params or {}) params["api_key"] = self.api_key