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::tuple::Tuple;
use crate::types::UNKNOWN;
use crate::types::{Type, BOOLEAN, DOUBLE, INT, STRING};
use pretty_assertions::assert_eq;
use rstest::rstest;
#[rstest]
#[case::simple_field_access(
"LET user = {name: 'Alice', age: 30} | LET result = user.name",
pipeline()
.let_cmd(|l| l.named_field("user", struct_literal().field("name", "Alice").field("age", 30)))
.let_cmd(|l| l.named_field("result", field(field_ref("user"), "name"))),
Struct::default()
.with_str("result", STRING)
.with_str("user", Struct::default().with_str("name", STRING).with_str("age", INT).into())
)]
#[case::nested_field_access(
"LET user = {address: {city: 'NYC', zip: 10001}} | LET city = user.address.city",
pipeline()
.let_cmd(|l| l.named_field("user", struct_literal().field("address", struct_literal().field("city", "NYC").field("zip", 10001))))
.let_cmd(|l| l.named_field("city", field(field(field_ref("user"), "address"), "city"))),
Struct::default()
.with_str("city", STRING)
.with_str("user", Struct::default().with_str("address", Struct::default().with_str("city", STRING).with_str("zip", INT).into()).into())
)]
#[case::field_access_with_expression(
"LET user = {age: 25} | LET next_age = user.age + 1",
pipeline()
.let_cmd(|l| l.named_field("user", struct_literal().field("age", 25)))
.let_cmd(|l| l.named_field("next_age", add(field(field_ref("user"), "age"), 1))),
Struct::default()
.with_str("next_age", INT)
.with_str("user", Struct::default().with_str("age", INT).into())
)]
#[case::simple_array_index(
"LET arr = [10, 20, 30] | LET first = arr[0]",
pipeline()
.let_cmd(|l| l.named_field("arr", array().element(10).element(20).element(30)))
.let_cmd(|l| l.named_field("first", index(field_ref("arr"), 0))),
Struct::default()
.with_str("first", INT)
.with_str("arr", Array::new(INT).into())
)]
#[case::negative_array_index(
"LET arr = [10, 20, 30] | LET last = arr[-1]",
pipeline()
.let_cmd(|l| l.named_field("arr", array().element(10).element(20).element(30)))
.let_cmd(|l| l.named_field("last", index(field_ref("arr"), negate(1)))),
Struct::default()
.with_str("last", INT)
.with_str("arr", Array::new(INT).into())
)]
#[case::variable_index(
"LET arr = [10, 20, 30] | LET idx = 1 | LET value = arr[idx]",
pipeline()
.let_cmd(|l| l.named_field("arr", array().element(10).element(20).element(30)))
.let_cmd(|l| l.named_field("idx", 1))
.let_cmd(|l| l.named_field("value", index(field_ref("arr"), field_ref("idx")))),
Struct::default()
.with_str("value", INT)
.with_str("idx", INT)
.with_str("arr", Array::new(INT).into())
)]
#[case::index_then_field(
"LET users = [{name: 'Alice'}, {name: 'Bob'}] | LET first_name = users[0].name",
pipeline()
.let_cmd(|l| l.named_field("users", array().element(struct_literal().field("name", "Alice")).element(struct_literal().field("name", "Bob"))))
.let_cmd(|l| l.named_field("first_name", field(index(field_ref("users"), 0), "name"))),
Struct::default()
.with_str("first_name", STRING)
.with_str("users", Array::new(Struct::default().with_str("name", STRING).into()).into())
)]
#[case::function_on_literal(
"LET result = len('hello')",
pipeline().let_cmd(|l| l.named_field("result", call("len").arg("hello"))),
Struct::default()
.with_str("result", INT)
)]
#[case::function_on_field(
"LET user = {name: 'Alice'} | LET name_length = len(user.name)",
pipeline()
.let_cmd(|l| l.named_field("user", struct_literal().field("name", "Alice")))
.let_cmd(|l| l.named_field("name_length", call("len").arg(field(field_ref("user"), "name")))),
Struct::default()
.with_str("name_length", INT)
.with_str("user", Struct::default().with_str("name", STRING).into())
)]
#[case::function_on_array_element(
"LET names = ['Alice', 'Bob'] | LET first_len = len(names[0])",
pipeline()
.let_cmd(|l| l.named_field("names", array().element("Alice").element("Bob")))
.let_cmd(|l| l.named_field("first_len", call("len").arg(index(field_ref("names"), 0)))),
Struct::default()
.with_str("first_len", INT)
.with_str("names", Array::new(STRING).into())
)]
#[case::nested_function_calls(
"LET x = 5 | LET result = abs(abs(x))",
pipeline()
.let_cmd(|l| l.named_field("x", 5))
.let_cmd(|l| l.named_field("result", call("abs").arg(call("abs").arg(field_ref("x"))))),
Struct::default()
.with_str("result", INT)
.with_str("x", INT)
)]
#[case::nested_array_access(
"LET matrix = [[1, 2], [3, 4]] | LET value = matrix[0][1]",
pipeline()
.let_cmd(|l| l.named_field("matrix", array().element(array().element(1).element(2)).element(array().element(3).element(4))))
.let_cmd(|l| l.named_field("value", index(index(field_ref("matrix"), 0), 1))),
Struct::default()
.with_str("value", INT)
.with_str("matrix", Array::new(Array::new(INT).into()).into())
)]
#[case::array_of_structs_access(
"LET users = [{name: 'Alice', scores: [90, 95]}, {name: 'Bob', scores: [85, 88]}] | LET alice_first_score = users[0].scores[0]",
pipeline()
.let_cmd(|l| l.named_field("users",
array()
.element(struct_literal().field("name", "Alice").field("scores", array().element(90).element(95)))
.element(struct_literal().field("name", "Bob").field("scores", array().element(85).element(88)))
))
.let_cmd(|l| l.named_field("alice_first_score", index(field(index(field_ref("users"), 0), "scores"), 0))),
Struct::default()
.with_str("alice_first_score", INT)
.with_str("users", Array::new(Struct::default().with_str("name", STRING).with_str("scores", Array::new(INT).into()).into()).into())
)]
#[case::arithmetic_on_indexed_values(
"LET nums = [10, 20, 30] | LET sum = nums[0] + nums[1]",
pipeline()
.let_cmd(|l| l.named_field("nums", array().element(10).element(20).element(30)))
.let_cmd(|l| l.named_field("sum", add(index(field_ref("nums"), 0), index(field_ref("nums"), 1)))),
Struct::default()
.with_str("sum", INT)
.with_str("nums", Array::new(INT).into())
)]
#[case::comparison_on_field_access(
"LET user = {age: 25} | LET is_adult = user.age >= 18",
pipeline()
.let_cmd(|l| l.named_field("user", struct_literal().field("age", 25)))
.let_cmd(|l| l.named_field("is_adult", gte(field(field_ref("user"), "age"), 18))),
Struct::default()
.with_str("is_adult", BOOLEAN)
.with_str("user", Struct::default().with_str("age", INT).into())
)]
#[case::function_with_multiple_complex_args(
"LET data = {values: [1, 2, 3]} | LET sum_first_two = data.values[0] + data.values[1]",
pipeline()
.let_cmd(|l| l.named_field("data", struct_literal().field("values", array().element(1).element(2).element(3))))
.let_cmd(|l| l.named_field("sum_first_two", add(
index(field(field_ref("data"), "values"), 0),
index(field(field_ref("data"), "values"), 1)
))),
Struct::default()
.with_str("sum_first_two", INT)
.with_str("data", Struct::default().with_str("values", Array::new(INT).into()).into())
)]
#[case::tuple_creation_and_indexing(
"LET t = (42, 'hello', true) | LET first = t.f0 | LET second = t.f1",
pipeline()
.let_cmd(|l| l.named_field("t", tuple().element(42).element("hello").element(true)))
.let_cmd(|l| l.named_field("first", field(field_ref("t"), "f0")))
.let_cmd(|l| l.named_field("second", field(field_ref("t"), "f1"))),
Struct::default()
.with_str("second", STRING)
.with_str("first", INT)
.with_str("t", Tuple::default().with(INT).with(STRING).with(BOOLEAN).into())
)]
#[case::empty_array_operations(
"LET empty = [] | LET has_items = len(empty) > 0",
pipeline()
.let_cmd(|l| l.named_field("empty", array()))
.let_cmd(|l| l.named_field("has_items", gt(call("len").arg(field_ref("empty")), 0))),
Struct::default()
.with_str("has_items", BOOLEAN)
.with_str("empty", Array::new(UNKNOWN).into())
)]
#[case::mixed_numeric_types(
"LET int_val = 10 | LET double_val = 3.14 AS double | LET result = int_val + double_val",
pipeline()
.let_cmd(|l| l.named_field("int_val", 10))
.let_cmd(|l| l.named_field("double_val", cast(decimal("3.14").unwrap(), Type::Double)))
.let_cmd(|l| l.named_field("result", add(field_ref("int_val"), field_ref("double_val")))),
Struct::default()
.with_str("result", DOUBLE)
.with_str("double_val", DOUBLE)
.with_str("int_val", INT)
)]
#[case::deeply_nested_struct(
"LET data = {level1: {level2: {level3: {value: 99}}}} | LET deep = data.level1.level2.level3.value",
pipeline()
.let_cmd(|l| l.named_field("data",
struct_literal().field("level1",
struct_literal().field("level2",
struct_literal().field("level3",
struct_literal().field("value", 99)
)
)
)
))
.let_cmd(|l| l.named_field("deep",
field(
field(
field(
field(field_ref("data"), "level1"),
"level2"
),
"level3"
),
"value"
)
)),
Struct::default()
.with_str("deep", INT)
.with_str("data",
Struct::default().with_str("level1",
Struct::default().with_str("level2",
Struct::default().with_str("level3",
Struct::default().with_str("value", INT).into()
).into()
).into()
).into()
)
)]
fn test_complex_expressions_in_pipelines(
#[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()?;
let actual_schema = typed.schema();
assert_eq!(actual_schema, expected_schema);
Ok(())
}