#!/usr/bin/env python3
"""Train separate neural network models for each ticker.

Each ticker gets its own LSTM model trained on MACD+RSI strategy-augmented data.
Per-ticker models capture instrument-specific patterns better than a single
global model trained on all tickers together.

Usage:
    python scripts/train_per_ticker.py --start 2021-01-01 --end 2024-01-01
    python scripts/train_per_ticker.py --tickers SBER,GAZP,LKOH
    python scripts/train_per_ticker.py --hidden-size 16 --epochs 50
"""

from __future__ import annotations

import argparse
import sys
import time
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.strategy_signals import StrategyConfig as StrategySignalConfig
from ai.trainer_v2 import MultiTaskTrainer, TrainerConfig

# Strategy tickers (PLZL excluded — worst performer)
STRATEGY_TICKERS = [
    "SBER", "NVTK", "PHOR", "ROSN", "VTBR", "LKOH",
    "GAZP", "ASTR", "MTSS", "SNGSP", "X5",
]

# Default architecture for per-ticker models (small, regularized)
DEFAULT_HIDDEN_SIZE = 16
DEFAULT_NUM_LAYERS = 1
DEFAULT_DROPOUT = 0.3
DEFAULT_BATCH_SIZE = 32
DEFAULT_EPOCHS = 60
DEFAULT_LEARNING_RATE = 0.001
DEFAULT_WEIGHT_DECAY = 1e-4


def parse_args():
    parser = argparse.ArgumentParser(
        description="Train per-ticker neural network models"
    )

    parser.add_argument(
        "--start", type=str, default="2021-01-01",
        help="Start date for training data (default: 2021-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: strategy tickers, no PLZL)",
    )
    parser.add_argument(
        "--hidden-size", type=int, default=DEFAULT_HIDDEN_SIZE,
        help=f"Hidden layer size (default: {DEFAULT_HIDDEN_SIZE} — small for per-ticker)",
    )
    parser.add_argument(
        "--num-layers", type=int, default=DEFAULT_NUM_LAYERS,
        help=f"Number of LSTM layers (default: {DEFAULT_NUM_LAYERS})",
    )
    parser.add_argument(
        "--dropout", type=float, default=DEFAULT_DROPOUT,
        help=f"Dropout rate (default: {DEFAULT_DROPOUT})",
    )
    parser.add_argument(
        "--batch-size", type=int, default=DEFAULT_BATCH_SIZE,
        help=f"Batch size (default: {DEFAULT_BATCH_SIZE})",
    )
    parser.add_argument(
        "--epochs", type=int, default=DEFAULT_EPOCHS,
        help=f"Max epochs per ticker (default: {DEFAULT_EPOCHS})",
    )
    parser.add_argument(
        "--learning-rate", type=float, default=DEFAULT_LEARNING_RATE,
        help=f"Learning rate (default: {DEFAULT_LEARNING_RATE})",
    )
    parser.add_argument(
        "--weight-decay", type=float, default=DEFAULT_WEIGHT_DECAY,
        help=f"Weight decay (default: {DEFAULT_WEIGHT_DECAY})",
    )
    parser.add_argument(
        "--tp-atr-mult", type=float, default=1.5,
        help="Take-profit ATR multiplier (default: 1.5)",
    )
    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-dir", type=str, default="models",
        help="Output directory for per-ticker models (default: models/)",
    )
    parser.add_argument(
        "--force", action="store_true", default=False,
        help="Retrain all tickers even if model exists",
    )
    parser.add_argument(
        "--context-window", type=int, default=24,
        help="Context window size in bars (default: 24)",
    )

    return parser.parse_args()


