import logging
from datetime import datetime, timedelta
import argparse
import pytz
import pandas as pd
import numpy as np
import requests
from firstock import firstock

# Suppress pandas future warning
pd.set_option('future.no_silent_downcasting', True)

# ==========================================
# REMOTE API INTEGRATION
# ==========================================
def _send_entry_request(url, params, endpoint_name):
    """Helper method to send a single request with error handling"""
    try:
        response = requests.get(url, params=params, timeout=5)
        response.raise_for_status()
        print(f"[{endpoint_name}] API Request successful for {params['instrument']} | Signal: {params['signal']} | Time: {params['tickTime']}")
    except requests.RequestException as e:
        print(f"[{endpoint_name}] API Request failed: {e}")

def create_trend_entry(tick_time_str, instrument, close_price, signal, lot_count=1, trigger_type=None):
    """Sends buy/sell signals to remote trading endpoints"""
    api_instrument = instrument
    
    if not api_instrument.startswith('NIFTY') and not api_instrument.startswith('SENSEX'):
        if api_instrument.replace('CE', '').replace('PE', '').isdigit():
            api_instrument = f'NIFTY{api_instrument}'
    
    try:
        close_price = int(round(float(close_price)))
    except (ValueError, TypeError):
        print(f"Invalid close_price: {close_price}")
        return
    
    remarks_val = str(trigger_type).replace('_', '-') if trigger_type else 'MB'

    params = {
        'tickTime': str(tick_time_str),
        'instrument': str(api_instrument),
        'closePrice': str(close_price),
        'signal': str(signal),
        'orderType': str(lot_count),
        'remarks': remarks_val
    }
    
    endpoints = [
        ("Multibagger", "http://143.244.141.41/php/createEntriesMft.php"),
        ("LastSupper", "http://139.59.6.25/php/createEntriesMft.php"),
        ("GoodFriday", "http://68.183.85.105/php/createEntriesMft.php"),
    ]
    
    for endpoint_name, url in endpoints:
        _send_entry_request(url, params, endpoint_name)


