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
31 changes: 31 additions & 0 deletions conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
from tests.llm.conftest import show_llm_summary_report
from holmes.core.tracing import readable_timestamp, get_active_branch_name
from tests.llm.utils.braintrust import get_braintrust_url
from unittest.mock import MagicMock, patch
import pytest


def pytest_addoption(parser):
Expand Down Expand Up @@ -126,3 +128,32 @@ def pytest_report_header(config):
# due to pytest quirks, we need to define this in the main conftest.py - when defined in the llm conftest.py it
# is SOMETIMES picked up and sometimes not, depending on how the test was invokedr
pytest_terminal_summary = show_llm_summary_report


@pytest.fixture(autouse=True)
def patch_supabase(monkeypatch):
monkeypatch.setattr("holmes.core.supabase_dal.ROBUSTA_ACCOUNT_ID", "test-cluster")
monkeypatch.setattr(
"holmes.core.supabase_dal.STORE_URL", "https://fakesupabaseref.supabase.co"
)
monkeypatch.setattr(
"holmes.core.supabase_dal.STORE_API_KEY",
"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJyb2xlIjoiYW5vbiIsImlhdCI6MTYzNTAwODQ4NywiZXhwIjoxOTUwNTg0NDg3fQ.l8IgkO7TQokGSc9OJoobXIVXsOXkilXl4Ak6SCX5qI8",
)
monkeypatch.setattr("holmes.core.supabase_dal.STORE_EMAIL", "mock_store_user")
monkeypatch.setattr(
"holmes.core.supabase_dal.STORE_PASSWORD", "mock_store_password"
)


@pytest.fixture(autouse=True, scope="session")
def storage_dal_mock():
with patch("holmes.config.SupabaseDal") as MockSupabaseDal:
mock_supabase_dal_instance = MagicMock()
MockSupabaseDal.return_value = mock_supabase_dal_instance
mock_supabase_dal_instance.sign_in.return_value = "mock_supabase_user_id"
mock_supabase_dal_instance.get_ai_credentials.return_value = (
"mock_account_id",
"mock_session_token",
)
yield mock_supabase_dal_instance
12 changes: 10 additions & 2 deletions holmes/clients/robusta_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,16 @@ class HolmesInfo(BaseModel):
latest_version: Optional[str] = None


class RobustaModelsResponse(BaseModel):
model_config = ConfigDict(extra="ignore")
models: List[str]
default_model: Optional[str] = None


@cache
def fetch_robusta_models(account_id, token) -> Optional[List[str]]:
def fetch_robusta_models(
account_id: str, token: str
) -> Optional[RobustaModelsResponse]:
try:
session_request = {"session_token": token, "account_id": account_id}
resp = requests.post(
Expand All @@ -25,7 +33,7 @@ def fetch_robusta_models(account_id, token) -> Optional[List[str]]:
)
resp.raise_for_status()
response_json = resp.json()
return response_json.get("models")
return RobustaModelsResponse(**response_json)
except Exception:
logging.exception("Failed to fetch robusta models")
return None
Expand Down
96 changes: 66 additions & 30 deletions holmes/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,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.tool_calling_llm import IssueInvestigator, ToolCallingLLM
from holmes.plugins.destinations.slack import SlackDestination
from holmes.plugins.sources.github import GitHubSource
Expand Down Expand Up @@ -135,6 +134,7 @@ class Config(RobustaBaseConfig):
_server_tool_executor: Optional[ToolExecutor] = None

_toolset_manager: Optional[ToolsetManager] = None
_default_robusta_model: Optional[str] = None

@property
def toolset_manager(self) -> ToolsetManager:
Expand Down Expand Up @@ -170,20 +170,28 @@ def configure_robusta_ai_model(self) -> None:
self._load_default_robusta_config()
return

