#!/usr/bin/env python

import sys
import time
import socket
import struct
import os
import tempfile

import paho.mqtt.client as mqtt

# All the data parsing is based on what I could gather from other scripts and from looking at the binary data.
# I did not find any helpful protocol description, but the data I get seems to be valid.

MULTICAST_IP = '239.12.255.254'  # fixed multicast IP of SMA Energy Meter/Home Manger
MULTICAST_PORT = 9522  # fixed port

DATA_START_OFFSET = 4
FORCE_PUBLISH_EVERY = 50

ENERGY_MAX = 10000000
POWER_MAX = 100000

dump_data = False
no_mqtt = False
tmp_path = os.path.join(tempfile.gettempdir(), 'sma_dump.bin')
serial_numbers = set()
print_offsets = False


def parse_args():
    global args
    import argparse

    parser = argparse.ArgumentParser(
        description='Listen to SMA Speedwire broadcast traffic and convert it to MQTT messages.',
        formatter_class=argparse.ArgumentDefaultsHelpFormatter
    )
    parser.add_argument('--topic', help='Topic for the MQTT message.', default='sma')
    parser.add_argument('--mqtt_client_id', help='Distinct client ID for the MQTT connection.', default='sma2mqtt')
    parser.add_argument('--mqtt_address', help='Address for the MQTT connection.', default='localhost')
    parser.add_argument('--mqtt_port', help='Port for the MQTT connection.', type=int, default=1883)
    parser.add_argument('--mqtt_username', help='User name for the MQTT connection.', default=None)
    parser.add_argument('--mqtt_password', help='Password name for the MQTT connection.', default=None)
    parser.add_argument('--no_mqtt', help='Don\'t connect to MQTT and just print the values.', action='store_true')
    parser.add_argument('--dump_data', help='Write the binary datagram to {TMP}/sma_dump.bin.', action='store_true')
    parser.add_argument('--serial_nr', help='Only watch packets for this serial number.', default=None)
    parser.add_argument('--force_print_serial', help='Print the serial number, even if only one was found.', action='store_true')
    parser.add_argument('--print_offsets', help='Print the offsets where the patterns were found.', action='store_true')

    return parser.parse_args()


args = parse_args()


class NotAnSmaPacket(Exception):
    pass


class MissingEndMarker(Exception):
    pass


class DataOutOfBounds(Exception):
    pass


class WrongSerialNr(Exception):
    pass


def find_int32_be(data, marker, offset):
    position = data.find(marker, offset) + len(marker)
    length = 4
    return position, int.from_bytes(data[position: position + length], byteorder='big')


def find_bigint64_be(data, marker, offset):
    position = data.find(marker, offset) + len(marker)  # the -1 is here to include the last \x00 into the 8 bytes
    length = 8
    return position, int.from_bytes(data[position: position + length], byteorder="big")


def validate_energy(kwh):
    if abs(kwh) > ENERGY_MAX:
        raise DataOutOfBounds


def validate_power(kw):
    if abs(kw) > POWER_MAX:
        raise DataOutOfBounds


find_funcs = {
    'bigint64_be': find_bigint64_be,
    'int32_be': find_int32_be,
}


validate_funcs = {
    'energy': validate_energy,
    'power': validate_power,
}


