#!/usr/bin/env python3
"""Finite educational models for ACK, recovery timers, strict buckets, and DRR.
No network I/O. Exact Fraction arithmetic; run with normal Python, not -O.
"""
from bisect import bisect_right
from collections import deque
from fractions import Fraction as F
from itertools import product
import json
import random

LIMIT = (1 << 62) - 1

def nat(value):
    if type(value) is not int or not 0 <= value <= LIMIT:
        raise ValueError('field is not a QUIC variable-length integer')
    return value

def decode_ack(largest, first, count, pairs):
    """Returns closed intervals; never allocates a possibly huge set of PNs."""
    largest, first, count = map(nat, (largest, first, count))
    if count != len(pairs):
        raise ValueError('truncated or surplus ACK ranges')
    low = largest - first
    if low < 0:
        raise ValueError('FRAME_ENCODING_ERROR: first range underflow')
    out = [(low, largest)]
    for gap, length in pairs:
        gap, length = nat(gap), nat(length)
        high = low - gap - 2
        low = high - length
        if low < 0:
            raise ValueError('FRAME_ENCODING_ERROR: subsequent underflow')
        out.append((low, high))
    return out

def expand(ranges):
    # Caller uses only the small finite experiment, never arbitrary network input.
    return {p for low, high in ranges for p in range(low, high + 1)}

def encode_ack(received):
    if not received:
        raise ValueError('an ACK needs at least one acknowledged PN')
    ordered = sorted({nat(x) for x in received}, reverse=True)
    high = low = ordered[0]
    intervals = []
    for p in ordered[1:]:
        if p == low - 1:
            low = p
        else:
            intervals.append((low, high))
            high = low = p
    intervals.append((low, high))
    pairs = [(intervals[i - 1][0] - h - 2, h - l)
             for i, (l, h) in enumerate(intervals) if i]
    return intervals[0][1], intervals[0][1] - intervals[0][0], len(pairs), pairs

def apply_ack(sent_by_space, acked_by_space, space, fields):
    ranges = decode_ack(*fields)
    got = expand(ranges)
    if not got <= sent_by_space[space]:
        raise ValueError('PROTOCOL_VIOLATION: acknowledged unsent packet')
    # All validation precedes state mutation.
    acked_by_space[space].update(got)

class ReceivedHistory:
    def __init__(self):
        self.minimum = 0
        self.seen = set()
        self.largest_ever = None
    def receive(self, pn):
        nat(pn)
        if pn < self.minimum or pn in self.seen:
            return False
        self.seen.add(pn)
        self.largest_ever = pn if self.largest_ever is None else max(self.largest_ever,pn)
        return True
    def retire_below(self, minimum):
        if minimum < self.minimum:
            raise ValueError('receive floor cannot decrease')
        self.minimum = minimum
        self.seen = {pn for pn in self.seen if pn >= minimum}
        # Largest processed PN survives even if every ACK range is removed.

def rtt_sample_eligible(frame_largest, newly_acked, ack_eliciting):
    return frame_largest in newly_acked and bool(newly_acked & ack_eliciting)

class Rtt:
    def __init__(self, max_delay=25):
        self.d = F(max_delay)
        self.minimum = self.s = self.v = self.latest = None
    def sample(self, raw, delay):
        raw, delay = F(raw), F(delay)
        if raw <= 0 or delay < 0:
            raise ValueError('invalid time sample')
        self.latest = raw
        if self.minimum is None:
            self.minimum = self.s = raw
            self.v = raw / 2
            return raw
        self.minimum = min(self.minimum, raw)
        delay = min(delay, self.d)  # Handshake is confirmed in this model.
        adjusted = raw - delay if raw - delay >= self.minimum else raw
        old_s = self.s
        self.v = 3 * self.v / 4 + abs(old_s - adjusted) / 4
        self.s = 7 * old_s / 8 + adjusted / 8
        return adjusted
    def pto(self, granularity=1, count=0):
        return (self.s + max(4 * self.v, F(granularity)) + self.d) * 2 ** count

