roglag / app /db.py
deploy
deploy ROGLAG threat console
c4f5819
Raw
History Blame Contribute Delete
8.27 kB
"""Async database layer (SQLAlchemy Core).
Portable across SQLite (default, zero-config self-host) and PostgreSQL
(docker-compose / production) using the same code. Stores classified log
events; alerts and stats are computed on read.
"""
from __future__ import annotations
from collections import Counter, defaultdict
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import (
JSON, BigInteger, Boolean, Column, Integer, MetaData, String, Table, Text,
func, select,
)
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from .classifier import CATEGORY_META
from .config import get_settings
metadata = MetaData()
# BigInteger on Postgres, but Integer on SQLite so it aliases rowid and
# autoincrements (SQLite only autoincrements INTEGER PRIMARY KEY).
PK_TYPE = BigInteger().with_variant(Integer, "sqlite")
log_events = Table(
"log_events", metadata,
Column("id", PK_TYPE, primary_key=True, autoincrement=True),
Column("ts", String(40), index=True), # event time (ISO8601)
Column("source", String(64), index=True),
Column("category", String(32), index=True), # detected category
Column("severity", String(16), index=True),
Column("mitre", String(32)),
Column("risk_score", Integer),
Column("src_ip", String(64), index=True),
Column("dst_ip", String(64)),
Column("username", String(64)),
Column("message", Text),
Column("reasons", Text),
Column("expected_category", String(32)), # original label if CSV provided one
Column("matches_label", Boolean),
Column("raw", JSON),
Column("batch_id", String(64), index=True),
Column("ingested_at", String(40), index=True),
)
_engine: AsyncEngine | None = None
def get_engine() -> AsyncEngine:
global _engine
if _engine is None:
settings = get_settings()
_engine = create_async_engine(settings.database_url, future=True, pool_pre_ping=True)
return _engine
async def init_db() -> None:
engine = get_engine()
async with engine.begin() as conn:
await conn.run_sync(metadata.create_all)
async def insert_events(events: list[dict[str, Any]]) -> int:
"""Insert enriched event rows. Returns count inserted."""
if not events:
return 0
rows = [_to_row(e) for e in events]
engine = get_engine()
async with engine.begin() as conn:
await conn.execute(log_events.insert(), rows)
return len(rows)
def _to_row(event: dict[str, Any]) -> dict[str, Any]:
return {
"ts": str(event.get("timestamp") or event.get("ts") or ""),
"source": event.get("source"),
"category": event.get("category"),
"severity": event.get("severity"),
"mitre": event.get("mitre"),
"risk_score": event.get("risk_score"),
"src_ip": event.get("src_ip"),
"dst_ip": event.get("dst_ip"),
"username": event.get("user") or event.get("username"),
"message": event.get("message"),
"reasons": "; ".join(event.get("reasons", [])) if isinstance(event.get("reasons"), list) else event.get("reasons"),
"expected_category": event.get("expected_category"),
"matches_label": event.get("matches_label"),
"raw": event.get("raw") or event,
"batch_id": event.get("batch_id"),
"ingested_at": event.get("ingested_at") or datetime.now(timezone.utc).isoformat(timespec="seconds"),
}
def _row_to_dict(row: Any) -> dict[str, Any]:
d = dict(row._mapping)
if isinstance(d.get("reasons"), str) and d["reasons"]:
d["reasons"] = d["reasons"].split("; ")
return d
async def fetch_events(limit: int = 100, offset: int = 0, category: str = "",
severity: str = "", source: str = "", q: str = "") -> dict[str, Any]:
engine = get_engine()
stmt = select(log_events)
if category:
stmt = stmt.where(log_events.c.category == category)
if severity:
stmt = stmt.where(log_events.c.severity == severity)
if source:
stmt = stmt.where(log_events.c.source == source)
if q:
stmt = stmt.where(log_events.c.message.ilike(f"%{q}%"))
count_stmt = select(func.count()).select_from(stmt.subquery())
page_stmt = stmt.order_by(log_events.c.id.desc()).limit(min(limit, 500)).offset(max(offset, 0))
async with engine.connect() as conn:
total = (await conn.execute(count_stmt)).scalar() or 0
rows = (await conn.execute(page_stmt)).fetchall()
return {"total": total, "limit": limit, "offset": offset,
"rows": [_row_to_dict(r) for r in rows]}
async def fetch_stats() -> dict[str, Any]:
engine = get_engine()
async with engine.connect() as conn:
total = (await conn.execute(select(func.count()).select_from(log_events))).scalar() or 0
cat_rows = (await conn.execute(
select(log_events.c.category, func.count()).group_by(log_events.c.category))).fetchall()
sev_rows = (await conn.execute(
select(log_events.c.severity, func.count()).group_by(log_events.c.severity))).fetchall()
src_rows = (await conn.execute(
select(log_events.c.source, func.count()).group_by(log_events.c.source))).fetchall()
ip_rows = (await conn.execute(
select(log_events.c.src_ip, func.count()).where(log_events.c.src_ip.isnot(None))
.group_by(log_events.c.src_ip).order_by(func.count().desc()).limit(10))).fetchall()
by_category = {k or "unknown": v for k, v in cat_rows}
risky = sum(v for k, v in by_category.items() if k != "benign")
return {
"total": total,
"risky": risky,
"risk_rate": round((risky / total) * 100, 2) if total else 0.0,
"by_category": by_category,
"by_severity": {k or "unknown": v for k, v in sev_rows},
"by_source": {k or "unknown": v for k, v in src_rows},
"top_ips": [[k, v] for k, v in ip_rows],
"categories": CATEGORY_META,
}
async def fetch_alerts(limit: int = 20) -> dict[str, Any]:
"""Group critical events into alerts (computed on read)."""
engine = get_engine()
stmt = select(log_events).where(log_events.c.severity == "critical").order_by(log_events.c.id.desc()).limit(2000)
async with engine.connect() as conn:
rows = [_row_to_dict(r) for r in (await conn.execute(stmt)).fetchall()]
groups: dict[str, list[dict[str, Any]]] = defaultdict(list)
for row in rows:
if row.get("category") == "priv_esc":
key = "|".join([str(row.get("category")), str(row.get("username")), str(row.get("src_ip"))])
else:
key = "|".join([str(row.get("category")), str(row.get("src_ip")), str(row.get("dst_ip")), str(row.get("username"))])
groups[key].append(row)
alerts = []
for key, items in groups.items():
items.sort(key=lambda r: str(r.get("ts", "")))
first, last = items[0], items[-1]
meta = CATEGORY_META.get(first.get("category"), CATEGORY_META["benign"])
alerts.append({
"title": meta["label"],
"category": first.get("category"),
"severity": "critical",
"mitre": meta["mitre"],
"count": len(items),
"first_seen": first.get("ts"),
"last_seen": last.get("ts"),
"src_ip": first.get("src_ip") or "-",
"dst_ip": first.get("dst_ip") or "-",
"user": first.get("username") or "-",
"message": last.get("message"),
"reasons": last.get("reasons"),
"recommended_action": meta["runbook"],
})
alerts.sort(key=lambda a: str(a["last_seen"]), reverse=True)
return {"total": len(alerts), "limit": limit, "alerts": alerts[:limit]}
async def search_rows(q: str, limit: int = 400) -> list[dict[str, Any]]:
"""Fetch candidate rows for keyword RAG (message LIKE any token)."""
engine = get_engine()
stmt = select(log_events).order_by(log_events.c.id.desc()).limit(limit)
if q:
stmt = select(log_events).where(log_events.c.message.ilike(f"%{q}%")).order_by(log_events.c.id.desc()).limit(limit)
async with engine.connect() as conn:
rows = (await conn.execute(stmt)).fetchall()
return [_row_to_dict(r) for r in rows]