import math
import numpy as np
from typing import Dict, List, Optional


class SalienceEngine:
    """
    SAL-v4 Salience Math Engine (S').
    Calculates the trajectory action S' based on Novelty, Retention, Meaning, Continuity, and Fatigue.
    """

    def __init__(self, weights: Optional[Dict[str, float]] = None):
        self.w = weights or {
            "w_A": 1.0,  # Novelty
            "w_R": 0.5,  # Retention
            "w_M": 2.0,  # Meaning/Payoff
            "k_phi": 0.1,  # Fatigue constant
            "lambda_t": 0.05,  # Time decay
            "lambda_thr": 1.0,  # Continuity threshold penalty
            "lambda_jerk": 0.5,  # Strategy jerk penalty
        }

    def calculate_l_sal(
        self, delta_a: float, r: float, m: float, c: float, phi: float, t: float
    ) -> float:
        """
        Calculates the instantaneous Lagrangian L_sal.
        """
        # Fatigue gate (Endocrine-style suppression)
        fatigue_gate = math.exp(-self.w["k_phi"] * phi)

        # Reward component
        reward = self.w["w_A"] * delta_a + self.w["w_R"] * r + self.w["w_M"] * m

        # Continuity-gated reward with time decay
        l_core = reward * c * fatigue_gate * math.exp(-self.w["lambda_t"] * t)

        return l_core

    def calculate_trajectory_action(self, trajectory: List[Dict]) -> float:
        """
        Calculates S' = sum(L_sal) - penalties.
        trajectory: List of steps, each with delta_a, r, m, c, phi, t, state_vec
        """
        s_prime = 0.0
        prev_c = 1.0
        prev_state_vec = None
        prev_velocity = None

        for i, step in enumerate(trajectory):
            # 1. Instantaneous L_sal
            l_t = self.calculate_l_sal(
                step["delta_a"], step["r"], step["m"], step["c"], step["phi"], step["t"]
            )

            # 2. Continuity Penalty (Throttling/Thrashing)
            penalty_thrash = self.w["lambda_thr"] * (step["c"] - prev_c) ** 2

            # 3. Jerk Penalty (Second-order change in strategy/state)
            penalty_jerk = 0.0
            if prev_state_vec is not None:
                current_velocity = step["state_vec"] - prev_state_vec
                if prev_velocity is not None:
                    acceleration = current_velocity - prev_velocity
                    penalty_jerk = (
                        self.w["lambda_jerk"] * np.linalg.norm(acceleration) ** 2
                    )
                prev_velocity = current_velocity
            prev_state_vec = step["state_vec"]
            prev_c = step["c"]

            s_prime += l_t - penalty_thrash - penalty_jerk

        return s_prime

    def estimate_uncertainty(self, s_prime: float, budget_spent: float) -> float:
        """
        Calculates the 'Escalation Trigger' signal.
        Low salience relative to cost = High uncertainty.
        """
        if budget_spent <= 0 or s_prime <= 0:
            return 0.0
        efficiency = s_prime / budget_spent
        efficiency = max(-10, min(10, efficiency))
        return 1.0 / (1.0 + math.exp(efficiency))
