"""Separate dephasing, readout errors and finite sampling in a two-atom experiment."""

from __future__ import annotations

import argparse
import json
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,
    Circuit,
    HybridProgram,
    LocalBackend,
    ObservableSet,
    PauliZ,
    PauliZZ,
    SimulationOptions,
    SitePattern,
    Waveform,
)
from cascaqit.simulators import NoiseChannel, NoiseModel


def readout_distribution(
    probabilities: dict[str, float], p01: float, p10: float
) -> dict[str, float]:
    """Apply independent classical bit errors to an exact state distribution."""
    distribution = {bits: 0.0 for bits in ("00", "01", "10", "11")}
    for source, probability in probabilities.items():
        for target in distribution:
            weight = probability
            for before, after in zip(source, target):
                error = p01 if before == "0" else p10
                weight *= error if before != after else 1 - error
            distribution[target] += weight
    return distribution


def experiment(
    output_dir: Path, *, shots: int = 128, time_steps: int = 32
) -> dict[str, Any]:
    """Save every execution before drawing the two language reports."""
    if shots < 1 or time_steps < 1:
        raise ValueError("Shots and time steps must be positive.")
    output_dir.mkdir(parents=True, exist_ok=True)
    duration, omega, detuning, rate = 0.6, 1.1, 0.7, 0.8
    p01, p10 = 0.08, 0.12
    analog = (
        AHSProgram(AtomRegister.line(count=2, spacing=5.0))
        .drive(
            rabi=Waveform.constant(omega, duration=duration),
            detuning=Waveform.constant(0.0, duration=duration),
            phase=0.0,
        )
        .local_detuning(
            waveform=Waveform.constant(detuning, duration=duration),
            pattern=SitePattern.from_mapping({"q0": 1.0, "q1": 0.25}),
        )
    )
    program = (
        HybridProgram("project.noisy.hybrid")
        .digital("prepare", Circuit(2).h(0).cx(0, 1))
        .analog("evolve", analog)
        .digital("readout_rotation", Circuit(2).cx(0, 1).h(0))
        .measure_all()
    )
    cases = (
        ("ideal", None),
        ("dephasing", NoiseModel("phase", (NoiseChannel.dephasing(rate),))),
        (
            "dephasing_readout",
            NoiseModel(
                "phase.readout",
                (NoiseChannel.dephasing(rate), NoiseChannel.readout(p01, p10=p10)),
            ),
        ),
    )
    rows: list[dict[str, Any]] = []
    raw: list[dict[str, Any]] = []
    for label, noise in cases:
        options = SimulationOptions(
            method="state_vector" if noise is None else "density_matrix",
            dtype="complex128",
            integrator="fixed_step_krylov",
            max_steps=time_steps,
        )
        for seed in (4881, 4882, 4883):
            result = (
                LocalBackend(analog_time_steps=time_steps)
                .run(
                    program,
                    shots=shots,
                    seed=seed,
                    noise=noise,
                    options=options,
                    observables=ObservableSet(
                        (PauliZ("q0"), PauliZ("q1"), PauliZZ("q0", "q1"))
                    ),
                )
                .result()
            )
            if result.probabilities is None or result.observable_batch is None:
                raise RuntimeError(
                    "The experiment requires probabilities and observables."
                )
            expected = readout_distribution(
                result.probabilities,
                p01 if label == "dephasing_readout" else 0.0,
                p10 if label == "dephasing_readout" else 0.0,
            )
            rows.append(
                {
                    "case": label,
                    "seed": seed,
                    "counts": result.counts,
                    "probabilities_before_readout": result.probabilities,
                    "expected_recorded_probabilities": expected,
                    "frequency_00": result.counts.get("00", 0) / shots,
                    "observables": result.observable_batch.to_dict(),
                    "state_chain": [
                        item.to_dict() for item in result.state_transitions()
                    ],
                    "method": result.metadata["simulation_plan"]["method_selected"],
                    "resource_estimate": result.metadata[
                        "simulation_resource_estimate"
                    ],
                    "options": options.to_dict(),
                    "noise_model": None if noise is None else noise.to_dict(),
                }
            )
            raw.append(result.to_dict())
    artifacts = {"data": str(output_dir / "comparison.json")}
    for language in ("en", "zh"):
        title = (
            "Probability and recorded frequency of 00"
            if language == "en"
            else "00 的概率与记录频率"
        )
        names = (
            ("Ideal", "Dephasing", "Dephasing + readout")
            if language == "en"
            else ("理想", "退相干", "退相干与读出误差")
        )
        plot = figure(
            title=title, x_range=list(names), y_range=(0, 1), width=880, height=430
        )
        first = rows[::3]
        plot.scatter(
            list(names),
            [r["probabilities_before_readout"]["00"] for r in first],
            size=14,
            marker="square",
            color="#2563eb",
            legend_label="Before readout" if language == "en" else "读出前概率",
        )
        plot.scatter(
            list(names),
            [r["expected_recorded_probabilities"]["00"] for r in first],
            size=16,
            marker="dash",
            color="#15803d",
            legend_label="Expected recorded" if language == "en" else "预期记录概率",
        )
        for index in range(3):
            plot.scatter(
                [(name, (index - 1) * 0.12) for name in names],
                [rows[case * 3 + index]["frequency_00"] for case in range(3)],
                size=8,
                color="#dc2626",
                legend_label="Sampled (3 seeds)"
                if language == "en"
                else "采样频率（三个种子）",
            )
        plot.legend.location = "bottom_left"
        path = output_dir / f"hybrid-{language}.html"
        path.write_text(file_html(column(plot), INLINE, title), encoding="utf-8")
        artifacts[f"report_{language}"] = str(path)
    payload = {
        "sdk_version": cascaqit.__version__,
        "program": program.to_dict(),
        "program_hash": program.stable_hash(),
        "duration_us": duration,
        "omega_rad_per_us": omega,
        "local_detuning_rad_per_us": [detuning, detuning * 0.25],
        "dephasing_rate_per_us": rate,
        "readout_p01": p01,
        "readout_p10": p10,
        "shots_per_run": shots,
        "total_shots": shots * len(rows),
        "time_steps": time_steps,
        "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__,
        "shots_per_run": shots,
        "total_shots": shots * len(rows),
        "cases": [
            {
                key: r[key]
                for key in (
                    "case",
                    "method",
                    "probabilities_before_readout",
                    "expected_recorded_probabilities",
                )
            }
            for r in rows[::3]
        ],
        "artifacts": artifacts,
    }


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


if __name__ == "__main__":
    main()
