hamelin_lib 0.15.4

Core library for Hamelin query language
Documentation
//! Tests for lambda expressions, transform, and filter functions.
//!
//! Lambda expressions use arrow syntax: `param -> body` or `(a, b) -> body`.
//! `test_parsed_*` tests parse from surface syntax and verify both AST structure
//! and type resolution. Builder-based tests remain for error/edge cases that
//! exercise the type checker directly (e.g., arity mismatches, type conflicts).

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;

// ======================== Lambda Parameter Shadowing Tests ========================
// Test that lambda parameters properly shadow outer bindings

#[rstest]
// Lambda parameter shadows outer binding
#[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,
) {
    // Create an environment where "x" already exists as a string
    let mut env = TypeEnvironment::default();
    env.bind(SimpleIdentifier::new("x").into(), STRING);

    // Despite "x" being a string in outer scope, the lambda param "x" is int
    // because it iterates over array(int), so x * 10 should work and return int
    let typed = type_check_expression(
        expr_builder.build(),
        ExpressionTypeCheckOptions::builder()
            .bindings(Arc::new(env))
            .build(),
    )
    .output;
    assert_eq!(*typed.resolved_type, expected_type);
}

// ======================== Transform Error Tests ========================

#[rstest]
// Non-lambda argument
#[case::transform_non_lambda(
    call("transform")
        .arg(array().element(1).element(2))
        .arg(1)
)]
// Lambda arity mismatch
#[case::transform_arity_mismatch(
    call("transform")
        .arg(array().element(1).element(2))
        .arg(lambda2("x", "y").body(add(field_ref("x"), field_ref("y"))))
)]
// Conflicting type hint
#[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"))))
)]
// Lambda body type error
#[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"
    );
}

// ======================== Lambda Parsing Tests ========================

#[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
        );
    }
}

// ======================== Filter Parsing and Type Tests ========================

#[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
        );
    }
}