From e1f95aa41e83b0372a3a67cd75e6c6ee6ce54740 Mon Sep 17 00:00:00 2001 From: discord9 Date: Fri, 4 Sep 2026 18:44:04 +0800 Subject: [PATCH] feat(flow): support durable incremental aggregate state Signed-off-by: discord9 --- src/common/function/src/aggrs/approximate.rs | 11 + .../function/src/aggrs/approximate/avg.rs | 545 +++++++++++++ src/common/meta/src/ddl/create_flow.rs | 21 + src/common/meta/src/ddl_manager.rs | 85 +- src/flow/src/batching_mode/engine.rs | 336 +++++++- src/flow/src/batching_mode/task.rs | 13 +- src/flow/src/batching_mode/task/ckpt.rs | 45 -- src/flow/src/batching_mode/task/inc.rs | 75 +- src/flow/src/batching_mode/task/test.rs | 609 ++------------ src/flow/src/batching_mode/utils.rs | 475 +++++++---- src/flow/src/batching_mode/utils/test.rs | 750 ++++++++++++++++-- src/flow/src/server.rs | 10 +- 12 files changed, 2115 insertions(+), 860 deletions(-) create mode 100644 src/common/function/src/aggrs/approximate/avg.rs diff --git a/src/common/function/src/aggrs/approximate.rs b/src/common/function/src/aggrs/approximate.rs index bfab225755..562723f672 100644 --- a/src/common/function/src/aggrs/approximate.rs +++ b/src/common/function/src/aggrs/approximate.rs @@ -18,6 +18,7 @@ use datatypes::arrow::datatypes::DataType; use crate::aggrs::aggr_wrapper::DeltaMergeWrapper; use crate::function_registry::FunctionRegistry; +pub mod avg; pub mod hll; pub mod uddsketch; pub mod welford; @@ -26,6 +27,16 @@ pub(crate) struct ApproximateFunction; impl ApproximateFunction { pub fn register(registry: &FunctionRegistry) { + let avg_merge = avg::AvgAccumulator::merge_udf_impl(); + registry.register_aggr(avg::AvgAccumulator::state_udf_impl()); + registry.register_aggr(avg_merge.clone()); + registry.register_aggr(AggregateUDF::new_from_impl(DeltaMergeWrapper::new( + avg_merge.clone(), + avg::AVG_STATE_NAME, + vec![DataType::Binary], + DataType::Binary, + ))); + let uddsketch_state = uddsketch::UddSketchState::state_udf_impl(); let uddsketch_merge = uddsketch::UddSketchState::merge_udf_impl(); let uddsketch_delta = AggregateUDF::new_from_impl(DeltaMergeWrapper::new( diff --git a/src/common/function/src/aggrs/approximate/avg.rs b/src/common/function/src/aggrs/approximate/avg.rs new file mode 100644 index 0000000000..49b0ed3274 --- /dev/null +++ b/src/common/function/src/aggrs/approximate/avg.rs @@ -0,0 +1,545 @@ +// Copyright 2023 Greptime Team +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use std::sync::Arc; + +use datafusion::arrow::array::{ArrayRef, Float64Array}; +use datafusion::arrow::compute::sum; +use datafusion::common::cast::{as_binary_array, as_primitive_array}; +use datafusion::common::not_impl_err; +use datafusion::error::{DataFusionError, Result as DfResult}; +use datafusion::logical_expr::function::AccumulatorArgs; +use datafusion::logical_expr::{Accumulator as DfAccumulator, AggregateUDF, Volatility}; +use datafusion::prelude::create_udaf; +use datafusion_common::ScalarValue; +use datatypes::arrow::datatypes::{DataType, Float64Type}; + +pub const AVG_STATE_NAME: &str = "avg_state"; +pub const AVG_MERGE_NAME: &str = "avg_merge"; + +const ENCODED_LEN: usize = 20; +const MAGIC: &[u8; 4] = b"AVG1"; + +/// The portable state used by the Float64 average aggregate functions. +#[derive(Debug, Clone, Copy, PartialEq)] +pub struct AvgState { + count: u64, + sum: f64, +} + +impl Default for AvgState { + fn default() -> Self { + Self { count: 0, sum: 0.0 } + } +} + +impl AvgState { + /// Returns the exact AVG1 representation of this state. + pub(crate) fn encode(&self) -> [u8; ENCODED_LEN] { + let mut encoded = [0; ENCODED_LEN]; + encoded[..4].copy_from_slice(MAGIC); + encoded[4..12].copy_from_slice(&self.count.to_le_bytes()); + encoded[12..20].copy_from_slice(&self.sum.to_bits().to_le_bytes()); + encoded + } + + /// Decodes and validates an AVG1 state. + pub fn decode(encoded: &[u8]) -> DfResult { + if encoded.len() != ENCODED_LEN || &encoded[..4] != MAGIC { + return Err(invalid_state()); + } + let count = decode_u64(encoded, 4); + let sum = f64::from_bits(decode_u64(encoded, 12)); + if count == 0 && sum.to_bits() != 0 { + return Err(invalid_state()); + } + Ok(Self { count, sum }) + } + + /// Returns the number of non-null input values in this state. + pub(crate) fn count(&self) -> u64 { + self.count + } + + /// Returns the average, or `None` for the canonical empty state. + pub fn average(&self) -> Option { + (self.count() != 0).then(|| self.sum / self.count() as f64) + } +} + +fn decode_u64(encoded: &[u8], offset: usize) -> u64 { + let mut bytes = [0; 8]; + bytes.copy_from_slice(&encoded[offset..offset + 8]); + u64::from_le_bytes(bytes) +} + +fn invalid_state() -> DataFusionError { + DataFusionError::Execution("Invalid AVG1 state".to_string()) +} + +fn count_overflow() -> DataFusionError { + DataFusionError::Execution("AVG count overflow".to_string()) +} + +#[derive(Debug, Clone, Copy)] +enum InputKind { + Float64, + Binary, +} + +/// Accumulates and merges AVG1 states. +#[derive(Debug)] +pub(crate) struct AvgAccumulator { + state: AvgState, + input: InputKind, +} + +impl Default for AvgAccumulator { + fn default() -> Self { + Self { + state: AvgState::default(), + input: InputKind::Float64, + } + } +} + +impl AvgAccumulator { + pub fn state_udf_impl() -> AggregateUDF { + create_udaf( + AVG_STATE_NAME, + vec![DataType::Float64], + Arc::new(DataType::Binary), + Volatility::Immutable, + Arc::new(Self::create_accumulator), + Arc::new(vec![DataType::Binary]), + ) + } + + pub fn merge_udf_impl() -> AggregateUDF { + create_udaf( + AVG_MERGE_NAME, + vec![DataType::Binary], + Arc::new(DataType::Binary), + Volatility::Immutable, + Arc::new(Self::create_accumulator), + Arc::new(vec![DataType::Binary]), + ) + } + + fn create_accumulator(args: AccumulatorArgs) -> DfResult> { + if args.is_distinct { + return not_impl_err!("AVG DISTINCT aggregations are not available"); + } + let input = match args.exprs[0].data_type(args.schema)? { + DataType::Float64 => InputKind::Float64, + DataType::Binary => InputKind::Binary, + data_type => return not_impl_err!("AVG functions do not support {data_type:?}"), + }; + Ok(Box::new(Self { + state: AvgState::default(), + input, + })) + } + + fn update_float64(&mut self, array: &ArrayRef) -> DfResult<()> { + let array = as_primitive_array::(array)?; + let mut count = self.state.count; + for _ in array.iter().flatten() { + count = count.checked_add(1).ok_or_else(count_overflow)?; + } + let sum = sum(array) + .map(|batch_sum| self.state.sum + batch_sum) + .unwrap_or(self.state.sum); + self.state = AvgState { count, sum }; + Ok(()) + } + + fn merge_states(&mut self, array: &ArrayRef) -> DfResult<()> { + let array = as_binary_array(array)?; + let states = array + .iter() + .flatten() + .map(AvgState::decode) + .collect::>>()?; + let count = states.iter().try_fold(self.state.count, |count, state| { + count.checked_add(state.count).ok_or_else(count_overflow) + })?; + let sums = states + .iter() + .filter(|state| state.count != 0) + .map(|state| Some(state.sum)) + .collect::>(); + let sum = sum(&Float64Array::from(sums)) + .map(|batch_sum| self.state.sum + batch_sum) + .unwrap_or(self.state.sum); + self.state = AvgState { count, sum }; + Ok(()) + } +} + +impl DfAccumulator for AvgAccumulator { + fn update_batch(&mut self, values: &[ArrayRef]) -> DfResult<()> { + let array = &values[0]; + match (self.input, array.data_type()) { + (InputKind::Float64, DataType::Float64) => self.update_float64(array), + (InputKind::Binary, DataType::Binary) => self.merge_states(array), + (_, data_type) => not_impl_err!("AVG input type does not match: {data_type:?}"), + } + } + + fn evaluate(&mut self) -> DfResult { + Ok(ScalarValue::Binary(Some(self.state.encode().to_vec()))) + } + + fn size(&self) -> usize { + std::mem::size_of::() + } + + fn state(&mut self) -> DfResult> { + Ok(vec![ScalarValue::Binary(Some( + self.state.encode().to_vec(), + ))]) + } + + fn merge_batch(&mut self, states: &[ArrayRef]) -> DfResult<()> { + self.merge_states(&states[0]) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use arrow::array::{BinaryArray, Float64Array}; + use datafusion_common::ScalarValue; + use datafusion_common::arrow::datatypes::DataType; + use datafusion_expr::TypeSignature; + use datafusion_physical_expr::aggregate::AggregateExprBuilder; + use datafusion_physical_expr::expressions::{Column, lit as physical_lit}; + + use super::*; + use crate::aggrs::aggr_wrapper::{aggr_delta_merge_func_name, aggr_state_func_name}; + use crate::function_registry::FUNCTION_REGISTRY; + + fn state(count: u64, sum: f64) -> Vec { + AvgState { count, sum }.encode().to_vec() + } + + #[test] + fn codec_golden_and_roundtrip() { + let empty = AvgState::default().encode(); + assert_eq!(empty.len(), ENCODED_LEN); + assert_eq!(empty.as_slice(), b"AVG1\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0"); + let mut accumulator = AvgAccumulator::default(); + accumulator + .update_batch(&[Arc::new(Float64Array::from(vec![Some(1.5)]))]) + .unwrap(); + let one = accumulator.state.encode(); + assert_eq!(one.len(), ENCODED_LEN); + assert_eq!( + one.as_slice(), + &[ + b'A', b'V', b'G', b'1', 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0xf8, 0x3f, + ] + ); + assert_eq!(&one[12..20], &1.5f64.to_bits().to_le_bytes()); + assert_eq!(AvgState::decode(&empty).unwrap().encode(), empty); + assert_eq!(AvgState::decode(&one).unwrap().encode(), one); + } + + #[test] + fn codec_rejects_malformed_states() { + assert!(AvgState::decode(b"").is_err()); + assert!(AvgState::decode(&[0; 19]).is_err()); + assert!(AvgState::decode(&[0; 21]).is_err()); + let mut avg2 = AvgState::default().encode(); + avg2[..4].copy_from_slice(b"AVG2"); + assert!(AvgState::decode(&avg2).is_err()); + let mut wrong_magic = AvgState::default().encode(); + wrong_magic[0] = b'X'; + assert!(AvgState::decode(&wrong_magic).is_err()); + for sum in [1.0, -0.0] { + assert!(AvgState::decode(&state(0, sum)).is_err()); + } + let mut count = AvgState { + count: 0x0102_0304_0506_0708, + sum: 0.0, + } + .encode(); + assert_eq!(&count[4..12], &0x0102_0304_0506_0708u64.to_le_bytes()); + count[4..12].reverse(); + assert_ne!( + AvgState::decode(&count).unwrap().count(), + 0x0102_0304_0506_0708 + ); + let mut sum = AvgState { count: 1, sum: 1.5 }.encode(); + assert_eq!(&sum[12..20], &1.5f64.to_bits().to_le_bytes()); + sum[12..20].reverse(); + assert_ne!(AvgState::decode(&sum).unwrap().average(), Some(1.5)); + } + + #[test] + fn codec_preserves_populated_float_bits() { + for bits in [ + 0.0f64.to_bits(), + (-0.0f64).to_bits(), + f64::INFINITY.to_bits(), + f64::NEG_INFINITY.to_bits(), + 0x7ff8_0000_0000_0001, + 0x7ff0_0000_0000_0001, + ] { + let encoded = state(1, f64::from_bits(bits)); + assert_eq!( + AvgState::decode(&encoded).unwrap().encode().as_slice(), + encoded + ); + } + } + + #[test] + fn distinct_is_rejected() { + let udf = AvgAccumulator::state_udf_impl(); + let schema = arrow_schema::Schema::empty(); + let expr = physical_lit(1.0f64); + let field = Arc::new(arrow_schema::Field::new("in", DataType::Float64, true)); + let args = AccumulatorArgs { + return_field: Arc::new(arrow_schema::Field::new("out", DataType::Binary, true)), + schema: &schema, + ignore_nulls: false, + order_bys: &[], + is_reversed: false, + name: AVG_STATE_NAME, + is_distinct: true, + exprs: std::slice::from_ref(&expr), + expr_fields: std::slice::from_ref(&field), + }; + assert!(udf.accumulator(args).is_err()); + } + + #[test] + fn state_counts_nulls_and_empty_is_canonical() { + let mut accumulator = AvgAccumulator::default(); + accumulator + .update_batch(&[Arc::new(Float64Array::from(vec![None, None]))]) + .unwrap(); + assert_eq!(accumulator.state.encode(), AvgState::default().encode()); + accumulator + .update_batch(&[Arc::new(Float64Array::from(vec![ + Some(1.0), + None, + Some(3.0), + Some(8.0), + ]))]) + .unwrap(); + assert_eq!(accumulator.state.count(), 3); + assert_eq!(accumulator.state.average(), Some(4.0)); + } + + #[test] + fn merge_preserves_populated_negative_zero_for_empty_input() { + let mut accumulator = AvgAccumulator { + state: AvgState { + count: 1, + sum: -0.0, + }, + input: InputKind::Binary, + }; + let expected = accumulator.state.encode(); + accumulator + .update_batch(&[Arc::new(BinaryArray::from(vec![ + None, + Some(AvgState::default().encode().as_slice()), + ]))]) + .unwrap(); + assert_eq!(accumulator.state.encode(), expected); + } + + #[test] + fn merge_ignores_nulls_and_merges_weighted_states() { + let mut accumulator = AvgAccumulator { + state: AvgState::default(), + input: InputKind::Binary, + }; + accumulator + .update_batch(&[Arc::new(BinaryArray::from(vec![ + Some(state(2, 4.0).as_slice()), + None, + Some(state(3, 15.0).as_slice()), + ]))]) + .unwrap(); + assert_eq!(accumulator.state.count(), 5); + assert_eq!(accumulator.state.average(), Some(19.0 / 5.0)); + let before = accumulator.state; + assert!( + accumulator + .update_batch(&[Arc::new(BinaryArray::from(vec![Some(&[][..])]))]) + .is_err() + ); + assert_eq!(accumulator.state, before); + } + + #[test] + fn overflow_does_not_mutate_update_or_merge() { + let mut update = AvgAccumulator { + state: AvgState { + count: u64::MAX, + sum: 1.0, + }, + input: InputKind::Float64, + }; + let before = update.state; + assert!( + update + .update_batch(&[Arc::new(Float64Array::from(vec![Some(2.0)]))]) + .is_err() + ); + assert_eq!(update.state, before); + + let mut merge = AvgAccumulator { + state: AvgState { + count: u64::MAX, + sum: 1.0, + }, + input: InputKind::Binary, + }; + let before = merge.state; + assert!( + merge + .update_batch(&[Arc::new(BinaryArray::from(vec![Some( + state(1, 2.0).as_slice() + )]))]) + .is_err() + ); + assert_eq!(merge.state, before); + } + + #[test] + fn registered_delta_merge_has_four_way_and_malformed_behavior() { + let udf = FUNCTION_REGISTRY + .get_aggr_func(&aggr_delta_merge_func_name(AVG_STATE_NAME)) + .unwrap(); + assert_eq!(udf.name(), "__avg_state_delta_merge"); + assert_eq!( + udf.signature().type_signature, + TypeSignature::Exact(vec![DataType::Binary, DataType::Binary]) + ); + let schema = Arc::new(arrow_schema::Schema::new(vec![ + arrow_schema::Field::new("delta", DataType::Binary, true), + arrow_schema::Field::new("persisted", DataType::Binary, true), + ])); + let expr = AggregateExprBuilder::new( + Arc::new(udf), + vec![ + Arc::new(Column::new("delta", 0)), + Arc::new(Column::new("persisted", 1)), + ], + ) + .schema(schema) + .alias("avg_delta_merge") + .build() + .unwrap(); + let delta = state(2, 3.0); + let persisted = state(2, 7.0); + for (left, right, expected) in [ + ( + Some(delta.as_slice()), + None, + AvgState { count: 2, sum: 3.0 }.encode(), + ), + ( + None, + Some(persisted.as_slice()), + AvgState { count: 2, sum: 7.0 }.encode(), + ), + (None, None, AvgState::default().encode()), + ( + Some(delta.as_slice()), + Some(persisted.as_slice()), + AvgState { + count: 4, + sum: 10.0, + } + .encode(), + ), + ] { + let mut accumulator = expr.create_accumulator().unwrap(); + accumulator + .update_batch(&[ + Arc::new(BinaryArray::from(vec![left])), + Arc::new(BinaryArray::from(vec![right])), + ]) + .unwrap(); + let ScalarValue::Binary(Some(actual)) = accumulator.evaluate().unwrap() else { + panic!("AVG delta merge state must be binary"); + }; + assert_eq!(actual.as_slice(), expected.as_slice()); + } + let mut accumulator = expr.create_accumulator().unwrap(); + assert!( + accumulator + .update_batch(&[ + Arc::new(BinaryArray::from(vec![Some(&[][..])])), + Arc::new(BinaryArray::from(vec![None])), + ]) + .is_err() + ); + let mut accumulator = expr.create_accumulator().unwrap(); + assert!( + accumulator + .update_batch(&[ + Arc::new(BinaryArray::from(vec![None])), + Arc::new(BinaryArray::from(vec![Some(&[][..])])), + ]) + .is_err() + ); + let mut accumulator = expr.create_accumulator().unwrap(); + accumulator + .update_batch(&[ + Arc::new(BinaryArray::from(vec![Some(delta.as_slice())])), + Arc::new(BinaryArray::from(vec![None])), + ]) + .unwrap(); + assert!( + accumulator + .update_batch(&[ + Arc::new(BinaryArray::from(vec![Some(delta.as_slice())])), + Arc::new(BinaryArray::from(vec![Some( + AvgState { + count: u64::MAX, + sum: 1.0 + } + .encode() + .as_slice() + )])), + ]) + .is_err() + ); + } + + #[test] + fn avg_registry_does_not_replace_native_state_registry() { + let avg = FUNCTION_REGISTRY.get_aggr_func(AVG_STATE_NAME).unwrap(); + let native = FUNCTION_REGISTRY + .get_aggr_func(&aggr_state_func_name("avg")) + .unwrap(); + assert_eq!( + avg.return_type(&[DataType::Float64]).unwrap(), + DataType::Binary + ); + assert!(matches!( + native.return_type(&[DataType::Float64]).unwrap(), + DataType::Struct(_) + )); + } +} diff --git a/src/common/meta/src/ddl/create_flow.rs b/src/common/meta/src/ddl/create_flow.rs index 2679e61e74..661158a293 100644 --- a/src/common/meta/src/ddl/create_flow.rs +++ b/src/common/meta/src/ddl/create_flow.rs @@ -568,6 +568,24 @@ pub fn validate_flow_options(flow_task: &CreateFlowTask) -> Result<()> { } } + if let Some(value) = flow_task + .flow_options + .get(FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY) + { + if value != FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE { + value + .parse::() + .map_err(|_| { + UnexpectedSnafu { + err_msg: format!( + "Invalid flow option {FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY}: {value}" + ), + } + .build() + })?; + } + } + defer_on_missing_source(flow_task)?; get_flow_type_from_options(flow_task)?; Ok(()) @@ -770,6 +788,9 @@ pub enum FlowType { pub const FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY: &str = "experimental_enable_incremental_read"; +/// Reserved internal value for Enterprise flows requiring exact sequence ranges. +pub const FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE: &str = + "__greptime_internal_exact_sequence_range"; impl FlowType { pub const BATCHING: &str = "batching"; diff --git a/src/common/meta/src/ddl_manager.rs b/src/common/meta/src/ddl_manager.rs index 8178d68961..72c6b1dd18 100644 --- a/src/common/meta/src/ddl_manager.rs +++ b/src/common/meta/src/ddl_manager.rs @@ -38,7 +38,10 @@ use crate::ddl::alter_logical_tables::AlterLogicalTablesProcedure; use crate::ddl::alter_table::{AlterTableProcedure, RegionRouteChanged, only_enables_skip_wal}; use crate::ddl::comment_on::CommentOnProcedure; use crate::ddl::create_database::{CreateDatabaseMetadataCommitterRef, CreateDatabaseProcedure}; -use crate::ddl::create_flow::CreateFlowProcedure; +use crate::ddl::create_flow::{ + CreateFlowProcedure, FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY, + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE, +}; use crate::ddl::create_logical_tables::CreateLogicalTablesProcedure; use crate::ddl::create_table::CreateTableProcedure; use crate::ddl::create_view::CreateViewProcedure; @@ -112,6 +115,8 @@ pub struct DdlManager { trigger_ddl_manager: Option, #[cfg(feature = "enterprise")] create_flow_handler: Option, + #[cfg(feature = "enterprise")] + drop_flow_handler: Option, } /// This trait is responsible for handling DDL tasks about triggers. e.g., @@ -160,6 +165,20 @@ pub trait CreateFlowHandler: Send + Sync { #[cfg(feature = "enterprise")] pub type CreateFlowHandlerRef = Arc; +/// Hook for classifying and handling DROP FLOW requests. +#[async_trait::async_trait] +pub trait DropFlowHandler: Send + Sync { + async fn drop_flow( + &self, + drop_flow_task: DropFlowTask, + procedure_manager: ProcedureManagerRef, + ddl_context: DdlContext, + procedure_context: ProcedureContext, + ) -> Result; +} + +pub type DropFlowHandlerRef = Arc; + macro_rules! procedure_loader_entry { ($procedure:ident) => { ( @@ -256,9 +275,17 @@ impl DdlManager { trigger_ddl_manager: None, #[cfg(feature = "enterprise")] create_flow_handler: None, + #[cfg(feature = "enterprise")] + drop_flow_handler: None, } } + #[cfg(feature = "enterprise")] + pub fn with_drop_flow_handler(mut self, drop_flow_handler: DropFlowHandlerRef) -> Self { + self.drop_flow_handler = Some(drop_flow_handler); + self + } + #[cfg(feature = "enterprise")] pub fn with_trigger_ddl_manager(mut self, trigger_ddl_manager: TriggerDdlManagerRef) -> Self { self.trigger_ddl_manager = Some(trigger_ddl_manager); @@ -992,6 +1019,17 @@ impl DdlManager { .await } DropFlow(drop_flow_task) => { + #[cfg(feature = "enterprise")] + if let Some(handler) = self.drop_flow_handler.as_ref() { + return handler + .drop_flow( + drop_flow_task, + self.procedure_manager.clone(), + self.ddl_context.clone(), + procedure_context, + ) + .await; + } handle_drop_flow_task(self, drop_flow_task, procedure_context).await } CreateView(create_view_task) => { @@ -1392,6 +1430,19 @@ async fn handle_create_flow_task( query_context: QueryContext, procedure_context: ProcedureContext, ) -> Result { + if create_flow_task + .flow_options + .get(FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY) + .is_some_and(|value| value == FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE) + { + return error::UnexpectedSnafu { + err_msg: format!( + "reserved flow option value for {FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY} is internal" + ), + } + .fail(); + } + #[cfg(feature = "enterprise")] if let Some(handler) = ddl_manager.create_flow_handler.as_ref() { return handler @@ -1578,7 +1629,11 @@ mod tests { use crate::ddl::create_database::{ AtomicCreateOutcome, CreateDatabaseMetadataCommitter, CreateDatabaseProcedure, }; - use crate::ddl::create_flow::CreateFlowProcedure; + #[cfg(feature = "enterprise")] + use crate::ddl::create_flow::{ + CreateFlowProcedure, FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY, + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE, + }; use crate::ddl::create_table::CreateTableProcedure; use crate::ddl::drop_table::DropTableProcedure; use crate::ddl::flow_meta::FlowMetadataAllocator; @@ -1901,6 +1956,32 @@ mod tests { ) } + #[cfg(feature = "enterprise")] + #[tokio::test] + async fn test_reserved_sequence_range_is_rejected_before_enterprise_handler() { + let handler = Arc::new(RecordingCreateFlowHandler::default()); + let ddl_manager = + build_soft_drop_test_ddl_manager().with_create_flow_handler(handler.clone()); + let mut task = test_create_flow_task(); + task.flow_options.insert( + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY.to_string(), + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE.to_string(), + ); + + let result = ddl_manager + .submit_ddl_task( + ExecutorContext { + query_context: Some(QueryContext::default()), + ..Default::default() + }, + SubmitDdlTaskRequest::new(DdlTask::new_create_flow(task)), + ) + .await; + + assert!(result.is_err()); + assert!(handler.tasks.lock().unwrap().is_empty()); + } + #[cfg(feature = "enterprise")] #[tokio::test] async fn test_create_flow_handler_dispatches_without_procedure() { diff --git a/src/flow/src/batching_mode/engine.rs b/src/flow/src/batching_mode/engine.rs index 943092d0fd..c3a21d73a5 100644 --- a/src/flow/src/batching_mode/engine.rs +++ b/src/flow/src/batching_mode/engine.rs @@ -21,7 +21,10 @@ use std::time::Duration; use api::v1::flow::DirtyWindowRequests; use catalog::CatalogManagerRef; use common_error::ext::BoxedError; -use common_meta::ddl::create_flow::{FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY, FlowType}; +use common_meta::ddl::create_flow::{ + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY, + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE, FlowType, +}; use common_meta::key::TableMetadataManagerRef; use common_meta::key::flow::FlowMetadataManagerRef; use common_meta::key::flow::flow_state::FlowStat; @@ -488,21 +491,24 @@ impl BatchingEngine { fn batch_opts_for_flow_options( &self, flow_options: &HashMap, + exact_sequence_range_required: bool, ) -> Result, Error> { let mut batch_opts = (*self.batch_opts).clone(); if let Some(enable_incremental_read) = flow_options.get(FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY) { - batch_opts.experimental_enable_incremental_read = enable_incremental_read - .parse::() - .map_err(|_| { + batch_opts.experimental_enable_incremental_read = if exact_sequence_range_required { + true + } else { + enable_incremental_read.parse::().map_err(|_| { InvalidQuerySnafu { reason: format!( "Invalid flow option {FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY}: {enable_incremental_read}" ), } .build() - })?; + })? + }; } Ok(Arc::new(batch_opts)) @@ -625,8 +631,14 @@ impl BatchingEngine { } ); - let batch_opts = self.batch_opts_for_flow_options(&flow_options)?; - + // The meta layer validates this reserved sentinel before it reaches the + // flownode. Derive the requirement once and pass it directly to task + // config; query-context extensions are irrelevant. + let exact_sequence_range_required = flow_options + .get(FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY) + .is_some_and(|value| value == FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE); + let batch_opts = + self.batch_opts_for_flow_options(&flow_options, exact_sequence_range_required)?; let mut source_table_names = Vec::with_capacity(2); for src_id in source_table_ids { // also check table option to see if ttl!=instant @@ -728,23 +740,42 @@ impl BatchingEngine { eval_schedule, }; - let task = BatchingTask::try_new(task_args)?; + let task = BatchingTask::try_new_with_exact_sequence_range_required( + task_args, + exact_sequence_range_required, + )?; let task_inner = task.clone(); let engine = self.query_engine.clone(); let frontend = self.frontend_client.clone(); + if task.config.exact_sequence_range_required { + ensure!( + task.sequence_range_capable().await?, + UnsupportedSnafu { + reason: format!( + "Flow {flow_id} requires exact sequence-range reads, but a source table lacks the Mito preserve_row_sequence capability" + ), + } + ); + } + // Create the sink before configuring persistence. A persistence-backed sink may // contain ordinary metadata columns supplied by `begin_attempt`, so strict plan/schema - // validation is deferred to execution for that path. OSS flows without a factory keep - // the existing creation-time validation. + // validation is deferred to execution only when persistence is actually created. Flows + // without a created collaborator keep the existing creation-time validation. let table = task.check_or_create_sink_table(&engine, &frontend).await?; let persistence = if let Some(factory) = &self.persistence_factory { let table_info = table.table_info(); let meta = &table_info.meta; - let effective_mode = if task.config.batch_opts.experimental_enable_incremental_read - && task.sequence_range_capable().await + let effective_mode = if task.config.exact_sequence_range_required { + crate::IncrementalMode::SequenceRange + } else if task.config.batch_opts.experimental_enable_incremental_read + && task + .sequence_range_capable() + .await + .is_ok_and(|capable| capable) { crate::IncrementalMode::SequenceRange } else { @@ -782,9 +813,11 @@ impl BatchingEngine { }; factory.create(context).await? } else { - task.validate_sink_table_schema(&engine).await?; None }; + if persistence.is_none() { + task.validate_sink_table_schema(&engine).await?; + } task.set_persistence(persistence).await?; let (start_tx, start_rx) = oneshot::channel(); @@ -1118,17 +1151,24 @@ impl FlowEngine for BatchingEngine { #[cfg(test)] mod tests { use api::v1::flow::{DirtyWindowRequest, TimeRange}; - use catalog::memory::new_memory_catalog_manager; + use catalog::RegisterTableRequest; + use catalog::memory::{MemoryCatalogManager, new_memory_catalog_manager}; use common_meta::key::TableMetadataManager; use common_meta::key::flow::FlowMetadataManager; use common_meta::key::table_route::TableRouteValue; use common_meta::key::test_utils::new_test_table_info_with_name; use common_meta::kv_backend::memory::MemoryKvBackend; + use common_recordbatch::RecordBatch; use common_time::timestamp::TimeUnit; + use datatypes::schema::{ColumnSchema, Schema}; + use datatypes::vectors::{TimestampMillisecondVector, UInt32Vector, VectorRef}; use query::options::QueryOptions; use session::context::QueryContext; use super::*; + use crate::batching_mode::persistence::{ + BatchingAttempt, BatchingPersistence, Factory, FactoryPlugin, RestoreOutcome, + }; use crate::test_utils::create_test_query_engine; struct DropNotify(Option>); @@ -1141,6 +1181,43 @@ mod tests { } } + struct TestPersistenceFactory { + create_persistence: bool, + } + + struct TestPersistence; + + #[async_trait::async_trait] + impl BatchingPersistence for TestPersistence { + async fn restore(&self) -> crate::Result { + Ok(RestoreOutcome::TrustedCheckpoint(BTreeMap::new())) + } + + async fn begin_attempt(&self) -> crate::Result { + Ok(BatchingAttempt::default()) + } + + async fn persist( + &self, + _attempt: BatchingAttempt, + _validated_checkpoints: BTreeMap, + ) -> crate::Result<()> { + Ok(()) + } + } + + #[async_trait::async_trait] + impl Factory for TestPersistenceFactory { + async fn create( + &self, + _context: PersistenceContext, + ) -> crate::Result>> { + Ok(self + .create_persistence + .then_some(Arc::new(TestPersistence) as Arc)) + } + } + async fn new_test_engine() -> BatchingEngine { let kv_backend = Arc::new(MemoryKvBackend::new()); let table_meta = Arc::new(TableMetadataManager::new(kv_backend.clone())); @@ -1161,20 +1238,227 @@ mod tests { ) } - #[tokio::test] - async fn test_flow_option_overrides_incremental_read_switch() { - let engine = new_test_engine().await; + async fn new_test_engine_with_persistence( + persistence_factory: Option, + ) -> BatchingEngine { + let kv_backend = Arc::new(MemoryKvBackend::new()); + let table_meta = Arc::new(TableMetadataManager::new(kv_backend.clone())); + table_meta.init().await.unwrap(); + let flow_meta = Arc::new(FlowMetadataManager::new(kv_backend)); + let query_engine = create_test_query_engine(); + let catalog_manager = query_engine.engine_state().catalog_manager().clone(); + let (frontend_client, _handler) = + FrontendClient::from_empty_grpc_handler(QueryOptions::default()); - let default_opts = engine.batch_opts_for_flow_options(&HashMap::new()).unwrap(); - assert!(!default_opts.experimental_enable_incremental_read); - - let enabled_opts = engine - .batch_opts_for_flow_options(&HashMap::from([( - FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY.to_string(), - "true".to_string(), - )])) + let engine = BatchingEngine::new_with_persistence( + Arc::new(frontend_client), + query_engine, + flow_meta, + table_meta, + catalog_manager, + BatchingModeOptions::default(), + persistence_factory, + ); + engine + .table_meta + .create_table_metadata( + new_test_table_info_with_name(1, "numbers_with_ts"), + TableRouteValue::physical(vec![]), + HashMap::new(), + ) + .await .unwrap(); - assert!(enabled_opts.experimental_enable_incremental_read); + engine + } + + fn register_sink_with_schema(engine: &BatchingEngine, name: &str, extended: bool) { + let mut columns = vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), false), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ]; + let mut vectors: Vec = vec![ + Arc::new(UInt32Vector::from_slice([1_u32])), + Arc::new(TimestampMillisecondVector::from_slice([0_i64])), + ]; + if extended { + columns.push(ColumnSchema::new( + "checkpoint", + ConcreteDataType::uint32_datatype(), + false, + )); + vectors.push(Arc::new(UInt32Vector::from_slice([1_u32]))); + } + let schema = Arc::new(Schema::new(columns)); + let recordbatch = RecordBatch::new(schema, vectors).unwrap(); + let table = table::test_util::MemTable::table(name, recordbatch); + let request = RegisterTableRequest { + catalog: "greptime".to_string(), + schema: "public".to_string(), + table_name: name.to_string(), + table_id: 9000, + table, + }; + engine + .catalog_manager + .as_any() + .downcast_ref::() + .unwrap() + .register_table_sync(request) + .unwrap(); + } + + fn flow_create_args(flow_id: FlowId, sink: &str) -> CreateFlowArgs { + CreateFlowArgs { + flow_id, + sink_table_name: [ + "greptime".to_string(), + "public".to_string(), + sink.to_string(), + ], + source_table_ids: vec![1], + create_if_not_exists: false, + or_replace: false, + expire_after: None, + eval_interval: Some(10), + comment: None, + sql: "SELECT number, ts FROM numbers_with_ts".to_string(), + flow_options: HashMap::new(), + query_ctx: Some(QueryContext::arc().as_ref().clone()), + eval_schedule: None, + } + } + + #[tokio::test] + async fn test_installed_persistence_factory_none_still_validates_sink_schema() { + let engine = new_test_engine_with_persistence(Some(FactoryPlugin(Arc::new( + TestPersistenceFactory { + create_persistence: false, + }, + )))) + .await; + register_sink_with_schema(&engine, "sink_factory_none", true); + + let result = engine + .create_flow_inner(flow_create_args(1, "sink_factory_none")) + .await; + + assert!( + result.is_err(), + "ordinary sink validation must reject mismatch" + ); + assert!(!engine.flow_exist_inner(1).await); + } + + #[tokio::test] + async fn test_matching_persistence_factory_allows_extended_sink_schema() { + let engine = new_test_engine_with_persistence(Some(FactoryPlugin(Arc::new( + TestPersistenceFactory { + create_persistence: true, + }, + )))) + .await; + register_sink_with_schema(&engine, "sink_factory_some", true); + + let result = engine + .create_flow_inner(flow_create_args(2, "sink_factory_some")) + .await + .unwrap(); + + assert_eq!(Some(2), result); + assert!(engine.flow_exist_inner(2).await); + engine.remove_flow_inner(2).await.unwrap(); + } + + #[tokio::test] + async fn test_no_persistence_factory_still_validates_sink_schema() { + let engine = new_test_engine_with_persistence(None).await; + register_sink_with_schema(&engine, "sink_no_factory", true); + + let result = engine + .create_flow_inner(flow_create_args(3, "sink_no_factory")) + .await; + + assert!( + result.is_err(), + "ordinary sink validation must reject mismatch" + ); + assert!(!engine.flow_exist_inner(3).await); + } + + #[tokio::test] + async fn test_flow_option_parser_matrix() { + let engine = new_test_engine().await; + let cases = [ + (None, false, false), + (Some("true"), true, false), + (Some("false"), false, false), + ( + Some(FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE), + true, + true, + ), + (Some("malformed"), false, false), + ]; + for (value, enabled, required) in cases { + let options = value + .map(|value| { + HashMap::from([( + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY.to_string(), + value.to_string(), + )]) + }) + .unwrap_or_default(); + let required_from_validated_sentinel = options + .get(FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY) + .is_some_and(|value| { + value == FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE + }); + match engine.batch_opts_for_flow_options(&options, required_from_validated_sentinel) { + Ok(opts) => { + assert_eq!(opts.experimental_enable_incremental_read, enabled); + assert_eq!(required_from_validated_sentinel, required); + } + Err(_) => assert!(!enabled && !required), + } + } + } + + #[tokio::test] + async fn test_forged_query_context_does_not_enable_exact_sequence_range() { + let engine = new_test_engine_with_persistence(Some(FactoryPlugin(Arc::new( + TestPersistenceFactory { + create_persistence: true, + }, + )))) + .await; + register_sink_with_schema(&engine, "forged_query_context", true); + let mut args = flow_create_args(4, "forged_query_context"); + let mut query_ctx = QueryContext::arc().as_ref().clone(); + query_ctx.set_extension("__old_forged_required_extension", "true"); + args.query_ctx = Some(query_ctx); + + engine.create_flow_inner(args).await.unwrap(); + let task = engine.runtime.read().await.tasks.get(&4).cloned().unwrap(); + assert!(!task.config.exact_sequence_range_required); + engine.remove_flow_inner(4).await.unwrap(); + } + + #[tokio::test] + async fn test_exact_sequence_range_capability_is_checked_before_task_startup() { + let engine = new_test_engine_with_persistence(None).await; + let mut args = flow_create_args(5, "exact_requires_capability"); + args.flow_options.insert( + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY.to_string(), + FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_SEQUENCE_RANGE.to_string(), + ); + + assert!(engine.create_flow_inner(args).await.is_err()); + assert!(!engine.flow_exist_inner(5).await); } #[test] diff --git a/src/flow/src/batching_mode/task.rs b/src/flow/src/batching_mode/task.rs index fde6c37c33..876749dea1 100644 --- a/src/flow/src/batching_mode/task.rs +++ b/src/flow/src/batching_mode/task.rs @@ -152,6 +152,7 @@ pub struct TaskConfig { pub catalog_manager: CatalogManagerRef, pub query_type: QueryType, pub batch_opts: Arc, + pub exact_sequence_range_required: bool, pub flow_eval_interval: Option, /// Typed schedule configuration, pre-parsed at task creation time. pub eval_schedule: Option, @@ -299,7 +300,11 @@ struct ExecuteOnceOutcome { impl BatchingTask { #[allow(clippy::too_many_arguments)] - pub fn try_new( + pub fn try_new(args: TaskArgs<'_>) -> Result { + Self::try_new_with_exact_sequence_range_required(args, false) + } + + pub fn try_new_with_exact_sequence_range_required( TaskArgs { flow_id, query, @@ -315,6 +320,7 @@ impl BatchingTask { flow_eval_interval, eval_schedule, }: TaskArgs<'_>, + exact_sequence_range_required: bool, ) -> Result { let mut state = TaskState::with_dirty_time_windows( query_ctx.clone(), @@ -339,6 +345,7 @@ impl BatchingTask { catalog_manager, output_schema: plan.schema().clone(), query_type: determine_query_type(query, &query_ctx)?, + exact_sequence_range_required, batch_opts, flow_eval_interval, eval_schedule, @@ -663,6 +670,7 @@ impl BatchingTask { async fn execute_logical_plan_unlocked( &self, + engine: &QueryEngineRef, frontend_client: &Arc, plan: &LogicalPlan, dirty_restore: &DirtyRestore, @@ -700,7 +708,7 @@ impl BatchingTask { // For incremental-mode SQL queries, attempt to rewrite the delta aggregate // plan into a safe delta-LEFT-JOIN-sink form before deciding on extensions. let incremental_plan = if coverage.is_incremental_delta() { - self.prepare_plan_for_incremental(&plan).await? + self.prepare_plan_for_incremental(engine, &plan).await? } else { None }; @@ -1431,6 +1439,7 @@ impl BatchingTask { }; let res = self .execute_logical_plan_unlocked( + engine, frontend_client, &new_query.plan, &new_query.dirty_restore, diff --git a/src/flow/src/batching_mode/task/ckpt.rs b/src/flow/src/batching_mode/task/ckpt.rs index c9e8716481..dc8f34c029 100644 --- a/src/flow/src/batching_mode/task/ckpt.rs +++ b/src/flow/src/batching_mode/task/ckpt.rs @@ -12,7 +12,6 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::error::Error as StdError; use std::time::Duration; use client::OutputWithMetrics; @@ -31,42 +30,6 @@ use crate::metrics::{ }; use crate::{Error, FlowId}; -/// Liveness guard: when a fenced repair query fails with a wrapped error whose -/// text indicates a stale snapshot fence (even when `StatusCode::RequestOutdated` -/// was lost through client layers), classify it as `SnapshotFenceExpired` to -/// break the retry loop and force a rebind of the fence high `H`. -/// -/// Long-term the structured `StatusCode` / retry hint path should be preserved -/// end-to-end; this text fallback is a narrow safety measure. -fn matches_stale_snapshot_fence_text(err: &Error) -> bool { - let markers = [ - "STALE_SNAPSHOT_FENCE", - "REBIND_SNAPSHOT_FENCE", - "snapshot upper bound stale", - ]; - // Check the top-level error Display and Debug. - let debug_str = format!("{:?}", err); - let display_str = err.to_string(); - for marker in &markers { - if debug_str.contains(marker) || display_str.contains(marker) { - return true; - } - } - // Walk the error source chain. - let mut source = err.source(); - while let Some(s) = source { - let debug_str = format!("{:?}", s); - let display_str = s.to_string(); - for marker in &markers { - if debug_str.contains(marker) || display_str.contains(marker) { - return true; - } - } - source = s.source(); - } - false -} - impl BatchingTask { /// Classify execution errors into checkpoint fallback reasons. A stale /// snapshot fence is special only for fenced repair chunks. @@ -80,14 +43,6 @@ impl BatchingTask { } else { FlowQueryFallbackReason::StaleCursor } - } else if matches!(coverage, QueryCoverage::FencedRepairChunk { .. }) - && matches_stale_snapshot_fence_text(err) - { - // Narrow text-based fallback for wrapped errors where the - // structured StatusCode::RequestOutdated was lost through - // frontend/client layers. Without this fenced repair will - // retry the same stale `given_seq` every refresh tick. - FlowQueryFallbackReason::SnapshotFenceExpired } else if matches!(coverage, QueryCoverage::IncrementalDelta) { FlowQueryFallbackReason::IncrementalQueryFailure } else { diff --git a/src/flow/src/batching_mode/task/inc.rs b/src/flow/src/batching_mode/task/inc.rs index d019d7380f..1a6a24ac30 100644 --- a/src/flow/src/batching_mode/task/inc.rs +++ b/src/flow/src/batching_mode/task/inc.rs @@ -18,6 +18,7 @@ use common_error::ext::BoxedError; use common_telemetry::debug; use common_telemetry::tracing::warn; use datafusion_expr::{DmlStatement, LogicalPlan}; +use query::QueryEngineRef; use query::options::{ FLOW_INCREMENTAL_AFTER_SEQS, FLOW_INCREMENTAL_MODE, FLOW_INCREMENTAL_MODE_MEMTABLE_ONLY, FLOW_SINK_TABLE_ID, @@ -36,6 +37,9 @@ use crate::batching_mode::utils::{ }; use crate::error::{ExternalSnafu, UnexpectedSnafu}; +// Kept local until the query-side extension enum exposes the exact scan mode. +const FLOW_INCREMENTAL_MODE_SEQUENCE_RANGE: &str = "sequence_range"; + impl BatchingTask { async fn sink_table_id(&self) -> Result { let table = self @@ -69,7 +73,7 @@ impl BatchingTask { /// table that cannot be resolved, is not the mito engine, or lacks the /// option — this returns `false` so the caller keeps the historical /// `memtable_only` mode instead of upgrading. - pub(crate) async fn sequence_range_capable(&self) -> bool { + pub(crate) async fn sequence_range_capable(&self) -> Result { for name in &self.config.source_table_names { let table = match self .config @@ -79,23 +83,16 @@ impl BatchingTask { { Ok(Some(table)) => table, Ok(None) => { - debug!( - "Flow {} source table {} not found; retaining memtable_only incremental mode", - self.config.flow_id, - name.join(".") - ); - return false; - } - Err(err) => { - warn!( - "Flow {} failed to resolve source table {} for sequence_range capability check; \ - retaining memtable_only incremental mode: {:?}", - self.config.flow_id, - name.join("."), - err - ); - return false; + return Err(UnexpectedSnafu { + reason: format!( + "Flow {} source table {} not found for sequence_range capability check", + self.config.flow_id, + name.join(".") + ), + } + .build()); } + Err(err) => Err(BoxedError::new(err)).context(ExternalSnafu)?, }; let info = table.table_info(); @@ -107,10 +104,10 @@ impl BatchingTask { .get(PRESERVE_ROW_SEQUENCE) .is_some_and(|value| value.eq_ignore_ascii_case("true")); if !preserves { - return false; + return Ok(false); } } - !self.config.source_table_names.is_empty() + Ok(!self.config.source_table_names.is_empty()) } /// For incremental-mode SQL queries, attempt to prepare an executable plan @@ -126,6 +123,7 @@ impl BatchingTask { /// incremental safe without a rewrite, so they return `Some(original_plan)`. pub(super) async fn prepare_plan_for_incremental( &self, + engine: &QueryEngineRef, plan: &LogicalPlan, ) -> Result, Error> { let is_incremental_sql = { @@ -208,6 +206,7 @@ impl BatchingTask { let rewritten_inner = match rewrite_incremental_aggregate_with_sink_merge( &inner_plan, &analysis, + engine, sink_table, &self.config.sink_table_name, None, @@ -271,12 +270,42 @@ impl BatchingTask { }; if let Some(checkpoints_json) = incremental_checkpoints_json { + // Select `sequence_range` only when every append-only source table + // proves the `preserve_row_sequence` capability; otherwise retain + // the historical `memtable_only` mode. The `sequence_range` scan + // keeps SSTs and reads the exact (checkpoint, scan-open snapshot] + // row-level delta; the engine fails closed when the capability + // does not hold at scan time. + let capable = match self.sequence_range_capable().await { + Ok(capable) => capable, + Err(err) => { + if self.config.exact_sequence_range_required { + return Err(err); + } + false + } + }; + if self.config.exact_sequence_range_required && !capable { + return Err(UnexpectedSnafu { + reason: format!( + "Flow {} requires exact sequence-range reads, but source capability was revoked", + self.config.flow_id + ), + } + .build()); + } let sink_table_id = self.sink_table_id().await?; + let incremental_mode = if capable { + debug!( + "Flow {} selected sequence_range incremental mode", + self.config.flow_id + ); + FLOW_INCREMENTAL_MODE_SEQUENCE_RANGE + } else { + FLOW_INCREMENTAL_MODE_MEMTABLE_ONLY + }; extensions.push((FLOW_SINK_TABLE_ID, sink_table_id.to_string())); - extensions.push(( - FLOW_INCREMENTAL_MODE, - FLOW_INCREMENTAL_MODE_MEMTABLE_ONLY.to_string(), - )); + extensions.push((FLOW_INCREMENTAL_MODE, incremental_mode.to_string())); extensions.push((FLOW_INCREMENTAL_AFTER_SEQS, checkpoints_json)); } diff --git a/src/flow/src/batching_mode/task/test.rs b/src/flow/src/batching_mode/task/test.rs index 2d0064be88..937376c1d0 100644 --- a/src/flow/src/batching_mode/task/test.rs +++ b/src/flow/src/batching_mode/task/test.rs @@ -13,8 +13,6 @@ // limitations under the License. use std::collections::{BTreeMap, BTreeSet, HashMap}; -use std::sync::Arc; -use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use catalog::RegisterTableRequest; use catalog::memory::MemoryCatalogManager; @@ -39,18 +37,14 @@ use query::options::{ use session::context::QueryContext; use snafu::ResultExt; use table::test_util::MemTable; -use table::{Table, TableRef}; -use tokio::sync::Notify; use super::*; -use crate::Result; use crate::batching_mode::checkpoint::{ CHECKPOINT_DECISION_ADVANCE, CHECKPOINT_DECISION_FALLBACK, CHECKPOINT_REASON_NONE, FlowCheckpointDecision, FlowQueryFallbackReason, }; use crate::batching_mode::eval_schedule::{FlowMissedTickPolicy, FlowScheduleConfig}; -use crate::batching_mode::persistence::{BatchingAttempt, BatchingPersistence, RestoreOutcome}; -use crate::batching_mode::state::{CheckpointMode, TaskStateCheckpointSnapshot}; +use crate::batching_mode::state::CheckpointMode; use crate::batching_mode::time_window::find_time_window_expr; use crate::test_utils::create_test_query_engine; @@ -91,6 +85,18 @@ async fn new_test_task_engine_and_plan_with_query_and_opts( query: &str, sink_table: &str, batch_opts: Arc, +) -> TestTaskParts { + new_test_task_engine_and_plan_with_query_and_opts_and_required( + query, sink_table, batch_opts, false, + ) + .await +} + +async fn new_test_task_engine_and_plan_with_query_and_opts_and_required( + query: &str, + sink_table: &str, + batch_opts: Arc, + exact_sequence_range_required: bool, ) -> TestTaskParts { let query_engine = create_test_query_engine(); let ctx = QueryContext::arc(); @@ -104,29 +110,32 @@ async fn new_test_task_engine_and_plan_with_query_and_opts( .unwrap(); let (_tx, rx) = tokio::sync::oneshot::channel(); - let task = BatchingTask::try_new(TaskArgs { - flow_id: 1, - query, - plan: plan.clone(), - time_window_expr: None, - expire_after: None, - sink_table_name: [ - "greptime".to_string(), - "public".to_string(), - sink_table.to_string(), - ], - source_table_names: vec![[ - "greptime".to_string(), - "public".to_string(), - "numbers_with_ts".to_string(), - ]], - query_ctx: ctx, - catalog_manager: query_engine.engine_state().catalog_manager().clone(), - shutdown_rx: rx, - batch_opts, - flow_eval_interval: None, - eval_schedule: None, - }) + let task = BatchingTask::try_new_with_exact_sequence_range_required( + TaskArgs { + flow_id: 1, + query, + plan: plan.clone(), + time_window_expr: None, + expire_after: None, + sink_table_name: [ + "greptime".to_string(), + "public".to_string(), + sink_table.to_string(), + ], + source_table_names: vec![[ + "greptime".to_string(), + "public".to_string(), + "numbers_with_ts".to_string(), + ]], + query_ctx: ctx, + catalog_manager: query_engine.engine_state().catalog_manager().clone(), + shutdown_rx: rx, + batch_opts, + flow_eval_interval: None, + eval_schedule: None, + }, + exact_sequence_range_required, + ) .unwrap(); TestTaskParts { @@ -397,50 +406,6 @@ fn flow_error_with_status(status_code: StatusCode) -> Error { .unwrap_err() } -/// Test-only error that carries a non-RequestOutdated status code but -/// displays a stale-snapshot-fence marker string, simulating the real-world -/// scenario where the structured status code is lost through frontend/client -/// wrapping layers. -#[derive(Debug)] -struct StaleFenceTextError { - code: StatusCode, - message: String, -} - -impl std::fmt::Display for StaleFenceTextError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.message) - } -} - -impl std::error::Error for StaleFenceTextError {} - -impl common_error::ext::ErrorExt for StaleFenceTextError { - fn status_code(&self) -> StatusCode { - self.code - } - fn as_any(&self) -> &dyn std::any::Any { - self - } -} - -impl common_error::ext::StackError for StaleFenceTextError { - fn debug_fmt(&self, _: usize, _: &mut Vec) {} - fn next(&self) -> Option<&dyn common_error::ext::StackError> { - None - } -} - -fn flow_error_with_code_and_text(code: StatusCode, text: &str) -> Error { - let inner = StaleFenceTextError { - code, - message: text.to_string(), - }; - Err::<(), _>(BoxedError::new(inner)) - .context(crate::error::ExternalSnafu) - .unwrap_err() -} - fn dirty_range(start: i64, end: i64) -> DirtyTimeWindows { let mut dirty = DirtyTimeWindows::default(); dirty.add_window( @@ -1453,78 +1418,6 @@ fn test_query_failure_reason_distinguishes_fenced_repair_stale_fence() { ); } -/// Wrapped errors carrying stale snapshot fence marker text in their -/// Display/Debug chain should be classified as `SnapshotFenceExpired` on -/// fenced repair coverage, even when the structured `StatusCode::RequestOutdated` -/// was lost through client layering. This prevents an infinite retry loop -/// where the fenced chunk re-sends the same stale `given_seq` every tick. -#[test] -fn test_query_failure_reason_text_fallback_stale_snapshot_fence() { - let high = BTreeMap::new(); - let fenced = QueryCoverage::FencedRepairChunk { high: high.clone() }; - - // STALE_SNAPSHOT_FENCE marker with a non-RequestOutdated status code - let err = flow_error_with_code_and_text( - StatusCode::Internal, - "gRPC error: STALE_SNAPSHOT_FENCE: snapshot upper bound stale, region: 1024/0", - ); - assert_eq!( - BatchingTask::query_failure_reason(&err, &fenced), - FlowQueryFallbackReason::SnapshotFenceExpired - ); - - // REBIND_SNAPSHOT_FENCE marker - let err = flow_error_with_code_and_text( - StatusCode::Internal, - "STALE_SNAPSHOT_FENCE ... retry_hint: REBIND_SNAPSHOT_FENCE", - ); - assert_eq!( - BatchingTask::query_failure_reason(&err, &fenced), - FlowQueryFallbackReason::SnapshotFenceExpired - ); - - // snapshot upper bound stale marker (the natural-language fragment) - let err = flow_error_with_code_and_text( - StatusCode::Internal, - "query failed: snapshot upper bound stale, consider rebinding", - ); - assert_eq!( - BatchingTask::query_failure_reason(&err, &fenced), - FlowQueryFallbackReason::SnapshotFenceExpired - ); - - // Fenced coverage with a generic wrapped error (no stale-fence marker) → - // still QueryFailure - let generic_err = - flow_error_with_code_and_text(StatusCode::Internal, "some transient network error"); - assert_eq!( - BatchingTask::query_failure_reason(&generic_err, &fenced), - FlowQueryFallbackReason::QueryFailure - ); - - // Non-fenced incremental coverage with stale-fence marker text must NOT - // classify as SnapshotFenceExpired; it should remain IncrementalQueryFailure. - let err = flow_error_with_code_and_text( - StatusCode::Internal, - "STALE_SNAPSHOT_FENCE blob in unexpected context", - ); - assert_eq!( - BatchingTask::query_failure_reason(&err, &QueryCoverage::IncrementalDelta), - FlowQueryFallbackReason::IncrementalQueryFailure - ); - - // Existing RequestOutdated behavior is unchanged. - let outdated_err = flow_error_with_status(StatusCode::RequestOutdated); - assert_eq!( - BatchingTask::query_failure_reason(&outdated_err, &fenced), - FlowQueryFallbackReason::SnapshotFenceExpired - ); - assert_eq!( - BatchingTask::query_failure_reason(&outdated_err, &QueryCoverage::IncrementalDelta), - FlowQueryFallbackReason::StaleCursor - ); -} - #[test] fn test_fenced_repair_coverage_produces_snapshot_seq_map_for_distributed_metadata_path() { // Covers the metadata boundary between QueryCoverage and the @@ -1672,82 +1565,6 @@ fn test_fenced_repair_transient_non_stale_failure_retries_same_high() { ); } -/// When `query_failure_reason` classifies a wrapped error as -/// `SnapshotFenceExpired` via the text-marker fallback (not via -/// `StatusCode::RequestOutdated`), the state machine must still -/// abandon the fenced repair and produce a `ScopedBaseRepair` plan -/// next, exactly like the structured-code path. -#[tokio::test] -async fn test_text_fallback_stale_fence_produces_scoped_base_repair() { - let TestTaskParts { - task, - query_engine, - .. - } = new_time_window_test_task_with_query( - "SELECT number, date_bin(INTERVAL '5 second', ts) AS time_window FROM numbers_with_ts GROUP BY time_window, number", - ) - .await; - let high = BTreeMap::from([(1_u64, 10_u64), (2_u64, 20_u64)]); - let filter = { - let mut state = task.state.write().unwrap(); - state - .dirty_time_windows - .add_window(Timestamp::new_second(10), Some(Timestamp::new_second(15))); - state - .dirty_time_windows - .add_window(Timestamp::new_second(100), Some(Timestamp::new_second(105))); - state.start_fenced_repair(high.clone()).unwrap(); - next_fenced_repair_filter(&mut state, 1) - }; - - // Construct a wrapped error that hits the text fallback (non-RequestOutdated - // status code with STALE_SNAPSHOT_FENCE marker text). - let err = flow_error_with_code_and_text( - StatusCode::Internal, - "STALE_SNAPSHOT_FENCE: snapshot upper bound stale, retry_hint: REBIND_SNAPSHOT_FENCE", - ); - let coverage = QueryCoverage::FencedRepairChunk { high }; - let reason = BatchingTask::query_failure_reason(&err, &coverage); - assert_eq!(reason, FlowQueryFallbackReason::SnapshotFenceExpired); - - { - let mut state = task.state.write().unwrap(); - let decision = BatchingTask::apply_query_failure_to_state( - &mut state, - std::time::Duration::from_millis(1), - &coverage, - reason, - ); - assert_eq!( - decision, - Some(FlowCheckpointDecision::FallbackToFullSnapshot { - previous_mode: CheckpointMode::FullSnapshot, - reason: FlowQueryFallbackReason::SnapshotFenceExpired, - }) - ); - assert!(state.pending_fenced_repair().is_none()); - - // Simulate the outer execution failure restore for the in-flight chunk. - state.restore_scoped_windows(&filter); - } - - let plan = task - .gen_query_with_time_window( - query_engine, - &aggregate_time_window_sink_schema(), - &[], - false, - Some(1), - ) - .await - .unwrap() - .expect("text-fallback stale fence should restore dirty windows for a fresh scoped repair"); - assert!( - matches!(plan.coverage, QueryCoverage::ScopedBaseRepair), - "next plan after text-fallback stale fence should be ScopedBaseRepair" - ); -} - #[test] fn test_checkpoint_decision_labels_are_stable() { let advance = FlowCheckpointDecision::AdvancedIncremental { @@ -1779,6 +1596,35 @@ fn test_checkpoint_decision_labels_are_stable() { ); } +#[tokio::test] +async fn test_exact_required_attempt_rejects_revoked_capability_without_extensions() { + let task = new_test_task_engine_and_plan_with_query_and_opts_and_required( + "SELECT number, ts FROM numbers_with_ts", + "exact_required_revoked", + incremental_batch_opts(), + true, + ) + .await + .into_task_and_plan() + .0; + + task.state + .write() + .unwrap() + .advance_checkpoints(HashMap::from([(1_u64, 10_u64)])); + let checkpoints_before = task.state.read().unwrap().checkpoints().clone(); + + let err = task + .build_flow_query_extensions(true, true) + .await + .unwrap_err(); + assert!(err.to_string().contains("requires exact sequence-range")); + assert_eq!( + task.state.read().unwrap().checkpoints(), + &checkpoints_before + ); +} + #[tokio::test] async fn test_build_flow_query_extensions_switches_with_checkpoint_mode() { let (task, _) = new_test_task_engine_and_plan_with_query( @@ -2348,7 +2194,10 @@ async fn test_prepare_plan_for_incremental_disables_on_non_aggregate() { CheckpointMode::Incremental ); - let incremental_plan = task.prepare_plan_for_incremental(&dml_plan).await.unwrap(); + let incremental_plan = task + .prepare_plan_for_incremental(&query_engine, &dml_plan) + .await + .unwrap(); assert!(incremental_plan.is_none()); let state = task.state.read().unwrap(); assert!(state.is_incremental_disabled()); @@ -2425,6 +2274,7 @@ async fn test_unsafe_incremental_plan_skip_restores_dirty_without_query() { let result = task .execute_logical_plan_unlocked( + &query_engine, &Arc::new(frontend_client), &dml_plan, &dirty_restore, @@ -2512,7 +2362,7 @@ async fn test_prepare_plan_for_incremental_group_by_without_merge_columns_uses_o .advance_checkpoints(HashMap::from([(1_u64, 10_u64)])); let incremental_plan = task - .prepare_plan_for_incremental(&dml_plan) + .prepare_plan_for_incremental(&query_engine, &dml_plan) .await .unwrap() .expect("plain GROUP BY is incremental-safe without a rewrite"); @@ -2557,7 +2407,10 @@ async fn test_auto_created_sql_aggregate_sink_reaches_incremental_safe() { .write() .unwrap() .advance_checkpoints(HashMap::from([(1_u64, 10_u64)])); - let incremental_plan = task.prepare_plan_for_incremental(&dml_plan).await.unwrap(); + let incremental_plan = task + .prepare_plan_for_incremental(&query_engine, &dml_plan) + .await + .unwrap(); let incremental_safe = incremental_plan.is_some(); assert!(incremental_safe); @@ -2687,309 +2540,3 @@ async fn test_insert_plan_matching_failure_restores_consumed_dirty_marker() { std::time::Duration::from_secs(5) ); } - -struct TestPersistence { - restores: AtomicUsize, - begins: AtomicUsize, - persists: AtomicUsize, - fail_persist: AtomicBool, -} - -struct BlockingPersistence { - persists: AtomicUsize, - fail: AtomicBool, - started: Notify, - release: Notify, -} - -#[async_trait::async_trait] -impl BatchingPersistence for BlockingPersistence { - async fn restore(&self) -> Result { - Ok(RestoreOutcome::TrustedCheckpoint(BTreeMap::new())) - } - - async fn begin_attempt(&self) -> Result { - Ok(BatchingAttempt { - ordinary_values: BTreeMap::new(), - }) - } - - async fn persist( - &self, - _attempt: BatchingAttempt, - _checkpoints: BTreeMap, - ) -> Result<()> { - self.persists.fetch_add(1, Ordering::SeqCst); - self.started.notify_one(); - self.release.notified().await; - if self.fail.load(Ordering::SeqCst) { - Err(crate::Error::External { - source: BoxedError::new(MockError::new(StatusCode::Internal)), - location: snafu::location!(), - }) - } else { - Ok(()) - } - } -} - -fn blocking_persistence(fail: bool) -> Arc { - Arc::new(BlockingPersistence { - persists: AtomicUsize::new(0), - fail: AtomicBool::new(fail), - started: Notify::new(), - release: Notify::new(), - }) -} - -fn install_persistence(task: &BatchingTask, persistence: Arc) { - *task.persistence.write().unwrap() = Some(persistence); -} - -fn candidate_transaction_states( - task: &BatchingTask, -) -> (TaskStateCheckpointSnapshot, TaskStateCheckpointSnapshot) { - let mut state = task.state.write().unwrap(); - state.advance_checkpoints(HashMap::from([(1, 10)])); - state.request_full_repair(); - let snapshot = state.checkpoint_snapshot(); - let mut candidate = snapshot.clone(); - candidate.checkpoint_mode = CheckpointMode::Incremental; - candidate.checkpoints = BTreeMap::from([(1, 20)]); - candidate.last_query_duration = Duration::from_millis(42); - candidate.last_exec_time_millis = Some(42); - state.restore_checkpoint_snapshot(snapshot.clone()); - (snapshot, candidate) -} - -fn test_attempt() -> BatchingAttempt { - BatchingAttempt { - ordinary_values: BTreeMap::new(), - } -} - -fn assert_candidate_state(task: &BatchingTask) { - let state = task.state.read().unwrap(); - assert_eq!(state.checkpoints(), &BTreeMap::from([(1, 20)])); - assert_eq!(state.last_query_duration(), Duration::from_millis(42)); - assert!(state.last_execution_time_millis().is_some()); -} - -#[tokio::test] -async fn test_checkpoint_persist_success_commits_candidate_after_release() { - let task = new_test_task_engine_and_plan_with_query( - "SELECT number, ts FROM numbers_with_ts", - "missing_sink", - ) - .await - .task; - let persistence = blocking_persistence(false); - install_persistence(&task, persistence.clone()); - let (snapshot, candidate) = candidate_transaction_states(&task); - let task_for_persist = task.clone(); - let attempt = test_attempt(); - let persist = tokio::spawn(async move { - task_for_persist - .persist_checkpoint_candidate( - snapshot.clone(), - candidate, - Some(&attempt), - DirtyRestore::Unscoped(DirtyTimeWindows::default()), - ) - .await - }); - persistence.started.notified().await; - { - let state = task.state.read().unwrap(); - assert_eq!(state.checkpoints(), &BTreeMap::from([(1, 10)])); - } - persistence.release.notify_one(); - persist.await.unwrap().unwrap(); - assert_candidate_state(&task); - assert_eq!(persistence.persists.load(Ordering::SeqCst), 1); -} - -#[tokio::test] -async fn test_checkpoint_persist_error_restores_pre_state_and_unions_dirty_once() { - let task = new_test_task_engine_and_plan_with_query( - "SELECT number, ts FROM numbers_with_ts", - "missing_sink", - ) - .await - .task; - let persistence = blocking_persistence(true); - install_persistence(&task, persistence.clone()); - let (snapshot, candidate) = candidate_transaction_states(&task); - let detached = dirty_range(1, 2); - let live = dirty_range(3, 4); - let task_for_persist = task.clone(); - let attempt = test_attempt(); - let persist = tokio::spawn(async move { - task_for_persist - .persist_checkpoint_candidate( - snapshot, - candidate, - Some(&attempt), - DirtyRestore::Unscoped(detached), - ) - .await - }); - persistence.started.notified().await; - task.state - .write() - .unwrap() - .dirty_time_windows - .add_dirty_windows(&live); - persistence.release.notify_one(); - assert!(persist.await.unwrap().is_err()); - let state = task.state.read().unwrap(); - assert_eq!(state.checkpoints(), &BTreeMap::from([(1, 10)])); - assert_eq!(state.dirty_time_windows.len(), 2); - assert_eq!(persistence.persists.load(Ordering::SeqCst), 1); -} - -#[tokio::test] -async fn test_full_repair_checkpoint_persist_handles_success_and_error() { - for fail in [false, true] { - let task = new_test_task_engine_and_plan_with_query( - "SELECT number, ts FROM numbers_with_ts", - "missing_sink", - ) - .await - .task; - let persistence = blocking_persistence(fail); - install_persistence(&task, persistence.clone()); - let (snapshot, candidate) = candidate_transaction_states(&task); - let dirty = dirty_range(1, 2); - let task_for_persist = task.clone(); - let attempt = test_attempt(); - let persist = tokio::spawn(async move { - task_for_persist - .persist_checkpoint_candidate( - snapshot, - candidate, - Some(&attempt), - DirtyRestore::FullRepair(dirty), - ) - .await - }); - persistence.started.notified().await; - task.state - .write() - .unwrap() - .dirty_time_windows - .add_dirty_windows(&dirty_range(1, 2)); - persistence.release.notify_one(); - let result = persist.await.unwrap(); - assert_eq!(result.is_err(), fail); - let state = task.state.read().unwrap(); - assert_eq!(state.dirty_time_windows.len(), 1); - assert_eq!(state.full_repair_required(), fail); - assert_eq!(persistence.persists.load(Ordering::SeqCst), 1); - } -} - -#[async_trait::async_trait] -impl BatchingPersistence for TestPersistence { - async fn restore(&self) -> Result { - self.restores.fetch_add(1, Ordering::SeqCst); - Ok(RestoreOutcome::TrustedCheckpoint(BTreeMap::from([(1, 2)]))) - } - async fn begin_attempt(&self) -> Result { - self.begins.fetch_add(1, Ordering::SeqCst); - Ok(BatchingAttempt { - ordinary_values: BTreeMap::new(), - }) - } - async fn persist( - &self, - _attempt: BatchingAttempt, - checkpoints: BTreeMap, - ) -> Result<()> { - assert_eq!(checkpoints, BTreeMap::from([(1, 2)])); - self.persists.fetch_add(1, Ordering::SeqCst); - if self.fail_persist.load(Ordering::SeqCst) { - Err(crate::Error::External { - source: BoxedError::new(MockError::new(StatusCode::Internal)), - location: snafu::location!(), - }) - } else { - Ok(()) - } - } -} - -#[tokio::test] -async fn test_persistence_restore_is_wired() { - let parts = new_test_task_engine_and_plan_with_query( - "SELECT number, ts FROM numbers_with_ts", - "missing_sink", - ) - .await; - let state = Arc::new(TestPersistence { - restores: AtomicUsize::new(0), - begins: AtomicUsize::new(0), - persists: AtomicUsize::new(0), - fail_persist: AtomicBool::new(false), - }); - let config = state.clone(); - parts.task.set_persistence(Some(config)).await.unwrap(); - assert_eq!(state.restores.load(Ordering::SeqCst), 1); - assert_eq!( - parts.task.state.read().unwrap().checkpoints(), - &BTreeMap::from([(1, 2)]) - ); -} - -#[tokio::test] -async fn test_full_repair_restore_is_sticky_and_unfiltered() { - let parts = new_time_window_test_task_with_query( - "SELECT number, date_bin(INTERVAL '5 second', ts) AS time_window FROM numbers_with_ts GROUP BY time_window, number", - ) - .await; - // Use a persistence implementation whose restore requests a full repair. - struct FullRepairPersistence; - #[async_trait::async_trait] - impl BatchingPersistence for FullRepairPersistence { - async fn restore(&self) -> Result { - Ok(RestoreOutcome::FullRepair) - } - async fn begin_attempt(&self) -> Result { - Ok(BatchingAttempt { - ordinary_values: BTreeMap::new(), - }) - } - async fn persist( - &self, - _attempt: BatchingAttempt, - _checkpoints: BTreeMap, - ) -> Result<()> { - Ok(()) - } - } - let config = Arc::new(FullRepairPersistence); - let sink = aggregate_time_window_sink_schema(); - parts - .task - .state - .write() - .unwrap() - .dirty_time_windows - .add_window(Timestamp::new_second(10), Some(Timestamp::new_second(15))); - parts.task.set_persistence(Some(config)).await.unwrap(); - let plan = parts - .task - .gen_query_with_time_window(parts.query_engine, &sink, &[], false, Some(1)) - .await - .unwrap() - .expect("full repair should always produce a plan"); - - assert!(matches!(plan.coverage, QueryCoverage::UnfilteredFull)); - assert!(matches!(plan.dirty_restore, DirtyRestore::FullRepair(_))); - let plan_text = plan.plan.to_string(); - assert!(!plan_text.contains("Filter:")); - assert!(!plan_text.contains("TimestampMillisecond(")); - // Full repair detaches the consumed dirty ownership from the live signal. - assert_eq!(parts.task.state.read().unwrap().dirty_time_windows.len(), 0); - assert!(parts.task.state.read().unwrap().full_repair_required()); -} diff --git a/src/flow/src/batching_mode/utils.rs b/src/flow/src/batching_mode/utils.rs index 51b5961a40..66256e3618 100644 --- a/src/flow/src/batching_mode/utils.rs +++ b/src/flow/src/batching_mode/utils.rs @@ -21,6 +21,7 @@ use catalog::CatalogManagerRef; use common_error::ext::BoxedError; use common_function::aggrs::aggr_wrapper::get_aggr_func; use common_telemetry::debug; +use datafusion::arrow::datatypes::DataType as ArrowDataType; use datafusion::datasource::DefaultTableSource; use datafusion::error::Result as DfResult; use datafusion::logical_expr::Expr; @@ -67,6 +68,9 @@ mod test; /// `max(numbers_with_ts.number)`. #[derive(Debug, Clone, PartialEq, Eq)] pub struct IncrementalAggregateMergeColumn { + /// Delta-plan field containing the aggregate result/state column. Repeated + /// projections intentionally share this field while using distinct outputs. + pub input_field_name: String, /// Final output/sink field name for the aggregate result/state column. /// pub output_field_name: String, @@ -74,9 +78,10 @@ pub struct IncrementalAggregateMergeColumn { } impl IncrementalAggregateMergeColumn { - /// Create a new merge column. + /// Create a new merge column whose delta and output fields have the same name. pub fn new(output_field_name: String, merge_op: IncrementalAggregateMergeOp) -> Self { Self { + input_field_name: output_field_name.clone(), output_field_name, merge_op, } @@ -93,6 +98,7 @@ pub enum IncrementalAggregateMergeOp { BitAnd, BitOr, BitXor, + AvgDeltaMerge, } /// Analysis result for an incremental aggregate plan. @@ -114,24 +120,6 @@ pub struct IncrementalAggregateAnalysis { pub unsupported_exprs: Vec, } -/// Recursively find all `Expr::Column` names inside an expression tree. -/// Only recurses into wrappers that are merge-transparent. -/// Non-transparent wrappers (e.g., `ScalarFunction`, `Negative`, `Cast`) are -/// intentionally not recursed into since their merge semantics would be -/// incorrect. -/// -/// `Cast`/`TryCast` are intentionally opaque: merging already-casted aggregate -/// outputs is not generally equivalent to casting the final merged aggregate. -fn find_column_names(expr: &Expr, names: &mut Vec) { - match expr { - Expr::Column(col) => { - names.push(col.name.clone()); - } - Expr::Alias(alias) => find_column_names(&alias.expr, names), - _ => {} - } -} - fn unqualified_col(name: impl Into) -> Expr { Expr::Column(Column::from_name(name.into())) } @@ -231,8 +219,10 @@ fn check_input_plan_shape(plan: &LogicalPlan) -> Result<(), String> { #[derive(Debug, Default)] struct OutputProjectionInfo { has_top_level_projection: bool, + /// Aggregate expression name and projected output field, in projection order. + aggregate_outputs: Vec<(String, String)>, + /// Original single-instance resolver mapping, retained for compatibility. output_aliases: HashMap, - duplicate_aggregate_aliases: BTreeSet, literal_columns: HashSet, output_field_names: Vec, } @@ -271,41 +261,36 @@ fn collect_output_projection_info(plan: &LogicalPlan) -> OutputProjectionInfo { for expr in &projection.expr { match expr { Expr::Alias(alias) => { - // Alias resolution has three cases: - // - 0 Column refs (e.g., literal `42 AS lit`): record literal output - // - 1 Column ref: record the mapping (e.g., `sum(x) AS total`) - // - >1 Column refs (e.g., `COALESCE(sum(x), sum(y))`): - // skip — ambiguous merge semantics + // Only a direct aggregate output column has the same + // merge semantics as the original resolver. In particular, + // do not mine aggregate columns through CAST/TryCast or + // other output wrappers. let alias_name = alias.name.clone(); - let mut col_names = Vec::new(); - find_column_names(&alias.expr, &mut col_names); - match col_names.len() { - 0 if is_passthrough_output_column(&alias_name, alias.expr.as_ref()) => { - projection_info.literal_columns.insert(alias_name); - } - 1 => { - if let Some(col_name) = col_names.into_iter().next() { - if let Some(existing_alias) = output_aliases.get(&col_name) { - if existing_alias != &alias_name { - projection_info.duplicate_aggregate_aliases.insert(format!( - "same aggregate output {col_name} is used by multiple aliases: {existing_alias}, {alias_name}" - )); - } - } else { - output_aliases.insert(col_name, alias_name); - } - } - } - _ => {} + if let Expr::Column(column) = alias.expr.as_ref() { + output_aliases + .entry(column.name.clone()) + .or_insert_with(|| alias_name.clone()); + projection_info + .aggregate_outputs + .push((column.name.clone(), alias_name)); + } else if let Expr::Alias(inner_alias) = alias.expr.as_ref() + && inner_alias.name.eq_ignore_ascii_case("count(*)") + && let Expr::Column(column) = inner_alias.expr.as_ref() + { + output_aliases + .entry(column.name.clone()) + .or_insert_with(|| alias_name.clone()); + projection_info + .aggregate_outputs + .push((column.name.clone(), alias_name)); + } else if is_passthrough_output_column(&alias_name, alias.expr.as_ref()) { + projection_info.literal_columns.insert(alias_name); } - - // If >1 column references detected (e.g., COALESCE(sum(x), sum(y))), - // intentionally skip alias mapping — the merge semantics are ambiguous. } Expr::Column(col) => { - output_aliases - .entry(col.name.clone()) - .or_insert(col.name.clone()); + projection_info + .aggregate_outputs + .push((col.name.clone(), col.name.clone())); } Expr::Literal(_, _) => { projection_info @@ -349,7 +334,10 @@ fn is_literal_or_cast_literal(expr: &Expr) -> bool { } } -fn merge_op_for_aggregate_expr(aggr_expr: &Expr) -> Result { +fn merge_op_for_aggregate_expr( + aggr_expr: &Expr, + input_schema: &DFSchema, +) -> Result { let Some(aggr_func) = get_aggr_func(aggr_expr) else { return Err(aggr_expr.to_string()); }; @@ -372,28 +360,50 @@ fn merge_op_for_aggregate_expr(aggr_expr: &Expr) -> Result Ok(IncrementalAggregateMergeOp::BitAnd), "bit_or" => Ok(IncrementalAggregateMergeOp::BitOr), "bit_xor" => Ok(IncrementalAggregateMergeOp::BitXor), + "avg_state" => match aggr_func.params.args.as_slice() { + [_] => Ok(IncrementalAggregateMergeOp::AvgDeltaMerge), + _ => Err(aggr_expr.to_string()), + }, + "avg_merge" => match aggr_func.params.args.as_slice() { + [arg] if arg.get_type(input_schema).ok() == Some(ArrowDataType::Binary) => { + Ok(IncrementalAggregateMergeOp::AvgDeltaMerge) + } + _ => Err(aggr_expr.to_string()), + }, _ => Err(aggr_expr.to_string()), } } -fn resolve_aggregate_output_field_name( +fn resolve_aggregate_output_fields( aggr_expr: &Expr, projection_info: &OutputProjectionInfo, output_field_name_set: &HashSet, -) -> Option { +) -> Vec<(String, String)> { // qualified_name() returns (Option, String) where the second // element is the unqualified column/alias name. This relies on // DataFusion's internal naming convention: aggregate expressions - // emit a column named after the aggregate itself (e.g. "SUM(x)"), - // which matches what the projection aliases reference. + // emit a column named after the aggregate itself (e.g. "SUM(x)"). + // Keep every matching projection occurrence because DataFusion can share + // one aggregate input field for identical expressions. let raw_name = aggr_expr.qualified_name().1; - if let Some(alias) = projection_info.output_aliases.get(&raw_name) { - Some(alias.clone()) - } else if !projection_info.has_top_level_projection && output_field_name_set.contains(&raw_name) - { - Some(raw_name) + if projection_info.has_top_level_projection { + let outputs = projection_info + .aggregate_outputs + .iter() + .filter(|(input_name, _)| input_name == &raw_name) + .cloned() + .collect::>(); + if outputs.len() > 1 { + outputs + } else if let Some(alias) = projection_info.output_aliases.get(&raw_name) { + vec![(raw_name, alias.clone())] + } else { + outputs + } + } else if output_field_name_set.contains(&raw_name) { + vec![(raw_name.clone(), raw_name)] } else { - None + vec![] } } @@ -479,7 +489,6 @@ pub fn analyze_incremental_aggregate_plan( .map(|name| format!("duplicate output field name: {name}")) .collect::>(); unsupported_exprs.push(reason); - unsupported_exprs.extend(projection_info.duplicate_aggregate_aliases.iter().cloned()); return Ok(Some(IncrementalAggregateAnalysis { group_key_names, merge_columns: vec![], @@ -513,27 +522,48 @@ pub fn analyze_incremental_aggregate_plan( aggregate, &group_key_names, )); - unsupported_exprs.extend(projection_info.duplicate_aggregate_aliases.iter().cloned()); for aggr_expr in aggr_exprs { - let merge_op = match merge_op_for_aggregate_expr(&aggr_expr) { + let merge_op = match merge_op_for_aggregate_expr(&aggr_expr, aggregate.input.schema()) { Ok(merge_op) => merge_op, Err(reason) => { unsupported_exprs.push(reason); continue; } }; - let Some(output_field_name) = resolve_aggregate_output_field_name( - &aggr_expr, - &projection_info, - &output_field_name_set, - ) else { + let aggregate_outputs = + resolve_aggregate_output_fields(&aggr_expr, &projection_info, &output_field_name_set); + if aggregate_outputs.is_empty() { unsupported_exprs.push(aggr_expr.to_string()); continue; + } + let Some((_, input_field_name)) = aggregate_outputs.first() else { + continue; }; - merge_columns.push(IncrementalAggregateMergeColumn::new( - output_field_name, - merge_op, - )); + // The old single-alias resolver selected the projected output name as + // the delta field. Keep that exact field for the shared input and only + // vary the final sink/output alias for repeated projections. + let input_field_name = input_field_name.clone(); + for (_, output_field_name) in aggregate_outputs { + merge_columns.push(IncrementalAggregateMergeColumn { + input_field_name: input_field_name.clone(), + output_field_name, + merge_op: merge_op.clone(), + }); + } + } + if projection_info.has_top_level_projection { + let output_positions = projection_info + .output_field_names + .iter() + .enumerate() + .map(|(position, name)| (name.as_str(), position)) + .collect::>(); + merge_columns.sort_by_key(|column| { + output_positions + .get(column.output_field_name.as_str()) + .copied() + .unwrap_or(usize::MAX) + }); } unsupported_exprs.extend( find_uncovered_output_fields(&projection_info, &group_key_names, &merge_columns) @@ -592,6 +622,7 @@ pub fn analyze_incremental_aggregate_plan( pub async fn rewrite_incremental_aggregate_with_sink_merge( delta_plan: &LogicalPlan, analysis: &IncrementalAggregateAnalysis, + engine: &QueryEngineRef, sink_table: TableRef, sink_table_name: &TableName, sink_dirty_filter: Option, @@ -626,6 +657,10 @@ pub async fn rewrite_incremental_aggregate_with_sink_merge( let delta_alias = "__flow_delta"; let sink_alias = "__flow_sink"; + let state_merge = analysis + .merge_columns + .iter() + .any(|column| matches!(column.merge_op, IncrementalAggregateMergeOp::AvgDeltaMerge)); let mut selected_columns = analysis.group_key_names.clone(); selected_columns.extend( analysis @@ -633,8 +668,18 @@ pub async fn rewrite_incremental_aggregate_with_sink_merge( .iter() .map(|c| c.output_field_name.clone()), ); - let mut delta_selected_columns = selected_columns.clone(); + let mut selected_column_names = HashSet::new(); + selected_columns.retain(|name| selected_column_names.insert(name.clone())); + let mut delta_selected_columns = analysis.group_key_names.clone(); + delta_selected_columns.extend( + analysis + .merge_columns + .iter() + .map(|c| c.input_field_name.clone()), + ); delta_selected_columns.extend(analysis.literal_columns.iter().cloned()); + let mut delta_selected_column_names = HashSet::new(); + delta_selected_columns.retain(|name| delta_selected_column_names.insert(name.clone())); let delta_selected_exprs = delta_selected_columns .iter() @@ -722,7 +767,6 @@ pub async fn rewrite_incremental_aggregate_with_sink_merge( .map(|c| qualified_column(sink_alias, c)) .collect::>(), ); - let joined = LogicalPlanBuilder::from(delta_selected) .join_detailed( sink_selected, @@ -747,21 +791,28 @@ pub async fn rewrite_incremental_aggregate_with_sink_merge( .iter() .map(|c| (&c.output_field_name, c)) .collect::>(); - let mut projection_exprs = Vec::with_capacity(analysis.output_field_names.len()); + let mut group_exprs = Vec::new(); + let mut state_aggr_exprs = Vec::new(); for output_field_name in &analysis.output_field_names { if group_key_names.contains(output_field_name) || literal_columns.contains(output_field_name) { - projection_exprs.push( - qualified_col(delta_alias, output_field_name.clone()).alias(output_field_name), - ); + let expr = + qualified_col(delta_alias, output_field_name.clone()).alias(output_field_name); + projection_exprs.push(expr.clone()); + group_exprs.push(expr); } else if let Some(merge_col) = merge_columns.get(output_field_name) { - projection_exprs.push(build_left_join_merge_expr( - delta_alias, - sink_alias, - merge_col, - )?); + if matches!( + merge_col.merge_op, + IncrementalAggregateMergeOp::AvgDeltaMerge + ) { + state_aggr_exprs.push(build_state_delta_merge_expr(engine, merge_col)?); + } else { + let expr = build_left_join_merge_expr(delta_alias, sink_alias, merge_col)?; + projection_exprs.push(expr.clone()); + group_exprs.push(expr); + } } else { return InvalidQuerySnafu { reason: format!( @@ -772,15 +823,72 @@ pub async fn rewrite_incremental_aggregate_with_sink_merge( } } - LogicalPlanBuilder::from(joined) - .project(projection_exprs) - .with_context(|_| DatafusionSnafu { - context: "Failed to build projection merge plan for incremental sink merge".to_string(), - })? - .build() - .with_context(|_| DatafusionSnafu { - context: "Failed to finalize incremental aggregate sink merge plan".to_string(), + if state_merge { + let aggregated = LogicalPlanBuilder::from(joined) + .aggregate(group_exprs, state_aggr_exprs) + .with_context(|_| DatafusionSnafu { + context: "Failed to aggregate state delta merge plan".to_string(), + })? + .build() + .with_context(|_| DatafusionSnafu { + context: "Failed to build state delta merge plan".to_string(), + })?; + let output_exprs = analysis + .output_field_names + .iter() + .cloned() + .map(unqualified_col) + .collect::>(); + LogicalPlanBuilder::from(aggregated) + .project(output_exprs) + .with_context(|_| DatafusionSnafu { + context: "Failed to project state delta merge plan".to_string(), + })? + .build() + .with_context(|_| DatafusionSnafu { + context: "Failed to finalize incremental aggregate sink merge plan".to_string(), + }) + } else { + LogicalPlanBuilder::from(joined) + .project(projection_exprs) + .with_context(|_| DatafusionSnafu { + context: "Failed to build projection merge plan for incremental sink merge" + .to_string(), + })? + .build() + .with_context(|_| DatafusionSnafu { + context: "Failed to finalize incremental aggregate sink merge plan".to_string(), + }) + } +} + +fn build_state_delta_merge_expr( + engine: &QueryEngineRef, + merge_col: &IncrementalAggregateMergeColumn, +) -> Result { + let Some(udaf) = engine + .engine_state() + .aggr_function("__avg_state_delta_merge") + .or_else(|| { + engine + .engine_state() + .session_state() + .aggregate_functions() + .get("__avg_state_delta_merge") + .map(|udaf| udaf.as_ref().clone()) }) + else { + return InvalidQuerySnafu { + reason: "Aggregate function __avg_state_delta_merge is not registered".to_string(), + } + .fail(); + }; + Ok(udaf + .call(vec![ + qualified_col("__flow_delta", merge_col.input_field_name.clone()), + qualified_col("__flow_sink", merge_col.output_field_name.clone()), + ]) + .alias(merge_col.output_field_name.clone())) } fn build_left_join_merge_expr( @@ -788,7 +896,7 @@ fn build_left_join_merge_expr( sink_alias: &str, merge_col: &IncrementalAggregateMergeColumn, ) -> Result { - let left = qualified_col(delta_alias, merge_col.output_field_name.clone()); + let left = qualified_col(delta_alias, merge_col.input_field_name.clone()); let right = qualified_col(sink_alias, merge_col.output_field_name.clone()); let merged = match merge_col.merge_op { IncrementalAggregateMergeOp::Sum => when(is_null(left.clone()), right.clone()) @@ -839,6 +947,12 @@ fn build_left_join_merge_expr( .with_context(|_| DatafusionSnafu { context: "Failed to build BIT_XOR merge expression".to_string(), })?, + IncrementalAggregateMergeOp::AvgDeltaMerge => { + return InvalidQuerySnafu { + reason: "state aggregate must be built with its delta UDAF".to_string(), + } + .fail(); + } }; Ok(merged.alias(merge_col.output_field_name.clone())) } @@ -1149,53 +1263,67 @@ impl ColumnMatcherRewriter { input_schema: &DFSchema, ) -> DfResult> { let original_exprs = exprs.clone(); - for column in self.schema.column_schemas() { - if let Some(value) = self.ordinary_values.get(&column.name) { - if value.data_type() != column.data_type.as_arrow_type() { - return Err(DataFusionError::Plan(format!( - "Configured batching metadata column {} has incompatible type", - column.name - ))); - } - if !exprs - .iter() - .any(|expr| expr.qualified_name().1 == column.name) - { - exprs.push(datafusion_expr::lit(value.clone()).alias(column.name.clone())); + self.validate_ordinary_values(&original_exprs)?; + let original_names = original_exprs + .iter() + .map(|expr| expr.qualified_name().1) + .collect::>(); + let duplicated_output_names = duplicate_names(&original_names); + if !duplicated_output_names.is_empty() { + return Err(DataFusionError::Plan(format!( + "Flow output schema contains duplicate column(s) {:?}. {}", + duplicated_output_names, + format_flow_sink_schema_mismatch(&original_exprs, self.schema.as_ref()) + ))); + } + + if self.allow_partial { + // Partial matching is intentionally name-based. Ordinary values are injected before + // it so they follow the same direct partial path as the other supplied columns. + for (idx, column) in self.schema.column_schemas().iter().enumerate() { + if let Some(value) = self.ordinary_values.get(&column.name) { + exprs.insert( + idx.min(exprs.len()), + datafusion_expr::lit(value.clone()).alias(column.name.clone()), + ); } } - } - if self.allow_partial { return self.modify_project_exprs_with_partial(exprs); } - let all_names = self + // Ordinary values are persistence-owned columns, not flow outputs. Remove them from the + // effective sink sequence while deciding whether the existing auto-column rules apply. + // This keeps those columns from hiding an auto-created update_at column that precedes them. + let effective_sink_columns = self .schema .column_schemas() .iter() - .map(|c| c.name.clone()) - .collect::>(); - // add columns if have different column count + .enumerate() + .filter(|(_, column)| !self.ordinary_values.contains_key(&column.name)) + .collect::>(); let query_col_cnt = exprs.len(); - let table_col_cnt = self.schema.column_schemas().len(); - debug!("query_col_cnt={query_col_cnt}, table_col_cnt={table_col_cnt}"); + let effective_sink_col_cnt = effective_sink_columns.len(); + debug!("query_col_cnt={query_col_cnt}, effective_sink_col_cnt={effective_sink_col_cnt}"); let placeholder_ts_expr = datafusion::logical_expr::lit(ScalarValue::TimestampMillisecond(Some(0), None)) .alias(AUTO_CREATED_PLACEHOLDER_TS_COL); - if query_col_cnt == table_col_cnt { - // still need to add alias, see below - } else if query_col_cnt + 1 == table_col_cnt { - let last_col_schema = self.schema.column_schemas().last().unwrap(); + if query_col_cnt == effective_sink_col_cnt { + // still need to add aliases, see below + } else if query_col_cnt + 1 == effective_sink_col_cnt { + let (_, last_col_schema) = effective_sink_columns.last().unwrap(); - // if time index column is auto created add it if last_col_schema.name == AUTO_CREATED_PLACEHOLDER_TS_COL - && self.schema.timestamp_index() == Some(table_col_cnt - 1) + && self.schema.timestamp_index() + == Some( + self.schema + .column_index_by_name(&last_col_schema.name) + .unwrap(), + ) { exprs.push(placeholder_ts_expr); } else if last_col_schema.data_type.is_timestamp() { - // is the update at column exprs.push(datafusion::prelude::now().alias(&last_col_schema.name)); } else { return Err(DataFusionError::Plan(format_flow_sink_schema_mismatch( @@ -1203,10 +1331,11 @@ impl ColumnMatcherRewriter { self.schema.as_ref(), ))); } - } else if query_col_cnt + 2 == table_col_cnt { - let mut col_iter = self.schema.column_schemas().iter().rev(); - let last_col_schema = col_iter.next().unwrap(); - let second_last_col_schema = col_iter.next().unwrap(); + } else if query_col_cnt + 2 == effective_sink_col_cnt { + let (_, last_col_schema) = effective_sink_columns.last().unwrap(); + let (_, second_last_col_schema) = effective_sink_columns + .get(effective_sink_col_cnt - 2) + .unwrap(); if second_last_col_schema.data_type.is_timestamp() { exprs.push(datafusion::prelude::now().alias(&second_last_col_schema.name)); } else { @@ -1217,7 +1346,12 @@ impl ColumnMatcherRewriter { } if last_col_schema.name == AUTO_CREATED_PLACEHOLDER_TS_COL - && self.schema.timestamp_index() == Some(table_col_cnt - 1) + && self.schema.timestamp_index() + == Some( + self.schema + .column_index_by_name(&last_col_schema.name) + .unwrap(), + ) { exprs.push(placeholder_ts_expr); } else { @@ -1233,7 +1367,74 @@ impl ColumnMatcherRewriter { ))); } - self.match_extra_output_columns(exprs, input_schema, &original_exprs, &all_names) + let exprs = self.match_extra_output_columns( + exprs, + input_schema, + &original_exprs, + &effective_sink_columns, + )?; + + // Put persistence-owned values back at their physical sink positions only after matching + // flow expressions against the effective sequence. + let mut exprs = exprs; + for (idx, column) in self.schema.column_schemas().iter().enumerate() { + if let Some(value) = self.ordinary_values.get(&column.name) { + exprs.insert( + idx.min(exprs.len()), + datafusion_expr::lit(value.clone()).alias(column.name.clone()), + ); + } + } + self.order_by_sink_schema(exprs, &original_exprs) + } + + fn order_by_sink_schema( + &self, + exprs: Vec, + original_exprs: &[Expr], + ) -> DfResult> { + let mut by_name = exprs + .into_iter() + .map(|expr| (expr.qualified_name().1, expr)) + .collect::>(); + let mut ordered = Vec::with_capacity(self.schema.column_schemas().len()); + for column in self.schema.column_schemas() { + if let Some(expr) = by_name.remove(&column.name) { + ordered.push(expr); + } + } + if !by_name.is_empty() || ordered.len() != self.schema.column_schemas().len() { + return Err(DataFusionError::Plan(format_flow_sink_schema_mismatch( + original_exprs, + self.schema.as_ref(), + ))); + } + Ok(ordered) + } + + fn validate_ordinary_values(&self, output_exprs: &[Expr]) -> DfResult<()> { + let output_names = output_exprs + .iter() + .map(|expr| expr.qualified_name().1) + .collect::>(); + for (name, value) in &self.ordinary_values { + let Some(column) = self.schema.column_schema_by_name(name) else { + return Err(DataFusionError::Plan(format!( + "Configured batching metadata column {name} does not exist in sink schema" + ))); + }; + if output_names.contains(name) { + return Err(DataFusionError::Plan(format!( + "Configured batching metadata column {name} collides with a flow output" + ))); + } + if value.data_type() != column.data_type.as_arrow_type() { + return Err(DataFusionError::Plan(format!( + "Configured batching metadata column {name} has incompatible type" + ))); + } + } + Ok(()) } /// Match flow output columns whose names are not in the sink schema by the same position only. @@ -1251,24 +1452,28 @@ impl ColumnMatcherRewriter { mut exprs: Vec, input_schema: &DFSchema, original_exprs: &[Expr], - all_names: &BTreeSet, + effective_sink_columns: &[(usize, &ColumnSchema)], ) -> DfResult> { let mut output_names = exprs .iter() .map(|expr| expr.qualified_name().1) .collect::>(); + let sink_names = effective_sink_columns + .iter() + .map(|(_, column)| column.name.as_str()) + .collect::>(); let output_name_set = output_names.iter().cloned().collect::>(); let extra_expr_indices = output_names .iter() .enumerate() - .filter_map(|(idx, name)| (!all_names.contains(name)).then_some(idx)) + .filter_map(|(idx, name)| (!sink_names.contains(name.as_str())).then_some(idx)) .collect::>(); - let missing_sink_indices = self - .schema - .column_schemas() + let missing_sink_indices = effective_sink_columns .iter() .enumerate() - .filter_map(|(idx, column)| (!output_name_set.contains(&column.name)).then_some(idx)) + .filter_map(|(idx, (_, column))| { + (!output_name_set.contains(&column.name)).then_some(idx) + }) .collect::>(); if extra_expr_indices.is_empty() && missing_sink_indices.is_empty() { @@ -1291,7 +1496,7 @@ impl ColumnMatcherRewriter { ))); } - let target_col_schema = &self.schema.column_schemas()[expr_idx]; + let (_, target_col_schema) = effective_sink_columns[expr_idx]; let expr_type = ConcreteDataType::from_arrow_type(&exprs[expr_idx].get_type(input_schema)?); if is_obviously_incompatible_positional_match(&expr_type, &target_col_schema.data_type) diff --git a/src/flow/src/batching_mode/utils/test.rs b/src/flow/src/batching_mode/utils/test.rs index 1356f87e8f..929b0f9da0 100644 --- a/src/flow/src/batching_mode/utils/test.rs +++ b/src/flow/src/batching_mode/utils/test.rs @@ -12,6 +12,7 @@ // See the License for the specific language governing permissions and // limitations under the License. +use std::collections::BTreeMap; use std::sync::Arc; use catalog::RegisterTableRequest; @@ -273,6 +274,81 @@ async fn test_sql_plan_convert() { ); } +#[test] +fn test_df_plan_to_sql_quotes_colon_table_name() { + // Prometheus-style table names contain ':' (e.g. + // `kube_pod_cpu_cores:sum`). The unparser dialect must quote them, + // otherwise the re-parsed SQL is invalid (`keyword: :`). + let table = single_row_u32_table("kube_pod_cpu_cores:sum", vec!["value"]); + let provider = Arc::new(DfTableProviderAdapter::new(table)); + let table_source = Arc::new(DefaultTableSource::new(provider)); + let table_ref = TableReference::full("catalog", "schema", "kube_pod_cpu_cores:sum"); + let plan = LogicalPlanBuilder::scan(table_ref, table_source, None) + .unwrap() + .build() + .unwrap(); + + let sql = df_plan_to_sql(&plan).unwrap(); + assert!( + sql.contains("\"kube_pod_cpu_cores:sum\""), + "expected quoted table name in {sql}" + ); + // The only occurrence of `cores:sum` must be inside the quoted identifier. + assert_eq!( + sql.matches("cores:sum").count(), + 1, + "colon should only appear inside quotes in {sql}" + ); +} + +#[test] +fn test_df_plan_to_sql_does_not_quote_plain_lowercase() { + let table = single_row_u32_table("plain_table", vec!["value"]); + let provider = Arc::new(DfTableProviderAdapter::new(table)); + let table_source = Arc::new(DefaultTableSource::new(provider)); + let table_ref = TableReference::full("catalog", "schema", "plain_table"); + let plan = LogicalPlanBuilder::scan(table_ref, table_source, None) + .unwrap() + .project(vec![datafusion_expr::col("value")]) + .unwrap() + .build() + .unwrap(); + + let sql = df_plan_to_sql(&plan).unwrap(); + assert!( + sql.contains("plain_table") && !sql.contains("\"plain_table\""), + "plain lowercase table should stay unquoted in {sql}" + ); + // `value` is not a reserved word, so the column stays unquoted. + assert!( + sql.contains("plain_table.value"), + "column unquoted in {sql}" + ); + assert!(!sql.contains('`'), "no backtick quoting in {sql}"); +} + +#[test] +fn test_df_plan_to_sql_quotes_digit_leading_table_name() { + // A table literally named `123metrics` starts with a digit and must be + // quoted, otherwise the re-parsed SQL is invalid. + let table = single_row_u32_table("123metrics", vec!["value"]); + let provider = Arc::new(DfTableProviderAdapter::new(table)); + let table_source = Arc::new(DefaultTableSource::new(provider)); + let table_ref = TableReference::full("catalog", "schema", "123metrics"); + let plan = LogicalPlanBuilder::scan(table_ref, table_source, None) + .unwrap() + .project(vec![datafusion_expr::col("value")]) + .unwrap() + .build() + .unwrap(); + + let sql = df_plan_to_sql(&plan).unwrap(); + assert!( + sql.contains("\"123metrics\""), + "expected digit-leading table name quoted in {sql}" + ); +} + #[tokio::test] async fn test_add_filter() { let testcases = vec![ @@ -668,8 +744,8 @@ async fn test_gen_plan_with_matching_schema_accepts_out_of_order_matching_names( output_names, vec![ "number".to_string(), - "ts".to_string(), - "time_window".to_string() + "time_window".to_string(), + "ts".to_string() ] ); assert!(duplicate_names(&output_names).is_empty()); @@ -844,6 +920,120 @@ async fn test_validate_sink_table_schema_rejects_existing_sink_missing_flow_colu assert!(err.contains("extra"), "{err}"); } +#[tokio::test] +async fn test_gen_plan_with_matching_schema_injects_attempt_columns_in_sink_order() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ColumnSchema::new("marker", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new("payload", ConcreteDataType::string_datatype(), true), + ColumnSchema::new("epoch", ConcreteDataType::uint64_datatype(), true), + ColumnSchema::new( + AUTO_CREATED_UPDATE_AT_TS_COL, + ConcreteDataType::timestamp_millisecond_datatype(), + true, + ), + ])); + let values = BTreeMap::from([ + ("marker".to_string(), ScalarValue::UInt32(Some(7))), + ( + "payload".to_string(), + ScalarValue::Utf8(Some("state".to_string())), + ), + ("epoch".to_string(), ScalarValue::UInt64(Some(9))), + ]); + let plan = gen_plan_with_matching_schema_and_values( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[], + false, + Some(&values), + ) + .await + .unwrap(); + let output_names = plan + .schema() + .fields() + .iter() + .map(|f| f.name().clone()) + .collect::>(); + assert_eq!( + output_names, + vec!["number", "ts", "marker", "payload", "epoch", "update_at"] + ); + let sql = df_plan_to_sql(&plan).unwrap(); + assert!(sql.contains("7 AS marker"), "{sql}"); + assert!(sql.contains("'state' AS payload"), "{sql}"); + assert!(sql.contains("9 AS epoch"), "{sql}"); + assert!(sql.contains("now() AS update_at"), "{sql}"); +} + +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_arbitrary_missing_attempt_column() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ColumnSchema::new("missing", ConcreteDataType::uint32_datatype(), true), + ])); + let err = gen_plan_with_matching_schema_and_values( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[], + false, + Some(&BTreeMap::new()), + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("missing sink columns"), "{err}"); +} + +#[tokio::test] +async fn test_gen_plan_with_matching_schema_no_attempt_still_rejects_missing_column() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ColumnSchema::new("state", ConcreteDataType::uint32_datatype(), true), + ])); + assert!( + gen_plan_with_matching_schema( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[], + false, + ) + .await + .is_err() + ); +} + #[tokio::test] async fn test_gen_plan_with_matching_schema_allow_partial_fills_nullable_columns() { let query_engine = create_test_query_engine(); @@ -1106,13 +1296,16 @@ async fn test_rewrite_incremental_aggregate_allows_alias_wrapped_scan() { let query_engine = create_test_query_engine(); let ctx = QueryContext::arc(); let sql = "SELECT max(n.number) AS number, n.ts FROM numbers_with_ts AS n GROUP BY n.ts"; - let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); + let plan = sql_to_df_plan(ctx, query_engine.clone(), sql, false) + .await + .unwrap(); let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap(); assert!(analysis.unsupported_exprs.is_empty()); let rewritten = rewrite_incremental_aggregate_with_sink_merge( &plan, &analysis, + &query_engine, single_row_u32_table("alias_wrapped_sink", vec!["ts", "number"]), &[ "greptime".to_string(), @@ -1364,6 +1557,7 @@ async fn test_analyze_incremental_aggregate_plan_allows_literal_outputs() { let rewritten = rewrite_incremental_aggregate_with_sink_merge( &plan, &analysis, + &query_engine, sink_table.clone(), &sink_table_name, None, @@ -1437,7 +1631,9 @@ async fn test_rewrite_incremental_aggregate_preserves_non_identifier_aliases() { let query_engine = create_test_query_engine(); let ctx = QueryContext::arc(); let sql = "SELECT max(number) AS \"max value\", number, 42 AS \"literal value\" FROM numbers_with_ts GROUP BY number"; - let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); + let plan = sql_to_df_plan(ctx, query_engine.clone(), sql, false) + .await + .unwrap(); let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap(); assert!(analysis.unsupported_exprs.is_empty()); assert_eq!( @@ -1449,6 +1645,7 @@ async fn test_rewrite_incremental_aggregate_preserves_non_identifier_aliases() { let rewritten = rewrite_incremental_aggregate_with_sink_merge( &plan, &analysis, + &query_engine, sink_table, &[ "greptime".to_string(), @@ -1522,24 +1719,34 @@ async fn test_datafusion_rejects_duplicate_output_names() { } #[tokio::test] -async fn test_analyze_incremental_aggregate_plan_rejects_same_aggregate_multiple_aliases() { - let query_engine = create_test_query_engine(); - let ctx = QueryContext::arc(); - let sql = "SELECT sum(number) AS a, sum(number) AS b, ts FROM numbers_with_ts GROUP BY ts"; - let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); +async fn test_analyze_incremental_aggregate_plan_supports_same_aggregate_multiple_aliases() { + let analysis = analyze_test_sql( + "SELECT sum(number) AS a, sum(number) AS b, ts FROM numbers_with_ts GROUP BY ts", + ) + .await; - let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap(); + assert!(analysis.unsupported_exprs.is_empty()); + assert_eq!(analysis.merge_columns.len(), 2); + assert_eq!(analysis.merge_columns[0].input_field_name, "a"); + assert_eq!(analysis.merge_columns[1].input_field_name, "a"); assert!( analysis - .unsupported_exprs + .merge_columns .iter() - .any(|expr| expr.contains("same aggregate output") - && expr.contains("a") - && expr.contains("b")), - "same aggregate with multiple aliases should be unsupported until explicit reproduction is implemented: {:?}", - analysis.unsupported_exprs + .all(|column| { column.merge_op == IncrementalAggregateMergeOp::Sum }) + ); + assert!( + analysis + .merge_columns + .iter() + .any(|column| column.output_field_name == "a") + ); + assert!( + analysis + .merge_columns + .iter() + .any(|column| column.output_field_name == "b") ); - assert!(analysis.merge_columns.is_empty()); } #[test] @@ -1583,6 +1790,94 @@ async fn test_analyze_incremental_aggregate_plan_rejects_avg() { assert!(!analysis.unsupported_exprs.is_empty()); } +#[tokio::test] +async fn test_analyze_incremental_aggregate_plan_supports_avg_state() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sql = "SELECT avg_state(number) AS avg_num, ts FROM numbers_with_ts GROUP BY ts"; + let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); + + let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap(); + assert!( + analysis.unsupported_exprs.is_empty(), + "avg_state should be supported: {:?}", + analysis.unsupported_exprs + ); + assert_eq!(analysis.merge_columns.len(), 1); + assert_eq!(analysis.merge_columns[0].output_field_name, "avg_num"); + assert_eq!( + analysis.merge_columns[0].merge_op, + IncrementalAggregateMergeOp::AvgDeltaMerge + ); +} + +#[tokio::test] +async fn test_analyze_incremental_aggregate_plan_supports_avg_merge() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sql = "SELECT avg_merge(avg_state(number)) AS avg_num, ts FROM numbers_with_ts GROUP BY ts"; + let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); + + let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap(); + assert!( + analysis.unsupported_exprs.is_empty(), + "avg_merge should be supported: {:?}", + analysis.unsupported_exprs + ); + assert_eq!(analysis.merge_columns.len(), 1); + assert_eq!( + analysis.merge_columns[0].merge_op, + IncrementalAggregateMergeOp::AvgDeltaMerge + ); +} + +#[tokio::test] +async fn test_analyze_incremental_aggregate_plan_supports_duplicate_avg_projections() { + let analysis = analyze_test_sql( + "SELECT avg_state(number) AS avg_num, avg_state(number + 1) AS avg_num_plus, ts FROM numbers_with_ts GROUP BY ts", + ) + .await; + + assert!(analysis.unsupported_exprs.is_empty()); + assert_eq!(analysis.merge_columns.len(), 2); + assert!( + analysis + .merge_columns + .iter() + .all(|column| { column.merge_op == IncrementalAggregateMergeOp::AvgDeltaMerge }) + ); + assert!( + analysis + .merge_columns + .iter() + .any(|column| column.output_field_name == "avg_num") + ); + assert!( + analysis + .merge_columns + .iter() + .any(|column| column.output_field_name == "avg_num_plus") + ); +} + +#[tokio::test] +async fn test_analyze_incremental_aggregate_plan_supports_avg_with_native_aggregate() { + let analysis = analyze_test_sql( + "SELECT avg_state(number) AS avg_num, sum(number) AS total, ts FROM numbers_with_ts GROUP BY ts", + ) + .await; + + assert!(analysis.unsupported_exprs.is_empty()); + assert_eq!(analysis.merge_columns.len(), 2); + assert!(analysis.merge_columns.iter().any(|column| { + column.output_field_name == "avg_num" + && column.merge_op == IncrementalAggregateMergeOp::AvgDeltaMerge + })); + assert!(analysis.merge_columns.iter().any(|column| { + column.output_field_name == "total" && column.merge_op == IncrementalAggregateMergeOp::Sum + })); +} + #[tokio::test] async fn test_analyze_incremental_aggregate_plan_rejects_distinct() { let query_engine = create_test_query_engine(); @@ -1640,6 +1935,7 @@ async fn test_rewrite_incremental_aggregate_with_left_join() { let rewritten = rewrite_incremental_aggregate_with_sink_merge( &plan, &analysis, + &query_engine, sink_table.clone(), &sink_table_name, None, @@ -1701,6 +1997,7 @@ async fn test_rewrite_incremental_aggregate_filters_sink_dirty_time_window() { let rewritten = rewrite_incremental_aggregate_with_sink_merge( &plan, &analysis, + &query_engine, sink_table.clone(), &sink_table_name, Some(sink_filter.clone()), @@ -1751,7 +2048,9 @@ async fn test_rewrite_incremental_aggregate_rejects_empty_group_keys() { let query_engine = create_test_query_engine(); let ctx = QueryContext::arc(); let sql = "SELECT max(number) AS number FROM numbers_with_ts"; - let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); + let plan = sql_to_df_plan(ctx, query_engine.clone(), sql, false) + .await + .unwrap(); let analysis = IncrementalAggregateAnalysis { group_key_names: vec![], merge_columns: vec![IncrementalAggregateMergeColumn::new( @@ -1772,6 +2071,7 @@ async fn test_rewrite_incremental_aggregate_rejects_empty_group_keys() { let err = rewrite_incremental_aggregate_with_sink_merge( &plan, &analysis, + &query_engine, sink_table, &sink_table_name, None, @@ -1790,7 +2090,9 @@ async fn test_rewrite_incremental_aggregate_preserves_raw_aggregate_field_name() let query_engine = create_test_query_engine(); let ctx = QueryContext::arc(); let sql = "SELECT max(number), number FROM numbers_with_ts GROUP BY number"; - let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); + let plan = sql_to_df_plan(ctx, query_engine.clone(), sql, false) + .await + .unwrap(); let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap(); assert!(analysis.unsupported_exprs.is_empty()); @@ -1804,6 +2106,7 @@ async fn test_rewrite_incremental_aggregate_preserves_raw_aggregate_field_name() let rewritten = rewrite_incremental_aggregate_with_sink_merge( &plan, &analysis, + &query_engine, sink_table.clone(), &sink_table_name, None, @@ -1841,9 +2144,7 @@ async fn test_null_cast() { let query_engine = create_test_query_engine(); let ctx = QueryContext::arc(); let sql = "SELECT NULL::DOUBLE FROM numbers_with_ts"; - let plan = sql_to_df_plan(ctx, query_engine.clone(), sql, false) - .await - .unwrap(); + let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap(); let _sub_plan = DFLogicalSubstraitConvertor {} .encode(&plan, DefaultSerializer) @@ -1986,77 +2287,352 @@ async fn test_gen_plan_with_matching_schema_last_non_null_rejects_extra_flow_col ); } -#[test] -fn test_df_plan_to_sql_quotes_colon_table_name() { - // Prometheus-style table names contain ':' (e.g. - // `kube_pod_cpu_cores:sum`). The unparser dialect must quote them, - // otherwise the re-parsed SQL is invalid (`keyword: :`). - let table = single_row_u32_table("kube_pod_cpu_cores:sum", vec!["value"]); - let provider = Arc::new(DfTableProviderAdapter::new(table)); - let table_source = Arc::new(DefaultTableSource::new(provider)); - let table_ref = TableReference::full("catalog", "schema", "kube_pod_cpu_cores:sum"); - let plan = LogicalPlanBuilder::scan(table_ref, table_source, None) - .unwrap() - .build() - .unwrap(); +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_unknown_attempt_column() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ])); + let values = BTreeMap::from([(String::from("unknown_attempt"), ScalarValue::Int64(Some(1)))]); + let err = gen_plan_with_matching_schema_and_values( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[0], + false, + Some(&values), + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("unknown_attempt"), "{err}"); + assert!(err.contains("does not exist in sink schema"), "{err}"); +} - let sql = df_plan_to_sql(&plan).unwrap(); - assert!( - sql.contains("\"kube_pod_cpu_cores:sum\""), - "expected quoted table name in {sql}" - ); - // The only occurrence of `cores:sum` must be inside the quoted identifier. +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_wrong_attempt_column_type() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new("attempt", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ])); + let values = BTreeMap::from([( + String::from("attempt"), + ScalarValue::Utf8(Some("one".into())), + )]); + let err = gen_plan_with_matching_schema_and_values( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[0], + false, + Some(&values), + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("attempt"), "{err}"); + assert!(err.contains("incompatible type"), "{err}"); +} + +#[tokio::test] +async fn test_gen_plan_with_matching_schema_matches_positional_alias_and_injects_attempt() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("renamed_number", ConcreteDataType::int64_datatype(), true), + ColumnSchema::new("attempt", ConcreteDataType::string_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ])); + let values = BTreeMap::from([( + String::from("attempt"), + ScalarValue::Utf8(Some("one".into())), + )]); + let plan = gen_plan_with_matching_schema_and_values( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[0], + false, + Some(&values), + ) + .await + .unwrap(); + let output_names = plan + .schema() + .fields() + .iter() + .map(|field| field.name().clone()) + .collect::>(); + assert_eq!(output_names, vec!["renamed_number", "attempt", "ts"]); +} + +#[tokio::test] +async fn test_gen_plan_with_matching_schema_injects_ordinary_columns_after_auto_update_at() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let mut sink_columns = (0..16) + .map(|idx| { + ColumnSchema::new( + format!("state_{idx}"), + ConcreteDataType::int32_datatype(), + true, + ) + }) + .collect::>(); + sink_columns.push(ColumnSchema::new( + "update_at", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + )); + sink_columns.extend([ + ColumnSchema::new( + "__ee_checkpoint_epoch", + ConcreteDataType::uint32_datatype(), + false, + ), + ColumnSchema::new( + "__ee_checkpoint_sequence", + ConcreteDataType::uint32_datatype(), + false, + ), + ColumnSchema::new( + "__ee_checkpoint_region", + ConcreteDataType::uint32_datatype(), + false, + ), + ]); + let sink_schema = Arc::new(Schema::new(sink_columns)); + let ordinary_values = BTreeMap::from([ + ( + "__ee_checkpoint_epoch".to_string(), + ScalarValue::UInt32(Some(1)), + ), + ( + "__ee_checkpoint_sequence".to_string(), + ScalarValue::UInt32(Some(2)), + ), + ( + "__ee_checkpoint_region".to_string(), + ScalarValue::UInt32(Some(3)), + ), + ]); + + let flow_exprs = (0..16) + .map(|idx| format!("number AS state_{idx}")) + .collect::>(); + let sql = format!("SELECT {} FROM numbers_with_ts", flow_exprs.join(", ")); + let plan = gen_plan_with_matching_schema_and_values( + &sql, + ctx, + query_engine, + sink_schema, + &[], + false, + Some(&ordinary_values), + ) + .await + .unwrap(); + + let output_names = plan + .schema() + .fields() + .iter() + .map(|field| field.name().clone()) + .collect::>(); assert_eq!( - sql.matches("cores:sum").count(), - 1, - "colon should only appear inside quotes in {sql}" + output_names, + vec![ + "state_0", + "state_1", + "state_2", + "state_3", + "state_4", + "state_5", + "state_6", + "state_7", + "state_8", + "state_9", + "state_10", + "state_11", + "state_12", + "state_13", + "state_14", + "state_15", + "update_at", + "__ee_checkpoint_epoch", + "__ee_checkpoint_sequence", + "__ee_checkpoint_region", + ] ); } -#[test] -fn test_df_plan_to_sql_does_not_quote_plain_lowercase() { - let table = single_row_u32_table("plain_table", vec!["value"]); - let provider = Arc::new(DfTableProviderAdapter::new(table)); - let table_source = Arc::new(DefaultTableSource::new(provider)); - let table_ref = TableReference::full("catalog", "schema", "plain_table"); - let plan = LogicalPlanBuilder::scan(table_ref, table_source, None) - .unwrap() - .project(vec![datafusion_expr::col("value")]) - .unwrap() - .build() - .unwrap(); - - let sql = df_plan_to_sql(&plan).unwrap(); - assert!( - sql.contains("plain_table") && !sql.contains("\"plain_table\""), - "plain lowercase table should stay unquoted in {sql}" - ); - // `value` is not a reserved word, so the column stays unquoted. - assert!( - sql.contains("plain_table.value"), - "column unquoted in {sql}" - ); - assert!(!sql.contains('`'), "no backtick quoting in {sql}"); +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_no_attempt_strict_mismatch() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ColumnSchema::new("attempt", ConcreteDataType::string_datatype(), true), + ])); + let err = gen_plan_with_matching_schema( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[0], + false, + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("does not match sink table schema"), "{err}"); + assert!(err.contains("attempt"), "{err}"); } -#[test] -fn test_df_plan_to_sql_quotes_digit_leading_table_name() { - // A table literally named `123metrics` starts with a digit and must be - // quoted, otherwise the re-parsed SQL is invalid. - let table = single_row_u32_table("123metrics", vec!["value"]); - let provider = Arc::new(DfTableProviderAdapter::new(table)); - let table_source = Arc::new(DefaultTableSource::new(provider)); - let table_ref = TableReference::full("catalog", "schema", "123metrics"); - let plan = LogicalPlanBuilder::scan(table_ref, table_source, None) - .unwrap() - .project(vec![datafusion_expr::col("value")]) - .unwrap() - .build() - .unwrap(); +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_unknown_attempt_column_in_partial_mode() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ColumnSchema::new("optional", ConcreteDataType::uint32_datatype(), true), + ])); + let values = BTreeMap::from([(String::from("unknown_attempt"), ScalarValue::Int64(Some(1)))]); + let err = gen_plan_with_matching_schema_and_values( + "SELECT number, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[0], + true, + Some(&values), + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("unknown_attempt"), "{err}"); + assert!(err.contains("does not exist in sink schema"), "{err}"); +} - let sql = df_plan_to_sql(&plan).unwrap(); - assert!( - sql.contains("\"123metrics\""), - "expected digit-leading table name quoted in {sql}" - ); +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_attempt_output_collision() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new("attempt", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ])); + let values = BTreeMap::from([(String::from("attempt"), ScalarValue::UInt32(Some(1)))]); + let err = gen_plan_with_matching_schema_and_values( + "SELECT number, number AS attempt, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[0], + false, + Some(&values), + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("collides with a flow output"), "{err}"); +} + +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_duplicate_original_outputs() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ])); + let err = gen_plan_with_matching_schema( + "SELECT * FROM numbers_with_ts AS lhs JOIN numbers_with_ts AS rhs ON lhs.ts = rhs.ts", + ctx, + query_engine, + sink_schema, + &[0], + false, + ) + .await + .unwrap_err(); + let diagnostic = format!("{err:?}"); + assert!(diagnostic.contains("duplicate column"), "{diagnostic}"); + assert!(diagnostic.contains("number"), "{diagnostic}"); +} + +#[tokio::test] +async fn test_gen_plan_with_matching_schema_rejects_attempt_output_collision_in_partial_mode() { + let query_engine = create_test_query_engine(); + let ctx = QueryContext::arc(); + let sink_schema = Arc::new(Schema::new(vec![ + ColumnSchema::new("number", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new("attempt", ConcreteDataType::uint32_datatype(), true), + ColumnSchema::new( + "ts", + ConcreteDataType::timestamp_millisecond_datatype(), + false, + ) + .with_time_index(true), + ])); + let values = BTreeMap::from([(String::from("attempt"), ScalarValue::UInt32(Some(1)))]); + let err = gen_plan_with_matching_schema_and_values( + "SELECT number, number AS attempt, ts FROM numbers_with_ts", + ctx, + query_engine, + sink_schema, + &[0], + true, + Some(&values), + ) + .await + .unwrap_err() + .to_string(); + assert!(err.contains("collides with a flow output"), "{err}"); } diff --git a/src/flow/src/server.rs b/src/flow/src/server.rs index e3742bb96a..fb15a687aa 100644 --- a/src/flow/src/server.rs +++ b/src/flow/src/server.rs @@ -327,7 +327,6 @@ pub struct FlownodeBuilder { /// receive a oneshot sender to send state size report state_report_handler: Option, frontend_client: Arc, - batching_persistence_factory: Option, } impl FlownodeBuilder { @@ -349,16 +348,9 @@ impl FlownodeBuilder { heartbeat_task: None, state_report_handler: None, frontend_client, - batching_persistence_factory: None, } } - /// Inject the optional batching persistence collaborator. - pub fn with_batching_persistence_factory(mut self, factory: FactoryPlugin) -> Self { - self.batching_persistence_factory = Some(factory); - self - } - pub fn with_heartbeat_task(self, heartbeat_task: HeartbeatTask) -> Self { let (sender, receiver) = SizeReportSender::new(); Self { @@ -420,7 +412,7 @@ impl FlownodeBuilder { self.table_meta.clone(), self.catalog_manager.clone(), self.opts.flow.batching_mode.clone(), - self.batching_persistence_factory.clone(), + self.plugins.get::(), )); let dual = Arc::new(FlowDualEngine::new( manager.clone(),