Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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: 18 additions & 3 deletions narwhals/_arrow/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,6 @@

from narwhals._arrow.series import ArrowSeries
from narwhals._compliant import EagerExpr
from narwhals._expression_parsing import ExprKind
from narwhals._expression_parsing import evaluate_output_names_and_aliases
from narwhals._expression_parsing import is_scalar_like
from narwhals.exceptions import ColumnNotFoundError
Expand All @@ -23,6 +22,7 @@

from narwhals._arrow.dataframe import ArrowDataFrame
from narwhals._arrow.namespace import ArrowNamespace
from narwhals._expression_parsing import ExprMetadata
from narwhals.utils import Version
from narwhals.utils import _FullContext

Expand Down Expand Up @@ -52,6 +52,21 @@ def __init__(
self._backend_version = backend_version
self._version = version
self._call_kwargs = call_kwargs or {}
self._metadata: ExprMetadata | None = None

def with_metadata(self, metadata: ExprMetadata) -> Self:
expr = self.__class__(
self._call,
function_name=self._function_name,
evaluate_output_names=self._evaluate_output_names,
alias_output_names=self._alias_output_names,
backend_version=self._backend_version,
version=self._version,
depth=self._depth,
call_kwargs=self._call_kwargs,
)
expr._metadata = metadata
return expr
Comment thread
dangotbanned marked this conversation as resolved.
Outdated

@classmethod
def from_column_names(
Expand Down Expand Up @@ -141,10 +156,10 @@ def shift(self: Self, n: int) -> Self:
def over(
self: Self,
partition_by: Sequence[str],
kind: ExprKind,
order_by: Sequence[str] | None,
) -> Self:
if partition_by and not is_scalar_like(kind):
assert self._metadata is not None # noqa: S101
if partition_by and not is_scalar_like(self._metadata.kind):
msg = "Only aggregation or literal operations are supported in grouped `over` context for PyArrow."
raise NotImplementedError(msg)

Expand Down
8 changes: 5 additions & 3 deletions narwhals/_compliant/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@
from narwhals._compliant.namespace import EagerNamespace
from narwhals._compliant.series import CompliantSeries
from narwhals._expression_parsing import ExprKind
from narwhals._expression_parsing import ExprMetadata
from narwhals.dtypes import DType
from narwhals.typing import TimeUnit
from narwhals.utils import Implementation
Expand Down Expand Up @@ -85,6 +86,7 @@ class CompliantExpr(Protocol38[CompliantFrameT, CompliantSeriesOrNativeExprT_co]
_alias_output_names: Callable[[Sequence[str]], Sequence[str]] | None
_depth: int
_function_name: str
_metadata: ExprMetadata | None

def __call__(
self, df: CompliantFrameT
Expand All @@ -105,6 +107,8 @@ def from_column_names(
@classmethod
def from_column_indices(cls, *column_indices: int, context: _FullContext) -> Self: ...

def with_metadata(self, metadata: ExprMetadata) -> Self: ...

def is_null(self) -> Self: ...
def abs(self) -> Self: ...
def all(self) -> Self: ...
Expand Down Expand Up @@ -165,9 +169,7 @@ def replace_strict(
*,
return_dtype: DType | type[DType] | None,
) -> Self: ...
def over(
self: Self, keys: Sequence[str], kind: ExprKind, order_by: Sequence[str] | None
) -> Self: ...
def over(self: Self, keys: Sequence[str], order_by: Sequence[str] | None) -> Self: ...
def sample(
self,
n: int | None,
Expand Down
17 changes: 16 additions & 1 deletion narwhals/_dask/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@

from narwhals._dask.dataframe import DaskLazyFrame
from narwhals._dask.namespace import DaskNamespace
from narwhals._expression_parsing import ExprMetadata
from narwhals.dtypes import DType
from narwhals.utils import Version
from narwhals.utils import _FullContext
Expand Down Expand Up @@ -66,6 +67,7 @@ def __init__(
self._backend_version = backend_version
self._version = version
self._call_kwargs = call_kwargs or {}
self._metadata: ExprMetadata | None = None

def __call__(self: Self, df: DaskLazyFrame) -> Sequence[dx.Series]:
return self._call(df)
Expand All @@ -78,6 +80,20 @@ def __narwhals_namespace__(self) -> DaskNamespace: # pragma: no cover

return DaskNamespace(backend_version=self._backend_version, version=self._version)

def with_metadata(self, metadata: ExprMetadata) -> Self:
expr = self.__class__(
self._call,
function_name=self._function_name,
evaluate_output_names=self._evaluate_output_names,
alias_output_names=self._alias_output_names,
backend_version=self._backend_version,
version=self._version,
depth=self._depth,
call_kwargs=self._call_kwargs,
)
expr._metadata = metadata
return expr

def broadcast(self, kind: Literal[ExprKind.AGGREGATION, ExprKind.LITERAL]) -> Self:
def func(df: DaskLazyFrame) -> list[dx.Series]:
return [result[0] for result in self(df)]
Expand Down Expand Up @@ -545,7 +561,6 @@ def null_count(self: Self) -> Self:
def over(
self: Self,
partition_by: Sequence[str],
kind: ExprKind,
order_by: Sequence[str] | None,
) -> Self:
# pandas is a required dependency of dask so it's safe to import this
Expand Down
14 changes: 14 additions & 0 deletions narwhals/_duckdb/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@

from narwhals._duckdb.dataframe import DuckDBLazyFrame
from narwhals._duckdb.namespace import DuckDBNamespace
from narwhals._expression_parsing import ExprMetadata
from narwhals.dtypes import DType
from narwhals.utils import Version
from narwhals.utils import _FullContext
Expand All @@ -58,6 +59,7 @@ def __init__(
self._alias_output_names = alias_output_names
self._backend_version = backend_version
self._version = version
self._metadata: ExprMetadata | None = None

def __call__(self: Self, df: DuckDBLazyFrame) -> Sequence[duckdb.Expression]:
return self._call(df)
Expand All @@ -72,6 +74,18 @@ def __narwhals_namespace__(self) -> DuckDBNamespace: # pragma: no cover
backend_version=self._backend_version, version=self._version
)

def with_metadata(self, metadata: ExprMetadata) -> Self:
expr = self.__class__(
self._call,
function_name=self._function_name,
evaluate_output_names=self._evaluate_output_names,
alias_output_names=self._alias_output_names,
backend_version=self._backend_version,
version=self._version,
)
expr._metadata = metadata
return expr

def broadcast(self, kind: Literal[ExprKind.AGGREGATION, ExprKind.LITERAL]) -> Self:
if kind is ExprKind.AGGREGATION:
msg = "Broadcasting aggregations is not yet supported for DuckDB."
Expand Down
19 changes: 17 additions & 2 deletions narwhals/_pandas_like/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
from typing import Sequence

from narwhals._compliant import EagerExpr
from narwhals._expression_parsing import ExprKind
from narwhals._expression_parsing import evaluate_output_names_and_aliases
from narwhals._expression_parsing import is_elementary_expression
from narwhals._pandas_like.group_by import PandasLikeGroupBy
Expand All @@ -18,6 +17,7 @@
if TYPE_CHECKING:
from typing_extensions import Self

from narwhals._expression_parsing import ExprMetadata
from narwhals._pandas_like.dataframe import PandasLikeDataFrame
from narwhals._pandas_like.namespace import PandasLikeNamespace
from narwhals.utils import Implementation
Expand Down Expand Up @@ -89,6 +89,7 @@ def __init__(
self._backend_version = backend_version
self._version = version
self._call_kwargs = call_kwargs or {}
self._metadata: ExprMetadata | None = None

def __narwhals_namespace__(self: Self) -> PandasLikeNamespace:
from narwhals._pandas_like.namespace import PandasLikeNamespace
Expand All @@ -99,6 +100,21 @@ def __narwhals_namespace__(self: Self) -> PandasLikeNamespace:

def __narwhals_expr__(self) -> None: ...

def with_metadata(self, metadata: ExprMetadata) -> Self:
expr = self.__class__(
self._call,
function_name=self._function_name,
evaluate_output_names=self._evaluate_output_names,
alias_output_names=self._alias_output_names,
backend_version=self._backend_version,
version=self._version,
depth=self._depth,
implementation=self._implementation,
call_kwargs=self._call_kwargs,
)
expr._metadata = metadata
return expr

@classmethod
def from_column_names(
cls: type[Self],
Expand Down Expand Up @@ -196,7 +212,6 @@ def shift(self: Self, n: int) -> Self:
def over(
self: Self,
partition_by: Sequence[str],
kind: ExprKind,
order_by: Sequence[str] | None,
) -> Self:
if not partition_by:
Expand Down
12 changes: 11 additions & 1 deletion narwhals/_polars/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from typing_extensions import Self

from narwhals._expression_parsing import ExprKind
from narwhals._expression_parsing import ExprMetadata
from narwhals.dtypes import DType
from narwhals.utils import Version

Expand All @@ -29,6 +30,7 @@ def __init__(
self._implementation = Implementation.POLARS
self._version = version
self._backend_version = backend_version
self._metadata: ExprMetadata | None = None

def __repr__(self: Self) -> str: # pragma: no cover
return "PolarsExpr"
Expand All @@ -38,6 +40,15 @@ def _from_native_expr(self: Self, expr: pl.Expr) -> Self:
expr, version=self._version, backend_version=self._backend_version
)

def with_metadata(self, metadata: ExprMetadata) -> Self:
expr = self.__class__(
self._native_expr,
backend_version=self._backend_version,
version=self._version,
)
expr._metadata = metadata
return expr

@classmethod
def _from_series(cls, series: Any) -> Self:
return cls(
Expand Down Expand Up @@ -108,7 +119,6 @@ def is_nan(self: Self) -> Self:
def over(
self: Self,
partition_by: Sequence[str],
kind: ExprKind,
order_by: Sequence[str] | None,
) -> Self:
if self._backend_version < (1, 9):
Expand Down
18 changes: 17 additions & 1 deletion narwhals/_spark_like/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from sqlframe.base.window import Window
from typing_extensions import Self

from narwhals._expression_parsing import ExprMetadata
from narwhals._spark_like.dataframe import SparkLikeLazyFrame
from narwhals._spark_like.namespace import SparkLikeNamespace
from narwhals._spark_like.typing import WindowFunction
Expand Down Expand Up @@ -60,6 +61,7 @@ def __init__(
self._version = version
self._implementation = implementation
self._window_function: WindowFunction | None = None
self._metadata: ExprMetadata | None = None

def __call__(self: Self, df: SparkLikeLazyFrame) -> Sequence[Column]:
return self._call(df)
Expand Down Expand Up @@ -123,6 +125,21 @@ def __narwhals_namespace__(self: Self) -> SparkLikeNamespace: # pragma: no cove
implementation=self._implementation,
)

def with_metadata(self, metadata: ExprMetadata) -> Self:
expr = self.__class__(
self._call,
function_name=self._function_name,
evaluate_output_names=self._evaluate_output_names,
alias_output_names=self._alias_output_names,
backend_version=self._backend_version,
version=self._version,
implementation=self._implementation,
)
if self._window_function is not None:
expr = expr._with_window_function(self._window_function)
expr._metadata = metadata
return expr

@classmethod
def from_column_names(
cls: type[Self],
Expand Down Expand Up @@ -499,7 +516,6 @@ def _n_unique(_input: Column) -> Column:
def over(
self: Self,
partition_by: Sequence[str],
kind: ExprKind,
order_by: Sequence[str] | None,
) -> Self:
if (window_function := self._window_function) is not None:
Expand Down
6 changes: 3 additions & 3 deletions narwhals/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -1583,9 +1583,9 @@ def over(
)

return self.__class__(
lambda plx: self._to_compliant_expr(plx).over(
flat_partition_by, order_by=order_by, kind=self._metadata.kind
),
lambda plx: self._to_compliant_expr(plx)
.with_metadata(self._metadata)
.over(flat_partition_by, order_by=order_by),
metadata,
)
Comment thread
dangotbanned marked this conversation as resolved.

Expand Down