#!/usr/bin/env python3
"""CLI script for training neural network trading model.

Usage:
    python scripts/train_neural_model.py --start 2023-01-01 --end 2024-01-01
    python scripts/train_neural_model.py --tickers SBER,GAZP,PLZL --model lstm
    python scripts/train_neural_model.py --epochs 200 --hidden-size 256
"""

from __future__ import annotations

import argparse
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).parent.parent))

import numpy as np
from loguru import logger

from ai.dataset_v2 import ALL_TICKERS, DatasetConfig, MultiTickerDataset
from ai.labeling import BarrierConfig
from ai.trainer_v2 import MultiTaskTrainer, TrainerConfig


def parse_args():
    parser = argparse.ArgumentParser(
        description="Train neural network for entry/SL/TP prediction"
    )
    
    parser.add_argument(
        "--start",
        type=str,
        default="2023-01-01",
        help="Start date for training data (default: 2023-01-01)",
    )
    parser.add_argument(
        "--end",
        type=str,
        default="2024-01-01",
        help="End date for training data (default: 2024-01-01)",
    )
    parser.add_argument(
        "--tickers",
        type=str,
        default=None,
        help="Comma-separated list of tickers (default: all 15 tickers)",
    )
    parser.add_argument(
        "--model",
        type=str,
        default="lstm",
        choices=["lstm", "transformer", "mlp"],
        help="Model architecture (default: lstm — uses proper per-bar sequence processing)",
    )
    parser.add_argument(
        "--epochs",
        type=int,
        default=100,
        help="Number of training epochs (default: 100)",
    )
    parser.add_argument(
        "--batch-size",
        type=int,
        default=128,
        help="Batch size (default: 128)",
    )
    parser.add_argument(
        "--hidden-size",
        type=int,
        default=64,
        help="Hidden layer size (default: 64 — reduced from 128 for regularization)",
    )
    parser.add_argument(
        "--num-layers",
        type=int,
        default=1,
        help="Number of LSTM/Transformer layers (default: 1 — reduced from 2)",
    )
    parser.add_argument(
        "--dropout",
        type=float,
        default=0.4,
        help="Dropout rate (default: 0.4 — increased from 0.2)",
    )
    parser.add_argument(
        "--weight-decay",
        type=float,
        default=1e-4,
        help="Weight decay for Adam optimizer (default: 1e-4 — increased from 1e-5)",
    )
    parser.add_argument(
        "--walk-forward",
        type=int,
        default=0,
        help="Number of walk-forward validation folds (default: 0 = disabled, 3-5 recommended)",
    )
    parser.add_argument(
        "--learning-rate",
        type=float,
        default=0.001,
        help="Learning rate (default: 0.001)",
    )
    parser.add_argument(
        "--context-window",
        type=int,
        default=24,
        help="Context window size in bars (default: 24)",
    )
    parser.add_argument(
        "--tp-atr-mult",
        type=float,
        default=1.5,
        help="Take-profit ATR multiplier (default: 1.5, breakeven WR=40%)",
    )
    parser.add_argument(
        "--sl-atr-mult",
        type=float,
        default=1.0,
        help="Stop-loss ATR multiplier (default: 1.0)",
    )
    parser.add_argument(
        "--max-holding-bars",
        type=int,
        default=48,
        help="Maximum holding period in bars (default: 48)",
    )
    parser.add_argument(
        "--output",
        type=str,
        default="models/neural_trader.pt",
        help="Output model path (default: models/neural_trader.pt)",
    )
    parser.add_argument(
        "--save-dataset",
        type=str,
        default=None,
        help="Save dataset to this path (optional)",
    )
    parser.add_argument(
        "--load-dataset",
        type=str,
        default=None,
        help="Load dataset from this path instead of generating (optional)",
    )
    parser.add_argument(
        "--sides",
        type=str,
        default="LONG,SHORT",
        help="Comma-separated sides to train on (default: LONG,SHORT)",
    )
    parser.add_argument(
        "--early-stopping",
        type=int,
        default=10,
        help="Early stopping patience (default: 10)",
    )
    parser.add_argument(
        "--auto-scale",
        action="store_true",
        default=True,
        help="Auto-tune architecture (hidden_size/layers/dropout/batch) to dataset size (default: True)",
    )
    parser.add_argument(
        "--no-auto-scale",
        action="store_false",
        dest="auto_scale",
        help="Disable auto-scaling, use explicit --hidden-size etc.",
    )
    parser.add_argument(
        "--use-strategy-signals",
        action="store_true",
        default=False,
        help="Augment labels with MACD+RSI strategy signals. Only bars where BOTH "
             "Triple Barrier AND strategy trigger are kept as entries (default: False)",
    )
    
    return parser.parse_args()


