"""Rebuild an order book by replaying update rows, and measure it against a real snapshot.

Usage:
    python orderbook_warmup.py binance_futures hour1.parquet [hour2.parquet ...]
    python orderbook_warmup.py bybit hour1.parquet [hour2.parquet ...] --json results.json

Give it consecutive hourly CryptoHFTData order book files (already plain Parquet).
The first snapshot in the stream anchors a *reference* book. A second *replayed* book
starts empty at that same moment. Both books then receive exactly the same update
messages, and at each checkpoint the script compares them. Because both books see the
same deltas, every difference is a level the replayed book has simply never seen change.

If the files contain no snapshot at all, the script still replays the updates from an
empty book and reports how many levels it holds at each checkpoint; there is just no
reference to compare against.

Requires: pip install pyarrow
"""
import argparse
import json
from decimal import Decimal
from itertools import groupby

CHECKPOINTS_S = (10, 30, 60, 120, 300, 600, 900, 1200, 1800, 2700, 3600, 5400)
TOP_N = (1, 5, 10, 20, 50)
BANDS_BPS = (10, 25, 50, 100)

UPDATE_KEY = (
    "symbol", "event_type", "received_time", "event_time", "transaction_time",
    "first_update_id", "final_update_id", "prev_final_update_id", "last_update_id",
)


# --------------------------------------------------------------------------- reading

def parquet_rows(paths):
    """Stream rows in physical file order; integer IDs stay exact."""
    import pyarrow.parquet as pq

    for path in paths:
        for batch in pq.ParquetFile(path).iter_batches(batch_size=65536):
            yield from batch.to_pylist()


def message_key(row):
    """Rows of one exchange message share these columns.

    Snapshot rows can carry per-row ``received_time`` (a snapshot may be written in two
    halves), so a snapshot is identified by its exchange timestamp and cursor instead.
    """
    if row["event_type"] == "snapshot":
        return ("snapshot", row["symbol"], row["event_time"], row["final_update_id"], row["last_update_id"])
    return tuple(row[k] for k in UPDATE_KEY)


def messages(rows):
    """Group *consecutive* rows into exchange messages. Never re-sort the rows first."""
    for _, group in groupby(rows, key=message_key):
        yield list(group)


# --------------------------------------------------------------------------- the book

class Book:
    """A plain L2 book: {side: {price: size}}."""

    def __init__(self):
        self.levels = {"bid": {}, "ask": {}}

    def apply(self, message):
        if message[0]["event_type"] == "snapshot":
            self.levels = {"bid": {}, "ask": {}}  # Clear once per snapshot, then load it.
        for row in message:
            side = row["side"]
            if side == "noop":  # Sequence-only marker, not a price level.
                continue
            price, size = Decimal(row["price"]), Decimal(row["quantity"])
            if size == 0:
                self.levels[side].pop(price, None)  # Deleting an unknown level is fine.
            else:
                self.levels[side][price] = size  # The NEW size at this price, not a delta.

    def top(self, side, n):
        return sorted(self.levels[side].items(), reverse=(side == "bid"))[:n]

    def mid(self):
        bids, asks = self.top("bid", 1), self.top("ask", 1)
        if not bids or not asks:
            return None
        return (bids[0][0] + asks[0][0]) / 2


# --------------------------------------------------------------------------- comparing

def compare(replayed, reference):
    """How does the replayed book differ from the reference at this instant?"""
    out = {"ref_levels": 0, "replayed_levels": 0, "missing": 0, "extra": 0, "wrong_size": 0}
    for side in ("bid", "ask"):
        ref, rep = reference.levels[side], replayed.levels[side]
        out["ref_levels"] += len(ref)
        out["replayed_levels"] += len(rep)
        out["missing"] += len(ref.keys() - rep.keys())
        out["extra"] += len(rep.keys() - ref.keys())
        out["wrong_size"] += sum(ref[p] != rep[p] for p in ref.keys() & rep.keys())
    for n in TOP_N:
        out[f"top{n}_equal"] = all(replayed.top(s, n) == reference.top(s, n) for s in ("bid", "ask"))
    mid = reference.mid()
    if mid:
        out["mid"] = str(mid)
        nearest = None
        for side in ("bid", "ask"):
            for price in reference.levels[side].keys() - replayed.levels[side].keys():
                distance = abs(price - mid) / mid * 10000
                nearest = distance if nearest is None or distance < nearest else nearest
        out["nearest_missing_bps"] = float(nearest) if nearest is not None else None
        for band in BANDS_BPS:
            out[f"missing_within_{band}bps"] = sum(
                1 for side in ("bid", "ask")
                for price in reference.levels[side].keys() - replayed.levels[side].keys()
                if abs(price - mid) / mid * 10000 <= band)
    return out


# --------------------------------------------------------------------------- sequencing

