import threading
import time
import json
import base64
import itertools
import requests
import websocket


#websocket.enableTrace(True)

class HolidayHackClient(threading.Thread):
    def __init__(self, url="1w_dq", stop_str="", debug=False):
        super().__init__()
        self.urls = {
            "1w_dq": "wss://signals.holidayhackchallenge.com/wire/dq",
            "spi_mosi": "wss://signals.holidayhackchallenge.com/wire/mosi",
            "spi_sck": "wss://signals.holidayhackchallenge.com/wire/sck",
            "i2c_sda": "wss://signals.holidayhackchallenge.com/wire/sda",
            "i2c_scl": "wss://signals.holidayhackchallenge.com/wire/scl"
        }
        self.wsapp = websocket.WebSocketApp(
            self.urls[url], 
            on_open=self.on_open, 
            on_message=self.on_message)
        self.msgs = []
        self.stop_str = stop_str
        self.stopped = False
        self.debug = debug
        
    def on_open(self, wsapp):
        print("WS Connected")

    def on_message(self, wsapp, message):
        if self.debug:
            print(message)
        try:
            msg_obj = json.loads(message)
        except json.JSONDecodeError as e:
            print("Could not JSON Decode Message: ")
            print(e)
            return
        if msg_obj.get("type", "") == "welcome":
            return
        self.msgs.append(msg_obj)
        if self.stop_str in message:
            self.close() 

    def close(self):
        if not self.stopped:
            self.wsapp.close()
            self.stopped = True

    def send(self, msg):
        self.wsapp.send_text(msg)

    def run(self):
        print("Run")
        self.wsapp.run_forever()
    
    def get_msgs(self):
        return self.msgs
    
def collect_msgs():
    all_msgs = {}
    urls = [
        ["1w_dq", "stop"],
        ["spi_mosi", "7990000"],
        ["spi_sck", "8000000"],
        ["i2c_sda", "2582000"],
        ["i2c_scl", "2580000"],
    ]
    for url in urls:
        try:
            client = HolidayHackClient(url[0], url[1], False)
            client.start()
            while not client.stopped:
                time.sleep(0.001)

            msgs = client.get_msgs()
            all_msgs[url[0]] = []
            # sort
            msgs = sorted(msgs, key=lambda d: d['t'])
            for msg in msgs:
                #print(msg)
                all_msgs[url[0]].append(msg)
        except KeyboardInterrupt as e:
            print("Closing connection")
            client.close()
        finally:
            client.join()
            print("Done")
    with open("all_msgs.json", "w") as f:
        json.dump(all_msgs, f, indent=2)

# Decodes to the following...
# read and decrypt the SPI bus data using the XOR key: icy
def decode_1w(msgs):
    msgs = msgs[4:]
    start = False
    bs = ""
    # Interpret the pulses into a string of bits.
    pulse_start = msgs[0]["t"]
    for x in range(1,len(msgs)):
        if msgs[x]["v"] == msgs[x-1]["v"]:
            continue
        if msgs[x]["v"] == 1 and msgs[x-1]["v"] == 0:
            # Rising Edge
            if msgs[x]["t"]-pulse_start < 15:
                # High Pulse: 1
                bs += "1"
            else:
                # Low Pulse: 0
                bs += "0"
        if msgs[x]["v"] == 0 and msgs[x-1]["v"] == 1:
            # Falling Edge 
            pulse_start = msgs[x]["t"]
    
    s = ""
    # Read the bits right to left, so LSB
    for x in range(0, len(bs), 8):
        bits = bs[x:x+8]
        bit_int = int(bits[::-1], 2)
        s+=chr(bit_int)
    print(s)

#read and decrypt the I2C bus data using the XOR key: bananza. the temperature sensor address is 0x3C
def decode_spi(data_msgs, clock_msgs):
    data_msgs = data_msgs[1:]
    bs = ""
    for msg in clock_msgs:
        if msg.get("marker", "") != "sample":
            continue
        v = value_at(data_msgs, msg["t"])
        bs += str(v)

    int_list = []
    # convert bit string into list of ints
    for x in range(0, len(bs), 8):
        bits = bs[x:x+8]
        bit_int = int(bits, 2)
        int_list.append(bit_int)

    key = b"icy"
    s = xor_decrypt(int_list, key)
    print(s.decode('ascii'))

def value_at(data_msgs, timestamp):
    for x in range(1, len(data_msgs)):
        cur = data_msgs[x]['t']
        prev = data_msgs[x-1]['t']
        if timestamp >= prev and timestamp < cur:
            #print(f"{prev} <= {timestamp} < {cur}")
            return data_msgs[x-1]["v"]
    #print(f"{timestamp} >= {data_msgs[-1]['t']}")
    return data_msgs[-1]['v']

def xor_decrypt(encrypted_data: bytes, key: bytes) -> bytes:
    decrypted_bytes = bytes(a ^ b for a, b in zip(encrypted_data, itertools.cycle(key)))
    return decrypted_bytes

def decode_i2c(data_msgs, clock_msgs):
    transactions = []
    transaction = []
    
    # Read the data, putting the bits in the correct indexes
    # based on the byteIndex,bitIndex fields
    for msg in data_msgs:
        if msg["marker"] == "start":
            transaction = []
        if msg["marker"] in ["address-bit", "data-bit"]:
            if len(transaction) <= msg["byteIndex"]:
                transaction.append(list("00000000"))
            transaction[msg["byteIndex"]][msg["bitIndex"]] = msg["v"]
        if msg["marker"] == "stop":
            transactions.append(transaction)

    # Construct transaction objects (address,data)
    transaction_objs = []
    for t in transactions:
        # right-most 7 bits are the address
        address = hex(int(''.join(map(str,t[0]))[:-1], 2))
        obj = {
            "address": address,
            "data": []
        }
        for x in range(1, len(t)):
            obj["data"].append(int(''.join(map(str,t[x])), 2))
        transaction_objs.append(obj)
    
    # We only care about address 0x3C
    for obj in transaction_objs:
        if obj["address"] != "0x3c":
            continue
        s = xor_decrypt(obj["data"], b"bananza")
        print(s.decode("ascii"))

if __name__ == "__main__":
    with open("all_msgs.json", "r") as f:
        msgs = json.load(f)

    decode_1w(msgs["1w_dq"])
    decode_spi(msgs["spi_mosi"], msgs["spi_sck"])
    decode_i2c(msgs["i2c_sda"], msgs["i2c_scl"])
