Skip to main content

query/query_engine/
default_serializer.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::sync::Arc;
16
17use common_error::ext::BoxedError;
18use common_function::function::FunctionContext;
19use common_function::function_registry::FUNCTION_REGISTRY;
20use common_query::error::RegisterUdfSnafu;
21use common_query::logical_plan::SubstraitPlanDecoder;
22use datafusion::catalog::CatalogProviderList;
23use datafusion::common::DataFusionError;
24use datafusion::error::Result;
25use datafusion::execution::context::SessionState;
26use datafusion::execution::registry::SerializerRegistry;
27use datafusion::execution::{FunctionRegistry, SessionStateBuilder};
28use datafusion::logical_expr::LogicalPlan;
29use datafusion_expr::UserDefinedLogicalNode;
30use greptime_proto::substrait_extension::MergeScan as PbMergeScan;
31use promql::functions::{
32    AbsentOverTime, AvgOverTime, Changes, CountOverTime, Delta, Deriv, DoubleExponentialSmoothing,
33    IDelta, Increase, LastOverTime, MaxOverTime, MinOverTime, MixedRange,
34    NativeHistogramAbsentOverTime, NativeHistogramAdd, NativeHistogramAggAvg,
35    NativeHistogramAggSum, NativeHistogramAvg, NativeHistogramAvgOverTime, NativeHistogramChanges,
36    NativeHistogramCount, NativeHistogramCountOverTime, NativeHistogramDelta,
37    NativeHistogramDivScalar, NativeHistogramDrop, NativeHistogramEq, NativeHistogramFraction,
38    NativeHistogramIDelta, NativeHistogramIRate, NativeHistogramIncrease,
39    NativeHistogramLastOverTime, NativeHistogramMulScalar, NativeHistogramNeg,
40    NativeHistogramNotEq, NativeHistogramPresentOverTime, NativeHistogramQuantile,
41    NativeHistogramRate, NativeHistogramResets, NativeHistogramScalarMul, NativeHistogramStddev,
42    NativeHistogramStdvar, NativeHistogramSub, NativeHistogramSum, NativeHistogramSumOverTime,
43    NativeHistogramToString, PredictLinear, PresentOverTime, PromqlFloatToString, QuantileOverTime,
44    Rate, Resets, Round, StddevOverTime, StdvarOverTime, SumOverTime, quantile_udaf,
45};
46use prost::Message;
47use session::context::QueryContextRef;
48use snafu::ResultExt;
49use substrait::extension_serializer::ExtensionSerializer;
50use substrait::{DFLogicalSubstraitConvertor, SubstraitPlan};
51
52use crate::dist_plan::MergeScanLogicalPlan;
53
54/// Extended [`substrait::extension_serializer::ExtensionSerializer`] but supports [`MergeScanLogicalPlan`] serialization.
55#[derive(Debug)]
56pub struct DefaultSerializer;
57
58impl SerializerRegistry for DefaultSerializer {
59    fn serialize_logical_plan(&self, node: &dyn UserDefinedLogicalNode) -> Result<Vec<u8>> {
60        if node.name() == MergeScanLogicalPlan::name() {
61            let merge_scan = node
62                .as_any()
63                .downcast_ref::<MergeScanLogicalPlan>()
64                .expect("Failed to downcast to MergeScanLogicalPlan");
65
66            let input = merge_scan.input();
67            let is_placeholder = merge_scan.is_placeholder();
68            let input = DFLogicalSubstraitConvertor
69                .encode(input, DefaultSerializer)
70                .map_err(|e| DataFusionError::External(Box::new(e)))?
71                .to_vec();
72
73            Ok(PbMergeScan {
74                is_placeholder,
75                input,
76            }
77            .encode_to_vec())
78        } else {
79            ExtensionSerializer.serialize_logical_plan(node)
80        }
81    }
82
83    fn deserialize_logical_plan(
84        &self,
85        name: &str,
86        bytes: &[u8],
87    ) -> Result<Arc<dyn UserDefinedLogicalNode>> {
88        if name == MergeScanLogicalPlan::name() {
89            // TODO(dennis): missing `session_state` to decode the logical plan in `MergeScanLogicalPlan`,
90            // so we only save the unoptimized logical plan for view currently.
91            Err(DataFusionError::Substrait(format!(
92                "Unsupported plan node: {name}"
93            )))
94        } else {
95            ExtensionSerializer.deserialize_logical_plan(name, bytes)
96        }
97    }
98}
99
100/// The datafusion `[LogicalPlan]` decoder.
101pub struct DefaultPlanDecoder {
102    session_state: SessionState,
103    query_ctx: QueryContextRef,
104}
105
106impl DefaultPlanDecoder {
107    pub fn new(
108        session_state: SessionState,
109        query_ctx: &QueryContextRef,
110    ) -> crate::error::Result<Self> {
111        Ok(Self {
112            session_state,
113            query_ctx: query_ctx.clone(),
114        })
115    }
116}
117
118#[async_trait::async_trait]
119impl SubstraitPlanDecoder for DefaultPlanDecoder {
120    async fn decode(
121        &self,
122        message: bytes::Bytes,
123        catalog_list: Arc<dyn CatalogProviderList>,
124        optimize: bool,
125    ) -> common_query::error::Result<LogicalPlan> {
126        // The session_state already has the `DefaultSerialzier` as `SerializerRegistry`.
127        let mut session_state = SessionStateBuilder::new_from_existing(self.session_state.clone())
128            .with_catalog_list(catalog_list)
129            .build();
130        // Substrait decoder will look up the UDFs in SessionState, so we need to register them
131        // Note: the query context must be passed to set the timezone
132        // We MUST register the UDFs after we build the session state, otherwise the UDFs will be lost
133        // if they have the same name as the default UDFs or their alias.
134        // e.g. The default UDF `to_char()` has an alias `date_format()`, if we register a UDF with the name `date_format()`
135        // before we build the session state, the UDF will be lost.
136        for func in FUNCTION_REGISTRY.scalar_functions() {
137            let udf = func.provide(FunctionContext {
138                query_ctx: self.query_ctx.clone(),
139                state: Default::default(),
140            });
141            session_state
142                .register_udf(Arc::new(udf))
143                .context(RegisterUdfSnafu { name: func.name() })?;
144        }
145
146        for func in FUNCTION_REGISTRY.aggregate_functions() {
147            let name = func.name().to_string();
148            session_state
149                .register_udaf(Arc::new(func))
150                .context(RegisterUdfSnafu { name })?;
151        }
152
153        let _ = session_state.register_udaf(quantile_udaf());
154
155        let _ = session_state.register_udf(Arc::new(IDelta::<false>::scalar_udf()));
156        let _ = session_state.register_udf(Arc::new(IDelta::<true>::scalar_udf()));
157        let _ = session_state.register_udf(Arc::new(Rate::scalar_udf()));
158        let _ = session_state.register_udf(Arc::new(Increase::scalar_udf()));
159        let _ = session_state.register_udf(Arc::new(Delta::scalar_udf()));
160        let _ = session_state.register_udf(Arc::new(Resets::scalar_udf()));
161        let _ = session_state.register_udf(Arc::new(Changes::scalar_udf()));
162        let _ = session_state.register_udf(Arc::new(Deriv::scalar_udf()));
163        let _ = session_state.register_udf(Arc::new(Round::scalar_udf()));
164        let _ = session_state.register_udf(Arc::new(AvgOverTime::scalar_udf()));
165        let _ = session_state.register_udf(Arc::new(MinOverTime::scalar_udf()));
166        let _ = session_state.register_udf(Arc::new(MaxOverTime::scalar_udf()));
167        let _ = session_state.register_udf(Arc::new(SumOverTime::scalar_udf()));
168        let _ = session_state.register_udf(Arc::new(CountOverTime::scalar_udf()));
169        let _ = session_state.register_udf(Arc::new(LastOverTime::scalar_udf()));
170        let _ = session_state.register_udf(Arc::new(AbsentOverTime::scalar_udf()));
171        let _ = session_state.register_udf(Arc::new(PresentOverTime::scalar_udf()));
172        let _ = session_state.register_udf(Arc::new(StddevOverTime::scalar_udf()));
173        let _ = session_state.register_udf(Arc::new(StdvarOverTime::scalar_udf()));
174        let _ = session_state.register_udf(Arc::new(QuantileOverTime::scalar_udf()));
175        let _ = session_state.register_udf(Arc::new(PredictLinear::scalar_udf()));
176        let double_exponential_smoothing_udf =
177            DoubleExponentialSmoothing::scalar_udf().with_aliases(["prom_holt_winters"]);
178        let _ = session_state.register_udf(Arc::new(double_exponential_smoothing_udf));
179
180        for udf in [
181            NativeHistogramAbsentOverTime::scalar_udf(),
182            NativeHistogramAdd::scalar_udf(),
183            NativeHistogramAvg::scalar_udf(),
184            NativeHistogramAvgOverTime::scalar_udf(),
185            NativeHistogramChanges::scalar_udf(),
186            NativeHistogramCount::scalar_udf(),
187            NativeHistogramCountOverTime::scalar_udf(),
188            NativeHistogramDelta::scalar_udf(),
189            NativeHistogramDivScalar::scalar_udf(),
190            NativeHistogramDrop::bool_false_udf(String::new(), None),
191            NativeHistogramDrop::bool_true_udf(String::new(), None),
192            NativeHistogramDrop::float_null_udf(String::new(), None),
193            NativeHistogramEq::scalar_udf(),
194            NativeHistogramFraction::scalar_udf(),
195            NativeHistogramIDelta::scalar_udf(),
196            NativeHistogramIRate::scalar_udf(),
197            NativeHistogramIncrease::scalar_udf(),
198            NativeHistogramLastOverTime::scalar_udf(),
199            MixedRange::float_udf(None),
200            MixedRange::histogram_udf(None),
201            NativeHistogramMulScalar::scalar_udf(),
202            NativeHistogramNeg::scalar_udf(),
203            NativeHistogramNotEq::scalar_udf(),
204            NativeHistogramPresentOverTime::scalar_udf(),
205            NativeHistogramQuantile::scalar_udf(),
206            NativeHistogramRate::scalar_udf(),
207            NativeHistogramResets::scalar_udf(),
208            NativeHistogramScalarMul::scalar_udf(),
209            NativeHistogramStddev::scalar_udf(),
210            NativeHistogramStdvar::scalar_udf(),
211            NativeHistogramSub::scalar_udf(),
212            NativeHistogramSum::scalar_udf(),
213            NativeHistogramSumOverTime::scalar_udf(),
214            NativeHistogramToString::scalar_udf(),
215            PromqlFloatToString::scalar_udf(),
216        ] {
217            let _ = session_state.register_udf(Arc::new(udf));
218        }
219        for udaf in [
220            NativeHistogramAggAvg::aggregate_udf(),
221            NativeHistogramAggSum::aggregate_udf(),
222        ] {
223            let _ = session_state.register_udaf(Arc::new(udaf));
224        }
225
226        let logical_plan = DFLogicalSubstraitConvertor
227            .decode(message, session_state)
228            .await
229            .map_err(BoxedError::new)
230            .context(common_query::error::DecodePlanSnafu)?;
231
232        if optimize {
233            self.session_state
234                .optimize(&logical_plan)
235                .map_err(Into::into)
236        } else {
237            Ok(logical_plan)
238        }
239    }
240}
241
242#[cfg(test)]
243mod tests {
244    use common_query::native_histogram::native_histogram_value_type;
245    use datafusion::catalog::TableProvider;
246    use datafusion::datasource::MemTable;
247    use datafusion::logical_expr::Extension;
248    use datafusion_expr::expr::ScalarFunction;
249    use datafusion_expr::{Expr, LogicalPlanBuilder, LogicalTableSource, col, lit};
250    use datatypes::arrow::datatypes::{
251        DataType as ArrowDataType, Field, Schema, SchemaRef, TimeUnit,
252    };
253    use datatypes::data_type::DataType;
254    use promql::extension_plan::RangeManipulate;
255    use session::context::QueryContext;
256
257    use super::*;
258    use crate::QueryEngineFactory;
259    use crate::dummy_catalog::DummyCatalogList;
260    use crate::optimizer::test_util::mock_table_provider;
261    use crate::options::QueryOptions;
262
263    fn mock_plan(schema: SchemaRef) -> LogicalPlan {
264        let table_source = LogicalTableSource::new(schema);
265        let projection = None;
266        let builder =
267            LogicalPlanBuilder::scan("devices", Arc::new(table_source), projection).unwrap();
268
269        builder
270            .filter(col("k0").eq(lit("hello")))
271            .unwrap()
272            .build()
273            .unwrap()
274    }
275
276    #[tokio::test]
277    async fn test_serializer_decode_plan() {
278        let catalog_list = catalog::memory::new_memory_catalog_manager().unwrap();
279        let factory = QueryEngineFactory::new(
280            catalog_list,
281            None,
282            None,
283            None,
284            None,
285            false,
286            QueryOptions::default(),
287        );
288
289        let engine = factory.query_engine();
290
291        let table_provider = Arc::new(mock_table_provider(1.into()));
292        let plan = mock_plan(table_provider.schema().clone());
293
294        let bytes = DFLogicalSubstraitConvertor
295            .encode(&plan, DefaultSerializer)
296            .unwrap();
297
298        let plan_decoder = engine
299            .engine_context(QueryContext::arc())
300            .new_plan_decoder()
301            .unwrap();
302        let catalog_list = Arc::new(DummyCatalogList::with_table_provider(table_provider));
303
304        let decode_plan = plan_decoder
305            .decode(bytes, catalog_list, false)
306            .await
307            .unwrap();
308
309        assert_eq!(
310            "Filter: devices.k0 = Utf8(\"hello\")
311  TableScan: devices",
312            decode_plan.to_string(),
313        );
314    }
315
316    #[tokio::test]
317    async fn test_serializer_decode_native_histogram_udf() {
318        let catalog_list = catalog::memory::new_memory_catalog_manager().unwrap();
319        let factory = QueryEngineFactory::new(
320            catalog_list,
321            None,
322            None,
323            None,
324            None,
325            false,
326            QueryOptions::default(),
327        );
328        let engine = factory.query_engine();
329        let schema = Arc::new(Schema::new(vec![Field::new(
330            "histogram",
331            native_histogram_value_type().as_arrow_type(),
332            true,
333        )]));
334        let plan = LogicalPlanBuilder::scan(
335            "devices",
336            Arc::new(LogicalTableSource::new(schema.clone())),
337            None,
338        )
339        .unwrap()
340        .aggregate(
341            Vec::<Expr>::new(),
342            vec![
343                Arc::new(NativeHistogramAggSum::aggregate_udf())
344                    .call(vec![col("histogram")])
345                    .alias("sum"),
346                Arc::new(NativeHistogramAggAvg::aggregate_udf())
347                    .call(vec![col("histogram")])
348                    .alias("avg"),
349            ],
350        )
351        .unwrap()
352        .project(vec![
353            Expr::ScalarFunction(ScalarFunction {
354                func: Arc::new(NativeHistogramCount::scalar_udf()),
355                args: vec![col("sum")],
356            }),
357            Expr::ScalarFunction(ScalarFunction {
358                func: Arc::new(NativeHistogramDrop::bool_false_udf(
359                    "ignored annotation".to_string(),
360                    None,
361                )),
362                args: vec![col("sum")],
363            }),
364            Expr::ScalarFunction(ScalarFunction {
365                func: Arc::new(NativeHistogramDrop::bool_true_udf(
366                    "ignored annotation".to_string(),
367                    None,
368                )),
369                args: vec![col("sum")],
370            }),
371            Expr::ScalarFunction(ScalarFunction {
372                func: Arc::new(NativeHistogramDrop::float_null_udf(
373                    "ignored annotation".to_string(),
374                    None,
375                )),
376                args: vec![col("sum")],
377            }),
378            Expr::ScalarFunction(ScalarFunction {
379                func: Arc::new(PromqlFloatToString::scalar_udf()),
380                args: vec![lit(2.0)],
381            }),
382        ])
383        .unwrap()
384        .build()
385        .unwrap();
386        let bytes = DFLogicalSubstraitConvertor
387            .encode(&plan, DefaultSerializer)
388            .unwrap();
389        let table_provider = Arc::new(MemTable::try_new(schema, vec![vec![]]).unwrap());
390        let plan_decoder = engine
391            .engine_context(QueryContext::arc())
392            .new_plan_decoder()
393            .unwrap();
394
395        let decoded = plan_decoder
396            .decode(
397                bytes,
398                Arc::new(DummyCatalogList::with_table_provider(table_provider)),
399                false,
400            )
401            .await
402            .unwrap();
403
404        let decoded = decoded.to_string();
405        assert!(decoded.contains("prom_native_histogram_count"));
406        assert!(decoded.contains("prom_native_histogram_drop_bool"));
407        assert!(decoded.contains("prom_native_histogram_keep_bool"));
408        assert!(decoded.contains("prom_native_histogram_drop_float"));
409        assert!(decoded.contains(NativeHistogramAggSum::name()));
410        assert!(decoded.contains(NativeHistogramAggAvg::name()));
411        assert!(decoded.contains(PromqlFloatToString::name()));
412
413        let schema = Arc::new(Schema::new(vec![
414            Field::new(
415                "timestamp",
416                ArrowDataType::Timestamp(TimeUnit::Millisecond, None),
417                false,
418            ),
419            Field::new("float", ArrowDataType::Float64, true),
420            Field::new(
421                "histogram",
422                native_histogram_value_type().as_arrow_type(),
423                true,
424            ),
425        ]));
426        let input = LogicalPlanBuilder::scan(
427            "devices",
428            Arc::new(LogicalTableSource::new(schema.clone())),
429            None,
430        )
431        .unwrap()
432        .build()
433        .unwrap();
434        let input = LogicalPlan::Extension(Extension {
435            node: Arc::new(
436                RangeManipulate::new(
437                    0,
438                    1000,
439                    1000,
440                    1000,
441                    "timestamp".to_string(),
442                    vec!["float".to_string(), "histogram".to_string()],
443                    input,
444                )
445                .unwrap(),
446            ),
447        });
448        let plan = LogicalPlanBuilder::from(input)
449            .project(vec![
450                Expr::ScalarFunction(ScalarFunction {
451                    func: Arc::new(MixedRange::float_udf(None)),
452                    args: vec![
453                        lit("last_over_time"),
454                        col("timestamp_range"),
455                        col("float"),
456                        col("histogram"),
457                    ],
458                })
459                .alias("mixed_float"),
460                Expr::ScalarFunction(ScalarFunction {
461                    func: Arc::new(MixedRange::histogram_udf(None)),
462                    args: vec![
463                        lit("last_over_time"),
464                        col("timestamp_range"),
465                        col("float"),
466                        col("histogram"),
467                    ],
468                })
469                .alias("mixed_histogram"),
470            ])
471            .unwrap()
472            .build()
473            .unwrap();
474        let bytes = DFLogicalSubstraitConvertor
475            .encode(&plan, DefaultSerializer)
476            .unwrap();
477        let table_provider = Arc::new(MemTable::try_new(schema, vec![vec![]]).unwrap());
478        let decoded = plan_decoder
479            .decode(
480                bytes,
481                Arc::new(DummyCatalogList::with_table_provider(table_provider)),
482                false,
483            )
484            .await
485            .unwrap()
486            .to_string();
487        assert!(decoded.contains("prom_mixed_range_float"));
488        assert!(decoded.contains("prom_mixed_range_histogram"));
489    }
490}