"""FDTD and EME radius sweep for a 90-degree X-cut TFLN Euler bend.

The X-cut crystal optic axis lies along global y, in the bend plane, matching
the official AnisotropicBendsEME notebook.  Each EME cell solves and propagates
16 modes with no mode sweep.  All X-cut tasks use a new result chain; no Z-cut
simulation data are reused.  Lengths are in micrometers and frequencies are
in Hz.

Commands:

    python tfln_euler_bend_radius_sweep_final.py build
    python tfln_euler_bend_radius_sweep_final.py estimate
    python tfln_euler_bend_radius_sweep_final.py run
    python tfln_euler_bend_radius_sweep_final.py analyze
"""

from __future__ import annotations

import argparse
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
import json
from pathlib import Path

import gdstk
import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import tidy3d as td
from scipy.integrate import cumulative_trapezoid
from scipy.interpolate import CubicSpline


# -----------------------------------------------------------------------------
# Sweep, device, and numerical parameters.
# -----------------------------------------------------------------------------

RADIUS_VALUES_UM = np.asarray([5.0, 7.5, 10.0, 15.0, 20.0])
NEW_RADIUS_VALUES_UM = RADIUS_VALUES_UM.copy()
WAVELENGTH_UM = 1.55
FREQUENCY_HZ = td.C_0 / WAVELENGTH_UM

# Literature-typical thin-film lithium-niobate rib geometry.
LN_THICKNESS_UM = 0.60
ETCH_DEPTH_UM = 0.30
SLAB_THICKNESS_UM = LN_THICKNESS_UM - ETCH_DEPTH_UM
RIDGE_TOP_WIDTH_UM = 0.90
SIDEWALL_FROM_HORIZONTAL_DEG = 70.0
SIDEWALL_ANGLE_RAD = np.deg2rad(90.0 - SIDEWALL_FROM_HORIZONTAL_DEG)

BEND_ANGLE_RAD = np.pi / 2
INPUT_LEAD_UM = 12.0
OUTPUT_LEAD_UM = 12.0

# Final EME settings after mode-count, bend-cell, and mode-plane checks.
EME_NUM_CELLS = 10
EME_NUM_MODES = 16

MODE_WINDOW_Y_UM = 10.0
MODE_WINDOW_Z_UM = 8.0
# Memory-safe X-cut mode plane.  The radial span remains large enough for the
# outward-shifted bend mode, while excess vertical cladding is removed.
EME_PLANE_Y_UM = 9.0
EME_PLANE_Z_UM = 6.0
TARGET_NEFF = 2.00
STEPS_PER_WAVELENGTH = 40
FDTD_RUN_TIME_QUALITY_FACTOR = 1.0
PATH_TOLERANCE_UM = 1e-4
PATH_SAMPLES = 4001

TASK_PREFIX = "tfln_euler_xcut_radius_sweep_1550_m16_r5_20"
OUTPUT_DIR = Path("results/radius_sweep_xcut_1550_m16_r5_20")
MANIFEST_PATH = OUTPUT_DIR / f"{TASK_PREFIX}_tasks.json"

# The crystal-cut change invalidates every completed Z-cut result.  A fresh
# manifest is required so the analysis cannot mix the two physical models.
REUSED_TASKS = {"fdtd": {}, "eme": {}}


@dataclass(frozen=True)
class EulerCenterline:
    """Sampled centerline and curvature for one symmetric Euler turn."""

    min_radius_um: float
    s: np.ndarray
    x: np.ndarray
    y: np.ndarray
    curvature: np.ndarray

    @property
    def length(self) -> float:
        return float(self.s[-1])

    @property
    def x_end(self) -> float:
        return float(self.x[-1])

    @property
    def y_end(self) -> float:
        return float(self.y[-1])


def radius_tag(radius_um: float) -> str:
    """Filename-safe radius label."""

    return f"{radius_um:g}".replace(".", "p")


