File size: 4,292 Bytes
3b2bae1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import unittest
from types import SimpleNamespace
from unittest.mock import patch

from chinatravel.data import load_datasets


class QueryLoaderTests(unittest.TestCase):
    def test_serialized_oracle_is_parsed_and_preserved_for_evaluation(self):
        record = {
            "uid": "uid-1",
            "hard_logic_py": "['result=True']",
            "nature_language": "test",
        }
        args = SimpleNamespace(
            splits="easy",
            lang="zh",
            oracle_translation=True,
        )

        with (
            patch.object(
                load_datasets,
                "_load_oracle_snapshot",
                return_value=[record],
            ),
            patch.object(
                load_datasets,
                "_configured_query_ids",
                return_value=["uid-1"],
            ),
        ):
            query_ids, records = load_datasets.load_query(args)

        self.assertEqual(query_ids, ["uid-1"])
        self.assertEqual(records["uid-1"]["hard_logic_py"], ["result=True"])

    def test_agent_facing_load_strips_oracle_fields(self):
        record = {
            "uid": "uid-1",
            "hard_logic_py": ["result=True"],
            "nature_language": "test",
        }
        args = SimpleNamespace(
            splits="easy",
            lang="zh",
            oracle_translation=False,
        )

        with (
            patch.object(
                load_datasets,
                "_load_huggingface_split",
                return_value=[record],
            ),
            patch.object(
                load_datasets,
                "_configured_query_ids",
                return_value=["uid-1"],
            ),
        ):
            _, records = load_datasets.load_query(args)

        self.assertNotIn("hard_logic_py", records["uid-1"])

    def test_all_supported_snapshots_have_complete_oracles(self):
        supported = {"zh": ("easy", "human"), "en": ("easy", "human")}
        for lang, splits in supported.items():
            for split in splits:
                with self.subTest(lang=lang, split=split):
                    args = SimpleNamespace(
                        splits=split,
                        lang=lang,
                        oracle_translation=True,
                    )
                    query_ids, records = load_datasets.load_query(args)
                    self.assertEqual(set(query_ids), set(records))
                    self.assertTrue(
                        all(
                            isinstance(record.get("hard_logic_py"), list)
                            for record in records.values()
                        )
                    )

    def test_human1000_english_data_is_not_treated_as_an_oracle(self):
        args = SimpleNamespace(
            splits="human1000",
            lang="en",
            oracle_translation=True,
        )

        with self.assertRaisesRegex(ValueError, "No en Oracle source"):
            load_datasets.load_query(args)

    def test_human1000_oracle_is_loaded_from_runtime_source(self):
        args = SimpleNamespace(
            splits="human1000",
            lang="zh",
            oracle_translation=True,
        )
        record = {
            "uid": "uid-1",
            "hard_logic_py": "['result=True']",
            "nature_language": "test",
        }

        with (
            patch.object(load_datasets, "_load_oracle_snapshot", return_value=None),
            patch.object(
                load_datasets,
                "_load_huggingface_split",
                return_value=[record],
            ) as remote_loader,
            patch.object(
                load_datasets,
                "_configured_query_ids",
                return_value=["uid-1"],
            ),
        ):
            query_ids, records = load_datasets.load_query(args)

        remote_loader.assert_called_once_with("human1000")
        self.assertEqual(query_ids, ["uid-1"])
        self.assertEqual(records["uid-1"]["hard_logic_py"], ["result=True"])

    def test_human1000_uses_oracle_benchmark_source(self):
        self.assertEqual(
            load_datasets.HUGGINGFACE_SPLIT_SOURCES["human1000"],
            ("LAMDA-NeSy/chinatravel_test", "test"),
        )


if __name__ == "__main__":
    unittest.main()