"""Visualize a trained PPO agent on Gymnasium's CartPole environment."""

from __future__ import annotations

import argparse
from pathlib import Path

import gymnasium as gym
import torch

from train import ActorCritic, resolve_device


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--checkpoint",
        type=Path,
        default=Path(__file__).with_name("ppo_cartpole.pt"),
    )
    parser.add_argument("--episodes", type=int, default=5)
    parser.add_argument("--seed", type=int, default=1)
    parser.add_argument("--device", choices=("cpu", "cuda", "auto"), default="cpu")
    return parser.parse_args()


def load_checkpoint(path: Path, device: torch.device) -> dict:
    if not path.is_file():
        raise FileNotFoundError(
            f"Checkpoint not found: {path}. Run train.py first or pass --checkpoint."
        )
    try:
        checkpoint = torch.load(path, map_location=device, weights_only=True)
    except TypeError:
        checkpoint = torch.load(path, map_location=device)
    required = {"model_state_dict", "env_id", "obs_dim", "action_dim", "hidden_size"}
    missing = required.difference(checkpoint)
    if missing:
        raise ValueError(f"Checkpoint is missing required fields: {sorted(missing)}")
    return checkpoint


def infer(args: argparse.Namespace) -> None:
    if args.episodes <= 0:
        raise ValueError("--episodes must be positive")

    device = resolve_device(args.device)
    checkpoint = load_checkpoint(args.checkpoint, device)
    model = ActorCritic(
        obs_dim=int(checkpoint["obs_dim"]),
        action_dim=int(checkpoint["action_dim"]),
        hidden_size=int(checkpoint["hidden_size"]),
    ).to(device)
    model.load_state_dict(checkpoint["model_state_dict"])
    model.eval()

    env = gym.make(checkpoint["env_id"], render_mode="human")
    try:
        for episode in range(1, args.episodes + 1):
            observation, _ = env.reset(seed=args.seed + episode - 1)
            episode_return = 0.0
            episode_length = 0
            done = False

            while not done:
                obs_tensor = torch.as_tensor(
                    observation, dtype=torch.float32, device=device
                ).unsqueeze(0)
                with torch.no_grad():
                    logits, _ = model(obs_tensor)
                    action = logits.argmax(dim=-1).item()
                observation, reward, terminated, truncated, _ = env.step(action)
                done = terminated or truncated
                episode_return += float(reward)
                episode_length += 1

            print(
                f"episode={episode:3d} return={episode_return:6.1f} "
                f"length={episode_length:3d}"
            )
    finally:
        env.close()


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