#!/usr/bin/env python3
"""
Neural network training script for Trading Strategy Tester.

Loads candle data from ALL instrument tables with ALL available timeframes,
normalizes to a single format, and trains both the TrendPredictor and
PatternRecognizer models.

Usage:
    python models/train_ml.py                  # Train with defaults
    python models/train_ml.py --epochs 100    # Custom epochs
    python models/train_ml.py --lr 0.0005     # Custom learning rate
    python models/train_ml.py --trend-only    # Train trend predictor only
    python models/train_ml.py --pattern-only  # Train pattern recognizer only
"""

from __future__ import annotations

import argparse
import json
import logging
import os
import sys
from datetime import datetime
from typing import List, Dict, Optional

# Ensure project root is on the path (required for absolute import style used in this project)
PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_DIR not in sys.path:
    sys.path.insert(0, PROJECT_DIR)

import torch
from torch.utils.data import Dataset, DataLoader

from core.trend_analysis import Candle, TrendDirection, detect_trend
from core.candle_patterns import PatternType, detect_all_patterns
from data.db_connector import DBConnector, DBConfig
from data.preprocessing import normalize_candles, remove_duplicates, fill_missing_candles

logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
logger = logging.getLogger(__name__)


def load_db_config() -> dict:
    """Load database configuration."""
    config_path = os.path.join(PROJECT_DIR, "config", "db_config.json")
    with open(config_path) as f:
        return json.load(f)


TF_SECONDS = {
    "M5": 300,
    "M15": 900,
    "H1": 3600,
    "D1": 86400,
    "W1": 604800,
}


def get_all_instruments(db: DBConnector) -> List[Dict[str, str]]:
    """Get all instruments from the database."""
    return db.instruments.get_all()


def get_timeframes_for_instrument(db: DBConnector, instrument: str) -> List[str]:
    """Get available timeframes for an instrument."""
    return db.instruments.get_timeframes(instrument)


def load_candles_for_training(
    db: DBConnector,
    instrument: str,
    timeframe: str,
    limit: int = 2000,
) -> List[Candle]:
    """
    Load candles for a given instrument/timeframe pair.

    Uses get_last_n to load the most recent candles up to the limit.
    """
    try:
        raw = db.candles.get_last_n(instrument, timeframe, limit)
        if not raw:
            logger.warning(f"No data for {instrument}_{timeframe}")
            return []

        candles = normalize_candles(raw)
        candles = remove_duplicates(candles)
        candles = fill_missing_candles(
            candles,
            expected_interval=TF_SECONDS.get(timeframe, 3600),
            max_gap=5,
        )

        logger.info(
            f"  {instrument}_{timeframe}: {len(candles)} candles loaded"
        )
        return candles

    except Exception as e:
        logger.error(f"Error loading {instrument}_{timeframe}: {e}")
        return []


def generate_trend_training_data(
    candles: List[Candle],
    sequence_length: int = 24,
) -> tuple:
    """
    Generate training sequences and labels for trend prediction.

    Uses a sliding window approach: each window of `sequence_length` candles
    produces one training sample. The label is the trend direction determined
    by the `detect_trend` function on that window.
    """
    sequences = []
    labels = []

    for i in range(len(candles) - sequence_length):
        window = candles[i:i + sequence_length]
        result = detect_trend(window, short_period=8, long_period=21)
        if result is not None:
            sequences.append(window)
            labels.append(result.direction.value)

    return sequences, labels


