#!/usr/bin/env python3
"""CS07b finite teaching model. Python 3.10+, standard library only.

Run: python foundations-multiplexed-stream-checker.py [--compact]
No network, TLS, complete QUIC/HTTP implementation, congestion recovery or I/O.
Explicit checks remain enabled under python -O. Packet commits are atomic.
"""
from __future__ import annotations
import argparse
import copy
from dataclasses import dataclass, field
import itertools
import json

MAX_OFFSET = (1 << 62) - 1

class ModelError(Exception):
    pass
class ContextError(ModelError):
    pass
class FlowControlError(ModelError):
    pass
class FinalSizeError(ModelError):
    pass
class ContentConflict(ModelError):
    pass
class FormatError(ModelError):
    pass
class ConnectionFailed(ModelError):
    pass

def require(ok, message):
    if not ok:
        raise AssertionError(message)

def number(value, label, maximum=MAX_OFFSET):
    if type(value) is not int or not 0 <= value <= maximum:
        raise FormatError(f"{label}: nonnegative bounded integer required")
    return value

@dataclass(frozen=True)
class Stream:
    sid: int
    offset: int
    data: bytes
    fin: bool = False

@dataclass(frozen=True)
class Reset:
    sid: int
    final_size: int

@dataclass(frozen=True)
class Packet:
    pn: int
    frames: tuple[Stream | Reset, ...]
    connection: str = "demo-connection"
    direction: str = "c2s"
    space: str = "application"

@dataclass
class StreamState:
    limit: int
    capacity: int = 64
    high: int = 0
    final_size: int | None = None
    frontier: int = 0
    output: bytearray = field(default_factory=bytearray)
    stopped: bool = False
    reset: bool = False
    complete: bool = False
    slots: list[int | None] = field(init=False)

    def __post_init__(self):
        self.slots = [None] * self.capacity