# ==========================================
# ROBUST VECTORIZED QQE MOD ENGINE
# ==========================================
class QQEModIndicator:
    """Vectorized QQE MOD Engine matching PineScript using pandas ewm"""
    
    def __init__(self, 
                 rsi_length_primary=14, rsi_smoothing_primary=5, qqe_factor_primary=4.236, threshold_primary=0.1,
                 rsi_length_secondary=14, rsi_smoothing_secondary=5, qqe_factor_secondary=1.61, threshold_secondary=0.1,
                 bollinger_length=50, bollinger_multiplier=0.35):
        
        self.rsi_len_p = rsi_length_primary
        self.rsi_smooth_p = rsi_smoothing_primary
        self.qqe_factor_p = qqe_factor_primary
        self.threshold_p = threshold_primary
        
        self.rsi_len_s = rsi_length_secondary
        self.rsi_smooth_s = rsi_smoothing_secondary
        self.qqe_factor_s = qqe_factor_secondary
        self.threshold_s = threshold_secondary
        
        self.bb_length = bollinger_length
        self.bb_mult = bollinger_multiplier

    def calculate_rma(self, series, length):
        """Matches PineScript ta.rma (Wilder's MA) using ewm alpha = 1 / length"""
        alpha = 1.0 / length
        return series.ewm(alpha=alpha, adjust=False).mean()

    def calculate_ema(self, series, period):
        """Matches PineScript ta.ema using ewm span = period"""
        return series.ewm(span=period, adjust=False).mean()

    def calculate_rsi_wilders(self, prices, period):
        """Matches PineScript ta.rsi using Wilder's MA"""
        delta = prices.diff()
        gain = delta.where(delta > 0, 0.0)
        loss = -delta.where(delta < 0, 0.0)
        
        avg_gain = self.calculate_rma(gain, period)
        avg_loss = self.calculate_rma(loss, period)
        
        rs = avg_gain / avg_loss
        rsi = 100.0 - (100.0 / (1.0 + rs))
        return rsi.fillna(50.0)

    def _calculate_qqe_single(self, source, rsi_length, smoothing_factor, qqe_factor):
        wilders_length = rsi_length * 2 - 1
        
        rsi = self.calculate_rsi_wilders(source, rsi_length)
        smoothed_rsi = self.calculate_ema(rsi, smoothing_factor)
        
        atr_rsi = (smoothed_rsi.shift(1) - smoothed_rsi).abs().fillna(0.0)
        smoothed_atr_rsi = self.calculate_ema(atr_rsi, wilders_length)
        dynamic_atr_rsi = smoothed_atr_rsi * qqe_factor
        
        n = len(source)
        long_band = np.zeros(n)
        short_band = np.zeros(n)
        trend_direction = np.zeros(n, dtype=int)
        qqe_trend_line = np.zeros(n)

        rsi_vals = smoothed_rsi.values
        dar_vals = dynamic_atr_rsi.values

        if n > 0:
            long_band[0] = rsi_vals[0] - dar_vals[0]
            short_band[0] = rsi_vals[0] + dar_vals[0]
            trend_direction[0] = 0
            qqe_trend_line[0] = long_band[0]

        for i in range(1, n):
            r_val = rsi_vals[i]
            r_prev = rsi_vals[i-1]
            d_val = dar_vals[i]

            lb_prev = long_band[i-1]
            sb_prev = short_band[i-1]
            t_prev = trend_direction[i-1]

            new_long = r_val - d_val
            new_short = r_val + d_val

            if r_prev > lb_prev and r_val > lb_prev:
                lb = max(lb_prev, new_long)
            else:
                lb = new_long

            if r_prev < sb_prev and r_val < sb_prev:
                sb = min(sb_prev, new_short)
            else:
                sb = new_short

            sb_prev_val = sb_prev
            sb_prev_prev_val = short_band[i-2] if i >= 2 else sb_prev
            cross_short = ((r_prev < sb_prev_prev_val) and (r_val >= sb_prev_val)) or ((r_prev > sb_prev_prev_val) and (r_val <= sb_prev_val))

            lb_prev_val = lb_prev
            lb_prev_prev_val = long_band[i-2] if i >= 2 else lb_prev
            cross_long = ((lb_prev_prev_val < r_prev) and (lb_prev_val >= r_val)) or ((lb_prev_prev_val > r_prev) and (lb_prev_val <= r_val))

            if cross_short:
                t = 1
            elif cross_long:
                t = -1
            else:
                t = t_prev

            long_band[i] = lb
            short_band[i] = sb
            trend_direction[i] = t
            qqe_trend_line[i] = lb if t == 1 else sb

        return pd.Series(qqe_trend_line, index=source.index), smoothed_rsi

    def calculate(self, close_prices, open_prices):
        df = pd.DataFrame(index=close_prices.index)

        primary_trend_line, primary_rsi = self._calculate_qqe_single(
            close_prices, self.rsi_len_p, self.rsi_smooth_p, self.qqe_factor_p
        )
        
        secondary_trend_line, secondary_rsi = self._calculate_qqe_single(
            close_prices, self.rsi_len_s, self.rsi_smooth_s, self.qqe_factor_s
        )
        
        df['primary_rsi'] = primary_rsi
        df['primary_rsi_sub50'] = primary_rsi - 50.0
        df['primary_qqe_tl'] = primary_trend_line
        
        df['secondary_rsi'] = secondary_rsi

        # New columns requested
        df['secondary_rsi_histogram'] = secondary_rsi - secondary_trend_line  # Backwards compatibility

        df['secondaryRsiHistogramValue'] = secondary_rsi
        h = df['secondaryRsiHistogramValue']
        df['hsSignal'] = (h.shift(3) < h.shift(2)) & (h.shift(2) < h.shift(1)) & (h.shift(1) > h)
        df['hbSignal'] = (h.shift(3) > h.shift(2)) & (h.shift(2) > h.shift(1)) & (h.shift(1) < h)


        df['secondary_trend_line'] = secondary_trend_line - 50.0

        # Standard QQE Mod Crossover Signals matching QMB / QMS tags
        df['rsi_cross_up'] = (df['primary_rsi_sub50'] > 0) & (df['primary_rsi_sub50'].shift(1) <= 0)
        df['rsi_cross_down'] = (df['primary_rsi_sub50'] < 0) & (df['primary_rsi_sub50'].shift(1) >= 0)

        df['hist_cross_up'] = (df['secondary_rsi_histogram'] > 0) & (df['secondary_rsi_histogram'].shift(1) <= 0)
        df['hist_cross_down'] = (df['secondary_rsi_histogram'] < 0) & (df['secondary_rsi_histogram'].shift(1) >= 0)

        df['QMB'] = df['hist_cross_up'] | df['rsi_cross_up']
        df['QMS'] = df['hist_cross_down'] | df['rsi_cross_down']

        df['qqe_long_signal'] = df['hbSignal']
        df['qqe_short_signal'] = df['hsSignal']



        '''df['strike_price'] = np.where(
            df['qqe_long_signal'],
            np.ceil(close_prices / 50.0) * 50,
            np.where(
                df['qqe_short_signal'],
                np.floor(close_prices / 50.0) * 50,
                np.nan
            )
        )'''

        df['strike_price'] = np.where(
            df['qqe_long_signal'],
            np.ceil(open_prices / 50.0) * 50,
            np.where(
                df['qqe_short_signal'],
                np.floor(open_prices / 50.0) * 50,
                np.nan
            )
        )
        
        return df


