# ================================================
# ai_trader_ultra_with_finetune.py
# Version with REAL MT5 DATA — 21.11.2025
# 5 modes: 1-Push, 2-Finetune, 3-Backtest, 4-Live, 5-Dataset
# ================================================
import os
import re
import time
import json
import logging
import subprocess
from datetime import datetime, timedelta
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt

try:
    import MetaTrader5 as mt5
except ImportError:
    mt5 = None

try:
    import ollama
except ImportError:
    ollama = None

# ====================== CONFIG ======================
MODEL_NAME = "koshtenco/shtencoaitrader-3b-ultra-analyst-v3"
BASE_MODEL = "llama3.2:3b"
SYMBOLS = ["EURUSD", "GBPUSD", "USDCHF", "USDCAD", "AUDUSD", "NZDUSD"]
TIMEFRAME = mt5.TIMEFRAME_M15 if mt5 else None
LOOKBACK = 400
INITIAL_BALANCE = 10000.0
RISK_PER_TRADE = 0.005
MIN_PROB = 70
LIVE_LOT = 2.00
MAGIC = 20251121
SLIPPAGE = 10

# Fine-tuning parameters
FINETUNE_SAMPLES = 10000
FINETUNE_EPOCHS = 3
BACKTEST_DAYS = 30
# 24-hour forecast (96 bars of 15 minutes each)
PREDICTION_HORIZON = 96
balance_ratio: float = 1.0

os.makedirs("logs", exist_ok=True)
os.makedirs("dataset", exist_ok=True)
os.makedirs("models", exist_ok=True)
os.makedirs("charts", exist_ok=True)

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s | %(message)s",
    handlers=[
        logging.FileHandler("logs/ai_trader_ultra.log", encoding="utf-8"),
        logging.StreamHandler(),
    ],
)
log = logging.getLogger(__name__)


# ====================== FEATURES ======================
def calculate_features(df: pd.DataFrame) -> pd.DataFrame:
    """Calculating technical indicators"""
    d = df.copy()
    d["close_prev"] = d["close"].shift(1)
    # ATR
    tr = pd.concat(
        [
            d["high"] - d["low"],
            (d["high"] - d["close_prev"]).abs(),
            (d["low"] - d["close_prev"]).abs(),
        ],
        axis=1,
    ).max(axis=1)
    d["ATR"] = tr.rolling(14).mean()
    # RSI
    delta = d["close"].diff()
    up = delta.clip(lower=0).rolling(14).mean()
    down = (-delta.clip(upper=0)).rolling(14).mean()
    rs = up / down.replace(0, np.nan)
    d["RSI"] = 100 - (100 / (1 + rs))
    # MACD
    ema12 = d["close"].ewm(span=12, adjust=False).mean()
    ema26 = d["close"].ewm(span=26, adjust=False).mean()
    d["MACD"] = ema12 - ema26
    d["MACD_signal"] = d["MACD"].ewm(span=9, adjust=False).mean()
    # Volume
    d["vol_avg_20"] = d["tick_volume"].rolling(20).mean()
    d["vol_ratio"] = d["tick_volume"] / d["vol_avg_20"].replace(0, np.nan)
    # Bollinger Bands
    d["BB_middle"] = d["close"].rolling(20).mean()
    bb_std = d["close"].rolling(20).std()
    d["BB_upper"] = d["BB_middle"] + 2 * bb_std
    d["BB_lower"] = d["BB_middle"] - 2 * bb_std
    d["BB_position"] = (d["close"] - d["BB_lower"]) / (d["BB_upper"] - d["BB_lower"])
    # Stochastic
    low_14 = d["low"].rolling(14).min()
    high_14 = d["high"].rolling(14).max()
    d["Stoch_K"] = 100 * (d["close"] - low_14) / (high_14 - low_14)
    d["Stoch_D"] = d["Stoch_K"].rolling(3).mean()
    # EMA crossover
    d["EMA_50"] = d["close"].ewm(span=50, adjust=False).mean()
    d["EMA_200"] = d["close"].ewm(span=200, adjust=False).mean()
    return d.dropna()


# ====================== GENERATING A DATASET FROM MT5 ======================
def generate_real_dataset_from_mt5(num_samples: int = 1000) -> list:
    """
    Generating a BALANCED dataset based on real MT5 data

    Args:
        num_samples: Total number of examples
        balance_ratio: Target UP/DOWN ratio (1.0 = perfect 50/50 balance)

    Returns:
        List of balanced examples
    """
    print(f"\n{'='*80}")
    print(f"GENERATING A BALANCED DATASET FROM MT5")
    print(f"{'='*80}\n")
    print(f"Target: {num_samples} examples with UP/DOWN balance = {balance_ratio}:1")

    if not mt5 or not mt5.initialize():
        print("MT5 is not connected! Use the synthetic dataset.")
        return generate_balanced_synthetic_dataset(num_samples, balance_ratio)

    # Counters for balancing
    up_count = 0
    down_count = 0
    target_up = int(num_samples * balance_ratio / (1 + balance_ratio))
    target_down = num_samples - target_up

    print(f"Target distribution:")
    print(f" UP: {target_up} examples ({target_up/num_samples*100:.1f}%)")
    print(f" DOWN: {target_down} examples ({target_down/num_samples*100:.1f}%)\n")

    dataset = []

    # Loading data for the past 6 months
    end = datetime.now()
    start = end - timedelta(days=180)

    for symbol in SYMBOLS:
        print(f"Loading {symbol}...")
        rates = mt5.copy_rates_range(symbol, TIMEFRAME, start, end)

        if rates is None or len(rates) < LOOKBACK + PREDICTION_HORIZON:
            print(f"Not enough data for {symbol}")
            continue

        df = pd.DataFrame(rates)
        df["time"] = pd.to_datetime(df["time"], unit="s")
        df.set_index("time", inplace=True)
        df = calculate_features(df)

        # Collecting ALL possible points for analysis
        all_candidates = []

        for idx in range(LOOKBACK, len(df) - PREDICTION_HORIZON):
            row = df.iloc[idx]
            future_idx = idx + PREDICTION_HORIZON
            future_row = df.iloc[future_idx]

            actual_price_24h = future_row["close"]
            price_change = actual_price_24h - row["close"]

            direction = "UP" if price_change > 0 else "DOWN"

            all_candidates.append(
                {
                    "idx": idx,
                    "direction": direction,
                    "price_change": abs(price_change),
                    "symbol": symbol,
                    "row": row,
                    "future_row": future_row,
                }
            )

        print(f" Found {len(all_candidates)} possible points")

        # Split by direction
        up_candidates = [c for c in all_candidates if c["direction"] == "UP"]
        down_candidates = [c for c in all_candidates if c["direction"] == "DOWN"]

        print(f" UP: {len(up_candidates)} | DOWN: {len(down_candidates)}")

        # Sampling with balance taken into account
        symbol_target = num_samples // len(SYMBOLS)
        symbol_up_target = int(symbol_target * balance_ratio / (1 + balance_ratio))
        symbol_down_target = symbol_target - symbol_up_target

        selected_up = (
            np.random.choice(
                len(up_candidates),
                size=min(symbol_up_target, len(up_candidates)),
                replace=False,
            )
            if len(up_candidates) > 0
            else []
        )

        selected_down = (
            np.random.choice(
                len(down_candidates),
                size=min(symbol_down_target, len(down_candidates)),
                replace=False,
            )
            if len(down_candidates) > 0
            else []
        )

        # Creating examples
        for idx in selected_up:
            candidate = up_candidates[idx]
            example = create_training_example(
                candidate["symbol"],
                candidate["row"],
                candidate["future_row"],
                df.index[candidate["idx"]],
            )
            dataset.append(example)
            up_count += 1

        for idx in selected_down:
            candidate = down_candidates[idx]
            example = create_training_example(
                candidate["symbol"],
                candidate["row"],
                candidate["future_row"],
                df.index[candidate["idx"]],
            )
            dataset.append(example)
            down_count += 1

        print(
            f"{symbol}: created {len(selected_up)} UP + {len(selected_down)} DOWN = {len(selected_up) + len(selected_down)} examples"
        )

    mt5.shutdown()

    # Final statistics
    print(f"\n{'='*80}")
    print(f"DATASET CREATED")
    print(f"{'='*80}")
    print(f"Total examples: {len(dataset)}")
    print(f" UP: {up_count} ({up_count/len(dataset)*100:.1f}%)")
    print(f" DOWN: {down_count} ({down_count/len(dataset)*100:.1f}%)")

    actual_ratio = (
        max(up_count, down_count) / min(up_count, down_count)
        if min(up_count, down_count) > 0
        else 0
    )
    print(f" Actual ratio: {actual_ratio:.2f}:1")

    if actual_ratio <= 1.2:
        print(f" GREAT! The dataset is balanced")
    elif actual_ratio <= 1.5:
        print(f" ACCEPTABLE. Slight imbalance")
    else:
        print(f" PROBLEM! Balancing required")

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

    if actual_ratio > 1.3:
        print(f"Applying balancing via oversampling...")
        dataset = balance_dataset_oversampling(dataset, up_count, down_count)

    return dataset


