#!/usr/bin/env python3
#-*- encoding: utf-8
#
# Author : Alexander Vogel (alexander.vogel.2305@gmail.com)
# Date   : 2026-07-16
# License: GNU General Public License v2
#
# Libexec: VMware Avi Load Balancer


import argparse
import requests
import sys
import urllib3

from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from requests import Response, Session
from typing import Dict, Any, Optional

urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning)


_MISSING = object()


def parse_arguments(argv):
    parser = argparse.ArgumentParser(description=__doc__)
    
    parser.add_argument("--server", type=str, required=True)
    parser.add_argument("--username", type=str, required=True)
    parser.add_argument("--password", type=str, required=True)
    parser.add_argument("--tenant", type=str, required=False)
    parser.add_argument("--version", type=str, required=True, help="Avi API schema version, for example 30.2.7")
    #parser.add_argument("--verify_ssl", type=bool, required=False)
    parser.add_argument("--sections", type=str, required=False, default="all")

    args = parser.parse_args(argv)
    
    return args


def dig(obj, *path, default=None, expected_type=None, remove_chars=None, ndigits=None):
    """
    Sicherer Zugriff auf verschachtelte Dicts und Listen.

    Optional kann der gefundene Wert in einen Zieltyp konvertiert und bei
    Float-Werten auf eine bestimmte Anzahl Nachkommastellen gerundet werden.

    Beispiele:
        dig(data, "value")
        dig(data, "value", expected_type=float)
        dig(data, "value", expected_type=float, remove_chars=["%"])
        dig(data, "value", expected_type=float, ndigits=3)
        dig(
            data,
            "value",
            expected_type=float,
            remove_chars=["%"],
            ndigits=2,
        )

    Args:
        obj:
            Das Startobjekt.

        *path:
            Keys oder Listenindizes.

        default:
            Rückgabewert bei fehlgeschlagenem Zugriff, Konvertierung
            oder Rundung.

        expected_type:
            Optionaler Zieltyp, z. B. int, float oder str.

        remove_chars:
            Optionale Liste von Zeichen oder Strings, die vor der
            Konvertierung entfernt werden.

        ndigits:
            Anzahl der Nachkommastellen für float-Werte.
            Wenn None, wird nicht gerundet.

    Returns:
        Der gefundene Wert, gegebenenfalls bereinigt, konvertiert und
        gerundet, oder `default`.
    """
    current = obj

    for key in path:
        try:
            if isinstance(current, Mapping):
                current = current.get(key, _MISSING)

            elif isinstance(current, Sequence) and not isinstance(
                current, (str, bytes)
            ):
                current = current[key]

            else:
                return default

        except (KeyError, IndexError, TypeError):
            return default

        if current is _MISSING:
            return default

    try:
        if expected_type is not None:
            if isinstance(current, str):
                current = current.strip()

                if remove_chars:
                    for char in remove_chars:
                        current = current.replace(char, "")

                current = current.strip()

            if not isinstance(current, expected_type):
                current = expected_type(current)

        if ndigits is not None and isinstance(current, float):
            current = round(current, ndigits)

    except (ValueError, TypeError):
        return default

    return current

 
