Spaces:
Running
Running
| 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() | |