#!/usr/bin/python3

import argparse
import collections
import itertools
import json
import logging
import subprocess
import sys
import time
from operator import itemgetter

import bcapi3dp
import bcutils3dp.discovery.constants
import bcutils3dp.discovery.sup
from bcutils3dp import bcsession

_logger = logging.getLogger("query")

# pylint: disable=no-member
_APT_STATE_CODE = {
    bcapi3dp.State_pb2.State.APT_STATE_MASTER: 1,
    bcapi3dp.State_pb2.State.APT_STATE_SLAVE: 2,
    bcapi3dp.State_pb2.State.APT_STATE_LINK: 3,
    bcapi3dp.State_pb2.State.APT_STATE_NONE: 0,
}


# _DUPLEX_CODE = {
#    bcapi3dp.State_pb2.State.Wired.DUPLEX_NONE: 0.0,
#    bcapi3dp.State_pb2.State.Wired.DUPLEX_HALF: 0.5,
#    bcapi3dp.State_pb2.State.Wired.DUPLEX_FULL: 1.0,
# }


_ROLES = {
    "view": bcsession.SESSION_ROLE_VIEW,
    "admin": bcsession.SESSION_ROLE_ADMIN,
    "co": bcsession.SESSION_ROLE_CO,
}

_BCUTILS_AUTH_DB = "/etc/3d-p/bcutils-common/authdb.json"


_TIME_ONLY_FIELDS = ["tw", "ca", "cb", "cr", "ct"]


def _calc_delta(now, previous, interval, keys):
    delta = now.copy()
    if previous is not None:
        for key in keys:
            # pylint: disable=invalid-name
            v = None
            try:
                v = now[key] - previous[key]
                if not key in _TIME_ONLY_FIELDS:
                    v = v / interval
            # pylint: disable=broad-exception-caught
            except Exception:
                pass
            if v is not None:
                delta[key] = v
    return delta


def _get_amplifiers(model, state):
    amplifier = {}
    for interface in state.wireless:
        model_wireless = None
        for mwval in model.wireless:
            if mwval.name == interface.name:
                model_wireless = mwval
                break
        if model_wireless is not None:
            for rwval in model.radiodb:
                if rwval.model == model_wireless.model:
                    amplifier[interface.name] = rwval.amplifier
    return amplifier


def _get_peers(state):
    output = collections.defaultdict(list)
    for iface in itertools.chain(state.wireless, state.wired):
        output[iface.name].extend(iface.peer)

    # special case for RPT peers
    # output['rpt'].extend(state.system.rptPeer)
    # special case for Live Trace peers
    # output['lt'].extend(state.instamesh.traces)
    return output


def _mac_to_string(mac_string):
    if isinstance(mac_string, bytes):
        return ":".join("%02x" % b for b in mac_string)
    return ":".join("%02x" % ord(b) for b in mac_string)


_COMMSTATS_DIFFS = ["rb", "tb", "rp", "tp", "rd", "pe", "rpp", "tpp", "rdp", "tdp"]


def _parse_commstats(stats, wired=False):
    commstats = {
        "rb": stats.rxBytes,
        "tb": stats.txBytes,
        "rp": stats.rxPackets,
        "tp": stats.txPackets,
        "rpp": stats.rxPausePackets,
        "tpp": stats.txPausePackets,
        "rdp": stats.rxDroppedPackets,
        "tdp": stats.txDroppedPackets,
    }
    if wired:
        commstats["pe"] = stats.pulseEvents
    else:
        commstats["rd"] = stats.radarDetections

    return commstats


_IMESH_DIFFS = [
    "ad",
    "ar",
    "ara",
    "aru",
    "at",
    "fd",
    "pd",
    "pm",
    "pr",
    "ps",
    "sfd",
    "tw",
    "ds",
    "dp",
    "ndd",
    "ndr",
    "ndra",
    "ndru",
    "ndt",
    "ur",
    "utf",
]

_INSTAMESH_PEER_TYPES = {
    # other: -1
    "LOCAL": 0,
    "WIRED_CLIENT": 1,
    "WIRELESS_CLIENT": 2,
    "APT_PEER": 3,
    "WIRELESS_PEER": 4,
    "RPT_PEER": 5,
    "WIRELESS_CLIENT_BRIDGE": 6,
}