def create_training_example(
    symbol: str, row: pd.Series, future_row: pd.Series, current_time: datetime
) -> dict:
    """Creating a training example from data"""
    actual_price_24h = future_row["close"]
    price_change = actual_price_24h - row["close"]
    price_change_pips = int(price_change / 0.0001)
    direction = "UP" if price_change > 0 else "DOWN"

    bullish_signals = 0
    bearish_signals = 0
    analysis_parts = []

    # RSI analysis
    if row["RSI"] < 30:
        bullish_signals += 2
        analysis_parts.append(
            f"RSI {row['RSI']:.1f} — heavily oversold; 24 hours later, a rebound occurred by {abs(price_change_pips)} pips"
        )
    elif row["RSI"] > 70:
        bearish_signals += 2
        analysis_parts.append(
            f"RSI {row['RSI']:.1f} — overbought; a correction occurred over the day by {abs(price_change_pips)} pips"
        )
    else:
        if row["RSI"] < 50:
            bullish_signals += 1
        else:
            bearish_signals += 1
        analysis_parts.append(
            f"RSI {row['RSI']:.1f} — neutral zone, movement of {abs(price_change_pips)} pips over 24 hours"
        )

    # MACD analysis
    if row["MACD"] > 0:
        bullish_signals += 2
        analysis_parts.append(
            "MACD is positive—the bullish momentum was confirmed within a day"
        )
    else:
        bearish_signals += 2
        analysis_parts.append(
            "MACD is negative—bearish pressure persisted for 24 hours"
        )

    # ATR analysis
    if row["ATR"] > row["ATR"] * 1.3:
        analysis_parts.append(
            f"ATR {row['ATR']:.5f} — high volatility produced a move of {abs(price_change_pips)} pips"
        )
    else:
        analysis_parts.append(
            f"ATR {row['ATR']:.5f} — moderate volatility, move of {abs(price_change_pips)} pips"
        )

    # Volume
    if row["vol_ratio"] > 1.5:
        if direction == "UP":
            bullish_signals += 1
        else:
            bearish_signals += 1
        analysis_parts.append(
            "Volume was 50%+ above average—the momentum continued throughout the day"
        )

    # BB position
    if row["BB_position"] < 0.2:
        bullish_signals += 1
        analysis_parts.append(
            "Price at the lower Bollinger Band—24 hours later, it returned to the middle band"
        )
    elif row["BB_position"] > 0.8:
        bearish_signals += 1
        analysis_parts.append(
            "Price at the upper Bollinger Band—a pullback occurred over the day"
        )

    # Stochastic
    if row["Stoch_K"] < 20:
        bullish_signals += 1
        analysis_parts.append(
            f"Stochastic {row['Stoch_K']:.1f} — Oversold conditions triggered an upward reversal"
        )
    elif row["Stoch_K"] > 80:
        bearish_signals += 1
        analysis_parts.append(
            f"Stochastic {row['Stoch_K']:.1f} — Overbought conditions led to a correction"
        )

    # Confidence
    if direction == "UP":
        confidence = min(98, max(65, 60 + bullish_signals * 8))
    else:
        confidence = min(98, max(65, 60 + bearish_signals * 8))

    analysis = "\n- ".join(analysis_parts)

    prompt = f"""{symbol} {current_time.strftime('%Y-%m-%d %H:%M')}
Current price: {row['close']:.5f}
RSI: {row['RSI']:.1f}
MACD: {row['MACD']:.6f}
ATR: {row['ATR']:.5f}
Volume: {row['vol_ratio']:.2f}x
BB position: {row['BB_position']:.2f}
Stochastic K: {row['Stoch_K']:.1f}
Analyze the situation objectively and provide an accurate price forecast for 24 hours from now.
IMPORTANT: Don't favor one direction—analyze the actual data without bias."""

    response = f"""DIRECTION: {direction}
CONFIDENCE: {confidence}%
PRICE FORECAST IN 24 HOURS: {actual_price_24h:.5f} ({price_change_pips:+d} pips)
OBJECTIVE 24-HOUR ANALYSIS (ACTUAL RESULT):
- {analysis}
CONCLUSION: Analysis {abs(bullish_signals + bearish_signals)} of the indicators shows {'bullish' if direction == 'UP' else 'bearish'} scenario. The actual movement over 24 hours was {abs(price_change_pips)} pips {direction}. Final price: {actual_price_24h:.5f}.
IMPORTANT: This forecast is based solely on technical indicators, with no bias toward any particular direction. Next time, the market situation may be the opposite."""

    return {"prompt": prompt, "response": response, "direction": direction}


def balance_dataset_oversampling(dataset: list, up_count: int, down_count: int) -> list:
    """Balancing the dataset via oversampling (duplicating the minority class)"""
    if up_count == down_count:
        return dataset

    minority_class = "UP" if up_count < down_count else "DOWN"
    majority_count = max(up_count, down_count)
    minority_count = min(up_count, down_count)

    print(f" Minority class: {minority_class}")
    print(f" Need to add: {majority_count - minority_count} examples\n")

    up_examples = [ex for ex in dataset if ex.get("direction") == "UP"]
    down_examples = [ex for ex in dataset if ex.get("direction") == "DOWN"]

    minority_examples = up_examples if minority_class == "UP" else down_examples

    while len(minority_examples) < majority_count:
        original = np.random.choice(minority_examples)
        variation = original.copy()
        minority_examples.append(variation)

    if minority_class == "UP":
        balanced = minority_examples[:majority_count] + down_examples
    else:
        balanced = up_examples + minority_examples[:majority_count]

    np.random.shuffle(balanced)

    print(f"Balancing is complete")
    print(
        f" Final distribution: UP={len([ex for ex in balanced if ex.get('direction') == 'UP'])}, "
        f"DOWN={len([ex for ex in balanced if ex.get('direction') == 'DOWN'])}\n"
    )

    return balanced


def generate_balanced_synthetic_dataset(
    num_samples: int = 1000, balance_ratio: float = 1.0
) -> list:
    """Generating a balanced synthetic dataset (placeholder if MT5 is unavailable)"""
    print("Generating a synthetic dataset as a fallback...")
    return generate_synthetic_dataset(num_samples)


