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()