|
5 | 5 | from pydantic import BaseModel, ConfigDict, Field, model_validator |
6 | 6 |
|
7 | 7 | 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 | +) |
9 | 15 | from .stt import BaseSTT as _BaseSTTCompat |
10 | 16 | from .tts import BaseTTS as _BaseTTSCompat |
11 | 17 |
|
@@ -497,28 +503,332 @@ def to_config(self) -> Dict[str, Any]: |
497 | 503 | return result |
498 | 504 |
|
499 | 505 |
|
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") |
504 | 508 |
|
| 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") |
505 | 528 |
|
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 |
510 | 542 |
|
| 543 | + def to_config(self) -> Dict[str, Any]: |
| 544 | + params: Dict[str, Any] = {"model": self.model, **(self.params or {})} |
511 | 545 |
|
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 |
516 | 552 |
|
| 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") |
517 | 590 |
|
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 |
522 | 832 |
|
523 | 833 |
|
524 | 834 | class SenseTimeAvatarOptions(BaseModel): |
|
0 commit comments