From 6b8253061949e01339845b64945ffc2350ac829f Mon Sep 17 00:00:00 2001 From: VsevolodX Date: Tue, 28 Jul 2026 13:58:04 -0700 Subject: [PATCH] update: fix kgrid metric --- .../providers/points_grid_data_provider.py | 17 ++++------------- .../context/test_points_grid_data_provider.py | 17 ++++++++++++++--- 2 files changed, 18 insertions(+), 16 deletions(-) diff --git a/src/py/mat3ra/wode/context/providers/points_grid_data_provider.py b/src/py/mat3ra/wode/context/providers/points_grid_data_provider.py index 85ae7eed..541abd36 100644 --- a/src/py/mat3ra/wode/context/providers/points_grid_data_provider.py +++ b/src/py/mat3ra/wode/context/providers/points_grid_data_provider.py @@ -10,6 +10,8 @@ DEFAULT_KPPRA = -1 +SCHEMA_FIELD_NAMES = set(PointsGridDataProviderSchema.model_fields) + # TODO: GlobalSetting for default KPPRA value class PointsGridDataProvider(PointsGridDataProviderSchema, ContextProvider): @@ -25,6 +27,7 @@ class PointsGridDataProvider(PointsGridDataProviderSchema, ContextProvider): shifts: List[float] = Field(default_factory=lambda: [0.0, 0.0, 0.0]) gridMetricType: GridMetricType = Field(default=GridMetricType.KPPRA) gridMetricValue: float = Field(default=DEFAULT_KPPRA) + preferGridMetric: bool = Field(default=False) @property def is_edited_key(self) -> str: @@ -32,19 +35,7 @@ def is_edited_key(self) -> str: @property def default_data(self) -> Dict[str, Any]: - data = { - "dimensions": self.dimensions, - "shifts": self.shifts, - "gridMetricType": self.grid_metric_type, - "divisor": self.divisor, - } - if self.grid_metric_value is not None: - data["gridMetricValue"] = self.grid_metric_value - if self.prefer_grid_metric is not None: - data["preferGridMetric"] = self.prefer_grid_metric - if self.reciprocal_vector_ratios is not None: - data["reciprocalVectorRatios"] = self.reciprocal_vector_ratios - return data + return self.model_dump(by_alias=True, exclude_none=True, include=SCHEMA_FIELD_NAMES) def get_reciprocal_vector_ratios(self, context: Optional[Dict[str, Any]] = None) -> Optional[List[float]]: effective_data = self._get_effective_data(context) diff --git a/tests/py/context/test_points_grid_data_provider.py b/tests/py/context/test_points_grid_data_provider.py index 9231a77f..449034eb 100644 --- a/tests/py/context/test_points_grid_data_provider.py +++ b/tests/py/context/test_points_grid_data_provider.py @@ -1,5 +1,8 @@ import pytest -from mat3ra.esse.models.context_providers_directory.points_grid_data_provider import GridMetricType +from mat3ra.esse.models.context_providers_directory.points_grid_data_provider import ( + GridMetricType, + PointsGridDataProviderSchema, +) from mat3ra.wode.context.providers import PointsGridDataProvider from mat3ra.wode.context.providers.points_grid_data_provider import DEFAULT_KPPRA # Test data constants @@ -16,8 +19,8 @@ "kgrid": { "dimensions": DIMENSIONS_CUSTOM, "shifts": SHIFTS_DEFAULT, - "divisor": DIVISOR_DEFAULT, "gridMetricType": GRID_METRIC_TYPE_DEFAULT, + "preferGridMetric": False, "gridMetricValue": DEFAULT_KPPRA, }, "isKgridEdited": True, @@ -27,8 +30,8 @@ "kgrid": { "dimensions": ["{{N_k}}", "{{N_k}}", "{{N_k}}"], "shifts": SHIFTS_DEFAULT, - "divisor": DIVISOR_DEFAULT, "gridMetricType": GRID_METRIC_TYPE_DEFAULT, + "preferGridMetric": False, "gridMetricValue": DEFAULT_KPPRA, "reciprocalVectorRatios": [1.0, 0.667, 0.5], }, @@ -87,6 +90,14 @@ def test_points_grid_data_provider_yield_data(init_params, expected_data): assert actual_data == expected_data +def test_default_data_conforms_to_esse_schema(): + """Emitted data must validate against ESSE and carry no extra keys -- guards field drift.""" + data = PointsGridDataProvider(dimensions=DIMENSIONS_CUSTOM).get_data() + + PointsGridDataProviderSchema.model_validate(data) # raises on missing/wrong required fields + assert set(data).issubset(set(PointsGridDataProviderSchema.model_fields)) # no subclass-only keys + + def test_points_grid_data_provider_get_reciprocal_vector_ratios_from_provider_data(): provider = PointsGridDataProvider(reciprocal_vector_ratios=[1.0, 0.667, 0.5])