mihir2007/Cyber-Risk
0
1"""2kaggle_loader.py3----------------4Ingests and maps real CVEs from the Kaggle CISA KEV & EPSS Enriched Dataset5(`data/cve.csv`) into the CRQ platform's PostgreSQL database.6 7Topology Mapping Logic:81. High-exploitability `NETWORK` attack vector CVEs with high EPSS scores9 are mapped to critical perimeter nodes:10 - `upi-gateway-prod-01`11 - `upi-gateway-prod-02`12 - `mobile-banking-gw-01`132. Database and privilege escalation CVEs (`LOCAL` / `NETWORK` with high CVSS)14 are mapped to backend crown jewels:15 - `core-banking-db-01`16 - `cust-db-primary`173. Representative CVEs are distributed across remaining enterprise assets by tier.184. Existing mock vulnerabilities are cleared and real CVEs are persisted in an19 atomic SQLAlchemy transaction.20"""21 22from __future__ import annotations23 24import logging25import os26import random27from pathlib import Path28from typing import Any29 30import numpy as np31import pandas as pd32from sqlalchemy.orm import Session33 34from models import Asset, NetworkEdge, Vulnerability35 36logger = logging.getLogger(__name__)37 38# Preferred paths for dataset discovery39DEFAULT_CSV_CANDIDATES = [40 Path("data") / "cve.csv",41 Path("Data") / "cve.csv",42 Path("data") / "cve_enriched.csv",43]44 45_PERIMETER_HOSTS = [46 "upi-gateway-prod-01",47 "upi-gateway-prod-02",48 "mobile-banking-gw-01",49]50 51_CROWN_JEWEL_HOSTS = [52 "core-banking-db-01",53 "cust-db-primary",54]55 56 57def find_dataset_path(custom_path: str | Path | None = None) -> Path:58 """Finds existing dataset file among common relative and absolute paths."""59 if custom_path:60 p = Path(custom_path)61 if p.exists():62 return p63 64 base_dir = Path(__file__).resolve().parent65 candidates = DEFAULT_CSV_CANDIDATES + [base_dir / c for c in DEFAULT_CSV_CANDIDATES]66 67 for cand in candidates:68 if cand.exists() and cand.is_file():69 return cand70 71 raise FileNotFoundError(72 f"Kaggle CVE dataset file not found. Checked: {[str(c) for c in candidates]}"73 )74 75 76def _resolve_column(df: pd.DataFrame, candidates: list[str]) -> str | None:77 """Returns the first matching column name present in DataFrame."""78 for col in candidates:79 if col in df.columns:80 return col81 return None82 83 84def load_and_normalize_cve_df(csv_path: Path) -> pd.DataFrame:85 """86 Loads `cve.csv` using Pandas and standardizes schema fields:87 - cve_id: string88 - cvss_score: float [0.0 - 10.0]89 - exploitability_score: float [0.0 - 10.0]90 - attack_vector: string ('NETWORK', 'LOCAL', 'ADJACENT_NETWORK', 'PHYSICAL')91 - epss_score: float [0.0 - 1.0]92 - cisa_kev: bool93 """94 logger.info("Reading Kaggle CVE dataset from %s", csv_path)95 df = pd.read_csv(csv_path, low_memory=False)96 97 col_id = _resolve_column(df, ["cve_id", "cveId", "CVE_ID", "cve"])98 col_cvss = _resolve_column(df, ["base_score", "cvss", "cvssV3_baseScore", "cvss_score"])99 col_exploit = _resolve_column(100 df, ["exploitability_score", "cvssV3_exploitabilityScore", "exploitabilityScore"]101 )102 col_av = _resolve_column(df, ["attack_vector", "attackVector", "attack_vec"])103 col_epss = _resolve_column(df, ["epss_score", "epss", "epssScore", "epss_perc"])104 col_kev = _resolve_column(df, ["cisa_kev", "known_exploited", "cisaKev", "cisa_known_exploited"])105 106 if not col_id:107 raise ValueError("Could not locate CVE identifier column in dataset.")108 109 normalized = pd.DataFrame()110 normalized["cve_id"] = df[col_id].astype(str).str.strip()111 112 if col_cvss:113 normalized["cvss_score"] = pd.to_numeric(df[col_cvss], errors="coerce").fillna(5.0)114 else:115 normalized["cvss_score"] = 5.0116 117 if col_exploit:118 normalized["exploitability_score"] = pd.to_numeric(119 df[col_exploit], errors="coerce"120 ).fillna(2.5)121 else:122 normalized["exploitability_score"] = 2.5123 124 if col_av:125 normalized["attack_vector"] = df[col_av].astype(str).str.upper().str.strip()126 else:127 normalized["attack_vector"] = "NETWORK"128 129 if col_epss:130 normalized["epss_score"] = pd.to_numeric(df[col_epss], errors="coerce").fillna(0.01)131 else:132 normalized["epss_score"] = 0.01133 134 if col_kev:135 kev_series = df[col_kev]136 if kev_series.dtype == bool:137 normalized["cisa_kev"] = kev_series.fillna(False)138 else:139 normalized["cisa_kev"] = (140 kev_series.astype(str)141 .str.lower()142 .isin(["true", "1", "yes", "y", "t"])143 )144 else:145 normalized["cisa_kev"] = False146 147 # Filter out empty or placeholder CVE IDs148 normalized = normalized[normalized["cve_id"].str.startswith("CVE-")].copy()149 normalized["cvss_score"] = normalized["cvss_score"].clip(0.0, 10.0).round(1)150 normalized["exploitability_score"] = normalized["exploitability_score"].clip(0.0, 10.0).round(1)151 normalized["epss_score"] = normalized["epss_score"].clip(0.0, 1.0).round(5)152 153 return normalized154 155 156def ensure_topology_nodes(db: Session) -> None:157 """Ensures crown jewel node `core-banking-db-01` and required topology edges exist."""158 core_db = db.query(Asset).filter(Asset.hostname == "core-banking-db-01").first()159 if not core_db:160 core_db = Asset(161 hostname="core-banking-db-01",162 asset_type="Customer DB",163 tier="Critical",164 business_unit="Retail Banking",165 revenue_per_minute=75_000.0,166 pii_records_count=350_000,167 financial_records_count=500_000,168 is_rbi_regulated=True,169 is_sebi_regulated=True,170 network_hops_from_internet=3,171 )172 db.add(core_db)173 db.flush()174 175 # Ensure edge exists in network_edges if table is used176 existing_edge = (177 db.query(NetworkEdge)178 .filter(179 NetworkEdge.source_node == "core-bank-switch-01",180 NetworkEdge.target_node == "core-banking-db-01",181 )182 .first()183 )184 if not existing_edge:185 db.add(186 NetworkEdge(187 source_node="core-bank-switch-01",188 target_node="core-banking-db-01",189 weight=1.5,190 protocol="SQL_NET",191 is_segmented=True,192 )193 )194 db.flush()195 196 197def ingest_kaggle_cves(198 db: Session, csv_path: str | Path | None = None199) -> dict[str, Any]:200 """201 Ingests Kaggle CVE dataset, performs topology mapping to assets, clears existing202 mock entries, and commits new vulnerabilities in an atomic transaction.203 """204 rng = random.Random(42)205 dataset_file = find_dataset_path(csv_path)206 df = load_and_normalize_cve_df(dataset_file)207 208 # Make sure all required asset nodes exist209 ensure_topology_nodes(db)210 211 assets = db.query(Asset).all()212 if not assets:213 raise RuntimeError("No assets found in database. Seed asset inventory first.")214 215 asset_by_host = {a.hostname: a for a in assets}216 217 # 1. Candidate pools218 # Perimeter: NETWORK attack vector, high EPSS, high exploitability219 perimeter_pool = df[220 (df["attack_vector"] == "NETWORK")221 & (df["epss_score"] >= 0.20)222 & (df["exploitability_score"] >= 2.8)223 ].sort_values(by=["epss_score", "cvss_score"], ascending=False)224 225 if len(perimeter_pool) < 30:226 # Fallback to general high EPSS network pool227 perimeter_pool = df[df["attack_vector"] == "NETWORK"].sort_values(228 by=["epss_score", "cvss_score"], ascending=False229 )230 231 # Backend Crown Jewels: Database / privilege escalation (LOCAL or NETWORK with high CVSS >= 8.5)232 crown_jewel_pool = df[233 (df["attack_vector"].isin(["LOCAL", "NETWORK"]))234 & (df["cvss_score"] >= 8.5)235 ].sort_values(by=["cvss_score", "epss_score"], ascending=False)236 237 # General pool for other assets238 general_pool = df.sort_values(by=["cvss_score", "epss_score"], ascending=False)239 240 used_cve_ids: set[str] = set()241 new_vuln_records: list[Vulnerability] = []242 mapped_assets: set[str] = set()243 244 def _pick_cves(pool_df: pd.DataFrame, count: int) -> list[dict[str, Any]]:245 selected: list[dict[str, Any]] = []246 for _, row in pool_df.iterrows():247 cve_id = row["cve_id"]248 if cve_id in used_cve_ids:249 continue250 used_cve_ids.add(cve_id)251 selected.append(252 {253 "cve_id": cve_id,254 "cvss_score": float(row["cvss_score"]),255 "epss_score": float(row["epss_score"]),256 "cisa_kev": bool(row["cisa_kev"]),257 "patch_available": True,258 "is_patched": rng.random() < 0.15, # 15% remediated259 }260 )261 if len(selected) >= count:262 break263 return selected264 265 # 2. Map Perimeter nodes266 for hostname in _PERIMETER_HOSTS:267 asset = asset_by_host.get(hostname)268 if not asset:269 continue270 cve_data_list = _pick_cves(perimeter_pool, count=6)271 for cve_data in cve_data_list:272 new_vuln_records.append(Vulnerability(asset_id=asset.id, **cve_data))273 mapped_assets.add(hostname)274 275 # 3. Map Crown Jewels276 for hostname in _CROWN_JEWEL_HOSTS:277 asset = asset_by_host.get(hostname)278 if not asset:279 continue280 cve_data_list = _pick_cves(crown_jewel_pool, count=6)281 for cve_data in cve_data_list:282 new_vuln_records.append(Vulnerability(asset_id=asset.id, **cve_data))283 mapped_assets.add(hostname)284 285 # 4. Map Remaining Assets by Tier286 for asset in assets:287 if asset.hostname in mapped_assets:288 continue289 290 if asset.tier == "Critical":291 count = 4292 pool = crown_jewel_pool293 elif asset.tier == "Medium":294 count = 2295 pool = general_pool[general_pool["cvss_score"] >= 6.0]296 else:297 count = 1298 pool = general_pool299 300 cve_data_list = _pick_cves(pool, count=count)301 for cve_data in cve_data_list:302 new_vuln_records.append(Vulnerability(asset_id=asset.id, **cve_data))303 mapped_assets.add(asset.hostname)304 305 # 5. Database transaction: delete existing mock entries and insert real CVEs306 try:307 db.query(Vulnerability).delete()308 db.add_all(new_vuln_records)309 db.commit()310 except Exception:311 db.rollback()312 raise313 314 # Invalidate optimizer caches so new CVE landscape is immediately reflected315 try:316 from optimizer import clear_optimizer_cache317 clear_optimizer_cache()318 except ImportError:319 pass320 321 result = {322 "status": "ok",323 "vulnerabilities_ingested": len(new_vuln_records),324 "mapped_assets": sorted(list(mapped_assets)),325 "message": (326 f"Successfully ingested {len(new_vuln_records)} real Kaggle CVEs "327 f"across {len(mapped_assets)} enterprise assets."328 ),329 }330 return result331 