def loss_scan(inflight, largest_acked, now, smoothed, latest, granularity=1):
    """Precondition: original sent PNs are contiguous (no sender-induced gaps).
    Missing inflight records may already have been acknowledged or removed.
    """
    delay = max(F(9, 8) * max(F(smoothed), F(latest)), F(granularity))
    lost, pending = [], []
    for pn, sent_at in inflight.items():
        if largest_acked is None or pn >= largest_acked:
            continue
        if pn <= largest_acked - 3 or F(now) - sent_at >= delay:
            lost.append(pn)
        else:
            pending.append(F(sent_at) + delay)
    return lost, min(pending) if pending else None

class Bucket:
    def __init__(self, capacity, rate):
        self.capacity, self.rate = F(capacity), F(rate)
        if self.capacity <= 0 or self.rate <= 0:
            raise ValueError('positive bucket capacity and rate required')
        self.tokens, self.time = self.capacity, F(0)
    def refill(self, time):
        time = F(time)
        if time < self.time:
            raise ValueError('time cannot move backwards')
        self.tokens = min(self.capacity, self.tokens + self.rate * (time - self.time))
        self.time = time
    def police(self, time, length):
        length = F(length)
        if length <= 0:
            raise ValueError('positive packet length required')
        self.refill(time)
        if self.tokens < length:
            return False
        self.tokens -= length
        return True
    def shape(self, arrivals):
        """Arrivals sorted by time; FIFO, instantaneous downstream acceptance."""
        out = []
        previous_arrival = F(0)
        for arrival, length in arrivals:
            arrival, length = F(arrival), F(length)
            if arrival < previous_arrival:
                raise ValueError('arrival order is not chronological')
            previous_arrival = arrival
            if not 0 < length <= self.capacity:
                raise ValueError('packet cannot fit strict bucket')
            self.refill(max(self.time, arrival))
            if self.tokens < length:
                self.refill(self.time + (length - self.tokens) / self.rate)
            assert self.tokens >= length
            self.tokens -= length
            out.append((self.time, length, self.tokens))
        return out

class Drr:
    def __init__(self, quantum):
        if not quantum or any(type(q) is not int or q <= 0 for q in quantum):
            raise ValueError('positive integer quantum required')
        self.q = list(quantum)
        self.queues = [deque() for _ in quantum]
        self.credit = [0] * len(quantum)
        self.active = deque()
        self.present = [False] * len(quantum)
        self.served = [0] * len(quantum)
    def enqueue(self, flow, lengths):
        lengths = list(lengths)
        if any(type(x) is not int or x <= 0 for x in lengths):
            raise ValueError('positive integer packet lengths required')
        if not lengths:
            return
        if not self.present[flow]:
            assert not self.queues[flow] and self.credit[flow] == 0
            self.active.append(flow)
            self.present[flow] = True
        self.queues[flow].extend(lengths)
    def visit(self):
        if not self.active:
            return None
        i = self.active.popleft()
        self.credit[i] += self.q[i]
        lengths = []
        while self.queues[i] and self.queues[i][0] <= self.credit[i]:
            length = self.queues[i].popleft()
            lengths.append(length)
            self.credit[i] -= length
            self.served[i] += length
        if self.queues[i]:
            assert 0 <= self.credit[i] < self.queues[i][0]
            self.active.append(i)
        else:
            self.credit[i] = 0
            self.present[i] = False
        assert len(set(self.active)) == len(self.active)
        assert {j for j, queue in enumerate(self.queues) if queue} == set(self.active)
        return {'flow': i, 'sent': lengths, 'credit': self.credit[i],
                'served': self.served.copy()}

def reject(call):
    try:
        call()
    except ValueError:
        return
    raise AssertionError('bad input was accepted')

