"""A B+tree index living in one flat file of fixed-size pages.

The point of the toy is the *page*, not the tree. Every read and write goes
through `Pager` a whole page at a time, so `pager.reads` counts exactly what a
real index cares about: how many blocks a lookup has to fetch. Keys and values
are fixed-width 64-bit integers and every node is one page, which is what
keeps the whole thing inside the LOC budget.

Page ids come from a monotonic allocator and nothing is ever freed, so the
same insertion order always produces a byte-identical file.
"""

import os
import struct

PAGE_SIZE = 1024
MAX_KEYS = 63                     # "order 63": a node holds at most 63 keys
HEADER = struct.Struct("!BHI")    # leaf flag, key count, leftmost child pid
SLOT = struct.Struct("!qq")       # the one record format: key, payload


class Pager:
    """The disk. One file; page id `p` is the bytes at offset p * PAGE_SIZE.

    There is no cache on purpose: `reads` is then a true count of page
    fetches, not of cache hits.
    """

    def __init__(self, path):
        if os.path.exists(path):
            os.remove(path)
        self.path = path
        self.f = open(path, "w+b")
        self.next_pid = 0
        self.reads = 0
        self.writes = 0

    def alloc(self):
        pid = self.next_pid
        self.next_pid += 1
        return pid

    def read_page(self, pid):
        self.reads += 1
        self.f.seek(pid * PAGE_SIZE)
        return self.f.read(PAGE_SIZE)

    def write_page(self, pid, data):
        self.writes += 1
        self.f.seek(pid * PAGE_SIZE)
        self.f.write(data)

    def close(self):
        self.f.close()


class Node:
    """One page, decoded. `kids` is values in a leaf, child page ids inside.

    An internal node holds n keys and n+1 children, which does not divide into
    equal records. Rather than invent a second record format, child 0 goes in
    the page header and slot i carries (keys[i], kids[i+1]).
    """

    __slots__ = ("leaf", "keys", "kids")

    def __init__(self, leaf, keys, kids):
        self.leaf = leaf
        self.keys = keys
        self.kids = kids

    def pack(self):
        out = [HEADER.pack(self.leaf, len(self.keys), 0 if self.leaf else self.kids[0])]
        for key, payload in zip(self.keys, self.kids if self.leaf else self.kids[1:]):
            out.append(SLOT.pack(key, payload))
        return b"".join(out).ljust(PAGE_SIZE, b"\x00")

    @staticmethod
    def unpack(buf):
        leaf, count, first = HEADER.unpack_from(buf, 0)
        keys, kids = [], [] if leaf else [first]
        for i in range(count):
            key, payload = SLOT.unpack_from(buf, HEADER.size + i * SLOT.size)
            keys.append(key)
            kids.append(payload)
        return Node(bool(leaf), keys, kids)


def scan(keys, key, leaf):
    """Linear scan for `key`. Returns (slot index, comparisons performed).

    In a leaf: the first slot whose key is >= `key` — where `key` belongs.
    Inside:    the child to descend into, so keys[i-1] <= key < keys[i].
    """
    i = 0
    while i < len(keys) and (keys[i] < key if leaf else keys[i] <= key):
        i += 1
    return i, min(i + 1, len(keys))


class BPlusTree:
    """Insert and point-lookup only. No delete, so `pages` never shrinks."""

    def __init__(self, pager):
        self.pager = pager
        self.root = pager.alloc()
        pager.write_page(self.root, Node(True, [], []).pack())
        self.height = 1
        self.comparisons = 0
        self.inserted = 0
        self.growth = []          # (insert number, new height) for every root split
        self.last_path = []       # page ids touched by the last get()

    def get(self, key):
        pid = self.root
        self.last_path = []
        while True:
            self.last_path.append(pid)
            node = Node.unpack(self.pager.read_page(pid))
            i, comps = scan(node.keys, key, node.leaf)
            self.comparisons += comps
            if not node.leaf:
                pid = node.kids[i]
                continue
            self.comparisons += 1
            if i < len(node.keys) and node.keys[i] == key:
                return node.kids[i]
            return None

    def insert(self, key, value):
        self.inserted += 1
        split = self._insert(self.root, key, value)
        if split is not None:
            sep, right = split
            left, self.root = self.root, self.pager.alloc()
            self.pager.write_page(self.root, Node(False, [sep], [left, right]).pack())
            self.height += 1
            self.growth.append((self.inserted, self.height))

    def _insert(self, pid, key, value):
        """Insert below `pid`. Returns None, or (separator, new right sibling)."""
        node = Node.unpack(self.pager.read_page(pid))
        i, comps = scan(node.keys, key, node.leaf)
        self.comparisons += comps
        if node.leaf:
            if i < len(node.keys) and node.keys[i] == key:
                node.kids[i] = value
            else:
                node.keys.insert(i, key)
                node.kids.insert(i, value)
        else:
            split = self._insert(node.kids[i], key, value)
            if split is None:
                return None       # nothing below changed shape: this page is untouched
            sep, right = split
            node.keys.insert(i, sep)
            node.kids.insert(i + 1, right)
        if len(node.keys) > MAX_KEYS:
            return self._split(pid, node)
        self.pager.write_page(pid, node.pack())
        return None

    def _split(self, pid, node):
        """Halve an overfull node in place; the new right half gets a fresh pid."""
        mid = len(node.keys) // 2
        sep = node.keys[mid]
        if node.leaf:
            # The separator is *copied* up: a B+tree keeps every key in a leaf.
            left = Node(True, node.keys[:mid], node.kids[:mid])
            right = Node(True, node.keys[mid:], node.kids[mid:])
        else:
            # The separator *moves* up; the children either side of it split.
            left = Node(False, node.keys[:mid], node.kids[:mid + 1])
            right = Node(False, node.keys[mid + 1:], node.kids[mid + 1:])
        rpid = self.pager.alloc()
        self.pager.write_page(pid, left.pack())
        self.pager.write_page(rpid, right.pack())
        return sep, rpid

    def stats(self):
        """Walk every allocated page. Nothing is ever freed, so page ids are dense."""
        leaves = keys = 0
        for pid in range(self.pager.next_pid):
            node = Node.unpack(self.pager.read_page(pid))
            if node.leaf:
                leaves += 1
                keys += len(node.keys)
        return {
            "height": self.height,
            "pages": self.pager.next_pid,
            "leaves": leaves,
            "fill": keys / (leaves * MAX_KEYS),
            "bytes": self.pager.next_pid * PAGE_SIZE,
        }
