1use 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#[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 Err(DataFusionError::Substrait(format!(
92 "Unsupported plan node: {name}"
93 )))
94 } else {
95 ExtensionSerializer.deserialize_logical_plan(name, bytes)
96 }
97 }
98}
99
100pub 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 let mut session_state = SessionStateBuilder::new_from_existing(self.session_state.clone())
128 .with_catalog_list(catalog_list)
129 .build();
130 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}