use crate::err::TranslationErrors;
use crate::tree::ast::expression::Expression;
use crate::tree::ast::ParseWithErrors;
use crate::tree::builder::*;
use crate::tree::options::ExpressionTypeCheckOptions;
use crate::type_check_expression;
use crate::types::array::Array;
use crate::types::range::Range;
use crate::types::struct_type::Struct;
use crate::types::tuple::Tuple;
use crate::types::{
Type, BOOLEAN, CALENDAR_INTERVAL, DOUBLE, INT, INTERVAL, STRING, TIMESTAMP, UNKNOWN,
};
use pretty_assertions::assert_eq;
use rstest::rstest;
#[rstest]
#[case::int_zero("0", int(0), INT)]
#[case::int_positive("42", int(42), INT)]
#[case::int_large("9999999", int(9999999), INT)]
#[case::bool_true("true", boolean(true), BOOLEAN)]
#[case::bool_false("false", boolean(false), BOOLEAN)]
#[case::str_single_quoted("'hello'", string("hello"), STRING)]
#[case::str_double_quoted(r#""world""#, string("world"), STRING)]
#[case::str_escaped_quote(r#"'can\'t'"#, string("can't"), STRING)]
#[case::double_positive("3.14 AS double", cast(double(3.14), Type::Double), DOUBLE)]
#[case::double_negative("-2.5 AS double", cast(negate(double(2.5)), Type::Double), DOUBLE)]
#[case::double_zero("0.0 AS double", cast(double(0.0), Type::Double), DOUBLE)]
#[case::null_literal("null", null(), UNKNOWN)]
#[case::op_add("1 + 2", add(1, 2), INT)]
#[case::op_subtract("5 - 3", subtract(5, 3), INT)]
#[case::op_multiply("4 * 2", multiply(4, 2), INT)]
#[case::op_divide("10 / 2", divide(10, 2), INT)]
#[case::op_modulo("7 % 3", modulo(7, 3), INT)]
#[case::op_add_double(
"(1.0 AS double) + (2.0 AS double)",
add(cast(double(1.0), Type::Double), cast(double(2.0), Type::Double)),
DOUBLE
)]
#[case::op_subtract_double(
"(10.5 AS double) - (2.5 AS double)",
subtract(cast(double(10.5), Type::Double), cast(double(2.5), Type::Double)),
DOUBLE
)]
#[case::type_promotion_int_double(
"42 + (3.14 AS double)",
add(42, cast(double(3.14), Type::Double)),
DOUBLE
)]
#[case::cmp_equal("1 == 2", eq(1, 2), BOOLEAN)]
#[case::cmp_not_equal("1 != 2", ne(1, 2), BOOLEAN)]
#[case::cmp_less_than("1 < 2", lt(1, 2), BOOLEAN)]
#[case::cmp_less_than_or_equal("1 <= 2", lte(1, 2), BOOLEAN)]
#[case::cmp_greater_than("1 > 2", gt(1, 2), BOOLEAN)]
#[case::cmp_greater_than_or_equal("1 >= 2", gte(1, 2), BOOLEAN)]
#[case::cmp_string_equal("'a' == 'b'", eq("a", "b"), BOOLEAN)]
#[case::cmp_string_not_equal("'a' != 'b'", ne("a", "b"), BOOLEAN)]
#[case::logic_and("true AND false", and(true, false), BOOLEAN)]
#[case::logic_or("true OR false", or(true, false), BOOLEAN)]
#[case::unary_negate_int("-42", negate(42), INT)]
#[case::unary_negate_double("-(3.14 AS double)", negate(cast(double(3.14), Type::Double)), DOUBLE)]
#[case::unary_not_true("NOT true", not(true), BOOLEAN)]
#[case::unary_not_false("NOT false", not(false), BOOLEAN)]
#[case::precedence_mult_before_add("1 + 2 * 3", add(1, multiply(2, 3)), INT)]
#[case::precedence_parens_override("(1 + 2) * 3", multiply(add(1, 2), 3), INT)]
#[case::precedence_complex(
"1 + 2 * 3 - 4 / 2",
subtract(add(1, multiply(2, 3)), divide(4, 2)),
INT
)]
#[case::paren_int("(42)", int(42), INT)]
#[case::paren_add("(1 + 2)", add(1, 2), INT)]
#[case::cast_to_string("42 AS string", cast(42, Type::String), STRING)]
#[case::cast_to_boolean("1 AS boolean", cast(1, Type::Boolean), BOOLEAN)]
#[case::cast_to_timestamp(
"'2024-01-01' AS timestamp",
cast("2024-01-01", Type::Timestamp),
Type::Timestamp
)]
#[case::interval_seconds("60sec", seconds(60), INTERVAL)]
#[case::interval_minutes("30min", minutes(30), INTERVAL)]
#[case::interval_hours("12hr", hours(12), INTERVAL)]
#[case::interval_days("7d", days(7), INTERVAL)]
#[case::interval_weeks("4w", weeks(4), INTERVAL)]
#[case::interval_months("6mon", months(6), CALENDAR_INTERVAL)]
#[case::interval_quarters("2q", quarters(2), CALENDAR_INTERVAL)]
#[case::interval_years("1y", years(1), CALENDAR_INTERVAL)]
#[case::interval_to_string("32m AS string", cast(minutes(32), Type::String), STRING)]
#[case::calendar_interval_to_string("6mon AS string", cast(months(6), Type::String), STRING)]
#[case::cast_int_to_int("42 AS int", cast(42, Type::Int), INT)]
#[case::cast_string_to_string("'hello' AS string", cast("hello", Type::String), STRING)]
#[case::cast_boolean_to_boolean("true AS boolean", cast(true, Type::Boolean), BOOLEAN)]
#[case::cast_double_to_double(
"(3.14 AS double) AS double",
cast(cast(double(3.14), Type::Double), Type::Double),
DOUBLE
)]
#[case::cast_int_to_double("42 AS double", cast(42, Type::Double), DOUBLE)]
#[case::cast_double_to_int(
"(3.14 AS double) AS int",
cast(cast(double(3.14), Type::Double), Type::Int),
INT
)]
#[case::cast_double_to_int_bare("3.14 AS int", cast(double(3.14), Type::Int), INT)]
#[case::cast_double_to_double_bare("3.14 AS double", cast(double(3.14), Type::Double), DOUBLE)]
#[case::cast_int_to_string_explicit("42 AS string", cast(42, Type::String), STRING)]
#[case::cast_double_to_string(
"(3.14 AS double) AS string",
cast(cast(double(3.14), Type::Double), Type::String),
STRING
)]
#[case::cast_boolean_to_string("true AS string", cast(true, Type::String), STRING)]
#[case::cast_timestamp_to_string(
"('2024-01-01' AS timestamp) AS string",
cast(cast("2024-01-01", Type::Timestamp), Type::String),
STRING
)]
#[case::cast_string_to_int("'42' AS int", cast("42", Type::Int), INT)]
#[case::cast_string_to_double("'3.14' AS double", cast("3.14", Type::Double), DOUBLE)]
#[case::cast_string_to_boolean_true("'true' AS boolean", cast("true", Type::Boolean), BOOLEAN)]
#[case::cast_string_to_timestamp_explicit(
"'2024-01-01' AS timestamp",
cast("2024-01-01", Type::Timestamp),
Type::Timestamp
)]
#[case::cast_int_to_boolean_explicit("1 AS boolean", cast(1, Type::Boolean), BOOLEAN)]
#[case::cast_int_zero_to_boolean("0 AS boolean", cast(0, Type::Boolean), BOOLEAN)]
#[case::cast_boolean_to_int("true AS int", cast(true, Type::Int), INT)]
#[case::cast_string_to_boolean_false("'false' AS boolean", cast("false", Type::Boolean), BOOLEAN)]
fn test_expressions(
#[case] input: &str,
#[case] expected: impl ExpressionBuilder,
#[case] expected_type: Type,
) -> Result<(), TranslationErrors> {
let actual = Expression::parse_result(input)?;
let expected_expr = expected.build();
assert_eq!(actual, expected_expr);
let typed = type_check_expression(actual, ExpressionTypeCheckOptions::builder().build())
.into_result()?;
assert_eq!(*typed.resolved_type, expected_type);
Ok(())
}
#[rstest]
#[case::cast_int_to_decimal("42 AS decimal")]
#[case::cast_double_to_decimal("(3.14 AS double) AS decimal")]
#[case::cast_bare_double_to_decimal("3.14 AS decimal")]
fn test_decimal_casts(#[case] input: &str) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
match &*typed.resolved_type {
Type::Decimal(_) => {} other => panic!("Expected Decimal type, got {:?}", other),
}
}
#[rstest]
#[case::numeric_inclusive_range("1..=5", Type::RangeInclusive(Range::new(INT)))]
#[case::interval_inclusive_range("-1h..=2h", Type::RangeInclusive(Range::new(INTERVAL)))]
#[case::prefix_inclusive_range("..=5", Type::RangeInclusive(Range::new(INT)))]
#[case::prefix_inclusive_interval("..=2h", Type::RangeInclusive(Range::new(INTERVAL)))]
#[case::exclusive_numeric_range("1..5", Range::new(INT).into())]
fn test_range_type_resolution(#[case] input: &str, #[case] expected_type: Type) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
assert_eq!(*typed.resolved_type, expected_type);
}
#[rstest]
#[case::exclusive_range("1..5", false)]
#[case::inclusive_range("1..=5", true)]
#[case::exclusive_prefix("..5", false)]
#[case::inclusive_prefix("..=5", true)]
fn test_range_type_variant(#[case] input: &str, #[case] expected_inclusive: bool) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
if expected_inclusive {
assert!(
matches!(&*typed.resolved_type, Type::RangeInclusive(_)),
"Expected RangeInclusive for '{input}', got {:?}",
typed.resolved_type
);
} else {
assert!(
matches!(&*typed.resolved_type, Type::Range(_)),
"Expected Range for '{input}', got {:?}",
typed.resolved_type
);
}
}
#[rstest]
#[case::int_to_timestamp("42 AS timestamp", "Cannot cast")]
#[case::timestamp_to_int("('2024-01-01' AS timestamp) AS int", "Cannot cast")]
#[case::array_to_string("[1, 2, 3] AS string", "Cannot cast")]
fn test_invalid_casts(#[case] input: &str, #[case] expected_error_fragment: &str) {
let expr =
Expression::parse_result(input).expect("Cast should parse successfully even if invalid");
let result =
type_check_expression(expr, ExpressionTypeCheckOptions::builder().build()).into_result();
assert!(
result.is_err(),
"Expected cast to fail type checking but it succeeded: {}",
input
);
let err = result.unwrap_err();
let err_string = err.to_string();
eprintln!("Error for '{}': {}", input, err_string);
assert!(
err_string.contains(expected_error_fragment),
"Error message '{}' does not contain expected fragment '{}'",
err_string,
expected_error_fragment
);
}
#[rstest]
#[case::len_string("len('hello')", INT)]
#[case::upper_string("upper('test')", STRING)]
#[case::lower_string("lower('TEST')", STRING)]
#[case::contains_string("contains('hello', 'ell')", BOOLEAN)]
#[case::starts_with_string("starts_with('hello', 'hel')", BOOLEAN)]
#[case::ends_with_string("ends_with('hello', 'lo')", BOOLEAN)]
#[case::abs_int("abs(-42)", INT)]
#[case::abs_double("abs(-3.14 AS double)", DOUBLE)]
#[case::round_double("round(3.14 AS double)", DOUBLE)]
#[case::ceil_double("ceil(3.14 AS double)", DOUBLE)]
#[case::floor_double("floor(3.14 AS double)", DOUBLE)]
#[case::sign_int("sign(-42)", INT)]
#[case::sign_double("sign(-3.14 AS double)", DOUBLE)]
#[case::len_array("len([1, 2, 3])", INT)]
#[case::array_distinct("array_distinct([1, 2, 2, 3])", Array::new(INT).into())]
#[case::len_string_polymorphic("len('abc')", INT)]
#[case::len_array_polymorphic("len([1, 2, 3, 4])", INT)]
fn test_function_return_types(#[case] input: &str, #[case] expected_type: Type) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
assert_eq!(*typed.resolved_type, expected_type);
}
#[rstest]
#[case::array_of_ints("[1, 2, 3]", Array::new(INT).into())]
#[case::array_of_strings("['a', 'b', 'c']", Array::new(STRING).into())]
#[case::array_of_booleans("[true, false, true]", Array::new(BOOLEAN).into())]
#[case::array_of_doubles("[(1.0 AS double), (2.0 AS double)]", Array::new(DOUBLE).into())]
#[case::empty_array("[]", Array::new(UNKNOWN).into())]
#[case::array_with_arithmetic("[1 + 1, 2 * 3, 4 - 1]", Array::new(INT).into())]
#[case::array_with_comparisons("[1 == 2, true, false]", Array::new(BOOLEAN).into())]
#[case::array_of_arrays(
"[[1, 2], [3, 4], [5, 6]]",
Array::new(Array::new(INT).into()).into()
)]
#[case::array_of_empty_arrays(
"[[], [], []]",
Array::new(Array::new(UNKNOWN).into()).into()
)]
#[case::tuple_int_string("(1, 'hello')", Tuple::new(vec![INT, STRING]).into())]
#[case::tuple_bool_int_string("(true, 42, 'world')", Tuple::new(vec![BOOLEAN, INT, STRING]).into())]
#[case::tuple_nested("(1, (2, 'inner'))", Tuple::new(vec![INT, Tuple::new(vec![INT, STRING]).into()]).into())]
#[case::struct_simple("{x: 1, y: 2}", Struct::default().with_str("x", INT).with_str("y", INT).into())]
#[case::struct_mixed("{name: 'Alice', age: 30}", Struct::default().with_str("name", STRING).with_str("age", INT).into())]
#[case::struct_nested("{point: {x: 1, y: 2}}", Struct::default().with_str("point", Struct::default().with_str("x", INT).with_str("y", INT).into()).into())]
fn test_collection_types(#[case] input: &str, #[case] expected_type: Type) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
assert_eq!(*typed.resolved_type, expected_type);
}
#[rstest]
#[case::array_index_int("[1, 2, 3][0]", INT)]
#[case::array_index_string("['a', 'b', 'c'][1]", STRING)]
#[case::map_lookup_int("map([('key1', 1), ('key2', 2)])['key1']", INT)]
#[case::map_lookup_string("map([('x', 'hello'), ('y', 'world')])['x']", STRING)]
#[case::variant_index_zero("(1 AS variant)[0]", Type::Variant)]
#[case::variant_index_one("('hello' AS variant)[1]", Type::Variant)]
fn test_index_access_types(#[case] input: &str, #[case] expected_type: Type) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
assert_eq!(*typed.resolved_type, expected_type);
}
#[rstest]
#[case::int_to_variant("42 AS variant", Type::Variant)]
#[case::double_to_variant("(3.14 AS double) AS variant", Type::Variant)]
#[case::string_to_variant("'hello' AS variant", Type::Variant)]
#[case::boolean_to_variant("true AS variant", Type::Variant)]
#[case::array_int_to_variant("[1, 2, 3] AS variant", Type::Variant)]
#[case::array_string_to_variant("['a', 'b'] AS variant", Type::Variant)]
#[case::array_array_to_variant("[[1, 2], [3, 4]] AS variant", Type::Variant)]
#[case::struct_to_variant("{x: 1, y: 'hello'} AS variant", Type::Variant)]
#[case::struct_nested_to_variant("{outer: {inner: 42}} AS variant", Type::Variant)]
#[case::map_to_variant("map('a': 1, 'b': 2) AS variant", Type::Variant)]
#[case::map_array_value_to_variant("map('x': [1, 2], 'y': [3, 4]) AS variant", Type::Variant)]
#[case::null_to_variant("null AS variant", Type::Variant)]
fn test_valid_casts_to_variant(#[case] input: &str, #[case] expected_type: Type) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
assert_eq!(*typed.resolved_type, expected_type);
}
#[rstest]
#[case::variant_to_int("(42 AS variant) AS int", INT)]
#[case::variant_to_double("((3.14 AS double) AS variant) AS double", DOUBLE)]
#[case::variant_to_string("('hello' AS variant) AS string", STRING)]
#[case::variant_to_boolean("(true AS variant) AS boolean", BOOLEAN)]
#[case::variant_to_timestamp("(42 AS variant) AS timestamp", TIMESTAMP)]
#[case::variant_to_array_int("([1, 2, 3] AS variant) AS array(int)", Array::new(INT).into())]
#[case::variant_to_array_string("(['a', 'b'] AS variant) AS array(string)", Array::new(STRING).into())]
#[case::variant_to_array_timestamp(
"(42 AS variant) AS array(timestamp)",
Array::new(TIMESTAMP).into()
)]
#[case::variant_to_nested_array(
"([[1, 2]] AS variant) AS array(array(int))",
Array::new(Array::new(INT).into()).into()
)]
#[case::variant_to_struct(
"({x: 1} AS variant) AS {x: int}",
Struct::default().with_str("x", INT).into()
)]
#[case::variant_to_struct_timestamp(
"(42 AS variant) AS {ts: timestamp}",
Struct::default().with_str("ts", TIMESTAMP).into()
)]
fn test_valid_casts_from_variant(#[case] input: &str, #[case] expected_type: Type) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
assert_eq!(*typed.resolved_type, expected_type);
}
#[rstest]
#[case::timestamp_to_variant("('2024-01-01' AS timestamp) AS variant", "Cannot cast")]
#[case::interval_to_variant("1d AS variant", "Cannot cast")]
#[case::calendar_interval_to_variant("1mon AS variant", "Cannot cast")]
#[case::array_timestamp_to_variant("[('2024-01-01' AS timestamp)] AS variant", "Cannot cast")]
#[case::array_interval_to_variant("[1d, 2d] AS variant", "Cannot cast")]
#[case::struct_timestamp_to_variant("{ts: ('2024-01-01' AS timestamp)} AS variant", "Cannot cast")]
#[case::struct_interval_to_variant("{dur: 1d} AS variant", "Cannot cast")]
#[case::map_timestamp_value_to_variant(
"map('x': ('2024-01-01' AS timestamp)) AS variant",
"Cannot cast"
)]
fn test_invalid_casts_to_variant(#[case] input: &str, #[case] expected_error_fragment: &str) {
let expr =
Expression::parse_result(input).expect("Cast should parse successfully even if invalid");
let result =
type_check_expression(expr, ExpressionTypeCheckOptions::builder().build()).into_result();
assert!(
result.is_err(),
"Expected cast to fail type checking but it succeeded: {}",
input
);
let err = result.unwrap_err();
let err_string = err.to_string();
eprintln!("Error for '{}': {}", input, err_string);
assert!(
err_string.contains(expected_error_fragment),
"Error message '{}' does not contain expected fragment '{}'",
err_string,
expected_error_fragment
);
}
#[rstest]
#[case::variant_to_interval("(42 AS variant) AS interval", "Cannot cast")]
#[case::variant_to_array_interval("(42 AS variant) AS array(interval)", "Cannot cast")]
#[case::variant_to_struct_interval("(42 AS variant) AS {dur: interval}", "Cannot cast")]
fn test_invalid_casts_from_variant(#[case] input: &str, #[case] expected_error_fragment: &str) {
let expr =
Expression::parse_result(input).expect("Cast should parse successfully even if invalid");
let result =
type_check_expression(expr, ExpressionTypeCheckOptions::builder().build()).into_result();
assert!(
result.is_err(),
"Expected cast to fail type checking but it succeeded: {}",
input
);
let err = result.unwrap_err();
let err_string = err.to_string();
eprintln!("Error for '{}': {}", input, err_string);
assert!(
err_string.contains(expected_error_fragment),
"Error message '{}' does not contain expected fragment '{}'",
err_string,
expected_error_fragment
);
}
#[rstest]
#[case::empty_array_to_array_int("[] AS array(int)", Array::new(INT).into())]
#[case::empty_array_to_array_string("[] AS array(string)", Array::new(STRING).into())]
#[case::empty_array_to_array_boolean("[] AS array(boolean)", Array::new(BOOLEAN).into())]
#[case::empty_array_to_nested_array("[] AS array(array(int))", Array::new(Array::new(INT).into()).into())]
#[case::null_to_array("null AS array(int)", Array::new(INT).into())]
fn test_empty_container_casts(#[case] input: &str, #[case] expected_type: Type) {
let typed = type_check_expression(
Expression::parse_result(input).unwrap(),
ExpressionTypeCheckOptions::builder().build(),
)
.into_result()
.unwrap();
assert_eq!(*typed.resolved_type, expected_type);
}