"""Runnable Fenwick tree examples and deterministic tests for the article.

The FenwickTree, count_inversions, and find_by_order definitions above the
sentinel are the implementations shown in the article. Running this file
checks that those definitions still match the article, then tests the article
source itself.

    python3 static/downloads/fenwick-tree-examples.py
"""

from collections.abc import Iterable, Sequence
import difflib
import random
import re
from pathlib import Path


class FenwickTree:
    """Point updates and half-open range sums over a fixed-length array."""

    def __init__(self, values: Iterable[int]):
        self.bit = [0, *values]
        self.n = len(self.bit) - 1

        # Forward each completed block to its containing parent.
        for i in range(1, self.n + 1):
            parent = i + (i & -i)
            if parent <= self.n:
                self.bit[parent] += self.bit[i]

    def add(self, index: int, delta: int) -> None:
        """Apply values[index] += delta."""
        if not 0 <= index < self.n:
            raise IndexError("index out of range")

        i = index + 1
        while i <= self.n:
            self.bit[i] += delta
            i += i & -i

    def prefix_sum(self, end: int) -> int:
        """Return sum(values[0:end]); end is exclusive."""
        if not 0 <= end <= self.n:
            raise IndexError("prefix endpoint out of range")

        total = 0
        i = end
        while i > 0:
            total += self.bit[i]
            i -= i & -i
        return total

    def range_sum(self, left: int, right: int) -> int:
        """Return sum(values[left:right])."""
        if not 0 <= left <= right <= self.n:
            raise IndexError("invalid range")

        return self.prefix_sum(right) - self.prefix_sum(left)

    def set(self, index: int, value: int) -> None:
        """Assign values[index] = value."""
        current = self.range_sum(index, index + 1)
        self.add(index, value - current)


def count_inversions(values: Sequence[int]) -> int:
    ranks = {
        value: i
        for i, value in enumerate(sorted(set(values)))
    }

    tree = FenwickTree([0] * len(ranks))
    inversions = 0

    for value in reversed(values):
        rank = ranks[value]

        # Ranks [0, rank) are strictly smaller.
        inversions += tree.prefix_sum(rank)
        tree.add(rank, 1)

    return inversions


def find_by_order(tree: FenwickTree, k: int) -> int:
    """
    Return the zero-based position containing the kth item.

    k is one-based.
    Requires nonnegative integer frequencies.
    """
    if not 1 <= k <= tree.prefix_sum(tree.n):
        raise ValueError("k must be between 1 and the total frequency")

    index = 0
    step = 1 << (tree.n.bit_length() - 1)

    while step:
        candidate = index + step

        if candidate <= tree.n and tree.bit[candidate] < k:
            index = candidate
            k -= tree.bit[candidate]

        step >>= 1

    return index


# END OF ARTICLE IMPLEMENTATIONS


class DifferenceFenwick:
    """Range additions and point queries using one Fenwick tree over D."""

    def __init__(self, values, fenwick_cls=FenwickTree):
        values = list(values)
        self.n = len(values)
        diff = [0] * self.n
        if self.n:
            diff[0] = values[0]
            for i in range(1, self.n):
                diff[i] = values[i] - values[i - 1]
        self.b1 = fenwick_cls(diff)

    def add_range(self, left, right, delta):
        if not 0 <= left <= right <= self.n:
            raise IndexError("invalid range")
        if left == right:
            return
        self.b1.add(left, delta)
        if right < self.n:
            self.b1.add(right, -delta)

    def point(self, index):
        if not 0 <= index < self.n:
            raise IndexError("index out of range")
        return self.b1.prefix_sum(index + 1)


class RangeAddFenwick:
    """Range additions and range sums using B1 = D[j] and B2 = j * D[j]."""

    def __init__(self, values, fenwick_cls=FenwickTree):
        values = list(values)
        self.n = len(values)
        diff = [0] * self.n
        if self.n:
            diff[0] = values[0]
            for i in range(1, self.n):
                diff[i] = values[i] - values[i - 1]
        self.b1 = fenwick_cls(diff)
        self.b2 = fenwick_cls([j * diff[j] for j in range(self.n)])

    def add_range(self, left, right, delta):
        if not 0 <= left <= right <= self.n:
            raise IndexError("invalid range")
        if left == right:
            return
        self.b1.add(left, delta)
        self.b2.add(left, left * delta)
        if right < self.n:
            self.b1.add(right, -delta)
            self.b2.add(right, right * -delta)

    def prefix_sum(self, end):
        if not 0 <= end <= self.n:
            raise IndexError("prefix endpoint out of range")
        return end * self.b1.prefix_sum(end) - self.b2.prefix_sum(end)

    def range_sum(self, left, right):
        if not 0 <= left <= right <= self.n:
            raise IndexError("invalid range")
        return self.prefix_sum(right) - self.prefix_sum(left)


