diff --git a/src/robusta/core/model/env_vars.py b/src/robusta/core/model/env_vars.py index 2f3fcbca2..d7260b7a7 100644 --- a/src/robusta/core/model/env_vars.py +++ b/src/robusta/core/model/env_vars.py @@ -136,6 +136,8 @@ def load_bool(env_var, default: bool): NAMESPACE_DATA_TTL = int(os.environ.get("NAMESPACE_DATA_TTL", 30 * 60)) # in seconds +NODE_IP_CACHE_TTL_SEC = int(os.environ.get("NODE_IP_CACHE_TTL_SEC", 15 * 60)) + PROCESSED_ALERTS_CACHE_TTL = int(os.environ.get("PROCESSED_ALERT_CACHE_TTL", 2 * 3600)) PROCESSED_ALERTS_CACHE_MAX_SIZE = int(os.environ.get("PROCESSED_ALERTS_CACHE_MAX_SIZE", 100_000)) diff --git a/src/robusta/integrations/prometheus/trigger.py b/src/robusta/integrations/prometheus/trigger.py index 79e1f450a..f79ef4cad 100644 --- a/src/robusta/integrations/prometheus/trigger.py +++ b/src/robusta/integrations/prometheus/trigger.py @@ -1,9 +1,11 @@ import logging +import time from typing import Any, Dict, List, NamedTuple, Optional, Type, Union from hikaru.model.rel_1_26 import DaemonSet, HorizontalPodAutoscaler, Job, Node, NodeList, StatefulSet from pydantic.main import BaseModel +from robusta.core.model.env_vars import NODE_IP_CACHE_TTL_SEC from robusta.core.model.events import ExecutionBaseEvent from robusta.core.playbooks.base_trigger import BaseTrigger, TriggerEvent from robusta.core.reporting.base import Finding @@ -130,15 +132,24 @@ class PrometheusAlertTriggers(BaseModel): class AlertEventBuilder: + _node_name_by_ip: Dict[str, str] = {} + _node_ip_cache_time: float = 0 + @classmethod - def __find_node_by_ip(cls, ip) -> Optional[Node]: + def __refresh_node_ip_cache(cls): nodes: NodeList = NodeList.listNode().obj - for node in nodes.items: - addresses = [a.address for a in node.status.addresses] - logging.info(f"node {node.metadata.name} has addresses {addresses}") - if ip in addresses: - return node - return None + cls._node_name_by_ip = { + address.address: node.metadata.name for node in nodes.items for address in node.status.addresses + } + cls._node_ip_cache_time = time.time() + + @classmethod + def __find_node_by_ip(cls, ip) -> Optional[Node]: + cache_expired = time.time() - cls._node_ip_cache_time > NODE_IP_CACHE_TTL_SEC + if cache_expired or ip not in cls._node_name_by_ip: + cls.__refresh_node_ip_cache() + node_name = cls._node_name_by_ip.get(ip) + return Node().read(name=node_name) if node_name else None @classmethod def __load_node(cls, alert: PrometheusAlert, node_name: str) -> Optional[Node]: