#!/usr/bin/env python3
"""Explicit teaching contracts for RED, integer-tick CoDel, classic ECN feedback.
No sockets, external input, clock settings, files or actual network changes.
RED uses section 7's PREVIOUS unmarked-arrivals convention, not Figure 2's
preincrement counter. CoDel state order follows RFC8289 section 5, with an
explicit ceiling control interval; this independently expressed reference is
not a kernel implementation. ECN is an already-negotiated, in-order/no-data-loss
feedback component with sufficient receiver credit; it halves integer byte
windows using floor division and rejects reductions below the >=4*M input range.
Only JSON to stdout. Explicit checks remain enabled in optimized mode.
"""
from fractions import Fraction as F
from collections import deque
from dataclasses import dataclass, replace
from math import isqrt
from copy import deepcopy
from itertools import product
import json, random


def require(ok, why):
    if not ok:
        raise ValueError(why)


def reject(call):
    try:
        call()
    except ValueError as e:
        return str(e)
    raise RuntimeError('invalid input was accepted')


class RED:
    def __init__(self, capacity=5, weight=F(1, 2), low=2, high=4, max_p=F(1, 2)):
        require(type(capacity) is int and capacity > 0, 'positive packet capacity')
        self.weight, self.low, self.high, self.max_p = map(F, (weight, low, high, max_p))
        require(0 < self.weight <= 1 and 0 <= self.low < self.high and 0 < self.max_p <= 1, 'RED parameters')
        self.capacity = capacity; self.q = deque(); self.avg = F(0); self.k = 0
        self.now = 0; self.idle_tick = 0

    def clock(self, now):
        require(type(now) is int and now >= self.now, 'monotone integer tick')
        self.now = now

    @staticmethod
    def adjusted(base, k):
        require(0 <= base <= 1 and type(k) is int and k >= 0, 'probability/count')
        if not base:
            return F(0)
        denominator = 1-k*base
        return F(1) if denominator <= 0 else min(F(1), base/denominator)

    def arrival(self, ident, now, uniform, ecn=False):
        uniform = F(uniform); require(0 <= uniform < 1, 'uniform tape in [0,1)')
        self.clock(now); before = len(self.q); old_avg = self.avg; old_k = self.k
        if self.q:
            self.avg = (1-self.weight)*self.avg+self.weight*len(self.q)
            idle_slots = 0
        else:
            require(self.idle_tick is not None, 'empty queue has idle anchor')
            idle_slots = now-self.idle_tick
            self.avg *= (1-self.weight)**idle_slots
            self.idle_tick = now  # never decay the same interval twice
        if self.avg < self.low:
            base = pa = F(0); selected = False; self.k = 0; region = 'LOW'
        elif self.avg >= self.high:
            base = pa = F(1); selected = True; self.k = 0; region = 'HIGH'
        else:
            base = self.max_p*(self.avg-self.low)/(self.high-self.low)
            pa = self.adjusted(base, self.k); selected = uniform < pa
            self.k = 0 if selected else self.k+1; region = 'MIDDLE'
        if len(self.q) >= self.capacity:
            action = 'TAIL_DROP'
        elif selected and not ecn:
            action = 'EARLY_DROP'
        else:
            action = 'MARK_CE' if selected else 'ENQUEUE'
            self.q.append(ident); self.idle_tick = None
        require(0 <= self.avg <= self.capacity and len(self.q) <= self.capacity, 'RED bounded state')
        return dict(id=ident, time=now, q_before=before, avg_before=old_avg, avg=self.avg,
                    old_k=old_k, k=self.k, base=base, pa=pa, selected=selected,
                    region=region, uniform=uniform, action=action, idle_slots=idle_slots,
                    queue=list(self.q))  # optional O(q) audit snapshot

    def dequeue(self, now):
        self.clock(now)
        if not self.q:
            return None
        p = self.q.popleft()
        if not self.q:
            self.idle_tick = now
        return p


@dataclass(frozen=True)
class Packet:
    ident: str
    size: int
    arrival: int


