diff --git a/gnm/shape/gnm_data_loader.py b/gnm/shape/gnm_data_loader.py index b260c58f..6722fd94 100644 --- a/gnm/shape/gnm_data_loader.py +++ b/gnm/shape/gnm_data_loader.py @@ -110,8 +110,8 @@ def _validate_gnm_data( has exactly the expected fields, the missing fields and the extra fields. """ expected_fields = gnm_data_schema.GNM_DATA_ATTRIBUTES - missing_fields = list(set(expected_fields) - set(data.keys())) - extra_fields = list(set(data.keys()) - set(expected_fields)) + missing_fields = sorted(set(expected_fields) - set(data.keys())) + extra_fields = sorted(set(data.keys()) - set(expected_fields)) return not missing_fields and not extra_fields, missing_fields, extra_fields diff --git a/gnm/shape/gnm_data_loader_test.py b/gnm/shape/gnm_data_loader_test.py index e51a3e8b..c6dd694b 100644 --- a/gnm/shape/gnm_data_loader_test.py +++ b/gnm/shape/gnm_data_loader_test.py @@ -14,12 +14,15 @@ """Tests for gnm_data_loader.""" +# pylint: disable=protected-access + from unittest import mock from absl.testing import absltest from absl.testing import parameterized from etils import epath from gnm.shape import gnm_data_loader +from gnm.shape import gnm_data_schema from gnm.shape.data.versions import gnm_specs from gnm.shape.data.versions import gnm_test_catalog @@ -45,6 +48,19 @@ def test_print_gnm_versions(self): class GNMModelLoadingTest(parameterized.TestCase): """Tests for loading GNM model files.""" + def test_validate_gnm_data_returns_sorted_diagnostics(self): + data = {field: object() for field in gnm_data_schema.GNM_DATA_ATTRIBUTES} + del data['identity_names'] + del data['joint_names'] + data['zz_extra'] = object() + data['aa_extra'] = object() + + valid, missing, extra = gnm_data_loader._validate_gnm_data(data) + + self.assertFalse(valid) + self.assertEqual(missing, ['identity_names', 'joint_names']) + self.assertEqual(extra, ['aa_extra', 'zz_extra']) + @parameterized.product( version=_MAINTAINED_MAJOR_GNM_VERSIONS, variant=gnm_test_catalog.ALL_VARIANTS,