"""Compare exact and sampled VQE objectives under a paired evaluation budget."""

from __future__ import annotations

import argparse
import json
from math import sqrt
from pathlib import Path
from typing import Any, Literal

import cascaqit
from cascaqit import (
    VQE,
    HamiltonianTerm,
    OptimizerConfig,
    PauliHamiltonian,
    PauliMeasurementConfig,
    PauliX,
    PauliZ,
    SampledSelectionConfig,
    SPSAConfig,
    VQESamplingBenchmarkConfig,
)


def experiment(
    output_dir: Path, *, repeats: int = 3, shots: int = 128, budget: int = 12
) -> dict[str, Any]:
    """Reuse executed benchmark data for both language reports and the JSON file."""
    hamiltonian = PauliHamiltonian(
        hamiltonian_id="project.finite-shot.hamiltonian",
        logical_order=("q0",),
        constant=0.1,
        terms=(
            HamiltonianTerm("x", 0.7, PauliX("q0")),
            HamiltonianTerm("z", -0.4, PauliZ("q0")),
        ),
    )
    vqe = VQE(hamiltonian, algorithm_id="project.finite-shot.vqe")
    config = VQESamplingBenchmarkConfig(
        repeats=repeats,
        objective_evaluation_budget=budget,
        optimizer=OptimizerConfig(
            method="SPSA",
            max_iterations=1,
            spsa=SPSAConfig(learning_rate=0.2, perturbation=0.15),
        ),
        measurement=PauliMeasurementConfig(
            shots_per_group=shots, allocation="coefficient_l1"
        ),
        sampled_selection=SampledSelectionConfig(
            candidate_count=2, repeats_per_candidate=2
        ),
        fixed_objective_repeats=2,
        adaptive_objective_repeats=2,
        adaptive_max_objective_repeats=3,
        adaptive_standard_error_target=0.08,
        final_shots=shots,
        confidence_level=0.95,
        root_seed=29,
    )
    result = vqe.benchmark_sampling(config)
    output_dir.mkdir(parents=True, exist_ok=True)
    artifacts = {"data": str(output_dir / "benchmark.json")}
    languages: tuple[Literal["en", "zh"], ...] = ("en", "zh")
    for language in languages:
        path = output_dir / f"vqe-{language}.html"
        result.report(
            path,
            language=language,
            title="Finite-shot VQE" if language == "en" else "有限采样 VQE",
        )
        artifacts[f"report_{language}"] = str(path)
    rows = []
    for run in result.runs:
        selection = run.result.sampled_selection
        rows.append(
            {
                "strategy": run.strategy,
                "repeat_index": run.repeat_index,
                "paired_seed": run.paired_seed,
                "selected_parameters": dict(
                    run.result.selected_evaluation.parameter_bind.values
                ),
                "selected_estimator_energy": run.selected_estimator_energy,
                "exact_selected_energy": run.exact_selected_energy,
                "estimator_error": run.estimator_error,
                "paired_exact_reference_gap": run.paired_exact_reference_gap,
                "objective_evaluations": run.objective_evaluation_count,
                "objective_backend_executions": run.objective_backend_execution_count,
                "objective_shots": run.objective_shots,
                "confirmation_backend_executions": (
                    run.confirmation_backend_execution_count
                ),
                "confirmation_shots": run.confirmation_shots,
                "final_sampling_backend_executions": (
                    run.final_sampling_backend_execution_count
                ),
                "final_sampling_shots": run.final_sampling_shots,
                "diagnostic_exact_backend_executions": (
                    run.diagnostic_exact_backend_execution_count
                ),
                "selection_status": None if selection is None else selection.status,
            }
        )
    payload = {
        "sdk_version": cascaqit.__version__,
        "hamiltonian": hamiltonian.to_dict(),
        "ground_energy": 0.1 - sqrt(0.65),
        "config": config.to_dict(),
        "rows": rows,
        "benchmark": result.to_dict(),
        "artifacts": artifacts,
    }
    Path(artifacts["data"]).write_text(
        json.dumps(payload, indent=2) + "\n", encoding="utf-8"
    )
    return {
        "sdk_version": cascaqit.__version__,
        "paired_repeats": repeats,
        "objective_evaluation_ceiling": budget,
        "ground_energy": payload["ground_energy"],
        "strategies": [
            {
                "strategy": item.strategy,
                "mean_exact_selected_energy": item.mean_exact_selected_energy,
                "mean_paired_exact_reference_gap": item.mean_paired_exact_reference_gap,
                "estimator_error_rmse": item.estimator_error_rmse,
            }
            for item in result.statistics
        ],
        "total_backend_executions": sum(
            row[key]
            for row in rows
            for key in (
                "objective_backend_executions",
                "confirmation_backend_executions",
                "final_sampling_backend_executions",
                "diagnostic_exact_backend_executions",
            )
        ),
        "total_shots": sum(
            row[key]
            for row in rows
            for key in ("objective_shots", "confirmation_shots", "final_sampling_shots")
        ),
        "artifacts": artifacts,
    }


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--output-dir", type=Path, default=Path("artifacts/finite-shot-vqe")
    )
    parser.add_argument("--repeats", type=int, default=3)
    parser.add_argument("--shots", type=int, default=128)
    parser.add_argument("--budget", type=int, default=12)
    args = parser.parse_args()
    print(
        json.dumps(
            experiment(
                args.output_dir,
                repeats=args.repeats,
                shots=args.shots,
                budget=args.budget,
            ),
            sort_keys=True,
        )
    )


if __name__ == "__main__":
    main()
