hamelin_lib 0.15.6

Core library for Hamelin query language
Documentation
//! Query tests with DEF statements (CTEs), JOIN, and LOOKUP commands
//!
//! These tests verify that queries with:
//! - Common Table Expressions (CTEs) using the DEF statement
//! - JOIN commands that merge schemas from environment chains
//! - LOOKUP commands that preserve aggregated fields
//!
//! Tests ensure:
//! 1. Parse correctly into Query AST nodes
//! 2. Match the builder-constructed equivalent
//! 3. Resolve types correctly with the provided catalog
//! 4. Produce the expected final schema
//! 5. JOIN/LOOKUP preserve fields from chained environments (base chain)

use std::sync::Arc;

use crate::err::TranslationErrors;
use crate::tree::ast::identifier::SimpleIdentifier;
use crate::tree::ast::query::Query;
use crate::tree::ast::ParseWithErrors;
use crate::tree::builder::*;
use crate::tree::tests::shared::*;
use crate::types::struct_type::Struct;
use crate::types::{Type, BOOLEAN, DOUBLE, INT, STRING, TIMESTAMP};
use pretty_assertions::assert_eq;
use rstest::rstest;

#[rstest]
#[case::simple_query(
    vec![("users", vec![("id", INT), ("name", STRING)])],
    "FROM users",
    query().main(pipeline().from(|f| f.table_reference("users"))).build(),
    Struct::default().with_str("id", INT).with_str("name", STRING)
)]
#[case::with_single_def(
    vec![("users", vec![("id", INT), ("name", STRING), ("active", BOOLEAN)])],
    "DEF active_users = FROM users | WHERE active == true;\nFROM active_users",
    query().main(pipeline().from(|f| f.table_reference("active_users")))
        .def_pipeline("active_users", pipeline().from(|f| f.table_reference("users")).where_cmd(eq(field_ref("active"), true)))
        .build(),
    Struct::default().with_str("id", INT).with_str("name", STRING).with_str("active", BOOLEAN)
)]
#[case::with_multiple_defs(
    vec![("users", vec![("id", INT), ("name", STRING), ("active", BOOLEAN), ("value", INT)])],
    "DEF active_users = FROM users | WHERE active == true;\nDEF high_value = FROM active_users | WHERE value > 1000;\nFROM high_value",
    query().main(pipeline().from(|f| f.table_reference("high_value")))
        .def_pipeline("active_users", pipeline().from(|f| f.table_reference("users")).where_cmd(eq(field_ref("active"), true)))
        .def_pipeline("high_value", pipeline().from(|f| f.table_reference("active_users")).where_cmd(gt(field_ref("value"), 1000)))
        .build(),
    Struct::default().with_str("id", INT).with_str("name", STRING).with_str("active", BOOLEAN).with_str("value", INT)
)]
#[case::with_complex_pipeline(
    vec![("sales", vec![("id", INT), ("year", INT), ("amount", INT), ("region", STRING)])],
    "DEF sales_2024 = FROM sales | WHERE year == 2024;\nFROM sales_2024 | AGG total = sum(amount) BY region",
    query().main(pipeline()
        .from(|f| f.table_reference("sales_2024"))
        .agg(|a| a.named_aggregate("total", call("sum").arg(field_ref("amount"))).group_by("region")))
        .def_pipeline("sales_2024", pipeline().from(|f| f.table_reference("sales")).where_cmd(eq(field_ref("year"), 2024)))
        .build(),
    Struct::default().with_str("region", STRING).with_str("total", INT)
)]
#[case::with_def_referencing_def(
    vec![("users", vec![("id", INT), ("name", STRING), ("active", BOOLEAN), ("verified", BOOLEAN)])],
    "DEF step1 = FROM users | WHERE active == true;\nDEF step2 = FROM step1 | WHERE verified == true;\nFROM step2",
    query().main(pipeline().from(|f| f.table_reference("step2")))
        .def_pipeline("step1", pipeline().from(|f| f.table_reference("users")).where_cmd(eq(field_ref("active"), true)))
        .def_pipeline("step2", pipeline().from(|f| f.table_reference("step1")).where_cmd(eq(field_ref("verified"), true)))
        .build(),
    Struct::default().with_str("id", INT).with_str("name", STRING).with_str("active", BOOLEAN).with_str("verified", BOOLEAN)
)]
fn test_queries_with_ctes(
    #[case] tables: Vec<(&str, Vec<(&str, Type)>)>,
    #[case] input: &str,
    #[case] expected_ast: Query,
    #[case] expected_output_fields: Struct,
) -> Result<(), TranslationErrors> {
    let (query, errors) = Query::parse_with_errors(input);
    if !errors.is_empty() {
        eprintln!("Parse errors:\n{}", errors);
    }
    assert!(errors.is_empty(), "Expected no parse errors");
    assert_eq!(query.to_string(), expected_ast.to_string());

    let catalog = catalog_with_tables(tables);
    let typed = type_check_with_catalog::<Query>(Arc::new(query), catalog).into_result()?;

    let final_schema = match typed.pipeline.kind.as_ref() {
        Ok(valid_pipeline) => valid_pipeline.final_schema.clone(),
        Err(e) => panic!("Type checking failed: {}", e),
    };

    assert_eq!(*final_schema.as_struct(), expected_output_fields);
    Ok(())
}
// ============================================================================
// JOIN/LOOKUP Tests with Environment Chains
// ============================================================================
//
// These tests specifically verify that JOIN and LOOKUP commands preserve
// fields from environment base chains. This was a bug where JOIN/LOOKUP
// would lose fields that were in the base chain (from SET commands after AGG).
//
// The bug manifested when:
// 1. Commands created an environment with a base chain (like SET extending AGG output)
// 2. JOIN/LOOKUP tried to merge this chained environment
// 3. The merge would only include immediate fields, losing base chain fields

