#!/usr/bin/env python3
"""R9 teaching checks: not a TCP stack, network client, or service deduplication protocol.
Python 3 standard library only. Execute to print the independently checkable ledger.
"""
from dataclasses import dataclass, field
from itertools import permutations
from fractions import Fraction
import ipaddress
import json

WIRE = bytes.fromhex('00 03 43 41 54 00 02 4f 4b')
CID = 'R9-one-use'

@dataclass
class Receiver:
    length: int
    width: int = 3
    cid: str = CID
    prefix: int = 0
    delivered: bytearray = field(default_factory=bytearray)
    seen: dict = field(default_factory=dict)
    work: int = 0
    ack_sends: int = 0

    def __post_init__(self):
        if self.length < 0 or self.width < 1:
            raise ValueError("R9 requires N >= 0 and m >= 1")

    def receive(self, cid, start, data):
        if cid != self.cid:
            return None
        # The receiver knows length and segmentation, not the sender's byte string.
        if self.width < 1 or start < 0 or start >= self.length or start % self.width:
            raise ValueError('invalid fixed segment boundary')
        end = min(start + self.width, self.length)
        if len(data) != end - start:
            raise ValueError("invalid fixed segment length")
        self.work += len(data)
        for offset, byte in enumerate(data, start):
            if offset in self.seen and self.seen[offset] != byte:
                raise ValueError("conflicting duplicate outside R9 model")
            self.seen[offset] = byte
        while self.prefix in self.seen:
            self.delivered.append(self.seen[self.prefix])
            self.prefix += 1
        assert bytes(self.delivered) == bytes(self.seen[k] for k in range(self.prefix))
        assert len(self.seen) <= self.length
        self.ack_sends += 1
        return (self.cid, self.prefix)

@dataclass
class Sender:
    source: bytes
    width: int = 3
    cid: str = CID
    una: int = 0
    nxt: int = 0
    data_sends: int = 0

    def __post_init__(self):
        if self.width < 1:
            raise ValueError("R9 requires m >= 1")

    def send_new(self):
        if self.nxt == len(self.source):
            return None
        start = self.nxt
        self.nxt = min(start + self.width, len(self.source))
        self.data_sends += 1
        return (self.cid, start, self.source[start:self.nxt])

    def timeout(self):
        if self.una == self.nxt:
            return None
        self.data_sends += 1
        return (self.cid, self.una,
                self.source[self.una:min(self.una + self.width, len(self.source))])

    def ack(self, ack):
        if ack is not None:
            cid, position = ack
            if cid == self.cid and self.una < position <= self.nxt:
                self.una = position
        assert 0 <= self.una <= self.nxt <= len(self.source)

class Parser:
    def __init__(self, maximum=16):
        self.maximum = maximum
        self.header = bytearray()
        self.body = bytearray()
        self.need = None
        self.error = False
        self.max_pending = 0

    def feed(self, data):
        if self.error:
            raise ValueError('terminal parse error')
        out = []
        for byte in data:
            if self.need is None:
                self.header.append(byte)
                if len(self.header) == 2:
                    self.need = int.from_bytes(self.header, 'big')
                    self.header.clear()
                    if self.need > self.maximum:
                        self.error = True
                        raise ValueError('length limit')
                    if self.need == 0:
                        out.append(b'')
                        self.need = None
            else:
                self.body.append(byte)
                if len(self.body) == self.need:
                    out.append(bytes(self.body))
                    self.body.clear()
                    self.need = None
            self.max_pending = max(self.max_pending, len(self.header) + len(self.body))
            assert len(self.body) <= self.maximum
        return out

    def eof(self):
        if self.error or self.header or self.need is not None:
            raise ValueError('truncated or invalid stream')
        return 'frame boundary'

    def state(self):
        return ('Header', self.header.hex()) if self.need is None else (
            'Body', self.need, self.body.decode('ascii'))

def split_by_mask(data, mask):
    start = 0
    for i in range(1, len(data)):
        if mask & (1 << (i - 1)):
            yield data[start:i]
            start = i
    yield data[start:]

def outcome(events, deadline=1000):
    """Classify the client report at deadline. success means a timely completed-attempt
    response; it does NOT promise exactly-once effects or no remaining attempts.
    A no_effect event is certified for the WHOLE logical request.
    A refusal of only one retry is deliberately represented as attempt_refused.
    """
    may_have_effect = False
    state = 'not_sent'
    for t, event in sorted(events):
        if t >= deadline:
            break  # boundary policy: only observations before D are timely
        if state in ('success', 'definite_failure'):
            break
        if event == 'send':
            may_have_effect = True
            state = 'pending'
        elif event == 'success':
            state = 'success'
        elif event == 'no_effect':
            state = 'definite_failure'
        elif event == 'local_reject' and not may_have_effect:
            state = 'definite_failure'
        elif event in ('disconnect', 'attempt_refused') and may_have_effect:
            state = 'unknown'
    if state in ('pending', 'unknown'):
        return 'unknown'
    if state == 'not_sent':
        return 'definite_failure'
    return state

