#!/usr/bin/env python3
"""Validate and summarize a length-prefixed W3CS recorder stream."""

from __future__ import annotations

import argparse
from collections import Counter
import importlib.util
from pathlib import Path
import statistics
import sys
import zlib


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)


def read_exact(source, size):
    data = source.read(size)
    if len(data) != size:
        raise protocol.ProtocolError("truncated length-prefixed message")
    return data


def percentile(values, fraction):
    if not values:
        return 0
    ordered = sorted(values)
    return ordered[min(len(ordered) - 1, int(len(ordered) * fraction))]


def inspect(path: Path):
    reassembler = protocol.Reassembler()
    messages = Counter()
    records = Counter()
    frame_bytes = Counter()
    completed_frames = 0
    compressed_frames = []
    packets = 0
    wire_bytes = 0
    with path.open("rb") as source:
        while True:
            raw_size = source.read(4)
            if not raw_size:
                break
            if len(raw_size) != 4:
                raise protocol.ProtocolError("truncated packet length")
            size = int.from_bytes(raw_size, "little")
            if size < protocol.HEADER.size or size > (
                    protocol.HEADER.size + protocol.MAX_FRAGMENT):
                raise protocol.ProtocolError(f"invalid packet size {size}")
            packet = read_exact(source, size)
            envelope = protocol.decode_envelope(packet)
            packets += 1
            wire_bytes += size + 4
            messages[envelope.kind.name] += 1
            if envelope.kind is protocol.Kind.FRAME:
                frame_bytes[envelope.frame] += size + 4
            complete = reassembler.push(packet)
            if complete is None:
                continue
            assembled, payload = complete
            if assembled.kind not in (protocol.Kind.FRAME,
                                      protocol.Kind.RESOURCE):
                continue
            decoded = protocol.decode_records(payload)
            records.update(record.opcode for record in decoded)
            if assembled.kind is protocol.Kind.FRAME:
                completed_frames += 1
                compressed_frames.append(len(zlib.compress(payload, level=1)))

    sizes = list(frame_bytes.values())
    print(f"packets={packets} wire_bytes={wire_bytes}")
    print(f"messages={dict(messages)} completed_frames={completed_frames}")
    print(f"records={dict(sorted(records.items()))}")
    if sizes:
        print("frame_wire_bytes "
              f"mean={statistics.fmean(sizes):.0f} "
              f"p50={percentile(sizes, .50)} "
              f"p95={percentile(sizes, .95)} "
              f"max={max(sizes)}")
    if compressed_frames:
        print("frame_zlib1_bytes "
              f"mean={statistics.fmean(compressed_frames):.0f} "
              f"p50={percentile(compressed_frames, .50)} "
              f"p95={percentile(compressed_frames, .95)} "
              f"max={max(compressed_frames)}")


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("stream", type=Path)
    args = parser.parse_args()
    inspect(args.stream)


if __name__ == "__main__":
    main()
