"""
YAML configuration, domain sizing, and pre-run checks.

The point of the pre-flight is that a 3D external-aero run is a 10-30 minute
commitment. Discovering afterwards that the geometry was rotated wrong, the
body was clipping a boundary, or the case needed four hours is the expensive
failure. Everything cheap enough to check in a second is checked before the
first timestep.

Runtime estimation
------------------
Wall time is predicted from a measured per-cell-per-step cost and the timestep
the CFL condition will actually pick. It is an estimate, not a promise -- but
it is derived from the same numbers the solver will use, so it is usually
within about 20%, which is enough to catch "this is a four-hour run" before
committing to it.
"""

import os

import numpy as np
import yaml

from .geometry import load_mesh, orient_mesh, place_in_domain, mesh_report
from .grid import Grid

# Measured on a Ryzen 9 5900X (12C/24T) against a 432k-cell 3D case: 127 ns
# per cell per step for the full step including the DCT pressure solve.
# Calibrated at 120 rather than the 2D figure (~80), because 3D pays for the
# w-momentum kernel and a third DCT axis, and under-promising a runtime is
# the failure that wastes an afternoon. Override per machine via
# run.cost_ns_per_cell_step.
DEFAULT_COST_NS = 120.0

DEFAULTS = {
    "geometry": {
        "path": None,
        "units": "mm",
        "scale": 1.0,
        "axis_order": "xyz",
        "rotate_deg": [0.0, 0.0, 0.0],
        "flip": [False, False, False],
    },
    "domain": {
        "cells_per_length": 40,
        "upstream": 1.5,
        "downstream": 4.0,
        "side": 1.5,
        "above": 2.0,
        "ground_clearance": None,
        "max_cells": 4_000_000,
    },
    "flow": {
        "speed": 30.0,
        "density": 1.225,
        "viscosity": 1.5e-5,
        "ground": True,
        "rolling_road": True,
        "sides": "farfield",
        # Subgrid model. "vreman" or null. Required above Re ~ 1e5: central
        # differencing supplies no dissipation of its own, so at vehicle
        # Reynolds numbers on an affordable grid there is nothing to remove
        # energy at the cutoff and the run goes unstable. Leave null only for
        # resolved low-Re verification cases.
        "sgs": "vreman",
    },
    "run": {
        "convective_times": 6.0,
        "end_time": None,
        "cfl": 0.4,
        "frames": 120,
        "save_fields": ["vorticity"],
        "output": "output/run",
        "time_budget_s": 1800,
        "cost_ns_per_cell_step": DEFAULT_COST_NS,
        "log_every": 200,
    },
    "view": {"elev": 22.0, "azim": -60.0},
}


def _merge(base, over):
    out = dict(base)
    for k, v in (over or {}).items():
        if isinstance(v, dict) and isinstance(out.get(k), dict):
            out[k] = _merge(out[k], v)
        else:
            out[k] = v
    return out


def load_config(path):
    with open(path, "r") as fh:
        user = yaml.safe_load(fh) or {}
    cfg = {k: _merge(v, user.get(k)) for k, v in DEFAULTS.items()}
    for k in user:
        if k not in cfg:
            cfg[k] = user[k]
    if not cfg["geometry"]["path"]:
        raise ValueError("config: geometry.path is required")
    return cfg


# ----------------------------------------------------------------------