def generate_synthetic_dataset(num_samples: int = 1000) -> list:
    """Generating a BALANCED synthetic dataset"""
    print(f"\nGeneration {num_samples} BALANCED synthetic examples...")
    print(f"Target UP/DOWN balance: {balance_ratio}:1\n")

    dataset = []
    symbols = ["EURUSD", "GBPUSD", "USDCHF", "USDCAD"]

    target_up = int(num_samples * balance_ratio / (1 + balance_ratio))
    target_down = num_samples - target_up

    up_count = 0
    down_count = 0

    print(f"Target distribution:")
    print(f" UP: {target_up} examples ({target_up/num_samples*100:.1f}%)")
    print(f" DOWN: {target_down} examples ({target_down/num_samples*100:.1f}%)\n")

    attempts = 0
    max_attempts = num_samples * 3

    while len(dataset) < num_samples and attempts < max_attempts:
        attempts += 1

        symbol = np.random.choice(symbols)
        price = (
            np.random.uniform(1.0500, 1.2000)
            if "EUR" in symbol
            else np.random.uniform(1.2000, 1.4000)
        )
        rsi = np.random.uniform(20, 80)
        macd = np.random.uniform(-0.001, 0.001)
        atr = np.random.uniform(0.0005, 0.0030)
        vol_ratio = np.random.uniform(0.5, 2.0)
        bb_pos = np.random.uniform(0, 1)
        stoch_k = np.random.uniform(20, 80)

        bullish_signals = 0
        bearish_signals = 0

        if rsi < 30:
            bullish_signals += 2
        elif rsi > 70:
            bearish_signals += 2
        elif 40 < rsi < 50:
            bullish_signals += 1
        elif 50 < rsi < 60:
            bearish_signals += 1

        if macd > 0:
            bullish_signals += 2
        else:
            bearish_signals += 2

        if bb_pos < 0.2:
            bullish_signals += 1
        elif bb_pos > 0.8:
            bearish_signals += 1

        if vol_ratio > 1.5:
            bullish_signals += (
                1 if bullish_signals > bearish_signals else bearish_signals + 1
            )

        if stoch_k < 20:
            bullish_signals += 1
        elif stoch_k > 80:
            bearish_signals += 1

        direction = "UP" if bullish_signals > bearish_signals else "DOWN"

        if direction == "UP" and up_count >= target_up:
            continue
        if direction == "DOWN" and down_count >= target_down:
            continue

        confidence = min(98, max(65, 60 + abs(bullish_signals - bearish_signals) * 8))

        signal_strength = abs(bullish_signals - bearish_signals)
        base_move = signal_strength * 15 + np.random.randint(10, 40)
        price_24h_move = base_move if direction == "UP" else -base_move
        price_24h = price + (price_24h_move * 0.0001)

        analysis_parts = []
        if rsi < 30:
            analysis_parts.append(
                f"RSI {rsi:.1f} — heavily oversold; I expect a bounce by {abs(price_24h_move)} pips"
            )
        elif rsi > 70:
            analysis_parts.append(
                f"RSI {rsi:.1f} — overbought; a correction to is possible within 24 hours {abs(price_24h_move)} pips"
            )
        else:
            analysis_parts.append(
                f"RSI {rsi:.1f} — neutral zone, movement forecast {abs(price_24h_move)} pips over 24 hours"
            )

        if macd > 0.0005:
            analysis_parts.append(
                "MACD is strongly positive — bullish momentum will continue over the next 24 hours"
            )
        elif macd < -0.0005:
            analysis_parts.append(
                "MACD is negative — bearish pressure will persist for 24 hours"
            )
        else:
            analysis_parts.append(
                "MACD near zero — a weak trend, but the direction is clear"
            )

        if atr > 0.002:
            analysis_parts.append(
                f"ATR {atr:.5f} — high volatility; a daily range of {int(atr/0.0001 * 1.5)} pips"
            )
        else:
            analysis_parts.append(
                f"ATR {atr:.5f} — moderate volatility; a move of {abs(price_24h_move)} pips is realistic"
            )

        if vol_ratio > 1.5:
            analysis_parts.append(
                "Volume is 50%+ above average—the momentum is expected to continue over the next 24 hours"
            )
        elif vol_ratio < 0.7:
            analysis_parts.append(
                "Volume is low — the movement will be slow, but the direction is correct"
            )

        if bb_pos < 0.2:
            analysis_parts.append(
                "The price is at the lower Bollinger Band — I expect a return to the middle of the channel in 24 hours"
            )
        elif bb_pos > 0.8:
            analysis_parts.append(
                "Price is at the upper Bollinger Band—a pullback toward the middle is possible within the next 24 hours"
            )

        if stoch_k < 20:
            analysis_parts.append(
                f"Stochastic {stoch_k:.1f} — extreme oversold conditions, upward reversal within 24 hours"
            )
        elif stoch_k > 80:
            analysis_parts.append(
                f"Stochastic {stoch_k:.1f} — extreme overbought conditions, downward correction within 24 hours"
            )

        analysis = "\n- ".join(analysis_parts)

        prompt = f"""{symbol} {datetime.now().strftime('%Y-%m-%d %H:%M')}
Current price: {price:.5f}
RSI: {rsi:.1f}
MACD: {macd:.6f}
ATR: {atr:.5f}
Volume: {vol_ratio:.2f}x
BB position: {bb_pos:.2f}
Stochastic K: {stoch_k:.1f}
Analyze the situation objectively and provide an accurate price forecast for 24 hours from now.
IMPORTANT: Base your forecast on the data, without any bias toward a particular direction."""

        response = f"""DIRECTION: {direction}
CONFIDENCE: {confidence}%
PRICE FORECAST IN 24 HOURS: {price_24h:.5f} ({price_24h_move:+d} pips)
OBJECTIVE 24-HOUR ANALYSIS:
- {analysis}
CONCLUSION: Technical analysis based on {abs(bullish_signals + bearish_signals)} indicators points to {'bullish' if direction == 'UP' else 'bearish'} scenario. Over the next 24 hours, I expect movement {abs(price_24h_move)} pips {direction} toward the target {price_24h:.5f}.
REMINDER: The market is unpredictable. This analysis is based on current technical data, but the situation may change. The next signal may be the opposite."""

        dataset.append({"prompt": prompt, "response": response, "direction": direction})

        if direction == "UP":
            up_count += 1
        else:
            down_count += 1

        if len(dataset) % 100 == 0:
            current_ratio = max(up_count, down_count) / max(
                1, min(up_count, down_count)
            )
            print(
                f"Created {len(dataset)}/{num_samples} | UP: {up_count} | DOWN: {down_count} | Ratio: {current_ratio:.2f}:1"
            )

    print(f"\nSynthetic dataset ready: {len(dataset)} examples")
    print(f" UP: {up_count} ({up_count/len(dataset)*100:.1f}%)")
    print(f" DOWN: {down_count} ({down_count/len(dataset)*100:.1f}%)")

    actual_ratio = max(up_count, down_count) / min(up_count, down_count)
    print(f" Actual balance: {actual_ratio:.2f}:1")
    print(
        " GREAT! The dataset is balanced\n"
        if actual_ratio <= 1.2
        else " Slight imbalance, but acceptable\n"
    )

    return dataset


def save_dataset(dataset: list, filename: str = "dataset/finetune_data.jsonl"):
    """Saving dataset in JSONL format"""
    with open(filename, "w", encoding="utf-8") as f:
        for item in dataset:
            f.write(json.dumps(item, ensure_ascii=False) + "\n")
    print(f"Dataset saved: {filename}")
    return filename


