#!/usr/bin/python3

import re
import requests
import subprocess
from argparse import ArgumentParser
from pathlib import Path


CONF_DIR = Path("/etc/nftables.conf.d")
SET_HEADER = re.compile(r"^\s*add\s+set\s+(\S+)\s+(\S+)\s+(\S+)\s*\{", re.MULTILINE)
URL_COMMENT = re.compile(r"#\s*url\s+(\S+)", re.IGNORECASE)


def extract_sets(path):
    """Yield (family, table, set_name, url_or_None) for every add set block in a file."""
    text = path.read_text()
    matches = list(SET_HEADER.finditer(text))
    for i, m in enumerate(matches):
        family, table, name = m.group(1), m.group(2), m.group(3)
        # Block boundaries are determined by match positions: the body of each set
        # is the text between the end of its header match and the start of the next
        # header match (or end of file). This avoids brace-counting entirely because
        # a new 'add set ... {' header unambiguously starts a new block regardless
        # of any braces that may appear in comments within the block.
        block_end = matches[i + 1].start() if i + 1 < len(matches) else len(text)
        block = text[m.end():block_end]
        url_match = URL_COMMENT.search(block)
        url = url_match.group(1) if url_match else None
        yield family, table, name, url


def handle_set(family, table, set_name, url):
    ref = f"{family} {table} {set_name}"
    elements = []
    if url:
        response = requests.get(url, timeout=30)
        response.raise_for_status()
        for line in response.text.splitlines():
            line = line.strip()
            if line and not line.startswith("#"):
                elements.append(line)

    script = f"flush set {ref}"
    if elements:
        script += f"\nadd element {ref} {{{','.join(elements)}}}"
    subprocess.run(["nft", "-f", "-"], input=script, text=True, check=True)


def all_sets():
    """Return dict of set_name -> (family, table, url) for all *.set files in CONF_DIR."""
    result = {}
    for path in CONF_DIR.glob("*.set"):
        for family, table, name, url in extract_sets(path):
            result[name] = (family, table, url)
    return result


def parse_args():
    parser = ArgumentParser(description="NFT set element reloader")
    parser.add_argument("sets", nargs="*", metavar="SET", help="NFT set name")
    return parser.parse_args()


if __name__ == "__main__":
    args = parse_args()
    sets = all_sets()
    targets = args.sets or sorted(sets)
    for set_name in targets:
        if set_name not in sets:
            raise SystemExit(f"Unknown set: {set_name}")
        family, table, url = sets[set_name]
        handle_set(family, table, set_name, url)
