#!/usr/bin/python3
import argparse
import collections
import glob
import json
import os
import pathlib
import re
import signal
import struct
import sys
import threading
import time
import uuid

import paho.mqtt.client as mqtt

_HOST = "localhost"
_PORT = 1883
_KEEPALIVE = 60

_DEFAULT_TOPICS = (
    "3d-p/#",
)

_MQTT_CLIENT = None

_UserData = collections.namedtuple("UserData", (
    "topics",
    "output_handler",
    
    "grep",
    "grep_any",
    "grep_exclude",
))


def sigint_handler(sig, frame):
    _MQTT_CLIENT.disconnect()


def on_connect(client, userdata, flags, rc):
    print("Connected to MQTT bus on {}:{}, connection-status: {}".format(_HOST, _PORT, rc), file=sys.stderr)
    for topic in (userdata.topics or _DEFAULT_TOPICS):
        print("Subscribing to '{}'...".format(topic), file=sys.stderr)
        client.subscribe(topic)


def _prepare_output_file(args):
    output_file_base = pathlib.Path(args.output_file)
    output_file_in_progress = output_file_base.with_suffix(".next")

    output_format = args.output_file_format
    OUTPUT_FORMAT_JSON = 'json'
    OUTPUT_FORMAT_LSJSON = 'linesep-json'
    if output_format not in (
        OUTPUT_FORMAT_JSON,
        OUTPUT_FORMAT_LSJSON,
    ):
        raise ValueError("Unsupported output-format: {}".format(output_format))

    import time
    iso8601 = args.output_file_timestamp_iso8601

    compress_file = args.output_file_compress
    if compress_file:
        import gzip

    output_file_rotate_interval = args.output_file_rotate_interval
    next_file_rotation = None
    if output_file_rotate_interval > 0:
        next_file_rotation = time.time() + output_file_rotate_interval
        print("Output-files will be rotated every {} seconds".format(output_file_rotate_interval), file=sys.stderr)

    output_file_rotate_size = args.output_file_rotate_size
    if output_file_rotate_size > 0:
        print("Output-files will be rotated upon exceeding {} bytes".format(output_file_rotate_size), file=sys.stderr)

    output_file_rotate_count = args.output_file_rotate_count
    if output_file_rotate_count > 0:
        print("Output-files will be limited to a maximum of {}".format(output_file_rotate_count), file=sys.stderr)

    def get_timestamp():
        if iso8601:
            return time.strftime("%Y-%m-%dT%H-%M-%SZ", time.gmtime())
        return str(int(time.time()))

    def parse_timestamp(timestamp):
        if iso8601:
            return time.mktime(time.strptime(timestamp, "%Y-%m-%dT%H-%M-%SZ"))
        return float(timestamp)

    _timestamped_file_re = re.compile(str(output_file_base) + r'\.(.+?)\.json(\.gz)?$')
    def start_file():
        nonlocal next_file_rotation
        
        if output_file_rotate_count > 0:
            glob_pattern = output_file_base.with_suffix('.*.json*')

            existing_files = []
            for candidate in glob.glob(str(glob_pattern)):
                match = _timestamped_file_re.match(candidate)
                if match:
                    try:
                        existing_files.append((parse_timestamp(match.group(1)), candidate))
                    except ValueError:  # unexpected timestamp format
                        continue

            existing_files.sort(reverse=True)
            while len(existing_files) > output_file_rotate_count:
                (_, candidate) = existing_files.pop()
                print("Removing old JSON file '{}'...".format(candidate), file=sys.stderr)
                try:
                    os.unlink(candidate)
                except Exception as e:
                    print("Failed to remove old JSON file: {}".format(e), file=sys.stderr)

        if next_file_rotation is not None:
            next_file_rotation = time.time() + output_file_rotate_interval
            
        print("Creating tempfile '{}'...".format(str(output_file_in_progress)), file=sys.stderr)
        with output_file_in_progress.open("w") as in_progress_file:
            in_progress_file.write('{}\n'.format(get_timestamp()))

    def finish_file():
        if output_file_in_progress.is_file():
            with output_file_in_progress.open('rb') as in_progress_file:
                timestamp = in_progress_file.readline().strip().decode('utf-8')

                if output_format == OUTPUT_FORMAT_JSON:
                    filetype_suffix = '.json'
                elif output_format == OUTPUT_FORMAT_LSJSON:
                    filetype_suffix = ''

                if compress_file:
                    destination_filename = output_file_base.with_suffix('.{}{}.gz'.format(timestamp, filetype_suffix))
                    output_file = gzip.open(destination_filename, "wb")
                else:
                    destination_filename = output_file_base.with_suffix('.{}{}'.format(timestamp, filetype_suffix))
                    output_file = destination_filename.open("wb")

                if output_format == OUTPUT_FORMAT_JSON:
                    print("Writing JSON data to '{}'...".format(destination_filename), file=sys.stderr)
                    output_file.write(b'[')
                    first_line = True
                    for line in in_progress_file:
                        if line:
                            if not line.startswith(b'{') or not line.endswith(b'}\n'):  # lazy, but sufficient, validation to detect corrupt lines
                                print("Skipping malformed JSON entry {}...".format(line), file=sys.stderr)
                                continue

                            if not first_line:
                                output_file.write(b',')
                            else:
                                first_line = False
                            output_file.write(line)
                    output_file.write(b']')
                elif output_format == OUTPUT_FORMAT_LSJSON:
                    print("Writing linebreak-delimited JSON data to '{}'...".format(destination_filename), file=sys.stderr)
                    for line in in_progress_file:
                        if line:
                            if not line.startswith(b'{') or not line.endswith(b'}\n'):  # lazy, but sufficient, validation to detect corrupt lines
                                print("Skipping malformed JSON entry {}...".format(line), file=sys.stderr)
                                continue

                            output_file.write(line)
                output_file.close()

            print("Clearing tempfile '{}'...".format(str(output_file_in_progress)), file=sys.stderr)
            try:
                output_file_in_progress.unlink()
            except Exception as e:
                print("Failed to clear tempfile: {}".format(e), file=sys.stderr)

    finish_file()  # clean up in the event that the system terminated in a dirty state
    start_file()  # start a new file

    def output_file(topic, payload):
        nonlocal next_file_rotation
        if next_file_rotation is not None and time.time() > next_file_rotation:
            finish_file()
            start_file()

        with output_file_in_progress.open("ab") as in_progress_file:
            in_progress_file.write(b'{"topic":"' + topic.encode('utf-8') + b'","payload":' + payload + b'}\n')
            
        if output_file_rotate_size > 0 and output_file_in_progress.stat().st_size >= output_file_rotate_size:
            finish_file()
            start_file()
    return (output_file, finish_file)


