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


ROOT = Path(__file__).resolve().parent


def load(name, path):
    spec = importlib.util.spec_from_file_location(name, path)
    module = importlib.util.module_from_spec(spec)
    sys.modules[name] = module
    spec.loader.exec_module(module)
    return module


protocol = load("test_w3cs_protocol", ROOT / "protocol.py")
relay_module = load("test_w3cs_relay", ROOT / "relay.py")


def native_stream(messages):
    output = bytearray()
    sequence = 1
    for kind, frame, payload in messages:
        packets = protocol.fragment(kind, payload, session=77,
                                    sequence=sequence, frame=frame)
        sequence += len(packets)
        for packet in packets:
            output.extend(struct.pack("<I", len(packet)))
            output.extend(packet)
    return io.BytesIO(output)


def unpack_batches(batches):
    reassembler = protocol.Reassembler()
    result = []
    for batch in batches:
        for packet in batch:
            completed = reassembler.push(packet)
            if completed:
                envelope, payload = completed
                result.append((envelope, protocol.decompress_payload(
                    payload, kind=envelope.kind, flags=envelope.flags)))
    return result


class RelayTest(unittest.TestCase):
    def test_resources_are_batched_before_frame(self):
        sent = []

        def reliable(packets):
            sent.append(("resource", packets))
            return True

        def frame(packets):
            sent.append(("frame", packets))
            return True

        resource_a = protocol.encode_records([protocol.Record(7, 0, b"a" * 5000)])
        resource_b = protocol.encode_records([protocol.Record(7, 0, b"b" * 5000)])
        frame_payload = protocol.encode_records([protocol.Record(33, 0, b"draw")])
        relay = relay_module.CommandStreamRelay(reliable, frame)
        relay.pump(native_stream([
            (protocol.Kind.RESOURCE, 0, resource_a),
            (protocol.Kind.RESOURCE, 0, resource_b),
            (protocol.Kind.FRAME, 5, frame_payload),
        ]))
        self.assertEqual([kind for kind, _ in sent], ["resource", "frame"])
        decoded_resources = unpack_batches([sent[0][1]])
        self.assertEqual(decoded_resources[0][1], resource_a + resource_b)
        decoded_frames = unpack_batches([sent[1][1]])
        dependency, = struct.unpack_from("<I", decoded_frames[0][1])
        self.assertGreater(dependency, 0)
        self.assertEqual(decoded_frames[0][1][4:], frame_payload)
        self.assertEqual(relay.stats.frames_sent, 1)

    def test_frame_drop_is_atomic(self):
        calls = []
        relay = relay_module.CommandStreamRelay(
            lambda packets: True,
            lambda packets: calls.append(len(packets)) and False)
        payload = protocol.encode_records([
            protocol.Record(33, 0, bytes(range(251)) * 100)])
        relay.pump(native_stream([(protocol.Kind.FRAME, 9, payload)]))
        self.assertEqual(len(calls), 1)
        self.assertGreater(calls[0], 0)
        self.assertEqual(relay.stats.frames_dropped, 1)
        self.assertEqual(relay.stats.frames_sent, 0)

    def test_reliable_backpressure_fails_closed(self):
        relay = relay_module.CommandStreamRelay(
            lambda packets: False, lambda packets: True)
        with self.assertRaises(relay_module.RelayError):
            relay.pump(native_stream([(
                protocol.Kind.RESOURCE, 0,
                protocol.encode_records([protocol.Record(7, 0, b"resource")]))]))


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