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
36 changes: 30 additions & 6 deletions src/typed_expression/validate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -115,12 +115,36 @@ impl Expression {

// Arithmetic operators
Expression::Add { left, right } => {
validate_binary_arithmetic(left, right, df_type, "addition", |l, r, et| {
TypedExpression::Add {
left: Arc::new(l),
right: Arc::new(r),
expression_type: et,
}
let typed_left = left.validate(df_type)?;
let typed_right = right.validate(df_type)?;
let left_type = typed_left.expression_type();
let right_type = typed_right.expression_type();
let left_dt = left_type.data_type();
let right_dt = right_type.data_type();

let is_addition = left_dt.is_numeric() && right_dt.is_numeric();
let is_concatenation = matches!(
(left_dt, right_dt),
(
DataType::String | DataType::Nothing,
DataType::String | DataType::Nothing
)
);

if !is_addition && !is_concatenation {
return Err(ValidationError::IncompatibleTypes {
operation: "addition".to_string(),
left_type: left_dt,
right_type: right_dt,
});
}

let result_type = promote_expression_types(left_type, right_type, "addition")?;
let result_dt = result_type.data_type();
Ok(TypedExpression::Add {
left: Arc::new(typed_left.cast_if_needed(result_dt)),
right: Arc::new(typed_right.cast_if_needed(result_dt)),
expression_type: result_type,
})
}

Expand Down
98 changes: 98 additions & 0 deletions tests/operators/test_add.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import pytest

from tabeline import Array, DataFrame, DataType
from tabeline.exceptions import IncompatibleTypesError
from tabeline.testing import assert_data_frames_equal

from .._types import (
Expand Down Expand Up @@ -123,3 +124,100 @@ def test_add_with_decimal_literal(original_dtype, expected_dtype, expression):
actual = df.transmute(result=expression)
expected = DataFrame(result=Array[expected_dtype](4.5, 6.5, None))
assert_data_frames_equal(actual, expected, absolute_tolerance=absolute_tolerance)


@pytest.mark.parametrize(
("left", "right", "answer"),
[
("a", "b", "ab"),
("", "x", "x"),
("x", "", "x"),
("x", None, None),
(None, "x", None),
(None, None, None),
],
)
def test_concatenate_strings(left, right, answer):
df = DataFrame(a=Array[DataType.String](left), b=Array[DataType.String](right))
actual = df.transmute(c="a + b")
expected = DataFrame(c=Array[DataType.String](answer))
assert_data_frames_equal(actual, expected)


@pytest.mark.parametrize(
("value", "answer"),
[
("hello", "hello world"),
("", " world"),
(None, None),
],
)
def test_concatenate_column_with_literal(value, answer):
df = DataFrame(a=Array[DataType.String](value))
actual = df.transmute(c="a + ' world'")
expected = DataFrame(c=Array[DataType.String](answer))
assert_data_frames_equal(actual, expected)


@pytest.mark.parametrize(
("value", "answer"),
[
("world", "hello world"),
("", "hello "),
(None, None),
],
)
def test_concatenate_literal_with_column(value, answer):
df = DataFrame(a=Array[DataType.String](value))
actual = df.transmute(c="'hello ' + a")
expected = DataFrame(c=Array[DataType.String](answer))
assert_data_frames_equal(actual, expected)


@pytest.mark.parametrize(
("string_value", "answer"),
[
("x", None),
(None, None),
],
)
def test_concatenate_string_with_nothing(string_value, answer):
df = DataFrame(a=Array[DataType.String](string_value), b=Array[DataType.Nothing](None))
actual = df.transmute(c="a + b")
expected = DataFrame(c=Array[DataType.Nothing](answer))
assert_data_frames_equal(actual, expected)


@pytest.mark.parametrize(
("string_value", "answer"),
[
("x", None),
(None, None),
],
)
def test_concatenate_nothing_with_string(string_value, answer):
df = DataFrame(a=Array[DataType.Nothing](None), b=Array[DataType.String](string_value))
actual = df.transmute(c="a + b")
expected = DataFrame(c=Array[DataType.Nothing](answer))
assert_data_frames_equal(actual, expected)


@pytest.mark.parametrize(
("left_values", "right_values", "left_type", "right_type"),
[
([1, 2, 3], [True, False, True], DataType.Integer64, DataType.Boolean),
([1, 2, 3], ["a", "b", "c"], DataType.Integer64, DataType.String),
(["a", "b", "c"], [1, 2, 3], DataType.String, DataType.Integer64),
(["a", "b", "c"], [True, False, True], DataType.String, DataType.Boolean),
([True, False, True], [True, False, True], DataType.Boolean, DataType.Boolean),
([True, False, True], [1, 2, 3], DataType.Boolean, DataType.Integer64),
([True, False, True], ["a", "b", "c"], DataType.Boolean, DataType.String),
],
)
def test_addition_rejects_incompatible_operand(left_values, right_values, left_type, right_type):
df = DataFrame(x=left_values, y=right_values)

with pytest.raises(IncompatibleTypesError) as exc_info:
df.mutate(z="x + y")

assert exc_info.value == IncompatibleTypesError("addition", left_type, right_type)
2 changes: 1 addition & 1 deletion tests/operators/test_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
@pytest.mark.parametrize(
("expression", "operation"),
[
("x + y", "addition"),
# No `x + y` because `+` also does string concatenation
("x - y", "subtraction"),
("x * y", "multiplication"),
("x / y", "division"),
Expand Down
Loading