class VmwareAviLoadBalancerAgent():
    def __init__(self, server: str, username: str, password: str, version: str, tenant: Optional[str] = None, verify_ssl: bool = False, timeout: int = 30) -> None:
        
        self.controller = f"https://{server}"
        self.username = username
        self.password = password
        self.tenant = tenant
        self.version = version
        self.verify_ssl = verify_ssl
        self.timeout = timeout

        self.session: Session = requests.Session()
        self.session.verify = verify_ssl


    @property
    def base_url(self) -> str:
        return f"{self.controller}"

    @property
    def api_base(self) -> str:
        return f"{self.base_url}/api"

    def _default_headers(self, include_csrf: bool = False) -> Dict[str, str]:
        headers: Dict[str, str] = {
            "Accept": "application/json",
            "Content-Type": "application/json",
            "X-Avi-Version": self.version,
        }

        if self.tenant:
            headers["X-Avi-Tenant"] = self.tenant

        if include_csrf:
            csrf = self.session.cookies.get("csrftoken")
            if csrf:
                headers["X-CSRFToken"] = csrf
                headers["Referer"] = self.base_url

        return headers


    def _raise_for_status(self, response: Response) -> None:
        if response.ok:
            return

        try:
            error_payload = response.json()
        except Exception:
            error_payload = response.text

        raise Exception(
            f"HTTP {response.status_code} bei {response.request.method} "
            f"{response.request.url}: {error_payload}"
        )

    def login(self) -> None:
        url = f"{self.base_url}/login"
        payload = {
            "username": self.username,
            "password": self.password,
        }

        response = self.session.post(
            url,
            data=payload,
            timeout=self.timeout,
            verify=self.verify_ssl,
        )
        self._raise_for_status(response)

        if "sessionid" not in self.session.cookies:
            raise Exception("Login erfolgreich, aber kein 'sessionid'-Cookie erhalten.")

    def logout(self) -> None:
        url = f"{self.base_url}/logout"
        response = self.session.post(
            url,
            headers=self._default_headers(include_csrf=True),
            timeout=self.timeout,
            verify=self.verify_ssl,
        )
        self._raise_for_status(response)

    def get(self, path: str, params: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
        url = f"{self.api_base}/{path.lstrip('/')}"
        response = self.session.get(
            url,
            headers=self._default_headers(include_csrf=False),
            params=params,
            timeout=self.timeout,
            verify=self.verify_ssl,
        )
        self._raise_for_status(response)
        return response.json()

    def post(self, path: str, data: Optional[Dict[str, Any]] = None, params: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
        url = f"{self.api_base}/{path.lstrip('/')}"

        response = self.session.post(
            url,
            headers=self._default_headers(include_csrf=True),
            params=params,
            json=data,
            timeout=self.timeout,
            verify=self.verify_ssl,
        )

        self._raise_for_status(response)
        return response.json()

    # Alerts
    def get_alerts(self):
        return self.get(f"/alert")

    # Certificates
    def get_certificates(self):
        return self.get(f"/sslkeyandcertificate")

    # Cluster and Nodes
    def get_cluster(self):
        return self.get(f"/cluster/runtime")
    
    # Cloud
    def get_cloudinventory(self):
        return self.get(f"/cloud-inventory")

    # pools
    def get_pool_inventory(self):
        params = {
            "step": 1, # Necessary to retrieve current data; without this parameter, average values ​​from the last 6 hours would be returned
        }
        return self.get(f"/pool-inventory", params=params)

    # ServiceEngine
    def get_serviceengine_cpu(self, uuid: str):
        return self.get(f"/serviceengine/{uuid}/cpu")

    def get_serviceengine_seagent(self, uuid: str):
        return self.get(f"/serviceengine/{uuid}/seagent")

    def get_serviceengines_inventory(self, metrics=[], limit=1):
        params = {
            "metric_id": ",".join(metrics),
            "limit": limit,
            "step": 1, # Necessary to retrieve current data; without this parameter, average values ​​from the last 6 hours would be returned
        }

        return self.get(f"/serviceengine-inventory", params=params)
    
    # VirtualService
    def get_virtualservices_inventory(self):
        params = {
            "step": 1, # Necessary to retrieve current data; without this parameter, average values ​​from the last 6 hours would be returned
        }
        return self.get(f"/virtualservice-inventory", params=params)

    # metrics
    def get_pool_metrics(self, metrics=[]):
        metric_data = {}
        
        params = {
            "include_statistics": False,
            "limit": 1,
            "metric_id": ",".join(metrics),
        }
        
        resp = self.get(f"/analytics/metrics/pool", params=params)

        # parse output
        for s in dig(resp, "results", 0, "series", default=[]):
            header = s.get("header", None)

            if header:
                if header['pool_uuid'] not in metric_data:
                    metric_data[header['pool_uuid']] = {}
                
                # replace . in name from metric, because is unfavorable for checkmk
                # 'l7_server.avg_application_response_time' -> 'l7_server_avg_application_response_time'
                header['name'] = header['name'].replace(".", "_")

                metric_data[header['pool_uuid']][header['name']] = round(dig(s, "data", 0, "value", default=0), 3)

        return metric_data
    

    def get_metric_collections(self, uuid: str, metric_requests=[]):
        metric_data = {}

        params = {
            "include_refs": False,
            "include_statistics": False,
            "pad_missing_data": False,
            "limit": 1,
        }

        data = {
            "metric_requests": metric_requests
        }
        
        resp = self.post("/analytics/metrics/collection", params=params, data=data)

        # parse data
        for series in dig(resp, "series", uuid, default=[]):

            name = dig(series, "header", "name")
            value = dig(series, "data", 0, "value")

            if name != None and value != None:
                # replace . in name from metric, because is unfavorable for checkmk
                # 'l7_server.avg_application_response_time' -> 'l7_server_avg_application_response_time'
                name = name.replace(".", "_")
                metric_data[name] = round(value, 3)

        return metric_data


def avi_time_diff_seconds(timestr: str, format: str = "%Y-%b-%d %H:%M:%S") -> float:
    if not timestr:
        return float("nan")

    try:
        dt = datetime.strptime(timestr, format)
        dt = dt.replace(tzinfo=timezone.utc)
        now = datetime.now(timezone.utc)

        return (now - dt).total_seconds()

    except Exception as e:
        return float("nan")


if __name__ == "__main__":

    args = parse_arguments(sys.argv[1:])
    
    # init api
    avi_agent = VmwareAviLoadBalancerAgent(
        server=args.server,
        username=args.username,
        password=args.password,
        tenant=args.tenant,
        version=args.version,
        verify_ssl=False,
    )

    # try to login
    try:
        avi_agent.login()

        # Alerts
        if 'alert' in args.sections or 'all' in args.sections:
            result = avi_agent.get_alerts()

            if len(result) > 0:
                sys.stdout.write(f"<<<vmware_avi_alert:sep(0)>>>\n")
                if result['count'] > 0:
                    for a in result['results']:
                        tmp_alert = {
                            "level": a['level'],
                            "state": a['state'],
                            "time": datetime.fromtimestamp(a['timestamp']).strftime("%Y-%m-%d %H:%M:%S"),
                            "summary": a['summary'],
                            "obj_name": a['obj_name'],
                        }
                        sys.stdout.write(f"{tmp_alert}\n")

        # Certificates
        if 'certificate' in args.sections or 'all' in args.sections:
            result = avi_agent.get_certificates()

            if len(result) > 0:
                sys.stdout.write(f"<<<vmware_avi_cert:sep(0)>>>\n")
                for cert in result['results']:
                    tmp_cert = {
                        "name": dig(cert, "name"),
                        "status": dig(cert, "status"),
                        "ocsp_error_status": dig(cert, "ocsp_error_status"),
                        "expire_status": dig(cert, "certificate", "expiry_status"),
                        "not_before": dig(cert, "certificate", "not_before"),
                        "not_after": dig(cert, "certificate", "not_after"),
                        "diff_not_after": ((datetime.strptime(cert['certificate']['not_after'], "%Y-%m-%d %H:%M:%S")) - datetime.now()).total_seconds(),
                        "self_signed": dig(cert, "certificate", "self_signed"),
                    }
                    sys.stdout.write(f"{tmp_cert}\n")

        # Cluster and Nodes
        if 'cluster' in args.sections or 'all' in args.sections:
            result = avi_agent.get_cluster()

            if len(result) > 0:
                sys.stdout.write(f"<<<vmware_avi_cluster:sep(0)>>>\n")
                tmp_cluster = {
                    "version": dig(result, "node_info", "version"),
                    "patch": dig(result, "node_info", "patch"),
                    "state": dig(result, "cluster_state", "state"),
                    "uptime": 0,
                    "services": [],
                }

                # uptime
                # The API returns up_since as a UTC timestamp without timezone information.
                # Example: "2026-07-13 15:26:11"
                if dig(result, "cluster_state", "up_since") != None:
                    up_since = datetime.strptime(result['cluster_state']['up_since'], "%Y-%m-%d %H:%M:%S").replace(tzinfo=timezone.utc)
                    tmp_cluster['uptime'] = (datetime.now(timezone.utc) - up_since).total_seconds()

                for s in result['service_states']:
                    tmp_cluster['services'].append({
                        "name": s['name'],
                        "state": s['state']
                    })

                sys.stdout.write(f"{tmp_cluster}\n")

                if len(result['node_states']) > 0:
                    sys.stdout.write(f"<<<vmware_avi_node:sep(0)>>>\n")
                    for n in result['node_states']:
                        tmp_node = {
                            "name": dig(n, "name"),
                            "state": dig(n, "state"),
                            "role": dig(n, "role"),
                            "uptime": 0,
                        }

                        # uptime
                        if dig(n, "up_since") != None:
                            up_since = datetime.strptime(n['up_since'], "%Y-%m-%d %H:%M:%S").replace(tzinfo=timezone.utc)
                            tmp_node['uptime'] = (datetime.now(timezone.utc) - up_since).total_seconds()

                        sys.stdout.write(f"{tmp_node}\n")

        # Clouds
        if 'cloud' in args.sections or 'all' in args.sections:
            result = avi_agent.get_cloudinventory()

            if len(result) > 0:
                sys.stdout.write(f"<<<vmware_avi_cloud:sep(0)>>>\n")

                for cloud in result['results']:

                    tmp_cloud = {
                        "name": dig(cloud, "config", "name"),
                        "state": dig(cloud, "status", "state"),
                    }
                    sys.stdout.write(f"{tmp_cloud}\n")

        # pools
        if "pool" in args.sections or "all" in args.sections:
            pool_metrics = [
                "l4_server.avg_total_rtt",                  # End to End Timing     Server RTT          s
                "l7_server.avg_application_response_time",  # End to End Timing     App Response        s
                "l4_server.avg_bandwidth",                  # Throughput            Throughput          bit/s
                "l4_server.max_open_conns",                 # Open Connections      Open Connections    #
                #"l4_server.avg_est_capacity",              # Estimated Capacity    Estimated Capacity  #
                #"l4_server.avg_available_capacity",        # Available Capacity    Available Capacity  #
                "l4_server.avg_complete_conns",             # New Connections       New Connections     /s
                "l4_server.avg_lossy_connections",          # New Connections       Lossy Connections   /s
                "l4_server.avg_errored_connections",        # New Connections       Bad Connectios      /s
                "l7_server.avg_complete_responses",         # Requests              Requests            /s
                "l7_server.avg_resp_4xx_errors",            # Requests              4xx                 /s
                "l7_server.avg_resp_5xx_errors",            # Requests              5xx                 /s
                "l7_server.pct_response_errors",            # Requests              Request Errors      /s
            ]

            result = avi_agent.get_pool_inventory()
            pool_metrics = avi_agent.get_pool_metrics(metrics=pool_metrics)

            if len(result['results']) > 0:
                sys.stdout.write(f"<<<vmware_avi_pool:sep(0)>>>\n")

                for pool in result['results']:
                    if dig(pool, "config", "uuid") == None:
                        continue

                    # num servers
                    num_servers_down = dig(pool, "runtime", "num_servers_enabled") - dig(pool, "runtime", "num_servers_up")
                    num_servers_disabled = dig(pool, "runtime", "num_servers") - dig(pool, "runtime", "num_servers_enabled")

                    tmp_pool = {
                        "name": dig(pool, "config", "name"),
                        "uuid": dig(pool, "config", "uuid"),
                        "enabled": dig(pool, "config", "enabled"),
                        "state": dig(pool, "runtime", "oper_status", "state"),
                        "num_servers": dig(pool, "runtime", "num_servers"),
                        "num_servers_enabled": dig(pool, "runtime", "num_servers_enabled"),
                        "num_servers_up": dig(pool, "runtime", "num_servers_up"),
                        "num_servers_down": num_servers_down,
                        "num_servers_disabled": num_servers_disabled,
                        "health_score": dig(pool, "health_score"),
                        "alert": dig(pool, "alert"),
                        "virtualservices": dig(pool, "virtualservices"),
                        "metrics": dig(pool_metrics, pool['config']['uuid'], default={}),
                    }
                    sys.stdout.write(f"{tmp_pool}\n")

        # ServiceEngines
        if 'service_engine' in args.sections or 'all' in args.sections:
            metrics = [
                "se_if.avg_bandwidth",
                "se_stats.avg_mem_usage",
                "se_if.avg_rx_pkts",
                "se_if.avg_tx_pkts",
                "se_stats.avg_connection_mem_usage",
                "se_stats.avg_dynamic_mem_usage",
                "se_stats.avg_ssl_session_cache",
                "se_stats.avg_persistent_table_usage",
                "se_stats.avg_packet_buffer_usage",
                "se_if.avg_rx_bytes",
                "se_if.avg_tx_bytes",
            ]

            result = avi_agent.get_serviceengines_inventory(metrics=metrics)

            # PIGGYBACK
            if len(result['results']) > 0:
                for se in result['results']:
                    if dig(se, "config", "name") == None:
                        continue

                    sys.stdout.write(f"<<<<{se['config']['name']}>>>>\n")
                    
                    # label
                    sys.stdout.write("<<<labels:sep(0)>>>\n")
                    sys.stdout.write("{\"vmware_avi/service_engine\": \"yes\"}\n")

                    # alert
                    sys.stdout.write(f"<<<vmware_avi_se_alert:sep(0)>>>\n")
                    alert = dig(se, "alert")
                    sys.stdout.write(f"{alert}\n")

                    # cpu
                    se_cpu = avi_agent.get_serviceengine_cpu(se['uuid'])
                    sys.stdout.write(f"<<<vmware_avi_se_cpu:sep(0)>>>\n")
                    cpu = {
                        "cpu_usage": dig(se_cpu, 0, "total_cpu_utilization", ndigits=3),
                    }
                    sys.stdout.write(f"{cpu}\n")

                    # diskusage
                    se_agent = avi_agent.get_serviceengine_seagent(se['uuid'])
                    sys.stdout.write(f"<<<vmware_avi_se_disk:sep(0)>>>\n")
                    disk_usage = {
                        "disk_usage": dig(se_agent, 0, "disk_space_usage", "percentage_used", expected_type=float, remove_chars=["%"]),
                    }
                    sys.stdout.write(f"{disk_usage}\n")

                    # healthscore
                    sys.stdout.write(f"<<<vmware_avi_se_health:sep(0)>>>\n")
                    health = se['health_score']
                    sys.stdout.write(f"{health}\n")

                    # heartbeat
                    sys.stdout.write(f"<<<vmware_avi_se_hb:sep(0)>>>\n")
                    hb = {
                        "hb_misses": dig(se, "runtime", "hb_status", "num_hb_misses"),
                        "hb_outstanding": dig(se, "runtime", "hb_status", "num_outstanding_hb"),
                        "last_hb_req_sent": avi_time_diff_seconds(se['runtime']['hb_status']['last_hb_req_sent']),
                        "last_hb_resp_recv": avi_time_diff_seconds(se['runtime']['hb_status']['last_hb_resp_recv']),
                    }
                    sys.stdout.write(f"{hb}\n")

                    # runtime
                    sys.stdout.write(f"<<<vmware_avi_se_runtime:sep(0)>>>\n")
                    runtime = {
                        "state": dig(se,"runtime", "oper_status", "state"),
                        "power_state": dig(se, "runtime", "power_state"),
                        "version": dig(se, "runtime", "version"),
                        "license_state": dig(se, "runtime", "license_state"),
                    }
                    sys.stdout.write(f"{runtime}\n")

                    # throughput
                    sys.stdout.write(f"<<<vmware_avi_se_if:sep(0)>>>\n")
                    interface = {
                        "throughput": dig(se, "metrics", "se_if.avg_bandwidth", "value", ndigits=3),
                        "rx_packets": dig(se, "metrics", "se_if.avg_rx_pkts", "value", ndigits=3),
                        "tx_packets": dig(se, "metrics", "se_if.avg_tx_pkts", "value", ndigits=3),
                        "rx_bits": dig(se, "metrics", "se_if.avg_rx_bytes", "value", ndigits=3) * 8,
                        "tx_bits": dig(se, "metrics", "se_if.avg_tx_bytes", "value", ndigits=3) * 8,
                    }
                    sys.stdout.write(f"{interface}\n")

                    # memory
                    sys.stdout.write(f"<<<vmware_avi_se_mem:sep(0)>>>\n")
                    mem = {
                        "mem_usage": dig(se, "metrics", "se_stats.avg_mem_usage", "value", expected_type=float, ndigits=3),
                        "con_mem_usage": dig(se, "metrics", "se_stats.avg_connection_mem_usage", "value", expected_type=float, ndigits=3),
                        "dyn_mem_usage": dig(se, "metrics", "se_stats.avg_dynamic_mem_usage", "value", expected_type=float, ndigits=3),
                    }
                    sys.stdout.write(f"{mem}\n")

                    # uptime
                    sys.stdout.write(f"<<<uptime>>>\n")
                    dt = datetime.strptime(se['runtime']['online_since'], "%Y-%b-%d %H:%M:%S")
                    sys.stdout.write(f"{(datetime.now() - dt).total_seconds()}\n")

                    sys.stdout.write("<<<<>>>>\n") # end of piggyback host

            # === END PIGGYBACK

        # VirtualServices
        if 'virtual_service' in args.sections or 'all' in args.sections:

            vs_metrics = [
                "l4_client.avg_total_rtt",                  # End to End Timing     Client RTT          s
                "l4_server.avg_total_rtt",                  # End to End Timing     Server RTT          s
                "l7_server.avg_application_response_time",  # End to End Timing     App Response        s
                "l7_client.avg_client_data_transfer_time",  # End to End Timing     Data Transfer       s
                "l4_client.avg_bandwidth",                  # Throughput            Throughput          bit/s
                "l4_client.max_open_conns",                 # Open Connections      Open Connections    #
                "l4_client.avg_complete_conns",             # Connections           Connections         /s
                "l4_client.avg_lossy_connections",          # Connections           Lossy Connections   /s
                "l4_client.avg_errored_connections",        # Connections           Bad Connectios      /s
                "l4_client.pct_connection_errors",          # Connections           Connection Errors   %
                "l7_client.avg_complete_responses",         # Requests              Requests            /s
                "l7_server.avg_resp_4xx_errors",            # Requests              4xx                 /s
                "l7_server.avg_resp_5xx_errors",            # Requests              5xx                 /s
                "l7_client.avg_resp_4xx_avi_errors",        # Requests              Avi 4xx             /s
                "l7_client.avg_resp_5xx_avi_errors",        # Requests              Avi 5xx             /s
                "l7_client.pct_response_errors",            # Requests              Request Errors      /s
            ]

            result = avi_agent.get_virtualservices_inventory()

            if len(result['results']) > 0:
                sys.stdout.write(f"<<<vmware_avi_vs:sep(0)>>>\n")
                for vs in result['results']:
                    
                    if dig(vs, "config", "uuid") == None:
                        continue

                    # get metric collections for virtual services
                    vs_metrics_data = avi_agent.get_metric_collections(
                        uuid=vs['config']['uuid'],
                        metric_requests = [
                            {"step": 300, "limit": 1, "entity_uuid": vs['config']['uuid'], "metric_id": ",".join(vs_metrics)},
                        ]
                    )
                    
                    tmp_vs = {
                        "name": dig(vs, "config", "name"),
                        "enabled": dig(vs, "config", "enabled"),
                        "state": dig(vs, "runtime", "oper_status", "state"),
                        "health_score": dig(vs, "health_score"),
                        "num_se_requested": dig(vs, "runtime", "vip_summary", 0, "num_se_requested"),
                        "num_se_assigned": dig(vs, "runtime", "vip_summary", 0, "num_se_assigned"),
                        "pools": dig(vs, "pools", default=[]),
                        "service_engine": dig(vs, "runtime", "vip_summary", 0, "service_engine"),
                        "metrics": vs_metrics_data,
                    }

                    sys.stdout.write(f"{tmp_vs}\n")

    finally:
        try:
            avi_agent.logout()
        except Exception as e:
            print(e)
