#!/usr/bin/env python3
"""
Binary Options Strategy Backtester — CLI entry point.

Workflow:
  1. python main.py --fetch        → download H1 data from DB
  2. python main.py --test         → run backtest on all 5 strategies
  3. python main.py --best         → show profitable strategies
  4. python main.py --report STR   → detailed report for a strategy
"""

import argparse
import logging
import sys
import os

os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'

from config import Config
from data_fetcher import DataFetcher
from indicators import compute_all_indicators
from strategies import STRATEGIES
from tester import StrategyTester

logging.basicConfig(level=logging.INFO, format='%(asctime)s  %(levelname)-7s %(message)s')
logger = logging.getLogger('main')


def cmd_fetch():
    fetcher = DataFetcher()
    config = Config()
    config.setup_dirs()

    data = fetcher.fetch_all()
    for instr, df in data.items():
        if not df.empty:
            df_with_ind = compute_all_indicators(df)
            path = os.path.join(config.DATA_DIR, f'{instr}_H1_indicators.csv')
            df_with_ind.to_csv(path)
            logger.info("Saved %s (%d rows) → %s", instr, len(df_with_ind), path)
        else:
            logger.warning("No data for %s", instr)


def cmd_test(args):
    import pandas as pd
    config = Config()
    config.setup_dirs()
    data = {}

    for instr in config.INSTRUMENTS:
        path = os.path.join(config.DATA_DIR, f'{instr}_H1_indicators.csv')
        if not os.path.exists(path):
            logger.error("File not found: %s. Run --fetch first.", path)
            sys.exit(1)
        df = pd.read_csv(path, index_col=0, parse_dates=True)
        data[instr] = df

    # Determine expiry bars to test
    if args.expiry is not None:
        expiry_list = [args.expiry]
    else:
        expiry_list = config.EXPIRY_BARS

    logger.info("Testing with expiry bars: %s", expiry_list)

    # Run all tests grouped by expiry
    all_results = {}  # expiry -> {key: report}
    tester = StrategyTester()

    for expiry in expiry_list:
        logger.info("=" * 50)
        logger.info("EXPIRY = %d hour(s)", expiry)
        logger.info("=" * 50)
        results = tester.run_all(data, STRATEGIES, expiry_bars_list=[expiry])
        all_results[expiry] = results

        path = os.path.join(config.REPORT_DIR, f'backtest_results_expiry{expiry}.json')
        tester.export_results(results, path)

        profitable = tester.filter_profitable(results)
        logger.info("Expiry %d: Profitable strategies: %d out of %d",
                    expiry, len(profitable), len(results))
        for k in profitable:
            logger.info("  %s", k)

    # Print summary table
    print_summary_table(all_results, expiry_list, config)


def print_summary_table(all_results: dict, expiry_list: list, config: Config):
    """Print a win_rate summary table for all strategy × instrument × expiry combos."""
    if not all_results:
        return

    print(f"\n{'=' * 100}")
    print("  WIN RATE SUMMARY — All Strategies × Instruments × Expirations")
    print(f"{'=' * 100}")

    # Build header
    header = f"  {'Strategy':20s} {'Instr':7s}"
    for exp in expiry_list:
        header += f"  exp{exp}_wr"
    header += "   Best"
    print(header)
    print("  " + "-" * 95)

    for instr in config.INSTRUMENTS:
        for name in config.STRATEGY_PARAMS:
            best_exp = None
            best_wr = -1
            row = f"  {name:20s} {instr:7s}"

            for exp in expiry_list:
                key = f'{instr}_{name}_exp{exp}'
                results = all_results.get(exp, {})
                report = results.get(key)
                if report and report.total_signals > 0:
                    wr = report.win_rate
                    row += f"  {wr:5.1f}%"
                    if wr > best_wr:
                        best_wr = wr
                        best_exp = exp
                else:
                    row += "      -"

            if best_exp is not None:
                row += f"   exp{best_exp} ({best_wr:.1f}%)"
            else:
                row += "   —"
            print(row)

    print(f"{'=' * 100}\n")