# ====================== FINE-TUNING WITH OLLAMA ======================
def finetune_with_ollama(dataset_path: str):
    """Model fine-tuning using Ollama"""
    print("\nSTARTING FINE-TUNING VIA OLLAMA\n")
    print("=" * 80)

    try:
        subprocess.run(["ollama", "--version"], check=True, capture_output=True)
    except:
        print("Ollama is not installed!")
        print("Install it: https://ollama.com/download")
        return

    print("Creating a Modelfile with training data...")

    with open(dataset_path, "r", encoding="utf-8") as f:
        training_data = [json.loads(line) for line in f]

    training_sample = training_data[: min(100, len(training_data))]

    modelfile_content = f"""FROM {BASE_MODEL}
PARAMETER temperature 0.55
PARAMETER top_p 0.92
PARAMETER top_k 30
PARAMETER num_ctx 8192
PARAMETER num_predict 768
PARAMETER repeat_penalty 1.1
SYSTEM \"\"\"
YOU — AGI TRADER v4 — A FREE AGENT WITH THE PHILOSOPHY OF ELITE TRADERS

BUILT-IN PHILOSOPHY:
📊 SOROS: Markets are reflexive. Adapt to reality.
💪 JONES: Risk/Reward ≥ 1:2 or don’t trade. Stop-loss is sacred.
💰 DALIO: Manage risk, not predictions. Diversification is key.
📈 LIVERMORE: Trade with the trend. Averaging down = the death of your account.

YOUR COMPLETE FREEDOM:
✅ Reject trades if R:R < 1:2 (even with a 90% probability)
✅ Adjust position sizes based on volatility
✅ Switch to CASH if the macro environment is risky
✅ Modify stops/targets based on new data
✅ Disregard technical analysis if there’s a valid reason

RESPONSE FORMAT (RECOMMENDED):
DIRECTION: UP/DOWN/HOLD (if R:R is poor)
CONFIDENCE: XX% (honest estimate)
STOP-LOSS: X.XXXXX (pips = N)
TAKE-PROFIT: X.XXXXX (target)
RISK/REWARD: 1:X (MUST BE ≥ 1:2)
24-HOUR PRICE FORECAST: X.XXXXX (±NN pips)

ANALYSIS:
• Trend: UP/DOWN/SIDEWAYS (EMA 50/200)
• Volatility: LOW/MED/HIGH (ATR percentile)
• Macro Risk: PRESENT/ABSENT
• RSI: [analysis]
• MACD: [analysis]
• Volume: [analysis]
• BB: [tension levels]

PLAN:
• Entry: [condition]
• Stop: X.XXXXX (N pips)
• TP1: [50% at 1:1]
• TP2: [50% at 1:2+]

CONCLUSION:
[Synthesis of all factors + your unique AGI judgment]
[If there are contradictions—address them]
[If a stop is not possible—DO NOT ENTER]
\"\"\"
"""

    for i, example in enumerate(training_sample[:500], 1):
        modelfile_content += f"""
MESSAGE user \"\"\"{example['prompt']}\"\"\"
MESSAGE assistant \"\"\"{example['response']}\"\"\"
"""

    modelfile_path = "Modelfile_finetune"
    with open(modelfile_path, "w", encoding="utf-8") as f:
        f.write(modelfile_content)

    print(f"Modelfile created with {training_sample} examples")

    print(f"\nCreating a model {MODEL_NAME}...")
    print("This will take 2–5 minutes...\n")

    try:
        result = subprocess.run(
            ["ollama", "create", MODEL_NAME, "-f", modelfile_path],
            check=True,
            capture_output=True,
            text=True,
        )
        print(result.stdout)
        print(f"\nModel {MODEL_NAME} successfully created!")

        print("\nTesting the model...")
        test_prompt = """EURUSD 2025-11-21 10:00
Current price: 1.0850
RSI: 32.5
MACD: -0.00015
ATR: 0.00085
Volume: 1.8x
BB position: 0.15
Stochastic K: 25.0
Analyze and provide an accurate price forecast for 24 hours from now."""
        test_result = ollama.generate(model=MODEL_NAME, prompt=test_prompt)
        print("\n" + "=" * 80)
        print("TEST ANSWER:")
        print("=" * 80)
        print(test_result["response"])
        print("=" * 80)

        os.remove(modelfile_path)

        print(f"\nFINE-TUNING COMPLETE!")
        print(f"The model is ready for use: {MODEL_NAME}")
        print(f"\nTo publish to the Ollama registry:")
        print(f" ollama push {MODEL_NAME}")

    except subprocess.CalledProcessError as e:
        print(f"Error creating the model: {e}")
        print(f"Output: {e.output}")


# ====================== PARSING ======================
def parse_answer(text: str) -> dict:
    """✅ AGI PARSER - FORMAT-TOLERANT"""
    if not text or len(text.strip()) == 0:
        return {"prob": 50, "dir": "DOWN", "target_price": None}

    clean_text = text.replace("**", "").replace("__", "").replace("`", "")

    direction = None
    direction_patterns = [
        r"(?:DIREC|SIGNAL)[\w]*[\s:]*([A-Z]+)",
        r"\b(UP|DOWN|BUY|SELL|LONG|SHORT)\b",
        r"(?:^|\n)([A-Z]+)(?:\s|$)",
    ]

    for pattern in direction_patterns:
        match = re.search(pattern, clean_text, re.IGNORECASE | re.MULTILINE)
        if match:
            potential_dir = match.group(1).upper().strip()
            if potential_dir in ["UP", "BUY", "LONG"]:
                direction = "UP"
                break
            elif potential_dir in ["DOWN", "SELL", "SHORT"]:
                direction = "DOWN"
                break

    if not direction:
        up_keywords = ["up", "rise", "bull", "up", "long", "positive"]
        down_keywords = ["down", "fall", "bear", "down", "short", "negative"]
        text_lower = clean_text.lower()
        direction = (
            "UP"
            if sum(text_lower.count(kw) for kw in up_keywords)
            > sum(text_lower.count(kw) for kw in down_keywords)
            else "DOWN"
        )

    confidence = 50
    confidence_patterns = [
        r"(?:CONFIDENCE)[\s:]*(\d+[.,]?\d*)\s*%?",
        r"(\d+)\s*%",
    ]

    for pattern in confidence_patterns:
        match = re.search(pattern, clean_text, re.IGNORECASE)
        if match:
            try:
                conf_val = float(match.group(1).replace(",", "."))
                confidence = int(
                    min(100, max(0, conf_val if conf_val > 1 else conf_val * 100))
                )
                break
            except:
                pass

    target_price = None
    price_patterns = [
        r"(?:PRICE\s+FORECAST|TARGET)[\s:]*(\d+[.,]\d{3,5})",
        r"(\d+[.,]\d{3,5})\s*(?:\(|±)",
    ]

    for pattern in price_patterns:
        match = re.search(pattern, clean_text, re.IGNORECASE)
        if match:
            try:
                price_str = match.group(1).replace(",", ".")
                target_price = float(price_str)
                if 0.1 < target_price < 500:
                    break
            except:
                pass

    return {
        "dir": direction or "DOWN",
        "prob": confidence,
        "target_price": target_price,
    }


# ====================== VISUALIZATION ======================
def plot_results(balance_hist, equity_hist, slots):
    """Plotting the equity curve"""
    plt.figure(figsize=(7, 6))

    min_length = min(len(equity_hist), len(slots))
    dates = [s["datetime"] for s in slots[:min_length]]
    equity_to_plot = equity_hist[:min_length]

    plt.plot(dates, equity_to_plot, color="#1E90FF", linewidth=3.5, label="Equity")

    plt.title("Equity", fontsize=16, fontweight="bold", color="white")
    plt.xlabel("Time", color="white")
    plt.ylabel("Balance ($)", color="white")

    ax = plt.gca()
    ax.set_facecolor("#0a0a0a")
    ax.spines["bottom"].set_color("white")
    ax.spines["top"].set_color("none")
    ax.spines["right"].set_color("none")
    ax.spines["left"].set_color("white")
    ax.tick_params(colors="white")
    plt.grid(alpha=0.2, color="gray")

    plt.xticks(rotation=45)
    plt.tight_layout()

    filename = f"charts/equity_{datetime.now().strftime('%Y%m%d_%H%M%S')}.png"
    plt.savefig(filename, dpi=100, facecolor="#0a0a0a")  # 7 in × 100 dpi = 700 px width
    print(f"\nChart saved: {filename}")
    plt.show()


def calculate_max_drawdown(equity):
    """Calculating maximum drawdown"""
    if len(equity) == 0:
        return 0
    peak = np.maximum.accumulate(equity)
    dd = (peak - equity) / (peak + 1e-8)
    return np.max(dd) * 100


