from typing import Any, Mapping

from core import clamp, require, require_unit_interval_closed_open, require_unit_interval_open_closed
from dynamics import DynamicsParams


def dynamics_params_from_u(
    u: float,
    mapping_cfg: Mapping[str, Any],
    dynamics_cfg: Mapping[str, Any],
    *,
    b_mu: float = 0.0,
    b_v: float = 0.0,
    mu_hat: float = 1.0,
) -> DynamicsParams:
    beta_min = float(require(mapping_cfg, "beta_min", "mapping"))
    beta_max = float(require(mapping_cfg, "beta_max", "mapping"))
    p_flip_min = float(require(mapping_cfg, "p_flip_min", "mapping"))
    p_flip_max = float(require(mapping_cfg, "p_flip_max", "mapping"))
    theta_base = float(require(dynamics_cfg, "theta", "dynamics"))

    if beta_min <= 0.0 or beta_max <= 0.0:
        raise ValueError("mapping.beta_min and mapping.beta_max must be > 0")

    uu = clamp(float(u), 0.0, 1.0)
    use_decomp = bool(mapping_cfg.get("use_constraint_decomp", False))
    if use_decomp:
        b_mu_c = clamp(float(b_mu), 0.0, 1.0)
        b_v_c = clamp(float(b_v), 0.0, 1.0)
        u_mu = clamp(uu * (1.0 - b_mu_c), 0.0, 1.0)
        u_v = clamp(uu * (1.0 - b_v_c), 0.0, 1.0)
    else:
        u_mu = uu
        u_v = uu

    beta = (1.0 - u_v) * beta_max + u_v * beta_min
    p_flip_drive = str(mapping_cfg.get("p_flip_drive", "v")).lower()
    if p_flip_drive == "v":
        u_flip = u_v
    elif p_flip_drive == "mu":
        u_flip = u_mu
    elif p_flip_drive == "max":
        u_flip = u_v if u_v >= u_mu else u_mu
    else:
        raise ValueError(f"Unsupported mapping.p_flip_drive: {p_flip_drive}")

    p_flip_escape = mapping_cfg.get("p_flip_escape")
    if p_flip_escape:
        u_above = require_unit_interval_closed_open(
            require(p_flip_escape, "u_above", "mapping.p_flip_escape"),
            "mapping.p_flip_escape.u_above",
        )

        if "b_mu_below" in p_flip_escape:
            b_mu_below = require_unit_interval_open_closed(
                require(p_flip_escape, "b_mu_below", "mapping.p_flip_escape"),
                "mapping.p_flip_escape.b_mu_below",
            )
            b_mu_c = clamp(float(b_mu), 0.0, 1.0)
            mu_gate = clamp((b_mu_below - b_mu_c) / b_mu_below, 0.0, 1.0)
        else:
            mu_hat_below = require_unit_interval_open_closed(
                require(p_flip_escape, "mu_hat_below", "mapping.p_flip_escape"),
                "mapping.p_flip_escape.mu_hat_below",
            )
            mu_hat_c = clamp(float(mu_hat), 0.0, 1.0)
            mu_gate = clamp((mu_hat_below - mu_hat_c) / mu_hat_below, 0.0, 1.0)
        u_gate = clamp((uu - u_above) / (1.0 - u_above), 0.0, 1.0)
        u_escape = float(mu_gate * u_gate)
        if u_escape > u_flip:
            u_flip = u_escape

    p_flip = p_flip_min + u_flip * (p_flip_max - p_flip_min)
    p_flip = clamp(p_flip, 0.0, 1.0)

    theta_mode = str(mapping_cfg.get("theta_mode", "fixed")).lower()
    if theta_mode == "fixed":
        theta = theta_base
    elif theta_mode == "linear":
        theta_max = float(mapping_cfg.get("theta_max", theta_base))
        theta_min = float(mapping_cfg.get("theta_min", -theta_max))
        theta = (1.0 - u_mu) * theta_max + u_mu * theta_min
    else:
        raise ValueError(f"Unsupported mapping.theta_mode: {theta_mode}")

    return DynamicsParams(beta=float(beta), p_flip=float(p_flip), theta=float(theta))
