고속 엔진 안전성 보강

This commit is contained in:
2026-05-06 21:56:59 +09:00
parent 2e38fc6454
commit a349cf34ed
4 changed files with 172 additions and 20 deletions
+16
View File
@@ -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 _undo_draw_action_c(self, UndoRecord* undo) except *
cdef void _recompute_score_caches(self) noexcept cdef void _recompute_score_caches(self) noexcept
cdef int _score_from_summary_c(self, int length, int handshakes, int numeric_sum) 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 _hand_index(self, int player, int slot)
cdef int _expedition_len_index(self, int player, int color) cdef int _expedition_len_index(self, int player, int color)
cdef int _expedition_index(self, int player, int color, int index) cdef int _expedition_index(self, int player, int color, int index)
@@ -4,6 +4,7 @@
from collections import Counter from collections import Counter
import random import random
from libc.string cimport memcpy
from libc.stdlib cimport free, malloc from libc.stdlib cimport free, malloc
from ..game import IllegalMoveError, LostCitiesConfig, config_from_mapping from ..game import IllegalMoveError, LostCitiesConfig, config_from_mapping
@@ -137,6 +138,10 @@ cdef class FastGameState:
config = config or LostCitiesConfig() config = config or LostCitiesConfig()
config.validate() config.validate()
encoded = [_encode_card_snapshot(card, config) for card in deck] 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)): if Counter(encoded) != Counter(_build_encoded_deck(config)):
raise ValueError("deck must contain exactly the cards defined by config") raise ValueError("deck must contain exactly the cards defined by config")
@@ -166,6 +171,10 @@ cdef class FastGameState:
cdef list cards cdef list cards
cards = [_encode_card_snapshot(card, config) for card in snapshot["deck"]] 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) state.deck_len = len(cards)
for index, card in enumerate(cards): for index, card in enumerate(cards):
state.deck[index] = <int>card state.deck[index] = <int>card
@@ -174,6 +183,11 @@ cdef class FastGameState:
cards = [ cards = [
_encode_card_snapshot(card, config) for card in snapshot["hands"][player] _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) state.hand_lens[player] = len(cards)
for index, card in enumerate(cards): for index, card in enumerate(cards):
state.hands[state._hand_index(player, index)] = <int>card state.hands[state._hand_index(player, index)] = <int>card
@@ -184,12 +198,22 @@ cdef class FastGameState:
_encode_card_snapshot(card, config) _encode_card_snapshot(card, config)
for card in snapshot["expeditions"][player][color] 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) state.expedition_lens[state._expedition_len_index(player, color)] = len(cards)
for index, card in enumerate(cards): for index, card in enumerate(cards):
state.expeditions[state._expedition_index(player, color, index)] = <int>card state.expeditions[state._expedition_index(player, color, index)] = <int>card
for color in range(state.n_colors): for color in range(state.n_colors):
cards = [_encode_card_snapshot(card, config) for card in snapshot["discards"][color]] 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) state.discard_lens[color] = len(cards)
for index, card in enumerate(cards): for index, card in enumerate(cards):
state.discards[state._discard_index(color, index)] = <int>card state.discards[state._discard_index(color, index)] = <int>card
@@ -275,26 +299,43 @@ cdef class FastGameState:
cpdef FastGameState clone(self): cpdef FastGameState clone(self):
cdef FastGameState other = FastGameState(self.config) cdef FastGameState other = FastGameState(self.config)
cdef int i
other.deck_len = self.deck_len other.deck_len = self.deck_len
for i in range(self.deck_len): memcpy(other.deck, self.deck, self.deck_len * sizeof(int))
other.deck[i] = self.deck[i] memcpy(other.hands, self.hands, 2 * self.hand_size * sizeof(int))
for i in range(2 * self.hand_size):
other.hands[i] = self.hands[i]
other.hand_lens[0] = self.hand_lens[0] other.hand_lens[0] = self.hand_lens[0]
other.hand_lens[1] = self.hand_lens[1] other.hand_lens[1] = self.hand_lens[1]
for i in range(2 * self.n_colors * self.cards_per_color): memcpy(
other.expeditions[i] = self.expeditions[i] other.expeditions,
for i in range(2 * self.n_colors): self.expeditions,
other.expedition_lens[i] = self.expedition_lens[i] 2 * self.n_colors * self.cards_per_color * sizeof(int),
other.last_numeric_ranks[i] = self.last_numeric_ranks[i] )
other.handshake_counts[i] = self.handshake_counts[i] memcpy(
other.numeric_sums[i] = self.numeric_sums[i] other.expedition_lens,
other.expedition_scores[i] = self.expedition_scores[i] self.expedition_lens,
for i in range(self.n_colors * self.cards_per_color): 2 * self.n_colors * sizeof(int),
other.discards[i] = self.discards[i] )
for i in range(self.n_colors): memcpy(
other.discard_lens[i] = self.discard_lens[i] 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[0] = self.total_scores[0]
other.total_scores[1] = self.total_scores[1] other.total_scores[1] = self.total_scores[1]
other.current_player = self.current_player other.current_player = self.current_player
@@ -312,7 +353,7 @@ cdef class FastGameState:
return mask return mask
for slot in range(self.hand_lens[self.current_player]): for slot in range(self.hand_lens[self.current_player]):
card = self.hands[self._hand_index(self.current_player, slot)] 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 mask[2 * slot + 1] = True
return mask return mask
@@ -453,14 +494,51 @@ cdef class FastGameState:
def validate_invariants(self): def validate_invariants(self):
self.config.validate() 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): if self.current_player not in (0, 1):
raise ValueError("current_player must be 0 or 1") raise ValueError("current_player must be 0 or 1")
if self.phase_id not in (_phase_card(), _phase_draw()): if self.phase_id not in (_phase_card(), _phase_draw()):
raise ValueError("invalid phase") 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: if self.pending_discarded_color >= self.n_colors:
raise ValueError("pending_discarded_color is out of range") 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: if self.hand_lens[0] > self.hand_size or self.hand_lens[1] > self.hand_size:
raise ValueError("hand exceeds 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( if Counter(_all_cards_from_snapshot(self.to_snapshot())) != Counter(
_build_encoded_deck(self.config) _build_encoded_deck(self.config)
): ):
@@ -698,6 +776,8 @@ cdef class FastGameState:
self.discard_lens[color] += 1 self.discard_lens[color] += 1
self.pending_discarded_color = color self.pending_discarded_color = color
self.phase_id = _phase_draw() 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(): if self.deck_len == 0 and not self._has_any_legal_draw():
self.terminal = True self.terminal = True
@@ -842,7 +922,7 @@ cdef class FastGameState:
score += self.bonus_amount score += self.bonus_amount
return score return score
cdef bint _has_any_legal_draw(self): cdef bint _has_any_legal_draw(self) noexcept:
cdef int color cdef int color
if self.deck_len > 0: if self.deck_len > 0:
return True return True
@@ -2,6 +2,7 @@ from __future__ import annotations
import random import random
import pytest
from coolrl_lost_cities.games.classic.game import Card, GameState, LostCitiesConfig from coolrl_lost_cities.games.classic.game import Card, GameState, LostCitiesConfig
from coolrl_lost_cities.games.classic.engines import FastGameState 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() 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: def test_fast_random_action_sequence_matches_game_state() -> None:
config = LostCitiesConfig( config = LostCitiesConfig(
n_colors=3, n_colors=3,