def main():
    s, r = Sender(WIRE), Receiver(len(WIRE))
    packets = [s.send_new() for _ in range(3)]
    rows = []
    for event, packet, deliver_ack in (
        ('P1 arrives', packets[1], True),
        ('P0 lost', None, False),
        ('P2 arrives', packets[2], True),
        ('retransmitted P0 arrives; ACK lost', 'timeout', False),
        ('network duplicate P1 arrives', packets[1], True),
    ):
        if packet == "timeout":
            packet = s.timeout()
        ack = r.receive(*packet) if packet else None
        if deliver_ack:
            s.ack(ack)
        assert s.una <= r.prefix
        rows.append({'event': event, 'receiver_r': r.prefix, 'sender_u': s.una,
                     'ack': ack[1] if ack else None})
    assert (s.una, r.prefix, s.data_sends) == (9, 9, 4)
    assert bytes(r.delivered) == WIRE
    s.ack((CID, 3)); s.ack((CID, 99)); s.ack(('old-flow', 9))
    assert s.una == 9
    assert r.receive('old-flow', 0, WIRE[:3]) is None
    assert Sender(b'').una == len(Receiver(0).delivered) == 0

    # Finite safety checks: every segment arrival order, duplicates, one dropped
    # initial segment, and later successful retransmission. Not a proof of liveness.
    safety_schedule_count = 0
    for order in permutations(range(3)):
        for dropped in (None, 0, 1, 2):
            sr, rr = Sender(WIRE), Receiver(len(WIRE))
            pp = [sr.send_new() for _ in range(3)]
            for i in order:
                if i != dropped:
                    for _ in range(2):
                        sr.ack(rr.receive(*pp[i]))
                        assert sr.una <= rr.prefix
            while sr.una < len(WIRE):
                sr.ack(rr.receive(*sr.timeout()))
            assert bytes(rr.delivered) == WIRE
            safety_schedule_count += 1

    parser = Parser()
    frame_rows = []
    recovered = []
    for chunk in (WIRE[:1], WIRE[1:4], WIRE[4:]):
        output = parser.feed(chunk)
        recovered.extend(output)
        frame_rows.append({'chunk_hex': chunk.hex(' '), 'state': parser.state(),
                           'output': [x.decode() for x in output]})
    assert recovered == [b'CAT', b'OK']
    parser.eof()
    partition_count = 0
    for mask in range(1 << (len(WIRE)-1)):
        p = Parser()
        found = []
        for chunk in split_by_mask(WIRE, mask):
            found.extend(p.feed(chunk))
        p.eof()
        assert found == [b'CAT', b'OK']
        partition_count += 1
    for end in range(len(WIRE) + 1):
        p = Parser(); p.feed(WIRE[:end])
        expected_boundary = end in (0, 5, 9)
        try:
            p.eof()
            assert expected_boundary
        except ValueError:
            assert not expected_boundary
    p = Parser(); assert p.feed(bytes.fromhex('0000000158')) == [b'', b'X']
    p.eof()
    p = Parser()
    try:
        p.feed(b'\xff\xff')
        raise AssertionError('oversize accepted')
    except ValueError:
        assert p.error and not p.body

    # Short writes preserve a prefix and a suffix; EAGAIN consumes nothing.
    full = b'ABCDEFGH'; sent = b''; pending = full; partial_rows = []
    for count in (3, None, 2, 3):
        if count is not None:
            sent += pending[:count]; pending = pending[count:]
        assert sent + pending == full
        partial_rows.append({'returned': 'EAGAIN' if count is None else count,
                             'pending': pending.decode()})
    assert sent == full and not pending

    # Exact rational arithmetic avoids decimal-to-binary rounding in the ledger.
    srtt, var = Fraction(4,5), Fraction(2,5)
    rtos = [srtt + 4*var]
    for sample in (Fraction(1), Fraction(7,10)):
        var = Fraction(3,4)*var + Fraction(1,4)*abs(srtt-sample)
        srtt = Fraction(7,8)*srtt + Fraction(1,8)*sample
        rtos.append(max(Fraction(1), srtt + max(Fraction(1,100), 4*var)))
    assert rtos == [Fraction(12,5), Fraction(89,40), Fraction(127,64)]

    routes = [('0.0.0.0/0','A'),('10.0.0.0/8','B'),('10.2.0.0/16','C'),
              ('10.2.3.0/24','D'),('10.2.3.128/25','E')]
    destinations = ['10.2.3.200','10.2.3.20','10.2.4.1','10.9.0.1','203.0.113.9']
    selected = []
    for dest in destinations:
        matches = [(ipaddress.ip_network(net).prefixlen, hop) for net,hop in routes
                   if ipaddress.ip_address(dest) in ipaddress.ip_network(net)]
        selected.append(max(matches)[1])
    assert selected == ['E','D','C','B','A']
    payload_size, mtu, header_size = 3000, 1500, 20
    chunk_limit = 8*((mtu-header_size)//8)
    assert mtu >= 68 and chunk_limit > 0
    fragments = []
    offset = 0
    while offset < payload_size:
        length = min(payload_size-offset, chunk_limit)
        more = int(offset+length < payload_size)
        fragments.append((offset,length,offset//8,more))
        offset += length
    fragment_total = sum(length+header_size for _,length,_,_ in fragments)
    assert sum(length for _,length,_,_ in fragments) == 3000
    assert all(offset * 8 == start and length+20 <= 1500
               for start,length,offset,_ in fragments)
    assert sum(length+20 for _,length,_,_ in fragments) == 3060
    joint_budget = max(0,min(1000,800)-(1600-1000))
    assert joint_budget == 200
    reno_threshold = max(6000/2,2*1000)
    assert reno_threshold == 3000
    cwnd, count = 4000, 0
    for ack_bytes in (1000,1000,1000,1000):
        count += ack_bytes
        if count >= cwnd:
            cwnd += 1000; count = 0
    assert (cwnd,count) == (5000,0)

    timeline = [('DNS',100),('connect',200),('send',50),('wait',200),
                ('backoff',150),('reconnect',100),('send',50)]
    elapsed = 0; deadline_rows = []
    for label,duration in timeline:
        elapsed += duration
        deadline_rows.append({'stage':label,'elapsed_ms':elapsed,'remaining_ms':max(0,1000-elapsed)})
    assert elapsed == 850
    assert min(400, 1000-elapsed) == 150
    assert max(0,1000-1010) == 0
    server_runs = {}
    schedules = {
        'request_lost': [(350,'send'),(360,'request_drop')],
        'executed_response_lost': [(350,'send'),(400,'apply'),(410,'response_drop')],
        'reconnected_late_response': [(350,'send'),(400,'apply'),(410,'response_drop'),
            (550,'disconnect'),(800,'reconnect'),(850,'send'),(900,'apply'),(1030,'success')],
        'success_with_other_attempt_later': [(350,'send'),(400,'apply'),(420,'send'),
            (450,'success'),(600,'apply'),(610,'response_drop')],
    }
    observations = {}
    for name, events in schedules.items():
        balance, applications = 100, 0
        observed, states = [], []
        for time, event in events:
            if event == 'apply':
                balance += 10
                applications += 1
            if event in ('send','disconnect','success'):
                observed.append((time,event))
            states.append({'time':time,'event':event,'balance':balance})
        observations[name] = observed
        server_runs[name] = {'applications':applications,'balance':balance,
                            'client_observations':observed,'server_trace':states}
    assert [x['balance'] for x in server_runs.values()] == [100,110,120,120]
    assert observations['request_lost'] == observations['executed_response_lost']
    observations.update({
        'retry_refusal_does_not_resolve_first': [(350,'send'),(550,'disconnect'),(850,'attempt_refused')],
        'timely_success': [(350,'send'),(900,'success')],
        'certified_logical_reject': [(350,'send'),(900,'no_effect')],
        'pre_send_validation': [(0,'local_reject')],
    })
    outcomes = {name:outcome(events) for name,events in observations.items()}
    assert outcomes == {
        'request_lost':'unknown', 'executed_response_lost':'unknown',
        'reconnected_late_response':'unknown', 'success_with_other_attempt_later':'success',
        'retry_refusal_does_not_resolve_first':'unknown', 'timely_success':'success',
        'certified_logical_reject':'definite_failure', 'pre_send_validation':'definite_failure'}
    assert server_runs['success_with_other_attempt_later']['server_trace'][3]['balance'] == 110
    assert server_runs['success_with_other_attempt_later']['balance'] == 120
    assert 10 + min(120,30) == 40  # negative DNS expiration
    assert 60-20 == 40 and 60-59 == 1  # remaining positive TTL
    print(json.dumps({'r9_trace':rows,'r9_data_sends':s.data_sends,'r9_ack_sends':r.ack_sends,
        'finite_safety_schedules':safety_schedule_count,'frame_trace':frame_rows,
        'all_short_read_partitions':partition_count,'short_writes':partial_rows,'rto_seconds':[float(x) for x in rtos],'rto_exact':[str(x) for x in rtos],
        'lpm_hops':selected,'ipv4_fragments':fragments,'ipv4_total_bytes':fragment_total,
        'receive_and_congestion_budget':joint_budget,'reno_threshold':reno_threshold,
        'reno_avoidance_after_four_acks':cwnd,'deadline_trace':deadline_rows,
        'outcomes_at_D':outcomes,'server_runs':server_runs,'status':'all assertions passed',
        'scope':'Finite teaching traces and parser splits, not TCP conformance or liveness verification'},
        ensure_ascii=False,indent=2))

if __name__ == '__main__':
    main()