rules = [
    {
        'name': 'total_w_buy',
        'pattern': b'\x00\x01\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
    {
        'name': 'kwh_buy',
        'pattern': b'\x00\x01\x08\x00',
        'type': 'bigint64_be',
        'value_type': 'energy',
        'divider': lambda v: round(v / 3600 / 1000, 8),
    },
    {
        'name': 'total_w_sell',
        'pattern': b'\x00\x02\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
    {
        'name': 'kwh_sell',
        'pattern': b'\x00\x02\x08\x00',
        'type': 'bigint64_be',
        'value_type': 'energy',
        'divider': lambda v: round(v / 3600 / 1000, 8),
    },
    {
        'name': 'l1_w_buy',
        'pattern': b'\x00\x15\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
    {
        'name': 'l1_w_sell',
        'pattern': b'\x00\x16\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
    {
        'name': 'l2_w_buy',
        'pattern': b'\x00\x29\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
    {
        'name': 'l2_w_sell',
        'pattern': b'\x00\x2A\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
    {
        'name': 'l3_w_buy',
        'pattern': b'\x00\x3d\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
    {
        'name': 'l3_w_sell',
        'pattern': b'\x00\x3e\x04\x00',
        'type': 'int32_be',
        'value_type': 'power',
        'divider': lambda v: v / 10,
    },
]

computed_values = {
    'l1_w': None,
    'l2_w': None,
    'l3_w': None,
    'total_w': None,
}

values_template = {rule['name']: None for rule in rules} | computed_values

last_values = values_template.copy()

pattern_offsets = {}

counter = 0


def setup_socket():
    mreq = struct.pack("4sl", socket.inet_aton(MULTICAST_IP), socket.INADDR_ANY)

    sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
    sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
    sock.setsockopt(socket.IPPROTO_IP, socket.IP_ADD_MEMBERSHIP, mreq)

    sock.bind(('', MULTICAST_PORT))
    return sock


def green(string):
    return f'\033[0;32m{string}\033[00m'


def red(string):
    return f'\033[0;31m{string}\033[00m'


def blue(string):
    return f'\033[0;34m{string}\033[00m'


def white(string):
    return f'\033[0;37m{string}\033[00m'


def color_value(number):
    if number < 0:
        return red(number)
    elif number > 0:
        return green(number)

    return white('-------')


def decode_speedwire(data):
    global counter, print_offsets

    if not data.startswith(b'SMA'):  # only handle packets that start with SMA
        raise NotAnSmaPacket

    if not data.endswith(b'\x00\x02\x0b\x05\x52\x00\x00\x00\x00'):
        raise MissingEndMarker

    susy_id = data[4: 7].hex()
    # print(susy_id.hex())
    serial_nr = data[20: 24].hex()
    if args.serial_nr and args.serial_nr != serial_nr:
        raise WrongSerialNr

    serial_numbers.add(serial_nr)

    # write binary dgram to disk
    if dump_data:
        with open(tmp_path, 'bw') as tmp_file:
            tmp_file.write(bytearray(data))

    values = {}

    for pattern in rules:
        position, value = find_funcs[pattern['type']](data, pattern['pattern'], DATA_START_OFFSET)
        values[pattern['name']] = pattern['divider'](value)

        if print_offsets:
            pattern_offsets[pattern['name']] = position

    if print_offsets:
        print_offsets = False
        print('Offsets:')
        for pattern in rules:
            hex_pattern = ':'.join([f'{i:02x}' for i in pattern['pattern']])
            print(f'''{hex_pattern}  {pattern_offsets[pattern['name']]: >3}  {pattern['name']}''')
        print()

    if counter == 0:
        serial_hdr = f'{"Serial": >8}' if len(serial_numbers) > 1 or args.force_print_serial else ''
        print(f'{serial_hdr}{"L1 W": >8} + {"L2 W": >8} + {"L3 W": >8} = {"TOTAL W": >9} | {"SELL kWh": >10} {"BUY kWh": >10}')
    counter += 1

    for name in computed_values.keys():
        values[f'{name}'] = values[f'{name}_sell'] if values[f'{name}_sell'] > values[f'{name}_buy'] else values[f'{name}_buy'] * -1

    # sometimes invalid large values get reported
    # check against the bounds and ignore runs with erroneous values
    for rule in rules:
        validate_funcs[rule['value_type']](values[rule['name']])

    # 20 results from the length of 10000.0 (7) + the length of the ANSI color characters (13), TODO preconvert to strings
    kwh_sell_str = f'{values["kwh_sell"]: >10.4f}'
    kwh_buy_str = f'{values["kwh_buy"]: >10.4f}'

    serial_str = blue(f'{serial_nr}') if len(serial_numbers) > 1 or args.force_print_serial else ''
    print(f'{serial_str}{color_value(values["l1_w"]): >20} + {color_value(values["l2_w"]): >20} + {color_value(values["l3_w"]): >20} = {color_value(values["total_w"]): >21} | {green(kwh_sell_str)} {red(kwh_buy_str)}')

    return serial_nr, values


def publish_values(mqtt_client, topic, values):
    global counter, last_values

    # publish all values regardless of changes every n runs
    if counter >= FORCE_PUBLISH_EVERY:
        last_values = values_template.copy()
        counter = 0

    # publish only values that have changed from the former run
    for k, v in values.items():
        if last_values[k] != v:
            if not no_mqtt:
                mqtt_client.publish(f'{topic}/{k}', v)
            last_values[k] = v


dump_data = args.dump_data
no_mqtt = args.no_mqtt
print_offsets = args.print_offsets

mqtt_client = mqtt.Client(args.mqtt_client_id)


def socket_loop():
    # setup the multicast socket to the Speedwire host
    sock = setup_socket()

    # loop that reads the Speedwire and publishes to MQTT
    while True:
        if sock:
            try:
                serial_nr, values = decode_speedwire(sock.recv(1024))

                if not args.serial_nr or args.serial_nr == serial_nr:
                    publish_values(mqtt_client, f'{args.topic}/{serial_nr}', values)

            # ignore errors for this run
            except NotAnSmaPacket:
                pass
            except MissingEndMarker:
                pass
            except DataOutOfBounds:
                pass
            except WrongSerialNr:
                pass


def main():
    if no_mqtt:
        while True:
            socket_loop()

    else:
        if args.mqtt_username and args.mqtt_password:
            mqtt_client.username_pw_set(args.mqtt_username, args.mqtt_password)

        try:
            while True:  # endless loops that reconnects if an mqtt or socket error occurs
                mqtt_client.loop_start()
                mqtt_client.connect(args.mqtt_address, args.mqtt_port)

                print('MQTT connecting', end='')
                # try to connect to the MQTT server
                for i in range(30):  # try for 3 seconds
                    if mqtt_client.is_connected():
                        break
                    print('.', end='')
                    sys.stdout.flush()
                    time.sleep(0.1)
                print('connected\n')

                if not mqtt_client.is_connected():
                    sys.exit(f'Could not connect to the MQTT server at {args.mqtt_address}:{args.mqtt_port}, please check your parameters.')

                socket_loop()

        except Exception as err:
            raise
            print(err)


if __name__ == '__main__':
    main()
