#!/usr/bin/env python3

import argparse
import configparser
import glob
import json
import logging
import os
import select
import socket
import subprocess
import threading
import time
import uuid

import paho.mqtt.client as mqtt
from pydbus import SystemBus
from ruamel import yaml

SYSTEM_CONFIG = '/etc/nerve/system.yaml'
GPIO_DIR = '/sys/class/gpio'
STATE_DIR = '/opt/system/var/lib/nerve-led-controller'

parser = argparse.ArgumentParser(description='Control leds on this local instance.')
parser.add_argument('--debug', default=False, action='store_true',
                    help="Print more output to stdout.")
parser.add_argument('--config', default='/etc/nerve/led-controller.yaml', metavar='PATH',
                    help="Set path to config file (default: %(default)s).")
parser.add_argument('--omb-config', metavar='PATH',
                    default='/opt/system/etc/nerve-omb/cfg/oblomb.properties',
                    help='Location of nerve-omb config file (default: %(default)s).')
parser.add_argument('--state-dir', default=STATE_DIR, metavar='PATH',
                    help='Location of the state dir (default: %(default)s).')
args = parser.parse_args()

# global variable to determine if the admin VM is up
if os.environ.get('LOGGING_INCLUDE_DATE') == '0':
    FORMAT = '[%(levelname)-8s] %(message)s'
else:
    FORMAT = '%(asctime)-15s [%(levelname)-8s] %(message)s'
if args.debug:
    logging.basicConfig(level='DEBUG', format=FORMAT)
else:
    logging.basicConfig(level='INFO', format=FORMAT)

log = logging.getLogger(__name__)


class ButtonController:
    def __init__(self, name):
        self.name = name
        self.path = os.path.join(GPIO_DIR, name, 'value')
        self.enable()

    def enable(self):
        log.info('%s: Enabling button.', self.name)
        path = os.path.join(GPIO_DIR, self.name, 'edge')
        with open(path, 'w') as stream:
            stream.write('both')

    def disable(self):
        log.info('%s: Disabling button.', self.name)
        path = os.path.join(GPIO_DIR, self.name, 'edge')
        with open(path, 'w') as stream:
            stream.write('none')

    def poll(self):
        with select.epoll() as p:
            fd = os.open(self.path, os.O_RDONLY)
            try:
                os.read(fd, 1)
                os.lseek(fd, 0, os.SEEK_END)
                p.register(fd, select.EPOLLET)
                p.poll()
            finally:
                os.close(fd)

    def is_pressed(self, button):
        with open(self.path, 'r') as stream:
            return stream.read().strip() == '0'


class LocalController:
    def __init__(self, name, led, dev='de01'):
        self.name = name
        self.led = led
        self.dev = dev

    def set_color(self, color):
        log.info('%s: Set LED to %s.', self.led, color)

        # get maximum brightnes
        max_path = '/sys/class/leds/%s:%s:%s/max_brightness' % (self.dev, color, self.led)
        with open(max_path) as stream:
            max_brightness = stream.read()

        path = '/sys/class/leds/%s:%s:%s/brightness' % (self.dev, color, self.led)
        with open(path, 'w') as stream:
            stream.write(max_brightness)

    def clear(self):
        log.info('%s: Turn off LED', self.led)
        for path in glob.glob('/sys/class/leds/%s:*:%s/brightness' % (self.dev, self.led)):
            with open(path, 'w') as stream:
                log.debug('%s -> 0', path)
                stream.write('0')


class DummyController(object):
    def __init__(self, name, *args, **kwargs):
        self.name = name
        self.args = args
        self.kwargs = kwargs
        self.led = kwargs.get('led', 'dummy')

    def set_color(self, color):
        log.info('%s: Set LED to %s.', self.led, color)

    def clear(self):
        log.info('%s: Turn off LED', self.led)


class LedThread(threading.Thread):
    def __init__(self, name, config, *args, **kwargs):
        self.led_name = name
        self.config = config

        cconfig = config.get('controller', {})
        cname = cconfig.get('class', 'DummyController')
        log.info('%s LED using controller "%s"', name, cname)

        # Configure the controller
        if cname == 'DummyController':
            self.controller = DummyController(name, **cconfig.get('kwargs', {}))
        else:
            self.controller = LocalController(name, **cconfig.get('kwargs', {}))

        super(LedThread, self).__init__(*args, **kwargs)

    def run(self):
        # Common method since we always use green = ok, off = not ok.
        self.old_status = False
        self.controller.clear()

        while True:
            try:
                status = self.status()

                if status != self.old_status:
                    self.old_status = status

                    try:
                        if status is True:
                            self.controller.set_color('green')
                        else:
                            self.controller.clear()
                    except Exception as e:
                        log.exception(e)

            finally:
                # make sure that this check runs at most once per second
                time.sleep(1)


class SocketMixin(object):
    def status(self):
        try:
            address = (self.config['host'], self.config['port'], )

            s = socket.create_connection(address, self.config.get('timeout', 1))
            s.close()

            return True
        except Exception:
            return False


