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
17 changes: 10 additions & 7 deletions src/graphty/planner.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,9 +202,8 @@ def run(self) -> pl.LazyFrame:
group_by: str | None = model_info.group_by

if group_by is None:
if self.model.model_fields:
return self.lazy_frame.select(*self._compile_exprs(model=self.model))
return self.lazy_frame
exprs: list[pl.Expr] = list(self._compile_exprs(model=self.model))
return self.lazy_frame.select(*exprs) if exprs else self.lazy_frame

return self.lazy_frame.group_by(group_by, maintain_order=True).agg(
*self._compile_exprs(model=self.model, group_context=True),
Expand Down Expand Up @@ -277,7 +276,8 @@ def _compile_exprs(
.alias(field_name)
)
else:
inner: pl.Expr = pl.col(model_info.alias_map[field_name])
dealiased_field_name: str = model_info.alias_map[field_name]
inner: pl.Expr = pl.col(dealiased_field_name)

agg: Aggregation = aggregation or Collect()
expr: pl.Expr = agg(inner)
Expand All @@ -292,13 +292,16 @@ def _compile_exprs(
)

else:
col: str = model_info.alias_map[field_name]
dealiased_field_name: str = model_info.alias_map[field_name]

if dealiased_field_name not in self._base_cols:
continue

if model_info.group_by is None:
yield pl.col(col)
yield pl.col(dealiased_field_name)
else:
reduction: Aggregation = aggregation or Reduce()
expr: pl.Expr = reduction(pl.col(col))
expr: pl.Expr = reduction(pl.col(dealiased_field_name))

yield (
expr
Expand Down
10 changes: 9 additions & 1 deletion tests/materializer/param.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
from collections.abc import Iterable
from dataclasses import dataclass
from re import Pattern


@dataclass
Expand All @@ -11,7 +13,13 @@ def __post_init__(self):
self.model_dump = self.bindings


@dataclass
class ExpectedException:
exception: Exception
match: str | Pattern | None = None


@dataclass
class Parameter:
kwargs: dict[str, object]
expected: Expected
expected: Expected | ExpectedException | Iterable[Expected | ExpectedException]
160 changes: 160 additions & 0 deletions tests/materializer/sad_path/test_sad_path_validation_error.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,160 @@
import pytest
from graphty import ConfigDict, ModelMaterializer
from pydantic import BaseModel, Field, ValidationError
from tests.materializer.param import Expected, ExpectedException, Parameter


class Model1(BaseModel):
dne: int


class Model2(BaseModel):
nested: Model1


class Model3(BaseModel):
model_config = ConfigDict(group_by="x")

x: int
dne: int


class Model4(BaseModel):
nested: Model3


class Model5(BaseModel):
model_config = ConfigDict(group_by="x")

x: int
nested: Model3


class Model6(BaseModel):
x: int = Field(alias="ALIAS")


class Model7(BaseModel):
nested: Model6


data = [
{"x": 1},
{"x": 2},
{"x": 3},
]


params: list[Parameter] = [
Parameter(
kwargs={"model": Model1, "data": data},
expected=[
Expected(bindings=[{"x": 1}, {"x": 2}, {"x": 3}]),
ExpectedException(
exception=ValidationError, match="1 validation error for Model1\ndne"
),
],
),
Parameter(
kwargs={"model": Model2, "data": data},
expected=[
Expected(
bindings=[
{"nested": {"x": 1}},
{"nested": {"x": 2}},
{"nested": {"x": 3}},
]
),
ExpectedException(
exception=ValidationError,
match="1 validation error for Model2\nnested.dne",
),
],
),
Parameter(
kwargs={"model": Model3, "data": data},
expected=[
Expected(bindings=[{"x": 1}, {"x": 2}, {"x": 3}]),
ExpectedException(
exception=ValidationError,
match="1 validation error for Model3\ndne",
),
],
),
Parameter(
kwargs={"model": Model4, "data": data},
expected=[
Expected(
bindings=[
{"nested": {"x": 1}},
{"nested": {"x": 2}},
{"nested": {"x": 3}},
]
),
ExpectedException(
exception=ValidationError,
match="1 validation error for Model4\nnested.dne",
),
],
),
Parameter(
kwargs={"model": Model5, "data": data},
expected=[
Expected(
bindings=[
{"x": 1, "nested": {"x": 1}},
{"x": 2, "nested": {"x": 2}},
{"x": 3, "nested": {"x": 3}},
]
),
ExpectedException(
exception=ValidationError,
match="1 validation error for Model5\nnested.dne",
),
],
),
Parameter(
kwargs={"model": Model6, "data": data},
expected=[
Expected(bindings=[{"x": 1}, {"x": 2}, {"x": 3}]),
ExpectedException(
exception=ValidationError,
match="1 validation error for Model6\nALIAS",
),
],
),
Parameter(
kwargs={"model": Model7, "data": data},
expected=[
Expected(
bindings=[
{"nested": {"x": 1}},
{"nested": {"x": 2}},
{"nested": {"x": 3}},
]
),
ExpectedException(
exception=ValidationError,
match="1 validation error for Model7\nnested.ALIAS",
),
],
),
]


@pytest.mark.parametrize("param", params)
def test_sad_path_empty_exprs_select(param):
materializer = ModelMaterializer(**param.kwargs)

for expected in param.expected:
match expected:
case Expected():
assert list(materializer.generate_bindings()) == expected.bindings
case ExpectedException():
with pytest.raises(
expected_exception=expected.exception, match=expected.match
):
list(materializer.generate_models())

case _:
assert False, "This should never happen."
Loading