def generate_pattern_training_data(
    candles: List[Candle],
    lookback: int = 5,
    min_patterns_per_window: int = 30,
) -> tuple:
    """
    Generate training sequences and labels for pattern recognition.

    Uses a sliding window approach. For each window of `lookback` candles,
    we check which of the 10 patterns are present using rule-based detection.
    """
    sequences = []
    pattern_labels = []

    pattern_index = {
        "bullish_engulfing": 0,
        "bearish_engulfing": 1,
        "morning_star": 2,
        "evening_star": 3,
        "hammer": 4,
        "inverted_hammer": 5,
        "shooting_star": 6,
        "doji": 7,
        "three_white_soldiers": 8,
        "three_black_crows": 9,
    }

    pattern_counts = {name: 0 for name in pattern_index}

    all_windows = []
    all_labels_raw = []

    for i in range(len(candles) - lookback):
        window = candles[i:i + lookback]
        signals = detect_all_patterns(window)

        label = [0] * 10
        for sig in signals:
            idx = pattern_index.get(sig.pattern.value)
            if idx is not None:
                label[idx] = 1

        all_windows.append(window)
        all_labels_raw.append(label)

        for name, idx in pattern_index.items():
            if label[idx] == 1:
                pattern_counts[name] += 1

    insufficient = [
        name for name, count in pattern_counts.items()
        if count < min_patterns_per_window
    ]
    if insufficient:
        logger.warning(
            f"Insufficient positive examples for patterns: {insufficient}. "
            f"Minimum {min_patterns_per_window} required per pattern. "
            f"Available: {pattern_counts}"
        )

    positive_windows = []
    positive_labels = []
    negative_windows = []
    negative_labels = []

    for window, label in zip(all_windows, all_labels_raw):
        if any(v == 1 for v in label):
            positive_windows.append(window)
            positive_labels.append(label)
        else:
            negative_windows.append(window)
            negative_labels.append(label)

    # Balance negative examples
    max_negatives = max(len(positive_windows) * 2, 100)
    if len(negative_windows) > max_negatives:
        step = len(negative_windows) // max_negatives
        negative_windows = negative_windows[::step][:max_negatives]
        negative_labels = negative_labels[::step][:max_negatives]

    logger.info(f"  Pattern training data: {len(positive_windows)} positive, "
                f"{len(negative_windows)} negative samples")
    logger.info(f"  Pattern counts: {pattern_counts}")

    sequences = positive_windows + negative_windows
    pattern_labels = positive_labels + negative_labels

    return sequences, pattern_labels


def train_trend_predictor(
    db: DBConnector,
    instruments: List[Dict[str, str]],
    sequence_length: int = 24,
    epochs: int = 50,
    lr: float = 0.001,
    min_candles: int = 100,
) -> Optional[dict]:
    """Train the TrendPredictor model using data from all instruments/timeframes."""
    from models.trend_predictor import TrendPredictor

    logger.info("=" * 60)
    logger.info("TREND PREDICTOR TRAINING")
    logger.info("=" * 60)

    all_candles: List[Candle] = []
    source_info = []

    for inst in instruments:
        name = inst["name"]
        timeframes = get_timeframes_for_instrument(db, name)

        for tf in timeframes:
            candles = load_candles_for_training(db, name, tf)
            if len(candles) >= min_candles:
                all_candles.extend(candles)
                source_info.append(f"  {name}_{tf}: {len(candles)} candles")

    logger.info(f"Total instruments/timeframes loaded: {len(source_info)}")
    for info in source_info:
        logger.info(info)

    if not all_candles:
        logger.error("No candle data available for training!")
        return None

    all_candles.sort(key=lambda c: c.timestamp)
    logger.info(f"Total candles for trend training: {len(all_candles)}")

    sequences, labels = generate_trend_training_data(all_candles, sequence_length)
    logger.info(f"Generated {len(sequences)} trend training sequences")

    if not sequences:
        logger.error("No valid training sequences generated!")
        return None

    label_counts = {}
    for l in labels:
        label_counts[l] = label_counts.get(l, 0) + 1
    logger.info(f"Label distribution: {label_counts}")

    predictor = TrendPredictor(sequence_length=sequence_length)
    if not predictor.is_available():
        logger.error("PyTorch not available, cannot train trend predictor")
        return None

    logger.info(f"Training TrendPredictor for {epochs} epochs (lr={lr})...")
    result = predictor.train(
        training_data=sequences,
        labels=labels,
        epochs=epochs,
        lr=lr,
        save=True,
    )

    logger.info(f"Trend predictor training complete: {result}")
    return result