models = fetch_robusta_models(
robusta_models = fetch_robusta_models(
self.account_id, self.session_token.get_secret_value()
)
if not models:
if not robusta_models or not robusta_models.models:
self._load_default_robusta_config()
return

for model in models:
for model in robusta_models.models:
logging.info(f"Loading Robusta AI model: {model}")
self._model_list[model] = {
"name": model,
"base_url": f"{ROBUSTA_API_ENDPOINT}/llm/{model}",
"is_robusta_model": True,
"model": "gpt-4o", # Robusta AI model is using openai like API.
Comment thread
moshemorad marked this conversation as resolved.
Outdated
}

if robusta_models.default_model:
logging.info(
f"Setting default Robusta AI model to: {robusta_models.default_model}"
)
self._default_robusta_model = robusta_models.default_model

except Exception:
logging.exception("Failed to get all robusta models")
# fallback to default behavior
Expand All @@ -193,9 +201,12 @@ def _load_default_robusta_config(self):
if self._should_load_robusta_ai() and self.api_key:
Comment thread
moshemorad marked this conversation as resolved.
Outdated
logging.info("Loading default Robusta AI model")
self._model_list[ROBUSTA_AI_MODEL_NAME] = {
"name": ROBUSTA_AI_MODEL_NAME,
"base_url": ROBUSTA_API_ENDPOINT,
"is_robusta_model": True,
"model": "gpt-4o",
}
Comment thread
moshemorad marked this conversation as resolved.
Outdated
self._default_robusta_model = ROBUSTA_AI_MODEL_NAME

def _should_load_robusta_ai(self) -> bool:
if not self.should_try_robusta_ai:
Expand Down Expand Up @@ -525,34 +536,59 @@ def create_slack_destination(self) -> "SlackDestination":
raise ValueError("--slack-channel must be specified")
return SlackDestination(self.slack_token.get_secret_value(), self.slack_channel)

def _get_llm(self, model_key: Optional[str] = None, tracer=None) -> "LLM":
api_key = self.api_key
model = self.model
def _get_model_params(self, model_key: Optional[str] = None) -> dict:
if not self._model_list:
logging.info("No model list setup, using config model")
return {}

if model_key:
model_params = self._model_list.get(model_key)
if model_params is not None:
logging.info(f"Using model: {model_key}")
return model_params.copy()

logging.error(f"Couldn't find model: {model_key} in model list")

if self._default_robusta_model:
model_params = self._model_list.get(self._default_robusta_model)
if model_params is not None:
logging.info(
f"Using default Robusta AI model: {self._default_robusta_model}"
)
return model_params.copy()

logging.error(
f"Couldn't find default Robusta AI model: {self._default_robusta_model} in model list"
)

first_model_params = next(iter(self._model_list.values())).copy()
logging.info("Using first model")
return first_model_params

def _get_llm(self, model_key: Optional[str] = None, tracer=None) -> "DefaultLLM":
model_params = self._get_model_params(model_key)
api_base = self.api_base
api_version = self.api_version
model_params = {}
if self._model_list:
# get requested model or the first credentials if no model requested.
model_params = (
self._model_list.get(model_key, {}).copy()
if model_key
else next(iter(self._model_list.values())).copy()
)
is_robusta_model = model_params.pop("is_robusta_model", False)
if is_robusta_model and self.api_key:
# we set here the api_key since it is being refresh when exprided and not as part of the model loading.
api_key = self.api_key.get_secret_value() # type: ignore
else:
api_key = model_params.pop("api_key", api_key)
model = model_params.pop("model", model)
# It's ok if the model does not have api base and api version, which are defaults to None.
# Handle both api_base and base_url - api_base takes precedence
model_api_base = model_params.pop("api_base", None)
model_base_url = model_params.pop("base_url", None)
api_base = model_api_base or model_base_url or api_base
api_version = model_params.pop("api_version", api_version)

return DefaultLLM(model, api_key, api_base, api_version, model_params, tracer) # type: ignore

is_robusta_model = model_params.pop("is_robusta_model", False)
if is_robusta_model and self.api_key:
# we set here the api_key since it is being refresh when exprided and not as part of the model loading.
api_key = self.api_key.get_secret_value() # type: ignore
else:
api_key = model_params.pop("api_key", None)

Comment thread
moshemorad marked this conversation as resolved.
Outdated
model = model_params.pop("model", self.model)
# It's ok if the model does not have api base and api version, which are defaults to None.
# Handle both api_base and base_url - api_base takes precedence
model_api_base = model_params.pop("api_base", None)
model_base_url = model_params.pop("base_url", None)
api_base = model_api_base or model_base_url or api_base
api_version = model_params.pop("api_version", api_version)
model_name = model_params.pop("name", None) or model_key or model

return DefaultLLM(
model, api_key, api_base, api_version, model_params, tracer, model_name
) # type: ignore

def get_models_list(self) -> List[str]:
if self._model_list:
Expand Down
4 changes: 3 additions & 1 deletion holmes/core/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,14 +72,16 @@ def __init__(
api_base: Optional[str] = None,
api_version: Optional[str] = None,
args: Optional[Dict] = None,
tracer=None,
tracer: Optional[Any] = None,
name: Optional[str] = None,
):
self.model = model
self.api_key = api_key
self.api_base = api_base
self.api_version = api_version
self.args = args or {}
self.tracer = tracer
self.name = name
Comment thread
moshemorad marked this conversation as resolved.
Outdated

self.check_llm(self.model, self.api_key, self.api_base, self.api_version)

Expand Down
19 changes: 18 additions & 1 deletion poetry.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@ pytest-cov = "^6.2.1"
types-python-dateutil = "^2.9.0.20250708"
pytest-dotenv = "^0.5.2"
pytest-shared-session-scope = "^0.4.0"
pytest-responses = "^0.5.1"

[build-system]
requires = ["poetry-core"]
Expand Down
15 changes: 11 additions & 4 deletions tests/config_class/test_config_api_base_version.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,6 @@ def test_config_get_llm_with_api_base_version():
"""Test that Config._get_llm passes api_base and api_version to DefaultLLM."""
config = Config(
model="test-model",
api_key="test-key",
api_base="https://test.api.base",
api_version="2023-12-01",
)
Expand All @@ -49,7 +48,7 @@ def test_config_get_llm_with_api_base_version():
# Check that DefaultLLM was called with the right positional arguments
call_args = mock_default_llm.call_args[0]
assert call_args[0] == "test-model"
assert call_args[1].get_secret_value() == "test-key" # api_key is SecretStr
assert call_args[1] is None
assert call_args[2] == "https://test.api.base"
assert call_args[3] == "2023-12-01"
assert call_args[4] == {}
Expand Down Expand Up @@ -85,7 +84,8 @@ def test_config_get_llm_with_model_list_api_base_version(monkeypatch, tmp_path):
"https://model.api.base",
"2024-02-01",
{},
None, # tracer
None, # tracer,
"test-model",
)
assert result == mock_llm_instance

