from concurrent import futures
import grpc

from firewall_filter_pb2 import FirewallFilterResponse, RefreshedAdresses  # type: ignore
import firewall_filter_pb2_grpc

import requests
from requests.exceptions import RequestException


class FirewallFilterService(firewall_filter_pb2_grpc.FirewallFilterServiceServicer):
    def Insert(self, request, context):
        try:
            firewall_rule = retrieve_current_firewall_rule(
                request.api_host,
                request.session_token,
                request.session_cookie,
                request.domain_id,
                request.gateway_policy_id,
            )
        except RequestException as e:
            return FirewallFilterResponse(
                status="ERROR",
                message=f"No fue posible recuperar regla de firewall actual. {e}",
            )

        for ip in request.ip_addresses:
            if ip in firewall_rule["rules"][0]["source_groups"]:
                continue  # Ya no se requiere añadir ip o segmento
            firewall_rule["rules"][0]["source_groups"].append(ip)

        firewall_rule["rules"][0]["_revision"] += 1

        try:
            r = requests.patch(
                f"{request.api_host}/policy/api/v1/infra/domains/{request.domain_id}/gateway-policies/{request.gateway_policy_id}",
                timeout=120,
                # proxies=dict(
                #     http="socks5://127.0.0.1:9000", https="socks5://127.0.0.1:9000"
                # ),
                json=firewall_rule,
                verify=False,
                headers={
                    "X-XSRF-TOKEN": request.session_token,
                    "Cookie": request.session_cookie,
                },
            )
            r.raise_for_status()
        except RequestException as e:
            return FirewallFilterResponse(
                status="ERROR",
                message=f"No fue posible actualizar regla de firewall con direcciones {firewall_rule['rules'][0]['source_groups']}. {e}",
            )

        return FirewallFilterResponse(
            status="OK",
            message=f"Listado de ips/segmentos de origen actualizado: {firewall_rule['rules'][0]['source_groups']}",
        )

    def Remove(self, request, context):
        try:
            firewall_rule = retrieve_current_firewall_rule(
                request.api_host,
                request.session_token,
                request.session_cookie,
                request.domain_id,
                request.gateway_policy_id,
            )
        except RequestException as e:
            return FirewallFilterResponse(
                status="ERROR",
                message=f"No fue posible recuperar regla de firewall actual. {e}",
            )

        for ip in request.ip_addresses:
            if ip not in firewall_rule["rules"][0]["source_groups"]:
                continue  # Ya no se requiere añadir ip o segmento
            firewall_rule["rules"][0]["source_groups"].remove(ip)

        try:
            r = requests.patch(
                f"{request.api_host}/policy/api/v1/infra/domains/{request.domain_id}/gateway-policies/{request.gateway_policy_id}",
                timeout=120,
                # proxies=dict(
                #     http="socks5://127.0.0.1:9000", https="socks5://127.0.0.1:9000"
                # ),
                json=firewall_rule,
                verify=False,
                headers={
                    "X-XSRF-TOKEN": request.session_token,
                    "Cookie": request.session_cookie,
                },
            )
            r.raise_for_status()
        except RequestException as e:
            return FirewallFilterResponse(
                status="ERROR",
                message=f"No fue posible actualizar regla de firewall con direcciones {firewall_rule['rules'][0]['source_groups']}. {e}",
            )

        return FirewallFilterResponse(
            status="OK",
            message=f"Listado de ips/segmentos de origen actualizado: {firewall_rule['rules'][0]['source_groups']}",
        )

    def Refresh(self, request, context):
        try:
            firewall_rule = retrieve_current_firewall_rule(
                request.api_host,
                request.session_token,
                request.session_cookie,
                request.domain_id,
                request.gateway_policy_id,
            )
        except RequestException as e:
            return FirewallFilterResponse(
                status="ERROR",
                message=f"No fue posible recuperar regla de firewall actual. {e}",
            )

        missing = []

        for ip in request.ip_addresses:
            if ip in firewall_rule["rules"][0]["source_groups"]:
                continue  # Ya no se requiere añadir ip o segmento
            firewall_rule["rules"][0]["source_groups"].remove(ip)
            missing.append(ip)

        return RefreshedAdresses(
            existing_vcd=firewall_rule["rules"][0]["source_groups"], missing_vcd=missing
        )


def retrieve_current_firewall_rule(
    api_host: str, session_token, session_cookie, domain_id, gateway_policy_id
):
    r = requests.request(
        method="GET",
        url=f"{api_host}/policy/api/v1/infra/domains/{domain_id}/gateway-policies/{gateway_policy_id}",
        # url=f"{api_host}/policy/api/v1/infra/domains/default/gateway-policies/tcp-blocking",
        timeout=120,
        # proxies=dict(http="socks5://127.0.0.1:9000", https="socks5://127.0.0.1:9000"),
        verify=False,
        headers={"X-XSRF-TOKEN": session_token, "Cookie": session_cookie},
    )
    r.raise_for_status()

    return r.json()


def serve():
    server = grpc.server(futures.ThreadPoolExecutor(max_workers=10))

    firewall_filter_pb2_grpc.add_FirewallFilterServiceServicer_to_server(
        FirewallFilterService(), server
    )

    with open("server.key", "rb") as fp:
        server_key = fp.read()
    with open("server.pem", "rb") as fp:
        server_cert = fp.read()

    creds = grpc.ssl_server_credentials([(server_key, server_cert)])
    server.add_insecure_port('[::]:50051')
    server.add_secure_port("[::]:8443", creds)
    server.start()
    server.wait_for_termination()


if __name__ == "__main__":
    serve()
