Skip to content
Closed
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
39 changes: 38 additions & 1 deletion downstream_dummy/tests/test_dummy.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
import pytest
from dummy import object_to_quantity
from pint import Quantity as PintQuantity

from openff.units import Quantity, Unit
from openff.units import Quantity, Unit, unit


def test_function_can_be_defined():
Expand Down Expand Up @@ -33,3 +34,39 @@ def dummy_function(
)
def test_object_to_quantity(input, output):
assert object_to_quantity(input) == output


def test_pydantic_model():
pydantic = pytest.importorskip("pydantic")

from typing import Annotated

from pydantic_pint import PydanticPintQuantity

from openff.units.units import Quantity

class DummyModel(pydantic.BaseModel):
quantity: Annotated[PintQuantity, PydanticPintQuantity("kilocalories_per_mole", ureg=unit)]

model = DummyModel(quantity=Quantity("1.0 * kilocalories_per_mole"))

assert DummyModel.model_validate(model.model_dump()).quantity == Quantity(
"1.0 * kilocalories_per_mole"
)


def test_pydantic_json():
pydantic = pytest.importorskip("pydantic")

from typing import Annotated

from pydantic_pint import PydanticPintQuantity

class DummyModel(pydantic.BaseModel):
quantity: Annotated[PintQuantity, PydanticPintQuantity("kilocalories_per_mole", ureg=unit)]

model = DummyModel(quantity="1.0 * kilocalories_per_mole")

assert DummyModel.model_validate_json(model.model_dump_json()).quantity == Quantity(
"1.0 * kilocalories_per_mole"
)
4 changes: 2 additions & 2 deletions openff/units/elements.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@

"""

from openff.units import Quantity, unit
from openff.units import Quantity

__all__ = [
"MASSES",
Expand All @@ -39,7 +39,7 @@
"""Mapping from atomic number to atomic mass"""
MASSES: dict[int, Quantity] = {
# https://github.com/hgrecco/pint/issues/1804
index + 1: Quantity(mass, unit.dalton)
index + 1: Quantity(mass, "dalton")
for index, mass in enumerate(
[
1.007947,
Expand Down
27 changes: 22 additions & 5 deletions openff/units/units.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,9 @@
import pint
from openff.utilities import requires_package
from pint import Measurement as _Measurement
from pint import Quantity as _Quantity
from pint import Quantity as PintQuantity
from pint import Unit as _Unit
from pydantic_pint import PydanticPintQuantity

from openff.units.utilities import get_defaults_path

Expand All @@ -24,15 +25,25 @@
"Unit",
"unit",
)
from typing import Annotated, Any, TypeVar

from pydantic import PlainSerializer


def ser_type(value: type[Any]) -> str:
return value.__name__


T = TypeVar("T")

JSONSerializableType = Annotated[type[T], PlainSerializer(ser_type, when_used="json-unless-none")]


class Unit(pint.UnitRegistry.Unit):
"""A unit of measure."""

pass


class Quantity(pint.UnitRegistry.Quantity):
class _Quantity(pint.UnitRegistry.Quantity):
"""A value with associated units."""

def __dask_tokenize__(self):
Expand Down Expand Up @@ -71,7 +82,7 @@ def _dask_finalize(results, func, args, units):


class UnitRegistry(pint.UnitRegistry):
_quantity_class = Quantity
_quantity_class = _Quantity
_unit_class = Unit
_measurement_class = Measurement

Expand All @@ -86,6 +97,12 @@ class UnitRegistry(pint.UnitRegistry):

pint.set_application_registry(DEFAULT_UNIT_REGISTRY)

Quantity = Annotated[
PintQuantity, PydanticPintQuantity("kilocalories_per_mole", ureg=DEFAULT_UNIT_REGISTRY)
]

DEFAULT_UNIT_REGISTRY._quantity_class = Quantity

Quantity.to_openmm = _to_openmm # type: ignore[attr-defined]

with warnings.catch_warnings():
Expand Down
Loading
Loading