"""Load sampled gameplay positions from SQLite or PostgreSQL.""" from __future__ import annotations import asyncio import json from urllib.parse import urlparse def is_postgresql_source(source: str) -> bool: return urlparse(source).scheme.lower() in {"postgres", "postgresql"} def _decode_json(value, default): if value is None: return default if isinstance(value, str): return json.loads(value) return value def _build_state(row, snake_rows) -> tuple[dict, dict] | None: board = _decode_json(row[1], {}) you = _decode_json(row[2], {}) if not board.get("snakes"): snakes = [] for snake_row in snake_rows: snake_id = snake_row[0] snake_name = snake_row[1] or (row[6] if snake_id == row[5] else snake_id) snakes.append({ "id": snake_id, "name": snake_name, "health": snake_row[2], "length": snake_row[3], "head": {"x": snake_row[4], "y": snake_row[5]}, "body": _decode_json(snake_row[6], []), "customizations": _decode_json(snake_row[7], {}), }) board = { "width": row[7], "height": row[8], "food": _decode_json(row[3], []), "hazards": _decode_json(row[4], []), "snakes": snakes, } if not you: you = next( (snake for snake in board.get("snakes", []) if snake.get("id") == row[5]), {}, ) if not you or not board.get("snakes"): return None return board, { "you": you, "game_id": row[9], "source": row[10] or "custom", "map": row[11] or "standard", "ruleset": { "name": row[12] or "standard", "version": row[13] or "v1.0.0", "settings": {}, }, "turn": int(row[14]), } async def _load_postgresql_states(dsn: str, samples: int, stride: int) -> list[tuple[dict, dict]]: try: import asyncpg except ImportError as exc: raise RuntimeError("asyncpg is required for PostgreSQL benchmark sources") from exc connection = await asyncpg.connect(dsn=dsn) try: max_id = int(await connection.fetchval("SELECT max(id) FROM turns") or 0) if max_id == 0: return [] query = """ SELECT t.id, t.board_state, t.you, t.food, t.hazards, g.your_snake_id, g.your_snake_name, g.width, g.height, g.game_id, g.source, g.map_name, g.ruleset_name, g.ruleset_version, t.turn FROM turns AS t JOIN games AS g ON g.game_id = t.game_id WHERE t.id >= $1 ORDER BY t.id LIMIT 1 """ snake_query = """ SELECT st.snake_id, COALESCE(gs.snake_name, st.snake_name), st.health, st.length, st.head_x, st.head_y, st.body, COALESCE(gs.customizations, '{}'::jsonb) FROM snake_turns AS st LEFT JOIN game_snakes AS gs ON gs.game_id = st.game_id AND gs.snake_id = st.snake_id WHERE st.game_id = $1 AND st.turn = $2 ORDER BY st.id """ states = [] next_id = max(1, max_id - (samples - 1) * stride) while len(states) < samples and next_id <= max_id: row = await connection.fetchrow(query, next_id) if row is None: break snake_rows = await connection.fetch(snake_query, row[9], row[14]) state = _build_state(row, snake_rows) if state is not None: states.append(state) next_id = int(row[0]) + stride return states finally: await connection.close() def load_postgresql_states(dsn: str, samples: int, stride: int) -> list[tuple[dict, dict]]: return asyncio.run(_load_postgresql_states(dsn, samples, stride))