File size: 3,601 Bytes
4be6a52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
import copy
import json

import pytest

from stackcraft.players import Decision, HeuristicPlayer, Observation, RandomPlayer
from stackcraft.replay import replay_states
from stackcraft.tournament import run_episode, run_tournament


def game_content(artifact: dict) -> dict:
    result = copy.deepcopy(artifact)
    result.pop("timing")
    for decision in result["decisions"]:
        decision.pop("decision_seconds")
    return result


@pytest.mark.parametrize("factory", [RandomPlayer, HeuristicPlayer])
def test_episode_repeatability_and_replay_parity(factory) -> None:
    episode = run_episode(factory(), seed=42, max_pieces=25)
    assert game_content(episode) == game_content(run_episode(factory(), 42, 25))
    decoded = json.loads(json.dumps(episode))
    final = replay_states(decoded["replay"])[-1]
    for name in ("lines", "score", "terminal"):
        assert episode["outcome"][name] == getattr(final, name)
    assert episode["outcome"]["pieces"] == final.piece_index
    assert episode["timing"]["count"] == len(episode["decisions"])
    assert not episode["errors"]


def test_tournament_pairs_seeds_and_fresh_players() -> None:
    calls = []

    def random_factory():
        calls.append(True)
        return RandomPlayer()

    tournament = run_tournament({"random": random_factory, "heuristic": HeuristicPlayer}, [5, 8], 5)
    assert len(calls) == 2
    for label in ("random", "heuristic"):
        group = [episode for episode in tournament["episodes"] if episode["player_id"] == label]
        assert [episode["seed"] for episode in group] == [5, 8]
        assert all(episode["outcome"]["cap_hit"] for episode in group)
    assert tournament["summary"]["random"]["error_rate"] == 0
    assert tournament["summary"]["heuristic"]["cap_hit_rate"] == 1
    for seed in (5, 8):
        pair = [e for e in tournament["episodes"] if e["seed"] == seed]
        streams = [
            [(d["observation"]["current"], d["observation"]["next_piece"]) for d in e["decisions"]]
            for e in pair
        ]
        assert streams[0] == streams[1]


class InvalidPlayer:
    name = "invalid"
    revision = "test"

    def choose(self, observation: Observation) -> Decision:
        return Decision("invalid")


class ThrowingPlayer(InvalidPlayer):
    def choose(self, observation: Observation) -> Decision:
        raise RuntimeError("unavailable")


@pytest.mark.parametrize(
    "player,kind", [(InvalidPlayer(), "invalid_decision"), (ThrowingPlayer(), "player_error")]
)
def test_failures_are_explicit_without_random_fallback(player, kind: str) -> None:
    episode = run_episode(player, 0)
    assert episode["outcome"]["status"] == "error"
    assert episode["outcome"]["pieces"] == 0
    assert episode["outcome"]["cap_hit"] is False
    assert episode["errors"][0]["kind"] == kind
    assert episode["replay"]["actions"] == []
    assert len(episode["decisions"]) == 1
    assert episode["timing"]["count"] == 1


def test_terminal_is_not_mislabeled_cap() -> None:
    episode = run_episode(RandomPlayer(), 0, 200)
    assert episode["outcome"]["terminal"]
    assert episode["outcome"]["status"] == "top_out"
    assert not episode["outcome"]["cap_hit"]


@pytest.mark.parametrize("seeds", [[], [1, 1], [True], [1.5]])
def test_bad_seed_pools_rejected(seeds) -> None:
    with pytest.raises(ValueError):
        run_tournament({"random": RandomPlayer}, seeds)


@pytest.mark.parametrize("limit", [0, -1, True, 1.1])
def test_bad_limits_rejected(limit) -> None:
    with pytest.raises(ValueError, match="positive integer"):
        run_episode(RandomPlayer(), 0, limit)