matiqorb/spades-agent-api
0
1from __future__ import annotations2 3from dataclasses import dataclass4import random5from typing import Dict, List, Optional, Tuple6 7SUITS = ("C", "D", "H", "S")8RANKS = tuple(range(2, 15))9 10 11@dataclass(frozen=True, order=True)12class Card:13 suit: str14 rank: int15 16 def to_str(self) -> str:17 rank_map = {11: "J", 12: "Q", 13: "K", 14: "A"}18 rank_str = rank_map.get(self.rank, str(self.rank))19 return f"{rank_str}{self.suit}"20 21 @staticmethod22 def from_str(value: str) -> "Card":23 suit = value[-1].upper()24 rank_token = value[:-1].upper()25 if rank_token == "A":26 rank = 1427 elif rank_token == "K":28 rank = 1329 elif rank_token == "Q":30 rank = 1231 elif rank_token == "J":32 rank = 1133 else:34 rank = int(rank_token)35 return Card(suit=suit, rank=rank)36 37 38def full_deck() -> List[Card]:39 return [Card(suit=suit, rank=rank) for suit in SUITS for rank in RANKS]40 41 42def card_to_index(card: Card) -> int:43 suit_idx = SUITS.index(card.suit)44 rank_idx = card.rank - 245 return suit_idx * 13 + rank_idx46 47 48def index_to_card(index: int) -> Card:49 suit_idx = index // 1350 rank_idx = index % 1351 return Card(suit=SUITS[suit_idx], rank=rank_idx + 2)52 53 54def team_for_player(player: int) -> int:55 return player % 256 57 58class SpadesGame:59 def __init__(self, seed: int = 0) -> None:60 self.seed = seed61 self.rng = random.Random(seed)62 self.scores = [0, 0]63 self.bags = [0, 0]64 self.dealer = 365 self.hand_id = 066 self.phase = "init"67 self.current_player = 068 self.hands: Dict[int, List[Card]] = {}69 self.bid_history: List[Optional[int]] = [None, None, None, None]70 self.trick: List[Tuple[int, Card]] = []71 self.trick_leader = 072 self.spades_broken = False73 self.tricks_won = [0, 0]74 self.history: List[Tuple[int, Card]] = []75 76 def reset_hand(self, seed: Optional[int] = None) -> None:77 if seed is not None:78 self.rng = random.Random(seed)79 self.hand_id += 180 self.phase = "bidding"81 self.bid_history = [None, None, None, None]82 self.trick = []83 self.tricks_won = [0, 0]84 self.spades_broken = False85 self.history = []86 self.dealer = (self.dealer + 1) % 487 self.current_player = (self.dealer + 1) % 488 deck = full_deck()89 self.rng.shuffle(deck)90 self.hands = {i: sorted(deck[i * 13 : (i + 1) * 13]) for i in range(4)}91 self.trick_leader = self.current_player92 93 def legal_bids(self, player: int) -> List[int]:94 if self.phase != "bidding" or player != self.current_player:95 return []96 return list(range(0, 14))97 98 def place_bid(self, player: int, bid: int) -> None:99 if bid not in self.legal_bids(player):100 raise ValueError("Illegal bid")101 self.bid_history[player] = bid102 if all(v is not None for v in self.bid_history):103 self.phase = "playing"104 self.current_player = (self.dealer + 1) % 4105 self.trick_leader = self.current_player106 return107 self.current_player = (self.current_player + 1) % 4108 109 def _lead_suit(self) -> Optional[str]:110 if not self.trick:111 return None112 return self.trick[0][1].suit113 114 def legal_cards(self, player: int) -> List[Card]:115 if self.phase != "playing" or player != self.current_player:116 return []117 hand = self.hands[player]118 lead_suit = self._lead_suit()119 if lead_suit is None:120 non_spades = [c for c in hand if c.suit != "S"]121 if self.spades_broken or not non_spades:122 return list(hand)123 return non_spades124 follow = [c for c in hand if c.suit == lead_suit]125 if follow:126 return follow127 return list(hand)128 129 def legal_card_mask(self, player: int) -> List[int]:130 mask = [0] * 52131 for card in self.legal_cards(player):132 mask[card_to_index(card)] = 1133 return mask134 135 def _resolve_trick_winner(self) -> int:136 lead_suit = self.trick[0][1].suit137 best_player, best_card = self.trick[0]138 for player, card in self.trick[1:]:139 if card.suit == "S" and best_card.suit != "S":140 best_player, best_card = player, card141 elif card.suit == best_card.suit and card.rank > best_card.rank:142 best_player, best_card = player, card143 elif best_card.suit != "S" and card.suit == lead_suit and card.rank > best_card.rank:144 best_player, best_card = player, card145 return best_player146 147 def play_card(self, player: int, card: Card) -> None:148 legal = self.legal_cards(player)149 if card not in legal:150 raise ValueError("Illegal card")151 self.hands[player].remove(card)152 self.trick.append((player, card))153 self.history.append((player, card))154 if card.suit == "S":155 self.spades_broken = True156 if len(self.trick) < 4:157 self.current_player = (self.current_player + 1) % 4158 return159 winner = self._resolve_trick_winner()160 self.tricks_won[team_for_player(winner)] += 1161 self.trick = []162 self.current_player = winner163 self.trick_leader = winner164 if sum(self.tricks_won) == 13:165 self._finish_hand()166 167 def _finish_hand(self) -> None:168 targets = [169 (self.bid_history[0] or 0) + (self.bid_history[2] or 0),170 (self.bid_history[1] or 0) + (self.bid_history[3] or 0),171 ]172 for team in (0, 1):173 tricks = self.tricks_won[team]174 target = targets[team]175 if tricks >= target:176 bags = tricks - target177 self.bags[team] += bags178 self.scores[team] += target * 10 + bags179 if self.bags[team] >= 10:180 self.scores[team] -= 100181 self.bags[team] -= 10182 else:183 self.scores[team] -= target * 10184 self.phase = "finished"185 186 def game_winner(self, target_score: int = 500) -> Optional[int]:187 if self.scores[0] >= target_score or self.scores[1] >= target_score:188 return 0 if self.scores[0] > self.scores[1] else 1189 return None190 191 def observation(self, player: int) -> Dict:192 return {193 "phase": self.phase,194 "player": player,195 "current_player": self.current_player,196 "hand": [c.to_str() for c in sorted(self.hands.get(player, []))],197 "bid_history": self.bid_history[:],198 "trick": [{"player": p, "card": c.to_str()} for p, c in self.trick],199 "spades_broken": self.spades_broken,200 "tricks_won": self.tricks_won[:],201 "scores": self.scores[:],202 "bags": self.bags[:],203 "legal_card_mask": self.legal_card_mask(player),204 "legal_cards": [c.to_str() for c in sorted(self.legal_cards(player))],205 "legal_bids": self.legal_bids(player),206 "dealer": self.dealer,207 "hand_id": self.hand_id,208 }209 