hamelin_lib 0.10.8

Core library for Hamelin query language
Documentation
//! Tests for ROWS command pipelines
//!
//! ROWS creates inline datasets from array literals of struct literals.
//! These tests validate parsing, AST structure, and schema inference from ROWS.

use crate::err::TranslationErrors;
use crate::tree::ast::pipeline::Pipeline;
use crate::tree::ast::ParseWithErrors;
use crate::tree::builder::*;
use crate::type_check;
use crate::types::struct_type::Struct;
use crate::types::{INT, STRING};
use pretty_assertions::assert_eq;
use rstest::rstest;

// ======================== ROWS Pipeline Tests ========================

#[rstest]
// Two-command pipelines
#[case::simple_where(
    r#"
    ROWS [
        {left: 'hello', right: 1},
        {left: 'world', right: 2}
    ]
    | WHERE right > 1
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("left", "hello").field("right", 1))
                .element(struct_literal().field("left", "world").field("right", 2)),
        )
        .where_cmd(gt(field_ref("right"), 1)),
    Struct::default()
        .with_str("left", STRING)
        .with_str("right", INT)
)]
#[case::simple_select(
    r#"
    ROWS [
        {name: 'Alice', age: 30},
        {name: 'Bob', age: 25}
    ]
    | SELECT name
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("name", "Alice").field("age", 30))
                .element(struct_literal().field("name", "Bob").field("age", 25)),
        )
        .select(|s| s.field("name")),
    Struct::default()
        .with_str("name", STRING)
)]
#[case::with_set(
    r#"
    ROWS [
        {x: 1, y: 2},
        {x: 3, y: 4}
    ]
    | SET sum = x + y
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("x", 1).field("y", 2))
                .element(struct_literal().field("x", 3).field("y", 4)),
        )
        .set_cmd(|l| l.named_field("sum", add(field_ref("x"), field_ref("y")))),
    Struct::default()
        .with_str("sum", INT)
        .with_str("x", INT)
        .with_str("y", INT)
)]
#[case::with_sort(
    r#"
    ROWS [
        {name: 'Charlie', score: 85},
        {name: 'Alice', score: 92},
        {name: 'Bob', score: 78}
    ]
    | SORT score DESC
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("name", "Charlie").field("score", 85))
                .element(struct_literal().field("name", "Alice").field("score", 92))
                .element(struct_literal().field("name", "Bob").field("score", 78)),
        )
        .sort(|s| s.desc(field_ref("score"))),
    Struct::default()
        .with_str("name", STRING)
        .with_str("score", INT)
)]
#[case::with_limit(
    r#"
    ROWS [
        {id: 1},
        {id: 2},
        {id: 3}
    ]
    | LIMIT 2
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("id", 1))
                .element(struct_literal().field("id", 2))
                .element(struct_literal().field("id", 3)),
        )
        .limit(2),
    Struct::default()
        .with_str("id", INT)
)]
// Multi-command pipelines
#[case::aggregation(
    r#"
    ROWS [
        {left: 'hello', right: 1},
        {left: 'world', right: 2},
        {left: 'hello', right: 1},
        {left: 'world', right: 2}
    ]
    | AGG count() BY left
    | SORT left
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("left", "hello").field("right", 1))
                .element(struct_literal().field("left", "world").field("right", 2))
                .element(struct_literal().field("left", "hello").field("right", 1))
                .element(struct_literal().field("left", "world").field("right", 2)),
        )
        .agg(|a| a.named_aggregate("count", call("count")).group_by("left"))
        .sort(|s| s.by(field_ref("left"))),
    Struct::default()
        .with_str("left", STRING)
        .with_str("count", INT)
)]
#[case::complex_aggregation(
    r#"
    ROWS [
        {left: 'hello', right: 1},
        {left: 'world', right: 2},
        {left: 'hello', right: 1},
        {left: 'world', right: 2}
    ]
    | AGG s = sum(right), c = count() BY left
    | SORT left
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("left", "hello").field("right", 1))
                .element(struct_literal().field("left", "world").field("right", 2))
                .element(struct_literal().field("left", "hello").field("right", 1))
                .element(struct_literal().field("left", "world").field("right", 2)),
        )
        .agg(|a| {
            a.named_aggregate("s", call("sum").arg(field_ref("right")))
                .named_aggregate("c", call("count"))
                .group_by("left")
        })
        .sort(|s| s.by(field_ref("left"))),
    Struct::default()
        .with_str("left", STRING)
        .with_str("s", INT)
        .with_str("c", INT)
)]
#[case::where_select_sort(
    r#"
    ROWS [
        {name: 'Alice', age: 30, active: true},
        {name: 'Bob', age: 25, active: false},
        {name: 'Charlie', age: 35, active: true}
    ]
    | WHERE active == true
    | SELECT name, age
    | SORT age DESC
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("name", "Alice").field("age", 30).field("active", true))
                .element(struct_literal().field("name", "Bob").field("age", 25).field("active", false))
                .element(struct_literal().field("name", "Charlie").field("age", 35).field("active", true)),
        )
        .where_cmd(eq(field_ref("active"), true))
        .select(|s| s.field("name").field("age"))
        .sort(|s| s.desc(field_ref("age"))),
    Struct::default()
        .with_str("name", STRING)
        .with_str("age", INT)
)]
// ROWS-only pipelines
#[case::rows_only(
    r#"
    ROWS [
        {name: 'Alice', age: 30},
        {name: 'Bob', age: 25}
    ]
    "#,
    pipeline()
        .rows(
            array()
                .element(struct_literal().field("name", "Alice").field("age", 30))
                .element(struct_literal().field("name", "Bob").field("age", 25)),
        ),
    Struct::default()
        .with_str("name", STRING)
        .with_str("age", INT)
)]
// empty_array case removed - ROWS requires at least one row to determine schema
#[case::single_element(
    "ROWS [{x: 42}]",
    pipeline().rows(array().element(struct_literal().field("x", 42))),
    Struct::default()
        .with_str("x", INT)
)]
fn test_rows_pipelines(
    #[case] query: &str,
    #[case] expected_ast: PipelineBuilder,
    #[case] expected_schema: Struct,
) -> Result<(), TranslationErrors> {
    let pipeline = Pipeline::parse_result(query)?;
    assert_eq!(pipeline, expected_ast.build());

    let typed = type_check(pipeline).into_result()?;
    let actual_schema = typed.schema();

    assert_eq!(actual_schema, expected_schema);
    Ok(())
}