#!/usr/bin/env python2.7
# -*- coding: utf-8 -*-
#
import os
import subprocess
import sys
import random

CONFIG_FILE = "/var/efw/geoipforward/config"
CHAIN       = "GEOIPFORWARD"


def check_and_create_chain():
    with open(os.devnull, 'wb') as devnull:
        ret = subprocess.call(["iptables", "-L", CHAIN], stdout=devnull, stderr=devnull)
    if ret != 0:
        subprocess.call(["iptables", "-N", CHAIN])

    with open(os.devnull, 'wb') as devnull:
        check_forward = subprocess.call(["iptables", "-C", "FORWARD", "-j", CHAIN],
                                        stdout=devnull, stderr=devnull)
    if check_forward != 0:
        subprocess.call(["iptables", "-I", "FORWARD", "1", "-j", CHAIN])

def flush_chain():
    subprocess.call(["iptables", "-F", CHAIN])

def insert_private_network_bypass():
    for subnet in ["192.168.0.0/16", "10.0.0.0/8", "172.16.0.0/12"]:
        subprocess.call(["iptables", "-A", CHAIN, "-s", subnet, "-j", "RETURN"])

def ipset_exists(name):
    with open(os.devnull, 'wb') as devnull:
        ret = subprocess.call(["ipset", "list", name], stdout=devnull, stderr=devnull)
    return ret == 0

def resolve_ipset_name(raw):
    if not raw:
        return None
    raw = raw.strip().lower()
    for candidate in [
        raw.replace("geoip_", "geo_"),
        "ipset_" + raw,
        raw,
        raw.replace("geoip_", ""),
    ]:
        if ipset_exists(candidate):
            return candidate
    return None

def is_geoip(val):
    return bool(val) and val.lower().startswith("geoip_")

def _run(cmd):
    subprocess.call(cmd)

def _add_log_rule(opts, action):
    prefix = "{}:{}".format(CHAIN, action)[:29]
    _run(["iptables", "-A", CHAIN] + opts + ["-j", "NFLOG", "--nflog-prefix", prefix])

def _emit(opts, action, logflag):
    """Emite LOG (se ativo) + regra de ação."""
    if logflag == "on":
        _add_log_rule(opts, action)
    _run(["iptables", "-A", CHAIN] + opts + ["-j", action])

def _emit_with_implicit_reject(match_opts, base_opts, action, logflag):
    _emit(match_opts, action, logflag)

    if action in ("ACCEPT", "ALLOW"):
        if logflag == "on":
            _add_log_rule(base_opts, "REJECT")
        _run(["iptables", "-A", CHAIN] + base_opts + ["-j", "REJECT"])

def build_pp(proto, port):
    opts = []
    if proto:
        opts.extend(["-p", proto])
    if port:
        opts.extend(["--dport", port])
    return opts

def build_geo_opt(ipset_name, direction, negate):
    if negate:
        return ["-m", "set", "!", "--match-set", ipset_name, direction]
    return ["-m", "set", "--match-set", ipset_name, direction]

