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
5 changes: 1 addition & 4 deletions dataframely/_pydantic.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,7 @@


def _dict_to_df(schema_type: type[BaseSchema], data: dict) -> pl.DataFrame:
return pl.from_dict(
data,
schema=schema_type.to_polars_schema(), # type: ignore[attr-defined]
)
return pl.from_dict(data, schema=pl.Schema(schema_type))


def _validate_df_schema(schema_type: type[_S], df: pl.DataFrame) -> DataFrame[_S]:
Expand Down
2 changes: 1 addition & 1 deletion dataframely/filter_result.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,7 @@ def _create_empty(cls, schema: type[S], with_casting_rules: bool) -> FailureInfo
rules = schema._validation_rules(with_cast=with_casting_rules)
lf = pl.LazyFrame(
schema={
**schema.to_polars_schema(), # type: ignore
**pl.Schema(schema),
**{rule: pl.Boolean for rule in rules},
}
)
Expand Down
9 changes: 0 additions & 9 deletions dataframely/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -1385,15 +1385,6 @@ def _validate_if_needed(

# ----------------------------- THIRD-PARTY PACKAGES ----------------------------- #

@classmethod
def to_polars_schema(cls) -> pl.Schema:
"""Obtain the polars schema for this schema.

Returns:
A :mod:`polars` schema that mirrors the schema defined by this class.
"""
return pl.Schema({name: col.dtype for name, col in cls.columns().items()})

@classmethod
def to_sqlalchemy_columns(cls, dialect: sa.Dialect) -> list[sa.Column]:
"""Obtain the SQLAlchemy column definitions for a particular dialect for this
Expand Down
1 change: 0 additions & 1 deletion docs/api/schema/conversion.rst
Original file line number Diff line number Diff line change
Expand Up @@ -7,4 +7,3 @@ Conversion
:toctree: _gen/

Schema.to_sqlalchemy_columns
Schema.to_polars_schema
4 changes: 2 additions & 2 deletions tests/collection/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,9 +56,9 @@ def test_cast() -> None:
"second": pl.LazyFrame({"a": [1, 2, 3], "b": [4, 5, 6]}),
},
)
assert collection.first.collect_schema() == MyFirstSchema.to_polars_schema()
assert collection.first.collect_schema() == pl.Schema(MyFirstSchema)
assert collection.second is not None
assert collection.second.collect_schema() == MySecondSchema.to_polars_schema()
assert collection.second.collect_schema() == pl.Schema(MySecondSchema)


@pytest.mark.parametrize(
Expand Down
6 changes: 3 additions & 3 deletions tests/collection/test_cast.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,16 +26,16 @@ def test_cast_valid(df_type: type[pl.DataFrame] | type[pl.LazyFrame]) -> None:
first = df_type({"a": [3]})
second = df_type({"a": [1]})
out = Collection.cast({"first": first, "second": second}) # type: ignore
assert out.first.collect_schema() == FirstSchema.to_polars_schema()
assert out.first.collect_schema() == pl.Schema(FirstSchema)
assert out.second is not None
assert out.second.collect_schema() == SecondSchema.to_polars_schema()
assert out.second.collect_schema() == pl.Schema(SecondSchema)


@pytest.mark.parametrize("df_type", [pl.DataFrame, pl.LazyFrame])
def test_cast_valid_optional(df_type: type[pl.DataFrame] | type[pl.LazyFrame]) -> None:
first = df_type({"a": [3]})
out = Collection.cast({"first": first}) # type: ignore
assert out.first.collect_schema() == FirstSchema.to_polars_schema()
assert out.first.collect_schema() == pl.Schema(FirstSchema)
assert out.second is None


Expand Down
5 changes: 3 additions & 2 deletions tests/collection/test_create_empty.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Copyright (c) QuantCo 2025-2026
# SPDX-License-Identifier: BSD-3-Clause

import polars as pl

import dataframely as dy

Expand All @@ -23,7 +24,7 @@ class MyCollection(dy.Collection):
def test_create_empty() -> None:
collection = MyCollection.create_empty()
assert collection.first.collect().height == 0
assert collection.first.collect_schema() == MyFirstSchema.to_polars_schema()
assert collection.first.collect_schema() == pl.Schema(MyFirstSchema)
assert collection.second is not None
assert collection.second.collect().height == 0
assert collection.second.collect_schema() == MySecondSchema.to_polars_schema()
assert collection.second.collect_schema() == pl.Schema(MySecondSchema)
2 changes: 1 addition & 1 deletion tests/columns/test_polars_schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,5 +9,5 @@

def test_polars_schema() -> None:
schema = create_schema("test", {"a": dy.Int32(nullable=False), "b": dy.Float32()})
pl_schema = schema.to_polars_schema()
pl_schema = pl.Schema(schema)
assert pl_schema == {"a": pl.Int32, "b": pl.Float32}
2 changes: 1 addition & 1 deletion tests/schema/test_cast.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ def test_cast_valid(
df = df_type(data)
out = MySchema.cast(df)
assert isinstance(out, df_type)
assert out.lazy().collect_schema() == MySchema.to_polars_schema()
assert out.lazy().collect_schema() == pl.Schema(MySchema)


def test_cast_invalid_schema_eager() -> None:
Expand Down
4 changes: 2 additions & 2 deletions tests/test_pydantic.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@ def df() -> pl.DataFrame:
"y": [4, 5, 6],
"comment": ["a", "b", "c"],
},
schema=Schema.to_polars_schema(),
schema=pl.Schema(Schema),
)


Expand All @@ -54,7 +54,7 @@ def invalid_df() -> pl.DataFrame:
"y": [4],
"comment": ["a"],
},
schema=Schema.to_polars_schema(),
schema=pl.Schema(Schema),
)


Expand Down