def build_case(cfg, verbose=True):
    """
    Turn a config into (grid, mesh, plan). Does not allocate solver fields.
    """
    gcfg, dcfg, fcfg, rcfg = (
        cfg["geometry"], cfg["domain"], cfg["flow"], cfg["run"]
    )

    mesh = load_mesh(gcfg["path"], units=gcfg["units"], scale=gcfg["scale"])
    mesh = orient_mesh(
        mesh,
        axis_order=gcfg["axis_order"],
        rotate_deg=gcfg["rotate_deg"],
        flip=gcfg["flip"],
    )

    size = mesh.bounds[1] - mesh.bounds[0]
    length = float(size[0])  # streamwise extent after orientation
    if length <= 0:
        raise ValueError("geometry has zero streamwise extent after orientation")

    lx = length * (1.0 + dcfg["upstream"] + dcfg["downstream"])
    ly = float(size[1]) * (1.0 + dcfg["above"])
    lz = float(size[2]) * (1.0 + 2.0 * dcfg["side"])

    h = length / dcfg["cells_per_length"]
    nx = max(int(round(lx / h)), 8)
    ny = max(int(round(ly / h)), 8)
    nz = max(int(round(lz / h)), 8)

    # Uniform-grid cost is brutal in 3D, so honour the cell cap by coarsening
    # rather than by silently running for hours.
    n_cells = nx * ny * nz
    if n_cells > dcfg["max_cells"]:
        f = (dcfg["max_cells"] / n_cells) ** (1.0 / 3.0)
        nx, ny, nz = (max(int(nx * f), 8), max(int(ny * f), 8), max(int(nz * f), 8))
        h = lx / nx

    grid = Grid(nx, ny, nz, lx, ly, lz, origin=(0.0, 0.0, 0.0))

    x_fraction = dcfg["upstream"] / (1.0 + dcfg["upstream"] + dcfg["downstream"])
    mesh = place_in_domain(
        mesh, grid,
        ground_clearance=dcfg["ground_clearance"],
        x_fraction=x_fraction,
    )

    plan = estimate_runtime(grid, cfg, length)
    plan["body_length_m"] = length
    plan["reynolds"] = fcfg["speed"] * length / fcfg["viscosity"]
    plan["cells_per_length"] = length / grid.dx
    return grid, mesh, plan


def estimate_runtime(grid, cfg, length):
    fcfg, rcfg = cfg["flow"], cfg["run"]
    u = fcfg["speed"]

    # The convective CFL limit binds, and it is set by the *peak* local speed,
    # not the freestream. Potential-flow intuition says ~1.4x U over a smooth
    # bluff body, and that is badly optimistic here: the actual timestep on the
    # Model S case came out 1.8x smaller than a 1.4x estimate predicted,
    # implying local peaks near 2.5x U. Sharp immersed features -- mirrors,
    # spoiler edges, the wheel arches -- accelerate flow far more than a smooth
    # body, and the CFL number is set by the single worst cell in the domain,
    # not by a representative one.
    #
    # 2.5x is calibrated against that measurement. Erring high is the right
    # direction: an estimate that promises 20 minutes and delivers 40 wastes
    # an afternoon, while one that promises 40 and delivers 30 costs nothing.
    u_peak = 2.5 * u
    conv_rate = u_peak / grid.dx + 0.5 * u_peak / grid.dy + 0.5 * u_peak / grid.dz
    dt = rcfg["cfl"] / conv_rate

    visc_dt = 0.5 / (
        fcfg["viscosity"]
        * (1 / grid.dx**2 + 1 / grid.dy**2 + 1 / grid.dz**2)
    )
    dt = min(dt, visc_dt)

    if rcfg["end_time"] is not None:
        t_end = float(rcfg["end_time"])
    else:
        t_end = rcfg["convective_times"] * grid.lx / u

    steps = int(np.ceil(t_end / dt))
    cost = rcfg.get("cost_ns_per_cell_step", DEFAULT_COST_NS) * 1e-9
    seconds = steps * grid.n_cells * cost

    bytes_per_field = grid.n_cells * 8
    n_fields = 16  # u,v,w,p, 3 rhs x2, 4 eta, div, scratch
    return {
        "dt": dt,
        "t_end": t_end,
        "steps": steps,
        "n_cells": grid.n_cells,
        "estimated_seconds": seconds,
        "estimated_memory_gb": bytes_per_field * n_fields / 1e9,
        "flow_through_times": t_end * u / grid.lx,
        "within_budget": seconds <= rcfg.get("time_budget_s", 1800),
    }


# ----------------------------------------------------------------------

