Skip to content

Commit 289e780

Browse files
plutolessclaude
andcommitted
refactor: collapse MLLM vendor configs into pydantic models
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
1 parent 13339ed commit 289e780

2 files changed

Lines changed: 128 additions & 146 deletions

File tree

src/agora_agent/agentkit/vendors/base.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -62,7 +62,7 @@ def to_config(self) -> Dict[str, Any]:
6262
"""Serialize the STT configuration to a dict for the REST API."""
6363

6464

65-
class BaseMLLM(ABC):
65+
class BaseMLLM(BaseModel, ABC):
6666
"""Abstract base class for all MLLM (multimodal LLM) vendor implementations.
6767
6868
When an MLLM is configured via :meth:`~agora_agent.agentkit.Agent.with_mllm`,
Lines changed: 127 additions & 145 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,14 @@
1-
import warnings
21
from typing import Any, Dict, List, Optional
32

4-
from pydantic import BaseModel, ConfigDict, Field
3+
from pydantic import ConfigDict, Field
54

65
from ...types.mllm_turn_detection import MllmTurnDetection
76
from .base import BaseMLLM
87

98
MllmTurnDetectionConfig = MllmTurnDetection
109

1110

12-
class OpenAIRealtimeOptions(BaseModel):
11+
class OpenAIRealtime(BaseMLLM):
1312
model_config = ConfigDict(extra="forbid")
1413

1514
api_key: str = Field(..., description="OpenAI API key")
@@ -26,49 +25,45 @@ class OpenAIRealtimeOptions(BaseModel):
2625
turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration")
2726
failure_message: Optional[str] = Field(default=None, description="Message played on failure")
2827

29-
class OpenAIRealtime(BaseMLLM):
30-
def __init__(self, **kwargs: Any):
31-
self.options = OpenAIRealtimeOptions(**kwargs)
32-
3328
def to_config(self) -> Dict[str, Any]:
3429
config: Dict[str, Any] = {
3530
"vendor": "openai",
36-
"api_key": self.options.api_key,
31+
"api_key": self.api_key,
3732
}
3833

39-
if self.options.url is not None:
40-
config["url"] = self.options.url
34+
if self.url is not None:
35+
config["url"] = self.url
4136
if (
42-
self.options.model is not None
43-
or self.options.params is not None
44-
or self.options.voice is not None
45-
or self.options.instructions is not None
46-
or self.options.input_audio_transcription is not None
37+
self.model is not None
38+
or self.params is not None
39+
or self.voice is not None
40+
or self.instructions is not None
41+
or self.input_audio_transcription is not None
4742
):
48-
params: Dict[str, Any] = {}
49-
if self.options.model is not None:
50-
params["model"] = self.options.model
51-
if self.options.params is not None:
52-
params.update(self.options.params)
53-
if self.options.voice is not None:
54-
params["voice"] = self.options.voice
55-
if self.options.instructions is not None:
56-
params["instructions"] = self.options.instructions
57-
if self.options.input_audio_transcription is not None:
58-
params["input_audio_transcription"] = self.options.input_audio_transcription
59-
config["params"] = params
60-
if self.options.greeting_message is not None:
61-
config["greeting_message"] = self.options.greeting_message
62-
if self.options.input_modalities is not None:
63-
config["input_modalities"] = self.options.input_modalities
64-
if self.options.output_modalities is not None:
65-
config["output_modalities"] = self.options.output_modalities
66-
if self.options.messages is not None:
67-
config["messages"] = self.options.messages
68-
if self.options.failure_message is not None:
69-
config["failure_message"] = self.options.failure_message
70-
if self.options.turn_detection is not None:
71-
config["turn_detection"] = self.options.turn_detection
43+
inner_params: Dict[str, Any] = {}
44+
if self.model is not None:
45+
inner_params["model"] = self.model
46+
if self.params is not None:
47+
inner_params.update(self.params)
48+
if self.voice is not None:
49+
inner_params["voice"] = self.voice
50+
if self.instructions is not None:
51+
inner_params["instructions"] = self.instructions
52+
if self.input_audio_transcription is not None:
53+
inner_params["input_audio_transcription"] = self.input_audio_transcription
54+
config["params"] = inner_params
55+
if self.greeting_message is not None:
56+
config["greeting_message"] = self.greeting_message
57+
if self.input_modalities is not None:
58+
config["input_modalities"] = self.input_modalities
59+
if self.output_modalities is not None:
60+
config["output_modalities"] = self.output_modalities
61+
if self.messages is not None:
62+
config["messages"] = self.messages
63+
if self.failure_message is not None:
64+
config["failure_message"] = self.failure_message
65+
if self.turn_detection is not None:
66+
config["turn_detection"] = self.turn_detection
7267

7368
return config
7469

@@ -77,7 +72,9 @@ def to_config(self) -> Dict[str, Any]:
7772
# is deprecated and reserved naming for future XaiSTT / XaiTTS cascading vendors.
7873

7974

80-
class XaiGrokOptions(BaseModel):
75+
class XaiGrok(BaseMLLM):
76+
"""xAI Grok MLLM vendor (`mllm.vendor`: ``xai``)."""
77+
8178
model_config = ConfigDict(extra="forbid")
8279

8380
api_key: str = Field(..., description="xAI API key")
@@ -93,46 +90,39 @@ class XaiGrokOptions(BaseModel):
9390
turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration")
9491
failure_message: Optional[str] = Field(default=None, description="Message played on failure")
9592

96-
97-
class XaiGrok(BaseMLLM):
98-
"""xAI Grok MLLM vendor (`mllm.vendor`: ``xai``)."""
99-
100-
def __init__(self, **kwargs: Any):
101-
self.options = XaiGrokOptions(**kwargs)
102-
10393
def to_config(self) -> Dict[str, Any]:
104-
params: Dict[str, Any] = dict(self.options.params or {})
105-
if self.options.voice is not None:
106-
params["voice"] = self.options.voice
107-
if self.options.language is not None:
108-
params["language"] = self.options.language
109-
if self.options.sample_rate is not None:
110-
params["sample_rate"] = self.options.sample_rate
94+
inner_params: Dict[str, Any] = dict(self.params or {})
95+
if self.voice is not None:
96+
inner_params["voice"] = self.voice
97+
if self.language is not None:
98+
inner_params["language"] = self.language
99+
if self.sample_rate is not None:
100+
inner_params["sample_rate"] = self.sample_rate
111101

112102
config: Dict[str, Any] = {
113103
"vendor": "xai",
114-
"api_key": self.options.api_key,
115-
"url": self.options.url,
116-
"params": params,
104+
"api_key": self.api_key,
105+
"url": self.url,
106+
"params": inner_params,
117107
}
118108

119-
if self.options.greeting_message is not None:
120-
config["greeting_message"] = self.options.greeting_message
121-
if self.options.input_modalities is not None:
122-
config["input_modalities"] = self.options.input_modalities
123-
if self.options.output_modalities is not None:
124-
config["output_modalities"] = self.options.output_modalities
125-
if self.options.messages is not None:
126-
config["messages"] = self.options.messages
127-
if self.options.failure_message is not None:
128-
config["failure_message"] = self.options.failure_message
129-
if self.options.turn_detection is not None:
130-
config["turn_detection"] = self.options.turn_detection
109+
if self.greeting_message is not None:
110+
config["greeting_message"] = self.greeting_message
111+
if self.input_modalities is not None:
112+
config["input_modalities"] = self.input_modalities
113+
if self.output_modalities is not None:
114+
config["output_modalities"] = self.output_modalities
115+
if self.messages is not None:
116+
config["messages"] = self.messages
117+
if self.failure_message is not None:
118+
config["failure_message"] = self.failure_message
119+
if self.turn_detection is not None:
120+
config["turn_detection"] = self.turn_detection
131121

132122
return config
133123

134124

135-
class VertexAIOptions(BaseModel):
125+
class VertexAI(BaseMLLM):
136126
model_config = ConfigDict(extra="forbid")
137127

138128
model: str = Field(..., description="Model name")
@@ -155,55 +145,51 @@ class VertexAIOptions(BaseModel):
155145
turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration")
156146
failure_message: Optional[str] = Field(default=None, description="Message played on failure")
157147

158-
class VertexAI(BaseMLLM):
159-
def __init__(self, **kwargs: Any):
160-
self.options = VertexAIOptions(**kwargs)
161-
162148
def to_config(self) -> Dict[str, Any]:
163149
# additional_params spread first so that explicit fields always win,
164150
# matching the TypeScript SDK.
165-
params: Dict[str, Any] = dict(self.options.additional_params or {})
166-
params["model"] = self.options.model
167-
params["project_id"] = self.options.project_id
168-
params["location"] = self.options.location
169-
params["adc_credentials_string"] = self.options.adc_credentials_string
170-
if self.options.instructions is not None:
171-
params["instructions"] = self.options.instructions
172-
if self.options.voice is not None:
173-
params["voice"] = self.options.voice
174-
if self.options.affective_dialog is not None:
175-
params["affective_dialog"] = self.options.affective_dialog
176-
if self.options.proactive_audio is not None:
177-
params["proactive_audio"] = self.options.proactive_audio
178-
if self.options.transcribe_agent is not None:
179-
params["transcribe_agent"] = self.options.transcribe_agent
180-
if self.options.transcribe_user is not None:
181-
params["transcribe_user"] = self.options.transcribe_user
182-
if self.options.http_options is not None:
183-
params["http_options"] = self.options.http_options
151+
inner_params: Dict[str, Any] = dict(self.additional_params or {})
152+
inner_params["model"] = self.model
153+
inner_params["project_id"] = self.project_id
154+
inner_params["location"] = self.location
155+
inner_params["adc_credentials_string"] = self.adc_credentials_string
156+
if self.instructions is not None:
157+
inner_params["instructions"] = self.instructions
158+
if self.voice is not None:
159+
inner_params["voice"] = self.voice
160+
if self.affective_dialog is not None:
161+
inner_params["affective_dialog"] = self.affective_dialog
162+
if self.proactive_audio is not None:
163+
inner_params["proactive_audio"] = self.proactive_audio
164+
if self.transcribe_agent is not None:
165+
inner_params["transcribe_agent"] = self.transcribe_agent
166+
if self.transcribe_user is not None:
167+
inner_params["transcribe_user"] = self.transcribe_user
168+
if self.http_options is not None:
169+
inner_params["http_options"] = self.http_options
184170

185171
config: Dict[str, Any] = {
186172
"vendor": "vertexai",
187-
"url": self.options.url if self.options.url is not None else "",
188-
"params": params,
173+
"url": self.url if self.url is not None else "",
174+
"params": inner_params,
189175
}
190-
if self.options.greeting_message is not None:
191-
config["greeting_message"] = self.options.greeting_message
192-
if self.options.input_modalities is not None:
193-
config["input_modalities"] = self.options.input_modalities
194-
if self.options.output_modalities is not None:
195-
config["output_modalities"] = self.options.output_modalities
196-
if self.options.messages is not None:
197-
config["messages"] = self.options.messages
198-
if self.options.failure_message is not None:
199-
config["failure_message"] = self.options.failure_message
200-
if self.options.turn_detection is not None:
201-
config["turn_detection"] = self.options.turn_detection
176+
if self.greeting_message is not None:
177+
config["greeting_message"] = self.greeting_message
178+
if self.input_modalities is not None:
179+
config["input_modalities"] = self.input_modalities
180+
if self.output_modalities is not None:
181+
config["output_modalities"] = self.output_modalities
182+
if self.messages is not None:
183+
config["messages"] = self.messages
184+
if self.failure_message is not None:
185+
config["failure_message"] = self.failure_message
186+
if self.turn_detection is not None:
187+
config["turn_detection"] = self.turn_detection
202188

203189
return config
204190

205191

206-
class GeminiLiveOptions(BaseModel):
192+
class GeminiLive(BaseMLLM):
207193
model_config = ConfigDict(extra="forbid")
208194

209195
api_key: str = Field(..., description="Google API key")
@@ -224,47 +210,43 @@ class GeminiLiveOptions(BaseModel):
224210
turn_detection: Optional[MllmTurnDetectionConfig] = Field(default=None, description="MLLM turn detection configuration")
225211
failure_message: Optional[str] = Field(default=None, description="Message played on failure")
226212

227-
class GeminiLive(BaseMLLM):
228-
def __init__(self, **kwargs: Any):
229-
self.options = GeminiLiveOptions(**kwargs)
230-
231213
def to_config(self) -> Dict[str, Any]:
232-
params: Dict[str, Any] = {}
233-
if self.options.additional_params is not None:
234-
params.update(self.options.additional_params)
235-
params["model"] = self.options.model
236-
if self.options.instructions is not None:
237-
params["instructions"] = self.options.instructions
238-
if self.options.voice is not None:
239-
params["voice"] = self.options.voice
240-
if self.options.affective_dialog is not None:
241-
params["affective_dialog"] = self.options.affective_dialog
242-
if self.options.proactive_audio is not None:
243-
params["proactive_audio"] = self.options.proactive_audio
244-
if self.options.transcribe_agent is not None:
245-
params["transcribe_agent"] = self.options.transcribe_agent
246-
if self.options.transcribe_user is not None:
247-
params["transcribe_user"] = self.options.transcribe_user
248-
if self.options.http_options is not None:
249-
params["http_options"] = self.options.http_options
214+
inner_params: Dict[str, Any] = {}
215+
if self.additional_params is not None:
216+
inner_params.update(self.additional_params)
217+
inner_params["model"] = self.model
218+
if self.instructions is not None:
219+
inner_params["instructions"] = self.instructions
220+
if self.voice is not None:
221+
inner_params["voice"] = self.voice
222+
if self.affective_dialog is not None:
223+
inner_params["affective_dialog"] = self.affective_dialog
224+
if self.proactive_audio is not None:
225+
inner_params["proactive_audio"] = self.proactive_audio
226+
if self.transcribe_agent is not None:
227+
inner_params["transcribe_agent"] = self.transcribe_agent
228+
if self.transcribe_user is not None:
229+
inner_params["transcribe_user"] = self.transcribe_user
230+
if self.http_options is not None:
231+
inner_params["http_options"] = self.http_options
250232

251233
config: Dict[str, Any] = {
252234
"vendor": "gemini",
253-
"api_key": self.options.api_key,
254-
"url": self.options.url if self.options.url is not None else "",
255-
"params": params,
235+
"api_key": self.api_key,
236+
"url": self.url if self.url is not None else "",
237+
"params": inner_params,
256238
}
257-
if self.options.greeting_message is not None:
258-
config["greeting_message"] = self.options.greeting_message
259-
if self.options.input_modalities is not None:
260-
config["input_modalities"] = self.options.input_modalities
261-
if self.options.output_modalities is not None:
262-
config["output_modalities"] = self.options.output_modalities
263-
if self.options.messages is not None:
264-
config["messages"] = self.options.messages
265-
if self.options.failure_message is not None:
266-
config["failure_message"] = self.options.failure_message
267-
if self.options.turn_detection is not None:
268-
config["turn_detection"] = self.options.turn_detection
239+
if self.greeting_message is not None:
240+
config["greeting_message"] = self.greeting_message
241+
if self.input_modalities is not None:
242+
config["input_modalities"] = self.input_modalities
243+
if self.output_modalities is not None:
244+
config["output_modalities"] = self.output_modalities
245+
if self.messages is not None:
246+
config["messages"] = self.messages
247+
if self.failure_message is not None:
248+
config["failure_message"] = self.failure_message
249+
if self.turn_detection is not None:
250+
config["turn_detection"] = self.turn_detection
269251

270252
return config

0 commit comments

Comments
 (0)