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
17 changes: 4 additions & 13 deletions src/py/mat3ra/wode/context/providers/points_grid_data_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@

DEFAULT_KPPRA = -1

SCHEMA_FIELD_NAMES = set(PointsGridDataProviderSchema.model_fields)


# TODO: GlobalSetting for default KPPRA value
class PointsGridDataProvider(PointsGridDataProviderSchema, ContextProvider):
Expand All @@ -25,26 +27,15 @@ 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:
return "isKgridEdited"

@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)
Expand Down
17 changes: 14 additions & 3 deletions tests/py/context/test_points_grid_data_provider.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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,
Expand All @@ -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],
},
Expand Down Expand Up @@ -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])

Expand Down
Loading