def task_name(solver: str, radius_um: float) -> str:
    return f"{TASK_PREFIX}_{solver}_radius_{radius_tag(radius_um)}um"


def build_euler_centerline(min_radius_um: float) -> EulerCenterline:
    """Construct a 0-to-maximum-to-0 curvature 90-degree Euler turn."""

    half_length = BEND_ANGLE_RAD * min_radius_um
    s = np.linspace(0.0, 2.0 * half_length, PATH_SAMPLES)
    curvature = np.where(
        s <= half_length,
        s / (min_radius_um * half_length),
        (2.0 * half_length - s) / (min_radius_um * half_length),
    )
    tangent_angle = cumulative_trapezoid(curvature, s, initial=0.0)
    x = cumulative_trapezoid(np.cos(tangent_angle), s, initial=0.0)
    y = cumulative_trapezoid(np.sin(tangent_angle), s, initial=0.0)
    return EulerCenterline(
        min_radius_um=float(min_radius_um),
        s=s,
        x=x,
        y=y,
        curvature=curvature,
    )


def ln_medium():
    """X-cut lithium niobate with its optic axis along global y."""

    return td.material_library["LiNbO3"]["Zelmon1997"](1)


def sio2_medium():
    try:
        return td.material_library["SiO2"]["Palik_NoLoss"]
    except KeyError:
        return td.material_library["SiO2"]["Palik_Lossless"]


def mode_sort_spec() -> td.ModeSortSpec:
    """Select the highest-index quasi-TE mode."""

    return td.ModeSortSpec(
        filter_key="TE_fraction",
        filter_reference=0.5,
        filter_order="over",
        sort_key="n_eff",
        sort_order="descending",
        track_freq=None,
    )


def mode_spec(
    *,
    bend_radius: float | None = None,
    for_eme: bool = False,
) -> td.ModeSpec:
    cls = td.EMEModeSpec if for_eme else td.ModeSpec
    kwargs = dict(
        num_modes=EME_NUM_MODES if for_eme else 2,
        target_neff=TARGET_NEFF,
        num_pml=(12, 12),
        bend_radius=bend_radius,
        bend_axis=1 if bend_radius is not None else None,
        sort_spec=mode_sort_spec(),
        precision="double",
    )
    if for_eme:
        # Preserve the physical crystal axes while the waveguide direction
        # rotates through the bend, as in the AnisotropicBendsEME notebook.
        kwargs["bend_medium_frame"] = "global"
        kwargs["increasing_mode_tolerance"] = 1e-3
        kwargs["interp_spec"] = None
    return cls(**kwargs)


def stack_structures(ridge_geometry: td.Geometry) -> list[td.Structure]:
    medium = ln_medium()
    slab = td.Structure(
        geometry=td.Box(
            center=(0.0, 0.0, -SLAB_THICKNESS_UM / 2),
            size=(td.inf, td.inf, SLAB_THICKNESS_UM),
        ),
        medium=medium,
        name="ln_slab",
    )
    ridge = td.Structure(
        geometry=ridge_geometry,
        medium=medium,
        name="ln_ridge",
    )
    return [slab, ridge]


def ridge_polygon(centerline: EulerCenterline) -> np.ndarray:
    """Create one smooth ridge polygon along the Euler path and both leads."""

    x_spline = CubicSpline(centerline.s, centerline.x)
    y_spline = CubicSpline(centerline.s, centerline.y)

    def curve(u: float) -> tuple[float, float]:
        s = float(np.clip(u, 0.0, 1.0) * centerline.length)
        return float(x_spline(s)), float(y_spline(s))

    def gradient(u: float) -> tuple[float, float]:
        s = float(np.clip(u, 0.0, 1.0) * centerline.length)
        return (
            float(x_spline(s, 1) * centerline.length),
            float(y_spline(s, 1) * centerline.length),
        )

    cell = gdstk.Cell(f"euler_r{radius_tag(centerline.min_radius_um)}")
    path = gdstk.RobustPath(
        (-INPUT_LEAD_UM, 0.0),
        RIDGE_TOP_WIDTH_UM,
        tolerance=PATH_TOLERANCE_UM,
        max_evals=100000,
        ends="flush",
    )
    path.segment((0.0, 0.0))
    path.parametric(
        curve,
        path_gradient=gradient,
        width=RIDGE_TOP_WIDTH_UM,
        relative=False,
    )
    path.segment((centerline.x_end, centerline.y_end + OUTPUT_LEAD_UM))
    cell.add(path)
    polygons = cell.get_polygons()
    if len(polygons) != 1:
        raise RuntimeError(f"Expected one ridge polygon, got {len(polygons)}.")
    vertices = np.asarray(polygons[0].points, dtype=float)
    if vertices.shape[0] < 50:
        raise RuntimeError(
            f"R={centerline.min_radius_um:g} um ridge has only "
            f"{vertices.shape[0]} vertices."
        )
    return vertices