def ack_tests():
    cases = 0
    for mask in range(1, 1 << 12):
        values = {p for p in range(12) if mask >> p & 1}
        fields = encode_ack(values)
        intervals = decode_ack(*fields)
        assert expand(intervals) == values
        assert all(intervals[i-1][0] > intervals[i][1]+1 for i in range(1,len(intervals)))
        assert sum(hi-lo+1 for lo, hi in intervals) == len(values)
        cases += 1
    malformed = 0
    for largest, first, gap, length in product(range(8), repeat=4):
        # Independent arithmetic condition for a two-range word:
        # its lowest PN is largest-first-gap-2-length.
        feasible = first <= largest and largest-first-gap-2-length >= 0
        try:
            ranges = decode_ack(largest, first, 1, [(gap,length)])
        except ValueError:
            assert not feasible
        else:
            assert feasible
            assert expand(ranges) == set(range(largest-first, largest+1)) | set(
                range(largest-first-gap-2-length, largest-first-gap-1))
        malformed += 1
    sent = {'Application': set(range(13)), 'Handshake': {12}}
    acknowledged = {s:set() for s in sent}
    fields = encode_ack({2,3,4,7,8,12})
    apply_ack(sent, acknowledged, 'Application', fields)
    saved = {s:v.copy() for s,v in acknowledged.items()}
    apply_ack(sent, acknowledged, 'Application', encode_ack({12}))
    assert acknowledged == saved
    apply_ack(sent, acknowledged, 'Handshake', encode_ack({12}))
    assert acknowledged['Application'] == saved['Application']
    saved = {s:v.copy() for s,v in acknowledged.items()}
    for bad in [(3,4,0,[]),(3,0,1,[(2,0)]),(12,0,1,[]),(13,0,0,[])]:
        reject(lambda:apply_ack(sent, acknowledged, 'Application', bad))
        assert acknowledged == saved
    sent = {'Application':{0,2}}; a = {'Application':set()}
    apply_ack(sent,a,'Application',encode_ack({2}))
    reject(lambda:apply_ack(sent,a,'Application',encode_ack({1})))
    reject(lambda:decode_ack(-1,0,0,[]))
    reject(lambda:decode_ack(LIMIT+1,0,0,[]))
    h=ReceivedHistory();ever=set();floor=0;rng=random.Random(5101)
    for _ in range(20000):
        if rng.randrange(6)==0:
            floor+=rng.randrange(4)
            h.retire_below(floor)
        else:
            pn=rng.randrange(max(20,floor+20))
            expected=pn>=floor and pn not in ever
            assert h.receive(pn)==expected
            if expected:ever.add(pn)
        assert h.seen=={p for p in ever if p>=floor}
        assert h.largest_ever==(max(ever) if ever else None)
    reject(lambda:h.retire_below(floor-1))
    return {'sets':cases,'two_range_inputs':malformed,'example_fields':fields,
            'receive_history_events':20000}

