Skip to content

Commit 8d10148

Browse files
committed
refactor: collapse LLM vendor configs; standalone copies for OpenAI-compatible LLMs
1 parent 2b49847 commit 8d10148

3 files changed

Lines changed: 748 additions & 281 deletions

File tree

src/agora_agent/agentkit/vendors/base.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
GoogleTTSSampleRate = Literal[8000, 16000, 22050, 24000, 44100, 48000]
1717

1818

19-
class BaseLLM(ABC):
19+
class BaseLLM(BaseModel, ABC):
2020
"""Abstract base class for all LLM vendor implementations.
2121
2222
Subclasses must implement :meth:`to_config` to return a dict that maps to

src/agora_agent/agentkit/vendors/cn.py

Lines changed: 327 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,13 @@
55
from pydantic import BaseModel, ConfigDict, Field, model_validator
66

77
from .avatar import BaseAvatar
8-
from .llm import OpenAI
8+
from .base import BaseLLM
9+
from .llm import (
10+
_OPENAI_MANAGED_MODELS,
11+
LlmGreetingConfigs,
12+
_dump_optional_model,
13+
_ensure_mcp_transport,
14+
)
915
from .stt import BaseSTT as _BaseSTTCompat
1016
from .tts import BaseTTS as _BaseTTSCompat
1117

@@ -497,28 +503,332 @@ def to_config(self) -> Dict[str, Any]:
497503
return result
498504

499505

500-
class AliyunLLM(OpenAI):
501-
def __init__(self, **kwargs: Any):
502-
kwargs["vendor"] = "aliyun"
503-
super().__init__(**kwargs)
506+
class AliyunLLM(BaseLLM):
507+
model_config = ConfigDict(extra="forbid")
504508

509+
api_key: Optional[str] = Field(default=None, description="OpenAI API key")
510+
model: str = Field(..., description="Model name")
511+
base_url: Optional[str] = Field(default=None, description="Custom base URL")
512+
temperature: Optional[float] = Field(default=None, ge=0.0, le=2.0)
513+
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
514+
max_tokens: Optional[int] = Field(default=None, gt=0)
515+
system_messages: Optional[List[Dict[str, Any]]] = Field(default=None)
516+
greeting_message: Optional[str] = Field(default=None)
517+
greeting_audio_url: Optional[str] = Field(default=None)
518+
failure_message: Optional[str] = Field(default=None)
519+
input_modalities: Optional[List[str]] = Field(default=None)
520+
params: Optional[Dict[str, Any]] = Field(default=None)
521+
headers: Optional[Dict[str, str]] = Field(default=None)
522+
output_modalities: Optional[List[str]] = Field(default=None)
523+
greeting_configs: Optional[LlmGreetingConfigs] = Field(default=None)
524+
template_variables: Optional[Dict[str, str]] = Field(default=None)
525+
vendor: Optional[str] = Field(default="aliyun")
526+
mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None)
527+
max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache")
505528

506-
class BytedanceLLM(OpenAI):
507-
def __init__(self, **kwargs: Any):
508-
kwargs["vendor"] = "bytedance"
509-
super().__init__(**kwargs)
529+
@model_validator(mode="after")
530+
def _validate_byok_params(self) -> "AliyunLLM":
531+
if not self.model:
532+
raise ValueError("OpenAI requires model")
533+
if self.api_key is not None and self.base_url is None:
534+
raise ValueError("OpenAI requires base_url when api_key is set")
535+
if self.api_key is None and self.base_url is not None:
536+
raise ValueError("OpenAI base_url is only valid when api_key is set")
537+
if self.api_key is None and self.model.strip().lower() not in _OPENAI_MANAGED_MODELS:
538+
raise ValueError("OpenAI requires api_key unless using a supported Agora-managed model")
539+
if self.api_key is None and self.vendor is not None:
540+
raise ValueError("OpenAI Agora-managed mode does not allow vendor")
541+
return self
510542

543+
def to_config(self) -> Dict[str, Any]:
544+
params: Dict[str, Any] = {"model": self.model, **(self.params or {})}
511545

512-
class DeepSeekLLM(OpenAI):
513-
def __init__(self, **kwargs: Any):
514-
kwargs["vendor"] = "deepseek"
515-
super().__init__(**kwargs)
546+
if self.max_tokens is not None:
547+
params["max_tokens"] = self.max_tokens
548+
if self.temperature is not None:
549+
params["temperature"] = self.temperature
550+
if self.top_p is not None:
551+
params["top_p"] = self.top_p
516552

553+
config: Dict[str, Any] = {
554+
"url": self.base_url or "https://api.openai.com/v1/chat/completions",
555+
"params": params,
556+
"style": "openai",
557+
"input_modalities": self.input_modalities or ["text"],
558+
}
559+
if self.api_key is not None:
560+
config["api_key"] = self.api_key
561+
if self.headers is not None:
562+
config["headers"] = self.headers
563+
564+
if self.system_messages is not None:
565+
config["system_messages"] = self.system_messages
566+
if self.greeting_message is not None:
567+
config["greeting_message"] = self.greeting_message
568+
if self.greeting_audio_url is not None:
569+
config["greeting_audio_url"] = self.greeting_audio_url
570+
if self.failure_message is not None:
571+
config["failure_message"] = self.failure_message
572+
if self.output_modalities is not None:
573+
config["output_modalities"] = self.output_modalities
574+
if self.greeting_configs is not None:
575+
config["greeting_configs"] = _dump_optional_model(self.greeting_configs)
576+
if self.template_variables is not None:
577+
config["template_variables"] = self.template_variables
578+
if self.vendor is not None:
579+
config["vendor"] = self.vendor
580+
if self.mcp_servers is not None:
581+
config["mcp_servers"] = _ensure_mcp_transport(self.mcp_servers)
582+
if self.max_history is not None:
583+
config["max_history"] = self.max_history
584+
585+
return config
586+
587+
588+
class BytedanceLLM(BaseLLM):
589+
model_config = ConfigDict(extra="forbid")
517590

518-
class TencentLLM(OpenAI):
519-
def __init__(self, **kwargs: Any):
520-
kwargs["vendor"] = "tencent"
521-
super().__init__(**kwargs)
591+
api_key: Optional[str] = Field(default=None, description="OpenAI API key")
592+
model: str = Field(..., description="Model name")
593+
base_url: Optional[str] = Field(default=None, description="Custom base URL")
594+
temperature: Optional[float] = Field(default=None, ge=0.0, le=2.0)
595+
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
596+
max_tokens: Optional[int] = Field(default=None, gt=0)
597+
system_messages: Optional[List[Dict[str, Any]]] = Field(default=None)
598+
greeting_message: Optional[str] = Field(default=None)
599+
greeting_audio_url: Optional[str] = Field(default=None)
600+
failure_message: Optional[str] = Field(default=None)
601+
input_modalities: Optional[List[str]] = Field(default=None)
602+
params: Optional[Dict[str, Any]] = Field(default=None)
603+
headers: Optional[Dict[str, str]] = Field(default=None)
604+
output_modalities: Optional[List[str]] = Field(default=None)
605+
greeting_configs: Optional[LlmGreetingConfigs] = Field(default=None)
606+
template_variables: Optional[Dict[str, str]] = Field(default=None)
607+
vendor: Optional[str] = Field(default="bytedance")
608+
mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None)
609+
max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache")
610+
611+
@model_validator(mode="after")
612+
def _validate_byok_params(self) -> "BytedanceLLM":
613+
if not self.model:
614+
raise ValueError("OpenAI requires model")
615+
if self.api_key is not None and self.base_url is None:
616+
raise ValueError("OpenAI requires base_url when api_key is set")
617+
if self.api_key is None and self.base_url is not None:
618+
raise ValueError("OpenAI base_url is only valid when api_key is set")
619+
if self.api_key is None and self.model.strip().lower() not in _OPENAI_MANAGED_MODELS:
620+
raise ValueError("OpenAI requires api_key unless using a supported Agora-managed model")
621+
if self.api_key is None and self.vendor is not None:
622+
raise ValueError("OpenAI Agora-managed mode does not allow vendor")
623+
return self
624+
625+
def to_config(self) -> Dict[str, Any]:
626+
params: Dict[str, Any] = {"model": self.model, **(self.params or {})}
627+
628+
if self.max_tokens is not None:
629+
params["max_tokens"] = self.max_tokens
630+
if self.temperature is not None:
631+
params["temperature"] = self.temperature
632+
if self.top_p is not None:
633+
params["top_p"] = self.top_p
634+
635+
config: Dict[str, Any] = {
636+
"url": self.base_url or "https://api.openai.com/v1/chat/completions",
637+
"params": params,
638+
"style": "openai",
639+
"input_modalities": self.input_modalities or ["text"],
640+
}
641+
if self.api_key is not None:
642+
config["api_key"] = self.api_key
643+
if self.headers is not None:
644+
config["headers"] = self.headers
645+
646+
if self.system_messages is not None:
647+
config["system_messages"] = self.system_messages
648+
if self.greeting_message is not None:
649+
config["greeting_message"] = self.greeting_message
650+
if self.greeting_audio_url is not None:
651+
config["greeting_audio_url"] = self.greeting_audio_url
652+
if self.failure_message is not None:
653+
config["failure_message"] = self.failure_message
654+
if self.output_modalities is not None:
655+
config["output_modalities"] = self.output_modalities
656+
if self.greeting_configs is not None:
657+
config["greeting_configs"] = _dump_optional_model(self.greeting_configs)
658+
if self.template_variables is not None:
659+
config["template_variables"] = self.template_variables
660+
if self.vendor is not None:
661+
config["vendor"] = self.vendor
662+
if self.mcp_servers is not None:
663+
config["mcp_servers"] = _ensure_mcp_transport(self.mcp_servers)
664+
if self.max_history is not None:
665+
config["max_history"] = self.max_history
666+
667+
return config
668+
669+
670+
class DeepSeekLLM(BaseLLM):
671+
model_config = ConfigDict(extra="forbid")
672+
673+
api_key: Optional[str] = Field(default=None, description="OpenAI API key")
674+
model: str = Field(..., description="Model name")
675+
base_url: Optional[str] = Field(default=None, description="Custom base URL")
676+
temperature: Optional[float] = Field(default=None, ge=0.0, le=2.0)
677+
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
678+
max_tokens: Optional[int] = Field(default=None, gt=0)
679+
system_messages: Optional[List[Dict[str, Any]]] = Field(default=None)
680+
greeting_message: Optional[str] = Field(default=None)
681+
greeting_audio_url: Optional[str] = Field(default=None)
682+
failure_message: Optional[str] = Field(default=None)
683+
input_modalities: Optional[List[str]] = Field(default=None)
684+
params: Optional[Dict[str, Any]] = Field(default=None)
685+
headers: Optional[Dict[str, str]] = Field(default=None)
686+
output_modalities: Optional[List[str]] = Field(default=None)
687+
greeting_configs: Optional[LlmGreetingConfigs] = Field(default=None)
688+
template_variables: Optional[Dict[str, str]] = Field(default=None)
689+
vendor: Optional[str] = Field(default="deepseek")
690+
mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None)
691+
max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache")
692+
693+
@model_validator(mode="after")
694+
def _validate_byok_params(self) -> "DeepSeekLLM":
695+
if not self.model:
696+
raise ValueError("OpenAI requires model")
697+
if self.api_key is not None and self.base_url is None:
698+
raise ValueError("OpenAI requires base_url when api_key is set")
699+
if self.api_key is None and self.base_url is not None:
700+
raise ValueError("OpenAI base_url is only valid when api_key is set")
701+
if self.api_key is None and self.model.strip().lower() not in _OPENAI_MANAGED_MODELS:
702+
raise ValueError("OpenAI requires api_key unless using a supported Agora-managed model")
703+
if self.api_key is None and self.vendor is not None:
704+
raise ValueError("OpenAI Agora-managed mode does not allow vendor")
705+
return self
706+
707+
def to_config(self) -> Dict[str, Any]:
708+
params: Dict[str, Any] = {"model": self.model, **(self.params or {})}
709+
710+
if self.max_tokens is not None:
711+
params["max_tokens"] = self.max_tokens
712+
if self.temperature is not None:
713+
params["temperature"] = self.temperature
714+
if self.top_p is not None:
715+
params["top_p"] = self.top_p
716+
717+
config: Dict[str, Any] = {
718+
"url": self.base_url or "https://api.openai.com/v1/chat/completions",
719+
"params": params,
720+
"style": "openai",
721+
"input_modalities": self.input_modalities or ["text"],
722+
}
723+
if self.api_key is not None:
724+
config["api_key"] = self.api_key
725+
if self.headers is not None:
726+
config["headers"] = self.headers
727+
728+
if self.system_messages is not None:
729+
config["system_messages"] = self.system_messages
730+
if self.greeting_message is not None:
731+
config["greeting_message"] = self.greeting_message
732+
if self.greeting_audio_url is not None:
733+
config["greeting_audio_url"] = self.greeting_audio_url
734+
if self.failure_message is not None:
735+
config["failure_message"] = self.failure_message
736+
if self.output_modalities is not None:
737+
config["output_modalities"] = self.output_modalities
738+
if self.greeting_configs is not None:
739+
config["greeting_configs"] = _dump_optional_model(self.greeting_configs)
740+
if self.template_variables is not None:
741+
config["template_variables"] = self.template_variables
742+
if self.vendor is not None:
743+
config["vendor"] = self.vendor
744+
if self.mcp_servers is not None:
745+
config["mcp_servers"] = _ensure_mcp_transport(self.mcp_servers)
746+
if self.max_history is not None:
747+
config["max_history"] = self.max_history
748+
749+
return config
750+
751+
752+
class TencentLLM(BaseLLM):
753+
model_config = ConfigDict(extra="forbid")
754+
755+
api_key: Optional[str] = Field(default=None, description="OpenAI API key")
756+
model: str = Field(..., description="Model name")
757+
base_url: Optional[str] = Field(default=None, description="Custom base URL")
758+
temperature: Optional[float] = Field(default=None, ge=0.0, le=2.0)
759+
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
760+
max_tokens: Optional[int] = Field(default=None, gt=0)
761+
system_messages: Optional[List[Dict[str, Any]]] = Field(default=None)
762+
greeting_message: Optional[str] = Field(default=None)
763+
greeting_audio_url: Optional[str] = Field(default=None)
764+
failure_message: Optional[str] = Field(default=None)
765+
input_modalities: Optional[List[str]] = Field(default=None)
766+
params: Optional[Dict[str, Any]] = Field(default=None)
767+
headers: Optional[Dict[str, str]] = Field(default=None)
768+
output_modalities: Optional[List[str]] = Field(default=None)
769+
greeting_configs: Optional[LlmGreetingConfigs] = Field(default=None)
770+
template_variables: Optional[Dict[str, str]] = Field(default=None)
771+
vendor: Optional[str] = Field(default="tencent")
772+
mcp_servers: Optional[List[Dict[str, Any]]] = Field(default=None)
773+
max_history: Optional[int] = Field(default=None, gt=0, description="Maximum number of conversation history messages to cache")
774+
775+
@model_validator(mode="after")
776+
def _validate_byok_params(self) -> "TencentLLM":
777+
if not self.model:
778+
raise ValueError("OpenAI requires model")
779+
if self.api_key is not None and self.base_url is None:
780+
raise ValueError("OpenAI requires base_url when api_key is set")
781+
if self.api_key is None and self.base_url is not None:
782+
raise ValueError("OpenAI base_url is only valid when api_key is set")
783+
if self.api_key is None and self.model.strip().lower() not in _OPENAI_MANAGED_MODELS:
784+
raise ValueError("OpenAI requires api_key unless using a supported Agora-managed model")
785+
if self.api_key is None and self.vendor is not None:
786+
raise ValueError("OpenAI Agora-managed mode does not allow vendor")
787+
return self
788+
789+
def to_config(self) -> Dict[str, Any]:
790+
params: Dict[str, Any] = {"model": self.model, **(self.params or {})}
791+
792+
if self.max_tokens is not None:
793+
params["max_tokens"] = self.max_tokens
794+
if self.temperature is not None:
795+
params["temperature"] = self.temperature
796+
if self.top_p is not None:
797+
params["top_p"] = self.top_p
798+
799+
config: Dict[str, Any] = {
800+
"url": self.base_url or "https://api.openai.com/v1/chat/completions",
801+
"params": params,
802+
"style": "openai",
803+
"input_modalities": self.input_modalities or ["text"],
804+
}
805+
if self.api_key is not None:
806+
config["api_key"] = self.api_key
807+
if self.headers is not None:
808+
config["headers"] = self.headers
809+
810+
if self.system_messages is not None:
811+
config["system_messages"] = self.system_messages
812+
if self.greeting_message is not None:
813+
config["greeting_message"] = self.greeting_message
814+
if self.greeting_audio_url is not None:
815+
config["greeting_audio_url"] = self.greeting_audio_url
816+
if self.failure_message is not None:
817+
config["failure_message"] = self.failure_message
818+
if self.output_modalities is not None:
819+
config["output_modalities"] = self.output_modalities
820+
if self.greeting_configs is not None:
821+
config["greeting_configs"] = _dump_optional_model(self.greeting_configs)
822+
if self.template_variables is not None:
823+
config["template_variables"] = self.template_variables
824+
if self.vendor is not None:
825+
config["vendor"] = self.vendor
826+
if self.mcp_servers is not None:
827+
config["mcp_servers"] = _ensure_mcp_transport(self.mcp_servers)
828+
if self.max_history is not None:
829+
config["max_history"] = self.max_history
830+
831+
return config
522832

523833

524834
class SenseTimeAvatarOptions(BaseModel):

0 commit comments

Comments
 (0)