"""Train a PyTorch PPO agent on Gymnasium's CartPole-v1 environment."""

from __future__ import annotations

import argparse
import random
from collections import deque
from pathlib import Path

import gymnasium as gym
import numpy as np
import torch
from torch import nn
from torch.distributions import Categorical


class ActorCritic(nn.Module):
    """Small actor-critic network for environments with discrete actions."""

    def __init__(self, obs_dim: int, action_dim: int, hidden_size: int = 64) -> None:
        super().__init__()
        self.backbone = nn.Sequential(
            nn.Linear(obs_dim, hidden_size),
            nn.Tanh(),
            nn.Linear(hidden_size, hidden_size),
            nn.Tanh(),
        )
        self.actor = nn.Linear(hidden_size, action_dim)
        self.critic = nn.Linear(hidden_size, 1)
        self.apply(self._init_weights)

    @staticmethod
    def _init_weights(module: nn.Module) -> None:
        if isinstance(module, nn.Linear):
            nn.init.orthogonal_(module.weight, gain=np.sqrt(2))
            nn.init.zeros_(module.bias)

    def forward(self, observations: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
        features = self.backbone(observations)
        return self.actor(features), self.critic(features).squeeze(-1)

    def distribution_and_value(
        self, observations: torch.Tensor
    ) -> tuple[Categorical, torch.Tensor]:
        logits, values = self(observations)
        return Categorical(logits=logits), values


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--env-id", default="CartPole-v1")
    parser.add_argument("--total-timesteps", type=int, default=100_000)
    parser.add_argument("--rollout-steps", type=int, default=2_048)
    parser.add_argument("--learning-rate", type=float, default=3e-4)
    parser.add_argument("--gamma", type=float, default=0.99)
    parser.add_argument("--gae-lambda", type=float, default=0.95)
    parser.add_argument("--clip-coef", type=float, default=0.2)
    parser.add_argument("--update-epochs", type=int, default=10)
    parser.add_argument("--minibatch-size", type=int, default=64)
    parser.add_argument("--value-coef", type=float, default=0.5)
    parser.add_argument("--entropy-coef", type=float, default=0.01)
    parser.add_argument("--max-grad-norm", type=float, default=0.5)
    parser.add_argument("--hidden-size", type=int, default=64)
    parser.add_argument("--seed", type=int, default=1)
    parser.add_argument("--device", choices=("cpu", "cuda", "auto"), default="cpu")
    parser.add_argument(
        "--checkpoint",
        type=Path,
        default=Path(__file__).with_name("ppo_cartpole.pt"),
    )
    return parser.parse_args()


def resolve_device(requested: str) -> torch.device:
    if requested == "auto":
        return torch.device("cuda" if torch.cuda.is_available() else "cpu")
    if requested == "cuda" and not torch.cuda.is_available():
        raise RuntimeError("CUDA was requested, but no CUDA device is available")
    return torch.device(requested)


def validate_args(args: argparse.Namespace) -> None:
    positive = {
        "total_timesteps": args.total_timesteps,
        "rollout_steps": args.rollout_steps,
        "update_epochs": args.update_epochs,
        "minibatch_size": args.minibatch_size,
        "hidden_size": args.hidden_size,
    }
    for name, value in positive.items():
        if value <= 0:
            raise ValueError(f"--{name.replace('_', '-')} must be positive")


def train(args: argparse.Namespace) -> Path:
    validate_args(args)
    random.seed(args.seed)
    np.random.seed(args.seed)
    torch.manual_seed(args.seed)

    device = resolve_device(args.device)
    env = gym.make(args.env_id)
    env.action_space.seed(args.seed)

    if not isinstance(env.observation_space, gym.spaces.Box) or len(
        env.observation_space.shape
    ) != 1:
        env.close()
        raise TypeError("PPO trainer requires a one-dimensional Box observation space")
    if not isinstance(env.action_space, gym.spaces.Discrete):
        env.close()
        raise TypeError("PPO trainer requires a Discrete action space")

    obs_dim = int(env.observation_space.shape[0])
    action_dim = int(env.action_space.n)
    model = ActorCritic(obs_dim, action_dim, args.hidden_size).to(device)
    optimizer = torch.optim.Adam(model.parameters(), lr=args.learning_rate, eps=1e-5)

    observation, _ = env.reset(seed=args.seed)
    global_step = 0
    episode_return = 0.0
    episode_length = 0
    episode_count = 0
    recent_returns: deque[float] = deque(maxlen=20)

    try:
        while global_step < args.total_timesteps:
            steps = min(args.rollout_steps, args.total_timesteps - global_step)
            observations: list[np.ndarray] = []
            actions: list[int] = []
            old_log_probs: list[float] = []
            rewards: list[float] = []
            values: list[float] = []
            next_values: list[float] = []
            terminated_flags: list[float] = []
            done_flags: list[float] = []

            for _ in range(steps):
                obs_tensor = torch.as_tensor(
                    observation, dtype=torch.float32, device=device
                ).unsqueeze(0)
                with torch.no_grad():
                    distribution, value = model.distribution_and_value(obs_tensor)
                    action = distribution.sample()
                    log_prob = distribution.log_prob(action)

                next_observation, reward, terminated, truncated, _ = env.step(
                    action.item()
                )
                done = terminated or truncated
                next_obs_tensor = torch.as_tensor(
                    next_observation, dtype=torch.float32, device=device
                ).unsqueeze(0)
                with torch.no_grad():
                    _, next_value = model(next_obs_tensor)

                observations.append(np.asarray(observation, dtype=np.float32))
                actions.append(action.item())
                old_log_probs.append(log_prob.item())
                rewards.append(float(reward))
                values.append(value.item())
                next_values.append(next_value.item())
                terminated_flags.append(float(terminated))
                done_flags.append(float(done))

                global_step += 1
                episode_return += float(reward)
                episode_length += 1
                observation = next_observation

                if done:
                    episode_count += 1
                    recent_returns.append(episode_return)
                    mean_return = float(np.mean(recent_returns))
                    print(
                        f"step={global_step:6d} episode={episode_count:4d} "
                        f"return={episode_return:6.1f} length={episode_length:3d} "
                        f"mean_20={mean_return:6.1f}"
                    )
                    observation, _ = env.reset()
                    episode_return = 0.0
                    episode_length = 0

            obs_batch = torch.as_tensor(
                np.asarray(observations), dtype=torch.float32, device=device
            )
            action_batch = torch.as_tensor(actions, dtype=torch.long, device=device)
            old_log_prob_batch = torch.as_tensor(
                old_log_probs, dtype=torch.float32, device=device
            )
            reward_batch = torch.as_tensor(rewards, dtype=torch.float32, device=device)
            value_batch = torch.as_tensor(values, dtype=torch.float32, device=device)
            next_value_batch = torch.as_tensor(
                next_values, dtype=torch.float32, device=device
            )
            terminated_batch = torch.as_tensor(
                terminated_flags, dtype=torch.float32, device=device
            )
            done_batch = torch.as_tensor(done_flags, dtype=torch.float32, device=device)

            advantages = torch.zeros_like(reward_batch)
            gae = torch.zeros((), dtype=torch.float32, device=device)
            for index in reversed(range(steps)):
                delta = (
                    reward_batch[index]
                    + args.gamma
                    * next_value_batch[index]
                    * (1.0 - terminated_batch[index])
                    - value_batch[index]
                )
                gae = (
                    delta
                    + args.gamma
                    * args.gae_lambda
                    * (1.0 - done_batch[index])
                    * gae
                )
                advantages[index] = gae
            returns = advantages + value_batch
            advantages = (advantages - advantages.mean()) / (
                advantages.std(unbiased=False) + 1e-8
            )

            for _ in range(args.update_epochs):
                permutation = torch.randperm(steps, device=device)
                for start in range(0, steps, args.minibatch_size):
                    indices = permutation[start : start + args.minibatch_size]
                    distribution, predicted_values = model.distribution_and_value(
                        obs_batch[indices]
                    )
                    new_log_probs = distribution.log_prob(action_batch[indices])
                    log_ratio = new_log_probs - old_log_prob_batch[indices]
                    ratio = log_ratio.exp()

                    policy_loss_unclipped = -advantages[indices] * ratio
                    policy_loss_clipped = -advantages[indices] * torch.clamp(
                        ratio, 1.0 - args.clip_coef, 1.0 + args.clip_coef
                    )
                    policy_loss = torch.maximum(
                        policy_loss_unclipped, policy_loss_clipped
                    ).mean()
                    value_loss = 0.5 * (
                        predicted_values - returns[indices]
                    ).pow(2).mean()
                    entropy = distribution.entropy().mean()
                    loss = (
                        policy_loss
                        + args.value_coef * value_loss
                        - args.entropy_coef * entropy
                    )

                    optimizer.zero_grad()
                    loss.backward()
                    nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm)
                    optimizer.step()
    finally:
        env.close()

    args.checkpoint.parent.mkdir(parents=True, exist_ok=True)
    training_config = vars(args).copy()
    training_config["checkpoint"] = str(args.checkpoint)
    torch.save(
        {
            "model_state_dict": model.state_dict(),
            "env_id": args.env_id,
            "obs_dim": obs_dim,
            "action_dim": action_dim,
            "hidden_size": args.hidden_size,
            "training_config": training_config,
            "global_step": global_step,
        },
        args.checkpoint,
    )
    print(f"Saved checkpoint to {args.checkpoint}")
    return args.checkpoint


if __name__ == "__main__":
    train(parse_args())
