Skip to main content

query/query_engine/
state.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::collections::HashMap;
16use std::fmt;
17use std::num::NonZeroUsize;
18use std::sync::{Arc, RwLock};
19
20use async_trait::async_trait;
21use catalog::CatalogManagerRef;
22use common_base::Plugins;
23use common_function::aggrs::aggr_wrapper::fix_order::FixStateUdafOrderingAnalyzer;
24use common_function::function_factory::ScalarFunctionFactory;
25use common_function::function_registry::FUNCTION_REGISTRY;
26use common_function::handlers::{
27    FlowServiceHandlerRef, ProcedureServiceHandlerRef, TableMutationHandlerRef,
28};
29use common_function::state::FunctionState;
30use common_stat::get_total_memory_bytes;
31use common_telemetry::warn;
32use datafusion::catalog::{Session, TableFunction};
33use datafusion::dataframe::DataFrame;
34use datafusion::error::Result as DfResult;
35use datafusion::execution::SessionStateBuilder;
36use datafusion::execution::context::{QueryPlanner, SessionConfig, SessionContext, SessionState};
37use datafusion::execution::memory_pool::{
38    FairSpillPool, GreedyMemoryPool, MemoryConsumer, MemoryLimit, MemoryPool, MemoryReservation,
39    TrackConsumersPool,
40};
41use datafusion::physical_optimizer::PhysicalOptimizerRule;
42use datafusion::physical_optimizer::optimizer::PhysicalOptimizer;
43use datafusion::physical_optimizer::sanity_checker::SanityCheckPlan;
44use datafusion::physical_plan::ExecutionPlan;
45use datafusion::physical_planner::{DefaultPhysicalPlanner, ExtensionPlanner, PhysicalPlanner};
46use datafusion_expr::{AggregateUDF, LogicalPlan as DfLogicalPlan, WindowUDF};
47use datafusion_optimizer::Analyzer;
48use datafusion_optimizer::analyzer::function_rewrite::ApplyFunctionRewrites;
49use datafusion_optimizer::optimizer::Optimizer;
50use partition::manager::PartitionRuleManagerRef;
51use promql::extension_plan::PromExtensionPlanner;
52use session::context::QueryContextRef;
53use table::TableRef;
54use table::table::adapter::DfTableProviderAdapter;
55
56use crate::QueryEngineContext;
57use crate::dist_plan::{
58    DistExtensionPlanner, DistPlannerAnalyzer, DistPlannerOptions, DynFilterRegistryManager,
59    MergeSortExtensionPlanner, RemoteDynFilterReceiverExtensionPlanner,
60    RemoteDynFilterRegistryLease,
61};
62use crate::metrics::{QUERY_MEMORY_POOL_REJECTED_TOTAL, QUERY_MEMORY_POOL_USAGE_BYTES};
63use crate::optimizer::ExtensionAnalyzerRule;
64use crate::optimizer::const_normalization::ConstNormalizationRule;
65use crate::optimizer::constant_term::MatchesConstantTermOptimizer;
66use crate::optimizer::count_nest_aggr::CountNestAggrRule;
67use crate::optimizer::count_wildcard::CountWildcardToTimeIndexRule;
68use crate::optimizer::enforce_sorting::EnforceSorting;
69use crate::optimizer::global_limit::EnsureGlobalLimitForFetch;
70use crate::optimizer::json_schema_concretize::JsonSchemaConcretizeRule;
71use crate::optimizer::json_type_concretize::JsonTypeConcretizeRule;
72use crate::optimizer::parallelize_scan::ParallelizeScan;
73use crate::optimizer::pass_distribution::PassDistribution;
74use crate::optimizer::promql_tsid_narrow_join::PromqlTsidNarrowJoin;
75use crate::optimizer::remove_duplicate::RemoveDuplicate;
76use crate::optimizer::scan_hint::ScanHintRule;
77use crate::optimizer::string_normalization::StringNormalizationRule;
78use crate::optimizer::transcribe_atat::TranscribeAtatRule;
79use crate::optimizer::type_conversion::TypeConversionRule;
80use crate::optimizer::windowed_sort::WindowedSortPhysicalRule;
81use crate::options::{QueryMemoryPoolPolicy, QueryOptions as QueryOptionsNew};
82use crate::query_engine::DefaultSerializer;
83use crate::query_engine::options::QueryOptions;
84use crate::query_engine::runtime::{
85    DefaultQueryRuntimeProvider, QueryRuntimeContext, QueryRuntimeProvider, QueryRuntimeProviderRef,
86};
87use crate::range_select::planner::RangeSelectPlanner;
88use crate::region_query::RegionQueryHandlerRef;
89
90/// Query engine global state
91#[derive(Clone)]
92pub struct QueryEngineState {
93    df_context: SessionContext,
94    catalog_manager: CatalogManagerRef,
95    dyn_filter_registry_manager: Arc<DynFilterRegistryManager>,
96    function_state: Arc<FunctionState>,
97    scalar_functions: Arc<RwLock<HashMap<String, ScalarFunctionFactory>>>,
98    aggr_functions: Arc<RwLock<HashMap<String, AggregateUDF>>>,
99    table_functions: Arc<RwLock<HashMap<String, Arc<TableFunction>>>>,
100    extension_rules: Vec<Arc<dyn ExtensionAnalyzerRule + Send + Sync>>,
101    plugins: Plugins,
102}
103
104impl fmt::Debug for QueryEngineState {
105    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
106        f.debug_struct("QueryEngineState")
107            .field("state", &self.df_context.state())
108            .finish()
109    }
110}
111
112impl QueryEngineState {
113    #[allow(clippy::too_many_arguments)]
114    pub fn new(
115        catalog_list: CatalogManagerRef,
116        partition_rule_manager: Option<PartitionRuleManagerRef>,
117        region_query_handler: Option<RegionQueryHandlerRef>,
118        table_mutation_handler: Option<TableMutationHandlerRef>,
119        procedure_service_handler: Option<ProcedureServiceHandlerRef>,
120        flow_service_handler: Option<FlowServiceHandlerRef>,
121        with_dist_planner: bool,
122        plugins: Plugins,
123        options: QueryOptionsNew,
124    ) -> Self {
125        Self::try_new(
126            catalog_list,
127            partition_rule_manager,
128            region_query_handler,
129            table_mutation_handler,
130            procedure_service_handler,
131            flow_service_handler,
132            with_dist_planner,
133            plugins,
134            options,
135        )
136        .expect("Failed to build query engine state")
137    }
138
139    #[allow(clippy::too_many_arguments)]
140    pub fn try_new(
141        catalog_list: CatalogManagerRef,
142        partition_rule_manager: Option<PartitionRuleManagerRef>,
143        region_query_handler: Option<RegionQueryHandlerRef>,
144        table_mutation_handler: Option<TableMutationHandlerRef>,
145        procedure_service_handler: Option<ProcedureServiceHandlerRef>,
146        flow_service_handler: Option<FlowServiceHandlerRef>,
147        with_dist_planner: bool,
148        plugins: Plugins,
149        options: QueryOptionsNew,
150    ) -> DfResult<Self> {
151        let total_memory = get_total_memory_bytes().max(0) as u64;
152        let memory_pool_size = options.memory_pool_size.resolve(total_memory) as usize;
153        let runtime_provider = plugins.get::<QueryRuntimeProviderRef>();
154        let mut session_config = SessionConfig::new().with_create_default_catalog_and_schema(false);
155        if options.parallelism > 0 {
156            session_config = session_config.with_target_partitions(options.parallelism);
157        }
158        if options.allow_query_fallback {
159            session_config
160                .options_mut()
161                .extensions
162                .insert(DistPlannerOptions {
163                    allow_query_fallback: true,
164                });
165        }
166
167        // todo(hl): This serves as a workaround for https://github.com/GreptimeTeam/greptimedb/issues/5659
168        // and we can add that check back once we upgrade datafusion.
169        session_config
170            .options_mut()
171            .execution
172            .skip_physical_aggregate_schema_check = true;
173
174        let runtime_context = QueryRuntimeContext::new(&options, memory_pool_size);
175        let default_runtime_provider = DefaultQueryRuntimeProvider;
176        default_runtime_provider.configure_session_config(runtime_context, &mut session_config);
177        if let Some(provider) = runtime_provider.as_ref() {
178            provider.configure_session_config(runtime_context, &mut session_config);
179        }
180        let runtime_builder = DefaultQueryRuntimeProvider::runtime_env_builder(runtime_context);
181        let runtime_env = match runtime_provider {
182            Some(provider) => provider.build_runtime_env(runtime_context, runtime_builder)?,
183            None => default_runtime_provider.build_runtime_env(runtime_context, runtime_builder)?,
184        };
185
186        // Apply extension rules
187        let mut extension_rules = Vec::new();
188
189        // The [`TypeConversionRule`] must be at first
190        extension_rules.insert(0, Arc::new(TypeConversionRule) as _);
191        extension_rules.push(Arc::new(CountNestAggrRule) as _);
192
193        // Apply the datafusion rules
194        let mut analyzer = Analyzer::new();
195        analyzer.rules.insert(0, Arc::new(TranscribeAtatRule));
196        analyzer.rules.insert(0, Arc::new(StringNormalizationRule));
197        analyzer
198            .rules
199            .insert(0, Arc::new(CountWildcardToTimeIndexRule));
200        analyzer.rules.push(Arc::new(ConstNormalizationRule));
201
202        // Add ApplyFunctionRewrites rule,
203        // Note we cannot use `analyzer.add_function_rewrite`
204        // because only rules are copied into session_state
205        analyzer.rules.insert(
206            0,
207            Arc::new(ApplyFunctionRewrites::new(
208                FUNCTION_REGISTRY.function_rewrites(),
209            )),
210        );
211        if with_dist_planner {
212            analyzer.rules.push(Arc::new(DistPlannerAnalyzer));
213            analyzer.rules.push(Arc::new(JsonSchemaConcretizeRule));
214        }
215        analyzer.rules.push(Arc::new(FixStateUdafOrderingAnalyzer));
216
217        // Note: Postgres oid-alias string coercion (`'x'::regclass`,
218        // `'public'::regnamespace`, ...) is handled by the
219        // `PostgresCompatibilityParser`'s built-in `RewriteRegCastToSubquery`
220        // rule at SQL-parse time, so no analyzer rule is registered here. The
221        // implicit `oid_col = 'name'` form is not resolved (it needs schema
222        // awareness the parser lacks); the few client queries that use it are
223        // handled by the parser's blacklist.
224
225        let mut optimizer = Optimizer::new();
226        optimizer.rules.push(Arc::new(ScanHintRule));
227        optimizer.rules.push(Arc::new(JsonTypeConcretizeRule));
228
229        // add physical optimizer
230        let mut physical_optimizer = PhysicalOptimizer::new();
231        // Change TableScan's partition right before enforcing distribution
232        physical_optimizer
233            .rules
234            .insert(5, Arc::new(ParallelizeScan));
235        // Pass distribution requirement to MergeScanExec to avoid unnecessary shuffling
236        physical_optimizer
237            .rules
238            .insert(6, Arc::new(PassDistribution));
239        // Prefer collecting narrow PromQL build sides over repartitioning wide label streams.
240        physical_optimizer
241            .rules
242            .insert(7, Arc::new(PromqlTsidNarrowJoin));
243        // Re-enforce sorting after custom rules update scan partitioning and distribution.
244        // Keep it immediately before DataFusion's default EnsureRequirements.
245        physical_optimizer.rules.insert(8, Arc::new(EnforceSorting));
246        // Add rule for windowed sort
247        physical_optimizer
248            .rules
249            .push(Arc::new(WindowedSortPhysicalRule));
250        // explicitly not do filter pushdown for windowed sort&part sort
251        // (notice that `PartSortExec` create another new dyn filter that need to be pushdown if want to use dyn filter optimization)
252        // benchmark shows it can cause performance regression due to useless filtering and extra shuffle.
253        // We can add a rule to do filter pushdown for windowed sort in the future if we find a way to avoid the performance regression.
254        physical_optimizer
255            .rules
256            .push(Arc::new(MatchesConstantTermOptimizer));
257        physical_optimizer
258            .rules
259            .push(Arc::new(EnsureGlobalLimitForFetch));
260        // Add rule to remove duplicate nodes generated by other rules. Run this in the last.
261        physical_optimizer.rules.push(Arc::new(RemoveDuplicate));
262        // Place SanityCheckPlan at the end of the list to ensure that it runs after all other rules.
263        Self::remove_physical_optimizer_rule(
264            &mut physical_optimizer.rules,
265            SanityCheckPlan {}.name(),
266        );
267        physical_optimizer.rules.push(Arc::new(SanityCheckPlan {}));
268
269        let session_state = SessionStateBuilder::new()
270            .with_config(session_config)
271            .with_runtime_env(runtime_env)
272            .with_default_features()
273            .with_analyzer_rules(analyzer.rules)
274            .with_serializer_registry(Arc::new(DefaultSerializer))
275            .with_query_planner(Arc::new(DfQueryPlanner::new(
276                catalog_list.clone(),
277                partition_rule_manager,
278                region_query_handler.clone(),
279                options.enable_per_region_metrics,
280            )))
281            .with_optimizer_rules(optimizer.rules)
282            .with_physical_optimizer_rules(physical_optimizer.rules)
283            .build();
284        let df_context = SessionContext::new_with_state(session_state);
285        register_function_aliases(&df_context);
286        register_pg_catalog_compat(&df_context);
287
288        Ok(Self {
289            df_context,
290            catalog_manager: catalog_list,
291            dyn_filter_registry_manager: Arc::new(DynFilterRegistryManager::default()),
292            function_state: Arc::new(FunctionState {
293                plugins: plugins.clone(),
294                table_mutation_handler,
295                procedure_service_handler,
296                flow_service_handler,
297            }),
298            aggr_functions: Arc::new(RwLock::new(HashMap::new())),
299            table_functions: Arc::new(RwLock::new(HashMap::new())),
300            extension_rules,
301            plugins,
302            scalar_functions: Arc::new(RwLock::new(HashMap::new())),
303        })
304    }
305
306    fn remove_physical_optimizer_rule(
307        rules: &mut Vec<Arc<dyn PhysicalOptimizerRule + Send + Sync>>,
308        name: &str,
309    ) {
310        rules.retain(|rule| rule.name() != name);
311    }
312
313    /// Optimize the logical plan by the extension analyzer rules.
314    pub fn optimize_by_extension_rules(
315        &self,
316        plan: DfLogicalPlan,
317        context: &QueryEngineContext,
318    ) -> DfResult<DfLogicalPlan> {
319        self.extension_rules
320            .iter()
321            .try_fold(plan, |acc_plan, rule| {
322                rule.analyze(acc_plan, context, self.session_state().config_options())
323            })
324    }
325
326    /// Run the full logical plan optimize phase for the given plan.
327    pub fn optimize_logical_plan(&self, plan: DfLogicalPlan) -> DfResult<DfLogicalPlan> {
328        self.session_state().optimize(&plan)
329    }
330
331    /// Retrieve the scalar function by name
332    pub fn scalar_function(&self, function_name: &str) -> Option<ScalarFunctionFactory> {
333        self.scalar_functions
334            .read()
335            .unwrap()
336            .get(function_name)
337            .cloned()
338    }
339
340    /// Retrieve scalar function names.
341    pub fn scalar_names(&self) -> Vec<String> {
342        self.scalar_functions
343            .read()
344            .unwrap()
345            .keys()
346            .cloned()
347            .collect()
348    }
349
350    /// Retrieve the aggregate function by name
351    pub fn aggr_function(&self, function_name: &str) -> Option<AggregateUDF> {
352        self.aggr_functions
353            .read()
354            .unwrap()
355            .get(function_name)
356            .cloned()
357    }
358
359    /// Retrieve aggregate function names.
360    pub fn aggr_names(&self) -> Vec<String> {
361        self.aggr_functions
362            .read()
363            .unwrap()
364            .keys()
365            .cloned()
366            .collect()
367    }
368
369    /// Retrieve table function by name
370    pub fn table_function(&self, function_name: &str) -> Option<Arc<TableFunction>> {
371        self.table_functions
372            .read()
373            .unwrap()
374            .get(function_name)
375            .cloned()
376    }
377
378    /// Retrieve table function names.
379    pub fn table_function_names(&self) -> Vec<String> {
380        self.table_functions
381            .read()
382            .unwrap()
383            .keys()
384            .cloned()
385            .collect()
386    }
387
388    /// Register an scalar function.
389    /// Will override if the function with same name is already registered.
390    pub fn register_scalar_function(&self, func: ScalarFunctionFactory) {
391        let name = func.name().to_string();
392        let x = self
393            .scalar_functions
394            .write()
395            .unwrap()
396            .insert(name.clone(), func);
397
398        if x.is_some() {
399            warn!("Already registered scalar function '{name}'");
400        }
401    }
402
403    /// Register an aggregate function.
404    ///
405    /// # Panics
406    /// Will panic if the function with same name is already registered.
407    ///
408    /// Panicking consideration: currently the aggregated functions are all statically registered,
409    /// user cannot define their own aggregate functions on the fly. So we can panic here. If that
410    /// invariant is broken in the future, we should return an error instead of panicking.
411    pub fn register_aggr_function(&self, func: AggregateUDF) {
412        let name = func.name().to_string();
413        let x = self
414            .aggr_functions
415            .write()
416            .unwrap()
417            .insert(name.clone(), func);
418        assert!(
419            x.is_none(),
420            "Already registered aggregate function '{name}'"
421        );
422    }
423
424    pub fn register_table_function(&self, func: Arc<TableFunction>) {
425        let name = func.name();
426        let x = self
427            .table_functions
428            .write()
429            .unwrap()
430            .insert(name.to_string(), func.clone());
431
432        if x.is_some() {
433            warn!("Already registered table function '{name}'");
434        }
435    }
436
437    /// Register a window function (UDWF) directly on the DataFusion SessionContext.
438    ///
439    /// This makes the function visible via `session_state.window_functions()`,
440    /// which is used by `DfContextProviderAdapter::get_window_meta`.
441    pub fn register_window_function(&self, func: WindowUDF) {
442        self.df_context.register_udwf(func);
443    }
444
445    pub fn catalog_manager(&self) -> &CatalogManagerRef {
446        &self.catalog_manager
447    }
448
449    pub fn dyn_filter_registry_manager(&self) -> Arc<DynFilterRegistryManager> {
450        self.dyn_filter_registry_manager.clone()
451    }
452
453    pub fn acquire_remote_dyn_filter_registry_lease(
454        &self,
455        query_ctx: &QueryContextRef,
456    ) -> Option<RemoteDynFilterRegistryLease> {
457        let query_id = query_ctx.remote_query_id_value()?;
458        Some(
459            self.dyn_filter_registry_manager
460                .clone()
461                .acquire_lease(query_id),
462        )
463    }
464
465    pub fn function_state(&self) -> Arc<FunctionState> {
466        self.function_state.clone()
467    }
468
469    /// Returns the [`TableMutationHandlerRef`] in state.
470    pub fn table_mutation_handler(&self) -> Option<&TableMutationHandlerRef> {
471        self.function_state.table_mutation_handler.as_ref()
472    }
473
474    /// Returns the [`ProcedureServiceHandlerRef`] in state.
475    pub fn procedure_service_handler(&self) -> Option<&ProcedureServiceHandlerRef> {
476        self.function_state.procedure_service_handler.as_ref()
477    }
478
479    pub(crate) fn disallow_cross_catalog_query(&self) -> bool {
480        self.plugins
481            .map::<QueryOptions, _, _>(|x| x.disallow_cross_catalog_query)
482            .unwrap_or(false)
483    }
484
485    pub fn session_state(&self) -> SessionState {
486        self.df_context.state()
487    }
488
489    /// Create a DataFrame for a table
490    pub fn read_table(&self, table: TableRef) -> DfResult<DataFrame> {
491        self.df_context
492            .read_table(Arc::new(DfTableProviderAdapter::new(table)))
493    }
494}
495
496struct DfQueryPlanner {
497    physical_planner: DefaultPhysicalPlanner,
498}
499
500impl fmt::Debug for DfQueryPlanner {
501    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
502        f.debug_struct("DfQueryPlanner").finish()
503    }
504}
505
506#[async_trait]
507impl QueryPlanner for DfQueryPlanner {
508    async fn create_physical_plan(
509        &self,
510        logical_plan: &DfLogicalPlan,
511        session: &dyn Session,
512    ) -> DfResult<Arc<dyn ExecutionPlan>> {
513        self.physical_planner
514            .create_physical_plan(logical_plan, session)
515            .await
516    }
517}
518
519/// MySQL-compatible scalar function aliases: (target_name, alias)
520const SCALAR_FUNCTION_ALIASES: &[(&str, &str)] = &[
521    ("upper", "ucase"),
522    ("lower", "lcase"),
523    ("ceil", "ceiling"),
524    ("substr", "mid"),
525    ("random", "rand"),
526];
527
528/// MySQL-compatible aggregate function aliases: (target_name, alias)
529const AGGREGATE_FUNCTION_ALIASES: &[(&str, &str)] =
530    &[("stddev_pop", "std"), ("var_pop", "variance")];
531
532/// Register function aliases.
533///
534/// This function adds aliases like `ucase` -> `upper`, `lcase` -> `lower`, etc.
535/// to make GreptimeDB more compatible with MySQL syntax.
536fn register_function_aliases(ctx: &SessionContext) {
537    let state = ctx.state();
538
539    for (target, alias) in SCALAR_FUNCTION_ALIASES {
540        if let Some(func) = state.scalar_functions().get(*target) {
541            let aliased = func.as_ref().clone().with_aliases([*alias]);
542            ctx.register_udf(aliased);
543        }
544    }
545
546    for (target, alias) in AGGREGATE_FUNCTION_ALIASES {
547        if let Some(func) = state.aggregate_functions().get(*target) {
548            let aliased = func.as_ref().clone().with_aliases([*alias]);
549            ctx.register_udaf(aliased);
550        }
551    }
552}
553
554/// Register the session-level (not function-registry) Postgres-compatibility
555/// extension sourced from `datafusion-pg-catalog`: widen Postgres `int4`
556/// bounds to `int8` for `generate_series`/`range` (their stock implementations
557/// require `int8` literals, and DF invokes them at planning time, before any
558/// analyzer rule can run).
559///
560/// The oid-alias type planner is supplied per SQL-planner context by
561/// [`DfContextProviderAdapter::get_type_planner`](crate::datafusion::planner::DfContextProviderAdapter::get_type_planner).
562/// Forward oid-alias name->oid resolution is handled at SQL-parse time by the
563/// `PostgresCompatibilityParser`'s built-in rewrite rule, not here.
564fn register_pg_catalog_compat(ctx: &SessionContext) {
565    datafusion_pg_catalog::pg_catalog::generate_series_arg_coercion::CoerceIntArgsToBigInt::widen(
566        ctx,
567    );
568}
569
570impl DfQueryPlanner {
571    fn new(
572        catalog_manager: CatalogManagerRef,
573        partition_rule_manager: Option<PartitionRuleManagerRef>,
574        region_query_handler: Option<RegionQueryHandlerRef>,
575        enable_per_region_metrics: bool,
576    ) -> Self {
577        let mut planners: Vec<Arc<dyn ExtensionPlanner + Send + Sync>> = vec![
578            Arc::new(PromExtensionPlanner),
579            Arc::new(RangeSelectPlanner),
580            Arc::new(RemoteDynFilterReceiverExtensionPlanner),
581        ];
582        if let (Some(region_query_handler), Some(partition_rule_manager)) =
583            (region_query_handler, partition_rule_manager)
584        {
585            planners.push(Arc::new(DistExtensionPlanner::new(
586                catalog_manager,
587                partition_rule_manager,
588                region_query_handler,
589                enable_per_region_metrics,
590            )));
591            planners.push(Arc::new(MergeSortExtensionPlanner {}));
592        }
593        Self {
594            physical_planner: DefaultPhysicalPlanner::with_extension_planners(planners),
595        }
596    }
597}
598
599/// A wrapper around a memory pool that records metrics.
600///
601/// This wrapper intercepts all memory pool operations and updates
602/// Prometheus metrics for monitoring query memory usage and rejections.
603///
604/// The inner pool is wrapped with `TrackConsumersPool` to preserve
605/// top-consumer error context on rejection.
606#[derive(Debug)]
607pub(super) struct MetricsMemoryPool {
608    inner: Arc<dyn MemoryPool>,
609}
610
611impl MetricsMemoryPool {
612    // Number of top memory consumers to report in OOM error messages
613    const TOP_CONSUMERS_TO_REPORT: usize = 5;
614
615    /// Create a new metrics-wrapped memory pool with the given size limit and
616    /// allocation policy.
617    pub(super) fn new(limit: usize, policy: QueryMemoryPoolPolicy) -> Self {
618        let top_n = NonZeroUsize::new(Self::TOP_CONSUMERS_TO_REPORT).unwrap();
619        let inner: Arc<dyn MemoryPool> = match policy {
620            QueryMemoryPoolPolicy::Greedy => {
621                Arc::new(TrackConsumersPool::new(GreedyMemoryPool::new(limit), top_n))
622            }
623            QueryMemoryPoolPolicy::Fair => {
624                Arc::new(TrackConsumersPool::new(FairSpillPool::new(limit), top_n))
625            }
626        };
627        Self { inner }
628    }
629
630    #[inline]
631    fn update_metrics(&self) {
632        QUERY_MEMORY_POOL_USAGE_BYTES.set(self.inner.reserved() as i64);
633    }
634}
635
636impl fmt::Display for MetricsMemoryPool {
637    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
638        write!(f, "{}(inner_pool: {})", self.name(), self.inner)
639    }
640}
641
642impl MemoryPool for MetricsMemoryPool {
643    fn name(&self) -> &str {
644        "metrics"
645    }
646
647    fn register(&self, consumer: &MemoryConsumer) {
648        self.inner.register(consumer);
649    }
650
651    fn unregister(&self, consumer: &MemoryConsumer) {
652        self.inner.unregister(consumer);
653    }
654
655    fn grow(&self, reservation: &MemoryReservation, additional: usize) {
656        self.inner.grow(reservation, additional);
657        self.update_metrics();
658    }
659
660    fn shrink(&self, reservation: &MemoryReservation, shrink: usize) {
661        self.inner.shrink(reservation, shrink);
662        self.update_metrics();
663    }
664
665    fn try_grow(
666        &self,
667        reservation: &MemoryReservation,
668        additional: usize,
669    ) -> datafusion_common::Result<()> {
670        let result = self.inner.try_grow(reservation, additional);
671        if result.is_err() {
672            QUERY_MEMORY_POOL_REJECTED_TOTAL.inc();
673        }
674        self.update_metrics();
675        result
676    }
677
678    fn reserved(&self) -> usize {
679        self.inner.reserved()
680    }
681
682    fn memory_limit(&self) -> MemoryLimit {
683        self.inner.memory_limit()
684    }
685}
686
687#[cfg(test)]
688mod tests {
689    use std::sync::atomic::{AtomicBool, Ordering};
690
691    use common_base::Plugins;
692    use common_base::memory_limit::MemoryLimit;
693    use common_base::readable_size::ReadableSize;
694    use datafusion::error::DataFusionError;
695    use datafusion::execution::memory_pool::{GreedyMemoryPool, MemoryLimit as DfMemoryLimit};
696    use datafusion::execution::runtime_env::{RuntimeEnv, RuntimeEnvBuilder};
697    use datafusion_common::config::SpillCompression;
698    use session::context::QueryContext;
699
700    use super::*;
701    use crate::options::{QueryOptions, QuerySpillCompression, QuerySpillMode};
702    use crate::query_engine::runtime::{
703        DefaultQueryRuntimeProvider, QueryRuntimeContext, QueryRuntimeProvider,
704        QueryRuntimeProviderRef,
705    };
706
707    fn new_query_engine_state() -> QueryEngineState {
708        new_query_engine_state_with(Plugins::default(), QueryOptions::default())
709    }
710
711    fn new_query_engine_state_with(plugins: Plugins, options: QueryOptions) -> QueryEngineState {
712        QueryEngineState::new(
713            catalog::memory::new_memory_catalog_manager().unwrap(),
714            None,
715            None,
716            None,
717            None,
718            None,
719            false,
720            plugins,
721            options,
722        )
723    }
724
725    struct TestRuntimeProvider {
726        build_called: AtomicBool,
727        configure_called: AtomicBool,
728    }
729
730    impl TestRuntimeProvider {
731        fn new() -> Self {
732            Self {
733                build_called: AtomicBool::new(false),
734                configure_called: AtomicBool::new(false),
735            }
736        }
737    }
738
739    impl QueryRuntimeProvider for TestRuntimeProvider {
740        fn configure_session_config(
741            &self,
742            ctx: QueryRuntimeContext<'_>,
743            config: &mut SessionConfig,
744        ) {
745            assert_eq!(ctx.resolved_memory_pool_size, 1024);
746            self.configure_called.store(true, Ordering::SeqCst);
747            *config = config.clone().with_target_partitions(7);
748        }
749
750        fn build_runtime_env(
751            &self,
752            ctx: QueryRuntimeContext<'_>,
753            builder: RuntimeEnvBuilder,
754        ) -> DfResult<Arc<RuntimeEnv>> {
755            assert_eq!(ctx.resolved_memory_pool_size, 1024);
756            self.build_called.store(true, Ordering::SeqCst);
757            builder
758                .with_memory_pool(Arc::new(GreedyMemoryPool::new(2048)))
759                .build()
760                .map(Arc::new)
761        }
762    }
763
764    struct ErrorRuntimeProvider;
765
766    impl QueryRuntimeProvider for ErrorRuntimeProvider {
767        fn build_runtime_env(
768            &self,
769            _ctx: QueryRuntimeContext<'_>,
770            _builder: RuntimeEnvBuilder,
771        ) -> DfResult<Arc<RuntimeEnv>> {
772            Err(DataFusionError::Execution("runtime provider error".into()))
773        }
774    }
775
776    #[test]
777    fn query_runtime_default_provider_keeps_bounded_memory_pool() {
778        let state = new_query_engine_state_with(
779            Plugins::default(),
780            QueryOptions {
781                memory_pool_size: MemoryLimit::Size(ReadableSize(1024)),
782                ..Default::default()
783            },
784        );
785
786        assert!(matches!(
787            state
788                .session_state()
789                .runtime_env()
790                .memory_pool
791                .memory_limit(),
792            DfMemoryLimit::Finite(1024)
793        ));
794    }
795
796    #[test]
797    fn query_runtime_provider_from_plugins_builds_runtime_env() {
798        let plugins = Plugins::default();
799        let provider = Arc::new(TestRuntimeProvider::new());
800        plugins.insert::<QueryRuntimeProviderRef>(provider.clone());
801
802        let state = new_query_engine_state_with(
803            plugins,
804            QueryOptions {
805                memory_pool_size: MemoryLimit::Size(ReadableSize(1024)),
806                ..Default::default()
807            },
808        );
809
810        assert!(provider.build_called.load(Ordering::SeqCst));
811        assert!(matches!(
812            state
813                .session_state()
814                .runtime_env()
815                .memory_pool
816                .memory_limit(),
817            DfMemoryLimit::Finite(2048)
818        ));
819    }
820
821    #[test]
822    fn query_runtime_provider_from_plugins_configures_session_config() {
823        let plugins = Plugins::default();
824        let provider = Arc::new(TestRuntimeProvider::new());
825        plugins.insert::<QueryRuntimeProviderRef>(provider.clone());
826
827        let state = new_query_engine_state_with(
828            plugins,
829            QueryOptions {
830                memory_pool_size: MemoryLimit::Size(ReadableSize(1024)),
831                experimental_spill_mode: QuerySpillMode::Custom,
832                experimental_spill_compression: QuerySpillCompression::Zstd,
833                ..Default::default()
834            },
835        );
836
837        assert!(provider.configure_called.load(Ordering::SeqCst));
838        assert_eq!(7, state.session_state().config().target_partitions());
839        assert_eq!(
840            SpillCompression::Zstd,
841            state
842                .session_state()
843                .config()
844                .options()
845                .execution
846                .spill_compression
847        );
848    }
849
850    #[test]
851    fn query_runtime_provider_error_is_returned_by_try_new() {
852        let plugins = Plugins::default();
853        plugins.insert::<QueryRuntimeProviderRef>(Arc::new(ErrorRuntimeProvider));
854
855        let err = match QueryEngineState::try_new(
856            catalog::memory::new_memory_catalog_manager().unwrap(),
857            None,
858            None,
859            None,
860            None,
861            None,
862            false,
863            plugins,
864            QueryOptions::default(),
865        ) {
866            Err(err) => err,
867            Ok(_) => panic!("expected runtime provider error"),
868        };
869
870        assert!(
871            matches!(err, DataFusionError::Execution(message) if message == "runtime provider error")
872        );
873    }
874
875    #[test]
876    fn query_engine_state_reuses_query_scoped_dyn_filter_registry_lease() {
877        let state = new_query_engine_state();
878        let query_ctx = QueryContext::arc();
879
880        let first = state
881            .acquire_remote_dyn_filter_registry_lease(&query_ctx)
882            .unwrap();
883        let second = state
884            .acquire_remote_dyn_filter_registry_lease(&query_ctx)
885            .unwrap();
886
887        assert!(first.ptr_eq(&second));
888        assert_eq!(state.dyn_filter_registry_manager().registry_count(), 1);
889        assert_eq!(
890            first.registry().query_id(),
891            query_ctx.remote_query_id_value().unwrap()
892        );
893    }
894
895    #[test]
896    fn query_engine_state_relies_on_query_context_remote_query_id_contract() {
897        let state = new_query_engine_state();
898        let query_ctx = QueryContext::arc();
899
900        assert!(query_ctx.remote_query_id_value().is_some());
901
902        let lease = state
903            .acquire_remote_dyn_filter_registry_lease(&query_ctx)
904            .unwrap();
905
906        assert_eq!(
907            lease.registry().query_id(),
908            query_ctx.remote_query_id_value().unwrap()
909        );
910        assert_eq!(state.dyn_filter_registry_manager().registry_count(), 1);
911    }
912
913    #[test]
914    fn query_engine_state_separates_registries_for_different_query_contexts() {
915        let state = new_query_engine_state();
916        let first_query_ctx = QueryContext::arc();
917        let second_query_ctx = QueryContext::arc();
918
919        let first = state
920            .acquire_remote_dyn_filter_registry_lease(&first_query_ctx)
921            .unwrap();
922        let second = state
923            .acquire_remote_dyn_filter_registry_lease(&second_query_ctx)
924            .unwrap();
925
926        assert!(!first.ptr_eq(&second));
927        assert_eq!(state.dyn_filter_registry_manager().registry_count(), 2);
928        assert_eq!(
929            first.registry().query_id(),
930            first_query_ctx.remote_query_id_value().unwrap()
931        );
932        assert_eq!(
933            second.registry().query_id(),
934            second_query_ctx.remote_query_id_value().unwrap()
935        );
936    }
937
938    /// Builds a runtime env through the default provider seam, mirroring what
939    /// [`QueryEngineState::try_new`] does.
940    fn build_runtime_env(options: &QueryOptions, memory_pool_size: usize) -> Arc<RuntimeEnv> {
941        let ctx = QueryRuntimeContext::new(options, memory_pool_size);
942        let builder = DefaultQueryRuntimeProvider::runtime_env_builder(ctx);
943        DefaultQueryRuntimeProvider
944            .build_runtime_env(ctx, builder)
945            .expect("Failed to build RuntimeEnv")
946    }
947
948    #[test]
949    fn test_build_runtime_env_custom_mode_with_path() {
950        // Use a temp directory managed manually so we don't need the `tempfile` crate.
951        let spill_dir = std::env::temp_dir().join(format!("df_spill_test_{}", std::process::id()));
952        let _ = std::fs::remove_dir_all(&spill_dir);
953        let opts = QueryOptions {
954            experimental_spill_mode: QuerySpillMode::Custom,
955            experimental_spill_path: Some(spill_dir.clone()),
956            experimental_spill_max_temp_directory_size: ReadableSize::gb(1),
957            ..Default::default()
958        };
959        let env = build_runtime_env(&opts, 0);
960
961        assert!(env.disk_manager.tmp_files_enabled());
962        let tmp_file = env.disk_manager.create_tmp_file("test spill");
963        assert!(tmp_file.is_ok());
964        assert!(spill_dir.exists());
965
966        let _ = std::fs::remove_dir_all(&spill_dir);
967    }
968
969    #[test]
970    fn test_build_runtime_env_disabled_mode() {
971        let opts = QueryOptions {
972            experimental_spill_mode: QuerySpillMode::Disabled,
973            ..Default::default()
974        };
975        let env = build_runtime_env(&opts, 0);
976
977        assert!(!env.disk_manager.tmp_files_enabled());
978        let result = env.disk_manager.create_tmp_file("test spill");
979        assert!(result.is_err());
980        if let Err(error) = result {
981            assert!(format!("{error}").contains("DiskManager is disabled"));
982        }
983    }
984
985    #[test]
986    fn test_metrics_memory_pool_policy_differs_with_multiple_spillable_consumers() {
987        for (policy, greedy) in [
988            (QueryMemoryPoolPolicy::Greedy, true),
989            (QueryMemoryPoolPolicy::Fair, false),
990        ] {
991            let pool: Arc<dyn MemoryPool> = Arc::new(MetricsMemoryPool::new(100, policy));
992            let first = MemoryConsumer::new("first")
993                .with_can_spill(true)
994                .register(&pool);
995            let second = MemoryConsumer::new("second")
996                .with_can_spill(true)
997                .register(&pool);
998
999            let first_result = first.try_grow(75);
1000            assert_eq!(first_result.is_ok(), greedy);
1001            if greedy {
1002                assert!(second.try_grow(26).is_err());
1003            } else {
1004                first.try_grow(50).unwrap();
1005                second.try_grow(50).unwrap();
1006            }
1007        }
1008    }
1009
1010    /// Builds a runtime with custom spill configuration and a bounded memory
1011    /// pool, then runs sort queries that probe spill-to-disk behaviour:
1012    ///
1013    /// - A pool large enough for merge chunks but still much smaller
1014    ///   than the total data.  Executes the physical plan directly so we can
1015    ///   walk the plan tree afterwards and sum `spill_count` / `spilled_rows` /
1016    ///   `spilled_bytes` on `SortExec` nodes.  All three must be > 0.
1017    #[tokio::test]
1018    async fn test_sort_spill_smoke_with_custom_runtime() {
1019        use arrow::array::Int32Array;
1020        use arrow::datatypes::{DataType, Field, Schema};
1021        use arrow::record_batch::RecordBatch;
1022        use datafusion::datasource::MemTable;
1023        use datafusion::physical_plan::collect as df_collect;
1024
1025        let spill_dir =
1026            std::env::temp_dir().join(format!("greptime_spill_smoke_{}", std::process::id()));
1027        let _ = std::fs::remove_dir_all(&spill_dir);
1028
1029        let opts = QueryOptions {
1030            experimental_spill_mode: QuerySpillMode::Custom,
1031            experimental_spill_path: Some(spill_dir.clone()),
1032            experimental_spill_max_temp_directory_size: ReadableSize::gb(1),
1033            experimental_memory_pool_policy: QueryMemoryPoolPolicy::Greedy,
1034            ..Default::default()
1035        };
1036
1037        let session_config = SessionConfig::new()
1038            .with_target_partitions(1)
1039            .with_sort_in_place_threshold_bytes(0)
1040            .with_sort_spill_reservation_bytes(64 * 1024);
1041
1042        let schema = Arc::new(Schema::new(vec![
1043            Field::new("id", DataType::Int32, false),
1044            Field::new("val", DataType::Int32, false),
1045        ]));
1046        let n_rows: i32 = 200_000;
1047        let batch_size: i32 = 5_000;
1048        let partitions = 1usize;
1049
1050        let mut table_partitions: Vec<Vec<RecordBatch>> = Vec::with_capacity(partitions);
1051        for _ in 0..partitions {
1052            let mut batches = Vec::new();
1053            let mut row_offset: i32 = 0;
1054            while row_offset < n_rows {
1055                let chunk_end = (row_offset + batch_size).min(n_rows);
1056                let chunk_len = (chunk_end - row_offset) as usize;
1057                let mut id_builder = Int32Array::builder(chunk_len);
1058                let mut val_builder = Int32Array::builder(chunk_len);
1059                for i in row_offset..chunk_end {
1060                    id_builder.append_value(i);
1061                    val_builder.append_value(n_rows - 1 - i);
1062                }
1063                batches.push(
1064                    RecordBatch::try_new(
1065                        Arc::clone(&schema),
1066                        vec![
1067                            Arc::new(id_builder.finish()),
1068                            Arc::new(val_builder.finish()),
1069                        ],
1070                    )
1071                    .unwrap(),
1072                );
1073                row_offset = chunk_end;
1074            }
1075            table_partitions.push(batches);
1076        }
1077
1078        fn build_ctx(
1079            session_config: &SessionConfig,
1080            runtime: &Arc<RuntimeEnv>,
1081            schema: &Arc<Schema>,
1082            partitions: &[Vec<RecordBatch>],
1083        ) -> SessionContext {
1084            let session_state = SessionStateBuilder::new()
1085                .with_config(session_config.clone())
1086                .with_runtime_env(Arc::clone(runtime))
1087                .with_default_features()
1088                .build();
1089            let ctx = SessionContext::new_with_state(session_state);
1090            let table = MemTable::try_new(Arc::clone(schema), partitions.to_vec()).unwrap();
1091            ctx.register_table("t", Arc::new(table)).unwrap();
1092            ctx
1093        }
1094
1095        {
1096            let mem_limit: usize = 512 * 1024; // 512 KB pool, 64 KB reserved for merge
1097
1098            let runtime = build_runtime_env(&opts, mem_limit);
1099            assert_eq!(runtime.memory_pool.reserved(), 0);
1100            let ctx = build_ctx(&session_config, &runtime, &schema, &table_partitions);
1101
1102            let df = ctx
1103                .sql("SELECT val FROM t ORDER BY val ASC")
1104                .await
1105                .expect("planning ORDER BY");
1106            let plan = df
1107                .create_physical_plan()
1108                .await
1109                .expect("creating physical plan");
1110            let task_ctx = ctx.task_ctx();
1111
1112            let batches = df_collect(Arc::clone(&plan), task_ctx)
1113                .await
1114                .expect("executing ORDER BY");
1115
1116            let (spill_count, spilled_rows, spilled_bytes) = sum_sort_spill_metrics(&plan);
1117            assert!(
1118                spill_count > 0,
1119                "expected SortExec spill_count > 0 (pool={} KB, reservation=64 KB), got 0",
1120                mem_limit / 1024,
1121            );
1122            assert!(
1123                spilled_rows > 0,
1124                "expected SortExec spilled_rows > 0, got 0 \
1125                 (spill_count={spill_count}, spilled_bytes={spilled_bytes})",
1126            );
1127            assert!(
1128                spilled_bytes > 0,
1129                "expected SortExec spilled_bytes > 0, got 0 \
1130                 (spill_count={spill_count}, spilled_rows={spilled_rows})",
1131            );
1132
1133            let vals: Vec<i32> = batches
1134                .iter()
1135                .flat_map(|b| {
1136                    let col = b.column(0).as_any().downcast_ref::<Int32Array>().unwrap();
1137                    (0..b.num_rows()).map(move |i| col.value(i))
1138                })
1139                .collect();
1140            assert_eq!(vals.len(), n_rows as usize);
1141            for w in vals.windows(2) {
1142                assert!(w[0] <= w[1], "sort order violation: {} > {}", w[0], w[1]);
1143            }
1144
1145            assert!(
1146                std::fs::read_dir(&spill_dir)
1147                    .ok()
1148                    .map(|mut entries| entries.any(|e| {
1149                        e.as_ref()
1150                            .map(|de| de.file_name().to_string_lossy().starts_with("datafusion-"))
1151                            .unwrap_or(false)
1152                    }))
1153                    .unwrap_or(false),
1154                "Expected 'datafusion-*' directory in spill path {:?} \
1155                 (DiskManager should create it on build).",
1156                spill_dir,
1157            );
1158
1159            assert_eq!(runtime.memory_pool.reserved(), 0);
1160        }
1161
1162        let _ = std::fs::remove_dir_all(&spill_dir);
1163    }
1164
1165    /// Walk a physical plan tree, summing spill metrics for every node whose
1166    /// name starts with `"SortExec"`.
1167    fn sum_sort_spill_metrics(plan: &Arc<dyn ExecutionPlan>) -> (usize, usize, usize) {
1168        let mut spill_count = 0usize;
1169        let mut spilled_rows = 0usize;
1170        let mut spilled_bytes = 0usize;
1171
1172        fn walk(plan: &Arc<dyn ExecutionPlan>, sc: &mut usize, sr: &mut usize, sb: &mut usize) {
1173            let name = plan.name();
1174            if name.starts_with("SortExec")
1175                && let Some(m) = plan.metrics()
1176            {
1177                *sc += m.spill_count().unwrap_or(0);
1178                *sr += m.spilled_rows().unwrap_or(0);
1179                *sb += m.spilled_bytes().unwrap_or(0);
1180            }
1181            for child in plan.children() {
1182                walk(child, sc, sr, sb);
1183            }
1184        }
1185
1186        walk(
1187            plan,
1188            &mut spill_count,
1189            &mut spilled_rows,
1190            &mut spilled_bytes,
1191        );
1192        (spill_count, spilled_rows, spilled_bytes)
1193    }
1194}