import pyvisa
import time
import numpy as np

# List of all valid timebases supported by OWON HDS2* series
TIMEBASES = [
    (2e-9, "2.0ns"), (5e-9, "5.0ns"), (10e-9, "10.0ns"), (20e-9, "20.0ns"),
    (50e-9, "50.0ns"), (100e-9, "100ns"), (200e-9, "200ns"), (500e-9, "500ns"),
    (1e-6, "1.0us"), (2e-6, "2.0us"), (5e-6, "5.0us"), (10e-6, "10us"),
    (20e-6, "20us"), (50e-6, "50us"), (100e-6, "100us"), (200e-6, "200us"),
    (500e-6, "500us"), (1e-3, "1.0ms"), (2e-3, "2.0ms"), (5e-3, "5.0ms"),
    (10e-3, "10ms"), (20e-3, "20ms"), (50e-3, "50ms"), (100e-3, "100ms"),
    (200e-3, "200ms"), (500e-3, "500ms"), (1.0, "1.0s"), (2.0, "2.0s"),
    (5.0, "5.0s"), (10.0, "10s"), (20.0, "20s"), (50.0, "50s"),
    (100.0, "100s"), (200.0, "200s"), (500.0, "500s"), (1000.0, "1000s")
]

# Supported scale tables according to probe attenuation
VOLT_SCALES = {
    "1X": [
        (0.010, "10.0mV"), (0.020, "20.0mV"), (0.050, "50.0mV"),
        (0.100, "100mV"),  (0.200, "200mV"),  (0.500, "500mV"),
        (1.000, "1.00V"),  (2.000, "2.00V"),  (5.000, "5.00V"),
        (10.00, "10.0V")
    ],
    "10X": [
        (0.100, "100mV"),  (0.200, "200mV"),  (0.500, "500mV"),
        (1.000, "1.00V"),  (2.000, "2.00V"),  (5.000, "5.00V"),
        (10.00, "10.0V"),  (20.00, "20.0V"),  (50.00, "50.0V"),
        (100.0, "100V")
    ],
    "100X": [
        (1.000, "1.00V"),  (2.000, "2.00V"),  (5.000, "5.00V"),
        (10.00, "10.0V"),  (20.00, "20.0V"),  (50.00, "50.0V"),
        (100.0, "100V"),   (200.0, "200V"),   (500.0, "500V"),
        (1000.0, "1.00kV")
    ],
    "1000X": [
        (10.00, "10.0V"),  (20.00, "20.0V"),  (50.00, "50.0V"),
        (100.0, "100V"),   (200.0, "200V"),   (500.0, "500V"),
        (1000.0, "1.00kV"),(2000.0, "2.00kV"),(5000.0, "5.00kV"),
        (10000.0, "10.0kV")
    ]
}

CSV_FILEPATH      = "csv/bode_data.csv"
LOG_ENABLE        = False
PROBE             = "1X"
GEN_PEAK          = 0.10 # 2*GEN_PEAK peak to peak voltage
FREQS_COUNT       = 10
MEAN_SAMPLE_COUNT = 3
curr_scale_idx    = 2 # 500 mV startup vertical height
scales            = VOLT_SCALES.get(PROBE)

def LOG(msg):
    if LOG_ENABLE:
        print(f"[LOG] {msg}")

def LOG_LOC(loc, msg):
    if LOG_ENABLE:
        print(f"[LOG]{loc} {msg}")

def set_best_timebase(frequency):
    period = 1.0 / frequency
    target_scale = period / 2.0  # Display width is roughly 10 divisions

    # Pick the closest supported value in the list
    best_match = min(TIMEBASES, key=lambda x: abs(x[0] - target_scale))
    LOG(f"HORizontal:SCALe {best_match[1]}")
    time.sleep(0.1)
    scope.write(f":HORizontal:SCALe {best_match[1]}")

