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
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

        #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 (can be 'twap'/'sig1_up'/'sig2_up'/'sig3_up'/'sig1_down'/'sig2_down'/'sig3_down')
        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.addLogDataInfo(f"Attempting login for {self.client_details[0]}")
            response = thefirstock.firstock_login(*self.client_details)
            if response.get("status") == "success":
                self.addLogDataInfo("Login successful")
            else:
                self.addLogDataInfo(f"Login failed: {response}")
                sys.exit()
        except Exception as e:
            self.addLogDataInfo(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("%d/%m/%Y %H:%M:%S"),
            end_time.strftime("%d/%m/%Y %H:%M:%S"),
            str(interval),
        )

        numeric_cols = ['intc', 'intvwap', 'oi', 'intoi']
        for col in numeric_cols:
            df[col] = pd.to_numeric(df[col], errors='coerce')

        df = df.sort_values(by='time', ascending=True)
        df['cost'] = (df['intc'] * df['oi']) / 10000000
        df['cost'] = df['cost'].round(2)
        return df

    def createTrendEntry(self, tickTimeStr, instrument, closePrice, signal, lotCount):
        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):
        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}")
        try:
            response = thefirstock.firstock_TimePriceSeries(
                userId=self.user_id,
                exchange=exchange,
                tradingSymbol=trading_symbol,
                startTime=start_time,
                endTime=end_time,
                interval=interval,
            )
            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': []}  # (instrument, buyPrice)

        position = None  # None, 'CE', 'PE'
        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']

            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

            # === 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 ===
            if position is None:
                if price > twap + self.trend_following_safeMargin and (last_exit_type != 'CE' or last_twap_direction == '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}")
                    position = 'CE'
                    entry_price = price
                    last_exit_type = 'NONE'
                    last_twap_direction = None

                elif price < twap - self.trend_following_safeMargin and (last_exit_type != 'PE' or last_twap_direction == '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}")
                    position = 'PE'
                    entry_price = price
                    last_exit_type = 'NONE'
                    last_twap_direction = None

            # === EXIT CE ===
            elif position == 'CE':
                if price < twap:
                    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 TWAP Reversal at {_time} | Price: {price} < TWAP: {twap:.2f} | P/L: ₹{pl}")
                        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

                elif price - entry_price >= self.trend_following_fixed_profit or price >= sig2_up:
                    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
                        if pl < 0:
                            num_negatives = num_negatives + 1
                        if price >= sig2_up:
                            self.addLogDataInfo(f"[EXIT] CE Profit Hit at {_time} | Price: {price} ≥ sig2_up: {sig2_up} | P/L: ₹{pl}")
                        else:
                            self.addLogDataInfo(f"[EXIT] CE Profit Hit at {_time} | Price: {price} ≥ Entry + ₹{self.trend_following_fixed_profit} | P/L: ₹{pl}")
                    self.open_positions_price['CE'].clear()
                    position = None
                    entry_price = None
                    last_exit_type = 'CE'
                    last_twap_direction = None

            # === EXIT PE ===
            elif position == 'PE':
                if price > twap:
                    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
                        if pl < 0:
                            num_negatives = num_negatives + 1
                        self.addLogDataInfo(f"[EXIT] PE TWAP Reversal at {_time} | Price: {price} > TWAP: {twap:.2f} | P/L: ₹{pl}")
                    self.open_positions_price['PE'].clear()
                    position = None
                    entry_price = None
                    last_exit_type = 'PE'
                    last_twap_direction = None

                elif entry_price - price >= self.trend_following_fixed_profit or price <= sig2_down:
                    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
                        if pl < 0:
                            num_negatives = num_negatives + 1
                        if price <= sig2_down:
                            self.addLogDataInfo(f"[EXIT] PE Profit Hit at {_time} | Price: {price} <= sig2_down: {sig2_down} | P/L: ₹{pl}")
                        else:
                            self.addLogDataInfo(f"[EXIT] PE Profit Hit at {_time} | Price: {price} ≤ Entry - ₹{self.trend_following_fixed_profit} | P/L: ₹{pl}")
                    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  # Current position type: None, 'CE', 'PE'
        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']
            


            ce_entry = df.loc[i, ce_entry_col]
            pe_entry = df.loc[i, pe_entry_col]

            ce_entry = ce_entry + self.reversal_entry_exit_safe_margin
            pe_entry = pe_entry - self.reversal_entry_exit_safe_margin

            ce_exit = df.loc[i, ce_exit_col]
            pe_exit = df.loc[i, pe_exit_col]

            ce_exit = ce_exit - self.reversal_entry_exit_safe_margin
            pe_exit = pe_exit + self.reversal_entry_exit_safe_margin

            #print(f"PE diff: {abs(price - pe_entry)} | CE diff: {abs(price - ce_entry)}")

            


            # === ENTRY: Price touches sig3_up → BUY PE ===
            if position is None and price >= pe_entry:
                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}")
                position = 'PE'
                entry_price = price

            # === ENTRY: Price touches sig3_down → BUY CE ===
            if position is None and price <= ce_entry:
                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}")
                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}")
                    print(f'ce_entry:{ce_entry:.2f}')
                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
        niftydf = self.process_symbol_data('NSE:Nifty 50', interval, start_time, end_time)

        df = niftydf[['time', 'intc', 'inth', 'intl']].rename(columns={
            'intc': 'close',
            'inth': 'high',
            'intl': 'low'
        }).copy()

        df['close'] = pd.to_numeric(df['close'], errors='coerce')
        df['high'] = pd.to_numeric(df['high'], errors='coerce')
        df['low'] = pd.to_numeric(df['low'], errors='coerce')
        df['time'] = pd.to_datetime(df['time'], format='%d-%m-%Y %H:%M:%S')
        df = df[df['time'].dt.time >= time(9, 15)].sort_values('time').reset_index(drop=True)

        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)


        #df = df[df['time'].dt.time >= time(9, 35)].sort_values('time').reset_index(drop=True)

        trend_df = pd.DataFrame()
        reverse_df = pd.DataFrame()

        # === Drop leading rows until std >= self.sigma_start_threshold ===
        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)


        #print(df)

        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 Realized 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 Realized 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"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)
