diff --git a/holmes/clients/robusta_client.py b/holmes/clients/robusta_client.py index 94950af64c..e99722aee5 100644 --- a/holmes/clients/robusta_client.py +++ b/holmes/clients/robusta_client.py @@ -1,4 +1,5 @@ -from typing import Optional +import logging +from typing import List, Optional import requests # type: ignore from functools import cache from pydantic import BaseModel, ConfigDict @@ -13,6 +14,23 @@ class HolmesInfo(BaseModel): latest_version: Optional[str] = None +@cache +def fetch_robusta_models(account_id, token) -> Optional[List[str]]: + try: + session_request = {"session_token": token, "account_id": account_id} + resp = requests.post( + f"{ROBUSTA_API_ENDPOINT}/api/llm/models", + json=session_request, + timeout=10, + ) + resp.raise_for_status() + response_json = resp.json() + return response_json.get("models") + except Exception: + logging.exception("Failed to fetch robusta models") + return None + + @cache def fetch_holmes_info() -> Optional[HolmesInfo]: try: diff --git a/holmes/common/env_vars.py b/holmes/common/env_vars.py index 6077635ef7..aa4602cd17 100644 --- a/holmes/common/env_vars.py +++ b/holmes/common/env_vars.py @@ -27,6 +27,7 @@ def load_bool(env_var, default: Optional[bool]) -> Optional[bool]: STORE_PASSWORD = os.environ.get("STORE_PASSWORD", "") HOLMES_POST_PROCESSING_PROMPT = os.environ.get("HOLMES_POST_PROCESSING_PROMPT", "") ROBUSTA_AI = load_bool("ROBUSTA_AI", None) +LOAD_ALL_ROBUSTA_MODELS = load_bool("LOAD_ALL_ROBUSTA_MODELS", True) ROBUSTA_API_ENDPOINT = os.environ.get("ROBUSTA_API_ENDPOINT", "https://api.robusta.dev") LOG_PERFORMANCE = os.environ.get("LOG_PERFORMANCE", None) diff --git a/holmes/config.py b/holmes/config.py index 93fc74605d..0613831f10 100644 --- a/holmes/config.py +++ b/holmes/config.py @@ -10,8 +10,14 @@ from pydantic import BaseModel, ConfigDict, FilePath, SecretStr +from holmes.clients.robusta_client import fetch_robusta_models from holmes.core.llm import DefaultLLM -from holmes.common.env_vars import ROBUSTA_AI, ROBUSTA_API_ENDPOINT, ROBUSTA_CONFIG_PATH +from holmes.common.env_vars import ( + ROBUSTA_AI, + LOAD_ALL_ROBUSTA_MODELS, + ROBUSTA_API_ENDPOINT, + ROBUSTA_CONFIG_PATH, +) from holmes.core.tools_utils.tool_executor import ToolExecutor from holmes.core.toolset_manager import ToolsetManager from holmes.plugins.runbooks import ( @@ -24,7 +30,6 @@ # Source plugin imports moved to their respective create methods to speed up startup if TYPE_CHECKING: from holmes.core.llm import LLM - from holmes.core.supabase_dal import SupabaseDal from holmes.core.tool_calling_llm import IssueInvestigator, ToolCallingLLM from holmes.plugins.destinations.slack import SlackDestination from holmes.plugins.sources.github import GitHubSource @@ -33,6 +38,7 @@ from holmes.plugins.sources.pagerduty import PagerDutySource from holmes.plugins.sources.prometheus.plugin import AlertManagerSource +from holmes.core.supabase_dal import SupabaseDal from holmes.core.config import config_path_dir from holmes.utils.definitions import RobustaConfig from holmes.utils.env import replace_env_vars_values @@ -73,6 +79,9 @@ class Config(RobustaBaseConfig): api_key: Optional[SecretStr] = ( None # if None, read from OPENAI_API_KEY or AZURE_OPENAI_ENDPOINT env var ) + account_id: Optional[str] = None + session_token: Optional[SecretStr] = None + model: Optional[str] = "gpt-4o" max_steps: int = 40 cluster_name: Optional[str] = None @@ -136,18 +145,51 @@ def toolset_manager(self) -> ToolsetManager: def model_post_init(self, __context: Any) -> None: self._model_list = parse_models_file(MODEL_LIST_FILE_LOCATION) - if self._should_load_robusta_ai(): - logging.info("Loading Robusta AI model") - self._model_list[ROBUSTA_AI_MODEL_NAME] = { - "base_url": ROBUSTA_API_ENDPOINT, - } + + if not self._should_load_robusta_ai(): + return + + self.configure_robusta_ai_model() def configure_robusta_ai_model(self) -> None: + try: + if not self.cluster_name or not LOAD_ALL_ROBUSTA_MODELS: + self._load_default_robusta_config() + return + + if not self.api_key: + dal = SupabaseDal(self.cluster_name) + self.load_robusta_api_key(dal) + + if not self.account_id or not self.session_token: + self._load_default_robusta_config() + return + + models = fetch_robusta_models( + self.account_id, self.session_token.get_secret_value() + ) + if not models: + self._load_default_robusta_config() + return + + for model in models: + logging.info(f"Loading Robusta AI model: {model}") + self._model_list[model] = { + "base_url": f"{ROBUSTA_API_ENDPOINT}/llm/{model}", + "is_robusta_model": True, + } + + except Exception: + logging.exception("Failed to get all robusta models") + # fallback to default behavior + self._load_default_robusta_config() + + def _load_default_robusta_config(self): if self._should_load_robusta_ai() and self.api_key: - logging.info("Loading Robusta AI model") + logging.info("Loading default Robusta AI model") self._model_list[ROBUSTA_AI_MODEL_NAME] = { "base_url": ROBUSTA_API_ENDPOINT, - "api_key": self.api_key.get_secret_value(), + "is_robusta_model": True, } def _should_load_robusta_ai(self) -> bool: @@ -485,7 +527,10 @@ def _get_llm(self, model_key: Optional[str] = None, tracer=None) -> "LLM": if model_key else next(iter(self._model_list.values())).copy() ) - api_key = model_params.pop("api_key", api_key) + if model_params.get("is_robusta_model") and self.api_key: + api_key = self.api_key.get_secret_value() + else: + api_key = model_params.pop("api_key", api_key) model = model_params.pop("model", model) return DefaultLLM(model, api_key, model_params, tracer) # type: ignore @@ -496,6 +541,13 @@ def get_models_list(self) -> List[str]: return json.dumps([self.model]) # type: ignore + def load_robusta_api_key(self, dal: SupabaseDal): + if ROBUSTA_AI: + account_id, token = dal.get_ai_credentials() + self.api_key = SecretStr(f"{account_id} {token}") + self.account_id = account_id + self.session_token = SecretStr(token) + class TicketSource(BaseModel): config: Config diff --git a/holmes/core/investigation.py b/holmes/core/investigation.py index f44be456a5..1440b63891 100644 --- a/holmes/core/investigation.py +++ b/holmes/core/investigation.py @@ -9,7 +9,6 @@ from holmes.core.supabase_dal import SupabaseDal from holmes.core.tracing import DummySpan, SpanType from holmes.utils.global_instructions import add_global_instructions_to_user_prompt -from holmes.utils.robusta import load_robusta_api_key from holmes.core.todo_manager import get_todo_manager from holmes.core.investigation_structured_output import ( @@ -28,7 +27,7 @@ def investigate_issues( model: Optional[str] = None, trace_span=DummySpan(), ) -> InvestigationResult: - load_robusta_api_key(dal=dal, config=config) + config.load_robusta_api_key(dal=dal) context = dal.get_issue_data(investigate_request.context.get("robusta_issue_id")) resource_instructions = dal.get_resource_instructions( @@ -82,7 +81,7 @@ def get_investigation_context( config: Config, request_structured_output_from_llm: Optional[bool] = None, ): - load_robusta_api_key(dal=dal, config=config) + config.load_robusta_api_key(dal=dal) ai = config.create_issue_investigator(dal=dal, model=investigate_request.model) raw_data = investigate_request.model_dump() diff --git a/holmes/utils/robusta.py b/holmes/utils/robusta.py deleted file mode 100644 index 350369b28f..0000000000 --- a/holmes/utils/robusta.py +++ /dev/null @@ -1,10 +0,0 @@ -from holmes.config import Config, ROBUSTA_AI_MODEL_NAME -from holmes.core.supabase_dal import SupabaseDal -from pydantic import SecretStr - - -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() diff --git a/server.py b/server.py index 91a2d90653..27ec89cf67 100644 --- a/server.py +++ b/server.py @@ -23,7 +23,6 @@ from litellm.exceptions import AuthenticationError from fastapi import FastAPI, HTTPException, Request from fastapi.responses import StreamingResponse -from holmes.utils.robusta import load_robusta_api_key from holmes.utils.stream import stream_investigate_formatter, stream_chat_formatter from holmes.common.env_vars import ( HOLMES_HOST, @@ -82,6 +81,7 @@ def init_logging(): def sync_before_server_start(): + config.load_robusta_api_key(dal=dal) try: update_holmes_status_in_db(dal, config) except Exception: @@ -184,7 +184,7 @@ def stream_investigate_issues(req: InvestigateRequest): @app.post("/api/workload_health_check") def workload_health_check(request: WorkloadHealthRequest): - load_robusta_api_key(dal=dal, config=config) + config.load_robusta_api_key(dal=dal) try: resource = request.resource workload_alerts: list[str] = [] @@ -251,7 +251,7 @@ def workload_health_conversation( request: WorkloadHealthChatRequest, ): try: - load_robusta_api_key(dal=dal, config=config) + config.load_robusta_api_key(dal=dal) ai = config.create_toolcalling_llm(dal=dal, model=request.model) global_instructions = dal.get_global_instructions_for_account() @@ -280,7 +280,7 @@ def workload_health_conversation( @app.post("/api/issue_chat") def issue_conversation(issue_chat_request: IssueChatRequest): try: - load_robusta_api_key(dal=dal, config=config) + config.load_robusta_api_key(dal=dal) ai = config.create_toolcalling_llm(dal=dal, model=issue_chat_request.model) global_instructions = dal.get_global_instructions_for_account() @@ -319,7 +319,7 @@ def already_answered(conversation_history: Optional[List[dict]]) -> bool: @app.post("/api/chat") def chat(chat_request: ChatRequest): try: - load_robusta_api_key(dal=dal, config=config) + config.load_robusta_api_key(dal=dal) ai = config.create_toolcalling_llm(dal=dal, model=chat_request.model) global_instructions = dal.get_global_instructions_for_account() diff --git a/tests/config_class/test_config_load_robusta_ai.py b/tests/config_class/test_config_load_robusta_ai.py index 8e1748c36d..ab1797f93e 100644 --- a/tests/config_class/test_config_load_robusta_ai.py +++ b/tests/config_class/test_config_load_robusta_ai.py @@ -1,83 +1,138 @@ from unittest.mock import patch +from pydantic import SecretStr from holmes.config import Config +def fake_load_robusta_api_key(config, _): + config.account_id = "mock-account" + config.session_token = SecretStr("mock-token") + config.api_key = SecretStr("mock-token") + + @patch("holmes.config.ROBUSTA_AI", True) -def test_cli_not_loading_robusta_ai(monkeypatch): +def test_cli_not_loading_robusta_ai(*, monkeypatch): config = Config.load_from_file(None) assert "Robusta" not in config._model_list @patch("holmes.config.ROBUSTA_AI", True) -def test_server_loads_robusta_ai_when_true(monkeypatch): +@patch("holmes.config.fetch_robusta_models", return_value=["Robusta/test"]) +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") +def test_server_loads_robusta_ai_when_true(mock_cluster, mock_fetch, *, monkeypatch): + def fake_loader(self, dal): + self.account_id = "mock-account" + self.session_token = SecretStr("mock-token") + self.api_key = SecretStr("mock-token") + + monkeypatch.setattr(Config, "load_robusta_api_key", fake_loader) config = Config.load_from_env() - assert "Robusta" in config._model_list + assert "Robusta/test" in config._model_list @patch("holmes.config.ROBUSTA_AI", None) -def test_server_loads_robusta_ai_when_not_exists_and_not_other_models(monkeypatch): +@patch("holmes.config.fetch_robusta_models", return_value=["Robusta/test"]) +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") +def test_server_loads_robusta_ai_when_not_exists_and_not_other_models( + mock_cluster, mock_fetch, *, monkeypatch +): + def fake_loader(self, dal): + self.account_id = "mock-account" + self.session_token = SecretStr("mock-token") + self.api_key = SecretStr("mock-token") + + monkeypatch.setattr(Config, "load_robusta_api_key", fake_loader) config = Config.load_from_env() - assert len(config._model_list) == 1 - assert "Robusta" in config._model_list + assert "Robusta/test" in config._model_list @patch("holmes.config.ROBUSTA_AI", False) -def test_server_not_loads_robusta_ai_when_false(monkeypatch): +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") +def test_server_not_loads_robusta_ai_when_false(mock_cluster, *, monkeypatch): config = Config.load_from_env() - assert len(config._model_list) == 0 assert "Robusta" not in config._model_list @patch("holmes.config.ROBUSTA_AI", True) +@patch("holmes.config.fetch_robusta_models", return_value=["Robusta/test"]) +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") @patch( "holmes.config.parse_models_file", return_value={"existing_model": {"base_url": "http://foo"}}, ) -def test_server_loads_robusta_ai_when_true_and_model_list_exists(monkeypatch): +def test_server_loads_robusta_ai_when_true_and_model_list_exists( + mock_parse, mock_cluster, mock_fetch, *, monkeypatch +): + def fake_loader(self, dal): + self.account_id = "mock-account" + self.session_token = SecretStr("mock-token") + self.api_key = SecretStr("mock-token") + + monkeypatch.setattr(Config, "load_robusta_api_key", fake_loader) config = Config.load_from_env() assert "existing_model" in config._model_list - assert "Robusta" in config._model_list + assert "Robusta/test" in config._model_list @patch("holmes.config.ROBUSTA_AI", False) +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") @patch( "holmes.config.parse_models_file", return_value={"existing_model": {"base_url": "http://foo"}}, ) -def test_server_not_loads_robusta_ai_when_false_and_model_list_exists(monkeypatch): +def test_server_not_loads_robusta_ai_when_false_and_model_list_exists( + mock_parse, mock_cluster, *, monkeypatch +): config = Config.load_from_env() assert "existing_model" in config._model_list assert "Robusta" not in config._model_list @patch("holmes.config.ROBUSTA_AI", None) +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") @patch( "holmes.config.parse_models_file", return_value={"existing_model": {"base_url": "http://foo"}}, ) -def test_server_not_loads_robusta_ai_when_no_env_var_and_model_list_exists(monkeypatch): +def test_server_not_loads_robusta_ai_when_no_env_var_and_model_list_exists( + mock_parse, mock_cluster, *, monkeypatch +): config = Config.load_from_env() assert "existing_model" in config._model_list assert "Robusta" not in config._model_list @patch("holmes.config.ROBUSTA_AI", True) -def test_server_loads_robusta_ai_when_model_var_exists(monkeypatch): +@patch("holmes.config.fetch_robusta_models", return_value=["Robusta/test"]) +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") +def test_server_loads_robusta_ai_when_model_var_exists( + mock_cluster, mock_fetch, *, monkeypatch +): monkeypatch.setenv("MODEL", "some_model") + + def fake_loader(self, dal): + self.account_id = "mock-account" + self.session_token = SecretStr("mock-token") + self.api_key = SecretStr("mock-token") + + monkeypatch.setattr(Config, "load_robusta_api_key", fake_loader) config = Config.load_from_env() - assert "Robusta" in config._model_list + assert "Robusta/test" in config._model_list @patch("holmes.config.ROBUSTA_AI", None) -def test_server_not_loads_robusta_ai_when_model_var_exists_and_no_env_var(monkeypatch): +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") +def test_server_not_loads_robusta_ai_when_model_var_exists_and_no_env_var( + mock_cluster, *, monkeypatch +): monkeypatch.setenv("MODEL", "some_model") config = Config.load_from_env() assert "Robusta" not in config._model_list @patch("holmes.config.ROBUSTA_AI", False) +@patch("holmes.config.Config._Config__get_cluster_name", return_value="test") def test_server_not_loads_robusta_ai_when_model_var_exists_and_false_env_var( - monkeypatch, + mock_cluster, *, monkeypatch ): monkeypatch.setenv("MODEL", "some_model") config = Config.load_from_env() diff --git a/tests/test_server_endpoints.py b/tests/test_server_endpoints.py index 673bb74283..3565a6b080 100644 --- a/tests/test_server_endpoints.py +++ b/tests/test_server_endpoints.py @@ -9,7 +9,7 @@ def client(): return TestClient(app) -@patch("server.load_robusta_api_key") +@patch("holmes.config.Config.load_robusta_api_key") @patch("holmes.config.Config.create_toolcalling_llm") @patch("holmes.core.supabase_dal.SupabaseDal.get_global_instructions_for_account") def test_api_chat_all_fields( @@ -78,7 +78,7 @@ def test_api_chat_all_fields( assert "pre_action_notification_text" in action -@patch("server.load_robusta_api_key") +@patch("holmes.config.Config.load_robusta_api_key") @patch("holmes.config.Config.create_toolcalling_llm") @patch("holmes.core.supabase_dal.SupabaseDal.get_global_instructions_for_account") def test_api_issue_chat_all_fields( @@ -140,7 +140,7 @@ def test_api_issue_chat_all_fields( assert "result" in tool_call -@patch("server.load_robusta_api_key") +@patch("holmes.config.Config.load_robusta_api_key") @patch("holmes.config.Config.create_toolcalling_llm") @patch("holmes.core.supabase_dal.SupabaseDal.get_global_instructions_for_account") def test_api_workload_health_chat( @@ -205,7 +205,7 @@ def test_api_workload_health_chat( assert "result" in tool_call -@patch("server.load_robusta_api_key") +@patch("holmes.config.Config.load_robusta_api_key") @patch("holmes.config.Config.create_toolcalling_llm") @patch("holmes.core.supabase_dal.SupabaseDal.get_global_instructions_for_account") @patch("holmes.core.supabase_dal.SupabaseDal.get_workload_issues")