#!/usr/bin/env python3 """Compare snake decisions and latency on deterministic synthetic positions. For outcome/win-rate tournaments use the local Battlesnake CLI. This harness is fast enough for CI and detects move disagreements, crashes, and latency changes. """ from __future__ import annotations import argparse import json import sys from pathlib import Path from statistics import mean, median from time import perf_counter sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from scripts.benchmark_snakes_from_db import load_states, percentile from scripts.snake_arena_scenarios import SCENARIOS, synthetic_states from server.GameBoard import GameBoard from snakes import SnakeBuilder def evaluate(name: str, states: list[tuple[dict, dict]]) -> tuple[list[str], dict]: moves: list[str] = [] durations: list[float] = [] duel_depths: list[int] = [] rollout_depths: list[int] = [] duel_nodes = 0 rollout_nodes = 0 cache_hits = 0 deadline_exits = 0 scenario_durations: dict[str, list[float]] = {} for index, (board_data, metadata) in enumerate(states): snake = SnakeBuilder.build(name) game_id = f"arena-{name}-{index}-{metadata['game_id']}" board = GameBoard( game_id=game_id, width=board_data["width"], height=board_data["height"], ruleset=metadata["ruleset"], source=metadata["source"], map=metadata["map"], snake_class=snake, ) board.read_game_data({ "game": { "id": game_id, "ruleset": metadata["ruleset"], "source": metadata["source"], "map": metadata["map"], "timeout": 500, }, "turn": metadata["turn"], "board": board_data, "you": metadata["you"], }) started = perf_counter() moves.append(snake.choose_move(board)) duration = (perf_counter() - started) * 1000.0 durations.append(duration) scenario = metadata.get("scenario", "recorded") scenario_durations.setdefault(scenario, []).append(duration) history = snake.get_history() if hasattr(snake, "get_history") else [] if history: thinking = history[-1] duel_depths.append(int(thinking.get("prism_duel_depth", thinking.get("minimax_depth_reached", 0)))) rollout_depths.append(int(thinking.get("prism_rollout_depth", 0))) duel_nodes += int(thinking.get("prism_duel_nodes", 0)) rollout_nodes += int(thinking.get("prism_rollout_nodes", 0)) cache_hits += int(thinking.get("prism_duel_cache_hits", 0)) cache_hits += int(thinking.get("prism_rollout_cache_hits", 0)) deadline_exits += int(thinking.get("prism_duel_deadline_exits", 0)) deadline_exits += int(thinking.get("prism_rollout_deadline_exits", 0)) return moves, { "snake": name, "positions": len(states), "mean_ms": mean(durations), "median_ms": median(durations), "p95_ms": percentile(durations, 0.95), "max_ms": max(durations), "mean_duel_depth": mean(duel_depths) if duel_depths else 0.0, "mean_rollout_depth": mean(rollout_depths) if rollout_depths else 0.0, "duel_nodes": duel_nodes, "rollout_nodes": rollout_nodes, "cache_hits": cache_hits, "deadline_exits": deadline_exits, "scenario_mean_ms": { scenario: mean(values) for scenario, values in scenario_durations.items() }, } def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--snake", action="append", default=[]) parser.add_argument("--database") parser.add_argument("--positions", type=int, default=100) parser.add_argument("--stride", type=int, default=997) parser.add_argument("--scenario", action="append", choices=sorted(SCENARIOS)) parser.add_argument("--json-output") args = parser.parse_args() states = ( load_states(args.database, max(1, args.positions), max(1, args.stride)) if args.database else synthetic_states(max(1, args.positions), args.scenario) ) if not states: raise SystemExit("No benchmark positions found") names = args.snake or ["ApexBattleSnake", "PrismBattleSnake_GPT_5_6_Sol"] move_sets: dict[str, list[str]] = {} reports: list[dict] = [] for name in names: moves, report = evaluate(name, states) move_sets[name] = moves reports.append(report) print( f"{name}: mean={report['mean_ms']:.3f} ms, " f"p95={report['p95_ms']:.3f} ms, max={report['max_ms']:.3f} ms, " f"duel-depth={report['mean_duel_depth']:.2f}, " f"rollout-depth={report['mean_rollout_depth']:.2f}, " f"nodes={report['duel_nodes'] + report['rollout_nodes']}, " f"cache-hits={report['cache_hits']}, deadline-exits={report['deadline_exits']}" ) baseline = names[0] disagreements = { name: sum(a != b for a, b in zip(move_sets[baseline], move_sets[name])) for name in names[1:] } if disagreements: print(f"Move disagreements versus {baseline}: {disagreements}") scenario_disagreements = {} for name in names[1:]: counts: dict[str, int] = {} for index, (baseline_move, candidate_move) in enumerate( zip(move_sets[baseline], move_sets[name]) ): if baseline_move != candidate_move: scenario = states[index][1].get("scenario", "recorded") counts[scenario] = counts.get(scenario, 0) + 1 scenario_disagreements[name] = counts if any(scenario_disagreements.values()): print(f"Disagreements by scenario: {scenario_disagreements}") payload = { "reports": reports, "baseline": baseline, "disagreements": disagreements, "scenario_disagreements": scenario_disagreements, } if args.json_output: Path(args.json_output).write_text(json.dumps(payload, indent=2) + "\n") if __name__ == "__main__": main()