import os import subprocess import tempfile import time import unittest from types import SimpleNamespace from unittest.mock import patch from app.rules import ( RuleActionResult, RuleManager, _parse_enabled_sources, _parse_source_catalog, ) class RuleManagerTests(unittest.TestCase): def make_manager(self, td): custom = os.path.join(td, "custom.rules") threshold = os.path.join(td, "threshold.config") local = os.path.join(td, "local.rules") open(custom, "w").close() open(threshold, "w").close() open(local, "w").close() cfg = SimpleNamespace( suricata_custom_rules=custom, suricata_threshold_config=threshold, suricata_local_rules=local, suricata_extra_rules_glob=os.path.join(td, "*.rules"), suricata_config="/etc/suricata/suricata.yaml", suricata_output_config="/opt/ids/suricata/ids-output.yaml", suricata_home_net="[192.168.0.0/16]", suricata_persist_lib_dir=os.path.join(td, "lib", "suricata"), ) return RuleManager(cfg, pid_provider=lambda: None, suricata_available=False) def test_validation_copies_managed_dataset_files_next_to_rules(self): from unittest.mock import patch import subprocess with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) manager.suricata_available = True dataset = os.path.join(td, "ti-ja4.lst") with open(dataset, "w", encoding="ascii") as handle: handle.write("dDEzX3Rlc3Q=\n") seen = {} def fake_run(cmd, **kwargs): rules_glob = cmd[cmd.index("-s") + 1] rules_dir = os.path.dirname(rules_glob) seen["dataset"] = open(os.path.join(rules_dir, "ti-ja4.lst"), encoding="ascii").read().strip() return subprocess.CompletedProcess(cmd, 0, stdout="ok") with patch("app.rules.subprocess.run", side_effect=fake_run): result = manager.validate("", "") self.assertTrue(result.ok) self.assertEqual(seen["dataset"], "dDEzX3Rlc3Q=") def test_scoped_suppression_uses_source_ip(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) captured = {} def replace(content): captured["content"] = content return RuleActionResult(True, "saved") manager.replace_threshold_config = replace result = manager.suppress_sid(1234567, "by_src", "192.0.2.10") self.assertTrue(result.ok) self.assertIn( "suppress gen_id 1, sig_id 1234567, track by_src, ip 192.0.2.10", captured["content"], ) def test_scoped_suppression_rejects_invalid_ip(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) result = manager.suppress_sid(1234567, "by_src", "not-an-ip") self.assertFalse(result.ok) def test_parses_official_source_catalog_output(self): output = """ Name: et/open Vendor: Proofpoint Summary: Emerging Threats Open ruleset License: MIT Tags: free, ids Name: oisf/trafficid Vendor: OISF Summary: Traffic identification rules License: MIT Parameters: code, token """ sources = _parse_source_catalog(output) self.assertEqual([item["name"] for item in sources], ["et/open", "oisf/trafficid"]) self.assertEqual(sources[0]["tags"], ["free", "ids"]) self.assertEqual(sources[1]["parameters"], ["code", "token"]) def test_parses_enabled_named_sources_only(self): output = """ From /etc/suricata/update.yaml: - https://rules.example.invalid/feed.rules Enabled sources: - oisf/trafficid - sslbl/ssl-fp-blacklist """ self.assertEqual( _parse_enabled_sources(output), {"oisf/trafficid", "sslbl/ssl-fp-blacklist"}, ) def test_accepts_single_segment_official_source_names(self): self.assertIsNotNone(RuleManager.SOURCE_NAME_RE.fullmatch("pawpatrules")) output = """ Enabled sources: - pawpatrules - oisf/trafficid """ self.assertEqual(_parse_enabled_sources(output), {"pawpatrules", "oisf/trafficid"}) def test_suricata_update_commands_use_persistent_data_directory(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) manager.suricata_available = True with patch("app.rules.subprocess.run") as run: run.return_value = subprocess.CompletedProcess([], 0, stdout="ok") manager._run_suricata_update(["enable-source", "oisf/trafficid"], timeout=10) command = run.call_args.args[0] self.assertEqual(command[:3], ["suricata-update", "enable-source", "oisf/trafficid"]) self.assertEqual(command[-2:], ["-D", os.path.join(td, "lib", "suricata")]) def test_source_queue_enables_many_then_rebuilds_once(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) manager.suricata_available = True enabled = [] rebuilds = [] manager.source_catalog = lambda: { "ok": True, "sources": [ {"name": "oisf/trafficid", "enabled": False, "parameters": []}, {"name": "sslbl/ssl-fp-blacklist", "enabled": False, "parameters": []}, ], } def fake_update(args, timeout): enabled.append(list(args)) return subprocess.CompletedProcess(args, 0, stdout="enabled") manager._run_suricata_update = fake_update manager._run_vendor_update_unlocked = lambda: (rebuilds.append(True) or RuleActionResult(True, "rebuilt")) result = manager.queue_sources(["oisf/trafficid", "sslbl/ssl-fp-blacklist"]) self.assertTrue(result.ok) deadline = time.time() + 2 while manager.source_queue_status()["status"] in {"queued", "running"} and time.time() < deadline: time.sleep(0.01) status = manager.source_queue_status() self.assertEqual(status["status"], "completed") self.assertEqual(status["completed"], 2) self.assertEqual(status["failed"], 0) self.assertEqual(len(rebuilds), 1) self.assertEqual( enabled, [ ["enable-source", "oisf/trafficid"], ["enable-source", "sslbl/ssl-fp-blacklist"], ], ) def test_manual_url_source_is_added_and_rules_rebuilt(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) manager.suricata_available = True calls = [] manager._run_suricata_update = lambda args, timeout: (calls.append(list(args)) or subprocess.CompletedProcess(args, 0, stdout="added")) manager._run_vendor_update_unlocked = lambda: RuleActionResult(True, "rebuilt") result = manager.add_manual_source("vendor/community", "https://rules.example.invalid/feed.rules") self.assertTrue(result.ok) self.assertEqual(calls, [["add-source", "vendor/community", "https://rules.example.invalid/feed.rules", "--no-checksum"]]) def test_manual_url_source_rejects_non_http_url(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) manager.suricata_available = True result = manager.add_manual_source("vendor/community", "file:///tmp/feed.rules") self.assertFalse(result.ok) def test_merged_rules_are_searchable_and_paginated(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) rules = os.path.join(td, "lib", "suricata", "rules") os.makedirs(rules, exist_ok=True) with open(os.path.join(rules, "suricata.rules"), "w", encoding="utf-8") as handle: handle.write('# generated\n') handle.write('alert tcp any any -> any 80 (msg:"one"; sid:1;)\n') handle.write('alert dns any any -> any any (msg:"two"; sid:2;)\n') handle.write('alert tcp any any -> any 443 (msg:"three"; sid:3;)\n') result = manager.merged_rules("tcp", offset=0, limit=1) self.assertTrue(result["ok"]) self.assertEqual(result["total_rules"], 3) self.assertEqual(result["matched"], 2) self.assertEqual(result["next_offset"], 1) self.assertIn('sid:1', result["content"]) self.assertEqual(manager.status()["vendor_rule_count"], 3) def test_adaptive_threshold_uses_global_limit_and_snapshot_is_persistent(self): with tempfile.TemporaryDirectory() as td: manager = self.make_manager(td) captured = {} def replace(content): captured["content"] = content return RuleActionResult(True, "saved") manager.replace_threshold_config = replace result = manager.add_threshold(2222, threshold_type="limit", track="by_src", count=3, seconds=60) self.assertTrue(result.ok) self.assertIn("threshold gen_id 1, sig_id 2222, type limit, track by_src, count 3, seconds 60", captured["content"]) with open(manager.config.suricata_custom_rules, "w", encoding="utf-8") as handle: handle.write('alert ip any any -> any any (msg:"snapshot"; sid:9900002;)\n') snap = manager.create_snapshot("unit") self.assertTrue(snap.ok) snapshots = manager.list_snapshots() self.assertEqual(len(snapshots), 1) self.assertTrue(snapshots[0]["id"].endswith(".tar.gz")) if __name__ == "__main__": unittest.main()