class SystemdStatusMixin:
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)

        bus = SystemBus()
        systemd = bus.get('org.freedesktop.systemd1')
        try:
            self.service = self.config['service']
        except KeyError:
            log.error('Could not get SystemD service from configuration.')
        else:
            path = systemd.LoadUnit(self.config['service'])
            self.unit = bus.get('.systemd1', path)

    def status(self):
        try:
            return self.unit.Get('org.freedesktop.systemd1.Unit', 'ActiveState') == 'active'
        except AttributeError:
            return False


class PlcThread(SocketMixin, LedThread):
    """Checks status of the Codesys runtime."""
    pass


class StatusThread(SystemdStatusMixin, LedThread):
    """Checks status of OVDM."""
    pass


class BlueFinThread(LedThread):
    def __init__(self, *args, **kwargs):
        self.state_path = os.path.join(kwargs.pop('state_dir'), 'status')
        self.was_pressed = False
        self.last_change = time.time()
        self.interfaces = kwargs.pop('interfaces')

        self.omb_config_path = kwargs.pop('omb_config_path')
        self.retrieve_serial()

        super(BlueFinThread, self).__init__(*args, **kwargs)

        self.button = None
        if 'button' in self.config:
            try:
                self.button = ButtonController(str(self.config['button']))
            except Exception as e:
                log.exception(e)

        if self.config.get('mqtt'):
            mqtt_config = self.config['mqtt']
            self.mqtt = mqtt.Client(client_id=mqtt_config.get('client_id', 'app/nlc'))
            self.mqtt.on_connect = self.mqtt_on_connect
            self.mqtt.on_message = self.mqtt_on_message
            self.mqtt.username_pw_set(mqtt_config['username'], mqtt_config['password'])

    def retrieve_serial(self):
        try:
            with open(self.omb_config_path) as stream:
                data = stream.read()

            config = configparser.ConfigParser()
            config.read_string('[ovdm]\n' + data)
            self.serial = config['ovdm']['gtw.id']
        except Exception:
            log.error('Could not retrieve serial from %s.', self.omb_config_path)

    def mqtt_on_connect(self, client, userdata, flags, rc):
        topic = "gtw/%s/nlc/req" % self.serial
        log.info('Subscribing to %s', topic)
        client.subscribe(topic)

    def mqtt_on_message(self, client, userdata, msg):
        log.debug('Received MQTT message on %s: %s', msg.topic, msg.payload)
        data = json.loads(msg.payload.decode('utf-8'))

        if data['name'] == 'network_state':
            cmd = data['params']['command']
            log.info('Received %s command via MQTT', cmd)
            params = {}
            message = ''
            code = 0

            if cmd == 'set_network_off':
                if self.connected:
                    try:
                        self.disconnect()
                    except Exception as e:
                        code = 2
                        message = str(e)
                else:
                    message = 'Node is already disconnected.'
                    code = 1

            elif cmd == 'set_network_on':
                if not self.connected:
                    try:
                        t = threading.Thread(target=self.connect)
                        t.start()
                    except Exception as e:
                        code = 2
                        message = str(e)
                else:
                    message = 'Node is already connected.'
                    code = 1

            elif cmd == 'get_network_state':
                params['status'] = 'on' if self.connected else 'off'
            else:
                log.error('Unknown command received: %s', cmd)
                return

            # update params for payload
            params['code'] = code
            if code != 0 and message:
                params['message'] = message

            payload = json.dumps({
                'name': data['name'],
                'sender': 'gtw/%s/nlc' % self.serial,
                'uid': data['uid'],
                'params': params,
            })

            log.debug('MQTT response to %s/rsp: %s', data['sender'], payload)
            self.mqtt.publish("%s/rsp" % data['sender'], payload)
        else:
            log.error('Unknown message name: %s', data['name'])

    def _ifdown(self, iface):
        try:
            cmd = ['ifdown', iface]
            log.info('# %s', ' '.join(cmd))
            subprocess.check_call(cmd)
        except subprocess.CalledProcessError as e:
            log.exception(e)

    def _ifup(self, iface):
        try:
            cmd = ['ifup', iface]
            log.info('# %s', ' '.join(cmd))
            subprocess.check_call(cmd)
        except subprocess.CalledProcessError as e:
            log.exception(e)

    def _set_ip_forward(self, state):
        """Disable IP forwarding.

        We have to completely disable IP-forwarding because the libvirt default network goes through the
        FORWARD chain and not the OUTPUT chain (see _firewall_cmd()).

        We cannot drop packages using firewall-cmd because libvirt handles it's forwarding configuration
        itself and firewall-cmd inserts its command after libvirts rules.
        """

        iface = self.config.get('wan_interface', 'br-wan')

        try:
            cmd = ['sysctl', '-w', 'net.ipv4.conf.%s.forwarding=%s' % (iface, state)]
            log.info('# %s', ' '.join(cmd))
            subprocess.check_call(cmd)
        except subprocess.CalledProcessError as e:
            log.exception(e)

        try:
            cmd = ['sysctl', '-w', 'net.ipv6.conf.%s.forwarding=%s' % (iface, state)]
            log.info('# %s', ' '.join(cmd))
            subprocess.check_call(cmd)
        except subprocess.CalledProcessError as e:
            log.exception(e)

    def _load_state(self):
        if os.path.exists(self.state_path):
            with open(self.state_path) as stream:
                try:
                    data = json.load(stream)
                    self.connected = data.get('connected', True)
                except ValueError as e:
                    log.exception(e)
                    self.connected = False
        else:
            log.info('%s not present, assuming we are online.', self.state_path)
            self.connected = True

    def _update_state(self, new_state):
        if os.path.exists(self.state_path):
            with open(self.state_path) as stream:
                try:
                    state = json.load(stream)
                except ValueError as e:
                    log.exception(e)
                    state = {}
        else:
            state = {}

        state['connected'] = self.connected = new_state

        with open(self.state_path, 'w') as stream:
            json.dump(state, stream)

    def mqtt_network_state_change(self, status):
        # notify LocalUI of state change
        payload = json.dumps({
            'name': 'network_state_change',
            'sender': 'gtw/%s/nlc' % self.serial,
            'uid': str(uuid.uuid4()),
            'params': {
                'status': status
            },
        })

        topic = 'gtw/%s/nlc/evt' % self.serial
        log.debug('MQTT event to %s: %s', topic, payload)
        self.mqtt.publish(topic, payload)

    def connect(self):
        log.info('Enable outgoing connections')
        #self._firewall_cmd('--remove-rule', 'ipv4')
        #self._firewall_cmd('--remove-rule', 'ipv6')
        #self._set_ip_forward('1')
        for iface in self.interfaces:
            self._ifup(iface)

        self._update_state(True)
        self.controller.set_color('blue')
        self.last_change = time.time()

    def disconnect(self):
        log.info('Disable outgoing connections')
        #self._firewall_cmd('--add-rule', 'ipv4')
        #self._firewall_cmd('--add-rule', 'ipv6')
        #self._set_ip_forward('0')
        for iface in self.interfaces:
            self._ifdown(iface)

        self._update_state(False)
        self.controller.clear()
        self.last_change = time.time()

    def run(self):
        self.mqtt.connect(self.config['mqtt']['host'], port=self.config['mqtt']['port'])
        self.mqtt.loop_start()

        self._load_state()
        if self.connected:
            log.info('Doing initial connect...')
            self.connect()
        else:
            log.info('Doing initial disconnect...')
            self.disconnect()

        while True:
            try:
                if self.button is None:
                    continue  # NOTE: still sleep for 1s in the finally-block

                self.button.poll()

                if self.config.get('send-udp-package'):
                    dest_ip = self.config.get('UDP_IP', '10.10.10.1')
                    dest_port = int(self.config.get('UDP_PORT', 15043))
                    dest_message = self.config.get('UDP_MESSAGE', b'foo')
                    log.warning('Sending UDP package to %s:%s.', dest_ip, dest_port)
                    sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
                    sock.sendto(dest_message, (dest_ip, dest_port))

                elif self.connected is True:
                    self.mqtt_network_state_change(status='off')
                    self.disconnect()
                else:
                    self.mqtt_network_state_change(status='on')
                    self.connect()
            except Exception:
                pass
            finally:
                # Button is disabled for a second after pressing it.
                # also serves as a rate limit when the above code throws an exception
                time.sleep(1)