class Receiver:
    """One authenticated connection, one direction, one PN space.

    Already-admitted 0/4 are c2s halves of client-initiated bidi streams.
    Fixed stream set and <=64 array caps are classroom resource assumptions.
    Protocol violations are terminal here; caller creates a fresh branch to test
    another counterexample. Snapshots retain output for audit, not production.
    """
    def __init__(self, limits=None, connection_limit=64, *, direction="c2s"):
        limits = {0: 64, 4: 64} if limits is None else limits
        self.connection = "demo-connection"
        if direction not in ("c2s", "s2c"):
            raise ContextError("unsupported direction")
        self.direction, self.space = direction, "application"
        self.streams = {}
        for sid, limit in limits.items():
            number(sid, "stream id")
            self.streams[sid] = StreamState(number(limit, "limit", 64))
        self.limit = number(connection_limit, "MAX_DATA")
        self.total = 0
        self.seen = set()
        self.failed = None

    def _live(self):
        if self.failed is not None:
            raise ConnectionFailed(self.failed)

    def _stream(self, sid):
        number(sid, "stream id")
        if sid not in self.streams:
            raise ContextError("stream is not admitted in this finite model")
        return self.streams[sid]

    def increase_connection(self, value):
        """Synchronized grant: receiver advertised it and sender learned it.

        Control propagation is omitted. This is NOT an inbound MAX_DATA
        changing this endpoint's same-direction receiving budget.
        """
        self._live()
        self.limit = max(self.limit, number(value, "MAX_DATA"))
        self.check()

    def increase_stream(self, sid, value):
        """Same synchronized-grant convention as increase_connection."""
        self._live()
        state = self._stream(sid)
        state.limit = max(state.limit, number(value, "MAX_STREAM_DATA", state.capacity))
        self.check()

    def stop_reading(self, sid):
        """Local choice; represents requesting STOP_SENDING, not delivering it."""
        self._live()
        state = self._stream(sid)
        if state.complete or state.reset:
            return
        state.stopped = True
        for i in range(state.frontier, state.capacity):
            state.slots[i] = None
        self.check()

    def _shape(self, frame):
        if not isinstance(frame, (Stream, Reset)):
            raise FormatError("unsupported frame")
        state = self._stream(frame.sid)
        if isinstance(frame, Reset):
            number(frame.final_size, "final size")
        else:
            number(frame.offset, "offset")
            if type(frame.data) is not bytes or type(frame.fin) is not bool:
                raise FormatError("bytes and boolean FIN required")
            number(frame.offset + len(frame.data), "end offset")
        return state

    def _account(self, state, end, final=None):
        # Final-size checks precede flow-control checks in this classroom model.
        if state.final_size is not None and end > state.final_size:
            raise FinalSizeError("data end exceeds recorded final size")
        if final is not None:
            if final < state.high:
                raise FinalSizeError("final size below prior highest end")
            if state.final_size is not None and final != state.final_size:
                raise FinalSizeError("final size changed")
        high = max(state.high, end)
        delta = high - state.high
        if high > state.limit or self.total + delta > self.limit:
            raise FlowControlError("stream or connection credit exceeded")
        # Stream limits never exceed the allocated classroom capacity.
        require(high <= state.capacity, "checked credit must fit finite array")
        state.high, self.total = high, self.total + delta
        if final is not None:
            state.final_size = final
        return delta

    def _frame(self, frame):
        state = self._shape(frame)
        if isinstance(frame, Reset):
            delta = self._account(state, frame.final_size, frame.final_size)
            if not state.complete:
                state.reset = True
                for i in range(state.frontier, state.capacity):
                    state.slots[i] = None
            return {"sid": frame.sid, "delta": delta, "new": "", "event": "reset"}
        end = frame.offset + len(frame.data)
        delta = self._account(state, end, end if frame.fin else None)
        if state.reset or state.stopped:
            return {"sid": frame.sid, "delta": delta, "new": "", "event": "discard-accounted"}
        for j, value in enumerate(frame.data, frame.offset):
            prior = state.slots[j]
            if prior is not None and prior != value:
                raise ContentConflict("same stream offset carries different bytes")
            state.slots[j] = value
        old = state.frontier
        while state.frontier < state.capacity and state.slots[state.frontier] is not None:
            state.output.append(state.slots[state.frontier])
            state.frontier += 1
        if state.final_size is not None and state.frontier == state.final_size:
            state.complete = True
        return {"sid": frame.sid, "delta": delta,
                "new": bytes(state.output[old:]).decode("latin1"), "event": "stream"}

    def receive(self, packet):
        self._live()
        if not isinstance(packet, Packet):
            raise FormatError("Packet required")
        if (packet.connection, packet.direction, packet.space) != (self.connection, self.direction, self.space):
            # Demultiplexing rejects other contexts; not a peer protocol error.
            raise ContextError("wrong connection/direction/packet number space")
        number(packet.pn, "packet number")
        if type(packet.frames) is not tuple:
            raise FormatError("frames must be a tuple")
        try:
            for frame in packet.frames:
                self._shape(frame)
            if packet.pn in self.seen:
                return {"pn": packet.pn, "duplicate_packet": True, "frames": [], "snapshot": self.snapshot()}
            trial = copy.deepcopy(self)
            events = [trial._frame(f) for f in packet.frames]
            trial.seen.add(packet.pn)
            trial.check()
        except ModelError as error:
            self.failed = type(error).__name__
            raise
        self.__dict__.update(trial.__dict__)
        return {"pn": packet.pn, "duplicate_packet": False, "frames": events, "snapshot": self.snapshot()}

    def check(self):
        require(self.total == sum(s.high for s in self.streams.values()), "connection sum drift")
        require(self.total <= self.limit, "connection credit invariant")
        for state in self.streams.values():
            require(0 <= state.frontier <= state.high <= state.limit <= state.capacity, "stream bounds")
            require(len(state.output) == state.frontier, "output/frontier mismatch")
            require(all(state.slots[j] == state.output[j] for j in range(state.frontier)), "prefix content")
            if not (state.reset or state.stopped):
                require(state.frontier == state.capacity or state.slots[state.frontier] is None, "unadvanced prefix")
            if state.final_size is not None:
                require(state.high == state.final_size, "known final size must be accounted")
            if state.complete:
                require(state.final_size == state.frontier and not state.reset, "normal EOF invariant")

    def snapshot(self):
        streams = {}
        for sid, s in self.streams.items():
            covered = [i for i, value in enumerate(s.slots) if value is not None]
            streams[str(sid)] = {"high": s.high, "frontier": s.frontier,
                "final_size": s.final_size, "limit": s.limit,
                "output": bytes(s.output).decode("latin1"), "covered": covered,
                "stopped": s.stopped, "reset": s.reset, "complete": s.complete}
        return {"streams": streams, "total": self.total, "limit": self.limit,
                "remaining": self.limit - self.total, "seen": sorted(self.seen)}


