#!/usr/bin/env python3
"""Read-only finite cache protocols and a separate teaching stride predictor.
Python 3.10+, standard library only. Public checks do not depend on assert.
"""
from dataclasses import dataclass
import itertools
import json
import random


def integer(x, lo, hi, name):
    if type(x) is not int or not lo <= x <= hi:
        raise ValueError(name)
    return x


def power(x, hi, name):
    integer(x, 1, hi, name)
    if x & (x - 1):
        raise ValueError(name)
    return x


@dataclass(frozen=True, eq=False)
class Request:
    serial: int
    address: int
    origin: str


@dataclass(frozen=True, eq=False)
class Transaction:
    serial: int
    block: int


@dataclass(frozen=True)
class Result:
    request: Request
    value: int


@dataclass(frozen=True)
class Admission:
    status: str
    request: object = None
    transaction: object = None
    results: tuple = ()
    reason: str = ''


@dataclass(eq=False)
class Way:
    block: object
    valid: bool
    owner: object
    stamp: int
    data: list


@dataclass(eq=False)
class Pending:
    token: Transaction
    set_index: int
    way_index: int
    primary: int
    targets: list
    received: list
    remaining: int


class _Cache:
    def __init__(self, memory, line_words=4, sets=2, ways=2, misses=4, targets=4):
        power(line_words, 256, 'line words')
        power(sets, 256, 'sets')
        integer(ways, 1, 16, 'ways')
        integer(misses, 1, 256, 'misses')
        integer(targets, 1, 256, 'targets')
        if type(memory) not in (list, tuple) or not 1 <= len(memory) <= (1 << 20):
            raise ValueError('memory')
        if len(memory) % line_words:
            raise ValueError('whole backing lines')
        for x in memory:
            integer(x, 0, (1 << 32) - 1, 'memory word')
        self.memory = tuple(memory)
        self.K, self.S, self.W, self.M, self.T = line_words, sets, ways, misses, targets
        self.lines = [[Way(None, False, None, 0, [0] * self.K)
                       for _ in range(self.W)] for _ in range(self.S)]
        self.pending = []
        self.request_serial = self.transaction_serial = self.clock = 0

    def _address(self, address):
        integer(address, 0, 4 * len(self.memory) - 4, 'address')
        if address % 4:
            raise ValueError('word alignment')
        word = address // 4
        return word // self.K, word % self.K

    def _request(self, address, origin):
        self.request_serial += 1
        return Request(self.request_serial, address, origin)

    def _touch(self, way):
        self.clock += 1
        way.stamp = self.clock

    def _lookup(self, token):
        if type(token) is not Transaction:
            raise ValueError('transaction identity')
        for i, rec in enumerate(self.pending):
            if rec.token is token:
                return i, rec
        raise ValueError('dead or foreign transaction')

    def read(self, address, origin='demand'):
        block, offset = self._address(address)
        if type(origin) is not str or origin not in ('demand', 'prefetch'):
            raise ValueError('origin')
        group = self.lines[block % self.S]
        for way in group:
            if way.valid and way.block == block:
                req = self._request(address, origin)
                self._touch(way)
                return Admission('HIT', req, results=(Result(req, way.data[offset]),))
        for rec in self.pending:
            if rec.token.block != block:
                continue
            if rec.received[offset]:
                req = self._request(address, origin)
                way = self.lines[rec.set_index][rec.way_index]
                return Admission('FORWARD', req, rec.token, (Result(req, way.data[offset]),))
            if len(rec.targets) == self.T:
                return Admission('RETRY', reason='target-full')
            req = self._request(address, origin)
            rec.targets.append(req)
            return Admission('MERGE', req, rec.token)
        if len(self.pending) == self.M:
            return Admission('RETRY', reason='mshr-full')
        eligible = [i for i, way in enumerate(group) if way.owner is None]
        if not eligible:
            return Admission('RETRY', reason='ways-reserved')
        invalid = [i for i in eligible if not group[i].valid]
        wi = invalid[0] if invalid else min(eligible, key=lambda i: (group[i].stamp, i))
        req = self._request(address, origin)
        self.transaction_serial += 1
        token = Transaction(self.transaction_serial, block)
        way = group[wi]
        # Unreceived old data is inaccessible, so no O(K) data clear is needed.
        way.valid, way.block, way.owner = False, block, token
        self.pending.append(Pending(token, block % self.S, wi, offset,
                                    [req], [False] * self.K, self.K))
        return Admission('MISS', req, token)

    def _install(self, index, rec):
        way = self.lines[rec.set_index][rec.way_index]
        way.valid, way.owner = True, None
        self._touch(way)
        del self.pending[index]

    def snapshot(self):
        """Diagnostic full copy, not a constant-cost cache event."""
        return (self.request_serial, self.transaction_serial, self.clock,
                tuple(tuple((w.block, w.valid, None if w.owner is None else w.owner.serial,
                             w.stamp, tuple(w.data)) for w in group) for group in self.lines),
                tuple((r.token.serial, r.token.block, r.set_index, r.way_index, r.primary,
                       tuple((q.serial, q.address, q.origin) for q in r.targets),
                       tuple(r.received), r.remaining) for r in self.pending))

    def check(self):
        """Expensive diagnostic invariants; read/complete/deliver do not call it."""
        need(len(self.pending) <= self.M, 'MSHR bound')
        need(len({r.token.block for r in self.pending}) == len(self.pending), 'one miss per block')
        owners = {(r.set_index, r.way_index): r for r in self.pending}
        need(len(owners) == len(self.pending), 'one owner per way')
        seen = set()
        for si, group in enumerate(self.lines):
            blocks = [w.block for w in group if w.valid]
            need(len(set(blocks)) == len(blocks), 'resident uniqueness')
            for wi, w in enumerate(group):
                if w.owner is None:
                    need((si, wi) not in owners, 'unowned record')
                    if w.valid:
                        need(w.block % self.S == si, 'resident mapping')
                        need(w.data == list(self.memory[w.block*self.K:(w.block+1)*self.K]), 'resident data')
                else:
                    r = owners[si, wi]
                    need(not w.valid and w.owner is r.token and w.block == r.token.block, 'reserved identity')
                    need(w.block % self.S == si and r.remaining == r.received.count(False), 'pending shape')
                    need(r.remaining > 0 and len(r.targets) <= self.T, 'pending lifetime')
                    for off, arrived in enumerate(r.received):
                        if arrived:
                            need(w.data[off] == self.memory[w.block*self.K+off], 'received data')
                    for q in r.targets:
                        need(q not in seen and q.address//(4*self.K) == w.block, 'target owner')
                        need(not r.received[(q.address//4)%self.K], 'returned target retained')
                        seen.add(q)
        return True


class WholeLineCache(_Cache):
    def complete(self, token):
        i, rec = self._lookup(token)
        way = self.lines[rec.set_index][rec.way_index]
        start = token.block * self.K
        way.data[:] = self.memory[start:start+self.K]
        out = tuple(Result(q, way.data[(q.address//4) % self.K]) for q in rec.targets)
        self._install(i, rec)
        return out


class BeatCache(_Cache):
    def deliver(self, token, offset):
        integer(offset, 0, self.K - 1, 'word offset')
        i, rec = self._lookup(token)
        if rec.received[offset]:
            raise ValueError('duplicate beat')
        way = self.lines[rec.set_index][rec.way_index]
        way.data[offset] = self.memory[token.block*self.K+offset]
        rec.received[offset] = True
        rec.remaining -= 1
        out, waiting = [], []
        for q in rec.targets:
            if (q.address//4) % self.K == offset:
                out.append(Result(q, way.data[offset]))
            else:
                waiting.append(q)
        rec.targets = waiting
        if rec.remaining == 0:
            self._install(i, rec)
        return tuple(out)


def critical_order(line_words, offset):
    power(line_words, 256, 'line words')
    integer(offset, 0, line_words - 1, 'word offset')
    return tuple((offset + i) % line_words for i in range(line_words))


@dataclass(frozen=True)
class History:
    pc_word: int
    last_block: int
    stride: int
    confidence: int


class StridePredictor:
    """Teaching saturating evidence count, NOT the Chen-Baer four-state FSM."""
    def __init__(self, blocks, line_words=4, entries=16, threshold=2, degree=1):
        integer(blocks, 1, 1 << 20, 'block domain')
        power(line_words, 256, 'line words')
        power(entries, 1024, 'predictor entries')
        integer(threshold, 2, 3, 'threshold')
        integer(degree, 1, 32, 'degree')
        if blocks * line_words > (1 << 30):
            raise ValueError('32-bit byte address domain')
        self.N, self.K, self.E, self.threshold, self.degree = blocks, line_words, entries, threshold, degree
        self.table = [None] * entries

    def observe(self, pc, address):
        integer(pc, 0, (1 << 32) - 4, 'PC')
        integer(address, 0, 4*self.N*self.K-4, 'address')
        if pc % 4 or address % 4:
            raise ValueError('alignment')
        pw, block = pc//4, address//(4*self.K)
        slot = pw % self.E
        old = self.table[slot]
        if old is None or old.pc_word != pw:
            self.table[slot] = History(pw, block, 0, 0)
            return ()
        delta = block - old.last_block
        if delta == 0:
            return ()
        count = min(3, old.confidence+1) if delta == old.stride else 1
        self.table[slot] = History(pw, block, delta, count)
        if count < self.threshold:
            return ()
        return tuple(b for i in range(1, self.degree+1)
                     if 0 <= (b := block+i*delta) < self.N)


def issue_hints(cache, blocks):
    """No predictor training here; each hint is an ordinary tagged cache read."""
    if not isinstance(cache, (WholeLineCache, BeatCache)) or type(blocks) not in (tuple, list):
        raise ValueError('hint batch')
    for b in blocks:
        integer(b, 0, len(cache.memory)//cache.K - 1, 'hint block')
    return tuple(cache.read(b*cache.K*4, 'prefetch') for b in blocks)


def need(ok, why):
    if not ok:
        raise RuntimeError(why)


def memory(words=64):
    return tuple((17*i+11) & ((1 << 32)-1) for i in range(words))


def public_examples():
    mem = memory()
    c = WholeLineCache(mem, 4, 2, 2, 2, 2)
    warm = c.read(0); c.complete(warm.transaction)
    first = c.read(16); hit = c.read(0); joined = c.read(20)
    before = c.snapshot(); retry = c.read(24)
    need(retry.reason == 'target-full' and c.snapshot() == before, 'target retry')
    second = c.read(32)
    full_hit = c.read(0)
    need(full_hit.status == 'HIT' and len(c.pending) == c.M, 'hit under full MSHRs')
    before = c.snapshot(); blocked = c.read(48)
    need(blocked.reason == 'mshr-full' and c.snapshot() == before, 'MSHR retry')
    late = c.complete(second.transaction); early = c.complete(first.transaction)
    c.check()
    whole = {'statuses': [first.status, hit.status, joined.status, retry.status, second.status, blocked.status],
             'reasons': [retry.reason, blocked.reason], 'return_addresses': [r.request.address for r in late+early],
             'return_values': [r.value for r in late+early], 'lower_transactions_including_warmup': c.transaction_serial, 'hit_while_full': full_hit.status}
    reserved = WholeLineCache(mem, 4, 1, 2, 3, 2)
    reserved.read(0); reserved.read(16)
    before = reserved.snapshot(); no_way = reserved.read(32)
    need(no_way.reason == 'ways-reserved' and reserved.snapshot() == before, 'global room but no way')
    whole['free_mshr_but_no_way'] = no_way.reason
    b = BeatCache(mem, 4, 1, 1, 2, 2)
    first = b.read(12); joined = b.read(0); token = first.transaction
    trace = []
    for off in critical_order(4, 3):
        replies = b.deliver(token, off)
        row = {'offset': off, 'returned_addresses': [r.request.address for r in replies],
               'live_transactions': len(b.pending), 'resident': b.lines[0][0].valid}
        if off == 3:
            forward = b.read(12); blocked = b.read(16)
            row.update(forward=forward.status, blocked=blocked.reason)
            need(forward.status == 'FORWARD' and blocked.reason == 'ways-reserved', 'partial lifetime')
        if off == 0:
            hint = issue_hints(b, (0,))[0]
            need(hint.status == 'FORWARD' and hint.results[0].value == 11, 'hint partial forwarding')
            row['hint_consumer'] = hint.status
        trace.append(row); b.check()
    pred = StridePredictor(16, 4, 4, 2, 2); predictions = []
    for block in (1, 3, 5, 5, 7, 2, 4, 6):
        hints = pred.observe(0x100, block*16)
        predictions.append({'block': block, 'history': pred.table[0].__dict__, 'hints': list(hints)})
    return {'whole_line': whole, 'beat_trace': trace, 'stride_trace': predictions,
            'latency_assumption': {'L': 20, 'tau': 2, 'natural_primary': 26, 'critical_primary': 20, 'full_line_both': 26}}


def consumer_trace(blocks, prefetch, sets=2, ways=2, domain=8):
    c = WholeLineCache(memory(domain), 1, sets, ways, 4, 4)
    p = StridePredictor(domain, 1, 4, 2, 1)
    counts = {'demand_hits': 0, 'demand_misses': 0, 'demand_merges': 0, 'prefetch_requests': 0, 'fills': 0}
    rows = []
    for block in blocks:
        adm = c.read(block*4)
        need(adm.status in ('HIT', 'MISS'), 'sequential consumer')
        counts['demand_hits' if adm.status == 'HIT' else 'demand_misses'] += 1
        replies = adm.results
        if adm.status == 'MISS':
            replies = c.complete(adm.transaction); counts['fills'] += 1
        need(len(replies) == 1 and replies[0].value == c.memory[block], 'demand value')
        hints = p.observe(0x100, block*4)
        consumed = issue_hints(c, hints) if prefetch else ()
        for request in consumed:
            counts['prefetch_requests'] += 1
            if request.status == 'MISS':
                c.complete(request.transaction); counts['fills'] += 1
        rows.append({'block': block, 'demand': adm.status, 'hints': list(hints),
                     'prefetch_status': [x.status for x in consumed]})
        c.check()
    return {'counts': counts, 'rows': rows,
            'service_assumption': 'Each issued fill completes before the next demand; no cycle or bandwidth speedup claim.'}


def tests():
    counts = {'event_sequences': 0, 'random_events': 0, 'accepted_loads': 0,
              'returned_loads': 0, 'predictor_observations': 0, 'refusals': 0}
    def run(cls, events, cfg=(2, 1, 2, 2, 2)):
        c = cls(memory(16), *cfg); waiting = set(); returned = set()
        def receive(batch):
            for r in batch:
                need(r.request in waiting and r.request not in returned, 'exactly once identity')
                need(r.value == c.memory[r.request.address//4], 'literal backing oracle')
                waiting.remove(r.request); returned.add(r.request); counts['returned_loads'] += 1
        for e in events:
            if e < 4:
                before = c.snapshot(); adm = c.read(e*4)
                if adm.status == 'RETRY':
                    need(c.snapshot() == before and adm.request is None, 'retry no changes')
                else:
                    waiting.add(adm.request); counts['accepted_loads'] += 1; receive(adm.results)
            elif c.pending:
                r = c.pending[(e-4) % len(c.pending)]
                if cls is WholeLineCache: receive(c.complete(r.token))
                else:
                    offsets = [i for i, v in enumerate(r.received) if not v]
                    receive(c.deliver(r.token, offsets[(e-4) % len(offsets)]))
            c.check()
            need(waiting == {q for r in c.pending for q in r.targets}, 'external pending ledger')
        while c.pending:
            r = c.pending[-1]
            if cls is WholeLineCache:receive(c.complete(r.token))
            else:
                for off in reversed(range(c.K)):
                    if not r.received[off]:receive(c.deliver(r.token, off))
        need(not waiting, 'fair drain'); c.check()
    for cls in (WholeLineCache, BeatCache):
        for events in itertools.product(range(6), repeat=5):
            run(cls, events); counts['event_sequences'] += 1
        rng = random.Random(25371)
        for _ in range(300):
            events = [rng.randrange(8) for _ in range(120)]
            run(cls, events, (rng.choice((1, 2, 4)), rng.choice((1, 2)), rng.choice((1, 2)), rng.choice((1, 2, 3)), rng.choice((1, 2, 3))))
            counts['random_events'] += len(events)
    # Predictor oracle stores the observed distinct-block sequence per current slot.
    for threshold in (2, 3):
        for seq in itertools.product(range(4), repeat=6):
            p = StridePredictor(8, 1, 2, threshold, 3); slots = {}
            for k, block in enumerate(seq):
                pc = (0, 4, 8)[k % 3]; slot = (pc//4) % 2
                if slot not in slots or slots[slot][0] != pc:
                    slots[slot] = (pc, [block]); expected = ()
                else:
                    path = slots[slot][1]
                    if path[-1] == block:expected = ()
                    else:
                        path.append(block); d = path[-1] - path[-2]; runlen = 1
                        for j in range(len(path)-2, 0, -1):
                            if path[j]-path[j-1] != d:break
                            runlen += 1
                        expected = tuple(block+i*d for i in range(1,4) if 0 <= block+i*d < 8) if runlen >= threshold else ()
                need(p.observe(pc, block*4) == expected, 'sequence-history prediction')
                counts['predictor_observations'] += 1
    def refuse(fn, obj=None):
        before = obj.snapshot() if isinstance(obj, _Cache) else tuple(obj.table) if obj else None
        try:fn()
        except ValueError:pass
        else:raise RuntimeError('missing rejection')
        after = obj.snapshot() if isinstance(obj, _Cache) else tuple(obj.table) if obj else None
        need(before == after, 'rejection atomicity'); counts['refusals'] += 1
    for cls in (WholeLineCache, BeatCache):
        c = cls(memory(16)); adm = c.read(0); t = adm.transaction
        action = (lambda tok: c.complete(tok)) if cls is WholeLineCache else (lambda tok: c.deliver(tok, 0))
        for bad in (None, Transaction(t.serial, t.block), WholeLineCache(memory(16)).read(0).transaction):refuse(lambda:action(bad), c)
        for address in (-4, 1, 64, True, 4.0):refuse(lambda:c.read(address), c)
        for origin in ('fetch', None, True):refuse(lambda:c.read(0, origin), c)
        if cls is WholeLineCache:c.complete(t)
        else:
            c.deliver(t,0); refuse(lambda:c.deliver(t,0), c)
            for off in (-1,4,True,0.0):refuse(lambda:c.deliver(t,off),c)
            for off in (1,2,3):c.deliver(t,off)
        refuse(lambda:action(t),c)
        refuse(lambda:issue_hints(c,(0,100)),c)
    for kw in ({'line_words':3},{'sets':0},{'ways':True},{'misses':0},{'targets':1.0}):refuse(lambda:WholeLineCache(memory(16),**kw))
    for mem in ([],[True]*4,[1,2,3],[1<<32]*4,'bad'):refuse(lambda:WholeLineCache(mem))
    p=StridePredictor(8)
    for pc,addr in ((True,0),(1,0),(-4,0),(1<<32,0),(0,1),(0,128),(0,True)):
        refuse(lambda:p.observe(pc,addr),p)
    for kw in ({'blocks':0},{'blocks':8,'entries':3},{'blocks':8,'threshold':1},{'blocks':8,'degree':True}):refuse(lambda:StridePredictor(**kw))
    return counts


def main():
    examples = public_examples()
    consumers = {'regular_plain': consumer_trace(range(8), False), 'regular_prefetch': consumer_trace(range(8), True),
                 'pollution_plain': consumer_trace((0,2,4,4), False, 2, 1),
                 'pollution_prefetch': consumer_trace((0,2,4,4), True, 2, 1)}
    need(consumers['regular_prefetch']['counts']['demand_hits'] == 5, 'regular benefit')
    need(consumers['pollution_prefetch']['counts']['fills'] == 5 and consumers['pollution_plain']['counts']['fills'] == 3, 'pollution cost')
    print(json.dumps({'status': 'PASS', 'counts': tests(), 'examples': examples, 'consumers': consumers}, ensure_ascii=False, indent=2, sort_keys=True))


if __name__ == '__main__':
    main()
