Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions gnm/shape/gnm_data_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
16 changes: 16 additions & 0 deletions gnm/shape/gnm_data_loader_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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,
Expand Down