def _article_path():
    return Path(__file__).resolve().parents[2] / "content" / "blog" / "fenwick-tree" / "index.md"


def _between(text, start, end):
    i = text.find(start)
    j = text.find(end, i + len(start))
    if i < 0 or j < 0:
        raise AssertionError("missing marker %r" % start)
    return text[i:j].strip()


def _article_fences(markdown):
    fences = re.findall(r"```python\n(.*?)```", markdown, flags=re.S)
    selected = []
    for fence in fences:
        if fence.startswith("from collections.abc import Iterable\n"):
            selected.append(fence)
        elif fence.startswith("from collections.abc import Sequence\n"):
            selected.append(fence)
        elif fence.startswith("def find_by_order("):
            selected.append(fence)
    if len(selected) != 3:
        raise AssertionError("expected 3 article implementations, found %s" % len(selected))
    return selected


def _require_same(label, actual, expected):
    if actual.strip() != expected.strip():
        diff = "".join(
            difflib.unified_diff(
                expected.strip().splitlines(True),
                actual.strip().splitlines(True),
                fromfile="article",
                tofile="download",
            )
        )
        raise AssertionError("%s does not match the article:\n%s" % (label, diff))


def _from(text, start):
    i = text.find(start)
    if i < 0:
        raise AssertionError("missing marker %r" % start)
    return text[i:].strip()


def load_article_implementations(markdown):
    """Execute the article's own Python blocks."""
    selected = _article_fences(markdown)
    local = Path(__file__).read_text().split("# END OF ARTICLE IMPLEMENTATIONS", 1)[0]
    _require_same(
        "FenwickTree",
        _between(local, "class FenwickTree:", "def count_inversions"),
        _from(selected[0], "class FenwickTree:"),
    )
    _require_same(
        "count_inversions",
        _between(local, "def count_inversions", "def find_by_order"),
        _from(selected[1], "def count_inversions"),
    )
    _require_same(
        "find_by_order",
        _from(local, "def find_by_order"),
        selected[2].strip(),
    )
    namespace = {}
    for fence in selected:
        exec(compile(fence, "content/blog/fenwick-tree/index.md", "exec"), namespace)
    return namespace


def _raises(exception, fn):
    try:
        fn()
    except exception:
        return
    raise AssertionError("expected %s" % exception.__name__)


def _selection_oracle(frequencies, k):
    if k < 1:
        raise ValueError("k must be between 1 and the total frequency")
    total = 0
    for index, frequency in enumerate(frequencies):
        total += frequency
        if total >= k:
            return index
    raise ValueError("k must be between 1 and the total frequency")


def _inversion_oracle(values):
    return sum(
        1
        for i in range(len(values))
        for j in range(i + 1, len(values))
        if values[i] > values[j]
    )


def _assert_matches_array(tree, values):
    for end in range(len(values) + 1):
        if tree.prefix_sum(end) != sum(values[:end]):
            raise AssertionError((end, tree.prefix_sum(end), sum(values[:end])))
    for left in range(len(values) + 1):
        for right in range(left, len(values) + 1):
            got = tree.range_sum(left, right)
            expected = sum(values[left:right])
            if got != expected:
                raise AssertionError((left, right, got, expected))


