Skip to content
Merged
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
21 changes: 17 additions & 4 deletions dpsynth/adapters/pydantic.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@ def _get_base_type(annotation: type[Any]) -> tuple[bool, type[Any]]:

def _numerical_attribute_from_field_info(
field_info: FieldInfo,
bounds: tuple[float, float] | None = None,
) -> domain.NumericalAttribute:
"""Infers a NumericalAttribute from a pydantic FieldInfo."""
# NumericalAttribute uses a convention where both the min_value and max_value
Expand All @@ -94,6 +95,8 @@ def _numerical_attribute_from_field_info(
upper_bound = math.nextafter(meta.lt, -math.inf)
case _:
continue
if bounds is not None:
lower_bound, upper_bound = bounds
if lower_bound is None or upper_bound is None:
raise ValueError("Must specify lower and upper bounds for numeric fields.")

Expand Down Expand Up @@ -132,9 +135,11 @@ def _categorical_attribute_from_field_info(

def infer_domain(
model_cls: type[pydantic.BaseModel],
*,
numerical_bounds: Mapping[str, tuple[float, float]] | None = None,
) -> dict[str, domain.AttributeType]:
"""Infers the domain of a pydantic model."""

numerical_bounds = numerical_bounds or {}
attributes: dict[str, domain.AttributeType] = {}
for name, meta in model_cls.model_fields.items():
_, base_type = _get_base_type(meta.annotation) # pyrefly: ignore[bad-argument-type]
Expand All @@ -143,10 +148,18 @@ def infer_domain(
is_model = is_class and issubclass(base_type, pydantic.BaseModel)
is_literal = typing.get_origin(base_type) is Literal
if is_model:
sub = infer_domain(base_type)
attributes.update({f"{name}.{k}": v for k, v in sub.items()})
prefix = f"{name}."
sub_bounds = {
k.removeprefix(prefix): v
for k, v in numerical_bounds.items()
if k.startswith(prefix)
}
sub = infer_domain(base_type, numerical_bounds=sub_bounds)
attributes.update({f"{prefix}{k}": v for k, v in sub.items()})
elif base_type in (int, float):
attributes[name] = _numerical_attribute_from_field_info(meta)
attributes[name] = _numerical_attribute_from_field_info(
meta, bounds=numerical_bounds.get(name)
)
elif base_type is str:
attributes[name] = domain.OpenSetCategoricalAttribute(
description=meta.description
Expand Down
35 changes: 35 additions & 0 deletions tests/adapters/pydantic_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,41 @@ def test_infer_domain_unsupported_type(self):
):
pydantic_api.infer_domain(ModelWithUnsupportedType)

def test_infer_domain_with_numerical_bounds(self):
class Nested(pydantic.BaseModel):
val: float | None = pydantic.Field(description="Signal value")

class Root(pydantic.BaseModel):
steps: Nested
hrv: Nested

schema = pydantic_api.infer_domain(
Root,
numerical_bounds={
"steps.val": (0.0, 50000.0),
"hrv.val": (0.0, 300.0),
},
)
self.assertEqual(
schema,
{
"steps.val": domain.NumericalAttribute(
min_value=0.0,
max_value=50000.0,
clip_to_range=False,
dtype="float",
description="Signal value",
),
"hrv.val": domain.NumericalAttribute(
min_value=0.0,
max_value=300.0,
clip_to_range=False,
dtype="float",
description="Signal value",
),
},
)

def test_open_set_str_description_and_optional_roundtrip(self):
class ModelWithDescriptions(pydantic.BaseModel):
timezone: str | None = pydantic.Field(description="IANA timezone")
Expand Down
Loading