def ridge_polyslab(vertices: np.ndarray) -> td.PolySlab:
    return td.PolySlab(
        axis=2,
        slab_bounds=(0.0, LN_THICKNESS_UM),
        vertices=vertices,
        sidewall_angle=SIDEWALL_ANGLE_RAD,
        reference_plane="top",
    )


def build_fdtd(centerline: EulerCenterline) -> td.Simulation:
    """Build the explicit curved FDTD reference for one bend radius."""

    structures = stack_structures(ridge_polyslab(ridge_polygon(centerline)))
    input_x = -3.0
    output_y = centerline.y_end + 3.0
    input_plane = td.Box(
        center=(input_x, 0.0, LN_THICKNESS_UM / 2),
        size=(0.0, MODE_WINDOW_Y_UM, MODE_WINDOW_Z_UM),
    )
    output_plane = td.Box(
        center=(centerline.x_end, output_y, LN_THICKNESS_UM / 2),
        size=(MODE_WINDOW_Y_UM, 0.0, MODE_WINDOW_Z_UM),
    )
    source = td.ModeSource(
        center=input_plane.center,
        size=input_plane.size,
        source_time=td.GaussianPulse(
            freq0=FREQUENCY_HZ,
            fwidth=0.1 * FREQUENCY_HZ,
        ),
        mode_spec=mode_spec(),
        direction="+",
        mode_index=0,
        name="source",
    )
    input_monitor_x = input_x + 1.0
    mode_in = td.ModeMonitor(
        center=(input_monitor_x, 0.0, LN_THICKNESS_UM / 2),
        size=input_plane.size,
        freqs=[FREQUENCY_HZ],
        mode_spec=mode_spec(),
        store_fields_direction="+",
        name="mode_in",
    )
    mode_out = td.ModeMonitor(
        center=output_plane.center,
        size=output_plane.size,
        freqs=[FREQUENCY_HZ],
        mode_spec=mode_spec(),
        store_fields_direction="+",
        name="mode_out",
    )

    x_min, x_max = -6.0, centerline.x_end + 4.0
    y_min, y_max = -4.0, centerline.y_end + 6.0
    z_min, z_max = -5.0, 4.0
    center = (
        (x_min + x_max) / 2,
        (y_min + y_max) / 2,
        (z_min + z_max) / 2,
    )
    return td.Simulation(
        center=center,
        size=(x_max - x_min, y_max - y_min, z_max - z_min),
        medium=sio2_medium(),
        structures=structures,
        sources=[source],
        monitors=[mode_in, mode_out],
        boundary_spec=td.BoundarySpec(
            x=td.Boundary.absorber(num_layers=80),
            y=td.Boundary.absorber(num_layers=80),
            z=td.Boundary.pml(num_layers=16),
        ),
        grid_spec=td.GridSpec.auto(
            wavelength=WAVELENGTH_UM,
            min_steps_per_wvl=STEPS_PER_WAVELENGTH,
        ),
        run_time=td.RunTimeSpec(quality_factor=FDTD_RUN_TIME_QUALITY_FACTOR),
        symmetry=(0, 0, 0),
    )


