import pandas as pd
import os
import numpy as np
from datetime import datetime

class CSVConverter:
    def __init__(self, input_filename="nifty_options_data_31OCT.csv", output_filename="nifty_consolidated_data.csv"):
        self.input_filename = input_filename
        self.output_filename = output_filename
        
    def read_original_csv(self):
        """Read the original CSV file with proper data type handling"""
        if not os.path.isfile(self.input_filename):
            print(f"Input file {self.input_filename} not found!")
            return None
        
        try:
            # Define data types for numeric columns to ensure proper reading
            dtype_spec = {
                'greeks_delta': 'float64',
                'greeks_theta': 'float64', 
                'greeks_gamma': 'float64',
                'greeks_vega': 'float64',
                'implied_volatility': 'float64',
                'last_price': 'float64',
                'oi': 'float64',
                'previous_close_price': 'float64',
                'volume': 'float64'
            }
            
            df = pd.read_csv(self.input_filename, dtype=dtype_spec)
            print(f"Successfully read {len(df)} records from {self.input_filename}")
            print(f"Data types:\n{df.dtypes}")
            return df
        except Exception as e:
            print(f"Error reading CSV file: {e}")
            return None
    
    def process_data(self, df):
        """Process the data and convert to consolidated format with proper data type handling"""
        if df is None or df.empty:
            print("No data to process")
            return []
        
        # Print sample data to debug
        print("\nSample of original data:")
        print(df[['timestamp', 'strike', 'optionType', 'greeks_delta', 'last_price']].head(10))
        
        # Group by timestamp to create one row per timestamp
        consolidated_rows = []
        
        for timestamp, group in df.groupby('timestamp'):
            print(f"\nProcessing timestamp: {timestamp}")
            
            # Get nifty_ltp (you'll need to modify this based on how it's stored)
            nifty_ltp = self.get_nifty_ltp_from_source(timestamp)
            
            row_data = {
                'timestamp': timestamp,
                'nifty_ltp': nifty_ltp
            }
            
            # Process each strike in this timestamp group
            strikes_processed = 0
            strikes_with_data = 0
            
            for strike in group['strike'].unique():
                strike_str = str(strike)
                signal_column_name = f"strike_{strike_str}"
                ce_price_column_name = f"ce_{strike_str}_ltp"
                pe_price_column_name = f"pe_{strike_str}_ltp"
                
                # Get CE and PE data for this strike
                strike_data = group[group['strike'] == strike]
                
                ce_row = strike_data[strike_data['optionType'] == 'CE']
                pe_row = strike_data[strike_data['optionType'] == 'PE']
                
                ce_delta = None
                pe_delta = None
                ce_price = None
                pe_price = None
                
                if not ce_row.empty:
                    ce_delta = ce_row['greeks_delta'].iloc[0]
                    ce_price = ce_row['last_price'].iloc[0]
                    # Handle NaN values
                    if pd.isna(ce_delta):
                        ce_delta = None
                        print(f"  Strike {strike_str} CE: delta is NaN")
                    if pd.isna(ce_price):
                        ce_price = None
                        print(f"  Strike {strike_str} CE: price is NaN")
                
                if not pe_row.empty:
                    pe_delta = pe_row['greeks_delta'].iloc[0]
                    pe_price = pe_row['last_price'].iloc[0]
                    # Handle NaN values
                    if pd.isna(pe_delta):
                        pe_delta = None
                        print(f"  Strike {strike_str} PE: delta is NaN")
                    if pd.isna(pe_price):
                        pe_price = None
                        print(f"  Strike {strike_str} PE: price is NaN")
                
                # Debug print for each strike
                print(f"  Strike {strike_str}: CE delta={ce_delta}, PE delta={pe_delta}, CE price={ce_price}, PE price={pe_price}")
                
                # Calculate signal with proper type checking
                if ce_delta is None or pe_delta is None:
                    signal = 'NA'
                    print(f"    -> Signal: NA (missing data)")
                elif ce_delta == 0 or pe_delta == 0:
                    signal = 'NA'
                    print(f"    -> Signal: NA (zero delta)")
                else:
                    # Ensure we're comparing numbers
                    try:
                        ce_delta_float = float(ce_delta)
                        pe_delta_float = float(pe_delta)
                        signal = 1 if ce_delta_float > (pe_delta_float * -1) else 0
                        strikes_with_data += 1
                        print(f"    -> Signal: {signal} (CE: {ce_delta_float}, PE: {pe_delta_float})")
                    except (ValueError, TypeError) as e:
                        signal = 'NA'
                        print(f"    -> Signal: NA (conversion error: {e})")
                
                # Add signal column
                row_data[signal_column_name] = signal
                
                # Add CE price column
                row_data[ce_price_column_name] = ce_price
                
                # Add PE price column
                row_data[pe_price_column_name] = pe_price
                
                strikes_processed += 1
            
            print(f"  Processed {strikes_processed} strikes, {strikes_with_data} with valid data")
            consolidated_rows.append(row_data)
        
        return consolidated_rows
    
    def save_consolidated_data(self, consolidated_rows):
        """Save the consolidated data to CSV"""
        if not consolidated_rows:
            print("No consolidated data to save")
            return False
        
        try:
            df_consolidated = pd.DataFrame(consolidated_rows)
            
            # Print the consolidated data for verification
            print("\nConsolidated data sample:")
            print(df_consolidated.head())
            
            # Print column names to verify new structure
            print(f"\nColumns in consolidated data ({len(df_consolidated.columns)} total):")
            for i, col in enumerate(df_consolidated.columns):
                print(f"  {i+1:2d}. {col}")
            
            # Check if file exists to determine whether to write header
            file_exists = os.path.isfile(self.output_filename)
            
            # Append to CSV (create new file if doesn't exist)
            df_consolidated.to_csv(self.output_filename, mode='a', header=not file_exists, index=False)
            print(f"Consolidated data saved to {self.output_filename} ({len(consolidated_rows)} records)")
            return True
            
        except Exception as e:
            print(f"Error saving consolidated data: {e}")
            return False
    
    def get_nifty_ltp_from_source(self, timestamp):
        """
        You'll need to implement this method based on how you store nifty_ltp
        This could read from another CSV, database, or API
        """
        # Placeholder - return a dummy value for testing
        # Replace this with your actual implementation
        return 25800.0  # Dummy value for testing
    
    def debug_data_quality(self, df):
        """Debug function to check data quality"""
        print("\n=== DATA QUALITY CHECK ===")
        print(f"Total records: {len(df)}")
        print(f"Unique timestamps: {df['timestamp'].nunique()}")
        print(f"Unique strikes: {df['strike'].nunique()}")
        
        # Check for missing values in greeks_delta
        missing_delta = df['greeks_delta'].isna().sum()
        print(f"Missing greeks_delta values: {missing_delta}")
        
        # Check for missing values in last_price
        missing_price = df['last_price'].isna().sum()
        print(f"Missing last_price values: {missing_price}")
        
        # Check data types and sample values
        print(f"\nSample greeks_delta values:")
        print(df['greeks_delta'].head(10))
        
        print(f"\nSample last_price values:")
        print(df['last_price'].head(10))
        
        # Check option type distribution
        print(f"\nOption type distribution:")
        print(df['optionType'].value_counts())
        
        # Check if we have both CE and PE for strikes
        strike_counts = df.groupby('strike')['optionType'].nunique()
        incomplete_strikes = strike_counts[strike_counts < 2]
        print(f"\nStrikes with incomplete data (missing CE or PE): {len(incomplete_strikes)}")
    
    def convert(self):
        """Main method to convert the CSV file"""
        print(f"Converting {self.input_filename} to {self.output_filename}...")
        
        # Read original CSV
        df = self.read_original_csv()
        if df is None:
            return
        
        # Run data quality check
        self.debug_data_quality(df)
        
        # Process data
        consolidated_rows = self.process_data(df)
        
        # Save consolidated data
        self.save_consolidated_data(consolidated_rows)

def main():
    """Main function to run the conversion"""
    converter = CSVConverter()
    converter.convert()

if __name__ == "__main__":
    main()