def _parse_instamesh(state):
    instamesh = {
        "ad": state.instamesh.arpDropped,
        "ar": state.instamesh.arpRequests,
        "ara": state.instamesh.arpRequestsAnswered,
        "aru": state.instamesh.arpRequestsUnicasted,
        "at": state.instamesh.arpTotal,
        "fd": state.instamesh.floodsDropped,
        "pd": state.instamesh.packetsDropped,
        "pm": state.instamesh.packetsMulticast,
        "pr": state.instamesh.packetsReceived,
        "ps": state.instamesh.packetsSent,
        "sfd": state.instamesh.sourceFloodsDropped,
        "tw": state.instamesh.timeWaited,
        "ds": state.instamesh.discoveriesSourced,
        "dp": state.instamesh.discoveriesPassed,
        "ndd": state.instamesh.ndDropped,
        "ndr": state.instamesh.ndRequests,
        "ndra": state.instamesh.ndRequestsAnswered,
        "ndru": state.instamesh.ndRequestsUnicasted,
        "ndt": state.instamesh.ndTotal,
        "ur": state.instamesh.undeliverablesReceived,
        "utf": state.instamesh.undeliverableTransmitFailures,
    }

    traces = []
    for trace in state.instamesh.traces:
        result = {
            "h": trace.host,
            "pt": _INSTAMESH_PEER_TYPES.get(
                bcapi3dp.Common_pb2.Trace.Path.PathType.Name(trace.path.type), -1
            ),
            "ei": trace.path.encapid,
        }
        if trace.path.cost != 0:
            result["co"] = trace.path.cost
        if trace.path.hopCost != 0:
            result["hco"] = trace.path.hopCost

        if trace.path.mac:
            result["hw"] = _mac_to_string(trace.path.mac)
        if trace.path.ip:
            result["ip"] = trace.path.ip
        if trace.path.name:
            result["n"] = trace.path.name

        if trace.path.channel:
            result.update(
                {
                    "ch": trace.path.channel,
                    "ra": trace.path.rate,
                    # Rajant has RSSI and Signal swapped and Signal is actually SNR
                    "rs": trace.path.signal,
                    "sn": trace.path.rssi,
                }
            )

        traces.append(result)
    instamesh["trc"] = traces

    return instamesh


_WIRELESS_DIFFS = ["ca", "cb", "cr", "ct"]


def _parse_interface_wireless(interface, amplifier):
    commstats = _parse_commstats(interface.stats, wired=False)
    commstats.update(
        {
            "it": 1,
            "hw": _mac_to_string(interface.mac),
            "ch": interface.channel,
            "in": interface.noise,
            "ia": amplifier.get(interface.name),
            "ir": interface.range,
            "txp": interface.txpower,
            "ca": interface.channelActiveTime,
            "cb": interface.channelBusyTime,
            "cr": interface.channelReceiveTime,
            "ct": interface.channelTransmitTime,
        }
    )
    return commstats


_WIRED_DIFFS = ["sc"]


def _parse_interface_wired(interface):
    commstats = _parse_commstats(interface.stats, wired=True)
    commstats.update(
        {
            "it": 0,
            "hw": _mac_to_string(interface.mac),
            "iu": bool(interface.linkup),
            "fd": interface.duplex == bcapi3dp.State_pb2.State.Wired.DUPLEX_FULL,
            "sp": interface.rate,
            "as": _APT_STATE_CODE[interface.aptState],
            "sc": interface.stateChanges,
        }
    )
    return commstats


def _parse_interfaces(state, amplifier):
    out = {
        interface.name: _parse_interface_wired(interface) for interface in state.wired
    }
    out.update(
        {
            interface.name: _parse_interface_wireless(interface, amplifier)
            for interface in state.wireless
        }
    )
    return out