def train_for_ticker(
    ticker: str,
    args,
    barrier_config: BarrierConfig,
) -> bool:
    """Train a single model for one ticker.

    Args:
        ticker: Ticker symbol
        args: CLI arguments
        barrier_config: Triple Barrier configuration

    Returns:
        True if training succeeded, False otherwise
    """
    model_path = Path(args.output_dir) / f"{ticker}_strategy.pt"

    if model_path.exists() and not args.force:
        logger.info(f"[{ticker}] Model exists, skipping: {model_path.name}")
        return True

    logger.info(f"{'='*60}")
    logger.info(f"[{ticker}] Training per-ticker model")
    logger.info(f"{'='*60}")

    # Build dataset for this single ticker with strategy augmentation
    dataset_config = DatasetConfig(
        tickers=[ticker],
        start_date=args.start,
        end_date=args.end,
        context_window=args.context_window,
        barrier_config=barrier_config,
        sides=["LONG", "SHORT"],
        use_strategy_signals=True,
        strategy_config=StrategySignalConfig(
            tp_atr_mult=args.tp_atr_mult,
            sl_atr_mult=args.sl_atr_mult,
            max_holding_bars=args.max_holding_bars,
        ),
    )

    try:
        dataset = MultiTickerDataset(dataset_config)
        loaded = dataset.load_all_tickers()

        if not loaded:
            logger.warning(f"[{ticker}] No data loaded")
            return False

        labels = dataset.generate_all_labels()

        if labels.empty:
            logger.warning(f"[{ticker}] No labels generated")
            return False

        X, y, metadata = dataset.build_dataset(labels)

        if X.size == 0:
            logger.warning(f"[{ticker}] Empty dataset")
            return False

    except Exception as e:
        logger.error(f"[{ticker}] Dataset generation failed: {e}")
        return False

    n_samples = X.shape[0]
    positive_rate = y[:, 0].mean()
    n_features = X.shape[1]

    logger.info(f"[{ticker}] Dataset: {n_samples} samples, {n_features} features")
    logger.info(f"[{ticker}] Positive rate: {positive_rate:.2%}")

    if n_samples < 20:
        logger.warning(f"[{ticker}] Too few samples ({n_samples}), skipping")
        return False

    # Auto-scale architecture to dataset size
    hidden_size = min(args.hidden_size, max(8, n_samples // 20))
    hidden_size = max(4, hidden_size)
    batch_size = min(args.batch_size, max(8, n_samples // 4))
    epochs = min(args.epochs, max(20, n_samples))

    logger.info(f"[{ticker}] Architecture: hidden={hidden_size}, "
                f"batch={batch_size}, epochs={epochs}")

    trainer_config = TrainerConfig(
        model_type="lstm",
        context_window=args.context_window,
        auto_scale=False,  # Manual scaling above
        hidden_size=hidden_size,
        num_layers=args.num_layers,
        dropout=args.dropout,
        learning_rate=args.learning_rate,
        batch_size=batch_size,
        epochs=epochs,
        early_stopping_patience=min(10, epochs // 3),
        weight_decay=args.weight_decay,
        walk_forward_n_folds=min(3, max(2, n_samples // 50)),
    )

    try:
        trainer = MultiTaskTrainer(trainer_config)

        if trainer_config.walk_forward_n_folds >= 2:
            metrics = trainer.walk_forward_fit(X, y)
            wf_auc = np.mean([r["best_auc"] for r in metrics.walk_forward_results])
        else:
            metrics = trainer.fit(X, y)
            wf_auc = metrics.entry_auc

        logger.info(f"[{ticker}] AUC={wf_auc:.4f}, "
                    f"best_epoch={metrics.best_epoch}, "
                    f"val_loss={metrics.val_loss:.4f}")

        model_path.parent.mkdir(parents=True, exist_ok=True)
        trainer.save(model_path)
        logger.info(f"[{ticker}] Model saved: {model_path.name} ({model_path.stat().st_size} bytes)")

        return True

    except Exception as e:
        logger.error(f"[{ticker}] Training failed: {e}")
        return False


def main():
    args = parse_args()

    tickers = args.tickers.split(",") if args.tickers else STRATEGY_TICKERS

    barrier_config = BarrierConfig(
        tp_atr_mult=args.tp_atr_mult,
        sl_atr_mult=args.sl_atr_mult,
        max_holding_bars=args.max_holding_bars,
    )

    logger.info("=" * 60)
    logger.info("Per-Ticker Neural Model Training")
    logger.info("=" * 60)
    logger.info(f"Period: {args.start} → {args.end}")
    logger.info(f"Tickers: {len(tickers)} ({', '.join(tickers)})")
    logger.info(f"Strategy: MACD+RSI augmented (use_strategy_signals=True)")
    logger.info(f"Architecture: hidden={args.hidden_size}, "
                f"layers={args.num_layers}, dropout={args.dropout}")
    logger.info(f"Output dir: {args.output_dir}/")
    logger.info(f"Force retrain: {args.force}")
    logger.info("-" * 60)

    results = {}
    t_start = time.time()

    for i, ticker in enumerate(tickers, 1):
        logger.info("")
        logger.info(f"[{i}/{len(tickers)}] Processing {ticker}...")
        t_ticker = time.time()

        ok = train_for_ticker(ticker, args, barrier_config)
        elapsed = time.time() - t_ticker
        results[ticker] = {"success": ok, "elapsed": f"{elapsed:.1f}s"}

    logger.info("")
    logger.info("=" * 60)
    logger.info("Training Summary")
    logger.info("=" * 60)

    successes = sum(1 for v in results.values() if v["success"])
    failures = sum(1 for v in results.values() if not v["success"])

    for ticker, res in results.items():
        status = "✅" if res["success"] else "❌"
        logger.info(f"  {status} {ticker}: {res['elapsed']}")

    logger.info("-" * 60)
    logger.info(f"Total: {successes} successful, {failures} failed "
                f"({len(tickers)} tickers, {time.time()-t_start:.0f}s total)")
    logger.info("=" * 60)

    return 0 if failures == 0 else 1


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