#! /usr/bin/env python
__author__ = "Ben Raanan"
__copyright__ = "Copyright 2021, MBARI"
__credits__ = ["MBARI"]
__license__ = "GPL-3.0"
__maintainer__ = "Ben Raanan"
__email__ = "byraanan at mbari.org"
__doc__ = '''

Defines an LCM publisher class for LRAUV backseat helper service.

@author: __author__
@license: __license__
'''
import logging
import time
import numpy

from LCM.TethysLcmTypes import LrauvLcmMessage, ByteArray, IntArray, FloatArray, DoubleArray, StringArray

logger = logging.getLogger("backseat")


class LcmPublisher(object):
    """
    LCM publisher class implementation
    """

    def __init__(self, lcm_instance, source_id=0):
        self.lc = lcm_instance  # initialized LCM object
        self.msg = LrauvLcmMessage()
        self.msg.source = source_id
        self.source_id = source_id
        self.seq_number = 0

        # List of supported types for add_variable method.
        # Note: python floats have double precision
        self.supported_byte = (bytes, bytearray, numpy.bytes_, numpy.byte, numpy.ubyte)
        self.supported_bool = (bool, numpy.bool_)  # bools are treated as bytes
        self.supported_int = (int, numpy.int16, numpy.int32, numpy.int64)
        self.supported_float = (numpy.float16, numpy.float32)
        self.supported_double = (float, numpy.float64)
        self.supported_string = (str, numpy.str_)

    def assign_value(self, array, val, valid=True):
        """
        Assign value to the array's data member and set dim and size appropriately.

        :param array: a supported TethysLcmType array object.
        :param val: value to be assigned
        :param valid: value validity flag
        :return: void
        """
        # handle bytes
        if isinstance(val, self.supported_byte):
            array.data = list(val)
            array.nDim = 1
            array.shape = [len(val)]

        # handle strings
        elif isinstance(val, self.supported_string):
            array.data = [val]
            array.nDim = 1
            array.shape = [len(array.data)]

        # handle iterable objects
        elif hasattr(val, '__iter__'):
            v = numpy.asarray(val)
            array.data = list(v.flat)
            array.nDim = v.ndim
            array.shape = v.shape

        # handle scalar data
        else:
            array.data = [val]
            array.nDim = 0
            array.shape = []

        array.size = len(array.data)
        array.valid = valid

    @staticmethod
    def get_name_index(array, name):
        """
        Finds the index of named array.

        :param array: TethysLcm array.
        :param name: name of variable to be found
        :return: index of named variable, or False if not found
        """
        for i, v in enumerate(array):
            if v.name == name:
                return i
        return None

    @staticmethod
    def get_type(val):
        """
        Determine the type of an element or elements (if homogeneous iterable obj)

        :param val: element or iterable obj
        :return: type of val or False if iterable obj is heterogeneous
        """
        try:
            # get an iterator to the flattened data structure
            # and check for consistent type:
            flat_iter = numpy.array(val).flat
            first_type = type(next(flat_iter))
            return first_type if all((type(x) is first_type) for x in flat_iter) else False
        except TypeError:
            # handle non-iterable data
            return type(val)

    def clear_msg(self):
        self.msg = LrauvLcmMessage()

    def timestamp(self, time_stamp_epoch_ms=None):
        """
        Set DataVector LCM message timestamp.

        :param time_stamp_epoch_ms: timestamp in epoch milliseconds
        :return: updates msg class member
        """
        if not time_stamp_epoch_ms:
            time_stamp_epoch_ms = int(time.time() * 1000)

        self.msg.epochMillisec = int(time_stamp_epoch_ms)

    def publish(self, channel_name):
        """
        LCM publish method for DataVectors type messages

        :param channel_name: string containing the LCM channel name
        :return: void
        """
        try:
            self.seq_number += 1
            self.msg.seqNo = self.seq_number
            if not self.msg.epochMillisec:
                self.timestamp()

            self.lc.publish(channel_name, self.msg.encode())
            self.msg.epochMillisec = None
            return True
        except:
            logger.error('Failed to publish LCM msg {}. '.format(self.seq_number), exc_info=True)
            return False

    def add_variable(self, name, val, unit='n/a'):
        """
        Create or add a member to a TethysLCM message array.

        :param name: member name str
        :param val: member value (can be iterable)
        :param unit: member unit str
        :return: updates class member msg
        """

        # determine the type of val
        type_ = self.get_type(val)

        # add the data to appropriate member
        if type_ in self.supported_bool or isinstance(val, self.supported_byte):
            self.add_byte(name, val, unit)

        elif type_ in self.supported_int:
            self.add_int(name, val, unit)

        elif type_ in self.supported_float:
            self.add_float(name, val, unit)

        elif type_ in self.supported_double:
            # Note: python floats have double precision
            self.add_double(name, val, unit)

        elif type_ in self.supported_string:
            self.add_str(name, val, unit)

        else:
            logger.error("Unsupported variable type for '{}' with val: {} {}".format(name, val, type(val)))

    def add_byte(self, name: str, val, unit: str):
        """
        Create or add a int valued member to nested ByteArray.

        :param name: member name str
        :param val: member value (can be iterable)
        :param unit: member unit str
        :return: updates class member msg
        """

        # get index to array member with matching name, None if not found
        idx_ = self.get_name_index(self.msg.byteArrays, name)

        if idx_ is not None:
            # name exists, update value
            self.assign_value(self.msg.byteArrays[idx_], val)
        else:
            # new member, create and populate
            new_array = ByteArray()
            new_array.name = name
            new_array.unit = unit
            self.assign_value(new_array, val)

            # update the LCM msg obj
            self.msg.byteArrays.append(new_array)
            self.msg.nByteArrays = len(self.msg.byteArrays)

    def add_int(self, name, val, unit):
        """
        Create or add a int valued member to nested IntArray.

        :param name: member name str
        :param val: member value (can be iterable)
        :param unit: member unit str
        :return: updates class member msg
        """

        # get index to array member with matching name, None if not found
        idx_ = self.get_name_index(self.msg.intArrays, name)

        if idx_ is not None:
            # name exists, update value
            self.assign_value(self.msg.intArrays[idx_], val)
        else:
            # new member, create and populate
            new_array = IntArray()
            new_array.name = name
            new_array.unit = unit
            self.assign_value(new_array, val)

            # update the LCM msg obj
            self.msg.intArrays.append(new_array)
            self.msg.nIntArrays = len(self.msg.intArrays)

    def add_float(self, name, val, unit):
        """
        Create or add a float32 valued member to nested FloatArray.

        :param name: member name str
        :param val: member value (can be iterable)
        :param unit: member unit str
        :return: updates class member msg
        """

        # get index to array member with matching name, None if not found
        idx_ = self.get_name_index(self.msg.floatArrays, name)

        if idx_ is not None:
            # name exists, update value
            self.assign_value(self.msg.floatArrays[idx_], val)

        else:
            # new member, create and populate
            new_array = FloatArray()
            new_array.name = name
            new_array.unit = unit
            self.assign_value(new_array, val)

            # update the LCM msg obj
            self.msg.floatArrays.append(new_array)
            self.msg.nFloatArrays = len(self.msg.floatArrays)

    def add_double(self, name, val, unit):
        """
        Create or add a double valued member to nested DoubleArray..

        :param name: member name str
        :param val: member value (can be iterable)
        :param unit: member unit str
        :return: updates class member msg
        """

        # get index to array member with matching name, None if not found
        idx_ = self.get_name_index(self.msg.doubleArrays, name)

        if idx_ is not None:
            # name exists, update value
            self.assign_value(self.msg.doubleArrays[idx_], val)

        else:
            # new member, create and populate
            new_array = DoubleArray()
            new_array.name = name
            new_array.unit = unit
            self.assign_value(new_array, val)

            # update the LCM msg obj
            self.msg.doubleArrays.append(new_array)
            self.msg.nDoubleArrays = len(self.msg.doubleArrays)

    def add_str(self, name, val, unit):
        """
        Create or add a str valued member to nested StringArray.

        :param name: member name str
        :param val: member value str
        :param unit: member unit str
        :return: updates class member msg
        """

        # get index to array member with matching name, None if not found
        idx_ = self.get_name_index(self.msg.stringArrays, name)

        if idx_ is not None:
            # name exists, update value
            self.assign_value(self.msg.stringArrays[idx_], val)
        else:
            # new member, create and populate
            new_array = StringArray()
            new_array.name = name
            new_array.unit = unit
            self.assign_value(new_array, val)

            # update the LCM msg obj
            self.msg.stringArrays.append(new_array)
            self.msg.nStringArrays = len(self.msg.stringArrays)
