#############################################################################
# Copyright (c) 2002-2021 MBARI
# Monterey Bay Aquarium Research Institute, all rights reserved.
#############################################################################
import unittest

import lcm

from test.TestUtils import *

lc = lcm.LCM()


class LCMPublishDataTest(unittest.TestCase):

    def test_publish_lcm(self):
        num_msgs = 10
        msg_pub_rate_sec = 0.01
        pub = LcmPublisher(lc)
        data = read_file()
        data.pop(0)  # drop first line

        channel_name = 'TEST'
        var_name = 'raw_data'

        print("\n")
        print("publishing {} msgs at {} Hz:".format(num_msgs, 1/msg_pub_rate_sec))
        for l in data[:num_msgs]:
            try:
                print("publishing msg # {}".format(pub.seq_number))
                pub.clear_msg()
                pub.add_variable(var_name, l)
                pub.timestamp()
                pub.publish(channel_name)
                time.sleep(msg_pub_rate_sec)
            except KeyboardInterrupt:
                break
        self.assertEqual(True, True)


class LCMPubTest(unittest.TestCase):

    def test_get_type(self):
        pub = LcmPublisher(lc)

        val = [1, [1, 2]]
        for v in val:
            self.assertEqual(pub.get_type(v), numpy.int64)

        val = [1.1, [1.1, 2.2]]
        for v in val:
            self.assertEqual(pub.get_type(v), numpy.float64)

        val = [numpy.float32(1.1), [numpy.float32(1.1), numpy.float32(2.2)]]
        for v in val:
            self.assertEqual(pub.get_type(v), numpy.float32)

        val = ['1.0', ['1.1', '2.2']]
        for v in val:
            self.assertEqual(pub.get_type(v), numpy.str_)

    def test_add_int(self):
        pub = LcmPublisher(lc)

        name = 'int_test'
        unit = 'n/a'

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            pub.add_int(name=name + str(i), unit=unit, val=i)

            self.assertEqual(pub.msg.nIntArrays, i + 1)
            pub.publish('add_int')

    def test_add_int_same_name(self):
        pub = LcmPublisher(lc)

        name = 'int_test'
        unit = 'n/a'

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            pub.add_int(name, unit=unit, val=i)

            self.assertEqual(len(pub.msg.intArrays), 1)
            self.assertEqual(pub.msg.intArrays[0].data[0], i)
            pub.publish('add_int_same_name')

    def test_add_int_array(self):
        pub = LcmPublisher(lc)

        name = 'int_test'
        unit = 'n/a'
        val = []

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            val.append(i)
            pub.add_int(name=name, unit=unit, val=val)

            self.assertEqual(pub.msg.intArrays[0].size, i + 1)
            pub.publish('add_int_same_array')

            pub.clear_msg()
            self.assertEqual(pub.msg.nIntArrays, 0)

    def test_add_float(self):
        pub = LcmPublisher(lc)

        name = 'float_test'
        unit = 'n/a'

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            pub.add_float(name + str(i), unit=unit, val=numpy.float32(i))

            self.assertEqual(pub.msg.nFloatArrays, i + 1)
            pub.publish('add_float')

    def test_add_float_same_name(self):
        pub = LcmPublisher(lc)

        name = 'float_test'
        unit = 'n/a'

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            pub.add_float(name, unit=unit, val=numpy.float32(i))
            
            self.assertEqual(len(pub.msg.floatArrays), 1)
            pub.publish('add_float_same_name')

    def test_add_float_array(self):
        pub = LcmPublisher(lc)

        name = 'float_test'
        unit = 'n/a'
        val = []

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            val.append(numpy.float32(i))
            pub.add_float(name, unit=unit, val=val)

            self.assertEqual(pub.msg.nFloatArrays, 1)
            self.assertEqual(pub.msg.floatArrays[0].size, i + 1)
            pub.publish('add_float_array')

            pub.clear_msg()
            self.assertEqual(pub.msg.nFloatArrays, 0)

    def test_add_double(self):
        pub = LcmPublisher(lc)

        name = 'double_test'
        unit = 'n/a'

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            pub.add_double(name + str(i), unit=unit, val=float(i))

            self.assertEqual(pub.msg.nDoubleArrays, i + 1)
            pub.publish('add_double')

    def test_add_double_same_name(self):
        pub = LcmPublisher(lc)

        name = 'double_test'
        unit = 'n/a'

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            pub.add_double(name, unit=unit, val=float(i))

            self.assertEqual(pub.msg.nDoubleArrays, 1)
            pub.publish('add_double_same_name')

    def test_add_double_array(self):
        pub = LcmPublisher(lc)

        name = 'double_test'
        unit = 'n/a'
        val = []

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            val.append(float(i))

            pub.add_double(name, unit=unit, val=val)
            self.assertEqual(pub.msg.doubleArrays[0].size, i + 1)
            self.assertEqual(len(pub.msg.doubleArrays), 1)


            pub.publish('add_double_array')

            pub.clear_msg()
            self.assertEqual(len(pub.msg.doubleArrays), 0)

    def test_add_string(self):
        pub = LcmPublisher(lc)

        name = 'string_test'
        unit = 'n/a'

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            pub.add_str(name=name+str(i), val='test' + str(i), unit=unit)

            self.assertEqual(pub.msg.nStringArrays, i+1)
            pub.publish('add_string')

    def test_add_string_array(self):
        pub = LcmPublisher(lc)

        name = 'string_test'
        unit = 'n/a'
        val = []

        print("")
        for i in range(10):
            pub.msg.epochMillisec = int(time.time() * 1000)
            val.append('test' + str(i))
            pub.add_str(name=name, val=val, unit=unit)

            self.assertEqual(pub.msg.nStringArrays, 1)
            self.assertEqual(pub.msg.stringArrays[0].size, i + 1)
            pub.publish('add_string_array')

            pub.clear_msg()
            self.assertEqual(pub.msg.nStringArrays, 0)

    def test_add_variable(self):
        pub = LcmPublisher(lc)

        name = 'var_test'
        unit = 'n/a'
        val = [['foo', 'bar'], [1, 2], [1.1, 2.2]]

        pub.add_variable(name, val[0], unit)
        self.assertEqual(pub.msg.stringArrays[0].size, 2)
        pub.clear_msg()

        pub.add_variable(name, val[1], unit)
        self.assertEqual(pub.msg.intArrays[0].size, 2)
        pub.clear_msg()

        pub.add_variable(name, val[2], unit)
        self.assertEqual(pub.msg.doubleArrays[0].size, 2)
        pub.clear_msg()

        print("\ntest mixed val: ", end="")
        pub.add_variable(name, [1.1, '2.2'], unit)
        self.assertEqual(pub.msg.nDoubleArrays, 0)
        pub.clear_msg()


if __name__ == '__main__':
    unittest.main()
