diff --git a/runtime/drivers/bigquery/dialect.go b/runtime/drivers/bigquery/dialect.go index 55e1f32eb252..18171fde60fc 100644 --- a/runtime/drivers/bigquery/dialect.go +++ b/runtime/drivers/bigquery/dialect.go @@ -70,6 +70,29 @@ func (d *dialect) OrderByAliasExpression(name string, desc bool) string { return res } +func (d *dialect) DimensionSelect(_ string, dim *runtimev1.MetricsViewSpec_Dimension) (dimSelect, unnestClause string, err error) { + expr, err := d.MetricsViewDimensionExpression(dim) + if err != nil { + return "", "", fmt.Errorf("failed to get dimension expression: %w", err) + } + alias := d.EscapeAlias(dim.Name) + if !dim.Unnest { + return fmt.Sprintf(`(%s) AS %s`, expr, alias), "", nil + } + unnestColName := d.EscapeIdentifier(drivers.TempName(fmt.Sprintf("unnested_%s_", dim.Name))) + return fmt.Sprintf(`%s AS %s`, unnestColName, alias), fmt.Sprintf(`, UNNEST(%s) AS %s`, expr, unnestColName), nil +} + +// LateralUnnest returns a comma join with UNNEST. BigQuery exposes each element directly under the alias, so there is no tuple to index into. +func (d *dialect) LateralUnnest(expr, _, colName string) (tbl string, tupleStyle, auto bool, err error) { + return fmt.Sprintf(`UNNEST(%s) AS %s`, expr, d.EscapeIdentifier(colName)), false, false, nil +} + +func (d *dialect) ArrayAnyExpression(arrExpr, elemAlias string) (open, elem, closing string, ok bool) { + elem = d.EscapeIdentifier(elemAlias) + return fmt.Sprintf("EXISTS (SELECT 1 FROM UNNEST(%s) AS %s WHERE ", arrExpr, elem), elem, ")", true +} + func (d *dialect) JoinOnExpression(lhs, rhs string) string { // BigQuery requires plain equality for FULL joins return fmt.Sprintf("coalesce(CAST(%s AS STRING), '__rill_sentinel__') = coalesce(CAST(%s AS STRING), '__rill_sentinel__')", lhs, rhs) diff --git a/runtime/drivers/bigquery/olap_test.go b/runtime/drivers/bigquery/olap_test.go index a003847d1c21..c4350e721386 100644 --- a/runtime/drivers/bigquery/olap_test.go +++ b/runtime/drivers/bigquery/olap_test.go @@ -2,12 +2,14 @@ package bigquery_test import ( "context" + "fmt" "testing" "time" "github.com/google/uuid" runtimev1 "github.com/rilldata/rill/proto/gen/rill/runtime/v1" "github.com/rilldata/rill/runtime/drivers" + "github.com/rilldata/rill/runtime/metricsview" "github.com/rilldata/rill/runtime/pkg/activity" "github.com/rilldata/rill/runtime/storage" "github.com/rilldata/rill/runtime/testruntime" @@ -121,6 +123,143 @@ func TestOLAP(t *testing.T) { } } +func TestUnnestDimension(t *testing.T) { + testmode.Expensive(t) + _, olap := acquireTestBigQuery(t) + + // Rows with overlapping, empty and NULL arrays, so joins that duplicate source rows and NULL handling are detectable. + // BigQuery arrays cannot contain NULL elements, and a NULL array is stored as an empty array. + name := "test_unnest_" + uuid.New().String()[:8] + table := "`rilldata.integration_test." + name + "`" + t.Cleanup(func() { + err := olap.Exec(context.Background(), &drivers.Statement{Query: "DROP TABLE IF EXISTS " + table}) + require.NoError(t, err) + }) + err := olap.Exec(t.Context(), &drivers.Statement{Query: "CREATE TABLE " + table + " AS SELECT id, tags FROM UNNEST(ARRAY>>[(1, ['a', 'b']), (2, ['b']), (3, ['c']), (4, ARRAY[]), (5, NULL)])"}) + require.NoError(t, err) + + mv := &runtimev1.MetricsViewSpec{ + Database: "rilldata", + DatabaseSchema: "integration_test", + Table: name, + Dimensions: []*runtimev1.MetricsViewSpec_Dimension{ + {Name: "tags", Column: "tags", Unnest: true}, + {Name: "id", Column: "id"}, + }, + Measures: []*runtimev1.MetricsViewSpec_Measure{ + {Name: "count", Expression: "count(*)", Type: runtimev1.MetricsViewSpec_MEASURE_TYPE_SIMPLE}, + }, + } + + // Same query shape as the executor's dimension validation. + dialect := olap.Dialect() + escapeTable := dialect.EscapeTable(mv.Database, mv.DatabaseSchema, mv.Table) + sel, unnestClause, err := dialect.DimensionSelect(escapeTable, mv.Dimensions[0]) + require.NoError(t, err) + err = olap.Exec(t.Context(), &drivers.Statement{Query: fmt.Sprintf("SELECT %s FROM %s %s GROUP BY 1", sel, escapeTable, unnestClause), DryRun: true}) + require.NoError(t, err) + + tagsFilter := func(op metricsview.Operator, val any) *metricsview.Expression { + return &metricsview.Expression{Condition: &metricsview.Condition{ + Operator: op, + Expressions: []*metricsview.Expression{{Name: "tags"}, {Value: val}}, + }} + } + count := func(where *metricsview.Expression) *metricsview.Query { + return &metricsview.Query{Measures: []metricsview.Measure{{Name: "count"}}, Where: where} + } + + tests := []struct { + name string + qry *metricsview.Query + want []map[string]any + }{ + { + name: "group by unnest dimension", + qry: &metricsview.Query{ + Dimensions: []metricsview.Dimension{{Name: "tags"}}, + Measures: []metricsview.Measure{{Name: "count"}}, + Sort: []metricsview.Sort{{Name: "tags"}}, + }, + want: []map[string]any{ + {"tags": "a", "count": int64(1)}, + {"tags": "b", "count": int64(2)}, + {"tags": "c", "count": int64(1)}, + }, + }, + { + // Row 1 matches both values but must be counted once. + name: "in filter counts each source row once", + qry: count(tagsFilter(metricsview.OperatorIn, []any{"a", "b"})), + want: []map[string]any{{"count": int64(2)}}, + }, + { + // Excludes rows containing 'a' even if they also contain other values; keeps the empty and NULL arrays. + name: "nin filter excludes rows containing any listed value", + qry: count(tagsFilter(metricsview.OperatorNin, []any{"a"})), + want: []map[string]any{{"count": int64(4)}}, + }, + { + name: "eq filter", + qry: count(tagsFilter(metricsview.OperatorEq, "b")), + want: []map[string]any{{"count": int64(2)}}, + }, + { + name: "neq filter excludes rows containing the value", + qry: count(tagsFilter(metricsview.OperatorNeq, "b")), + want: []map[string]any{{"count": int64(3)}}, + }, + { + name: "eq filter with no match", + qry: count(tagsFilter(metricsview.OperatorEq, "missing")), + want: []map[string]any{{"count": int64(0)}}, + }, + { + name: "filter combined with group by on another dimension", + qry: &metricsview.Query{ + Dimensions: []metricsview.Dimension{{Name: "id"}}, + Measures: []metricsview.Measure{{Name: "count"}}, + Where: tagsFilter(metricsview.OperatorIn, []any{"a", "b"}), + Sort: []metricsview.Sort{{Name: "id"}}, + }, + want: []map[string]any{ + {"id": int64(1), "count": int64(1)}, + {"id": int64(2), "count": int64(1)}, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.qry.MetricsView = name + ast, err := metricsview.NewAST(mv, allowAllSecurity{}, tt.qry, dialect) + require.NoError(t, err) + sql, args, err := ast.SQL() + require.NoError(t, err) + require.Equal(t, tt.want, queryRows(t, olap, sql, args)) + }) + } +} + +func queryRows(t *testing.T, olap drivers.OLAPStore, query string, args []any) []map[string]any { + rows, err := olap.Query(t.Context(), &drivers.Statement{Query: query, Args: args}) + require.NoError(t, err) + defer rows.Close() + var res []map[string]any + for rows.Next() { + row := make(map[string]any) + require.NoError(t, rows.MapScan(row)) + res = append(res, row) + } + require.NoError(t, rows.Err()) + return res +} + +type allowAllSecurity struct{} + +func (allowAllSecurity) CanAccessField(string) bool { return true } +func (allowAllSecurity) RowFilter() string { return "" } +func (allowAllSecurity) QueryFilter() *runtimev1.Expression { return nil } + func TestEmptyRows(t *testing.T) { testmode.Expensive(t) _, olap := acquireTestBigQuery(t) diff --git a/runtime/drivers/clickhouse/dialect.go b/runtime/drivers/clickhouse/dialect.go index b32f9622624b..42f418040707 100644 --- a/runtime/drivers/clickhouse/dialect.go +++ b/runtime/drivers/clickhouse/dialect.go @@ -80,9 +80,9 @@ func (d *dialect) AutoUnnest(expr string) string { return fmt.Sprintf("arrayJoin(%s)", expr) } -func (d *dialect) RequiresArrayContainsForInOperator() bool { return true } - -func (d *dialect) GetArrayContainsFunction() (string, error) { return "hasAny", nil } +func (d *dialect) ArrayContainsAnyExpression(arrExpr, valuesExpr string) (expr string, ok bool) { + return fmt.Sprintf("hasAny(%s, [%s])", arrExpr, valuesExpr), true +} func (d *dialect) CastToDataType(typ runtimev1.Type_Code) (string, error) { switch typ { diff --git a/runtime/drivers/databricks/dialect.go b/runtime/drivers/databricks/dialect.go index b9b216f59f6a..c93965ef2072 100644 --- a/runtime/drivers/databricks/dialect.go +++ b/runtime/drivers/databricks/dialect.go @@ -71,14 +71,26 @@ func (d *dialect) DimensionSelect(escapeTable string, dim *runtimev1.MetricsView return sel, fmt.Sprintf(` LATERAL VIEW EXPLODE(%s) %s AS %s`, dim.Expression, unnestTableName, unnestColName), nil } +// LateralUnnest uses tuple style so the element is referenced as tableAlias.colName; an unqualified colName is ambiguous when it matches the source column. func (d *dialect) LateralUnnest(expr, tableAlias, colName string) (tbl string, tupleStyle, auto bool, err error) { - return fmt.Sprintf(`LATERAL VIEW EXPLODE(%s) %s AS %s`, expr, tableAlias, d.EscapeIdentifier(colName)), false, false, nil + return fmt.Sprintf(`LATERAL VIEW EXPLODE(%s) %s AS %s`, expr, tableAlias, d.EscapeIdentifier(colName)), true, false, nil } func (d *dialect) UnnestSQLSuffix(tbl string) string { return fmt.Sprintf(" %s", tbl) } +// ArrayAnyExpression wraps EXISTS in COALESCE: it returns NULL rather than false when no element matches and some element is NULL, +// which would otherwise make negated filters drop the row. +func (d *dialect) ArrayAnyExpression(arrExpr, elemAlias string) (open, elem, closing string, ok bool) { + return fmt.Sprintf("COALESCE(EXISTS(%s, %s -> ", arrExpr, elemAlias), elemAlias, "), FALSE)", true +} + +// ArrayContainsAnyExpression wraps arrays_overlap in COALESCE for the same reason as ArrayAnyExpression. +func (d *dialect) ArrayContainsAnyExpression(arrExpr, valuesExpr string) (expr string, ok bool) { + return fmt.Sprintf("COALESCE(arrays_overlap(%s, array(%s)), FALSE)", arrExpr, valuesExpr), true +} + func (d *dialect) DateTruncExpr(dim *runtimev1.MetricsViewSpec_Dimension, grain runtimev1.TimeGrain, tz string, firstDayOfWeek, firstMonthOfYear int) (string, error) { if tz == "UTC" || tz == "Etc/UTC" { tz = "" diff --git a/runtime/drivers/databricks/olap_test.go b/runtime/drivers/databricks/olap_test.go index 11f7b7c0d5c4..116978fc16ab 100644 --- a/runtime/drivers/databricks/olap_test.go +++ b/runtime/drivers/databricks/olap_test.go @@ -1,6 +1,11 @@ package databricks_test import ( + "context" + "fmt" + "github.com/google/uuid" + runtimev1 "github.com/rilldata/rill/proto/gen/rill/runtime/v1" + "github.com/rilldata/rill/runtime/metricsview" "strings" "testing" "time" @@ -170,6 +175,148 @@ func TestQuerySchema(t *testing.T) { require.Equal(t, "string_col", schema.Fields[1].Name) } +func TestUnnestDimension(t *testing.T) { + t.Skip("skipping due to inactive Databricks account") + testmode.Expensive(t) + _, olap := acquireTestDatabricks(t) + + // Rows with overlapping, empty, NULL-element and NULL arrays, so joins that duplicate source rows and NULL handling are detectable. + // The table is created in the DSN's current schema. + name := "test_unnest_" + uuid.New().String()[:8] + t.Cleanup(func() { + err := olap.Exec(context.Background(), &drivers.Statement{Query: "DROP TABLE IF EXISTS " + name}) + require.NoError(t, err) + }) + err := olap.Exec(t.Context(), &drivers.Statement{Query: "CREATE TABLE " + name + " AS SELECT CAST(id AS BIGINT) AS id, tags FROM VALUES (1, array('a', 'b')), (2, array('b')), (3, array('c')), (4, CAST(array() AS ARRAY)), (5, array('c', NULL)), (6, CAST(NULL AS ARRAY)) AS t(id, tags)"}) + require.NoError(t, err) + + mv := &runtimev1.MetricsViewSpec{ + Table: name, + Dimensions: []*runtimev1.MetricsViewSpec_Dimension{ + {Name: "tags", Column: "tags", Unnest: true}, + {Name: "id", Column: "id"}, + }, + Measures: []*runtimev1.MetricsViewSpec_Measure{ + {Name: "count", Expression: "count(*)", Type: runtimev1.MetricsViewSpec_MEASURE_TYPE_SIMPLE}, + }, + } + + // Same query shape as the executor's dimension validation. + dialect := olap.Dialect() + escapeTable := dialect.EscapeTable(mv.Database, mv.DatabaseSchema, mv.Table) + sel, unnestClause, err := dialect.DimensionSelect(escapeTable, mv.Dimensions[0]) + require.NoError(t, err) + err = olap.Exec(t.Context(), &drivers.Statement{Query: fmt.Sprintf("SELECT %s FROM %s %s GROUP BY 1", sel, escapeTable, unnestClause), DryRun: true}) + require.NoError(t, err) + + tagsFilter := func(op metricsview.Operator, val any) *metricsview.Expression { + return &metricsview.Expression{Condition: &metricsview.Condition{ + Operator: op, + Expressions: []*metricsview.Expression{{Name: "tags"}, {Value: val}}, + }} + } + count := func(where *metricsview.Expression) *metricsview.Query { + return &metricsview.Query{Measures: []metricsview.Measure{{Name: "count"}}, Where: where} + } + + tests := []struct { + name string + qry *metricsview.Query + want []map[string]any + }{ + { + name: "group by unnest dimension", + qry: &metricsview.Query{ + Dimensions: []metricsview.Dimension{{Name: "tags"}}, + Measures: []metricsview.Measure{{Name: "count"}}, + Sort: []metricsview.Sort{{Name: "tags"}}, + }, + want: []map[string]any{ + {"tags": "a", "count": int64(1)}, + {"tags": "b", "count": int64(2)}, + {"tags": "c", "count": int64(2)}, + {"tags": nil, "count": int64(1)}, + }, + }, + { + // Row 1 matches both values but must be counted once. + name: "in filter counts each source row once", + qry: count(tagsFilter(metricsview.OperatorIn, []any{"a", "b"})), + want: []map[string]any{{"count": int64(2)}}, + }, + { + // Excludes rows containing 'a' even if they also contain other values; keeps the empty, NULL-element and NULL arrays. + name: "nin filter excludes rows containing any listed value", + qry: count(tagsFilter(metricsview.OperatorNin, []any{"a"})), + want: []map[string]any{{"count": int64(5)}}, + }, + { + name: "eq filter", + qry: count(tagsFilter(metricsview.OperatorEq, "b")), + want: []map[string]any{{"count": int64(2)}}, + }, + { + // Keeps the NULL-element and NULL arrays: a NULL comparison must not be treated as a match. + name: "neq filter excludes rows containing the value", + qry: count(tagsFilter(metricsview.OperatorNeq, "b")), + want: []map[string]any{{"count": int64(4)}}, + }, + { + name: "ilike filter", + qry: count(tagsFilter(metricsview.OperatorIlike, "%B%")), + want: []map[string]any{{"count": int64(2)}}, + }, + { + name: "eq filter with no match", + qry: count(tagsFilter(metricsview.OperatorEq, "missing")), + want: []map[string]any{{"count": int64(0)}}, + }, + { + name: "filter combined with group by on another dimension", + qry: &metricsview.Query{ + Dimensions: []metricsview.Dimension{{Name: "id"}}, + Measures: []metricsview.Measure{{Name: "count"}}, + Where: tagsFilter(metricsview.OperatorIn, []any{"a", "b"}), + Sort: []metricsview.Sort{{Name: "id"}}, + }, + want: []map[string]any{ + {"id": int64(1), "count": int64(1)}, + {"id": int64(2), "count": int64(1)}, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.qry.MetricsView = mv.Table + ast, err := metricsview.NewAST(mv, allowAllSecurity{}, tt.qry, dialect) + require.NoError(t, err) + sql, args, err := ast.SQL() + require.NoError(t, err) + require.Equal(t, tt.want, queryRows(t, olap, sql, args)) + }) + } +} + +func queryRows(t *testing.T, olap drivers.OLAPStore, query string, args []any) []map[string]any { + rows, err := olap.Query(t.Context(), &drivers.Statement{Query: query, Args: args}) + require.NoError(t, err) + defer rows.Close() + var res []map[string]any + for rows.Next() { + row := make(map[string]any) + require.NoError(t, rows.MapScan(row)) + res = append(res, row) + } + require.NoError(t, rows.Err()) + return res +} + +type allowAllSecurity struct{} + +func (allowAllSecurity) CanAccessField(string) bool { return true } +func (allowAllSecurity) RowFilter() string { return "" } +func (allowAllSecurity) QueryFilter() *runtimev1.Expression { return nil } + func acquireTestDatabricks(t *testing.T) (drivers.Handle, drivers.OLAPStore) { cfg := testruntime.AcquireConnector(t, "databricks") conn, err := drivers.Open("databricks", "", "default", cfg, storage.MustNew(t.TempDir(), nil), activity.NewNoopClient(), zap.NewNop()) diff --git a/runtime/drivers/dialect.go b/runtime/drivers/dialect.go index 445f2ee5323a..559900d1d0c9 100644 --- a/runtime/drivers/dialect.go +++ b/runtime/drivers/dialect.go @@ -47,10 +47,19 @@ type Dialect interface { GetCastExprForLike() string SupportsRegexMatch() bool GetRegexMatchFunction() (string, error) - RequiresArrayContainsForInOperator() bool - GetArrayContainsFunction() (string, error) + // ArrayContainsAnyExpression returns an expression that is true if the array arrExpr contains any of the comma-separated valuesExpr. + // ok is false if the dialect has no such expression, in which case the condition is evaluated against the unnested elements instead. + ArrayContainsAnyExpression(arrExpr, valuesExpr string) (expr string, ok bool) DimensionSelect(escapeTable string, dim *runtimev1.MetricsViewSpec_Dimension) (dimSelect, unnestClause string, err error) + // LateralUnnest returns the join clause that unnests expr. If tupleStyle is false the element is referenced by colName alone, + // and the dialect must implement ArrayAnyExpression since it cannot be referenced from a correlated subquery. LateralUnnest(expr, tableAlias, colName string) (tbl string, tupleStyle, auto bool, err error) + // UnnestedColumn returns the expression for the array element exposed by LateralUnnest in tuple style. + UnnestedColumn(tableAlias, colName string) string + // ArrayAnyExpression returns fragments for a condition that is true if any element of arrExpr satisfies it. + // The condition on a single element is written between open and close and references the element as elem. + // ok is false if the dialect has no such expression, in which case a correlated EXISTS subquery over LateralUnnest is used where possible. + ArrayAnyExpression(arrExpr, elemAlias string) (open, elem, closing string, ok bool) UnnestSQLSuffix(tbl string) string // AutoUnnest wraps an expression so the dialect unnests it automatically (used when LateralUnnest reports auto == true). AutoUnnest(expr string) string @@ -218,6 +227,14 @@ func (b *BaseDialect) LateralUnnest(expr, tableAlias, colName string) (tbl strin return fmt.Sprintf(`LATERAL UNNEST(%s) %s(%s)`, expr, tableAlias, b.escapeIdentifier(colName)), true, false, nil } +func (b *BaseDialect) UnnestedColumn(tableAlias, colName string) string { + return b.EscapeMember(tableAlias, colName) +} + +func (b *BaseDialect) ArrayAnyExpression(_, _ string) (open, elem, closing string, ok bool) { + return "", "", "", false +} + func (b *BaseDialect) UnnestSQLSuffix(tbl string) string { return fmt.Sprintf(", %s", tbl) } @@ -226,12 +243,8 @@ func (b *BaseDialect) AutoUnnest(expr string) string { return expr } -func (b *BaseDialect) RequiresArrayContainsForInOperator() bool { - return false -} - -func (b *BaseDialect) GetArrayContainsFunction() (string, error) { - return "", fmt.Errorf("array contains not supported for %s dialect", b.String()) +func (b *BaseDialect) ArrayContainsAnyExpression(_, _ string) (expr string, ok bool) { + return "", false } func (b *BaseDialect) MetricsViewDimensionExpression(dimension *runtimev1.MetricsViewSpec_Dimension) (string, error) { diff --git a/runtime/drivers/duckdb/dialect.go b/runtime/drivers/duckdb/dialect.go index 3e0cc075ab24..c7fa2ea15977 100644 --- a/runtime/drivers/duckdb/dialect.go +++ b/runtime/drivers/duckdb/dialect.go @@ -38,9 +38,9 @@ func (d *dialect) EscapeTable(db, schema, table string) string { return d.EscapeIdentifier(table) } -func (d *dialect) RequiresArrayContainsForInOperator() bool { return true } - -func (d *dialect) GetArrayContainsFunction() (string, error) { return "list_has_any", nil } +func (d *dialect) ArrayContainsAnyExpression(arrExpr, valuesExpr string) (expr string, ok bool) { + return fmt.Sprintf("list_has_any(%s, [%s])", arrExpr, valuesExpr), true +} func (d *dialect) OrderByExpression(name string, desc bool) string { res := d.EscapeIdentifier(name) diff --git a/runtime/drivers/snowflake/dialect.go b/runtime/drivers/snowflake/dialect.go index a4174c45e59b..486ba532aecb 100644 --- a/runtime/drivers/snowflake/dialect.go +++ b/runtime/drivers/snowflake/dialect.go @@ -43,6 +43,42 @@ func (d *dialect) OrderByAliasExpression(name string, desc bool) string { return res } +func (d *dialect) DimensionSelect(_ string, dim *runtimev1.MetricsViewSpec_Dimension) (dimSelect, unnestClause string, err error) { + expr, err := d.MetricsViewDimensionExpression(dim) + if err != nil { + return "", "", fmt.Errorf("failed to get dimension expression: %w", err) + } + alias := d.EscapeAlias(dim.Name) + if !dim.Unnest { + return fmt.Sprintf(`(%s) AS %s`, expr, alias), "", nil + } + unnestCol := drivers.TempName(fmt.Sprintf("unnested_%s_", dim.Name)) + tableAlias := drivers.TempName("tbl") + tbl, _, _, err := d.LateralUnnest(expr, tableAlias, unnestCol) + if err != nil { + return "", "", err + } + return fmt.Sprintf(`%s AS %s`, d.UnnestedColumn(tableAlias, unnestCol), alias), ", " + tbl, nil +} + +// LateralUnnest aliases every FLATTEN output column so the element is addressable as tableAlias.colName. +// FLATTEN cannot be wrapped in an inline view because Snowflake does not resolve the outer array column inside it. +func (d *dialect) LateralUnnest(expr, tableAlias, colName string) (tbl string, tupleStyle, auto bool, err error) { + return fmt.Sprintf(`LATERAL FLATTEN(INPUT => %s) %s (seq, key, path, index, %s, this)`, expr, tableAlias, d.EscapeIdentifier(colName)), true, false, nil +} + +// UnnestedColumn casts the element to VARCHAR. FLATTEN yields VARIANT elements for semi-structured arrays, which the driver returns as JSON-encoded text. +func (d *dialect) UnnestedColumn(tableAlias, colName string) string { + return d.EscapeMember(tableAlias, colName) + "::VARCHAR" +} + +// ArrayAnyExpression uses FILTER because Snowflake rejects correlated FLATTEN inside EXISTS subqueries. +// It also serves IN filters: ARRAYS_OVERLAP does not accept structured arrays and compares raw VARIANT elements, which would not match the VARCHAR values shown by UnnestedColumn. +// FILTER(NULL, ...) is NULL, so the result is coalesced to FALSE to keep rows with a NULL array under negated filters. +func (d *dialect) ArrayAnyExpression(arrExpr, elemAlias string) (open, elem, closing string, ok bool) { + return fmt.Sprintf("COALESCE(ARRAY_SIZE(FILTER(%s, %s -> ", arrExpr, elemAlias), elemAlias + "::VARCHAR", ")) > 0, FALSE)", true +} + func (d *dialect) DateTruncExpr(dim *runtimev1.MetricsViewSpec_Dimension, grain runtimev1.TimeGrain, tz string, firstDayOfWeek, firstMonthOfYear int) (string, error) { if tz == "UTC" || tz == "Etc/UTC" { tz = "" diff --git a/runtime/drivers/snowflake/olap_test.go b/runtime/drivers/snowflake/olap_test.go index 0ab31bca2b0f..e10be5a9bee9 100644 --- a/runtime/drivers/snowflake/olap_test.go +++ b/runtime/drivers/snowflake/olap_test.go @@ -1,12 +1,17 @@ package snowflake_test import ( + "context" "encoding/json" + "fmt" "strings" "testing" "time" + "github.com/google/uuid" + runtimev1 "github.com/rilldata/rill/proto/gen/rill/runtime/v1" "github.com/rilldata/rill/runtime/drivers" + "github.com/rilldata/rill/runtime/metricsview" "github.com/rilldata/rill/runtime/pkg/activity" "github.com/rilldata/rill/runtime/storage" "github.com/rilldata/rill/runtime/testruntime" @@ -184,6 +189,149 @@ func TestDryRun(t *testing.T) { require.NoError(t, err) } +// TestUnnestDimension creates a table in the DSN's current database and schema, so the DSN must point at a writable schema. +func TestUnnestDimension(t *testing.T) { + t.Skip("skipping due to inactive Snowflake account") + testmode.Expensive(t) + _, olap := acquireTestSnowflake(t) + + // Rows with overlapping, empty, NULL-element and NULL arrays, so joins that duplicate source rows and NULL handling are detectable. + // The driver returns NUMBER columns as strings. + name := "test_unnest_" + uuid.New().String()[:8] + t.Cleanup(func() { + err := olap.Exec(context.Background(), &drivers.Statement{Query: "DROP TABLE IF EXISTS " + name}) + require.NoError(t, err) + }) + err := olap.Exec(t.Context(), &drivers.Statement{Query: "CREATE TABLE " + name + " AS SELECT 1 AS id, ARRAY_CONSTRUCT('a', 'b') AS tags UNION ALL SELECT 2, ARRAY_CONSTRUCT('b') UNION ALL SELECT 3, ARRAY_CONSTRUCT('c') UNION ALL SELECT 4, ARRAY_CONSTRUCT() UNION ALL SELECT 5, ARRAY_CONSTRUCT('c', NULL) UNION ALL SELECT 6, NULL::ARRAY"}) + require.NoError(t, err) + + mv := &runtimev1.MetricsViewSpec{ + Table: name, + Dimensions: []*runtimev1.MetricsViewSpec_Dimension{ + {Name: "tags", Column: "tags", Unnest: true}, + {Name: "id", Column: "id"}, + }, + Measures: []*runtimev1.MetricsViewSpec_Measure{ + {Name: "count", Expression: "count(*)", Type: runtimev1.MetricsViewSpec_MEASURE_TYPE_SIMPLE}, + }, + } + + // Same query shape as the executor's dimension validation. + dialect := olap.Dialect() + escapeTable := dialect.EscapeTable(mv.Database, mv.DatabaseSchema, mv.Table) + sel, unnestClause, err := dialect.DimensionSelect(escapeTable, mv.Dimensions[0]) + require.NoError(t, err) + err = olap.Exec(t.Context(), &drivers.Statement{Query: fmt.Sprintf("SELECT %s FROM %s %s GROUP BY 1", sel, escapeTable, unnestClause), DryRun: true}) + require.NoError(t, err) + + tagsFilter := func(op metricsview.Operator, val any) *metricsview.Expression { + return &metricsview.Expression{Condition: &metricsview.Condition{ + Operator: op, + Expressions: []*metricsview.Expression{{Name: "tags"}, {Value: val}}, + }} + } + count := func(where *metricsview.Expression) *metricsview.Query { + return &metricsview.Query{Measures: []metricsview.Measure{{Name: "count"}}, Where: where} + } + + tests := []struct { + name string + qry *metricsview.Query + want []map[string]any + }{ + { + name: "group by unnest dimension", + qry: &metricsview.Query{ + Dimensions: []metricsview.Dimension{{Name: "tags"}}, + Measures: []metricsview.Measure{{Name: "count"}}, + Sort: []metricsview.Sort{{Name: "tags"}}, + }, + want: []map[string]any{ + {"tags": "a", "count": "1"}, + {"tags": "b", "count": "2"}, + // FLATTEN does not emit a row for a NULL element. + {"tags": "c", "count": "2"}, + }, + }, + { + // Row 1 matches both values but must be counted once. + name: "in filter counts each source row once", + qry: count(tagsFilter(metricsview.OperatorIn, []any{"a", "b"})), + want: []map[string]any{{"count": "2"}}, + }, + { + // Excludes rows containing 'a' even if they also contain other values; keeps the empty, NULL-element and NULL arrays. + name: "nin filter excludes rows containing any listed value", + qry: count(tagsFilter(metricsview.OperatorNin, []any{"a"})), + want: []map[string]any{{"count": "5"}}, + }, + { + name: "eq filter", + qry: count(tagsFilter(metricsview.OperatorEq, "b")), + want: []map[string]any{{"count": "2"}}, + }, + { + // Keeps the NULL-element and NULL arrays: a NULL comparison must not be treated as a match. + name: "neq filter excludes rows containing the value", + qry: count(tagsFilter(metricsview.OperatorNeq, "b")), + want: []map[string]any{{"count": "4"}}, + }, + { + name: "ilike filter", + qry: count(tagsFilter(metricsview.OperatorIlike, "%B%")), + want: []map[string]any{{"count": "2"}}, + }, + { + name: "eq filter with no match", + qry: count(tagsFilter(metricsview.OperatorEq, "missing")), + want: []map[string]any{{"count": "0"}}, + }, + { + name: "filter combined with group by on another dimension", + qry: &metricsview.Query{ + Dimensions: []metricsview.Dimension{{Name: "id"}}, + Measures: []metricsview.Measure{{Name: "count"}}, + Where: tagsFilter(metricsview.OperatorIn, []any{"a", "b"}), + Sort: []metricsview.Sort{{Name: "id"}}, + }, + want: []map[string]any{ + {"id": "1", "count": "1"}, + {"id": "2", "count": "1"}, + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.qry.MetricsView = name + ast, err := metricsview.NewAST(mv, allowAllSecurity{}, tt.qry, dialect) + require.NoError(t, err) + sql, args, err := ast.SQL() + require.NoError(t, err) + require.Equal(t, tt.want, queryRows(t, olap, sql, args)) + }) + } +} + +func queryRows(t *testing.T, olap drivers.OLAPStore, query string, args []any) []map[string]any { + rows, err := olap.Query(t.Context(), &drivers.Statement{Query: query, Args: args}) + require.NoError(t, err) + defer rows.Close() + var res []map[string]any + for rows.Next() { + row := make(map[string]any) + require.NoError(t, rows.MapScan(row)) + res = append(res, row) + } + require.NoError(t, rows.Err()) + return res +} + +type allowAllSecurity struct{} + +func (allowAllSecurity) CanAccessField(string) bool { return true } +func (allowAllSecurity) RowFilter() string { return "" } +func (allowAllSecurity) QueryFilter() *runtimev1.Expression { return nil } + func acquireTestSnowflake(t *testing.T) (drivers.Handle, drivers.OLAPStore) { cfg := testruntime.AcquireConnector(t, "snowflake") conn, err := drivers.Open("snowflake", "", "default", cfg, storage.MustNew(t.TempDir(), nil), activity.NewNoopClient(), zap.NewNop()) diff --git a/runtime/metricsview/ast.go b/runtime/metricsview/ast.go index 49b823dde57c..782a887c8dd6 100644 --- a/runtime/metricsview/ast.go +++ b/runtime/metricsview/ast.go @@ -266,7 +266,7 @@ func NewAST(mv *runtimev1.MetricsViewSpec, sec MetricsViewSecurity, qry *Query, } else { ast.unnests = append(ast.unnests, tblWithAlias) if tupleStyle { - f.Expr = ast.Dialect.EscapeMember(unnestAlias, f.Name) + f.Expr = ast.Dialect.UnnestedColumn(unnestAlias, f.Name) } else { f.Expr = ast.Dialect.EscapeMember("", f.Name) } diff --git a/runtime/metricsview/ast_unnest_test.go b/runtime/metricsview/ast_unnest_test.go new file mode 100644 index 000000000000..1bb1d95686ce --- /dev/null +++ b/runtime/metricsview/ast_unnest_test.go @@ -0,0 +1,181 @@ +package metricsview + +import ( + "testing" + + runtimev1 "github.com/rilldata/rill/proto/gen/rill/runtime/v1" + "github.com/rilldata/rill/runtime/drivers" + "github.com/rilldata/rill/runtime/drivers/bigquery" + "github.com/rilldata/rill/runtime/drivers/clickhouse" + "github.com/rilldata/rill/runtime/drivers/databricks" + "github.com/rilldata/rill/runtime/drivers/druid" + "github.com/rilldata/rill/runtime/drivers/duckdb" + "github.com/rilldata/rill/runtime/drivers/snowflake" + "github.com/stretchr/testify/require" +) + +func TestUnnestSQL(t *testing.T) { + mv := &runtimev1.MetricsViewSpec{ + Table: "test_table", + Dimensions: []*runtimev1.MetricsViewSpec_Dimension{ + {Name: "tags", Column: "tags", Unnest: true}, + {Name: "city", Column: "city"}, + }, + Measures: []*runtimev1.MetricsViewSpec_Measure{ + {Name: "count", Expression: "count(*)", Type: runtimev1.MetricsViewSpec_MEASURE_TYPE_SIMPLE}, + }, + } + sec := skipMetricsViewSecurity{} + + tagsEqA := &Expression{Condition: &Condition{ + Operator: OperatorEq, + Expressions: []*Expression{{Name: "tags"}, {Value: "a"}}, + }} + + tests := []struct { + name string + dialect drivers.Dialect + dims []Dimension + where *Expression + wantSQL string + wantArgs []any + }{ + { + name: "bigquery: group by unnest dim", + dialect: bigquery.DialectBigQuery, + dims: []Dimension{{Name: "tags"}}, + wantSQL: "SELECT (`tags`) AS `tags`, (count(*)) AS `count` FROM `test_table`, UNNEST(`tags`) AS `tags` GROUP BY 1", + }, + { + name: "bigquery: filter on unnest dim not in select", + dialect: bigquery.DialectBigQuery, + dims: []Dimension{{Name: "city"}}, + where: tagsEqA, + wantSQL: "SELECT (`city`) AS `city`, (count(*)) AS `count` FROM `test_table` WHERE EXISTS (SELECT 1 FROM UNNEST(`tags`) AS `t0` WHERE ((`t0`) = ?)) GROUP BY 1", + wantArgs: []any{"a"}, + }, + { + name: "snowflake: group by unnest dim", + dialect: snowflake.DialectSnowflake, + dims: []Dimension{{Name: "tags"}}, + wantSQL: `SELECT (t0.tags::VARCHAR) AS "tags", (count(*)) AS "count" FROM test_table, LATERAL FLATTEN(INPUT => tags) t0 (seq, key, path, index, tags, this) GROUP BY 1`, + }, + { + name: "snowflake: filter on unnest dim not in select", + dialect: snowflake.DialectSnowflake, + dims: []Dimension{{Name: "city"}}, + where: tagsEqA, + wantSQL: `SELECT (city) AS "city", (count(*)) AS "count" FROM test_table WHERE COALESCE(ARRAY_SIZE(FILTER(tags, t0 -> ((t0::VARCHAR) = ?))) > 0, FALSE) GROUP BY 1`, + wantArgs: []any{"a"}, + }, + { + name: "snowflake: nin filter on unnest dim not in select", + dialect: snowflake.DialectSnowflake, + dims: []Dimension{{Name: "city"}}, + where: &Expression{Condition: &Condition{ + Operator: OperatorNin, + Expressions: []*Expression{{Name: "tags"}, {Value: []any{"a", "b"}}}, + }}, + wantSQL: `SELECT (city) AS "city", (count(*)) AS "count" FROM test_table WHERE NOT COALESCE(ARRAY_SIZE(FILTER(tags, t0 -> ((t0::VARCHAR) IN (?,?)))) > 0, FALSE) GROUP BY 1`, + wantArgs: []any{"a", "b"}, + }, + { + name: "databricks: group by unnest dim", + dialect: databricks.DialectDatabricks, + dims: []Dimension{{Name: "tags"}}, + wantSQL: "SELECT (`t0`.`tags`) AS `tags`, (count(*)) AS `count` FROM `test_table` LATERAL VIEW EXPLODE(`tags`) t0 AS `tags` GROUP BY 1", + }, + { + name: "databricks: filter on unnest dim not in select", + dialect: databricks.DialectDatabricks, + dims: []Dimension{{Name: "city"}}, + where: tagsEqA, + wantSQL: "SELECT (`city`) AS `city`, (count(*)) AS `count` FROM `test_table` WHERE COALESCE(EXISTS(`tags`, t0 -> ((t0) = ?)), FALSE) GROUP BY 1", + wantArgs: []any{"a"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + qry := &Query{ + MetricsView: "test", + Dimensions: tt.dims, + Measures: []Measure{{Name: "count"}}, + Where: tt.where, + } + + ast, err := NewAST(mv, sec, qry, tt.dialect) + require.NoError(t, err) + + sql, args, err := ast.SQL() + require.NoError(t, err) + require.Equal(t, tt.wantSQL, sql) + require.Equal(t, tt.wantArgs, args) + }) + } +} + +// Subquery filters remain supported for scalar dimensions and existing general unnest paths. +func TestUnnestSubqueryFilterSQL(t *testing.T) { + mv := &runtimev1.MetricsViewSpec{ + Table: "test_table", + Dimensions: []*runtimev1.MetricsViewSpec_Dimension{ + {Name: "tags", Column: "tags", Unnest: true}, + {Name: "city", Column: "city"}, + }, + Measures: []*runtimev1.MetricsViewSpec_Measure{ + {Name: "count", Expression: "count(*)", Type: runtimev1.MetricsViewSpec_MEASURE_TYPE_SIMPLE}, + }, + } + base := drivers.NewBaseDialect(drivers.DialectNamePostgres, drivers.DoubleQuotesEscapeIdentifier, drivers.DoubleQuotesEscapeIdentifier) + tests := []struct { + dialect drivers.Dialect + wantErr string + }{ + {duckdb.DialectDuckDB, "the right value must be a list of values for an array IN condition"}, + {clickhouse.DialectClickhouse, "the right value must be a list of values for an array IN condition"}, + {databricks.DialectDatabricks, "the right value must be a list of values for an array IN condition"}, + {snowflake.DialectSnowflake, `dialect snowflake does not support subquery filters on unnest dimension "tags"`}, + {bigquery.DialectBigQuery, `dialect bigquery does not support subquery filters on unnest dimension "tags"`}, + {druid.DialectDruid, ""}, + {&base, ""}, + } + for _, tt := range tests { + for _, op := range []Operator{OperatorIn, OperatorNin} { + for _, shape := range []struct { + name string + dim string + dims []Dimension + }{ + {"unselected unnest dimension", "tags", []Dimension{{Name: "city"}}}, + {"selected unnest dimension", "tags", []Dimension{{Name: "tags"}}}, + {"scalar dimension", "city", []Dimension{{Name: "city"}}}, + } { + t.Run(tt.dialect.String()+"/"+string(op)+"/"+shape.name, func(t *testing.T) { + where := &Expression{Condition: &Condition{ + Operator: op, + Expressions: []*Expression{ + {Name: shape.dim}, + {Subquery: &Subquery{ + Dimension: Dimension{Name: shape.dim}, + Measures: []Measure{{Name: "count"}}, + Having: &Expression{Condition: &Condition{Operator: OperatorGt, Expressions: []*Expression{{Name: "count"}, {Value: 10}}}}, + }}, + }, + }} + qry := &Query{MetricsView: "test", Dimensions: shape.dims, Measures: []Measure{{Name: "count"}}, Where: where} + ast, err := NewAST(mv, skipMetricsViewSecurity{}, qry, tt.dialect) + if shape.name == "unselected unnest dimension" && tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + return + } + require.NoError(t, err) + sql, args, err := ast.SQL() + require.NoError(t, err) + require.Contains(t, sql, " IN (SELECT "+tt.dialect.EscapeAlias(shape.dim)+" FROM (") + require.Equal(t, []any{10}, args) + }) + } + } + } +} diff --git a/runtime/metricsview/astexpr.go b/runtime/metricsview/astexpr.go index a3eb41e39dd9..a626274f829f 100644 --- a/runtime/metricsview/astexpr.go +++ b/runtime/metricsview/astexpr.go @@ -119,7 +119,7 @@ func (b *sqlExprBuilder) writeSubquery(sub *Subquery) error { // Output: (SELECT FROM ()) b.writeString("(SELECT ") - b.writeString(b.ast.Dialect.EscapeIdentifier(sub.Dimension.Name)) + b.writeString(b.ast.Dialect.EscapeAlias(sub.Dimension.Name)) b.writeString(" FROM (") b.writeString(sql) b.writeString("))") @@ -277,10 +277,12 @@ func (b *sqlExprBuilder) writeBinaryCondition(exprs []*Expression, op Operator) return b.writeBinaryConditionInner(nil, right, leftExpr, op) } - // For IN/NIN on unnest dimensions backed by DuckDB or ClickHouse, use native array-contains - // functions (list_has_any / hasAny). This avoids double-counting when a row's array contains multiple matching values. - if (op == OperatorIn || op == OperatorNin) && b.ast.Dialect.RequiresArrayContainsForInOperator() { - return b.writeArrayContainsCondition(leftExpr, right, op == OperatorNin) + // For IN/NIN on unnest dimensions, prefer a native array-contains expression over an unnest join where the dialect supports it. + // It avoids scanning the unnested rows and double-counting rows whose array contains multiple matching values. + if op == OperatorIn || op == OperatorNin { + if handled, err := b.writeArrayContainsCondition(leftExpr, right, op == OperatorNin); handled || err != nil { + return err + } } // Generate unnest join @@ -294,16 +296,19 @@ func (b *sqlExprBuilder) writeBinaryCondition(exprs []*Expression, op Operator) leftExpr = b.ast.Dialect.AutoUnnest(leftExpr) return b.writeBinaryConditionInner(nil, right, leftExpr, op) } - var unnestColAlias string - if tupleStyle { - unnestColAlias = b.ast.Dialect.EscapeMember(unnestTableAlias, left.Name) - } else { - unnestColAlias = b.ast.Dialect.EscapeAlias(left.Name) + // A filter on an unnest dimension that is not selected should match each source row once, even if several of its elements match. + // Prefer the dialect's native any-element expression, otherwise use a correlated EXISTS subquery over the unnest join. + open, elem, closing, ok := b.ast.Dialect.ArrayAnyExpression(leftExpr, unnestTableAlias) + if ok && right.Subquery != nil { + return fmt.Errorf("dialect %s does not support subquery filters on unnest dimension %q", b.ast.Dialect, left.Name) } - - if !tupleStyle { // if tupleStyle, then we cannot refer to the column by table alias - b.ast.unnests = append(b.ast.unnests, unnestFrom) - return b.writeBinaryConditionInner(nil, right, unnestColAlias, op) + if !ok && !tupleStyle { + return fmt.Errorf("dialect %s cannot filter on unnest dimension %q: it must support tuple-style unnest or an array any-element expression", b.ast.Dialect, left.Name) + } + if !ok { + open = "EXISTS (SELECT 1 FROM " + unnestFrom + " WHERE " + elem = b.ast.Dialect.UnnestedColumn(unnestTableAlias, left.Name) + closing = ")" } // Need to move "NOT" to outside of the subquery @@ -320,18 +325,15 @@ func (b *sqlExprBuilder) writeBinaryCondition(exprs []*Expression, op Operator) not = true } - // Output: [NOT] EXISTS (SELECT 1 FROM WHERE ) if not { b.writeString("NOT ") } - b.writeString("EXISTS (SELECT 1 FROM ") - b.writeString(unnestFrom) - b.writeString(" WHERE ") - err = b.writeBinaryConditionInner(nil, right, unnestColAlias, op) + b.writeString(open) + err = b.writeBinaryConditionInner(nil, right, elem, op) if err != nil { return err } - b.writeByte(')') + b.writeString(closing) return nil } @@ -659,46 +661,38 @@ func (b *sqlExprBuilder) writeInConditionForValues(left *Expression, leftOverrid return nil } -func (b *sqlExprBuilder) writeArrayContainsCondition(leftExpr string, right *Expression, not bool) error { - vals, ok := right.Value.([]any) - if !ok { - return fmt.Errorf("the right value must be a list of values for an array IN condition") - } - - if len(vals) == 0 { +// writeArrayContainsCondition writes a native array-contains condition for a list of values. +// It returns false without writing anything if the dialect has no such expression. +func (b *sqlExprBuilder) writeArrayContainsCondition(leftExpr string, right *Expression, not bool) (bool, error) { + vals, isList := right.Value.([]any) + if isList && len(vals) == 0 { if not { b.writeString("TRUE") } else { b.writeString("FALSE") } - return nil + return true, nil } - b.writeByte('(') - if not { - b.writeString("NOT ") + // NULL values in the list are not handled separately: ClickHouse's hasAny matches them, while DuckDB's list_has_any and Databricks' arrays_overlap ignore them. + // There is no reliable way to check for NULL elements; leftExpr IS NULL checks for a NULL array, not NULL elements. + expr, ok := b.ast.Dialect.ArrayContainsAnyExpression("("+leftExpr+")", strings.TrimSuffix(strings.Repeat("?,", len(vals)), ",")) + if !ok { + return false, nil } - arrayContainsFunc, err := b.ast.Dialect.GetArrayContainsFunction() - if err != nil { - return err + if !isList { + return false, fmt.Errorf("the right value must be a list of values for an array IN condition") } - b.writeString(arrayContainsFunc) + b.writeByte('(') - b.writeParenthesizedString(leftExpr) - b.writeString(", [") - // not handling NULL values separately as clickhouse hasAny function takes care of it however, duckdb ignores null values in the list_has_any function, but there is no reliable way to make it work, - // but even using leftExpr IS NULL does not solve the issue as it checks for null array rather than null values in the array. - for i, val := range vals { - if i > 0 { - b.writeByte(',') - } - b.writeString("?") - b.args = append(b.args, val) + if not { + b.writeString("NOT ") } - b.writeString("])") + b.writeString(expr) b.writeByte(')') + b.args = append(b.args, vals...) - return nil + return true, nil } func (b *sqlExprBuilder) writeByte(v byte) { diff --git a/runtime/metricsview/astexpr_test.go b/runtime/metricsview/astexpr_test.go index 5abe4b87c431..2ef354ed5758 100644 --- a/runtime/metricsview/astexpr_test.go +++ b/runtime/metricsview/astexpr_test.go @@ -6,7 +6,9 @@ import ( runtimev1 "github.com/rilldata/rill/proto/gen/rill/runtime/v1" "github.com/rilldata/rill/runtime/drivers" "github.com/rilldata/rill/runtime/drivers/clickhouse" + "github.com/rilldata/rill/runtime/drivers/databricks" "github.com/rilldata/rill/runtime/drivers/duckdb" + "github.com/rilldata/rill/runtime/drivers/snowflake" "github.com/stretchr/testify/require" ) @@ -138,6 +140,59 @@ func TestArrayContainsCondition(t *testing.T) { wantSQL: `(hasAny(("tags"), [?,?]))`, wantArgs: []any{nil, "a"}, }, + { + name: "databricks: in on unnest dim uses arrays_overlap", + dialect: databricks.DialectDatabricks, + where: &Expression{Condition: &Condition{ + Operator: OperatorIn, + Expressions: []*Expression{ + {Name: "tags"}, + {Value: []any{"a", "b"}}, + }, + }}, + wantSQL: "(COALESCE(arrays_overlap((`tags`), array(?,?)), FALSE))", + wantArgs: []any{"a", "b"}, + }, + { + name: "databricks: nin on unnest dim uses NOT arrays_overlap", + dialect: databricks.DialectDatabricks, + where: &Expression{Condition: &Condition{ + Operator: OperatorNin, + Expressions: []*Expression{ + {Name: "tags"}, + {Value: []any{"a", "b"}}, + }, + }}, + wantSQL: "(NOT COALESCE(arrays_overlap((`tags`), array(?,?)), FALSE))", + wantArgs: []any{"a", "b"}, + }, + { + name: "databricks: in on unnest dim already in select falls back to normal IN", + dialect: databricks.DialectDatabricks, + dims: []Dimension{{Name: "tags"}}, + where: &Expression{Condition: &Condition{ + Operator: OperatorIn, + Expressions: []*Expression{ + {Name: "tags"}, + {Value: []any{"a", "b"}}, + }, + }}, + wantSQL: "((`t0`.`tags`) IN (?,?))", + wantArgs: []any{"a", "b"}, + }, + { + name: "snowflake: in on unnest dim uses FILTER", + dialect: snowflake.DialectSnowflake, + where: &Expression{Condition: &Condition{ + Operator: OperatorIn, + Expressions: []*Expression{ + {Name: "tags"}, + {Value: []any{"a", "b"}}, + }, + }}, + wantSQL: "COALESCE(ARRAY_SIZE(FILTER(tags, t2 -> ((t2::VARCHAR) IN (?,?)))) > 0, FALSE)", + wantArgs: []any{"a", "b"}, + }, { name: "duckdb: in on non-unnest dim uses normal IN", dialect: duckdb.DialectDuckDB,