def cursor(row, exchange):
    if exchange == "binance_futures":
        return row["last_update_id"] if row["event_type"] == "snapshot" else row["final_update_id"]
    return row["last_update_id"]  # Bybit cross-sequence ``seq``.


class Sequencer:
    """Checks that updates form an unbroken chain from the anchor onwards."""

    def __init__(self, exchange):
        self.exchange = exchange
        self.previous = None
        self.bridged = exchange != "binance_futures"

    def anchor(self, snapshot_row):
        self.previous = cursor(snapshot_row, self.exchange)
        self.bridged = self.exchange != "binance_futures"

    def accept(self, row):
        """Return False to skip a stale update, True to apply it; raise on a gap."""
        if self.previous is None:
            self.previous = cursor(row, self.exchange)
            return True
        if self.exchange == "binance_futures":
            if not self.bridged:  # First update after a REST snapshot: Binance bridge rule.
                if row["final_update_id"] < self.previous:
                    return False
                if not (row["first_update_id"] <= self.previous <= row["final_update_id"]
                        or row["prev_final_update_id"] == self.previous):
                    raise ValueError("cannot bridge the snapshot to the update stream")
                self.bridged = True
            elif row["prev_final_update_id"] != self.previous:
                raise ValueError(f"Binance chain broken: pu={row['prev_final_update_id']} "
                                 f"expected {self.previous}")
        elif row["last_update_id"] <= self.previous:
            raise ValueError("Bybit seq did not advance; check for a reset or reordering")
        self.previous = cursor(row, self.exchange)
        return True


# --------------------------------------------------------------------------- experiment

def run(exchange, message_iter, checkpoints=CHECKPOINTS_S):
    replayed, reference = Book(), None
    seq = Sequencer(exchange)
    start_ms = None
    pending = list(checkpoints)
    results = []
    updates = 0

    def record(label, elapsed_s, book=None):
        row = {"label": label, "elapsed_s": round(elapsed_s, 3), "updates": updates}
        if reference is not None:
            row.update(compare(replayed, book or reference))
        else:
            row["replayed_levels"] = sum(len(v) for v in replayed.levels.values())
        results.append(row)

    for message in message_iter:
        row = message[0]
        if start_ms is None:
            start_ms = row["event_time"]
        elapsed = (row["event_time"] - start_ms) / 1000

        if row["event_type"] == "snapshot":
            if reference is None:
                # The first snapshot anchors the experiment. Anything replayed before it is
                # discarded so both books start at exactly this instant.
                reference, replayed = Book(), Book()
                reference.apply(message)
                seq.anchor(row)
                start_ms, updates, pending, results = row["event_time"], 0, list(checkpoints), []
                continue  # The replayed book deliberately stays empty.
            # A later snapshot: an exact comparison is only fair at the same cursor.
            later = Book()
            later.apply(message)
            if cursor(row, exchange) == seq.previous:
                record("next snapshot (same cursor)", elapsed, later)
                if reference is not None:
                    check = compare(reference, later)
                    results[-1]["reference_matches_next_snapshot"] = (
                        check["missing"] == check["extra"] == check["wrong_size"] == 0)
            else:
                results.append({"label": "next snapshot has a different cursor (gap or reconnect); stopped",
                                "elapsed_s": round(elapsed, 3), "updates": updates})
            return results

        if not seq.accept(row):
            continue
        replayed.apply(message)
        if reference is not None:
            reference.apply(message)
        updates += 1
        while pending and elapsed >= pending[0]:
            target = pending.pop(0)
            record(f"{target}s", elapsed)
            results[-1]["target_s"] = target

    if updates:
        record("end of data", elapsed)
    return results


def print_results(results):
    for r in results:
        if "replayed_levels" not in r:
            print(f"{r['label']} after {r['updates']:,} updates")
            continue
        if "ref_levels" not in r:
            print(f"{r['label']:>34} | replayed book holds {r['replayed_levels']:,} levels "
                  f"after {r['updates']:,} updates (no snapshot to compare against)")
            continue
        tops = " ".join(f"top{n}={'ok' if r[f'top{n}_equal'] else 'DIFF'}" for n in TOP_N)
        nearest = r.get("nearest_missing_bps")
        print(f"{r['label']:>34} | levels ref={r['ref_levels']:,} replayed={r['replayed_levels']:,} "
              f"missing={r['missing']:,} extra={r['extra']} wrong_size={r['wrong_size']} | {tops} | "
              f"nearest missing level: {'none' if nearest is None else f'{nearest:.1f} bps from mid'}")


def main():
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("exchange", choices=("binance_futures", "bybit"))
    parser.add_argument("files", nargs="+", help="consecutive hourly *_orderbook.parquet files")
    parser.add_argument("--json", help="also write the checkpoint results to this file")
    args = parser.parse_args()
    results = run(args.exchange, messages(parquet_rows(args.files)))
    print_results(results)
    if args.json:
        with open(args.json, "w") as handle:
            json.dump(results, handle, indent=1)


if __name__ == "__main__":
    main()