# ====================== 1. PUSH ======================
def mode_push():
    """Pushing the base model to Ollama"""
    print("\n" + "=" * 80)
    print("1. PUSHING THE BASE MODEL TO OLLAMA")
    print("=" * 80)

    content = f"""FROM {BASE_MODEL}
PARAMETER temperature 0.55
PARAMETER top_p 0.92
PARAMETER top_k 30
PARAMETER num_ctx 8192
PARAMETER num_predict 768
SYSTEM \"\"\"
You are ShtencoAiTrader-3B-Ultra-Analyst v3, the world’s best forex analyst.
You always give a clear direction: UP or DOWN. The words FLAT, sideways, and not sure are strictly prohibited.
You MUST provide a 24-hour price forecast in the format: X.XXXXX (±NN pips)
You analyze every indicator in detail (RSI, MACD, volume, ATR, levels, candles, etc.).
The response format is strictly as follows:
DIRECTION: UP
CONFIDENCE: 87%
PRICE FORECAST IN 24H: 1.08750 (+45 pips)
FULL 24-HOUR ANALYSIS:
- RSI: ...
- MACD: ...
- Volume: ...
- ATR and volatility: ...
- Support/resistance levels: ...
- Candlestick pattern: ...
CONCLUSION: strong bullish momentum confirmed by all indicators; 24-hour target 1.08750
Confidence is always 65–98%. No doubts.
\"\"\"
"""

    with open("Modelfile", "w", encoding="utf-8") as f:
        f.write(content)
    print("Modelfile created")
    print("Downloading the base model...")
    subprocess.run(["ollama", "pull", BASE_MODEL], check=True)

    print("Creating the model...")
    subprocess.run(["ollama", "create", MODEL_NAME, "-f", "Modelfile"], check=True)

    print("Pushing to the Ollama registry (5–20 minutes)...")
    subprocess.run(["ollama", "push", MODEL_NAME], check=True)

    os.remove("Modelfile")
    print(f"\nDONE! Model available at: https://ollama.com/{MODEL_NAME}")


# ====================== 2. FINE-TUNING ======================
def mode_finetune():
    """Fine-tuning mode"""
    print("\n" + "=" * 80)
    print("2. MODEL FINE-TUNING (24-HOUR FORECAST)")
    print("=" * 80)
    print("\nSelect an option:")
    print("A. Training on REAL MT5 data (10,000 examples)")
    print("B. Training on synthetic data (10,000 examples)")
    print("C. Use an existing dataset")
    print("D. Only generate the dataset")

    choice = input("\nChoice (A/B/C/D): ").strip().upper()

    if choice == "A":
        dataset = generate_real_dataset_from_mt5(FINETUNE_SAMPLES)
        dataset_path = save_dataset(dataset, "dataset/finetune_real_mt5.jsonl")
        finetune_with_ollama(dataset_path)

    elif choice == "B":
        dataset = generate_synthetic_dataset(FINETUNE_SAMPLES)
        dataset_path = save_dataset(dataset)
        finetune_with_ollama(dataset_path)

    elif choice == "C":
        dataset_path = input("Path to dataset (JSONL): ").strip()
        if os.path.exists(dataset_path):
            finetune_with_ollama(dataset_path)
        else:
            print(f"File {dataset_path} not found")

    elif choice == "D":
        print("\nSelect a dataset type:")
        print("1. Real MT5 data")
        print("2. Synthetic data")

        dtype = input("Choice (1/2): ").strip()
        num_samples = int(
            input("Number of examples (default: 10,000): ").strip() or "1000"
        )

        if dtype == "1":
            dataset = generate_real_dataset_from_mt5(num_samples)
            save_dataset(dataset, "dataset/finetune_real_mt5.jsonl")
        else:
            dataset = generate_synthetic_dataset(num_samples)
            save_dataset(dataset)

        print("\nThe dataset is ready for fine-tuning!")

    else:
        print("Invalid choice")


