prediqai / RULE /aml_engine /action_engine.py
ganesh-vilje's picture
Deploy to Hugging Face Main
f8f02c0
Raw
History Blame Contribute Delete
6.09 kB
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")