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 = 'SENSEX' + 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 getSensexStrikePriceGivenCurrentValue(self, closePrice, type):
        strike_prices = []
        try:
            price = float(closePrice)
            remainder = price % 100

            base_strike = int(price + (100 - remainder)) if remainder > 50 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 += 100
        elif base_strike > price and type == 'CE':
            base_strike -= 100

        for step in range(self.order_allowed_steps):
            if type == 'CE':
                strike_prices.append(base_strike - 100 * step)
            else:
                strike_prices.append(base_strike + 100 * step)

        return strike_prices


    def fetch_all_data(self, elapsed: int):
        elapsed = int(elapsed)
        elapsedStart = elapsed + 1

        # Handle Monday special case
        if datetime.now(self.ist).weekday() == 0:
            elapsedStart = elapsed + 3

        # Set time range
        start_time = datetime.now(self.ist).replace(hour=14, minute=0, second=0) - timedelta(days=elapsedStart)
        end_time = datetime.now(self.ist).replace(hour=15, minute=31, second=0) - timedelta(days=elapsed)

        interval = self.tick_interval
        sensexdf = self.process_symbol_data('BSE:SENSEX', interval, start_time, end_time)

        df = sensexdf[['time', 'intc', 'inth', 'intl']].rename(columns={
            'intc': 'close',
            'inth': 'high',
            'intl': 'low'
        }).copy()

        # Convert types
        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 using ta library ===
        period = 14
        adx = ADXIndicator(high=df['high'], low=df['low'], close=df['close'], window=period)

        df['+DI'] = adx.adx_pos()
        df['-DI'] = adx.adx_neg()
        df['ADX'] = adx.adx()
        df['histogram'] = df['+DI'] - df['-DI']

        df = df[df['histogram'].notna()].reset_index(drop=True)


        # Filter for current day and after 9:15 AM
        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 ===
        signals = []
        signal_types = []
        state = 'neutral'
        final_closePrice = df.iloc[-1]['close'] if len(df) > 0 else 0.0

        for i in range(len(df)):
            _time = df.loc[i, 'time']
            hist = df.loc[i, 'histogram']
            final_closePrice = df.loc[i, 'close']

            #print(_time, hist)

            if state == 'neutral':
                if hist > 35.3:
                    signals.append(i)
                    signal_types.append('PE')  # SELL signal
                    state = 'waiting_sell_exit'
                elif hist < -35.3:
                    signals.append(i)
                    signal_types.append('CE')  # BUY signal
                    state = 'waiting_buy_exit'
            elif state == 'waiting_sell_exit' and hist < 15:
                state = 'neutral'
            elif state == 'waiting_buy_exit' and hist > -15:
                state = 'neutral'

        # Create signals DataFrame
        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']
            opposite_type = 'PE' if signal_type == 'CE' else 'CE'

            # Close opposite positions
            for open_instr in open_positions[opposite_type].copy():
                instrument_key = open_instr.split('-')[0]
                action = 'SELL'

                matches = [(i, bp) for i, (iname, bp) in enumerate(self.open_positions_price[opposite_type]) 
                        if iname == instrument_key]

                for i, buy_price in matches:
                    pl = round((closePrice - buy_price) * self.lotCount, 2) if opposite_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[opposite_type] = [
                    (iname, bp) for (iname, bp) in self.open_positions_price[opposite_type]
                    if iname != instrument_key
                ]
                open_positions[opposite_type].discard(open_instr)

            # Open new positions
            strikePrices = self.getSensexStrikePriceGivenCurrentValue(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))

        # Final summary
        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)
