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 = 20
        self.max_profit = 200
        self.coolingPeriod = 14
        self.min_body_size = 10  # Minimum body size for trade entry
        
        # 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_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 Heikin Ashi color changes and candle characteristics."""
        self.logger.info(f"Input DataFrame columns: {list(df.columns)}, rows: {len(df)}")
        
        minimum_columns = ['time', 'open', 'high', 'low', '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
        
        if len(df) == 0:
            self.logger.error("Input DataFrame is empty")
            return df
        
        df = df.copy()
        df['time'] = pd.to_datetime(df['time'])
        numeric_cols = ['open', 'high', 'low', 'close']
        df[numeric_cols] = df[numeric_cols].apply(pd.to_numeric, errors='coerce')
        
        # Calculate standard Heikin Ashi candles
        df['ha_close'] = (df['open'] + df['high'] + df['low'] + df['close']) / 4
        df['ha_open'] = np.nan
        df['ha_high'] = np.nan
        df['ha_low'] = np.nan
        
        # First candle
        if len(df) > 0:
            df.loc[0, 'ha_open'] = (df.loc[0, 'open'] + df.loc[0, 'close']) / 2
            df.loc[0, 'ha_high'] = max(df.loc[0, 'high'], df.loc[0, 'ha_open'], df.loc[0, 'ha_close'])
            df.loc[0, 'ha_low'] = min(df.loc[0, 'low'], df.loc[0, 'ha_open'], df.loc[0, 'ha_close'])
        
        # Subsequent candles (recursive)
        for i in range(1, len(df)):
            df.loc[i, 'ha_open'] = (df.loc[i-1, 'ha_open'] + df.loc[i-1, 'ha_close']) / 2
            df.loc[i, 'ha_high'] = max(df.loc[i, 'high'], df.loc[i, 'ha_open'], df.loc[i, 'ha_close'])
            df.loc[i, 'ha_low'] = min(df.loc[i, 'low'], df.loc[i, 'ha_open'], df.loc[i, 'ha_close'])
        
        # Determine color and candle characteristics
        df['color'] = np.where(df['ha_close'] > df['ha_open'], 'green', 'red')
        df['body_size'] = abs(df['ha_close'] - df['ha_open'])
        df['lower_shadow'] = np.where(df['color'] == 'green', df['ha_open'] - df['ha_low'], df['ha_close'] - df['ha_low'])
        df['upper_shadow'] = np.where(df['color'] == 'green', df['ha_high'] - df['ha_close'], df['ha_high'] - df['ha_open'])
        
        # 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
        
        # Initialize columns
        df['signal'] = 'NEUTRAL'
        df['position'] = 'NONE'
        df['pl'] = 0.0
        df['candle_strength'] = 'NORMAL'
        
        prev_color = None
        
        for i, row in df_trading_day.iterrows():
            try:
                current_time = row['time']
                close = row['close']
                current_color = row['color']
                body_size = row['body_size']
                lower_shadow = row['lower_shadow']
                upper_shadow = row['upper_shadow']
                
                # Determine candle strength
                if current_color == 'green' and lower_shadow == 0:
                    df.loc[i, 'candle_strength'] = 'STRONG_UPTREND'
                elif current_color == 'red' and upper_shadow == 0:
                    df.loc[i, 'candle_strength'] = 'STRONG_DOWNTREND'
                else:
                    df.loc[i, 'candle_strength'] = 'NORMAL'
                
                if prev_color is not None:
                    # Check for color change to trigger exit
                    if prev_color != current_color:
                        exit_reason = 'ColorChange'
                        if self.positions['CE'] is not None:
                            ce_strike, ce_price, ce_time = self.positions['CE']
                            ce_instrument = f"NIFTY{ce_strike}CE"
                            pl = (close - ce_price) * self.lotCount
                            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.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
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'SELL', 
                                'instrument': ce_instrument,
                                'pl': pl,
                                'exit_reason': exit_reason
                            })
                        if self.positions['PE'] is not None:
                            pe_strike, pe_price, pe_time = self.positions['PE']
                            pe_instrument = f"NIFTY{pe_strike}PE"
                            pl = (pe_price - close) * self.lotCount
                            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.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
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'SELL', 
                                'instrument': pe_instrument,
                                'pl': pl,
                                'exit_reason': exit_reason
                            })
                    
                    # Check for entry conditions with body size filter
                    if df.loc[i, 'candle_strength'] == 'STRONG_UPTREND' and body_size >= self.min_body_size and self.positions['CE'] is None:
                        strike = self.get_strike_price(close, 'CE')
                        if strike is not None:
                            ce_instrument = f"NIFTY{strike}CE"
                            self.logger.info(
                                f"{current_time} | CE Buy     | "
                                f"price: {close:.2f} | Strength: {df.loc[i, 'candle_strength']} | Body Size: {body_size:.2f}"
                            )
                            self.create_order(current_time, ce_instrument, close, 'BUY', self.lotCount)
                            self.positions['CE'] = (strike, close, current_time)
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'BUY', 
                                'instrument': ce_instrument
                            })
                    
                    elif df.loc[i, 'candle_strength'] == 'STRONG_DOWNTREND' and body_size >= self.min_body_size and self.positions['PE'] is None:
                        strike = self.get_strike_price(close, 'PE')
                        if strike is not None:
                            pe_instrument = f"NIFTY{strike}PE"
                            self.logger.info(
                                f"{current_time} | PE Buy     | "
                                f"price: {close:.2f} | Strength: {df.loc[i, 'candle_strength']} | Body Size: {body_size:.2f}"
                            )
                            self.create_order(current_time, pe_instrument, close, 'BUY', self.lotCount)
                            self.positions['PE'] = (strike, close, current_time)
                            self.signals.append({
                                'time': current_time, 
                                'price': close, 
                                'type': 'BUY', 
                                'instrument': pe_instrument
                            })
                
                df.loc[i, 'signal'] = current_color.upper() if current_color else 'NEUTRAL'
                df.loc[i, 'position'] = 'CE' if self.positions['CE'] is not None else 'PE' if self.positions['PE'] is not None else 'NONE'
                df.loc[i, 'pl'] = self.total_pl
                prev_color = current_color
                
            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=15, minute=30, 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")