def packet(pn, *frames, **context):
    return Packet(pn, tuple(frames), **context)


def expect_error(kind, action):
    try:
        action()
    except kind as error:
        return {"error": type(error).__name__, "message": str(error)}
    except Exception as error:
        raise AssertionError(f"expected {kind.__name__}, got {type(error).__name__}") from error
    raise AssertionError(f"expected {kind.__name__}, no error")


def primary_trace():
    r = Receiver()
    sequence = [packet(11, Stream(4, 0, b"XY")), packet(12, Stream(0, 2, b"cd")),
                packet(13, Stream(0, 0, b"ab")), packet(14, Stream(0, 0, b"ab"))]
    trace = [r.receive(p) for p in sequence]
    fronts = [(row["snapshot"]["streams"]["0"]["frontier"], row["snapshot"]["streams"]["4"]["frontier"]) for row in trace]
    require(fronts == [(0, 2), (0, 2), (4, 2), (4, 2)], "main fronts")
    same_pn = r.receive(sequence[0])
    require(same_pn["duplicate_packet"], "equal PN was not suppressed")
    before = r.snapshot()
    late = r.receive(packet(10, Stream(0, 0, b"ab")))
    require(not late["duplicate_packet"], "late unseen lower PN falsely classified duplicate")
    require(r.total == before["total"] and all(x["new"] == "" for x in late["frames"]), "late original duplicated bytes")
    return {"lost": {"pn": 10, "sid": 0, "offset": 0, "data": "ab"}, "arrivals": trace, "equal_packet_duplicate": same_pn, "late_original": late}


def h2_data(sid, data):
    number(sid, "HTTP/2 stream id", (1 << 31) - 1)
    require(sid != 0, "HTTP/2 DATA cannot use connection stream 0")
    return len(data).to_bytes(3, "big") + b"\x00\x00" + sid.to_bytes(4, "big") + data


def tcp_h2_trace():
    # HEADERS and connection setup already consumed; supports only DATA subset.
    encoded = [h2_data(1, b"ab"), h2_data(3, b"XY"), h2_data(1, b"cd")]
    wire = b"".join(encoded)
    slots, tcp_frontier, parsed = [None] * len(wire), 0, 0
    outputs = {1: b"", 3: b""}
    history = []
    for start, block in [(11, encoded[1]), (22, encoded[2]), (0, encoded[0]), (0, encoded[0])]:
        for j, byte in enumerate(block, start):
            if slots[j] is not None:
                require(slots[j] == byte, "TCP duplicate conflict")
            slots[j] = byte
        while tcp_frontier < len(slots) and slots[tcp_frontier] is not None:
            tcp_frontier += 1
        available = bytes(slots[:tcp_frontier])
        while parsed + 9 <= len(available):
            size = int.from_bytes(available[parsed:parsed + 3], "big")
            if parsed + 9 + size > len(available):
                break
            require(available[parsed + 3:parsed + 5] == b"\0\0", "only unpadded DATA accepted")
            sid = int.from_bytes(available[parsed + 5:parsed + 9], "big")
            require(sid in outputs, "pre-opened stream required")
            outputs[sid] += available[parsed + 9:parsed + 9 + size]
            parsed += 9 + size
        history.append({"received_range": [start, start + len(block)], "tcp_frontier": tcp_frontier,
                        "outputs": {str(k): v.decode() for k, v in outputs.items()}})
    require([x["tcp_frontier"] for x in history] == [0, 0, 33, 33], "TCP hole model")
    require(outputs == {1: b"abcd", 3: b"XY"}, "H2 final content")
    return {"mapping": {"QUIC 0": "HTTP/2 1", "QUIC 4": "HTTP/2 3"},
            "data_frames_hex": [p.hex() for p in encoded], "trace": history}