#[rstest]
// Test LOOKUP preserves fields from AGG + SET chain
#[case::lookup_after_agg_and_set(
    vec![
        ("events", vec![("timestamp", TIMESTAMP), ("user", STRING), ("action", STRING)]),
        ("users", vec![("user", STRING), ("email", STRING)])
    ],
    "FROM events
| AGG total = count() BY user
| SET is_active = total > 10
| LOOKUP u = users ON user == u.user",
    Struct::default()
        .with_str("is_active", BOOLEAN)
        .with_str("user", STRING)
        .with_str("total", INT)
        .with_str("u", Struct::default().with_str("user", STRING).with_str("email", STRING).into())
)]
// Test LOOKUP preserves AGG GROUP BY fields through multiple SETs
#[case::lookup_preserves_group_by_through_multiple_sets(
    vec![
        ("sales", vec![("product", STRING), ("region", STRING), ("amount", INT)]),
        ("products", vec![("product", STRING), ("category", STRING)])
    ],
    "FROM sales
| AGG total_sales = sum(amount) BY product, region
| SET high_value = total_sales > 1000
| SET status = if(high_value, 'hot', 'cold')
| LOOKUP p = products ON product == p.product",
    Struct::default()
        .with_str("status", STRING)
        .with_str("high_value", BOOLEAN)
        .with_str("product", STRING)
        .with_str("region", STRING)
        .with_str("total_sales", INT)
        .with_str("p", Struct::default().with_str("product", STRING).with_str("category", STRING).into())
)]
// Test LOOKUP preserves aggregate fields
#[case::lookup_preserves_aggregate_fields(
    vec![
        ("orders", vec![("customer_id", INT), ("order_date", TIMESTAMP), ("amount", INT)]),
        ("customers", vec![("customer_id", INT), ("name", STRING), ("tier", STRING)])
    ],
    "FROM orders
| AGG total_orders = count(), total_amount = sum(amount), avg_amount = avg(amount) BY customer_id
| LOOKUP c = customers ON customer_id == c.customer_id",
    Struct::default()
        .with_str("customer_id", INT)
        .with_str("total_orders", INT)
        .with_str("total_amount", INT)
        .with_str("avg_amount", DOUBLE)
        .with_str("c", Struct::default().with_str("customer_id", INT).with_str("name", STRING).with_str("tier", STRING).into())
)]
// Test JOIN preserves fields from SET chain
#[case::join_after_multiple_sets(
    vec![
        ("events", vec![("id", INT), ("user_id", INT), ("timestamp", TIMESTAMP)]),
        ("users", vec![("user_id", INT), ("name", STRING)])
    ],
    "FROM events