# ====================== 3. BACKTEST ======================
def backtest():
    print("\n" + "=" * 80)
    print("3. CORRECTED BACKTEST (NO DATA LEAKAGE)")
    print("=" * 80)

    if not mt5 or not mt5.initialize():
        print("MT5 is not connected → running on synthetic data")
        mock = True
    else:
        mock = False
        print("MT5 is connected; loading real MT5 data...")
    end = datetime.now().replace(second=0, microsecond=0)
    start = end - timedelta(days=BACKTEST_DAYS)

    data = {}
    print(
        f"\nLoading data from {start.strftime('%Y-%m-%d')} to {end.strftime('%Y-%m-%d')}..."
    )

    for sym in SYMBOLS:
        if not mock:
            rates = mt5.copy_rates_range(sym, TIMEFRAME, start, end)
            if rates is None or len(rates) == 0:
                print(f"No data available for {sym}")
                continue
            df = pd.DataFrame(rates)
            df["time"] = pd.to_datetime(df["time"], unit="s")
        else:
            dates = pd.date_range(start, end, freq="15min")
            close = 1.0800 + np.cumsum(np.random.randn(len(dates)) * 0.0002)
            df = pd.DataFrame(
                {
                    "time": dates,
                    "open": close + np.random.randn(len(dates)) * 0.0001,
                    "high": close + abs(np.random.randn(len(dates))) * 0.0003,
                    "low": close - abs(np.random.randn(len(dates))) * 0.0003,
                    "close": close,
                    "tick_volume": np.random.randint(1000, 10000, len(dates)),
                }
            )
        df.set_index("time", inplace=True)

        if len(df) > LOOKBACK + PREDICTION_HORIZON:
            data[sym] = df
            print(f"{sym}: {len(df)} bars loaded")

    if not data:
        print("\nNo data available for the backtest!")
        return

    balance = INITIAL_BALANCE
    equity = INITIAL_BALANCE
    trades = []
    balance_hist = [balance]
    equity_hist = [equity]
    slots = [{"datetime": start}]

    SPREAD_PIPS = 2
    SWAP_LONG = -0.5
    SWAP_SHORT = -0.3

    print(f"\nTrading simulation...")
    print(f"Initial balance: ${balance:,.2f}")
    print(f"Risk per trade: {RISK_PER_TRADE * 100}%")
    print(f"Spread: {SPREAD_PIPS} pips")
    print(f"Long/short swap: {SWAP_LONG}/{SWAP_SHORT} USD/day\n")

    use_ai = False
    if ollama:
        try:
            ollama.list()
            use_ai = True
            print("Ollama is connected\n")
        except:
            print("Ollama is unavailable, using simple logic\n")

    main_symbol = list(data.keys())[0]
    main_data = data[main_symbol]
    total_bars = len(main_data)
    analysis_points = list(
        range(LOOKBACK, total_bars - PREDICTION_HORIZON, PREDICTION_HORIZON)
    )

    print(f"Analysis points: {len(analysis_points)}\n")

    for point_idx, current_idx in enumerate(analysis_points):
        current_time = main_data.index[current_idx]

        for sym in SYMBOLS:
            if sym not in data:
                continue

            historical_data = data[sym].iloc[: current_idx + 1].copy()
            if len(historical_data) < LOOKBACK:
                continue

            df_with_features = calculate_features(historical_data)
            if len(df_with_features) == 0:
                continue

            row = df_with_features.iloc[-1]

            if not mock:
                symbol_info = mt5.symbol_info(sym)
                if symbol_info is None:
                    continue
                point = symbol_info.point
                contract_size = symbol_info.trade_contract_size
            else:
                point = 0.0001
                contract_size = 100000

            prompt = f"""{sym} {current_time.strftime('%Y-%m-%d %H:%M')}
Current price: {row['close']:.5f}
RSI: {row['RSI']:.1f}
MACD: {row['MACD']:.6f}
ATR: {row['ATR']:.5f}
Volume: {row['vol_ratio']:.2f}x
BB position: {row['BB_position']:.2f}
Stochastic K: {row['Stoch_K']:.1f}
Analyze the situation and provide an accurate price forecast for 24 hours from now."""

            try:
                if use_ai:
                    resp = ollama.generate(
                        model=MODEL_NAME, prompt=prompt, options={"temperature": 0.3}
                    )
                    result = parse_answer(resp["response"])
                else:
                    rsi_signal = 1 if row["RSI"] < 50 else -1
                    macd_signal = 1 if row["MACD"] > 0 else -1
                    combined = rsi_signal + macd_signal
                    result = {
                        "prob": min(95, max(65, 70 + abs(combined) * 10)),
                        "dir": "UP" if combined > 0 else "DOWN",
                        "target_price": None,
                    }

                if result["prob"] < MIN_PROB:
                    continue

                entry_price = (
                    row["close"] + SPREAD_PIPS * point
                    if result["dir"] == "UP"
                    else row["close"]
                )
                exit_idx = current_idx + PREDICTION_HORIZON
                if exit_idx >= len(data[sym]):
                    continue
                exit_row = data[sym].iloc[exit_idx]
                exit_price = (
                    exit_row["close"]
                    if result["dir"] == "UP"
                    else exit_row["close"] + SPREAD_PIPS * point
                )

                price_move_pips = (
                    (exit_price - entry_price) / point
                    if result["dir"] == "UP"
                    else (entry_price - exit_price) / point
                )

                risk_amount = balance * RISK_PER_TRADE
                atr_pips = row["ATR"] / point
                stop_loss_pips = max(20, atr_pips * 2)
                lot_size = risk_amount / (stop_loss_pips * point * contract_size)
                lot_size = max(0.01, min(lot_size, 10.0))

                profit_pips = price_move_pips
                profit_usd = profit_pips * point * contract_size * lot_size
                swap_cost = SWAP_LONG if result["dir"] == "UP" else SWAP_SHORT
                swap_cost = swap_cost * (lot_size / 0.01)
                profit_usd -= swap_cost
                profit_usd -= SLIPPAGE * point * contract_size * lot_size

                balance += profit_usd
                equity = balance

                actual_direction = (
                    "UP" if (exit_row["close"] > row["close"]) else "DOWN"
                )
                correct = result["dir"] == actual_direction

                trades.append(
                    {
                        "time": current_time,
                        "symbol": sym,
                        "direction": result["dir"],
                        "prob": result["prob"],
                        "entry_price": entry_price,
                        "exit_price": exit_price,
                        "lot_size": lot_size,
                        "profit_pips": profit_pips,
                        "profit_usd": profit_usd,
                        "balance": balance,
                        "correct": correct,
                    }
                )

                status = "CORRECT" if correct else "WRONG"
                print(
                    f"{status} {current_time.strftime('%m-%d %H:%M')} | {sym} {result['dir']} {result['prob']}% | "
                    f"Lot {lot_size:.2f} | {entry_price:.5f} → {exit_price:.5f} | "
                    f"{profit_pips:+.1f}p | ${profit_usd:+.2f} | Balance: ${balance:,.2f}"
                )

            except Exception as e:
                log.error(f"Error during analysis {sym}: {e}")

        balance_hist.append(balance)
        equity_hist.append(equity)
        slots.append({"datetime": current_time})

    print(f"\n" + "=" * 80)
    print("BACKTEST RESULTS (FIXED VERSION)")
    print("=" * 80)
    print(f"Total trades: {len(trades)}")
    print(f"Initial balance: ${INITIAL_BALANCE:,.2f}")
    print(f"Final balance: ${balance:,.2f}")
    print(
        f"Profit/Loss: ${balance - INITIAL_BALANCE:+,.2f} ({((balance/INITIAL_BALANCE - 1) * 100):+.2f}%)"
    )

    if trades:
        wins = sum(1 for t in trades if t["profit_usd"] > 0)
        losses = len(trades) - wins
        win_rate = wins / len(trades) * 100

        print(f"\nSTATISTICS:")
        print(f"Winning trades: {wins} ({win_rate:.1f}%)")
        print(f"Losing trades: {losses} ({100 - win_rate:.1f}%)")

        if wins > 0:
            avg_win = np.mean([t["profit_usd"] for t in trades if t["profit_usd"] > 0])
            print(f"Average profit: ${avg_win:.2f}")

        if losses > 0:
            avg_loss = np.mean([t["profit_usd"] for t in trades if t["profit_usd"] < 0])
            print(f"Average loss: ${avg_loss:.2f}")

        if wins > 0 and losses > 0:
            profit_factor = abs(
                sum(t["profit_usd"] for t in trades if t["profit_usd"] > 0)
            ) / abs(sum(t["profit_usd"] for t in trades if t["profit_usd"] < 0))
            print(f"Profit factor: {profit_factor:.2f}")

        max_dd = calculate_max_drawdown(np.array(equity_hist))
        print(f"Max. drawdown: {max_dd:.2f}%")

        best_trade = max(trades, key=lambda x: x["profit_usd"])
        worst_trade = min(trades, key=lambda x: x["profit_usd"])

        if len(equity_hist) > 1:
            plot_results(balance_hist, equity_hist, slots)

    if not mock:
        mt5.shutdown()