def credit_start(connection_limit=8):
    r = Receiver({0: 6, 4: 6}, connection_limit)
    # Only two actual bytes on s0; gap [0,2), H0=4. s4 has three bytes.
    r.receive(packet(1, Stream(0, 2, b"cd"), Stream(4, 0, b"XYZ")))
    require((r.streams[0].high, r.streams[4].high, r.total) == (4, 3, 7), "4+3 fixture")
    return r


def credit_trace():
    r = credit_start()
    rows = [{"event": "4+3 initial", "snapshot": r.snapshot()}]
    r.receive(packet(2, Stream(0, 2, b"cd")))
    rows.append({"event": "same offsets resent", "snapshot": r.snapshot()})
    r.increase_connection(10)
    rows.append({"event": "MAX_DATA 10", "snapshot": r.snapshot()})
    r.receive(packet(3, Reset(0, 6)))
    rows.append({"event": "RESET 0 final_size 6", "snapshot": r.snapshot()})
    r.receive(packet(4, Stream(0, 4, b"ef")))
    rows.append({"event": "late old bytes [4,6)", "snapshot": r.snapshot()})
    r.increase_connection(8)
    r.increase_stream(4, 3)
    rows.append({"event": "smaller credits ignored", "snapshot": r.snapshot()})
    require(r.total == 9 and r.limit == 10 and r.streams[4].limit == 6, "final accounting")
    require(r.streams[0].output == b"" and r.streams[0].reset, "reset gap not fabricated")
    return rows