def test_article_examples(FenwickTree, count_inversions, find_by_order):
    values = [3, 2, 5, 1, 4, 6, 2, 7]
    tree = FenwickTree(values)
    if tree.bit != [0, 3, 5, 5, 11, 4, 10, 2, 30]:
        raise AssertionError(tree.bit)
    if tree.prefix_sum(7) != 23:
        raise AssertionError(tree.prefix_sum(7))
    if tree.range_sum(2, 6) != 16:
        raise AssertionError(tree.range_sum(2, 6))
    tree.add(2, 4)
    if values != [3, 2, 5, 1, 4, 6, 2, 7]:
        raise AssertionError(values)
    if [tree.bit[3], tree.bit[4], tree.bit[8]] != [9, 15, 34]:
        raise AssertionError(tree.bit)
    if tree.range_sum(2, 6) != 20:
        raise AssertionError(tree.range_sum(2, 6))
    tree.set(4, 10)
    if tree.range_sum(2, 6) != 26:
        raise AssertionError(tree.range_sum(2, 6))
    if tree.prefix_sum(0) != 0 or tree.range_sum(3, 3) != 0:
        raise AssertionError("empty prefix or empty range")

    checked = FenwickTree([3, -2, 5])
    if checked.prefix_sum(0) != 0:
        raise AssertionError("prefix 0")
    if checked.range_sum(0, 3) != 6 or checked.range_sum(1, 1) != 0:
        raise AssertionError("shown range sums")
    before = checked.bit.copy()
    if checked.range_sum(1, 3) != 3 or checked.bit != before:
        raise AssertionError("query mutated the tree")
    checked.add(1, -4)
    if checked.range_sum(0, 3) != 2:
        raise AssertionError(checked.range_sum(0, 3))
    checked.set(0, 10)
    if checked.range_sum(0, 3) != 9:
        raise AssertionError(checked.range_sum(0, 3))

    if count_inversions([3, 1, 2]) != 2 or count_inversions([2, 2, 1]) != 2:
        raise AssertionError("shown inversion counts")

    frequencies = FenwickTree([2, 0, 3, 1])
    if frequencies.bit[4] != 6 or frequencies.bit[2] != 2 or frequencies.bit[3] != 3:
        raise AssertionError(frequencies.bit)
    if [find_by_order(frequencies, k) for k in (1, 3, 6)] != [0, 2, 3]:
        raise AssertionError("shown frequency selection")
    if (12 & -12) != 4:
        raise AssertionError("lowbit illustration")


def test_arrays(FenwickTree):
    cases = [
        [],
        [0],
        [5],
        [-3],
        [1, 2, 3, 4, 5],
        [3, -2, 5],
        [10 ** 100, -(10 ** 100), 10 ** 100 + 7],
        list(range(-4, 9)),
    ]
    for values in cases:
        tree = FenwickTree(list(values))
        _assert_matches_array(tree, list(values))

    consumed = FenwickTree(iter([4, -5, 6]))
    _assert_matches_array(consumed, [4, -5, 6])
    tuple_input = (8, 1, 1)
    tree = FenwickTree(tuple_input)
    tree.add(1, 3)
    if tuple_input != (8, 1, 1):
        raise AssertionError("tuple input was modified")
    _assert_matches_array(tree, [8, 4, 1])

    original = [3, 2, 5]
    tree = FenwickTree(original)
    tree.add(0, -8)
    tree.set(2, 100)
    if original != [3, 2, 5]:
        raise AssertionError(original)
    original[0] = 99
    if tree.range_sum(0, 1) != -5:
        raise AssertionError("tree aliased the input list")


def test_invalid_bounds(FenwickTree):
    tree = FenwickTree([1, 2, 3])
    before = tree.bit.copy()
    for call in (
        lambda: tree.add(-1, 1),
        lambda: tree.add(3, 1),
        lambda: tree.prefix_sum(-1),
        lambda: tree.prefix_sum(4),
        lambda: tree.range_sum(2, 1),
        lambda: tree.range_sum(-1, 1),
        lambda: tree.range_sum(0, 4),
        lambda: tree.set(-1, 1),
        lambda: tree.set(3, 1),
    ):
        _raises(IndexError, call)
    if tree.bit != before:
        raise AssertionError("invalid call mutated the tree")

    empty = FenwickTree([])
    if empty.prefix_sum(0) != 0 or empty.range_sum(0, 0) != 0:
        raise AssertionError("empty tree")
    _raises(IndexError, lambda: empty.add(0, 1))
    _raises(IndexError, lambda: empty.prefix_sum(1))