def rtt_and_loss_tests():
    for largest_new, some_new_ack_eliciting in product([False,True],repeat=2):
        newly={7} if largest_new else set()
        if some_new_ack_eliciting:newly.add(5)
        assert rtt_sample_eligible(7,newly,{5})==(largest_new and some_new_ack_eliciting)
    # The largest newly acknowledged packet may itself be non-ack-eliciting.
    assert rtt_sample_eligible(7,{5,7},{5})
    assert not rtt_sample_eligible(7,{5},{5})
    est = Rtt(); rows=[]
    for raw, delay in [(100,10),(140,25),(90,20)]:
        used=est.sample(raw,delay)
        rows.append([raw,est.minimum,used,est.s,est.v])
    assert rows[-1] == [90,F(90),F(90),F(6425,64),F(1085,32)]
    assert est.pto() == F(16705,64)
    other=Rtt();other.sample(100,10);other.sample(140,40)
    assert (other.s,other.v)==(F(815,8),F(165,4))
    rng=random.Random(5102);sample_cases=0
    for _ in range(1000):
        d=rng.randint(0,30);e=Rtt(d); raw_history=[]; s=v=None
        for j in range(10):
            raw=F(rng.randint(1,200));delay=F(rng.randint(0,60));raw_history.append(raw)
            m=min(raw_history)
            if j==0: s=raw;v=raw/2
            else:
                adjusted=raw
                limited=min(delay,d)
                if raw-limited>=m:adjusted-=limited
                v=(3*v+abs(s-adjusted))/4
                s=(7*s+adjusted)/8
            e.sample(raw,delay)
            assert (e.minimum,e.s,e.v)==(m,s,v)
            sample_cases+=1
    inflight={pn:F(10*(pn-20)) for pn in range(20,26)}
    del inflight[23]
    states=[]
    for now in [F(110),F(245,2),F(265,2)]:
        lost,deadline=loss_scan(inflight,23,now,100,80)
        for pn in lost: del inflight[pn]
        states.append({'time':now,'lost':lost,'next_loss_time':deadline,'inflight':sorted(inflight)})
    assert [r['lost'] for r in states]==[[20],[21],[22]]
    assert [r['next_loss_time'] for r in states]==[F(245,2),F(265,2),None]
    assert inflight=={24:F(40),25:F(50)}
    assert loss_scan(inflight,23,10000,100,80)==([],None)
    assert loss_scan({20:F(0)},None,10000,100,80)==([],None)
    shuffled={22:F(20),20:F(0),25:F(50),21:F(10),24:F(40)}
    got,deadline=loss_scan(shuffled,23,125,100,80)
    assert set(got)=={20,21} and deadline==F(265,2)
    assert got==[20,21]  # Output follows encountered lost entries, not a sort.
    reverse={21:F(10),25:F(50),20:F(0),22:F(20)}
    got,deadline=loss_scan(reverse,23,125,100,80)
    assert got==[21,20] and deadline==F(265,2)
    p=F(100)+4*20+25
    first=50+p
    before=inflight.copy();inflight[26]=first
    assert all(inflight[k]==v for k,v in before.items())
    assert first==255 and first+2*p==665
    thresholds=0
    for latest, smooth, now, acked in product([1,8,40],[1,8,40],range(0,80,5),range(6)):
        flight={p:F(p*4) for p in range(6) if p!=acked and p*4<=now}
        lost,_=loss_scan(flight,acked,now,smooth,latest)
        expected=[]
        for pn,st in flight.items():
            if pn>=acked:continue
            gap=acked-pn
            age=F(now)-st
            if gap>=3 or (age>=1 and 8*age>=9*max(latest,smooth)):
                expected.append(pn)
        assert lost==expected
        thresholds+=1
    return {'rtt_rows':rows,'rtt_samples':sample_cases,'base_pto_ms':est.pto(),
            'loss_trace':states,'threshold_snapshots':thresholds,
            'probe_times_ms':[first,first+2*p]}

def envelope(events, capacity, rate):
    # Check all event-bounded intervals, including immediately before each burst.
    # Maximal byte counts at fixed endpoints occur at these positions.
    batches={}
    for time,length in events:batches[time]=batches.get(time,F(0))+length
    ordered=sorted(batches)
    checks=0
    for i,start in enumerate(ordered):
        total=F(0)
        for end in ordered[i:]:
            total+=batches[end]
            assert total<=capacity+rate*(end-start)
            checks+=1
    return checks

def bucket_tests():
    arrivals=[(0,1500),(0,1000),(0,500)]
    shaped=Bucket(2000,1000).shape(arrivals)
    assert [x[0] for x in shaped]==[0,F(1,2),1]
    p=Bucket(2000,1000)
    policed=[p.police(t,n) for t,n in arrivals]
    assert policed==[True,False,True] and p.tokens==0
    reject(lambda:Bucket(1000,1000).shape([(0,1500)]))
    combined=Bucket(2000,1000).shape([(0,n) for n in [500,500,500,500,1500,500]])
    assert [x[0] for x in combined]==[0,0,0,0,F(3,2),2]
    rng=random.Random(5103);intervals=0
    for _ in range(1200):
        capacity=rng.randint(1,30);rate=rng.randint(1,10)
        arrival=F(0);arrivals=[]
        for j in range(15):
            arrival+=F(rng.randint(0,3),2)
            arrivals.append((arrival,rng.randint(1,capacity)))
        output=Bucket(capacity,rate).shape(arrivals)
        assert all(output[i][0]>=arrivals[i][0] for i in range(len(output)))
        assert all(output[i][0]<=output[i+1][0] for i in range(len(output)-1))
        # Replay actual admission events with a meter, separately from waiting logic.
        replay=Bucket(capacity,rate)
        for time,length,_ in output: assert replay.police(time,length)
        intervals+=envelope([(t,n) for t,n,_ in output],F(capacity),F(rate))
        meter=Bucket(capacity,rate);accepted=[]
        for t,n in arrivals:
            prior=meter.tokens
            prior_time=meter.time
            ok=meter.police(t,n)
            expected=min(F(capacity),prior+rate*(t-prior_time))
            assert meter.tokens==expected-(n if ok else 0)
            if ok:accepted.append((t,n))
        intervals+=envelope(accepted,F(capacity),F(rate))
    return {'shaper':shaped,'policer_accepted':policed,'combined':combined,
            'random_sequences':1200,'event_interval_checks':intervals}

