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
import math
#from thefirstock import thefirstock
from ta.trend import ADXIndicator
from firstock import firstock
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 StockDataFetcher:

    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

        self.trendTransactionsCounter = 0
        self.reverseTransactionsCounter = 0

        self.total_pl = 0.0
        self.total_trend_pl = 0.0
        self.total_reverse_pl = 0.0

       
        self.lotCount = 1                            # number of lots at each strike price. Control via client
        self.reverseLotCount = 4
        self.tick_interval = 1
        self.order_allowed_steps = 1                 # number of steps allowed from strike price to enter orders

        # Moving Average Parameters (from Pine Script)
        self.ma_length = 20
        self.ma_type = 2  # 1=SMA, 2=EMA, 3=WMA, 4=HullMA, 5=VWMA, 6=RMA, 7=TEMA
        self.ma_length2 = 50  # Second MA length
        self.ma_type2 = 1     # Second MA type
        self.color_smoothing = 2
        self.use_second_ma = True
        self.show_crosses = True

        #Trend following Orders -related params
        self.enable_trend_following_orders = True
        self.trend_follow_sigma_start_threshold = 5
        self.trend_following_safeMargin = 5.0
        self.allowed_trend_follow_negatives = 10
        self.trend_following_fixed_profit = 100.0

        # Trend reversal Entry and Exit Signal Levels
        self.enable_trend_reversal_orders = True
        self.trend_reverse_sigma_start_threshold = 5
        self.reversal_entry_exit_safe_margin = 5.0
        self.ce_entry_level = 'sig3_down'
        self.ce_exit_level = 'sig3_up'
        self.pe_entry_level = 'sig3_up'
        self.pe_exit_level = 'sig3_down'
        
        self.debug_data = ''
        
    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 addLogDataDebug(self, text):
        self.logger.debug(f'{text}')
        self.debug_data += text + '\n'

    def addLogDataInfo(self, text):
        self.logger.info(f'{text}')
        self.debug_data += text + '\n'

    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 = 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 process_symbol_data(self, symbol: str, interval: int, start_time: datetime, end_time: datetime):
        exchange, trading_symbol = symbol.split(":")
        df = self.fetch_time_price_series(
            exchange, trading_symbol,
            start_time.strftime("%H:%M:%S %d-%m-%Y"),
            end_time.strftime("%H:%M:%S %d-%m-%Y"),
            str(interval),
        )

        df = df.sort_values(by='time', ascending=True)
        df = df.drop(['volume', 'oi'], axis=1)

        return df

    def calculate_moving_average(self, series: pd.Series, length: int, ma_type: int) -> pd.Series:
        """
        Calculate moving averages based on Pine Script logic
        1=SMA, 2=EMA, 3=WMA, 4=HullMA, 5=VWMA, 6=RMA, 7=TEMA
        """
        if ma_type == 1:  # SMA
            return series.rolling(window=length).mean()
        elif ma_type == 2:  # EMA
            return series.ewm(span=length, adjust=False).mean()
        elif ma_type == 3:  # WMA
            weights = np.arange(1, length + 1)
            def wma_func(x):
                return np.dot(x, weights) / weights.sum() if len(x) == length else np.nan
            return series.rolling(window=length).apply(wma_func, raw=True)
        elif ma_type == 4:  # Hull MA
            wma_half = series.rolling(window=length//2).apply(
                lambda x: np.dot(x, np.arange(1, len(x)+1)) / np.arange(1, len(x)+1).sum() if len(x) == length//2 else np.nan, 
                raw=True
            )
            wma_full = series.rolling(window=length).apply(
                lambda x: np.dot(x, np.arange(1, len(x)+1)) / np.arange(1, len(x)+1).sum() if len(x) == length else np.nan, 
                raw=True
            )
            hull_series = 2 * wma_half - wma_full
            return hull_series.rolling(window=int(np.sqrt(length))).apply(
                lambda x: np.dot(x, np.arange(1, len(x)+1)) / np.arange(1, len(x)+1).sum() if len(x) == int(np.sqrt(length)) else np.nan, 
                raw=True
            )
        elif ma_type == 5:  # VWMA (Volume Weighted - requires volume data)
            # For simplicity, using close price as both price and volume proxy
            return series.rolling(window=length).mean()
        elif ma_type == 6:  # RMA (Relative Moving Average)
            return series.ewm(alpha=1/length, adjust=False).mean()
        elif ma_type == 7:  # TEMA
            ema1 = series.ewm(span=length, adjust=False).mean()
            ema2 = ema1.ewm(span=length, adjust=False).mean()
            ema3 = ema2.ewm(span=length, adjust=False).mean()
            return 3 * (ema1 - ema2) + ema3
        else:
            return series.rolling(window=length).mean()  # Default to SMA

    def calculate_ma_direction(self, ma_series: pd.Series, smoothing: int) -> pd.Series:
        """Calculate MA direction for color coding"""
        return ma_series >= ma_series.shift(smoothing)

    def detect_ma_cross(self, ma1: pd.Series, ma2: pd.Series) -> pd.Series:
        """Detect crosses between two moving averages"""
        cross_up = (ma1 > ma2) & (ma1.shift(1) <= ma2.shift(1))
        cross_down = (ma1 < ma2) & (ma1.shift(1) >= ma2.shift(1))
        return cross_up | cross_down

    def createTrendEntry(self, tickTimeStr, instrument, closePrice, signal, lotCount):
        return
        instrument = 'NIFTY' + str(instrument)
        self.trendTransactionsCounter += 1

        if self.createEntries:
            url = "http://143.244.141.41/php/createEntries.php"
            tickTimeStr += ':' + str(self.trendTransactionsCounter).zfill(3)

            params = {
                'tickTime': str(tickTimeStr),
                'instrument': str(instrument),
                'closePrice': str(closePrice),
                'signal': str(signal),
                'orderType': str(lotCount)
            }

            response = requests.get(url, params=params)
            if response.status_code != 200:
                self.addLogDataDebug(f"Request failed with status code: {response.status_code}")
                if response.text != '':
                    self.addLogDataDebug(response.text + "\n")

    def createReverseEntry(self, tickTimeStr, instrument, closePrice, signal, lotCount):
        return
        instrument = 'NIFTY' + str(instrument)
        self.reverseTransactionsCounter += 1

        if self.createEntries:
            url = "http://143.244.141.41/php/createEntries.php"
            tickTimeStr += ':' + str(self.reverseTransactionsCounter).zfill(6)

            params = {
                'tickTime': str(tickTimeStr),
                'instrument': str(instrument),
                'closePrice': str(closePrice),
                'signal': str(signal),
                'orderType': str(lotCount)
            }

            response = requests.get(url, params=params)
            if response.status_code != 200:
                self.addLogDataDebug(f"Request failed with status code: {response.status_code}")
                if response.text != '':
                    self.addLogDataDebug(response.text + "\n")

    def fetch_time_price_series(self, exchange: str, trading_symbol: str, start_time: str, end_time: str, interval: str) -> pd.DataFrame:
        #self.addLogDataDebug(f"Fetching data for {trading_symbol} from {start_time} to {end_time}. interval:{interval}")

        print(f"userId: {self.user_id}, exchange: {exchange}, tradingSymbol: {trading_symbol}, startTime: {start_time}, endTime: {end_time}, interval: {interval}")

        try:
            response = firstock.timePriceSeries(
                userId=self.user_id,
                exchange=exchange,
                tradingSymbol=trading_symbol,
                startTime=start_time,
                endTime=end_time,
                interval=str(interval) + "mi"
            )  

            if response.get("status") == "success":
                return pd.DataFrame(response.get("data", []))
            self.addLogDataDebug(f"Fetch failed: {response}")
            return pd.DataFrame()
        except Exception as e:
            self.addLogDataDebug(f"Error fetching series: {e}")
            return pd.DataFrame()

    def getNiftyStrikePriceGivenCurrentValue(self, closePrice, type):
        strike_prices = []
        try:
            price = float(closePrice)
            remainder = price % 50

            base_strike = int(price + (50 - remainder)) if remainder > 25 else int(price - remainder)

        except (ValueError, TypeError):
            raise ValueError("Invalid closePrice input. Must be a number or numeric string.")

        if base_strike < price and type == 'PE':
            base_strike += 50
        elif base_strike > price and type == 'CE':
            base_strike -= 50

        for step in range(self.order_allowed_steps):
            if type == 'CE':
                strike_prices.append(base_strike - 50 * step)
            else:
                strike_prices.append(base_strike + 50 * step)

        return strike_prices
    
    def handle_trend_following_orders(self, df: pd.DataFrame):
        self.open_positions_price = {'CE': [], 'PE': []}

        position = None
        entry_price = None
        last_exit_type = 'NONE'
        last_twap_direction = None
        num_negatives = 0
        first_entry_allowed = False

        for i in range(1, len(df)):
            _time = df.loc[i, 'time']
            price = df.loc[i, 'close']
            twap = df.loc[i, 'twap']
            high = df.loc[i, 'high']
            low = df.loc[i, 'low']
            sig2_up = df.loc[i, 'sig2_up']
            sig2_down = df.loc[i, 'sig2_down']
            
            # Get MA values and signals
            ma1 = df.loc[i, 'ma1']
            ma2 = df.loc[i, 'ma2'] if self.use_second_ma else None
            ma_cross = df.loc[i, 'ma_cross'] if self.use_second_ma and self.show_crosses else False
            ma_direction = df.loc[i, 'ma_direction']

            if not first_entry_allowed:
                std = df.loc[i, 'std']
                twap_deviation = abs(price - twap)
                if twap_deviation <= 0.5 * std:
                    first_entry_allowed = True
                    self.addLogDataInfo(f"[INFO] First entry allowed at {_time} | Price: {price} | TWAP: {twap:.2f} | Std: {std:.2f}")
                else:
                    continue

            if num_negatives >= self.allowed_trend_follow_negatives:
                continue

            # === MA Cross Signal ===
            if ma_cross and self.show_crosses:
                self.addLogDataInfo(f"[MA CROSS] Detected at {_time} | Price: {price} | MA1: {ma1:.2f} | MA2: {ma2:.2f}")

            # === TWAP Cutover Reset Detection ===
            if last_exit_type in ['CE', 'PE']:
                if (high > twap and price < twap) or (low < twap and price > twap):
                    last_exit_type = 'NONE'
                    self.addLogDataInfo(f"[RESET] TWAP crossover at {_time} | High: {high} / Low: {low} vs TWAP: {twap:.2f}")

            # === TWAP Direction Change Tracking ===
            prev_price = df.loc[i - 1, 'close']
            prev_twap = df.loc[i - 1, 'twap']
            if prev_price < prev_twap and price > twap:
                last_twap_direction = 'UP'
            elif prev_price > prev_twap and price < twap:
                last_twap_direction = 'DOWN'

            # === ENTRY with MA Filter ===
            if position is None:
                # Bullish entry: Price above TWAP AND MA trending up
                if (price > twap + self.trend_following_safeMargin and 
                    (last_exit_type != 'CE' or last_twap_direction == 'UP') and
                    ma_direction):  # MA trending up
                    
                    strikePrices = self.getNiftyStrikePriceGivenCurrentValue(price, 'CE')
                    for strike in strikePrices:
                        instrument = f"{strike}CE"
                        self.createTrendEntry(str(_time), instrument, price, 'BUY', self.lotCount)
                        self.open_positions_price['CE'].append((instrument, price))
                    self.addLogDataInfo(f"[ENTRY] BUY CE at {_time} | Price: {price} > TWAP: {twap:.2f} | MA Up: {ma_direction}")
                    position = 'CE'
                    entry_price = price
                    last_exit_type = 'NONE'
                    last_twap_direction = None

                # Bearish entry: Price below TWAP AND MA trending down
                elif (price < twap - self.trend_following_safeMargin and 
                      (last_exit_type != 'PE' or last_twap_direction == 'DOWN') and
                      not ma_direction):  # MA trending down
                    
                    strikePrices = self.getNiftyStrikePriceGivenCurrentValue(price, 'PE')
                    for strike in strikePrices:
                        instrument = f"{strike}PE"
                        self.createTrendEntry(str(_time), instrument, price, 'BUY', self.lotCount)
                        self.open_positions_price['PE'].append((instrument, price))
                    self.addLogDataInfo(f"[ENTRY] BUY PE at {_time} | Price: {price} < TWAP: {twap:.2f} | MA Down: {not ma_direction}")
                    position = 'PE'
                    entry_price = price
                    last_exit_type = 'NONE'
                    last_twap_direction = None

            # === EXIT CE ===
            elif position == 'CE':
                exit_condition = (price < twap) or (ma_cross and ma1 < ma2)  # Added MA cross down condition
                profit_condition = (price - entry_price >= self.trend_following_fixed_profit) or (price >= sig2_up)
                
                if exit_condition or profit_condition:
                    for instr, buy_price in self.open_positions_price['CE']:
                        pl = round((price - buy_price) * self.lotCount, 2)
                        self.createTrendEntry(str(_time), instr, price, 'SELL', self.lotCount)
                        self.total_trend_pl += pl
                        self.addLogDataInfo(f"[EXIT] CE at {_time} | Price: {price} | P/L: ₹{pl} | "
                                          f"Reason: {'TWAP/MACross' if exit_condition else 'Profit'}")
                        if pl < 0:
                            num_negatives = num_negatives + 1
                    self.open_positions_price['CE'].clear()
                    position = None
                    entry_price = None
                    last_exit_type = 'CE'
                    last_twap_direction = None

            # === EXIT PE ===
            elif position == 'PE':
                exit_condition = (price > twap) or (ma_cross and ma1 > ma2)  # Added MA cross up condition
                profit_condition = (entry_price - price >= self.trend_following_fixed_profit) or (price <= sig2_down)
                
                if exit_condition or profit_condition:
                    for instr, buy_price in self.open_positions_price['PE']:
                        pl = round((buy_price - price) * self.lotCount, 2)
                        self.createTrendEntry(str(_time), instr, price, 'SELL', self.lotCount)
                        self.total_trend_pl += pl
                        self.addLogDataInfo(f"[EXIT] PE at {_time} | Price: {price} | P/L: ₹{pl} | "
                                          f"Reason: {'TWAP/MACross' if exit_condition else 'Profit'}")
                        if pl < 0:
                            num_negatives = num_negatives + 1
                    self.open_positions_price['PE'].clear()
                    position = None
                    entry_price = None
                    last_exit_type = 'PE'
                    last_twap_direction = None

    def handle_trend_reversal_orders(self, df: pd.DataFrame):
        self.open_positions_price = {'CE': [], 'PE': []}
        position = None
        entry_price = None

        ce_entry_col = self.ce_entry_level
        ce_exit_col = self.ce_exit_level
        pe_entry_col = self.pe_entry_level
        pe_exit_col = self.pe_exit_level

        for i in range(0, len(df)):
            _time = df.loc[i, 'time']
            price = df.loc[i, 'close']
            twap = df.loc[i, 'twap']
            
            # Get MA values for additional confirmation
            ma1 = df.loc[i, 'ma1']
            ma2 = df.loc[i, 'ma2'] if self.use_second_ma else None
            ma_direction = df.loc[i, 'ma_direction']

            ce_entry = df.loc[i, ce_entry_col] + self.reversal_entry_exit_safe_margin
            pe_entry = df.loc[i, pe_entry_col] - self.reversal_entry_exit_safe_margin
            ce_exit = df.loc[i, ce_exit_col] - self.reversal_entry_exit_safe_margin
            pe_exit = df.loc[i, pe_exit_col] + self.reversal_entry_exit_safe_margin

            # === ENTRY: Price touches sig3_up → BUY PE (with MA confirmation) ===
            if position is None and price >= pe_entry and not ma_direction:
                strikePrices = self.getNiftyStrikePriceGivenCurrentValue(price, 'PE')
                for strike in strikePrices:
                    instrument = f"{strike}PE"
                    self.createReverseEntry(str(_time), instrument, price, 'BUY', self.reverseLotCount)
                    self.open_positions_price['PE'].append((instrument, price))
                self.addLogDataInfo(f"[ENTRY] BUY PE at {_time} | Price: {price} ≥ {pe_entry_col}: {pe_entry} | MA Down: {not ma_direction}")
                position = 'PE'
                entry_price = price

            # === ENTRY: Price touches sig3_down → BUY CE (with MA confirmation) ===
            if position is None and price <= ce_entry and ma_direction:
                strikePrices = self.getNiftyStrikePriceGivenCurrentValue(price, 'CE')
                for strike in strikePrices:
                    instrument = f"{strike}CE"
                    self.createReverseEntry(str(_time), instrument, price, 'BUY', self.reverseLotCount)
                    self.open_positions_price['CE'].append((instrument, price))
                self.addLogDataInfo(f"[ENTRY] BUY CE at {_time} | Price: {price} ≤ {ce_entry_col}: {ce_entry} | MA Up: {ma_direction}")
                position = 'CE'
                entry_price = price

            # === EXIT PE ===
            if position == 'PE' and price <= pe_exit:
                for instr, buy_price in self.open_positions_price['PE']:
                    pl = round((buy_price - price) * self.reverseLotCount, 2)
                    self.createReverseEntry(str(_time), instr, price, 'SELL', self.reverseLotCount)
                    self.total_reverse_pl += pl
                    self.addLogDataInfo(f"[EXIT] PE SELL at {_time} | Price: {price} ≤ {pe_exit_col}: {pe_exit:.2f} | P/L: ₹{pl}")
                self.open_positions_price['PE'].clear()
                position = None
                entry_price = None

            # === EXIT CE ===
            if position == 'CE' and price >= ce_exit:
                for instr, buy_price in self.open_positions_price['CE']:
                    pl = round((price - buy_price) * self.reverseLotCount, 2)
                    self.createReverseEntry(str(_time), instr, price, 'SELL', self.reverseLotCount)
                    self.total_reverse_pl += pl
                    self.addLogDataInfo(f"[EXIT] CE SELL at {_time} | Price: {price} ≥ {ce_exit_col}: {ce_exit:.2f} | P/L: ₹{pl}")
                self.open_positions_price['CE'].clear()
                position = None
                entry_price = None

    def fetch_all_data(self, elapsed: int):
        now_ist = datetime.now(self.ist)
        start_time = now_ist.replace(hour=9, minute=15, second=0) - timedelta(days=elapsed)
        end_time = now_ist.replace(hour=15, minute=31, second=0) - timedelta(days=elapsed)

        interval = self.tick_interval
        df = self.process_symbol_data('NSE:NIFTY', interval, start_time, end_time)

        #### start from here

        # Calculate Moving Averages
        df['ma1'] = self.calculate_moving_average(df['close'], self.ma_length, self.ma_type)
        
        if self.use_second_ma:
            df['ma2'] = self.calculate_moving_average(df['close'], self.ma_length2, self.ma_type2)
            df['ma_cross'] = self.detect_ma_cross(df['ma1'], df['ma2'])
        
        df['ma_direction'] = self.calculate_ma_direction(df['ma1'], self.color_smoothing)

        # Existing calculations
        df['source'] = (df['close'] + df['high'] + df['low'] + df['close']) / 4
        df['twap'] = df['source'].expanding().mean()
        df['std'] = df['source'].expanding().std(ddof=0)
        df['sig1_up'] = (df['twap'] + df['std']).round(2)
        df['sig1_down'] = (df['twap'] - df['std']).round(2)
        df['sig2_up'] = (df['twap'] + 2 * df['std']).round(2)
        df['sig2_down'] = (df['twap'] - 2 * df['std']).round(2)
        df['sig3_up'] = (df['twap'] + 3 * df['std']).round(2)
        df['sig3_down'] = (df['twap'] - 3 * df['std']).round(2)
        df['sig3_up_diff'] = (df['sig3_up'] - df['close']).round(2)
        df['sig3_down_diff'] = (df['close'] - df['sig3_down']).round(2)

        trend_df = pd.DataFrame()
        reverse_df = pd.DataFrame()

        if self.enable_trend_following_orders:
            start_idx = df[df['std'] >= self.trend_follow_sigma_start_threshold].first_valid_index()
            if start_idx is not None:
                trend_df = df.loc[start_idx:].reset_index(drop=True)

        if self.enable_trend_reversal_orders:
            start_idx = df[df['std'] >= self.trend_reverse_sigma_start_threshold].first_valid_index()
            if start_idx is not None:
                reverse_df = df.loc[start_idx:].reset_index(drop=True)

        if not df.empty:
            if self.enable_trend_following_orders:
                self.addLogDataInfo("")
                self.addLogDataInfo("Handling trend following orders...")   
                self.handle_trend_following_orders(trend_df)
                self.addLogDataInfo(f"Total Trend Following P/L: ₹{round(self.total_trend_pl, 2)}")
                self.total_pl = self.total_pl + self.total_trend_pl

            if self.enable_trend_reversal_orders:
                self.addLogDataInfo("")
                self.addLogDataInfo("Handling trend reversal orders...")   
                self.handle_trend_reversal_orders(reverse_df)
                self.addLogDataInfo(f"Total Trend Reversal P/L: ₹{round(self.total_reverse_pl, 2)}")
                self.total_pl = self.total_pl + self.total_reverse_pl

        self.addLogDataInfo("")
        self.addLogDataInfo(f"Final close price: {df.iloc[-1]['close']}")
        self.addLogDataInfo(f"MA1 ({self.ma_type}): {df.iloc[-1]['ma1']:.2f}")
        if self.use_second_ma:
            self.addLogDataInfo(f"MA2 ({self.ma_type2}): {df.iloc[-1]['ma2']:.2f}")
        self.addLogDataInfo(f"Total Realized P/L: ₹{round(self.total_pl, 2)}")
        self.addLogDataInfo("\n")

if __name__ == "__main__":
    client_details = get_client_details()
    elapsed = float(sys.argv[1]) if len(sys.argv) > 1 else 0
    createEntries = len(sys.argv) > 2
    stock_fetcher = StockDataFetcher(client_details, createEntries)
    stock_fetcher.login()
    stock_fetcher.fetch_all_data(elapsed)