434 lines
19 KiB
Python
434 lines
19 KiB
Python
"""Evaluate clustering quality by comparing against ground-truth labels.
|
||
|
||
Usage:
|
||
set DJANGO_ALLOW_ASYNC_UNSAFE=true
|
||
runtime\python\python.exe scripts\eval_clustering.py
|
||
"""
|
||
import os, sys, json, asyncio, traceback
|
||
os.environ['DJANGO_ALLOW_ASYNC_UNSAFE'] = 'true'
|
||
os.environ.setdefault('DJANGO_SETTINGS_MODULE', 'tianxuan.settings')
|
||
|
||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||
import django; django.setup()
|
||
|
||
import numpy as np
|
||
import polars as pl
|
||
from sklearn.metrics import adjusted_rand_score, normalized_mutual_info_score
|
||
from collections import Counter, defaultdict
|
||
|
||
from analysis.models import AnalysisRun, ClusterResult, EntityProfile, ClusterFeature
|
||
from analysis.session_store import SessionStore
|
||
from analysis.data_loader import load_csv_directory
|
||
from analysis.entity_detector import detect_entity_column
|
||
from analysis.entity_aggregator import aggregate_by_entity
|
||
from analysis.tool_registry import _handle_run_clustering, _handle_extract_features
|
||
|
||
CSV_PATH = 'data/wnl_converted.csv'
|
||
|
||
def main():
|
||
print("=" * 70)
|
||
print("CLUSTERING QUALITY EVALUATION — WNL Converted Dataset")
|
||
print("=" * 70)
|
||
|
||
# ── Step 1: Load CSV ────────────────────────────────────────────────
|
||
print("\n[1] Loading CSV...")
|
||
lf, schema, row_count, file_count, memory_mb = load_csv_directory(CSV_PATH)
|
||
print(f" Rows: {row_count}, Files: {file_count}, Memory: {memory_mb:.1f} MB")
|
||
|
||
store = SessionStore()
|
||
store.drop_all()
|
||
ds_id = 'eval_ds'
|
||
store.store_dataset(ds_id, lf, schema=schema, metadata={
|
||
'row_count': row_count, 'file_count': file_count, 'csv_glob': CSV_PATH,
|
||
})
|
||
|
||
# ── Step 2: Detect entity column ────────────────────────────────────
|
||
print("\n[2] Detecting entity column...")
|
||
result = detect_entity_column(ds_id)
|
||
used_col = result.get('recommended')
|
||
candidates = result.get('candidates', [])
|
||
print(f" Detected: {used_col} (from {len(candidates)} candidates)")
|
||
|
||
# ── Step 3: Aggregate by entity ──────────────────────────────────────
|
||
print("\n[3] Aggregating by entity...")
|
||
agg_lf, feature_columns = aggregate_by_entity(lf, used_col, schema)
|
||
df_agg = agg_lf.collect(streaming=True)
|
||
entity_ds_id = 'eval_entity'
|
||
store.store_dataset(entity_ds_id, agg_lf, schema=dict(zip(feature_columns, ['']*len(feature_columns))),
|
||
metadata={'row_count': len(df_agg), 'entity_column': used_col, 'csv_glob': CSV_PATH})
|
||
print(f" Entities: {len(df_agg)}, Features: {len(feature_columns)}")
|
||
|
||
# ── Step 4: Clustering ──────────────────────────────────────────────
|
||
print("\n[4] Running clustering (HDBSCAN)...")
|
||
user_features = [
|
||
c for c in feature_columns
|
||
if not c.startswith('_') and c not in ('first_seen', 'last_seen', 'src_ip')
|
||
][:10]
|
||
print(f" Features: {user_features}")
|
||
|
||
clust_result = asyncio.run(_handle_run_clustering(
|
||
dataset_id=entity_ds_id,
|
||
cluster_columns=user_features,
|
||
algorithm='hdbscan', params={}, random_state=42,
|
||
))
|
||
|
||
if 'error' in clust_result:
|
||
print(f" ERROR: {clust_result['error']}")
|
||
return
|
||
|
||
cluster_id = clust_result.get('cluster_result_id', '')
|
||
n_clusters = clust_result.get('n_clusters', 0)
|
||
quality = clust_result.get('quality_metrics', {})
|
||
print(f" Clusters: {n_clusters}")
|
||
print(f" Noise ratio: {quality.get('noise_ratio', 'N/A')}")
|
||
print(f" Silhouette: {quality.get('silhouette_score', 'N/A')}")
|
||
print(f" Davies-Bouldin: {quality.get('davies_bouldin_score', 'N/A')}")
|
||
print(f" CH Score: {quality.get('calinski_harabasz_score', 'N/A')}")
|
||
print(f" Cluster sizes: {quality.get('cluster_sizes', {})}")
|
||
|
||
# ── Step 5: Feature extraction ──────────────────────────────────────
|
||
print("\n[5] Extracting distinguishing features...")
|
||
feat_result = asyncio.run(_handle_extract_features(
|
||
dataset_id=entity_ds_id, cluster_result_id=cluster_id,
|
||
top_k=10, method='zscore', save_to_db=True,
|
||
))
|
||
if 'error' in feat_result:
|
||
print(f" Feature extraction error: {feat_result['error']}")
|
||
else:
|
||
print(f" Features extracted: {feat_result.get('n_features', 0)}")
|
||
|
||
# ── Get cluster labels from SessionStore ───────────────────────────
|
||
print("\n[6] Reading cluster labels from session store...")
|
||
cluster_entry = store.get_cluster_result(cluster_id)
|
||
if cluster_entry is None:
|
||
print(" ERROR: Cluster result not found in SessionStore!")
|
||
return
|
||
|
||
labels_array = np.array(cluster_entry['labels'])
|
||
n_entities = len(labels_array)
|
||
print(f" Labels array size: {n_entities}")
|
||
|
||
# ── Map entities to ground-truth labels ────────────────────────────
|
||
print("\n[7] Mapping entities to ground-truth labels...")
|
||
# The entity column is :ipd (dst IP)
|
||
# The original CSV has labels per-flow. After aggregation by :ipd,
|
||
# each entity may have multiple flows with different labels.
|
||
# We need the majority label per entity.
|
||
|
||
# Read original data
|
||
df_raw = pl.read_csv(CSV_PATH, schema_overrides={'label': pl.Utf8})
|
||
|
||
# Get entity values from aggregated data
|
||
entity_values = df_agg[used_col].to_list()
|
||
|
||
# Build entity -> label mapping from raw data (majority label per entity)
|
||
raw_entities = df_raw[used_col].to_list()
|
||
raw_labels = df_raw['label'].to_list()
|
||
|
||
entity_label_counter = defaultdict(Counter)
|
||
for ent, lbl in zip(raw_entities, raw_labels):
|
||
if lbl and str(lbl).strip():
|
||
entity_label_counter[ent][str(lbl).strip()] += 1
|
||
|
||
# Get majority label for each entity in the aggregated set
|
||
true_labels = []
|
||
for ev in entity_values:
|
||
if ev in entity_label_counter:
|
||
true_labels.append(entity_label_counter[ev].most_common(1)[0][0])
|
||
else:
|
||
true_labels.append('Unknown')
|
||
|
||
true_labels = np.array(true_labels)
|
||
print(f" Mapped {len(true_labels)} entities to ground-truth labels")
|
||
|
||
# Label distribution (ground truth among entities)
|
||
print("\n Ground truth label distribution among entities:")
|
||
gt_counts = Counter(true_labels)
|
||
for lbl, cnt in gt_counts.most_common():
|
||
print(f" {lbl:>20s}: {cnt:4d} ({cnt/len(true_labels)*100:5.1f}%)")
|
||
|
||
# ── Purity Score ────────────────────────────────────────────────────
|
||
print("\n[8] Computing purity score...")
|
||
# For each cluster, find the most common true label
|
||
total_correct = 0
|
||
cluster_label_map = {}
|
||
for label in sorted(set(labels_array)):
|
||
mask = labels_array == label
|
||
cluster_entities = true_labels[mask]
|
||
if len(cluster_entities) == 0:
|
||
continue
|
||
most_common = Counter(cluster_entities).most_common(1)[0]
|
||
cluster_label_map[int(label)] = {
|
||
'dominant_label': most_common[0],
|
||
'dominant_count': most_common[1],
|
||
'size': len(cluster_entities),
|
||
}
|
||
total_correct += most_common[1]
|
||
|
||
purity = total_correct / len(labels_array)
|
||
print(f" Purity: {purity:.4f} ({total_correct}/{len(labels_array)} correct)")
|
||
|
||
# ── Adjusted Rand Index ─────────────────────────────────────────────
|
||
print("\n[9] Computing Adjusted Rand Index (ARI)...")
|
||
# Need numeric labels for ARI - map truth labels to integers
|
||
unique_true = sorted(set(true_labels))
|
||
true_to_int = {lbl: i for i, lbl in enumerate(unique_true)}
|
||
true_numeric = np.array([true_to_int[lbl] for lbl in true_labels])
|
||
|
||
# Noise points (-1) cause issues for some metrics, filter them
|
||
valid_mask = labels_array >= 0
|
||
if valid_mask.sum() > 1:
|
||
ari = adjusted_rand_score(true_numeric[valid_mask], labels_array[valid_mask])
|
||
nmi = normalized_mutual_info_score(true_numeric[valid_mask], labels_array[valid_mask])
|
||
print(f" ARI (excluding noise): {ari:.4f}")
|
||
print(f" NMI (excluding noise): {nmi:.4f}")
|
||
else:
|
||
print(" (Not enough non-noise points to compute ARI/NMI)")
|
||
ari = float('nan')
|
||
|
||
# With noise (-1) - some implementations treat -1 as just another label
|
||
ari_with_noise = adjusted_rand_score(true_numeric, labels_array)
|
||
print(f" ARI (with noise): {ari_with_noise:.4f}")
|
||
|
||
# ── Detailed per-cluster analysis ───────────────────────────────────
|
||
print("\n" + "=" * 70)
|
||
print("PER-CLUSTER PROFILES")
|
||
print("=" * 70)
|
||
|
||
# Get feature stats (z-scores) from feature extraction
|
||
features_list = feat_result.get('features', [])
|
||
|
||
# Also compute per-cluster feature means and z-scores from scratch
|
||
numeric_cols = [c for c in df_agg.columns
|
||
if df_agg[c].dtype in (pl.Float32, pl.Float64, pl.Int32, pl.Int64,
|
||
pl.UInt32, pl.UInt64)
|
||
and not c.startswith('_')]
|
||
|
||
# Global stats
|
||
global_means = {}
|
||
global_stds = {}
|
||
for col in numeric_cols:
|
||
s = df_agg[col].drop_nulls()
|
||
if len(s) > 0:
|
||
global_means[col] = float(s.mean())
|
||
global_stds[col] = float(s.std()) if s.std() is not None and s.std() > 1e-10 else 1.0
|
||
|
||
# Build feature index for z-scores
|
||
feature_by_cluster = defaultdict(dict)
|
||
for f in features_list:
|
||
feature_by_cluster[f['cluster_label']][f['feature_name']] = f
|
||
|
||
# Per cluster report
|
||
# First build cross-tab: cluster x service
|
||
cluster_service = defaultdict(Counter)
|
||
for lbl, true_lbl in zip(labels_array, true_labels):
|
||
cluster_service[int(lbl)][true_lbl] += 1
|
||
|
||
# Sort clusters by size
|
||
cluster_sizes = Counter(labels_array)
|
||
sorted_clusters = sorted(cluster_sizes.keys())
|
||
|
||
for label in sorted_clusters:
|
||
size = cluster_sizes[label]
|
||
pct = size / n_entities * 100
|
||
is_noise = label == -1
|
||
|
||
print(f"\n Cluster {label}: {size} entities ({pct:.1f}%){' [NOISE]' if is_noise else ''}")
|
||
|
||
# Feature z-scores (top 5)
|
||
cluster_feats = feature_by_cluster.get(label, {})
|
||
sorted_feats = sorted(cluster_feats.items(),
|
||
key=lambda x: abs(x[1].get('distinguishing_score', 0)), reverse=True)
|
||
if sorted_feats:
|
||
for feat_name, feat_data in sorted_feats[:5]:
|
||
score = feat_data.get('distinguishing_score', 0)
|
||
mean_val = feat_data.get('mean', 0)
|
||
arrow = '+' if score >= 0 else ''
|
||
print(f" {feat_name:35s} {mean_val:>10.4f} ({arrow}{score:.2f}σ)")
|
||
else:
|
||
# Compute z-scores from scratch
|
||
label_mask = labels_array == label
|
||
for col in numeric_cols[:8]:
|
||
cluster_vals = df_agg[col].filter(pl.Series('_mask', label_mask)).drop_nulls()
|
||
if len(cluster_vals) > 0 and col in global_means and col in global_stds:
|
||
c_mean = float(cluster_vals.mean())
|
||
z = (c_mean - global_means[col]) / max(global_stds[col], 1e-10)
|
||
arrow = '+' if z >= 0 else ''
|
||
print(f" {col:35s} {c_mean:>10.4f} ({arrow}{z:.2f}σ)")
|
||
|
||
# Dominant services
|
||
print(f" {'':─<50s}")
|
||
svc_counts = cluster_service.get(label, Counter())
|
||
if svc_counts:
|
||
total = sum(svc_counts.values())
|
||
dominant_services = svc_counts.most_common(5)
|
||
svc_str = ', '.join(f"{svc}({cnt/total*100:.0f}%)" for svc, cnt in dominant_services)
|
||
print(f" Services: {svc_str}")
|
||
else:
|
||
print(f" Services: (none)")
|
||
|
||
# Classification
|
||
print(f" {'':─<50s}")
|
||
if is_noise:
|
||
classification = "NOISE — unassigned by HDBSCAN"
|
||
else:
|
||
classification = classify_cluster(label, cluster_feats, svc_counts)
|
||
print(f" → {classification}")
|
||
|
||
# ── Summary ──────────────────────────────────────────────────────────
|
||
print("\n" + "=" * 70)
|
||
print("SUMMARY")
|
||
print("=" * 70)
|
||
print(f" Total entities: {n_entities}")
|
||
print(f" Number of clusters (excl. noise): {n_clusters}")
|
||
print(f" Noise count: {cluster_sizes.get(-1, 0)}")
|
||
print(f" Noise ratio: {quality.get('noise_ratio', 'N/A')}")
|
||
print(f" Silhouette score: {quality.get('silhouette_score', 'N/A')}")
|
||
print(f" Purity: {purity:.4f}")
|
||
print(f" ARI (excl. noise): {ari:.4f}" if not np.isnan(ari) else " ARI: N/A")
|
||
print(f" NMI (excl. noise): {nmi:.4f}" if not np.isnan(ari) else " NMI: N/A")
|
||
print()
|
||
|
||
# Confusion matrix: cluster x service
|
||
print(" Confusion Matrix (Cluster × Service):")
|
||
all_services = sorted(set(true_labels))
|
||
header = f" {'Cluster':>8s}" + ''.join(f"{s:>14s}" for s in all_services[:8])
|
||
if len(all_services) > 8:
|
||
header += f" ... ({len(all_services)-8} more)"
|
||
print(header)
|
||
for label in sorted_clusters:
|
||
row = cluster_service.get(label, Counter())
|
||
row_str = f" {label:>8d}" + ''.join(f"{row.get(s, 0):>14d}" for s in all_services[:8])
|
||
if len(all_services) > 8:
|
||
extras = sum(row.get(s, 0) for s in all_services[8:])
|
||
row_str += f" (+{extras})"
|
||
print(row_str)
|
||
|
||
# ── Save report ─────────────────────────────────────────────────────
|
||
report = {
|
||
'run_info': {
|
||
'dataset': CSV_PATH,
|
||
'total_rows': row_count,
|
||
'n_entities': n_entities,
|
||
'n_clusters': n_clusters,
|
||
},
|
||
'quality_metrics': {
|
||
'silhouette': quality.get('silhouette_score'),
|
||
'davies_bouldin': quality.get('davies_bouldin_score'),
|
||
'calinski_harabasz': quality.get('calinski_harabasz_score'),
|
||
'noise_ratio': quality.get('noise_ratio'),
|
||
'purity': round(purity, 4),
|
||
'ari': round(ari, 4) if not np.isnan(ari) else None,
|
||
'nmi': round(nmi, 4) if not np.isnan(ari) else None,
|
||
},
|
||
'cluster_sizes': quality.get('cluster_sizes', {}),
|
||
'cluster_service_map': {
|
||
str(k): dict(v.most_common()) for k, v in sorted(cluster_service.items())
|
||
},
|
||
'purity_breakdown': cluster_label_map,
|
||
}
|
||
|
||
report_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
||
'logs', 'clustering_eval_report.json')
|
||
os.makedirs(os.path.dirname(report_path), exist_ok=True)
|
||
with open(report_path, 'w') as f:
|
||
json.dump(report, f, indent=2, default=str)
|
||
print(f"\n Report saved to: {report_path}")
|
||
|
||
# ── Also save DB results properly ───────────────────────────────────
|
||
print("\n Saving results to DB...")
|
||
run = AnalysisRun.objects.create(
|
||
csv_glob=CSV_PATH, status='completed',
|
||
total_flows=row_count, entity_count=n_entities,
|
||
cluster_count=n_clusters, entity_column=used_col,
|
||
)
|
||
|
||
# Save ClusterResult records
|
||
for label in sorted_clusters:
|
||
size = cluster_sizes[label]
|
||
proportion = size / n_entities if n_entities > 0 else 0
|
||
ClusterResult.objects.create(
|
||
run=run, cluster_label=int(label),
|
||
size=size, proportion=proportion,
|
||
noise_ratio=quality.get('noise_ratio'),
|
||
silhouette_score=quality.get('silhouette_score'),
|
||
)
|
||
print(f" Saved {len(sorted_clusters)} ClusterResult records")
|
||
|
||
# Save EntityProfile records with correct cluster labels
|
||
saved_ep = 0
|
||
for i, ev in enumerate(entity_values):
|
||
if i >= len(labels_array):
|
||
break
|
||
lbl = labels_array[i]
|
||
cr = ClusterResult.objects.filter(run=run, cluster_label=int(lbl)).first()
|
||
EntityProfile.objects.create(
|
||
run=run, entity_value=str(ev),
|
||
cluster_label=int(lbl), cluster=cr,
|
||
feature_json={},
|
||
)
|
||
saved_ep += 1
|
||
print(f" Saved {saved_ep} EntityProfile records")
|
||
print(f" Run ID: #{run.id}")
|
||
|
||
print("\n" + "=" * 70)
|
||
print("EVALUATION COMPLETE")
|
||
print("=" * 70)
|
||
|
||
|
||
def classify_cluster(label, features, services):
|
||
"""Classify a cluster as NORMAL, SUSPICIOUS, or MALICIOUS based on TLS features and services."""
|
||
if not services:
|
||
return "UNKNOWN"
|
||
|
||
top_services = services.most_common(3)
|
||
service_names = [s[0] for s in top_services]
|
||
|
||
# Check for suspicious TLS signals in feature z-scores
|
||
high_sni_missing = False
|
||
low_tls_modern = False
|
||
high_non_std_port = False
|
||
low_recoverable = False
|
||
|
||
for feat_name, feat_data in features.items():
|
||
score = feat_data.get('distinguishing_score', 0)
|
||
if 'sni_missing' in feat_name and score > 1.5:
|
||
high_sni_missing = True
|
||
elif 'tls_modern' in feat_name and score < -1.5:
|
||
low_tls_modern = True
|
||
elif 'non_standard_port' in feat_name and score > 1.5:
|
||
high_non_std_port = True
|
||
elif 'recoverable' in feat_name and score < -1.0:
|
||
low_recoverable = True
|
||
|
||
# Also look at services
|
||
suspicious_services = {'Unknown'}
|
||
|
||
# If cluster has high concentration of one service, that's normal behavior
|
||
dominant_pct = services.most_common(1)[0][1] / sum(services.values())
|
||
|
||
if dominant_pct > 0.7:
|
||
return f"NORMAL — mostly {services.most_common(1)[0][0]} ({dominant_pct:.0%})"
|
||
|
||
# Check for suspicious feature patterns
|
||
suspicious_signals = []
|
||
if high_sni_missing:
|
||
suspicious_signals.append('high sni_missing')
|
||
if low_tls_modern:
|
||
suspicious_signals.append('low tls_modern')
|
||
if high_non_std_port:
|
||
suspicious_signals.append('non_std_ports')
|
||
if low_recoverable:
|
||
suspicious_signals.append('low_recoverability')
|
||
|
||
if len(suspicious_signals) >= 2:
|
||
return f"SUSPICIOUS — {', '.join(suspicious_signals)}; services: {', '.join(service_names)}"
|
||
elif len(suspicious_signals) >= 1:
|
||
return f"QUERY — {', '.join(suspicious_signals)}; dominant: {services.most_common(1)[0][0]}"
|
||
else:
|
||
return f"NORMAL — services: {', '.join(service_names)}"
|
||
|
||
|
||
if __name__ == '__main__':
|
||
main()
|