#!/usr/bin/env python2.7
# -*- coding: utf-8 -*-

import os
import subprocess
import sys
import random

CONFIG_FILE = "/var/efw/geoipinput/config"
CHAIN       = "GEOIPINPUT"


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_input = subprocess.call(["iptables", "-C", "INPUT", "-j", CHAIN],
                                      stdout=devnull, stderr=devnull)
    if check_input != 0:
        subprocess.call(["iptables", "-A", "INPUT", "-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(src_raw):
    if not src_raw:
        return None

    src_raw = src_raw.strip().lower()

    for candidate in [
        src_raw.replace("geoip_", "geo_"), 
        "ipset_" + src_raw,
        src_raw,
        src_raw.replace("geoip_", ""),
    ]:
        if ipset_exists(candidate):
            return candidate

    return None


def build_base_opts(interface, proto, port):
    opts = []
    if interface and interface.lower() != "any":
        opts.extend(["-i", interface])
    if proto:
        opts.extend(["-p", proto])
    if port:
        opts.extend(["--dport", port])
    return opts


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

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

def _add_accept(base_opts, set_opt, logflag):
    if logflag == "on":
        _add_log_rule(base_opts, set_opt, "ACCEPT")
    _run(["iptables", "-A", CHAIN] + base_opts + set_opt + ["-j", "ACCEPT"])

def _add_block(base_opts, set_opt, action, logflag):
    iptables_action = "ACCEPT" if action == "ALLOW" else action
    if logflag == "on":
        _add_log_rule(base_opts, set_opt, iptables_action)
    _run(["iptables", "-A", CHAIN] + base_opts + set_opt + ["-j", iptables_action])

def apply_rule(proto, interface, src_geoip, port, action, logflag, negateflag):

    src_list  = [s.strip() for s in src_geoip.split('&')] if src_geoip else [""]
    port_list = [p.strip() for p in port.split('&')]      if port      else [""]

    action = action.upper()
    if action == "ALLOW":
        action = "ACCEPT"

    for src in src_list:
        ipset_name = None
        is_ip = False
        
        if src:
            if src.lower().startswith("geoip_"):
                ipset_name = resolve_ipset_name(src)
                if not ipset_name:
                    print >> sys.stderr, \
                        "Warning: ipset not found for '{}'. Rule ignored.".format(src)
                    continue
            else:
                is_ip = True

        for p in port_list:
            base_opts = build_base_opts(interface, proto, p)

            if action == "ACCEPT":
                if ipset_name:
                    set_accept = ["-m", "set", "--match-set", ipset_name, "src"]
                elif is_ip:
                    set_accept = ["-s", src]
                else:
                    set_accept = []

                _add_accept(base_opts, set_accept, logflag)

                if not is_ip:
                    _add_block(base_opts, [], "REJECT", logflag)

            else:
                if ipset_name:
                    if negateflag == "on":
                        set_opt = ["-m", "set", "!", "--match-set", ipset_name, "src"]
                    else:
                        set_opt = ["-m", "set", "--match-set", ipset_name, "src"]
                elif is_ip:
                    if negateflag == "on":
                        set_opt = ["!", "-s", src]
                    else:
                        set_opt = ["-s", src]
                else:
                    set_opt = []

                _add_block(base_opts, set_opt, action, logflag)


def check_random_set_unhealthy():
    """
    Sorteia um GeoIP ativo no arquivo de configuração e valida a saúde do ipset.
    Retorna True se estiver descarregado (menos de 10 linhas via wc -l ou inexistente).
    Retorna False se estiver saudável (populado).
    """
    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[3].strip()
            
            if status == "on" and src:
                for sub_src in src.split('&'):
                    sub_src = sub_src.strip()
                    if sub_src.lower().startswith("geoip_"):
                        active_geoips.append(sub_src)

    if not active_geoips:
        # Sem regras de GeoIP ativas para testar, assume saudável para evitar loop
        return False

    # Escolha aleatória baseado na sua estratégia
    sampled_geo = random.choice(active_geoips)
    ipset_name = resolve_ipset_name(sampled_geo)

    if not ipset_name:
        return True

    try:
        # Lógica idêntica do seu segundo script (wc -l < 10)
        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()   
            interface = parts[2].strip()   
            src       = 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", interface, src, port, action, logflag, negate)
                apply_rule("udp", interface, src, port, action, logflag, negate)
            else:
                apply_rule(proto_norm, interface, src, port, action, logflag, negate)

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

if __name__ == "__main__":
    main()