"""Stdlib-only tests (no pytest): plain asserts in functions called from a
__main__ block. Run: `python3 test_lsm.py`.

The headline test locks in the resurrection from the demo, so the claim in
commentary.html section 6 cannot rot silently. The rest are unit checks on
the pieces that produce it: tombstones, newest-wins reads, and the read
amplification numbers.

Each test gets a fresh scratch directory under `test-data/`, wiped first.
"""

import os
import shutil

from lsm import DEL, SET, LSMTree, SSTable

ROOT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "test-data")


def fresh(name, flush_threshold=2):
    path = os.path.join(ROOT, name)
    shutil.rmtree(path, ignore_errors=True)
    return LSMTree(path, flush_threshold=flush_threshold)


def demo_trace(tree):
    """The exact six writes demo.py makes: apple is deleted one segment
    after it is written, so its value and its tombstone end up in different
    segment files."""
    tree.put("apple", "red")
    tree.put("banana", "yellow")   # seals seg0000: apple=red, banana=yellow
    tree.delete("apple")
    tree.put("cherry", "dark")     # seals seg0001: apple=<tombstone>, cherry
    tree.put("date", "brown")
    tree.put("elder", "black")     # seals seg0002: date, elder
    return tree


def test_trace_builds_three_segments():
    tree = demo_trace(fresh("trace"))
    assert len(tree.segments) == 3
    assert [s.name for s in tree.segments] == [
        "seg0000.sst", "seg0001.sst", "seg0002.sst"]
    assert dict(tree.segments[0].items()) == {
        "apple": (SET, "red"), "banana": (SET, "yellow")}
    assert dict(tree.segments[1].items()) == {
        "apple": (DEL, ""), "cherry": (SET, "dark")}
    assert tree.get("apple") is None


def test_partial_compaction_dropping_tombstones_resurrects_the_key():
    """THE AHA. Merge the newest 2 of 3 segments and drop tombstones: the
    tombstone in seg0001 is destroyed while apple=red survives in seg0000,
    which no longer has anything newer contradicting it."""
    tree = demo_trace(fresh("resurrect"))
    assert tree.get("apple") is None
    tree.compact(2, drop_tombstones=True)
    assert tree.get("apple") == "red"


def test_partial_compaction_keeping_tombstones_is_safe():
    tree = demo_trace(fresh("keep"))
    tree.compact(2, drop_tombstones=False)
    assert dict(tree.segments[-1].items())["apple"] == (DEL, "")
    assert tree.get("apple") is None


def test_full_compaction_dropping_tombstones_is_safe():
    """Dropping is safe exactly when the merge covers every segment that
    could hold the key -- here, all three."""
    tree = demo_trace(fresh("full"))
    tree.compact(3, drop_tombstones=True)
    assert len(tree.segments) == 1
    assert "apple" not in dict(tree.segments[0].items())
    assert tree.get("apple") is None


def test_compaction_is_the_only_thing_that_reclaims_space():
    tree = demo_trace(fresh("space"))
    before = tree.store_bytes()
    tree.compact(3, drop_tombstones=True)
    assert tree.store_bytes() < before
    # 6 written entries collapse to 4 live keys in one segment.
    assert sum(1 for _ in tree.segments[0].items()) == 4


def test_deleting_every_key_grows_the_store():
    """A delete is a write, so an empty store is bigger than a full one."""
    tree = fresh("grow")
    for key, value in [("apple", "red"), ("banana", "yellow"),
                       ("cherry", "dark"), ("date", "brown")]:
        tree.put(key, value)
    written = tree.store_bytes()
    for key in ["apple", "banana", "cherry", "date"]:
        tree.delete(key)
    assert tree.store_bytes() > written
    assert (written, tree.store_bytes()) == (63, 108)
    assert all(tree.get(k) is None for k in ["apple", "banana", "cherry", "date"])


def test_a_miss_reads_every_segment():
    """Read amplification: the newest key costs one segment, a key that was
    never written costs all of them."""
    tree = fresh("amplify")
    for i in range(10):
        tree.put(f"k{i:02d}", f"v{i:02d}")
    assert len(tree.segments) == 5

    assert tree.get("k09") == "v09"
    assert (tree.last_get_segments, tree.last_get_lines) == (1, 2)
    assert tree.get("k00") == "v00"
    assert (tree.last_get_segments, tree.last_get_lines) == (5, 5)
    assert tree.get("zzz") is None
    assert (tree.last_get_segments, tree.last_get_lines) == (5, 10)


def test_newest_segment_shadows_older_ones():
    tree = fresh("shadow")
    tree.put("apple", "red")
    tree.put("banana", "yellow")
    tree.put("apple", "green")
    tree.put("cherry", "dark")
    assert tree.get("apple") == "green"
    # Both versions are still on disk; the read just never reaches the old one.
    assert dict(tree.segments[0].items())["apple"] == (SET, "red")
    assert tree.last_get_segments == 1


def test_memtable_is_read_before_any_segment():
    tree = fresh("memtable")
    tree.put("apple", "red")
    tree.put("banana", "yellow")   # flushed
    tree.put("apple", "green")     # still in the memtable
    assert tree.get("apple") == "green"
    assert tree.last_get_segments == 0


def test_segments_are_written_in_key_order():
    tree = fresh("sorted")
    tree.put("zulu", "z")
    tree.put("alpha", "a")
    keys = [k for k, _ in tree.segments[0].items()]
    assert keys == ["alpha", "zulu"]
    assert keys == sorted(keys)


def test_sstable_file_is_plain_sorted_text():
    tree = demo_trace(fresh("text"))
    with open(tree.segments[1].path) as f:
        assert f.read() == "apple\tdel\t\ncherry\tset\tdark\n"
    assert isinstance(tree.segments[1], SSTable)


TESTS = [
    test_trace_builds_three_segments,
    test_partial_compaction_dropping_tombstones_resurrects_the_key,
    test_partial_compaction_keeping_tombstones_is_safe,
    test_full_compaction_dropping_tombstones_is_safe,
    test_compaction_is_the_only_thing_that_reclaims_space,
    test_deleting_every_key_grows_the_store,
    test_a_miss_reads_every_segment,
    test_newest_segment_shadows_older_ones,
    test_memtable_is_read_before_any_segment,
    test_segments_are_written_in_key_order,
    test_sstable_file_is_plain_sorted_text,
]


if __name__ == "__main__":
    shutil.rmtree(ROOT, ignore_errors=True)
    failed = 0
    for test in TESTS:
        try:
            test()
        except AssertionError as exc:
            failed += 1
            print(f"FAIL  {test.__name__}: {exc}")
        else:
            print(f"PASS  {test.__name__}")
    shutil.rmtree(ROOT, ignore_errors=True)
    if failed:
        print(f"\n{failed} of {len(TESTS)} tests FAILED")
        raise SystemExit(1)
    print(f"\nAll {len(TESTS)} tests PASSED")