def build_eme(centerline: EulerCenterline) -> td.EMESimulation:
    """Build the straight-coordinate EME representation for one radius."""

    large = 1e6
    ridge_vertices = np.asarray(
        [
            [-large, -RIDGE_TOP_WIDTH_UM / 2],
            [large, -RIDGE_TOP_WIDTH_UM / 2],
            [large, RIDGE_TOP_WIDTH_UM / 2],
            [-large, RIDGE_TOP_WIDTH_UM / 2],
        ],
        dtype=float,
    )
    structures = stack_structures(ridge_polyslab(ridge_vertices))
    straight_mode = mode_spec(for_eme=True)

    bend_lengths = np.full(EME_NUM_CELLS, centerline.length / EME_NUM_CELLS)
    bend_s = (np.arange(EME_NUM_CELLS) + 0.5) * bend_lengths[0]
    sampled_curvature = np.interp(bend_s, centerline.s, centerline.curvature)
    # Negative is the Tidy3D convention for this counter-clockwise +x to +y turn.
    local_radii = -1.0 / np.maximum(sampled_curvature, 1e-12)
    radial_plane_y = min(EME_PLANE_Y_UM, 1.8 * centerline.min_radius_um)
    # The bent-coordinate plane must not cross its curvature center.  The
    # 9 um radial plane matches the official anisotropic-bend notebook and
    # leaves clearance from the R=5 um case's 5.56 um cell-center radius.
    if np.min(np.abs(local_radii)) <= radial_plane_y / 2:
        raise RuntimeError(
            "EME radial plane reaches the bend-coordinate singularity."
        )
    bend_modes = [
        mode_spec(bend_radius=float(local_radius), for_eme=True)
        for local_radius in local_radii
    ]
    mode_specs = [straight_mode, *bend_modes, straight_mode]
    lengths = np.concatenate(
        ([INPUT_LEAD_UM], bend_lengths, [OUTPUT_LEAD_UM])
    )
    eme_grid = td.EMEExplicitGrid(
        boundaries=np.cumsum(lengths)[:-1],
        mode_specs=mode_specs,
    )
    total_length = float(lengths.sum())
    return td.EMESimulation(
        center=(total_length / 2, 0.0, 0.0),
        size=(total_length, radial_plane_y, EME_PLANE_Z_UM),
        medium=sio2_medium(),
        structures=structures,
        axis=0,
        freqs=[FREQUENCY_HZ],
        eme_grid_spec=eme_grid,
        grid_spec=td.GridSpec.auto(
            wavelength=WAVELENGTH_UM,
            min_steps_per_wvl=STEPS_PER_WAVELENGTH,
        ),
        store_port_modes=False,
        sweep_spec=None,
        constraint="passive",
    )


def build_models(
    radii_um=RADIUS_VALUES_UM,
) -> tuple[dict[str, td.Simulation], dict[str, td.EMESimulation]]:
    """Build FDTD and EME models for the selected bend radii."""

    radii = _validated_radius_selection(radii_um)
    fdtd_models: dict[str, td.Simulation] = {}
    eme_models: dict[str, td.EMESimulation] = {}
    for radius_um in radii:
        centerline = build_euler_centerline(float(radius_um))
        fdtd_models[task_name("fdtd", radius_um)] = build_fdtd(centerline)
        eme_models[task_name("eme", radius_um)] = build_eme(centerline)
    return fdtd_models, eme_models


