#!/usr/bin/env python3

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

from pydbus import SystemBus
from ruamel import yaml

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).")
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__)

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


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

    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

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

        # Configure the controller
        if cname == 'DummyController':
            self.controller = DummyController(name, **config['controller'].get('kwargs', {}))
        else:
            self.controller = LocalController(name, **config['controller'].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')
        path = systemd.LoadUnit(self.config['service'])
        self.unit = bus.get('.systemd1', path)

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


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(STATE_DIR, 'status')
        self.was_pressed = False
        self.last_change = time.time()
        self.interfaces = kwargs.pop('interfaces')

        super(BlueFinThread, self).__init__(*args, **kwargs)
        self.button = ButtonController(str(self.config['button']))

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

    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 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):
        while True:
            try:
                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.warn('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.disconnect()
                else:
                    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()
    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.warn('%s: No such file or directory.')

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


plc_thread = PlcThread('plc', config['plc'])
plc_thread.daemon = True
plc_thread.start()

status_thread = StatusThread('status', config['status'])
status_thread.daemon = True
status_thread.start()

fin_thread = BlueFinThread('fin', config['blue-fin'], interfaces=interfaces)
fin_thread.daemon = True
fin_thread.start()

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