from __future__ import annotations """ aml_engine/action_engine.py - Alert generation layer. When a rule evaluation mask returns True for any rows in the DataFrame, the ActionEngine generates structured Alert objects for each triggered row. Alert generation is vectorized where possible: - Only rows where mask=True are processed - Field value extraction is done via .loc[] bulk indexing - No row-by-row Python loops for the full dataset """ import logging from datetime import datetime, timezone from typing import Dict, List, Optional import pandas as pd from RULE.aml_engine.alert_models import Alert from RULE.aml_engine.rule_models import AMLRule logger = logging.getLogger(__name__) # Default transaction ID column priority list TRANSACTION_ID_COLUMNS = [ "transaction_id", "message_id", "Transaction_Id", "hit_id", "hit.hit_id", "mt103.transaction_reference_number", "mt103.reference", ] class ActionEngine: """ Generates Alert objects from rule evaluation boolean masks. Usage: engine = ActionEngine() alerts = engine.generate_alerts(df, rule, mask) """ def __init__( self, id_column: Optional[str] = None, timestamp_utc: Optional[datetime] = None, ): self.id_column = id_column self._timestamp = timestamp_utc def generate_alerts( self, df: pd.DataFrame, rule: AMLRule, mask: pd.Series, ) -> List[Alert]: """ Generate Alert objects for all rows where mask is True. Args: df: Normalized DataFrame (after DatasetAdapter). rule: The AMLRule that was evaluated. mask: Boolean Series from RuleEngine.evaluate_rule(). Returns: List of Alert objects (one per triggered row). """ triggered_df = df.loc[mask] if triggered_df.empty: return [] id_col = self._detect_id_column(df) triggered_fields = self._extract_triggered_fields(rule) ts = self._get_timestamp() alerts: List[Alert] = [] for row_index, row in triggered_df.iterrows(): transaction_id = self._get_transaction_id(row, id_col, row_index) trigger_values = self._extract_trigger_values(row, triggered_fields) alert = Alert( rule_id=rule.rule_id, rule_name=rule.rule_name, transaction_id=transaction_id, decision=rule.action.decision, severity=rule.severity, auto_block=rule.action.auto_block, comment=rule.action.comment_template, sla_hours=rule.action.sla_hours, regulatory_action=rule.action.regulatory_action, signal_type=rule.action.signal_type, priority=rule.action.priority, hard_stop=rule.action.hard_stop, triggered_fields=triggered_fields, trigger_values=trigger_values, timestamp=ts, row_index=int(row_index), ) alerts.append(alert) logger.info("Rule '%s': generated %d alerts.", rule.rule_id, len(alerts)) return alerts def generate_all_alerts( self, df: pd.DataFrame, rules: List[AMLRule], masks: Dict[str, pd.Series], ) -> List[Alert]: """ Generate alerts for all rules at once. Args: df: Normalized DataFrame. rules: List of AMLRule objects. masks: Dict {rule_id: bool_mask} from RuleEngine.evaluate_all(). Returns: Combined flat list of all Alert objects. """ rule_map = {r.rule_id: r for r in rules} all_alerts: List[Alert] = [] for rule_id, mask in masks.items(): if rule_id not in rule_map: continue rule = rule_map[rule_id] alerts = self.generate_alerts(df, rule, mask) all_alerts.extend(alerts) return all_alerts # ── Private helpers ──────────────────────────────────────────────────────── def _detect_id_column(self, df: pd.DataFrame) -> Optional[str]: """Find the best transaction ID column in the DataFrame.""" if self.id_column and self.id_column in df.columns: return self.id_column for col in TRANSACTION_ID_COLUMNS: if col in df.columns: return col return None @staticmethod def _get_transaction_id(row: pd.Series, id_col: Optional[str], row_index) -> Optional[str]: """Extract transaction ID from the row.""" if id_col and id_col in row.index: val = row[id_col] if pd.notna(val): return str(val) return f"row_{row_index}" @staticmethod def _extract_triggered_fields(rule: AMLRule) -> List[str]: """Get list of field names involved in rule conditions.""" fields = [c.field for c in rule.conditions] if rule.secondary_conditions: fields.extend(c.field for c in rule.secondary_conditions) return list(dict.fromkeys(fields)) @staticmethod def _extract_trigger_values( row: pd.Series, fields: List[str] ) -> Dict[str, object]: """Extract field values from a row for the triggered fields.""" values = {} for field in fields: if field in row.index: val = row[field] if pd.notna(val): values[field] = val if not hasattr(val, 'item') else val.item() else: values[field] = None return values def _get_timestamp(self) -> str: """Return ISO UTC timestamp string.""" if self._timestamp: return self._timestamp.isoformat() + "Z" return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z")