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

The headline test spawns a real child process and really SIGKILLs it, so
it pins the claim the commentary is built on: with fsync switched off,
kill -9 still returns every single record.
"""

import os
import signal
import subprocess
import sys
import tempfile

from wal_kv import WalKV, simulate_power_loss, _frame, _encode, _PUT

DIR = tempfile.mkdtemp(prefix="wal-kv-test-")


def path(name):
    return os.path.join(DIR, name)


def test_sigkill_without_fsync_loses_nothing():
    """The aha: fsync off, real SIGKILL, and not one record is missing."""
    p = path("kill.log")
    child = (
        "import os, signal, sys; sys.path.insert(0, %r);"
        "from wal_kv import WalKV;"
        "kv = WalKV(%r, flush=True, fsync=False);"
        "[kv.put(b'k%%02d' %% i, b'v%%02d' %% i) for i in range(5)];"
        "os.kill(os.getpid(), signal.SIGKILL)"
        % (os.path.dirname(os.path.abspath(__file__)), p)
    )
    proc = subprocess.run([sys.executable, "-c", child], capture_output=True)

    assert proc.returncode == -signal.SIGKILL, proc.returncode
    kv = WalKV(p)
    assert kv.replayed == 5, kv.replayed
    assert kv.get(b"k04") == b"v04"
    kv.close()


def test_power_loss_without_fsync_loses_everything():
    """Same bytes, different failure: nothing was ever pushed to disk."""
    p = path("power.log")
    kv = WalKV(p, flush=True, fsync=False)
    for i in range(5):
        kv.put(b"k%02d" % i, b"v%02d" % i)
    kv.close()

    simulate_power_loss(p, 0)  # fsync was never called, so zero durable bytes
    recovered = WalKV(p)
    assert recovered.replayed == 0, recovered.replayed
    assert recovered.data == {}
    recovered.close()


def test_torn_tail_is_discarded_and_the_rest_survives():
    """A crash mid-append leaves half a record; replay stops, keeps the rest."""
    p = path("torn.log")
    kv = WalKV(p, fsync=True)
    kv.put(b"alpha", b"one")
    kv.put(b"beta", b"two")
    kv.close()

    full = os.path.getsize(p)
    with open(p, "r+b") as f:
        f.truncate(full - 4)  # chop the last record in half

    recovered = WalKV(p)
    assert recovered.replayed == 1, recovered.replayed
    assert recovered.get(b"alpha") == b"one"
    assert recovered.get(b"beta") is None
    assert recovered.discarded > 0, recovered.discarded
    recovered.close()


def test_corrupt_payload_is_caught_by_the_checksum():
    """Framing intact, bytes rotten -- only the crc can tell."""
    p = path("corrupt.log")
    kv = WalKV(p, fsync=True)
    kv.put(b"alpha", b"one")
    kv.put(b"beta", b"two")
    kv.close()

    with open(p, "r+b") as f:
        f.seek(os.path.getsize(p) - 1)
        f.write(b"X")  # same length, different content

    recovered = WalKV(p)
    assert recovered.replayed == 1, recovered.replayed
    assert recovered.get(b"beta") is None
    recovered.close()


def test_delete_is_a_tombstone_that_survives_replay():
    """Without a delete record, replay would resurrect the key."""
    p = path("tomb.log")
    kv = WalKV(p, fsync=True)
    kv.put(b"alpha", b"one")
    kv.delete(b"alpha")
    kv.close()

    recovered = WalKV(p)
    assert recovered.replayed == 2, recovered.replayed  # put AND delete
    assert recovered.get(b"alpha") is None
    recovered.close()


def test_replay_applies_records_in_order():
    """Last write wins, because the log is ordered and replay respects it."""
    p = path("order.log")
    kv = WalKV(p, fsync=True)
    kv.put(b"alpha", b"first")
    kv.put(b"alpha", b"second")
    kv.close()

    recovered = WalKV(p)
    assert recovered.get(b"alpha") == b"second"
    recovered.close()


def test_trailing_garbage_does_not_resurrect_as_a_record():
    """Bytes appended after a clean log are junk, not data."""
    p = path("garbage.log")
    kv = WalKV(p, fsync=True)
    kv.put(b"alpha", b"one")
    kv.close()

    with open(p, "ab") as f:
        f.write(_frame(_encode(_PUT, b"ghost", b"boo"))[:6])  # half a header

    recovered = WalKV(p)
    assert recovered.replayed == 1, recovered.replayed
    assert recovered.get(b"ghost") is None
    assert recovered.discarded == 6, recovered.discarded
    recovered.close()


TESTS = [
    test_sigkill_without_fsync_loses_nothing,
    test_power_loss_without_fsync_loses_everything,
    test_torn_tail_is_discarded_and_the_rest_survives,
    test_corrupt_payload_is_caught_by_the_checksum,
    test_delete_is_a_tombstone_that_survives_replay,
    test_replay_applies_records_in_order,
    test_trailing_garbage_does_not_resurrect_as_a_record,
]


if __name__ == "__main__":
    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__}")
    if failed:
        print(f"\n{failed} of {len(TESTS)} tests FAILED")
        raise SystemExit(1)
    print(f"\nAll {len(TESTS)} tests PASSED")
