use std::sync::Arc;
use crate::tree::ast::identifier::SimpleIdentifier;
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::transform_multiply(
call("transform")
.arg(array().element(1).element(2).element(3))
.arg(lambda1("x").body(multiply(field_ref("x"), 2))),
Array::new(INT).into()
)]
#[case::transform_string_to_int(
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::transform_add(
call("transform")
.arg(array().element(10).element(20).element(30))
.arg(lambda1("n").body(add(field_ref("n"), 5))),
Array::new(INT).into()
)]
#[case::transform_upper(
call("transform")
.arg(array().element("hello").element("world"))
.arg(lambda1("s").body(call("upper").arg(field_ref("s")))),
Array::new(STRING).into()
)]
#[case::transform_nested_arrays(
call("transform")
.arg(array().element(array().element(1).element(2)).element(array().element(3)))
.arg(lambda1("x").body(call("len").arg(field_ref("x")))),
Array::new(INT).into()
)]
fn test_transform(#[case] expr_builder: impl ExpressionBuilder, #[case] expected_type: Type) {
let typed = type_check_expression(
expr_builder.build(),
ExpressionTypeCheckOptions::builder().build(),
)
.output;
assert_eq!(*typed.resolved_type, expected_type);
}
#[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"
);
}