def test_randomized_updates(FenwickTree):
    rng = random.Random(20261001)
    for n in (0, 1, 2, 5, 13, 31):
        values = [rng.randint(-20, 20) for _ in range(n)]
        tree = FenwickTree(values)
        for _ in range(40):
            if n and rng.randrange(2) == 0:
                index = rng.randrange(n)
                if rng.randrange(2) == 0:
                    delta = rng.randint(-15, 15)
                    tree.add(index, delta)
                    values[index] += delta
                else:
                    value = rng.randint(-15, 15)
                    tree.set(index, value)
                    values[index] = value
            end = rng.randint(0, n)
            if tree.prefix_sum(end) != sum(values[:end]):
                raise AssertionError("randomized prefix")
            left = rng.randint(0, n)
            right = rng.randint(left, n)
            before = tree.bit.copy()
            if tree.range_sum(left, right) != sum(values[left:right]):
                raise AssertionError("randomized range")
            if tree.bit != before:
                raise AssertionError("randomized query mutated the tree")


def test_frequency_selection(FenwickTree, find_by_order):
    samples = [
        [2, 0, 3, 1],
        [1],
        [0, 0, 4],
        [5, 0, 0, 0],
        [0, 0, 0, 1],
    ]
    for frequencies in samples:
        _check_frequencies(FenwickTree, find_by_order, frequencies)

    rng = random.Random(20261001)
    for _ in range(20):
        frequencies = [rng.randint(0, 4) for _ in range(12)]
        _check_frequencies(FenwickTree, find_by_order, frequencies)

    frequencies = [2, 0, 3, 1]
    tree = FenwickTree(frequencies)
    tree.add(2, -3)
    frequencies[2] -= 3
    _check_frequencies(FenwickTree, find_by_order, frequencies, tree)
    tree.add(0, -2)
    frequencies[0] -= 2
    _check_frequencies(FenwickTree, find_by_order, frequencies, tree)
    if 0 in (find_by_order(tree, k) for k in range(1, sum(frequencies) + 1)):
        raise AssertionError("removed position was selected")


def _check_frequencies(FenwickTree, find_by_order, frequencies, tree=None):
    if tree is None:
        tree = FenwickTree(frequencies)
    total = sum(frequencies)
    for k in range(1, total + 1):
        if find_by_order(tree, k) != _selection_oracle(frequencies, k):
            raise AssertionError((frequencies, k))
    before = tree.bit.copy()
    for k in (0, -1, total + 1):
        _raises(ValueError, lambda k=k: find_by_order(tree, k))
    if tree.bit != before:
        raise AssertionError("invalid selection mutated the tree")
    if total == 0:
        return
    first = find_by_order(tree, 1)
    last = find_by_order(tree, total)
    if frequencies[first] <= 0 or frequencies[last] <= 0:
        raise AssertionError("endpoint is not a positive frequency")
    for index, frequency in enumerate(frequencies):
        if frequency == 0 and index in (
            find_by_order(tree, k) for k in range(1, total + 1)
        ):
            raise AssertionError("zero-frequency position selected")


def test_inversions(count_inversions):
    samples = [
        [],
        [1],
        [3, 1, 2],
        [2, 2, 1],
        [1, 2, 3],
        [3, 2, 1],
        [-2, -2, -5, 0],
        [5, -1, 5, -1],
    ]
    for values in samples:
        if count_inversions(values) != _inversion_oracle(values):
            raise AssertionError(values)
    rng = random.Random(20261001)
    for _ in range(15):
        values = [rng.randint(-5, 5) for _ in range(12)]
        if count_inversions(values) != _inversion_oracle(values):
            raise AssertionError(values)