def drr_tests():
    d=Drr([1000,1000]);d.enqueue(0,[500]*30);d.enqueue(1,[1500]*30)
    trace=[d.visit() for _ in range(6)]
    assert [x['served'] for x in trace[1::2]]==[[1000,0],[2000,1500],[3000,3000]]
    assert [x['credit'] for x in trace[1::2]]==[1000,500,0]
    selected=[(x['flow'],length) for x in trace for length in x['sent']]
    assert selected[:6]==[(0,500),(0,500),(0,500),(0,500),(1,1500),(0,500)]
    weighted=Drr([1000,2000]);weighted.enqueue(0,[500]*30);weighted.enqueue(1,[1500]*30)
    wt=[weighted.visit() for _ in range(6)]
    assert wt[-1]['served']==[3000,6000]
    empty=Drr([1000]);empty.enqueue(0,[500]);e=empty.visit()
    assert e['credit']==0 and not empty.active
    empty.enqueue(0,[1500]);e=empty.visit()
    assert not e['sent'] and e['credit']==1000
    assert empty.visit()['sent']==[1500]
    tiny=Drr([1]);tiny.enqueue(0,[1500]);idle=0
    for _ in range(1500):
        record=tiny.visit()
        if not record['sent']:idle+=1
    assert idle==1499 and tiny.served==[1500]
    rng=random.Random(5104);visits=0
    for case in range(1000):
        n=rng.randint(2,5);quanta=[rng.randint(1,15) for _ in range(n)]
        engine=Drr(quanta);prefixes=[];maxima=[]
        for i in range(n):
            lengths=[rng.randint(1,10) for _ in range(400)]
            maxima.append(max(lengths));engine.enqueue(i,lengths)
            totals=[0]
            for value in lengths:totals.append(totals[-1]+value)
            prefixes.append(totals)
        for round_number in range(1,21):
            for i in range(n):
                event=engine.visit();assert event['flow']==i
                budget=round_number*quanta[i]
                at=bisect_right(prefixes[i],budget)-1
                assert at<len(prefixes[i])-1 # All cases remain backlogged.
                assert engine.served[i]==prefixes[i][at]
                assert engine.credit[i]==budget-prefixes[i][at]
                assert 0<=engine.credit[i]<maxima[i]
                visits+=1
            for i in range(n):
                for j in range(i):
                    difference=abs(F(engine.served[i],quanta[i])-F(engine.served[j],quanta[j]))
                    assert difference<max(F(maxima[i],quanta[i]),F(maxima[j],quanta[j]))
    return {'three_round_trace':trace,'weighted_final':wt[-1]['served'],
            'small_quantum_empty_visits':idle,'backlogged_cases':1000,
            'prefix_sum_oracle_visits':visits,'selected_first_six':selected[:6]}

def serial(value):
    if isinstance(value,F):
        return int(value) if value.denominator==1 else str(value)
    raise TypeError(type(value).__name__)

def main():
    if not __debug__:
        raise RuntimeError('Run without -O: this checker requires active assertions.')
    result={'status':'PASS','scope':'Finite teaching models; not a QUIC implementation or certification',
            'ack':ack_tests(),'recovery':rtt_and_loss_tests(),
            'bucket':bucket_tests(),'drr':drr_tests()}
    print(json.dumps(result,default=serial,ensure_ascii=False,indent=2))

if __name__=='__main__':main()
