525 lines
21 KiB
Python
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)
|