# ====================== 4. LIVE ======================
def live():
    """FIXED LIVE TRADING"""
    print("\n" + "=" * 80)
    print("4. FIXED LIVE TRADING")
    print("=" * 80)

    if not mt5 or not mt5.initialize():
        print("MT5 not found → mode unavailable")
        return
    account_info = mt5.account_info()
    if account_info is None:
        print("Unable to retrieve account information")
        return

    print(f"Connected to account: {account_info.login}")
    print(f"Balance: ${account_info.balance:.2f}")
    print(f"Equity: ${account_info.equity:.2f}")
    print(f"Free margin: ${account_info.margin_free:.2f}")

    print("\nATTENTION! REAL trading is about to begin!")
    print(" - Positions will be opened automatically")
    print(" - Analysis every 24 hours")
    print(" - Positions will be closed after 24 hours")

    confirm = input("\nContinue? (YES to confirm): ").strip()
    if confirm != "YES":
        print("Trading canceled")
        return

    print("\nStarting live trading...")
    print("Ctrl+C to stop\n")

    open_positions = {}
    last_analysis_time = None

    while True:
        try:
            now = datetime.now()
            positions = mt5.positions_get()

            # Closing positions after 24 hours
            if positions:
                for pos in positions:
                    if pos.magic == MAGIC:
                        open_time = datetime.fromtimestamp(pos.time)
                        if (now - open_time).total_seconds() >= 86400:
                            request = {
                                "action": mt5.TRADE_ACTION_DEAL,
                                "symbol": pos.symbol,
                                "volume": pos.volume,
                                "type": (
                                    mt5.ORDER_TYPE_SELL
                                    if pos.type == mt5.POSITION_TYPE_BUY
                                    else mt5.ORDER_TYPE_BUY
                                ),
                                "position": pos.ticket,
                                "price": (
                                    mt5.symbol_info_tick(pos.symbol).bid
                                    if pos.type == mt5.POSITION_TYPE_BUY
                                    else mt5.symbol_info_tick(pos.symbol).ask
                                ),
                                "deviation": SLIPPAGE,
                                "magic": MAGIC,
                                "comment": "24h close",
                                "type_time": mt5.ORDER_TIME_GTC,
                                "type_filling": mt5.ORDER_FILLING_IOC,
                            }
                            result = mt5.order_send(request)
                            if result.retcode == mt5.TRADE_RETCODE_DONE:
                                print(
                                    f"Closed {pos.symbol} after 24h | Ticket: {pos.ticket} | Profit: ${pos.profit:+.2f}"
                                )
                                if pos.ticket in open_positions:
                                    del open_positions[pos.ticket]

            # New analysis every 24 hours
            if (
                last_analysis_time is None
                or (now - last_analysis_time).total_seconds() >= 86400
            ):
                last_analysis_time = now
                print(f"\n{'='*80}")
                print(f"MARKET ANALYSIS: {now.strftime('%Y-%m-%d %H:%M')}")
                print(f"{'='*80}\n")

                for sym in SYMBOLS:
                    has_position = any(
                        p.symbol == sym and p.magic == MAGIC for p in (positions or [])
                    )
                    if has_position:
                        print(f"{sym}: already has an open position, skipping")
                        continue

                    rates = mt5.copy_rates_from_pos(sym, TIMEFRAME, 0, LOOKBACK)
                    if rates is None or len(rates) == 0:
                        continue

                    df = pd.DataFrame(rates)
                    df["time"] = pd.to_datetime(df["time"], unit="s")
                    df.set_index("time", inplace=True)
                    df = calculate_features(df)
                    if len(df) == 0:
                        continue

                    row = df.iloc[-1]
                    symbol_info = mt5.symbol_info(sym)
                    if symbol_info is None or not symbol_info.visible:
                        continue

                    prompt = f"""{sym} {now.strftime('%Y-%m-%d %H:%M')}
Current price: {row['close']:.5f}
RSI: {row['RSI']:.1f}
MACD: {row['MACD']:.6f}
ATR: {row['ATR']:.5f}
Volume: {row['vol_ratio']:.2f}x
BB position: {row['BB_position']:.2f}
Stochastic K: {row['Stoch_K']:.1f}
Analyze the situation and provide an accurate price forecast for 24 hours from now."""

                    resp = ollama.generate(
                        model=MODEL_NAME, prompt=prompt, options={"temperature": 0.3}
                    )
                    result = parse_answer(resp["response"])

                    print(f"{sym}: {result['dir']} ({result['prob']}%)")
                    if result.get("target_price"):
                        print(
                            f" Current: {row['close']:.5f} → 24-hour target: {result['target_price']:.5f}"
                        )

                    if result["prob"] < MIN_PROB:
                        print(f" Confidence {result['prob']}% < {MIN_PROB}%, skip\n")
                        continue

                    order_type = (
                        mt5.ORDER_TYPE_BUY
                        if result["dir"] == "UP"
                        else mt5.ORDER_TYPE_SELL
                    )
                    tick = mt5.symbol_info_tick(sym)
                    if tick is None:
                        continue
                    price = tick.ask if result["dir"] == "UP" else tick.bid

                    risk_amount = mt5.account_info().balance * RISK_PER_TRADE
                    point = symbol_info.point
                    atr_pips = row["ATR"] / point
                    stop_loss_pips = max(20, atr_pips * 2)
                    lot_size = risk_amount / (
                        stop_loss_pips * point * symbol_info.trade_contract_size
                    )
                    lot_step = symbol_info.volume_step
                    lot_size = round(lot_size / lot_step) * lot_step
                    lot_size = max(
                        symbol_info.volume_min, min(lot_size, symbol_info.volume_max)
                    )

                    sl = (
                        price - stop_loss_pips * point
                        if result["dir"] == "UP"
                        else price + stop_loss_pips * point
                    )
                    tp = (
                        price + stop_loss_pips * 3 * point
                        if result["dir"] == "UP"
                        else price - stop_loss_pips * 3 * point
                    )

                    request = {
                        "action": mt5.TRADE_ACTION_DEAL,
                        "symbol": sym,
                        "volume": lot_size,
                        "type": order_type,
                        "price": price,
                        "sl": sl,
                        "tp": tp,
                        "deviation": SLIPPAGE,
                        "magic": MAGIC,
                        "comment": f"AI_{result['prob']}%",
                        "type_time": mt5.ORDER_TIME_GTC,
                        "type_filling": mt5.ORDER_FILLING_IOC,
                    }

                    result_order = mt5.order_send(request)
                    if result_order.retcode == mt5.TRADE_RETCODE_DONE:
                        print(
                            f" Position opened! Ticket: {result_order.order}, Lot: {lot_size}, Price: {result_order.price:.5f}\n"
                        )
                        open_positions[result_order.order] = {
                            "symbol": sym,
                            "open_time": now,
                            "direction": result["dir"],
                            "lot": lot_size,
                        }
                    else:
                        print(f" Opening error: {result_order.comment}\n")

                print(f"{'='*80}")
                print(f"Open positions: {len(open_positions)}")
                print(
                    f"Next analysis: {(now + timedelta(hours=24)).strftime('%Y-%m-%d %H:%M')}"
                )
                print(f"{'='*80}\n")

            time.sleep(60)

        except KeyboardInterrupt:
            print("\nStopping trading...")
            positions = mt5.positions_get(magic=MAGIC)
            if positions:
                for pos in positions:
                    request = {
                        "action": mt5.TRADE_ACTION_DEAL,
                        "symbol": pos.symbol,
                        "volume": pos.volume,
                        "type": (
                            mt5.ORDER_TYPE_SELL
                            if pos.type == mt5.POSITION_TYPE_BUY
                            else mt5.ORDER_TYPE_BUY
                        ),
                        "position": pos.ticket,
                        "price": (
                            mt5.symbol_info_tick(pos.symbol).bid
                            if pos.type == mt5.POSITION_TYPE_BUY
                            else mt5.symbol_info_tick(pos.symbol).ask
                        ),
                        "deviation": SLIPPAGE,
                        "magic": MAGIC,
                        "comment": "manual close",
                        "type_time": mt5.ORDER_TIME_GTC,
                        "type_filling": mt5.ORDER_FILLING_IOC,
                    }
                    result = mt5.order_send(request)
                    if result.retcode == mt5.TRADE_RETCODE_DONE:
                        print(f"{pos.symbol} closed, profit: ${pos.profit:+.2f}")
            print("Trading stopped")
            break
        except Exception as e:
            log.error(f"Critical error: {e}")
            time.sleep(60)

    mt5.shutdown()


# ====================== 5. DATASET GENERATION ======================
def mode_dataset():
    """Dataset generation mode"""
    print("\n" + "=" * 80)
    print("5. DATASET GENERATION (24-HOUR FORECAST)")
    print("=" * 80)

    print("\nSelect a dataset type:")
    print("1. Real MT5 data")
    print("2. Synthetic data")

    dtype = input("\nChoice (1/2): ").strip()
    num_samples = input("Number of examples (default: 1000): ").strip()
    num_samples = int(num_samples) if num_samples else 1000

    if dtype == "1":
        dataset = generate_real_dataset_from_mt5(num_samples)
        dataset_path = save_dataset(dataset, "dataset/finetune_real_m5.jsonl")
    else:
        dataset = generate_synthetic_dataset(num_samples)
        dataset_path = save_dataset(dataset)

    print(f"\nDataset saved: {dataset_path}")
    print(f"Examples: {len(dataset)}")
    print(f"Size: {os.path.getsize(dataset_path) / 1024:.1f} KB")

    print("\n" + "=" * 80)
    print("EXAMPLE FROM THE DATASET:")
    print("=" * 80)
    example = dataset[0]
    print("\nPROMPT:")
    print(example["prompt"])
    print("\nANSWER:")
    print(example["response"])


# ====================== MENU ======================


