File size: 5,051 Bytes
4be6a52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
import copy
import hashlib
import json

import pytest

from stackcraft.data import (
    DatasetConfig,
    audit_dataset,
    canonical_json,
    generate_dataset,
    observation_hash,
    write_dataset,
)

COMMIT = "412be98cf90bc7e3e60cd39d8bf9d479bcf65d27"


@pytest.fixture
def bundle():
    return generate_dataset(
        DatasetConfig(
            train_seeds=(10000, 10001, 10002), validation_seeds=(20000, 20001, 20002), max_pieces=2
        ),
        COMMIT,
    )


def refresh_hash(records, manifest, split):
    data = "".join(canonical_json(row) + "\n" for row in records[split])
    manifest["splits"][split]["sha256"] = hashlib.sha256(data.encode()).hexdigest()
    manifest["splits"][split]["records"] = len(records[split])


def test_dataset_is_reproducible_and_never_collects_test_episodes(bundle) -> None:
    assert bundle == generate_dataset(DatasetConfig(**bundle.manifest["config"]), COMMIT)
    assert audit_dataset(bundle.records, bundle.manifest) == {
        split: len(rows) for split, rows in bundle.records.items()
    }
    assert bundle.manifest["test_trajectories_generated"] is False
    assert bundle.manifest["seed_pools"]["reserved_test"] == list(range(30000, 30200))
    assert all(
        row["seed"] not in range(30000, 30200) for rows in bundle.records.values() for row in rows
    )
    assert set(bundle.manifest["source_hashes"]) == {
        "engine.py",
        "pieces.py",
        "schema.py",
        "players/__init__.py",
        "expert.py",
        "data.py",
    }


def test_observation_hash_ignores_colors_not_geometry(bundle) -> None:
    observation = copy.deepcopy(bundle.records["train"][1]["observation"])
    colored = copy.deepcopy(observation)
    colored["board"] = [[7 if cell else 0 for cell in row] for row in colored["board"]]
    assert observation_hash(observation) == observation_hash(colored)
    colored["board"][0][0] = 1
    assert observation_hash(observation) != observation_hash(colored)


def test_collector_deduplicates_identical_starting_observations() -> None:
    # At least two of 60 starts share one of the 49 possible current/preview pairs.
    result = generate_dataset(
        DatasetConfig(
            train_seeds=tuple(range(10000, 10060)), validation_seeds=(20000,), max_pieces=1
        ),
        COMMIT,
    )
    assert result.manifest["deduplication"]["exclusions"]["train"]["within_split"] > 0
    hashes = [row["observation_hash"] for rows in result.records.values() for row in rows]
    assert len(hashes) == len(set(hashes))


@pytest.mark.parametrize(
    "field,value", [("action_id", "not-legal"), ("seed", 30000), ("source_commit", "bad")]
)
def test_audit_rejects_invalid_provenance_or_labels_even_after_rehash(bundle, field, value) -> None:
    rows, manifest = copy.deepcopy(bundle.records), copy.deepcopy(bundle.manifest)
    rows["train"][0][field] = value
    refresh_hash(rows, manifest, "train")
    with pytest.raises(ValueError):
        audit_dataset(rows, manifest)


def test_audit_rejects_hidden_fields(bundle) -> None:
    rows, manifest = copy.deepcopy(bundle.records), copy.deepcopy(bundle.manifest)
    rows["train"][0]["observation"]["seed"] = 10000
    refresh_hash(rows, manifest, "train")
    with pytest.raises(ValueError, match="hidden fields"):
        audit_dataset(rows, manifest)


def test_audit_detects_cross_split_identical_observation(bundle) -> None:
    rows, manifest = copy.deepcopy(bundle.records), copy.deepcopy(bundle.manifest)
    source = copy.deepcopy(rows["train"][0])
    target = rows["validation"][0]
    for field in ("observation", "observation_hash", "action_values", "action_id"):
        target[field] = source[field]
    refresh_hash(rows, manifest, "validation")
    with pytest.raises(ValueError, match="duplicate"):
        audit_dataset(rows, manifest)


def test_written_hashes_match_and_artifacts_are_not_overwritten(bundle, tmp_path) -> None:
    write_dataset(bundle, tmp_path)
    for split in bundle.records:
        assert (
            hashlib.sha256((tmp_path / f"{split}.jsonl").read_bytes()).hexdigest()
            == bundle.manifest["splits"][split]["sha256"]
        )
    assert json.loads((tmp_path / "manifest.json").read_text()) == bundle.manifest
    with pytest.raises(FileExistsError):
        write_dataset(bundle, tmp_path)


@pytest.mark.parametrize(
    "kwargs",
    [
        {"train_seeds": (0,)},
        {"train_seeds": (1000,)},
        {"validation_seeds": (10000,)},
        {"train_seeds": (True,)},
        {"max_pieces": 0},
        {"behavior_cycle": ("unknown",)},
    ],
)
def test_dataset_pool_and_config_validation(kwargs) -> None:
    with pytest.raises(ValueError):
        DatasetConfig(**kwargs)


def test_working_tree_provenance_is_explicit() -> None:
    result = generate_dataset(
        DatasetConfig(train_seeds=(10000,), validation_seeds=(20000,), max_pieces=1),
        COMMIT + "+working-tree",
    )
    assert result.manifest["source_commit"] == COMMIT + "+working-tree"
    assert result.manifest["source_hashes"]["data.py"]