use std::sync::Arc;
use crate::err::TranslationErrors;
use crate::tree::ast::expression::Expression;
use crate::tree::ast::pipeline::Pipeline;
use crate::tree::ast::query::Query;
use crate::tree::ast::ParseWithErrors;
use crate::tree::options::ExpressionTypeCheckOptions;
use crate::tree::tests::shared::catalog_with_tables;
use crate::type_check;
use crate::type_check_expression;
use crate::type_check_with_provider;
use crate::types::{Type, BOOLEAN, DOUBLE, INT, STRING};
use rstest::rstest;
#[rstest]
#[case::select_nonexistent_field(
vec![("name", STRING), ("age", INT)],
"SELECT email",
1
)]
#[case::select_multiple_nonexistent_fields(
vec![("name", STRING)],
"SELECT age, email, phone",
1
)]
#[case::select_mix_valid_invalid(
vec![("name", STRING), ("age", INT)],
"SELECT name, email",
1
)]
#[case::where_nonexistent_field(
vec![("name", STRING), ("age", INT)],
"WHERE email == 'test@example.com'",
1
)]
#[case::where_nonexistent_in_and(
vec![("name", STRING)],
"WHERE name == 'Alice' AND age > 18",
1
)]
#[case::where_both_sides_nonexistent(
vec![("name", STRING)],
"WHERE email == phone",
1
)]
#[case::set_nonexistent_source(
vec![("name", STRING)],
"SET full_name = email",
1
)]
#[case::set_nonexistent_in_expression(
vec![("age", INT)],
"SET result = age + missing_field",
1
)]
#[case::set_multiple_missing(
vec![("x", INT)],
"SET result = x + y * z",
1
)]
#[case::agg_nonexistent_field(
vec![("amount", INT)],
"AGG total = sum(missing_field)",
1
)]
#[case::agg_nonexistent_group_by(
vec![("amount", INT)],
"AGG total = sum(amount) BY region",
1
)]
#[case::agg_both_missing(
vec![("price", INT)],
"AGG total = sum(amount) BY region",
1
)]
#[case::agg_bare_column_with_by(
vec![("amount", DOUBLE), ("name", STRING)],
"AGG amount BY name",
1
)]
#[case::agg_bare_column_without_by(
vec![("amount", DOUBLE)],
"AGG amount",
1
)]
#[case::sort_nonexistent_field(
vec![("name", STRING)],
"SORT age DESC",
1
)]
#[case::sort_multiple_nonexistent(
vec![("name", STRING)],
"SORT age DESC, email ASC",
1
)]
#[case::sort_mix_valid_invalid(
vec![("name", STRING), ("age", INT)],
"SORT name ASC, email DESC",
1
)]
#[case::drop_nonexistent_field(
vec![("name", STRING), ("age", INT)],
"DROP email",
1
)]
#[case::drop_multiple_nonexistent(
vec![("name", STRING)],
"DROP age, email",
1
)]
#[case::drop_mix_valid_invalid(
vec![("name", STRING), ("age", INT)],
"DROP name, email",
1
)]
#[case::undefined_function(
vec![("name", STRING)],
"SELECT result = undefined_function(name)",
1
)]
#[case::function_with_nonexistent_field(
vec![("name", STRING)],
"SELECT result = len(email)",
1
)]
#[case::dropped_field_in_select(
vec![("name", STRING), ("age", INT), ("email", STRING)],
"DROP email | SELECT name, email",
1
)]
#[case::dropped_field_in_where(
vec![("name", STRING), ("age", INT), ("active", BOOLEAN)],
"DROP active | WHERE active == true",
1
)]
#[case::non_grouped_field_after_agg(
vec![("region", STRING), ("category", STRING), ("amount", INT)],
"AGG total = sum(amount) BY region | SELECT region, category",
1
)]
#[case::original_field_after_agg(
vec![("region", STRING), ("amount", INT)],
"AGG total = sum(amount) BY region | SELECT amount",
1
)]
fn test_error_cases_with_set(
#[case] input_fields: Vec<(&str, Type)>,
#[case] command: &str,
#[case] min_errors: usize,
) -> Result<(), TranslationErrors> {
let set_clause = input_fields
.iter()
.map(|(name, typ)| {
let value = match typ {
t if t == &STRING => "'dummy'",
t if t == &INT => "0",
t if t == &DOUBLE => "0.0",
t if t == &BOOLEAN => "true",
_ => "'unknown'",
};
format!("{} = {}", name, value)
})
.collect::<Vec<_>>()
.join(", ");
let query = format!("SET {} | {}", set_clause, command);
let pipeline = Pipeline::parse_result(&query)?;
let errors = type_check(pipeline).errors;
assert!(errors.len() >= min_errors);
Ok(())
}
#[rstest]
#[case::from_nonexistent_table(
vec![],
"FROM missing_table",
1
)]
#[case::from_empty_catalog(
vec![],
"FROM users",
1
)]
#[case::access_replaced_field_after_struct_literal(
vec![("events", vec![("id", INT)])],
"FROM events | SET x.a = 1 | SET x = {y: 2} | SET z = x.a",
1
)]
#[case::access_deep_replaced_field_after_struct_literal(
vec![("events", vec![("id", INT)])],
"FROM events | SET x.a.b = 1 | SET x = {y: 2} | SELECT x.a.b",
1
)]
#[case::access_replaced_field_in_where(
vec![("events", vec![("id", INT)])],
"FROM events | SET x.a = 1 | SET x = {y: 2} | WHERE x.a > 0",
1
)]
#[case::error_in_first_command(
vec![("users", vec![("name", STRING), ("age", INT)])],
"FROM users | SELECT email",
1
)]
#[case::error_in_middle_command(
vec![("users", vec![("name", STRING), ("age", INT)])],
"FROM users | WHERE email == 'test' | SELECT name",
1
)]
#[case::error_in_last_command(
vec![("users", vec![("name", STRING), ("age", INT)])],
"FROM users | WHERE age > 18 | SELECT name, email",
1
)]
#[case::multiple_errors_in_pipeline(
vec![("users", vec![("name", STRING)])],
"FROM users | WHERE age > 18 | SELECT name, email",
2
)]
fn test_error_cases_with_from(
#[case] tables: Vec<(&str, Vec<(&str, Type)>)>,
#[case] query: &str,
#[case] min_errors: usize,
) -> Result<(), TranslationErrors> {
let pipeline = Pipeline::parse_result(query)?;
let catalog = catalog_with_tables(tables);
let errors = type_check_with_provider(pipeline, catalog).errors;
assert!(errors.len() >= min_errors);
Ok(())
}
#[rstest]
#[case::string_add_boolean("'hello' + true", "could not find a matching function")]
#[case::string_multiply_int("'hello' * 5", "could not find a matching function")]
#[case::boolean_add_int("true + 5", "could not find a matching function")]
#[case::string_subtract("'hello' - 'world'", "could not find a matching function")]
#[case::len_int("len(42)", "matching function definition")]
#[case::len_boolean("len(true)", "matching function definition")]
#[case::upper_int("upper(42)", "matching function definition")]
#[case::abs_string("abs('hello')", "matching function definition")]
#[case::round_string("round('3.14')", "matching function definition")]
#[case::unknown_function("unknown_func(42)", "no function by this name")]
#[case::undefined_function("not_a_function('test')", "no function by this name")]
#[case::mixed_array_int_string("[1, 'two', 3]", "could not determine array element type")]
#[case::mixed_array_bool_int("[true, 1, false]", "could not determine array element type")]
#[case::flatten_non_array_of_array("flatten([1, 2, 3])", "matching function definition")]
#[case::flatten_non_array("flatten('not an array')", "matching function definition")]
#[case::in_array_type_mismatch("'hello' IN [1, 2]", "type")]
#[case::not_in_array_type_mismatch("'hello' NOT IN [1, 2]", "type")]
#[case::in_tuple_type_mismatch("'hello' IN (1, 2)", "type")]
#[case::in_tuple_mixed_types("'hello' IN (1, 'foo')", "type")]
#[case::in_range_type_mismatch("'hello' IN 1..2", "type")]
#[case::not_in_range_type_mismatch("'hello' NOT IN 1..2", "type")]
#[case::range_int_string("1..'hello'", "matching function definition")]
#[case::range_int_double("1..1.5", "matching function definition")]
#[case::range_int_interval("1..-1h", "matching function definition")]
#[case::range_inclusive_int_string("1..='hello'", "matching function definition")]
#[case::range_inclusive_int_interval("1..=-1h", "matching function definition")]
#[case::if_non_bool_condition("if('hello', 0, 0)", "matching function definition")]
#[case::if_non_bool_condition_two_args("if('hello', 0)", "matching function definition")]
#[case::if_branch_type_mismatch("if(true, 'hello', 1)", "matching function definition")]
#[case::case_branch_type_mismatch("case(true: 'hello', false: 0)", "type")]
#[case::case_non_bool_condition(
"case('hello': 'hello', false: 'hello')",
"matching function definition"
)]
#[case::case_all_non_bool(
"case('hello': 'hello', 'hello': 'hello')",
"matching function definition"
)]
#[case::coalesce_type_mismatch("coalesce(1, 'hello')", "type")]
#[case::cidr_contains_int_first_arg(
"cidr_contains(123, '192.168.1.100')",
"matching function definition"
)]
#[case::cidr_contains_int_second_arg(
"cidr_contains('192.168.1.0/24', 456)",
"matching function definition"
)]
#[case::is_ipv4_int_arg("is_ipv4(123)", "matching function definition")]
#[case::is_ipv6_int_arg("is_ipv6(456)", "matching function definition")]
fn test_type_errors(#[case] input: &str, #[case] expected_error: &str) {
let result = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result();
assert!(
result.is_err(),
"Expected error but got success for: {}",
input
);
let err = result.unwrap_err();
let err_string = err.to_string();
eprintln!("Error for '{}': {}", input, err_string);
assert!(
err_string.contains(expected_error),
"Error '{}' does not contain expected fragment '{}'",
err_string,
expected_error
);
}
#[rstest]
#[case::select_duplicate_same_type(
vec![("x", INT), ("y", INT)],
"SELECT a = x, a = y",
1
)]
#[case::select_duplicate_diff_type(
vec![("x", INT), ("y", STRING)],
"SELECT a = x, a = y",
1
)]
#[case::select_nested_struct_conflict(
vec![("x", INT), ("y", STRING)],
"SET result = {x: x, y: y} | SELECT result, result.x = 'conflict'",
1
)]
#[case::set_duplicate_same_type(
vec![("x", INT), ("y", INT)],
"SET a = x, a = y",
1
)]
#[case::set_duplicate_diff_type(
vec![("x", INT), ("y", STRING)],
"SET a = x, a = y",
1
)]
#[case::window_duplicate_same_type(
vec![("x", INT)],
"WINDOW a = count(), a = sum(x)",
1
)]
#[case::agg_duplicate_same_type(
vec![("x", INT), ("y", STRING)],
"AGG a = count(), a = sum(x) BY y",
1
)]
#[case::agg_duplicate_groupby_conflict(
vec![("x", INT), ("y", STRING)],
"AGG x = count() BY x",
1
)]
#[case::set_prefix_scalar_then_child(
vec![("x", INT), ("y", INT)],
"SET a = x, a.b = y",
1
)]
#[case::set_prefix_child_then_scalar(
vec![("x", INT), ("y", INT)],
"SET a.b = x, a = y",
1
)]
#[case::set_prefix_deep_nesting(
vec![("x", INT), ("y", INT)],
"SET a.b = x, a.b.c = y",
1
)]
#[case::window_prefix_clash(
vec![("x", INT)],
"WINDOW a = count(), a.b = sum(x)",
1
)]
#[case::select_prefix_scalar_then_child(
vec![("x", INT), ("y", INT)],
"SELECT a = x, a.b = y",
1
)]
#[case::parse_duplicate_fields(
vec![("msg", STRING)],
"PARSE msg '*-*' a, a",
1
)]
#[case::struct_literal_duplicate_field(
vec![("x", INT), ("y", INT)],
"SET s = {a: x, a: y}",
1
)]
fn test_duplicate_field_errors(
#[case] input_fields: Vec<(&str, Type)>,
#[case] pipeline: &str,
#[case] min_errors: usize,
) -> Result<(), TranslationErrors> {
let set_clause = input_fields
.iter()
.map(|(name, typ)| {
let value = match typ {
t if t == &STRING => "'dummy'",
t if t == &INT => "0",
t if t == &DOUBLE => "0.0",
t if t == &BOOLEAN => "true",
_ => "'unknown'",
};
format!("{} = {}", name, value)
})
.collect::<Vec<_>>()
.join(", ");
let query = format!("SET {} | {}", set_clause, pipeline);
let parsed = Pipeline::parse_result(&query)?;
let errors = type_check(parsed).errors;
assert!(
errors.len() >= min_errors,
"Expected at least {} errors, got {}",
min_errors,
errors.len()
);
Ok(())
}
#[rstest]
#[case::append_then_select(
vec![("events", vec![("id", INT), ("name", STRING)])],
"FROM events | APPEND output_table | SELECT id",
"APPEND command must be the final command"
)]
#[case::append_then_where(
vec![("events", vec![("id", INT), ("name", STRING)])],
"FROM events | APPEND output_table | WHERE id > 0",
"APPEND command must be the final command"
)]
#[case::append_in_middle(
vec![("events", vec![("id", INT), ("name", STRING)])],
"FROM events | APPEND output_table | SET x = 1 | SELECT id",
"APPEND command must be the final command"
)]
fn test_misplaced_side_effect_commands(
#[case] tables: Vec<(&str, Vec<(&str, Type)>)>,
#[case] query: &str,
#[case] expected_error: &str,
) -> Result<(), TranslationErrors> {
let pipeline = Pipeline::parse_result(query)?;
let catalog = catalog_with_tables(tables);
let errors = type_check_with_provider::<Pipeline>(Arc::new(pipeline), catalog).errors;
eprintln!("Errors: {}", errors);
assert!(
!errors.is_empty(),
"Expected at least one error for misplaced side-effect command"
);
let error_string = errors.to_string();
assert!(
error_string.contains(expected_error),
"Error '{}' does not contain expected fragment '{}'",
error_string,
expected_error
);
Ok(())
}
#[test]
fn agg_allows_scalar_def_outside_aggregate() {
let input = "DEF multiplier = 2;\nSET amount = 1, region = 'r'\n| AGG total = sum(amount) * multiplier BY region";
let (query, parse_errors) = Query::parse_with_errors(input);
assert!(parse_errors.is_empty(), "{parse_errors:#?}");
let result = type_check::<Query>(Arc::new(query));
assert!(result.errors.is_empty(), "{}", result.errors);
}
#[test]
fn agg_allows_lambda_param_in_body_under_aggregate() {
let input = "SET items = [1, 2], grp = 1\n| AGG result = array_agg(transform(items, x -> x + 1)) BY grp";
let (query, parse_errors) = Query::parse_with_errors(input);
assert!(parse_errors.is_empty(), "{parse_errors:#?}");
let result = type_check::<Query>(Arc::new(query));
assert!(result.errors.is_empty(), "{}", result.errors);
}
#[test]
fn agg_allows_outer_lambda_param_in_nested_lambda_body_under_aggregate() {
let input = "SET items = [1, 2], grp = 1\n| AGG result = array_agg(transform(items, outer -> transform([outer], inner -> outer + inner))) BY grp";
let (query, parse_errors) = Query::parse_with_errors(input);
assert!(parse_errors.is_empty(), "{parse_errors:#?}");
let result = type_check::<Query>(Arc::new(query));
assert!(result.errors.is_empty(), "{}", result.errors);
}
#[rstest]
#[case::in_empty_tuple("'hello' IN ()")]
fn test_expression_parse_errors(#[case] input: &str) {
assert!(
Expression::parse_result(input).is_err(),
"Expected parse error for: {}",
input
);
}
#[rstest]
#[case::garbage_input("!@#$%")]
#[case::empty_input("")]
fn test_parse_error_query_produces_error_pipeline(#[case] input: &str) {
let (query, _parse_errors) = Query::parse_with_errors(input);
let typed = type_check::<Query>(Arc::new(query)).output;
assert!(
typed.pipeline.kind.is_err(),
"Expected pipeline.kind to be Err for parse-error query '{}', but got Ok",
input
);
}