CoolFace
Apppublic

mihir2007/Cyber-Risk

sourceHugging Faceupdated 16d agoView on Hugging Face
0likes
kaggle_loader.py331 linesDownload Raw Back to root
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