def _parse_peer(channel, peer):
    return {
        "hw": _mac_to_string(peer.mac),
        "co": peer.cost,
        "ei": peer.encapId,
        "en": bool(peer.enabled),
        "ch": channel,
        "ip": peer.ipv4Address,
        "ra": peer.rate / 10.0,
        # Rajant has RSSI and Signal swapped and Signal is actually SNR
        "rs": peer.signal,
        "sn": peer.rssi,
        "ag": peer.age,
        "tp": peer.txpower,
    }


def _get_encap_id(peer):
    try:
        return peer.encapId
    # pylint: disable=bare-except
    except:
        return peer.encapid


def _parse_peers(state, _amplifier):
    iface_channels = dict((i.name, i.channel) for i in state.wireless)

    return {
        interface: {
            _get_encap_id(peer): _parse_peer(
                iface_channels.get(interface, 0),
                peer,
            )
            for peer in interface_peers
        }
        for (interface, interface_peers) in _get_peers(state).items()
        if interface != "rpt"
    }


def _parse_system(state):
    out = {
        "ei": state.system.encapId,
        #'p': state.system.platform,
        "u": state.system.uptime,
        "i": state.system.idle,
        "ru": bool(state.system.running),
        "bu": bool(state.system.bridgeup),
        "fm": state.system.freeMemory * 1024,
        "l": bool(state.system.locked),
        "bc": int(state.system.bootCounter / 2),
        "t": state.system.temperature / 100.0,
    }

    if state.system.reboot:
        out["rb"] = True
    if state.system.isRebooting:
        out["irb"] = True

    for volt_sensor in state.system.sensors.voltage:
        voltage = volt_sensor.value.current / 100.0
        if volt_sensor.name == "battery":
            out["vb"] = voltage
        elif volt_sensor.name == "input":
            out["vi"] = voltage

    return out


_INTERFACE_DIFFS = _WIRED_DIFFS + _WIRELESS_DIFFS + _COMMSTATS_DIFFS


def _resolve_link_name(link_index):
    """get interface name from index"""
    try:
        output = subprocess.run(
            ["ip", "-j", "link"], capture_output=True, check=False
        )
        for link in json.loads(output.stdout):
            if int(link["ifindex"]) == int(link_index):
                return link["ifname"]
    # pylint: disable=bare-except
    except:
        pass
    return None

class _DiscoveryManager:
    """discover breadcrumbs"""

    def __init__(self):
        self._target = None

    def discover(self):
        """discover breadcrumb device"""
        _logger.info("Attempting to discover locally attached BreadCrumb...")
        # find first directly connected crumb and use that
        sup = bcutils3dp.discovery.sup.SupDiscovery(
            interface=None,
            service=bcutils3dp.discovery.constants.SERVICE_V11,
            local=True,
            scantime=5000,
            maxhits=0,
        )
        for crumb in sup.execute():
            if not crumb.local:
                _logger.debug("Ignoring non-local BreadCrumb %s", crumb.serial)
                continue
            _logger.info("Discovered local BreadCrumb %s", crumb.serial)
            # Discovery sometimes ip4, ip6 and link local ip6
            if len(crumb.source_info) >= 4 and crumb.source_info[3]:
                link_name = _resolve_link_name(crumb.source_info[3])
                if link_name:
                    # use link name if applicable
                    self._target = f"{crumb.source}%{link_name}"
                else:
                    self._target = f"{crumb.source}%{crumb.source_info[3]}"
            else:
                self._target = crumb.source
            break
        return self._target


_DISCOVERY_MANAGER = _DiscoveryManager()


