"""Compare three local compilation routes for the same three-node MIS problem."""

from __future__ import annotations

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

import cascaqit
from cascaqit import LocalBackend, MockNeutralAtomTarget, SimulationOptions, visualize
from cascaqit.problems import GraphProblemIR, ProblemCompiler, decode_graph_bitstring


def experiment(
    output_dir: Path, *, shots: int = 128, seed: int = 4702
) -> dict[str, Any]:
    """Keep problem, candidates, parameter budgets and report inputs together."""
    if shots < 1:
        raise ValueError("shots must be positive.")
    output_dir.mkdir(parents=True, exist_ok=True)
    graph = GraphProblemIR.from_edges(
        problem_id="lesson.project.mis-routes",
        positions={"a": (0.0, 0.0), "b": (6.0, 0.0), "c": (12.0, 0.0)},
        edges=(("a", "b"), ("b", "c")),
    )
    candidates = [decode_graph_bitstring(graph, f"{i:03b}") for i in range(8)]
    feasible = [c for c in candidates if c["is_independent"]]
    optimum = max(c["selection_size"] for c in feasible)
    target = MockNeutralAtomTarget.local_ahs_v0_1()
    backend = LocalBackend(target=target, analog_time_steps=32)
    options = SimulationOptions(
        dtype="complex128", integrator="fixed_step_krylov", max_steps=32
    )
    parameter_sets = {
        "digital": (
            {"gamma_0": 0.16, "beta_0": 0.24},
            {"gamma_0": 0.28, "beta_0": -0.18},
        ),
        "hybrid": (
            {"gamma_0": 0.16, "beta_0": 0.24},
            {"gamma_0": 0.28, "beta_0": -0.18},
        ),
        "analog": (
            {"anneal_time": 0.4, "omega_max": 1.0},
            {"anneal_time": 0.7, "omega_max": 1.4},
        ),
    }
    executions = {}
    modes = {}
    for mode, parameters in parameter_sets.items():
        algorithm = "qaa" if mode == "analog" else "qaoa"
        compiled = ProblemCompiler().compile(
            graph, mode=mode, algorithm=algorithm, target=target
        )
        execution = compiled.optimize(
            parameter_sets=parameters,
            shots=shots,
            seed=seed,
            backend=backend,
            options=None if mode == "digital" else options,
        )
        counts = execution.result.counts
        decoded = [
            (decode_graph_bitstring(graph, bits), count)
            for bits, count in counts.items()
        ]
        feasible_shots = sum(count for item, count in decoded if item["is_independent"])
        optimal_shots = sum(
            count
            for item, count in decoded
            if item["is_independent"] and item["selection_size"] == optimum
        )
        history_shots = sum(item.shots for item in execution.parameter_history)
        modes[mode] = {
            "algorithm": algorithm,
            "problem_hash": execution.problem_hash,
            "logical_order": list(execution.logical_order),
            "parameter_sets": parameters,
            "evaluation_count": len(execution.parameter_history),
            "history_shots": history_shots,
            "selected_result_shots": sum(counts.values()),
            "selected_evaluation_index": execution.selected_evaluation_index,
            "objective_value": execution.objective_value,
            "best_observed_bitstring": execution.best_observed_candidate.bitstring,
            "best_observed_feasible": execution.best_observed_candidate.feasible,
            "counts": counts,
            "feasible_frequency": feasible_shots / shots,
            "optimal_frequency": optimal_shots / shots,
        }
        executions[mode] = execution
    artifacts = {"data": str(output_dir / "routes.json")}
    for language in ("en", "zh"):
        path = output_dir / f"mis-{language}.html"
        visualize(
            executions,
            output=path,
            language=language,
            title="MIS: Digital / Hybrid / Analog",
        )
        artifacts[f"report_{language}"] = str(path)
    summary = {
        "sdk_version": cascaqit.__version__,
        "seed": seed,
        "shots_per_evaluation": shots,
        "classical_maximum_size": optimum,
        "classical_optimal_bitstrings": [
            c["bitstring"] for c in feasible if c["selection_size"] == optimum
        ],
        "modes": modes,
        "artifacts": artifacts,
    }
    raw = {
        **summary,
        "graph": graph.to_dict(),
        "simulation_options": options.to_dict(),
        "executions": {mode: run.to_dict() for mode, run in executions.items()},
    }
    Path(artifacts["data"]).write_text(
        json.dumps(raw, indent=2) + "\n", encoding="utf-8"
    )
    return summary


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output-dir", type=Path, default=Path("artifacts/mis-routes"))
    parser.add_argument("--shots", type=int, default=128)
    parser.add_argument("--seed", type=int, default=4702)
    args = parser.parse_args()
    print(
        json.dumps(
            experiment(args.output_dir, shots=args.shots, seed=args.seed),
            sort_keys=True,
        )
    )


if __name__ == "__main__":
    main()
