| Current Path : /bin/ |
| Current File : //bin/modsec-live |
#!/usr/bin/imh-python3.13
"""
Manage dynamic ModSec rules
Usage:
modsec-live ip add TARGET CIDR [--expiry HOURS]
modsec-live ip remove TARGET CIDR
modsec-live ip reset TARGET
modsec-live ip list
modsec-live ip generate [--output-dir PATH]
"""
import argparse
import contextlib
import hashlib
import ipaddress
import os
import re
import socket
import sqlite3
import sys
import time
from pathlib import Path
DEFAULT_DB = "/etc/apache2/conf.d/imh-modsec/dynamic_blocks.sqlite"
DEFAULT_OUTPUT_DIR = "/etc/apache2/conf.d/imh-modsec"
DEFAULT_EXPIRY_HOURS = 24
LOCK_TIMEOUT = 2.0
INLINE_THRESHOLD = 100
TARGET_NAME_RE = r"^[a-z0-9._-]+$"
class LockError(Exception):
pass
@contextlib.contextmanager
def lock(name: str = "modsec_live"):
"""Abstract UNIX socket lock with retry.
Waits up to LOCK_TIMEOUT seconds for the lock to free.
"""
key = hashlib.sha256(name.encode("utf-8")).hexdigest()
lock_socket = socket.socket(
socket.AF_UNIX, socket.SOCK_DGRAM
)
deadline = time.monotonic() + LOCK_TIMEOUT
while True:
try:
lock_socket.bind(f"\0{key}")
break
except OSError:
if time.monotonic() >= deadline:
lock_socket.close()
raise LockError(
f"Could not acquire lock within {LOCK_TIMEOUT}s"
)
time.sleep(0.1)
try:
yield
finally:
lock_socket.close()
# -- Lua generation --
VALID_STATUSES = {403, 429}
def cidr_to_octets(cidr: str, status: int) -> dict:
"""Convert CIDR to pre-computed octet ranges for Lua matching."""
net = ipaddress.ip_network(cidr, strict=False)
start = list(net.network_address.packed)
end = list(net.broadcast_address.packed)
octets = [(start[i], end[i]) for i in range(4)]
return {"name": str(net), "status": status, "octets": octets}
def _render_entry(entry: dict) -> str:
octs = ", ".join(
f"{{{lo},{hi}}}" for lo, hi in entry["octets"]
)
status_part = ""
if entry["status"] != 429:
status_part = f'status = {entry["status"]}, '
return (
f' {{ name = "{entry["name"]}", '
f"{status_part}"
f"octets = {{ {octs} }} }},"
)
def render_inline(blocks: dict[str, list[dict]]) -> str:
"""Render all blocks into a single index file with inline data."""
lines = [
"-- Generated by modsec-live ip generate",
"-- Do not edit manually.",
"return {",
]
for key in sorted(blocks):
lines.append(f' ["{key}"] = {{')
for entry in blocks[key]:
lines.append(_render_entry(entry))
lines.append(" },")
lines.append("}")
return "\n".join(lines) + "\n"
def render_index(targets: list[str]) -> str:
"""Render an index file pointing to per-target files."""
lines = [
"-- Generated by modsec-live ip generate",
"-- Do not edit manually.",
"return {",
]
for key in sorted(targets):
lines.append(f' ["{key}"] = "{key}",')
lines.append("}")
return "\n".join(lines) + "\n"
def render_target(entries: list[dict]) -> str:
"""Render a single target's block list."""
lines = [
"-- Generated by modsec-live ip generate",
"-- Do not edit manually.",
"return {",
]
for entry in entries:
lines.append(_render_entry(entry))
lines.append("}")
return "\n".join(lines) + "\n"
# -- Database --
class StateDB:
def __init__(self, path: str):
self.conn = sqlite3.connect(path)
self.conn.execute("""
CREATE TABLE IF NOT EXISTS ip_blocks (
target TEXT NOT NULL,
cidr TEXT NOT NULL,
status INTEGER NOT NULL DEFAULT 429,
created_at REAL NOT NULL,
expires_at REAL NOT NULL,
PRIMARY KEY (target, cidr)
)
""")
self.conn.commit()
def close(self):
self.conn.close()
def purge_expired(self):
self.conn.execute(
"DELETE FROM ip_blocks WHERE expires_at <= ?",
(time.time(),),
)
self.conn.commit()
def ip_add(
self, target: str, cidr: str, status: int,
expires_at: float,
) -> None:
self.conn.execute(
"INSERT OR REPLACE INTO ip_blocks "
"(target, cidr, status, created_at, expires_at) "
"VALUES (?, ?, ?, ?, ?)",
(target, cidr, status, time.time(), expires_at),
)
self.conn.commit()
def ip_remove(self, target: str, cidr: str) -> int:
cur = self.conn.execute(
"DELETE FROM ip_blocks "
"WHERE target = ? AND cidr = ?",
(target, cidr),
)
self.conn.commit()
return cur.rowcount
def ip_reset(self, target: str) -> int:
cur = self.conn.execute(
"DELETE FROM ip_blocks WHERE target = ?",
(target,),
)
self.conn.commit()
return cur.rowcount
def active_ip_blocks(self) -> list[tuple[str, str, int]]:
"""Return (target, cidr, status) for non-expired blocks."""
self.purge_expired()
return self.conn.execute(
"SELECT target, cidr, status FROM ip_blocks "
"WHERE expires_at > ? ORDER BY target, cidr",
(time.time(),),
).fetchall()
def all_ip_blocks(
self,
) -> list[tuple[str, str, int, float, float]]:
"""Return (target, cidr, status, created_at, expires_at)."""
return self.conn.execute(
"SELECT target, cidr, status, created_at, expires_at "
"FROM ip_blocks ORDER BY target, cidr"
).fetchall()
# -- IP subcommands --
def ip_generate(db_path: str, output_dir: str):
db = StateDB(db_path)
rows = db.active_ip_blocks()
db.close()
blocks: dict[str, list[dict]] = {}
for target, cidr, row_status in rows:
try:
entry = cidr_to_octets(cidr, row_status)
except ValueError as e:
print(
f"Warning: invalid CIDR '{cidr}' for "
f"'{target}': {e}",
file=sys.stderr,
)
continue
blocks.setdefault(target, []).append(entry)
out = Path(output_dir)
total = sum(len(v) for v in blocks.values())
index_path = out / "dynamic_ip_blocks.lua"
if total <= INLINE_THRESHOLD:
# Clean up per-target files first (safe: old index still works)
ip_blocks_dir = out / "ip_blocks"
if ip_blocks_dir.is_dir():
for f in ip_blocks_dir.iterdir():
f.unlink()
# Index last: atomically switch to inline mode
_write_if_changed(index_path, render_inline(blocks))
else:
ip_blocks_dir = out / "ip_blocks"
ip_blocks_dir.mkdir(exist_ok=True)
# Write per-target files first (safe: old index still works)
active_files = set()
for target, entries in blocks.items():
target_path = ip_blocks_dir / f"{target}.lua"
active_files.add(target_path.name)
_write_if_changed(
target_path, render_target(entries)
)
# Index second: atomically switch to split mode
_write_if_changed(
index_path, render_index(list(blocks.keys()))
)
# Stale cleanup last (safe: index no longer references them)
for f in ip_blocks_dir.iterdir():
if f.name not in active_files:
f.unlink()
print(
f"Generated {index_path}: "
f"{len(blocks)} targets, {total} ranges"
f"{' (split)' if total > INLINE_THRESHOLD else ''}"
)
def _write_if_changed(path: Path, content: str):
if path.exists() and path.read_text() == content:
return
# Atomic write: temp file + rename avoids partial reads by Lua
tmp = path.with_suffix(".tmp")
tmp.write_text(content)
tmp.chmod(0o644)
os.replace(tmp, path)
def ip_add(args):
try:
net = ipaddress.ip_network(args.cidr, strict=False)
except ValueError as e:
print(f"Invalid CIDR '{args.cidr}': {e}", file=sys.stderr)
sys.exit(1)
cidr = str(net)
target = args.target.lower()
_validate_target(target)
if args.expiry <= 0:
print("Expiry must be positive", file=sys.stderr)
sys.exit(1)
if args.expiry > 24 * 365:
print("Expiry seems unreasonably long", file=sys.stderr)
sys.exit(1)
expires = time.time() + (args.expiry * 60 * 60)
if args.status not in VALID_STATUSES:
print(
f"Invalid status {args.status}, "
f"must be one of {sorted(VALID_STATUSES)}",
file=sys.stderr,
)
sys.exit(1)
db = StateDB(args.db)
db.ip_add(target, cidr, args.status, expires)
db.close()
print(
f"Added {cidr} for {target} "
f"(status {args.status}, expires in {args.expiry}h)"
)
ip_generate(args.db, args.output_dir)
def _validate_target(name: str):
if not re.match(TARGET_NAME_RE, name):
print(
f"Invalid target name '{name}': "
f"must match {TARGET_NAME_RE}",
file=sys.stderr,
)
sys.exit(1)
def ip_remove(args):
try:
net = ipaddress.ip_network(args.cidr, strict=False)
except ValueError as e:
print(f"Invalid CIDR '{args.cidr}': {e}", file=sys.stderr)
sys.exit(1)
cidr = str(net)
target = args.target.lower()
_validate_target(target)
db = StateDB(args.db)
removed = db.ip_remove(target, cidr)
db.close()
if removed:
print(f"Removed {cidr} for {target}")
ip_generate(args.db, args.output_dir)
else:
print(f"No matching rule found for {target} {cidr}")
def ip_reset(args):
target = args.target.lower()
_validate_target(target)
db = StateDB(args.db)
count = db.ip_reset(target)
db.close()
print(f"Removed {count} rule(s) for {target}")
if count:
ip_generate(args.db, args.output_dir)
def ip_list(args):
db = StateDB(args.db)
rows = db.all_ip_blocks()
db.close()
if not rows:
print("No rules")
return
now = time.time()
for target, cidr, row_status, _, expires in rows:
remaining = expires - now
if remaining <= 0:
expiry_str = "EXPIRED"
else:
hours = remaining / (60 * 60)
if hours >= 1:
expiry_str = f"{hours:.1f}h remaining"
else:
expiry_str = f"{remaining / 60:.0f}m remaining"
print(
f" {target:30s} {cidr:20s} "
f"{row_status} {expiry_str}"
)
# -- CLI --
def main():
if os.geteuid() != 0:
print("Error: must be run as root", file=sys.stderr)
sys.exit(1)
parser = argparse.ArgumentParser(
description="Manage dynamic ModSec rules"
)
parser.add_argument(
"--db",
default=DEFAULT_DB,
help=f"Path to sqlite database (default: {DEFAULT_DB})",
)
sub = parser.add_subparsers(
dest="group", title="rule types", metavar="COMMAND"
)
sub.required = False
# -- ip subcommand group --
ip_parser = sub.add_parser("ip", help="IP block rules")
ip_parser.add_argument(
"--output-dir",
default=DEFAULT_OUTPUT_DIR,
help=f"Output directory (default: {DEFAULT_OUTPUT_DIR})",
)
ip_sub = ip_parser.add_subparsers(
dest="command", required=True
)
ip_sub.add_parser(
"generate", help="Generate dynamic_ip_blocks.lua from DB"
)
p_add = ip_sub.add_parser("add", help="Add a block rule")
p_add.add_argument(
"target", help="Domain name or linux username"
)
p_add.add_argument(
"cidr", help="CIDR range (e.g. 10.0.0.0/8)"
)
p_add.add_argument(
"--expiry",
type=float,
default=DEFAULT_EXPIRY_HOURS,
help=f"Hours until expiry (default: {DEFAULT_EXPIRY_HOURS})",
)
p_add.add_argument(
"--status",
type=int,
default=429,
help="HTTP status code (default: 429)",
)
p_rm = ip_sub.add_parser("remove", help="Remove a block rule")
p_rm.add_argument(
"target", help="Domain name or linux username"
)
p_rm.add_argument("cidr", help="CIDR range to remove")
p_reset = ip_sub.add_parser(
"reset", help="Remove all rules for a target"
)
p_reset.add_argument(
"target", help="Domain name or linux username"
)
ip_sub.add_parser("list", help="List all rules")
args = parser.parse_args()
if not args.group:
parser.print_help()
sys.exit(1)
with lock():
match args.group, args.command:
case "ip", "generate":
ip_generate(args.db, args.output_dir)
case "ip", "add":
ip_add(args)
case "ip", "remove":
ip_remove(args)
case "ip", "reset":
ip_reset(args)
case "ip", "list":
ip_list(args)
if __name__ == "__main__":
main()