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;
#[rstest]
#[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())
)]
#[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())
)]
#[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())
)]
#[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())
)]
#[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())
)]
#[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())
)]
#[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)
)]
#[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())
)]
#[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(())
}
#[rstest]
#[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(())
}