Skip to content

Commit 26706d7

Browse files
feat(agentkit): add GenericAvatar and session-aware avatar validation
Adds the GenericAvatar vendor wrapper and extends avatar validation helpers for generic and RTC-backed avatars. Session-derived fields such as agora_appid, agora_channel, and agora_token can now be validated after AgentSession enrichment.
1 parent 9df782b commit 26706d7

2 files changed

Lines changed: 76 additions & 1 deletion

File tree

src/agora_agent/agentkit/avatar_types.py

Lines changed: 34 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,21 @@ def is_anam_avatar(config: typing.Dict[str, typing.Any]) -> bool:
1717
return config.get("vendor") == "anam"
1818

1919

20-
def validate_avatar_config(config: typing.Dict[str, typing.Any]) -> None:
20+
def is_generic_avatar(config: typing.Dict[str, typing.Any]) -> bool:
21+
return config.get("vendor") == "generic"
22+
23+
24+
def is_rtc_avatar(config: typing.Dict[str, typing.Any]) -> bool:
25+
params = config.get("params", {})
26+
return isinstance(params, dict) and bool(params.get("agora_uid")) and (
27+
is_heygen_avatar(config) or is_live_avatar_avatar(config) or is_generic_avatar(config)
28+
)
29+
30+
31+
def validate_avatar_config(
32+
config: typing.Dict[str, typing.Any],
33+
require_session_fields: bool = False,
34+
) -> None:
2135
"""Validates avatar configuration at runtime.
2236
2337
Parameters
@@ -45,6 +59,8 @@ def validate_avatar_config(config: typing.Dict[str, typing.Any]) -> None:
4559
f"Invalid quality for {label}: {params.get('quality')}. "
4660
f"Must be one of: {', '.join(valid_qualities)}"
4761
)
62+
if require_session_fields and not params.get("agora_token"):
63+
raise ValueError(f"{label} avatar requires agora_token after session enrichment")
4864
elif is_akool_avatar(config):
4965
params = config.get("params", {})
5066
if not params.get("api_key"):
@@ -53,6 +69,23 @@ def validate_avatar_config(config: typing.Dict[str, typing.Any]) -> None:
5369
params = config.get("params", {})
5470
if not params.get("api_key"):
5571
raise ValueError("Anam avatar requires api_key")
72+
elif is_generic_avatar(config):
73+
params = config.get("params", {})
74+
if not params.get("api_key"):
75+
raise ValueError("Generic avatar requires api_key")
76+
if not params.get("api_base_url"):
77+
raise ValueError("Generic avatar requires api_base_url")
78+
if not params.get("avatar_id"):
79+
raise ValueError("Generic avatar requires avatar_id")
80+
if not params.get("agora_uid"):
81+
raise ValueError("Generic avatar requires agora_uid")
82+
if require_session_fields:
83+
if not params.get("agora_token"):
84+
raise ValueError("Generic avatar requires agora_token after session enrichment")
85+
if not params.get("agora_appid"):
86+
raise ValueError("Generic avatar requires agora_appid after session enrichment")
87+
if not params.get("agora_channel"):
88+
raise ValueError("Generic avatar requires agora_channel after session enrichment")
5689

5790

5891
def validate_tts_sample_rate(

src/agora_agent/agentkit/vendors/avatar.py

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -132,6 +132,48 @@ def to_config(self) -> Dict[str, Any]:
132132
return {"enable": enable, "vendor": "liveavatar", "params": params}
133133

134134

135+
class GenericAvatarOptions(BaseModel):
136+
model_config = ConfigDict(extra="forbid")
137+
138+
api_key: str = Field(..., description="Generic avatar provider API key")
139+
api_base_url: str = Field(..., description="Avatar provider API base URL")
140+
avatar_id: str = Field(..., description="Avatar ID")
141+
agora_uid: str = Field(..., description="Agora UID for the avatar video stream")
142+
agora_appid: Optional[str] = Field(default=None, description="Agora App ID; filled by AgentSession when omitted")
143+
agora_token: Optional[str] = Field(default=None, description="RTC token; generated by AgentSession when omitted")
144+
agora_channel: Optional[str] = Field(default=None, description="Agora channel; filled by AgentSession when omitted")
145+
enable: Optional[bool] = Field(default=None, description="Enable avatar (default: true)")
146+
additional_params: Optional[Dict[str, Any]] = Field(default=None, description="Additional vendor-specific parameters")
147+
148+
class GenericAvatar(BaseAvatar):
149+
def __init__(self, **kwargs: Any):
150+
self.options = GenericAvatarOptions(**kwargs)
151+
152+
@property
153+
def required_sample_rate(self) -> int:
154+
return 0
155+
156+
def to_config(self) -> Dict[str, Any]:
157+
params: Dict[str, Any] = {
158+
"api_key": self.options.api_key,
159+
"api_base_url": self.options.api_base_url,
160+
"avatar_id": self.options.avatar_id,
161+
"agora_uid": self.options.agora_uid,
162+
}
163+
164+
if self.options.agora_appid is not None:
165+
params["agora_appid"] = self.options.agora_appid
166+
if self.options.agora_token is not None:
167+
params["agora_token"] = self.options.agora_token
168+
if self.options.agora_channel is not None:
169+
params["agora_channel"] = self.options.agora_channel
170+
if self.options.additional_params is not None:
171+
params = {**self.options.additional_params, **params}
172+
173+
enable = self.options.enable if self.options.enable is not None else True
174+
return {"enable": enable, "vendor": "generic", "params": params}
175+
176+
135177
class AnamAvatarOptions(BaseModel):
136178
model_config = ConfigDict(extra="forbid")
137179

0 commit comments

Comments
 (0)