diff --git a/docs/src/js/interfaces/MaterializedViewDefinition.md b/docs/src/js/interfaces/MaterializedViewDefinition.md index a1ffb6e2b..68bcb1415 100644 --- a/docs/src/js/interfaces/MaterializedViewDefinition.md +++ b/docs/src/js/interfaces/MaterializedViewDefinition.md @@ -8,8 +8,8 @@ The query that defines a materialized view, as stored: `SELECT columns FROM [ns.]table [, function(args) AS alias | , UNNEST(column) AS alias] -[WHERE predicate] [LIMIT n]`. A Function in `FROM` position yields one row -per element it returns. +[WHERE predicate] [GROUP BY expr, ...] [LIMIT n]`. A Function in `FROM` +position yields one row per element it returns. ## Properties diff --git a/nodejs/__test__/materialized_view.test.ts b/nodejs/__test__/materialized_view.test.ts index 352628208..5e5c08745 100644 --- a/nodejs/__test__/materialized_view.test.ts +++ b/nodejs/__test__/materialized_view.test.ts @@ -54,7 +54,7 @@ describe("materialized views", () => { // A newer writer's layout is reported, never guessed at. for (const newer of [ - `{"format":2,"query":${JSON.stringify(query)}}`, + `{"format":3,"query":${JSON.stringify(query)}}`, '{"kind":"select_v3","source_table":"people"}', ]) { expect(() => read(newer)).toThrow(/cannot refresh/); diff --git a/nodejs/lancedb/materialized_view.ts b/nodejs/lancedb/materialized_view.ts index 15ae52274..51f8cfead 100644 --- a/nodejs/lancedb/materialized_view.ts +++ b/nodejs/lancedb/materialized_view.ts @@ -7,14 +7,17 @@ import { Table } from "./table"; /** Schema metadata key holding a materialized view's definition. */ export const DEFINITION_META_KEY = "mv.definition"; -/** The stored layout this version reads: `{"format": 1, "query": ""}`. */ -export const DEFINITION_FORMAT = 1; +/** + * The newest stored layout this version reads: `{"format": N, "query": ""}`, + * format 2 being a query with `GROUP BY`. + */ +export const DEFINITION_FORMAT = 2; /** * The query that defines a materialized view, as stored: * `SELECT columns FROM [ns.]table [, function(args) AS alias | , UNNEST(column) AS alias] - * [WHERE predicate] [LIMIT n]`. A Function in `FROM` position yields one row - * per element it returns. + * [WHERE predicate] [GROUP BY expr, ...] [LIMIT n]`. A Function in `FROM` + * position yields one row per element it returns. */ export interface MaterializedViewDefinition { /** The defining query, in the canonical spelling the server stores. */ diff --git a/python/python/lancedb/materialized_view.py b/python/python/lancedb/materialized_view.py index c487eb044..a836210cb 100644 --- a/python/python/lancedb/materialized_view.py +++ b/python/python/lancedb/materialized_view.py @@ -29,9 +29,10 @@ SelectArg = Union[ ] -DEFINITION_FORMAT = 1 -"""The stored layout this version reads: ``{"format": 1, "query": ""}``. -A ``kind`` key beside it is for readers older than the format number.""" +DEFINITION_FORMAT = 2 +"""The newest stored layout this version reads: ``{"format": N, "query": ""}``, +format 2 being a query with ``GROUP BY``. A ``kind`` key beside it is for +readers older than the format number.""" @dataclass @@ -40,7 +41,7 @@ class MaterializedViewDefinition: SELECT columns FROM [ns.]table [, function(args) AS alias | , UNNEST(column) AS alias] - [WHERE predicate] [LIMIT n] + [WHERE predicate] [GROUP BY expr, ...] [LIMIT n] A Function in ``FROM`` position yields one row per element it returns. """ diff --git a/python/python/tests/test_materialized_views.py b/python/python/tests/test_materialized_views.py index cfe2405a4..c030ac1df 100644 --- a/python/python/tests/test_materialized_views.py +++ b/python/python/tests/test_materialized_views.py @@ -522,7 +522,7 @@ def test_stored_queries_and_legacy_layouts_are_read(): # A newer writer's layout is reported, never guessed at. for newer in ( - {"format": 2, "query": query}, + {"format": 3, "query": query}, {"kind": "select_v3", "source_table": "people"}, ): with pytest.raises(NotImplementedError, match="cannot refresh"): diff --git a/rust/lancedb/src/materialized_view.rs b/rust/lancedb/src/materialized_view.rs index 74d831c82..067bde13d 100644 --- a/rust/lancedb/src/materialized_view.rs +++ b/rust/lancedb/src/materialized_view.rs @@ -9,6 +9,7 @@ //! query this version cannot maintain reads back as unrefreshable, not as a //! plain table. Queries, indexes and search work on the view unchanged. +mod grouped; mod query; pub mod refresh; @@ -79,13 +80,21 @@ const EMBEDDING_FUNCTIONS_META_KEY: &str = "embedding_functions"; /// produces, which is what lets a query embed its own text. const COLUMN_DEFINITIONS_META_KEY: &str = "lancedb::column_definitions"; -/// The layout this version writes under [`DEFINITION_META_KEY`]: -/// `{"format": 1, "query": ""}`, the query as -/// [`MaterializedViewDefinition::to_sql`] renders it. A reader refuses a -/// newer format rather than guess at it. The layout also carries -/// `"kind": "query"`, which readers older than the format number report -/// as an unrefreshable view instead of failing to read the metadata. -pub const DEFINITION_FORMAT: u64 = 1; +/// The newest layout this version reads under [`DEFINITION_META_KEY`]: +/// `{"format": N, "query": ""}`, the query as +/// [`MaterializedViewDefinition::to_sql`] renders it. Format 2 is a query +/// with `GROUP BY`; every other query is written as format 1, so a reader +/// that predates grouping reports a grouped view as unrefreshable rather +/// than failing to parse it. A reader refuses a newer format rather than +/// guess at it. The layout also carries `"kind": "query"`, which readers +/// older than the format number report as an unrefreshable view instead of +/// failing to read the metadata. +pub const DEFINITION_FORMAT: u64 = 2; + +/// The format `definition` is written in; see [`DEFINITION_FORMAT`]. +fn format_of(definition: &MaterializedViewDefinition) -> u64 { + if definition.is_grouped() { 2 } else { 1 } +} /// Legacy `kind` tag of the structured layout written before /// [`DEFINITION_FORMAT`] existed; still read, never written. @@ -240,6 +249,9 @@ pub struct MaterializedViewDefinition { pub projections: Vec, /// SQL predicate selecting the rows the view holds. pub filter: Option, + /// `GROUP BY` expressions. A grouped view holds one row per group and is + /// recomputed in full whenever its source changes. + pub group_by: Vec, /// Cap on the number of rows the view holds, in materialization order. pub limit: Option, } @@ -250,7 +262,7 @@ impl MaterializedViewDefinition { /// ```sql /// SELECT , ... /// FROM [ns.]table [, function(args) AS alias | , UNNEST(column) AS alias] - /// [WHERE predicate] [LIMIT n] + /// [WHERE predicate] [GROUP BY expr, ...] [LIMIT n] /// ``` /// /// A Function in `FROM` position yields one row per element it returns; @@ -283,6 +295,11 @@ impl MaterializedViewDefinition { definition_to_metadata(self) } + /// Whether the query has a `GROUP BY`. + pub fn is_grouped(&self) -> bool { + !self.group_by.is_empty() + } + /// Whether the query is `SELECT *`. pub fn selects_star(&self) -> bool { matches!(self.projections.as_slice(), [p] if *p == ViewProjection::star()) @@ -359,7 +376,7 @@ struct LegacyDefinition { pub(crate) fn definition_to_metadata(definition: &MaterializedViewDefinition) -> Result { Ok(serde_json::json!({ "kind": "query", - "format": DEFINITION_FORMAT, + "format": format_of(definition), "query": definition.to_sql(), }) .to_string()) @@ -389,6 +406,13 @@ pub fn read_definition(metadata: &HashMap) -> Result) -> Result err, })?; + if definition.is_grouped() { + return grouped::plan(source_schema, definition, filter); + } // Projections are typed against the source, or for an unnested view // against the source with the list column replaced by its element // under the alias, where `c.chunk` is an ordinary nested path. @@ -616,34 +644,7 @@ pub(crate) fn plan( } if let Some(filter) = filter.as_deref() { - let expr = planner - .parse_filter(filter) - .map_err(|e| Error::InvalidInput { - message: format!("invalid view filter: {e}"), - })?; - ensure_immutable(&expr, |message| Error::InvalidInput { - message: format!("invalid view filter: {message}"), - })?; - let filter_inputs = resolve_inputs(&source_schema, &expr, |message| Error::InvalidInput { - message: format!("invalid view filter: {message}"), - })?; - // A committed filter has to be usable as a predicate. - let read_schema = project_schema(&source_schema, &filter_inputs); - let data_type = Planner::new(read_schema.clone()) - .create_physical_expr(&expr) - .map_err(|e| Error::InvalidInput { - message: format!("invalid view filter: {e}"), - })? - .data_type(read_schema.as_ref()) - .map_err(|e| Error::InvalidInput { - message: format!("invalid view filter: {e}"), - })?; - if data_type != DataType::Boolean { - return Err(Error::InvalidInput { - message: format!("view filter must be a boolean predicate, not {data_type}"), - }); - } - inputs.extend(filter_inputs); + inputs.extend(plan_filter(&source_schema, filter)?); } let definition = MaterializedViewDefinition { @@ -655,6 +656,7 @@ pub(crate) fn plan( .map(|(output, expression)| ViewProjection { output, expression }) .collect(), filter, + group_by: Vec::new(), limit, }; let mut inputs: Vec = inputs @@ -671,6 +673,39 @@ pub(crate) fn plan( }) } +/// Check that `filter` is an immutable boolean predicate over `source_schema`, +/// returning the columns it reads. +pub(crate) fn plan_filter(source_schema: &SchemaRef, filter: &str) -> Result> { + let expr = Planner::new(source_schema.clone()) + .parse_filter(filter) + .map_err(|e| Error::InvalidInput { + message: format!("invalid view filter: {e}"), + })?; + ensure_immutable(&expr, |message| Error::InvalidInput { + message: format!("invalid view filter: {message}"), + })?; + let filter_inputs = resolve_inputs(source_schema, &expr, |message| Error::InvalidInput { + message: format!("invalid view filter: {message}"), + })?; + // A committed filter has to be usable as a predicate. + let read_schema = project_schema(source_schema, &filter_inputs); + let data_type = Planner::new(read_schema.clone()) + .create_physical_expr(&expr) + .map_err(|e| Error::InvalidInput { + message: format!("invalid view filter: {e}"), + })? + .data_type(read_schema.as_ref()) + .map_err(|e| Error::InvalidInput { + message: format!("invalid view filter: {e}"), + })?; + if data_type != DataType::Boolean { + return Err(Error::InvalidInput { + message: format!("view filter must be a boolean predicate, not {data_type}"), + }); + } + Ok(filter_inputs) +} + /// The schema a projection over an unnested view is planned against: the /// source's, with the list column replaced by its element type under the /// alias. @@ -1046,6 +1081,15 @@ impl PreparedDeclaration { /// read: the column the view projects it to, if any, otherwise an /// internal projection added here, named by [`input_column_name`]. pub fn input_column(&mut self, source_column: &str) -> Result { + // A grouped view's rows are groups; no source row carries a value into one. + if self.definition.is_grouped() { + return Err(Error::InvalidInput { + message: format!( + "a computed column on a grouped view cannot read source column \ + '{source_column}'; declare it on a view over the grouped view" + ), + }); + } if let Some(output) = self.lineage.get(source_column).and_then(|o| o.first()) { return Ok(output.clone()); } @@ -1357,6 +1401,7 @@ pub async fn prepare_declaration( .collect(), }, filter: filter.map(str::to_string), + group_by: Vec::new(), limit, }; prepare_definition(source, definition).await @@ -1383,6 +1428,26 @@ pub async fn prepare_declaration( /// # Ok(()) /// # } /// ``` +/// +/// A grouped view holds one row per group and is recomputed in full on +/// every refresh after a source change: +/// +/// ``` +/// # #![recursion_limit = "256"] +/// use lancedb::materialized_view::{MaterializedViewDefinition, prepare_definition}; +/// +/// # async fn declare(events: &lancedb::Table) -> Result<(), Box> { +/// let definition = MaterializedViewDefinition::from_sql( +/// "SELECT kind, count(*) AS n, array_agg(id) AS ids FROM events GROUP BY kind", +/// )?; +/// let view = prepare_definition(events, definition) +/// .await? +/// .create("events_by_kind") +/// .await?; +/// view.refresh().execute().await?; +/// # Ok(()) +/// # } +/// ``` pub async fn prepare_definition( source: &Table, definition: MaterializedViewDefinition, @@ -2067,6 +2132,7 @@ mod tests { }, ], filter: Some("age >= 18".into()), + group_by: Vec::new(), limit: Some(10), lateral: None, } @@ -3110,6 +3176,7 @@ mod tests { expression: "name".to_string(), }], filter: None, + group_by: Vec::new(), limit: None, } } @@ -3187,8 +3254,8 @@ mod tests { #[test] fn a_newer_format_is_reported_not_guessed() { assert_eq!( - read(r#"{"format":2,"query":"SELECT name FROM people"}"#).unwrap(), - Some(StoredDefinition::Newer { format: "2".into() }) + read(r#"{"format":3,"query":"SELECT name FROM people"}"#).unwrap(), + Some(StoredDefinition::Newer { format: "3".into() }) ); assert_eq!( read(r#"{"kind":"join"}"#).unwrap(), @@ -3201,6 +3268,7 @@ mod tests { r#"{"format":"one"}"#, r#"{"format":1}"#, r#"{"format":1,"query":"SELECT name FROM people GROUP BY name"}"#, + r#"{"format":2,"query":"SELECT name FROM people ORDER BY name"}"#, ] { assert!(read(stored).is_err(), "{stored}"); } @@ -3399,7 +3467,7 @@ mod tests { // The stored definition is the query alone; bindings live beside it. let stored: serde_json::Value = serde_json::from_str(&schema.metadata()[DEFINITION_META_KEY]).unwrap(); - assert_eq!(stored["format"], DEFINITION_FORMAT); + assert_eq!(stored["format"], 1); assert_eq!(view.definition().projections.len(), 2); assert_eq!(view.table().count_rows(None).await.unwrap(), 0); assert_eq!(conn.open_materialized_view("v").await.unwrap().name(), "v"); diff --git a/rust/lancedb/src/materialized_view/grouped.rs b/rust/lancedb/src/materialized_view/grouped.rs new file mode 100644 index 000000000..fc0dca430 --- /dev/null +++ b/rust/lancedb/src/materialized_view/grouped.rs @@ -0,0 +1,261 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +//! Views with `GROUP BY`: one row per group, planned and executed by +//! DataFusion, and recomputed in full whenever the source changes, since a +//! group spans fragments. + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; + +use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema, SchemaRef}; +use datafusion::catalog::default_table_source::provider_as_source; +use datafusion::datasource::empty::EmptyTable; +use datafusion::execution::SessionStateBuilder; +use datafusion::execution::context::SessionState; +use datafusion::physical_plan::SendableRecordBatchStream; +use datafusion::physical_plan::stream::RecordBatchStreamAdapter; +use datafusion::prelude::SessionContext; +use datafusion_common::TableReference; +use datafusion_common::config::ConfigOptions; +use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion}; +use datafusion_expr::planner::{ContextProvider, ExprPlanner}; +use datafusion_expr::{ + AggregateUDF, HigherOrderUDF, LogicalPlan, ScalarUDF, TableSource, WindowUDF, +}; +use datafusion_sql::planner::SqlToRel; +use datafusion_sql::sqlparser::dialect::GenericDialect; +use datafusion_sql::sqlparser::parser::Parser; +use futures::StreamExt; +use lance::Dataset; +use lance_core::ROW_ID; +use lance_datafusion::exec::SessionContextExt; + +use super::refresh::to_view_batch; +use super::{ + MaterializedViewDefinition, Planned, SOURCE_ROW_ID_COLUMN, ViewProjection, ensure_immutable, + plan_filter, query, without_declarations, +}; +use crate::{Error, Result}; + +/// The name the source is planned under; never visible to the user. +const SOURCE: &str = "__source"; + +fn session() -> SessionState { + SessionStateBuilder::new().with_default_features().build() +} + +/// The query DataFusion runs: the view's projections plus the group's +/// smallest source row id, which stands as the row's provenance. The filter +/// is not here; the lance scan applies it, as for any other view. +fn sql(definition: &MaterializedViewDefinition) -> String { + let items: Vec = definition + .projections + .iter() + .map(|p| format!("{} AS {}", p.expression, query::ident_sql(&p.output))) + .collect(); + format!( + "SELECT {}, min({ROW_ID}) AS {ROW_ID} FROM {SOURCE} GROUP BY {}", + items.join(", "), + definition.group_by.join(", ") + ) +} + +struct Provider<'a> { + state: &'a SessionState, + source: Arc, +} + +impl ContextProvider for Provider<'_> { + fn get_table_source( + &self, + name: TableReference, + ) -> datafusion_common::Result> { + if name.schema().is_none() && name.table() == SOURCE { + Ok(self.source.clone()) + } else { + datafusion_common::plan_err!("a grouped view reads only its source, not '{name}'") + } + } + fn get_expr_planners(&self) -> &[Arc] { + self.state.expr_planners() + } + fn get_function_meta(&self, name: &str) -> Option> { + self.state.scalar_functions().get(name).cloned() + } + fn get_higher_order_meta(&self, name: &str) -> Option> { + self.state.higher_order_functions().get(name).cloned() + } + fn get_aggregate_meta(&self, name: &str) -> Option> { + self.state.aggregate_functions().get(name).cloned() + } + fn get_window_meta(&self, name: &str) -> Option> { + self.state.window_functions().get(name).cloned() + } + fn get_variable_type(&self, _: &[String]) -> Option { + None + } + fn options(&self) -> &ConfigOptions { + self.state.config_options() + } + fn udf_names(&self) -> Vec { + self.state.scalar_functions().keys().cloned().collect() + } + fn higher_order_function_names(&self) -> Vec { + self.state + .higher_order_functions() + .keys() + .cloned() + .collect() + } + fn udaf_names(&self) -> Vec { + self.state.aggregate_functions().keys().cloned().collect() + } + fn udwf_names(&self) -> Vec { + self.state.window_functions().keys().cloned().collect() + } +} + +fn logical_plan( + state: &SessionState, + source: Arc, + definition: &MaterializedViewDefinition, +) -> Result { + let invalid = |e: &dyn std::fmt::Display| Error::InvalidInput { + message: format!("invalid grouped view: {e}"), + }; + let statement = Parser::parse_sql(&GenericDialect {}, &sql(definition)) + .map_err(|e| invalid(&e))? + .pop() + .ok_or_else(|| invalid(&"empty query"))?; + SqlToRel::new(&Provider { state, source }) + .sql_statement_to_plan(statement) + .map_err(|e| invalid(&e)) +} + +/// Plan a grouped definition against the source schema. `filter` is the +/// canonical filter, already spelled the way [`super::plan`] stores it. +pub(super) fn plan( + source_schema: SchemaRef, + definition: &MaterializedViewDefinition, + filter: Option, +) -> Result { + query::check_grouping(definition)?; + let mut inputs = match &filter { + Some(filter) => plan_filter(&source_schema, filter)?, + None => Vec::new(), + }; + let mut projections: Vec = Vec::with_capacity(definition.projections.len()); + for p in &definition.projections { + if p.output == SOURCE_ROW_ID_COLUMN || p.output == ROW_ID { + return Err(Error::InvalidInput { + message: format!("view column name '{}' is reserved", p.output), + }); + } + if projections.iter().any(|seen| seen.output == p.output) { + return Err(Error::ColumnAlreadyExists { + name: p.output.clone(), + }); + } + let expression = + query::canonical_expr(&p.expression).map_err(|e| Error::InvalidExpression { + column: p.output.clone(), + message: e.to_string(), + })?; + projections.push(ViewProjection { + output: p.output.clone(), + expression, + }); + } + let group_by = definition + .group_by + .iter() + .map(|key| query::canonical_expr(key)) + .collect::>>()?; + let definition = MaterializedViewDefinition { + projections, + filter, + group_by, + ..definition.clone() + }; + + let mut fields = source_schema.fields().to_vec(); + fields.push(Arc::new(ArrowField::new(ROW_ID, DataType::UInt64, false))); + let empty = provider_as_source(Arc::new(EmptyTable::new(Arc::new(ArrowSchema::new( + fields, + ))))); + let planned = logical_plan(&session(), empty, &definition)?; + + let mut exprs = Vec::new(); + planned + .apply(|node| { + exprs.extend(node.expressions()); + Ok(TreeNodeRecursion::Continue) + }) + .map_err(|e| Error::InvalidInput { + message: format!("invalid grouped view: {e}"), + })?; + for expr in &exprs { + ensure_immutable(expr, |message| Error::InvalidInput { + message: format!("invalid grouped view: {message}"), + })?; + inputs.extend( + expr.column_refs() + .into_iter() + .map(|c| c.name.clone()) + .filter(|name| name != ROW_ID && source_schema.field_with_name(name).is_ok()), + ); + } + inputs.sort(); + inputs.dedup(); + + let fields = planned + .schema() + .fields() + .iter() + .filter(|f| f.name() != ROW_ID) + .map(|f| without_declarations(f)) + .collect(); + Ok(Planned { + definition, + fields, + lineage: HashMap::new(), + inputs, + }) +} + +/// Every group of `source`, as view rows in `schema`. +pub(super) async fn stream( + source: &Dataset, + definition: &MaterializedViewDefinition, + inputs: &[String], + schema: SchemaRef, + rows_written: Arc, +) -> Result { + let mut scanner = source.scan(); + scanner.with_row_id(); + if let Some(filter) = &definition.filter { + scanner.filter(filter)?; + } + scanner.project(inputs)?; + let scan: SendableRecordBatchStream = scanner.try_into_stream().await?.into(); + + let state = session(); + let ctx = SessionContext::new_with_state(state.clone()); + let source = ctx.read_one_shot(scan)?.into_view(); + let planned = logical_plan(&state, provider_as_source(source), definition)?; + let groups = ctx + .execute_logical_plan(planned) + .await? + .execute_stream() + .await?; + + let out_schema = schema.clone(); + let mapped = groups.map(move |batch| { + let batch = to_view_batch(&batch?, &out_schema)?; + rows_written.fetch_add(batch.num_rows() as u64, Ordering::Relaxed); + Ok(batch) + }); + Ok(Box::pin(RecordBatchStreamAdapter::new(schema, mapped))) +} diff --git a/rust/lancedb/src/materialized_view/query.rs b/rust/lancedb/src/materialized_view/query.rs index 35af77f0b..c1a987b2c 100644 --- a/rust/lancedb/src/materialized_view/query.rs +++ b/rust/lancedb/src/materialized_view/query.rs @@ -18,8 +18,9 @@ //! query: fail closed, never materialize it at the wrong cardinality. use datafusion_sql::sqlparser::ast::{ - Expr, FunctionArg, FunctionArgExpr, JoinOperator, LimitClause, ObjectName, ObjectNamePart, - Query, SelectItem, SetExpr, Statement, TableFactor, TableFunctionArgs, TableWithJoins, Value, + Expr, FunctionArg, FunctionArgExpr, GroupByExpr, JoinOperator, LimitClause, ObjectName, + ObjectNamePart, Query, SelectItem, SetExpr, Statement, TableFactor, TableFunctionArgs, + TableWithJoins, Value, }; use datafusion_sql::sqlparser::dialect::GenericDialect; use datafusion_sql::sqlparser::keywords::{ @@ -40,7 +41,8 @@ fn invalid(message: impl Into) -> Error { } const SHAPE: &str = "a materialized view is defined by `SELECT columns FROM table \ - [, function(args) AS alias | , UNNEST(column) AS alias] [WHERE predicate] [LIMIT n]`"; + [, function(args) AS alias | , UNNEST(column) AS alias] [WHERE predicate] \ + [GROUP BY expr, ...] [LIMIT n]`"; /// Whether `name` must be delimited to read back as this identifier: bare, /// the parser would take it as a keyword, or a different spelling. @@ -394,14 +396,59 @@ fn extract(query: &Query) -> Result { Some(_) => return Err(invalid("view limit must be a plain `LIMIT n`")), }; - Ok(MaterializedViewDefinition { + if select.having.is_some() { + return Err(invalid("HAVING is not supported in a view query")); + } + let group_by = match &select.group_by { + GroupByExpr::Expressions(exprs, modifiers) + if modifiers.is_empty() + && !exprs.iter().any(|e| { + matches!(e, Expr::Rollup(_) | Expr::Cube(_) | Expr::GroupingSets(_)) + }) => + { + exprs.iter().map(|e| e.to_string()).collect() + } + _ => { + return Err(invalid( + "a view groups by a plain `GROUP BY expr, ...` list", + )); + } + }; + + let definition = MaterializedViewDefinition { source_table, source_namespace, lateral, projections, filter: select.selection.as_ref().map(|e| e.to_string()), + group_by, limit, - }) + }; + check_grouping(&definition)?; + Ok(definition) +} + +/// The clauses a grouped view cannot combine with: each would have to be +/// maintained per group rather than per source row. +pub fn check_grouping(definition: &MaterializedViewDefinition) -> Result<()> { + if !definition.is_grouped() { + return Ok(()); + } + if definition.lateral.is_some() { + return Err(invalid( + "GROUP BY cannot be combined with a FROM-position item; group in one view \ + and expand in a view over it", + )); + } + if definition.selects_star() { + return Err(invalid( + "a grouped view selects its keys and aggregates, not `*`", + )); + } + if definition.limit.is_some() { + return Err(invalid("LIMIT is not supported together with GROUP BY")); + } + Ok(()) } /// The canonical query for `definition`; [`parse`] reads it back equal. @@ -451,6 +498,9 @@ pub fn render(definition: &MaterializedViewDefinition) -> String { if let Some(filter) = &definition.filter { sql.push_str(&format!(" WHERE {filter}")); } + if definition.is_grouped() { + sql.push_str(&format!(" GROUP BY {}", definition.group_by.join(", "))); + } if let Some(limit) = definition.limit { sql.push_str(&format!(" LIMIT {limit}")); } @@ -489,6 +539,10 @@ mod tests { "SELECT id, c.text FROM docs, chunk(upper(body)) AS c", ), ("SELECT * FROM `select`.t", "SELECT * FROM `select`.t"), + ( + "select k, array_agg(id) as ids, count(*) n from t where x > 1 group by k", + "SELECT k, array_agg(id) AS ids, count(*) AS n FROM t WHERE x > 1 GROUP BY k", + ), ( "SELECT meta.title AS title, id AS id FROM t", "SELECT meta.title, id FROM t", @@ -547,7 +601,6 @@ mod tests { #[test] fn unsupported_clauses_are_refused() { for sql in [ - "SELECT id FROM t GROUP BY id", "SELECT id FROM t ORDER BY id", "SELECT DISTINCT id FROM t", "SELECT id FROM t LIMIT 5 OFFSET 2", @@ -567,6 +620,12 @@ mod tests { "SELECT *, id FROM t", "WITH q AS (SELECT 1) SELECT id FROM t", "SELECT id FROM t HAVING id > 1", + "SELECT k, count(*) AS n FROM t GROUP BY k HAVING count(*) > 1", + "SELECT k FROM t GROUP BY ALL", + "SELECT k FROM t GROUP BY ROLLUP (k)", + "SELECT * FROM t GROUP BY k", + "SELECT k FROM t GROUP BY k LIMIT 5", + "SELECT k, c.x FROM t, UNNEST(xs) AS c GROUP BY k", ] { assert!(parse(sql).is_err(), "{sql}"); } diff --git a/rust/lancedb/src/materialized_view/refresh.rs b/rust/lancedb/src/materialized_view/refresh.rs index 0fc9b99c5..39e1d8b9c 100644 --- a/rust/lancedb/src/materialized_view/refresh.rs +++ b/rust/lancedb/src/materialized_view/refresh.rs @@ -297,6 +297,22 @@ pub(crate) async fn execute_refresh( }); } + // A group spans fragments: there is no per-fragment increment. + if definition.is_grouped() { + return rebuild( + view_native, + &view_ds, + &source_ds, + source_version, + source_ts, + definition, + &inputs, + None, + persist.is_some(), + expected_incarnation, + ) + .await; + } let watermark = watermark.filter(|_| view_intact); match plan_increment( &source_ds, @@ -629,9 +645,17 @@ fn same_meaning( || stored.limit != replanned.limit || stored.projections.len() != replanned.projections.len() || stored.filter.is_some() != replanned.filter.is_some() + || stored.group_by.len() != replanned.group_by.len() { return false; } + // Lance's planner does not read aggregates; a grouped definition is + // compared in its canonical spelling, which planning produced. + if stored.is_grouped() { + return stored.projections == replanned.projections + && stored.group_by == replanned.group_by + && stored.filter == replanned.filter; + } let schema = match unnest { None => source_schema.clone(), Some(unnest) => match super::flattened_schema(source_schema, unnest) { @@ -1288,6 +1312,9 @@ async fn compute_stream( schema: SchemaRef, rows_written: Arc, ) -> Result { + if definition.is_grouped() { + return super::grouped::stream(source, definition, inputs, schema, rows_written).await; + } let RowScope { fragments, created_after, @@ -1363,31 +1390,42 @@ async fn compute_stream( None => batch, Some(expanded) => expanded.apply(&batch)?, }; - let mut columns = Vec::with_capacity(out_schema.fields().len()); - for field in out_schema.fields() { - if computed_column_from_field(field).is_some() { - columns.push(new_null_array(field.data_type(), batch.num_rows())); - continue; - } - let name = if field.name() == SOURCE_ROW_ID_COLUMN { - ROW_ID - } else { - field.name() - }; - let column = batch.column_by_name(name).ok_or_else(|| { - DataFusionError::Internal(format!( - "view column '{}' is not produced by the view's definition", - field.name() - )) - })?; - columns.push(column.clone()); - } + let batch = to_view_batch(&batch, &out_schema)?; rows_written.fetch_add(batch.num_rows() as u64, Ordering::Relaxed); - Ok(RecordBatch::try_new(out_schema.clone(), columns)?) + Ok(batch) }); Ok(Box::pin(RecordBatchStreamAdapter::new(schema, mapped))) } +/// `batch` in the view's physical schema: the row id becomes the +/// provenance column and computed columns are written NULL for their owner +/// to fill. +pub(super) fn to_view_batch( + batch: &RecordBatch, + out_schema: &SchemaRef, +) -> datafusion::common::Result { + let mut columns = Vec::with_capacity(out_schema.fields().len()); + for field in out_schema.fields() { + if computed_column_from_field(field).is_some() { + columns.push(new_null_array(field.data_type(), batch.num_rows())); + continue; + } + let name = if field.name() == SOURCE_ROW_ID_COLUMN { + ROW_ID + } else { + field.name() + }; + let column = batch.column_by_name(name).ok_or_else(|| { + DataFusionError::Internal(format!( + "view column '{}' is not produced by the view's definition", + field.name() + )) + })?; + columns.push(column.clone()); + } + Ok(RecordBatch::try_new(out_schema.clone(), columns)?) +} + /// The post-scan half of an unnested view's refresh: the scan reads /// `raw_inputs` plus the row id, and each batch is unnested on the list /// column, filtered, then projected by expressions typed against @@ -3641,6 +3679,7 @@ mod tests { }, ], filter: None, + group_by: Vec::new(), limit: None, lateral: None, }; @@ -3675,6 +3714,7 @@ mod tests { expression: "x".into(), }], filter: None, + group_by: Vec::new(), limit: None, lateral: None, }; @@ -4304,4 +4344,149 @@ mod tests { 3 ); } + + async fn grouped_source(conn: &Connection) -> Table { + conn.create_table( + "src", + record_batch!(("k", Int32, [1, 1, 2, 3]), ("x", Int32, [10, 11, 20, 30])).unwrap(), + ) + .write_options(crate::materialized_view::tests::stable_row_ids()) + .execute() + .await + .unwrap() + } + + async fn declare(source: &Table, name: &str, sql: &str) -> Result { + crate::materialized_view::prepare_definition( + source, + MaterializedViewDefinition::from_sql(sql)?, + ) + .await? + .create(name) + .await + } + + /// Rows of `table` as `columns` rendered to text, sorted. + async fn rows(table: &Table, columns: &[&str]) -> Vec { + let batches = table + .query() + .select(Select::columns(columns)) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let mut rows = Vec::new(); + for batch in &batches { + let formatters: Vec<_> = batch + .columns() + .iter() + .map(|c| { + arrow_cast::display::ArrayFormatter::try_new(c.as_ref(), &Default::default()) + .unwrap() + }) + .collect(); + for row in 0..batch.num_rows() { + let cells: Vec = formatters + .iter() + .map(|f| f.value(row).to_string()) + .collect(); + rows.push(cells.join(" ")); + } + } + rows.sort(); + rows + } + + #[tokio::test] + async fn a_grouped_view_holds_one_row_per_group() { + let conn = connect("memory://").execute().await.unwrap(); + let source = grouped_source(&conn).await; + let view = declare( + &source, + "groups", + "SELECT k, array_agg(x) AS xs, sum(x) AS total FROM src WHERE x < 30 GROUP BY k", + ) + .await + .unwrap(); + assert_eq!( + view.definition().to_sql(), + "SELECT k, array_agg(x) AS xs, sum(x) AS total FROM src WHERE x < 30 GROUP BY k" + ); + + let result = view.refresh().execute().await.unwrap(); + assert_eq!(result.mode, RefreshMode::Rebuild); + assert_eq!(result.rows_written, 2); + assert_eq!(rows(view.table(), &["k", "total"]).await, ["1 21", "2 20"]); + let unchanged = view.refresh().execute().await.unwrap(); + assert_eq!(unchanged.mode, RefreshMode::NoOp); + + // Any source change recomputes every group. + source + .add(record_batch!(("k", Int32, [2]), ("x", Int32, [5])).unwrap()) + .execute() + .await + .unwrap(); + let result = view.refresh().execute().await.unwrap(); + assert_eq!(result.mode, RefreshMode::Rebuild); + assert_eq!(rows(view.table(), &["k", "total"]).await, ["1 21", "2 25"]); + + // A view over the groups expands them again. + let expanded = declare( + view.table(), + "elements", + "SELECT k, e AS x FROM groups, UNNEST(xs) AS e", + ) + .await + .unwrap(); + expanded.refresh().execute().await.unwrap(); + assert_eq!( + rows(expanded.table(), &["k", "x"]).await, + ["1 10", "1 11", "2 20", "2 5"] + ); + } + + /// Provenance is one source row per group, so the view's row-id keys + /// stay unique. + #[tokio::test] + async fn a_grouped_view_records_one_source_row_per_group() { + let conn = connect("memory://").execute().await.unwrap(); + let source = grouped_source(&conn).await; + let view = declare( + &source, + "groups", + "SELECT k, count(*) AS n FROM src GROUP BY k", + ) + .await + .unwrap(); + view.refresh().execute().await.unwrap(); + let ids = rows(view.table(), &[SOURCE_ROW_ID_COLUMN]).await; + assert_eq!(ids.len(), 3); + assert_eq!(ids.iter().collect::>().len(), 3, "{ids:?}"); + } + + #[tokio::test] + async fn a_grouped_view_is_planned_at_declaration() { + let conn = connect("memory://").execute().await.unwrap(); + let source = grouped_source(&conn).await; + for sql in [ + // x is neither a key nor aggregated + "SELECT k, x FROM src GROUP BY k", + "SELECT k, random() AS r FROM src GROUP BY k", + "SELECT k, count(*) AS n FROM src GROUP BY missing", + "SELECT k, count(*) AS __source_row_id FROM src GROUP BY k", + "SELECT k, count(*) AS k FROM src GROUP BY k", + ] { + assert!(declare(&source, "v", sql).await.is_err(), "{sql}"); + } + let mut grouped = crate::materialized_view::prepare_definition( + &source, + MaterializedViewDefinition::from_sql("SELECT k, count(*) AS n FROM src GROUP BY k") + .unwrap(), + ) + .await + .unwrap(); + assert!(grouped.input_column("x").is_err()); + } } diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 0c51f50ae..6bbb49adf 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -2144,6 +2144,7 @@ impl BaseTable for RemoteTable { }) .collect(), filter: response.filter, + group_by: Vec::new(), limit: response.limit, }, };