From c1c50267b508c4157bf8a87cf7b06d6e86078c82 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Wed, 6 May 2026 22:35:51 +0900 Subject: [PATCH] =?UTF-8?q?FastGameState=EB=A5=BC=20=EA=B8=B0=EB=B3=B8=20?= =?UTF-8?q?=EA=B2=8C=EC=9E=84=20=EC=83=81=ED=83=9C=EB=A1=9C=20=EC=A0=84?= =?UTF-8?q?=ED=99=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../games/classic/__init__.py | 2 +- .../games/classic/bots/base.py | 4 +- .../games/classic/bots/heuristic.py | 5 +- .../games/classic/bots/passive.py | 4 +- .../games/classic/engines/fast.pxd | 10 +- .../games/classic/engines/fast.pyx | 210 ++++-- src/coolrl_lost_cities/games/classic/env.py | 3 +- .../games/classic/evaluation.py | 3 +- src/coolrl_lost_cities/games/classic/game.pyx | 629 +----------------- .../games/classic/interfaces.py | 2 +- .../games/classic/pygame_pvp.py | 3 +- .../games/classic/snapshots.py | 3 +- tests/__init__.py | 1 + tests/games/__init__.py | 1 + tests/games/classic/__init__.py | 1 + .../games/classic/engines/test_fast_parity.py | 213 +++--- tests/games/classic/helpers.py | 39 ++ tests/games/classic/test_bots.py | 148 +++-- tests/games/classic/test_canonical_state.py | 17 +- tests/games/classic/test_env.py | 26 +- tests/games/classic/test_masks.py | 8 +- tests/games/classic/test_rules.py | 53 +- 22 files changed, 477 insertions(+), 908 deletions(-) create mode 100644 tests/__init__.py create mode 100644 tests/games/__init__.py create mode 100644 tests/games/classic/__init__.py create mode 100644 tests/games/classic/helpers.py diff --git a/src/coolrl_lost_cities/games/classic/__init__.py b/src/coolrl_lost_cities/games/classic/__init__.py index 3dab106..58db0ae 100644 --- a/src/coolrl_lost_cities/games/classic/__init__.py +++ b/src/coolrl_lost_cities/games/classic/__init__.py @@ -5,6 +5,7 @@ from .bots import ( available_bot_names, build_bot, ) +from .engines import FastGameState as GameState from .env import LostCitiesEnv from .evaluation import ( GameResult, @@ -17,7 +18,6 @@ from .evaluation import ( play_match, ) from .game import ( - GameState, IllegalMoveError, LostCitiesConfig, classic_config, diff --git a/src/coolrl_lost_cities/games/classic/bots/base.py b/src/coolrl_lost_cities/games/classic/bots/base.py index af96884..7efd22b 100644 --- a/src/coolrl_lost_cities/games/classic/bots/base.py +++ b/src/coolrl_lost_cities/games/classic/bots/base.py @@ -1,6 +1,6 @@ from __future__ import annotations -from ..game import GameState +from ..engines import FastGameState as GameState from ..interfaces import BotInput, Snapshot try: @@ -10,7 +10,7 @@ except ImportError as exc: # pragma: no cover def legal_from_obs(obs_or_state: BotInput) -> np.ndarray: - if isinstance(obs_or_state, GameState): + if isinstance(obs_or_state, GameState) or hasattr(obs_or_state, "legal_mask"): return np.asarray(obs_or_state.legal_mask(), dtype=bool) if isinstance(obs_or_state, Snapshot): return np.asarray(obs_or_state.legal_mask, dtype=bool) diff --git a/src/coolrl_lost_cities/games/classic/bots/heuristic.py b/src/coolrl_lost_cities/games/classic/bots/heuristic.py index 5c14295..0df4af7 100644 --- a/src/coolrl_lost_cities/games/classic/bots/heuristic.py +++ b/src/coolrl_lost_cities/games/classic/bots/heuristic.py @@ -4,7 +4,8 @@ import logging from dataclasses import dataclass from functools import lru_cache -from ..game import Card, GameState, LostCitiesConfig +from ..engines import FastGameState as GameState +from ..game import Card, LostCitiesConfig from ..interfaces import BotInput, LostCitiesBot from .base import first_legal, legal_from_obs @@ -151,7 +152,7 @@ class SafeHeuristicBot(LostCitiesBot): self.params = params or SafeHeuristicParams() def act(self, obs_or_state: BotInput) -> int: - if not isinstance(obs_or_state, GameState): + if not isinstance(obs_or_state, GameState) and not hasattr(obs_or_state, "legal_mask"): LOGGER.debug( "SafeHeuristicBot fallback to first legal: input_type=%s", type(obs_or_state).__name__, diff --git a/src/coolrl_lost_cities/games/classic/bots/passive.py b/src/coolrl_lost_cities/games/classic/bots/passive.py index 14529be..ccda7bc 100644 --- a/src/coolrl_lost_cities/games/classic/bots/passive.py +++ b/src/coolrl_lost_cities/games/classic/bots/passive.py @@ -1,6 +1,6 @@ from __future__ import annotations -from ..game import GameState +from ..engines import FastGameState as GameState from ..interfaces import BotInput, Snapshot from .base import first_legal, legal_from_obs @@ -9,7 +9,7 @@ class PassiveDiscardBot: """Baseline that avoids opening expeditions whenever discarding is legal.""" def act(self, obs_or_state: BotInput) -> int: - if isinstance(obs_or_state, GameState): + if isinstance(obs_or_state, GameState) or hasattr(obs_or_state, "legal_mask"): return self._act_phase_local( obs_or_state.phase, obs_or_state.legal_mask(), diff --git a/src/coolrl_lost_cities/games/classic/engines/fast.pxd b/src/coolrl_lost_cities/games/classic/engines/fast.pxd index 7e4e40d..ac8b8ea 100644 --- a/src/coolrl_lost_cities/games/classic/engines/fast.pxd +++ b/src/coolrl_lost_cities/games/classic/engines/fast.pxd @@ -30,13 +30,13 @@ cdef class FastGameState: cdef int cards_per_color cdef int stride - cdef int* deck + cdef int* deck_cards cdef int deck_len - cdef int* hands + cdef int* hand_cards cdef int hand_lens[2] - cdef int* expeditions + cdef int* expedition_cards cdef int* expedition_lens - cdef int* discards + cdef int* discard_cards cdef int* discard_lens cdef int* last_numeric_ranks cdef int* handshake_counts @@ -61,6 +61,7 @@ cdef class FastGameState: cpdef list legal_draw_mask(self) cpdef list legal_mask(self) cpdef list unified_legal_mask(self) + cpdef object unified_legal_mask_np(self) cpdef list legal_actions(self) cpdef list unified_legal_actions(self) cpdef int from_unified_action(self, int action_id) @@ -106,3 +107,4 @@ cdef class FastGameState: cdef int _card_color(self, int card) cdef int _card_rank(self, int card) cdef object _card_snapshot(self, int card) + cdef object _card_obj(self, int card) diff --git a/src/coolrl_lost_cities/games/classic/engines/fast.pyx b/src/coolrl_lost_cities/games/classic/engines/fast.pyx index 46fecbc..11c78d1 100644 --- a/src/coolrl_lost_cities/games/classic/engines/fast.pyx +++ b/src/coolrl_lost_cities/games/classic/engines/fast.pyx @@ -7,7 +7,7 @@ import random from libc.string cimport memcpy from libc.stdlib cimport free, malloc, realloc -from ..game import IllegalMoveError, LostCitiesConfig, config_from_mapping +from ..game import Card, IllegalMoveError, LostCitiesConfig, config_from_mapping cdef inline int _phase_card(): @@ -20,11 +20,11 @@ cdef inline int _phase_draw(): cdef class FastGameState: def __cinit__(self): - self.deck = NULL - self.hands = NULL - self.expeditions = NULL + self.deck_cards = NULL + self.hand_cards = NULL + self.expedition_cards = NULL self.expedition_lens = NULL - self.discards = NULL + self.discard_cards = NULL self.discard_lens = NULL self.last_numeric_ranks = NULL self.handshake_counts = NULL @@ -38,16 +38,16 @@ cdef class FastGameState: self._configure(config) def __dealloc__(self): - if self.deck != NULL: - free(self.deck) - if self.hands != NULL: - free(self.hands) - if self.expeditions != NULL: - free(self.expeditions) + if self.deck_cards != NULL: + free(self.deck_cards) + if self.hand_cards != NULL: + free(self.hand_cards) + if self.expedition_cards != NULL: + free(self.expedition_cards) if self.expedition_lens != NULL: free(self.expedition_lens) - if self.discards != NULL: - free(self.discards) + if self.discard_cards != NULL: + free(self.discard_cards) if self.discard_lens != NULL: free(self.discard_lens) if self.last_numeric_ranks != NULL: @@ -75,13 +75,13 @@ cdef class FastGameState: self.cards_per_color = self.n_ranks + self.n_handshakes self.stride = self.n_ranks + 1 - self.deck = malloc(self.total_cards * sizeof(int)) - self.hands = malloc(2 * self.hand_size * sizeof(int)) - self.expeditions = malloc( + self.deck_cards = malloc(self.total_cards * sizeof(int)) + self.hand_cards = malloc(2 * self.hand_size * sizeof(int)) + self.expedition_cards = malloc( 2 * self.n_colors * self.cards_per_color * sizeof(int) ) self.expedition_lens = malloc(2 * self.n_colors * sizeof(int)) - self.discards = malloc(self.n_colors * self.cards_per_color * sizeof(int)) + self.discard_cards = malloc(self.n_colors * self.cards_per_color * sizeof(int)) self.discard_lens = malloc(self.n_colors * sizeof(int)) self.last_numeric_ranks = malloc(2 * self.n_colors * sizeof(int)) self.handshake_counts = malloc(2 * self.n_colors * sizeof(int)) @@ -92,11 +92,11 @@ cdef class FastGameState: self.undo_stack_capacity * sizeof(UndoRecord) ) if ( - self.deck == NULL - or self.hands == NULL - or self.expeditions == NULL + self.deck_cards == NULL + or self.hand_cards == NULL + or self.expedition_cards == NULL or self.expedition_lens == NULL - or self.discards == NULL + or self.discard_cards == NULL or self.discard_lens == NULL or self.last_numeric_ranks == NULL or self.handshake_counts == NULL @@ -159,11 +159,11 @@ cdef class FastGameState: cdef FastGameState state = cls(config) state.deck_len = len(encoded) for i, card in enumerate(encoded): - state.deck[i] = card + state.deck_cards[i] = card for _ in range(config.hand_size): for player in range(2): state.deck_len -= 1 - state.hands[state._hand_index(player, state.hand_lens[player])] = state.deck[ + state.hand_cards[state._hand_index(player, state.hand_lens[player])] = state.deck_cards[ state.deck_len ] state.hand_lens[player] += 1 @@ -186,7 +186,7 @@ cdef class FastGameState: ) state.deck_len = len(cards) for index, card in enumerate(cards): - state.deck[index] = card + state.deck_cards[index] = card for player in range(2): cards = [ @@ -199,7 +199,7 @@ cdef class FastGameState: ) state.hand_lens[player] = len(cards) for index, card in enumerate(cards): - state.hands[state._hand_index(player, index)] = card + state.hand_cards[state._hand_index(player, index)] = card for player in range(2): for color in range(state.n_colors): @@ -214,7 +214,7 @@ cdef class FastGameState: ) state.expedition_lens[state._expedition_len_index(player, color)] = len(cards) for index, card in enumerate(cards): - state.expeditions[state._expedition_index(player, color, index)] = card + state.expedition_cards[state._expedition_index(player, color, index)] = card for color in range(state.n_colors): cards = [_encode_card_snapshot(card, config) for card in snapshot["discards"][color]] @@ -225,7 +225,7 @@ cdef class FastGameState: ) state.discard_lens[color] = len(cards) for index, card in enumerate(cards): - state.discards[state._discard_index(color, index)] = card + state.discard_cards[state._discard_index(color, index)] = card state.current_player = int(snapshot.get("current_player", 0)) state.phase = snapshot.get("phase", "card") @@ -263,13 +263,58 @@ cdef class FastGameState: def action_size(self): return self.card_action_size + self.draw_action_size + @property + def deck(self): + return [self._card_obj(self.deck_cards[i]) for i in range(self.deck_len)] + + @property + def hands(self): + return [ + [ + self._card_obj(self.hand_cards[self._hand_index(player, i)]) + for i in range(self.hand_lens[player]) + ] + for player in range(2) + ] + + @property + def expeditions(self): + return [ + [ + [ + self._card_obj( + self.expedition_cards[ + self._expedition_index(player, color, i) + ] + ) + for i in range( + self.expedition_lens[ + self._expedition_len_index(player, color) + ] + ) + ] + for color in range(self.n_colors) + ] + for player in range(2) + ] + + @property + def discards(self): + return [ + [ + self._card_obj(self.discard_cards[self._discard_index(color, i)]) + for i in range(self.discard_lens[color]) + ] + for color in range(self.n_colors) + ] + def to_snapshot(self): return { "config": self.config.to_snapshot(), - "deck": [self._card_snapshot(self.deck[i]) for i in range(self.deck_len)], + "deck": [self._card_snapshot(self.deck_cards[i]) for i in range(self.deck_len)], "hands": [ [ - self._card_snapshot(self.hands[self._hand_index(player, i)]) + self._card_snapshot(self.hand_cards[self._hand_index(player, i)]) for i in range(self.hand_lens[player]) ] for player in range(2) @@ -278,7 +323,7 @@ cdef class FastGameState: [ [ self._card_snapshot( - self.expeditions[self._expedition_index(player, color, i)] + self.expedition_cards[self._expedition_index(player, color, i)] ) for i in range( self.expedition_lens[ @@ -292,7 +337,7 @@ cdef class FastGameState: ], "discards": [ [ - self._card_snapshot(self.discards[self._discard_index(color, i)]) + self._card_snapshot(self.discard_cards[self._discard_index(color, i)]) for i in range(self.discard_lens[color]) ] for color in range(self.n_colors) @@ -309,13 +354,13 @@ cdef class FastGameState: cpdef FastGameState clone(self): cdef FastGameState other = FastGameState(self.config) other.deck_len = self.deck_len - memcpy(other.deck, self.deck, self.deck_len * sizeof(int)) - memcpy(other.hands, self.hands, 2 * self.hand_size * sizeof(int)) + memcpy(other.deck_cards, self.deck_cards, self.deck_len * sizeof(int)) + memcpy(other.hand_cards, self.hand_cards, 2 * self.hand_size * sizeof(int)) other.hand_lens[0] = self.hand_lens[0] other.hand_lens[1] = self.hand_lens[1] memcpy( - other.expeditions, - self.expeditions, + other.expedition_cards, + self.expedition_cards, 2 * self.n_colors * self.cards_per_color * sizeof(int), ) memcpy( @@ -324,8 +369,8 @@ cdef class FastGameState: 2 * self.n_colors * sizeof(int), ) memcpy( - other.discards, - self.discards, + other.discard_cards, + self.discard_cards, self.n_colors * self.cards_per_color * sizeof(int), ) memcpy(other.discard_lens, self.discard_lens, self.n_colors * sizeof(int)) @@ -361,7 +406,7 @@ cdef class FastGameState: if self.terminal: return mask for slot in range(self.hand_lens[self.current_player]): - card = self.hands[self._hand_index(self.current_player, slot)] + card = self.hand_cards[self._hand_index(self.current_player, slot)] mask[2 * slot] = self._can_play_encoded_card_c(self.current_player, card) mask[2 * slot + 1] = True return mask @@ -510,6 +555,48 @@ cdef class FastGameState: cpdef int last_numeric_rank(self, int player, int color): return self.last_numeric_ranks[self._expedition_len_index(player, color)] + def has_numeric(self, int player, int color): + return self.last_numeric_rank(player, color) > 0 + + def can_play_card(self, int player, object card): + return self.can_play_encoded_card(player, _encode_card_snapshot(card, self.config)) + + def hand_slots(self, player=None): + cdef int p = self.current_player if player is None else int(player) + cdef list hand = [] + cdef int i + for i in range(self.hand_lens[p]): + hand.append(self._card_obj(self.hand_cards[self._hand_index(p, i)])) + while len(hand) < self.hand_size: + hand.append(None) + return hand + + def sort_hands(self): + self.sort_hand(0) + self.sort_hand(1) + + def sort_hand(self, player=None): + cdef int p = self.current_player if player is None else int(player) + cdef int i + cdef int j + cdef int key + cdef int current + for i in range(1, self.hand_lens[p]): + key = self.hand_cards[self._hand_index(p, i)] + j = i - 1 + while j >= 0 and self.hand_cards[self._hand_index(p, j)] > key: + current = self.hand_cards[self._hand_index(p, j)] + self.hand_cards[self._hand_index(p, j + 1)] = current + j -= 1 + self.hand_cards[self._hand_index(p, j + 1)] = key + + cpdef object unified_legal_mask_np(self): + try: + import numpy as np + except ImportError as exc: # pragma: no cover + raise RuntimeError("numpy is required for unified_legal_mask_np") from exc + return np.asarray(self.unified_legal_mask(), dtype=bool) + cpdef int expedition_score(self, int player, int color): return self.expedition_scores[self._expedition_len_index(player, color)] @@ -552,7 +639,7 @@ cdef class FastGameState: seen_numeric = False last_rank = 0 for index in range(length): - card = self.expeditions[self._expedition_index(player, color, index)] + card = self.expedition_cards[self._expedition_index(player, color, index)] if self._card_color(card) != color: raise ValueError("expedition contains wrong color") rank = self._card_rank(card) @@ -595,7 +682,7 @@ cdef class FastGameState: return True return self._can_play_encoded_card_c( self.current_player, - self.hands[self._hand_index(self.current_player, slot)], + self.hand_cards[self._hand_index(self.current_player, slot)], ) if action_id < 0 or action_id >= 1 + self.n_colors: return False @@ -616,7 +703,7 @@ cdef class FastGameState: return 0 if self.phase_id == _phase_card(): for slot in range(self.hand_lens[self.current_player]): - card = self.hands[self._hand_index(self.current_player, slot)] + card = self.hand_cards[self._hand_index(self.current_player, slot)] if self._can_play_encoded_card_c(self.current_player, card): out_actions[count] = 2 * slot count += 1 @@ -645,7 +732,7 @@ cdef class FastGameState: return 0 if self.phase_id == _phase_card(): for slot in range(self.hand_lens[self.current_player]): - card = self.hands[self._hand_index(self.current_player, slot)] + card = self.hand_cards[self._hand_index(self.current_player, slot)] if self._can_play_encoded_card_c(self.current_player, card): out_actions[count] = 2 * slot count += 1 @@ -697,7 +784,7 @@ cdef class FastGameState: undo.total_score_before = self.total_scores[self.current_player] if self.phase_id == _phase_card(): slot = action_id // 2 - card = self.hands[self._hand_index(self.current_player, slot)] + card = self.hand_cards[self._hand_index(self.current_player, slot)] color = self._card_color(card) cache_index = self._expedition_len_index(self.current_player, color) undo.slot = slot @@ -709,11 +796,11 @@ cdef class FastGameState: undo.numeric_sum_before = self.numeric_sums[cache_index] undo.expedition_score_before = self.expedition_scores[cache_index] elif action_id == 0: - undo.card = self.deck[self.deck_len - 1] + undo.card = self.deck_cards[self.deck_len - 1] else: color = action_id - 1 undo.color = color - undo.card = self.discards[self._discard_index(color, self.discard_lens[color] - 1)] + undo.card = self.discard_cards[self._discard_index(color, self.discard_lens[color] - 1)] cdef void _apply_action_with_undo_c(self, int action_id, UndoRecord* undo) except *: self._fill_undo_c(action_id, undo) @@ -802,7 +889,7 @@ cdef class FastGameState: cdef int slot = action_id // 2 cdef bint play = action_id % 2 == 0 cdef int player = self.current_player - cdef int card = self.hands[self._hand_index(player, slot)] + cdef int card = self.hand_cards[self._hand_index(player, slot)] cdef int color = self._card_color(card) cdef int rank = self._card_rank(card) cdef int i @@ -810,12 +897,12 @@ cdef class FastGameState: cdef int old_score cdef int new_score for i in range(slot, self.hand_lens[player] - 1): - self.hands[self._hand_index(player, i)] = self.hands[self._hand_index(player, i + 1)] + self.hand_cards[self._hand_index(player, i)] = self.hand_cards[self._hand_index(player, i + 1)] self.hand_lens[player] -= 1 if play: length_index = self._expedition_len_index(player, color) old_score = self.expedition_scores[length_index] - self.expeditions[self._expedition_index(player, color, self.expedition_lens[length_index])] = card + self.expedition_cards[self._expedition_index(player, color, self.expedition_lens[length_index])] = card self.expedition_lens[length_index] += 1 if rank == 0: self.handshake_counts[length_index] += 1 @@ -830,7 +917,7 @@ cdef class FastGameState: self.expedition_scores[length_index] = new_score self.total_scores[player] += new_score - old_score else: - self.discards[self._discard_index(color, self.discard_lens[color])] = card + self.discard_cards[self._discard_index(color, self.discard_lens[color])] = card self.discard_lens[color] += 1 self.pending_discarded_color = color self.phase_id = _phase_draw() @@ -845,12 +932,12 @@ cdef class FastGameState: cdef int color if action_id == 0: self.deck_len -= 1 - card = self.deck[self.deck_len] + card = self.deck_cards[self.deck_len] else: color = action_id - 1 self.discard_lens[color] -= 1 - card = self.discards[self._discard_index(color, self.discard_lens[color])] - self.hands[self._hand_index(player, self.hand_lens[player])] = card + card = self.discard_cards[self._discard_index(color, self.discard_lens[color])] + self.hand_cards[self._hand_index(player, self.hand_lens[player])] = card self.hand_lens[player] += 1 self.pending_discarded_color = -1 self.turn_count += 1 @@ -882,7 +969,7 @@ cdef class FastGameState: if play: length_index = self._expedition_len_index(player, color) self.expedition_lens[length_index] -= 1 - moved = self.expeditions[self._expedition_index(player, color, self.expedition_lens[length_index])] + moved = self.expedition_cards[self._expedition_index(player, color, self.expedition_lens[length_index])] self.last_numeric_ranks[length_index] = undo.last_numeric_before self.handshake_counts[length_index] = undo.handshake_count_before self.numeric_sums[length_index] = undo.numeric_sum_before @@ -890,12 +977,12 @@ cdef class FastGameState: self.total_scores[player] = undo.total_score_before else: self.discard_lens[color] -= 1 - moved = self.discards[self._discard_index(color, self.discard_lens[color])] + moved = self.discard_cards[self._discard_index(color, self.discard_lens[color])] if moved != card: raise ValueError("undo card mismatch") for i in range(self.hand_lens[player], slot, -1): - self.hands[self._hand_index(player, i)] = self.hands[self._hand_index(player, i - 1)] - self.hands[self._hand_index(player, slot)] = card + self.hand_cards[self._hand_index(player, i)] = self.hand_cards[self._hand_index(player, i - 1)] + self.hand_cards[self._hand_index(player, slot)] = card self.hand_lens[player] += 1 self.current_player = player self.phase_id = _phase_card() @@ -912,15 +999,15 @@ cdef class FastGameState: cdef int moved cdef int color self.hand_lens[player] -= 1 - moved = self.hands[self._hand_index(player, self.hand_lens[player])] + moved = self.hand_cards[self._hand_index(player, self.hand_lens[player])] if moved != card: raise ValueError("undo draw mismatch") if action_id == 0: - self.deck[self.deck_len] = card + self.deck_cards[self.deck_len] = card self.deck_len += 1 else: color = action_id - 1 - self.discards[self._discard_index(color, self.discard_lens[color])] = card + self.discard_cards[self._discard_index(color, self.discard_lens[color])] = card self.discard_lens[color] += 1 self.current_player = player self.phase_id = _phase_draw() @@ -949,7 +1036,7 @@ cdef class FastGameState: length = self.expedition_lens[cache_index] for card_index in range(length): rank = self._card_rank( - self.expeditions[ + self.expedition_cards[ self._expedition_index(player, color, card_index) ] ) @@ -1016,6 +1103,9 @@ cdef class FastGameState: cdef object _card_snapshot(self, int card): return {"color": self._card_color(card), "rank": self._card_rank(card)} + cdef object _card_obj(self, int card): + return Card(self._card_color(card), self._card_rank(card)) + def _build_encoded_deck(config): deck = [] diff --git a/src/coolrl_lost_cities/games/classic/env.py b/src/coolrl_lost_cities/games/classic/env.py index 88e1b04..67e337c 100644 --- a/src/coolrl_lost_cities/games/classic/env.py +++ b/src/coolrl_lost_cities/games/classic/env.py @@ -1,6 +1,7 @@ from __future__ import annotations -from .game import GameState, IllegalMoveError, LostCitiesConfig +from .engines import FastGameState as GameState +from .game import IllegalMoveError, LostCitiesConfig try: import numpy as np diff --git a/src/coolrl_lost_cities/games/classic/evaluation.py b/src/coolrl_lost_cities/games/classic/evaluation.py index 4a0acc6..0f3ce1a 100644 --- a/src/coolrl_lost_cities/games/classic/evaluation.py +++ b/src/coolrl_lost_cities/games/classic/evaluation.py @@ -10,7 +10,8 @@ from typing import Any import numpy as np from .bots import available_bot_names, build_bot -from .game import GameState, LostCitiesConfig, classic_config +from .engines import FastGameState as GameState +from .game import LostCitiesConfig, classic_config from .interfaces import LostCitiesBot BotFactory = Callable[[int | None], LostCitiesBot] diff --git a/src/coolrl_lost_cities/games/classic/game.pyx b/src/coolrl_lost_cities/games/classic/game.pyx index 8589977..9ca637e 100644 --- a/src/coolrl_lost_cities/games/classic/game.pyx +++ b/src/coolrl_lost_cities/games/classic/game.pyx @@ -1,23 +1,16 @@ # cython: language_level=3, boundscheck=False, wraparound=False, cdivision=True, initializedcheck=False -"""Deprecated compatibility rules engine for Lost Cities classic. +"""Core Lost Cities classic types. -This module remains the public engine while the replacement fast engine is -developed under ``coolrl_lost_cities.games.classic.engines``. New traversal, -simulation, and training work should target the fast engine once it exists. +``GameState`` is provided by the C-array fast engine. """ -from collections import Counter from dataclasses import dataclass, fields -import random from typing import Any, Literal -import numpy as np - cimport cython Phase = Literal["card", "draw"] -DEPRECATED_ENGINE = True class IllegalMoveError(ValueError): @@ -190,621 +183,6 @@ def build_deck(config): return deck -def _card_counter(cards): - return Counter(cards) - - -def _cards_from_snapshot(data): - if not isinstance(data, list): - raise ValueError(f"expected card list snapshot, got {type(data).__name__}") - return [Card.from_snapshot(card) for card in data] - - -def _cards_to_snapshot(cards): - return [card.to_snapshot() for card in cards] - - -cdef class GameState: - cdef public object config - cdef public list deck - cdef public list hands - cdef public list expeditions - cdef public list discards - cdef public int current_player - cdef public str phase - cdef public object pending_discarded_color - cdef public int turn_count - cdef public bint terminal - - def __init__( - self, - config, - deck=None, - hands=None, - expeditions=None, - discards=None, - int current_player=0, - phase="card", - pending_discarded_color=None, - int turn_count=0, - bint terminal=False, - ): - self.config = config - self.deck = list(deck) if deck is not None else [] - self.hands = hands if hands is not None else [[], []] - self.expeditions = expeditions if expeditions is not None else [ - [[] for _ in range(config.n_colors)], - [[] for _ in range(config.n_colors)], - ] - self.discards = discards if discards is not None else [ - [] for _ in range(config.n_colors) - ] - self.current_player = current_player - self.phase = phase - self.pending_discarded_color = pending_discarded_color - self.turn_count = turn_count - self.terminal = terminal - - @classmethod - def new_game(cls, config=None, *, seed=None): - config = config or LostCitiesConfig() - config.validate() - rng = random.Random(config.seed if seed is None else seed) - deck = build_deck(config) - rng.shuffle(deck) - return cls.new_game_from_deck(deck, config) - - @classmethod - def new_game_from_deck(cls, deck, config=None): - config = config or LostCitiesConfig() - config.validate() - cards = [Card.from_snapshot(card) for card in deck] - if _card_counter(cards) != _card_counter(build_deck(config)): - raise ValueError("deck must contain exactly the cards defined by config") - - state = cls.empty(config) - state.deck = list(cards) - cdef int player - for _ in range(config.hand_size): - for player in range(2): - state.hands[player].append(state.deck.pop()) - state.validate_invariants() - return state - - @classmethod - def empty(cls, config=None): - config = config or LostCitiesConfig() - config.validate() - return cls( - config=config, - deck=[], - hands=[[], []], - expeditions=[ - [[] for _ in range(config.n_colors)], - [[] for _ in range(config.n_colors)], - ], - discards=[[] for _ in range(config.n_colors)], - ) - - @classmethod - def from_snapshot(cls, snapshot, *, validate=True): - config = config_from_mapping(snapshot["config"]) - phase = snapshot.get("phase", "card") - if phase not in ("card", "draw"): - raise ValueError(f"invalid phase: {phase!r}") - - state = cls( - config=config, - deck=_cards_from_snapshot(snapshot["deck"]), - hands=[ - _cards_from_snapshot(snapshot["hands"][0]), - _cards_from_snapshot(snapshot["hands"][1]), - ], - expeditions=[ - [ - _cards_from_snapshot(color_cards) - for color_cards in snapshot["expeditions"][0] - ], - [ - _cards_from_snapshot(color_cards) - for color_cards in snapshot["expeditions"][1] - ], - ], - discards=[ - _cards_from_snapshot(color_cards) - for color_cards in snapshot["discards"] - ], - current_player=int(snapshot.get("current_player", 0)), - phase=phase, - pending_discarded_color=snapshot.get("pending_discarded_color"), - turn_count=int(snapshot.get("turn_count", 0)), - terminal=bool(snapshot.get("terminal", False)), - ) - if state.pending_discarded_color is not None: - state.pending_discarded_color = int(state.pending_discarded_color) - if validate: - state.validate_invariants() - return state - - def to_snapshot(self): - return { - "config": self.config.to_snapshot(), - "deck": _cards_to_snapshot(self.deck), - "hands": [_cards_to_snapshot(hand) for hand in self.hands], - "expeditions": [ - [_cards_to_snapshot(expedition) for expedition in player_expeditions] - for player_expeditions in self.expeditions - ], - "discards": [_cards_to_snapshot(discard) for discard in self.discards], - "current_player": self.current_player, - "phase": self.phase, - "pending_discarded_color": self.pending_discarded_color, - "turn_count": self.turn_count, - "terminal": self.terminal, - } - - cpdef GameState clone(self): - cdef GameState other = GameState.__new__(GameState) - other.config = self.config - other.deck = list(self.deck) - other.hands = [list(self.hands[0]), list(self.hands[1])] - other.expeditions = [ - [list(exp) for exp in self.expeditions[0]], - [list(exp) for exp in self.expeditions[1]], - ] - other.discards = [list(pile) for pile in self.discards] - other.current_player = self.current_player - other.phase = self.phase - other.pending_discarded_color = self.pending_discarded_color - other.turn_count = self.turn_count - other.terminal = self.terminal - return other - - @property - def card_action_size(self): - return self.config.card_action_size - - @property - def draw_action_size(self): - return self.config.draw_action_size - - @property - def action_size(self): - return self.config.action_size - - def sort_hands(self): - cdef int player - for player in range(2): - self.sort_hand(player) - - def sort_hand(self, player=None): - cdef int p = self.current_player if player is None else int(player) - self.hands[p].sort(key=_card_sort_key) - - def hand_slots(self, player=None): - cdef int p = self.current_player if player is None else int(player) - cdef list hand = self.hands[p] - cdef int hand_size = self.config.hand_size - cdef int n = len(hand) - cdef int i - cdef list out = [] - for i in range(hand_size): - if i < n: - out.append(hand[i]) - else: - out.append(None) - return out - - cpdef int last_numeric_rank(self, int player, int color): - cdef list expedition = self.expeditions[player][color] - cdef int best = 0 - cdef int n = len(expedition) - cdef int i - cdef Card card - for i in range(n): - card = expedition[i] - if card.rank == 0: - continue - if card.rank > best: - best = card.rank - return best - - def has_numeric(self, int player, int color): - return self.last_numeric_rank(player, color) > 0 - - cpdef bint can_play_card(self, int player, Card card): - cdef int n_colors = self.config.n_colors - cdef int n_ranks = self.config.n_ranks - if card.color < 0 or card.color >= n_colors: - return False - if card.rank < 0 or card.rank > n_ranks: - return False - cdef int last_numeric = self.last_numeric_rank(player, card.color) - if card.rank == 0: - return last_numeric == 0 - return card.rank > last_numeric - - cpdef list legal_card_mask(self): - cdef int size = self.card_action_size - cdef list mask = [False] * size - if self.terminal: - return mask - cdef list hand = self.hands[self.current_player] - cdef int hand_size = self.config.hand_size - cdef int n = len(hand) - cdef int slot - cdef Card card - for slot in range(hand_size): - if slot >= n: - continue - card = hand[slot] - mask[2 * slot] = self.can_play_card(self.current_player, card) - mask[2 * slot + 1] = True - return mask - - cpdef list legal_draw_mask(self): - cdef int size = self.draw_action_size - cdef list mask = [False] * size - if self.terminal: - return mask - mask[0] = len(self.deck) > 0 - cdef int n_colors = self.config.n_colors - cdef int color - cdef object pending = self.pending_discarded_color - for color in range(n_colors): - mask[1 + color] = ( - len(self.discards[color]) > 0 - and (pending is None or color != pending) - ) - return mask - - cpdef list legal_mask(self): - if self.phase == "card": - return self.legal_card_mask() - return self.legal_draw_mask() - - cpdef list unified_legal_mask(self): - cdef int draw_size = self.draw_action_size - cdef int card_size = self.card_action_size - cdef list result - if self.phase == "card": - result = self.legal_card_mask() - result.extend([False] * draw_size) - return result - result = [False] * card_size - result.extend(self.legal_draw_mask()) - return result - - cpdef object unified_legal_mask_np(self): - cdef int n_colors = self.config.n_colors - cdef int hand_size = self.config.hand_size - cdef int card_action_size = 2 * hand_size - cdef int draw_action_size = 1 + n_colors - cdef int total = card_action_size + draw_action_size - - mask_arr = np.zeros(total, dtype=bool) - cdef unsigned char[::1] view = mask_arr.view(np.uint8) - if self.terminal: - return mask_arr - - cdef int slot, n, color - cdef Card card - cdef list hand - cdef int p = self.current_player - cdef object pending - - if self.phase == "card": - hand = self.hands[p] - n = len(hand) - for slot in range(hand_size): - if slot >= n: - continue - card = hand[slot] - if self.can_play_card(p, card): - view[2 * slot] = 1 - view[2 * slot + 1] = 1 - else: - pending = self.pending_discarded_color - if len(self.deck) > 0: - view[card_action_size] = 1 - for color in range(n_colors): - if ( - len(self.discards[color]) > 0 - and (pending is None or color != pending) - ): - view[card_action_size + 1 + color] = 1 - return mask_arr - - def to_unified_action(self, int action_id, phase=None): - cdef str p = self.phase if phase is None else phase - if p == "card": - if action_id < 0 or action_id >= self.card_action_size: - raise IllegalMoveError(f"card action {action_id} is out of range") - return action_id - if action_id < 0 or action_id >= self.draw_action_size: - raise IllegalMoveError(f"draw action {action_id} is out of range") - return self.card_action_size + action_id - - cpdef int from_unified_action(self, int action_id): - if action_id < 0 or action_id >= self.action_size: - raise IllegalMoveError(f"action {action_id} is out of range") - if self.phase == "card": - if action_id >= self.card_action_size: - raise IllegalMoveError( - f"card action {action_id} is illegal during card phase" - ) - return action_id - if action_id < self.card_action_size: - raise IllegalMoveError( - f"card action {action_id} is illegal during draw phase" - ) - return action_id - self.card_action_size - - cpdef apply_action(self, int action_id): - if self.terminal: - raise IllegalMoveError("game is already terminal") - cdef list mask = self.legal_mask() - if action_id < 0 or action_id >= len(mask) or not mask[action_id]: - raise IllegalMoveError( - f"illegal action {action_id} in phase {self.phase} " - f"for player {self.current_player}" - ) - if self.phase == "card": - self._apply_card_action(action_id) - else: - self._apply_draw_action(action_id) - - cpdef apply_unified_action(self, int action_id): - self.apply_action(self.from_unified_action(action_id)) - - cpdef object apply_action_with_undo(self, int action_id): - if self.terminal: - raise IllegalMoveError("game is already terminal") - cdef list mask = self.legal_mask() - if action_id < 0 or action_id >= len(mask) or not mask[action_id]: - raise IllegalMoveError( - f"illegal action {action_id} in phase {self.phase} " - f"for player {self.current_player}" - ) - cdef object undo - if self.phase == "card": - undo = self._card_action_undo(action_id) - self._apply_card_action(action_id) - else: - undo = self._draw_action_undo(action_id) - self._apply_draw_action(action_id) - return undo - - cpdef object apply_unified_action_with_undo(self, int action_id): - return self.apply_action_with_undo(self.from_unified_action(action_id)) - - cpdef undo_action(self, object undo): - cdef str phase = undo[0] - if phase == "card": - self._undo_card_action(undo) - return - if phase == "draw": - self._undo_draw_action(undo) - return - raise ValueError(f"invalid undo phase: {phase!r}") - - cdef object _card_action_undo(self, int action_id): - cdef int slot = action_id // 2 - cdef bint play = action_id % 2 == 0 - cdef Card card = self.hands[self.current_player][slot] - return ( - "card", - self.current_player, - action_id, - self.pending_discarded_color, - self.terminal, - slot, - play, - card, - ) - - cdef object _draw_action_undo(self, int action_id): - cdef Card card - cdef list source - if action_id == 0: - source = self.deck - else: - source = self.discards[action_id - 1] - card = source[len(source) - 1] - return ( - "draw", - self.current_player, - action_id, - self.pending_discarded_color, - self.terminal, - self.turn_count, - card, - ) - - cdef void _apply_card_action(self, int action_id) except *: - cdef int slot = action_id // 2 - cdef bint play = action_id % 2 == 0 - cdef Card card = self.hands[self.current_player].pop(slot) - if play: - self.expeditions[self.current_player][card.color].append(card) - else: - self.discards[card.color].append(card) - self.pending_discarded_color = card.color - self.phase = "draw" - cdef int n_colors = self.config.n_colors - cdef int color - cdef object pending = self.pending_discarded_color - cdef bint any_legal_draw = False - if len(self.deck) == 0: - for color in range(n_colors): - if len(self.discards[color]) > 0 and (pending is None or color != pending): - any_legal_draw = True - break - if not any_legal_draw: - self.terminal = True - - cdef void _apply_draw_action(self, int action_id) except *: - cdef Card card - cdef int color - if action_id == 0: - card = self.deck.pop() - else: - color = action_id - 1 - card = self.discards[color].pop() - self.hands[self.current_player].append(card) - self.pending_discarded_color = None - self.turn_count += 1 - if len(self.deck) == 0: - self.terminal = True - return - self.current_player = 1 - self.current_player - self.phase = "card" - - cdef void _undo_card_action(self, object undo) except *: - cdef int player = undo[1] - cdef object pending_before = undo[3] - cdef bint terminal_before = undo[4] - cdef int slot = undo[5] - cdef bint play = undo[6] - cdef Card card = undo[7] - cdef Card moved - if play: - moved = self.expeditions[player][card.color].pop() - else: - moved = self.discards[card.color].pop() - if moved != card: - raise ValueError("undo card mismatch") - self.hands[player].insert(slot, card) - self.current_player = player - self.phase = "card" - self.pending_discarded_color = pending_before - self.terminal = terminal_before - - cdef void _undo_draw_action(self, object undo) except *: - cdef int player = undo[1] - cdef int action_id = undo[2] - cdef object pending_before = undo[3] - cdef bint terminal_before = undo[4] - cdef int turn_count_before = undo[5] - cdef Card card = undo[6] - cdef Card moved = self.hands[player].pop() - if moved != card: - raise ValueError("undo draw mismatch") - if action_id == 0: - self.deck.append(card) - else: - self.discards[action_id - 1].append(card) - self.current_player = player - self.phase = "draw" - self.pending_discarded_color = pending_before - self.turn_count = turn_count_before - self.terminal = terminal_before - - cpdef int expedition_score(self, int player, int color): - return score_expedition(self.expeditions[player][color], self.config) - - cpdef int total_score(self, int player): - cdef int total = 0 - cdef int color - cdef int n_colors = self.config.n_colors - for color in range(n_colors): - total += score_expedition(self.expeditions[player][color], self.config) - return total - - cpdef int score_diff(self, int player=0): - cdef int other = 1 - player - return self.total_score(player) - self.total_score(other) - - def validate_invariants(self): - self.config.validate() - if self.current_player not in (0, 1): - raise ValueError("current_player must be 0 or 1") - if self.phase not in ("card", "draw"): - raise ValueError(f"invalid phase: {self.phase!r}") - if len(self.hands) != 2: - raise ValueError("hands must contain two players") - if len(self.expeditions) != 2: - raise ValueError("expeditions must contain two players") - if len(self.discards) != self.config.n_colors: - raise ValueError("discard pile count must match n_colors") - - all_cards = [] - all_cards.extend(self.deck) - for player, hand in enumerate(self.hands): - if len(hand) > self.config.hand_size: - raise ValueError(f"hand {player} exceeds hand_size") - all_cards.extend(hand) - - for player, expeditions in enumerate(self.expeditions): - if len(expeditions) != self.config.n_colors: - raise ValueError("expedition color count must match n_colors") - for color, expedition in enumerate(expeditions): - self._validate_expedition(player, color, expedition) - all_cards.extend(expedition) - for discard in self.discards: - all_cards.extend(discard) - - for card in all_cards: - self._validate_card(card) - if _card_counter(all_cards) != _card_counter(build_deck(self.config)): - raise ValueError("card conservation failed") - - if self.phase == "card" and self.pending_discarded_color is not None: - raise ValueError("pending_discarded_color must be None during card phase") - if self.pending_discarded_color is not None: - color = self.pending_discarded_color - if color < 0 or color >= self.config.n_colors: - raise ValueError("pending_discarded_color is out of range") - if not self.discards[color]: - raise ValueError("pending discard color must have a discard pile card") - - any_legal = any(self.unified_legal_mask()) - if self.terminal and any_legal: - raise ValueError("terminal state must have no legal actions") - if not self.terminal and not any_legal: - raise ValueError("non-terminal state must have at least one legal action") - - def _validate_card(self, Card card): - if card.color < 0 or card.color >= self.config.n_colors: - raise ValueError(f"card color out of range: {card}") - if card.rank < 0 or card.rank > self.config.n_ranks: - raise ValueError(f"card rank out of range: {card}") - - def _validate_expedition(self, int player, int color, list expedition): - cdef bint seen_numeric = False - cdef int last_numeric = 0 - cdef Card card - for card in expedition: - if card.color != color: - raise ValueError( - f"player {player} expedition {color} contains wrong color" - ) - if card.rank == 0: - if seen_numeric: - raise ValueError( - f"player {player} expedition {color} has handshake after number" - ) - continue - seen_numeric = True - if card.rank <= last_numeric: - raise ValueError( - f"player {player} expedition {color} is not strictly increasing" - ) - last_numeric = card.rank - - def __reduce__(self): - # support pickle via snapshot round-trip - return (_rebuild_game_state, (self.to_snapshot(),)) - - -def _rebuild_game_state(snapshot): - return GameState.from_snapshot(snapshot, validate=False) - - -def _card_sort_key(Card card): - return (card.color, card.rank) - - cpdef int score_expedition(list expedition, config): cdef int n = len(expedition) if n == 0: @@ -824,3 +202,6 @@ cpdef int score_expedition(list expedition, config): if n >= config.bonus_threshold: score += config.bonus_amount return score + + +from .engines.fast import FastGameState as GameState diff --git a/src/coolrl_lost_cities/games/classic/interfaces.py b/src/coolrl_lost_cities/games/classic/interfaces.py index 4d188d3..13ac9f5 100644 --- a/src/coolrl_lost_cities/games/classic/interfaces.py +++ b/src/coolrl_lost_cities/games/classic/interfaces.py @@ -2,7 +2,7 @@ from __future__ import annotations from typing import Protocol, TypeAlias, runtime_checkable -from .game import GameState +from .engines import FastGameState as GameState from .snapshots import Snapshot BotInput: TypeAlias = dict | GameState | Snapshot diff --git a/src/coolrl_lost_cities/games/classic/pygame_pvp.py b/src/coolrl_lost_cities/games/classic/pygame_pvp.py index d593062..6ec4114 100644 --- a/src/coolrl_lost_cities/games/classic/pygame_pvp.py +++ b/src/coolrl_lost_cities/games/classic/pygame_pvp.py @@ -13,7 +13,8 @@ from pathlib import Path from typing import Any, Literal from .bots import DEFAULT_BOT, LostCitiesBot, available_bot_names, build_bot -from .game import Card, GameState, LostCitiesConfig, classic_config +from .engines import FastGameState as GameState +from .game import Card, LostCitiesConfig, classic_config from .resources import theme_path from .snapshots import Snapshot, snapshot_from_state, snapshot_summary diff --git a/src/coolrl_lost_cities/games/classic/snapshots.py b/src/coolrl_lost_cities/games/classic/snapshots.py index 9ba4fa0..40d111c 100644 --- a/src/coolrl_lost_cities/games/classic/snapshots.py +++ b/src/coolrl_lost_cities/games/classic/snapshots.py @@ -2,7 +2,8 @@ from __future__ import annotations from dataclasses import dataclass -from .game import Card, GameState, LostCitiesConfig, score_expedition +from .engines import FastGameState as GameState +from .game import Card, LostCitiesConfig, score_expedition @dataclass diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/games/__init__.py b/tests/games/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/games/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/games/classic/__init__.py b/tests/games/classic/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/games/classic/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/games/classic/engines/test_fast_parity.py b/tests/games/classic/engines/test_fast_parity.py index 8202ebd..99596ff 100644 --- a/tests/games/classic/engines/test_fast_parity.py +++ b/tests/games/classic/engines/test_fast_parity.py @@ -13,7 +13,7 @@ def _card(color: int, rank: int) -> dict[str, int]: return {"color": color, "rank": rank} -def _classic_snapshot( +def _snapshot( *, deck: list[dict[str, int]] | None = None, hands: list[list[dict[str, int]]] | None = None, @@ -59,25 +59,26 @@ def _classic_snapshot( } -def test_fast_new_game_from_deck_matches_game_state_snapshot() -> None: +def test_public_game_state_alias_matches_fast_new_game_from_deck_snapshot() -> None: config = LostCitiesConfig() deck = build_deck(config) - classic = GameState.new_game_from_deck(deck, config) - fast = FastGameState.new_game_from_deck(deck, config) + assert GameState is FastGameState + left = GameState.new_game_from_deck(deck, config) + right = FastGameState.new_game_from_deck(deck, config) - assert fast.to_snapshot() == classic.to_snapshot() - fast.validate_invariants() + assert right.to_snapshot() == left.to_snapshot() + right.validate_invariants() def test_fast_snapshot_roundtrip_preserves_snapshot() -> None: config = LostCitiesConfig(seed=11) - classic = GameState.new_game(config) - fast = FastGameState.from_snapshot(classic.to_snapshot()) + left = GameState.new_game(config) + right = FastGameState.from_snapshot(left.to_snapshot()) - assert fast.to_snapshot() == classic.to_snapshot() - restored = FastGameState.from_snapshot(fast.to_snapshot()) - assert restored.to_snapshot() == fast.to_snapshot() + assert right.to_snapshot() == left.to_snapshot() + restored = FastGameState.from_snapshot(right.to_snapshot()) + assert restored.to_snapshot() == right.to_snapshot() def test_fast_from_snapshot_rejects_oversized_regions_before_write() -> None: @@ -125,39 +126,39 @@ def test_fast_validate_invariants_rejects_bad_expedition_order() -> None: FastGameState.from_snapshot(snapshot) -def test_fast_pending_discard_matches_game_state() -> None: - snapshot = _classic_snapshot( +def test_fast_pending_discard_sequence_is_deterministic() -> None: + snapshot = _snapshot( hands=[ [_card(0, 1)], [_card(1, 1)], ], deck=[_card(2, 1), _card(3, 1)], ) - classic = GameState.from_snapshot(snapshot) - fast = FastGameState.from_snapshot(snapshot) + left = GameState.from_snapshot(snapshot) + right = FastGameState.from_snapshot(snapshot) - classic.apply_action(1) - fast.apply_action(1) - assert fast.to_snapshot() == classic.to_snapshot() - assert fast.legal_draw_mask() == classic.legal_draw_mask() - assert fast.legal_draw_mask()[1] is False + left.apply_action(1) + right.apply_action(1) + assert right.to_snapshot() == left.to_snapshot() + assert right.legal_draw_mask() == left.legal_draw_mask() + assert right.legal_draw_mask()[1] is False - classic.apply_action(0) - fast.apply_action(0) - classic.apply_action(1) - fast.apply_action(1) - classic.apply_action(0) - fast.apply_action(0) - classic.apply_action(1) - fast.apply_action(1) + left.apply_action(0) + right.apply_action(0) + left.apply_action(1) + right.apply_action(1) + left.apply_action(0) + right.apply_action(0) + left.apply_action(1) + right.apply_action(1) - assert fast.to_snapshot() == classic.to_snapshot() - assert fast.legal_draw_mask() == classic.legal_draw_mask() - assert fast.legal_draw_mask()[1] is True + assert right.to_snapshot() == left.to_snapshot() + assert right.legal_draw_mask() == left.legal_draw_mask() + assert right.legal_draw_mask()[1] is True -def test_fast_terminal_edges_match_game_state() -> None: - last_draw_snapshot = _classic_snapshot( +def test_fast_terminal_edges_are_deterministic() -> None: + last_draw_snapshot = _snapshot( deck=[_card(1, 1)], hands=[ [_card(0, 1)], @@ -168,16 +169,16 @@ def test_fast_terminal_edges_match_game_state() -> None: last_draw_snapshot["deck"] = [last_draw_snapshot["deck"][-1]] for card in remaining_deck: last_draw_snapshot["discards"][card["color"]].append(card) - classic = GameState.from_snapshot(last_draw_snapshot) - fast = FastGameState.from_snapshot(last_draw_snapshot) + left = GameState.from_snapshot(last_draw_snapshot) + right = FastGameState.from_snapshot(last_draw_snapshot) - classic.apply_action(1) - fast.apply_action(1) - classic.apply_action(0) - fast.apply_action(0) + left.apply_action(1) + right.apply_action(1) + left.apply_action(0) + right.apply_action(0) - assert fast.to_snapshot() == classic.to_snapshot() - assert fast.terminal is True + assert right.to_snapshot() == left.to_snapshot() + assert right.terminal is True defensive_snapshot = { "config": LostCitiesConfig().to_snapshot(), @@ -191,18 +192,18 @@ def test_fast_terminal_edges_match_game_state() -> None: "turn_count": 0, "terminal": False, } - classic = GameState.from_snapshot(defensive_snapshot, validate=False) - fast = FastGameState.from_snapshot(defensive_snapshot, validate=False) + left = GameState.from_snapshot(defensive_snapshot, validate=False) + right = FastGameState.from_snapshot(defensive_snapshot, validate=False) - classic.apply_action(1) - fast.apply_action(1) + left.apply_action(1) + right.apply_action(1) - assert fast.to_snapshot() == classic.to_snapshot() - assert fast.terminal is True + assert right.to_snapshot() == left.to_snapshot() + assert right.terminal is True -def test_fast_last_numeric_legality_matches_game_state() -> None: - handshake_snapshot = _classic_snapshot( +def test_fast_last_numeric_legality_edges() -> None: + handshake_snapshot = _snapshot( hands=[ [_card(0, 1)], [], @@ -212,12 +213,12 @@ def test_fast_last_numeric_legality_matches_game_state() -> None: [[], [], [], [], []], ], ) - classic = GameState.from_snapshot(handshake_snapshot) - fast = FastGameState.from_snapshot(handshake_snapshot) - assert fast.legal_card_mask() == classic.legal_card_mask() - assert fast.legal_card_mask()[0] is True + left = GameState.from_snapshot(handshake_snapshot) + right = FastGameState.from_snapshot(handshake_snapshot) + assert right.legal_card_mask() == left.legal_card_mask() + assert right.legal_card_mask()[0] is True - numeric_snapshot = _classic_snapshot( + numeric_snapshot = _snapshot( hands=[ [_card(0, 0), _card(0, 3), _card(0, 5)], [], @@ -227,16 +228,16 @@ def test_fast_last_numeric_legality_matches_game_state() -> None: [[], [], [], [], []], ], ) - classic = GameState.from_snapshot(numeric_snapshot) - fast = FastGameState.from_snapshot(numeric_snapshot) - assert fast.legal_card_mask() == classic.legal_card_mask() - assert fast.legal_card_mask()[0] is False - assert fast.legal_card_mask()[2] is False - assert fast.legal_card_mask()[4] is True + left = GameState.from_snapshot(numeric_snapshot) + right = FastGameState.from_snapshot(numeric_snapshot) + assert right.legal_card_mask() == left.legal_card_mask() + assert right.legal_card_mask()[0] is False + assert right.legal_card_mask()[2] is False + assert right.legal_card_mask()[4] is True -def test_fast_score_cache_and_undo_match_game_state() -> None: - snapshot = _classic_snapshot( +def test_fast_score_cache_and_undo_restore_snapshot() -> None: + snapshot = _snapshot( hands=[ [_card(0, 7)], [], @@ -261,26 +262,26 @@ def test_fast_score_cache_and_undo_match_game_state() -> None: [[], [], [], [], []], ], ) - classic = GameState.from_snapshot(snapshot) - fast = FastGameState.from_snapshot(snapshot) - before = fast.to_snapshot() + left = GameState.from_snapshot(snapshot) + right = FastGameState.from_snapshot(snapshot) + before = right.to_snapshot() - assert fast.expedition_score(0, 0) == classic.expedition_score(0, 0) - assert fast.total_score(0) == classic.total_score(0) + assert right.expedition_score(0, 0) == left.expedition_score(0, 0) + assert right.total_score(0) == left.total_score(0) - undo = fast.apply_action_with_undo(0) - classic.apply_action(0) - assert fast.to_snapshot() == classic.to_snapshot() - assert fast.expedition_score(0, 0) == classic.expedition_score(0, 0) - assert fast.total_score(0) == classic.total_score(0) + undo = right.apply_action_with_undo(0) + left.apply_action(0) + assert right.to_snapshot() == left.to_snapshot() + assert right.expedition_score(0, 0) == left.expedition_score(0, 0) + assert right.total_score(0) == left.total_score(0) - fast.undo_action(undo) - assert fast.to_snapshot() == before - assert fast.total_score(0) == GameState.from_snapshot(before).total_score(0) + right.undo_action(undo) + assert right.to_snapshot() == before + assert right.total_score(0) == GameState.from_snapshot(before).total_score(0) def test_fast_discard_draw_push_pop_restores_snapshot() -> None: - snapshot = _classic_snapshot( + snapshot = _snapshot( hands=[[], [_card(1, 1)]], discards=[[_card(0, 1)], [], [], [], []], phase="draw", @@ -294,62 +295,60 @@ def test_fast_discard_draw_push_pop_restores_snapshot() -> None: assert state.to_snapshot() == before -def test_fast_random_action_sequence_matches_game_state() -> None: +def test_fast_random_action_sequence_is_deterministic() -> None: config = LostCitiesConfig() for seed in range(48): - classic = GameState.new_game(config, seed=seed) - fast = FastGameState.new_game(config, seed=seed) + left = GameState.new_game(config, seed=seed) + right = FastGameState.new_game(config, seed=seed) rng = random.Random(seed ^ 0xF457) steps = 0 while True: - assert fast.to_snapshot() == classic.to_snapshot() - assert fast.unified_legal_mask() == classic.unified_legal_mask() - assert fast.unified_legal_actions() == [ - index for index, is_legal in enumerate(classic.unified_legal_mask()) if is_legal + assert right.to_snapshot() == left.to_snapshot() + assert right.unified_legal_mask() == left.unified_legal_mask() + assert right.unified_legal_actions() == [ + index for index, is_legal in enumerate(left.unified_legal_mask()) if is_legal ] - assert fast.score_diff(0) == classic.score_diff(0) - if classic.terminal: + assert right.score_diff(0) == left.score_diff(0) + if left.terminal: break - legal = [ - index for index, is_legal in enumerate(classic.unified_legal_mask()) if is_legal - ] + legal = [index for index, is_legal in enumerate(left.unified_legal_mask()) if is_legal] action = rng.choice(legal) - classic.apply_unified_action(action) - fast.apply_unified_action(action) + left.apply_unified_action(action) + right.apply_unified_action(action) steps += 1 assert steps < 1000 -def test_fast_random_bot_self_play_matches_game_state() -> None: +def test_fast_random_bot_self_play_is_deterministic() -> None: config = LostCitiesConfig() for seed in range(32): - classic = GameState.new_game(config, seed=seed) - fast = FastGameState.new_game(config, seed=seed) - classic_bots = [RandomBot(seed=seed * 2), RandomBot(seed=seed * 2 + 1)] - fast_bots = [RandomBot(seed=seed * 2), RandomBot(seed=seed * 2 + 1)] + left = GameState.new_game(config, seed=seed) + right = FastGameState.new_game(config, seed=seed) + left_bots = [RandomBot(seed=seed * 2), RandomBot(seed=seed * 2 + 1)] + right_bots = [RandomBot(seed=seed * 2), RandomBot(seed=seed * 2 + 1)] steps = 0 while True: - assert fast.to_snapshot() == classic.to_snapshot() - if classic.terminal: + assert right.to_snapshot() == left.to_snapshot() + if left.terminal: break - player = classic.current_player - assert fast.current_player == player - classic_action = classic_bots[player].act(classic) - fast_action = fast_bots[player].act({"legal_mask": fast.legal_mask()}) - assert fast_action == classic_action + player = left.current_player + assert right.current_player == player + left_action = left_bots[player].act(left) + right_action = right_bots[player].act({"legal_mask": right.legal_mask()}) + assert right_action == left_action - classic.apply_action(classic_action) - fast.apply_action(fast_action) + left.apply_action(left_action) + right.apply_action(right_action) steps += 1 assert steps < 1000 - assert fast.total_score(0) == classic.total_score(0) - assert fast.total_score(1) == classic.total_score(1) - assert fast.score_diff(0) == classic.score_diff(0) + assert right.total_score(0) == left.total_score(0) + assert right.total_score(1) == left.total_score(1) + assert right.score_diff(0) == left.score_diff(0) def test_fast_apply_undo_restores_every_legal_action() -> None: diff --git a/tests/games/classic/helpers.py b/tests/games/classic/helpers.py new file mode 100644 index 0000000..55bb79a --- /dev/null +++ b/tests/games/classic/helpers.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from coolrl_lost_cities.games.classic.game import Card, GameState, LostCitiesConfig + + +def make_state( + config: LostCitiesConfig | None = None, + *, + deck: list[Card] | None = None, + hands: list[list[Card]] | None = None, + expeditions: list[list[list[Card]]] | None = None, + discards: list[list[Card]] | None = None, + current_player: int = 0, + phase: str = "card", + pending_discarded_color: int | None = None, + turn_count: int = 0, + terminal: bool = False, + validate: bool = False, +) -> GameState: + config = config or LostCitiesConfig() + return GameState.from_snapshot( + { + "config": config.to_snapshot(), + "deck": deck or [], + "hands": hands or [[], []], + "expeditions": expeditions + or [ + [[] for _ in range(config.n_colors)], + [[] for _ in range(config.n_colors)], + ], + "discards": discards or [[] for _ in range(config.n_colors)], + "current_player": current_player, + "phase": phase, + "pending_discarded_color": pending_discarded_color, + "turn_count": turn_count, + "terminal": terminal, + }, + validate=validate, + ) diff --git a/tests/games/classic/test_bots.py b/tests/games/classic/test_bots.py index 735f5b2..a4085cf 100644 --- a/tests/games/classic/test_bots.py +++ b/tests/games/classic/test_bots.py @@ -7,6 +7,14 @@ from coolrl_lost_cities.games.classic.bots import ( ) from coolrl_lost_cities.games.classic.bots.heuristic import draw_from_discard_action from coolrl_lost_cities.games.classic.evaluation import play_game_for_evaluation +from tests.games.classic.helpers import make_state + + +def _expeditions(config: LostCitiesConfig) -> list[list[list[Card]]]: + return [ + [[] for _ in range(config.n_colors)], + [[] for _ in range(config.n_colors)], + ] def test_builtin_bots_implement_lost_cities_bot() -> None: @@ -31,15 +39,26 @@ def test_safe_heuristic_opponent_value_ignores_hidden_hand() -> None: bot = SafeHeuristicBot() discard_card = Card(color=0, rank=6) - state_a = GameState.empty(config) - state_a.expeditions[1][0] = [Card(color=0, rank=0), Card(color=0, rank=4)] - state_a.discards[0] = [discard_card] - state_a.hands[1] = [Card(color=0, rank=5)] + expeditions_a = _expeditions(config) + expeditions_a[1][0] = [Card(color=0, rank=0), Card(color=0, rank=4)] + state_a = make_state( + config, + hands=[[], [Card(color=0, rank=5)]], + expeditions=expeditions_a, + discards=[[discard_card], []], + ) - state_b = GameState.empty(config) - state_b.expeditions[1][0] = [Card(color=0, rank=0), Card(color=0, rank=4)] - state_b.discards[0] = [discard_card] - state_b.hands[1] = [Card(color=0, rank=5), Card(color=0, rank=7), Card(color=0, rank=8)] + expeditions_b = _expeditions(config) + expeditions_b[1][0] = [Card(color=0, rank=0), Card(color=0, rank=4)] + state_b = make_state( + config, + hands=[ + [], + [Card(color=0, rank=5), Card(color=0, rank=7), Card(color=0, rank=8)], + ], + expeditions=expeditions_b, + discards=[[discard_card], []], + ) value_a = bot._card_value_for_opponent( state=state_a, @@ -62,13 +81,21 @@ def test_safe_heuristic_started_expedition_value_ignores_invalid_lower_followup( bot = SafeHeuristicBot() high_card = Card(color=0, rank=8) - base_state = GameState.empty(config) - base_state.expeditions[0][0] = [Card(color=0, rank=4)] - base_state.hands[0] = [high_card] + base_expeditions = _expeditions(config) + base_expeditions[0][0] = [Card(color=0, rank=4)] + base_state = make_state( + config, + hands=[[high_card], []], + expeditions=base_expeditions, + ) - lower_followup_state = GameState.empty(config) - lower_followup_state.expeditions[0][0] = [Card(color=0, rank=4)] - lower_followup_state.hands[0] = [Card(color=0, rank=5), high_card] + lower_expeditions = _expeditions(config) + lower_expeditions[0][0] = [Card(color=0, rank=4)] + lower_followup_state = make_state( + config, + hands=[[Card(color=0, rank=5), high_card], []], + expeditions=lower_expeditions, + ) base_value = bot._started_expedition_play_value( state=base_state, @@ -92,12 +119,15 @@ def test_safe_heuristic_draws_playable_discard_instead_of_deck() -> None: config = LostCitiesConfig(n_colors=2, n_ranks=8, hand_size=3) bot = SafeHeuristicBot() - state = GameState.empty(config) - state.current_player = 0 - state.phase = "draw" - state.expeditions[0][0] = [Card(color=0, rank=4)] - state.discards[0] = [Card(color=0, rank=6)] - state.deck = [Card(color=1, rank=8)] + expeditions = _expeditions(config) + expeditions[0][0] = [Card(color=0, rank=4)] + state = make_state( + config, + deck=[Card(color=1, rank=8)], + expeditions=expeditions, + discards=[[Card(color=0, rank=6)], []], + phase="draw", + ) assert bot._act_draw(state) == draw_from_discard_action(0) @@ -106,20 +136,23 @@ def test_safe_heuristic_can_draw_discard_to_deny_opponent_when_losing() -> None: config = LostCitiesConfig(n_colors=2, n_ranks=8, hand_size=4) bot = SafeHeuristicBot() - state = GameState.empty(config) - state.current_player = 0 - state.phase = "draw" - state.deck = [Card(color=1, rank=8), Card(color=1, rank=7)] - state.hands[0] = [Card(color=0, rank=0), Card(color=0, rank=7)] - state.expeditions[0][1] = [Card(color=1, rank=8)] - state.expeditions[1][0] = [ + expeditions = _expeditions(config) + expeditions[0][1] = [Card(color=1, rank=8)] + expeditions[1][0] = [ Card(color=0, rank=0), Card(color=0, rank=5), Card(color=0, rank=6), Card(color=0, rank=7), Card(color=0, rank=8), ] - state.discards[0] = [Card(color=0, rank=6)] + state = make_state( + config, + deck=[Card(color=1, rank=8), Card(color=1, rank=7)], + hands=[[Card(color=0, rank=0), Card(color=0, rank=7)], []], + expeditions=expeditions, + discards=[[Card(color=0, rank=6)], []], + phase="draw", + ) assert state.score_diff(0) < 0 assert bot._act_draw(state) == draw_from_discard_action(0) @@ -152,16 +185,17 @@ def test_safe_heuristic_classic_self_play_opens_expeditions() -> None: def test_safe_heuristic_avoids_opening_weak_fifth_color() -> None: config = LostCitiesConfig(n_colors=5, n_ranks=8, hand_size=8) bot = SafeHeuristicBot() - state = GameState.empty(config) - state.current_player = 0 - state.phase = "card" - - state.expeditions[0][0] = [Card(color=0, rank=4)] - state.expeditions[0][1] = [Card(color=1, rank=4)] - state.expeditions[0][2] = [Card(color=2, rank=5)] - state.expeditions[0][3] = [Card(color=3, rank=6)] + expeditions = _expeditions(config) + expeditions[0][0] = [Card(color=0, rank=4)] + expeditions[0][1] = [Card(color=1, rank=4)] + expeditions[0][2] = [Card(color=2, rank=5)] + expeditions[0][3] = [Card(color=3, rank=6)] weak_open = Card(color=4, rank=4) - state.hands[0] = [weak_open, Card(color=4, rank=7), Card(color=0, rank=6)] + state = make_state( + config, + hands=[[weak_open, Card(color=4, rank=7), Card(color=0, rank=6)], []], + expeditions=expeditions, + ) state.sort_hand(0) assert ( @@ -180,11 +214,16 @@ def test_safe_heuristic_avoids_opening_weak_fifth_color() -> None: def test_safe_heuristic_prefers_followup_on_started_expedition() -> None: config = LostCitiesConfig(n_colors=3, n_ranks=8, hand_size=5) bot = SafeHeuristicBot() - state = GameState.empty(config) - state.current_player = 0 - state.phase = "card" - state.expeditions[0][0] = [Card(color=0, rank=4)] - state.hands[0] = [Card(color=0, rank=6), Card(color=1, rank=4), Card(color=1, rank=7)] + expeditions = _expeditions(config) + expeditions[0][0] = [Card(color=0, rank=4)] + state = make_state( + config, + hands=[ + [Card(color=0, rank=6), Card(color=1, rank=4), Card(color=1, rank=7)], + [], + ], + expeditions=expeditions, + ) state.sort_hand(0) action = bot._act_card(state) @@ -197,15 +236,20 @@ def test_safe_heuristic_prefers_followup_on_started_expedition() -> None: def test_safe_heuristic_avoids_unopened_discard_draw_after_four_opens() -> None: config = LostCitiesConfig(n_colors=5, n_ranks=8, hand_size=8) bot = SafeHeuristicBot() - state = GameState.empty(config) - state.current_player = 0 - state.phase = "draw" - state.deck = [Card(color=0, rank=8), Card(color=1, rank=8)] - state.expeditions[0][0] = [Card(color=0, rank=4)] - state.expeditions[0][1] = [Card(color=1, rank=4)] - state.expeditions[0][2] = [Card(color=2, rank=5)] - state.expeditions[0][3] = [Card(color=3, rank=6)] - state.hands[0] = [Card(color=4, rank=4), Card(color=4, rank=7)] - state.discards[4] = [Card(color=4, rank=5)] + expeditions = _expeditions(config) + expeditions[0][0] = [Card(color=0, rank=4)] + expeditions[0][1] = [Card(color=1, rank=4)] + expeditions[0][2] = [Card(color=2, rank=5)] + expeditions[0][3] = [Card(color=3, rank=6)] + discards = [[] for _ in range(config.n_colors)] + discards[4] = [Card(color=4, rank=5)] + state = make_state( + config, + deck=[Card(color=0, rank=8), Card(color=1, rank=8)], + hands=[[Card(color=4, rank=4), Card(color=4, rank=7)], []], + expeditions=expeditions, + discards=discards, + phase="draw", + ) assert bot._act_draw(state) == 0 diff --git a/tests/games/classic/test_canonical_state.py b/tests/games/classic/test_canonical_state.py index d099fad..7ea4d6a 100644 --- a/tests/games/classic/test_canonical_state.py +++ b/tests/games/classic/test_canonical_state.py @@ -6,6 +6,7 @@ import pytest from coolrl_lost_cities.games.classic.game import Card, GameState, LostCitiesConfig import coolrl_lost_cities.games.classic as classic +from tests.games.classic.helpers import make_state FIXTURE_DIR = Path(classic.__file__).resolve().parent / "fixtures" @@ -85,17 +86,21 @@ def test_snapshot_roundtrip_preserves_json_state() -> None: def test_validate_invariants_detects_card_loss() -> None: state = GameState.new_game(LostCitiesConfig(seed=7)) - state.deck.pop() + snapshot = state.to_snapshot() + snapshot["deck"].pop() + broken = GameState.from_snapshot(snapshot, validate=False) with pytest.raises(ValueError, match="card conservation"): - state.validate_invariants() + broken.validate_invariants() def test_validate_invariants_detects_bad_expedition_order() -> None: - state = GameState.new_game(LostCitiesConfig(seed=8)) - card = state.deck.pop() - state.expeditions[0][card.color].extend([Card(card.color, 2), Card(card.color, 1)]) - state.deck.extend([Card(card.color, 2), Card(card.color, 1)]) + config = LostCitiesConfig(seed=8) + state = make_state( + config, + deck=GameState.new_game(config).deck, + expeditions=[[[Card(0, 2), Card(0, 1)], [], [], [], []], [[], [], [], [], []]], + ) with pytest.raises(ValueError, match="strictly increasing"): state.validate_invariants() diff --git a/tests/games/classic/test_env.py b/tests/games/classic/test_env.py index afe2734..756a2c4 100644 --- a/tests/games/classic/test_env.py +++ b/tests/games/classic/test_env.py @@ -1,7 +1,8 @@ import numpy as np -from coolrl_lost_cities.games.classic.game import Card, GameState, LostCitiesConfig +from coolrl_lost_cities.games.classic.game import Card, LostCitiesConfig from coolrl_lost_cities.games.classic.env import LostCitiesEnv +from tests.games.classic.helpers import make_state def test_env_observation_uses_fixed_unified_mask() -> None: @@ -24,11 +25,12 @@ def test_env_observation_uses_fixed_unified_mask() -> None: def test_env_step_accepts_legacy_draw_action_ids() -> None: config = LostCitiesConfig() env = LostCitiesEnv(config) - env.state = GameState.empty(config) - env.state.hands[0] = [Card(0, 1)] - env.state.hands[1] = [Card(1, 1)] - env.state.deck = [Card(2, 1), Card(2, 2)] - env.state.phase = "draw" + env.state = make_state( + config, + deck=[Card(2, 1), Card(2, 2)], + hands=[[Card(0, 1)], [Card(1, 1)]], + phase="draw", + ) obs, reward, done, _ = env.step(0) @@ -50,11 +52,13 @@ def test_terminal_reward_is_relative_to_actor() -> None: bonus_threshold=99, ) env = LostCitiesEnv(config) - env.state = GameState.empty(config) - env.state.current_player = 1 - env.state.phase = "draw" - env.state.deck = [Card(1, 1)] - env.state.expeditions[1][0] = [Card(0, 1)] + env.state = make_state( + config, + deck=[Card(1, 1)], + expeditions=[[[], []], [[Card(0, 1)], []]], + current_player=1, + phase="draw", + ) _, reward, done, _ = env.step(config.card_action_size) diff --git a/tests/games/classic/test_masks.py b/tests/games/classic/test_masks.py index a29eec8..7d0e1d9 100644 --- a/tests/games/classic/test_masks.py +++ b/tests/games/classic/test_masks.py @@ -1,6 +1,7 @@ from coolrl_lost_cities.games.classic.game import Card, GameState, LostCitiesConfig from coolrl_lost_cities.games.classic.bots import RandomBot +from tests.games.classic.helpers import make_state def test_legal_mask_has_action_in_nonterminal_phases() -> None: @@ -12,8 +13,7 @@ def test_legal_mask_has_action_in_nonterminal_phases() -> None: def test_empty_hand_slots_are_masked() -> None: - state = GameState.empty(LostCitiesConfig()) - state.hands[0] = [Card(0, 1)] + state = make_state(hands=[[Card(0, 1)], []]) mask = state.legal_card_mask() assert mask[0] is True assert mask[1] is True @@ -21,9 +21,7 @@ def test_empty_hand_slots_are_masked() -> None: def test_empty_discard_pile_draw_is_illegal() -> None: - state = GameState.empty(LostCitiesConfig()) - state.phase = "draw" - state.deck = [Card(0, 1)] + state = make_state(deck=[Card(0, 1)], phase="draw") mask = state.legal_draw_mask() assert mask[0] is True assert all(mask[1 + color] is False for color in range(state.config.n_colors)) diff --git a/tests/games/classic/test_rules.py b/tests/games/classic/test_rules.py index 6bde94a..4d3d0e4 100644 --- a/tests/games/classic/test_rules.py +++ b/tests/games/classic/test_rules.py @@ -7,6 +7,8 @@ from coolrl_lost_cities.games.classic.game import ( build_deck, ) +from tests.games.classic.helpers import make_state + def test_deck_generation_count() -> None: config = LostCitiesConfig(n_colors=3, n_ranks=5, n_handshakes=1, hand_size=5) @@ -23,25 +25,27 @@ def test_initial_hands_remove_cards_from_deck() -> None: def test_play_must_be_ascending() -> None: config = LostCitiesConfig() - state = GameState.empty(config) - state.hands[0] = [Card(0, 2)] - state.expeditions[0][0] = [Card(0, 4)] + state = make_state( + config, + hands=[[Card(0, 2)], []], + expeditions=[[[Card(0, 4)], [], [], [], []], [[], [], [], [], []]], + ) assert state.legal_card_mask()[0] is False def test_handshake_after_number_forbidden() -> None: config = LostCitiesConfig() - state = GameState.empty(config) - state.hands[0] = [Card(1, 0)] - state.expeditions[0][1] = [Card(1, 1)] + state = make_state( + config, + hands=[[Card(1, 0)], []], + expeditions=[[[], [Card(1, 1)], [], [], []], [[], [], [], [], []]], + ) assert state.legal_card_mask()[0] is False def test_cannot_draw_just_discarded_color() -> None: config = LostCitiesConfig() - state = GameState.empty(config) - state.hands[0] = [Card(2, 2)] - state.deck = [Card(0, 1)] + state = make_state(config, deck=[Card(0, 1)], hands=[[Card(2, 2)], []]) state.apply_action(1) mask = state.legal_draw_mask() assert mask[1 + 2] is False @@ -49,9 +53,7 @@ def test_cannot_draw_just_discarded_color() -> None: def test_drawing_just_discarded_color_is_rejected() -> None: config = LostCitiesConfig() - state = GameState.empty(config) - state.hands[0] = [Card(2, 2)] - state.deck = [Card(0, 1)] + state = make_state(config, deck=[Card(0, 1)], hands=[[Card(2, 2)], []]) state.apply_action(1) @@ -61,10 +63,11 @@ def test_drawing_just_discarded_color_is_rejected() -> None: def test_discarded_color_can_be_drawn_after_turn_advances() -> None: config = LostCitiesConfig() - state = GameState.empty(config) - state.hands[0] = [Card(2, 2)] - state.hands[1] = [Card(0, 1)] - state.deck = [Card(1, 1), Card(1, 2)] + state = make_state( + config, + deck=[Card(1, 1), Card(1, 2)], + hands=[[Card(2, 2)], [Card(0, 1)]], + ) state.apply_action(1) state.apply_action(0) assert state.current_player == 1 @@ -75,10 +78,11 @@ def test_discarded_color_can_be_drawn_after_turn_advances() -> None: def test_discarded_card_is_removed_when_drawn_later() -> None: config = LostCitiesConfig() - state = GameState.empty(config) - state.hands[0] = [Card(2, 2)] - state.hands[1] = [Card(0, 1)] - state.deck = [Card(1, 1), Card(1, 2)] + state = make_state( + config, + deck=[Card(1, 1), Card(1, 2)], + hands=[[Card(2, 2)], [Card(0, 1)]], + ) state.apply_action(1) assert state.discards[2] == [Card(2, 2)] @@ -95,9 +99,7 @@ def test_discarded_card_is_removed_when_drawn_later() -> None: def test_deck_exhaustion_ends_after_last_deck_draw() -> None: config = LostCitiesConfig() - state = GameState.empty(config) - state.hands[0] = [Card(0, 1)] - state.deck = [Card(1, 1)] + state = make_state(config, deck=[Card(1, 1)], hands=[[Card(0, 1)], []]) state.apply_action(1) state.apply_action(0) assert state.terminal is True @@ -106,10 +108,7 @@ def test_deck_exhaustion_ends_after_last_deck_draw() -> None: def test_card_phase_can_end_game_when_no_draw_sources_exist() -> None: config = LostCitiesConfig(n_colors=3, n_ranks=5, n_handshakes=1, hand_size=5) - state = GameState.empty(config) - state.hands[0] = [Card(0, 1)] - state.hands[1] = [Card(1, 1)] - state.deck = [] + state = make_state(config, hands=[[Card(0, 1)], [Card(1, 1)]]) state.apply_action(1) assert state.phase == "draw" assert state.terminal is True