| 23 | |
| 24 | @dataclass(slots=True) |
| 25 | class LLMConfig: |
| 26 | default_provider: str = "openai" |
| 27 | default_model: str = "" |
| 28 | temperature: float = 0.0 |
| 29 | providers: dict[str, LLMProviderConfig] = field(default_factory=dict) |
| 30 | |
| 31 | def get_provider(self, provider_name: str) -> LLMProviderConfig: |
| 32 | source = self.providers.get(provider_name, LLMProviderConfig()) |
| 33 | return LLMProviderConfig( |
| 34 | api_base=source.api_base, |
| 35 | api_key=source.api_key, |
| 36 | temperature=self.temperature if source.temperature is None else source.temperature, |
| 37 | ) |
| 38 | |
| 39 | def as_legacy_provider_config(self) -> ModelProviderConfig: |
| 40 | openai = self.get_provider("openai") |
| 41 | nvidia = self.get_provider("nvidia") |
| 42 | together = self.get_provider("together") |
| 43 | return ModelProviderConfig( |
| 44 | openai_base=openai.api_base, |
| 45 | openai_key=openai.api_key, |
| 46 | nvidia_key=nvidia.api_key, |
| 47 | together_key=together.api_key, |
| 48 | temperature=self.temperature, |
| 49 | ) |
| 50 | |
| 51 | |
| 52 | @dataclass(slots=True) |