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

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.transactionsCounter = 0
        self.lotCount = 1

        self.tick_interval = 1
        self.order_allowed_steps = 1

        self.total_pl = 0.0

        self.debug_data = ''
        self.open_positions_price = {'CE': [], 'PE': []}  # (instrument, buyPrice)

    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 createEntry(self, tickTimeStr, instrument, closePrice, signal, lotCount):
        instrument = 'NIFTY' + str(instrument)
        self.transactionsCounter += 1

        if self.createEntries:
            url = "http://143.244.141.41/php/createMoneyHeistEntry.php"
            tickTimeStr += ':' + str(self.transactionsCounter).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 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 fetch_all_data(self, elapsed: int):
        elapsed = int(elapsed)
        elapsedStart = elapsed + 1

        now_ist = datetime.now(self.ist)
        if now_ist.weekday() == 0:
            elapsedStart = elapsed + 3

        start_time = now_ist.replace(hour=14, minute=0, second=0) - timedelta(days=elapsedStart)
        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.sort_values('time').reset_index(drop=True)
        df = df[df['time'].dt.time >= time(9, 15)]
        df = df[df['time'].dt.time <= time(15, 30)]
        df = df.reset_index(drop=True)

        # === ADX Calculation ===
        adx = ADXIndicator(high=df['high'], low=df['low'], close=df['close'], window=14)
        df['+DI'] = adx.adx_pos()
        df['-DI'] = adx.adx_neg()
        df['ADX'] = adx.adx()
        df['histogram'] = df['+DI'] - df['-DI']

        # === PMO Calculation ===
        smoothing_period = 35
        double_smoothing_period = 20
        signal_period = 10

        roc = df['close'].pct_change(periods=1) * 100
        pmo_raw = roc.ewm(span=smoothing_period, adjust=False).mean()
        df['PMO'] = pmo_raw.ewm(span=double_smoothing_period, adjust=False).mean()
        df['PMO_SIGNAL'] = df['PMO'].ewm(span=signal_period, adjust=False).mean()

        df = df[df['histogram'].notna()].reset_index(drop=True)

        target_date = pd.to_datetime(end_time).date()
        df = df[df['time'].dt.date == target_date]
        df = df[df['time'].dt.time >= time(9, 15)]
        df = df.reset_index(drop=True)

        # === Signal Generation with 5-bar histogram confirmation + PMO ===
        signals = []
        signal_types = []
        open_ce = False
        open_pe = False
        last_exit_was_ce = False
        last_exit_was_pe = False
        waiting_for_zero_cross = False
        final_closePrice = df.iloc[-1]['close'] if len(df) > 0 else 0.0

        for i in range(4, len(df)):
            _time = df.loc[i, 'time']
            final_closePrice = df.loc[i, 'close']
            hist_values = df.loc[i-4:i, 'histogram'].values
            hist = hist_values[-1]
            pmo_val = df.loc[i, 'PMO']
            pmo_signal_val = df.loc[i, 'PMO_SIGNAL']

            if waiting_for_zero_cross:
                if (last_exit_was_ce and hist < 0) or (last_exit_was_pe and hist > 0):
                    waiting_for_zero_cross = False
                    last_exit_was_ce = False
                    last_exit_was_pe = False

            if not open_ce and not open_pe and not waiting_for_zero_cross:
                if all(h > 15 for h in hist_values) and pmo_val > 0:
                    signals.append(i)
                    signal_types.append('CE')
                    open_ce = True
                elif all(h < -15 for h in hist_values) and pmo_val < 0:
                    signals.append(i)
                    signal_types.append('PE')
                    open_pe = True

            elif open_ce:
                if hist >= 25 or hist < 0:
                    print (hist < 0)
                    print(pmo_val < pmo_signal_val)
                    signals.append(i)
                    signal_types.append('CE_EXIT')
                    open_ce = False
                    waiting_for_zero_cross = True
                    last_exit_was_ce = True

            elif open_pe:
                if hist <= -25 or hist > 0:
                    print (hist > 0)
                    print(pmo_val > pmo_signal_val)
                    signals.append(i)
                    signal_types.append('PE_EXIT')
                    open_pe = False
                    waiting_for_zero_cross = True
                    last_exit_was_pe = True

        df_signals = pd.DataFrame()
        if signals:
            df_signals = df.loc[signals].reset_index(drop=True)
            df_signals['signal'] = signal_types

        self.addLogDataInfo("ADX Histogram Signals:")
        open_positions = {'CE': set(), 'PE': set()}

        for index, row in df_signals.iterrows():
            timeStr = str(row['time'])
            closePrice = row['close']
            signal_type = row['signal']

            if 'EXIT' in signal_type:
                base_type = signal_type.replace('_EXIT', '')
                for open_instr in open_positions[base_type].copy():
                    instrument_key = open_instr.split('-')[0]
                    action = 'SELL'

                    matches = [(i, bp) for i, (iname, bp) in enumerate(self.open_positions_price[base_type]) 
                            if iname == instrument_key]

                    for i, buy_price in matches:
                        pl = round((closePrice - buy_price) * self.lotCount, 2) if base_type == 'CE' else round((buy_price - closePrice) * self.lotCount, 2)
                        self.addLogDataInfo(f"closeEntry <= time: {timeStr}, instrument: {instrument_key}, "
                                            f"Price: {closePrice}, action: {action}, P/L: {pl}")
                        self.createEntry(timeStr, instrument_key, closePrice, action, self.lotCount)
                        self.total_pl += pl

                    self.open_positions_price[base_type] = [
                        (iname, bp) for (iname, bp) in self.open_positions_price[base_type]
                        if iname != instrument_key
                    ]
                    open_positions[base_type].discard(open_instr)

            else:
                strikePrices = self.getNiftyStrikePriceGivenCurrentValue(closePrice, signal_type)
                for strikePrice in strikePrices:
                    instrument = f"{strikePrice}{signal_type}"
                    action = 'BUY'
                    self.addLogDataInfo(f"createEntry => time: {timeStr}, instrument: {instrument}, "
                                        f"Price: {closePrice}, action: {action}")
                    self.createEntry(timeStr, instrument, closePrice, action, self.lotCount)

                    entry = f"{instrument}-{index}"
                    open_positions[signal_type].add(entry)
                    self.open_positions_price[signal_type].append((instrument, closePrice))

        self.addLogDataInfo("")
        self.addLogDataInfo(f"current close price: {final_closePrice}")
        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)