class InterfaceStats:
    """Provides access to interface stats"""
    _METRIC_INTERFACE = 0
    _METRIC_INSTAMESH = 1
    _METRIC_PEERS = 2
    _METRIC_SYSTEM = 3

    _target = None  #: target IP of breadcrumb to connect to
    _port = 2300  #: target port of breadcrumb to connect to
    _role = None  #: the breadcrumb auth role being used
    _passphrase = None  #: the passphrase for the auth role
    _session = None  #: the bcapi session
    _session_timeout = 1.0  #: the socket timeout to use for the bcapi session
    _model = None  #: the bcapi provided model info of the connected breadcrumb
    _amplifier = {}  #: A dictionary of amplifier levels by interface

    _first_poll = True
    _last_polls = {}
    _last_poll_time = None

    def __init__(self, settings):
        breadcrumb_settings = settings.get("breadcrumb")
        if breadcrumb_settings:
            self._target = breadcrumb_settings.get("target_ip")
            self._session_timeout = breadcrumb_settings.get("session_timeout", 1.0)
            self._role = _ROLES.get(breadcrumb_settings.get("role"))

        if not self._target:
            _logger.warning(
                "No explicit target BreadCrumb specified; will discover. "
                "Specifying a BreadCrumb is highly recommended!"
            )
            self._target = _DISCOVERY_MANAGER.discover()

        if not self._role:
            _logger.warning(
                "No explicit BreadCrumb role specified; assuming View. "
                "Specifying a role is highly recommended!"
            )
            self._role = bcsession.SESSION_ROLE_VIEW

        # read bcutils auth db for the crumb
        auth_db = {}
        with open(_BCUTILS_AUTH_DB, encoding="utf-8") as bcutils_auth_db:
            dbdata = json.load(bcutils_auth_db)
            if not self._target or (dbdata.get(self._target, None) is None):
                auth_db = dbdata.get("*", {})
            else:
                auth_db = dbdata.get(self._target)
        role_string = (
            "view"
            if not breadcrumb_settings.get("role")
            else breadcrumb_settings.get("role")
        )
        self._passphrase = auth_db.get(role_string)

    def _create_session(self):
        session = bcsession.BcSession()
        session.start(
            self._target,
            port=self._port,
            role=self._role,
            passphrase=self._passphrase,
            timeout=self._session_timeout,
        )
        return session

    def destroy_session(self):
        """destroy breadcrumb session"""
        if self._session is not None:
            try:
                self._session.stop()
            # pylint: disable=broad-exception-caught
            except Exception:
                pass
            finally:
                self._session = None

    def get_session(self):
        """obtain breadcrumb session"""
        if self._session is not None:
            return self._session

        try:
            self._session = self._create_session()

            # fetch model and build map of interface amplifiers
            bcm = bcapi3dp.Message_pb2.BCMessage()
            bcm.state.CopyFrom(bcapi3dp.State_pb2.State())
            bcm.model.CopyFrom(bcapi3dp.ModelDatabase_pb2.BcModel())
            self._session.sendmsg(bcm)
            bcmsg = self._session.recvmsg()
            self._model = bcmsg.model
            self._amplifier = _get_amplifiers(self._model, bcmsg.state)
        # pylint: disable=broad-exception-caught
        except Exception:
            _logger.exception("Unable to start BreadCrumb session")
            self._session = None
            time.sleep(1)
        finally:
            return self._session

    def _try_send_receive(self, message):
        # no message to send?
        if message is None:
            return None

        # no session means we can't send anything
        session = self.get_session()
        if session is None:
            return None

        try:
            session.sendmsg(message)
            return session.recvmsg()
        # pylint: disable=broad-exception-caught
        except Exception as exc:
            _logger.error("Unable to communicate with BreadCrumb: %s", exc, exc_info=1)
            self.destroy_session()
            return None

    def get_state(self):
        """returns current breadcrumb state"""
        try:
            # get state from the crumb
            bcm = bcapi3dp.Message_pb2.BCMessage()
            bcm.state.CopyFrom(bcapi3dp.State_pb2.State())
            return self._try_send_receive(bcm).state
        # pylint: disable=broad-exception-caught
        except Exception as exc:
            # error; return None
            _logger.error("Unable to get BreadCrumb state: %s", exc, exc_info=1)
            return None

    def _get_last_peer_poll(self, iface, encapid):
        lpp = self._last_polls[InterfaceStats._METRIC_PEERS]
        if not iface in lpp:
            return None
        return lpp[iface].get(encapid)

    _MASK = (
        "qtx",
        "qrx",
    )

    def gather(self):
        """gather interface stats"""
        state = self.get_state()
        if state is None:
            raise EnvironmentError("Unable to get BreadCrumb state")
        polls = {
            InterfaceStats._METRIC_INTERFACE: _parse_interfaces(state, self._amplifier),
            InterfaceStats._METRIC_INSTAMESH: _parse_instamesh(state),
            InterfaceStats._METRIC_PEERS: _parse_peers(state, self._amplifier),
            InterfaceStats._METRIC_SYSTEM: _parse_system(state),
        }

        now = time.time()

        # if this is first poll, we can't emit a full picture, so just pass
        if self._first_poll is True:
            self._first_poll = False
            self._last_polls = polls
            self._last_poll_time = now
            return (None, True, InterfaceStats._MASK)

        # determine poll interval time
        poll_interval = now - self._last_poll_time

        # calculate deltas/s for each kind of metric over poll interval

        # calculate CPU utilisation
        system_metrics = polls[InterfaceStats._METRIC_SYSTEM].copy()
        system_metrics_uptime = (
            system_metrics.pop("u")
            - self._last_polls[InterfaceStats._METRIC_SYSTEM]["u"]
        )
        system_metrics_idle = (
            system_metrics.pop("i")
            - self._last_polls[InterfaceStats._METRIC_SYSTEM]["i"]
        )
        if system_metrics_uptime > 0.0:
            system_metrics["cu"] = (
                system_metrics_uptime - system_metrics_idle
            ) / system_metrics_uptime
        else:
            system_metrics["cu"] = 0.0

        interface_stats = dict(
            (
                if_name,
                _calc_delta(
                    interface_now,
                    self._last_polls[InterfaceStats._METRIC_INTERFACE][if_name],
                    poll_interval,
                    _INTERFACE_DIFFS,
                ),
            )
            for (if_name, interface_now) in polls[
                InterfaceStats._METRIC_INTERFACE
            ].items()
        )

        formatted_interface_stats = dict(
            (
                name,
                {
                    #'am': self._amplifier.get(name, 0),
                    "st": iface,
                    "pr": [],
                },
            )
            for (name, iface) in interface_stats.items()
        )

        for iface_name, peers in polls[InterfaceStats._METRIC_PEERS].items():
            fis = formatted_interface_stats.get(iface_name)
            if fis is None:
                fis = formatted_interface_stats[iface_name] = {}

            fis["pr"] = [
                _calc_delta(
                    peer_now,
                    self._get_last_peer_poll(iface_name, encapid),
                    poll_interval,
                    _COMMSTATS_DIFFS,
                )
                for (encapid, peer_now) in peers.items()
            ]
            fis["pr"].sort(key=itemgetter("co"))

        instamesh_stats = _calc_delta(
            polls[InterfaceStats._METRIC_INSTAMESH],
            self._last_polls[InterfaceStats._METRIC_INSTAMESH],
            poll_interval,
            _IMESH_DIFFS,
        )

        self._last_poll_time = now
        self._last_polls = polls

        metric = {
            "if": formatted_interface_stats,
            "im": instamesh_stats,
            "sy": system_metrics,
        }

        return (metric, True, InterfaceStats._MASK)


