Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 19 additions & 1 deletion holmes/clients/robusta_client.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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:
Comment thread
Avi-Robusta marked this conversation as resolved.
logging.exception("Failed to fetch robusta models")
return None

Comment thread
Avi-Robusta marked this conversation as resolved.

@cache
def fetch_holmes_info() -> Optional[HolmesInfo]:
try:
Expand Down
1 change: 1 addition & 0 deletions holmes/common/env_vars.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
72 changes: 62 additions & 10 deletions holmes/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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

Comment thread
Avi-Robusta marked this conversation as resolved.
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
Comment thread
moshemorad marked this conversation as resolved.

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}")
Comment thread
Avi-Robusta marked this conversation as resolved.
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:
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
5 changes: 2 additions & 3 deletions holmes/core/investigation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand All @@ -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(
Expand Down Expand Up @@ -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()
Expand Down
10 changes: 0 additions & 10 deletions holmes/utils/robusta.py

This file was deleted.

10 changes: 5 additions & 5 deletions server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Comment thread
Avi-Robusta marked this conversation as resolved.
except Exception:
Expand Down Expand Up @@ -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] = []
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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()
Expand Down
Loading