hamelin_lib 0.7.13

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, STRING};
use pretty_assertions::assert_eq;
use rstest::rstest;

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

#[rstest]
// Arithmetic broadcasting - array on left
#[case::broadcast_multiply_array_left(
    "LET result = [1, 2, 3] * 10",
    pipeline().let_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(
    "LET result = [1, 2, 3] + 10",
    pipeline().let_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(
    "LET result = [10, 20, 30] - 5",
    pipeline().let_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(
    "LET result = [10, 20, 30] / 2",
    pipeline().let_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(
    "LET result = [10, 21, 32] % 3",
    pipeline().let_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(
    "LET result = 10 * [1, 2, 3]",
    pipeline().let_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(
    "LET result = 10 + [1, 2, 3]",
    pipeline().let_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(
    "LET result = -[1, 2, 3]",
    pipeline().let_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(
    "LET result = NOT [true, false]",
    pipeline().let_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(
    "LET result = [1, 5, 10] > 3",
    pipeline().let_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(
    "LET result = [1, 5, 10] < 3",
    pipeline().let_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(
    "LET result = [1, 2, 3] == 2",
    pipeline().let_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(
    "LET result = [1, 2, 3] != 2",
    pipeline().let_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(
    "LET result = [1, 5, 10] >= 5",
    pipeline().let_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(
    "LET result = [1, 5, 10] <= 5",
    pipeline().let_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(
    "LET result = ['hello', 'world'] + '!'",
    pipeline().let_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(
    "LET result = upper(['hello', 'world'])",
    pipeline().let_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(
    "LET result = lower(['HELLO', 'WORLD'])",
    pipeline().let_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(
    "LET result = abs([-1, -2, 3])",
    pipeline().let_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(
    "LET arr = [1, 2, 3] | LET result = arr * 10",
    pipeline()
        .let_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
        .let_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(
    "LET x = 10 | LET result = [1, 2, 3] * x",
    pipeline()
        .let_cmd(|l| l.named_field("x", 10))
        .let_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(
    "LET result = [[1, 2] * 3, [3, 4] * 3] + [1]",
    pipeline().let_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(
    "LET arr = null AS array(int) | LET result = arr * 10",
    pipeline()
        .let_cmd(|l| l.named_field("arr", cast(null(), Array::new(INT).into())))
        .let_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(
    "LET arr = [1, 2, 3] | LET result = len(arr)",
    pipeline()
        .let_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
        .let_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(
    "LET result = [1, 2] + [3, 4]",
    pipeline().let_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(
    "LET result = array_distinct([1, 2, 2, 3])",
    pipeline().let_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(
    "LET result = slice([1, 2, 3, 4], 2, 3)",
    pipeline().let_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(())
}