# ==========================================
# DATA FETCHER CLASS
# ==========================================
class StockDataFetcher:
    def __init__(self, client_details):
        self.user_id = client_details[0]
        self.client_details = client_details
        self.ist = pytz.timezone('Asia/Kolkata')
        self.setup_logger()

    def setup_logger(self):
        logging.basicConfig(
            level=logging.INFO,
            format='%(asctime)s - %(levelname)s - %(message)s',
            datefmt='%Y-%m-%d %H:%M:%S'
        )
        self.logger = logging.getLogger(__name__)

    def login(self):
        try:
            response = firstock.login(*self.client_details)
            if response.get("status") == "success":
                self.logger.info("Login successful")
                return True
            else:
                self.logger.error(f"Login failed: {response}")
                return False
        except Exception as e:
            self.logger.error(f"Login error: {e}")
            return False

    def fetch_time_price_series(self, exchange, trading_symbol, start_time, end_time, interval):
        try:
            response = firstock.timePriceSeries(
                userId=self.user_id,
                exchange=exchange,
                tradingSymbol=trading_symbol,
                startTime=start_time,
                endTime=end_time,
                interval=interval
            )
            
            if response.get("status") == "success":
                data = response.get("data", [])
                if data:
                    df = pd.DataFrame(data)
                    if 'time' in df.columns:
                        df['datetime'] = pd.to_datetime(df['time'], format='%H:%M:%S %d-%m-%Y')
                    self.logger.info(f"Fetched {len(df)} records")
                    return df
            return pd.DataFrame()
            
        except Exception as e:
            self.logger.error(f"Error fetching data: {e}")
            return pd.DataFrame()

    def fetch_data(self, symbol, interval_minutes=5, count=1000, elapsed=0):
        now_ist = datetime.now(self.ist)
        target_date = now_ist.date() - timedelta(days=elapsed)
        
        if elapsed == 0:
            end_time = now_ist
            if now_ist.hour >= 15 and now_ist.minute >= 31:
                end_time = now_ist.replace(hour=15, minute=30, second=0)
        else:
            end_time = datetime.combine(target_date, datetime.min.time()).replace(hour=15, minute=30, second=0)
            end_time = self.ist.localize(end_time) if end_time.tzinfo is None else end_time
        
        start_time = end_time - timedelta(days=25)
        start_time = start_time.replace(hour=9, minute=15, second=0)
        
        interval_str = f"{interval_minutes}mi"
        start_str = start_time.strftime("%H:%M:%S %d-%m-%Y")
        end_str = end_time.strftime("%H:%M:%S %d-%m-%Y")
        
        translated_symbol = 'NSE:NIFTY'
        exchange, trading_symbol = translated_symbol.split(":")
        
        self.logger.info(f"Fetching {interval_minutes}min data for {symbol} up to target date: {target_date}")
        
        df = self.fetch_time_price_series(exchange, trading_symbol, start_str, end_str, interval_str)
        
        if not df.empty:
            df = df.sort_values(by='datetime', ascending=True)
            df = df[(df['datetime'].dt.hour >= 9) & (df['datetime'].dt.hour <= 15)]
            df = df[(df['datetime'].dt.hour != 9) | (df['datetime'].dt.minute >= 15)]
            df = df.tail(count)
            
        return df


