"""Vectorised envs, in-process or across worker processes.

Auto-resets on episode end. A truncated episode returns its final observation
in `final_obs` so PPO can bootstrap the value instead of pretending the world
ended at the timer.
"""

import multiprocessing as mp
import numpy as np

from env import RobotEnv


class SyncVecEnv:
    def __init__(self, spec_path, n, seed=0, randomise=None, cmd_scale=1.0,
                 task=None, layout="union"):
        self.envs = [RobotEnv(spec_path, seed=seed + i, randomise=randomise,
                              cmd_scale=cmd_scale, task=task, layout=layout)
                     for i in range(n)]
        self.task_name = self.envs[0].task_name
        self.n = n
        self.obs_dim = self.envs[0].obs_dim
        self.act_dim = self.envs[0].act_dim

    def set_cmd_scale(self, s):
        for e in self.envs:
            e.cmd_scale = s

    def reset(self):
        return np.stack([e.reset() for e in self.envs])

    def step(self, actions):
        obs = np.empty((self.n, self.obs_dim))
        rew = np.empty(self.n)
        term = np.zeros(self.n, dtype=bool)
        trunc = np.zeros(self.n, dtype=bool)
        final = {}
        for i, e in enumerate(self.envs):
            o, r, f, tr, info = e.step(actions[i])
            rew[i], term[i], trunc[i] = r, f, tr
            if f or tr:
                if tr and not f:
                    final[i] = o
                o = e.reset()
            obs[i] = o
        return obs, rew, term, trunc, final

    def close(self):
        pass


def _worker(remote, spec_path, n, seed, randomise, cmd_scale, task, layout):
    venv = SyncVecEnv(spec_path, n, seed=seed, randomise=randomise,
                      cmd_scale=cmd_scale, task=task, layout=layout)
    try:
        while True:
            cmd, payload = remote.recv()
            if cmd == "reset":
                remote.send(venv.reset())
            elif cmd == "step":
                remote.send(venv.step(payload))
            elif cmd == "cmd_scale":
                venv.set_cmd_scale(payload)
            elif cmd == "close":
                break
    except (EOFError, KeyboardInterrupt):
        pass
    finally:
        remote.close()


class AsyncVecEnv:
    def __init__(self, spec_path, n_workers, envs_per_worker, seed=0,
                 randomise=None, cmd_scale=1.0, task=None, layout="union"):
        ctx = mp.get_context("spawn")
        self.n_workers = n_workers
        self.k = envs_per_worker
        self.n = n_workers * envs_per_worker
        self.remotes, self.procs = [], []
        for w in range(n_workers):
            parent, child = ctx.Pipe()
            p = ctx.Process(target=_worker,
                            args=(child, spec_path, envs_per_worker,
                                  seed + w * 1000, randomise, cmd_scale, task,
                                  layout),
                            daemon=True)
            p.start()
            child.close()
            self.remotes.append(parent)
            self.procs.append(p)
        probe = RobotEnv(spec_path, seed=seed, task=task, layout=layout)
        self.obs_dim, self.act_dim = probe.obs_dim, probe.act_dim
        self.task_name = probe.task_name

    def set_cmd_scale(self, s):
        for r in self.remotes:
            r.send(("cmd_scale", s))

    def reset(self):
        for r in self.remotes:
            r.send(("reset", None))
        return np.concatenate([r.recv() for r in self.remotes])

    def step(self, actions):
        chunks = actions.reshape(self.n_workers, self.k, -1)
        for r, c in zip(self.remotes, chunks):
            r.send(("step", c))
        obs, rew, term, trunc, final = [], [], [], [], {}
        for w, r in enumerate(self.remotes):
            o, rw, t, tr, f = r.recv()
            obs.append(o); rew.append(rw); term.append(t); trunc.append(tr)
            for i, v in f.items():
                final[w * self.k + i] = v
        return (np.concatenate(obs), np.concatenate(rew),
                np.concatenate(term), np.concatenate(trunc), final)

    def close(self):
        for r in self.remotes:
            try:
                r.send(("close", None))
            except (BrokenPipeError, OSError):
                pass
        for p in self.procs:
            p.join(timeout=2)


def make_vec(spec_path, n_envs, n_workers=1, seed=0, randomise=None,
             cmd_scale=1.0, task=None, layout="union"):
    if n_workers <= 1:
        return SyncVecEnv(spec_path, n_envs, seed=seed, randomise=randomise,
                          cmd_scale=cmd_scale, task=task, layout=layout)
    if n_envs % n_workers:
        raise ValueError("n_envs must divide evenly across workers")
    return AsyncVecEnv(spec_path, n_workers, n_envs // n_workers, seed=seed,
                       randomise=randomise, cmd_scale=cmd_scale, task=task,
                       layout=layout)
