Spaces:
Runtime error
Runtime error
| from __future__ import annotations | |
| import unittest | |
| from unittest.mock import patch | |
| from fastapi.testclient import TestClient | |
| from app import app | |
| from caro5_bot import BotRequest, select_move | |
| from caro5_bot.features import ( | |
| CH_CURRENT_PLAYER, | |
| CH_CURRENT_STONES, | |
| CH_NORMAL, | |
| CH_OPPONENT_STONES, | |
| FEATURE_CHANNELS, | |
| PLAY_ONLY_FEATURE_CHANNELS, | |
| encode_feature_planes, | |
| encode_play_only_feature_planes, | |
| ) | |
| from caro5_bot.model import ModelMetadata, OnnxMctsEvaluator, inspect_model | |
| from caro5_bot.mcts import MctsConfig, run_mcts | |
| from caro5_bot.tactical import Board, BotMove, check_win | |
| def cell(player: int) -> dict[str, object]: | |
| return {"playerId": str(player), "symbol": "X" if player == 0 else "O", "timestamp": 0} | |
| class TacticalBotTest(unittest.TestCase): | |
| def test_check_win_respects_no_overlines(self) -> None: | |
| board = Board({(x, 7): 0 for x in range(2, 8)}, ("X", "O"), 15) | |
| self.assertTrue(check_win(board, BotMove(5, 7), 0, no_overlines=False)) | |
| self.assertFalse(check_win(board, BotMove(5, 7), 0, no_overlines=True)) | |
| def test_selects_immediate_winning_move(self) -> None: | |
| board = {f"{x}:7": cell(0) for x in range(3, 7)} | |
| response = select_move(BotRequest(board=board, current_player=0)) | |
| self.assertEqual(response.source, "alpha_beta") | |
| self.assertIn((response.x, response.y), {(2, 7), (7, 7)}) | |
| def test_blocks_opponent_immediate_win(self) -> None: | |
| board = {f"{x}:7": cell(1) for x in range(3, 7)} | |
| response = select_move(BotRequest(board=board, current_player=0)) | |
| self.assertEqual(response.source, "alpha_beta") | |
| self.assertIn((response.x, response.y), {(2, 7), (7, 7)}) | |
| def test_rejects_when_no_legal_moves_remain(self) -> None: | |
| board = {f"{x}:{y}": cell((x + y) % 2) for y in range(5) for x in range(5)} | |
| with self.assertRaisesRegex(ValueError, "No legal moves"): | |
| select_move(BotRequest(board=board, current_player=0, rules={"boardSize": 5})) | |
| def test_expert_profile_uses_hybrid_search_for_strategy(self) -> None: | |
| response = select_move(BotRequest(board={}, current_player=0, profile="expert", seed=7)) | |
| self.assertEqual(response.source, "hybrid") | |
| self.assertEqual(response.profile, "expert") | |
| self.assertEqual((response.x, response.y), (7, 7)) | |
| def test_legacy_balanced_profile_normalizes_to_normal(self) -> None: | |
| response = select_move(BotRequest(board={}, current_player=0, profile="balanced", seed=7)) | |
| self.assertEqual(response.profile, "normal") | |
| self.assertEqual(response.source, "hybrid") | |
| def test_ai_profile_falls_back_without_model(self) -> None: | |
| with patch.dict("os.environ", {}, clear=True): | |
| response = select_move(BotRequest(board={}, current_player=0, profile="ai", seed=7)) | |
| self.assertEqual(response.profile, "ai") | |
| self.assertEqual(response.source, "hybrid") | |
| self.assertFalse(response.model_available) | |
| class MctsTest(unittest.TestCase): | |
| def test_mcts_uses_evaluator_priors_for_root_choice(self) -> None: | |
| board = Board({(6, 7): 0, (8, 7): 1}, ("X", "O"), 15) | |
| class CenterEvaluator: | |
| def evaluate(self, _board: Board, _root_player: int, _turn: int, legal_moves: list[BotMove]) -> tuple[float, list[float]]: | |
| priors = [10.0 if (move.x, move.y) == (7, 7) else 0.01 for move in legal_moves] | |
| return 0.25, priors | |
| result = run_mcts( | |
| board=board, | |
| root_player=0, | |
| no_overlines=False, | |
| config=MctsConfig(simulations=0, candidate_radius=1, max_candidates=18), | |
| seed=11, | |
| evaluator=CenterEvaluator(), | |
| ) | |
| self.assertIsNotNone(result) | |
| self.assertEqual((result.move.x, result.move.y), (7, 7)) | |
| class OnnxEvaluatorTest(unittest.TestCase): | |
| def test_onnx_evaluator_masks_policy_to_legal_moves_and_flips_value(self) -> None: | |
| board = Board({(6, 7): 0, (8, 7): 1}, ("X", "O"), 15) | |
| legal_moves = [BotMove(7, 6), BotMove(7, 7), BotMove(7, 8)] | |
| test_case = self | |
| class FakeModel: | |
| metadata = ModelMetadata(board_size=15, feature_channels=7, policy_size=225) | |
| def evaluate(self, features: list[float]) -> tuple[list[float], float]: | |
| self_feature_count = 7 * 15 * 15 | |
| test_case.assertEqual(len(features), self_feature_count) | |
| logits = [-10.0] * 225 | |
| logits[7 * 15 + 7] = 5.0 | |
| return logits, -0.4 | |
| evaluator = OnnxMctsEvaluator(FakeModel(), no_overlines=True) # type: ignore[arg-type] | |
| value, priors = evaluator.evaluate(board, root_player=0, turn=1, legal_moves=legal_moves) | |
| self.assertAlmostEqual(value, 0.4) | |
| self.assertEqual(len(priors), len(legal_moves)) | |
| self.assertEqual(max(range(len(priors)), key=lambda index: priors[index]), 1) | |
| class FeatureEncoderTest(unittest.TestCase): | |
| def test_encodes_full_feature_planes(self) -> None: | |
| board = Board({(7, 7): 0, (8, 7): 1}, ("X", "O"), 15) | |
| planes = encode_feature_planes(board, current_player=0, rules={"noOverlines": True}, phase="play", last_move=BotMove(8, 7)) | |
| area = 15 * 15 | |
| self.assertEqual(len(planes), FEATURE_CHANNELS * area) | |
| self.assertEqual(planes[CH_CURRENT_STONES * area + 7 * 15 + 7], 1) | |
| self.assertEqual(planes[CH_OPPONENT_STONES * area + 7 * 15 + 8], 1) | |
| self.assertEqual(planes[CH_NORMAL * area], 1) | |
| self.assertTrue(all(value == 0 for value in planes[CH_CURRENT_PLAYER * area: (CH_CURRENT_PLAYER + 1) * area])) | |
| def test_encodes_play_only_feature_planes(self) -> None: | |
| board = Board({(7, 7): 1}, ("X", "O"), 15) | |
| planes = encode_play_only_feature_planes(board, current_player=0) | |
| self.assertEqual(len(planes), PLAY_ONLY_FEATURE_CHANNELS * 15 * 15) | |
| class ModelLoadTest(unittest.TestCase): | |
| def test_model_is_unavailable_without_env_paths(self) -> None: | |
| with patch.dict("os.environ", {}, clear=True): | |
| state = inspect_model(None, None) | |
| self.assertFalse(state.available) | |
| self.assertIn("CARO5_MODEL_ONNX", state.error or "") | |
| def test_model_is_unavailable_when_files_are_missing(self) -> None: | |
| state = inspect_model("/tmp/no-such-model.onnx", "/tmp/no-such-metadata.json") | |
| self.assertFalse(state.available) | |
| self.assertIn("model file not found", state.error or "") | |
| class BotApiTest(unittest.TestCase): | |
| def setUp(self) -> None: | |
| self.client = TestClient(app) | |
| def test_status_endpoint_reports_profiles(self) -> None: | |
| with patch.dict("os.environ", {}, clear=True): | |
| response = self.client.get("/api/bot/status") | |
| self.assertEqual(response.status_code, 200) | |
| payload = response.json() | |
| self.assertTrue(payload["available"]) | |
| self.assertFalse(payload["modelAvailable"]) | |
| self.assertIn("CARO5_MODEL_ONNX", payload["modelError"]) | |
| self.assertEqual(payload["defaultProfile"], "normal") | |
| self.assertIn("normal", payload["profiles"]) | |
| self.assertIn("expert", payload["profiles"]) | |
| self.assertIn("ai", payload["profiles"]) | |
| def test_move_endpoint_returns_legal_move(self) -> None: | |
| response = self.client.post( | |
| "/api/bot/move", | |
| json={ | |
| "board": {f"{x}:7": cell(1) for x in range(3, 7)}, | |
| "currentPlayer": 0, | |
| "playerSymbols": ["X", "O"], | |
| "rules": {"boardSize": 15, "noOverlines": False}, | |
| "profile": "normal", | |
| "seed": 123, | |
| }, | |
| ) | |
| self.assertEqual(response.status_code, 200) | |
| payload = response.json() | |
| self.assertTrue(payload["legal"]) | |
| self.assertEqual(payload["profile"], "normal") | |
| self.assertEqual(payload["seed"], 123) | |
| self.assertGreaterEqual(payload["elapsedMs"], 0) | |
| self.assertFalse(payload["modelAvailable"]) | |
| self.assertIn((payload["x"], payload["y"]), {(2, 7), (7, 7)}) | |
| self.assertEqual(payload["move"], {"x": payload["x"], "y": payload["y"]}) | |
| def test_move_endpoint_accepts_legacy_balanced_profile(self) -> None: | |
| response = self.client.post( | |
| "/api/bot/move", | |
| json={ | |
| "board": {}, | |
| "currentPlayer": 0, | |
| "rules": {"boardSize": 15}, | |
| "profile": "balanced", | |
| "seed": 123, | |
| }, | |
| ) | |
| self.assertEqual(response.status_code, 200) | |
| self.assertEqual(response.json()["profile"], "normal") | |
| if __name__ == "__main__": | |
| unittest.main() | |