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;
#[rstest]
#[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())
)]
#[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())
)]
#[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())
)]
#[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())
)]
#[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())
)]
#[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())
)]
#[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)
)]
#[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())
)]
#[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(())
}
#[rstest]
#[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(())
}
#[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);
}
#[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(())
}
#[rstest]
#[case::field_access_nested("ROWS [{outer: [[{name: 'a'}]]}] | SET x = outer.name")]
#[case::binary_op_nested("SET x = [[1, 2], [3, 4]] * 3")]
#[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);
}