Files
routeros-suricata-tzsp/app/rules.py
T
2026-08-14 11:33:01 +02:00

525 lines
21 KiB
Python

from __future__ import annotations
import glob
import ipaddress
import os
import re
import shutil
import signal
import subprocess
import tempfile
import threading
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable
from .config import Config
@dataclass(frozen=True)
class RuleActionResult:
ok: bool
message: str
class RuleManager:
MAX_RULE_BYTES = 512 * 1024
MAX_THRESHOLD_BYTES = 256 * 1024
SOURCE_INDEX_URL = "https://www.openinfosecfoundation.org/rules/index.yaml"
DEFAULT_SOURCE = "et/open"
SOURCE_NAME_RE = re.compile(r"^[A-Za-z0-9_.+-]+/[A-Za-z0-9_.+-]+$")
def __init__(
self,
config: Config,
pid_provider: Callable[[], int | None],
suricata_available: bool = True,
) -> None:
self.config = config
self.pid_provider = pid_provider
self.suricata_available = suricata_available
self._lock = threading.RLock()
self._operation_lock = threading.RLock()
self._update_lock = threading.Lock()
self._last_result = "not changed"
self._ensure_files()
def _ensure_files(self) -> None:
for path in (self.config.suricata_custom_rules, self.config.suricata_threshold_config):
Path(path).parent.mkdir(parents=True, exist_ok=True)
Path(path).touch(exist_ok=True)
def status(self) -> dict:
custom = self._read(self.config.suricata_custom_rules)
builtin = self._read(self.config.suricata_local_rules)
threshold = self._read(self.config.suricata_threshold_config)
with self._lock:
last_result = self._last_result
vendor_rules = "/var/lib/suricata/rules/suricata.rules"
source_index = _first_existing_path(
"/var/lib/suricata/update/cache/index.yaml",
"/var/lib/suricata/rules/cache/index.yaml",
)
return {
"available": self.suricata_available,
"custom_rules_path": self.config.suricata_custom_rules,
"extra_rules_glob": self.config.suricata_extra_rules_glob,
"threshold_config_path": self.config.suricata_threshold_config,
"builtin_rule_count": _count_rules(builtin),
"custom_rule_count": _count_rules(custom),
"managed_rule_files": len(glob.glob(self.config.suricata_extra_rules_glob)),
"threshold_entry_count": _count_config_entries(threshold),
"suppressed_sids": _suppressed_sids(threshold),
"vendor_rules_path": vendor_rules,
"vendor_rules_size_bytes": _file_size(vendor_rules),
"vendor_rules_updated_at": _file_mtime_iso(vendor_rules),
"source_index_updated_at": _file_mtime_iso(source_index) if source_index else None,
"source_index_url": self.SOURCE_INDEX_URL,
"last_result": last_result,
}
def content(self) -> dict:
return {
"custom_rules": self._read(self.config.suricata_custom_rules),
"threshold_config": self._read(self.config.suricata_threshold_config),
"status": self.status(),
}
def replace_custom_rules(self, content: str) -> RuleActionResult:
return self._replace_and_reload(
self.config.suricata_custom_rules,
content,
self.MAX_RULE_BYTES,
"custom rules",
)
def replace_threshold_config(self, content: str) -> RuleActionResult:
return self._replace_and_reload(
self.config.suricata_threshold_config,
content,
self.MAX_THRESHOLD_BYTES,
"threshold configuration",
)
def suppress_sid(
self,
sid: int,
track: str | None = None,
ip: str | None = None,
) -> RuleActionResult:
sid = int(sid)
if sid <= 0:
return RuleActionResult(False, "SID must be a positive integer")
track = (track or "").strip().lower()
if track in {"", "global"}:
line = f"suppress gen_id 1, sig_id {sid}"
label = f"SID {sid}"
elif track in {"by_src", "by_dst"}:
if not ip:
return RuleActionResult(False, "IP is required for scoped suppression")
try:
network = ipaddress.ip_network(str(ip).strip(), strict=False)
except ValueError:
return RuleActionResult(False, "invalid suppression IP/network")
ip_text = str(network.network_address) if network.prefixlen == network.max_prefixlen else str(network)
line = f"suppress gen_id 1, sig_id {sid}, track {track}, ip {ip_text}"
label = f"SID {sid} {track} {ip_text}"
else:
return RuleActionResult(False, "track must be global, by_src or by_dst")
with self._operation_lock:
current = self._read(self.config.suricata_threshold_config)
existing = {item.strip().casefold() for item in current.splitlines() if item.strip()}
if line.casefold() in existing:
return RuleActionResult(True, f"{label} is already suppressed")
if current and not current.endswith("\n"):
current += "\n"
current += line + "\n"
return self.replace_threshold_config(current)
def update_vendor_rules(self) -> RuleActionResult:
if not self.suricata_available:
return RuleActionResult(False, "Suricata rule updates are unavailable in this mode")
if not self._update_lock.acquire(blocking=False):
return RuleActionResult(False, "a Suricata rule-source operation is already running")
try:
return self._run_vendor_update_unlocked()
finally:
self._update_lock.release()
def source_catalog(self) -> dict:
if not self.suricata_available:
return {
"ok": False,
"error": "Suricata rule sources are unavailable in this mode",
"sources": [],
}
catalog = self._run_suricata_update(["list-sources", "--free"], timeout=60)
if catalog.returncode != 0:
return {
"ok": False,
"error": _command_tail(catalog.stdout, "could not list rule sources"),
"sources": [],
}
enabled_proc = self._run_suricata_update(["list-sources", "--enabled"], timeout=30)
enabled = _parse_enabled_sources(enabled_proc.stdout or "") if enabled_proc.returncode == 0 else set()
sources = _parse_source_catalog(catalog.stdout or "")
for source in sources:
source["default"] = source["name"] == self.DEFAULT_SOURCE
source["enabled"] = source["default"] or source["name"] in enabled
source["can_toggle"] = not source["default"] and not bool(source.get("parameters"))
return {
"ok": True,
"catalog": "OISF suricata-update source index",
"catalog_url": self.SOURCE_INDEX_URL,
"free_only": True,
"sources": sources,
"enabled_sources": sorted(
{source["name"] for source in sources if source.get("enabled")}
),
"status": self.status(),
}
def refresh_source_catalog(self) -> RuleActionResult:
if not self.suricata_available:
return RuleActionResult(False, "Suricata rule sources are unavailable in this mode")
if not self._update_lock.acquire(blocking=False):
return RuleActionResult(False, "a Suricata rule-source operation is already running")
try:
proc = self._run_suricata_update(["update-sources"], timeout=120)
if proc.returncode == 0:
result = RuleActionResult(True, _command_tail(proc.stdout, "OISF source catalog refreshed"))
else:
result = RuleActionResult(False, _command_tail(proc.stdout, "OISF source catalog refresh failed"))
with self._lock:
self._last_result = result.message
return result
finally:
self._update_lock.release()
def set_source_enabled(self, source_name: str, enabled: bool) -> RuleActionResult:
source_name = str(source_name or "").strip()
if not self.SOURCE_NAME_RE.fullmatch(source_name):
return RuleActionResult(False, "invalid rule source name")
if source_name == self.DEFAULT_SOURCE:
if enabled:
return RuleActionResult(True, "ET/Open is the default suricata-update source and is already active")
return RuleActionResult(False, "ET/Open is the default source and cannot be disabled from this panel")
if not self.suricata_available:
return RuleActionResult(False, "Suricata rule sources are unavailable in this mode")
if not self._update_lock.acquire(blocking=False):
return RuleActionResult(False, "a Suricata rule-source operation is already running")
try:
catalog = self.source_catalog()
if not catalog.get("ok"):
return RuleActionResult(False, str(catalog.get("error") or "could not read source catalog"))
source = next((item for item in catalog.get("sources", []) if item.get("name") == source_name), None)
if source is None:
return RuleActionResult(False, "source is not present in the current OISF catalog")
if enabled and source.get("parameters"):
params = ", ".join(source["parameters"])
return RuleActionResult(False, f"source requires parameters ({params}); configure it manually with suricata-update")
if bool(source.get("enabled")) == bool(enabled):
return RuleActionResult(True, f"{source_name} is already {'enabled' if enabled else 'disabled'}")
verb = "enable-source" if enabled else "disable-source"
proc = self._run_suricata_update([verb, source_name], timeout=60)
if proc.returncode != 0:
result = RuleActionResult(False, _command_tail(proc.stdout, f"could not {verb} {source_name}"))
else:
updated = self._run_vendor_update_unlocked()
if updated.ok:
result = RuleActionResult(
True,
f"{source_name} {'enabled' if enabled else 'disabled'}; {updated.message}",
)
else:
result = RuleActionResult(
False,
f"{source_name} {'enabled' if enabled else 'disabled'}, but rules were not rebuilt: {updated.message}",
)
with self._lock:
self._last_result = result.message
return result
finally:
self._update_lock.release()
def _run_vendor_update_unlocked(self) -> RuleActionResult:
try:
proc = subprocess.run(
["/opt/ids/scripts/update-rules.sh"],
check=False,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
timeout=300,
)
except (OSError, subprocess.TimeoutExpired) as exc:
result = RuleActionResult(False, f"vendor rule update could not run: {exc}")
else:
tail = _command_tail(proc.stdout, "vendor rules updated")
if proc.returncode == 0:
result = RuleActionResult(True, tail)
else:
result = RuleActionResult(False, f"vendor rule update failed: {tail}")
with self._lock:
self._last_result = result.message
return result
@staticmethod
def _run_suricata_update(args: list[str], timeout: int) -> subprocess.CompletedProcess:
try:
return subprocess.run(
["suricata-update", *args],
check=False,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
timeout=timeout,
)
except (OSError, subprocess.TimeoutExpired) as exc:
return subprocess.CompletedProcess(
["suricata-update", *args],
127,
stdout=f"suricata-update could not run: {exc}",
)
def reload(self) -> RuleActionResult:
if not self.suricata_available:
return RuleActionResult(False, "Suricata is not available in this mode")
pid = self.pid_provider()
if not pid:
return RuleActionResult(False, "Suricata process is not running")
try:
os.kill(int(pid), signal.SIGUSR2)
except OSError as exc:
result = RuleActionResult(False, f"reload failed: {exc}")
else:
result = RuleActionResult(True, f"rule reload requested for Suricata PID {pid}")
with self._lock:
self._last_result = result.message
return result
def validate(self, custom_rules: str, threshold_config: str) -> RuleActionResult:
if not self.suricata_available:
return RuleActionResult(False, "Suricata validation is unavailable in web-only development mode")
with tempfile.TemporaryDirectory(prefix="suricata-rules-test-") as td:
rules_dir = os.path.join(td, "rules")
threshold_path = os.path.join(td, "threshold.config")
log_dir = os.path.join(td, "log")
os.mkdir(rules_dir)
os.mkdir(log_dir)
custom_real = os.path.realpath(self.config.suricata_custom_rules)
copied = set()
for source in glob.glob(self.config.suricata_extra_rules_glob):
if os.path.realpath(source) == custom_real or not os.path.isfile(source):
continue
name = os.path.basename(source)
shutil.copyfile(source, os.path.join(rules_dir, name))
copied.add(name)
local_name = os.path.basename(self.config.suricata_local_rules) or "local.rules"
if local_name not in copied and os.path.isfile(self.config.suricata_local_rules):
shutil.copyfile(self.config.suricata_local_rules, os.path.join(rules_dir, local_name))
self._write(os.path.join(rules_dir, os.path.basename(self.config.suricata_custom_rules) or "custom.rules"), custom_rules)
self._write(threshold_path, threshold_config)
cmd = [
"suricata",
"-T",
"-c",
self.config.suricata_config,
"-l",
log_dir,
"-s",
os.path.join(rules_dir, "*.rules"),
"--set",
f"vars.address-groups.HOME_NET={self.config.suricata_home_net}",
"--set",
f"threshold-file={threshold_path}",
]
try:
proc = subprocess.run(
cmd,
check=False,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
timeout=45,
)
except (OSError, subprocess.TimeoutExpired) as exc:
return RuleActionResult(False, f"validation could not run: {exc}")
if proc.returncode == 0:
return RuleActionResult(True, "Suricata configuration and rules validated")
output = (proc.stdout or "").strip().splitlines()
tail = " | ".join(output[-8:])
if len(tail) > 1200:
tail = tail[-1200:]
return RuleActionResult(False, f"Suricata validation failed: {tail or 'unknown error'}")
def _replace_and_reload(
self,
path: str,
content: str,
max_bytes: int,
label: str,
) -> RuleActionResult:
if not isinstance(content, str):
return RuleActionResult(False, f"{label} must be text")
if len(content.encode("utf-8")) > max_bytes:
return RuleActionResult(False, f"{label} exceeds {max_bytes} bytes")
with self._operation_lock:
custom = content if path == self.config.suricata_custom_rules else self._read(self.config.suricata_custom_rules)
threshold = content if path == self.config.suricata_threshold_config else self._read(self.config.suricata_threshold_config)
validation = self.validate(custom, threshold)
if not validation.ok:
with self._lock:
self._last_result = validation.message
return validation
self._atomic_write(path, content)
reload_result = self.reload()
if reload_result.ok:
result = RuleActionResult(True, f"{label} saved; {reload_result.message}")
else:
result = RuleActionResult(False, f"{label} saved but {reload_result.message}")
with self._lock:
self._last_result = result.message
return result
@staticmethod
def _read(path: str) -> str:
try:
with open(path, "r", encoding="utf-8") as handle:
return handle.read()
except FileNotFoundError:
return ""
@staticmethod
def _write(path: str, content: str) -> None:
with open(path, "w", encoding="utf-8") as handle:
handle.write(content)
if content and not content.endswith("\n"):
handle.write("\n")
@classmethod
def _atomic_write(cls, path: str, content: str) -> None:
directory = os.path.dirname(path) or "."
os.makedirs(directory, exist_ok=True)
fd, tmp = tempfile.mkstemp(prefix=".rules-", dir=directory, text=True)
try:
with os.fdopen(fd, "w", encoding="utf-8") as handle:
handle.write(content)
if content and not content.endswith("\n"):
handle.write("\n")
handle.flush()
os.fsync(handle.fileno())
os.chmod(tmp, 0o644)
os.replace(tmp, path)
finally:
try:
os.unlink(tmp)
except FileNotFoundError:
pass
_ANSI_RE = re.compile(r"\x1b\[[0-9;]*m")
def _strip_ansi(value: str) -> str:
return _ANSI_RE.sub("", value or "")
def _parse_source_catalog(output: str) -> list[dict]:
sources: list[dict] = []
current: dict | None = None
for raw in _strip_ansi(output).splitlines():
line = raw.strip()
if line.startswith("Name:"):
if current and current.get("name"):
sources.append(current)
current = {
"name": line.split(":", 1)[1].strip(),
"vendor": "",
"summary": "",
"license": "",
"tags": [],
"parameters": [],
}
continue
if current is None or ":" not in line:
continue
key, value = (part.strip() for part in line.split(":", 1))
key = key.lower()
if key in {"vendor", "summary", "license", "subscription", "deprecated", "obsolete"}:
current[key] = value
elif key in {"tags", "parameters", "replaces"}:
current[key] = [part.strip() for part in value.split(",") if part.strip()]
if current and current.get("name"):
sources.append(current)
return sources
def _parse_enabled_sources(output: str) -> set[str]:
result: set[str] = set()
for raw in _strip_ansi(output).splitlines():
match = re.match(r"^\s*-\s+([A-Za-z0-9_.+-]+/[A-Za-z0-9_.+-]+)\s*$", raw)
if match:
result.add(match.group(1))
return result
def _command_tail(output: str | None, fallback: str) -> str:
lines = [line.strip() for line in _strip_ansi(output or "").splitlines() if line.strip()]
tail = " | ".join(lines[-8:])
if len(tail) > 1400:
tail = tail[-1400:]
return tail or fallback
def _first_existing_path(*paths: str) -> str | None:
return next((path for path in paths if os.path.isfile(path)), None)
def _file_size(path: str) -> int:
try:
return os.path.getsize(path)
except OSError:
return 0
def _file_mtime_iso(path: str | None) -> str | None:
if not path:
return None
try:
timestamp = os.path.getmtime(path)
except OSError:
return None
return datetime.fromtimestamp(timestamp, timezone.utc).isoformat()
def _count_rules(content: str) -> int:
return sum(
1
for line in content.splitlines()
if line.strip() and not line.lstrip().startswith("#")
)
def _count_config_entries(content: str) -> int:
return sum(
1
for line in content.splitlines()
if line.strip() and not line.lstrip().startswith("#")
)
def _suppressed_sids(content: str) -> list[int]:
result: set[int] = set()
for match in re.finditer(r"^\s*suppress\s+gen_id\s+1\s*,\s*sig_id\s+(\d+)", content, re.I | re.M):
result.add(int(match.group(1)))
return sorted(result)