#!/usr/bin/env python3
"""Compare two forward-test (dry-run) Freqtrade databases over the same window.

Built for the bot3 GrokCombinedD48 vs bot1 D48 evaluation
(docs/strategy-iterations.md 07-11 节: 前向观察 7~14 天后比较信号数、PF、
止损率、平均盈亏和持仓时间).  The window defaults to [first trade open in
db-b, now].  Trades already open at the window boundary are reported separately;
unequal carried-in exposure invalidates a clean entry-filter A/B because portfolio
limits can suppress otherwise valid signals.

Stdlib only; pipe to the deployment host:

    ssh <host> python3 - \
        --db-a /data/freqtrade/user_data/tradesv3-v5.sqlite --label-a D48 \
        --db-b /data/freqtrade/user_data/tradesv3-grok-combined-d48-forward.sqlite \
        --label-b GrokCombined \
        < freqtrade/scripts/forward_compare.py
"""

from __future__ import annotations

import argparse
import sqlite3
from datetime import datetime, timezone

STOP_REASONS = ("stop_loss", "stoploss_on_exchange")


def normalize_since(value: str) -> str:
    """Normalize user-provided ISO timestamps to SQLite's UTC text format."""
    parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
    if parsed.tzinfo is not None:
        parsed = parsed.astimezone(timezone.utc).replace(tzinfo=None)
    return parsed.isoformat(sep=" ", timespec="microseconds")


def summarize(db: str, since: str) -> dict:
    con = sqlite3.connect(f"file:{db}?mode=ro", uri=True)
    con.row_factory = sqlite3.Row
    rows = list(con.execute("SELECT * FROM trades WHERE open_date >= ?", (since,)))
    carried = list(
        con.execute(
            """
            SELECT * FROM trades
            WHERE open_date < ? AND (close_date IS NULL OR close_date >= ?)
            """,
            (since, since),
        )
    )
    closed = [r for r in rows if not r["is_open"]]
    wins = [r for r in closed if (r["close_profit"] or 0) > 0]
    losses = [r for r in closed if (r["close_profit"] or 0) <= 0]
    gross_win = sum(r["close_profit_abs"] or 0 for r in wins)
    gross_loss = abs(sum(r["close_profit_abs"] or 0 for r in losses))

    def hold_hours(r) -> float:
        open_dt = datetime.fromisoformat(r["open_date"])
        end = datetime.fromisoformat(r["close_date"]) if r["close_date"] else datetime.utcnow()
        return (end - open_dt).total_seconds() / 3600

    return {
        "signals": len(rows),
        "open": len(rows) - len(closed),
        "closed": len(closed),
        "carried_in": len(carried),
        "carried_in_long": sum(1 for r in carried if not r["is_short"]),
        "carried_in_short": sum(1 for r in carried if r["is_short"]),
        "carried_in_closed": sum(1 for r in carried if not r["is_open"]),
        "carried_in_net_usdt": sum(
            r["close_profit_abs"] or 0 for r in carried if not r["is_open"]
        ),
        "carried_in_pairs": [
            f"{r['pair']}:{'short' if r['is_short'] else 'long'}" for r in carried
        ],
        "winrate": len(wins) / len(closed) if closed else None,
        "net_usdt": sum(r["close_profit_abs"] or 0 for r in closed),
        "pf": (gross_win / gross_loss) if gross_loss else None,
        "avg_profit_pct": (
            sum(r["close_profit"] or 0 for r in closed) / len(closed) * 100 if closed else None
        ),
        # stoploss_on_exchange also fires after trailing moved the stop into
        # profit, so a losing close is what identifies a genuine hard stop.
        "stop_rate": (
            sum(
                1 for r in closed
                if r["exit_reason"] in STOP_REASONS and (r["close_profit"] or 0) < 0
            ) / len(closed)
            if closed else None
        ),
        "avg_hold_h_closed": (
            sum(hold_hours(r) for r in closed) / len(closed) if closed else None
        ),
    }


def fmt(value, pattern: str) -> str:
    return pattern.format(value) if value is not None else "-"


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--db-a", required=True)
    parser.add_argument("--db-b", required=True)
    parser.add_argument("--label-a", default="A")
    parser.add_argument("--label-b", default="B")
    parser.add_argument(
        "--since",
        help="ISO date/datetime; default = first trade open_date in db-b",
    )
    args = parser.parse_args()

    since = args.since
    if not since:
        con_b = sqlite3.connect(f"file:{args.db_b}?mode=ro", uri=True)
        since = con_b.execute("SELECT min(open_date) FROM trades").fetchone()[0]
        if not since:
            print("db-b has no trades yet; nothing to compare")
            return 0
    since = normalize_since(since)

    now = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M UTC")
    print(f"window: {since} -> {now}")
    a = summarize(args.db_a, since)
    b = summarize(args.db_b, since)

    metrics = [
        ("signals (trades opened)", "signals", "{}"),
        ("still open", "open", "{}"),
        ("closed", "closed", "{}"),
        ("carried in at window start", "carried_in", "{}"),
        ("  carried long", "carried_in_long", "{}"),
        ("  carried short", "carried_in_short", "{}"),
        ("carried-in trades total USDT", "carried_in_net_usdt", "{:+.2f}"),
        ("winrate", "winrate", "{:.1%}"),
        ("net USDT", "net_usdt", "{:+.2f}"),
        ("profit factor", "pf", "{:.2f}"),
        ("avg profit %/trade", "avg_profit_pct", "{:+.2f}"),
        ("hard-stop rate", "stop_rate", "{:.1%}"),
        ("avg hold h (closed)", "avg_hold_h_closed", "{:.1f}"),
    ]
    width = max(len(m[0]) for m in metrics)
    print(f"{'':<{width}}  {args.label_a:>14} {args.label_b:>14}")
    for label, key, pattern in metrics:
        print(f"{label:<{width}}  {fmt(a[key], pattern):>14} {fmt(b[key], pattern):>14}")
    if a["carried_in_pairs"] or b["carried_in_pairs"]:
        print(f"carried-in {args.label_a}: {', '.join(a['carried_in_pairs']) or '-'}")
        print(f"carried-in {args.label_b}: {', '.join(b['carried_in_pairs']) or '-'}")
    if (
        a["carried_in"] != b["carried_in"]
        or a["carried_in_long"] != b["carried_in_long"]
        or a["carried_in_short"] != b["carried_in_short"]
    ):
        print(
            "WARNING: unequal carried-in exposure — this is not a clean entry-filter A/B; "
            "position limits can suppress otherwise valid signals."
        )
    if min(a["closed"], b["closed"]) < 10:
        print("NOTE: <10 closed trades on at least one side — directional read only.")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