def test_range_updates(FenwickTree):
    def check(values, updates):
        array = list(values)
        one = DifferenceFenwick(array, FenwickTree)
        two = RangeAddFenwick(array, FenwickTree)
        for left, right, delta in updates:
            one.add_range(left, right, delta)
            two.add_range(left, right, delta)
            for index in range(left, right):
                array[index] += delta
        for index, value in enumerate(array):
            if one.point(index) != value:
                raise AssertionError((index, one.point(index), value))
        for end in range(len(array) + 1):
            formula = end * two.b1.prefix_sum(end) - two.b2.prefix_sum(end)
            if two.prefix_sum(end) != sum(array[:end]) or formula != sum(array[:end]):
                raise AssertionError((end, two.prefix_sum(end), formula, sum(array[:end])))
        for left in range(len(array) + 1):
            for right in range(left, len(array) + 1):
                if two.range_sum(left, right) != sum(array[left:right]):
                    raise AssertionError((left, right))

    check([3, -2, 5, 7], [])
    check([3, -2, 5, 7], [(1, 1, 9), (4, 4, -3), (0, 4, 2), (1, 4, -5), (2, 3, 8)])
    check([], [(0, 0, 4)])
    check([10 ** 50], [(0, 1, -(10 ** 40)), (1, 1, 3)])

    tree = RangeAddFenwick([1, 2, 3, 4], FenwickTree)
    before = (tree.b1.bit.copy(), tree.b2.bit.copy())
    tree.add_range(4, 4, 11)
    tree.add_range(2, 2, -6)
    if (tree.b1.bit, tree.b2.bit) != before:
        raise AssertionError("empty interval changed a tree")
    for call in (
        lambda: tree.add_range(-1, 1, 1),
        lambda: tree.add_range(1, 0, 1),
        lambda: tree.add_range(0, 5, 1),
        lambda: tree.prefix_sum(-1),
        lambda: tree.prefix_sum(5),
        lambda: tree.range_sum(3, 2),
    ):
        _raises(IndexError, call)

    rng = random.Random(20261001)
    values = [rng.randint(-10, 10) for _ in range(9)]
    updates = []
    for _ in range(25):
        left = rng.randint(0, 9)
        right = rng.randint(left, 9)
        updates.append((left, right, rng.randint(-12, 12)))
    updates.append((9, 9, 5))
    updates.append((0, 9, -4))
    check(values, updates)


def test_article_source_constraints(markdown):
    for token in (
        ":chatgpt-content-reference",
        "sandbox:/mnt/data",
        "&#x20;",
        "*from*",
        "*import*",
        "*self*",
        "\\_\\_init\\_\\_",
    ):
        if token in markdown:
            raise AssertionError("unexpected export artifact: %s" % token)
    identity = "range_sum(left, right) = prefix_sum(right) - prefix_sum(left)"
    if "```python\n%s\n```" % identity in markdown:
        raise AssertionError("identity is still in a Python block")
    if "```text\n%s\n```" % identity not in markdown:
        raise AssertionError("identity text block missing")
    if "## Range Updates: Difference Arrays and Two Trees" not in markdown:
        raise AssertionError("range-update heading missing")
    if "sandbox:" in markdown:
        raise AssertionError("sandbox URL remains")
    if "/downloads/fenwick-tree-examples.py" not in markdown:
        raise AssertionError("downloadable example link missing")


def run_tests(FenwickTree, count_inversions, find_by_order):
    test_article_examples(FenwickTree, count_inversions, find_by_order)
    test_arrays(FenwickTree)
    test_invalid_bounds(FenwickTree)
    test_randomized_updates(FenwickTree)
    test_frequency_selection(FenwickTree, find_by_order)
    test_inversions(count_inversions)
    test_range_updates(FenwickTree)


def main():
    path = _article_path()
    if not path.is_file():
        raise SystemExit("article source not found: %s" % path)
    markdown = path.read_text()
    test_article_source_constraints(markdown)
    namespace = load_article_implementations(markdown)
    run_tests(namespace["FenwickTree"], namespace["count_inversions"], namespace["find_by_order"])
    values = [3, 2, 5, 1, 4, 6, 2, 7]
    tree = namespace["FenwickTree"](values)
    print(tree.prefix_sum(7))
    print(tree.range_sum(2, 6))
    tree.add(2, 4)
    print(tree.range_sum(2, 6))
    tree.set(4, 10)
    print(tree.range_sum(2, 6))
    print(tree.prefix_sum(0))
    print(tree.range_sum(3, 3))
    print(namespace["count_inversions"]([3, 1, 2]))
    print(namespace["count_inversions"]([2, 2, 1]))
    frequencies = namespace["FenwickTree"]([2, 0, 3, 1])
    print(namespace["find_by_order"](frequencies, 1))
    print(namespace["find_by_order"](frequencies, 3))
    print(namespace["find_by_order"](frequencies, 6))
    print("all checks passed")


if __name__ == "__main__":
    main()
