Files
2026-06-27 21:24:57 -04:00

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"),
)