""" Policy Rollback — auto-rollback if shadow performance degrades. Enables: - Detecting when a promoted policy performs worse than baseline - Automatically reverting to previous best - Preventing bad policies from reaching live """ from __future__ import annotations import time from dataclasses import dataclass, field from typing import Any, List, Mapping, Optional, Tuple from malkhut.state import FulfilmentPolicyParams from malkhut.training.registry import PolicyRegistry, PolicyStage, PolicyRecord @dataclass(frozen=True, slots=True) class RollbackEvent: """Record of a rollback event.""" ts_ns: int rolled_back_version: str rolled_back_to: str reason: str performance_drop: float class PolicyRollback: """ Auto-rollback if shadow performance degrades. Monitors active policy performance and reverts to previous best if performance drops below threshold. """ def __init__( self, registry: PolicyRegistry, degradation_threshold: float = -5.0, # bps min_shadow_steps: int = 10, ) -> None: self._registry = registry self._degradation_threshold = degradation_threshold self._min_shadow_steps = min_shadow_steps self._shadow_scores: List[float] = [] self._rollback_events: List[RollbackEvent] = [] def record_shadow_score(self, score: float) -> None: """Record a shadow performance score.""" self._shadow_scores.append(score) def check_rollback(self) -> Optional[RollbackEvent]: """ Check if rollback is needed. Returns RollbackEvent if rollback should happen, None otherwise. """ if len(self._shadow_scores) < self._min_shadow_steps: return None # Compare recent average to baseline recent = self._shadow_scores[-self._min_shadow_steps:] avg_recent = sum(recent) / len(recent) # Get baseline score (first policy in registry) baseline = self._registry.get_by_stage(PolicyStage.ACTIVE) if not baseline: return None baseline_record = baseline[0] if baseline else None if not baseline_record: return None # Check degradation if avg_recent < self._degradation_threshold: # Find previous best to rollback to previous = self._registry.get_by_stage(PolicyStage.RETIRED) if previous: rollback_to = previous[0].version else: rollback_to = "baseline" # Perform rollback current = baseline_record[0] if isinstance(baseline_record, list) else baseline_record self._registry.retire(current.version, "auto_rollback_degradation") event = RollbackEvent( ts_ns=time.time_ns(), rolled_back_version=current.version, rolled_back_to=rollback_to, reason=f"performance_drop_{avg_recent:.2f}", performance_drop=avg_recent - self._degradation_threshold, ) self._rollback_events.append(event) return event return None @property def rollback_events(self) -> List[RollbackEvent]: return list(self._rollback_events) @property def shadow_score_count(self) -> int: return len(self._shadow_scores)