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.sma_period = 259
        self.activationZonePoints = 5
        self.stopLossPoints = 10
        self.maxProfit = 75  # Fixed profit target for exit

        self.max_transactions_per_day = 4

        self.start_trail_threshold = 10
        self.fixed_trail = 5
        self.k = 3.0

        self.total_pl = 0.0

        self.debug_data = ''
        self.open_positions_price = {'CE': [], 'PE': []}  # (instrument, buyPrice, max_pl)

    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)
        now_ist = datetime.now(self.ist)
        elapsedStart = elapsed + 7 
        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)

        self.addLogDataDebug(f"Data range start_time: {start_time}, end_time: {end_time}")

        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)

        # === Simple Moving Averages ===
        df['sma'] = df['close'].ewm(span=self.sma_period, adjust=False).mean()

        # === 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']

        df['inZoneDiff'] = (df['close'] - df['sma']).round(2)
        df['inZone'] = (abs(df['close'] - df['sma']) <= self.activationZonePoints).astype(int)

        # === Filter for last day ===
        target_date = df['time'].iloc[-1].date()
        df = df[df['time'].dt.date == target_date].reset_index(drop=True)

        # === Signal Generation and Immediate Processing ===
        open_positions = {'CE': set(), 'PE': set()}
        final_closePrice = df.iloc[-1]['close'] if len(df) > 0 else 0.0
        stop_trading = False
        activation_zone_hit = False

        self.addLogDataInfo("ADX Histogram Signals:")

        for i in range(1, len(df)):


            if self.max_transactions_per_day <= 0:
                self.addLogDataInfo("Max trades for the day reached. Stopping further trades.")
                stop_trading = True

            if stop_trading:
                break

            _time = df.loc[i, 'time']
            closePrice = df.loc[i, 'close']
            final_closePrice = closePrice
            histogram = df.loc[i, 'histogram']
            inZone = df.loc[i, 'inZone']
            inZoneDiff = df.loc[i, 'inZoneDiff']
            sma = df.loc[i, 'sma']
            tick_time = _time.time()
            timeStr = str(_time)

            # Check if price is in activation zone
            if inZone == 1:
                if activation_zone_hit == False:
                    self.addLogDataInfo(f"Activation zone hit at time: {_time}, closePrice: {closePrice}, inZoneDiff: {inZoneDiff}, sma: {sma}")
                activation_zone_hit = True
                
            elif abs(closePrice - sma) > 50:
                if activation_zone_hit:
                    self.addLogDataInfo(f"Price moved out of activation zone at time: {_time}, closePrice: {closePrice}, inZoneDiff: {inZoneDiff}, sma: {sma}")
                activation_zone_hit = False
                
            # Check for 14:58 exit condition
            if tick_time >= time(14, 58):
                for base_type in ['CE', 'PE']:
                    if self.open_positions_price[base_type]:
                        self.addLogDataInfo(f"Forcing {base_type} exit at 14:58")
                        for open_instr in open_positions[base_type].copy():
                            instrument_key = open_instr.split('-')[0]
                            action = 'SELL'
                            matches = [(j, bp, mpl) for j, (iname, bp, mpl) in enumerate(self.open_positions_price[base_type]) 
                                       if iname == instrument_key]
                            for j, buy_price, max_pl 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].pop(j)
                                self.max_transactions_per_day -= 1
                            open_positions[base_type].discard(open_instr)
                        stop_trading = True
                continue

            # Check for exit conditions
            if self.open_positions_price['CE']:
                new_list = []
                for instrument, buy_price, max_pl in self.open_positions_price['CE']:
                    pl = (closePrice - buy_price) * self.lotCount
                    pl = round(pl, 2)
                    #self.addLogDataInfo(f"---- Checking CE exit: time={timeStr} instrument={instrument}, close_price={closePrice:.2f}, buy_price={buy_price:.2f}, pl={pl}, stop_level={stop_level:.2f}")
                    updated_max_pl = max(max_pl, pl)
                    exit_position = False
                    if pl >= self.maxProfit or (inZone == 0 and histogram < 0 and inZoneDiff < 0) or pl <= -self.stopLossPoints:
                        exit_position = True
                    elif updated_max_pl > self.start_trail_threshold:
                        trail = self.fixed_trail + self.k * math.sqrt(updated_max_pl)
                        stop_level = updated_max_pl - trail
                        if pl < stop_level:
                            exit_position = True
                            self.addLogDataInfo(f"SRATS exit triggered: pl={pl}, stop_level={stop_level}")

                    if not exit_position:
                        new_list.append((instrument, buy_price, updated_max_pl))
                    else:
                        action = 'SELL'
                        self.addLogDataInfo(f"closeEntry <= time: {timeStr}, instrument: {instrument}, "
                                            f"Price: {closePrice}, action: {action}, P/L: {pl}")
                        self.createEntry(timeStr, instrument, closePrice, action, self.lotCount)
                        self.total_pl += pl
                        open_positions['CE'].discard(f"{instrument}-{i}")
                        self.max_transactions_per_day -= 1
                self.open_positions_price['CE'] = new_list

            if self.open_positions_price['PE']:
                new_list = []
                for instrument, buy_price, max_pl in self.open_positions_price['PE']:
                    pl = (buy_price - closePrice) * self.lotCount
                    pl = round(pl, 2)
                    #self.addLogDataInfo(f"---- Checking PE exit: time={timeStr} instrument={instrument}, close_price={closePrice:.2f}, buy_price={buy_price:.2f}, pl={pl}, stop_level={stop_level:.2f}")
                    updated_max_pl = max(max_pl, pl)
                    exit_position = False
                    if pl >= self.maxProfit or (inZone == 0 and histogram > 0 and inZoneDiff > 0) or pl <= -self.stopLossPoints:
                        exit_position = True
                    elif updated_max_pl > self.start_trail_threshold:
                        trail = self.fixed_trail + self.k * math.sqrt(updated_max_pl)
                        stop_level = updated_max_pl - trail
                        if pl < stop_level:
                            exit_position = True
                            self.addLogDataInfo(f"SRATS exit triggered: pl={pl}, stop_level={stop_level}")

                    if not exit_position:
                        new_list.append((instrument, buy_price, updated_max_pl))
                    else:
                        action = 'SELL'
                        self.addLogDataInfo(f"closeEntry <= time: {timeStr}, instrument: {instrument}, "
                                            f"Price: {closePrice}, action: {action}, P/L: {pl}")
                        self.createEntry(timeStr, instrument, closePrice, action, self.lotCount)
                        self.total_pl += pl
                        open_positions['PE'].discard(f"{instrument}-{i}")
                        self.max_transactions_per_day -= 1
                self.open_positions_price['PE'] = new_list

            # Check for entry conditions
            if not self.open_positions_price['CE'] and not self.open_positions_price['PE'] and activation_zone_hit:
                if histogram > 0 and inZoneDiff > 0:
                    self.addLogDataInfo(f"Generated CE signal at time: {_time}, closePrice: {closePrice}, histogram: {histogram}, inZoneDiff: {inZoneDiff}, sma: {sma}")
                    strikePrices = self.getNiftyStrikePriceGivenCurrentValue(closePrice, 'CE')
                    for strikePrice in strikePrices:
                        instrument = f"{strikePrice}CE"
                        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}-{i}"
                        open_positions['CE'].add(entry)
                        self.open_positions_price['CE'].append((instrument, closePrice, 0))
                        stop_level = 0
                        self.max_transactions_per_day -= 1
                elif histogram < 0 and inZoneDiff < 0:
                    self.addLogDataInfo(f"Generated PE signal at time: {_time}, closePrice: {closePrice}, histogram: {histogram}, inZoneDiff: {inZoneDiff}, sma: {sma}")
                    strikePrices = self.getNiftyStrikePriceGivenCurrentValue(closePrice, 'PE')
                    for strikePrice in strikePrices:
                        instrument = f"{strikePrice}PE"
                        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}-{i}"
                        open_positions['PE'].add(entry)
                        self.open_positions_price['PE'].append((instrument, closePrice, 0))
                        stop_level = 0
                        self.max_transactions_per_day -= 1

        self.addLogDataInfo("")
        self.addLogDataInfo(f"current close price: {final_closePrice} @ {timeStr}")
        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)