"""Driver-Seegmiller backward-facing-step SST validation journal."""

from __future__ import annotations

import sys
from pathlib import Path

import numpy as np

CASE_DIR = Path(__file__).resolve().parent
sys.path.insert(0, str(CASE_DIR.parents[3] / "server"))

from pytorsocae import TorsoCAESession

H = 1.0
DEPTH = 0.5 * H
UPSTREAM_LENGTH = 4.0 * H
DOWNSTREAM_LENGTH = 25.0 * H
TOTAL_HEIGHT = 9.0 * H

RE_H = 36_000.0
U_REF = 1.0
RHO = 1.0
MU = RHO * U_REF * H / RE_H
BOUNDARY_LAYER_THICKNESS = 1.5 * H
INLET_TURBULENCE_INTENSITY = 0.00061
INLET_TURBULENT_VISCOSITY_RATIO = 0.009
WALL_FIRST_CELL_HEIGHT = 1.5e-2 * H
SST_VARIANT = "sst_2003m"


def _minimum_expression(lhs: str, rhs: str) -> str:
    return f"(0.5*(({lhs})+({rhs})-abs(({lhs})-({rhs}))))"


def _inlet_velocity_expression() -> str:
    lower = f"(((y-{H})/{BOUNDARY_LAYER_THICKNESS})**(1.0/7.0))"
    upper = f"((({TOTAL_HEIGHT}-y)/{BOUNDARY_LAYER_THICKNESS})**(1.0/7.0))"
    return f"{U_REF}*{_minimum_expression(_minimum_expression(lower, '1.0'), _minimum_expression(upper, '1.0'))}"


def _geometric_axis(start: float, stop: float, cells: int, first: float) -> np.ndarray:
    """Return a one-sided geometrically graded axis clustered at ``start``."""
    length = float(stop - start)
    if cells <= 0 or not 0.0 < first <= length / cells:
        raise ValueError("Geometric axis requires cells > 0 and first <= uniform spacing")
    if np.isclose(first * cells, length):
        return np.linspace(start, stop, cells + 1)

    def covered(ratio: float) -> float:
        return first * np.expm1(cells * np.log(ratio)) / (ratio - 1.0)

    lo, hi = 1.0, 2.0
    while covered(hi) < length:
        hi *= 2.0
    for _ in range(80):
        mid = 0.5 * (lo + hi)
        if covered(mid) < length:
            lo = mid
        else:
            hi = mid
    ratio = 0.5 * (lo + hi)
    widths = first * ratio ** np.arange(cells, dtype=np.float64)
    widths *= length / float(np.sum(widths))
    axis = start + np.concatenate(([0.0], np.cumsum(widths)))
    axis[-1] = stop
    return axis


def _symmetric_geometric_axis(
    start: float,
    stop: float,
    cells: int,
    first: float,
) -> np.ndarray:
    """Return an even-cell axis clustered equally at both ends."""
    if cells <= 0 or cells % 2:
        raise ValueError("Symmetric geometric axis requires a positive even cell count")
    midpoint = 0.5 * (start + stop)
    lower = _geometric_axis(start, midpoint, cells // 2, first)
    upper = start + stop - lower[::-1]
    return np.concatenate((lower, upper[1:]))


def _structured_blocks() -> list[dict]:
    x_upstream = np.linspace(-UPSTREAM_LENGTH, 0.0, 25)
    x_downstream = np.concatenate(
        (
            np.linspace(0.0, 8.0 * H, 73),
            np.linspace(8.0 * H, 10.0 * H, 9)[1:],
            np.linspace(10.0 * H, DOWNSTREAM_LENGTH, 17)[1:],
        )
    )
    y_lower = _geometric_axis(0.0, H, 32, WALL_FIRST_CELL_HEIGHT)
    y_upper = _symmetric_geometric_axis(H, TOTAL_HEIGHT, 48, WALL_FIRST_CELL_HEIGHT)
    z = np.asarray([0.0, DEPTH], dtype=np.float64)

    def block(x, y, surface_tags):
        return {
            "x": x.tolist(),
            "y": y.tolist(),
            "z": z.tolist(),
            "surface_tags": surface_tags,
            "volume_tag": 1,
        }

    return [
        block(
            x_upstream,
            y_upper,
            {"xmin": 1, "xmax": 90, "ymin": 2, "ymax": 4, "zmin": 5, "zmax": 3},
        ),
        block(
            x_downstream,
            y_lower,
            {"xmin": 6, "xmax": 7, "ymin": 8, "ymax": 91, "zmin": 5, "zmax": 3},
        ),
        block(
            x_downstream,
            y_upper,
            {"xmin": 92, "xmax": 7, "ymin": 93, "ymax": 4, "zmin": 5, "zmax": 3},
        ),
    ]


def build_session() -> TorsoCAESession:
    session = TorsoCAESession()
    session.build_multiblock_structured_mesh(
        _structured_blocks(),
        mesh_id="mesh_0",
        name="BFS Structured Hex Mesh",
        volume_name="BFS Channel",
    )

    session.solid("BFS Channel").material(rho=RHO, mu=MU, n_seeds=40)
    session.set_physics("cfd", submodel="rans_k_omega_sst", backend="dolfinx")
    session.set_model_options(
        sst_variant=SST_VARIANT,
        turbulence_intensity=INLET_TURBULENCE_INTENSITY,
        turbulent_viscosity_ratio=INLET_TURBULENT_VISCOSITY_RATIO,
        wall_treatment="wall_function",
    )

    session.surface(1).bc(
        "inlet_velocity",
        values=[_inlet_velocity_expression(), "0.0", "0.0"],
    )
    session.surface(2).bc("wall")
    session.surface(3).bc("symmetry", values=[0.0], components=["z"])
    session.surface(4).bc("wall")
    session.surface(5).bc("symmetry", values=[0.0], components=["z"])
    session.surface(6).bc("wall")
    session.surface(7).bc("outlet_pressure", values=[0.0])
    session.surface(8).bc("wall")

    session.set_solver_options(
        algo="auto",
        precond="schur",
        tol=1.0e-7,
        max_iter=800,
        # Driver-Seegmiller is a steady mean-flow validation. pseudo_dt is a
        # continuation parameter, not a physical turbulence-resolution step.
        steady_rans=True,
        pseudo_dt=0.05,
        pseudo_dt_max=0.1,
        pseudo_dt_growth=1.01,
        num_steps=1500,
        max_inner_iter=1,
        inner_tol=1.0e-3,
        steady_tol=5.0e-4,
        steady_min_steps=100,
        steady_consecutive_steps=5,
        live_viz_interval=100,
        n_cores=4,
        device="cpu",
    )
    return session


def main() -> dict:
    return build_session().compute(mesh_ids=["mesh_0"])


if __name__ == "__main__":
    print(main())
