use std::sync::Arc;
use crate::err::TranslationErrors;
use crate::tree::ast::expression::Expression;
use crate::tree::ast::identifier::SimpleIdentifier;
use crate::tree::ast::ParseWithErrors;
use crate::tree::builder::*;
use crate::tree::options::ExpressionTypeCheckOptions;
use crate::tree::typed_ast::environment::TypeEnvironment;
use crate::type_check_expression;
use crate::types::array::Array;
use crate::types::{Type, INT, STRING};
use pretty_assertions::assert_eq;
use rstest::rstest;
#[rstest]
#[case::lambda_shadows_outer(
call("transform")
.arg(array().element(1).element(2).element(3))
.arg(lambda1("x").body(multiply(field_ref("x"), 10))),
Array::new(INT).into()
)]
fn test_lambda_shadowing(
#[case] expr_builder: impl ExpressionBuilder,
#[case] expected_type: Type,
) {
let mut env = TypeEnvironment::default();
env.bind(SimpleIdentifier::new("x").into(), STRING);
let typed = type_check_expression(
expr_builder.build(),
ExpressionTypeCheckOptions::builder()
.bindings(Arc::new(env))
.build(),
)
.output;
assert_eq!(*typed.resolved_type, expected_type);
}
#[rstest]
#[case::transform_non_lambda(
call("transform")
.arg(array().element(1).element(2))
.arg(1)
)]
#[case::transform_arity_mismatch(
call("transform")
.arg(array().element(1).element(2))
.arg(lambda2("x", "y").body(add(field_ref("x"), field_ref("y"))))
)]
#[case::transform_conflicting_hint(
call("transform")
.arg(array().element(1).element(2))
.arg(lambda().param_with_type("x", Arc::new(STRING)).body(call("len").arg(field_ref("x"))))
)]
#[case::transform_lambda_body_error(
call("transform")
.arg(array().element(1))
.arg(lambda1("x").body(call("upper").arg(field_ref("x"))))
)]
#[case::transform_lambda_body_error_deep(
call("transform")
.arg(array().element(1).element(2).element(3))
.arg(lambda1("e").body(add(add(field_ref("e"), "a"), "b")))
)]
fn test_transform_errors(#[case] expr_builder: impl ExpressionBuilder) {
let errors = type_check_expression(
expr_builder.build(),
ExpressionTypeCheckOptions::builder().build(),
)
.errors;
assert!(
!errors.is_empty(),
"Expected type-check errors but got none"
);
}
#[rstest]
#[case::single_param(
"transform([1, 2, 3], x -> x * 2)",
call("transform")
.arg(array().element(1).element(2).element(3))
.arg(lambda1("x").body(multiply(field_ref("x"), 2))),
Array::new(INT).into()
)]
#[case::single_param_addition(
"transform([10, 20], n -> n + 5)",
call("transform")
.arg(array().element(10).element(20))
.arg(lambda1("n").body(add(field_ref("n"), 5))),
Array::new(INT).into()
)]
#[case::lambda_with_function_call(
"transform(['hello', 'world'], s -> upper(s))",
call("transform")
.arg(array().element("hello").element("world"))
.arg(lambda1("s").body(call("upper").arg(field_ref("s")))),
Array::new(STRING).into()
)]
#[case::lambda_with_len(
"transform(['a', 'bb', 'ccc'], s -> len(s))",
call("transform")
.arg(array().element("a").element("bb").element("ccc"))
.arg(lambda1("s").body(call("len").arg(field_ref("s")))),
Array::new(INT).into()
)]
#[case::lambda_body_with_pair(
"transform([1, 2], x -> x : x * 2)",
call("transform")
.arg(array().element(1).element(2))
.arg(lambda1("x").body(pair(field_ref("x"), multiply(field_ref("x"), 2)))),
Array::new(
crate::types::tuple::Tuple::new(vec![INT, INT]).into()
).into()
)]
fn test_parsed_transform(
#[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::parenthesized_single_param(
"transform([1, 2, 3], (x) -> x * 2)",
Array::new(INT).into()
)]
#[case::multi_param_tuple(
"transform([1, 2, 3], (x) -> x + 1)",
Array::new(INT).into()
)]
fn test_parsed_lambda_type_only(
#[case] input: &str,
#[case] expected_type: Type,
) -> Result<(), TranslationErrors> {
let actual = Expression::parse_result(input)?;
let typed = type_check_expression(actual, ExpressionTypeCheckOptions::builder().build())
.into_result()?;
assert_eq!(*typed.resolved_type, expected_type);
Ok(())
}
#[rstest]
#[case::lambda_parse_error("transform([1], 42 -> x)")]
fn test_parsed_lambda_errors(#[case] input: &str) {
let expr = Expression::parse_result(input);
if let Ok(expr) = expr {
let result = type_check_expression(expr, ExpressionTypeCheckOptions::builder().build())
.into_result();
assert!(
result.is_err(),
"Expected error but got success for: {}",
input
);
}
}
#[rstest]
#[case::filter_integers(
"filter([1, 2, 3, 4], x -> x > 2)",
Array::new(INT).into()
)]
#[case::filter_strings(
"filter(['hello', 'hi', 'hey'], s -> len(s) > 2)",
Array::new(STRING).into()
)]
#[case::filter_with_equality(
"filter([1, 2, 3], x -> x == 2)",
Array::new(INT).into()
)]
fn test_parsed_filter(
#[case] input: &str,
#[case] expected_type: Type,
) -> Result<(), TranslationErrors> {
let actual = Expression::parse_result(input)?;
let typed = type_check_expression(actual, ExpressionTypeCheckOptions::builder().build())
.into_result()?;
assert_eq!(*typed.resolved_type, expected_type);
Ok(())
}
#[rstest]
#[case::filter_non_boolean_lambda("filter([1, 2, 3], x -> x * 2)")]
#[case::filter_non_lambda("filter([1, 2, 3], 42)")]
fn test_filter_errors(#[case] input: &str) {
let expr = Expression::parse_result(input);
if let Ok(expr) = expr {
let result = type_check_expression(expr, ExpressionTypeCheckOptions::builder().build())
.into_result();
assert!(
result.is_err(),
"Expected error but got success for: {}",
input
);
}
}