import numpy as np
import pandas as pd
import time

# Import the specific device classes from stlab
from stlab.devices.RS_ZND import RS_ZND
from stlab.devices.Vaunix_Att import Vaunix_Att

# Import your custom DataFolder class
from dexplore.data_folder import DataFolder 

def main():
    print("Starting...")
    # 1. Initialize DataFolder
    base_data_directory = "./measurement_data" 
    script_name = "./third_att_sweep_test.py"
    dfol = DataFolder(base_data_directory, script_name)

    # 2. Connect to the Instruments using their specific classes
    # Replace with your actual VISA/USB addresses
    vna = RS_ZND("TCPIP0::172.19.20.21::inst0::INSTR", reset=True, verb=False)
    att_small = Vaunix_Att(23671)
    att_small.SetAttenuation(630)
    attenuator = Vaunix_Att(19417)

    # Delete any default traces/windows hanging around
    vna.write("CALC:PAR:DEL:ALL")

    # Define the S21 measurement and name the trace 'Trc1'
    vna.write("CALC:PAR:SDEF 'Trc1', 'S21'")

    # Feed 'Trc1' to the display on Window 1
    vna.write("DISP:WIND1:TRAC:EFE 'Trc1'")

    # 3. Configure the VNA
    freq_start = 4e9 #Hz
    freq_stop = 8e9 #Hz
    ifbw = 1000 #Hz
    vna_power = 0 #dBm


    # meas_points = min(2000, (freq_stop - freq_start)/5e6)
    meas_points = min(5000, (freq_stop - freq_start)/1e8)
    vna.SetRange(freq_start, freq_stop)
    vna.SetPoints(meas_points)
    vna.SetIFBW(ifbw)
    vna.SetPower(vna_power)

    # attenuator.SetAttenuation(100)
    start=0.0
    end = 0.0
    # time.sleep(0.1)
    # start_time = time.time()
    # output = vna.MeasureScreen_pd()
    # end_time = time.time()
    # print(f"Time taken: {(end_time - start_time):.2f} s")
    # print(output.keys())

    # # # 4. Define our sweep parameters
    att_start = 0
    att_stop = 630
    att_density = 1/10
    att_points = int((att_stop - att_start) * att_density + 1)
    attenuations = np.linspace(att_start, att_stop, att_points)  
    dataset_name = "20260703_VNA_sweep_1_8Ghz_Vaunix_Att"
    
    loop_name_1 = "Attenuation (dB)"
    loop_values_1 = attenuations/10


    metadata = {
    "Description": "attenuation sweep",
    "frequency": f"{freq_start} to {freq_stop} Hz",
    "vna_power": vna_power,
    "ifbw": ifbw,
    "att_start": att_start,
    "att_stop": att_stop,
    "att_points": att_points,
    }

    print(f"Starting 2D Sweep. Data will be saved to: {dfol.folder_full_path}")

    # 5. Start the nested loop
    for i, atten in enumerate(attenuations):
        print(f"Measuring step {i}/{len(attenuations)}: Attenuation = {atten/10} dB")
        
        # --- Set the Attenuator ---
        attenuator.SetAttenuation(atten) 
        time.sleep(0.1) # Settle time
        
        # --- Trigger VNA Measurement & Fetch Data ---
        # Using the dedicated class, fetching data is usually a single method call
        # that returns frequencies and complex data directly
        if(i==0):
            start= time.time()
        output = vna.MeasureScreen_pd(N_averages=1) # Slightly modified this one to be able to take averages
        vna.write("DISP:WIND1:TRAC:Y:AUTO ONCE") # Auto-scale the VNA display (not required for measurement but for my amusement)
        if(i==0):
            end = time.time()
            print(f"Estimated time: {((end-start)/60 * len(attenuations)):2f} minutes")
        # --- Save to DataFolder ---    
        if not dataset_name in dfol.datasets.keys():
            data_names = output.keys()[1:]       
            sweep_name = output.keys()[0]        
            sweep_values = output[sweep_name]
            
            dfol.create_stlab_dataset(
                dataset_name=dataset_name, 
                data_names=data_names, 
                trace_name=sweep_name, 
                trace_values=sweep_values,
                loop1_name=loop_name_1, 
                loop1_values=loop_values_1
            )  
            
        dfol.datasets[dataset_name].attrs.update(metadata)
        dfol.add_stlab_trace(dataset_name, output, loop1_index=i)


    dfol.save_data()
    print("Measurement loop complete!")


    attenuator.close()
    vna.close()

if __name__ == "__main__":
    main()