"""PPO in NumPy: diagonal-Gaussian policy, GAE, clipped surrogate."""

import json
import numpy as np

from nets import MLP, Adam, RunningNorm

LOG2PI = float(np.log(2 * np.pi))


class Policy:
    def __init__(self, obs_dim, act_dim, hidden=(64, 64), init_log_std=-0.5, seed=0):
        rng = np.random.default_rng(seed)
        self.pi = MLP([obs_dim, *hidden, act_dim], out_gain=0.01, rng=rng)
        self.vf = MLP([obs_dim, *hidden, 1], out_gain=1.0, rng=rng)
        self.log_std = np.full(act_dim, float(init_log_std))
        self.norm = RunningNorm(obs_dim)
        self.obs_dim, self.act_dim = obs_dim, act_dim

    def act(self, obs, rng, deterministic=False):
        x = self.norm(obs)
        mu = self.pi.forward(x)
        v = self.vf.forward(x)[:, 0]
        if deterministic:
            return mu, v, np.zeros(len(obs))
        std = np.exp(self.log_std)
        a = mu + std * rng.normal(size=mu.shape)
        return a, v, self.logp(mu, a)

    def logp(self, mu, a):
        std = np.exp(self.log_std)
        return -0.5 * np.sum(((a - mu) / std) ** 2 + 2 * self.log_std + LOG2PI, axis=1)

    def value(self, obs):
        return self.vf.forward(self.norm(obs))[:, 0]

    def save(self, path, meta=None):
        with open(path, "w") as f:
            json.dump({
                "pi": self.pi.state(), "vf": self.vf.state(),
                "log_std": self.log_std.tolist(), "norm": self.norm.state(),
                "obs_dim": self.obs_dim, "act_dim": self.act_dim,
                "meta": meta or {},
            }, f)

    @classmethod
    def load(cls, path):
        with open(path) as f:
            s = json.load(f)
        o = cls(s["obs_dim"], s["act_dim"])
        o.pi = MLP.load(s["pi"])
        o.vf = MLP.load(s["vf"])
        o.log_std = np.array(s["log_std"])
        o.norm = RunningNorm.load(s["norm"])
        o.meta = s.get("meta", {})
        return o


def gae(rew, val, term, trunc, last_val, gamma, lam):
    T, N = rew.shape
    adv = np.zeros((T, N))
    running = np.zeros(N)
    next_val = last_val
    for t in reversed(range(T)):
        done = term[t] | trunc[t]
        nv = np.where(term[t], 0.0, next_val)
        delta = rew[t] + gamma * nv - val[t]
        running = delta + gamma * lam * np.where(done, 0.0, running)
        adv[t] = running
        next_val = val[t]
    return adv, adv + val


class PPO:
    def __init__(self, policy, lr=3e-4, clip=0.2, epochs=8, minibatches=4,
                 ent_coef=0.004, vf_coef=0.5, target_kl=0.02, max_grad_norm=0.5):
        self.p = policy
        self.opt_pi = Adam(policy.pi.params() + [policy.log_std], lr=lr)
        self.opt_vf = Adam(policy.vf.params(), lr=lr)
        self.clip, self.epochs, self.mb = clip, epochs, minibatches
        self.ent_coef, self.vf_coef = ent_coef, vf_coef
        self.target_kl, self.max_grad_norm = target_kl, max_grad_norm

    def update(self, obs, act, logp_old, adv, ret, rng):
        n = len(obs)
        x = self.p.norm(obs)
        adv = (adv - adv.mean()) / (adv.std() + 1e-8)
        idx = np.arange(n)
        mb_size = n // self.mb
        stats = {"kl": 0.0, "clipfrac": 0.0, "pg": 0.0, "vf": 0.0, "n": 0}

        for _ in range(self.epochs):
            rng.shuffle(idx)
            for s in range(0, n, mb_size):
                b = idx[s:s + mb_size]
                xb, ab = x[b], act[b]

                cache = []
                mu = self.p.pi.forward(xb, cache)
                std = np.exp(self.p.log_std)
                z = (ab - mu) / std
                logp = -0.5 * np.sum(z ** 2 + 2 * self.p.log_std + LOG2PI, axis=1)
                ratio = np.exp(np.clip(logp - logp_old[b], -20, 20))
                a_b = adv[b]

                unclipped = ratio * a_b
                clipped = np.clip(ratio, 1 - self.clip, 1 + self.clip) * a_b
                use_unclipped = unclipped <= clipped
                # d(-min)/d(logp) = -ratio*A where the unclipped branch wins.
                dlogp = np.where(use_unclipped, -ratio * a_b, 0.0) / len(b)

                dmu = dlogp[:, None] * (z / std)
                dlog_std = np.sum(dlogp[:, None] * (z ** 2 - 1.0), axis=0)
                dlog_std -= self.ent_coef * np.ones(self.p.act_dim)  # entropy bonus

                gW, gb = self.p.pi.backward(cache, dmu)
                self.opt_pi.step(gW + gb + [dlog_std], self.max_grad_norm)

                vcache = []
                v = self.p.vf.forward(xb, vcache)[:, 0]
                dv = (self.vf_coef * 2.0 * (v - ret[b]) / len(b))[:, None]
                vgW, vgb = self.p.vf.backward(vcache, dv)
                self.opt_vf.step(vgW + vgb, self.max_grad_norm)

                approx_kl = float(np.mean((ratio - 1) - np.log(ratio + 1e-12)))
                stats["kl"] += approx_kl
                stats["clipfrac"] += float(np.mean(np.abs(ratio - 1) > self.clip))
                stats["pg"] += float(-np.mean(np.minimum(unclipped, clipped)))
                stats["vf"] += float(np.mean((v - ret[b]) ** 2))
                stats["n"] += 1

            if stats["kl"] / max(1, stats["n"]) > self.target_kl * 1.5:
                break

        k = max(1, stats["n"])
        return {m: stats[m] / k for m in ("kl", "clipfrac", "pg", "vf")}
