330 lines
14 KiB
Python
Executable File
330 lines
14 KiB
Python
Executable File
"""
|
|
utils/api.py — async data layer for NVD CVE API v2 and CISA KEV feed.
|
|
|
|
Responsibilities:
|
|
* Single shared aiohttp.ClientSession (created on bot startup, closed on shutdown).
|
|
* Rate-limit-friendly: on-disk TTL caching + retry/backoff for NVD 403/429/5xx.
|
|
* Normalises raw API JSON into simple dataclasses (CVERecord, KEVEntry) so the
|
|
cogs/embeds never touch the messy upstream schema directly.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
from dataclasses import dataclass, field
|
|
from datetime import datetime, timedelta, timezone
|
|
from typing import Any
|
|
|
|
import aiohttp
|
|
|
|
from config import settings
|
|
from utils.helpers import (
|
|
JSONCache,
|
|
parse_nvd_datetime,
|
|
severity_from_score,
|
|
)
|
|
|
|
log = logging.getLogger("vulnforge.api")
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────
|
|
# Data models
|
|
# ──────────────────────────────────────────────────────────────
|
|
@dataclass
|
|
class CVERecord:
|
|
"""A normalised view of a single NVD CVE entry."""
|
|
|
|
cve_id: str
|
|
description: str
|
|
published: datetime | None
|
|
last_modified: datetime | None
|
|
cvss_score: float | None
|
|
cvss_severity: str
|
|
cvss_vector: str | None
|
|
cvss_version: str | None
|
|
cwes: list[str] = field(default_factory=list)
|
|
references: list[str] = field(default_factory=list)
|
|
status: str = "Unknown"
|
|
|
|
@property
|
|
def nvd_url(self) -> str:
|
|
return f"{settings.nvd_cve_web}/{self.cve_id}"
|
|
|
|
|
|
@dataclass
|
|
class KEVEntry:
|
|
"""A normalised CISA Known-Exploited-Vulnerability entry."""
|
|
|
|
cve_id: str
|
|
vendor: str
|
|
product: str
|
|
name: str
|
|
date_added: str
|
|
short_description: str
|
|
required_action: str
|
|
due_date: str
|
|
ransomware: str
|
|
|
|
@property
|
|
def date_added_dt(self) -> datetime | None:
|
|
try:
|
|
return datetime.strptime(self.date_added, "%Y-%m-%d").replace(
|
|
tzinfo=timezone.utc
|
|
)
|
|
except (ValueError, TypeError):
|
|
return None
|
|
|
|
|
|
# ──────────────────────────────────────────────────────────────
|
|
# Client
|
|
# ──────────────────────────────────────────────────────────────
|
|
class SecurityAPI:
|
|
"""Async client for NVD + CISA KEV with caching and retry/backoff."""
|
|
|
|
def __init__(self) -> None:
|
|
self._session: aiohttp.ClientSession | None = None
|
|
self._cache = JSONCache(settings.cache_dir)
|
|
# In-memory KEV index so per-CVE lookups don't reparse the whole feed.
|
|
self._kev_index: dict[str, KEVEntry] = {}
|
|
self._kev_loaded_at: datetime | None = None
|
|
self._kev_lock = asyncio.Lock()
|
|
|
|
# ── lifecycle ────────────────────────────────────────────
|
|
async def start(self) -> None:
|
|
if self._session is None or self._session.closed:
|
|
headers = {"User-Agent": settings.user_agent}
|
|
if settings.nvd_api_key:
|
|
headers["apiKey"] = settings.nvd_api_key
|
|
timeout = aiohttp.ClientTimeout(total=settings.http_timeout)
|
|
self._session = aiohttp.ClientSession(headers=headers, timeout=timeout)
|
|
log.info("aiohttp session started (NVD key: %s)", bool(settings.nvd_api_key))
|
|
|
|
async def close(self) -> None:
|
|
if self._session and not self._session.closed:
|
|
await self._session.close()
|
|
log.info("aiohttp session closed")
|
|
|
|
def _require_session(self) -> aiohttp.ClientSession:
|
|
if self._session is None or self._session.closed:
|
|
raise RuntimeError("SecurityAPI session not started — call start() first.")
|
|
return self._session
|
|
|
|
# ── low-level GET with retry/backoff ─────────────────────
|
|
async def _get_json(
|
|
self, url: str, params: dict[str, Any] | None = None
|
|
) -> dict[str, Any] | None:
|
|
session = self._require_session()
|
|
delay = 2.0
|
|
for attempt in range(1, settings.request_retries + 1):
|
|
try:
|
|
async with session.get(url, params=params) as resp:
|
|
if resp.status == 200:
|
|
return await resp.json()
|
|
if resp.status in (403, 429, 503, 502, 500):
|
|
log.warning(
|
|
"GET %s → %s (attempt %s/%s), backing off %.1fs",
|
|
url, resp.status, attempt, settings.request_retries, delay,
|
|
)
|
|
await asyncio.sleep(delay)
|
|
delay *= 2
|
|
continue
|
|
log.error("GET %s → unexpected status %s", url, resp.status)
|
|
return None
|
|
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
|
log.warning(
|
|
"GET %s failed (attempt %s/%s): %s",
|
|
url, attempt, settings.request_retries, exc,
|
|
)
|
|
await asyncio.sleep(delay)
|
|
delay *= 2
|
|
log.error("GET %s exhausted retries", url)
|
|
return None
|
|
|
|
# ── NVD: single CVE ──────────────────────────────────────
|
|
async def fetch_cve(self, cve_id: str) -> CVERecord | None:
|
|
"""Fetch a single CVE by id (cached for cve_cache_ttl)."""
|
|
cache_key = f"cve_{cve_id}"
|
|
cached = self._cache.get(cache_key, settings.cve_cache_ttl)
|
|
if cached is not None:
|
|
return self._parse_cve_item(cached)
|
|
|
|
data = await self._get_json(settings.nvd_base_url, {"cveId": cve_id})
|
|
if not data:
|
|
return None
|
|
vulns = data.get("vulnerabilities") or []
|
|
if not vulns:
|
|
return None
|
|
item = vulns[0]
|
|
self._cache.set(cache_key, item)
|
|
return self._parse_cve_item(item)
|
|
|
|
# ── NVD: recent CVEs (last N days) ───────────────────────
|
|
async def fetch_recent_cves(self, days: int = 7) -> list[CVERecord]:
|
|
"""
|
|
Fetch CVEs published in the last `days` days.
|
|
|
|
NVD limits the pubStartDate/pubEndDate window to 120 days and paginates
|
|
with resultsPerPage / startIndex. We page through everything in the
|
|
window and cache the assembled list.
|
|
"""
|
|
cache_key = f"recent_{days}d"
|
|
cached = self._cache.get(cache_key, settings.cve_cache_ttl)
|
|
if cached is not None:
|
|
return [self._parse_cve_item(i) for i in cached]
|
|
|
|
end = datetime.now(timezone.utc)
|
|
start = end - timedelta(days=days)
|
|
iso = "%Y-%m-%dT%H:%M:%S.000"
|
|
results_per_page = 2000
|
|
start_index = 0
|
|
collected: list[dict[str, Any]] = []
|
|
|
|
while True:
|
|
params = {
|
|
"pubStartDate": start.strftime(iso),
|
|
"pubEndDate": end.strftime(iso),
|
|
"resultsPerPage": results_per_page,
|
|
"startIndex": start_index,
|
|
}
|
|
data = await self._get_json(settings.nvd_base_url, params)
|
|
if not data:
|
|
break
|
|
vulns = data.get("vulnerabilities") or []
|
|
collected.extend(vulns)
|
|
total = int(data.get("totalResults", 0))
|
|
start_index += results_per_page
|
|
if start_index >= total or not vulns:
|
|
break
|
|
# Be gentle on the API between pages.
|
|
await asyncio.sleep(0.7 if settings.nvd_api_key else 6.0)
|
|
|
|
self._cache.set(cache_key, collected)
|
|
log.info("Fetched %s CVEs from the last %s days", len(collected), days)
|
|
return [self._parse_cve_item(i) for i in collected]
|
|
|
|
# ── CISA KEV feed ────────────────────────────────────────
|
|
async def _load_kev(self, force: bool = False) -> dict[str, KEVEntry]:
|
|
"""Load + index the KEV feed (cached on disk + in memory)."""
|
|
async with self._kev_lock:
|
|
fresh = (
|
|
self._kev_index
|
|
and self._kev_loaded_at
|
|
and datetime.now(timezone.utc) - self._kev_loaded_at
|
|
< timedelta(seconds=settings.kev_cache_ttl)
|
|
)
|
|
if fresh and not force:
|
|
return self._kev_index
|
|
|
|
payload = None if force else self._cache.get("kev_feed", settings.kev_cache_ttl)
|
|
if payload is None:
|
|
payload = await self._get_json(settings.cisa_kev_url)
|
|
if payload:
|
|
self._cache.set("kev_feed", payload)
|
|
|
|
index: dict[str, KEVEntry] = {}
|
|
for v in (payload or {}).get("vulnerabilities", []):
|
|
entry = KEVEntry(
|
|
cve_id=v.get("cveID", "").upper(),
|
|
vendor=v.get("vendorProject", "Unknown"),
|
|
product=v.get("product", "Unknown"),
|
|
name=v.get("vulnerabilityName", ""),
|
|
date_added=v.get("dateAdded", ""),
|
|
short_description=v.get("shortDescription", ""),
|
|
required_action=v.get("requiredAction", ""),
|
|
due_date=v.get("dueDate", ""),
|
|
ransomware=v.get("knownRansomwareCampaignUse", "Unknown"),
|
|
)
|
|
if entry.cve_id:
|
|
index[entry.cve_id] = entry
|
|
|
|
self._kev_index = index
|
|
self._kev_loaded_at = datetime.now(timezone.utc)
|
|
log.info("KEV feed loaded: %s entries", len(index))
|
|
return index
|
|
|
|
async def get_kev_entry(self, cve_id: str) -> KEVEntry | None:
|
|
index = await self._load_kev()
|
|
return index.get(cve_id.upper())
|
|
|
|
async def is_in_kev(self, cve_id: str) -> bool:
|
|
return (await self.get_kev_entry(cve_id)) is not None
|
|
|
|
async def latest_kev(self, limit: int = 10) -> list[KEVEntry]:
|
|
"""Return the most recently added KEV entries (newest first)."""
|
|
index = await self._load_kev()
|
|
entries = sorted(
|
|
index.values(),
|
|
key=lambda e: e.date_added or "",
|
|
reverse=True,
|
|
)
|
|
return entries[:limit]
|
|
|
|
async def kev_added_since(self, days: int = 7) -> list[KEVEntry]:
|
|
"""KEV entries whose dateAdded falls within the last `days` days."""
|
|
index = await self._load_kev()
|
|
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
|
|
out: list[KEVEntry] = []
|
|
for e in index.values():
|
|
dt = e.date_added_dt
|
|
if dt and dt >= cutoff:
|
|
out.append(e)
|
|
out.sort(key=lambda e: e.date_added or "", reverse=True)
|
|
return out
|
|
|
|
# ── parsing ──────────────────────────────────────────────
|
|
@staticmethod
|
|
def _parse_cve_item(item: dict[str, Any]) -> CVERecord:
|
|
"""Convert a raw NVD `vulnerabilities[]` item into a CVERecord."""
|
|
cve = item.get("cve", {})
|
|
cve_id = cve.get("id", "UNKNOWN")
|
|
|
|
# Description: prefer English.
|
|
description = ""
|
|
for d in cve.get("descriptions", []):
|
|
if d.get("lang") == "en":
|
|
description = d.get("value", "")
|
|
break
|
|
if not description and cve.get("descriptions"):
|
|
description = cve["descriptions"][0].get("value", "")
|
|
|
|
# CVSS: prefer v3.1 → v3.0 → v2 (highest available).
|
|
score = severity = vector = version = None
|
|
metrics = cve.get("metrics", {})
|
|
for key in ("cvssMetricV31", "cvssMetricV30", "cvssMetricV2"):
|
|
arr = metrics.get(key)
|
|
if arr:
|
|
data = arr[0].get("cvssData", {})
|
|
score = data.get("baseScore")
|
|
vector = data.get("vectorString")
|
|
version = data.get("version")
|
|
severity = (
|
|
data.get("baseSeverity")
|
|
or arr[0].get("baseSeverity")
|
|
or severity_from_score(score)
|
|
)
|
|
break
|
|
|
|
cwes: list[str] = []
|
|
for weakness in cve.get("weaknesses", []):
|
|
for desc in weakness.get("description", []):
|
|
val = desc.get("value")
|
|
if val and val.upper() != "NVD-CWE-NOINFO" and val not in cwes:
|
|
cwes.append(val)
|
|
|
|
refs = [r.get("url") for r in cve.get("references", []) if r.get("url")]
|
|
|
|
return CVERecord(
|
|
cve_id=cve_id,
|
|
description=description or "No description available.",
|
|
published=parse_nvd_datetime(cve.get("published")),
|
|
last_modified=parse_nvd_datetime(cve.get("lastModified")),
|
|
cvss_score=score,
|
|
cvss_severity=(severity or severity_from_score(score)).upper(),
|
|
cvss_vector=vector,
|
|
cvss_version=version,
|
|
cwes=cwes,
|
|
references=refs,
|
|
status=cve.get("vulnStatus", "Unknown"),
|
|
)
|