Expand Down Expand Up @@ -122,6 +122,7 @@ def test_config_get_llm_model_list_overrides_config_values(monkeypatch, tmp_path
"2024-03-01", # from model list
{},
None, # tracer
"test-model",
)


Expand Down Expand Up @@ -156,6 +157,7 @@ def test_config_get_llm_model_list_defaults_to_config_values(monkeypatch, tmp_pa
"2023-01-01", # from config
{},
None, # tracer
"test-model",
)


Expand Down Expand Up @@ -202,6 +204,7 @@ def test_config_get_llm_with_non_none_model_list_first_model_fallback(
"2024-01-01", # from first model
{},
None, # tracer
"gpt-4",
)


Expand Down Expand Up @@ -255,7 +258,8 @@ def test_config_get_llm_with_specific_model_from_model_list(monkeypatch, tmp_pat
"https://openai.api.base", # from openai-gpt35 model
"2024-04-01", # from openai-gpt35 model
{},
None, # tracer
None, # tracer,
"openai-gpt35",
)


Expand Down Expand Up @@ -289,6 +293,7 @@ def test_config_get_llm_with_base_url_only(monkeypatch, tmp_path):
"2024-01-01",
{},
None, # tracer
"test-model",
)


Expand Down Expand Up @@ -323,6 +328,7 @@ def test_config_get_llm_api_base_overrides_base_url(monkeypatch, tmp_path):
"2024-01-01",
{},
None, # tracer
"test-model",
)


Expand Down Expand Up @@ -358,6 +364,7 @@ def test_config_get_llm_neither_api_base_nor_base_url_uses_config(
"2024-01-01",
{},
None, # tracer
"test-model",
)


Expand Down
Loading
Loading