#!/usr/bin/env python3
"""Bounded relay from the native recorder stream to WebRTC packet batches."""

from __future__ import annotations

from dataclasses import dataclass
import importlib.util
import io
from pathlib import Path
import struct
import sys
import threading


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)

LENGTH = struct.Struct("<I")
FRAME_DEPENDENCY = struct.Struct("<I")
MAX_NATIVE_PACKET = protocol.HEADER.size + protocol.MAX_FRAGMENT
RESOURCE_BATCH_LIMIT = 2 * 1024 * 1024


@dataclass
class RelayStats:
    native_packets: int = 0
    resource_messages: int = 0
    resource_batches: int = 0
    frame_messages: int = 0
    frames_sent: int = 0
    frames_dropped: int = 0
    input_bytes: int = 0
    output_bytes: int = 0


class RelayError(RuntimeError):
    pass


class CommandStreamRelay:
    """Reassemble, batch, compress, and fragment one recorder session.

    ``send_reliable`` and ``send_frame`` receive a complete packet list. They
    must return false without sending any packet when their channel cannot
    accept the complete list. This keeps frame loss atomic.
    """

    def __init__(self, send_reliable, send_frame, *, resource_batch_limit=None):
        self.send_reliable = send_reliable
        self.send_frame = send_frame
        self.resource_batch_limit = resource_batch_limit or RESOURCE_BATCH_LIMIT
        self.native = protocol.Reassembler()
        self.session = 0
        self.sequence = 1
        self.last_resource_sequence = 0
        self.resource_payloads: list[bytes] = []
        self.resource_bytes = 0
        self.stats = RelayStats()
        self.stop_event = threading.Event()

    def _packets(self, kind, payload, *, frame=0, compress=True,
                 extra_flags=protocol.Flag(0)):
        flags = protocol.Flag(extra_flags)
        if compress:
            compression_flags, payload = protocol.compress_payload(payload)
            flags |= compression_flags
        packets = protocol.fragment(
            kind, payload, session=self.session, sequence=self.sequence,
            frame=frame, flags=flags)
        self.sequence += len(packets)
        return packets

    def _send(self, callback, packets, *, frame=False):
        if callback(packets):
            self.stats.output_bytes += sum(len(packet) for packet in packets)
            return True
        if frame:
            self.stats.frames_dropped += 1
            return False
        raise RelayError("reliable command channel exceeded its bounded queue")

    def flush_resources(self):
        if not self.resource_payloads:
            return
        payload = b"".join(self.resource_payloads)
        packets = self._packets(protocol.Kind.RESOURCE, payload)
        self._send(self.send_reliable, packets)
        self.last_resource_sequence = self.sequence - 1
        self.stats.resource_batches += 1
        self.resource_payloads.clear()
        self.resource_bytes = 0

    def feed_packet(self, packet: bytes):
        self.stats.native_packets += 1
        self.stats.input_bytes += len(packet)
        completed = self.native.push(packet)
        if completed is None:
            return
        envelope, payload = completed
        if not self.session:
            self.session = envelope.session
        if envelope.session != self.session:
            raise RelayError("recorder session changed without reconnect")
        if envelope.kind is protocol.Kind.RESOURCE:
            if (self.resource_payloads
                    and self.resource_bytes + len(payload)
                    > self.resource_batch_limit):
                self.flush_resources()
            self.resource_payloads.append(payload)
            self.resource_bytes += len(payload)
            self.stats.resource_messages += 1
            return
        if envelope.kind is protocol.Kind.FRAME:
            self.flush_resources()
            self.stats.frame_messages += 1
            payload = FRAME_DEPENDENCY.pack(self.last_resource_sequence) + payload
            packets = self._packets(protocol.Kind.FRAME, payload,
                                    frame=envelope.frame,
                                    extra_flags=envelope.flags
                                    & protocol.Flag.KEYFRAME)
            if self._send(self.send_frame, packets, frame=True):
                self.stats.frames_sent += 1
            return
        self.flush_resources()
        packets = self._packets(envelope.kind, payload, frame=envelope.frame,
                                compress=False)
        self._send(self.send_reliable, packets)

    def pump(self, stream: io.BufferedReader):
        while not self.stop_event.is_set():
            raw_size = stream.read(LENGTH.size)
            if not raw_size:
                break
            if len(raw_size) != LENGTH.size:
                raise RelayError("truncated recorder packet length")
            size, = LENGTH.unpack(raw_size)
            if size < protocol.HEADER.size or size > MAX_NATIVE_PACKET:
                raise RelayError("invalid recorder packet length")
            packet = stream.read(size)
            if len(packet) != size:
                raise RelayError("truncated recorder packet")
            self.feed_packet(packet)
        self.flush_resources()

    def stop(self):
        self.stop_event.set()


def main():
    import argparse

    parser = argparse.ArgumentParser()
    parser.add_argument("capture", type=Path)
    args = parser.parse_args()
    reliable, frames = [], []
    relay = CommandStreamRelay(
        lambda packets: not reliable.append(packets),
        lambda packets: not frames.append(packets))
    with args.capture.open("rb") as stream:
        relay.pump(stream)
    print(relay.stats)
    print(f"reliable_batches={len(reliable)} frame_batches={len(frames)}")


if __name__ == "__main__":
    main()