def preflight(cfg, grid, mesh, plan, strict=True, frontal_area=None):
    """
    Validate geometry and domain. Returns (ok, dict of findings).

    Errors block the run; warnings are printed and the run proceeds.

    frontal_area, when given, is the rasterised area of the immersed body and
    overrides the silhouette estimate. Pass it whenever it is available: the
    silhouette fills enclosed holes, which for a part with real through-gaps
    can overstate the area severalfold, and blockage warnings computed from
    the wrong area are worse than no warning at all.
    """
    errors, warnings = [], []
    rep = mesh_report(
        mesh, os.path.basename(cfg["geometry"]["path"]),
        frontal_area=frontal_area is None,
    )
    if frontal_area is not None:
        rep["frontal_area_m2"] = frontal_area

    if rep["faces"] == 0:
        errors.append("mesh has no faces")
    if not rep["watertight"]:
        warnings.append(
            "mesh is not watertight -- the voxel flood fill may leak into "
            "the interior. Check the solid-fraction preview before trusting "
            "forces."
        )
    if not rep["winding_consistent"]:
        warnings.append("mesh winding is inconsistent; normals may be flipped")

    lo, hi = mesh.bounds
    if lo[0] < 0 or hi[0] > grid.lx or lo[1] < 0 or hi[1] > grid.ly \
            or lo[2] < 0 or hi[2] > grid.lz:
        errors.append(
            f"geometry extends outside the domain: bounds {lo.round(3)} to "
            f"{hi.round(3)} vs domain (0,0,0)-({grid.lx:.3f},{grid.ly:.3f},"
            f"{grid.lz:.3f})"
        )

    size = hi - lo
    cpl = size[0] / grid.dx
    if cpl < 20:
        warnings.append(
            f"only {cpl:.0f} cells along the body. Below about 20 the "
            f"immersed boundary cannot resolve separation and drag will be "
            f"unreliable."
        )

    frontal = rep["frontal_area_m2"]
    blockage = 100.0 * frontal / (grid.ly * grid.lz)
    if blockage > 10:
        warnings.append(
            f"blockage is {blockage:.1f}% (frontal area vs domain cross "
            f"section). Above ~5% the confined freestream accelerates and "
            f"drag is overpredicted; widen domain.side / domain.above."
        )
    elif blockage > 5:
        warnings.append(f"blockage {blockage:.1f}%; expect a few percent of "
                        f"drag overprediction")

    if plan["estimated_memory_gb"] > 24:
        warnings.append(
            f"estimated {plan['estimated_memory_gb']:.1f} GB of field storage"
        )
    if not plan["within_budget"]:
        warnings.append(
            f"estimated runtime {plan['estimated_seconds']/60:.0f} min exceeds "
            f"the {cfg['run']['time_budget_s']/60:.0f} min budget. Reduce "
            f"domain.cells_per_length or run.convective_times."
        )

    ok = not errors or not strict
    return ok, {"errors": errors, "warnings": warnings, "mesh": rep,
                "blockage_pct": blockage}


def describe(cfg, grid, mesh, plan, checks):
    """Human-readable pre-flight summary."""
    rep = checks["mesh"]
    L = plan["body_length_m"]
    lines = [
        "=" * 68,
        f"geometry   {cfg['geometry']['path']}",
        f"           {rep['faces']:,} faces, watertight={rep['watertight']}, "
        f"size {np.round(rep['size_m'], 3).tolist()} m",
        f"           frontal area {rep['frontal_area_m2']:.4f} m^2, "
        f"blockage {checks['blockage_pct']:.1f}%",
        f"domain     {grid}",
        f"           {plan['cells_per_length']:.0f} cells along the body "
        f"(L = {L:.3f} m)",
        f"flow       U = {cfg['flow']['speed']} m/s, rho = "
        f"{cfg['flow']['density']} kg/m^3, nu = {cfg['flow']['viscosity']:g} m^2/s",
        f"           Re_L = {plan['reynolds']:.3g}",
        f"schedule   dt ~ {plan['dt']:.3e} s, t_end = {plan['t_end']:.3f} s, "
        f"{plan['steps']:,} steps",
        f"           {plan['flow_through_times']:.1f} domain flow-throughs",
        f"estimate   {plan['estimated_seconds']/60:.1f} min wall clock, "
        f"{plan['estimated_memory_gb']:.2f} GB fields",
        "=" * 68,
    ]
    for w in checks["warnings"]:
        lines.append(f"WARNING: {w}")
    for e in checks["errors"]:
        lines.append(f"ERROR:   {e}")
    return "\n".join(lines)
