diff --git a/src/typed_expression/validate.rs b/src/typed_expression/validate.rs index 36ae73a..4c88217 100644 --- a/src/typed_expression/validate.rs +++ b/src/typed_expression/validate.rs @@ -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, }) } diff --git a/tests/operators/test_add.py b/tests/operators/test_add.py index 63958f6..675bc02 100644 --- a/tests/operators/test_add.py +++ b/tests/operators/test_add.py @@ -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 ( @@ -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) diff --git a/tests/operators/test_validation.py b/tests/operators/test_validation.py index 09547ef..00ff4e7 100644 --- a/tests/operators/test_validation.py +++ b/tests/operators/test_validation.py @@ -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"),