def counterexamples():
    out = {}
    r = credit_start()
    before = r.snapshot()
    out["reset_exceeds_connection"] = expect_error(FlowControlError, lambda: r.receive(packet(2, Reset(0, 6))))
    require(r.snapshot() == before, "failed reset modified state")
    out["protocol_error_is_terminal"] = expect_error(ConnectionFailed, lambda: r.receive(packet(3, Stream(4, 0, b"XYZ"))))
    r = credit_start()
    out["reset_below_highest_end"] = expect_error(FinalSizeError, lambda: r.receive(packet(2, Reset(0, 3))))
    r = credit_start(20)
    out["stream_limit_independent"] = expect_error(FlowControlError, lambda: r.receive(packet(2, Stream(0, 6, b"g"))))
    r = credit_start()
    before = r.snapshot()
    out["combined_packet_credit"] = expect_error(FlowControlError, lambda: r.receive(packet(2, Stream(0, 4, b"e"), Stream(4, 3, b"Z"))))
    require(r.snapshot() == before, "multi-frame failure partially committed")
    out["combined_packet_credit"]["atomic_snapshot_preserved"] = True
    r = Receiver()
    r.receive(packet(1, Stream(0, 0, b"ab")))
    before = r.snapshot()
    out["conflicting_overlap"] = expect_error(ContentConflict, lambda: r.receive(packet(2, Stream(4, 0, b"XY"), Stream(0, 1, b"X"))))
    require(r.snapshot() == before, "content-conflict partial delivery escaped")
    r = Receiver()
    before = r.snapshot()
    out["within_packet_conflict"] = expect_error(ContentConflict, lambda: r.receive(packet(2, Stream(0, 0, b"ab"), Stream(0, 1, b"X"))))
    require(r.snapshot() == before, "same-packet conflict partially committed")
    r = Receiver()
    r.receive(packet(1, Stream(0, 0, bytes([0, 255, 128]))))
    require(bytes(r.streams[0].output) == bytes([0, 255, 128]), "arbitrary byte payload corrupted")
    out["non_ascii_bytes"] = {"output_hex": bytes(r.streams[0].output).hex()}
    for name, frame in [("data_beyond_final", Stream(0, 4, b"x")), ("fin_disagrees", Stream(0, 0, b"", True)), ("reset_disagrees", Reset(0, 5))]:
        r = Receiver()
        r.receive(packet(1, Stream(0, 2, b"cd", True)))
        out[name] = expect_error(FinalSizeError, lambda f=frame: r.receive(packet(2, f)))
    r = Receiver()
    out["unknown_final_size"] = {"final_size": r.streams[0].final_size, "complete": r.streams[0].complete}
    require(not r.streams[0].complete, "no data is not EOF")
    for name, context in [("wrong_direction", {"direction": "s2c"}), ("wrong_connection", {"connection": "other"}), ("wrong_packet_space", {"space": "handshake"})]:
        r = Receiver()
        before = r.snapshot()
        out[name] = expect_error(ContextError, lambda c=context: r.receive(packet(0, Stream(0, 0, b"a"), **c)))
        require(r.snapshot() == before, "wrong context modified state")
    r = Receiver()
    out["end_overflow"] = expect_error(FormatError, lambda: r.receive(packet(0, Stream(0, MAX_OFFSET, b"x"))))
    r = Receiver()
    out["negative_offset"] = expect_error(FormatError, lambda: r.receive(packet(0, Stream(0, -1, b"a"))))
    return out


def termination_and_stop():
    fin = Receiver()
    rows = [fin.receive(packet(1, Stream(0, 2, b"cd", True)))]
    require(fin.streams[0].final_size == 4 and not fin.streams[0].complete, "FIN incorrectly bypassed hole")
    rows.append(fin.receive(packet(2, Stream(0, 0, b"ab"))))
    require(fin.streams[0].complete and fin.streams[0].output == b"abcd", "FIN completion")
    rows.append(fin.receive(packet(3, Reset(0, 4))))
    require(fin.streams[0].complete and not fin.streams[0].reset, "late consistent reset revoked completed output")
    empty = Receiver()
    empty.receive(packet(1, Stream(0, 0, b"", True)))
    require(empty.streams[0].complete and empty.total == 0, "empty FIN")
    stopped = Receiver({0: 6, 4: 6}, 10)
    stopped.receive(packet(1, Stream(0, 2, b"cd")))
    stopped.stop_reading(0)
    stop_rows = [{"event": "stop_reading", "snapshot": stopped.snapshot()}]
    stop_rows.append(stopped.receive(packet(2, Stream(0, 4, b"e"))))
    stop_rows.append(stopped.receive(packet(3, Reset(0, 6))))
    stop_rows.append(stopped.receive(packet(4, Stream(0, 0, b"abcdef", True))))
    require(stopped.total == 6 and stopped.streams[0].output == b"", "stop failed to account or discard")
    reverse = Receiver(direction="s2c")
    reverse.receive(packet(3, Stream(0, 0, b"OK", True), direction="s2c"))
    require(reverse.streams[0].complete and reverse.total == 2, "reverse half damaged")
    return {"fin_before_hole": rows, "empty_fin": empty.snapshot(), "stop_then_reset": stop_rows, "independent_reverse_half": reverse.snapshot()}


