"""Scan pulse duration for independent atoms and two finite interaction strengths."""

from __future__ import annotations

import argparse
import json
from math import pi
from pathlib import Path
from typing import Any

from bokeh.embed import file_html
from bokeh.layouts import column
from bokeh.plotting import figure
from bokeh.resources import INLINE

import cascaqit
from cascaqit import (
    AHSProgram,
    AtomRegister,
    LocalBackend,
    SimulationOptions,
    VanDerWaalsInteraction,
    Waveform,
)


def experiment(
    output_dir: Path, *, points: int = 8, time_steps: int = 32, shots: int = 256
) -> dict[str, Any]:
    """Retain each result and draw both language reports from the same rows."""
    if not 2 <= points <= 256 or time_steps < 1 or shots < 1:
        raise ValueError("Use 2..256 points, positive time steps and positive shots.")
    output_dir.mkdir(parents=True, exist_ok=True)
    omega, c6, cutoff, seed = 2 * pi, 312_500.0, 12.0, 4701
    configurations = (
        ("independent", 5.0, False),
        ("near", 5.0, True),
        ("far", 10.0, True),
    )
    options = SimulationOptions(
        method="state_vector",
        blockade_mode="full",
        dtype="complex128",
        integrator="fixed_step_krylov",
        max_steps=time_steps,
    )
    rows: list[dict[str, Any]] = []
    raw: list[dict[str, Any]] = []
    for label, spacing, enabled in configurations:
        for index in range(points):
            duration = (index + 1) / points
            program = (
                AHSProgram(
                    AtomRegister.line(count=2, spacing=spacing),
                    interaction=VanDerWaalsInteraction(
                        c6=c6,
                        cutoff_radius=cutoff,
                        enabled=enabled,
                    ),
                )
                .drive(
                    rabi=Waveform.constant(omega, duration=duration),
                    detuning=Waveform.constant(0.0, duration=duration),
                    phase=0.0,
                )
                .measure()
            )
            result = (
                LocalBackend(analog_time_steps=time_steps)
                .run(
                    program,
                    shots=shots,
                    seed=seed + len(rows),
                    options=options,
                )
                .result()
            )
            if result.probabilities is None or sum(result.counts.values()) != shots:
                raise RuntimeError("The scan did not return a complete distribution.")
            probabilities = result.probabilities
            rows.append(
                {
                    "case": label,
                    "spacing_um": spacing,
                    "duration_us": duration,
                    "interaction_rad_per_us": c6 / spacing**6 if enabled else 0.0,
                    "probability_q0": probabilities.get("10", 0)
                    + probabilities.get("11", 0),
                    "probability_11": probabilities.get("11", 0),
                    "frequency_11": result.counts.get("11", 0) / shots,
                    "seed": seed + len(rows),
                    "counts": result.counts,
                }
            )
            raw.append(result.to_dict())
    artifacts = {"data": str(output_dir / "scan.json")}
    labels = {
        "en": (
            "Pulse duration (us)",
            "Excitation probability of atom 0",
            "Double excitation",
            ("Independent", "5 um, finite interaction", "10 um, finite interaction"),
        ),
        "zh": (
            "脉冲时长 (us)",
            "第 0 个原子的激发概率",
            "双激发概率",
            ("独立原子", "5 um，有限相互作用", "10 um，有限相互作用"),
        ),
    }
    for language, (xlabel, first_title, second_title, legends) in labels.items():
        plots = []
        for field, title in (
            ("probability_q0", first_title),
            ("probability_11", second_title),
        ):
            plot = figure(
                title=title,
                x_axis_label=xlabel,
                y_axis_label="P",
                width=880,
                height=330,
                y_range=(-0.03, 1.03),
            )
            for (label, _, _), color, legend in zip(
                configurations, ("#2563eb", "#dc2626", "#15803d"), legends
            ):
                subset = [row for row in rows if row["case"] == label]
                xs, ys = [r["duration_us"] for r in subset], [r[field] for r in subset]
                plot.line(xs, ys, color=color, legend_label=legend, line_width=2)
                plot.scatter(xs, ys, color=color, size=6)
            plot.legend.click_policy = "hide"
            plots.append(plot)
        path = output_dir / f"rabi-{language}.html"
        path.write_text(file_html(column(*plots), INLINE, "Rabi"), encoding="utf-8")
        artifacts[f"report_{language}"] = str(path)
    payload = {
        "sdk_version": cascaqit.__version__,
        "omega_rad_per_us": omega,
        "c6_rad_um6_per_us": c6,
        "cutoff_um": cutoff,
        "points_per_case": points,
        "time_steps": time_steps,
        "shots_per_point": shots,
        "total_shots": len(rows) * shots,
        "simulation_options": options.to_dict(),
        "rows": rows,
        "raw_results": raw,
        "artifacts": artifacts,
    }
    Path(artifacts["data"]).write_text(
        json.dumps(payload, indent=2) + "\n", encoding="utf-8"
    )
    return {
        "sdk_version": cascaqit.__version__,
        "points_per_case": points,
        "time_steps": time_steps,
        "shots_per_point": shots,
        "total_shots": len(rows) * shots,
        "interaction_rad_per_us": {
            label: c6 / spacing**6 if enabled else 0.0
            for label, spacing, enabled in configurations
        },
        "middle_points": [r for r in rows if r["duration_us"] == 0.5],
        "artifacts": artifacts,
    }


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--output-dir", type=Path, default=Path("artifacts/rabi-interactions")
    )
    parser.add_argument("--points", type=int, default=8)
    parser.add_argument("--time-steps", type=int, default=32)
    parser.add_argument("--shots", type=int, default=256)
    args = parser.parse_args()
    print(
        json.dumps(
            experiment(
                args.output_dir,
                points=args.points,
                time_steps=args.time_steps,
                shots=args.shots,
            ),
            sort_keys=True,
        )
    )


if __name__ == "__main__":
    main()
