고속 엔진 안전성 보강
This commit is contained in:
@@ -0,0 +1,16 @@
|
||||
# Fast Engine Follow-up Optimizations
|
||||
|
||||
The current fast engine exposes Python wrappers for testing and debugging, but
|
||||
serious traversal code should use the Cython `fast.pxd` API directly.
|
||||
|
||||
Deferred work:
|
||||
|
||||
1. Add an internal undo stack with `push_action()` / `pop_action()` so Python
|
||||
callers can avoid tuple allocation when they need nested search.
|
||||
2. Keep traversal legal-action generation caller-buffer based. Consider a
|
||||
reusable Python-wrapper action buffer only if wrapper profiling shows
|
||||
`legal_actions()` allocation is material.
|
||||
3. Consider direct NumPy or feature-buffer output for RL pipelines instead of
|
||||
building Python lists and converting later.
|
||||
4. Consider a single contiguous allocation for state arrays after profiling the
|
||||
simpler separate-allocation layout.
|
||||
@@ -88,7 +88,7 @@ cdef class FastGameState:
|
||||
cdef void _undo_draw_action_c(self, UndoRecord* undo) except *
|
||||
cdef void _recompute_score_caches(self) noexcept
|
||||
cdef int _score_from_summary_c(self, int length, int handshakes, int numeric_sum) noexcept
|
||||
cdef bint _has_any_legal_draw(self)
|
||||
cdef bint _has_any_legal_draw(self) noexcept
|
||||
cdef int _hand_index(self, int player, int slot)
|
||||
cdef int _expedition_len_index(self, int player, int color)
|
||||
cdef int _expedition_index(self, int player, int color, int index)
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
from collections import Counter
|
||||
import random
|
||||
|
||||
from libc.string cimport memcpy
|
||||
from libc.stdlib cimport free, malloc
|
||||
|
||||
from ..game import IllegalMoveError, LostCitiesConfig, config_from_mapping
|
||||
@@ -137,6 +138,10 @@ cdef class FastGameState:
|
||||
config = config or LostCitiesConfig()
|
||||
config.validate()
|
||||
encoded = [_encode_card_snapshot(card, config) for card in deck]
|
||||
if len(encoded) != int(config.deck_size):
|
||||
raise ValueError(
|
||||
f"deck length must be {config.deck_size}, got {len(encoded)}"
|
||||
)
|
||||
if Counter(encoded) != Counter(_build_encoded_deck(config)):
|
||||
raise ValueError("deck must contain exactly the cards defined by config")
|
||||
|
||||
@@ -166,6 +171,10 @@ cdef class FastGameState:
|
||||
cdef list cards
|
||||
|
||||
cards = [_encode_card_snapshot(card, config) for card in snapshot["deck"]]
|
||||
if len(cards) > state.total_cards:
|
||||
raise ValueError(
|
||||
f"deck snapshot exceeds capacity {state.total_cards}: {len(cards)}"
|
||||
)
|
||||
state.deck_len = len(cards)
|
||||
for index, card in enumerate(cards):
|
||||
state.deck[index] = <int>card
|
||||
@@ -174,6 +183,11 @@ cdef class FastGameState:
|
||||
cards = [
|
||||
_encode_card_snapshot(card, config) for card in snapshot["hands"][player]
|
||||
]
|
||||
if len(cards) > state.hand_size:
|
||||
raise ValueError(
|
||||
f"hand {player} snapshot exceeds hand_size "
|
||||
f"{state.hand_size}: {len(cards)}"
|
||||
)
|
||||
state.hand_lens[player] = len(cards)
|
||||
for index, card in enumerate(cards):
|
||||
state.hands[state._hand_index(player, index)] = <int>card
|
||||
@@ -184,12 +198,22 @@ cdef class FastGameState:
|
||||
_encode_card_snapshot(card, config)
|
||||
for card in snapshot["expeditions"][player][color]
|
||||
]
|
||||
if len(cards) > state.cards_per_color:
|
||||
raise ValueError(
|
||||
f"expedition {player}/{color} snapshot exceeds capacity "
|
||||
f"{state.cards_per_color}: {len(cards)}"
|
||||
)
|
||||
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)] = <int>card
|
||||
|
||||
for color in range(state.n_colors):
|
||||
cards = [_encode_card_snapshot(card, config) for card in snapshot["discards"][color]]
|
||||
if len(cards) > state.cards_per_color:
|
||||
raise ValueError(
|
||||
f"discard {color} snapshot exceeds capacity "
|
||||
f"{state.cards_per_color}: {len(cards)}"
|
||||
)
|
||||
state.discard_lens[color] = len(cards)
|
||||
for index, card in enumerate(cards):
|
||||
state.discards[state._discard_index(color, index)] = <int>card
|
||||
@@ -275,26 +299,43 @@ cdef class FastGameState:
|
||||
|
||||
cpdef FastGameState clone(self):
|
||||
cdef FastGameState other = FastGameState(self.config)
|
||||
cdef int i
|
||||
other.deck_len = self.deck_len
|
||||
for i in range(self.deck_len):
|
||||
other.deck[i] = self.deck[i]
|
||||
for i in range(2 * self.hand_size):
|
||||
other.hands[i] = self.hands[i]
|
||||
memcpy(other.deck, self.deck, self.deck_len * sizeof(int))
|
||||
memcpy(other.hands, self.hands, 2 * self.hand_size * sizeof(int))
|
||||
other.hand_lens[0] = self.hand_lens[0]
|
||||
other.hand_lens[1] = self.hand_lens[1]
|
||||
for i in range(2 * self.n_colors * self.cards_per_color):
|
||||
other.expeditions[i] = self.expeditions[i]
|
||||
for i in range(2 * self.n_colors):
|
||||
other.expedition_lens[i] = self.expedition_lens[i]
|
||||
other.last_numeric_ranks[i] = self.last_numeric_ranks[i]
|
||||
other.handshake_counts[i] = self.handshake_counts[i]
|
||||
other.numeric_sums[i] = self.numeric_sums[i]
|
||||
other.expedition_scores[i] = self.expedition_scores[i]
|
||||
for i in range(self.n_colors * self.cards_per_color):
|
||||
other.discards[i] = self.discards[i]
|
||||
for i in range(self.n_colors):
|
||||
other.discard_lens[i] = self.discard_lens[i]
|
||||
memcpy(
|
||||
other.expeditions,
|
||||
self.expeditions,
|
||||
2 * self.n_colors * self.cards_per_color * sizeof(int),
|
||||
)
|
||||
memcpy(
|
||||
other.expedition_lens,
|
||||
self.expedition_lens,
|
||||
2 * self.n_colors * sizeof(int),
|
||||
)
|
||||
memcpy(
|
||||
other.discards,
|
||||
self.discards,
|
||||
self.n_colors * self.cards_per_color * sizeof(int),
|
||||
)
|
||||
memcpy(other.discard_lens, self.discard_lens, self.n_colors * sizeof(int))
|
||||
memcpy(
|
||||
other.last_numeric_ranks,
|
||||
self.last_numeric_ranks,
|
||||
2 * self.n_colors * sizeof(int),
|
||||
)
|
||||
memcpy(
|
||||
other.handshake_counts,
|
||||
self.handshake_counts,
|
||||
2 * self.n_colors * sizeof(int),
|
||||
)
|
||||
memcpy(other.numeric_sums, self.numeric_sums, 2 * self.n_colors * sizeof(int))
|
||||
memcpy(
|
||||
other.expedition_scores,
|
||||
self.expedition_scores,
|
||||
2 * self.n_colors * sizeof(int),
|
||||
)
|
||||
other.total_scores[0] = self.total_scores[0]
|
||||
other.total_scores[1] = self.total_scores[1]
|
||||
other.current_player = self.current_player
|
||||
@@ -312,7 +353,7 @@ cdef class FastGameState:
|
||||
return mask
|
||||
for slot in range(self.hand_lens[self.current_player]):
|
||||
card = self.hands[self._hand_index(self.current_player, slot)]
|
||||
mask[2 * slot] = self.can_play_encoded_card(self.current_player, card)
|
||||
mask[2 * slot] = self._can_play_encoded_card_c(self.current_player, card)
|
||||
mask[2 * slot + 1] = True
|
||||
return mask
|
||||
|
||||
@@ -453,14 +494,51 @@ cdef class FastGameState:
|
||||
|
||||
def validate_invariants(self):
|
||||
self.config.validate()
|
||||
cdef int player
|
||||
cdef int color
|
||||
cdef int index
|
||||
cdef int length
|
||||
cdef int card
|
||||
cdef int rank
|
||||
cdef int last_rank
|
||||
cdef bint seen_numeric
|
||||
if self.current_player not in (0, 1):
|
||||
raise ValueError("current_player must be 0 or 1")
|
||||
if self.phase_id not in (_phase_card(), _phase_draw()):
|
||||
raise ValueError("invalid phase")
|
||||
if self.deck_len < 0 or self.deck_len > self.total_cards:
|
||||
raise ValueError("deck length out of range")
|
||||
if self.pending_discarded_color >= self.n_colors:
|
||||
raise ValueError("pending_discarded_color is out of range")
|
||||
if self.hand_lens[0] > self.hand_size or self.hand_lens[1] > self.hand_size:
|
||||
raise ValueError("hand exceeds hand_size")
|
||||
for color in range(self.n_colors):
|
||||
if self.discard_lens[color] < 0 or self.discard_lens[color] > self.cards_per_color:
|
||||
raise ValueError("discard length out of range")
|
||||
for player in range(2):
|
||||
if self.hand_lens[player] < 0:
|
||||
raise ValueError("hand length out of range")
|
||||
for color in range(self.n_colors):
|
||||
length = self.expedition_lens[self._expedition_len_index(player, color)]
|
||||
if length < 0 or length > self.cards_per_color:
|
||||
raise ValueError("expedition length out of range")
|
||||
seen_numeric = False
|
||||
last_rank = 0
|
||||
for index in range(length):
|
||||
card = self.expeditions[self._expedition_index(player, color, index)]
|
||||
if self._card_color(card) != color:
|
||||
raise ValueError("expedition contains wrong color")
|
||||
rank = self._card_rank(card)
|
||||
if rank < 0 or rank > self.n_ranks:
|
||||
raise ValueError("card rank out of range")
|
||||
if rank == 0:
|
||||
if seen_numeric:
|
||||
raise ValueError("expedition has handshake after number")
|
||||
else:
|
||||
seen_numeric = True
|
||||
if rank <= last_rank:
|
||||
raise ValueError("expedition is not strictly increasing")
|
||||
last_rank = rank
|
||||
if Counter(_all_cards_from_snapshot(self.to_snapshot())) != Counter(
|
||||
_build_encoded_deck(self.config)
|
||||
):
|
||||
@@ -698,6 +776,8 @@ cdef class FastGameState:
|
||||
self.discard_lens[color] += 1
|
||||
self.pending_discarded_color = color
|
||||
self.phase_id = _phase_draw()
|
||||
# Defensive terminal branch for externally constructed states where the
|
||||
# deck was already empty before the card phase action.
|
||||
if self.deck_len == 0 and not self._has_any_legal_draw():
|
||||
self.terminal = True
|
||||
|
||||
@@ -842,7 +922,7 @@ cdef class FastGameState:
|
||||
score += self.bonus_amount
|
||||
return score
|
||||
|
||||
cdef bint _has_any_legal_draw(self):
|
||||
cdef bint _has_any_legal_draw(self) noexcept:
|
||||
cdef int color
|
||||
if self.deck_len > 0:
|
||||
return True
|
||||
|
||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import random
|
||||
|
||||
import pytest
|
||||
from coolrl_lost_cities.games.classic.game import Card, GameState, LostCitiesConfig
|
||||
|
||||
from coolrl_lost_cities.games.classic.engines import FastGameState
|
||||
@@ -46,6 +47,61 @@ def test_fast_snapshot_roundtrip_preserves_snapshot() -> None:
|
||||
assert restored.to_snapshot() == fast.to_snapshot()
|
||||
|
||||
|
||||
def test_fast_from_snapshot_rejects_oversized_regions_before_write() -> None:
|
||||
config = _small_config()
|
||||
state = GameState.new_game(config, seed=3)
|
||||
|
||||
deck_snapshot = state.to_snapshot()
|
||||
deck_snapshot["deck"] = deck_snapshot["deck"] + [
|
||||
{"color": 0, "rank": 1},
|
||||
{"color": 0, "rank": 2},
|
||||
{"color": 1, "rank": 1},
|
||||
]
|
||||
with pytest.raises(ValueError, match="deck snapshot exceeds capacity"):
|
||||
FastGameState.from_snapshot(deck_snapshot)
|
||||
|
||||
hand_snapshot = state.to_snapshot()
|
||||
hand_snapshot["hands"][0] = hand_snapshot["hands"][0] + [{"color": 0, "rank": 1}]
|
||||
with pytest.raises(ValueError, match="hand 0 snapshot exceeds hand_size"):
|
||||
FastGameState.from_snapshot(hand_snapshot)
|
||||
|
||||
expedition_snapshot = state.to_snapshot()
|
||||
expedition_snapshot["expeditions"][0][0] = [
|
||||
{"color": 0, "rank": 1},
|
||||
{"color": 0, "rank": 2},
|
||||
{"color": 0, "rank": 1},
|
||||
]
|
||||
with pytest.raises(ValueError, match="expedition 0/0 snapshot exceeds capacity"):
|
||||
FastGameState.from_snapshot(expedition_snapshot)
|
||||
|
||||
discard_snapshot = state.to_snapshot()
|
||||
discard_snapshot["discards"][0] = [
|
||||
{"color": 0, "rank": 1},
|
||||
{"color": 0, "rank": 2},
|
||||
{"color": 0, "rank": 1},
|
||||
]
|
||||
with pytest.raises(ValueError, match="discard 0 snapshot exceeds capacity"):
|
||||
FastGameState.from_snapshot(discard_snapshot)
|
||||
|
||||
|
||||
def test_fast_validate_invariants_rejects_bad_expedition_order() -> None:
|
||||
config = _small_config()
|
||||
snapshot = FastGameState.new_game(config, seed=4).to_snapshot()
|
||||
snapshot["deck"].extend(
|
||||
[
|
||||
{"color": 0, "rank": 2},
|
||||
{"color": 0, "rank": 1},
|
||||
]
|
||||
)
|
||||
snapshot["expeditions"][0][0] = [
|
||||
{"color": 0, "rank": 2},
|
||||
{"color": 0, "rank": 1},
|
||||
]
|
||||
|
||||
with pytest.raises(ValueError, match="expedition is not strictly increasing"):
|
||||
FastGameState.from_snapshot(snapshot)
|
||||
|
||||
|
||||
def test_fast_random_action_sequence_matches_game_state() -> None:
|
||||
config = LostCitiesConfig(
|
||||
n_colors=3,
|
||||
|
||||
Reference in New Issue
Block a user