def set_best_vert_scale():
    global curr_scale_idx

    read_val_str = scope.query(":MEASurement:CH1:MAX?").strip()
    LOG_LOC("[set_best_vert_scale]", f"read_val_str == {read_val_str}")

    while read_val_str.startswith(">"):
        LOG_LOC("[set_best_vert_scale]['>']", f"startswith('>')")
        if curr_scale_idx <= 9:
            curr_scale_idx = curr_scale_idx + 1
            new_scale = scales[curr_scale_idx][1]
            LOG_LOC("[set_best_vert_scale]['>']", f"Changing vertical scale to {new_scale}")
            scope.write(f":CH1:SCALe {new_scale}")

            if freq < 10:
                # Roll mode
                time.sleep(10)
            else:
                time.sleep(1)

            # Read again with new scale
            read_val_str = scope.query(":MEASurement:CH1:MAX?").strip()
            LOG_LOC("[set_best_vert_scale]['>'][read again]", f"read_val_str == {read_val_str}")
        else:
            read_val_str.lstrip(">")

    while read_val_str.startswith("<"):
        LOG_LOC("[set_best_vert_scale]['<']", f"startswith('<')")
        if curr_scale_idx > 0:
            curr_scale_idx = curr_scale_idx - 1
            new_scale = scales[curr_scale_idx][1]
            # print(f"Changing vertical scale to {new_scale}")
            scope.write(f":CH1:SCALe {new_scale}")

            if freq < 10:
                # Roll mode
                time.sleep(10)
            else:
                time.sleep(1)

            # Read again with new scale
            read_val_str = scope.query(":MEASurement:CH1:MAX?").strip()
            LOG_LOC("[set_best_vert_scale]['<']", f"read_val_str == {read_val_str}")
        else:
            read_val_str.lstrip("<")

    read_val = float(read_val_str)
    LOG_LOC("[set_best_vert_scale]", f"read_val == {read_val_str}")

    if curr_scale_idx > 0 and read_val < 2*scales[curr_scale_idx - 1][0]:
        curr_scale_idx = curr_scale_idx - 1
        new_scale = scales[curr_scale_idx][1]
        # print(f"Changing vertical scale to {new_scale}")
        scope.write(f":CH1:SCALe {new_scale}")

        if freq < 10:
            # Roll mode
            time.sleep(10)

        # Read again with new scale
        read_val_str = scope.query(":MEASurement:CH1:MAX?").strip()
        LOG_LOC("[set_best_vert_scale][read_val < 2*next_scale", f"read_val_str == {read_val_str}")

        while read_val_str.startswith(">") or read_val_str.startswith("<"):
            read_val_str = scope.query(":MEASurement:CH1:MAX?").strip()

        read_val = float(read_val_str)

# example: 5 decades (10^0 to 10^5) at 5 points/decade -> 5 * 5 + 1 = 26 points
# freqs = np.concatenate((np.array([1]), np.logspace(4, 5, num=FREQS_COUNT)))
freqs = np.logspace(0, 5, num=FREQS_COUNT)
freqs = np.unique([int(f) for f in freqs])
print(freqs)
amps  = []

rm = pyvisa.ResourceManager('@py')

try:
    scope = rm.open_resource("ASRL/dev/ttyUSB0::INSTR")
except:
    print("[ERROR] Device not found")
    exit(0)

try:
    data = scope.read_raw()
    print("Pending:", data)
except pyvisa.errors.VisaIOError as e:
    print("No pending SCPI data")

print("Identification: ", scope.query("*IDN?"), end="")

# Configure scope
scope.write(f":CH1:SCALe {scales[curr_scale_idx][1]}")
scope.write(":CH1:COUPling AC")
scope.write(f":CH1:PROBE {PROBE}")
scope.write(f":FUNCtion:AMPLitude {GEN_PEAK*2}")
scope.write(":FUNCtion SINE")

print(f"Generator set to function SINE")
# print("    freq,  amp/peak")

for freq in freqs:
    # print(f"Probing {freq:2.2f} Hz")
    scope.write(f":FUNCtion:FREQuency {freq:2.2f}")

    # print(scope.query(":FUNCtion:FREQuency?"))
    set_best_timebase(freq)

    # Roll mode
    if freq < 10:
        scope.write(":CH1:COUPling DC")
        time.sleep(10)
    else:
        scope.write(":CH1:COUPling AC")
        time.sleep(0.5)

    amp_arr = []
    for i in range(MEAN_SAMPLE_COUNT):
        set_best_vert_scale()
        read_val_str = scope.query(":MEASurement:CH1:MAX?").strip()
        LOG_LOC(f"[main_loop]", f"read_val_str == {read_val_str}")

        while read_val_str.startswith(">") or read_val_str.startswith("<"):
            set_best_vert_scale()
            read_val_str = scope.query(":MEASurement:CH1:MAX?").strip()
            LOG_LOC(f"[main_loop]", f"read_val_str == {read_val_str}")

        read_val = float(read_val_str)

        # print(f"{read_val:<3.3f}")
        amp_arr.append(read_val)

        time.sleep(0.5) # time between mean samples

    amp_mean = (sum(amp_arr) / MEAN_SAMPLE_COUNT)
    amp_mean = 0.001 if amp_mean < 0.001 else amp_mean

    amp_norm = amp_mean / GEN_PEAK
    amps.append(amp_norm)

    print(f"{freq:>8}, {amp_norm:<3.3f}")

scope.close()

# print(f"freqs = [{' '.join(map(str, freqs))}]")
# print(f"amps  = [{' '.join(map(str, amps))}]")

data = np.column_stack((freqs, amps))
np.savetxt(f"{CSV_FILEPATH}",
           data,
           delimiter=",",
           header="Frequency,Amplitude",
           comments="",
           fmt=["%d", "%.3f"])