def train_pattern_recognizer(
    db: DBConnector,
    instruments: List[Dict[str, str]],
    lookback: int = 5,
    epochs: int = 30,
    lr: float = 0.001,
    min_candles: int = 50,
) -> Optional[dict]:
    """Train the PatternRecognizer model using data from all instruments/timeframes."""
    from models.pattern_recognizer import PatternRecognizer

    logger.info("=" * 60)
    logger.info("PATTERN RECOGNIZER TRAINING")
    logger.info("=" * 60)

    all_candles: List[Candle] = []
    source_info = []

    for inst in instruments:
        name = inst["name"]
        timeframes = get_timeframes_for_instrument(db, name)

        for tf in timeframes:
            candles = load_candles_for_training(db, name, tf)
            if len(candles) >= min_candles:
                all_candles.extend(candles)
                source_info.append(f"  {name}_{tf}: {len(candles)} candles")

    logger.info(f"Total instruments/timeframes loaded: {len(source_info)}")
    for info in source_info:
        logger.info(info)

    if not all_candles:
        logger.error("No candle data available for training!")
        return None

    all_candles.sort(key=lambda c: c.timestamp)
    logger.info(f"Total candles for pattern training: {len(all_candles)}")

    sequences, pattern_labels = generate_pattern_training_data(
        all_candles, lookback, min_patterns_per_window=30
    )
    logger.info(f"Generated {len(sequences)} pattern training sequences")

    if not sequences:
        logger.error("No valid training sequences generated!")
        return None

    recognizer = PatternRecognizer(lookback=lookback)
    if not recognizer.is_available():
        logger.error("PyTorch not available, cannot train pattern recognizer")
        return None

    logger.info(f"Training PatternRecognizer for {epochs} epochs (lr={lr})...")
    result = recognizer.train(
        training_data=sequences,
        pattern_labels=pattern_labels,
        epochs=epochs,
        lr=lr,
        save=True,
    )

    logger.info(f"Pattern recognizer training complete: {result}")
    return result


def main():
    parser = argparse.ArgumentParser(
        description="Train ML models on all instrument/timeframe data"
    )
    parser.add_argument(
        "--epochs", type=int, default=50,
        help="Number of training epochs (default: 50)"
    )
    parser.add_argument(
        "--lr", type=float, default=0.001,
        help="Learning rate (default: 0.001)"
    )
    parser.add_argument(
        "--trend-only", action="store_true",
        help="Train only the trend predictor"
    )
    parser.add_argument(
        "--pattern-only", action="store_true",
        help="Train only the pattern recognizer"
    )
    parser.add_argument(
        "--min-candles", type=int, default=100,
        help="Minimum candles per source to include (default: 100)"
    )
    parser.add_argument(
        "--sequence-length", type=int, default=24,
        help="Sequence length for trend predictor (default: 24)"
    )
    parser.add_argument(
        "--lookback", type=int, default=5,
        help="Lookback period for pattern recognizer (default: 5)"
    )
    args = parser.parse_args()

    # Load database config and connect
    db_config = load_db_config()
    db = DBConnector(DBConfig(**db_config))

    if not db.health_check():
        logger.error("Database connection failed!")
        sys.exit(1)

    logger.info("Database connection established")

    # Get all instruments
    instruments = get_all_instruments(db)
    logger.info(f"Found {len(instruments)} instruments: "
                f"{[inst['name'] for inst in instruments]}")

    results = {}

    # Train trend predictor
    if not args.pattern_only:
        result = train_trend_predictor(
            db, instruments,
            sequence_length=args.sequence_length,
            epochs=args.epochs,
            lr=args.lr,
            min_candles=args.min_candles,
        )
        if result:
            results["trend_predictor"] = result

    # Train pattern recognizer
    if not args.trend_only:
        result = train_pattern_recognizer(
            db, instruments,
            lookback=args.lookback,
            epochs=args.epochs,
            lr=args.lr,
            min_candles=args.min_candles,
        )
        if result:
            results["pattern_recognizer"] = result

    db.close()

    # Summary
    logger.info("=" * 60)
    logger.info("TRAINING COMPLETE")
    logger.info("=" * 60)
    for model_name, result in results.items():
        logger.info(f"{model_name}: {result}")

    if not results:
        logger.warning("No models were trained. Check data availability.")
        sys.exit(1)


if __name__ == "__main__":
    main()