Skip to main content

query/
planner.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::any::Any;
16use std::borrow::Cow;
17use std::collections::{HashMap, HashSet};
18use std::str::FromStr;
19use std::sync::Arc;
20
21use arrow_schema::DataType;
22use async_trait::async_trait;
23use catalog::table_source::DfTableSourceProvider;
24use common_error::ext::BoxedError;
25use common_query::promql_annotations::promql_annotation_collector;
26use common_telemetry::tracing;
27use datafusion::common::{DFSchema, plan_err};
28use datafusion::execution::SessionStateBuilder;
29use datafusion::execution::context::SessionState;
30use datafusion::sql::planner::PlannerContext;
31use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion};
32use datafusion_common::{ScalarValue, ToDFSchema};
33use datafusion_expr::expr::{Exists, InSubquery};
34use datafusion_expr::{
35    Analyze, Explain, ExplainFormat, Expr as DfExpr, LogicalPlan, LogicalPlanBuilder, PlanType,
36    ToStringifiedPlan, col,
37};
38use datafusion_sql::planner::{ParserOptions, SqlToRel};
39use log_query::LogQuery;
40use promql_parser::parser::EvalStmt;
41use session::context::QueryContextRef;
42use snafu::{ResultExt, ensure};
43use sql::CteContent;
44use sql::ast::Expr as SqlExpr;
45use sql::statements::explain::ExplainStatement;
46use sql::statements::query::Query;
47use sql::statements::statement::Statement;
48use sql::statements::tql::Tql;
49
50use crate::error::{
51    CteColumnSchemaMismatchSnafu, PlanSqlSnafu, QueryPlanSnafu, Result, SqlSnafu,
52    UnimplementedSnafu,
53};
54use crate::log_query::planner::LogQueryPlanner;
55use crate::parser::{DEFAULT_LOOKBACK_STRING, PromQuery, QueryLanguageParser, QueryStatement};
56use crate::promql::planner::PromPlanner;
57use crate::query_engine::{DefaultPlanDecoder, QueryEngineState};
58use crate::range_select::plan_rewrite::RangePlanRewriter;
59use crate::{DfContextProviderAdapter, QueryEngineContext};
60
61#[async_trait]
62pub trait LogicalPlanner: Send + Sync {
63    async fn plan(&self, stmt: &QueryStatement, query_ctx: QueryContextRef) -> Result<LogicalPlan>;
64
65    async fn plan_logs_query(
66        &self,
67        query: LogQuery,
68        query_ctx: QueryContextRef,
69    ) -> Result<LogicalPlan>;
70
71    fn optimize(&self, plan: LogicalPlan) -> Result<LogicalPlan>;
72
73    fn as_any(&self) -> &dyn Any;
74}
75
76pub struct DfLogicalPlanner {
77    engine_state: Arc<QueryEngineState>,
78    session_state: SessionState,
79}
80
81impl DfLogicalPlanner {
82    pub fn new(engine_state: Arc<QueryEngineState>) -> Self {
83        let session_state = engine_state.session_state();
84        Self {
85            engine_state,
86            session_state,
87        }
88    }
89
90    /// Derive a [`SessionState`] whose [`ExecutionProps`] includes
91    /// `query_execution_start_time` if a scheduled time extension is present
92    /// in the query context.
93    fn derive_session_state_with_scheduled_time(
94        &self,
95        query_ctx: &QueryContextRef,
96    ) -> Result<SessionState> {
97        let extensions = query_ctx.extensions();
98        match crate::options::parse_scheduled_time_datetime(&extensions)? {
99            Some(dt) => {
100                let execution_props = self
101                    .session_state
102                    .execution_props()
103                    .clone()
104                    .with_query_execution_start_time(dt);
105                Ok(
106                    SessionStateBuilder::new_from_existing(self.session_state.clone())
107                        .with_execution_props(execution_props)
108                        .build(),
109                )
110            }
111            None => Ok(self.session_state.clone()),
112        }
113    }
114
115    /// Basically the same with `explain_to_plan` in DataFusion, but adapted to Greptime's
116    /// `plan_sql` to support Greptime Statements.
117    async fn explain_to_plan(
118        &self,
119        explain: &ExplainStatement,
120        query_ctx: QueryContextRef,
121    ) -> Result<LogicalPlan> {
122        let plan = self.plan_sql(&explain.statement, query_ctx).await?;
123        if matches!(plan, LogicalPlan::Explain(_)) {
124            return plan_err!("Nested EXPLAINs are not supported").context(PlanSqlSnafu);
125        }
126
127        let verbose = explain.verbose;
128        let analyze = explain.analyze;
129        let format = explain.format.map(|f| f.to_string());
130
131        let plan = Arc::new(plan);
132        let schema = LogicalPlan::explain_schema();
133        let schema = ToDFSchema::to_dfschema_ref(schema)?;
134
135        if verbose && format.is_some() {
136            return plan_err!("EXPLAIN VERBOSE with FORMAT is not supported").context(PlanSqlSnafu);
137        }
138
139        if analyze {
140            // notice format is already set in query context, so can be ignore here
141            Ok(LogicalPlan::Analyze(Analyze {
142                verbose,
143                input: plan,
144                schema,
145            }))
146        } else {
147            let stringified_plans = vec![plan.to_stringified(PlanType::InitialLogicalPlan)];
148
149            // default to configuration value
150            let options = self.session_state.config().options();
151            let format = format
152                .map(|x| ExplainFormat::from_str(&x))
153                .transpose()?
154                .unwrap_or_else(|| options.explain.format.clone());
155
156            Ok(LogicalPlan::Explain(Explain {
157                verbose,
158                explain_format: format,
159                plan,
160                stringified_plans,
161                schema,
162                logical_optimization_succeeded: false,
163            }))
164        }
165    }
166
167    #[tracing::instrument(skip_all)]
168    #[async_recursion::async_recursion]
169    async fn plan_sql(&self, stmt: &Statement, query_ctx: QueryContextRef) -> Result<LogicalPlan> {
170        let mut planner_context = PlannerContext::new();
171        let mut stmt = Cow::Borrowed(stmt);
172        let mut is_tql_cte = false;
173
174        // handle explain before normal processing so we can explain Greptime Statements
175        if let Statement::Explain(explain) = stmt.as_ref() {
176            return self.explain_to_plan(explain, query_ctx).await;
177        }
178
179        // Check for hybrid CTEs before normal processing
180        if self.has_hybrid_ctes(stmt.as_ref()) {
181            let stmt_owned = stmt.into_owned();
182            let mut query = match stmt_owned {
183                Statement::Query(query) => query.as_ref().clone(),
184                _ => unreachable!("has_hybrid_ctes should only return true for Query statements"),
185            };
186            self.plan_query_with_hybrid_ctes(&query, query_ctx.clone(), &mut planner_context)
187                .await?;
188
189            // remove the processed TQL CTEs from the query
190            query.hybrid_cte = None;
191            stmt = Cow::Owned(Statement::Query(Box::new(query)));
192            is_tql_cte = true;
193        }
194
195        let mut df_stmt = stmt.as_ref().try_into().context(SqlSnafu)?;
196
197        // TODO(LFC): Remove this when Datafusion supports **both** the syntax and implementation of "explain with format".
198        if let datafusion::sql::parser::Statement::Statement(
199            box datafusion::sql::sqlparser::ast::Statement::Explain { .. },
200        ) = &mut df_stmt
201        {
202            UnimplementedSnafu {
203                operation: "EXPLAIN with FORMAT using raw datafusion planner",
204            }
205            .fail()?;
206        }
207
208        let scheduled_state = self.derive_session_state_with_scheduled_time(&query_ctx)?;
209        let table_provider = DfTableSourceProvider::new(
210            self.engine_state.catalog_manager().clone(),
211            self.engine_state.disallow_cross_catalog_query(),
212            query_ctx.clone(),
213            Arc::new(DefaultPlanDecoder::new(
214                scheduled_state.clone(),
215                &query_ctx,
216            )?),
217            scheduled_state
218                .config_options()
219                .sql_parser
220                .enable_ident_normalization,
221        );
222
223        let context_provider = DfContextProviderAdapter::try_new(
224            self.engine_state.clone(),
225            scheduled_state.clone(),
226            Some(&df_stmt),
227            query_ctx.clone(),
228        )
229        .await?;
230
231        let config_options = self.session_state.config().options();
232        let parser_options = &config_options.sql_parser;
233        let parser_options = ParserOptions {
234            map_string_types_to_utf8view: false,
235            ..parser_options.into()
236        };
237
238        let sql_to_rel = SqlToRel::new_with_options(&context_provider, parser_options);
239
240        // this IF is to handle different version of ASTs
241        let result = if is_tql_cte {
242            let Statement::Query(query) = stmt.into_owned() else {
243                unreachable!("is_tql_cte should only be true for Query statements");
244            };
245            let sqlparser_stmt = sqlparser::ast::Statement::Query(Box::new(query.inner));
246            sql_to_rel
247                .sql_statement_to_plan_with_context(sqlparser_stmt, &mut planner_context)
248                .context(PlanSqlSnafu)?
249        } else {
250            sql_to_rel
251                .statement_to_plan(df_stmt)
252                .context(PlanSqlSnafu)?
253        };
254
255        common_telemetry::debug!("Logical planner, statement to plan result: {result}");
256        let plan = RangePlanRewriter::new(table_provider, query_ctx.clone())
257            .rewrite(result)
258            .await?;
259
260        // Optimize logical plan by extension rules
261        let context = QueryEngineContext::new(scheduled_state, query_ctx);
262        let plan = self
263            .engine_state
264            .optimize_by_extension_rules(plan, &context)?;
265        common_telemetry::debug!("Logical planner, optimize result: {plan}");
266
267        Ok(plan)
268    }
269
270    /// Generate a relational expression from a SQL expression
271    #[tracing::instrument(skip_all)]
272    pub(crate) async fn sql_to_expr(
273        &self,
274        sql: SqlExpr,
275        schema: &DFSchema,
276        normalize_ident: bool,
277        query_ctx: QueryContextRef,
278    ) -> Result<DfExpr> {
279        let scheduled_state = self.derive_session_state_with_scheduled_time(&query_ctx)?;
280        let context_provider = DfContextProviderAdapter::try_new(
281            self.engine_state.clone(),
282            scheduled_state,
283            None,
284            query_ctx,
285        )
286        .await?;
287
288        let config_options = self.session_state.config().options();
289        let parser_options = &config_options.sql_parser;
290        let parser_options: ParserOptions = ParserOptions {
291            map_string_types_to_utf8view: false,
292            enable_ident_normalization: normalize_ident,
293            ..parser_options.into()
294        };
295
296        let sql_to_rel = SqlToRel::new_with_options(&context_provider, parser_options);
297
298        Ok(sql_to_rel.sql_to_expr(sql, schema, &mut PlannerContext::new())?)
299    }
300
301    #[tracing::instrument(skip_all)]
302    async fn plan_pql(&self, stmt: &EvalStmt, query_ctx: QueryContextRef) -> Result<LogicalPlan> {
303        let mut scheduled_state = self.derive_session_state_with_scheduled_time(&query_ctx)?;
304        let promql_annotations = query_ctx.remote_query_id().map(promql_annotation_collector);
305        if let Some(collector) = &promql_annotations {
306            scheduled_state
307                .config_mut()
308                .options_mut()
309                .extensions
310                .insert(collector.clone());
311        }
312        let plan_decoder = Arc::new(DefaultPlanDecoder::new(
313            scheduled_state.clone(),
314            &query_ctx,
315        )?);
316        let table_provider = DfTableSourceProvider::new(
317            self.engine_state.catalog_manager().clone(),
318            self.engine_state.disallow_cross_catalog_query(),
319            query_ctx.clone(),
320            plan_decoder,
321            scheduled_state
322                .config_options()
323                .sql_parser
324                .enable_ident_normalization,
325        );
326        let plan = PromPlanner::stmt_to_plan_with_annotations(
327            table_provider,
328            stmt,
329            &self.engine_state,
330            promql_annotations,
331        )
332        .await
333        .map_err(BoxedError::new)
334        .context(QueryPlanSnafu)?;
335
336        let context = QueryEngineContext::new(scheduled_state, query_ctx);
337        Ok(self
338            .engine_state
339            .optimize_by_extension_rules(plan, &context)?)
340    }
341
342    #[tracing::instrument(skip_all)]
343    fn optimize_logical_plan(&self, plan: LogicalPlan) -> Result<LogicalPlan> {
344        Ok(self.engine_state.optimize_logical_plan(plan)?)
345    }
346
347    /// Check if a statement contains hybrid CTEs (mix of SQL and TQL)
348    fn has_hybrid_ctes(&self, stmt: &Statement) -> bool {
349        if let Statement::Query(query) = stmt {
350            query
351                .hybrid_cte
352                .as_ref()
353                .map(|hybrid_cte| !hybrid_cte.cte_tables.is_empty())
354                .unwrap_or(false)
355        } else {
356            false
357        }
358    }
359
360    /// Plan a query with hybrid CTEs using DataFusion's native PlannerContext
361    async fn plan_query_with_hybrid_ctes(
362        &self,
363        query: &Query,
364        query_ctx: QueryContextRef,
365        planner_context: &mut PlannerContext,
366    ) -> Result<()> {
367        let hybrid_cte = query.hybrid_cte.as_ref().unwrap();
368
369        for cte in &hybrid_cte.cte_tables {
370            match &cte.content {
371                CteContent::Tql(tql) => {
372                    // Plan TQL and register in PlannerContext
373                    let mut logical_plan = self.tql_to_logical_plan(tql, query_ctx.clone()).await?;
374                    if !cte.columns.is_empty() {
375                        let schema = logical_plan.schema();
376                        let schema_fields = schema.fields().to_vec();
377                        ensure!(
378                            schema_fields.len() == cte.columns.len(),
379                            CteColumnSchemaMismatchSnafu {
380                                cte_name: cte.name.value.clone(),
381                                original: schema_fields
382                                    .iter()
383                                    .map(|field| field.name().clone())
384                                    .collect::<Vec<_>>(),
385                                expected: cte
386                                    .columns
387                                    .iter()
388                                    .map(|column| column.to_string())
389                                    .collect::<Vec<_>>(),
390                            }
391                        );
392                        let aliases = cte
393                            .columns
394                            .iter()
395                            .zip(schema_fields.iter())
396                            .map(|(column, field)| col(field.name()).alias(column.to_string()));
397                        logical_plan = LogicalPlanBuilder::from(logical_plan)
398                            .project(aliases)
399                            .context(PlanSqlSnafu)?
400                            .build()
401                            .context(PlanSqlSnafu)?;
402                    }
403
404                    // Wrap in SubqueryAlias to ensure proper table qualification for CTE
405                    logical_plan = LogicalPlan::SubqueryAlias(
406                        datafusion_expr::SubqueryAlias::try_new(
407                            Arc::new(logical_plan),
408                            cte.name.value.clone(),
409                        )
410                        .context(PlanSqlSnafu)?,
411                    );
412
413                    planner_context.insert_cte(&cte.name.value, logical_plan);
414                }
415                CteContent::Sql(_) => {
416                    // SQL CTEs should have been moved to the main query's WITH clause
417                    // during parsing, so we shouldn't encounter them here
418                    unreachable!("SQL CTEs should not be in hybrid_cte.cte_tables");
419                }
420            }
421        }
422
423        Ok(())
424    }
425
426    /// Convert TQL to LogicalPlan directly
427    async fn tql_to_logical_plan(
428        &self,
429        tql: &Tql,
430        query_ctx: QueryContextRef,
431    ) -> Result<LogicalPlan> {
432        match tql {
433            Tql::Eval(eval) => {
434                // Convert TqlEval to PromQuery then to QueryStatement::Promql
435                let prom_query = PromQuery {
436                    query: eval.query.clone(),
437                    start: eval.start.clone(),
438                    end: eval.end.clone(),
439                    step: eval.step.clone(),
440                    lookback: eval
441                        .lookback
442                        .clone()
443                        .unwrap_or_else(|| DEFAULT_LOOKBACK_STRING.to_string()),
444                    alias: eval.alias.clone(),
445                };
446                let stmt = QueryLanguageParser::parse_promql(&prom_query, &query_ctx)?;
447
448                self.plan(&stmt, query_ctx).await
449            }
450            Tql::Explain(_) => UnimplementedSnafu {
451                operation: "TQL EXPLAIN in CTEs",
452            }
453            .fail(),
454            Tql::Analyze(_) => UnimplementedSnafu {
455                operation: "TQL ANALYZE in CTEs",
456            }
457            .fail(),
458        }
459    }
460
461    /// Extracts cast types for all placeholders in a logical plan.
462    /// Returns a map where each placeholder ID is mapped to:
463    /// - Some(DataType) if the placeholder is cast to a specific type
464    /// - None if the placeholder exists but has no cast
465    ///
466    /// Example: `$1::TEXT` returns `{"$1": Some(DataType::Utf8)}`
467    ///
468    /// This function walks through all expressions in the logical plan,
469    /// including subqueries, to identify placeholders and their cast types.
470    fn extract_placeholder_cast_types(
471        plan: &LogicalPlan,
472    ) -> Result<HashMap<String, Option<DataType>>> {
473        let mut placeholder_types = HashMap::new();
474        let mut casted_placeholders = HashSet::new();
475
476        Self::extract_from_plan(plan, &mut placeholder_types, &mut casted_placeholders)?;
477
478        Ok(placeholder_types)
479    }
480
481    fn extract_from_plan(
482        plan: &LogicalPlan,
483        placeholder_types: &mut HashMap<String, Option<DataType>>,
484        casted_placeholders: &mut HashSet<String>,
485    ) -> Result<()> {
486        plan.apply(|node| {
487            for expr in node.expressions() {
488                let _ = expr.apply(|e| {
489                    // Handle casted placeholders
490                    if let DfExpr::Cast(cast) = e
491                        && let DfExpr::Placeholder(ph) = &*cast.expr
492                    {
493                        placeholder_types.insert(ph.id.clone(), Some(cast.data_type.clone()));
494                        casted_placeholders.insert(ph.id.clone());
495                    }
496
497                    // Handle arrow_cast(Placeholder, 'type_string') generated by SQL rewriter
498                    if let DfExpr::ScalarFunction(scalar_func) = e
499                        && scalar_func.name() == "arrow_cast"
500                        && scalar_func.args.len() == 2
501                        && let DfExpr::Placeholder(ph) = &scalar_func.args[0]
502                        && let DfExpr::Literal(ScalarValue::Utf8(Some(type_str)), _) =
503                            &scalar_func.args[1]
504                        && let Ok(data_type) = type_str.parse::<DataType>()
505                    {
506                        placeholder_types.insert(ph.id.clone(), Some(data_type));
507                        casted_placeholders.insert(ph.id.clone());
508                    }
509
510                    // Handle bare (non-casted) placeholders
511                    if let DfExpr::Placeholder(ph) = e
512                        && !casted_placeholders.contains(&ph.id)
513                        && !placeholder_types.contains_key(&ph.id)
514                    {
515                        placeholder_types.insert(ph.id.clone(), None);
516                    }
517
518                    // Recurse into subquery plans embedded in expressions
519                    match e {
520                        DfExpr::Exists(Exists { subquery, .. })
521                        | DfExpr::InSubquery(InSubquery { subquery, .. })
522                        | DfExpr::ScalarSubquery(subquery) => {
523                            Self::extract_from_plan(
524                                &subquery.subquery,
525                                placeholder_types,
526                                casted_placeholders,
527                            )?;
528                        }
529                        _ => {}
530                    }
531
532                    Ok(TreeNodeRecursion::Continue)
533                });
534            }
535            Ok(TreeNodeRecursion::Continue)
536        })?;
537        Ok(())
538    }
539
540    fn infer_limit_placeholder_types(
541        plan: &LogicalPlan,
542        placeholder_types: &mut HashMap<String, Option<DataType>>,
543    ) -> Result<()> {
544        plan.apply(|node| {
545            if let LogicalPlan::Limit(limit) = node {
546                for expr in limit.skip.iter().chain(limit.fetch.iter()) {
547                    expr.apply(|e| {
548                        if let DfExpr::Placeholder(ph) = e {
549                            placeholder_types
550                                .entry(ph.id.clone())
551                                .and_modify(|existing| {
552                                    if existing.is_none() {
553                                        *existing = Some(DataType::Int64);
554                                    }
555                                })
556                                .or_insert(Some(DataType::Int64));
557                        }
558
559                        Ok(TreeNodeRecursion::Continue)
560                    })?;
561                }
562            }
563
564            Ok(TreeNodeRecursion::Continue)
565        })?;
566
567        Ok(())
568    }
569
570    /// Gets inferred parameter types from a logical plan.
571    /// Returns a map where each parameter ID is mapped to:
572    /// - Some(DataType) if the parameter type could be inferred
573    /// - None if the parameter type could not be inferred
574    ///
575    /// This function first uses DataFusion's `get_parameter_types()` to infer types.
576    /// If any parameters have `None` values (i.e., DataFusion couldn't infer their types),
577    /// it falls back to using `extract_placeholder_cast_types()` to detect explicit casts
578    /// and applies context-specific inference such as LIMIT/OFFSET placeholders.
579    ///
580    /// This is because datafusion can only infer types for a limited cases.
581    ///
582    /// Example: For query `WHERE $1::TEXT AND $2`, DataFusion may not infer `$2`'s type,
583    /// but this function will return `{"$1": Some(DataType::Utf8), "$2": None}`.
584    pub fn get_inferred_parameter_types(
585        plan: &LogicalPlan,
586    ) -> Result<HashMap<String, Option<DataType>>> {
587        let mut param_types = plan.get_parameter_types().context(PlanSqlSnafu)?;
588
589        let has_none = param_types.values().any(|v| v.is_none());
590
591        if has_none {
592            let cast_types = Self::extract_placeholder_cast_types(plan)?;
593
594            for (id, opt_type) in cast_types {
595                param_types
596                    .entry(id)
597                    .and_modify(|existing| {
598                        if existing.is_none() {
599                            *existing = opt_type.clone();
600                        }
601                    })
602                    .or_insert(opt_type);
603            }
604
605            Self::infer_limit_placeholder_types(plan, &mut param_types)?;
606        }
607
608        Ok(param_types)
609    }
610}
611
612#[async_trait]
613impl LogicalPlanner for DfLogicalPlanner {
614    #[tracing::instrument(skip_all)]
615    async fn plan(&self, stmt: &QueryStatement, query_ctx: QueryContextRef) -> Result<LogicalPlan> {
616        match stmt {
617            QueryStatement::Sql(stmt) => self.plan_sql(stmt, query_ctx).await,
618            QueryStatement::Promql(stmt, _alias) => self.plan_pql(stmt, query_ctx).await,
619        }
620    }
621
622    async fn plan_logs_query(
623        &self,
624        query: LogQuery,
625        query_ctx: QueryContextRef,
626    ) -> Result<LogicalPlan> {
627        let plan_decoder = Arc::new(DefaultPlanDecoder::new(
628            self.session_state.clone(),
629            &query_ctx,
630        )?);
631        let table_provider = DfTableSourceProvider::new(
632            self.engine_state.catalog_manager().clone(),
633            self.engine_state.disallow_cross_catalog_query(),
634            query_ctx,
635            plan_decoder,
636            self.session_state
637                .config_options()
638                .sql_parser
639                .enable_ident_normalization,
640        );
641
642        let mut planner = LogQueryPlanner::new(table_provider, self.session_state.clone());
643        planner
644            .query_to_plan(query)
645            .await
646            .map_err(BoxedError::new)
647            .context(QueryPlanSnafu)
648    }
649
650    fn optimize(&self, plan: LogicalPlan) -> Result<LogicalPlan> {
651        self.optimize_logical_plan(plan)
652    }
653
654    fn as_any(&self) -> &dyn Any {
655        self
656    }
657}
658
659#[cfg(test)]
660mod tests {
661    use std::sync::Arc;
662
663    use arrow_schema::DataType;
664    use catalog::RegisterTableRequest;
665    use catalog::memory::MemoryCatalogManager;
666    use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
667    use datatypes::prelude::ConcreteDataType;
668    use datatypes::schema::{ColumnSchema, Schema};
669    use session::context::QueryContext;
670    use store_api::metric_engine_consts::{
671        DATA_SCHEMA_TABLE_ID_COLUMN_NAME, DATA_SCHEMA_TSID_COLUMN_NAME, LOGICAL_TABLE_METADATA_KEY,
672        METRIC_ENGINE_NAME,
673    };
674    use table::metadata::{TableInfoBuilder, TableMetaBuilder};
675    use table::test_util::EmptyTable;
676
677    use super::*;
678    use crate::parser::{PromQuery, QueryLanguageParser};
679    use crate::{QueryEngineFactory, QueryEngineRef};
680
681    async fn create_test_engine() -> QueryEngineRef {
682        let columns = vec![
683            ColumnSchema::new("id", ConcreteDataType::int32_datatype(), false),
684            ColumnSchema::new("name", ConcreteDataType::string_datatype(), true),
685        ];
686        let schema = Arc::new(Schema::new(columns));
687        let table_meta = TableMetaBuilder::empty()
688            .schema(schema)
689            .primary_key_indices(vec![0])
690            .value_indices(vec![1])
691            .next_column_id(1024)
692            .build()
693            .unwrap();
694        let table_info = TableInfoBuilder::new("test", table_meta).build().unwrap();
695        let table = EmptyTable::from_table_info(&table_info);
696
697        crate::tests::new_query_engine_with_table(table)
698    }
699
700    fn create_promql_test_engine() -> QueryEngineRef {
701        let catalog_manager = MemoryCatalogManager::with_default_setup();
702        let physical_table_name = "phy";
703        let physical_table_id = 999u32;
704
705        let physical_schema = Arc::new(Schema::new(vec![
706            ColumnSchema::new(
707                DATA_SCHEMA_TABLE_ID_COLUMN_NAME.to_string(),
708                ConcreteDataType::uint32_datatype(),
709                false,
710            ),
711            ColumnSchema::new(
712                DATA_SCHEMA_TSID_COLUMN_NAME.to_string(),
713                ConcreteDataType::uint64_datatype(),
714                false,
715            ),
716            ColumnSchema::new("tag_0", ConcreteDataType::string_datatype(), false),
717            ColumnSchema::new("tag_1", ConcreteDataType::string_datatype(), false),
718            ColumnSchema::new(
719                "timestamp",
720                ConcreteDataType::timestamp_millisecond_datatype(),
721                false,
722            )
723            .with_time_index(true),
724            ColumnSchema::new("field_0", ConcreteDataType::float64_datatype(), true),
725        ]));
726        let physical_meta = TableMetaBuilder::empty()
727            .schema(physical_schema)
728            .primary_key_indices(vec![0, 1, 2, 3])
729            .value_indices(vec![4, 5])
730            .engine(METRIC_ENGINE_NAME.to_string())
731            .next_column_id(1024)
732            .build()
733            .unwrap();
734        let physical_info = TableInfoBuilder::default()
735            .table_id(physical_table_id)
736            .name(physical_table_name)
737            .meta(physical_meta)
738            .build()
739            .unwrap();
740        catalog_manager
741            .register_table_sync(RegisterTableRequest {
742                catalog: DEFAULT_CATALOG_NAME.to_string(),
743                schema: DEFAULT_SCHEMA_NAME.to_string(),
744                table_name: physical_table_name.to_string(),
745                table_id: physical_table_id,
746                table: EmptyTable::from_table_info(&physical_info),
747            })
748            .unwrap();
749
750        let mut options = table::requests::TableOptions::default();
751        options.extra_options.insert(
752            LOGICAL_TABLE_METADATA_KEY.to_string(),
753            physical_table_name.to_string(),
754        );
755        let logical_schema = Arc::new(Schema::new(vec![
756            ColumnSchema::new("tag_0", ConcreteDataType::string_datatype(), false),
757            ColumnSchema::new("tag_1", ConcreteDataType::string_datatype(), false),
758            ColumnSchema::new(
759                "timestamp",
760                ConcreteDataType::timestamp_millisecond_datatype(),
761                false,
762            )
763            .with_time_index(true),
764            ColumnSchema::new("field_0", ConcreteDataType::float64_datatype(), true),
765        ]));
766        let logical_meta = TableMetaBuilder::empty()
767            .schema(logical_schema)
768            .primary_key_indices(vec![0, 1])
769            .value_indices(vec![3])
770            .engine(METRIC_ENGINE_NAME.to_string())
771            .options(options)
772            .next_column_id(1024)
773            .build()
774            .unwrap();
775        let logical_info = TableInfoBuilder::default()
776            .table_id(1024)
777            .name("some_metric")
778            .meta(logical_meta)
779            .build()
780            .unwrap();
781        catalog_manager
782            .register_table_sync(RegisterTableRequest {
783                catalog: DEFAULT_CATALOG_NAME.to_string(),
784                schema: DEFAULT_SCHEMA_NAME.to_string(),
785                table_name: "some_metric".to_string(),
786                table_id: 1024,
787                table: EmptyTable::from_table_info(&logical_info),
788            })
789            .unwrap();
790
791        QueryEngineFactory::new(
792            catalog_manager,
793            None,
794            None,
795            None,
796            None,
797            false,
798            crate::options::QueryOptions::default(),
799        )
800        .query_engine()
801    }
802
803    async fn parse_sql_to_plan(sql: &str) -> LogicalPlan {
804        let stmt = QueryLanguageParser::parse_sql(sql, &QueryContext::arc()).unwrap();
805        let engine = create_test_engine().await;
806        engine
807            .planner()
808            .plan(&stmt, QueryContext::arc())
809            .await
810            .unwrap()
811    }
812
813    async fn parse_promql_to_plan(query: &str) -> LogicalPlan {
814        let engine = create_promql_test_engine();
815        let query_ctx = QueryContext::arc();
816        let stmt = QueryLanguageParser::parse_promql(
817            &PromQuery {
818                query: query.to_string(),
819                start: "0".to_string(),
820                end: "10".to_string(),
821                step: "5s".to_string(),
822                lookback: "300s".to_string(),
823                alias: None,
824            },
825            &query_ctx,
826        )
827        .unwrap();
828
829        engine.planner().plan(&stmt, query_ctx).await.unwrap()
830    }
831
832    #[tokio::test]
833    async fn test_extract_placeholder_cast_types_multiple() {
834        let plan = parse_sql_to_plan(
835            "SELECT $1::INT, $2::TEXT, $3, $4::INTEGER FROM test WHERE $5::FLOAT > 0",
836        )
837        .await;
838        let types = DfLogicalPlanner::extract_placeholder_cast_types(&plan).unwrap();
839
840        assert_eq!(types.len(), 5);
841        assert_eq!(types.get("$1"), Some(&Some(DataType::Int32)));
842        assert_eq!(types.get("$2"), Some(&Some(DataType::Utf8)));
843        assert_eq!(types.get("$3"), Some(&None));
844        assert_eq!(types.get("$4"), Some(&Some(DataType::Int32)));
845        assert_eq!(types.get("$5"), Some(&Some(DataType::Float32)));
846    }
847
848    #[tokio::test]
849    async fn test_get_inferred_parameter_types_fallback_for_udf_args() {
850        // datafusion is not able to infer type for scalar function arguments
851        let plan = parse_sql_to_plan(
852            "SELECT parse_ident($1), parse_ident($2::TEXT) FROM test WHERE id > $3",
853        )
854        .await;
855        let types = DfLogicalPlanner::get_inferred_parameter_types(&plan).unwrap();
856
857        assert_eq!(types.len(), 3);
858
859        let type_1 = types.get("$1").unwrap();
860        let type_2 = types.get("$2").unwrap();
861        let type_3 = types.get("$3").unwrap();
862
863        assert!(type_1.is_none(), "Expected $1 to be None");
864        assert_eq!(type_2, &Some(DataType::Utf8));
865        assert_eq!(type_3, &Some(DataType::Int32));
866    }
867
868    #[tokio::test]
869    async fn test_get_inferred_parameter_types_limit_offset() {
870        let plan = parse_sql_to_plan("SELECT id FROM test LIMIT $1 OFFSET $2").await;
871        let types = DfLogicalPlanner::get_inferred_parameter_types(&plan).unwrap();
872
873        assert_eq!(types.get("$1"), Some(&Some(DataType::Int64)));
874        assert_eq!(types.get("$2"), Some(&Some(DataType::Int64)));
875    }
876
877    #[tokio::test]
878    async fn test_plan_pql_applies_extension_rules() {
879        for inner_agg in ["count", "sum", "avg", "min", "max", "stddev", "stdvar"] {
880            let plan = parse_promql_to_plan(&format!(
881                "sum(irate(some_metric[1h])) / scalar(count({inner_agg}(some_metric) by (tag_0)))"
882            ))
883            .await;
884            let plan_str = plan.display_indent_schema().to_string();
885            assert!(plan_str.contains("Distinct:"), "{inner_agg}: {plan_str}");
886        }
887    }
888
889    #[tokio::test]
890    async fn test_plan_pql_filters_null_only_groups_for_non_count_inner_aggs() {
891        let count_plan = parse_promql_to_plan("scalar(count(count(some_metric) by (tag_0)))").await;
892        let count_plan_str = count_plan.display_indent_schema().to_string();
893        assert!(
894            !count_plan_str.contains("field_0 IS NOT NULL"),
895            "{count_plan_str}"
896        );
897
898        for inner_agg in ["sum", "avg", "min", "max", "stddev", "stdvar"] {
899            let plan = parse_promql_to_plan(&format!(
900                "scalar(count({inner_agg}(some_metric) by (tag_0)))"
901            ))
902            .await;
903            let plan_str = plan.display_indent_schema().to_string();
904            assert!(
905                plan_str.contains("field_0 IS NOT NULL"),
906                "{inner_agg}: {plan_str}"
907            );
908        }
909    }
910
911    #[tokio::test]
912    async fn test_plan_pql_skips_extension_rules_for_non_direct_or_unsupported_inner_agg() {
913        for query in [
914            "sum(irate(some_metric[1h])) / scalar(count(sum(irate(some_metric[1h])) by (tag_0)))",
915            "sum(irate(some_metric[1h])) / scalar(count(group(some_metric) by (tag_0)))",
916        ] {
917            let plan = parse_promql_to_plan(query).await;
918            let plan_str = plan.display_indent_schema().to_string();
919            assert!(!plan_str.contains("Distinct:"), "{query}: {plan_str}");
920        }
921    }
922
923    #[tokio::test]
924    async fn test_plan_sql_does_not_apply_nested_count_rule() {
925        let plan = parse_sql_to_plan(
926            "SELECT id, count(inner_count) \
927             FROM ( \
928                 SELECT id, count(name) AS inner_count \
929                 FROM test \
930                 GROUP BY id \
931                 ORDER BY id \
932                 LIMIT 1000000 \
933             ) t \
934             GROUP BY id \
935             ORDER BY id",
936        )
937        .await;
938
939        let plan_str = plan.display_indent_schema().to_string();
940        assert!(!plan_str.contains("Distinct:"), "{plan_str}");
941    }
942
943    #[tokio::test]
944    async fn test_get_inferred_parameter_types_subquery() {
945        let plan = parse_sql_to_plan(
946            r#"SELECT * FROM test WHERE id = (SELECT id FROM test CROSS JOIN (SELECT parse_ident($1::TEXT) AS parts) p LIMIT 1)"#,
947        ).await;
948        let types = DfLogicalPlanner::get_inferred_parameter_types(&plan).unwrap();
949
950        assert_eq!(types.len(), 1);
951        let type_1 = types.get("$1").unwrap();
952        assert_eq!(type_1, &Some(DataType::Utf8));
953    }
954
955    #[tokio::test]
956    async fn test_get_inferred_parameter_types_insert() {
957        let plan = parse_sql_to_plan("INSERT INTO test (id, name) VALUES ($1, $2), ($3, $4)").await;
958        let types = DfLogicalPlanner::get_inferred_parameter_types(&plan).unwrap();
959
960        assert_eq!(types.len(), 4);
961        assert_eq!(types.get("$1"), Some(&Some(DataType::Int32)));
962        assert_eq!(types.get("$2"), Some(&Some(DataType::Utf8)));
963        assert_eq!(types.get("$3"), Some(&Some(DataType::Int32)));
964        assert_eq!(types.get("$4"), Some(&Some(DataType::Utf8)));
965    }
966
967    #[tokio::test]
968    async fn test_get_inferred_parameter_types_arrow_cast() {
969        let plan = parse_sql_to_plan("SELECT $1::INT64, $2::FLOAT64, $3::INT16, $4::INT32, $5::UINT8, $6::UINT16, $7::UINT32").await;
970        let types = DfLogicalPlanner::get_inferred_parameter_types(&plan).unwrap();
971
972        assert_eq!(types.get("$1"), Some(&Some(DataType::Int64)));
973        assert_eq!(types.get("$2"), Some(&Some(DataType::Float64)));
974        assert_eq!(types.get("$3"), Some(&Some(DataType::Int16)));
975        assert_eq!(types.get("$4"), Some(&Some(DataType::Int32)));
976        assert_eq!(types.get("$5"), Some(&Some(DataType::UInt8)));
977        assert_eq!(types.get("$6"), Some(&Some(DataType::UInt16)));
978        assert_eq!(types.get("$7"), Some(&Some(DataType::UInt32)));
979
980        let plan = parse_sql_to_plan("SELECT $1::INT8, $2::FLOAT8, $3::INT2, $4::INT8").await;
981        let types = DfLogicalPlanner::get_inferred_parameter_types(&plan).unwrap();
982
983        assert_eq!(types.get("$1"), Some(&Some(DataType::Int64)));
984        assert_eq!(types.get("$2"), Some(&Some(DataType::Float64)));
985        assert_eq!(types.get("$3"), Some(&Some(DataType::Int16)));
986    }
987}