def _output_harvest(topic, payload):
    topic = topic.encode('utf-8')
    sys.stdout.buffer.write(b''.join([
        struct.pack('>H', len(topic)),  # two unsigned, big-endian bytes for the topic-length (unsigned short)
        topic,  # the topic itself, as UTF-8-encoded data
        struct.pack('>I', len(payload)),  # four unsigned, big-endian bytes for the payload-length (unsigned int)
        payload,  # the payload itself, as UTF-8-encoded data
    ]))
    sys.stdout.buffer.flush()


def _output_console(topic, payload):
    print()  # add some whitespace
    print(topic)

    try:  # present it as JSON
        payload = json.loads(payload)
    except Exception:  # not JSON
        try:  # present it as text
            payload_decoded = payload.decode('utf-8')
        except Exception:  # not representable as text
            print(payload)
        else:
            print(payload_decoded)
    else:
        print(json.dumps(payload, indent=4, sort_keys=True))


def on_message(client, userdata, msg):
    # check to make sure all matches are satisfied
    for regex in userdata.grep:
        if not regex.search(msg.payload):
            return
            
    # check to make sure all exclusions are satisfied
    for regex in userdata.grep_exclude:
        if regex.search(msg.payload):
            return
            
    # check to make sure any match is satisfied
    if userdata.grep_any:
        for regex in userdata.grep_any:
            if regex.search(msg.payload):
                break
        else:
            return

    userdata.output_handler(msg.topic, msg.payload)


if __name__ == '__main__':
    parser = argparse.ArgumentParser(description='Observe or extract data from the local MQTT bus')
    
    parser.add_argument('topics', metavar='TOPIC', type=str, nargs='*',
                    help="topics to monitor; uses a set of defaults if unspecified")
    
    parser.add_argument('--grep', type=str, nargs='*',
                    help="a regular expression to apply to each payload; the message is only displayed if it matches all instances; may be repeated")
    parser.add_argument('--grep-any', type=str, nargs='*',
                    help="a regular expression to apply to each payload; the message is only displayed if it matches any instance; may be repeated")
    parser.add_argument('--grep-exclude', type=str, nargs='*',
                    help="a regular expression to apply to each payload; the message is only displayed if it does not match any instance; may be repeated")

    parser.add_argument("--harvest", action="store_true",
                        help="format output for consumption by another process; see documentation in-script for details")

    parser.add_argument("--output-file", type=str, default=None,
                        help="write output to a file instead of stdout")
    parser.add_argument("--output-file-format", choices=['json', 'linesep-json'], default='json',
                        help="the output-format to use")
    parser.add_argument("--output-file-rotate-interval", type=int, default=0,
                        help="the number of seconds after which to rotate the output file; defaults to not rotating")
    parser.add_argument("--output-file-rotate-size", type=int, default=0,
                        help="the number of bytes after which to rotate the output file; defaults to not rotating")
    parser.add_argument("--output-file-rotate-count", type=int, default=0,
                        help="the maximum number of output files to keep; defaults to not having a limit")
    parser.add_argument("--output-file-timestamp-iso8601", action="store_true",
                        help="use ISO8601 timestamps as part of filenames, rather than UNIX timestamps")
    parser.add_argument("--output-file-compress", action="store_true",
                        help="compress output files using gzip")
    
    args = parser.parse_args()

    cleanup_handler = lambda: None
    if args.output_file:
        (output_handler, cleanup_handler) = _prepare_output_file(args)
    elif args.harvest:
        output_handler = _output_harvest
    else:
        output_handler = _output_console

    _MQTT_CLIENT = mqtt.Client(
        client_id="mqtt-watch_{}".format(str(uuid.uuid4())),
        clean_session=True,
        userdata=_UserData(
            args.topics,
            output_handler,
            
            [re.compile(v.encode('utf-8')) for v in args.grep or ()],
            [re.compile(v.encode('utf-8')) for v in args.grep_any or ()],
            [re.compile(v.encode('utf-8')) for v in args.grep_exclude or ()],
        ),
    )
    _MQTT_CLIENT.on_connect = on_connect
    _MQTT_CLIENT.on_message = on_message
    
    _MQTT_CLIENT.connect(_HOST, _PORT, _KEEPALIVE)
    
    signal.signal(signal.SIGINT, sigint_handler)
    print("To end execution, use ^C\n", file=sys.stderr)
    
    _MQTT_CLIENT.loop_forever()
    cleanup_handler()  # ensure that all captured data is properly serialised
