79 lines
2.7 KiB
Python
79 lines
2.7 KiB
Python
from __future__ import annotations
|
|
|
|
import ipaddress
|
|
from dataclasses import dataclass
|
|
from typing import Any
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Decision:
|
|
should_block: bool
|
|
target: str | None
|
|
reason: str
|
|
|
|
|
|
class PolicyEngine:
|
|
def __init__(
|
|
self,
|
|
auto_block: bool,
|
|
max_severity: int,
|
|
monitored_networks: str,
|
|
never_block: str,
|
|
) -> None:
|
|
self.auto_block = auto_block
|
|
self.max_severity = max_severity
|
|
self.monitored = _parse_networks(monitored_networks)
|
|
self.never_block = _parse_networks(never_block)
|
|
|
|
def evaluate(self, event: dict[str, Any]) -> Decision:
|
|
alert = event.get("alert") or {}
|
|
try:
|
|
severity = int(alert.get("severity", 999))
|
|
except (TypeError, ValueError):
|
|
return Decision(False, None, "missing or invalid severity")
|
|
|
|
if severity > self.max_severity:
|
|
return Decision(False, None, f"severity {severity} is below block threshold")
|
|
|
|
src = _ip(event.get("src_ip"))
|
|
dst = _ip(event.get("dest_ip"))
|
|
if src is None or dst is None:
|
|
return Decision(False, None, "alert has no usable IPv4/IPv6 endpoints")
|
|
|
|
src_local = self._is_monitored(src)
|
|
dst_local = self._is_monitored(dst)
|
|
if src_local == dst_local:
|
|
return Decision(False, None, "cannot identify one remote endpoint")
|
|
|
|
target = dst if src_local else src
|
|
if not target.is_global:
|
|
return Decision(False, str(target), "remote endpoint is not globally routable")
|
|
if self._is_never_block(target):
|
|
return Decision(False, str(target), "remote endpoint is on NEVER_BLOCK list")
|
|
if not self.auto_block:
|
|
return Decision(False, str(target), "observation mode: AUTO_BLOCK=false")
|
|
return Decision(True, str(target), f"severity {severity} matched automatic block policy")
|
|
|
|
def _is_monitored(self, address: ipaddress._BaseAddress) -> bool:
|
|
return any(address in network for network in self.monitored if network.version == address.version)
|
|
|
|
def _is_never_block(self, address: ipaddress._BaseAddress) -> bool:
|
|
return any(address in network for network in self.never_block if network.version == address.version)
|
|
|
|
|
|
def _parse_networks(value: str) -> list[ipaddress._BaseNetwork]:
|
|
result = []
|
|
for item in (value or "").split(","):
|
|
item = item.strip()
|
|
if not item:
|
|
continue
|
|
result.append(ipaddress.ip_network(item, strict=False))
|
|
return result
|
|
|
|
|
|
def _ip(value: Any) -> ipaddress._BaseAddress | None:
|
|
try:
|
|
return ipaddress.ip_address(str(value))
|
|
except ValueError:
|
|
return None
|