caro5 / tests /test_tactical_bot.py
Pedro de Carvalho
Update app
f35ee73
Raw
History Blame Contribute Delete
8.74 kB
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()