hamelin_lib 0.9.2

Core library for Hamelin query language
Documentation
//! Tests for broadcasting functionality
//!
//! Broadcasting allows scalar operations to be applied element-wise to arrays.
//! For example, `[1, 2, 3] * 10` broadcasts and returns `array(int)`.
//!
//! Note: The AST comparison tests the parsed structure (before broadcast coercion).
//! The type/schema comparison verifies that broadcasting produces the correct result type.
//!
//! See also: `lambda.rs` for transform() function tests.

use crate::err::TranslationErrors;
use crate::tree::ast::pipeline::Pipeline;
use crate::tree::ast::ParseWithErrors;
use crate::tree::builder::*;
use crate::type_check;
use crate::types::array::Array;
use crate::types::struct_type::Struct;
use crate::types::{BOOLEAN, INT, INTERVAL, STRING, TIMESTAMP};
use pretty_assertions::assert_eq;
use rstest::rstest;

// ======================== Broadcasting Tests ========================

#[rstest]
// Arithmetic broadcasting - array on left
#[case::broadcast_multiply_array_left(
    "SET result = [1, 2, 3] * 10",
    pipeline().set_cmd(|l| l.named_field("result", multiply(array().element(1).element(2).element(3), 10))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::broadcast_add_array_left(
    "SET result = [1, 2, 3] + 10",
    pipeline().set_cmd(|l| l.named_field("result", add(array().element(1).element(2).element(3), 10))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::broadcast_subtract_array_left(
    "SET result = [10, 20, 30] - 5",
    pipeline().set_cmd(|l| l.named_field("result", subtract(array().element(10).element(20).element(30), 5))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::broadcast_divide_array_left(
    "SET result = [10, 20, 30] / 2",
    pipeline().set_cmd(|l| l.named_field("result", divide(array().element(10).element(20).element(30), 2))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::broadcast_modulo_array_left(
    "SET result = [10, 21, 32] % 3",
    pipeline().set_cmd(|l| l.named_field("result", modulo(array().element(10).element(21).element(32), 3))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
// Arithmetic broadcasting - array on right
#[case::broadcast_multiply_array_right(
    "SET result = 10 * [1, 2, 3]",
    pipeline().set_cmd(|l| l.named_field("result", multiply(10, array().element(1).element(2).element(3)))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::broadcast_add_array_right(
    "SET result = 10 + [1, 2, 3]",
    pipeline().set_cmd(|l| l.named_field("result", add(10, array().element(1).element(2).element(3)))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
// Unary prefix broadcasting
#[case::broadcast_negate_array(
    "SET result = -[1, 2, 3]",
    pipeline().set_cmd(|l| l.named_field("result", negate(array().element(1).element(2).element(3)))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::broadcast_not_array(
    "SET result = NOT [true, false]",
    pipeline().set_cmd(|l| l.named_field("result", not(array().element(true).element(false)))),
    Struct::default().with_str("result", Array::new(BOOLEAN).into())
)]
// Comparison broadcasting - returns array of booleans
#[case::broadcast_greater_than(
    "SET result = [1, 5, 10] > 3",
    pipeline().set_cmd(|l| l.named_field("result", gt(array().element(1).element(5).element(10), 3))),
    Struct::default().with_str("result", Array::new(BOOLEAN).into())
)]
#[case::broadcast_less_than(
    "SET result = [1, 5, 10] < 3",
    pipeline().set_cmd(|l| l.named_field("result", lt(array().element(1).element(5).element(10), 3))),
    Struct::default().with_str("result", Array::new(BOOLEAN).into())
)]
#[case::broadcast_equal(
    "SET result = [1, 2, 3] == 2",
    pipeline().set_cmd(|l| l.named_field("result", eq(array().element(1).element(2).element(3), 2))),
    Struct::default().with_str("result", Array::new(BOOLEAN).into())
)]
#[case::broadcast_not_equal(
    "SET result = [1, 2, 3] != 2",
    pipeline().set_cmd(|l| l.named_field("result", ne(array().element(1).element(2).element(3), 2))),
    Struct::default().with_str("result", Array::new(BOOLEAN).into())
)]
#[case::broadcast_greater_than_or_equal(
    "SET result = [1, 5, 10] >= 5",
    pipeline().set_cmd(|l| l.named_field("result", gte(array().element(1).element(5).element(10), 5))),
    Struct::default().with_str("result", Array::new(BOOLEAN).into())
)]
#[case::broadcast_less_than_or_equal(
    "SET result = [1, 5, 10] <= 5",
    pipeline().set_cmd(|l| l.named_field("result", lte(array().element(1).element(5).element(10), 5))),
    Struct::default().with_str("result", Array::new(BOOLEAN).into())
)]
// String broadcasting
#[case::broadcast_string_concat(
    "SET result = ['hello', 'world'] + '!'",
    pipeline().set_cmd(|l| l.named_field("result", add(array().element("hello").element("world"), "!"))),
    Struct::default().with_str("result", Array::new(STRING).into())
)]
// Function broadcasting
#[case::broadcast_upper_function(
    "SET result = upper(['hello', 'world'])",
    pipeline().set_cmd(|l| l.named_field("result", call("upper").arg(array().element("hello").element("world")))),
    Struct::default().with_str("result", Array::new(STRING).into())
)]
#[case::broadcast_lower_function(
    "SET result = lower(['HELLO', 'WORLD'])",
    pipeline().set_cmd(|l| l.named_field("result", call("lower").arg(array().element("HELLO").element("WORLD")))),
    Struct::default().with_str("result", Array::new(STRING).into())
)]
#[case::broadcast_abs_function(
    "SET result = abs([-1, -2, 3])",
    pipeline().set_cmd(|l| l.named_field("result", call("abs").arg(array().element(negate(1)).element(negate(2)).element(3)))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
// Broadcasting with variables
#[case::broadcast_with_variable_array(
    "SET arr = [1, 2, 3] | SET result = arr * 10",
    pipeline()
        .set_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
        .set_cmd(|l| l.named_field("result", multiply(field_ref("arr"), 10))),
    Struct::default()
        .with_str("result", Array::new(INT).into())
        .with_str("arr", Array::new(INT).into())
)]
#[case::broadcast_with_variable_scalar(
    "SET x = 10 | SET result = [1, 2, 3] * x",
    pipeline()
        .set_cmd(|l| l.named_field("x", 10))
        .set_cmd(|l| l.named_field("result", multiply(array().element(1).element(2).element(3), field_ref("x")))),
    Struct::default()
        .with_str("result", Array::new(INT).into())
        .with_str("x", INT)
)]
// Nested array broadcast: [[1, 2] * 3, [3, 4] * 3] + [1]
// Inner broadcasts: [1, 2] * 3 -> array(int), [3, 4] * 3 -> array(int)
// Outer: array(array(int)) + array(int) broadcasts left -> array(array(int))
#[case::nested_array_broadcast(
    "SET result = [[1, 2] * 3, [3, 4] * 3] + [1]",
    pipeline().set_cmd(|l| l.named_field("result", add(
        array()
            .element(multiply(array().element(1).element(2), 3))
            .element(multiply(array().element(3).element(4), 3)),
        array().element(1),
    ))),
    Struct::default().with_str("result", Array::new(Array::new(INT).into()).into())
)]
// Null array broadcasting - null arrays should still type-check correctly
#[case::broadcast_null_array_multiply(
    "SET arr = null AS array(int) | SET result = arr * 10",
    pipeline()
        .set_cmd(|l| l.named_field("arr", cast(null(), Array::new(INT).into())))
        .set_cmd(|l| l.named_field("result", multiply(field_ref("arr"), 10))),
    Struct::default()
        .with_str("result", Array::new(INT).into())
        .with_str("arr", Array::new(INT).into())
)]
fn test_broadcasting(
    #[case] query: &str,
    #[case] expected_ast: PipelineBuilder,
    #[case] expected_schema: Struct,
) -> Result<(), TranslationErrors> {
    let pipeline = Pipeline::parse_result(query)?;
    assert_eq!(pipeline, expected_ast.build());

    let typed = type_check(pipeline).into_result()?;
    assert_eq!(typed.schema(), expected_schema);

    Ok(())
}

