Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 18 additions & 7 deletions src/agora_agent/agentkit/vendors/avatar.py
Original file line number Diff line number Diff line change
@@ -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")
Expand All @@ -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
Expand All @@ -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")
Expand Down Expand Up @@ -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
Expand All @@ -115,14 +120,16 @@ 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")
avatar_id: Optional[str] = Field(default=None, description="Avatar ID")
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
Expand All @@ -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")
Expand All @@ -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
Expand All @@ -179,14 +188,16 @@ 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")
avatar_id: str = Field(..., description="Anam avatar ID")
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
Expand Down
Loading
Loading