Spaces:
Runtime error
Runtime error
| 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 | |
| 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}" | |
| 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)) | |
| 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") | |