// ======================== Non-Broadcasting Tests ========================
// These verify that some operations correctly DON'T broadcast

#[rstest]
// Array operations that work directly on arrays (no broadcasting)
#[case::array_len_no_broadcast(
    "SET arr = [1, 2, 3] | SET result = len(arr)",
    pipeline()
        .set_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
        .set_cmd(|l| l.named_field("result", call("len").arg(field_ref("arr")))),
    Struct::default()
        .with_str("result", INT)
        .with_str("arr", Array::new(INT).into())
)]
#[case::array_concat_no_broadcast(
    "SET result = [1, 2] + [3, 4]",
    pipeline().set_cmd(|l| l.named_field("result", add(array().element(1).element(2), array().element(3).element(4)))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::array_distinct_no_broadcast(
    "SET result = array_distinct([1, 2, 2, 3])",
    pipeline().set_cmd(|l| l.named_field("result", call("array_distinct").arg(array().element(1).element(2).element(2).element(3)))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::array_slice_no_broadcast(
    "SET result = slice([1, 2, 3, 4], 2, 3)",
    pipeline().set_cmd(|l| l.named_field("result", call("slice").arg(array().element(1).element(2).element(3).element(4)).arg(2).arg(3))),
    Struct::default().with_str("result", Array::new(INT).into())
)]
fn test_no_broadcast(
    #[case] query: &str,
    #[case] expected_ast: PipelineBuilder,
    #[case] expected_schema: Struct,
) -> Result<(), TranslationErrors> {
    let pipeline = Pipeline::parse_result(query)?;
    assert_eq!(pipeline, expected_ast.build());

    let typed = type_check(pipeline).into_result()?;
    assert_eq!(typed.schema(), expected_schema);

    Ok(())
}

// ======================== Type Error Tests ========================
// These verify that incompatible array concatenations produce type errors

#[rstest]
#[case::interval_array_plus_timestamp_array("SET result = [1h] + [now()]")]
#[case::empty_interval_array_plus_timestamp_array("SET result = [] AS array(interval) + [now()]")]
#[case::timestamp_array_plus_interval_array("SET result = [now()] + [1h]")]
#[case::int_array_plus_string_array("SET result = [1, 2] + ['a', 'b']")]
#[case::boolean_array_plus_int_array("SET result = [true] + [1]")]
fn test_incompatible_array_concat(#[case] query: &str) {
    let pipeline = Pipeline::parse_result(query).expect("Should parse successfully");

    let result = type_check(pipeline).into_result();
    assert!(
        result.is_err(),
        "Expected type error for incompatible array concat but it succeeded: {}",
        query
    );

    let err = result.unwrap_err();
    eprintln!("Error for '{}': {}", query, err);
}

// ======================== Valid Array Concat Tests ========================

#[rstest]
#[case::interval_plus_interval(
    "SET result = [1h] + [2h]",
    Struct::default().with_str("result", Array::new(INTERVAL).into())
)]
#[case::timestamp_plus_timestamp(
    "SET result = [now()] + [now()]",
    Struct::default().with_str("result", Array::new(TIMESTAMP).into())
)]
#[case::int_plus_int(
    "SET result = [1, 2] + [3, 4]",
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::empty_unknown_plus_int(
    "SET result = [] + [1, 2]",
    Struct::default().with_str("result", Array::new(INT).into())
)]
#[case::empty_typed_plus_same_type(
    "SET result = [] AS array(interval) + [1h]",
    Struct::default().with_str("result", Array::new(INTERVAL).into())
)]
fn test_compatible_array_concat(
    #[case] query: &str,
    #[case] expected_schema: Struct,
) -> Result<(), TranslationErrors> {
    let pipeline = Pipeline::parse_result(query)?;
    let typed = type_check(pipeline).into_result()?;
    assert_eq!(typed.schema(), expected_schema);
    Ok(())
}