with open(args.config) as stream:
    config = yaml.safe_load(stream)

if os.path.exists(SYSTEM_CONFIG):
    with open(SYSTEM_CONFIG) as stream:
        system_config = yaml.safe_load(stream)
    #interfaces = system_config.get('localui', {}).get('interfaces', {}).keys()
    # intermediate state: only WAN interface is configured, all other interfaces remain unchanged
    interfaces = ['wan']
    rtvm_ip = system_config.get('codesys', {}).get('address')

    for key, value in config.items():
        if value.get('host') == 'rtvm':
            value['host'] = rtvm_ip

else:
    log.warning('%s: No such file or directory.', SYSTEM_CONFIG)
    interfaces = []  # connect/disconnect will not manage any interfaces


state_dir = os.path.abspath(args.state_dir)
omb_config = os.path.abspath(args.omb_config)


if not os.path.exists(state_dir):
    os.makedirs(state_dir)


# Display status of CODESYS runtime
if 'plc' in config:
    plc_thread = PlcThread('plc', config['plc'])
    plc_thread.daemon = True
    plc_thread.start()

# Display (SystemD) status of nerve-ovdm
if 'status' in config:
    status_thread = StatusThread('status', config['status'])
    status_thread.daemon = True
    status_thread.start()

# Handles blue fins, button is embedded
fin_thread = BlueFinThread('fin', config['blue-fin'], interfaces=interfaces,
                           state_dir=state_dir, omb_config_path=omb_config)
fin_thread.daemon = True
fin_thread.start()

while True:
    try:
        time.sleep(10)
    except KeyboardInterrupt:
        break