def apply_rule(proto, src, dst, port, action, logflag, negateflag):
    src_list  = [s.strip() for s in src.split('&')] if src else [""]
    dst_list  = [d.strip() for d in dst.split('&')] if dst else [""]
    port_list = [p.strip() for p in port.split('&')] if port else [""]

    action = action.upper()
    
    negate = (negateflag == "on")

    for src_val in src_list:
        for dst_val in dst_list:
            for p in port_list:
                pp = build_pp(proto, p)

                if is_geoip(src_val):
                    ipset_name = resolve_ipset_name(src_val)
                    if not ipset_name:
                        print >> sys.stderr, \
                            "Warning: ipset not found for '{}'. Rule ignored.".format(src_val)
                        continue

                    geo_opt  = build_geo_opt(ipset_name, "src", negate)
                    dst_opt  = ["-d", dst_val] if dst_val else []

                    match_opts = geo_opt + dst_opt + pp
                    base_opts  = dst_opt + pp

                    _emit_with_implicit_reject(match_opts, base_opts, action, logflag)

                elif is_geoip(dst_val):
                    ipset_name = resolve_ipset_name(dst_val)
                    if not ipset_name:
                        print >> sys.stderr, \
                            "Warning: ipset not found for '{}'. Rule ignored.".format(dst_val)
                        continue

                    geo_opt  = build_geo_opt(ipset_name, "dst", negate)
                    src_opt  = ["-s", src_val] if src_val else []

                    match_opts = src_opt + geo_opt + pp
                    base_opts  = src_opt + pp

                    _emit_with_implicit_reject(match_opts, base_opts, action, logflag)

                else:
                    src_opt = ["-s", src_val] if src_val else []
                    dst_opt = ["-d", dst_val] if dst_val else []
                    _emit(src_opt + dst_opt + pp, action, logflag)


def check_random_set_unhealthy():
    if not os.path.exists(CONFIG_FILE):
        return True

    active_geoips = []
    with open(CONFIG_FILE, 'r') as f:
        for raw_line in f:
            line = raw_line.strip()
            if not line or line.startswith('#'):
                continue
            parts = line.split(',')
            if len(parts) < 9:
                continue
            
            status = parts[0].strip()
            src    = parts[2].strip()
            dst    = parts[3].strip()
            
            if status == "on":
                for check_val in [src, dst]:
                    if check_val:
                        for sub_val in check_val.split('&'):
                            sub_val = sub_val.strip()
                            if sub_val.lower().startswith("geoip_"):
                                active_geoips.append(sub_val)

    if not active_geoips:

        return False

    sampled_geo = random.choice(active_geoips)
    ipset_name = resolve_ipset_name(sampled_geo)

    if not ipset_name:
        return True

    try:

        out = subprocess.check_output("ipset list %s | wc -l" % ipset_name, shell=True)
        if int(out.strip()) < 10:
            return True
    except Exception:
        return True

    return False


def main():
    if check_random_set_unhealthy():
        ingestor_cmd = ['/usr/local/bin/geoip-set-ingestor']
    else:
        ingestor_cmd = ['/usr/local/bin/geoip-set-ingestor', '--boot']

    with open(os.devnull, 'wb') as devnull:
        if not sys.stdout.isatty():
            ret = subprocess.call(ingestor_cmd, stdout=devnull, stderr=devnull)
        else:
            ret = subprocess.call(ingestor_cmd)

    if ret != 0:
        sys.exit(1)

    if not os.path.exists(CONFIG_FILE):
        if sys.stdout.isatty():
            print >> sys.stderr, "Error: configuration file not found: {}".format(CONFIG_FILE)
        sys.exit(1)

    check_and_create_chain()
    flush_chain()
    insert_private_network_bypass()

    with open(CONFIG_FILE, 'r') as f:
        for raw_line in f:
            line = raw_line.strip()
            if not line or line.startswith('#'):
                continue

            parts = line.split(',')
            if len(parts) < 9:
                continue

            status  = parts[0].strip()
            proto   = parts[1].strip()
            src     = parts[2].strip()
            dst     = parts[3].strip()
            port    = parts[4].strip()
            action  = parts[5].strip()
            logflag = parts[7].strip()
            negate  = parts[8].strip()

            if status != "on":
                continue

            proto_norm = proto.lower()
            if proto_norm in ("any", ""):
                proto_norm = ""

            if proto_norm == "tcp&udp":
                apply_rule("tcp", src, dst, port, action, logflag, negate)
                apply_rule("udp", src, dst, port, action, logflag, negate)
            else:
                apply_rule(proto_norm, src, dst, port, action, logflag, negate)

    if sys.stdout.isatty():
        print "GEOIPFORWARD firewall rules applied successfully."

if __name__ == "__main__":
    main()