// ======================== Nested-Array Broadcast Rejection ========================
// Broadcasting applies exactly one level. On a value of type Array<Array<T>>,
// both field-access (`.field`) and function-call (`op`, unary prefix) broadcasts
// must error at typecheck — a single `.` or single operator is never an implicit
// two-level broadcast. Testing them side-by-side keeps the two flavors aligned.

#[rstest]
// Field-access broadcast on Array<Array<Struct>>: would require a second
// implicit broadcast layer, which is not allowed.
#[case::field_access_nested("ROWS [{outer: [[{name: 'a'}]]}] | SET x = outer.name")]
// Binary function-call broadcast on Array<Array<Int>>: no matching overload
// for `*(array(array(int)), int)`.
#[case::binary_op_nested("SET x = [[1, 2], [3, 4]] * 3")]
// Unary prefix broadcast on Array<Array<Int>>: same reasoning as the binary case.
#[case::unary_op_nested("SET x = -[[1, 2], [3, 4]]")]
fn test_nested_array_broadcast_rejected(#[case] query: &str) {
    let pipeline = Pipeline::parse_result(query).expect("Should parse successfully");

    let result = type_check(pipeline).into_result();
    assert!(
        result.is_err(),
        "Expected type error for nested-array broadcast but it succeeded: {}",
        query
    );

    let err = result.unwrap_err();
    eprintln!("Error for '{}': {}", query, err);
}