178 lines
6.7 KiB
Python
178 lines
6.7 KiB
Python
# SPDX-License-Identifier: GPL-3.0-or-later
|
|
"""Tests for the read-only web dashboard: data layer + auth (real server)."""
|
|
import json
|
|
import ssl
|
|
import tempfile
|
|
import threading
|
|
import unittest
|
|
import urllib.error
|
|
import urllib.request
|
|
from pathlib import Path
|
|
|
|
from enodia_sentinel import web
|
|
from enodia_sentinel import incident
|
|
from enodia_sentinel.alert import Alert, Severity
|
|
from enodia_sentinel.config import Config
|
|
|
|
|
|
def _make_cfg(tmp: Path) -> Config:
|
|
c = Config()
|
|
c.log_dir = tmp
|
|
return c
|
|
|
|
|
|
def _write_alert(tmp: Path, name: str, severity: str, sigs):
|
|
alerts = [{"signature": s, "sid": 100010, "severity": severity} for s in sigs]
|
|
(tmp / f"{name}.json").write_text(json.dumps({
|
|
"time": "2026-05-31T00:00:00-07:00", "host": "woofbox",
|
|
"severity": severity, "alerts": alerts,
|
|
}))
|
|
(tmp / f"{name}.log").write_text(f"=== ENODIA SENTINEL ALERT ===\n{severity}\n")
|
|
|
|
|
|
def _incident_alert(sig: str, pids=()):
|
|
return Alert(Severity.CRITICAL, sig, f"k:{sig}", f"{sig} detail",
|
|
tuple(pids), sid=100010, classtype="test")
|
|
|
|
|
|
class TestDataLayer(unittest.TestCase):
|
|
def setUp(self):
|
|
self.dir = tempfile.TemporaryDirectory()
|
|
self.tmp = Path(self.dir.name)
|
|
self.cfg = _make_cfg(self.tmp)
|
|
|
|
def tearDown(self):
|
|
self.dir.cleanup()
|
|
|
|
def test_list_and_status(self):
|
|
_write_alert(self.tmp, "alert-20260531-000001", "CRITICAL", ["reverse_shell"])
|
|
_write_alert(self.tmp, "alert-20260531-000002", "HIGH", ["new_listener"])
|
|
alerts = web.list_alerts(self.cfg)
|
|
self.assertEqual(len(alerts), 2)
|
|
self.assertEqual(alerts[0]["name"], "alert-20260531-000002.log") # newest first
|
|
st = web.daemon_status(self.cfg)
|
|
self.assertEqual(st["total_alerts"], 2)
|
|
self.assertEqual(st["counts"]["CRITICAL"], 1)
|
|
self.assertFalse(st["running"]) # no live pidfile
|
|
|
|
def test_get_alert_and_traversal(self):
|
|
_write_alert(self.tmp, "alert-20260531-000001", "CRITICAL", ["x"])
|
|
got = web.get_alert(self.cfg, "alert-20260531-000001.log")
|
|
self.assertIn("text", got)
|
|
self.assertIn("json", got)
|
|
# path traversal / bad names rejected
|
|
self.assertIsNone(web.get_alert(self.cfg, "../../etc/passwd"))
|
|
self.assertIsNone(web.get_alert(self.cfg, "events.log"))
|
|
|
|
def test_tail_events(self):
|
|
(self.tmp / "events.log").write_text("l1\nl2\nl3\n")
|
|
self.assertEqual(web.tail_events(self.cfg, 2), ["l2", "l3"])
|
|
|
|
def test_incident_and_response_plan_data(self):
|
|
_write_alert(self.tmp, "alert-20260531-000001", "CRITICAL", ["reverse_shell"])
|
|
iid = incident.record(self.cfg, "alert-20260531-000001.log",
|
|
[_incident_alert("reverse_shell", [4242])],
|
|
lineage={4242}, when=1000.0, host="h")
|
|
incs = web.list_incidents(self.cfg)
|
|
self.assertEqual(incs[0]["id"], iid)
|
|
inc = web.get_incident(self.cfg, iid)
|
|
self.assertEqual(inc["incident"]["id"], iid)
|
|
self.assertEqual(len(inc["timeline"]), 1)
|
|
plan = web.response_plan(self.cfg, iid)
|
|
self.assertEqual(plan["incident_id"], iid)
|
|
self.assertEqual(plan["mode"], "dry-run")
|
|
|
|
def test_posture_report_shape(self):
|
|
def runner(_cfg):
|
|
yield Alert(Severity.HIGH, "ssh_root_login", "k",
|
|
"root SSH login enabled", sid=100040,
|
|
classtype="host-posture")
|
|
yield Alert(Severity.MEDIUM, "sudo_nopasswd", "k2",
|
|
"passwordless sudo", sid=100043,
|
|
classtype="host-posture")
|
|
|
|
report = web.posture_report(self.cfg, runner=runner)
|
|
self.assertEqual(report["count"], 2)
|
|
self.assertEqual(report["counts"], {"HIGH": 1, "MEDIUM": 1})
|
|
self.assertEqual(report["findings"][0]["signature"], "ssh_root_login")
|
|
|
|
|
|
class TestNetworkHelpers(unittest.TestCase):
|
|
def test_is_loopback(self):
|
|
self.assertTrue(web.is_loopback("127.0.0.1"))
|
|
self.assertFalse(web.is_loopback("100.64.1.2"))
|
|
|
|
def test_resolve_bind_explicit(self):
|
|
c = Config()
|
|
c.web_bind = "100.64.1.2"
|
|
self.assertEqual(web.resolve_bind(c), "100.64.1.2")
|
|
|
|
def test_ensure_token_persists(self):
|
|
with tempfile.TemporaryDirectory() as d:
|
|
c = _make_cfg(Path(d))
|
|
t1 = web.ensure_token(c)
|
|
t2 = web.ensure_token(c)
|
|
self.assertTrue(t1)
|
|
self.assertEqual(t1, t2) # stable across calls
|
|
|
|
def test_self_signed_tls_material_is_created(self):
|
|
with tempfile.TemporaryDirectory() as d:
|
|
c = _make_cfg(Path(d))
|
|
cert, key = web.ensure_tls_cert(c, "127.0.0.1")
|
|
self.assertTrue(cert.is_file())
|
|
self.assertTrue(key.is_file())
|
|
|
|
|
|
class TestAuth(unittest.TestCase):
|
|
"""Spin up the real server on loopback and check token enforcement."""
|
|
|
|
def setUp(self):
|
|
self.dir = tempfile.TemporaryDirectory()
|
|
self.tmp = Path(self.dir.name)
|
|
_write_alert(self.tmp, "alert-20260531-000001", "CRITICAL", ["reverse_shell"])
|
|
self.cfg = _make_cfg(self.tmp)
|
|
self.cfg.web_bind = "127.0.0.1"
|
|
self.cfg.web_port = 0 # ephemeral
|
|
self.cfg.web_token = "secret-token"
|
|
self.httpd, _bind, _tok = web.build_server(self.cfg)
|
|
self.assertTrue(Path(self.httpd.tls_cert).is_file())
|
|
self.port = self.httpd.server_address[1]
|
|
self.t = threading.Thread(target=self.httpd.serve_forever, daemon=True)
|
|
self.t.start()
|
|
self.ctx = ssl._create_unverified_context()
|
|
|
|
def tearDown(self):
|
|
self.httpd.shutdown()
|
|
self.httpd.server_close()
|
|
self.dir.cleanup()
|
|
|
|
def _get(self, path, token=None):
|
|
url = f"https://127.0.0.1:{self.port}{path}"
|
|
req = urllib.request.Request(url)
|
|
if token:
|
|
req.add_header("Authorization", f"Bearer {token}")
|
|
return urllib.request.urlopen(req, timeout=4, context=self.ctx)
|
|
|
|
def test_unauthorized_without_token(self):
|
|
with self.assertRaises(urllib.error.HTTPError) as cm:
|
|
self._get("/api/status")
|
|
self.assertEqual(cm.exception.code, 401)
|
|
|
|
def test_authorized_with_token(self):
|
|
resp = self._get("/api/status", token="secret-token")
|
|
data = json.loads(resp.read())
|
|
self.assertEqual(data["total_alerts"], 1)
|
|
|
|
def test_token_via_query_param(self):
|
|
resp = self._get("/api/alerts?token=secret-token")
|
|
self.assertEqual(resp.status, 200)
|
|
|
|
def test_posture_endpoint(self):
|
|
resp = self._get("/api/posture", token="secret-token")
|
|
data = json.loads(resp.read())
|
|
self.assertIn("findings", data)
|
|
self.assertIn("count", data)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|