def _new_manifest() -> dict[str, object]:
    """Initialize a clean X-cut sweep manifest."""

    return {
        "task_prefix": TASK_PREFIX,
        "wavelength_um": WAVELENGTH_UM,
        "radii_um": RADIUS_VALUES_UM.tolist(),
        "crystal_cut": "X-cut",
        "optic_axis_global": "y",
        "settings": {
            "steps_per_wavelength": STEPS_PER_WAVELENGTH,
            "fdtd_run_time_quality_factor": FDTD_RUN_TIME_QUALITY_FACTOR,
            "eme_bend_cells": EME_NUM_CELLS,
            "eme_modes_solved_and_propagated": EME_NUM_MODES,
            "eme_plane_um": [EME_PLANE_Y_UM, EME_PLANE_Z_UM],
            "eme_bend_medium_frame": "global",
        },
        "tasks": json.loads(json.dumps(REUSED_TASKS)),
        "reused_radii_um": [],
        "new_radii_um": NEW_RADIUS_VALUES_UM.tolist(),
    }


def upload_and_estimate(
    radii_um=NEW_RADIUS_VALUES_UM,
    manifest_path: Path = MANIFEST_PATH,
) -> dict[str, object]:
    """Upload missing selected tasks and estimate them without starting compute."""

    from tidy3d import web

    radii = _validated_radius_selection(radii_um)
    manifest = (
        json.loads(manifest_path.read_text())
        if manifest_path.exists()
        else _new_manifest()
    )
    fdtd_models, eme_models = build_models(radii)
    models = {"fdtd": fdtd_models, "eme": eme_models}
    new_estimates: dict[str, dict[str, float]] = {"fdtd": {}, "eme": {}}
    aggregate = 0.0
    for solver in ("fdtd", "eme"):
        missing_names = {
            task_name(solver, radius_um): models[solver][task_name(solver, radius_um)]
            for radius_um in radii
            if radius_tag(radius_um) not in manifest["tasks"][solver]
        }
        if not missing_names:
            continue
        batch = web.Batch(
            simulations=missing_names,
            folder_name="TFLN Euler radius sweep",
            verbose=True,
        )
        batch.upload()
        batch_total = float(batch.estimate_cost(verbose=True))
        aggregate += batch_total
        for radius_um in radii:
            name = task_name(solver, radius_um)
            if name not in missing_names:
                continue
            job = batch.jobs[name]
            cost = float(job.estimate_cost(verbose=False))
            manifest["tasks"][solver][radius_tag(radius_um)] = {
                "task_name": name,
                "task_id": job.task_id,
                "estimated_cost_flexcredits": cost,
            }
            new_estimates[solver][radius_tag(radius_um)] = cost
    manifest["new_task_estimates_flexcredits"] = new_estimates
    manifest["aggregate_new_estimated_cost_flexcredits"] = aggregate
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    manifest_path.write_text(json.dumps(manifest, indent=2))
    print(f"New-task aggregate estimate: {aggregate:.3f} FlexCredits")
    print(f"Task manifest: {manifest_path}")
    return manifest


def _validated_radius_selection(radii_um) -> np.ndarray:
    radii = np.asarray(radii_um, dtype=float)
    if radii.ndim != 1 or radii.size == 0:
        raise ValueError("Select at least one radius.")
    if np.unique(radii).size != radii.size:
        raise ValueError("Radius selections must not contain duplicates.")
    unavailable = [
        float(radius)
        for radius in radii
        if not np.any(np.isclose(radius, RADIUS_VALUES_UM, rtol=0.0, atol=1e-12))
    ]
    if unavailable:
        raise ValueError(f"Radii not present in the uploaded sweep: {unavailable}.")
    return radii


def run_uploaded_tasks(
    radii_um=NEW_RADIUS_VALUES_UM,
    manifest_path: Path = MANIFEST_PATH,
) -> dict[str, str]:
    """Start only the explicitly selected, previously estimated tasks."""

    from tidy3d import web

    radii = _validated_radius_selection(radii_um)
    manifest = json.loads(manifest_path.read_text())
    selected = []
    for solver in ("fdtd", "eme"):
        for radius_um in radii:
            item = manifest["tasks"][solver][radius_tag(radius_um)]
            selected.append((solver, float(radius_um), item["task_id"]))

    def start_one(item: tuple[str, float, str]) -> tuple[str, str]:
        solver, radius_um, task_id = item
        status = str(web.get_info(task_id).status).lower()
        if status == "draft":
            web.start(task_id)
            status = "started"
        elif status not in {"queued", "pre", "running", "post", "success"}:
            raise RuntimeError(
                f"Cannot start {solver} R={radius_um:g} um from status '{status}'."
            )
        key = f"{solver}_radius_{radius_tag(radius_um)}um"
        url = f"https://tidy3d.simulation.cloud/workbench?taskId={task_id}"
        print(f"{key}: {status} · {url}")
        return key, url

    with ThreadPoolExecutor(max_workers=len(selected)) as executor:
        urls = dict(executor.map(start_one, selected))
    manifest["started_radii_um"] = radii.tolist()
    manifest_path.write_text(json.dumps(manifest, indent=2))
    return urls


