hamelin_lib 0.7.13

Core library for Hamelin query language
Documentation
//! Tests for lambda expressions and transform function
//!
//! Lambda expressions are anonymous functions written as `param -> body`.
//! The `transform()` function applies a lambda element-wise to an array.
//!
//! These tests use builders since lambda syntax isn't in the parser yet.

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;

// ======================== Transform Function Tests ========================

#[rstest]
// Basic transform with arithmetic
#[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()
)]
// Transform changing element type (string -> int via len)
#[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()
)]
// Transform with addition
#[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()
)]
// Transform with function call
#[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()
)]
// Transform over nested arrays (array(array(int)) -> array(int))
#[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);
}

// ======================== 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"
    );
}