def __get_settings():

    parser = argparse.ArgumentParser(
        description="interactive monitor for BreadCrumb state",
    )
    parser.add_argument(
        "--target-ip",
        dest="target_ip",
        type=str,
        default=None,
        help="IP address of the BreadCrumb to query",
    )
    parser.add_argument(
        "--session-timeout",
        dest="session_timeout",
        type=float,
        default=1.0,
        help="how long to wait for a session",
    )
    parser.add_argument(
        "--role",
        dest="role",
        type=str,
        choices=tuple(_ROLES.keys()),
        default="view",
        help="authentication role to use",
    )

    args = parser.parse_args()

    output = {}
    if args.target_ip is not None:
        output["target_ip"] = args.target_ip
    if args.session_timeout is not None:
        output["session_timeout"] = args.session_timeout
    if args.role is not None:
        output["role"] = args.role

    return {"breadcrumb": output}


if __name__ == "__main__":
    all_interface_stats = InterfaceStats(settings=__get_settings())

    sys.stderr.write("to generate data, send a linebreak\n")
    while True:
        try:
            sys.stdin.readline()  # doesn't matter what it is
        except EOFError:  # stdin was closed
            break

        sys.stdout.write(json.dumps(all_interface_stats.gather()) + "\n")
        sys.stdout.flush()
