from dataclasses import dataclass
from typing import Any, Mapping

from core import clamp, map_risk_to_target_u, require, require_unit_interval_open_closed


@dataclass
class ControllerState:
    u: float
    b_ewma: float


@dataclass
class ControllerOutput:
    u: float
    u_target: float
    b_eff: float
    risk_eff: float


class SmoothThermostatController:
    def __init__(self, cfg: Mapping[str, Any]):
        self.gamma = require_unit_interval_open_closed(require(cfg, "gamma", "controller"), "controller.gamma")
        self.pi_cfg = require(cfg, "pi", "controller")
        self.b_eff_mode = str(require(cfg, "b_eff_mode", "controller")).lower()

        u0 = float(require(cfg, "u0", "controller"))
        self.slew_cfg = cfg.get("slew", None)
        self.b_ewma_alpha = require_unit_interval_open_closed(cfg.get("b_ewma_alpha", 0.2), "controller.b_ewma_alpha")

        self.state = ControllerState(u=clamp(u0, 0.0, 1.0), b_ewma=0.5)

    def step(self, b: float) -> ControllerOutput:
        b_raw = clamp(float(b), 0.0, 1.0)
        self.state.b_ewma = (1.0 - self.b_ewma_alpha) * self.state.b_ewma + self.b_ewma_alpha * b_raw

        if self.b_eff_mode == "raw":
            b_eff = b_raw
        elif self.b_eff_mode == "ewma":
            b_eff = self.state.b_ewma
        else:
            raise ValueError(f"Unsupported controller.b_eff_mode: {self.b_eff_mode}")

        risk_eff = clamp(1.0 - b_eff, 0.0, 1.0)
        u_target = clamp(map_risk_to_target_u(risk_eff, self.pi_cfg), 0.0, 1.0)

        u_prev = float(self.state.u)
        u_next = u_prev + self.gamma * (u_target - u_prev)

        if self.slew_cfg is not None:
            max_delta = float(require(self.slew_cfg, "max_delta", "controller.slew"))
            du = clamp(u_next - u_prev, -max_delta, max_delta)
            u_next = u_prev + du

        u_next = clamp(u_next, 0.0, 1.0)
        self.state.u = float(u_next)

        return ControllerOutput(u=float(self.state.u), u_target=float(u_target), b_eff=float(b_eff), risk_eff=float(risk_eff))