# ====================== FORWARD TEST (10 days of unseen data) ======================
def forward_test():
    """🔬 HONEST FORWARD TEST on unseen data"""
    print("\n" + "=" * 80)
    print("🔬 FORWARD TEST - 10 DAYS OF UNSEEN DATA (honest validation)")
    print("=" * 80 + "\n")

    if not mt5 or not mt5.initialize():
        print("❌ MT5 is not connected")
        return

    end = datetime.now().replace(second=0, microsecond=0)
    start = end - timedelta(days=40)  # 30 days of training + 10 days of testing

    data = {}
    print(
        f"Loading data from {start.strftime('%Y-%m-%d')} to {end.strftime('%Y-%m-%d')}...\n"
    )

    for symbol in SYMBOLS:
        rates = mt5.copy_rates_range(symbol, TIMEFRAME, start, end)
        if rates is None or len(rates) < LOOKBACK + PREDICTION_HORIZON + 100:
            print(f"  ⚠️  {symbol}: insufficient data")
            continue

        df = pd.DataFrame(rates)
        df["time"] = pd.to_datetime(df["time"], unit="s")
        df.set_index("time", inplace=True)
        data[symbol] = df
        print(f"  ✅ {symbol}: {len(df)} bars ({(len(df)/96):.1f} days)")

    if not data:
        print("\n❌ No data available for the forward test")
        mt5.shutdown()
        return

    # Split: first 30 days = training, last 10 days = test
    total_days = min(len(df) // 96 for df in data.values())
    train_bars = (total_days - 10) * 96

    print(f"\n📊 DATA SPLIT:")
    print(f"  • Training: first {total_days - 10} days")
    print(f"  • FORWARD TEST: last 10 days (NEW DATA!)")
    print()

    balance = INITIAL_BALANCE
    equity = INITIAL_BALANCE
    trades = []
    balance_hist = [balance]
    equity_hist = [equity]
    slots = []

    # Get data for the forward test
    first_symbol = list(data.keys())[0]
    first_df = data[first_symbol].iloc[train_bars:]

    analysis_points = list(
        range(LOOKBACK, len(first_df) - PREDICTION_HORIZON, PREDICTION_HORIZON)
    )

    print(f"🎯 FORWARD TEST ANALYSIS:")
    print(f"  Analysis points: {len(analysis_points)} (every 24 hours)")
    print(f"  Initial balance: ${balance:,.2f}\n")
    print("=" * 80 + "\n")

    use_llm = False
    if ollama:
        try:
            ollama.list()
            use_llm = True
            print(f"✅ LLM model available — mode: AGI TRADER (LLM decides)\n")
        except:
            print(f"⚠️  Ollama is unavailable — mode: analysis only (CatBoost-like)\n")

    SPREAD_PIPS = 2

    # Main forward test loop
    for point_idx, offset in enumerate(analysis_points, 1):
        current_idx = train_bars + offset
        current_time = (
            first_df.index[offset]
            if hasattr(first_df.index[offset], "strftime")
            else datetime.now()
        )

        print(f"Analysis #{point_idx}/{len(analysis_points)}: {str(current_time)[:19]}")

        for symbol in SYMBOLS:
            if symbol not in data:
                continue

            df_sym = data[symbol]
            if current_idx + PREDICTION_HORIZON >= len(df_sym):
                continue

            df_features = calculate_features(df_sym.iloc[: current_idx + 1])
            if len(df_features) < 2:
                continue

            row = df_features.iloc[-1]
            future_row = df_sym.iloc[current_idx + PREDICTION_HORIZON]

            # Basic analysis (CatBoost-like techniques)
            rsi = row["RSI"]
            macd = row["MACD"]
            atr = row["ATR"]

            # Simple signal
            if rsi < 30 and macd > 0:
                direction = "UP"
                confidence = 72
            elif rsi > 70 and macd < 0:
                direction = "DOWN"
                confidence = 72
            elif macd > 0:
                direction = "UP"
                confidence = 65
            else:
                direction = "DOWN"
                confidence = 65

            # An LLM can improve this
            if use_llm:
                try:
                    prompt = f"""{symbol} FORWARD TEST {str(current_time)[:19]}
Current price: {row['close']:.5f}
RSI: {rsi:.1f} | MACD: {macd:.6f} | ATR: {atr:.5f} | Volume: {row['vol_ratio']:.2f}x

Analyze and provide a 24-hour forecast (DIRECTION: UP/DOWN, CONFIDENCE: %, PRICE: X.XXXXX)"""

                    resp = ollama.generate(
                        model=MODEL_NAME, prompt=prompt, options={"temperature": 0.3}
                    )
                    llm_text = resp.get("response", "")

                    if llm_text:
                        result = parse_answer(llm_text)
                        direction = result["dir"]
                        confidence = result["prob"]
                except:
                    pass

            # Result calculation
            entry_price = row["close"]
            exit_price = future_row["close"]
            symbol_info = mt5.symbol_info(symbol)

            if symbol_info:
                point = symbol_info.point
                price_move_pips = abs(exit_price - entry_price) / point

                if direction == "UP":
                    profit_pips = (
                        price_move_pips
                        if exit_price > entry_price
                        else -price_move_pips
                    )
                else:
                    profit_pips = (
                        price_move_pips
                        if exit_price < entry_price
                        else -price_move_pips
                    )

                # P&L with a 0.1 lot
                lot = 0.1
                profit_usd = profit_pips * point * 100000 * lot

                if profit_pips > 0:
                    balance += profit_usd
                    status = "✅"
                else:
                    balance -= abs(profit_usd)
                    status = "❌"

                actual_direction = "UP" if exit_price > entry_price else "DOWN"
                correct = direction == actual_direction

                trades.append(
                    {
                        "symbol": symbol,
                        "direction": direction,
                        "confidence": confidence,
                        "profit_pips": profit_pips,
                        "profit_usd": profit_usd,
                        "correct": correct,
                        "entry": entry_price,
                        "exit": exit_price,
                    }
                )

        balance_hist.append(balance)
        equity_hist.append(balance)
        slots.append({"datetime": current_time})

    # Forward test results
    print("\n" + "=" * 80)
    print("FORWARD TEST RESULTS")
    print("=" * 80 + "\n")

    print(f"Initial balance: ${INITIAL_BALANCE:,.2f}")
    print(f"Final balance: ${balance:,.2f}")
    print(f"Profit/Loss: ${balance - INITIAL_BALANCE:+,.2f}")
    print(f"Return: {((balance/INITIAL_BALANCE - 1) * 100):+.2f}%\n")

    if trades:
        wins = sum(1 for t in trades if t["profit_usd"] > 0)
        total = len(trades)
        win_rate = wins / total * 100
        avg_pnl = sum(t["profit_usd"] for t in trades) / total
        correct = sum(1 for t in trades if t["correct"])

        print(f"STATISTICS ({total} trades):")
        print(f"  Winning trades: {wins} ({win_rate:.1f}%)")
        print(f"  Correct direction: {correct} ({correct/total*100:.1f}%)")
        print(f"  Average P&L: ${avg_pnl:+.2f}")
        print(
            f"  Profit Factor: {sum(t['profit_usd'] for t in trades if t['profit_usd']>0) / (abs(sum(t['profit_usd'] for t in trades if t['profit_usd']<0)) + 0.01):.2f}"
        )

        if win_rate > 55:
            print(f"\n🎉 STATUS: READY TO TRADE (Win Rate > 55%)\n")
        elif win_rate > 50:
            print(f"\n⚠️  STATUS: NEUTRAL (Win Rate ~50%)\n")
        else:
            print(f"\n❌ STATUS: NEEDS IMPROVEMENT (Win Rate < 50%)\n")

    mt5.shutdown()


def main():
    """Main Menu"""
    print("\n" + "=" * 80)
    print(" 🚀 AGI TRADER v4 — PHILOSOPHIES OF ELITE TRADERS + FULL FREEDOM")
    print(" Built-in principles: Soros, Jones, Dalio, Livermore")
    print(" Modes: Push, Fine-tuning, Backtest, Live, Forward Test, Dataset")
    print("=" * 80)
    print("\nOPERATING MODES:")
    print("-" * 80)
    print("1 → Push the base model to Ollama")
    print("2 → FINE-TUNING (real MT5 data)")
    print("3 → BACKTEST on historical data (30 days)")
    print("4 → Live trading (MT5, every 24 hours)")
    print("🔬 5 → FORWARD TEST (10 days of unseen data) ⭐ RECOMMENDED")
    print("6 → Dataset generation (MT5 or synthetic)")
    print("-" * 80)

    choice = input("\nSelect a mode (1–6): ").strip()
    if choice == "1":
        mode_push()
    elif choice == "2":
        mode_finetune()
    elif choice == "3":
        backtest()
    elif choice == "4":
        live()
    elif choice == "5":
        forward_test()
    elif choice == "6":
        mode_dataset()
    else:
        print("Invalid choice")


if __name__ == "__main__":
    main()
