198 lines
6.6 KiB
Python
198 lines
6.6 KiB
Python
from __future__ import annotations
|
|
|
|
import gymnasium as gym
|
|
import numpy as np
|
|
|
|
from blokus_gym import (
|
|
STANDARD_PIECES,
|
|
ActionMaskWrapper,
|
|
BlokusEnv,
|
|
BlokusMultiAgentEnv,
|
|
GreedyBot,
|
|
)
|
|
|
|
|
|
class TestBlokusEnv:
|
|
def test_env_creation(self):
|
|
env = BlokusEnv(num_players=4, board_size=20)
|
|
assert env.num_players == 4
|
|
assert env.board_size == 20
|
|
|
|
def test_observation_space(self):
|
|
env = BlokusEnv(num_players=4, board_size=20)
|
|
obs, _ = env.reset(seed=42)
|
|
assert "board" in obs
|
|
assert "pieces" in obs
|
|
assert "corners" in obs
|
|
assert obs["board"].shape == (20, 20)
|
|
assert obs["pieces"].shape == (21,)
|
|
assert obs["corners"].shape == (20, 20)
|
|
|
|
def test_action_space(self):
|
|
env = BlokusEnv(num_players=4, board_size=20)
|
|
assert env.action_space.shape == ()
|
|
assert env.action_space.n > 0
|
|
|
|
def test_reset_returns_obs_and_info(self):
|
|
env = BlokusEnv(num_players=4, board_size=20)
|
|
obs, info = env.reset(seed=42)
|
|
assert isinstance(obs, dict)
|
|
assert isinstance(info, dict)
|
|
assert "action_mask" in info
|
|
|
|
def test_action_mask_valid(self):
|
|
env = BlokusEnv(num_players=4, board_size=20)
|
|
obs, info = env.reset(seed=42)
|
|
mask = info["action_mask"]
|
|
assert mask.dtype == bool
|
|
assert mask.sum() > 0 # Should have valid moves
|
|
|
|
def test_step_with_valid_action(self):
|
|
env = BlokusEnv(num_players=4, board_size=20)
|
|
obs, info = env.reset(seed=42)
|
|
valid_actions = np.where(info["action_mask"])[0]
|
|
action = valid_actions[0]
|
|
obs, reward, terminated, truncated, info = env.step(action)
|
|
assert isinstance(reward, float)
|
|
assert isinstance(terminated, bool)
|
|
assert isinstance(truncated, bool)
|
|
|
|
def test_step_with_invalid_action(self):
|
|
env = BlokusEnv(num_players=4, board_size=20)
|
|
obs, info = env.reset(seed=42)
|
|
# Use an action that's likely invalid
|
|
obs, reward, terminated, truncated, info = env.step(env.action_space.n - 1)
|
|
assert reward < 0 # Penalty for invalid action
|
|
assert terminated # Should end episode
|
|
|
|
def test_two_player_env(self):
|
|
env = BlokusEnv(num_players=2, board_size=14)
|
|
obs, info = env.reset(seed=42)
|
|
assert env.num_players == 2
|
|
|
|
def test_three_player_env(self):
|
|
env = BlokusEnv(num_players=3, board_size=20)
|
|
obs, info = env.reset(seed=42)
|
|
assert env.num_players == 3
|
|
|
|
def test_custom_pieces(self):
|
|
custom = [p for p in STANDARD_PIECES if p.size < 5]
|
|
env = BlokusEnv(num_players=2, board_size=10, pieces=custom)
|
|
obs, info = env.reset(seed=42)
|
|
assert obs["pieces"].shape == (len(custom),)
|
|
|
|
def test_greedy_bot_opponent(self):
|
|
env = BlokusEnv(num_players=2, board_size=7, bot_type=GreedyBot)
|
|
obs, info = env.reset(seed=42)
|
|
assert env.bots[1] is not None
|
|
|
|
def test_game_over(self):
|
|
env = BlokusEnv(num_players=2, board_size=5, max_steps=10)
|
|
obs, info = env.reset(seed=42)
|
|
terminated = False
|
|
truncated = False
|
|
steps = 0
|
|
while not (terminated or truncated) and steps < 50:
|
|
valid = np.where(info["action_mask"])[0]
|
|
if len(valid) == 0:
|
|
break
|
|
obs, reward, terminated, truncated, info = env.step(valid[0])
|
|
steps += 1
|
|
|
|
def test_seed_reproducibility(self):
|
|
env1 = BlokusEnv(num_players=2, board_size=7)
|
|
env2 = BlokusEnv(num_players=2, board_size=7)
|
|
obs1, _ = env1.reset(seed=42)
|
|
obs2, _ = env2.reset(seed=42)
|
|
np.testing.assert_array_equal(obs1["board"], obs2["board"])
|
|
|
|
def test_render_text(self):
|
|
env = BlokusEnv(num_players=2, board_size=7, render_mode="ansi")
|
|
obs, info = env.reset(seed=42)
|
|
result = env.render()
|
|
assert isinstance(result, str)
|
|
|
|
|
|
class TestActionMaskWrapper:
|
|
def test_wrapper_creation(self):
|
|
env = BlokusEnv(num_players=2, board_size=7)
|
|
env = ActionMaskWrapper(env)
|
|
obs, info = env.reset(seed=42)
|
|
assert "observation" in obs
|
|
assert "action_mask" in obs
|
|
|
|
def test_wrapper_step(self):
|
|
env = BlokusEnv(num_players=2, board_size=7)
|
|
env = ActionMaskWrapper(env)
|
|
obs, info = env.reset(seed=42)
|
|
valid = np.where(obs["action_mask"])[0]
|
|
obs, reward, terminated, truncated, info = env.step(valid[0])
|
|
assert "observation" in obs
|
|
assert "action_mask" in obs
|
|
|
|
|
|
class TestMultiAgentEnv:
|
|
def test_env_creation(self):
|
|
env = BlokusMultiAgentEnv(num_players=2, board_size=7)
|
|
assert env.num_players == 2
|
|
assert len(env.agents) == 2
|
|
|
|
def test_reset(self):
|
|
env = BlokusMultiAgentEnv(num_players=2, board_size=7)
|
|
obs, info = env.reset(seed=42)
|
|
assert "player_0" in obs
|
|
|
|
def test_step(self):
|
|
env = BlokusMultiAgentEnv(num_players=2, board_size=7)
|
|
obs, info = env.reset(seed=42)
|
|
mask = env.get_action_mask("player_0")
|
|
valid = np.where(mask)[0]
|
|
obs, rewards, terminations, truncations, infos = env.step(valid[0])
|
|
assert isinstance(rewards, dict)
|
|
assert isinstance(terminations, dict)
|
|
|
|
def test_observation_space(self):
|
|
env = BlokusMultiAgentEnv(num_players=2, board_size=7)
|
|
obs, _ = env.reset(seed=42)
|
|
for agent in env.agents:
|
|
if agent in obs:
|
|
assert "board" in obs[agent]
|
|
|
|
def test_action_space(self):
|
|
env = BlokusMultiAgentEnv(num_players=2, board_size=7)
|
|
env.reset(seed=42)
|
|
space = env.action_space("player_0")
|
|
assert space.n > 0
|
|
|
|
def test_four_player(self):
|
|
env = BlokusMultiAgentEnv(num_players=4, board_size=7)
|
|
obs, info = env.reset(seed=42)
|
|
assert len(env.agents) == 4
|
|
|
|
|
|
class TestRegisteredEnvs:
|
|
def test_blokus_v0(self):
|
|
env = gym.make("Blokus-v0")
|
|
obs, info = env.reset(seed=42)
|
|
assert obs["board"].shape == (20, 20)
|
|
|
|
def test_blokus_duo_v0(self):
|
|
env = gym.make("BlokusDuo-v0")
|
|
obs, info = env.reset(seed=42)
|
|
assert obs["board"].shape == (14, 14)
|
|
|
|
def test_blokus_junior_v0(self):
|
|
env = gym.make("BlokusJunior-v0")
|
|
obs, info = env.reset(seed=42)
|
|
assert obs["board"].shape == (14, 14)
|
|
|
|
def test_blokus_simple_v0(self):
|
|
env = gym.make("BlokusSimple-v0")
|
|
obs, info = env.reset(seed=42)
|
|
assert obs["board"].shape == (7, 7)
|
|
|
|
def test_blokus_greedy_v0(self):
|
|
env = gym.make("BlokusGreedy-v0")
|
|
obs, info = env.reset(seed=42)
|
|
assert obs["board"].shape == (20, 20)
|