#!/usr/bin/env python3
"""Vaunix phase shifter sweep with S21 measurement on a Rohde & Schwarz ZND VNA."""
from __future__ import annotations

import datetime
import os
import time

import numpy as np

from dexplore.data_folder import DataFolder
from stlab.devices.RS_ZND import RS_ZND
from stlab.devices.Vaunix_Phase import Vaunix_Phase

vna_addr = "TCPIP0::172.19.20.21::inst0::INSTR"
freq_start = 4e9
freq_stop = 8e9
ifbw = 1000
vna_points_setting = 0
vna_points_max = 5000
auto_points_span_divisor = 5e7
averages = 1
vna_power = 0

phase_start = 0
phase_stop = 360
phase_step = 1
phase_loop_label = "Phase Shift (deg)"

data_dir = "./measurement_data"
dataset_name = ""
description = ""
settle_s = 0.01


def phase_values() -> np.ndarray:
    n = int((phase_stop - phase_start) / phase_step + 1)
    return np.linspace(phase_start, phase_stop, n)


def vna_points() -> int:
    if vna_points_setting > 0:
        return vna_points_setting
    span = freq_stop - freq_start
    return int(min(vna_points_max, span / auto_points_span_divisor))


def ds_name() -> str:
    if dataset_name:
        return dataset_name
    stamp = datetime.date.today().strftime("%Y%m%d")
    return f"{stamp}_phase"


def metadata(n_phase: int) -> dict:
    return {
        "Name": ds_name(),
        "Description": description or "Vaunix phase sweep",
        "frequency": f"{freq_start:.2e} - {freq_stop:.2e} Hz",
        "ifbw": ifbw,
        "vna_points": vna_points(),
        "averages": averages,
        "vna_power": vna_power,
        "ph_start": phase_start,
        "ph_stop": phase_stop,
        "ph_points": n_phase,
    }


def setup_s21(vna: RS_ZND) -> None:
    vna.write("CALC:PAR:DEL:ALL")
    vna.write("CALC:PAR:SDEF 'Trc1', 'S21'")
    vna.write("DISP:WIND1:TRAC:EFE 'Trc1'")


def configure_vna(vna: RS_ZND) -> None:
    setup_s21(vna)
    vna.SetRange(freq_start, freq_stop)
    vna.SetPoints(vna_points())
    vna.SetIFBW(ifbw)
    vna.SetPower(vna_power)


def main() -> None:
    phases = phase_values()
    meta = metadata(len(phases))
    name = ds_name()
    addr = os.environ.get("VNA_ADDR") or vna_addr
    out_dir = os.environ.get("SWEEP_DATA_DIR") or data_dir

    vna = RS_ZND(addr, reset=True, verb=False)
    phase = Vaunix_Phase()
    dfol = DataFolder(out_dir, __file__)
    created = False

    configure_vna(vna)
    print(f"Starting sweep -> {dfol.folder_full_path}")
    print(f"Dataset: {name}")

    t0 = time.time()
    for i, deg in enumerate(phases):
        print(f"  {phase_loop_label}: step {i + 1}/{len(phases)} = {deg}")
        phase.SetPhase(deg)
        time.sleep(settle_s)
        out = vna.MeasureScreen_pd(N_averages=averages)
        vna.write("DISP:WIND1:TRAC:Y:AUTO ONCE")
        if i == 0:
            est_min = (time.time() - t0) / 60 * len(phases)
            print(f"  Estimated time: {est_min:.1f} min")
        if not created:
            dfol.create_stlab_dataset(
                dataset_name=name,
                data_names=out.keys()[1:],
                trace_name=out.keys()[0],
                trace_values=out[out.keys()[0]],
                loop1_name=phase_loop_label,
                loop1_values=phases,
                loop2_name="None",
                loop2_values=[0.0],
            )
            dfol.datasets[name].attrs.update(meta)
            created = True
        dfol.add_stlab_trace(name, out, loop1_index=i, loop2_index=0)

    dfol.save_data()
    print(f"  Elapsed: {(time.time() - t0) / 60:.1f} min")
    print("Measurement complete.")

    phase.close()
    vna.close()


if __name__ == "__main__":
    main()
