diff --git a/src/aggregation/accessor_helpers.rs b/src/aggregation/accessor_helpers.rs index 7668cba55..2766ce6e1 100644 --- a/src/aggregation/accessor_helpers.rs +++ b/src/aggregation/accessor_helpers.rs @@ -56,23 +56,35 @@ pub(crate) fn get_numeric_or_date_column_types() -> &'static [ColumnType] { ] } +/// Outcome of looking a field name up in the [`ValueSourceRegistry`]. +enum RegisteredSource { + /// The field name is not registered: use the physical fast field. + NotRegistered, + /// The registered source has a type the aggregation cannot consume in this segment. + /// + /// It is treated as if the field had no value in the segment. In particular, it does not + /// fall back to a physical fast field with the same name. + Absent, + Source(Box), +} + fn resolve_registered_source( reader: &SegmentReader, value_sources: &ValueSourceRegistry, field_name: &str, allowed_column_types_opt: Option<&[ColumnType]>, -) -> crate::Result>> { +) -> crate::Result { let Some(provider) = value_sources.get(field_name) else { - return Ok(None); + return Ok(RegisteredSource::NotRegistered); }; - let source = provider.for_segment(reader)?; + let source = provider.for_segment(reader, allowed_column_types_opt)?; let column_type = source.column_type(); if let Some(allowed_column_types) = allowed_column_types_opt { if !allowed_column_types.contains(&column_type) { - return Ok(None); + return Ok(RegisteredSource::Absent); } } - Ok(Some(source)) + Ok(RegisteredSource::Source(source)) } pub(crate) fn get_value_source( @@ -81,23 +93,23 @@ pub(crate) fn get_value_source( field_name: &str, allowed_column_types: Option<&[ColumnType]>, ) -> crate::Result> { - if let Some(registered) = - resolve_registered_source(reader, value_sources, field_name, allowed_column_types)? - { - return Ok(registered); + match resolve_registered_source(reader, value_sources, field_name, allowed_column_types)? { + RegisteredSource::Source(registered) => return Ok(registered), + RegisteredSource::Absent => {} + RegisteredSource::NotRegistered => { + if let Some(source) = + open_physical_sources(reader, field_name, allowed_column_types, true)?.pop() + { + return Ok(source); + } + } } - let ff_fields = reader.fast_fields(); - let (column, column_type) = ff_fields - .u64_lenient_for_type(allowed_column_types, field_name)? - .unwrap_or_else(|| { - ( - Column::build_empty_column(reader.num_docs()), - ColumnType::U64, - ) - }); // The empty-column shim stays physical on purpose: several fast paths check // `as_column()` and would otherwise degrade for a merely absent field. - Ok(Box::new((column, column_type))) + Ok(Box::new(( + Column::build_empty_column(reader.num_docs()), + ColumnType::U64, + ))) } pub(crate) fn get_dynamic_columns( @@ -125,22 +137,60 @@ pub(crate) fn get_all_value_sources( fallback_type: ColumnType, ) -> crate::Result>> { // A registered source shadows the physical type fan-out entirely. - if let Some(registered) = - resolve_registered_source(reader, value_sources, field_name, allowed_column_types)? - { - return Ok(vec![registered]); + let mut sources: Vec> = + match resolve_registered_source(reader, value_sources, field_name, allowed_column_types)? { + RegisteredSource::Source(registered) => return Ok(vec![registered]), + RegisteredSource::Absent => Vec::new(), + RegisteredSource::NotRegistered => { + open_physical_sources(reader, field_name, allowed_column_types, false)? + } + }; + if sources.is_empty() { + sources.push(Box::new(( + Column::build_empty_column(reader.num_docs()), + fallback_type, + ))); } - let ff_fields = reader.fast_fields(); - let mut ff_field_with_type: Vec<(Column, ColumnType)> = - ff_fields.u64_lenient_for_type_all(allowed_column_types, field_name)?; - if ff_field_with_type.is_empty() { - ff_field_with_type.push((Column::build_empty_column(reader.num_docs()), fallback_type)); + Ok(sources) +} + +/// Opens the fast-field columns of `field_name` whose type is allowed, in columnar order. +/// +/// Text columns are opened as `StrColumn`, so that the source carries its dictionary. Other +/// columns use their monotonic `u64` mapping. +/// +/// If `first_only` is true, at most the first allowed column is returned. +fn open_physical_sources( + reader: &SegmentReader, + field_name: &str, + allowed_column_types: Option<&[ColumnType]>, + first_only: bool, +) -> crate::Result>> { + let column_handles: Vec = + reader.fast_fields().dynamic_column_handles(field_name)?; + let mut sources: Vec> = Vec::with_capacity(column_handles.len()); + for handle in column_handles { + let column_type = handle.column_type(); + if let Some(allowed_column_types) = allowed_column_types { + if !allowed_column_types.contains(&column_type) { + continue; + } + } + if column_type == ColumnType::Str { + let DynamicColumn::Str(str_column) = handle.open()? else { + return Err(crate::TantivyError::InternalError(format!( + "the text column of `{field_name}` could not be opened as a text column" + ))); + }; + sources.push(Box::new(str_column)); + } else if let Some(column) = handle.open_u64_lenient()? { + sources.push(Box::new((column, column_type))); + } else { + continue; + } + if first_only { + break; + } } - Ok(ff_field_with_type - .into_iter() - .map(|(column, column_type)| { - let source: Box = Box::new((column, column_type)); - source - }) - .collect()) + Ok(sources) } diff --git a/src/aggregation/agg_data.rs b/src/aggregation/agg_data.rs index 6dae08c26..79fd020dc 100644 --- a/src/aggregation/agg_data.rs +++ b/src/aggregation/agg_data.rs @@ -48,6 +48,21 @@ impl AggregationsSegmentCtx { column_block_accessor: ColumnBlockAccessor::default(), } } + + /// Charges the memory grown by the value sources while loading blocks to the aggregation + /// limits. + /// + /// Must be called after each call driving the collection (collect, flush). + pub(crate) fn charge_value_source_memory(&mut self) -> crate::Result<()> { + let pending_memory_consumption = + self.column_block_accessor.take_pending_memory_consumption(); + if pending_memory_consumption == 0 { + return Ok(()); + } + self.context + .limits + .add_memory_consumed(pending_memory_consumption as u64) + } } /// A node of the per-segment aggregation request tree. @@ -316,6 +331,23 @@ fn require_physical_column( }) } +/// Rejects a field name shadowed by a registered value source. +/// +/// For aggregations that read physical columns and dictionaries directly. Without this check, +/// they would silently read the physical field with the same name. +fn reject_registered_source( + value_sources: &ValueSourceRegistry, + field_name: &str, + agg_kind: &str, +) -> crate::Result<()> { + if value_sources.get(field_name).is_some() { + return Err(crate::TantivyError::InvalidArgument(format!( + "{agg_kind} does not support the computed value source `{field_name}`" + ))); + } + Ok(()) +} + fn build_nodes( agg_name: &str, req: &Aggregation, @@ -565,6 +597,7 @@ fn build_composite_node( ) -> crate::Result { let mut composite_accessors = Vec::with_capacity(req.sources.len()); for source in &req.sources { + reject_registered_source(&context.value_sources, source.field(), "composite")?; let source_after_key_opt = req.after.get(source.name()).map(|k| &k.0); let source_accessor = CompositeSourceAccessors::build_for_source(reader, source, source_after_key_opt)?; @@ -598,6 +631,7 @@ fn build_multi_terms_nodes( let mut accessors_by_field = Vec::with_capacity(req.terms.len()); for field_def in &req.terms { let field_name = &field_def.field; + reject_registered_source(value_sources, field_name, "multi_terms")?; let str_dict_column = reader.fast_fields().str(field_name)?; // multi_terms resolves missing values through `ColumnIndex::has_value` per document, and // exposes its columns on a public struct, so it stays physical-only. @@ -891,7 +925,6 @@ fn build_terms_or_cardinality_nodes( ) -> crate::Result> { let mut nodes = Vec::new(); - let str_dict_column = reader.fast_fields().str(field_name)?; let value_sources = &context.value_sources; let include_bytes = matches!(req, TermsOrCardinalityRequest::Terms(_)); @@ -971,19 +1004,24 @@ fn build_terms_or_cardinality_nodes( // When excluding, the behavior could be to include non-string values continue; } - let str_col = str_dict_column - .as_ref() - .expect("str_dict_column must exist for string column"); - allowed_term_ids = build_allowed_term_ids_for_str( - str_col, - &req.include, - &req.exclude, - missing.is_some(), - )?; + if let Some(str_col) = accessor.as_str_column() { + allowed_term_ids = build_allowed_term_ids_for_str( + str_col, + &req.include, + &req.exclude, + missing.is_some(), + )?; + } else if accessor.term_dictionary().is_some() { + // Filters are resolved by searching the sstable dictionary. + return Err(crate::TantivyError::InvalidArgument(format!( + "terms aggregation with `include` / `exclude` requires a physical \ + text field, but `{field_name}` is a computed value source" + ))); + } + // Otherwise, the source has no value: there is nothing to filter. }; AggNodeData::Terms(TermsAggReqData { accessor, - str_dict_column: str_dict_column.clone(), missing_value_for_accessor, name: agg_name.to_string(), req: TermsAggregationInternal::from_req(req), @@ -993,19 +1031,8 @@ fn build_terms_or_cardinality_nodes( }) } TermsOrCardinalityRequest::Cardinality(ref req) => { - // `str_dict_column` is computed once per field; for JSON paths - // with mixed types it's `Some` even on the numeric req_data. - // Cardinality only consults it for the str column path, so - // gate by column_type to avoid driving non-str collectors - // through the coupon-cache path. - let str_dict_column_for_req = if column_type == ColumnType::Str { - str_dict_column.clone() - } else { - None - }; AggNodeData::Cardinality(CardinalityAggReqData { accessor, - str_dict_column: str_dict_column_for_req, missing_value_for_accessor, name: agg_name.to_string(), req: req.clone(), diff --git a/src/aggregation/bucket/term_agg/mod.rs b/src/aggregation/bucket/term_agg/mod.rs index 5ebb75d35..dd89879e7 100644 --- a/src/aggregation/bucket/term_agg/mod.rs +++ b/src/aggregation/bucket/term_agg/mod.rs @@ -4,8 +4,7 @@ use std::net::Ipv6Addr; use columnar::column_values::CompactSpaceU64Accessor; use columnar::{ - ColumnType, Dictionary, MonotonicallyMappableToU128, MonotonicallyMappableToU64, - NumericalValue, StrColumn, + ColumnType, Dictionary, MonotonicallyMappableToU128, MonotonicallyMappableToU64, NumericalValue, }; use common::{BitSet, TinySet}; use rustc_hash::FxHashMap; @@ -26,7 +25,7 @@ use crate::aggregation::intermediate_agg_result::{ IntermediateKey, IntermediateTermBucketEntry, IntermediateTermBucketResult, }; use crate::aggregation::segment_agg_result::{BucketIdProvider, SegmentAggregationCollector}; -use crate::aggregation::{format_date, BucketId, Key, ValueSource}; +use crate::aggregation::{format_date, BucketId, Key, TermOrdDictionary, ValueSource}; use crate::error::DataCorruption; use crate::TantivyError; @@ -38,8 +37,6 @@ mod flattened_term_histogram; pub(crate) struct TermsAggReqData { /// The column accessor to access the fast field values. pub(crate) accessor: Box, - /// The string dictionary column if the field is of type text. - pub(crate) str_dict_column: Option, /// The missing value as u64 value. pub(crate) missing_value_for_accessor: Option, /// Used to build the correct nested result when we have an empty result. @@ -1333,12 +1330,12 @@ where let column_type = term_req.accessor.column_type(); if column_type == ColumnType::Str { + // A text source without a dictionary has no value (e.g. the empty-column shim). let fallback_dict = Dictionary::empty(); - let term_dict = term_req - .str_dict_column - .as_ref() - .map(|el| el.dictionary()) - .unwrap_or_else(|| &fallback_dict); + let term_dict: &dyn TermOrdDictionary = term_req + .accessor + .term_dictionary() + .unwrap_or(&fallback_dict); // Collect into a map to dedup by key, then flush into `out`. Two cases need it: a real // term may equal the `missing` placeholder, and the min_doc_count==0 fill must skip @@ -1377,7 +1374,7 @@ where let mut intermediate_entry_it = intermediate_entries.into_iter(); - term_dict.sorted_ords_to_term_cb(&term_ids[..], |term| { + term_dict.sorted_ords_to_term_cb(&term_ids[..], &mut |term| { let intermediate_entry = intermediate_entry_it.next().unwrap(); dict.insert( IntermediateKey::Str( @@ -1387,9 +1384,14 @@ where ); })?; - if term_req.req.min_doc_count == 0 { + // Only physical text fields have a dictionary of all the segment's terms to stream. + // `min_doc_count == 0` is rejected for computed text sources at build time. + if let (0, Some(str_column)) = ( + term_req.req.min_doc_count, + term_req.accessor.as_str_column(), + ) { // TODO: Handle rev streaming for descending sorting by keys - let mut stream = term_dict.stream()?; + let mut stream = str_column.dictionary().stream()?; let empty_sub_aggregation = IntermediateAggregationResults::empty_from_req(&term_req.sub_aggregations); while stream.advance() { diff --git a/src/aggregation/collector.rs b/src/aggregation/collector.rs index b3bdef927..2f17a3cb6 100644 --- a/src/aggregation/collector.rs +++ b/src/aggregation/collector.rs @@ -174,14 +174,12 @@ impl SegmentCollector for AggregationSegmentCollector { return; } self.agg_collector.push(0, doc); - match self + let result = self .agg_collector .check_flush_local(&mut self.aggs_with_accessor) - { - Ok(_) => {} - Err(e) => { - self.error = Some(e); - } + .and_then(|_| self.aggs_with_accessor.charge_value_source_memory()); + if let Err(e) = result { + self.error = Some(e); } } fn collect_block(&mut self, docs: &[DocId]) { @@ -189,15 +187,13 @@ impl SegmentCollector for AggregationSegmentCollector { return; } - match self.agg_collector.get_sub_agg_collector().collect( - 0, - docs, - &mut self.aggs_with_accessor, - ) { - Ok(_) => {} - Err(e) => { - self.error = Some(e); - } + let result = self + .agg_collector + .get_sub_agg_collector() + .collect(0, docs, &mut self.aggs_with_accessor) + .and_then(|_| self.aggs_with_accessor.charge_value_source_memory()); + if let Err(e) = result { + self.error = Some(e); } } @@ -206,6 +202,7 @@ impl SegmentCollector for AggregationSegmentCollector { return Err(err); } self.agg_collector.flush(&mut self.aggs_with_accessor)?; + self.aggs_with_accessor.charge_value_source_memory()?; let mut sub_aggregation_res = IntermediateAggregationResults::default(); self.agg_collector diff --git a/src/aggregation/metric/cardinality/mod.rs b/src/aggregation/metric/cardinality/mod.rs index c164ad388..f137cb7e7 100644 --- a/src/aggregation/metric/cardinality/mod.rs +++ b/src/aggregation/metric/cardinality/mod.rs @@ -17,7 +17,7 @@ mod term_ord_accumulator; use std::hash::Hash; -use columnar::{ColumnType, StrColumn}; +use columnar::ColumnType; use common::BitSet; use datasketches::hll::{Coupon, HllSketch, HllType, HllUnion}; pub(crate) use numeric_collector::SegmentNumericCardinalityCollector; @@ -101,8 +101,6 @@ pub struct CardinalityAggregationReq { pub(crate) struct CardinalityAggReqData { /// The column accessor to access the fast field values. pub(crate) accessor: Box, - /// The string dictionary column if the field is of type string. - pub(crate) str_dict_column: Option, /// The missing value normalized to the internal u64 representation of the field type. pub(crate) missing_value_for_accessor: Option, /// The name of the aggregation. diff --git a/src/aggregation/metric/cardinality/str_collector.rs b/src/aggregation/metric/cardinality/str_collector.rs index 4bfa61d6c..f4e15e7a2 100644 --- a/src/aggregation/metric/cardinality/str_collector.rs +++ b/src/aggregation/metric/cardinality/str_collector.rs @@ -11,7 +11,7 @@ use std::fmt::Debug; use std::io; -use columnar::{ColumnType, Dictionary}; +use columnar::ColumnType; use datasketches::hll::Coupon; use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet}; @@ -22,7 +22,7 @@ use crate::aggregation::intermediate_agg_result::{ IntermediateAggregationResult, IntermediateAggregationResults, IntermediateMetricResult, }; use crate::aggregation::segment_agg_result::SegmentAggregationCollector; -use crate::aggregation::*; +use crate::aggregation::{TermOrdDictionary, *}; /// A CouponCache is here to cache the mapping term ordinal -> coupon (see above). /// The idea is that we do not want to fetch terms associated to several term ordinals, @@ -122,7 +122,7 @@ impl Debug for SegmentStrCardinalityCollector { /// Returns a mapping from term_ord to the hash (coupon) of the associated term. fn build_coupon_cache( buckets: &[Option], - dictionary: &Dictionary, + dictionary: &dyn TermOrdDictionary, missing_value_opt: Option<&Key>, ) -> io::Result { // Pass 1 computes the capacity hint, pass 2 inserts. @@ -137,11 +137,11 @@ fn build_coupon_cache( let mut term_ords: Vec = term_ords_set.into_iter().collect(); term_ords.sort_unstable(); - term_ords.pop_if(|highest_term_ord| *highest_term_ord >= dictionary.num_terms() as u64); + term_ords.pop_if(|highest_term_ord| *highest_term_ord >= dictionary.num_terms()); let mut coupons: Vec = Vec::with_capacity(term_ords.len()); let all_term_ords_found: bool = - dictionary.sorted_ords_to_term_cb(&term_ords, |term_bytes| { + dictionary.sorted_ords_to_term_cb(&term_ords, &mut |term_bytes| { let coupon: Coupon = Coupon::from_value(term_bytes); coupons.push(coupon); })?; @@ -225,9 +225,9 @@ impl SegmentAggregationCollector ) -> crate::Result<()> { self.prepare_max_bucket(bucket_id, agg_data)?; let req_data = &self.req_data; - let Some(str_dict_column) = &req_data.str_dict_column else { + let Some(term_dictionary) = req_data.accessor.term_dictionary() else { return Err(crate::TantivyError::InternalError( - "a str cardinality collector requires a str dictionary column".to_string(), + "a str cardinality collector requires a term dictionary".to_string(), )); }; // Strings are dictionary encoded. Fetching the terms associated to strings @@ -239,7 +239,7 @@ impl SegmentAggregationCollector if self.coupon_cache.is_none() { self.coupon_cache = Some(build_coupon_cache( &self.buckets, - str_dict_column.dictionary(), + term_dictionary, req_data.req.missing.as_ref(), )?); } diff --git a/src/aggregation/mod.rs b/src/aggregation/mod.rs index 81b03551c..bf0a9004c 100644 --- a/src/aggregation/mod.rs +++ b/src/aggregation/mod.rs @@ -147,7 +147,7 @@ use std::fmt::Display; use std::sync::Arc; pub(crate) use value_source::ColumnBlockAccessor; -pub use value_source::{ValueSource, ValueSourceProvider, ValueSourceRegistry}; +pub use value_source::{TermOrdDictionary, ValueSource, ValueSourceProvider, ValueSourceRegistry}; #[cfg(test)] mod agg_tests; diff --git a/src/aggregation/value_source/block_accessor.rs b/src/aggregation/value_source/block_accessor.rs index c9d57c95b..9cfce0569 100644 --- a/src/aggregation/value_source/block_accessor.rs +++ b/src/aggregation/value_source/block_accessor.rs @@ -38,6 +38,8 @@ pub(crate) struct ColumnBlockAccessor { cardinality: Cardinality, /// Whether any document has multiple values in the loaded block. multivalued: bool, + /// Memory grown by the sources while loading blocks. + pending_memory_consumption: usize, } impl ColumnBlockAccessor { @@ -48,12 +50,15 @@ impl ColumnBlockAccessor { /// inflate document counts. #[inline] fn fetch_block(&mut self, docs: &[DocId], source: &mut S) { + let memory_before = source.memory_consumption(); self.cardinality = source.load_block( docs, &mut self.val_cache, &mut self.docid_cache, &mut self.row_id_cache, ); + self.pending_memory_consumption += + source.memory_consumption().saturating_sub(memory_before); // Full/Optional cannot repeat documents; Full may leave docid_cache stale. self.multivalued = self.cardinality.is_multivalue() && self.docid_cache.windows(2).any(|pair| pair[0] == pair[1]); @@ -63,6 +68,14 @@ impl ColumnBlockAccessor { ); } + /// Returns the memory grown by the sources since the last call, and resets it. + /// + /// `load_block` cannot fail, so the growth is accumulated here and must be charged to the + /// aggregation limits by the driver of the collection. + pub(crate) fn take_pending_memory_consumption(&mut self) -> usize { + std::mem::replace(&mut self.pending_memory_consumption, 0) + } + /// Fetches a block from a column known to be full (hence we pass the ColumnValue Object /// directly). /// diff --git a/src/aggregation/value_source/mod.rs b/src/aggregation/value_source/mod.rs index 8b937e684..3341c43a9 100644 --- a/src/aggregation/value_source/mod.rs +++ b/src/aggregation/value_source/mod.rs @@ -5,9 +5,10 @@ mod value_source_registry; pub(crate) mod tests; use std::borrow::Borrow; +use std::io; pub(crate) use block_accessor::ColumnBlockAccessor; -use columnar::{Cardinality, Column, ColumnType, ColumnValues, RowId}; +use columnar::{Cardinality, Column, ColumnType, ColumnValues, Dictionary, RowId, StrColumn}; pub use value_source_registry::{ValueSourceProvider, ValueSourceRegistry}; use crate::DocId; @@ -56,6 +57,75 @@ pub trait ValueSource: std::fmt::Debug { let column = self.as_column()?; Some((column.min_value(), column.max_value())) } + + /// Returns the physical text column, if this source is backed by one. + /// + /// For the paths that need the full sstable dictionary (regex search, streaming all of the + /// terms) rather than resolving ords. `Some` implies that `term_dictionary()` is `Some` and + /// that its ords are sorted with the terms. + fn as_str_column(&self) -> Option<&StrColumn> { + None + } + + /// Dictionary resolving the term ords returned by `load_block`, for `Str` sources. + /// + /// Every `Str` source with values must return its dictionary. `None` is only acceptable for + /// a source without any value (e.g. the empty-column shim of an absent field). + /// + /// Contract: the ords returned by earlier `load_block` calls remain valid. The dictionary + /// may grow as more blocks are loaded. + fn term_dictionary(&self) -> Option<&dyn TermOrdDictionary> { + None + } + + /// Heap memory owned by the source, in bytes. + /// + /// Sources can grow while loading blocks (e.g. a dictionary built on the fly). The growth + /// observed across `load_block` calls is charged to the aggregation memory limits. + fn memory_consumption(&self) -> usize { + 0 + } +} + +/// Resolves the term ords of a `Str` [`ValueSource`] back into terms. +pub trait TermOrdDictionary { + /// Returns true if the order of the ords matches the lexicographic order of the terms. + /// + /// Aggregations relying on that property (e.g. a terms aggregation ordered by `_key`) must + /// check it. + fn ords_sorted_with_terms(&self) -> bool; + + /// Number of terms in the dictionary. Valid ords are `0..num_terms()`. + fn num_terms(&self) -> u64; + + /// Calls `callback` with the term associated with each ord, in the order of `sorted_ords`. + /// + /// Precondition: `sorted_ords` is sorted in ascending order. + /// + /// Returns false if an ord was not found in the dictionary. + fn sorted_ords_to_term_cb( + &self, + sorted_ords: &[u64], + callback: &mut dyn FnMut(&[u8]), + ) -> io::Result; +} + +impl TermOrdDictionary for Dictionary { + fn ords_sorted_with_terms(&self) -> bool { + true + } + + fn num_terms(&self) -> u64 { + Dictionary::num_terms(self) as u64 + } + + fn sorted_ords_to_term_cb( + &self, + sorted_ords: &[u64], + callback: &mut dyn FnMut(&[u8]), + ) -> io::Result { + Dictionary::sorted_ords_to_term_cb(self, sorted_ords, callback) + } } // Lenient columns have erased their logical type; the tuple retains it alongside the values. @@ -93,6 +163,38 @@ impl> + std::fmt::Debug> ValueSource for (ColumnRe } } +// A physical text column: the values are the term ords of its dictionary. +impl ValueSource for StrColumn { + #[inline] + fn column_type(&self) -> ColumnType { + ColumnType::Str + } + + #[inline] + fn load_block( + &mut self, + docs: &[DocId], + values: &mut Vec, + docids: &mut Vec, + row_ids: &mut Vec, + ) -> Cardinality { + (self.ords(), ColumnType::Str).load_block(docs, values, docids, row_ids) + } + + #[inline] + fn as_column(&self) -> Option<&Column> { + Some(self.ords()) + } + + fn as_str_column(&self) -> Option<&StrColumn> { + Some(self) + } + + fn term_dictionary(&self) -> Option<&dyn TermOrdDictionary> { + Some(self.dictionary()) + } +} + /// `docs` has to be sorted ascending and free of duplicates. #[inline] fn load_full_column_values( diff --git a/src/aggregation/value_source/tests.rs b/src/aggregation/value_source/tests.rs index 63589e138..4b4921993 100644 --- a/src/aggregation/value_source/tests.rs +++ b/src/aggregation/value_source/tests.rs @@ -29,7 +29,11 @@ impl ValueSource for Constant { pub(crate) struct ConstantProvider(pub u64); impl ValueSourceProvider for ConstantProvider { - fn for_segment(&self, _reader: &SegmentReader) -> crate::Result> { + fn for_segment( + &self, + _reader: &SegmentReader, + _allowed_column_types: Option<&[ColumnType]>, + ) -> crate::Result> { Ok(Box::new(Constant(self.0))) } } @@ -106,6 +110,169 @@ fn test_registered_source_as_sub_aggregation_of_terms() { assert_eq!(buckets[1]["s"]["value"], 1.0); } +/// A constant source of an arbitrary type, whose memory grows by `memory_growth_per_block` bytes +/// on each loaded block. +#[derive(Debug)] +struct TypedConstant { + column_type: ColumnType, + value: u64, + memory_growth_per_block: usize, + memory_consumption: usize, +} + +impl ValueSource for TypedConstant { + fn column_type(&self) -> ColumnType { + self.column_type + } + + fn load_block( + &mut self, + docs: &[DocId], + values: &mut Vec, + _docids: &mut Vec, + _row_ids: &mut Vec, + ) -> Cardinality { + self.memory_consumption += self.memory_growth_per_block; + values.clear(); + values.resize(docs.len(), self.value); + Cardinality::Full + } + + fn memory_consumption(&self) -> usize { + self.memory_consumption + } +} + +struct TypedConstantProvider { + column_type: ColumnType, + value: u64, + memory_growth_per_block: usize, +} + +impl ValueSourceProvider for TypedConstantProvider { + fn for_segment( + &self, + _reader: &SegmentReader, + _allowed_column_types: Option<&[ColumnType]>, + ) -> crate::Result> { + Ok(Box::new(TypedConstant { + column_type: self.column_type, + value: self.value, + memory_growth_per_block: self.memory_growth_per_block, + memory_consumption: 0, + })) + } +} + +fn registry_with(name: &str, provider: TypedConstantProvider) -> ValueSourceRegistry { + let mut registry = ValueSourceRegistry::default(); + registry.register(name, Arc::new(provider)); + registry +} + +fn try_run_agg( + index: &crate::Index, + aggs: serde_json::Value, + context: crate::aggregation::AggContextParams, +) -> crate::Result { + use crate::aggregation::agg_req::Aggregations; + use crate::aggregation::AggregationCollector; + use crate::query::AllQuery; + + let aggs: Aggregations = serde_json::from_value(aggs).unwrap(); + let collector = AggregationCollector::from_aggs(aggs, context); + let searcher = index.reader().unwrap().searcher(); + let result = searcher.search(&AllQuery, &collector)?; + Ok(serde_json::to_value(result).unwrap()) +} + +#[test] +fn test_registered_source_with_disallowed_type_does_not_fall_back_to_physical_field() { + let index = index_with_scores(&[10, 20, 30, 40]); + // `score` is shadowed by a bool source, which stats cannot consume. + let registry = registry_with( + "score", + TypedConstantProvider { + column_type: ColumnType::Bool, + value: 1, + memory_growth_per_block: 0, + }, + ); + let result = run_agg_with_registry( + &index, + serde_json::json!({ "s": { "stats": { "field": "score" } } }), + registry, + ); + // The field is treated as absent, rather than read from the physical `score` column. + assert_eq!(result["s"]["count"], 0); +} + +#[test] +fn test_composite_and_multi_terms_reject_registered_source() { + let index = index_with_scores(&[10, 20]); + let context = crate::aggregation::AggContextParams::default().with_value_sources(Arc::new( + registry_with( + "score", + TypedConstantProvider { + column_type: ColumnType::U64, + value: 1, + memory_growth_per_block: 0, + }, + ), + )); + let composite = serde_json::json!({ + "c": { + "composite": { "size": 10, "sources": [{ "s": { "terms": { "field": "score" } } }] } + } + }); + let err = try_run_agg(&index, composite, context.clone()).unwrap_err(); + assert!( + err.to_string().contains("composite does not support"), + "{err}" + ); + let multi_terms = serde_json::json!({ + "m": { "multi_terms": { "terms": [{ "field": "score" }] } } + }); + let err = try_run_agg(&index, multi_terms, context).unwrap_err(); + assert!( + err.to_string().contains("multi_terms does not support"), + "{err}" + ); +} + +#[test] +fn test_value_source_memory_growth_is_charged_to_limits() { + use crate::aggregation::AggregationLimitsGuard; + use crate::tokenizer::TokenizerManager; + + let index = index_with_scores(&[10, 20, 30, 40]); + let aggs = serde_json::json!({ "s": { "sum": { "field": "computed" } } }); + let context_with_growth = |memory_growth_per_block: usize| { + let limits = AggregationLimitsGuard::new(Some(1_000_000), None); + crate::aggregation::AggContextParams::new(limits, TokenizerManager::default()) + .with_value_sources(Arc::new(registry_with( + "computed", + TypedConstantProvider { + column_type: ColumnType::U64, + value: 1, + memory_growth_per_block, + }, + ))) + }; + let result = try_run_agg(&index, aggs.clone(), context_with_growth(0)).unwrap(); + assert_eq!(result["s"]["value"], 4.0); + let err = try_run_agg(&index, aggs, context_with_growth(10_000_000)).unwrap_err(); + assert!( + matches!( + err, + crate::TantivyError::AggregationError( + crate::aggregation::AggregationError::MemoryExceeded { .. } + ) + ), + "{err}" + ); +} + #[test] fn test_is_contiguous() { assert!(!is_contiguous(&[])); diff --git a/src/aggregation/value_source/value_source_registry.rs b/src/aggregation/value_source/value_source_registry.rs index 89d6cde88..b280480eb 100644 --- a/src/aggregation/value_source/value_source_registry.rs +++ b/src/aggregation/value_source/value_source_registry.rs @@ -3,13 +3,27 @@ use std::collections::HashMap; use std::sync::Arc; +use columnar::ColumnType; + use super::ValueSource; use crate::SegmentReader; /// Creates a value source for each segment. pub trait ValueSourceProvider: Send + Sync + 'static { /// Binds this definition to a single segment. - fn for_segment(&self, reader: &SegmentReader) -> crate::Result>; + /// + /// `allowed_column_types` is the set of column types the aggregation can consume, if it is + /// restricted. Providers can use it to steer type resolution. + /// + /// - A definition that can never produce one of these types (independently of the segment) + /// should return an error. + /// - A source whose type is not allowed for this specific segment is accepted, and treated by + /// the aggregation as if the field had no value in this segment. + fn for_segment( + &self, + reader: &SegmentReader, + allowed_column_types: Option<&[ColumnType]>, + ) -> crate::Result>; } /// Named computed sources available to an aggregation request. @@ -59,7 +73,7 @@ mod tests { let searcher = index.reader().unwrap().searcher(); let value_source_provider = registry.get("computed").unwrap(); let mut value_source = value_source_provider - .for_segment(searcher.segment_reader(0u32)) + .for_segment(searcher.segment_reader(0u32), None) .unwrap(); let mut values = Vec::new(); let mut doc_ids = Vec::new();