def main():
    args = parse_args()
    
    logger.info("=" * 60)
    logger.info("Neural Trading Model Training")
    logger.info("=" * 60)
    
    tickers = args.tickers.split(",") if args.tickers else ALL_TICKERS
    sides = args.sides.split(",")
    
    logger.info(f"Period: {args.start} → {args.end}")
    logger.info(f"Tickers: {len(tickers)} ({', '.join(tickers[:5])}{'...' if len(tickers) > 5 else ''})")
    logger.info(f"Sides: {sides}")
    logger.info(f"Model: {args.model}")
    logger.info(f"Context window: {args.context_window} bars")
    logger.info(f"Barriers: TP={args.tp_atr_mult}×ATR, SL={args.sl_atr_mult}×ATR")
    logger.info("-" * 60)
    
    if args.load_dataset:
        logger.info(f"Loading dataset from {args.load_dataset}")
        dataset = MultiTickerDataset()
        X, y, metadata = dataset.load_dataset(Path(args.load_dataset))
    else:
        barrier_config = BarrierConfig(
            tp_atr_mult=args.tp_atr_mult,
            sl_atr_mult=args.sl_atr_mult,
            max_holding_bars=args.max_holding_bars,
        )
        
        dataset_config = DatasetConfig(
            tickers=tickers,
            start_date=args.start,
            end_date=args.end,
            context_window=args.context_window,
            barrier_config=barrier_config,
            sides=sides,
            use_strategy_signals=args.use_strategy_signals,
        )
        
        logger.info("Loading data from database...")
        dataset = MultiTickerDataset(dataset_config)
        loaded = dataset.load_all_tickers()
        
        if not loaded:
            logger.error("No data loaded. Check database connection.")
            return 1
        
        logger.info(f"Loaded {len(loaded)} tickers")
        
        logger.info("Generating Triple Barrier labels...")
        labels = dataset.generate_all_labels()
        
        if labels.empty:
            logger.error("No labels generated.")
            return 1
        
        logger.info(f"Generated {len(labels)} labels")
        
        logger.info("Building dataset...")
        X, y, metadata = dataset.build_dataset(labels)
        
        if X.size == 0:
            logger.error("Empty dataset.")
            return 1
        
        if args.save_dataset:
            logger.info(f"Saving dataset to {args.save_dataset}")
            dataset.save_dataset(X, y, metadata, Path(args.save_dataset))
    
    logger.info(f"Dataset: {X.shape[0]} samples, {X.shape[1]} features")
    logger.info(f"Entry signal positive rate: {y[:, 0].mean():.2%}")
    logger.info("-" * 60)
    
    trainer_config = TrainerConfig(
        model_type=args.model,
        context_window=args.context_window,
        auto_scale=args.auto_scale,
        hidden_size=args.hidden_size,
        num_layers=args.num_layers,
        dropout=args.dropout,
        learning_rate=args.learning_rate,
        batch_size=args.batch_size,
        epochs=args.epochs,
        early_stopping_patience=args.early_stopping,
        weight_decay=args.weight_decay,
        walk_forward_n_folds=args.walk_forward,
    )
    
    logger.info("Training model...")
    logger.info(f"  Architecture: {args.model}")
    logger.info(f"  Auto-scale: {args.auto_scale}")
    logger.info(f"  Hidden size: {args.hidden_size} (may be overridden by auto-scale)")
    logger.info(f"  Layers: {args.num_layers}")
    logger.info(f"  Dropout: {args.dropout}")
    logger.info(f"  Weight decay: {args.weight_decay}")
    logger.info(f"  Learning rate: {args.learning_rate}")
    logger.info(f"  Batch size: {args.batch_size}")
    logger.info(f"  Max epochs: {args.epochs}")
    if args.walk_forward > 0:
        logger.info(f"  Walk-forward folds: {args.walk_forward}")
    logger.info("-" * 60)
    
    trainer = MultiTaskTrainer(trainer_config)
    if args.walk_forward > 0:
        metrics = trainer.walk_forward_fit(X, y)
    else:
        metrics = trainer.fit(X, y)
    
    logger.info("=" * 60)
    logger.info("Training Results")
    logger.info("=" * 60)
    logger.info(f"Best epoch: {metrics.best_epoch}")
    logger.info(f"Train loss: {metrics.train_loss:.4f}")
    logger.info(f"Val loss: {metrics.val_loss:.4f}")
    logger.info("-" * 60)
    logger.info("Entry Signal Metrics:")
    logger.info(f"  Accuracy:  {metrics.entry_accuracy:.4f}")
    logger.info(f"  Precision: {metrics.entry_precision:.4f}")
    logger.info(f"  Recall:    {metrics.entry_recall:.4f}")
    logger.info(f"  F1 Score:  {metrics.entry_f1:.4f}")
    logger.info(f"  AUC:       {metrics.entry_auc:.4f}")
    logger.info("-" * 60)
    logger.info("SL/TP Regression Metrics:")
    logger.info(f"  SL MSE: {metrics.sl_mse:.4f}")
    logger.info(f"  SL MAE: {metrics.sl_mae:.4f}")
    logger.info(f"  TP MSE: {metrics.tp_mse:.4f}")
    logger.info(f"  TP MAE: {metrics.tp_mae:.4f}")
    logger.info("-" * 60)
    if metrics.walk_forward_results:
        logger.info("Walk-Forward Results:")
        wf_aucs = [r["best_auc"] for r in metrics.walk_forward_results]
        logger.info(f"  Folds: {len(wf_aucs)}, Avg AUC: {np.mean(wf_aucs):.4f}, "
                     f"Min: {min(wf_aucs):.4f}, Max: {max(wf_aucs):.4f}")
        for r in metrics.walk_forward_results:
            logger.info(f"  Fold {r['fold']}: AUC={r['best_auc']:.4f}, "
                        f"train={r['train_samples']}, val={r['val_samples']}")
    logger.info("=" * 60)
    
    output_path = Path(args.output)
    output_path.parent.mkdir(parents=True, exist_ok=True)
    trainer.save(output_path)
    
    logger.info(f"Model saved to: {output_path}")
    logger.info("")
    logger.info("Usage example:")
    logger.info(f"  from ai.inference_v2 import NeuralPredictor")
    logger.info(f"  predictor = NeuralPredictor(model_path=Path('{output_path}'))")
    logger.info(f"  signals = predictor.predict_signals(df, side='LONG')")
    
    return 0


if __name__ == "__main__":
    sys.exit(main())
