import logging
import sys
import time
from datetime import datetime, timedelta, time
from typing import List
import pytz
import pandas as pd
import requests
import json
from thefirstock import thefirstock
import numpy as np

def get_client_details():
    try:
        url = 'http://143.244.141.41/php/getUserDetails.php'
        response = requests.get(url)
        response.raise_for_status()
        data = response.json()
        if not data.get('success', False):
            raise ValueError("getUserDetails endpoint returned unsuccessful response")
        user_data = data['data']
        return [user_data['field1'], user_data['field2'], user_data['field3'], user_data['field4'], user_data['field5']]
    except requests.exceptions.RequestException as e:
        print(f"HTTP Request failed: {e}")
        return None
    except (json.JSONDecodeError, KeyError) as e:
        print(f"Failed to parse response: {e}")
        return None

class TWAPTrader:
    def __init__(self, client_details: List[str], createEntries: bool = False):
        self.client_details = client_details
        self.user_id = client_details[0]
        self.ist = pytz.timezone('Asia/Kolkata')
        self.logger = self.setup_logger()
        self.createEntries = createEntries
        
        # Trading parameters
        self.timeframe_period = 21
        self.smoothing_period = 14  
        self.tick_interval = 1
        self.lotCount = 1
        self.numberOfPrevDays = 3
        self.max_loss = 200
        self.max_profit = 200
        self.coolingPeriod = 14
        
        # Track positions and P/L
        self.positions = {'CE': None, 'PE': None}
        self.total_pl = 0.0
        self.trade_history = []
        self.signals = []

    class ISTFormatter(logging.Formatter):
        def formatTime(self, record, datefmt=None):
            ist = pytz.timezone('Asia/Kolkata')
            record_time = datetime.fromtimestamp(record.created, tz=ist)
            return record_time.strftime(datefmt or '%Y-%m-%d %H:%M:%S')
        
    def setup_logger(self) -> logging.Logger:
        logger = logging.getLogger(__name__)
        logger.setLevel(logging.DEBUG)
        if not logger.handlers:
            formatter = self.ISTFormatter('%(asctime)s - %(levelname)s - %(message)s')
            console_handler = logging.StreamHandler()
            console_handler.setFormatter(formatter)
            logger.addHandler(console_handler)
        logging.getLogger().handlers.clear()
        return logger

    def login(self):
        try:
            self.logger.info(f"Attempting login for {self.client_details[0]}")
            response = thefirstock.firstock_login(*self.client_details)
            if response.get("status") == "success":
                self.logger.info("Login successful")
            else:
                self.logger.error(f"Login failed: {response}")
                sys.exit()
        except Exception as e:
            self.logger.error(f"Login error: {e}")
            sys.exit()

    def create_order(self, tickTimeStr, instrument, price, action, lotCount):
        if self.createEntries:
            url = "http://143.244.141.41/php/createEntries.php"
            params = {
                'tickTime': str(tickTimeStr),
                'instrument': str(instrument),
                'closePrice': str(price),
                'signal': str(action),
                'orderType': str(lotCount)
            }
            try:
                response = requests.get(url, params=params)
                if response.status_code != 200:
                    self.logger.error(f"Order failed with status {response.status_code}: {response.text}")
            except Exception as e:
                self.logger.error(f"Order request failed: {e}")

    def get_strike_price(self, current_price, option_type):
        try:
            price = float(current_price)
            remainder = price % 50
            if remainder > 25:
                return int(price + (50 - remainder))
            else:
                return int(price - remainder)
        except (ValueError, TypeError):
            self.logger.error(f"Invalid price for strike calculation: {current_price}")
            return None

    def calculate_twap(self, df: pd.DataFrame) -> pd.DataFrame:
        numeric_cols = ['close', 'high', 'low', 'open']
        for col in numeric_cols:
            df[col] = pd.to_numeric(df[col], errors='coerce')
        
        df['ohlc4'] = (df['open'] + df['high'] + df['low'] + df['close']) / 4
        df['date'] = pd.to_datetime(df['time']).dt.date
        df['cum_sum'] = df.groupby('date')['ohlc4'].cumsum()
        df['cum_count'] = df.groupby('date').cumcount() + 1
        df['a_twap'] = df['cum_sum'] / df['cum_count']
        df['price_accum'] = df['ohlc4'].rolling(window=self.timeframe_period, min_periods=1).sum()
        df['weight'] = df.index.to_series().rolling(window=self.timeframe_period, min_periods=1).count() - 1
        df['twap'] = df['price_accum'] / (df['weight'] + 1)
        df['ma'] = df['twap'].rolling(window=self.smoothing_period, min_periods=1).mean()
        df['ema9'] = df['ohlc4'].ewm(span=9, adjust=False, min_periods=1).mean()
        return df
    
    def calculate_day_twap(self, df: pd.DataFrame) -> pd.DataFrame:
        return df

    def calculate_win_rate(self):
        if not self.trade_history:
            return 0.0
        profitable_trades = sum(1 for trade in self.trade_history if trade['pl'] > 0)
        total_trades = len(self.trade_history)
        return (profitable_trades / total_trades * 100) if total_trades > 0 else 0.0

    def process_data(self, df: pd.DataFrame, trading_date: datetime.date) -> pd.DataFrame:
        """Process trading data to generate signals and manage positions based on EMA9-TWAP difference.
        Prevent trading for 30 ticks after a negative P/L for an individual trade."""
        # Log DataFrame info for debugging
        self.logger.info(f"Input DataFrame columns: {list(df.columns)}, rows: {len(df)}")
        
        # Validate minimum required columns
        minimum_columns = ['time', 'close']
        if not all(col in df.columns for col in minimum_columns):
            self.logger.error(f"Missing minimum required columns: {set(minimum_columns) - set(df.columns)}")
            return df
        
        # Ensure DataFrame has enough rows
        if len(df) == 0:
            self.logger.error("Input DataFrame is empty")
            return df
        
        df = df.copy()
        df['close'] = df['close'].astype(float)
        df['time'] = pd.to_datetime(df['time'])
        
        # Calculate required indicators if not present
        if 'twap' not in df.columns:
            self.logger.info("Calculating TWAP")
            df = self.calculate_twap(df)
            if 'twap' not in df.columns:
                self.logger.error("Failed to calculate TWAP column")
                return df
        
        if 'a_twap' not in df.columns:
            self.logger.info("Calculating day TWAP")
            df = self.calculate_day_twap(df)
            if 'a_twap' not in df.columns:
                self.logger.error("Failed to calculate a_twap column")
                return df
        
        if 'ma' not in df.columns:
            self.logger.info(f"Calculating {self.smoothing_period}-period simple moving average")
            df['ma'] = df['close'].rolling(window=self.smoothing_period).mean()
        
        if 'ema9' not in df.columns:
            self.logger.info("Calculating 9-period exponential moving average")
            df['ema9'] = df['close'].ewm(span=9, adjust=False).mean()
        
        # Calculate ema_twap_diff
        df['ema_twap_diff'] = df['ema9'] - df['twap']
        
        # Validate all required columns
        required_columns = ['time', 'close', 'twap', 'ma', 'ema9', 'a_twap', 'ema_twap_diff']
        missing_cols = set(required_columns) - set(df.columns)
        if missing_cols:
            self.logger.error(f"Missing required columns after calculations: {missing_cols}")
            return df
        
        # Check for NaN values in critical columns
        nan_counts = df[required_columns].isna().sum()
        if nan_counts.any():
            self.logger.error(f"NaN values detected in columns: {nan_counts[nan_counts > 0].to_dict()}")
            return df
        
        # Ensure numeric columns
        numeric_cols = ['close', 'twap', 'ma', 'ema9', 'a_twap', 'ema_twap_diff']
        df[numeric_cols] = df[numeric_cols].apply(pd.to_numeric, errors='coerce')
        
        # Check for sufficient data points
        if len(df) < (self.smoothing_period + self.timeframe_period):
            self.logger.error(f"Not enough data points for {self.timeframe_period}-tick TWAP with {self.smoothing_period}-tick smoothing")
            return df
        
        # Filter by trading date
        df_trading_day = df[df['time'].dt.date == trading_date].copy()
        if df_trading_day.empty:
            self.logger.error(f"No data available for trading date {trading_date}")
            return df
        
        # Verify ema_twap_diff in filtered DataFrame
        if 'ema_twap_diff' not in df_trading_day.columns:
            self.logger.error("ema_twap_diff column missing in df_trading_day")
            return df
        if df_trading_day['ema_twap_diff'].isna().any():
            self.logger.error(f"NaN values in ema_twap_diff for trading date {trading_date}")
            return df
        
        # Initialize columns
        df['signal'] = 'NEUTRAL'
        df['position'] = 'NONE'
        df['pl'] = 0.0
        
        current_trend = None
        profit_tracker = 0
        prev_ema_twap_diff = 0
        ema_twap_peak = float('-inf')  # Track max ema_twap_diff for CE positions
        ema_twap_trough = float('inf')  # Track min ema_twap_diff for PE positions
        diff_threshold = 5.0  # Exit if ema_twap_diff moves 5 points in opposite direction
        
        # Initialize cooldown variables if not already set
        if not hasattr(self, 'cooldown_active'):
            self.cooldown_active = False
            self.cooldown_ticks = 0
        
        for i, row in df_trading_day.iterrows():
            try:
                current_time = row['time']
                close = row['close']
                twap = row['twap']
                ma = row['ma']
                ema9 = row['ema9']
                a_twap = row['a_twap']
                ema_twap_diff = row['ema_twap_diff']
                
                # Handle cooldown logic
                if self.cooldown_active:
                    self.cooldown_ticks -= 1
                   # self.logger.debug(f"Cooldown active, {self.cooldown_ticks} ticks remaining")
                    if self.cooldown_ticks <= 0:
                        self.cooldown_active = False
                        self.logger.info(f"Cooldown period ended at {current_time}, resuming trading")
                
                # Update peak/trough for ema_twap_diff
                if self.positions['CE'] is not None:
                    ema_twap_peak = max(ema_twap_peak, ema_twap_diff)
                elif self.positions['PE'] is not None:
                    ema_twap_trough = min(ema_twap_trough, ema_twap_diff)
                
                # Determine trend
                new_trend = None
                can_enter = False
                if ema_twap_diff > 0:
                    new_trend = 'UP'
                    can_enter = (self.positions['CE'] is None) or (prev_ema_twap_diff <= 0)
                elif ema_twap_diff < 0:
                    new_trend = 'DOWN'
                    can_enter = (self.positions['PE'] is None) or (prev_ema_twap_diff >= 0)
                
                df.loc[i, 'signal'] = new_trend if new_trend else 'NEUTRAL'
                
                # Handle PE position exit
                if self.positions['PE'] is not None:
                    pe_strike, pe_price, pe_time = self.positions['PE']
                    pl = (pe_price - close) * self.lotCount
                    exit_reason = None
                    
                    if profit_tracker < pl:
                        profit_tracker = pl
                    
                    if pl < (profit_tracker * 0.9) and profit_tracker > 200:
                        exit_reason = 'ProfitTrack'
                    elif pl < -self.max_loss:
                        exit_reason = 'SL'
                    elif pl > self.max_profit:
                        exit_reason = 'PF'
                    elif ema_twap_diff > (ema_twap_trough + diff_threshold):
                        exit_reason = 'EMATWAPDiff'
                    
                    if exit_reason:
                        self.logger.info(
                            f"{current_time} | PE Close   | "
                            f"sell: {close:.2f} | Buy: {pe_price:.2f} | "
                            f"pl: {pl:.2f} | max_profit: {profit_tracker:.2f} | "
                            f"exit_reason: {exit_reason}"
                        )
                        pe_instrument = f"NIFTY{pe_strike}PE"
                        self.create_order(current_time, pe_instrument, close, 'SELL', self.lotCount)
                        self.total_pl += pl
                        self.trade_history.append({'pl': pl, 'exit_reason': exit_reason})
                        self.positions['PE'] = None
                        ema_twap_trough = float('inf')  # Reset trough
                        self.signals.append({
                            'time': current_time, 
                            'price': close, 
                            'type': 'SELL', 
                            'instrument': pe_instrument,
                            'pl': pl,
                            'exit_reason': exit_reason
                        })
                        # Check if trade P/L is negative to trigger cooldown
                        if pl < 0:
                            self.cooldown_active = True
                            self.cooldown_ticks = self.coolingPeriod 
                            self.logger.info(f"Negative trade P/L ({pl:.2f}) for PE, entering 30-tick cooldown at {current_time}")
                
                # Handle CE position exit
                if self.positions['CE'] is not None:
                    ce_strike, ce_price, ce_time = self.positions['CE']
                    pl = (close - ce_price) * self.lotCount
                    exit_reason = None
                    
                    if profit_tracker < pl:
                        profit_tracker = pl
                    
                    if pl < (profit_tracker * 0.9) and profit_tracker > 200:
                        exit_reason = 'ProfitTrack'
                    elif pl < -self.max_loss:
                        exit_reason = 'SL'
                    elif pl > self.max_profit:
                        exit_reason = 'PF'
                    elif ema_twap_diff < (ema_twap_peak - diff_threshold):
                        exit_reason = 'EMATWAPDiff'
                    
                    if exit_reason:
                        self.logger.info(
                            f"{current_time} | CE Close   | "
                            f"sell: {close:.2f} | Buy: {ce_price:.2f} | "
                            f"pl: {pl:.2f} | max_profit: {profit_tracker:.2f} | "
                            f"exit_reason: {exit_reason}"
                        )
                        ce_instrument = f"NIFTY{ce_strike}CE"
                        self.create_order(current_time, ce_instrument, close, 'SELL', self.lotCount)
                        self.total_pl += pl
                        self.trade_history.append({'pl': pl, 'exit_reason': exit_reason})
                        self.positions['CE'] = None
                        ema_twap_peak = float('-inf')  # Reset peak
                        self.signals.append({
                            'time': current_time, 
                            'price': close, 
                            'type': 'SELL', 
                            'instrument': ce_instrument,
                            'pl': pl,
                            'exit_reason': exit_reason
                        })
                        # Check if trade P/L is negative to trigger cooldown
                        if pl < 0:
                            self.cooldown_active = True
                            self.cooldown_ticks = self.coolingPeriod
                            self.logger.info(f"Negative trade P/L ({pl:.2f}) for CE, entering 30-tick cooldown at {current_time}")
                
                # Handle trend changes and position entry only if not in cooldown
                if not self.cooldown_active and new_trend is not None and new_trend != current_trend and can_enter:
                    if new_trend == 'UP':
                        strike = self.get_strike_price(close, 'CE')
                        if strike is not None:
                            self.logger.info(
                                f"{current_time} | UPTREND    | "
                                f"Close: {close:.2f} | TWAP: {twap:.2f} | "
                                f"EMA9-TWAP: {ema_twap_diff:.2f}"
                            )
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'TREND', 
                                'trend': 'UP'
                            })
                            
                            if self.positions['PE'] is not None:
                                pe_strike, pe_price, pe_time = self.positions['PE']
                                pe_instrument = f"NIFTY{pe_strike}PE"
                                self.create_order(current_time, pe_instrument, close, 'SELL', self.lotCount)
                                pl = (pe_price - close) * self.lotCount
                                exit_reason = 'Cutover'
                                self.logger.info(
                                    f"{current_time} | PE Close   | "
                                    f"sell: {close:.2f} | Buy: {pe_price:.2f} | "
                                    f"pl: {pl:.2f} | exit_reason: {exit_reason}"
                                )
                                self.total_pl += pl
                                self.trade_history.append({'pl': pl, 'exit_reason': exit_reason})
                                self.positions['PE'] = None
                                ema_twap_trough = float('inf')  # Reset trough
                                self.signals.append({
                                    'time': current_time, 
                                    'price': close, 
                                    'type': 'SELL', 
                                    'instrument': pe_instrument,
                                    'pl': pl,
                                    'exit_reason': exit_reason
                                })
                                # Check if trade P/L is negative to trigger cooldown
                                if pl < 0:
                                    self.cooldown_active = True
                                    self.cooldown_ticks = self.coolingPeriod
                                    self.logger.info(f"Negative trade P/L ({pl:.2f}) for PE cutover, entering 30-tick cooldown at {current_time}")
                                    current_trend = None
                                    continue
                            
                            ce_instrument = f"NIFTY{strike}CE"
                            self.create_order(current_time, ce_instrument, close, 'BUY', self.lotCount)
                            self.positions['CE'] = (strike, close, current_time)
                            ema_twap_peak = ema_twap_diff  # Initialize peak
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'BUY', 
                                'instrument': ce_instrument
                            })
                            profit_tracker = 0
                    
                    elif new_trend == 'DOWN':
                        strike = self.get_strike_price(close, 'PE')
                        if strike is not None:
                            self.logger.info(
                                f"{current_time} | DOWNTREND  | "
                                f"Close: {close:.2f} | TWAP: {twap:.2f} | "
                                f"EMA9-TWAP: {ema_twap_diff:.2f}"
                            )
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'TREND', 
                                'trend': 'DOWN'
                            })
                            
                            if self.positions['CE'] is not None:
                                ce_strike, ce_price, ce_time = self.positions['CE']
                                ce_instrument = f"NIFTY{ce_strike}CE"
                                self.create_order(current_time, ce_instrument, close, 'SELL', self.lotCount)
                                pl = (close - ce_price) * self.lotCount
                                exit_reason = 'Cutover'
                                self.logger.info(
                                    f"{current_time} | CE Close   | "
                                    f"sell: {close:.2f} | Buy: {ce_price:.2f} | "
                                    f"pl: {pl:.2f} | exit_reason: {exit_reason}"
                                )
                                self.total_pl += pl
                                self.trade_history.append({'pl': pl, 'exit_reason': exit_reason})
                                self.positions['CE'] = None
                                ema_twap_peak = float('-inf')  # Reset peak
                                self.signals.append({
                                    'time': current_time, 
                                    'price': close, 
                                    'type': 'SELL', 
                                    'instrument': ce_instrument,
                                    'pl': pl,
                                    'exit_reason': exit_reason
                                })
                                # Check if trade P/L is negative to trigger cooldown
                                if pl < 0:
                                    self.cooldown_active = True
                                    self.cooldown_ticks = self.coolingPeriod
                                    self.logger.info(f"Negative trade P/L ({pl:.2f}) for CE cutover, entering 30-tick cooldown at {current_time}")
                                    current_trend = None
                                    continue
                            
                            pe_instrument = f"NIFTY{strike}PE"
                            self.create_order(current_time, pe_instrument, close, 'BUY', self.lotCount)
                            self.positions['PE'] = (strike, close, current_time)
                            ema_twap_trough = ema_twap_diff  # Initialize trough
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'BUY', 
                                'instrument': pe_instrument
                            })
                            profit_tracker = 0
                    
                    current_trend = new_trend
                
                df.loc[i, 'pl'] = self.total_pl
                df.loc[i, 'position'] = 'CE' if self.positions['CE'] is not None else 'PE' if self.positions['PE'] is not None else 'NONE'
                prev_ema_twap_diff = ema_twap_diff
                
            except Exception as e:
                self.logger.error(f"Error processing data at index {i}: {str(e)}")
        
        self.df_processed = df
        win_rate = self.calculate_win_rate()
        self.logger.info(f"Total P/L for {trading_date}: ₹{self.total_pl:.2f}")
        self.logger.info(f"Win Rate: {win_rate:.2f}% (Profitable trades: {sum(1 for t in self.trade_history if t['pl'] > 0)} / Total trades: {len(self.trade_history)})")
        
        return self.df_processed
    
    def fetch_and_process_data(self, elapsed: int):
        now_ist = datetime.now(self.ist)
        start_time = now_ist.replace(hour=9, minute=15, second=0) - timedelta(days=elapsed + self.numberOfPrevDays)
        end_time = now_ist.replace(hour=14, minute=51, second=0) - timedelta(days=elapsed)
        
        try:
            response = thefirstock.firstock_TimePriceSeries(
                userId=self.user_id,
                exchange='NSE',
                tradingSymbol='Nifty 50',
                startTime=start_time.strftime("%d/%m/%Y %H:%M:%S"),
                endTime=end_time.strftime("%d/%m/%Y %H:%M:%S"),
                interval=str(self.tick_interval)
            )
            
            if response.get("status") == "success":
                df = pd.DataFrame(response.get("data", []))
                if not df.empty:
                    df = df[['time', 'intc', 'inth', 'intl', 'into']].rename(columns={
                        'intc': 'close',
                        'inth': 'high',
                        'intl': 'low',
                        'into': 'open'
                    })
                    df['time'] = pd.to_datetime(df['time'], format='%d-%m-%Y %H:%M:%S')
                    df = df.sort_values(by='time', ascending=True)
                    df = df[(df['time'].dt.time >= time(9, 15)) & (df['time'].dt.time <= time(15, 30))]
                    df.reset_index(drop=True, inplace=True)
                    
                    self.total_pl = 0.0
                    self.trade_history = []
                    self.process_data(df, end_time.date())
                    win_rate = self.calculate_win_rate()
                    self.logger.info(f"Final Summary for {end_time.date()}:")
                    self.logger.info(f"Total P/L: ₹{self.total_pl:.2f}")
                    self.logger.info(f"Win Rate: {win_rate:.2f}% (Profitable trades: {sum(1 for t in self.trade_history if t['pl'] > 0)} / Total trades: {len(self.trade_history)})")
                else:
                    self.logger.error("No data returned from API for the specified period")
            else:
                self.logger.error(f"API request failed: {response}")
                
        except Exception as e:
            self.logger.error(f"Error in fetch_and_process_data: {e}")

if __name__ == "__main__":
    client_details = get_client_details()
    if client_details:
        elapsed = float(sys.argv[1]) if len(sys.argv) > 1 else 0
        
        createEntries = len(sys.argv) > 2
        trader = TWAPTrader(client_details, createEntries)
        trader.login()
        trader.fetch_and_process_data(elapsed)
    else:
        print("Failed to get client details")