diff --git a/holmes/config.py b/holmes/config.py index 1b10347ada..71b8c6cdc1 100644 --- a/holmes/config.py +++ b/holmes/config.py @@ -142,6 +142,14 @@ def model_post_init(self, __context: Any) -> None: "base_url": ROBUSTA_API_ENDPOINT, } + def configure_robusta_ai_model(self) -> None: + if self._should_load_robusta_ai() and self.api_key: + logging.info("Loading Robusta AI model") + self._model_list[ROBUSTA_AI_MODEL_NAME] = { + "base_url": ROBUSTA_API_ENDPOINT, + "api_key": self.api_key.get_secret_value(), + } + def _should_load_robusta_ai(self) -> bool: if not self.should_try_robusta_ai: return False @@ -479,12 +487,6 @@ def _get_llm(self, model_key: Optional[str] = None, tracer=None) -> "LLM": ) api_key = model_params.pop("api_key", api_key) model = model_params.pop("model", model) - if ( - not api_key - and "robusta.dev" in model_params.get("base_url", "") - and self.api_key - ): - api_key = self.api_key.get_secret_value() return DefaultLLM(model, api_key, model_params, tracer) # type: ignore diff --git a/holmes/utils/robusta.py b/holmes/utils/robusta.py index 57c320fa30..350369b28f 100644 --- a/holmes/utils/robusta.py +++ b/holmes/utils/robusta.py @@ -7,3 +7,4 @@ def load_robusta_api_key(dal: SupabaseDal, config: Config): if ROBUSTA_AI_MODEL_NAME in config._model_list: account_id, token = dal.get_ai_credentials() config.api_key = SecretStr(f"{account_id} {token}") + config.configure_robusta_ai_model()