def analyze_with_qqe_mod(fetcher, symbol='nifty', interval_minutes=5, count=1000, elapsed=0):
    df = fetcher.fetch_data(symbol, interval_minutes, count, elapsed)
    
    if df.empty:
        print("No data fetched")
        return pd.DataFrame(), {}

    # ==========================================
    # PRINT CURRENT TIME & LAST TWO RAW CANDLES
    # ==========================================
    current_time_ist = datetime.now(fetcher.ist).strftime('%Y-%m-%d %H:%M:%S')
    print("\n" + "="*85)
    print(f"CURRENT TIME (IST): {current_time_ist}")
    print("="*85)
    print("LAST TWO RAW CANDLES:")
    print(df.tail(2).to_string(index=False))
    print("="*85 + "\n")
    
    qqe_mod = QQEModIndicator()
    qqe_data = qqe_mod.calculate(df['close'], df['open'])
    
    result_df = pd.concat([df, qqe_data], axis=1)
    
    target_date = datetime.now(fetcher.ist).date() - timedelta(days=elapsed)
    result_df_target = result_df[result_df['datetime'].dt.date == target_date].copy()
    
    if result_df_target.empty:
        print(f"No data available for target date {target_date}.")
        return pd.DataFrame(), {}

    # ==========================================
    # DEBUG TRACE PRINTING
    # ==========================================

    print("\n" + "~"*110)
    print(f"DEBUG TRACE REPORT FOR TARGET DATE: {target_date}")
    print("~"*110)
    print(f"{'TIME':<10} | {'OPEN':<8} | {'CLOSE':<8} | {'PRIM_RSI-50':<12} | {'HISTOGRAM':<12} | {'SEC_HIST_VAL':<14} | {'LONG_SIG':<9} | {'SHORT_SIG':<9}")
    print("-" * 100)
    for _, r in result_df_target.iterrows():
        t_str = r['datetime'].strftime('%H:%M')
        o_val = f"{r['open']:.2f}"
        c_val = f"{r['close']:.2f}"
        p_rsi = f"{r['primary_rsi_sub50']:.2f}"
        hist = f"{r['secondary_rsi_histogram']:.2f}"
        sec_hist_val = f"{r['secondaryRsiHistogramValue']:.2f}"
        l_sig = str(r['qqe_long_signal'])
        s_sig = str(r['qqe_short_signal'])
        
        marker = " 🎯" if r['qqe_long_signal'] or r['qqe_short_signal'] else ""
        print(f"{t_str:<10} | {o_val:<8} | {c_val:<8} | {p_rsi:<12} | {hist:<12} | {sec_hist_val:<14} | {l_sig:<9} | {s_sig:<9}{marker}")
    print("~"*110 + "\n")

    long_signals = result_df_target[result_df_target['qqe_long_signal'] == True]
    short_signals = result_df_target[result_df_target['qqe_short_signal'] == True]

    signal_data = {
        'long_signals': long_signals,
        'short_signals': short_signals,
        'current_long': result_df_target['qqe_long_signal'].iloc[-1] if len(result_df_target) > 0 else False,
        'current_short': result_df_target['qqe_short_signal'].iloc[-1] if len(result_df_target) > 0 else False,
        'total_long': len(long_signals),
        'total_short': len(short_signals)
    }
    
    return result_df_target, signal_data