def multipacket_migration():
    r = Receiver()
    # Packet 20 carrying s0:ab and s4:X is lost. Y alone cannot fill s4.
    rows = [r.receive(packet(21, Stream(0, 2, b"cd"), Stream(4, 1, b"Y")))]
    require((r.streams[0].frontier, r.streams[4].frontier) == (0, 0), "mixed packet hole")
    rows.append(r.receive(packet(22, Stream(4, 0, b"X"))))
    require((r.streams[0].frontier, r.streams[4].frontier) == (0, 2), "cross-stream barrier invented")
    rows.append(r.receive(packet(23, Stream(0, 0, b"ab"))))
    require((r.streams[0].frontier, r.streams[4].frontier) == (4, 2), "mixed packet repair")
    return {"lost_packet": 20, "lost_frames": [{"sid": 0, "offset": 0, "data": "ab"}, {"sid": 4, "offset": 0, "data": "X"}], "trace": rows}


def exhaustive_tests():
    originals = {0: b"abcd", 4: b"XY"}
    segments = [Stream(0, 0, b"ab"), Stream(0, 2, b"cd"), Stream(4, 0, b"XY")]
    # Every 5-event word over the three fragments, then repair all missing data.
    count = 0
    for word in itertools.product(range(3), repeat=5):
        r = Receiver()
        for pn, index in enumerate(word + (0, 1, 2)):
            p = packet(pn, segments[index])
            r.receive(p)
            if pn % 2 == 0:
                before = r.snapshot()
                r.receive(p)
                require(r.snapshot() == before, "duplicate packet changes state")
            for sid, original in originals.items():
                state = r.streams[sid]
                require(bytes(state.output) == original[:state.frontier], "exhaustive prefix safety")
        require(r.streams[0].output == b"abcd" and r.streams[4].output == b"XY", "repair didn't complete")
        count += 1
    # All orderings of four split/overlapping, same-value intervals.
    intervals = [Stream(0, 0, b"abc"), Stream(0, 1, b"bcd"), Stream(4, 0, b"X"), Stream(4, 1, b"Y")]
    overlap_count = 0
    for order in itertools.permutations(intervals):
        r = Receiver()
        for pn, frame in enumerate(order):
            r.receive(packet(pn, frame))
        require(r.streams[0].output == b"abcd" and r.streams[4].output == b"XY", "overlap resegmentation")
        overlap_count += 1
    # Independent arithmetic oracle for sparse one-byte ends and final sizes.
    ledger_cases = 0
    for h0, h4, cap in itertools.product(range(5), range(5), range(10)):
        for final in range(7):
            if h0 + h4 > cap:
                continue
            r = Receiver({0: 6, 4: 6}, cap)
            for sid, h in [(0, h0), (4, h4)]:
                if h:
                    r.receive(packet(sid, Stream(sid, h - 1, b"x")))
            expected = FinalSizeError if final < h0 else FlowControlError if final + h4 > cap else None
            if expected:
                expect_error(expected, lambda: r.receive(packet(99, Reset(0, final))))
            else:
                r.receive(packet(99, Reset(0, final)))
                require(r.total == final + h4, "arithmetic reset oracle")
            ledger_cases += 1
    return {"five_event_words_then_repair": count, "overlap_orders": overlap_count, "reset_arithmetic_cases": ledger_cases,
            "scope": "finite tests, not a proof of real QUIC conformance or infinite liveness"}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--compact", action="store_true", help="omit verbose per-event ledgers")
    args = parser.parse_args()
    result = {"model": "CS07b fixed authenticated connection, one direction and PN space",
        "primary": primary_trace(), "tcp_http2_comparison": tcp_h2_trace(),
        "independent_credit_4_plus_3": credit_trace(), "counterexamples": counterexamples(),
        "termination": termination_and_stop(), "multi_frame_migration": multipacket_migration(),
        "finite_test_counts": exhaustive_tests(), "status": "passed"}
    if args.compact:
        result = {"status": result["status"], "finite_test_counts": result["finite_test_counts"],
                  "counterexamples": sorted(result["counterexamples"]),
                  "main_frontiers": [[0, 2], [0, 2], [4, 2], [4, 2]], "credit_after_reset": 9}
    print(json.dumps(result, ensure_ascii=False, indent=2))

if __name__ == "__main__":
    main()
