#!/usr/bin/env python3
"""L19: fully pipelined issue resources, DAG lists and modulo schedules.

Positive integer latency is time until result availability, NOT resource occupancy.
Completions at t precede launches at t. Standard library, stdout only, no assert.
"""
from collections import deque, Counter
from itertools import product
from copy import deepcopy
import json
import random


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


def check_graph(nodes, edges, capacities, cyclic=False):
    require(type(nodes) is dict and all(type(x) is str and x for x in nodes), "bad nodes")
    require(type(capacities) is dict and all(type(r) is str and type(c) is int and c > 0 for r, c in capacities.items()), "bad capacities")
    for data in nodes.values():
        require(type(data) is dict and set(data) == {"latency", "resource"}, "bad operation data")
        require(type(data["latency"]) is int and data["latency"] > 0, "latency must be positive integer")
        require(data["resource"] in capacities, "unknown resource")
    require(type(edges) in (tuple, list), "bad edge list")
    for edge in edges:
        require(type(edge) in (tuple, list) and len(edge) == 3, "bad dependence")
        u, v, distance = edge
        require(u in nodes and v in nodes and type(distance) is int and distance >= 0, "bad dependence endpoint/distance")
        require(cyclic or distance == 0, "DAG scheduling cannot consume iteration distances")
    # The within-iteration graph must be a DAG; carried cycles may be present.
    pred = {v: [] for v in nodes}
    succ = {v: [] for v in nodes}
    for u, v, d in edges:
        if d == 0:
            pred[v].append(u)
            succ[u].append(v)
    degree = {v: len(pred[v]) for v in nodes}
    ready = deque(v for v in nodes if degree[v] == 0)
    order = []
    while ready:
        u = ready.popleft()
        order.append(u)
        for v in succ[u]:
            degree[v] -= 1
            if degree[v] == 0:
                ready.append(v)
    require(len(order) == len(nodes), "zero-distance dependence cycle")
    return order, pred, succ


def schedule_check(nodes, edges, capacities, starts, ii=None):
    check_graph(nodes, edges, capacities, cyclic=ii is not None)
    require(type(starts) is dict and set(starts) == set(nodes), "schedule must cover exactly all operations")
    require(all(type(t) is int and t >= 0 for t in starts.values()), "bad start time")
    if ii is not None:
        require(type(ii) is int and ii > 0, "bad initiation interval")
    errors, slots = [], {}
    for u, v, d in edges:
        slack = starts[v] + (ii * d if ii is not None else 0) - starts[u] - nodes[u]["latency"]
        if slack < 0:
            errors.append({"kind": "dependence", "edge": [u, v, d], "slack": slack})
    for v in nodes:
        key = (nodes[v]["resource"], starts[v] if ii is None else starts[v] % ii)
        slots.setdefault(key, []).append(v)
    for (resource, slot), occupants in slots.items():
        if len(occupants) > capacities[resource]:
            errors.append({"kind": "issue_capacity", "resource": resource, "slot": slot, "operations": occupants})
    return {"valid": not errors, "errors": errors,
            "slots": [{"resource": r, "slot": t, "operations": occupants} for (r, t), occupants in slots.items()],
            "last_completion": max((starts[v] + nodes[v]["latency"] for v in nodes), default=0)}


def list_schedule(nodes, edges, capacities):
    order, pred, succ = check_graph(nodes, edges, capacities)
    rank = {v: i for i, v in enumerate(nodes)}
    height = {}
    for u in reversed(order):
        height[u] = nodes[u]["latency"] + max((height[v] for v in succ[u]), default=0)
    starts, trace, time = {}, [], 0
    while len(starts) < len(nodes):
        ready = [v for v in nodes if v not in starts and all(u in starts and starts[u] + nodes[u]["latency"] <= time for u in pred[v])]
        used, issued = Counter(), []
        for v in sorted(ready, key=lambda v: (-height[v], rank[v])):
            resource = nodes[v]["resource"]
            if used[resource] < capacities[resource]:
                starts[v] = time
                used[resource] += 1
                issued.append(v)
        trace.append({"time": time, "ready": ready, "issued": issued})
        if len(starts) == len(nodes):
            break
        next_events = [starts[v] + nodes[v]["latency"] for v in starts if starts[v] + nodes[v]["latency"] > time]
        if any(v not in starts for v in ready):
            next_events.append(time + 1)  # A fully pipelined issue port resets next cycle.
        require(next_events, "impossible stalled DAG")
        time = min(next_events)
    certificate = schedule_check(nodes, edges, capacities, starts)
    require(certificate["valid"], "list algorithm emitted an invalid schedule")
    return {"starts": starts, "height": height, "trace": trace, "certificate": certificate}


