import os
import time
import pytz
import pandas as pd
import numpy as np
np.NaN = np.nan  # Add this line before importing pandas_ta
import pandas_ta as ta
import logging
from datetime import datetime, timedelta
from typing import List
from thefirstock import thefirstock
import requests
import json


class StockDataFetcher:
    def __init__(self, client_details: List[str]):
        self.client_details = client_details
        self.user_id = client_details[0]
        self.ist = pytz.timezone('Asia/Kolkata')
        self.logger = self.setup_logger()
        self.logger.info("**************************")
        self.total = {}

    def setup_logger(self) -> logging.Logger:
        logger = logging.getLogger(__name__)
        logger.setLevel(logging.INFO)
        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

    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 login(self):
        try:
            response = thefirstock.firstock_login(*self.client_details)
            if response.get("status") == "success":
                self.logger.info("Login successful")
            else:
                self.logger.error(f"Login failed: {response}")
        except Exception as e:
            self.logger.error(f"Login error: {e}")

    def fetch_time_price_series(
        self, exchange: str, trading_symbol: str, start_time: str, end_time: str, interval: str
    ) -> pd.DataFrame:
        
        self.logger.info(f"----------------------------------")
        self.logger.info(f"Fetching data for {trading_symbol}")
        self.logger.info(f"----------------------------------")
        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", []))
            else:
                self.logger.error(f"Fetch failed: {response}")
                return pd.DataFrame()
        except Exception as e:
            self.logger.error(f"Error fetching series: {e}")
            return pd.DataFrame()
        
    def calculate_pivot_levels(self, ohlc):
        open_price = float(ohlc['open'])
        high = float(ohlc['high'])
        low = float(ohlc['low'])
        close = float(ohlc['close'])
        
        # Calculate Pivot Point
        pivot = (high + close + low + open_price) / 4

        # Calculate Resistance Levels
        r1 = (2 * pivot) - low
        r2 = pivot + (r1 - ((2 * pivot) - high))
        r3 = high + 2 * (pivot - low)

        # Calculate Support Levels
        s1 = (2 * pivot) - high
        s2 = pivot - (r1 - s1)
        s3 = low - 2 * (high - pivot)

        pivot_data = {
                'open' :open_price,
                'high' :high,
                'low'  :low,
                'close':close,
                'PP': round(pivot),
                'R1': round(r1),
                'R2': round(r2),
                'R3': round(r3),
                'S1': round(s1),
                'S2': round(s2),
                'S3': round(s3)
            }

        return pivot_data

        

    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),
        )

        # Preprocess timestamps and ensure data integrity
        df['time'] = pd.to_datetime(df['time'], errors='coerce', dayfirst=True)
        df['time'] = df['time'].dt.tz_localize(self.ist)

        

        numeric_cols = ['intc', 'intv', 'into', 'inth', 'intl', '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 = df.iloc[0:].reset_index(drop=True)
        df['intoi'] = df['intoi'] / 75
        df['oi'] = df['oi'] / 75
        print(df[:32])
        mydf = df[['time', 'into', 'inth', 'intl', 'intc', 'intoi']]

        mydf = mydf.copy()

        first_intc = mydf['into'].iloc[0]
        mydf['priceChange'] = mydf['intc'] - first_intc

        mydf.at[0, 'intoi'] = 0
        mydf['sumOI'] = mydf['intoi'].cumsum()

        graph_data = {
        'time': mydf['time'].dt.strftime('%Y-%m-%d %H:%M:%S').tolist(),
        'priceChange': mydf['priceChange'].tolist(),
        'sumOI': mydf['sumOI'].tolist(),
        }

        pivot_data = {}

        if len(df) > 5:
            ohlc = {}
            ohlc['open'] = mydf['into'].iloc[0]
            ohlc['high'] = mydf.head(5)['inth'].max()
            ohlc['low'] = mydf.head(5)['intl'].min()
            ohlc['close'] = mydf['intc'].iloc[4]

            pivot_data = self.calculate_pivot_levels(ohlc)


        return graph_data, pivot_data


    def fetch_all_data(self):
        symbol = "NFO:NIFTY24APR25F"

        for interval in [3]:
            start_time = datetime.now(self.ist).replace(hour=9, minute=15, second=0)
            #start_time = start_time - timedelta(days=5)
            end_time = datetime.now(self.ist).replace(hour=15, minute=30, second=0)
            graph_data, pivot_data = self.process_symbol_data(symbol, interval, start_time, end_time)
            time.sleep(1)
            return graph_data, pivot_data  # Return the graph data instead of printing it


if __name__ == "__main__":
    client_details = ['RA1383', 'O9i8u7y6$$', '13061983', 'RA1383_API', '85c819b0c188c44bc714b87645aef3d6']
    
    stock_fetcher = StockDataFetcher(client_details)
    stock_fetcher.login()
    graph_data, pivot_data = stock_fetcher.fetch_all_data()
    print(json.dumps({
                "graphData": graph_data,
                "pivotData": pivot_data
            }))
