Skip to content

Commit 2b49847

Browse files
committed
refactor: collapse STT vendor configs into pydantic models
1 parent 551fb63 commit 2b49847

4 files changed

Lines changed: 133 additions & 200 deletions

File tree

src/agora_agent/agentkit/vendors/base.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from abc import ABC, abstractmethod
22
from typing import Any, Dict, Optional
33

4+
from pydantic import BaseModel
45
from typing_extensions import Literal
56

67
# Supported sample rates across all TTS providers.
@@ -49,7 +50,7 @@ def sample_rate(self) -> Optional[int]:
4950
"""The configured sample rate in Hz, or ``None`` if not explicitly set."""
5051

5152

52-
class BaseSTT(ABC):
53+
class BaseSTT(BaseModel, ABC):
5354
"""Abstract base class for all STT vendor implementations.
5455
5556
Subclasses must implement :meth:`to_config` to return a dict that maps to

src/agora_agent/agentkit/vendors/cn.py

Lines changed: 47 additions & 75 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010
from .tts import BaseTTS as _BaseTTSCompat
1111

1212

13-
class TencentSTTOptions(BaseModel):
13+
class TencentSTT(_BaseSTTCompat):
1414
model_config = ConfigDict(extra="forbid")
1515

1616
key: str = Field(..., description="Tencent ASR secret key")
@@ -20,36 +20,28 @@ class TencentSTTOptions(BaseModel):
2020
voice_id: str = Field(..., description="Tencent ASR voice id")
2121
additional_params: Optional[Dict[str, Any]] = Field(default=None)
2222

23-
24-
class TencentSTT(_BaseSTTCompat):
25-
def __init__(self, **kwargs: Any):
26-
self.options = TencentSTTOptions(**kwargs)
27-
2823
def to_config(self) -> Dict[str, Any]:
29-
params: Dict[str, Any] = dict(self.options.additional_params or {})
24+
params: Dict[str, Any] = dict(self.additional_params or {})
3025
params.update(
3126
{
32-
"key": self.options.key,
33-
"app_id": self.options.app_id,
34-
"secret": self.options.secret,
35-
"engine_model_type": self.options.engine_model_type,
36-
"voice_id": self.options.voice_id,
27+
"key": self.key,
28+
"app_id": self.app_id,
29+
"secret": self.secret,
30+
"engine_model_type": self.engine_model_type,
31+
"voice_id": self.voice_id,
3732
}
3833
)
3934
return {"vendor": "tencent", "params": params}
4035

4136

4237
class FengmingSTT(_BaseSTTCompat):
43-
def __init__(self, **kwargs: Any):
44-
if kwargs:
45-
unexpected = ", ".join(sorted(kwargs))
46-
raise TypeError(f"FengmingSTT does not accept parameters: {unexpected}")
38+
model_config = ConfigDict(extra="forbid")
4739

4840
def to_config(self) -> Dict[str, Any]:
4941
return {"vendor": "fengming"}
5042

5143

52-
class XfyunSTTOptions(BaseModel):
44+
class XfyunSTT(_BaseSTTCompat):
5345
model_config = ConfigDict(extra="forbid")
5446

5547
api_key: Optional[str] = Field(default=None, description="Xfyun ASR API key")
@@ -58,28 +50,23 @@ class XfyunSTTOptions(BaseModel):
5850
language: Optional[str] = Field(default=None, description="Xfyun ASR language")
5951
additional_params: Optional[Dict[str, Any]] = Field(default=None)
6052

61-
62-
class XfyunSTT(_BaseSTTCompat):
63-
def __init__(self, **kwargs: Any):
64-
self.options = XfyunSTTOptions(**kwargs)
65-
6653
def to_config(self) -> Dict[str, Any]:
67-
params: Dict[str, Any] = dict(self.options.additional_params or {})
68-
if self.options.api_key is not None:
69-
params["api_key"] = self.options.api_key
70-
if self.options.app_id is not None:
71-
params["app_id"] = self.options.app_id
72-
if self.options.api_secret is not None:
73-
params["api_secret"] = self.options.api_secret
74-
if self.options.language is not None:
75-
params["language"] = self.options.language
54+
params: Dict[str, Any] = dict(self.additional_params or {})
55+
if self.api_key is not None:
56+
params["api_key"] = self.api_key
57+
if self.app_id is not None:
58+
params["app_id"] = self.app_id
59+
if self.api_secret is not None:
60+
params["api_secret"] = self.api_secret
61+
if self.language is not None:
62+
params["language"] = self.language
7663
return {
7764
"vendor": "xfyun",
7865
"params": params,
7966
}
8067

8168

82-
class XfyunBigModelSTTOptions(BaseModel):
69+
class XfyunBigModelSTT(_BaseSTTCompat):
8370
model_config = ConfigDict(extra="forbid")
8471

8572
api_key: Optional[str] = Field(default=None, description="Xfyun BigModel ASR API key")
@@ -89,30 +76,25 @@ class XfyunBigModelSTTOptions(BaseModel):
8976
language: Optional[str] = Field(default=None, description="Xfyun BigModel ASR language")
9077
additional_params: Optional[Dict[str, Any]] = Field(default=None)
9178

92-
93-
class XfyunBigModelSTT(_BaseSTTCompat):
94-
def __init__(self, **kwargs: Any):
95-
self.options = XfyunBigModelSTTOptions(**kwargs)
96-
9779
def to_config(self) -> Dict[str, Any]:
98-
params: Dict[str, Any] = dict(self.options.additional_params or {})
99-
if self.options.api_key is not None:
100-
params["api_key"] = self.options.api_key
101-
if self.options.app_id is not None:
102-
params["app_id"] = self.options.app_id
103-
if self.options.api_secret is not None:
104-
params["api_secret"] = self.options.api_secret
105-
if self.options.language_name is not None:
106-
params["language_name"] = self.options.language_name
107-
if self.options.language is not None:
108-
params["language"] = self.options.language
80+
params: Dict[str, Any] = dict(self.additional_params or {})
81+
if self.api_key is not None:
82+
params["api_key"] = self.api_key
83+
if self.app_id is not None:
84+
params["app_id"] = self.app_id
85+
if self.api_secret is not None:
86+
params["api_secret"] = self.api_secret
87+
if self.language_name is not None:
88+
params["language_name"] = self.language_name
89+
if self.language is not None:
90+
params["language"] = self.language
10991
return {
11092
"vendor": "xfyun_bigmodel",
11193
"params": params,
11294
}
11395

11496

115-
class XfyunDialectSTTOptions(BaseModel):
97+
class XfyunDialectSTT(_BaseSTTCompat):
11698
model_config = ConfigDict(extra="forbid")
11799

118100
app_id: Optional[str] = Field(default=None, description="Xfyun Dialect ASR app id")
@@ -121,28 +103,23 @@ class XfyunDialectSTTOptions(BaseModel):
121103
language: Optional[str] = Field(default=None, description="Xfyun Dialect ASR language")
122104
additional_params: Optional[Dict[str, Any]] = Field(default=None)
123105

124-
125-
class XfyunDialectSTT(_BaseSTTCompat):
126-
def __init__(self, **kwargs: Any):
127-
self.options = XfyunDialectSTTOptions(**kwargs)
128-
129106
def to_config(self) -> Dict[str, Any]:
130-
params: Dict[str, Any] = dict(self.options.additional_params or {})
131-
if self.options.app_id is not None:
132-
params["app_id"] = self.options.app_id
133-
if self.options.access_key_id is not None:
134-
params["access_key_id"] = self.options.access_key_id
135-
if self.options.access_key_secret is not None:
136-
params["access_key_secret"] = self.options.access_key_secret
137-
if self.options.language is not None:
138-
params["language"] = self.options.language
107+
params: Dict[str, Any] = dict(self.additional_params or {})
108+
if self.app_id is not None:
109+
params["app_id"] = self.app_id
110+
if self.access_key_id is not None:
111+
params["access_key_id"] = self.access_key_id
112+
if self.access_key_secret is not None:
113+
params["access_key_secret"] = self.access_key_secret
114+
if self.language is not None:
115+
params["language"] = self.language
139116
return {
140117
"vendor": "xfyun_dialect",
141118
"params": params,
142119
}
143120

144121

145-
class MicrosoftSTTOptions(BaseModel):
122+
class MicrosoftSTT(_BaseSTTCompat):
146123
model_config = ConfigDict(extra="forbid")
147124

148125
key: str = Field(..., description="Azure subscription key")
@@ -151,20 +128,15 @@ class MicrosoftSTTOptions(BaseModel):
151128
phrase_list: Optional[List[str]] = Field(default=None, description="Microsoft ASR phrase list")
152129
additional_params: Optional[Dict[str, Any]] = Field(default=None)
153130

154-
155-
class MicrosoftSTT(_BaseSTTCompat):
156-
def __init__(self, **kwargs: Any):
157-
self.options = MicrosoftSTTOptions(**kwargs)
158-
159131
def to_config(self) -> Dict[str, Any]:
160-
params: Dict[str, Any] = dict(self.options.additional_params or {})
132+
params: Dict[str, Any] = dict(self.additional_params or {})
161133
params.update({
162-
"key": self.options.key,
163-
"region": self.options.region,
164-
"language": self.options.language,
134+
"key": self.key,
135+
"region": self.region,
136+
"language": self.language,
165137
})
166-
if self.options.phrase_list is not None:
167-
params["phrase_list"] = self.options.phrase_list
138+
if self.phrase_list is not None:
139+
params["phrase_list"] = self.phrase_list
168140
return {
169141
"vendor": "microsoft",
170142
"params": params,

0 commit comments

Comments
 (0)