def cmd_best(args):
    import json
    config = Config()

    if args.expiry is not None:
        # Show best for a specific expiry
        path = os.path.join(config.REPORT_DIR, f'backtest_results_expiry{args.expiry}.json')
        if not os.path.exists(path):
            logger.error("No results for expiry %d. Run --test --expiry %d first.",
                         args.expiry, args.expiry)
            sys.exit(1)
        _show_best_file(path, config, f"expiry={args.expiry}h")
    else:
        # Show best for all expiries
        found_any = False
        for exp in config.EXPIRY_BARS:
            path = os.path.join(config.REPORT_DIR, f'backtest_results_expiry{exp}.json')
            if os.path.exists(path):
                found_any = True
                _show_best_file(path, config, f"expiry={exp}h")
        if not found_any:
            logger.error("No results found. Run --test first.")
            sys.exit(1)


def _show_best_file(path: str, config: Config, label: str):
    import json
    with open(path) as f:
        data = json.load(f)

    cfg = config.BACKTEST
    print(f"\n{'=' * 90}")
    print(f"  BEST — {label}")
    print(f"  (win_rate ≥ {cfg['min_win_rate']*100:.0f}%, pf ≥ {cfg['min_profit_factor']}, trades ≥ {cfg['min_trades']})")
    print(f"{'=' * 90}")

    profitable = []
    for k, v in data.items():
        if (v['win_rate'] >= cfg['min_win_rate'] * 100
                and v['profit_factor'] >= cfg['min_profit_factor']
                and v['total_signals'] >= cfg['min_trades']):
            profitable.append((k, v))

    if not profitable:
        print("  No strategies passed the thresholds.")
        return

    profitable.sort(key=lambda x: x[1]['profit_factor'], reverse=True)
    for k, v in profitable:
        exp = v.get('expiry_bars', '?')
        print(f"  {k:35s}  exp={exp}  wr={v['win_rate']:5.1f}%  pf={v['profit_factor']:5.2f}  "
              f"dd={v['max_drawdown_pct']:5.1f}%  pnl={v['total_pnl']:+.0f}  N={v['total_signals']}")


def cmd_report(args):
    import json
    config = Config()

    expiry = args.expiry if args.expiry is not None else 1
    path = os.path.join(config.REPORT_DIR, f'backtest_results_expiry{expiry}.json')
    if not os.path.exists(path):
        logger.error("No results for expiry %d. Run --test --expiry %d first.", expiry, expiry)
        sys.exit(1)

    with open(path) as f:
        data = json.load(f)

    strategy_key = args.report
    if strategy_key not in data:
        logger.error("Strategy '%s' not found. Available: %s", strategy_key, list(data.keys()))
        sys.exit(1)

    v = data[strategy_key]
    print(f"\n{'=' * 60}")
    print(f"  {strategy_key}  (expiry={v.get('expiry_bars', '?')})")
    print(f"{'=' * 60}")
    print(f"  Signals:    {v['total_signals']}")
    print(f"  Wins:       {v['wins']}")
    print(f"  Losses:     {v['losses']}")
    print(f"  Win rate:   {v['win_rate']:.1f}%")
    print(f"  Profit factor: {v['profit_factor']:.2f}")
    print(f"  Max DD:     {v['max_drawdown_pct']:.1f}%")
    print(f"  Total PnL:  {v['total_pnl']:+.0f}")
    print(f"  Max cons W: {v['max_consecutive_wins']}")
    print(f"  Max cons L: {v['max_consecutive_losses']}")
    print(f"  Avg conf W: {v['avg_confidence_win']:.1f}")
    print(f"  Avg conf L: {v['avg_confidence_loss']:.1f}")
    if 'equity_snapshot' in v:
        print(f"  Last equity: {v['equity_snapshot']}")


def main():
    parser = argparse.ArgumentParser(description='Binary Options Strategy Backtester')
    parser.add_argument('--fetch', action='store_true', help='Download H1 data from DB and compute indicators')
    parser.add_argument('--test', action='store_true', help='Run backtest on all 5 strategies')
    parser.add_argument('--best', action='store_true', help='Show profitable strategies')
    parser.add_argument('--report', type=str, metavar='KEY', help='Detailed report for strategy (e.g. BITCOIN_ema_rsi_trend)')
    parser.add_argument('--expiry', type=int, default=None, metavar='N',
                        help='Expiry in H1 bars (1-5). Default: test all. For --best/--report: filter by expiry.')
    args = parser.parse_args()

    if args.fetch:
        cmd_fetch()
    elif args.test:
        cmd_test(args)
    elif args.best:
        cmd_best(args)
    elif args.report:
        cmd_report(args)
    else:
        parser.print_help()


if __name__ == '__main__':
    import pandas as pd
    main()
