import requests
import logging
import pytz
import sys
import json
import time
import pandas as pd
import numpy as np
from datetime import datetime, timedelta
from thefirstock import thefirstock


client_user_details_url = 'http://143.244.141.41/php/getUserDetails.php'

class dataFetcher:
    def __init__(self, type):
        self.logger = self.setup_logger()
        self.client_details = self.getClientDetails()
        self.user_id = self.client_details[0]
        self.ist = pytz.timezone('Asia/Kolkata')

        self.debug_data = ''
        self.tick_interval = 1
        self.step_value = 50 
        self.no_of_steps = 20
        self.sleep_duration_bw_calls = 1

        if type == 'Weekly':
            self.start_time = datetime(2025, 6, 26, 9, 15, 0, tzinfo=self.ist)
            self.end_time   = datetime(2025, 7, 3, 15, 31, 0, tzinfo=self.ist)
            self.base_strike_price = 25550
            self.expiry_date = '03JUL25'
            self.expiryType = 'Weekly'

        else:
            self.no_of_steps = 10
            self.tick_interval = 5
            self.start_time = datetime(2025, 6, 27, 9, 15, 0, tzinfo=self.ist)
            self.end_time   = datetime(2025, 7, 24, 15, 35, 0, tzinfo=self.ist)
            self.base_strike_price = 25300
            self.expiry_date = '24JUL25'
            self.expiryType = 'Monthly'
        
        

    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 = self.debug_data + text + '\n'

    def addLogDataInfo(self, text):
        self.logger.info(f'{text}')
        self.debug_data = self.debug_data + text + '\n'

    def getClientDetails(self):
        try:
            response = requests.get(client_user_details_url)
            response.raise_for_status()  # Raise exception for HTTP errors
            data = response.json()
            if not data.get('success', False):
                raise ValueError("getUserDetails endpoint returned unsuccessful response")
            user_data = data['data']
            client_details = [
                user_data['field1'],    # userId
                user_data['field2'],  # password
                user_data['field3'],      # TOTP
                user_data['field4'],    # apiKey
                user_data['field5'] # vendorCode
            ]
            return client_details
        except requests.exceptions.RequestException as e:
            self.addLogDataDebug(f"HTTP Request failed: {e}")
            return None
        except (json.JSONDecodeError, KeyError) as e:
            self.addLogDataDebug(f"Failed to parse response: {e}")
            return None

    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 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 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['time'] = pd.to_datetime(df['time'], format='%d-%m-%Y %H:%M:%S')
        df = df.sort_values(by='time', ascending=True)

        df['cost'] = (df['intc'] * df['oi']) / 10000000
        df['cost'] = df['cost'].round(2)

        return df

    def fetch_data(self):

        range_delta = self.step_value * self.no_of_steps
        final_dfs = []
        for strike in range(self.base_strike_price - range_delta, self.base_strike_price + range_delta + 1, self.step_value):
            pe_symbol = f"NFO:NIFTY{self.expiry_date}P{strike}"
            ce_symbol = f"NFO:NIFTY{self.expiry_date}C{strike}"

            # Retry fetching PE data until successful
            no_of_retries_pending = 3
            while True:
                try:
                    time.sleep(self.sleep_duration_bw_calls)
                    pe_df = self.process_symbol_data(pe_symbol, self.tick_interval, self.start_time, self.end_time)
                    break  # Exit loop on success
                except Exception as e:
                    no_of_retries_pending = no_of_retries_pending - 1
                    self.addLogDataInfo(f"Error fetching PE data for strikePrice: {strike}. Retrying... Exception: {e}")
                    time.sleep(2*self.sleep_duration_bw_calls)
                    if no_of_retries_pending == 0:
                        sys.exit()

            # Retry fetching CE data until successful
            no_of_retries_pending = 3
            while True:
                try:
                    time.sleep(self.sleep_duration_bw_calls)
                    ce_df = self.process_symbol_data(ce_symbol, self.tick_interval, self.start_time, self.end_time)
                    break  # Exit loop on success
                except Exception as e:
                    no_of_retries_pending = no_of_retries_pending - 1
                    self.addLogDataInfo(f"Error fetching CE data for strikePrice: {strike}. Retrying... Exception: {e}")
                    time.sleep(2*self.sleep_duration_bw_calls)
                    if no_of_retries_pending == 0:
                        sys.exit()


            merged_df = pd.merge(pe_df, ce_df, on='time', how='outer', suffixes=(f'_{strike}_pe', f'_{strike}_ce'))
            merged_df.fillna(0, inplace=True)
            merged_df[f'diff_{strike}'] =  merged_df[f'cost_{strike}_ce'] - merged_df[f'cost_{strike}_pe']
            result_df = merged_df[['time', f'cost_{strike}_ce', f'cost_{strike}_pe', f'diff_{strike}', f'intc_{strike}_pe', f'intc_{strike}_ce']]
            final_dfs.append(result_df)

        if final_dfs:
            combined_df = pd.concat(final_dfs, axis=1)
            # Remove duplicate time columns
            combined_df = combined_df.loc[:,~combined_df.columns.duplicated()]

            # Get list of columns that start with 'diff_'
            diff_cols = [col for col in combined_df.columns if col.startswith('diff_')]

            # Calculate sum only for diff_ columns
            combined_df = combined_df.copy()
            combined_df['sumDiff'] = combined_df[diff_cols].sum(axis=1)

            if combined_df.iloc[-1].isna().any():
                combined_df = combined_df.iloc[:-1]


            no_of_retries_pending = 3
            while True:
                try:
                    time.sleep(self.sleep_duration_bw_calls)
                    niftydf = self.process_symbol_data('NSE:Nifty 50', self.tick_interval, self.start_time, self.end_time)
                    niftydf = niftydf.rename(columns={'intc': 'niftyPrice'})
                    break  # Exit loop on success
                except Exception as e:
                    no_of_retries_pending = no_of_retries_pending - 1
                    self.addLogDataInfo(f"Error fetching data for Nifty 50. Retrying... Exception: {e}")
                    time.sleep(2*self.sleep_duration_bw_calls)
                    if no_of_retries_pending == 0:
                        sys.exit()

           
            diff_df = pd.concat([
                        combined_df[['time']],
                        combined_df[[col for col in combined_df.columns if col.startswith('diff_')]],
                        combined_df['sumDiff']
                    ], axis=1)
            
            diff_df = pd.merge(diff_df, niftydf[['time', 'niftyPrice']], on='time', how='left')

            other_cols = [col for col in diff_df.columns if col not in ['time', 'niftyPrice', 'sumDiff']]
            diff_df = diff_df[['time', 'niftyPrice', 'sumDiff'] + other_cols]

            # Step 4: Remove 'diff_' prefix from all column names (except 'time', 'niftyPrice', 'sumDiff')
            diff_df.columns = [
                col.replace('diff_', '') if col not in ['time', 'niftyPrice', 'sumDiff'] else col
                for col in diff_df.columns
            ]

            # Ensure numeric and clean data
            diff_df['niftyPrice'] = pd.to_numeric(diff_df['niftyPrice'], errors='coerce')

            # Drop rows where niftyPrice is NaN or inf
            diff_df = diff_df[np.isfinite(diff_df['niftyPrice'])]

            # Now safe to convert
            diff_df['basePrice'] = (np.round(diff_df['niftyPrice'] / 50) * 50).astype(int)


            # Step 1: Extract all price columns (excluding special columns)
            price_columns = [col for col in diff_df.columns if col not in ['time', 'niftyPrice', 'basePrice', 'sumDiff']]

            # Step 2: Convert price columns to integers for indexing
            price_column_ints = sorted([int(col) for col in price_columns])
            price_column_strs = [str(p) for p in price_column_ints]

            # Step 3: Function to calculate signCount for each row
            def compute_sign_count(row):
                base = int(row['basePrice'])
                if base not in price_column_ints:
                    return np.nan  # basePrice not in columns
                
                base_index = price_column_ints.index(base)
                
                # Get indices for ±10 around base
                start_idx = max(0, base_index - 10)
                end_idx = min(len(price_column_ints), base_index + 11)  # +11 because slicing is exclusive

                surrounding_cols = [str(price_column_ints[i]) for i in range(start_idx, end_idx)]
                values = row[surrounding_cols].values

                pos_count = np.nansum(values > 0)
                neg_count = np.nansum(values < 0)

                return pos_count - neg_count

            # Step 4: Apply function row-wise
            diff_df['signCount'] = diff_df.apply(compute_sign_count, axis=1)

            diff_df.to_csv(f'/var/www/html/optionData/{self.expiryType}_{self.expiry_date}.csv', index=False)
            #diff_df.to_csv(f'niftyFull_{self.expiry_date}.csv', index=False)
        


        

        pass

if __name__ == "__main__":


    type = float(sys.argv[1]) if len(sys.argv) > 1 else 0 

    if type == 0:
        type = 'Weekly'
    else:
        type = 'Monthly'

    data_fetcher = dataFetcher(type)
    data_fetcher.login()
    data_fetcher.fetch_data()