File size: 4,948 Bytes
d491dc1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
import asyncio
import random
import datetime
import pandas as pd
from dataclasses import dataclass
from typing import Optional, Any, Dict

from backend.core.logger import get_logger

logger = get_logger(__name__)

@dataclass 
class LoopState:
    iteration: int = 0
    last_psi: float = 0.0
    status: str = "GREEN"
    consecutive_amber: int = 0
    stop_requested: bool = False

@dataclass
class PSIComputeResult:
    score: float
    feature_drifts: Dict[str, Any]
    data_source: str

class PSIMonitorLoop:
    """
    Self-running PSI drift monitoring loop.
    Loop design:
        Trigger:  every 60 seconds
        Generator: compute PSI on latest 200 sales events vs. training reference
        Verifier:  PSI < 0.10 → GREEN, < 0.20 → AMBER, >= 0.20 → RED + retrain signal
        Stop rule: stop_requested flag or 3 consecutive RED readings
    """
    INTERVAL_SECONDS = 60
    MAX_CONSECUTIVE_RED = 3

    def __init__(self, db_session_factory, safeguards, event_bus=None):
        self.db_factory = db_session_factory
        self.safeguards = safeguards
        self.event_bus = event_bus
        self.state = LoopState()

    async def run(self):
        logger.info("[PSI Loop] Starting autonomous PSI monitoring loop...")
        while not self.state.stop_requested:
            try:
                # Generator: Compute PSI
                psi_result = await self._compute_psi()
                
                # Verifier: Determine status
                status = self._verify(psi_result)
                self.state.status = status
                self.state.last_psi = psi_result.score
                
                # Side effect & Stop rule evaluation
                if status == "RED":
                    self.state.consecutive_amber += 1
                    logger.warning(f"[PSI Loop] High drift detected (PSI={psi_result.score:.4f}, RED status #{self.state.consecutive_amber})")
                    if self.state.consecutive_amber >= self.MAX_CONSECUTIVE_RED:
                        logger.warn("[PSI Loop] Triggering retrain signal due to 3 consecutive RED readings.")
                        if self.event_bus:
                            await self.event_bus.publish("retrain_requested", {
                                "reason": "3 consecutive RED PSI readings",
                                "psi": psi_result.score,
                            })
                        self.state.consecutive_amber = 0
                else:
                    self.state.consecutive_amber = 0

                self.state.iteration += 1
            except Exception as e:
                logger.error(f"[PSI Loop] Error in monitor loop: {e}")

            await asyncio.sleep(self.INTERVAL_SECONDS)

    async def _compute_psi(self) -> PSIComputeResult:
        if not self.db_factory:
            return PSIComputeResult(score=0.0412, feature_drifts={}, data_source="synthetic")

        db = self.db_factory()
        try:
            from backend.db.models import SalesEvent
            sales_events = db.query(SalesEvent).order_by(SalesEvent.created_at.desc()).limit(200).all()
            
            if len(sales_events) < 30:
                from ml_core.demand_simulation import generate_training_data
                X, _, _, _, _ = generate_training_data(n_samples=100)
                prod_df = pd.DataFrame({
                    'weather_temp': X[:, 0],
                    'weather_rain': X[:, 1],
                    'time_elapsed_sec': X[:, 2]
                })
                data_source = "synthetic"
            else:
                prod_df = pd.DataFrame([{
                    'weather_temp': getattr(e, 'weather_temp', None),
                    'weather_rain': getattr(e, 'weather_rain', None),
                    'time_elapsed_sec': getattr(e, 'time_elapsed_sec', None)
                } for e in sales_events if getattr(e, 'weather_temp', None) is not None])
                if len(prod_df) < 30:
                    from ml_core.demand_simulation import generate_training_data
                    X, _, _, _, _ = generate_training_data(n_samples=100)
                    prod_df = pd.DataFrame({
                        'weather_temp': X[:, 0],
                        'weather_rain': X[:, 1],
                        'time_elapsed_sec': X[:, 2]
                    })
                    data_source = "synthetic"
                else:
                    data_source = "real"

            drifts = self.safeguards.calculate_drift_metrics(prod_df)
            max_score = max([v.get("psi", 0.0) for v in drifts.values()]) if drifts else 0.04
            return PSIComputeResult(score=max_score, feature_drifts=drifts, data_source=data_source)
        finally:
            db.close()

    def _verify(self, result: PSIComputeResult) -> str:
        if result.score < 0.10:
            return "GREEN"
        elif result.score < 0.20:
            return "AMBER"
        return "RED"