#!/usr/bin/env python3 """ Cluster malware samples by similarity into family groups. Takes a similarity matrix (from similarity_analyzer.py) or computes one from a sample directory, then applies hierarchical or threshold-based clustering to group related samples. Usage: python3 cluster_samples.py --similarity-matrix matrix.json --threshold 0.6 python3 cluster_samples.py --samples-dir ./samples/ --threshold 0.5 python3 cluster_samples.py --similarity-matrix matrix.json --method hierarchical """ from __future__ import annotations import argparse import json import os import sys from datetime import datetime, timezone from pathlib import Path # --------------------------------------------------------------------------- # Optional dependency imports # --------------------------------------------------------------------------- try: import numpy as np HAS_NUMPY = True except ImportError: HAS_NUMPY = False try: from scipy.cluster.hierarchy import linkage, fcluster from scipy.spatial.distance import squareform HAS_SCIPY = True except ImportError: HAS_SCIPY = False # --------------------------------------------------------------------------- # Clustering algorithms (no external dependencies) # --------------------------------------------------------------------------- def threshold_cluster(matrix: list, names: list, threshold: float) -> list: """ Simple threshold-based clustering. Two samples are in the same cluster if their similarity >= threshold. Uses union-find for transitive closure. """ n = len(names) parent = list(range(n)) def find(x) -> dict: while parent[x] != x: parent[x] = parent[parent[x]] x = parent[x] return x def union(a, b) -> None: ra, rb = find(a), find(b) if ra != rb: parent[ra] = rb for i in range(n): for j in range(i + 1, n): if matrix[i][j] >= threshold: union(i, j) # Build clusters clusters_map = {} for i in range(n): root = find(i) clusters_map.setdefault(root, []).append(i) clusters = [] for cluster_id, (root, members) in enumerate(sorted(clusters_map.items())): # Compute intra-cluster average similarity if len(members) > 1: similarities = [] for ii in range(len(members)): for jj in range(ii + 1, len(members)): similarities.append(matrix[members[ii]][members[jj]]) avg_sim = sum(similarities) / len(similarities) else: avg_sim = 1.0 clusters.append({ "cluster_id": cluster_id, "size": len(members), "samples": [names[m] for m in members], "sample_indices": members, "average_similarity": round(avg_sim, 4), "confidence": _similarity_to_confidence(avg_sim), }) return clusters def hierarchical_cluster( matrix: list, names: list, threshold: float, linkage_method: str = "average" ) -> list: """ Hierarchical agglomerative clustering using scipy. Falls back to threshold clustering if scipy is unavailable. """ if not HAS_SCIPY or not HAS_NUMPY: print( "Warning: scipy/numpy not available, falling back to threshold clustering", file=sys.stderr, ) return threshold_cluster(matrix, names, threshold) n = len(names) sim_matrix = np.array(matrix) # Convert similarity to distance (1 - similarity) dist_matrix = 1.0 - sim_matrix np.fill_diagonal(dist_matrix, 0.0) # Ensure symmetry and non-negativity dist_matrix = (dist_matrix + dist_matrix.T) / 2.0 dist_matrix = np.maximum(dist_matrix, 0.0) # Convert to condensed form for scipy condensed = squareform(dist_matrix, checks=False) # Perform hierarchical clustering Z = linkage(condensed, method=linkage_method) # Cut the dendrogram at the distance threshold (1 - similarity threshold) dist_threshold = 1.0 - threshold labels = fcluster(Z, t=dist_threshold, criterion="distance") # Build cluster output clusters_map = {} for i, label in enumerate(labels): clusters_map.setdefault(int(label), []).append(i) clusters = [] for cluster_id, (label, members) in enumerate(sorted(clusters_map.items())): if len(members) > 1: similarities = [] for ii in range(len(members)): for jj in range(ii + 1, len(members)): similarities.append(matrix[members[ii]][members[jj]]) avg_sim = sum(similarities) / len(similarities) else: avg_sim = 1.0 clusters.append({ "cluster_id": cluster_id, "size": len(members), "samples": [names[m] for m in members], "sample_indices": members, "average_similarity": round(avg_sim, 4), "confidence": _similarity_to_confidence(avg_sim), }) return clusters def _similarity_to_confidence(similarity: float) -> str: """Convert average similarity score to a confidence label.""" if similarity >= 0.8: return "high" elif similarity >= 0.5: return "medium" elif similarity >= 0.3: return "low" else: return "very low" # --------------------------------------------------------------------------- # Input loading # --------------------------------------------------------------------------- def load_similarity_matrix(path: str) -> dict: """Load a similarity matrix from JSON.""" try: with open(path, "r", encoding="utf-8") as f: data = json.load(f) except FileNotFoundError: print(f"Error: File not found: {path}", file=sys.stderr) sys.exit(1) except json.JSONDecodeError as e: print(f"Error: Invalid JSON: {e}", file=sys.stderr) sys.exit(1) if "matrix" not in data or "sample_names" not in data: print( "Error: Input JSON must contain 'matrix' and 'sample_names' keys", file=sys.stderr, ) sys.exit(1) return data def compute_matrix_from_samples(samples_dir: str, metrics: list) -> dict: """Compute similarity matrix by importing from similarity_analyzer.""" # Import sibling module script_dir = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, script_dir) try: from similarity_analyzer import build_similarity_matrix except ImportError: print( "Error: Could not import similarity_analyzer.py from scripts directory", file=sys.stderr, ) sys.exit(1) sample_paths = [] for entry in sorted(os.listdir(samples_dir)): full_path = os.path.join(samples_dir, entry) if os.path.isfile(full_path): sample_paths.append(os.path.abspath(full_path)) if len(sample_paths) < 2: print("Error: Need at least 2 samples to cluster", file=sys.stderr) sys.exit(1) return build_similarity_matrix(sample_paths, metrics) # --------------------------------------------------------------------------- # Output formatting # --------------------------------------------------------------------------- def format_clusters_report(clusters: list, matrix_data: dict) -> dict: """Format clustering results into a structured report.""" report = { "timestamp": datetime.now(tz=timezone.utc).isoformat(), "total_samples": matrix_data.get("sample_count", 0), "total_clusters": len(clusters), "singleton_clusters": sum(1 for c in clusters if c["size"] == 1), "multi_sample_clusters": sum(1 for c in clusters if c["size"] > 1), "largest_cluster_size": max(c["size"] for c in clusters) if clusters else 0, "clusters": clusters, "summary": [], } # Generate human-readable summary for cluster in clusters: if cluster["size"] > 1: report["summary"].append( f"Cluster {cluster['cluster_id']}: {cluster['size']} samples " f"(avg similarity: {cluster['average_similarity']:.2f}, " f"confidence: {cluster['confidence']}) - " f"{', '.join(cluster['samples'][:5])}" + ("..." if len(cluster["samples"]) > 5 else "") ) singletons = [c["samples"][0] for c in clusters if c["size"] == 1] if singletons: report["summary"].append( f"Unclustered singletons ({len(singletons)}): " + ", ".join(singletons[:10]) + ("..." if len(singletons) > 10 else "") ) return report # --------------------------------------------------------------------------- # CLI # --------------------------------------------------------------------------- def main() -> None: parser = argparse.ArgumentParser( description="Cluster malware samples by code similarity", formatter_class=argparse.RawDescriptionHelpFormatter, epilog=( "Examples:\n" " %(prog)s --similarity-matrix matrix.json --threshold 0.6\n" " %(prog)s --samples-dir ./samples/ --threshold 0.5\n" " %(prog)s --similarity-matrix matrix.json --method hierarchical\n" ), ) input_group = parser.add_mutually_exclusive_group(required=True) input_group.add_argument( "--input", "--similarity-matrix", "-m", dest="similarity_matrix", help="Path to similarity matrix JSON (from similarity_analyzer.py)", ) input_group.add_argument( "--samples-dir", "-d", help="Directory of samples to compute similarity and cluster", ) parser.add_argument( "--threshold", "-t", type=float, default=0.6, help="Similarity threshold for clustering (0.0-1.0, default: 0.6)", ) parser.add_argument( "--method", choices=["threshold", "hierarchical"], default="threshold", help="Clustering method (default: threshold)", ) parser.add_argument( "--linkage", choices=["single", "complete", "average", "ward"], default="average", help="Linkage method for hierarchical clustering (default: average)", ) parser.add_argument( "--metrics", nargs="+", default=["all"], help="Metrics for on-the-fly matrix computation (default: all)", ) parser.add_argument( "--output", "-o", help="Output file path (default: stdout)", ) parser.add_argument( "--format", default="json", choices=["json", "text", "csv"], help="Output format (default: json)", ) args = parser.parse_args() # Validate threshold if not 0.0 <= args.threshold <= 1.0: print("Error: Threshold must be between 0.0 and 1.0", file=sys.stderr) sys.exit(1) # Load or compute similarity matrix if args.similarity_matrix: matrix_data = load_similarity_matrix(args.similarity_matrix) else: matrix_data = compute_matrix_from_samples(args.samples_dir, args.metrics) matrix = matrix_data["matrix"] names = matrix_data["sample_names"] print( f"Clustering {len(names)} samples with threshold={args.threshold}, " f"method={args.method}", file=sys.stderr, ) # Perform clustering if args.method == "hierarchical": clusters = hierarchical_cluster( matrix, names, args.threshold, args.linkage ) else: clusters = threshold_cluster(matrix, names, args.threshold) # Format output report = format_clusters_report(clusters, matrix_data) output_text = json.dumps(report, indent=2, default=str) # Write output if args.output: Path(args.output).parent.mkdir(parents=True, exist_ok=True) Path(args.output).write_text(output_text, encoding="utf-8") print(f"Clustering results written to {args.output}", file=sys.stderr) else: print(output_text) # Print summary to stderr print(f"\n--- Clustering Summary ---", file=sys.stderr) print(f"Total samples: {report['total_samples']}", file=sys.stderr) print(f"Clusters found: {report['total_clusters']}", file=sys.stderr) print(f"Multi-sample clusters: {report['multi_sample_clusters']}", file=sys.stderr) print(f"Singletons: {report['singleton_clusters']}", file=sys.stderr) for line in report["summary"]: print(f" {line}", file=sys.stderr) if __name__ == "__main__": main()