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

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 = 5
        self.lotCount = 1
        self.no_of_past_data_days = 4  # Number of past days to fetch data for
        self.max_loss = 200  # Maximum loss threshold for closing positions
        self.max_profit = 20
        
        # Track positions and P/L
        self.positions = {'CE': None, 'PE': None}  # {'strike': (entry_price, entry_time)}
        self.total_pl = 0.0

    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):
        """Create BUY/SELL order"""
        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):
        """Calculate nearest strike price"""
        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:

        # Convert numeric columns
        numeric_cols = ['close', 'high', 'low', 'open']
        for col in numeric_cols:
            df[col] = pd.to_numeric(df[col], errors='coerce')
        
        # Calculate OHLC4 (average price)
        df['ohlc4'] = (df['open'] + df['high'] + df['low'] + df['close']) / 4

        
        df['date'] = pd.to_datetime(df['time']).dt.date



        # Anchored TWAP: cumulative average for each day
        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']

        # Calculate TWAP using 59-tick window
        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)

       
        # Apply 13-tick smoothing
        df['ma'] = df['twap'].rolling(window=self.smoothing_period, min_periods=1).mean()
        
        return df
    
    def calculate_day_twap(self, df: pd.DataFrame) -> pd.DataFrame:

        
        return df
    

    def process_data(self, df: pd.DataFrame):
        """Main trading logic that only logs trend flips"""
        df = self.calculate_twap(df)

        df = self.calculate_day_twap(df)



        
        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
        
        current_trend = None  # None, 'UP', or 'DOWN'
        profit_tracker =0
        
        for i in range(self.timeframe_period, len(df)):
            try:
                current_time = df.loc[i, 'time']
                close = float(df.loc[i, 'close'])
                twap = float(df.loc[i, 'twap'])
                ma = float(df.loc[i, 'ma'])
                prev_close = float(df.loc[i-1, 'close'])
                prev_twap = float(df.loc[i-1, 'twap'])
                prev_ma = float(df.loc[i-1, 'ma'])
                a_twap = float(df.loc[i, 'a_twap'])
                
                valid_up = (
                    (prev_close < prev_twap) and 
                    (close > twap) and 
                    (close > ma) and
                    (ma > twap) )
                
                valid_down = (
                    (prev_close > prev_twap) and 
                    (close < twap) and 
                    (close < ma) and
                    (ma < twap) )
                

                valid_up = ((prev_close < prev_twap) and
                    (close > twap) )
                
                valid_down = ((prev_close > prev_twap) and
                    (close < twap) )
                
                   
                           
                # Determine new trend state
                new_trend = None
                if valid_up:
                    new_trend = 'UP'
                elif valid_down:
                    new_trend = 'DOWN'

                '''
                print(
                        f"time: {current_time}, "
                        f"prev_close: {prev_close:.2f}, prev_twap: {prev_twap:.2f}, "
                        f"close: {close:.2f}, twap: {twap:.2f}, ma: {ma:.2f}, "
                        f"(prev_close < prev_twap): {prev_close < prev_twap}, "
                        f"(close > twap): {close > twap}, "
                        f"(close > ma): {close > ma}, "
                        f"(ma > twap): {ma > twap}, "
                        f"(prev_close > prev_twap): {prev_close > prev_twap}, "
                        f"(close < twap): {close < twap}, "
                        f"(close < ma): {close < ma}, "
                        f"(ma < twap): {ma < twap}, "
                        f"valid_up: {valid_up}, valid_down: {valid_down}, "
                        f"trend: {new_trend}"
                    )
                '''
                

                if self.positions['PE'] is not None:
                    pe_strike, pe_price, pe_time = self.positions['PE']
                    pl = (pe_price - close) * self.lotCount
                    force_close = False

                    if profit_tracker < pl:
                        profit_tracker = pl
                        #self.logger.info(f"{current_time} | max_profit: {profit_tracker:.2f}")
                    else:
                        if pl < (profit_tracker * 0.9) and profit_tracker > 200:
                            force_close = True



                    if pl < -self.max_loss or pl > self.max_profit or force_close or close > twap:
                        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}"
                        )
                        pe_instrument = f"NIFTY{pe_strike}PE"
                        self.create_order(current_time, pe_instrument, close, 'SELL', self.lotCount)
                        self.total_pl += pl
                        self.positions['PE'] = None


                if self.positions['CE'] is not None:
                    ce_strike, ce_price, ce_time = self.positions['CE']
                    pl = (close - ce_price) * self.lotCount
                    force_close = False

                    if profit_tracker < pl:
                        profit_tracker = pl
                        #self.logger.info(f"{current_time} | max_profit: {profit_tracker:.2f}")
                    else:
                        if pl < (profit_tracker * 0.9) and profit_tracker > 200:
                            force_close = True

                    if pl < -self.max_loss or pl > self.max_profit or force_close or close < twap:
                        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}"
                                )
                        ce_instrument = f"NIFTY{ce_strike}CE"
                        self.create_order(current_time, ce_instrument, close, 'SELL', self.lotCount)
                        self.total_pl += pl
                        self.positions['CE'] = None



                # Only log and act on trend changes
                if new_trend is not None and new_trend != current_trend:
                    if new_trend == 'UP':
                        
                        strike = self.get_strike_price(close, 'CE')
                        if strike is not None and a_twap < close :
                            self.logger.info(
                                    f"{current_time} | UPTREND    | "
                                    f"Close: {close:.2f} > TWAP: {twap:.2f} | "
                                    f"MA: {ma:.2f} > TWAP: {twap:.2f}"
                                )

                            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
                                self.logger.info(
                                    f"{current_time} | PE Close      | "
                                    f"sell: {close:.2f} | Buy: {pe_price:.2f} | "
                                    f"pl: {pl:.2f} "
                                )
                                self.total_pl += pl
                                self.positions['PE'] = None

                            ce_instrument = f"NIFTY{strike}CE"
                            self.create_order(current_time, ce_instrument, close, 'BUY', self.lotCount)
                            self.positions['CE'] = (strike, close, current_time)
                       
                            profit_tracker = 0
                                    


                      
                            
                            
                    
                    elif new_trend == 'DOWN':
                        strike = self.get_strike_price(close, 'PE')
                        if strike is not None  and a_twap > close :

                            self.logger.info(
                                    f"{current_time} | DOWNTREND  | "
                                    f"Close: {close:.2f} < TWAP: {twap:.2f} | "
                                    f"MA: {ma:.2f} < TWAP: {twap:.2f}"
                                )

                            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
                                self.logger.info(
                                            f"{current_time} | CE Close   | "
                                            f"sell: {close:.2f} | Buy: {ce_price:.2f} | "
                                            f"pl: {pl:.2f} "
                                        )
                                self.total_pl += pl
                                self.positions['CE'] = None

                            pe_instrument = f"NIFTY{strike}PE"
                            self.create_order(current_time, pe_instrument, close, 'BUY', self.lotCount)
                            self.positions['PE'] = (strike, close, current_time)
                            profit_tracker = 0


                    current_trend = new_trend
                            
            except Exception as e:
                self.logger.error(f"Error processing data at index {i}: {e}")
                

    def fetch_and_process_data(self, elapsed: int):
        """Fetch data and run trading strategy"""
        now_ist = datetime.now(self.ist)
        start_time = now_ist.replace(hour=9, minute=15, second=0) - timedelta(days=elapsed + self.no_of_past_data_days)
        end_time = now_ist.replace(hour=15, minute=31, second=0) - timedelta(days=elapsed)
        
        try:
            # Fetch Nifty data
            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'
                    })
                    # Convert time and filter market hours
                    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.process_data(df)
                    self.logger.info(f"Total P/L: ₹{self.total_pl:.2f}")
                else:
                    self.logger.error("No data returned from API")
            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")