290 lines
11 KiB
Python
Executable File
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()]
|