feat(checker): add arithmetic binary ops on columns

This commit is contained in:
HEL
2026-07-03 00:42:00 +02:00
parent 74c07c9afb
commit ece2e3a6a3
+30 -3
View File
@@ -69,8 +69,7 @@ class ColumnMethodRegistry(MethodRegistry[Call]):
new_column = ColumnType(type=new_inner_type) new_column = ColumnType(type=new_inner_type)
return new_column return new_column
@method("add", "__add__") def _element_wise(self, call: Call, method: str) -> Type:
def add(self, call: Call) -> Type:
# TODO: support add with scalar # TODO: support add with scalar
# Build signature with new column type and generic operand # Build signature with new column type and generic operand
@@ -87,7 +86,7 @@ class ColumnMethodRegistry(MethodRegistry[Call]):
required=True, required=True,
), ),
], ],
returns=self._element_binary_op(call, "__add__"), returns=self._element_binary_op(call, method),
), ),
) )
@@ -105,6 +104,34 @@ class ColumnMethodRegistry(MethodRegistry[Call]):
return result.result return result.result
@method("add", "__add__")
def add(self, call: Call) -> Type:
return self._element_wise(call, "__add__")
@method("sub", "__sub__")
def sub(self, call: Call) -> Type:
return self._element_wise(call, "__sub__")
@method("mul", "__mul__")
def mul(self, call: Call) -> Type:
return self._element_wise(call, "__mul__")
@method("div", "truediv", "__truediv__")
def truediv(self, call: Call) -> Type:
return self._element_wise(call, "__truediv__")
@method("floordiv", "__floordiv__")
def floordiv(self, call: Call) -> Type:
return self._element_wise(call, "__floordiv__")
@method("mod", "__mod__")
def mod(self, call: Call) -> Type:
return self._element_wise(call, "__mod__")
@method("pow", "__pow__")
def pow(self, call: Call) -> Type:
return self._element_wise(call, "__pow__")
@method() @method()
def mean(self, call: Call) -> Type: def mean(self, call: Call) -> Type:
signature = Function( signature = Function(