Your IP : 216.73.217.62


Current Path : /bin/
Upload File :
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()