"""PolicyEngine — applies static rules + learned heuristics + live overrides to scheduler decisions. Constructor accepts an optional ``db`` (anything with an ``events`` dict-like attribute). If ``db`` is None, the heuristics layer falls back to the in-process reputation store (``_reputation``) which is fed by :meth:`update_reputation`. Learning loop (F-04 remediation): :meth:`update_reputation` records per-node success/failure outcomes and feeds them back into :meth:`_heuristics` so the scheduler can prefer nodes that historically succeed and de-prioritise nodes that overheat or fail repeatedly. The reputation store is an exponentially-decaying moving average — old observations count for less, so a node that recovered from a bad spell isn't penalised forever. """ from __future__ import annotations import logging import threading import time from typing import Any, Dict, Optional log = logging.getLogger(__name__) # Decay factor for the moving average. New observation contributes # ``_DECAY`` of the new value, old history contributes ``1 - _DECAY``. _DECAY: float = 0.3 class PolicyEngine: """Static rules + learned heuristics + live overrides.""" def __init__(self, db: Optional[Any] = None) -> None: self.db = db self.rules: list = [] self.overrides: Dict[str, Any] = {} # Per-node reputation store (populated by update_reputation). # shape: {node_name: {"success": float, "failure": float, # "thermal_spike": float, "last_seen": float, # "samples": int}} self._reputation: Dict[str, Dict[str, float]] = {} self._lock = threading.Lock() # ------------------------------------------------------------------ # Rule registration # ------------------------------------------------------------------ def add_rule(self, rule: Any) -> None: self.rules.append(rule) def set_override(self, key: str, value: Any) -> None: self.overrides[key] = value def clear_override(self, key: str) -> None: self.overrides.pop(key, None) # ------------------------------------------------------------------ # Learning loop (F-04 remediation) # ------------------------------------------------------------------ def update_reputation( self, node: str, success: bool, duration: float = 0.0, thermal_spike: bool = False, ) -> Dict[str, float]: """Feed an execution outcome back into the reputation store. Parameters ---------- node Node identifier (typically ``node["name"]``). success ``True`` if the action completed with rc=0. duration Wall-clock seconds the action took. Used to track slow nodes via the rolling average. thermal_spike ``True`` if the node's temperature rose > 0.15°C during the action (overheating risk). Returns the updated reputation dict for ``node`` so callers can log / emit it on the bus. """ if not node: return {} with self._lock: cur = self._reputation.setdefault(node, { "success": 0.5, "failure": 0.0, "thermal_spike": 0.0, "avg_duration": 0.0, "samples": 0, "last_seen": 0.0, }) # Exponentially-weighted moving average cur["success"] = (_DECAY * (1.0 if success else 0.0) + (1 - _DECAY) * cur["success"]) cur["failure"] = (_DECAY * (0.0 if success else 1.0) + (1 - _DECAY) * cur["failure"]) cur["thermal_spike"] = (_DECAY * (1.0 if thermal_spike else 0.0) + (1 - _DECAY) * cur["thermal_spike"]) if duration > 0: cur["avg_duration"] = (_DECAY * duration + (1 - _DECAY) * cur["avg_duration"]) cur["samples"] = cur["samples"] + 1 cur["last_seen"] = time.time() return dict(cur) def get_reputation(self, node: str) -> Dict[str, float]: """Return the current reputation snapshot for ``node`` (or empty).""" with self._lock: return dict(self._reputation.get(node, {})) def all_reputations(self) -> Dict[str, Dict[str, float]]: """Return a deep-ish snapshot of all known node reputations.""" with self._lock: return {n: dict(v) for n, v in self._reputation.items()} # ------------------------------------------------------------------ # Apply policy to node selection # ------------------------------------------------------------------ def evaluate(self, action: Dict[str, Any], target: str, node: Dict[str, Any]) -> float: score_modifier: float = 0.0 for rule in self.rules: try: score_modifier += rule.apply(action, target, node) except Exception: log.debug("policy rule raised", exc_info=True) score_modifier += self._heuristics(action, target, node) score_modifier += self._overrides(action, target, node) return score_modifier # ------------------------------------------------------------------ # Heuristics (learning layer) # ------------------------------------------------------------------ def _heuristics(self, action: Dict[str, Any], target: str, node: Dict[str, Any]) -> float: score: float = 0.0 name = node.get("name") if isinstance(node, dict) else None if not name: return 0.0 # Layer 1: reputation store (always available — even with db=None) with self._lock: rep = self._reputation.get(name) if rep: # Reward historical success, penalise failure + thermal spikes score += rep["success"] * 20.0 score -= rep["failure"] * 50.0 score -= rep["thermal_spike"] * 15.0 # Penalise slow nodes (only after we have a few samples) if rep["samples"] >= 3 and rep["avg_duration"] > 0: # Normalise: 60s baseline, 1 point per 60s over baseline score -= max(0.0, (rep["avg_duration"] - 60.0) / 60.0) * 5.0 # Layer 2: db-backed history (preserved for backward compat) if self.db is not None and hasattr(self.db, "events"): try: history = self.db.events except Exception: history = None if history is not None: success_count = 0 if isinstance(history, dict): iterable = (e for sess in history.values() for e in (sess or [])) else: iterable = history or [] for event in iterable: if not isinstance(event, dict): continue if event.get("type") in ("action_end", "task_update"): d = event.get("data", event) if (d.get("action") == action.get("name") and d.get("state") == "done" and d.get("node") == name): success_count += 1 score += success_count * 2 return score # ------------------------------------------------------------------ # Live overrides # ------------------------------------------------------------------ def _overrides(self, action: Dict[str, Any], target: str, node: Dict[str, Any]) -> float: score: float = 0.0 forced = self.overrides.get("force_node") if forced and node.get("name") == forced: score += 1000 avoid = self.overrides.get("avoid_node") if avoid and node.get("name") == avoid: score -= 1000 return score