import pytest from generals.agents import ExpanderAgent, HunterAgent, RandomAgent from general_bots_training.opponents import ModelOpponent, SelfPlayOpponent, make_opponent @pytest.mark.parametrize( ("name", "expected_type"), [ ("random", RandomAgent), ("expander", ExpanderAgent), ("hunter", HunterAgent), ], ) def test_make_opponent(name, expected_type): assert isinstance(make_opponent(name), expected_type) def test_make_opponent_is_case_insensitive(): assert isinstance(make_opponent("HUNTER"), HunterAgent) def test_make_opponent_supports_self_play_aliases(): assert isinstance(make_opponent("self_play"), SelfPlayOpponent) assert isinstance(make_opponent("self-play"), SelfPlayOpponent) def test_make_opponent_rejects_unknown_name(): with pytest.raises(ValueError, match="unknown opponent"): make_opponent("turtle") def test_make_opponent_model_requires_checkpoint(): with pytest.raises(ValueError, match="requires a checkpoint path"): make_opponent("model") def test_make_opponent_model_loads_checkpoint(tmp_path): import equinox as eqx import jax.random as jrandom from general_bots_training.network import PolicyValueNetwork network = PolicyValueNetwork(jrandom.PRNGKey(0), in_channels=14) checkpoint_path = tmp_path / "model.eqx" eqx.tree_serialise_leaves(checkpoint_path, network) opponent = make_opponent(f"model:{checkpoint_path}") assert isinstance(opponent, ModelOpponent) def test_make_opponent_rejects_checkpoint_on_non_model(): with pytest.raises(ValueError, match="does not take a checkpoint"): make_opponent("random:some.eqx")