#!/usr/bin/env python3
"""DISK-32 educational models. Python standard library; no disk-device access.

Outputs deterministic JSON. Invalid contracts raise ValueError. Assertions are
proof checks; optimized Python is deliberately refused rather than skipping them.
"""
import itertools
import json
import random
from fractions import Fraction


def lookup(extents, b):
    """Sorted predecessor search followed by half-open end check."""
    lo, hi = 0, len(extents)
    while lo < hi:
        mid = (lo + hi) // 2
        if extents[mid][0] <= b:
            lo = mid + 1
        else:
            hi = mid
    if lo == 0:
        return None
    l, n, p = extents[lo - 1]
    return p + b - l if b < l + n else None


def normalize(extents):
    result = []
    for l, n, p in sorted(extents):
        if n <= 0:
            raise ValueError('nonpositive extent')
        if result:
            a, k, q = result[-1]
            if a + k > l:
                raise ValueError('logical overlap')
            if a + k == l and q + k == p:
                result[-1] = (a, k + n, q)
                continue
        result.append((l, n, p))
    return result


def replace(extents, a, c, new_physical):
    """Representation-only replacement. Caller must own destination blocks."""
    if not 0 <= a < c or new_physical < 0:
        raise ValueError('invalid replacement')
    result = []
    for l, n, p in extents:
        end = l + n
        if end <= a or l >= c:
            result.append((l, n, p))
            continue
        if l < a:
            result.append((l, a - l, p))
        if end > c:
            result.append((c, end - c, p + c - l))
    result.append((a, c - a, new_physical))
    return normalize(result)


def runs(bits):
    """Independent bitmap-to-maximal-interval oracle."""
    out = []
    start = None
    for i, bit in enumerate(list(bits) + [False]):
        if bit and start is None:
            start = i
        if not bit and start is not None:
            out.append((start, i))
            start = None
    return out


class RangePool:
    def __init__(self, free):
        self.by_address = set(free)
        self.by_length = {(b - a, a) for a, b in free}
        self.owned = {}
        self.next_owner = 0

    def snapshot(self):
        return (frozenset(self.by_address), frozenset(self.by_length),
                tuple(sorted(self.owned.items())), self.next_owner)

    def _remove(self, a, b):
        self.by_address.remove((a, b))
        self.by_length.remove((b - a, a))

    def _insert(self, a, b):
        if a < b:
            self.by_address.add((a, b))
            self.by_length.add((b - a, a))

    def allocate(self, n, policy):
        if n <= 0 or policy not in ('first', 'best'):
            raise ValueError('positive size and known policy required')
        if policy == 'first':
            found = next(((a, b) for a, b in sorted(self.by_address) if b - a >= n), None)
        else:
            key = next(((length, a) for length, a in sorted(self.by_length) if length >= n), None)
            found = None if key is None else (key[1], key[1] + key[0])
        if found is None:
            return None
        a, b = found
        self._remove(a, b)
        self._insert(a + n, b)
        token = self.next_owner
        self.next_owner += 1
        self.owned[token] = (a, a + n)
        return token, (a, a + n)

    def release(self, token):
        if token not in self.owned:
            raise ValueError('unknown or already released owner')
        a, b = self.owned.pop(token)
        left = next(((x, y) for x, y in self.by_address if y == a), None)
        right = next(((x, y) for x, y in self.by_address if x == b), None)
        if left:
            self._remove(*left)
            a = left[0]
        if right:
            self._remove(*right)
            b = right[1]
        self._insert(a, b)


