ForcePilot/backend/package/yuxi/channel/monitoring/diagnostic_heartbeat.py

474 lines
16 KiB
Python
Raw Normal View History

import asyncio
import gc
import logging
import time
from dataclasses import dataclass, field
logger = logging.getLogger(__name__)
DEFAULT_HEARTBEAT_INTERVAL = 30
STUCK_SESSION_TIMEOUT = 600
MAX_CONSECUTIVE_FAILURES = 5
MB = 1024 * 1024
DEFAULT_RSS_WARNING_MB = 1536
DEFAULT_RSS_CRITICAL_MB = 3072
DEFAULT_DATA_WARNING_MB = 1024
DEFAULT_DATA_CRITICAL_MB = 2048
DEFAULT_GC_OBJECTS_WARNING = 500000
DEFAULT_GC_SAMPLE_INTERVAL = 5
DEFAULT_RSS_GROWTH_WARNING_MB = 512
DEFAULT_RSS_GROWTH_CRITICAL_MB = 1024
DEFAULT_GROWTH_WINDOW_S = 600
DEFAULT_PRESSURE_COOLDOWN_S = 300
@dataclass
class MemorySnapshot:
rss_mb: float = 0.0
vms_mb: float = 0.0
data_mb: float = 0.0
gc_objects: int = 0
gc_g0_count: int = 0
gc_g1_count: int = 0
gc_g2_count: int = 0
timestamp: float = 0.0
@dataclass
class MemoryPressureEvent:
level: str # "warning" | "critical"
reason: str # "rss_threshold" | "data_threshold" | "rss_growth"
rss_mb: float = 0.0
threshold_mb: float = 0.0
rss_growth_mb: float = 0.0
window_s: float = 0.0
timestamp: float = 0.0
@dataclass
class HeartbeatReport:
active_sessions: int = 0
stuck_sessions: int = 0
stuck_details: list[dict] = field(default_factory=list)
tool_loop_sessions: int = 0
tool_loop_details: list[dict] = field(default_factory=list)
memory: MemorySnapshot = field(default_factory=MemorySnapshot)
active_tasks: int = 0
uptime_seconds: float = 0.0
process_start_timestamp: float = 0.0
timestamp: float = 0.0
warnings: list[str] = field(default_factory=list)
memory_pressure: MemoryPressureEvent | None = None
def to_dict(self) -> dict:
result = {
"active_sessions": self.active_sessions,
"stuck_sessions": self.stuck_sessions,
"stuck_details": self.stuck_details,
"tool_loop_sessions": self.tool_loop_sessions,
"tool_loop_details": self.tool_loop_details,
"memory": {
"rss_mb": self.memory.rss_mb,
"vms_mb": self.memory.vms_mb,
"data_mb": self.memory.data_mb,
"gc_objects": self.memory.gc_objects,
"gc_g0_count": self.memory.gc_g0_count,
"gc_g1_count": self.memory.gc_g1_count,
"gc_g2_count": self.memory.gc_g2_count,
"timestamp": self.memory.timestamp,
},
"active_tasks": self.active_tasks,
"uptime_seconds": self.uptime_seconds,
"process_start_timestamp": self.process_start_timestamp,
"timestamp": self.timestamp,
"warnings": self.warnings,
}
if self.memory_pressure:
result["memory_pressure"] = {
"level": self.memory_pressure.level,
"reason": self.memory_pressure.reason,
"rss_mb": self.memory_pressure.rss_mb,
"threshold_mb": self.memory_pressure.threshold_mb,
"rss_growth_mb": self.memory_pressure.rss_growth_mb,
"window_s": self.memory_pressure.window_s,
"timestamp": self.memory_pressure.timestamp,
}
return result
class DiagnosticHeartbeat:
"""诊断心跳30s 间隔检测卡住 session、内存快照 + 内存压力。
集成到 FastAPI lifespan 后启动可被 health endpoint 消费
"""
def __init__(
self,
interval: int = DEFAULT_HEARTBEAT_INTERVAL,
stuck_timeout: int = STUCK_SESSION_TIMEOUT,
rss_warning_mb: float = DEFAULT_RSS_WARNING_MB,
rss_critical_mb: float = DEFAULT_RSS_CRITICAL_MB,
data_warning_mb: float = DEFAULT_DATA_WARNING_MB,
data_critical_mb: float = DEFAULT_DATA_CRITICAL_MB,
rss_growth_warning_mb: float = DEFAULT_RSS_GROWTH_WARNING_MB,
rss_growth_critical_mb: float = DEFAULT_RSS_GROWTH_CRITICAL_MB,
growth_window_s: float = DEFAULT_GROWTH_WINDOW_S,
pressure_cooldown_s: float = DEFAULT_PRESSURE_COOLDOWN_S,
gc_objects_warning: int = DEFAULT_GC_OBJECTS_WARNING,
gc_sample_interval: int = DEFAULT_GC_SAMPLE_INTERVAL,
):
self._interval = interval
self._stuck_timeout = stuck_timeout
self._rss_warning_mb = rss_warning_mb
self._rss_critical_mb = rss_critical_mb
self._data_warning_mb = data_warning_mb
self._data_critical_mb = data_critical_mb
self._rss_growth_warning_mb = rss_growth_warning_mb
self._rss_growth_critical_mb = rss_growth_critical_mb
self._growth_window_s = growth_window_s
self._pressure_cooldown_s = pressure_cooldown_s
self._gc_objects_warning = gc_objects_warning
self._gc_sample_interval = gc_sample_interval
self._running = False
self._task: asyncio.Task | None = None
self._last_report: HeartbeatReport | None = None
self._last_memory: MemorySnapshot | None = None
self._last_pressure_at: dict[str, float] = {}
self._start_time: float | None = None
self._start_monotonic: float | None = None
self._consecutive_failures = 0
self._cycle_count = 0
@property
def last_report(self) -> HeartbeatReport | None:
return self._last_report
@staticmethod
def _resolve_start_time() -> float:
try:
import psutil
return psutil.Process().create_time()
except Exception:
return time.time()
async def start(self):
self._running = True
self._start_time = self._resolve_start_time()
self._start_monotonic = time.monotonic()
self._task = asyncio.create_task(self._loop())
logger.info(
"DiagnosticHeartbeat started (interval=%ds, stuck_timeout=%ds)",
self._interval,
self._stuck_timeout,
)
async def stop(self):
self._running = False
if self._task:
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
self._task = None
logger.info("DiagnosticHeartbeat stopped")
async def _loop(self):
while self._running:
await asyncio.sleep(self._interval)
try:
self._last_report = await self._collect()
self._consecutive_failures = 0
await self._push_metrics(self._last_report)
if self._last_report.warnings:
for w in self._last_report.warnings:
logger.warning("DiagnosticHeartbeat: %s", w)
if self._last_report.memory_pressure:
mp = self._last_report.memory_pressure
logger.warning(
"DiagnosticHeartbeat: 内存压力 %s (reason=%s, rss=%.0fMB, threshold=%.0fMB)",
mp.level,
mp.reason,
mp.rss_mb,
mp.threshold_mb,
)
except Exception:
self._consecutive_failures += 1
logger.exception(
"DiagnosticHeartbeat: collect failed (consecutive=%d/%d)",
self._consecutive_failures,
MAX_CONSECUTIVE_FAILURES,
)
if self._consecutive_failures >= MAX_CONSECUTIVE_FAILURES:
logger.critical(
"DiagnosticHeartbeat: %d consecutive collect failures, heartbeat may be broken",
self._consecutive_failures,
)
async def _collect(self) -> HeartbeatReport:
now = time.monotonic()
self._cycle_count += 1
memory = await self._capture_memory(now)
self._last_memory = memory
snapshots = _get_all_snapshots()
stuck_details, stuck_count = await self._detect_stuck_sessions(now, snapshots)
active_tasks = len(asyncio.all_tasks())
active_sessions = _count_active_sessions_from(snapshots)
uptime_seconds = now - self._start_monotonic if self._start_monotonic else 0.0
process_start_timestamp = self._start_time or 0.0
memory_pressure = self._detect_memory_pressure(memory, now)
warnings: list[str] = []
if stuck_count > 0:
warnings.append(f"检测到 {stuck_count} 个卡住 session")
if memory.rss_mb > self._rss_warning_mb:
warnings.append(f"内存 RSS {memory.rss_mb:.0f}MB > {self._rss_warning_mb:.0f}MB")
if memory.data_mb > self._data_warning_mb:
warnings.append(f"内存 DATA {memory.data_mb:.0f}MB > {self._data_warning_mb:.0f}MB")
if memory.gc_objects > self._gc_objects_warning:
warnings.append(f"GC 对象数 {memory.gc_objects} > {self._gc_objects_warning}")
if memory_pressure:
level_label = "严重" if memory_pressure.level == "critical" else ""
warnings.append(f"内存压力{level_label}: {memory_pressure.reason} (RSS={memory_pressure.rss_mb:.0f}MB)")
return HeartbeatReport(
active_sessions=active_sessions,
stuck_sessions=stuck_count,
stuck_details=stuck_details,
tool_loop_sessions=0,
tool_loop_details=[],
memory=memory,
active_tasks=active_tasks,
uptime_seconds=uptime_seconds,
process_start_timestamp=process_start_timestamp,
timestamp=now,
warnings=warnings,
memory_pressure=memory_pressure,
)
async def _push_metrics(self, report: HeartbeatReport):
try:
from yuxi.channel.monitoring.metrics import (
active_sessions,
active_tasks,
memory_gc_objects,
memory_rss_mb,
stuck_sessions,
)
active_sessions.set(report.active_sessions)
stuck_sessions.set(report.stuck_sessions)
memory_rss_mb.set(report.memory.rss_mb)
memory_gc_objects.set(report.memory.gc_objects)
active_tasks.set(report.active_tasks)
except Exception:
pass
def _detect_memory_pressure(self, current: MemorySnapshot, now: float) -> MemoryPressureEvent | None:
pressure = _check_threshold_pressure(
current,
rss_warning_mb=self._rss_warning_mb,
rss_critical_mb=self._rss_critical_mb,
)
if pressure is None:
pressure = _check_data_threshold_pressure(
current,
data_warning_mb=self._data_warning_mb,
data_critical_mb=self._data_critical_mb,
)
if pressure is None:
pressure = _check_growth_pressure(
self._last_memory,
current,
now,
growth_window_s=self._growth_window_s,
rss_growth_warning_mb=self._rss_growth_warning_mb,
rss_growth_critical_mb=self._rss_growth_critical_mb,
)
if pressure is None:
return None
key = pressure.reason
last_at = self._last_pressure_at.get(key, 0)
if now - last_at < self._pressure_cooldown_s:
return None
self._last_pressure_at[key] = now
return pressure
async def _detect_stuck_sessions(self, now: float, snapshots: dict) -> tuple[list[dict], int]:
stuck: list[dict] = []
for key, snap in snapshots.items():
if snap.last_event_at and (now - snap.last_event_at) > self._stuck_timeout:
state_value = _get_state_value(snap.state)
if snap.connected and state_value in ("running", "retrying"):
stuck.append(
{
"key": key,
"state": state_value,
"last_event_at": snap.last_event_at,
"idle_seconds": round(now - snap.last_event_at, 1),
}
)
return stuck, len(stuck)
async def _capture_memory(self, now: float) -> MemorySnapshot:
snapshot = MemorySnapshot(timestamp=now)
try:
import psutil
proc = psutil.Process()
mem = proc.memory_info()
snapshot.rss_mb = round(mem.rss / 1024 / 1024, 2)
snapshot.vms_mb = round(mem.vms / 1024 / 1024, 2)
snapshot.data_mb = round(getattr(mem, "data", 0) / 1024 / 1024, 2)
except Exception:
pass
gc_counts = gc.get_count()
snapshot.gc_g0_count = gc_counts[0]
snapshot.gc_g1_count = gc_counts[1]
snapshot.gc_g2_count = gc_counts[2]
if self._cycle_count % self._gc_sample_interval == 0:
snapshot.gc_objects = len(gc.get_objects())
return snapshot
def get_memory_trend(self) -> dict:
if self._last_memory:
return {
"rss_mb": self._last_memory.rss_mb,
"vms_mb": self._last_memory.vms_mb,
"data_mb": self._last_memory.data_mb,
"gc_objects": self._last_memory.gc_objects,
"timestamp": self._last_memory.timestamp,
}
return {}
def reset(self):
self._last_memory = None
self._last_pressure_at.clear()
self._last_report = None
self._cycle_count = 0
def _check_threshold_pressure(
memory: MemorySnapshot,
*,
rss_warning_mb: float,
rss_critical_mb: float,
) -> MemoryPressureEvent | None:
if memory.rss_mb >= rss_critical_mb:
return MemoryPressureEvent(
level="critical",
reason="rss_threshold",
rss_mb=memory.rss_mb,
threshold_mb=rss_critical_mb,
timestamp=memory.timestamp,
)
if memory.rss_mb >= rss_warning_mb:
return MemoryPressureEvent(
level="warning",
reason="rss_threshold",
rss_mb=memory.rss_mb,
threshold_mb=rss_warning_mb,
timestamp=memory.timestamp,
)
return None
def _check_data_threshold_pressure(
memory: MemorySnapshot,
*,
data_warning_mb: float,
data_critical_mb: float,
) -> MemoryPressureEvent | None:
if memory.data_mb >= data_critical_mb:
return MemoryPressureEvent(
level="critical",
reason="data_threshold",
rss_mb=memory.rss_mb,
threshold_mb=data_critical_mb,
timestamp=memory.timestamp,
)
if memory.data_mb >= data_warning_mb:
return MemoryPressureEvent(
level="warning",
reason="data_threshold",
rss_mb=memory.rss_mb,
threshold_mb=data_warning_mb,
timestamp=memory.timestamp,
)
return None
def _check_growth_pressure(
previous: MemorySnapshot | None,
current: MemorySnapshot,
now: float,
*,
growth_window_s: float,
rss_growth_warning_mb: float,
rss_growth_critical_mb: float,
) -> MemoryPressureEvent | None:
if previous is None or previous.rss_mb <= 0:
return None
window_s = now - previous.timestamp
if window_s <= 0 or window_s > growth_window_s:
return None
rss_growth_mb = current.rss_mb - previous.rss_mb
if rss_growth_mb <= 0:
return None
if rss_growth_mb >= rss_growth_critical_mb:
return MemoryPressureEvent(
level="critical",
reason="rss_growth",
rss_mb=current.rss_mb,
threshold_mb=rss_growth_critical_mb,
rss_growth_mb=round(rss_growth_mb, 2),
window_s=round(window_s, 1),
timestamp=current.timestamp,
)
if rss_growth_mb >= rss_growth_warning_mb:
return MemoryPressureEvent(
level="warning",
reason="rss_growth",
rss_mb=current.rss_mb,
threshold_mb=rss_growth_warning_mb,
rss_growth_mb=round(rss_growth_mb, 2),
window_s=round(window_s, 1),
timestamp=current.timestamp,
)
return None
def _get_state_value(state) -> str:
return state.value if hasattr(state, "value") else str(state)
def _get_all_snapshots() -> dict:
try:
from yuxi.channel.runtime.manager import gateway # 延迟导入以避免循环依赖Python import 缓存保证无重复加载开销
return gateway.get_all_snapshots()
except Exception:
logger.warning("DiagnosticHeartbeat: get_all_snapshots failed", exc_info=True)
return {}
def _count_active_sessions_from(snapshots: dict) -> int:
return sum(1 for s in snapshots.values() if _get_state_value(s.state) == "running")
diagnostic_heartbeat = DiagnosticHeartbeat()