| SET day = 1
| SET is_recent = day < 7
| JOIN u = users ON user_id == u.user_id",
    Struct::default()
        .with_str("is_recent", BOOLEAN)
        .with_str("day", INT)
        .with_str("id", INT)
        .with_str("user_id", INT)
        .with_str("timestamp", TIMESTAMP)
        .with_str("u", Struct::default().with_str("user_id", INT).with_str("name", STRING).into())
)]
// Test JOIN after AGG preserves GROUP BY and aggregate fields
#[case::join_after_agg_preserves_all_fields(
    vec![
        ("logs", vec![("level", STRING), ("message", STRING)]),
        ("levels", vec![("level", STRING), ("severity", INT)])
    ],
    "FROM logs
| AGG count = count() BY level
| JOIN l = levels ON level == l.level",
    Struct::default()
        .with_str("level", STRING)
        .with_str("count", INT)
        .with_str("l", Struct::default().with_str("level", STRING).with_str("severity", INT).into())
)]
// Test LOOKUP in WHERE clause can reference AGG fields
#[case::lookup_then_where_references_agg_fields(
    vec![
        ("metrics", vec![("host", STRING), ("value", INT)]),
        ("hosts", vec![("host", STRING), ("max_value", INT)])
    ],
    "FROM metrics
| AGG total = sum(value), avg = avg(value) BY host
| LOOKUP h = hosts ON host == h.host
| WHERE total > h.max_value",
    Struct::default()
        .with_str("host", STRING)
        .with_str("total", INT)
        .with_str("avg", DOUBLE)
        .with_str("h", Struct::default().with_str("host", STRING).with_str("max_value", INT).into())
)]
// Test chained JOIN/LOOKUP operations preserve all fields
#[case::multiple_lookups_preserve_all_fields(
    vec![
        ("orders", vec![("order_id", INT), ("customer_id", INT), ("product_id", INT)]),
        ("customers", vec![("customer_id", INT), ("name", STRING)]),
        ("products", vec![("product_id", INT), ("category", STRING)])
    ],
    "FROM orders
| LOOKUP c = customers ON customer_id == c.customer_id
| LOOKUP p = products ON product_id == p.product_id",
    Struct::default()
        .with_str("order_id", INT)
        .with_str("customer_id", INT)
        .with_str("product_id", INT)
        .with_str("c", Struct::default().with_str("customer_id", INT).with_str("name", STRING).into())
        .with_str("p", Struct::default().with_str("product_id", INT).with_str("category", STRING).into())
)]
// Test LOOKUP after AGG with many GROUP BY fields
#[case::lookup_after_agg_many_group_by_fields(
    vec![
        ("sales", vec![("year", INT), ("quarter", INT), ("region", STRING), ("product", STRING), ("revenue", INT)]),
        ("targets", vec![("year", INT), ("quarter", INT), ("region", STRING), ("target", INT)])
    ],
    "FROM sales
| AGG total_revenue = sum(revenue) BY year, quarter, region, product
| LOOKUP t = targets
  ON year == t.year
  AND quarter == t.quarter
  AND region == t.region",
    Struct::default()
        .with_str("year", INT)
        .with_str("quarter", INT)
        .with_str("region", STRING)
        .with_str("product", STRING)
        .with_str("total_revenue", INT)
        .with_str("t", Struct::default().with_str("year", INT).with_str("quarter", INT).with_str("region", STRING).with_str("target", INT).into())
)]
// Test JOIN with complex field references from environment chain
#[case::join_with_complex_on_clause_from_chain(
    vec![
        ("events", vec![("user_id", INT), ("timestamp", TIMESTAMP), ("score", INT)]),
        ("users", vec![("user_id", INT), ("email", STRING)])
    ],
    "FROM events
