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()