"""
Unified network performance metrics collector.
"""
import asyncio
import os
import time
import queue
import threading
from typing import Any, Dict
from common.npu_device_info.collector_for_npu_device_info import NPUDeviceInfoCollector
from logger.setup_logger import setup_logger
from network_performance_exporter.collectors.collector_for_rdma import RDMACollector
from network_performance_exporter.collectors.collector_for_npu_roce import NPURoCECollector
from network_performance_exporter.collectors.collector_for_disk import DiskCollector
from network_performance_exporter.config import Config
from network_performance_exporter.exporters.prometheus_exporter import PrometheusExporter
from network_performance_exporter.exporters.nats_exporter import NatsExporter
from utils.utils import detect_devices
module_name = os.getenv("MODULE_NAME", "default")
logger = setup_logger()
class NetworkPerformanceCollector:
"""Unified collector for network performance metrics (coordinator)."""
def __init__(self, config: Config, prometheus_exporter: PrometheusExporter,
nats_exporter: NatsExporter, update_interval: int = 15):
self.config = config
self.prometheus_exporter = prometheus_exporter
self.nats_exporter = nats_exporter
self.update_interval = update_interval
self._started_collectors = []
self._start_barrier = None
def bootstrap(self) -> None:
"""初始化并启动所有必要的 Collector(根据可用设备)。"""
devices = detect_devices()
has_npu = devices.get("npu", False)
logger.info("Device detection result: has_npu=%s", has_npu)
if has_npu:
self._bootstrap_npu_collectors()
else:
logger.info("NPU not detected, skipping NPU RoCE Collector")
self._create_rdma_and_disk_collectors()
self._setup_barrier()
def _bootstrap_npu_collectors(self) -> None:
"""Bootstrap NPU-related collectors."""
logger.info("NPU detected, creating NPU RoCE Collector...")
npu_roce_collector = NPURoCECollector(self.config.node_name, self.update_interval)
self._started_collectors.append(("npu_roce", npu_roce_collector))
npu_device_collector = NPUDeviceInfoCollector()
try:
npu_device_collector.collect_info()
except Exception as e:
logger.debug("Initial NPU device collection failed (thread will retry): %s", str(e))
npu_device_collector.start()
def _create_rdma_and_disk_collectors(self) -> None:
"""Create RDMA and Disk collectors."""
rdma_collector = RDMACollector(self.config.node_name, self.update_interval)
disk_collector = DiskCollector(self.config.node_name, self.update_interval)
self._started_collectors.append(("rdma", rdma_collector))
self._started_collectors.append(("disk", disk_collector))
logger.info("All collectors created: %d collectors", len(self._started_collectors))
def _setup_barrier(self) -> None:
barrier = threading.Barrier(len(self._started_collectors) + 1)
self._start_barrier = barrier
for _, collector in self._started_collectors:
collector.set_start_barrier(barrier)
logger.info("Barrier created with %d participants", len(self._started_collectors) + 1)
def start(self):
"""Start all sub-collectors simultaneously."""
for name, collector in self._started_collectors:
collector.start()
logger.info("Started collector: %s", name)
def stop(self):
"""Stop all sub-collectors."""
for name, collector in self._started_collectors:
collector.stop()
logger.info("Stopped collector: %s", name)
async def collect_device_data(self, collector_name: str, collector) -> Dict[str, Any]:
"""
Collect data asynchronously from a specific collector.
"""
try:
result = await asyncio.to_thread(collector.result_queue.get,
timeout=self.update_interval)
logger.debug("Collected data for %s: %s", collector_name, result)
return result
except queue.Empty:
logger.warning("Timeout occurred while collecting data for %s", collector_name)
return {}
except Exception as e:
logger.error("Collector %s encountered an error: %s", collector_name, str(e))
return {}
async def collect_one_round(self) -> Dict[str, Any]:
"""Collect data from all collectors in parallel."""
tasks = [self.collect_device_data(name, collector)
for name, collector in self._started_collectors]
results = await asyncio.gather(*tasks)
return {name: result for (name, _), result in zip(self._started_collectors, results)}
async def run_forever(self, stop_event: asyncio.Event = None):
"""
Continuous collection and publishing loop.
Synchronizes with sub-collectors at startup using barrier.
Args:
stop_event: Optional asyncio.Event to signal when to stop
"""
loop = asyncio.get_running_loop()
await loop.run_in_executor(None, self._start_barrier.wait)
logger.info("Main program synchronized with collectors, starting collection loop")
while True:
if stop_event and stop_event.is_set():
logger.info("Stop signal received, exiting collection loop")
break
try:
start_time = time.time()
metrics_data = await self.collect_one_round()
try:
self.prometheus_exporter.update(metrics_data)
except Exception as e:
logger.error("Failed to update Prometheus exporter: %s", str(e))
if self.nats_exporter.is_connected():
try:
await self.nats_exporter.publish(metrics_data)
except Exception as e:
logger.error("Failed to publish to NATS: %s", str(e))
elapsed = time.time() - start_time
remaining = self.update_interval - elapsed
if remaining > 0:
await asyncio.sleep(remaining)
except Exception as e:
logger.error("Error in collection loop: %s", str(e))
await asyncio.sleep(self.update_interval)