import ctypes
import contextlib
import io
import socket
import struct
import threading
import unittest
from unittest import mock

import msc
import mscb


class FakeSubmaster:
    def __init__(self):
        self.socket = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
        self.socket.bind(("127.0.0.1", 0))
        self.socket.settimeout(0.1)
        self.port = self.socket.getsockname()[1]
        self.requests = []
        self.running = True
        self.thread = threading.Thread(target=self._run, daemon=True)

    def __enter__(self):
        self.thread.start()
        return self

    def __exit__(self, *_args):
        self.running = False
        self.thread.join(timeout=1)
        self.socket.close()

    def _run(self):
        while self.running:
            try:
                packet, peer = self.socket.recvfrom(1500)
            except socket.timeout:
                continue
            size, sequence, flags, version = struct.unpack("!HHBB", packet[:6])
            request = packet[6:]
            self.requests.append((flags, request))
            assert size == len(request)
            assert version == mscb.MSCB_PROTOCOL_VERSION
            response = self._respond(flags, request)
            if response is not None:
                header = struct.pack("!HHBB", len(response), sequence, 0, mscb.MSCB_PROTOCOL_VERSION)
                self.socket.sendto(header + response, peer)

    @staticmethod
    def _respond(flags, request):
        if flags & mscb.RS485_FLAG_NO_ACK:
            return None
        if flags & mscb.RS485_FLAG_CMD:
            if request[0] == mscb.MCMD_ECHO:
                return bytes((mscb.MCMD_ECHO, 5, 0x12, 0x34))
            if request[0] == mscb.MCMD_TOKEN:
                return bytes((mscb.MCMD_ACK,))
        if request[0] == mscb.MCMD_PING16:
            return bytes((mscb.MCMD_ACK,))
        if request[0] != mscb.MCMD_ADDR_NODE16:
            return b"\xff"

        command = request[4]
        if command == mscb.MCMD_GET_INFO:
            raw = bytearray(84)
            raw[0:2] = bytes((5, 2))
            raw[2:4] = (0x1234).to_bytes(2, "big")
            raw[4:6] = (0x0102).to_bytes(2, "big")
            raw[6:8] = (0x2345).to_bytes(2, "big")
            raw[8:16] = b"FAKENODE"
            raw[30:32] = (256).to_bytes(2, "big")
            response = bytes((mscb.MCMD_ACK + 7, len(raw))) + raw
            return response + bytes((mscb.crc8(response),))
        if command == mscb.MCMD_GET_INFO + 1:
            raw = bytes((2, 24, 0, 0, 0)) + b"Voltage\0".ljust(16, b"\0")
            response = bytes((mscb.MCMD_ACK + 7, len(raw))) + raw
            return response + bytes((mscb.crc8(response),))
        if command == mscb.MCMD_READ + 1:
            response = bytes((mscb.MCMD_ACK + 2, 0x12, 0x34))
            return response + bytes((mscb.crc8(response),))
        if mscb.MCMD_WRITE_ACK < command < mscb.MCMD_WRITE_ACK + 8:
            return bytes((mscb.MCMD_ACK, request[-1]))
        return b"\xff"


class PurePythonProtocolTests(unittest.TestCase):
    def tearDown(self):
        for fd in list(mscb._connections):
            mscb.mscb_exit(fd)

    def test_crc_matches_reference_table(self):
        self.assertEqual(mscb.crc8(b""), 0)
        self.assertEqual(mscb.crc8(b"\x01"), 0x5E)
        self.assertEqual(mscb.crc8(b"123456789"), 0xA1)

    def test_init_info_read_write_and_close_over_udp(self):
        with FakeSubmaster() as server:
            device = f"127.0.0.1:{server.port}"
            fd = mscb.mscb_init(device, len(device), b"secret", 0)
            self.assertGreater(fd, 0)
            self.assertEqual(mscb.mscb_ping(fd, 0x1234, 0, 0), mscb.MSCB_SUCCESS)

            info = msc.MSCB_INFO()
            self.assertEqual(mscb.mscb_info(fd, 0x1234, ctypes.byref(info)), mscb.MSCB_SUCCESS)
            self.assertEqual(info.node_address, 0x1234)
            self.assertEqual(info.group_address, 0x0102)
            self.assertEqual(info.buf_size, 256)
            self.assertEqual(bytes(info.node_name), b"FAKENODE")

            variable = msc.MSCB_INFO_VAR()
            self.assertEqual(
                mscb.mscb_info_variable(fd, 0x1234, 0, ctypes.byref(variable)),
                mscb.MSCB_SUCCESS,
            )
            self.assertEqual(variable.width, 2)
            self.assertEqual(bytes(variable.name), b"Voltage")

            size = ctypes.c_int(2)
            value = ctypes.create_string_buffer(2)
            self.assertEqual(mscb.mscb_read(fd, 0x1234, 0, value, ctypes.byref(size)), mscb.MSCB_SUCCESS)
            self.assertEqual(value.raw, b"\x34\x12")

            output = ctypes.c_ushort(0x1234)
            self.assertEqual(mscb.mscb_write(fd, 0x1234, 0, ctypes.byref(output), 2), mscb.MSCB_SUCCESS)
            write_request = next(request for _flags, request in server.requests
                                 if len(request) > 4 and request[4] == mscb.MCMD_WRITE_ACK + 3)
            self.assertEqual(write_request[6:8], b"\x12\x34")

            self.assertEqual(mscb.mscb_exit(fd), mscb.MSCB_SUCCESS)

    def test_scan_resolves_mscb_names_and_reports_echo_information(self):
        class FakeScanSocket:
            def __init__(self, *_args, **_kwargs):
                self.pending = None

            def sendto(self, packet, peer):
                self.pending = (packet, peer)
                return len(packet)

            def settimeout(self, _timeout):
                pass

            def recvfrom(self, _size):
                if self.pending is None:
                    raise socket.timeout
                request, peer = self.pending
                self.pending = None
                _size, sequence, _flags, _version = struct.unpack("!HHBB", request[:6])
                uptime = 90061
                payload = bytes((mscb.MCMD_ECHO, 5, 0x12, 0x34)) + uptime.to_bytes(4, "big")
                header = struct.pack("!HHBB", len(payload), sequence, 0, mscb.MSCB_PROTOCOL_VERSION)
                return header + payload, peer

            def close(self):
                pass

        def resolve(hostname):
            if hostname == "MSCB007":
                return "192.0.2.7"
            raise socket.gaierror

        output = io.StringIO()
        with mock.patch.object(mscb.socket, "socket", FakeScanSocket), \
             mock.patch.object(mscb.socket, "gethostbyname", side_effect=resolve), \
             contextlib.redirect_stdout(output):
            mscb.mscb_scan_udp()

        text = output.getvalue()
        self.assertIn("Found MSCB007, PV 5, Rev. 0x1234", text)
        self.assertIn("UT 1d 01h 01m 01s", text)


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