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