def _fdtd_transmission(path: Path) -> float:
    import h5py

    with h5py.File(path, "r") as file:
        data = file["data"]
        amp_groups = [
            data[name]["amps"]
            for name in sorted(data.keys(), key=int)
            if "amps" in data[name]
        ]
        if len(amp_groups) != 2:
            raise RuntimeError("Expected two FDTD modal monitor groups.")
        input_amp = amp_groups[0]["__xarray_dataarray_variable__"][0, 0, 0]
        output_amp = amp_groups[1]["__xarray_dataarray_variable__"][0, 0, 0]
    return float(np.abs(output_amp) ** 2 / np.abs(input_amp) ** 2)


def _eme_transmission(path: Path) -> float:
    import h5py

    with h5py.File(path, "r") as file:
        s21 = np.asarray(file["smatrix/S21/__xarray_dataarray_variable__"][:])
    if s21.shape != (1, 1, EME_NUM_MODES, EME_NUM_MODES):
        raise RuntimeError(f"Unexpected EME S21 shape: {s21.shape}.")
    return float(np.abs(s21[0, 0, 0, 0]) ** 2)


def analyze(
    manifest_path: Path = MANIFEST_PATH,
    output_dir: Path = OUTPUT_DIR,
    radii_um=RADIUS_VALUES_UM,
) -> dict[str, object]:
    """Download completed data and plot transmission and actual cost."""

    from tidy3d import web

    radii = _validated_radius_selection(radii_um)
    manifest = json.loads(manifest_path.read_text())
    output_dir.mkdir(parents=True, exist_ok=True)
    transmission = {"fdtd": [], "eme": []}
    actual_cost = {"fdtd": [], "eme": []}
    for solver in ("fdtd", "eme"):
        for radius_um in radii:
            item = manifest["tasks"][solver][radius_tag(radius_um)]
            path = Path(
                item.get("result_path", output_dir / f"{item['task_name']}.hdf5")
            )
            if not path.exists():
                path = output_dir / f"{item['task_name']}.hdf5"
                web.download(item["task_id"], path=str(path), verbose=True)
                item["result_path"] = str(path)
            extractor = _fdtd_transmission if solver == "fdtd" else _eme_transmission
            transmission[solver].append(extractor(path))
            cost = float(web.real_cost(item["task_id"]))
            actual_cost[solver].append(cost)
            item["actual_cost_flexcredits"] = cost

    fdtd_t = np.asarray(transmission["fdtd"])
    eme_t = np.asarray(transmission["eme"])
    fdtd_cost = np.asarray(actual_cost["fdtd"])
    eme_cost = np.asarray(actual_cost["eme"])
    signed_error_percent = (eme_t - fdtd_t) / fdtd_t * 100.0
    fdtd_loss_db = 10.0 * np.log10(np.clip(fdtd_t, 1e-20, None))
    eme_loss_db = 10.0 * np.log10(np.clip(eme_t, 1e-20, None))
    manifest["settings"]["eme_plane_um"] = [EME_PLANE_Y_UM, EME_PLANE_Z_UM]

    plt.rcParams.update(
        {
            "font.family": "DejaVu Sans",
            "font.size": 14,
            "axes.labelsize": 16,
            "legend.fontsize": 14,
            "xtick.labelsize": 14,
            "ytick.labelsize": 14,
        }
    )
    fig, ax = plt.subplots(figsize=(8.4, 5.4), constrained_layout=True)
    fig.patch.set_facecolor("white")
    blue = "#164E8C"
    red = "#C9364A"
    ax.set_facecolor("#FBFCFE")
    ax.grid(True, color="#B8C2CC", alpha=0.35, linewidth=0.8)
    ax.spines["top"].set_visible(False)
    ax.spines["right"].set_visible(False)
    ax.axhline(1.0, color="#64748B", linewidth=1.0, alpha=0.75)

    eme_line, = ax.plot(
        radii,
        eme_t,
        linewidth=2.7,
        marker="s",
        markersize=6.5,
        color=blue,
        zorder=2,
        label="EME",
    )
    fdtd_line, = ax.plot(
        radii,
        fdtd_t,
        linestyle=(0, (5, 3)),
        linewidth=2.5,
        marker="o",
        markersize=7,
        color=red,
        zorder=4,
        label="FDTD",
    )
    all_transmission = np.concatenate((fdtd_t, eme_t))
    margin = max(0.02, 0.05 * float(np.ptp(all_transmission)))
    ax.set_ylim(
        max(0.0, float(all_transmission.min()) - margin),
        min(1.02, max(1.0, float(all_transmission.max()) + margin)),
    )
    ax.set_xlabel("Effective radius (µm)")
    ax.set_ylabel("Fundamental-mode transmission, T")
    ax.set_xticks(radii)
    ax.legend(
        handles=[fdtd_line, eme_line],
        loc="lower right",
        frameon=True,
        framealpha=0.95,
    )
    transmission_plot_path = output_dir / f"{TASK_PREFIX}_transmission_vs_radius.png"
    fig.savefig(transmission_plot_path, dpi=220, bbox_inches="tight")
    plt.close(fig)

    fig, ax = plt.subplots(figsize=(8.4, 5.4), constrained_layout=True)
    fig.patch.set_facecolor("white")
    ax.set_facecolor("#FBFCFE")
    ax.grid(True, color="#B8C2CC", alpha=0.35, linewidth=0.8)
    ax.spines["top"].set_visible(False)
    ax.spines["right"].set_visible(False)

    eme_cost_line, = ax.plot(
        radii,
        eme_cost,
        linewidth=2.7,
        marker="s",
        markersize=6.5,
        color=blue,
        zorder=2,
        label="EME",
    )
    fdtd_cost_line, = ax.plot(
        radii,
        fdtd_cost,
        linestyle=(0, (5, 3)),
        linewidth=2.5,
        marker="o",
        markersize=7,
        color=red,
        zorder=4,
        label="FDTD",
    )
    for radius_um, cost in zip(radii, fdtd_cost):
        ax.annotate(
            f"{cost:.2f}",
            (radius_um, cost),
            xytext=(0, 8),
            textcoords="offset points",
            ha="center",
            va="bottom",
            fontsize=12,
            color=red,
        )
    for radius_um, cost in zip(radii, eme_cost):
        ax.annotate(
            f"{cost:.3f}",
            (radius_um, cost),
            xytext=(0, 8),
            textcoords="offset points",
            ha="center",
            va="bottom",
            fontsize=12,
            color=blue,
        )
    ax.set_ylim(0.0, 1.12 * float(max(fdtd_cost.max(), eme_cost.max())))
    ax.set_xlabel("Effective radius (µm)")
    ax.set_ylabel("Actual cost (FlexCredits)")
    ax.set_xticks(radii)
    ax.legend(
        handles=[fdtd_cost_line, eme_cost_line],
        loc="upper left",
        frameon=True,
        framealpha=0.95,
    )
    cost_plot_path = output_dir / f"{TASK_PREFIX}_actual_flexcredit_cost_vs_radius.png"
    fig.savefig(cost_plot_path, dpi=220, bbox_inches="tight")
    plt.close(fig)

    results = {
        "wavelength_um": WAVELENGTH_UM,
        "radius_um": radii.tolist(),
        "crystal_cut": "X-cut",
        "optic_axis_global": "y",
        "fdtd": {
            "run_time_quality_factor": FDTD_RUN_TIME_QUALITY_FACTOR,
            "transmission": fdtd_t.tolist(),
            "bend_loss_db_10log10_transmission": fdtd_loss_db.tolist(),
        },
        "eme": {
            "num_bend_cells": EME_NUM_CELLS,
            "num_modes_solved_and_propagated": EME_NUM_MODES,
            "plane_um": [EME_PLANE_Y_UM, EME_PLANE_Z_UM],
            "bend_medium_frame": "global",
            "transmission": eme_t.tolist(),
            "bend_loss_db_10log10_transmission": eme_loss_db.tolist(),
        },
        "signed_relative_error_percent": signed_error_percent.tolist(),
        "max_absolute_relative_error_percent": float(
            np.max(np.abs(signed_error_percent))
        ),
        "actual_cost_flexcredits": {
            "fdtd": fdtd_cost.tolist(),
            "eme": eme_cost.tolist(),
            "aggregate": float(fdtd_cost.sum() + eme_cost.sum()),
        },
        "tasks": manifest["tasks"],
    }
    json_path = output_dir / f"{TASK_PREFIX}_transmission_and_cost_vs_radius.json"
    json_path.write_text(json.dumps(results, indent=2))
    manifest_path.write_text(json.dumps(manifest, indent=2))
    print(f"Transmission plot: {transmission_plot_path}")
    print(f"Actual-cost plot: {cost_plot_path}")
    print(f"Data: {json_path}")
    return results


