import tkinter as tk
from tkinter import ttk, messagebox, filedialog
import visa
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from matplotlib.backends.backend_tkagg import FigureCanvasTkAgg

class GPIBApp:
    def __init__(self, root):
        self.root = root
        self.root.title("GPIB Instrument Control")
        self.setup_gpib()
        self.setup_ui()
        self.initialize_variables()

    def setup_gpib(self):
        self.rm = visa.ResourceManager()
        self.instrument = self.rm.open_resource('GPIB0::8::INSTR')
        self.instrument.timeout = 5000  # 5 seconds timeout
        self.instrument.write("*RST")  # Reset the instrument

    def setup_ui(self):
        # Left Column
        left_frame = ttk.Frame(self.root)
        left_frame.grid(row=0, column=0, padx=10, pady=10)

        # Coupling
        ttk.Label(left_frame, text="Coupling").grid(row=0, column=0)
        self.coupling = tk.StringVar(value="AC<55")
        ttk.Combobox(left_frame, textvariable=self.coupling, values=["AC<55", "AC+DC", "AC>55"]).grid(row=0, column=1)

        # Time/Cycle Settings
        self.time_cycle = tk.StringVar(value="Time")
        ttk.Radiobutton(left_frame, text="Time", variable=self.time_cycle, value="Time").grid(row=1, column=0)
        ttk.Radiobutton(left_frame, text="Cycles", variable=self.time_cycle, value="Cycles").grid(row=1, column=1)
        self.speed = tk.StringVar(value="FAST")
        self.cycles = tk.IntVar(value=0)
        ttk.Combobox(left_frame, textvariable=self.speed, values=["FAST", "MEDIUM", "SLOW", "VERY SLOW"]).grid(row=2, column=0)
        ttk.Entry(left_frame, textvariable=self.cycles).grid(row=2, column=1)

        # Delay Time
        ttk.Label(left_frame, text="Delay Time").grid(row=3, column=0)
        self.delay_time = tk.DoubleVar(value=0.0)
        ttk.Entry(left_frame, textvariable=self.delay_time).grid(row=3, column=1)

        # Output Waveform
        ttk.Label(left_frame, text="Output Waveform").grid(row=4, column=0)
        self.waveform = tk.StringVar(value="SINE")
        ttk.Combobox(left_frame, textvariable=self.waveform, values=["SINE", "TRIANGLE", "SQUARE"]).grid(row=4, column=1)

        # Amplitude
        ttk.Label(left_frame, text="Amplitude (V)").grid(row=5, column=0)
        self.amplitude = tk.DoubleVar(value=1.0)
        ttk.Entry(left_frame, textvariable=self.amplitude).grid(row=5, column=1)

        # DC Offset
        self.dc_offset_enabled = tk.BooleanVar(value=False)
        ttk.Checkbutton(left_frame, text="DC Offset", variable=self.dc_offset_enabled).grid(row=6, column=0)
        self.dc_offset = tk.DoubleVar(value=0.0)
        ttk.Entry(left_frame, textvariable=self.dc_offset).grid(row=6, column=1)

        # Sweep Frequency
        ttk.Label(left_frame, text="Sweep Start Freq (Hz)").grid(row=7, column=0)
        self.sweep_start = tk.DoubleVar(value=10.0)
        ttk.Entry(left_frame, textvariable=self.sweep_start).grid(row=7, column=1)
        ttk.Label(left_frame, text="Sweep Stop Freq (Hz)").grid(row=8, column=0)
        self.sweep_stop = tk.DoubleVar(value=2200000.0)
        ttk.Entry(left_frame, textvariable=self.sweep_stop).grid(row=8, column=1)

        # Sweep Type
        self.sweep_type = tk.StringVar(value="logarithmic")
        ttk.Radiobutton(left_frame, text="Logarithmic", variable=self.sweep_type, value="logarithmic").grid(row=9, column=0)
        ttk.Radiobutton(left_frame, text="Linear", variable=self.sweep_type, value="linear").grid(row=9, column=1)
        self.steps = tk.IntVar(value=10)
        ttk.Entry(left_frame, textvariable=self.steps).grid(row=10, column=0)

        # Ratio
        ttk.Label(left_frame, text="Ratio").grid(row=11, column=0)
        self.ratio = tk.StringVar(value="CH2/CH1")
        ttk.Combobox(left_frame, textvariable=self.ratio, values=["CH2/CH1", "CH1/CH2", "CH1/Output", "CH2/Output"]).grid(row=11, column=1)

        # Loop
        self.loop = tk.BooleanVar(value=False)
        ttk.Checkbutton(left_frame, text="Loop", variable=self.loop).grid(row=12, column=0)

        # Buttons
        ttk.Button(left_frame, text="AT START FREQ", command=self.at_start_freq).grid(row=13, column=0)
        ttk.Button(left_frame, text="START SWEEP", command=self.start_sweep).grid(row=13, column=1)
        ttk.Button(left_frame, text="STOP", command=self.stop).grid(row=14, column=0)

        # Right Column
        right_frame = ttk.Frame(self.root)
        right_frame.grid(row=0, column=1, padx=10, pady=10)

        # Save to File
        ttk.Button(right_frame, text="Save to File", command=self.save_to_file).grid(row=0, column=0)
        self.autosave = tk.BooleanVar(value=False)
        ttk.Checkbutton(right_frame, text="Autosave", variable=self.autosave).grid(row=0, column=1)

        # Calibrate
        ttk.Button(right_frame, text="Calibrate", command=self.calibrate).grid(row=1, column=0)
        self.compensate = tk.BooleanVar(value=False)
        ttk.Checkbutton(right_frame, text="Compensate", variable=self.compensate).grid(row=1, column=1)

        # Labels
        ttk.Label(right_frame, text="Frequency [Hz]").grid(row=2, column=0)
        self.freq_label = ttk.Label(right_frame, text="0.0")
        self.freq_label.grid(row=2, column=1)
        ttk.Label(right_frame, text="CH1 [Vrms]").grid(row=3, column=0)
        self.ch1_label = ttk.Label(right_frame, text="0.0")
        self.ch1_label.grid(row=3, column=1)
        ttk.Label(right_frame, text="CH2 [Vrms]").grid(row=4, column=0)
        self.ch2_label = ttk.Label(right_frame, text="0.0")
        self.ch2_label.grid(row=4, column=1)
        ttk.Label(right_frame, text="Gain [dB]").grid(row=5, column=0)
        self.gain_label = ttk.Label(right_frame, text="0.0")
        self.gain_label.grid(row=5, column=1)
        ttk.Label(right_frame, text="Phase [deg]").grid(row=6, column=0)
        self.phase_label = ttk.Label(right_frame, text="0.0")
        self.phase_label.grid(row=6, column=1)

        # Graph
        self.figure, self.ax = plt.subplots()
        self.canvas = FigureCanvasTkAgg(self.figure, master=right_frame)
        self.canvas.get_tk_widget().grid(row=7, column=0, columnspan=2)

        # Slider
        self.slider = ttk.Scale(right_frame, from_=0, to=100, orient="horizontal")
        self.slider.grid(row=8, column=0, columnspan=2)
        self.slider.bind("<Motion>", self.update_marker)

    def initialize_variables(self):
        self.sweep_data = []
        self.calibration_data = []
        self.is_sweeping = False

    def at_start_freq(self):
        self.set_output_parameters()
        self.instrument.write("OUTPUT,ON")
        self.instrument.write(f"FREQUE,{self.sweep_start.get()}")
        self.read_data()
        self.instrument.write("OUTPUT,OFF")

    def start_sweep(self):
        self.is_sweeping = True
        self.set_output_parameters()
        self.instrument.write("OUTPUT,ON")
        frequencies = self.calculate_frequencies()
        for freq in frequencies:
            if not self.is_sweeping:
                break
            self.instrument.write(f"FREQUE,{freq}")
            self.read_data()
        self.instrument.write("OUTPUT,OFF")
        if self.autosave.get():
            self.save_to_file()

    def stop(self):
        self.is_sweeping = False

    def calibrate(self):
        self.compensate.set(False)
        self.start_sweep()
        self.calibration_data = self.sweep_data.copy()
        pd.DataFrame(self.calibration_data).to_csv("Calibration.csv", index=False)

    def save_to_file(self):
        file_path = filedialog.asksaveasfilename(defaultextension=".csv")
        if file_path:
            pd.DataFrame(self.sweep_data).to_csv(file_path, index=False)

    def read_data(self):
        raw_data = self.instrument.query("GAINPH?")
        data = list(map(float, raw_data.split(',')))
        freq, ch1, ch2, phase_ch1, phase_ch2 = data
        if self.compensate.get():
            ch1 -= self.calibration_data[0][1]
            ch2 -= self.calibration_data[0][2]
            phase_ch1 -= self.calibration_data[0][3]
            phase_ch2 -= self.calibration_data[0][4]
        self.sweep_data.append([freq, ch1, ch2, phase_ch1, phase_ch2])
        self.update_labels(freq, ch1, ch2, phase_ch1, phase_ch2)
        self.update_graph()

    def update_labels(self, freq, ch1, ch2, phase_ch1, phase_ch2):
        self.freq_label.config(text=f"{freq:.2f}")
        self.ch1_label.config(text=f"{ch1:.2f}")
        self.ch2_label.config(text=f"{ch2:.2f}")
        gain, phase = self.calculate_gain_phase(ch1, ch2, phase_ch1, phase_ch2)
        self.gain_label.config(text=f"{gain:.2f}")
        self.phase_label.config(text=f"{phase:.2f}")

    def calculate_gain_phase(self, ch1, ch2, phase_ch1, phase_ch2):
        ratio = self.ratio.get()
        if ratio == "CH2/CH1":
            gain = 20 * np.log10(ch2 / ch1)
            phase = phase_ch2 - phase_ch1
        elif ratio == "CH1/CH2":
            gain = 20 * np.log10(ch1 / ch2)
            phase = phase_ch1 - phase_ch2
        elif ratio == "CH1/Output":
            gain = 20 * np.log10(ch1 / self.amplitude.get())
            phase = phase_ch1
        elif ratio == "CH2/Output":
            gain = 20 * np.log10(ch2 / self.amplitude.get())
            phase = phase_ch2
        return gain, phase

    def update_graph(self):
        self.ax.clear()
        frequencies = [data[0] for data in self.sweep_data]
        gains = [self.calculate_gain_phase(data[1], data[2], data[3], data[4])[0] for data in self.sweep_data]
        phases = [self.calculate_gain_phase(data[1], data[2], data[3], data[4])[1] for data in self.sweep_data]
        self.ax.plot(frequencies, gains, label="Gain [dB]")
        self.ax.plot(frequencies, phases, label="Phase [deg]")
        self.ax.legend()
        self.canvas.draw()

    def update_marker(self, event):
        marker_freq = self.slider.get()
        self.marker_freq_label.config(text=f"{marker_freq:.2f}")
        # Update marker gain and phase labels based on the marker frequency

    def calculate_frequencies(self):
        start = self.sweep_start.get()
        stop = self.sweep_stop.get()
        steps = self.steps.get()
        if self.sweep_type.get() == "logarithmic":
            return np.logspace(np.log10(start), np.log10(stop), steps)
        else:
            return np.linspace(start, stop, steps)

    def set_output_parameters(self):
        self.instrument.write(f"WAVEFO,{self.waveform.get()}")
        self.instrument.write(f"AMPLIT,{self.amplitude.get()}")
        self.instrument.write(f"OFFSET,{self.dc_offset.get() if self.dc_offset_enabled.get() else 0}")
        self.instrument.write(f"SPEED,{self.speed.get()}")
        self.instrument.write(f"CYCLES,{self.cycles.get()}")
        self.instrument.write(f"DELAY,{self.delay_time.get()}")

if __name__ == "__main__":
    root = tk.Tk()
    app = GPIBApp(root)
    root.mainloop()