diff --git a/snuba/web/rpc/storage_routing/load_retriever.py b/snuba/web/rpc/storage_routing/load_retriever.py index 99c64aceae..968d5aa5fa 100644 --- a/snuba/web/rpc/storage_routing/load_retriever.py +++ b/snuba/web/rpc/storage_routing/load_retriever.py @@ -1,6 +1,9 @@ +from __future__ import annotations + import inspect import json from collections.abc import Callable +from dataclasses import dataclass, fields from functools import wraps from typing import Any @@ -18,25 +21,26 @@ ) +@dataclass(kw_only=True, slots=True, frozen=True) class LoadInfo: - cluster_load: float - concurrent_queries: int - - def __init__(self, cluster_load: float, concurrent_queries: int) -> None: - self.cluster_load = cluster_load - self.concurrent_queries = concurrent_queries + cluster_load: float = -1.0 + concurrent_queries: float = -1.0 + cgroup_user_time_normalized: float = -1.0 + disk_inflight_ops: float = -1.0 + memory_allocated: float = -1.0 + part_mutation: float = -1.0 - def to_dict(self) -> dict[str, float | int]: - return { - "cluster_load": self.cluster_load, - "concurrent_queries": self.concurrent_queries, - } + def to_dict(self) -> dict[str, float]: + return {f.name: getattr(self, f.name) for f in fields(self)} @classmethod - def from_dict(cls, load_info_dict: dict[str, float | int]) -> "LoadInfo": + def from_dict(cls, load_info_dict: dict[str, float | int | None]) -> LoadInfo: return cls( - cluster_load=load_info_dict["cluster_load"], - concurrent_queries=int(load_info_dict["concurrent_queries"]), + **{ + f.name: float(v) + for f in fields(cls) + if (v := load_info_dict.get(f.name)) is not None + } ) @@ -82,76 +86,92 @@ def wrapper(*args: Any, **kwargs: Any) -> LoadInfo: def get_cluster_loadinfo( storage_set_key: StorageSetKey = StorageSetKey.EVENTS_ANALYTICS_PLATFORM, ) -> LoadInfo: + cluster_name = None try: cluster = get_cluster(storage_set_key) cluster_name = str(cluster.get_clickhouse_cluster_name()) if cluster.is_single_node(): - cluster_load_query = """ - SELECT - toFloat32(value)/ (SELECT - max(toInt32(replaceAll(metric, 'OSNiceTimeCPU', ''))) + 1 as num_cpus - FROM system.asynchronous_metrics - WHERE metric LIKE '%OSNiceTimeCPU%') * 100 as normalized_load - FROM system.asynchronous_metrics - WHERE metric = 'LoadAverage1' - """ - concurrent_queries_query = """ - SELECT - count() - FROM system.processes + metrics_from = """ + SELECT hostName() AS host, metric, toFloat64(value) AS value FROM system.asynchronous_metrics + UNION ALL + SELECT hostName() AS host, metric, toFloat64(value) AS value FROM system.metrics """ else: - cluster_load_query = f""" - SELECT - max(load_average.value / cpu_counts.num_cpus * 100) AS max_normalized_load - FROM ( - SELECT - hostName() AS host, - value, - metric - FROM clusterAllReplicas('{cluster.get_clickhouse_cluster_name()}', 'system', asynchronous_metrics) - WHERE metric = 'LoadAverage1' - ) AS load_average - JOIN ( + metrics_from = f""" + SELECT hostName() AS host, metric, toFloat64(value) AS value + FROM clusterAllReplicas('{cluster_name}', 'system', asynchronous_metrics) + UNION ALL + SELECT hostName() AS host, metric, toFloat64(value) AS value + FROM clusterAllReplicas('{cluster_name}', 'system', metrics) + """ + + # maxIf() returns 0 when the condition never matches; 0 is a real idle value + # (Query/PartMutation). max(if(..., NULL)) stays NULL. + # https://clickhouse.com/docs/sql-reference/aggregate-functions/combinators#-if + cluster_load_query = f""" SELECT - hostName() AS host, - max(toInt32(replaceAll(metric, 'OSNiceTimeCPU', ''))) + 1 AS num_cpus - FROM clusterAllReplicas('{cluster.get_clickhouse_cluster_name()}', 'system', asynchronous_metrics) - WHERE metric LIKE 'OSNiceTimeCPU%' - GROUP BY host - ) AS cpu_counts - ON load_average.host = cpu_counts.host - """ - concurrent_queries_query = f""" - SELECT sum(count) AS concurrent_queries + max(cluster_load) AS cluster_load, + max(concurrent_queries) AS concurrent_queries, + max(cgroup_user_time_normalized) AS cgroup_user_time_normalized, + max(disk_inflight_ops) AS disk_inflight_ops, + max(memory_allocated) AS memory_allocated, + max(part_mutation) AS part_mutation + FROM ( + SELECT + ifNull( + max(if(metric = 'LoadAverage1', value, NULL)) + / (max(if( + startsWith(metric, 'OSNiceTimeCPU'), + toInt32OrZero(replaceAll(metric, 'OSNiceTimeCPU', '')), + NULL + )) + 1) + * 100, + -1 + ) AS cluster_load, + max(if(metric = 'Query', value, NULL)) AS concurrent_queries, + max(if(metric = 'CGroupUserTimeNormalized', value, NULL)) AS cgroup_user_time_normalized, + -- Actual metric is BlockInFlightOps_ + max(if(startsWith(metric, 'BlockInFlightOps'), value, NULL)) AS disk_inflight_ops, + max(if(metric = 'MemoryTracking', value, NULL)) AS memory_allocated, + max(if(metric = 'PartMutation', value, NULL)) AS part_mutation FROM ( - SELECT count() AS count - FROM clusterAllReplicas('{cluster.get_clickhouse_cluster_name()}', 'system', 'processes') + {metrics_from} ) - """ + WHERE metric IN ( + 'LoadAverage1', 'CGroupUserTimeNormalized', + 'Query', 'MemoryTracking', 'PartMutation' + ) + OR startsWith(metric, 'OSNiceTimeCPU') + OR startsWith(metric, 'BlockInFlightOps') + GROUP BY host + ) + """ - cluster_load = float( + row = ( cluster.get_query_connection(ClickhouseClientSettings.INTERNAL) .execute(cluster_load_query) - .results[0][0] + .results[0] ) - concurrent_queries = int( - cluster.get_query_connection(ClickhouseClientSettings.INTERNAL) - .execute(concurrent_queries_query) - .results[0][0] + load_info = LoadInfo.from_dict( + { + "cluster_load": row[0], + "concurrent_queries": row[1], + "cgroup_user_time_normalized": row[2], + "disk_inflight_ops": row[3], + "memory_allocated": row[4], + "part_mutation": row[5], + } ) - load_info = LoadInfo(cluster_load=cluster_load, concurrent_queries=concurrent_queries) - metrics.gauge("cluster_load", load_info.cluster_load, tags={"cluster_name": cluster_name}) - metrics.gauge( - "concurrent_queries", - load_info.concurrent_queries, - tags={"cluster_name": cluster_name}, - ) + tags = {"cluster_name": cluster_name} + for name, value in load_info.to_dict().items(): + metrics.gauge(name, value, tags=tags) return load_info except Exception as e: - metrics.increment("get_cluster_loadinfo_failure", tags={"cluster_name": cluster_name}) + metrics.increment( + "get_cluster_loadinfo_failure", tags={"cluster_name": cluster_name or "unknown"} + ) sentry_sdk.capture_exception(e) - return LoadInfo(cluster_load=-1.0, concurrent_queries=-1) + return LoadInfo() diff --git a/tests/web/rpc/v1/routing_strategies/test_cluster_loadinfo.py b/tests/web/rpc/v1/routing_strategies/test_cluster_loadinfo.py index 90f850a9b4..78027ae9bb 100644 --- a/tests/web/rpc/v1/routing_strategies/test_cluster_loadinfo.py +++ b/tests/web/rpc/v1/routing_strategies/test_cluster_loadinfo.py @@ -2,16 +2,54 @@ import pytest -from snuba.web.rpc.storage_routing.load_retriever import get_cluster_loadinfo +from snuba.web.rpc.storage_routing.load_retriever import LoadInfo, get_cluster_loadinfo + +# Always present on CH. CGroupUserTimeNormalized / disk_inflight_ops are +# host-dependent and may be -1 without failing the probe. +_REQUIRED_FIELDS = ( + "cluster_load", + "concurrent_queries", + "memory_allocated", + "part_mutation", +) +_ALL_FIELDS = _REQUIRED_FIELDS + ( + "cgroup_user_time_normalized", + "disk_inflight_ops", +) + + +def _assert_probe_ok(load_info: LoadInfo) -> None: + assert load_info is not None + for field in _REQUIRED_FIELDS: + assert getattr(load_info, field) != -1, field + + +def _assert_probe_failed(load_info: LoadInfo) -> None: + assert load_info is not None + for field in _ALL_FIELDS: + assert getattr(load_info, field) == -1, field + + +def test_from_dict_old_cache_shape() -> None: + load_info = LoadInfo.from_dict({"cluster_load": 1.5, "concurrent_queries": 3}) + assert load_info.cluster_load == 1.5 + assert load_info.concurrent_queries == 3.0 + assert load_info.cgroup_user_time_normalized == -1 + assert load_info.disk_inflight_ops == -1 + assert load_info.memory_allocated == -1 + assert load_info.part_mutation == -1 + + +def test_from_dict_ignores_unknown_keys() -> None: + load_info = LoadInfo.from_dict({"cluster_load": 2.0, "not_a_field": 99, "also_unknown": None}) + assert load_info.cluster_load == 2.0 + assert load_info.concurrent_queries == -1 @pytest.mark.redis_db @pytest.mark.clickhouse_db def test_get_cluster_load() -> None: - load_info = get_cluster_loadinfo() - assert load_info is not None - assert load_info.cluster_load != -1.0 - assert load_info.concurrent_queries != -1 + _assert_probe_ok(get_cluster_loadinfo()) @pytest.mark.redis_db @@ -33,10 +71,7 @@ def test_get_cluster_loadinfo_if_cache_fails() -> None: mock_redis.side_effect = Exception("Test error") with patch("snuba.redis.get_redis_client") as mock_redis_client: mock_redis_client.return_value = mock_redis - load_info = get_cluster_loadinfo() - assert load_info is not None - assert load_info.cluster_load != -1.0 - assert load_info.concurrent_queries != -1 + _assert_probe_ok(get_cluster_loadinfo()) @pytest.mark.redis_db @@ -44,7 +79,4 @@ def test_get_cluster_loadinfo_if_cache_fails() -> None: def test_get_cluster_load_error_handling() -> None: with patch("snuba.clickhouse.connect.ClickhouseConnectPool.execute") as mock_execute: mock_execute.side_effect = Exception("Test error") - load_info = get_cluster_loadinfo() - assert load_info is not None - assert load_info.cluster_load == -1.0 - assert load_info.concurrent_queries == -1 + _assert_probe_failed(get_cluster_loadinfo())