#!/usr/bin/env python3
"""Replay maker fills with a stale-quote cancellation rule.

The production simulator only expires quotes after 30 minutes.  This tool uses
the independently captured public CLOB stream to ask a narrower question: how
much of the realised maker P&L survives if a quote is cancelled when a later
book snapshot shows that the external best bid has moved below our price?

`price_change` was not persisted before adaptive-maker v1, so this replay is a
conservative diagnostic, not proof that the cancel could always beat a trade.
"""

from __future__ import annotations

import argparse
import bisect
import json
import sqlite3
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime


def ts(value: str) -> float:
    return datetime.fromisoformat(value.replace("Z", "+00:00")).timestamp()


@dataclass
class Event:
    time: float
    kind: str
    price: float | None
    size: float | None
    side: str | None
    best_bid: float | None


def load_events(path: str) -> tuple[dict[str, list[Event]], dict[str, list[float]]]:
    conn = sqlite3.connect(f"file:{path}?mode=ro", uri=True)
    events: dict[str, list[Event]] = defaultdict(list)
    query = """SELECT received_ts, token_id, event_type, price, size, side,
                      payload_json
               FROM clob_l2_events
               WHERE token_id IS NOT NULL
               ORDER BY token_id, received_ts, id"""
    for received, token, kind, price, size, side, payload in conn.execute(query):
        best_bid = None
        if kind == "book":
            bids = json.loads(payload).get("bids", [])
            prices = [float(x["price"]) for x in bids if float(x.get("size", 0)) > 0]
            best_bid = max(prices, default=None)
        events[token].append(Event(ts(received), kind, price, size, side, best_bid))
    conn.close()
    return events, {token: [x.time for x in rows] for token, rows in events.items()}


def orders(path: str, since: str):
    conn = sqlite3.connect(f"file:{path}?mode=ro", uri=True)
    conn.row_factory = sqlite3.Row
    query = """SELECT m.*, s.signal_price, s.edge, s.model_prob,
                      s.model_prob_corr, s.reference_bid, s.reference_ask,
                      s.best_bid_size, s.best_ask_size, s.kind, s.city,
                      CASE WHEN m.side='NO'
                           THEN (s.bucket != st.winning_bucket)
                           ELSE (s.bucket = st.winning_bucket) END payout
               FROM maker_sim_orders m
               JOIN snapshots s ON s.id=m.source_snapshot_id
               JOIN settlements st ON st.event_slug=m.event_slug
               WHERE m.created_ts >= ? AND s.signal_price IS NOT NULL
               ORDER BY m.created_ts"""
    rows = conn.execute(query, (since,)).fetchall()
    conn.close()
    return rows


def replay(row, stream: list[Event], times: list[float], cancel_ticks: int | None):
    start, end = ts(row["created_ts"]), ts(row["expires_ts"])
    at_price = 0.0
    filled = 0.0
    cancelled = False
    cancel_time = None
    first_fill = None
    i = bisect.bisect_left(times, start)
    while i < len(stream) and stream[i].time <= end:
        event = stream[i]
        if event.kind == "book" and cancel_ticks is not None and event.best_bid is not None:
            floor = row["price"] - cancel_ticks * 0.01
            if event.best_bid < floor - 1e-9:
                cancelled, cancel_time = True, event.time
                break
        elif (
            event.kind == "last_trade_price"
            and (event.side or "").upper() == "SELL"
            and event.price is not None
            and event.size is not None
            and event.price <= row["price"] + 1e-9
        ):
            if event.price < row["price"] - 1e-9:
                gained = event.size
            else:
                next_at = at_price + event.size
                gained = max(0.0, next_at - row["queue_ahead"]) - max(
                    0.0, at_price - row["queue_ahead"]
                )
                at_price = next_at
            gained = min(gained, row["shares"] - filled)
            if gained > 0:
                first_fill = first_fill or event.time
                filled += gained
            if filled + 1e-9 >= row["shares"]:
                break
        i += 1
    return filled, cancelled, cancel_time, first_fill


def summarize(label: str, values: list[tuple], split: float | None = None):
    selected = values if split is None else [x for x in values if x[0] >= split]
    qty = sum(x[1] for x in selected)
    stake = sum(x[1] * x[2] for x in selected)
    pnl = sum(x[1] * (x[3] - x[2]) for x in selected)
    active = sum(x[1] > 0 for x in selected)
    cancelled = sum(x[4] for x in selected)
    roi = 100 * pnl / stake if stake else 0
    print(
        f"{label:16s} orders={len(selected):6d} fills={active:5d} "
        f"qty={qty:10.1f} stake=${stake:10.2f} pnl=${pnl:9.2f} "
        f"roi={roi:7.2f}% cancelled={cancelled:5d}"
    )


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--db", required=True)
    parser.add_argument("--l2-db", required=True)
    parser.add_argument("--since", default="2026-08-03T06:33:24+00:00")
    parser.add_argument("--test-since", default="2026-08-07T00:00:00+00:00")
    args = parser.parse_args()

    streams, event_times = load_events(args.l2_db)
    rows = orders(args.db, args.since)
    test_since = ts(args.test_since)
    print(f"settled orders={len(rows)}, tokens={len(streams)}")
    for ticks in [None, 0, 1, 2, 3]:
        values = []
        for row in rows:
            stream = streams.get(row["token_id"], [])
            times = event_times.get(row["token_id"], [])
            filled, cancelled, _, _ = replay(row, stream, times, ticks)
            values.append(
                (ts(row["created_ts"]), filled, row["price"], row["payout"], cancelled)
            )
        name = "no-cancel" if ticks is None else f"cancel-{ticks}t"
        summarize(name, values)
        summarize(name + " test", values, test_since)


if __name__ == "__main__":
    main()