| SET is_high_score = score > 100
| AGG total_score = sum(score), event_count = count() BY user_id, is_high_score
| JOIN u = users ON user_id == u.user_id",
    Struct::default()
        .with_str("user_id", INT)
        .with_str("is_high_score", BOOLEAN)
        .with_str("total_score", INT)
        .with_str("event_count", INT)
        .with_str("u", Struct::default().with_str("user_id", INT).with_str("email", STRING).into())
)]
fn test_join_lookup_with_environment_chains(
    #[case] tables: Vec<(&str, Vec<(&str, Type)>)>,
    #[case] input: &str,
    #[case] expected_output_fields: Struct,
) -> Result<(), TranslationErrors> {
    let (query, errors) = Query::parse_with_errors(input);
    if !errors.is_empty() {
        eprintln!("Parse errors:\n{}", errors);
    }
    assert!(errors.is_empty(), "Expected no parse errors");

    let catalog = catalog_with_tables(tables);
    let typed = type_check_with_catalog(query, catalog).into_result()?;

    let final_schema = match typed.pipeline.kind.as_ref() {
        Ok(valid_pipeline) => valid_pipeline.final_schema.clone(),
        Err(e) => panic!("Type checking failed: {}", e),
    };

    assert_eq!(*final_schema.as_struct(), expected_output_fields);
    Ok(())
}

// ============================================================================
// Query.datasets() Tests
// ============================================================================
//
// These tests verify that Query::datasets() correctly extracts all external
// table references from a query, excluding tabular DEF names (pipeline CTEs).
// Scalar DEF names must not be excluded — they do not shadow table references.

#[rstest]
#[case::simple_from(
    "FROM users",
    vec!["users"]
)]
#[case::from_with_schema(
    "FROM schema:users",
    vec!["schema:users"]
)]
#[case::multiple_from_tables(
    "FROM users, orders",
    vec!["users", "orders"]
)]
#[case::union_tables(
    "FROM users | UNION orders",
    vec!["users", "orders"]
)]
#[case::join_extracts_both_tables(
    "FROM users | JOIN o = orders ON id == o.id",
    vec!["users", "orders"]
)]
#[case::lookup_extracts_both_tables(
    "FROM users | LOOKUP o = orders ON id == o.id",
    vec!["users", "orders"]
)]
#[case::with_clause_excludes_cte(
    "DEF active = FROM users | WHERE active == true;\nFROM active",
    vec!["users"]
)]
#[case::with_clause_multiple_ctes(
    "DEF a = FROM users;\nDEF b = FROM a | JOIN o = orders ON id == o.id;\nFROM b",
    vec!["users", "orders"]
)]
#[case::cte_in_join_excluded(
    "DEF temp = FROM users;\nFROM orders | JOIN t = temp ON id == t.id",
    vec!["users", "orders"]
)]
#[case::complex_query_all_sources(
    "DEF enriched = FROM events | JOIN u = users ON user_id == u.id;\nFROM enriched | LOOKUP p = products ON product_id == p.id",
    vec!["events", "users", "products"]
)]
#[case::append_extracts_target(
    "FROM users | APPEND orders",
    vec!["users", "orders"]
)]
#[case::append_cte_excluded(
    "DEF temp = FROM users;\nFROM orders | APPEND temp",
    vec!["users", "orders"]
)]
#[case::match_extracts_pattern_tables(
    "MATCH a=logins b=logouts",
    vec!["logins", "logouts"]
)]
#[case::match_nested_patterns(
    "MATCH (a=logins b=logouts)+",
    vec!["logins", "logouts"]
)]
#[case::match_cte_excluded(
    "DEF temp = FROM users;\nMATCH a=temp b=logins",
    vec!["users", "logins"]
)]
#[case::scalar_def_name_does_not_exclude_same_named_table(
    "DEF x = 42;\nFROM x",
    vec!["x"]
)]
fn test_query_datasets(#[case] input: &str, #[case] expected_tables: Vec<&str>) {
    let (query, errors) = Query::parse_with_errors(input);
    if !errors.is_empty() {
        eprintln!("Parse errors:\n{}", errors);
    }
    assert!(errors.is_empty(), "Expected no parse errors");

    let default_space = SimpleIdentifier::new("test");
    let mut table_names: Vec<String> = query
        .datasets(Some(&default_space))
        .expect("datasets() should succeed")
        .into_iter()
        .map(|id| id.to_string())
        .collect();
    table_names.sort_unstable();

    let mut expected: Vec<String> = expected_tables
        .iter()
        .map(|s| {
            if s.contains(':') {
                (*s).to_string()
            } else {
                format!("test:{s}")
            }
        })
        .collect();
    expected.sort_unstable();

    assert_eq!(table_names, expected, "Mismatch for query: {}", input);
}