From 5d3062a4ee2760c161d4c0fce0ecda8e6022a5d2 Mon Sep 17 00:00:00 2001 From: Brandon Nguyen <58405975+Bratah123@users.noreply.github.com> Date: Thu, 16 Jul 2026 00:27:44 -0700 Subject: [PATCH] FEAT: fix tournament deck validator --- spirit/network/message_names.py | 1 + spirit/packets/handlers/tournaments.py | 116 +++++++++++++++++++---- tests/test_tournament_deck_validation.py | 77 +++++++++++++++ 3 files changed, 178 insertions(+), 16 deletions(-) create mode 100644 tests/test_tournament_deck_validation.py diff --git a/spirit/network/message_names.py b/spirit/network/message_names.py index f45af09..2d13c99 100644 --- a/spirit/network/message_names.py +++ b/spirit/network/message_names.py @@ -334,6 +334,7 @@ class OutboundMsg(str, Enum): TOURNAMENT_QUEUE_LEFT = "TournamentQueueLeft" TOURNAMENT_QUEUE_LEFT_FAILED = "TournamentQueueLeftFailed" JOIN_TOURNAMENT_FAILED = "JoinTournamentFailed" + JOIN_TOURNAMENT_FAILED_INVALID_DECK = "JoinTournamentFailedInvalidDeck" TOURNAMENT_STARTED = "TournamentStarted" TOURNAMENT_ROUND_UPDATED = "TournamentRoundUpdated" TOURNAMENT_NEXT_ROUND_STARTING = "TournamentNextRoundStarting" diff --git a/spirit/packets/handlers/tournaments.py b/spirit/packets/handlers/tournaments.py index 4e59d6b..8382d8e 100644 --- a/spirit/packets/handlers/tournaments.py +++ b/spirit/packets/handlers/tournaments.py @@ -4,6 +4,9 @@ import uuid from spirit.network.message_names import InboundMsg, OutboundMsg from spirit.database.async_utils import run_db from spirit.database import tournament_data +from spirit.database.player_data import get_owned_counts +from spirit.game import rules +from spirit.game.format_manager import FormatManager from spirit.game.tournament_manager import ( TournamentManager, STATE_OPEN, STATE_ENTRY_CLOSED, STATE_RESOLVED, STATE_HIDDEN, now_ms, @@ -34,6 +37,37 @@ def serializable_deck(deck_data: dict, name: str = "Tournament Deck") -> dict: } +def tournament_format_value(tournament, legacy: bool = False): + """Returns the configured format identifier for one tournament flow.""" + definition = getattr(tournament, "definition", {}) or {} + if legacy: + return definition.get("format") or "Unlimited" + game = definition.get("game") or {} + return game.get("format") or definition.get("format") or "Unlimited" + + +def validate_tournament_deck(deck: dict, tournament, owned_counts=None, + legacy: bool = False): + """Validates a tournament deck and returns (results, format_error).""" + format_value = tournament_format_value(tournament, legacy=legacy) + format_guid = FormatManager().resolve_format_guid(format_value) + if format_guid is None: + return [], f"The {format_value} tournament format is not available on this server." + return rules.validate_deck( + deck, [format_guid], owned_counts=owned_counts), None + + +def validation_error_text(results: list) -> str: + """Returns the first human-readable deck validation failure.""" + if results: + details = results[0].get("results") or [] + if details: + explanation = details[0].get("explanation") or {} + if explanation.get("id"): + return str(explanation["id"]) + return "This deck is not valid for this tournament." + + def progress_dict(entry: dict) -> dict: """AsyncTournamentProgress wire shape from a tournament_data entry dict.""" return { @@ -71,6 +105,11 @@ class TournamentHandler(BaseHandler): return serializable_deck(d["deck_data"], d.get("name") or "Tournament Deck") return None + async def _validate_tournament_deck(self, deck: dict, tournament, legacy: bool = False): + owned = await run_db(get_owned_counts, self._account_id()) + return validate_tournament_deck( + deck, tournament, owned_counts=owned, legacy=legacy) + # ------------------------------------------------------------- listing @handle(InboundMsg.GET_ACTIVE_ASYNC_TOURNAMENTS) @@ -105,6 +144,12 @@ class TournamentHandler(BaseHandler): "reason": {"id": text}, }, request_id) + async def _join_invalid_deck(self, results: list, request_id: int = 0): + await self.send({ + "messageName": OutboundMsg.JOIN_TOURNAMENT_FAILED_INVALID_DECK.value, + "deckValidationResult": results, + }, request_id) + @handle(InboundMsg.GET_ACTIVE_TOURNAMENTS) async def handle_get_active_tournaments_legacy(self, message, request_id, flags): # Zero active entries safely renders the maintenance/"no events" panel. @@ -141,8 +186,14 @@ class TournamentHandler(BaseHandler): return await self._join_failed("You are already in a tournament.", request_id) deck_json = self._resolve_deck_json(message.get("deck")) - if not deck_json or not deck_json.get("piles", {}).get("deck"): - return await self._join_failed("tournament.error.invalid_deck", request_id) + if not deck_json: + return await self._join_failed("The selected deck could not be found.", request_id) + validation, format_error = await self._validate_tournament_deck( + deck_json, tournament, legacy=True) + if format_error: + return await self._join_failed(format_error, request_id) + if not validation or not validation[0]["valid"]: + return await self._join_invalid_deck(validation, request_id) fees = tournament.legacy_entry_fees() error = await run_db(tournament_data.charge_fees, account_id, fees) @@ -251,7 +302,21 @@ class TournamentHandler(BaseHandler): OutboundMsg.JOIN_ASYNC_TOURNAMENT_ERROR.value, "This tournament is not open for entries.", request_id) - deck_json = self._resolve_deck_json(deck_id) or {} + deck_json = self._resolve_deck_json(deck_id) + if not deck_json: + return await self._error( + OutboundMsg.JOIN_ASYNC_TOURNAMENT_ERROR.value, + "The selected deck could not be found.", request_id) + validation, format_error = await self._validate_tournament_deck( + deck_json, tournament) + if format_error: + return await self._error( + OutboundMsg.JOIN_ASYNC_TOURNAMENT_ERROR.value, + format_error, request_id) + if not validation or not validation[0]["valid"]: + return await self._error( + OutboundMsg.JOIN_ASYNC_TOURNAMENT_ERROR.value, + validation_error_text(validation), request_id) entry, error = await run_db( tournament_data.create_entry, self._account_id(), tournament.tournament_id, tournament.definition, currency, deck_json) @@ -298,23 +363,42 @@ class TournamentHandler(BaseHandler): run = tournament.run_config if deck_id and run.get("allowDeckSwitching", True): deck_json = self._resolve_deck_json(deck_id) - if deck_json: - await run_db(tournament_data.update_entry_deck, - entry_id, self._account_id(), deck_json) - entry["deck"] = deck_json - await self.send({ - "messageName": OutboundMsg.ASYNC_TOURNAMENT_DECK_UPDATED.value, - "tournamentID": entry["tournament_id"], - "entryID": entry_id, - "deck": deck_json, - "limitedCollection": [], - }, 0) + if not deck_json: + return await self._error( + OutboundMsg.START_ASYNC_TOURNAMENT_GAME_ERROR.value, + "The selected deck could not be found.", request_id) + validation, format_error = await self._validate_tournament_deck( + deck_json, tournament) + if format_error: + return await self._error( + OutboundMsg.START_ASYNC_TOURNAMENT_GAME_ERROR.value, + format_error, request_id) + if not validation or not validation[0]["valid"]: + return await self._error( + OutboundMsg.START_ASYNC_TOURNAMENT_GAME_ERROR.value, + validation_error_text(validation), request_id) + await run_db(tournament_data.update_entry_deck, + entry_id, self._account_id(), deck_json) + entry["deck"] = deck_json + await self.send({ + "messageName": OutboundMsg.ASYNC_TOURNAMENT_DECK_UPDATED.value, + "tournamentID": entry["tournament_id"], + "entryID": entry_id, + "deck": deck_json, + "limitedCollection": [], + }, 0) deck = entry.get("deck") or {} - if not deck.get("piles", {}).get("deck"): + validation, format_error = await self._validate_tournament_deck( + deck, tournament) + if format_error: return await self._error( OutboundMsg.START_ASYNC_TOURNAMENT_GAME_ERROR.value, - "No valid deck for this tournament run.", request_id) + format_error, request_id) + if not validation or not validation[0]["valid"]: + return await self._error( + OutboundMsg.START_ASYNC_TOURNAMENT_GAME_ERROR.value, + validation_error_text(validation), request_id) # Complete the RPC, then ride the normal matchmaking pipeline # (MatchQueueEntered -> ConfirmReadyForMatch -> MatchFound). diff --git a/tests/test_tournament_deck_validation.py b/tests/test_tournament_deck_validation.py new file mode 100644 index 0000000..f5d1fcc --- /dev/null +++ b/tests/test_tournament_deck_validation.py @@ -0,0 +1,77 @@ +import unittest +import uuid + +from spirit.game.attributes import AttrID, CardType, DeckFormat, PokemonStage +from spirit.game import rules +from spirit.game.scripts.cards import loader as card_loader +from spirit.game.tournament_manager import TournamentDef +from spirit.packets.handlers.tournaments import validate_tournament_deck + + +class TournamentDeckValidationTests(unittest.TestCase): + @classmethod + def setUpClass(cls): + card_loader.load_all() + cls.swsh_basic = next( + card for card in card_loader.cards + if card.key == "SWSH8" + and card.get_attribute_value(AttrID.CARD_TYPE) == CardType.POKEMON.value + and card.get_attribute_value(AttrID.STAGE, 0) == PokemonStage.BASIC.value + ) + cls.bw_card = next(card for card in card_loader.cards if card.key == "BW1") + cls.water = next( + card for card in card_loader.cards + if card.key == "Free_Energy" + and rules.card_display_name(card) == "Water Energy" + ) + + def make_deck(self, pile_name="deck", include_bw=False): + cards = [self.swsh_basic.guid] * 4 + if include_bw: + cards.append(self.bw_card.guid) + cards.extend([self.water.guid] * (60 - len(cards))) + return { + "deckID": str(uuid.uuid4()), + "deckName": "Tournament Test", + "piles": {pile_name: cards}, + } + + def tournament(self, definition): + return TournamentDef(str(uuid.uuid4()), definition, True) + + def test_legacy_join_accepts_server_and_client_pile_names(self): + tournament = self.tournament({"format": "Modified"}) + for pile_name in ("deck", "CakePile"): + results, error = validate_tournament_deck( + self.make_deck(pile_name), tournament, legacy=True) + self.assertIsNone(error) + self.assertTrue(results[0]["valid"], pile_name) + + def test_legacy_join_enforces_tournament_format(self): + tournament = self.tournament({"format": "Modified"}) + results, error = validate_tournament_deck( + self.make_deck("CakePile", include_bw=True), tournament, legacy=True) + self.assertIsNone(error) + self.assertFalse(results[0]["valid"]) + self.assertIn( + "DeckContainsBannedCards", + {detail["failureType"] for detail in results[0]["results"]}, + ) + + def test_async_game_format_accepts_a_guid(self): + tournament = self.tournament({ + "format": "Unlimited", + "game": {"format": DeckFormat.STANDARD.value}, + }) + deck = self.make_deck("CakePile", include_bw=True) + async_results, async_error = validate_tournament_deck(deck, tournament) + legacy_results, legacy_error = validate_tournament_deck( + deck, tournament, legacy=True) + self.assertIsNone(async_error) + self.assertIsNone(legacy_error) + self.assertFalse(async_results[0]["valid"]) + self.assertTrue(legacy_results[0]["valid"]) + + +if __name__ == "__main__": + unittest.main()