# ==========================================
# TRADE EXECUTION ENGINE WITH REMOTE API
# ==========================================
def process_trades_with_sl_tp(result_df, target_pts=30, sl_pts=15):
    if result_df.empty:
        return pd.DataFrame()

    if 'high' not in result_df.columns:
        result_df['high'] = result_df['close']
    if 'low' not in result_df.columns:
        result_df['low'] = result_df['close']

    trades = []
    active_trade = None

    for idx, row in result_df.iterrows():
        dt = row['datetime']
        time_obj = dt.time()
        close_price = float(row['close'])
        high_price = float(row['high'])
        low_price = float(row['low'])
        tick_time_str = dt.strftime("%Y-%m-%d %H:%M:%S")

        if active_trade is not None and time_obj >= datetime.strptime("15:15", "%H:%M").time():
            active_trade['exit_time'] = dt
            active_trade['exit_price'] = close_price
            
            pnl = (close_price - active_trade['entry_price']) if active_trade['signal_type'] == 'LONG' else (active_trade['entry_price'] - close_price)
            active_trade['pnl_pts'] = round(pnl, 2)
            active_trade['status'] = 'CLOSED AT 15:15 ⏰'
            
            create_trend_entry(
                tick_time_str=tick_time_str,
                instrument=active_trade['instrument'],
                close_price=close_price,
                signal="SELL",
                lot_count=1,
                trigger_type="TIME-EXIT-1515"
            )

            trades.append(active_trade)
            active_trade = None
            continue

        if active_trade is not None:
            entry_price = active_trade['entry_price']
            signal_type = active_trade['signal_type']
            
            hit_tp = False
            hit_sl = False
            exit_price = 0.0
            
            if signal_type == 'LONG':
                target_level = entry_price + target_pts
                sl_level = entry_price - sl_pts
                
                if high_price >= target_level:
                    hit_tp = True
                    exit_price = target_level
                elif low_price <= sl_level:
                    hit_sl = True
                    exit_price = sl_level
                    
            elif signal_type == 'SHORT':
                target_level = entry_price - target_pts
                sl_level = entry_price + sl_pts
                
                if low_price <= target_level:
                    hit_tp = True
                    exit_price = target_level
                elif high_price >= sl_level:
                    hit_sl = True
                    exit_price = sl_level
            
            if hit_tp or hit_sl:
                trigger_remark = "TARGET-HIT" if hit_tp else "SL-HIT"
                active_trade['exit_time'] = dt
                active_trade['exit_price'] = exit_price
                active_trade['status'] = 'TARGET HIT 🎯' if hit_tp else 'SL HIT ❌'
                active_trade['pnl_pts'] = target_pts if hit_tp else -sl_pts
                
                create_trend_entry(
                    tick_time_str=tick_time_str,
                    instrument=active_trade['instrument'],
                    close_price=exit_price,
                    signal="SELL",
                    lot_count=1,
                    trigger_type=trigger_remark
                )

                trades.append(active_trade)
                active_trade = None

        is_long = row.get('qqe_long_signal', False)
        is_short = row.get('qqe_short_signal', False)

        if is_long or is_short:
            incoming_signal = 'LONG' if is_long else 'SHORT'
            
            if active_trade is not None and active_trade['signal_type'] != incoming_signal:
                active_trade['exit_time'] = dt
                active_trade['exit_price'] = close_price
                
                pnl = (close_price - active_trade['entry_price']) if active_trade['signal_type'] == 'LONG' else (active_trade['entry_price'] - close_price)
                active_trade['pnl_pts'] = round(pnl, 2)
                active_trade['status'] = 'SIGNAL REVERSAL 🔄'
                
                create_trend_entry(
                    tick_time_str=tick_time_str,
                    instrument=active_trade['instrument'],
                    close_price=close_price,
                    signal="SELL",
                    lot_count=1,
                    trigger_type="signal-change"
                )

                trades.append(active_trade)
                active_trade = None

            if active_trade is None and time_obj < datetime.strptime("15:00", "%H:%M").time():
                option_type = 'CE' if is_long else 'PE'
                strike_val = int(row.get('strike_price')) if pd.notnull(row.get('strike_price')) else 0
                
                instrument_symbol = f"{strike_val}{option_type}"
                
                if incoming_signal == 'LONG':
                    target_price = close_price + target_pts
                    sl_price = close_price - sl_pts
                else:
                    target_price = close_price - target_pts
                    sl_price = close_price + sl_pts

                active_trade = {
                    'entry_time': dt,
                    'signal_type': incoming_signal,
                    'option_type': option_type,
                    'instrument': instrument_symbol,
                    'strike': strike_val,
                    'entry_price': close_price,
                    'target_price': target_price,
                    'sl_price': sl_price,
                    'exit_time': None,
                    'exit_price': None,
                    'status': 'OPEN ⏳',
                    'pnl_pts': 0.0
                }

                create_trend_entry(
                    tick_time_str=tick_time_str,
                    instrument=instrument_symbol,
                    close_price=close_price,
                    signal="BUY",
                    lot_count=1,
                    trigger_type=f"QQE-MOD-{incoming_signal}"
                )

    if active_trade is not None and time_obj >= datetime.strptime("15:15", "%H:%M").time():
        last_row = result_df.iloc[-1]
        exit_dt = last_row['datetime']
        exit_price = float(last_row['close'])
        
        pnl = (exit_price - active_trade['entry_price']) if active_trade['signal_type'] == 'LONG' else (active_trade['entry_price'] - exit_price)
        
        active_trade['exit_time'] = exit_dt
        active_trade['exit_price'] = exit_price
        active_trade['pnl_pts'] = round(pnl, 2)
        active_trade['status'] = 'CLOSED AT EOD 🕒'
        
        create_trend_entry(
            tick_time_str=exit_dt.strftime("%Y-%m-%d %H:%M:%S"),
            instrument=active_trade['instrument'],
            close_price=exit_price,
            signal="SELL",
            lot_count=1,
            trigger_type="EOD-CLOSE"
        )
        trades.append(active_trade)

    return pd.DataFrame(trades)


