ChinaTravel / tests /test_query_loader.py
Cbphcr's picture
Publish ChinaTravel static evaluator entry
3b2bae1
Raw
History Blame Contribute Delete
4.29 kB
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()