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

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 = 2

        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

        if datetime.now(self.ist).weekday() == 0:
            elapsedStart = elapsed + 3

        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
        niftydf = self.process_symbol_data('NSE:Nifty 50', interval, start_time, end_time)
        df = niftydf[['time', 'intc']].rename(columns={'intc': 'close'}).copy()
        df['time'] = pd.to_datetime(df['time'], format='%d-%m-%Y %H:%M:%S')
        df = df.sort_values('time').reset_index(drop=True)

        smoothing_period = 35
        double_smoothing_period = 20
        signal_period = 10
        overbought = 0.74
        oversold = -0.74

        df['ROC_1'] = df['close'].pct_change() * 100
        df['Smoothed_ROC'] = df['ROC_1'].ewm(span=smoothing_period, adjust=False).mean()
        df['PMO'] = df['Smoothed_ROC'].ewm(span=double_smoothing_period, adjust=False).mean()
        df['PMO_Signal'] = df['PMO'].ewm(span=signal_period, adjust=False).mean()

        target_date = pd.to_datetime(end_time).date()
        df['PMO'] = df['PMO'] * 100
        df['PMO_Signal'] = df['PMO_Signal'] * 100
        df = df[df['time'].dt.date == target_date]
        df = df[df['time'].dt.time >= time(9, 15)]
        df = df.reset_index(drop=True)
        

        signals = []
        signal_types = []
        state = 'neutral'
        final_closePrice = 0.0

        for i in range(1, len(df)):
            #prev_pmo = df.loc[i-1, 'PMO']
            curr_pmo = df.loc[i, 'PMO']
            curr_signal = df.loc[i, 'PMO_Signal']
            final_closePrice = df.loc[i, 'close']

            if state == 'neutral':
                if curr_pmo >= overbought and curr_signal >= overbought and curr_pmo < curr_signal:    #entry for overbought cutover
                #if prev_pmo < overbought and curr_pmo >= overbought and curr_pmo > curr_signal:         #entry for initial overbought cross
                    signals.append(i)
                    signal_types.append('PE')
                    state = 'waiting_ob_exit'
                elif curr_pmo <= oversold and curr_signal <= oversold and curr_pmo > curr_signal:      #entry for oversold cutover
                #elif prev_pmo > oversold and curr_pmo <= oversold and curr_pmo < curr_signal:           #entry for initial oversold cross
                    signals.append(i)
                    signal_types.append('CE')
                    state = 'waiting_os_exit'
            elif state == 'waiting_ob_exit' and curr_pmo < overbought:
                state = 'neutral'
            elif state == 'waiting_os_exit' and curr_pmo > oversold:
                state = 'neutral'

        df_signals = df.loc[signals].reset_index(drop=True)
        df_signals['signal'] = signal_types

        self.addLogDataInfo("Crossover 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'

            for open_instr in open_positions[opposite_type]:
                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 if opposite_type == 'CE'
                               else (buy_price - closePrice) * self.lotCount, 2)

                    self.addLogDataInfo(f"closeEntry  <= time: {timeStr}, instrument: {instrument_key}, 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].clear()

            strikePrices = self.getNiftyStrikePriceGivenCurrentValue(closePrice, signal_type)
            for strikePrice in strikePrices:
                instrument = str(strikePrice) + signal_type
                action = 'BUY'
                self.addLogDataInfo(f"createEntry => time: {timeStr}, instrument: {instrument}, Price: {closePrice}, action: {action}")
                self.createEntry(timeStr, instrument, closePrice, action, self.lotCount)

                entry = instrument + '-' + str(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)
