#!/usr/bin/env python3
"""
query_db.py — Query the sweep database and surface top-N rows by any metric.

Primary use case: find parameter sets with strong OOS performance that weren't
the IS winner (e.g., lower IS composite but positive OOS P&L), and export them
as candidate CSVs to evaluate for production use or cross-asset analysis.

Usage:
    python3 tools/query_db.py --asset COINBASE_BTCUSD --tf 4H --top 20
    python3 tools/query_db.py --asset COINBASE_BTCUSD --tf 4H --sort calmar_ratio --top 10
    python3 tools/query_db.py --list-columns
    python3 tools/query_db.py --asset COINBASE_BTCUSD --top 50 --min-trades 15 --out results/candidates.csv
"""

import argparse
import os
import sqlite3
import sys

import pandas as pd

DB_PATH = "results/sweep_database.db"

# Columns shown in the summary table (not all 80 columns)
SUMMARY_COLS = [
    "id", "asset", "timeframe", "composite_score", "calmar_ratio",
    "sortino_ratio", "total_pnl_pct", "max_drawdown", "total_trades",
    "pct_in_market", "subperiod_consistent", "iteration_number",
]

SORTABLE_METRICS = [
    "composite_score", "calmar_ratio", "sortino_ratio", "sharpe_ratio",
    "total_pnl_pct", "pnl_dd_ratio", "total_trades", "pct_in_market",
]


def load_db(asset=None, tf=None, min_trades=None, subperiod_only=False):
    if not os.path.exists(DB_PATH):
        sys.exit(f"DB not found: {DB_PATH}")
    conn = sqlite3.connect(DB_PATH)
    query = "SELECT * FROM sweep_results WHERE 1=1"
    params = []
    if asset:
        query += " AND asset = ?"
        params.append(asset)
    if tf:
        query += " AND timeframe = ?"
        params.append(tf)
    if min_trades:
        query += " AND total_trades >= ?"
        params.append(min_trades)
    if subperiod_only:
        query += " AND subperiod_consistent = 1"
    df = pd.read_sql_query(query, conn, params=params)
    conn.close()
    return df


def main():
    parser = argparse.ArgumentParser(
        description="Query sweep_database.db and surface top-N rows by any metric."
    )
    parser.add_argument("--asset", help="Filter by asset (e.g. COINBASE_BTCUSD)")
    parser.add_argument("--tf", "--timeframe", dest="tf", help="Filter by timeframe (e.g. 4H)")
    parser.add_argument(
        "--sort", default="composite_score",
        help=f"Column to sort by descending (default: composite_score). "
             f"Common options: {', '.join(SORTABLE_METRICS)}"
    )
    parser.add_argument("--top", type=int, default=20, help="Number of rows to return (default: 20)")
    parser.add_argument("--min-trades", type=int, help="Minimum total_trades filter")
    parser.add_argument(
        "--subperiod-only", action="store_true",
        help="Only include rows where subperiod_consistent=1 (Calmar>0 in both IS halves)"
    )
    parser.add_argument(
        "--out", help="Write full param rows (all columns) to this CSV path"
    )
    parser.add_argument(
        "--list-columns", action="store_true",
        help="Print all available column names and exit"
    )
    parser.add_argument(
        "--show-weights", action="store_true",
        help="Include i_w_* weight columns in the printed table"
    )
    args = parser.parse_args()

    if args.list_columns:
        if not os.path.exists(DB_PATH):
            sys.exit(f"DB not found: {DB_PATH}")
        conn = sqlite3.connect(DB_PATH)
        cols = [r[1] for r in conn.execute("PRAGMA table_info(sweep_results)").fetchall()]
        conn.close()
        print("\n".join(cols))
        return

    df = load_db(
        asset=args.asset,
        tf=args.tf,
        min_trades=args.min_trades,
        subperiod_only=args.subperiod_only,
    )

    if df.empty:
        print("No rows matched the filters.")
        return

    if args.sort not in df.columns:
        sys.exit(
            f"Sort column '{args.sort}' not found in DB. "
            f"Use --list-columns to see available columns."
        )

    df_sorted = df.sort_values(args.sort, ascending=False).head(args.top)

    # Print summary
    combos = df.groupby(["asset", "timeframe"]).size()
    print(f"\nDB total rows matching filters: {len(df):,}")
    print(f"Combos: {', '.join(f'{a}/{t}({n})' for (a, t), n in combos.items())}")
    print(f"\nTop {args.top} by {args.sort}:\n")

    display_cols = SUMMARY_COLS.copy()
    if args.show_weights:
        display_cols += [c for c in df.columns if c.startswith("i_w_")]

    # Only show columns that exist in this DB
    display_cols = [c for c in display_cols if c in df_sorted.columns]
    pd.set_option("display.max_columns", None)
    pd.set_option("display.width", 200)
    pd.set_option("display.float_format", "{:.4f}".format)
    print(df_sorted[display_cols].to_string(index=False))

    if args.out:
        df_sorted.to_csv(args.out, index=False)
        print(f"\nFull rows written to: {args.out}")
        print(f"  {len(df_sorted)} rows × {len(df_sorted.columns)} columns")


if __name__ == "__main__":
    main()