def stripe(b, members=4, chunk=2, capacity=8):
    if b < 0 or members <= 0 or chunk <= 0 or capacity <= 0 or capacity % chunk or b >= members * capacity:
        raise ValueError('invalid stripe parameters')
    k, o = divmod(b, chunk)
    return k % members, (k // members) * chunk + o


def mirror(b, capacity=8):
    group, row = stripe(b, members=2, capacity=capacity)
    return (2 * group, 2 * group + 1), row


def parity_location(b, capacity=8):
    if b < 0 or capacity <= 0 or b >= 3 * capacity:
        raise ValueError('logical block outside array')
    row, slot = divmod(b, 3)
    p = (3 - row) % 4
    return [d for d in range(4) if d != p][slot], row, p


def check_extent():
    initial = [(0, 3, 8), (5, 4, 20)]
    changed = replace(initial, 6, 8, 12)
    assert changed == [(0, 3, 8), (5, 1, 20), (6, 2, 12), (8, 1, 23)]
    assert lookup(initial, 1553 // 256) * 256 + 1553 % 256 == 5393
    assert lookup(initial, 1025 // 256) is None
    assert replace(initial, 6, 8, 21) == initial
    cases = 0
    rng = random.Random(3101)
    for mask in range(256):
        physical = list(range(8, 16))
        rng.shuffle(physical)
        oldmap = {i: physical[i] for i in range(8) if mask >> i & 1}
        ext = normalize([(i, 1, p) for i, p in oldmap.items()])
        for a in range(8):
            for c in range(a + 1, 9):
                out = replace(ext, a, c, 100)
                oracle = dict(oldmap)
                oracle.update({i: 100 + i - a for i in range(a, c)})
                actual = {i: lookup(out, i) for i in range(8) if lookup(out, i) is not None}
                assert actual == oracle
                assert len(set(actual.values())) == len(actual)
                assert normalize(out) == out
                cases += 1
    return dict(byte_1553_device_byte=5393, replacement=changed, exhaustive_replacements=cases)


def check_ranges():
    cases = 0
    for mask in range(1 << 11):
        bits = [False] + [bool(mask >> i & 1) for i in range(11)]
        free = runs(bits)
        for n in range(1, 13):
            for policy in ('first', 'best'):
                pool = RangePool(free)
                before = pool.snapshot()
                choices = [(a, b) for a, b in free if b - a >= n]
                chosen = min(choices, key=(lambda x: x[0]) if policy == 'first' else
                             (lambda x: (x[1] - x[0], x[0])), default=None)
                got = pool.allocate(n, policy)
                afterbits = list(bits)
                if chosen is None:
                    assert got is None and pool.snapshot() == before
                else:
                    a, b = chosen
                    assert got[1] == (a, a + n)
                    afterbits[a:a + n] = [False] * n
                assert pool.by_address == set(runs(afterbits))
                assert pool.by_length == {(b - a, a) for a, b in runs(afterbits)}
                cases += 1
    rng = random.Random(3102)
    pool = RangePool([(1, 32)])
    transitions = 0
    for _ in range(6000):
        before = pool.snapshot()
        if pool.owned and rng.random() < .46:
            pool.release(rng.choice(list(pool.owned)))
        else:
            if pool.allocate(rng.randrange(1, 15), rng.choice(('first', 'best'))) is None:
                assert pool.snapshot() == before
        occupied = set()
        for a, b in pool.owned.values():
            assert not occupied.intersection(range(a, b))
            occupied.update(range(a, b))
        bits = [i != 0 and i not in occupied for i in range(32)]
        assert pool.by_address == set(runs(bits))
        assert pool.by_length == {(b - a, a) for a, b in runs(bits)}
        saved = pool.snapshot()
        try:
            pool.release(-1)
        except ValueError:
            pass
        else:
            raise AssertionError('invalid release accepted')
        assert saved == pool.snapshot()
        transitions += 1
    examples = {}
    for seq in [(5, 7), (5, 3, 6)]:
        for policy in ('first', 'best'):
            pool = RangePool([(1, 9), (12, 18)])
            examples[str(seq) + ':' + policy] = [pool.allocate(n, policy) for n in seq]
    assert examples['(5, 7):first'][-1] is None
    assert examples['(5, 7):best'][-1][1] == (1, 8)
    assert examples['(5, 3, 6):first'][-1][1] == (12, 18)
    assert examples['(5, 3, 6):best'][-1] is None
    pool = RangePool([(1, 9), (12, 18)])
    tokens = [pool.allocate(n, 'first')[0] for n in [2, 3, 3]]
    for token in [tokens[0], tokens[2], tokens[1]]:
        pool.release(token)
    assert pool.by_address == {(1, 9), (12, 18)}
    return dict(exhaustive_policy_cases=cases, mixed_transitions=transitions, examples=examples,
                merge_restores=[[1, 9], [12, 18]])


def read_file(durable):
    imap = durable[durable[0]]
    inode = durable[imap['F']]
    return [durable[a] for a in inode]


def check_lfs():
    old = {0: 4, 1: 'first', 2: 'old', 3: [1, 2], 4: {'F': 3}}
    appended = {8: 'new', 9: [1, 8], 10: {'F': 9}}
    valid = broken = 0
    for mask in range(8):
        durable = dict(old)
        durable.update({a: appended[a] for bit, a in enumerate([8, 9, 10]) if mask >> bit & 1})
        assert read_file(durable) == ['first', 'old']
        valid += 1
        unsafe = dict(durable)
        unsafe[0] = 10
        try:
            got = read_file(unsafe)
        except KeyError:
            broken += 1
        else:
            assert mask == 7 and got == ['first', 'new']
    for selected in (4, 10):
        durable = dict(old)
        durable.update(appended)
        durable[0] = selected
        assert read_file(durable) == ['first', 'old' if selected == 4 else 'new']
        valid += 1
    assert broken == 7
    return dict(safe_durable_stage_states=valid, premature_root_broken_subsets=broken,
                byte_300_new_location=[8, 44], batch_MB_for_90_percent=9)


def check_cleaning():
    cases = 0
    for mask in range(256):
        current = {i: i if mask >> i & 1 else 100 + i for i in range(8)}
        disk = {a: ('object', i) for i, a in current.items()}
        for i in range(8):
            disk.setdefault(i, ('old', i))
        before = {i: disk[a] for i, a in current.items()}
        live = [i for i in range(8) if current[i] == i]
        updated = dict(current)
        for i in live:
            disk[200 + i] = disk[i]
            updated[i] = 200 + i
        # Publish only after all copies. Old contents cannot yet be reclaimed.
        for i in range(8):
            del disk[i]
        assert {i: disk[a] for i, a in updated.items()} == before
        assert not set(updated.values()).intersection(range(8))
        L = len(live)
        assert 8 - L >= 0
        if L < 8:
            assert Fraction(L + 8 - L, 8 - L) == 1 / (1 - Fraction(L, 8))
            assert Fraction(8 + L + 8 - L, 8 - L) == 2 / (1 - Fraction(L, 8))
        cases += 1
    return dict(all_live_masks=cases, main_live_slots=[1, 7], copied=2, reclaimed=8,
                net_free=6, data_write_amplification=str(Fraction(8, 6)),
                total_data_io_per_new_user_block=str(Fraction(16, 6)),
                cleaner_only_io_per_net_free_block=str(Fraction(10, 6)),
                excludes=['summary', 'inode', 'imap', 'checkpoint', 'free-index', 'flush'],
                six_live_variant=dict(net_free=2, write_amplification=4, total_io_ratio=8))


def check_raid():
    map_cases = 0
    for C in range(2, 65, 2):
        pairs = [stripe(b, capacity=C) for b in range(4 * C)]
        assert len(set(pairs)) == 4 * C
        assert set(pairs) == set(itertools.product(range(4), range(C)))
        copies = [x for b in range(2 * C) for x in
                  [(d, mirror(b, C)[1]) for d in mirror(b, C)[0]]]
        assert len(set(copies)) == 4 * C
        assert set(copies) == set(itertools.product(range(4), range(C)))
        for b in range(3 * C):
            d, r, p = parity_location(b, C)
            assert d != p and 0 <= r < C
        actual = {(parity_location(b, C)[0], parity_location(b, C)[1]) for b in range(3 * C)}
        expected = {(d, r) for r in range(C) for d in range(4) if d != (3 - r) % 4}
        assert actual == expected
        map_cases += 4 * C + 2 * C + 3 * C
    failure_states = []
    for mask in range(16):
        failed = {i for i in range(4) if mask >> i & 1}
        possible = all(any(d not in failed for d in mirror(b)[0]) for b in range(16))
        assert possible == (not {0, 1} <= failed and not {2, 3} <= failed)
        failure_states.append(dict(failed=sorted(failed), full_recovery=possible))
    assert sum(x['full_recovery'] for x in failure_states if len(x['failed']) == 2) == 4
    assert stripe(13) == (2, 3) and mirror(13) == ((0, 1), 7)
    assert parity_location(8) == (3, 2, 1)
    for fn, invalid in [(stripe, -1), (stripe, 32), (mirror, 16), (parity_location, 24)]:
        try:
            fn(invalid)
        except ValueError:
            pass
        else:
            raise AssertionError('out-of-capacity block accepted')
    # Exhaust all 4-bit codewords, each missing position, and every one-block update.
    erasures = updates = 0
    for x, y, z in itertools.product(range(16), repeat=3):
        p = x ^ y ^ z
        codeword = [x, y, z, p]
        for missing in range(4):
            recover = 0
            for i, value in enumerate(codeword):
                if i != missing:
                    recover ^= value
            assert recover == codeword[missing]
            erasures += 1
        for new in range(16):
            assert p ^ x ^ new == new ^ y ^ z
            # An off-diagonal update corrupts unchanged y exactly by delta.
            assert p ^ new ^ z == y ^ x ^ new
            assert (p ^ new ^ z == y) == (new == x)
            updates += 1
    table = []
    x, y, z, p, new, newp = 0x3C, 0xA5, 0x5A, 0xC3, 0x0F, 0xF0
    assert p == x ^ y ^ z and newp == new ^ y ^ z
    for dnew, pnew in [(False, False), (True, False), (False, True), (True, True)]:
        dx, dp = (new if dnew else x), (newp if pnew else p)
        recovered = dp ^ dx ^ z
        assert recovered == (0xA5 if dnew == pnew else 0x96)
        table.append(dict(data=f'{dx:02X}', parity=f'{dp:02X}', recovered_Y=f'{recovered:02X}'))
    # Redo guard: no committed log implies no home writes; a committed log repairs
    # every member-write subset. The log itself and all members are available.
    redo_cases = 0
    for mask in range(4):
        members = [new if mask & 1 else x, y, z, newp if mask & 2 else p]
        for repetition in (1, 2):
            for _ in range(repetition):
                members[0], members[3] = new, newp
            assert members[0] ^ members[1] ^ members[2] == members[3]
            assert members[3] ^ members[0] ^ members[2] == y
            redo_cases += 1
    return dict(mapping_cases=map_cases, failure_sets=failure_states,
                four_bit_single_erasures=erasures, four_bit_updates=updates,
                block_13_raid0=[2, 3], block_13_raid10=[[0, 1], 7],
                block_8_parity_layout=[3, 2, 1], write_hole=table,
                redo_available_log_and_members_cases=redo_cases,
                io_counts=[dict(changed=k, read_modify_write=2*k+2,
                                reconstruct_write=(3-k)+(k+1)) for k in (1, 2, 3)])


def main():
    if not __debug__:
        raise RuntimeError('Run without -O: proof assertions must not be disabled.')
    result = dict(model='DISK-32 independent snapshots; no hardware I/O',
                  extents=check_extent(), free_ranges=check_ranges(),
                  lfs=check_lfs(), cleaning=check_cleaning(), raid=check_raid())
    print(json.dumps(result, ensure_ascii=False, indent=2))


if __name__ == '__main__':
    main()
