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