def recurrence_cycles(nodes, edges):
    """Enumerate simple directed cycles, uniquely by their least declaration rank.

    This deliberately small reference is exponential in the graph size.
    Parallel edge records remain distinct; duplicate witnesses do not affect max.
    """
    rank = {v: i for i, v in enumerate(nodes)}
    succ = {v: [] for v in nodes}
    for e in edges:
        succ[e[0]].append(e)
    result = []
    for start in nodes:
        def visit(u, seen, path):
            for edge in succ[u]:
                _, v, _ = edge
                if v == start:
                    cycle = path + [edge]
                    latency = sum(nodes[e[0]]["latency"] for e in cycle)
                    distance = sum(e[2] for e in cycle)
                    require(distance > 0, "zero-distance cycle")
                    result.append({"edges": cycle, "latency_sum": latency, "distance_sum": distance,
                                   "lower_bound": (latency + distance - 1) // distance})
                elif rank[v] > rank[start] and v not in seen:
                    visit(v, seen | {v}, path + [edge])
        visit(start, {start}, [])
    return result


def lower_bounds(nodes, edges, capacities):
    check_graph(nodes, edges, capacities, cyclic=True)
    counts = Counter(x["resource"] for x in nodes.values())
    resource = [{"resource": r, "uses": n, "capacity": capacities[r],
                 "lower_bound": (n + capacities[r] - 1) // capacities[r]} for r, n in counts.items()]
    cycles = recurrence_cycles(nodes, edges)
    res = max((x["lower_bound"] for x in resource), default=0)
    rec = max((x["lower_bound"] for x in cycles), default=0)
    return {"resource": resource, "cycles": cycles, "resource_bound": res, "recurrence_bound": rec,
            "combined": max(1, res, rec)}


def modulo_search(nodes, edges, capacities, ii, horizon, budget=100000):
    """Search only integer starts in [0,horizon], at a GIVEN positive ii.

    Lower-bound violation is a global impossibility certificate. Exhaustion of a
    finite window is not global infeasibility. Budget exhaustion is inconclusive.
    """
    order, _, _ = check_graph(nodes, edges, capacities, cyclic=True)
    require(type(ii) is int and ii > 0 and type(horizon) is int and horizon >= 0, "bad search window")
    require(type(budget) is int and budget >= 0, "bad search budget")
    bounds = lower_bounds(nodes, edges, capacities)
    if ii < bounds["combined"]:
        return {"status": "infeasible_by_bound", "ii": ii, "bounds": bounds, "attempts": 0, "starts": None}
    starts, slots = {}, Counter()
    attempts = 0
    stopped = False
    answer = None

    def search(index):
        nonlocal attempts, stopped, answer
        if index == len(order):
            answer = dict(starts)
            return True
        v = order[index]
        resource = nodes[v]["resource"]
        for time in range(horizon + 1):
            if attempts == budget:
                stopped = True
                return False
            attempts += 1
            slot = resource, time % ii
            if slots[slot] >= capacities[resource]:
                continue
            starts[v] = time
            legal = all(starts[w] + d * ii >= starts[u] + nodes[u]["latency"]
                        for u, w, d in edges if u in starts and w in starts)
            if legal:
                slots[slot] += 1
                if search(index + 1):
                    return True
                slots[slot] -= 1
                if slots[slot] == 0:
                    del slots[slot]  # Keep only active placements, not visited empty slots.
            del starts[v]
            if stopped:
                return False
        return False

    search(0)
    status = "feasible" if answer is not None else "search_limit" if stopped else "no_schedule_in_window"
    result = {"status": status, "ii": ii, "horizon": horizon, "attempts": attempts, "starts": answer, "bounds": bounds}
    if answer is not None:
        result["certificate"] = schedule_check(nodes, edges, capacities, answer, ii)
        require(result["certificate"]["valid"], "search emitted invalid schedule")
    return result


def expand(nodes, edges, capacities, offsets, ii, trips):
    require(type(trips) is int and trips >= 0, "bad trip count")
    require(schedule_check(nodes, edges, capacities, offsets, ii)["valid"], "invalid periodic schedule")
    expanded, expanded_edges, starts, identities, initial_edges = {}, [], {}, {}, []
    def name(v, i):
        return v + "[" + str(i) + "]"
    for i in range(trips):
        for v in nodes:
            label = name(v, i)
            expanded[label] = dict(nodes[v])
            starts[label] = offsets[v] + i * ii
            identities[label] = (v, i)
        for u, v, d in edges:
            if i >= d:
                expanded_edges.append((name(u, i - d), name(v, i), 0))
            else:
                initial_edges.append({"producer": [u, i - d], "consumer": [v, i]})
    check = schedule_check(expanded, expanded_edges, capacities, starts)
    require(check["valid"], "finite expansion invalid")
    return {"nodes": expanded, "edges": expanded_edges, "starts": starts, "identities": identities,
            "initial_edges": initial_edges, "certificate": check}


def prefix_graph(add_latency=1):
    nodes = {"A": {"latency": 2, "resource": "MEM"}, "B": {"latency": 3, "resource": "MUL"},
             "C": {"latency": add_latency, "resource": "ALU"}, "D": {"latency": 1, "resource": "MEM"}}
    edges = (("A", "B", 0), ("B", "C", 0), ("C", "D", 0), ("C", "C", 1))
    return nodes, edges, {"MEM": 1, "MUL": 1, "ALU": 1}


def _execute_prefix_events(expanded, x, factor, initial, overlapping_fault=False):
    # This private fault switch deliberately violates the public disjointness
    # contract, to display why the omitted D_i -> A_(i+1) edge matters.
    n = len(x)
    memory = list(x) + [0] if overlapping_fault else None
    output = [None] * n
    values = {("C", -1): initial}
    completions, launches = {}, {}
    for label, time in expanded["starts"].items():
        launches.setdefault(time, []).append(label)
    times = sorted(set(launches) | {expanded["starts"][v] + expanded["nodes"][v]["latency"] for v in expanded["nodes"]})
    trace = []
    for time in times:
        events = []
        # Values and stores committed at the boundary are available to this cycle.
        for label, value in completions.get(time, []):
            op, i = expanded["identities"][label]
            if op == "D":
                output[i] = value
                if memory is not None:
                    memory[i + 1] = value
                events.append({"complete": label, "write_Y_index": i, "value": value})
            else:
                values[op, i] = value
                events.append({"complete": label, "value": value})
        for label in launches.get(time, []):
            op, i = expanded["identities"][label]
            if op == "A":
                value = memory[i] if memory is not None else x[i]
            elif op == "B":
                require(("A", i) in values, "multiply operand not available")
                value = values["A", i] * factor
            elif op == "C":
                require(("B", i) in values and ("C", i - 1) in values, "accumulator operands not available")
                value = values["B", i] + values["C", i - 1]
            else:
                require(op == "D" and ("C", i) in values, "store operand not available")
                value = values["C", i]
            end = time + expanded["nodes"][label]["latency"]
            completions.setdefault(end, []).append((label, value))
            events.append({"issue": label, "sampled_or_computed_value": value, "available_at": end})
        trace.append({"time": time, "events": events})
    require(all(type(y) is int for y in output), "missing store completion")
    return {"output": output, "trace": trace, "last_completion": expanded["certificate"]["last_completion"], "shared_memory_fault": memory}


def execute_prefix(offsets, ii, x, factor=2, initial=0, add_latency=1):
    require(type(x) in (list, tuple) and all(type(a) is int for a in x), "bad input array")
    require(type(factor) is int and type(initial) is int, "bad integer scalar")
    nodes, edges, capacities = prefix_graph(add_latency)
    expanded = expand(nodes, edges, capacities, offsets, ii, len(x))
    # X and Y are distinct objects by construction; neither aliases any temporary.
    return _execute_prefix_events(expanded, tuple(x), factor, initial)


def execute_list_example(nodes, starts, x):
    require(type(x) is int, "bad scalar")
    expressions = {"A": ("mul", "x", "x"), "B": ("add", "x", 1), "C": ("sub", "x", 1),
                   "D": ("mul", "C", 2), "E": ("add", "A", "D"), "F": ("add", "D", 3)}
    values, completions, trace = {"x": x}, {}, []
    times = sorted(set(starts.values()) | {starts[v] + nodes[v]["latency"] for v in nodes})
    for time in times:
        for v, value in completions.get(time, []):
            values[v] = value
        issued = []
        for v in nodes:
            if starts[v] != time:
                continue
            op, a, b = expressions[v]
            require((type(a) is int or a in values) and (type(b) is int or b in values), "operand not complete")
            a, b = a if type(a) is int else values[a], b if type(b) is int else values[b]
            value = a + b if op == "add" else a - b if op == "sub" else a * b
            completions.setdefault(time + nodes[v]["latency"], []).append((v, value))
            issued.append({"operation": v, "value": value})
        trace.append({"time": time, "issued": issued})
    return {"values": {v: values[v] for v in nodes}, "trace": trace}


def greedy_counterexample():
    names = "ABCDEF"
    latencies = (4, 3, 1, 1, 1, 1)
    resources = ("M", "M", "M", "A", "M", "M")
    nodes = {v: {"latency": lat, "resource": r} for v, lat, r in zip(names, latencies, resources)}
    edges = (("C", "D", 0), ("A", "E", 0), ("D", "E", 0), ("D", "F", 0))
    caps = {"M": 1, "A": 1}
    result = list_schedule(nodes, edges, caps)
    better = dict(zip(names, (0, 2, 1, 2, 4, 3)))
    check = schedule_check(nodes, edges, caps, better)
    require(result["certificate"]["last_completion"] == 6 and check["valid"] and check["last_completion"] == 5, "greedy counterexample")
    first = execute_list_example(nodes, result["starts"], 4)
    second = execute_list_example(nodes, better, 4)
    require(first["values"] == second["values"] == {"A": 16, "B": 5, "C": 3, "D": 6, "E": 22, "F": 9}, "list values")
    for x in range(-50, 51):
        require(execute_list_example(nodes, result["starts"], x)["values"] == execute_list_example(nodes, better, x)["values"], "list semantic regression")
    return {"nodes": nodes, "edges": edges, "greedy": result, "better_starts": better, "better_certificate": check,
            "greedy_execution": first, "better_execution": second}


def self_test():
    nodes, edges, caps = prefix_graph()
    sequential = list_schedule(nodes, tuple(e for e in edges if e[2] == 0), caps)
    require(sequential["starts"] == {"A": 0, "B": 2, "C": 5, "D": 6}, "single iteration schedule")
    main = modulo_search(nodes, edges, caps, 2, 7)
    offsets = main["starts"]
    require(main["status"] == "feasible" and offsets == {"A": 0, "B": 2, "C": 5, "D": 7}, "main modulo schedule")
    actual = execute_prefix(offsets, 2, [1, 2, 3, 4])
    serial = execute_prefix(sequential["starts"], 7, [1, 2, 3, 4])
    require(actual["output"] == serial["output"] == [2, 6, 12, 20], "prefix values")
    require(actual["last_completion"] == 14 and serial["last_completion"] == 28, "finite makespan")
    colliding = schedule_check(nodes, edges, caps, sequential["starts"], 2)
    require(not colliding["valid"] and colliding["errors"][0]["kind"] == "issue_capacity", "cross-iteration collision")
    no_window = modulo_search(nodes, edges, caps, 2, 6)
    limited = modulo_search(nodes, edges, caps, 2, 7, budget=1)
    impossible = modulo_search(nodes, edges, caps, 1, 20)
    require(no_window["status"] == "no_schedule_in_window" and limited["status"] == "search_limit" and impossible["status"] == "infeasible_by_bound", "search outcomes")
    slow_nodes, slow_edges, slow_caps = prefix_graph(3)
    recurrence_reject = modulo_search(slow_nodes, slow_edges, slow_caps, 2, 20)
    require(recurrence_reject["status"] == "infeasible_by_bound" and recurrence_reject["bounds"]["recurrence_bound"] == 3, "recurrence bound")
    slow = modulo_search(slow_nodes, slow_edges, slow_caps, 3, 8)
    require(slow["status"] == "feasible", "recurrence-respecting schedule")
    require(execute_prefix(slow["starts"], 3, [1, 2, 3, 4], add_latency=3)["output"] == actual["output"], "slow accumulator")
    gap_nodes = {"P": {"latency": 2, "resource": "R"}, "Q": {"latency": 2, "resource": "R"}}
    gap_edges = (("P", "Q", 0), ("Q", "P", 2))
    gap_bounds = lower_bounds(gap_nodes, gap_edges, {"R": 1})
    gap_two = modulo_search(gap_nodes, gap_edges, {"R": 1}, 2, 8)
    gap_three = modulo_search(gap_nodes, gap_edges, {"R": 1}, 3, 8)
    require(gap_bounds["combined"] == 2 and gap_two["status"] == "no_schedule_in_window", "MII gap finite search")
    require(gap_three["starts"] == {"P": 0, "Q": 2}, "MII gap II=3")
    # The page gives a separate global proof: II=2 forces Q-P=2, hence a
    # same-residue collision. The bounded search status does NOT claim that proof.
    expanded = expand(nodes, edges, caps, offsets, 2, 4)
    fault = _execute_prefix_events(expanded, [1, 2, 3, 4], 2, 0, overlapping_fault=True)
    serial_expanded = expand(nodes, edges, caps, sequential["starts"], 7, 4)
    source_overlap = _execute_prefix_events(serial_expanded, [1, 2, 3, 4], 2, 0, overlapping_fault=True)
    require(source_overlap["output"] == [2, 6, 18, 54] and fault["output"] == [2, 6, 12, 20], "alias boundary")
    fixed_edges = edges + (("D", "A", 1),)
    alias_bounds = lower_bounds(nodes, fixed_edges, caps)
    require(alias_bounds["recurrence_bound"] == 7, "alias adds real carried recurrence")
    require(not schedule_check(nodes, fixed_edges, caps, offsets, 2)["valid"], "alias dependence rejected")
    rng = random.Random(19019)
    prefix_cases = 0
    for n in range(21):
        for _ in range(20):
            x = [rng.randrange(-20, 21) for _ in range(n)]
            c, acc = rng.randrange(-4, 5), rng.randrange(-10, 11)
            expected, value = [], acc
            for a in x:
                value += c * a
                expected.append(value)
            got = execute_prefix(offsets, 2, x, c, acc)
            require(got["output"] == expected, "random prefix")
            require(got["last_completion"] == (0 if n == 0 else 8 + 2 * (n - 1)), "startup/drain length")
            prefix_cases += 1
    dag_cases = 0
    for _ in range(1000):
        count = rng.randrange(1, 13)
        p = {str(i): {"latency": rng.randrange(1, 8), "resource": rng.choice(("M", "A"))} for i in range(count)}
        es = [(str(i), str(j), 0) for i in range(count) for j in range(i + 1, count) if rng.random() < .2]
        r = list_schedule(p, es, {"M": 1, "A": 2})
        require(r["certificate"]["valid"] and len(r["trace"]) <= 2 * count, "DAG schedule/event bound")
        dag_cases += 1
    # Exhaust every offset in a tiny window, independently of search pruning.
    modulo_cases = 0
    for _ in range(180):
        p = {str(i): {"latency": rng.randrange(1, 4), "resource": rng.choice(("M", "A"))} for i in range(3)}
        es = [(str(i), str(j), 0) for i in range(3) for j in range(i + 1, 3) if rng.random() < .35]
        es += [(str(i), str(j), 1) for i in range(3) for j in range(3) if rng.random() < .15]
        cap, ii, horizon = {"M": 1, "A": 1}, rng.randrange(1, 5), 4
        possible = any(schedule_check(p, es, cap, dict(zip(p, times)), ii)["valid"] for times in product(range(horizon + 1), repeat=3))
        result = modulo_search(p, es, cap, ii, horizon)
        require((result["status"] == "feasible") == possible, "exhaustive search equivalence")
        require(result["status"] != "search_limit", "tiny search budget")
        modulo_cases += 1
    return {"status": "PASS", "list_counterexample": greedy_counterexample(), "single_iteration": sequential,
            "modulo": main, "periodic_execution": actual, "serial_execution": serial,
            "first_iteration_collides_later": colliding, "search_outcomes": {"short_window": no_window, "limited": limited, "ii_one": impossible},
            "recurrence_change": {"ii_two_rejected": recurrence_reject, "ii_three": slow},
            "alias_fault": {"source_output": source_overlap["output"], "wrong_periodic_output": fault["output"], "repaired_graph_bounds": alias_bounds},
            "mii_not_sufficient": {"bounds": gap_bounds, "ii_two_window": gap_two, "ii_three": gap_three},
            "initial_edges": expanded["initial_edges"],
            "regressions": {"prefix_inputs": prefix_cases, "random_DAGs": dag_cases, "exhaustive_modulo_windows": modulo_cases},
            "zero_trips": execute_prefix(offsets, 2, []), "one_trip": execute_prefix(offsets, 2, [3])}


if __name__ == "__main__":
    print(json.dumps(self_test(), ensure_ascii=False, indent=2, sort_keys=True))
