"""The oracles are the training signal, so they get tested harder than the model.""" from __future__ import annotations import os import random import sys from pathlib import Path ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) import pytest import envs.maze as mz import envs.snake as sk @pytest.mark.parametrize("topology", mz.TOPOLOGIES) @pytest.mark.parametrize("size", [7, 11, 21]) def test_maze_is_solvable_and_oracle_is_optimal(size, topology): state = mz.generate_maze(size, topology, seed=size * 13 + len(topology)) distances = mz.bfs_distances(size, state.walls, state.goal) start = distances[state.position] assert start < mz.UNREACHABLE, "generator must not emit an unsolvable maze" # Following the oracle must reach the goal in exactly the shortest distance. walker = state.clone() for _ in range(start): probs = mz.optimal_action_distribution(walker) action = max(probs, key=probs.get) assert walker.step(action) assert walker.solved() assert walker.steps == start def test_maze_oracle_ties_are_all_shortest(): state = mz.generate_maze(15, "loops", 4) distances = mz.bfs_distances(15, state.walls, state.goal) probs = mz.optimal_action_distribution(state) best = min(distances[(state.position[0] + mz.DIRECTIONS[a][0], state.position[1] + mz.DIRECTIONS[a][1])] for a in probs if probs[a] > 0) for action, mass in probs.items(): if mass > 0: cell = (state.position[0] + mz.DIRECTIONS[action][0], state.position[1] + mz.DIRECTIONS[action][1]) assert distances[cell] == best def test_maze_questions_agree_with_the_board(): state = mz.generate_maze(11, "tree", 99) profile = mz.question_profile(state) for action in mz.ACTIONS: dr, dc = mz.DIRECTIONS[action] cell = (state.position[0] + dr, state.position[1] + dc) expected = float(state.free(cell)) assert profile["targets"][f"clear_{action}"]["probs"]["true"] == expected def test_snake_oracle_survives_far_longer_than_greedy(): """The whole point of rewriting the Snake target: greedy food-chasing traps.""" def run(policy, seed): state = sk.new_game(12, seed) while state.alive and state.steps < 1500: actions = state.candidate_actions() if not actions: break state.step(policy(state, actions)) return state.eaten def greedy(state, actions): best, pick = None, actions[0] for action in actions: info = sk.evaluate_action(state, action) if not info["safe"]: continue if best is None or info["food_distance"] < best: best, pick = info["food_distance"], action return pick def oracle(state, actions): probs, _ = sk.oracle_action_distribution(state) return max(probs, key=probs.get) if probs else actions[0] seeds = range(6) g = sum(run(greedy, s) for s in seeds) / 6 o = sum(run(oracle, s) for s in seeds) / 6 assert o > g, f"survival oracle {o} should beat greedy {g}" def test_snake_oracle_never_picks_an_immediately_fatal_move(): rng = random.Random(0) for seed in range(40): state = sk.random_state(12, rng.randrange(6, 60), rng) if state is None: continue probs, _ = sk.oracle_action_distribution(state) safe = [a for a in state.candidate_actions() if sk.evaluate_action(state, a)["safe"]] if not safe: continue # nothing survives; any answer is equally wrong for action, mass in probs.items(): if mass > 0: assert action in safe def test_random_state_is_a_valid_snake(): rng = random.Random(7) for _ in range(30): state = sk.random_state(10, rng.randrange(4, 40), rng) if state is None: continue assert len(set(state.body)) == len(state.body), "body must not self-overlap" for a, b in zip(state.body, state.body[1:]): assert abs(a[0] - b[0]) + abs(a[1] - b[1]) == 1, "body must be connected" assert state.food not in set(state.body) assert all(0 <= c < 10 for cell in state.body for c in cell) def test_room_level_is_relative_to_length(): # The bug this replaced: an absolute board fraction made every sampled # state the top level, so the label carried no information at all. assert sk.room_level(space=5, length=20) == 0 assert sk.room_level(space=200, length=20) == 4 assert sk.room_level(space=200, length=200) == 0 def test_mazes_are_identical_across_processes(): """Board generation must not depend on the interpreter's hash salt. ``str.__hash__`` is salted per process, so seeding the generator with ``hash(topology)`` handed every run a different maze for the same ``(size, topology, seed)``. Training only resampled, but probe.py and play.py regenerate their boards, so their numbers silently stopped being comparable between runs -- and nobody cloning the repository could reproduce a single published figure. Two explicit, different hash seeds are enough to catch a regression; the default salt is already random. """ import subprocess import sys code = ("import sys; sys.path.insert(0, %r); from envs import maze as mz; " "print([(mz.generate_maze(s, t, 900_101).goal, " "len(mz.generate_maze(s, t, 900_101).walls)) " "for s in (9, 11) for t in mz.TOPOLOGIES])" % str(ROOT)) outs = set() for salt in ("0", "1", "12345"): env = {**os.environ, "PYTHONHASHSEED": salt} got = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, env=env, cwd=str(ROOT)) assert got.returncode == 0, got.stderr outs.add(got.stdout.strip()) assert len(outs) == 1, f"maze depends on the hash salt: {outs}" @pytest.mark.parametrize("seed", [1439, 2428, 3869]) def test_known_degenerate_seeds_now_generate(seed): """These three used to raise, and nothing up the stack caught it. random_obstacle picks its density before it knows whether the result percolates, so a small fraction of seeds leave a largest component too small to use. generate_maze raised RuntimeError; the sampler did not handle it; so a single unlucky draw killed a training run after an hour of compute. Generation retries with a thinner density instead. """ state = mz.generate_maze(7, "random_obstacle", seed) free = 7 * 7 - len(state.walls) assert free >= 4 assert state.position != state.goal reachable = mz.bfs_distances(7, state.walls, state.position) assert reachable[state.goal] < mz.UNREACHABLE def test_generation_is_total_over_the_sampled_seed_space(): """No seed in the range training draws from may fail to produce a board.""" for size in (7, 9, 11): for topology in mz.TOPOLOGIES: for seed in range(400): state = mz.generate_maze(size, topology, seed) assert size * size - len(state.walls) >= 4 def test_frozen_splits_are_byte_reproducible(tmp_path): """The frozen splits exist so every run is scored on identical states. That guarantee was vacuous while generate_maze depended on the hash salt: the committed files were still fixed, but nobody -- including a later run of this repository -- could rebuild them. Cheap to check; 580 states take about seven seconds. """ import hashlib import subprocess import sys committed = ROOT / "data" / "frozen" if not (committed / "val.jsonl").exists(): pytest.skip("frozen splits not built") got = subprocess.run( [sys.executable, str(ROOT / "scripts" / "build_dataset.py"), "--out-dir", str(tmp_path)], capture_output=True, text=True, cwd=str(ROOT)) assert got.returncode == 0, got.stderr def digest(path): return hashlib.sha256(path.read_bytes()).hexdigest() for name in ("val", "test", "ood"): assert digest(tmp_path / f"{name}.jsonl") == digest(committed / f"{name}.jsonl"), ( f"{name}.jsonl is not reproducible from build_dataset.py")