import importlib.util
from pathlib import Path
import sys
import unittest


ROOT = Path(__file__).resolve().parent
SPEC = importlib.util.spec_from_file_location("w3cs_protocol", ROOT / "protocol.py")
protocol = importlib.util.module_from_spec(SPEC)
sys.modules[SPEC.name] = protocol
SPEC.loader.exec_module(protocol)


class ProtocolTest(unittest.TestCase):
    def test_envelope_round_trip(self):
        original = protocol.Envelope(protocol.Kind.HELLO, protocol.Flag.LAST,
                                     7, 9, 0, 0, 1, b"hello")
        self.assertEqual(original, protocol.decode_envelope(
            protocol.encode_envelope(original)))

    def test_corruption_is_rejected(self):
        packet = bytearray(protocol.fragment(
            protocol.Kind.FRAME, b"draw", session=1, sequence=2, frame=3)[0])
        packet[-1] ^= 1
        with self.assertRaises(protocol.ProtocolError):
            protocol.decode_envelope(packet)

    def test_out_of_order_reassembly(self):
        payload = bytes(range(251)) * 200
        packets = protocol.fragment(protocol.Kind.FRAME, payload, session=4,
                                    sequence=10, frame=99)
        reassembler = protocol.Reassembler()
        result = None
        for packet in reversed(packets):
            result = reassembler.push(packet) or result
        self.assertIsNotNone(result)
        self.assertEqual(payload, result[1])

    def test_conflicting_duplicate_is_rejected(self):
        packets = protocol.fragment(protocol.Kind.FRAME,
                                    b"x" * (protocol.MAX_FRAGMENT + 1),
                                    session=4, sequence=10, frame=99)
        reassembler = protocol.Reassembler()
        reassembler.push(packets[0])
        changed = bytearray(packets[0])
        changed[-1] ^= 1
        # Repair the checksum so this exercises duplicate conflict detection.
        envelope = protocol.decode_envelope(packets[0])
        altered = protocol.encode_envelope(protocol.Envelope(
            envelope.kind, envelope.flags, envelope.session, envelope.sequence,
            envelope.frame, envelope.fragment_index, envelope.fragment_count,
            bytes(changed[protocol.HEADER.size:])))
        with self.assertRaises(protocol.ProtocolError):
            reassembler.push(altered)

    def test_frame_fragments_expire(self):
        now = [10.0]
        clock = lambda: now[0]
        packets = protocol.fragment(protocol.Kind.FRAME,
                                    b"x" * (protocol.MAX_FRAGMENT + 1),
                                    session=1, sequence=1, frame=1)
        reassembler = protocol.Reassembler(clock)
        self.assertIsNone(reassembler.push(packets[0]))
        now[0] += protocol.FRAME_TTL_SECONDS + 0.001
        reassembler.expire()
        self.assertFalse(reassembler._pending)

    def test_records_round_trip_and_unknown_opcode(self):
        records = [protocol.Record(1, 0, b"a"),
                   protocol.Record(250, 3, b"future")]
        self.assertEqual(records, protocol.decode_records(
            protocol.encode_records(records)))

    def test_resource_messages_do_not_share_reassembly_state(self):
        first = protocol.fragment(protocol.Kind.RESOURCE,
                                  b"a" * (protocol.MAX_FRAGMENT + 1),
                                  session=1, sequence=10)
        second = protocol.fragment(protocol.Kind.RESOURCE,
                                   b"b" * (protocol.MAX_FRAGMENT + 1),
                                   session=1, sequence=20)
        reassembler = protocol.Reassembler()
        self.assertIsNone(reassembler.push(first[0]))
        self.assertIsNone(reassembler.push(second[0]))
        self.assertEqual(b"a" * (protocol.MAX_FRAGMENT + 1),
                         reassembler.push(first[1])[1])
        self.assertEqual(b"b" * (protocol.MAX_FRAGMENT + 1),
                         reassembler.push(second[1])[1])

    def test_fragment_size_is_bounded(self):
        packets = protocol.fragment(protocol.Kind.RESOURCE,
                                    b"z" * (protocol.MAX_FRAGMENT * 3 + 5),
                                    session=1, sequence=1)
        self.assertEqual(4, len(packets))
        self.assertTrue(all(len(packet) <= protocol.HEADER.size
                            + protocol.MAX_FRAGMENT for packet in packets))

    def test_compressed_payload_round_trip(self):
        original = b"repeated wc3 state " * 1000
        flags, compressed = protocol.compress_payload(original)
        self.assertEqual(flags, protocol.Flag.COMPRESSED)
        self.assertEqual(original, protocol.decompress_payload(
            compressed, kind=protocol.Kind.FRAME, flags=flags))

    def test_compressed_payload_limit_is_checked_before_inflate(self):
        header = protocol.COMPRESSED_HEADER.pack(
            protocol.MAX_FRAME + 1, int(protocol.Compression.DEFLATE))
        with self.assertRaises(protocol.ProtocolError):
            protocol.decompress_payload(header + b"bad",
                                        kind=protocol.Kind.FRAME,
                                        flags=protocol.Flag.COMPRESSED)


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