airflow-repo/dags/anomaly_detection_dag.py
2026-07-01 14:40:11 +03:30

392 lines
No EOL
14 KiB
Python

from airflow import DAG
from airflow.decorators import task
from airflow.providers.amazon.aws.hooks.s3 import S3Hook
from airflow.providers.postgres.hooks.postgres import PostgresHook
from datetime import datetime, timedelta, timezone
import pandas as pd
import requests
import os
import json
import uuid
import logging
logger = logging.getLogger(__name__)
# --- Configuration & Constants ---
MINIO_BUCKET = "fraud-features"
AWS_CONN_ID = "minio_conn"
POSTGRES_CONN_ID = "postgres_conn"
BATCH_INTERVAL_HOURS = int(os.environ.get('BATCH_INTERVAL_HOURS', 6))
default_args = {
'owner': 'data_engineering',
'depends_on_past': False,
'retries': 1,
'retry_delay': timedelta(minutes=5),
}
with DAG(
dag_id='session_anomaly_detection_pipeline',
default_args=default_args,
description='Batch fraud detection pipeline using PostHog, MinIO, MLFlow, and Postgres',
schedule_interval=timedelta(hours=BATCH_INTERVAL_HOURS),
start_date=datetime(2026, 6, 20, tzinfo=timezone.utc),
catchup=False,
tags=['fraud', 'posthog', 'mlflow'],
) as dag:
@task
def extract_features_to_minio(**kwargs) -> dict:
logger.info("=== Starting feature extraction task ===")
try:
run_id = kwargs["run_id"]
logger.info("Run ID: %s", run_id)
# ------------------------------------------------------------------
# 1. Determine extraction window
# ------------------------------------------------------------------
logger.info("Connecting to PostgreSQL...")
pg_hook = PostgresHook(postgres_conn_id=POSTGRES_CONN_ID)
logger.info("Fetching previous successful batch...")
last_run = pg_hook.get_first(
"SELECT MAX(window_end) FROM batch_runs WHERE status = 'SUCCESS'"
)
last_window_end = last_run[0] if last_run and last_run[0] else None
logger.info("Last successful window_end: %s", last_window_end)
interval_hours = BATCH_INTERVAL_HOURS
logger.info("Batch interval: %s hours", interval_hours)
if last_window_end:
start_date = last_window_end
end_date = start_date + timedelta(hours=interval_hours)
else:
end_date = datetime.now(timezone.utc)
start_date = end_date - timedelta(hours=interval_hours)
start_str = start_date.strftime("%Y-%m-%d %H:%M:%S")
end_str = end_date.strftime("%Y-%m-%d %H:%M:%S")
logger.info("Extraction window: %s -> %s", start_str, end_str)
# ------------------------------------------------------------------
# 2. Build HogQL query
# ------------------------------------------------------------------
dag_dir = os.path.dirname(os.path.abspath(__file__))
sql_file_path = os.path.join(dag_dir, "sql", "session_features.hql")
logger.info("Reading HogQL template from %s", sql_file_path)
with open(sql_file_path, "r") as f:
template = f.read()
query = template.format(
start_date=start_str,
end_date=end_str,
)
logger.info("Generated HogQL query:")
logger.info("\n%s", query)
# ------------------------------------------------------------------
# 3. Query PostHog
# ------------------------------------------------------------------
POSTHOG_HOST = os.environ.get("POSTHOG_HOST")
POSTHOG_PROJECT_ID = os.environ.get("POSTHOG_PROJECT_ID", "1")
POSTHOG_API_KEY = os.environ.get("POSTHOG_API_KEY")
logger.info("POSTHOG_API_KEY: %s", POSTHOG_API_KEY)
logger.info("Sending request to PostHog...")
logger.info("Host: %s", POSTHOG_HOST)
logger.info("Project ID: %s", POSTHOG_PROJECT_ID)
response = requests.post(
f"{POSTHOG_HOST}/api/projects/{POSTHOG_PROJECT_ID}/query/",
headers={
"Authorization": f"Bearer {POSTHOG_API_KEY}",
"Content-Type": "application/json",
},
json={
"query": {
"kind": "HogQLQuery",
"query": query,
}
},
timeout=300,
)
logger.info("PostHog status code: %s", response.status_code)
if not response.ok:
logger.error("PostHog response:\n%s", response.text)
response.raise_for_status()
data = response.json()
logger.info("Successfully parsed JSON response.")
# ------------------------------------------------------------------
# 4. Convert to DataFrame
# ------------------------------------------------------------------
columns = data.get("columns", [])
results = data.get("results", [])
logger.info("Columns: %s", columns)
logger.info("Rows returned: %d", len(results))
df = pd.DataFrame(results, columns=columns)
logger.info(
"DataFrame shape: %s x %s",
df.shape[0],
df.shape[1],
)
# ------------------------------------------------------------------
# 5. Save parquet
# ------------------------------------------------------------------
local_path = f"/tmp/features_{run_id}.parquet"
logger.info("Writing parquet to %s", local_path)
df.to_parquet(local_path, index=False)
logger.info(
"Parquet size: %.2f MB",
os.path.getsize(local_path) / (1024 * 1024),
)
# ------------------------------------------------------------------
# 6. Upload to MinIO
# ------------------------------------------------------------------
logger.info("Connecting to MinIO...")
s3_hook = S3Hook(aws_conn_id=AWS_CONN_ID)
s3_key = f"batches/{run_id}/session_features.parquet"
logger.info(
"Uploading to bucket=%s key=%s",
MINIO_BUCKET,
s3_key,
)
s3_hook.load_file(
filename=local_path,
key=s3_key,
bucket_name=MINIO_BUCKET,
replace=True,
)
logger.info("Upload complete.")
os.remove(local_path)
logger.info("Temporary parquet deleted.")
logger.info("=== Feature extraction completed successfully ===")
return {
"s3_uri": f"s3://{MINIO_BUCKET}/{s3_key}",
"window_start": start_date.isoformat(),
"window_end": end_date.isoformat(),
"record_count": len(df),
}
except Exception:
logger.exception("Feature extraction task failed!")
raise
@task
def call_mlflow_inference(extraction_result: dict, **kwargs) -> str:
"""
Task 2: Reads features from MinIO, sends payload to local MLflow endpoint,
saves predictions (including SHAP values) back to MinIO.
"""
logger.info("=== Starting MLflow inference task ===")
run_id = kwargs['run_id']
features_s3_uri = extraction_result.get("s3_uri") if isinstance(extraction_result, dict) else extraction_result
s3_hook = S3Hook(aws_conn_id=AWS_CONN_ID)
bucket = features_s3_uri.split("/")[2]
key = "/".join(features_s3_uri.split("/")[3:])
logger.info("Downloading features from %s", features_s3_uri)
local_features_path = s3_hook.download_file(key=key, bucket_name=bucket, local_path="/tmp")
df_features = pd.read_parquet(local_features_path)
if df_features.empty:
logger.info("DataFrame is empty. Skipping inference.")
os.remove(local_features_path)
return features_s3_uri
logger.info("Loaded %d rows for inference.", len(df_features))
payload = {"dataframe_split": df_features.to_dict(orient="split")}
mlflow_url = os.environ.get("MLFLOW_API_URL", "http://host.docker.internal:5001/invocations")
response = requests.post(
mlflow_url,
data=json.dumps(payload),
headers={"Content-Type": "application/json"},
timeout=120
)
if not response.ok:
logger.error("MLflow API error: %s - %s", response.status_code, response.text)
response.raise_for_status()
predictions_data = response.json().get("predictions", [])
# Extract model outputs into the dataframe
df_features['prediction'] = [p.get('prediction') for p in predictions_data]
df_features['anomaly_score'] = [p.get('anomaly_score') for p in predictions_data]
df_features['decision_score'] = [p.get('decision_score') for p in predictions_data]
df_features['is_anomaly'] = df_features['prediction'] == -1
# Extract SHAP values (returned as stringified JSON by your model)
df_features['shap_values'] = [p.get('shap_values', '{}') for p in predictions_data]
local_preds_path = f"/tmp/predictions_{run_id}.parquet"
df_features.to_parquet(local_preds_path, index=False)
preds_s3_key = f"batches/{run_id}/predictions.parquet"
s3_hook.load_file(
filename=local_preds_path,
key=preds_s3_key,
bucket_name=MINIO_BUCKET,
replace=True
)
os.remove(local_features_path)
os.remove(local_preds_path)
return f"s3://{MINIO_BUCKET}/{preds_s3_key}"
@task
def load_predictions_to_postgres(predictions_s3_uri: str, extraction_result: dict, **kwargs):
logger.info("=== Loading predictions into PostgreSQL ===")
run_id = kwargs["run_id"]
s3_hook = S3Hook(aws_conn_id=AWS_CONN_ID)
pg_hook = PostgresHook(postgres_conn_id=POSTGRES_CONN_ID)
bucket = predictions_s3_uri.split("/")[2]
key = "/".join(predictions_s3_uri.split("/")[3:])
local_path = s3_hook.download_file(
key=key,
bucket_name=bucket,
local_path="/tmp",
)
df = pd.read_parquet(local_path)
logger.info("Loaded %d prediction rows", len(df))
window_start = extraction_result["window_start"]
window_end = extraction_result["window_end"]
total_sessions = len(df)
total_anomalies = int(df["is_anomaly"].sum()) if total_sessions else 0
contamination_rate = total_anomalies / total_sessions if total_sessions else 0.0
mean_score = float(df["anomaly_score"].mean()) if total_sessions else None
max_score = float(df["anomaly_score"].max()) if total_sessions else None
# ---------------------------
# ISOLATE FEATURE COLUMNS
# ---------------------------
# Define which columns are NOT part of the JSONB feature payload
metadata_cols = {
'session_id', 'user_id', 'prediction',
'anomaly_score', 'decision_score', 'is_anomaly', 'shap_values'
}
# Everything else is a feature
feature_cols = [col for col in df.columns if col not in metadata_cols]
from psycopg2.extras import Json, execute_values
records = []
for _, row in df.iterrows():
# 1. Build a dict of features for this specific row (dropping nulls safely)
row_features = {
col: row[col]
for col in feature_cols
if pd.notna(row[col])
}
# 2. Parse the SHAP values back into a dict (since MLflow returned stringified JSON)
shap_raw = row.get("shap_values")
shap_dict = json.loads(shap_raw) if isinstance(shap_raw, str) else (shap_raw or {})
records.append((
None, # batch_id placeholder
row["session_id"],
row.get("user_id"),
window_start,
row["anomaly_score"],
bool(row["is_anomaly"]),
Json(row_features), # Automatically adapts dict to JSONB
Json(shap_dict) # Automatically adapts dict to JSONB
))
insert_batch_sql = """
INSERT INTO batch_runs (
dag_run_id, window_start, window_end, mlflow_model_version, status
) VALUES (%s,%s,%s,%s,'SUCCESS')
RETURNING batch_id;
"""
prediction_sql = """
INSERT INTO session_predictions (
batch_id, session_id, user_id, session_start_time,
anomaly_score, is_anomaly, session_features, shap_values
) VALUES %s;
"""
update_batch_sql = """
UPDATE batch_runs
SET total_sessions_processed=%s, total_anomalies_detected=%s,
contamination_rate=%s, mean_anomaly_score=%s, max_anomaly_score=%s
WHERE batch_id=%s;
"""
conn = pg_hook.get_conn()
try:
with conn:
with conn.cursor() as cur:
cur.execute(
insert_batch_sql,
(run_id, window_start, window_end, "v1.0.0"),
)
batch_id = cur.fetchone()[0]
# Replace the 'None' placeholder with the actual batch_id
records = [(batch_id, *r[1:]) for r in records]
execute_values(cur, prediction_sql, records)
cur.execute(
update_batch_sql,
(total_sessions, total_anomalies, contamination_rate,
mean_score, max_score, batch_id),
)
logger.info("Inserted %d predictions for batch %s", total_sessions, batch_id)
except Exception:
conn.rollback()
with conn:
with conn.cursor() as cur:
cur.execute("UPDATE batch_runs SET status='FAILED' WHERE dag_run_id=%s;", (run_id,))
raise
finally:
conn.close()
os.remove(local_path)
# --- Pipeline Orchestration ---
features_uri = extract_features_to_minio()
predictions_uri = call_mlflow_inference(features_uri)
load_predictions_to_postgres(predictions_uri, features_uri)