class CoDel:
    def __init__(self, capacity=10000, maxpacket=1000, target=5, interval=100):
        require(all(type(v) is int and v > 0 for v in (capacity, maxpacket, target, interval)), 'positive CoDel parameters')
        require(capacity >= maxpacket, 'capacity covers maximum packet')
        self.capacity = capacity; self.maxpacket = maxpacket; self.target = target; self.interval = interval
        self.q = deque(); self.bytes = 0; self.first_above = None; self.dropping = False
        self.count = 0; self.lastcount = 0; self.drop_next = 0; self.now = 0

    def clock(self, now):
        require(type(now) is int and now >= self.now, 'monotone integer tick')
        self.now = now

    def spacing(self, count):
        require(type(count) is int and count >= 1, 'positive control count')
        return 1+isqrt((self.interval*self.interval-1)//count)

    def enqueue(self, ident, size, now):
        require(type(size) is int and 0 < size <= self.maxpacket, 'packet length outside contract')
        self.clock(now)
        if self.bytes+size > self.capacity:
            return 'TAIL_DROP'
        self.q.append(Packet(ident, size, now)); self.bytes += size
        return 'ENQUEUE'

    def candidate(self, now, observations):
        before = self.first_above
        if not self.q:
            self.first_above = None
            observations.append(dict(id=None, remaining_bytes=0, above_before=before,
                                     first_above=None, eligible=False, reason='EMPTY'))
            return None, False
        p = self.q.popleft(); self.bytes -= p.size; sojourn = now-p.arrival
        eligible = False
        if sojourn < self.target or self.bytes <= self.maxpacket:
            self.first_above = None; reason = 'LOW_DELAY_OR_BACKLOG'
        elif self.first_above is None:
            self.first_above = now+self.interval; reason = 'START_INTERVAL'
        else:
            eligible = now >= self.first_above; reason = 'INTERVAL_ELAPSED' if eligible else 'WAIT_INTERVAL'
        observations.append(dict(id=p.ident, sojourn=sojourn, remaining_bytes=self.bytes,
                                 above_before=before, first_above=self.first_above,
                                 eligible=eligible, reason=reason))
        return p, eligible

    def state(self):
        return dict(first_above=self.first_above, dropping=self.dropping, count=self.count,
                    lastcount=self.lastcount, drop_next=self.drop_next, bytes=self.bytes)

    def dequeue(self, now):
        self.clock(now); before = self.state(); observations = []; drops = []; decisions = []
        p, eligible = self.candidate(now, observations)
        if self.dropping:
            if not eligible:
                self.dropping = False
            while self.dropping and now >= self.drop_next:
                deadline = self.drop_next; old_count = self.count
                drops.append(p.ident); self.count += 1
                p, eligible = self.candidate(now, observations)
                if eligible:
                    self.drop_next = deadline+self.spacing(self.count)
                else:
                    self.dropping = False
                decisions.append(dict(kind='CATCH_UP', dropped=drops[-1], comparison_deadline=deadline,
                                      old_count=old_count, count=self.count, eligible_next=eligible,
                                      next_deadline=self.drop_next))
        elif eligible:
            drops.append(p.ident); p, eligible = self.candidate(now, observations)
            delta = self.count-self.lastcount; old_deadline = self.drop_next
            reuse = delta > 1 and now-old_deadline < 16*self.interval
            self.count = delta if reuse else 1
            self.dropping = True; self.drop_next = now+self.spacing(self.count); self.lastcount = self.count
            decisions.append(dict(kind='ENTER', dropped=drops[-1], delta=delta, old_deadline=old_deadline,
                                  reuse=bool(reuse), count=self.count, next_deadline=self.drop_next))
        return dict(time=now, before=before, observations=observations, drops=drops,
                    delivered=None if p is None else p.ident, decisions=decisions, after=self.state(),
                    queue=[x.ident for x in self.q])  # optional O(q) audit snapshot


def router_signal(codepoint, selected, room):
    require(codepoint in ('Not-ECT', 'ECT0', 'ECT1', 'CE'), 'unknown ECN codepoint')
    if not room:
        return 'DROP', None
    if selected and codepoint == 'Not-ECT':
        return 'DROP', None
    return 'FORWARD', 'CE' if selected else codepoint


class Receiver:
    def __init__(self):
        self.next = 0; self.echo = False

    def data(self, start, size, ce=False, cwr=False):
        require(type(start) is int and type(size) is int and start == self.next and size > 0,
                'receiver requires new in-order data')
        self.next += size
        if cwr:
            self.echo = False
        if ce:
            self.echo = True
        return dict(ack=self.next, ece=bool(self.echo))


class Sender:
    def __init__(self, mss=1000, window=8000):
        require(type(mss) is int and mss > 0 and type(window) is int and window >= mss, 'window model')
        self.mss = mss; self.window = window; self.u = self.n = 0
        self.guard = None; self.pending_cwr = False; self.reductions = 0

    def send(self, size):
        require(type(size) is int and 0 < size <= self.mss, 'positive segment at most MSS')
        require(self.n-self.u+size <= self.window, 'no congestion-window credit')
        p = dict(start=self.n, size=size, cwr=self.pending_cwr)
        self.n += size; self.pending_cwr = False
        return p

    def ack(self, ack, ece):
        require(type(ack) is int and 0 <= ack <= self.n, 'ACK beyond sent prefix')
        before = self.window; old_u = self.u; responded = False
        if ack > self.u:
            if ece and (self.guard is None or ack > self.guard):
                require(self.window >= 4*self.mss, 'handoff: small-window ECN needs complete TCP response')
                self.window //= 2; self.guard = self.n; self.pending_cwr = True
                self.reductions += 1; responded = True
            self.u = ack
        return dict(ack=ack, ece=bool(ece), old_u=old_u, u=self.u, sent=self.n,
                    window_before=before, window=self.window, responded=responded,
                    guard=self.guard, pending_cwr=self.pending_cwr, flight=self.n-self.u,
                    new_credit=max(0, self.window-(self.n-self.u)))


def main_trace():
    red = RED(); draws = [F(9,10)]*7; draws[4] = F(3,10)
    rr = [red.arrival(str(i+1), 0, draws[i]) for i in range(7)]
    require(list(red.q) == ['1','2','3','4','6'], 'RED queue trace')
    require([r['avg'] for r in rr] == [F(0),F(1,2),F(5,4),F(17,8),F(49,16),F(113,32),F(273,64)], 'RED averages')
    require(rr[4]['pa'] == F(17,47) and rr[5]['pa'] == F(49,128), 'RED exact adjusted probabilities')
    marked = RED(); marked_trace = [marked.arrival(str(i+1), 0, draws[i], ecn=True) for i in range(7)]
    c = CoDel()
    for ident in 'ABCDEFGHIJ': require(c.enqueue(ident, 1000, 0) == 'ENQUEUE', 'initial capacity')
    require(c.enqueue('overflow',1000,0)=='TAIL_DROP' and c.count==0, 'tail drop not controller drop')
    trace = [c.dequeue(now) for now in (5,104,105,405,406)]
    require(trace[3]['drops'] == list('EFGH') and trace[3]['delivered']=='I', 'catch-up candidates')
    require([d['comparison_deadline'] for d in trace[3]['decisions']] == [205,276,334,384], 'compare old deadline before each update')
    branches = []
    for start in (500,2500):
        fresh = deepcopy(c)
        for i in range(10): fresh.enqueue('N'+str(i),1000,start)
        first = fresh.dequeue(start+5); second = fresh.dequeue(start+105)
        branches.append(dict(arrival=start,first=first,second=second))
    require(branches[0]['second']['after']['count']==4 and branches[0]['second']['after']['drop_next']==655, 'recent count reuse')
    require(branches[1]['second']['after']['count']==1 and branches[1]['second']['after']['drop_next']==2705, 'long idle reset')
    sender=Sender();receiver=Receiver();packets=[sender.send(1000)for _ in range(8)];events=[]
    lost=receiver.data(**packets[0],ce=True); require(lost==dict(ack=1000,ece=True),'lost first echo')
    for i in range(1,5):
        ack=receiver.data(**packets[i]);events.append(sender.ack(**ack))
        if i<4:reject(lambda:sender.send(1000))
    ninth=sender.send(1000);require(ninth['cwr'] and ninth['start']==8000,'first new data carries CWR')
    for i in range(5,8):events.append(sender.ack(**receiver.data(**packets[i])))
    events.append(sender.ack(**receiver.data(**ninth,ce=True)))
    tenth=sender.send(1000);require(tenth['cwr'],'second response gets new CWR')
    events.append(sender.ack(**receiver.data(**tenth)))
    late=sender.ack(8000,True);require(sender.u==10000 and sender.reductions==2,'late ACK cannot react or retreat')
    return dict(red_drop=rr,red_ecn=marked_trace,codel=trace,codel_reentry=branches,
                ecn=dict(lost_ack=lost,events=events,ninth=ninth,tenth=tenth,late_ack=late))


def tests():
    rng=random.Random(141009);counts=dict(red_events=0,probability_intervals=0,codel_events=0,spacing_cases=0,receiver_sequences=0,sender_events=0)
    # Exact conditional probability masses, independent of random empirical rates.
    for n in range(1,201):
        alive=F(1);masses=[]
        for k in range(n):
            pa=RED.adjusted(F(1,n),k);masses.append(alive*pa);alive*=1-pa
        require(alive==0 and masses==[F(1,n)]*n,'uniform next-mark law')
        counts['probability_intervals']+=n
    for run in range(160):
        r=RED(capacity=9);items=[];t=0
        for step in range(100):
            t+=rng.randrange(4)
            if rng.randrange(3)==0:
                expected=items.pop(0)if items else None;require(r.dequeue(t)==expected,'RED literal FIFO')
            else:
                id=str(step);row=r.arrival(id,t,F(rng.randrange(1000),1000),ecn=bool(run%2))
                if row['action']in ('ENQUEUE','MARK_CE'):items.append(id)
                require(row['queue']==items and len(items)<=9,'RED ownership/capacity')
            counts['red_events']+=1
    for interval in range(1,160):
        for c in range(1,100):
            x=CoDel(interval=interval).spacing(c)
            require(c*x*x>=interval*interval and c*(x-1)*(x-1)<interval*interval,'minimal integer interval')
            counts['spacing_cases']+=1
    for run in range(220):
        c=CoDel(capacity=6000,target=5,interval=31);accepted=[];gone=[];t=0
        for step in range(120):
            t+=rng.randrange(12)
            if rng.randrange(2):
                id=str(step);size=rng.randrange(1,1001)
                if c.enqueue(id,size,t)=='ENQUEUE':accepted.append(id)
            else:
                row=c.dequeue(t);gone.extend(row['drops'])
                if row['delivered']is not None:gone.append(row['delivered'])
            require(gone+[p.ident for p in c.q]==accepted,'CoDel single removal FIFO conservation')
            require(c.bytes==sum(p.size for p in c.q) and 0<=c.bytes<=c.capacity,'CoDel byte account')
            counts['codel_events']+=1
    for bits in product((False,True),repeat=12):
        r=Receiver();expected=False
        for i in range(6):
            ce,cwr=bits[2*i:2*i+2];expected=ce or (expected and not cwr)
            require(r.data(i,1,ce,cwr)['ece']==expected,'receiver latch truth table')
        counts['receiver_sequences']+=1
    for run in range(200):
        s=Sender(mss=1,window=128);n=0;last_guard=-1;cuts=0
        for step in range(80):
            if rng.randrange(3)==0 and s.n-s.u<s.window:
                s.send(1)
            elif s.n:
                ack=rng.randrange(s.n+1);ece=bool(rng.randrange(2));old=s.u
                if ack>old and ece and(s.guard is None or ack>s.guard)and s.window<4:
                    reject(lambda:s.ack(ack,ece));continue
                row=s.ack(ack,ece)
                require(s.u>=old and s.u<=s.n,'ACK prefix monotonic')
                if row['responded']:
                    require(ack>last_guard,'distinct data-window response');last_guard=s.guard;cuts+=1
                require(s.reductions==cuts,'response accounting')
            counts['sender_events']+=1
    for code,selected,room in product(('Not-ECT','ECT0','ECT1','CE'),(False,True),(False,True)):
        action,out=router_signal(code,selected,room)
        require((action=='DROP')==(not room or(selected and code=='Not-ECT')),'router admission truth table')
        if code=='CE'and action=='FORWARD':require(out=='CE','never erase CE')
    return counts


def encode(value):
    if isinstance(value,F):return str(value)
    raise TypeError(type(value).__name__)


if __name__=='__main__':
    print(json.dumps(dict(status='PASS',demonstration=main_trace(),checks=tests()),default=encode,ensure_ascii=False,indent=2))