def print_build_summary(
    fdtd_models: dict[str, td.Simulation],
    eme_models: dict[str, td.EMESimulation],
) -> None:
    print("Models built locally.")
    print("  crystal: X-cut LiNbO3, optic axis along global y")
    print(f"  wavelength: {WAVELENGTH_UM:.4f} um")
    print(f"  radii: {RADIUS_VALUES_UM.tolist()} um")
    print(f"  mesh: wavelength/{STEPS_PER_WAVELENGTH}")
    print(f"  FDTD runtime quality factor: {FDTD_RUN_TIME_QUALITY_FACTOR:g}")
    print(f"  FDTD tasks: {len(fdtd_models)}")
    print(f"  EME tasks: {len(eme_models)}, N={EME_NUM_CELLS}, M={EME_NUM_MODES}")
    print(f"  EME plane: up to {EME_PLANE_Y_UM:g} x {EME_PLANE_Z_UM:g} um")
    for name, simulation in fdtd_models.items():
        size = tuple(round(value, 3) for value in simulation.size)
        print(f"  {name} domain: {size} um")
    mode_counts = {
        spec.num_modes
        for simulation in eme_models.values()
        for spec in simulation.eme_grid.mode_specs
    }
    print(f"  EME mode counts present: {sorted(mode_counts)}")


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "command",
        nargs="?",
        choices=("build", "estimate", "run", "analyze"),
        default="build",
    )
    parser.add_argument(
        "--radii",
        nargs="+",
        type=float,
        default=None,
        help=(
            "Selected radii. Defaults to all five radii for every command."
        ),
    )
    args = parser.parse_args()
    radii = (
        args.radii
        if args.radii is not None
        else (
            NEW_RADIUS_VALUES_UM.tolist()
            if args.command in {"estimate", "run"}
            else RADIUS_VALUES_UM.tolist()
        )
    )
    if args.command == "build":
        fdtd_models, eme_models = build_models(radii)
        print_build_summary(fdtd_models, eme_models)
    elif args.command == "estimate":
        upload_and_estimate(radii_um=radii)
    elif args.command == "run":
        run_uploaded_tasks(radii_um=radii)
    else:
        analyze(radii_um=radii)


if __name__ == "__main__":
    main()
