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