def print_trade_execution_report(trades_df):
    print("\n" + "="*85)
    print("QQE MOD TRADING EXECUTION REPORT")
    print("="*85)
    
    if trades_df.empty:
        print("No QQE MOD trades triggered for the selected date.")
        return

    for _, trade in trades_df.iterrows():
        print(f"Signal: {trade['signal_type']:5} | Symbol: NIFTY{trade['instrument']}")
        print(f"  Entry Time : {trade['entry_time'].strftime('%H:%M:%S')} @ {trade['entry_price']:.2f}")
        print(f"  Target Level: {trade['target_price']:.2f} | SL Level: {trade['sl_price']:.2f}")
        print(f"  Exit Time  : {trade['exit_time'].strftime('%H:%M:%S')} @ {trade['exit_price']:.2f}")
        print(f"  Result     : {trade['status']} (PnL: {trade['pnl_pts']:+.2f} pts)")
        print("-" * 85)

    total_pnl = trades_df['pnl_pts'].sum()
    print(f"Total Trades: {len(trades_df)} | Total PnL: {total_pnl:+.2f} Index Points")
    print("="*85)


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="QQE MOD Trading Script")
    parser.add_argument("--elapsed", type=int, default=0, help="Elapsed time parameter (defaults to 0 if not provided)",)
    args = parser.parse_args()
    
    client_details = ['DB1485', 'ABcd$1234', '14121985', 'DB1485_API', 'c18634e5002598698626e3590b52c520']
    
    fetcher = StockDataFetcher(client_details)
    if not fetcher.login():
        exit(1)
    
    symbol = 'nifty'
    result_df, signals = analyze_with_qqe_mod(fetcher, symbol, interval_minutes=5, count=1000, elapsed=args.elapsed)
    
    if not result_df.empty:
        trades_df = process_trades_with_sl_tp(result_df, target_pts=300, sl_pts=150)
        print_trade_execution_report(trades_df)
    else:
        print("No data fetched or processed for the specified elapsed days.")