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

290 lines
11 KiB
Python
Executable File

"""
utils/helpers.py — small, dependency-light utilities.
Contains:
* Severity helpers (label → colour, ordering).
* CVE-ID validation/normalisation.
* Datetime parsing/formatting for NVD timestamps.
* A tiny on-disk JSON cache with TTL for rate-limit friendliness.
* The async SQLite store for bot config + user notes.
"""
from __future__ import annotations
import json
import logging
import re
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import aiosqlite
from config import settings
log = logging.getLogger("vulnforge.helpers")
# CVE IDs look like: CVE-2024-12345 (4-digit year, 4+ digit sequence).
_CVE_RE = re.compile(r"^CVE-\d{4}-\d{4,}$", re.IGNORECASE)
# ──────────────────────────────────────────────────────────────
# Severity helpers
# ──────────────────────────────────────────────────────────────
SEVERITY_ORDER = {"CRITICAL": 4, "HIGH": 3, "MEDIUM": 2, "LOW": 1, "NONE": 0, "UNKNOWN": 0}
def severity_from_score(score: float | None) -> str:
"""Map a CVSS base score (0-10) to a qualitative severity label."""
if score is None:
return "UNKNOWN"
if score >= 9.0:
return "CRITICAL"
if score >= 7.0:
return "HIGH"
if score >= 4.0:
return "MEDIUM"
if score > 0.0:
return "LOW"
return "NONE"
def severity_color(severity: str) -> int:
"""Return the cyberpunk palette colour int for a severity label."""
mapping = {
"CRITICAL": settings.color_critical,
"HIGH": settings.color_high,
"MEDIUM": settings.color_medium,
"LOW": settings.color_low,
"NONE": settings.color_info,
"UNKNOWN": settings.color_dark,
}
return mapping.get(severity.upper(), settings.color_info)
def severity_emoji(severity: str) -> str:
"""A glanceable emoji marker per severity tier."""
return {
"CRITICAL": "🟥",
"HIGH": "🟧",
"MEDIUM": "🟨",
"LOW": "🟦",
"NONE": "⬜",
"UNKNOWN": "⬛",
}.get(severity.upper(), "⬛")
# ──────────────────────────────────────────────────────────────
# CVE-ID + datetime helpers
# ──────────────────────────────────────────────────────────────
def is_valid_cve_id(cve_id: str) -> bool:
"""True when the string matches the canonical CVE-YYYY-NNNN format."""
return bool(_CVE_RE.match(cve_id.strip()))
def normalize_cve_id(cve_id: str) -> str:
"""Upper-case + strip a CVE id for consistent lookups/keys."""
return cve_id.strip().upper()
def parse_nvd_datetime(raw: str | None) -> datetime | None:
"""Parse NVD ISO timestamps (e.g. '2024-05-01T12:00:00.000') as UTC."""
if not raw:
return None
cleaned = raw.replace("Z", "+00:00")
for fmt in ("%Y-%m-%dT%H:%M:%S.%f%z", "%Y-%m-%dT%H:%M:%S%z"):
try:
return datetime.strptime(cleaned, fmt)
except ValueError:
continue
try:
dt = datetime.fromisoformat(cleaned)
return dt if dt.tzinfo else dt.replace(tzinfo=timezone.utc)
except ValueError:
log.debug("Could not parse NVD datetime: %s", raw)
return None
def discord_timestamp(dt: datetime | None, style: str = "R") -> str:
"""Render a Discord dynamic timestamp tag, or 'N/A' if dt is None."""
if dt is None:
return "N/A"
return f"<t:{int(dt.timestamp())}:{style}>"
def truncate(text: str, limit: int) -> str:
"""Trim text to `limit` chars (Discord field/desc safe), adding an ellipsis."""
if text is None:
return ""
return text if len(text) <= limit else text[: limit - 1].rstrip() + "…"
# ──────────────────────────────────────────────────────────────
# JSON file cache with TTL
# ──────────────────────────────────────────────────────────────
class JSONCache:
"""A minimal on-disk JSON cache keyed by filename, with TTL expiry."""
def __init__(self, cache_dir: Path) -> None:
self._dir = cache_dir
self._dir.mkdir(parents=True, exist_ok=True)
def _path(self, key: str) -> Path:
safe = re.sub(r"[^A-Za-z0-9_.-]", "_", key)
return self._dir / f"{safe}.json"
def get(self, key: str, ttl: int) -> Any | None:
"""Return cached payload if present and younger than ttl seconds."""
path = self._path(key)
if not path.exists():
return None
try:
if time.time() - path.stat().st_mtime > ttl:
return None
with path.open("r", encoding="utf-8") as fh:
return json.load(fh)
except (OSError, json.JSONDecodeError) as exc:
log.warning("Cache read failed for %s: %s", key, exc)
return None
def set(self, key: str, payload: Any) -> None:
"""Persist payload to disk (best-effort; failures are logged only)."""
path = self._path(key)
try:
with path.open("w", encoding="utf-8") as fh:
json.dump(payload, fh)
except OSError as exc:
log.warning("Cache write failed for %s: %s", key, exc)
# ──────────────────────────────────────────────────────────────
# SQLite store (config + user notes)
# ──────────────────────────────────────────────────────────────
class Store:
"""
Async SQLite wrapper.
Tables:
guild_config(guild_id PK, digest_channel_id, notify_role_id, updated_at)
cve_notes(id PK, guild_id, user_id, cve_id, note, created_at)
kev_seen(cve_id PK, date_added) -- tracks KEV entries we've recorded
"""
def __init__(self, db_path: Path) -> None:
self._db_path = db_path
async def init(self) -> None:
"""Create tables if they do not exist."""
self._db_path.parent.mkdir(parents=True, exist_ok=True)
async with aiosqlite.connect(self._db_path) as db:
await db.execute(
"""
CREATE TABLE IF NOT EXISTS guild_config (
guild_id INTEGER PRIMARY KEY,
digest_channel_id INTEGER,
notify_role_id INTEGER,
updated_at TEXT
)
"""
)
await db.execute(
"""
CREATE TABLE IF NOT EXISTS cve_notes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
guild_id INTEGER,
user_id INTEGER NOT NULL,
cve_id TEXT NOT NULL,
note TEXT NOT NULL,
created_at TEXT NOT NULL
)
"""
)
await db.execute(
"""
CREATE TABLE IF NOT EXISTS kev_seen (
cve_id TEXT PRIMARY KEY,
date_added TEXT
)
"""
)
await db.commit()
log.info("SQLite store initialised at %s", self._db_path)
# ── guild config ─────────────────────────────────────────
async def set_guild_config(
self,
guild_id: int,
digest_channel_id: int | None,
notify_role_id: int | None,
) -> None:
async with aiosqlite.connect(self._db_path) as db:
await db.execute(
"""
INSERT INTO guild_config (guild_id, digest_channel_id, notify_role_id, updated_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(guild_id) DO UPDATE SET
digest_channel_id = excluded.digest_channel_id,
notify_role_id = excluded.notify_role_id,
updated_at = excluded.updated_at
""",
(
guild_id,
digest_channel_id,
notify_role_id,
datetime.now(timezone.utc).isoformat(),
),
)
await db.commit()
async def get_guild_config(self, guild_id: int) -> dict[str, Any] | None:
async with aiosqlite.connect(self._db_path) as db:
db.row_factory = aiosqlite.Row
async with db.execute(
"SELECT * FROM guild_config WHERE guild_id = ?", (guild_id,)
) as cur:
row = await cur.fetchone()
return dict(row) if row else None
async def all_guild_configs(self) -> list[dict[str, Any]]:
async with aiosqlite.connect(self._db_path) as db:
db.row_factory = aiosqlite.Row
async with db.execute(
"SELECT * FROM guild_config WHERE digest_channel_id IS NOT NULL"
) as cur:
return [dict(r) for r in await cur.fetchall()]
# ── user notes ───────────────────────────────────────────
async def add_note(
self, guild_id: int | None, user_id: int, cve_id: str, note: str
) -> None:
async with aiosqlite.connect(self._db_path) as db:
await db.execute(
"""
INSERT INTO cve_notes (guild_id, user_id, cve_id, note, created_at)
VALUES (?, ?, ?, ?, ?)
""",
(
guild_id,
user_id,
normalize_cve_id(cve_id),
note,
datetime.now(timezone.utc).isoformat(),
),
)
await db.commit()
async def get_notes(self, cve_id: str, limit: int = 5) -> list[dict[str, Any]]:
async with aiosqlite.connect(self._db_path) as db:
db.row_factory = aiosqlite.Row
async with db.execute(
"""
SELECT * FROM cve_notes
WHERE cve_id = ?
ORDER BY created_at DESC
LIMIT ?
""",
(normalize_cve_id(cve_id), limit),
) as cur:
return [dict(r) for r in await cur.fetchall()]