1use super::query_engine::{
4 count_to_f64, execute_query_with_candidates, normalize_scores, parse_optional_filter,
5 score_to_f32, sort_docs,
6};
7use super::{Collection, CollectionSnapshot};
8use crate::config::IoBackend;
9use crate::doc::Doc;
10use crate::error::{Error, Result};
11use crate::iterator::DocIterator;
12use crate::multi_query::{MultiQuery, RerankMethod};
13use crate::query::{GroupBySearchQuery, SearchQuery};
14use crate::stats::{IndexUsage, QueryKind, QueryObservation};
15use serde_json::Value;
16use std::cmp::Ordering;
17use std::collections::{BTreeMap, HashMap, HashSet};
18
19#[derive(Debug, Default)]
20struct MultiQueryTelemetry {
21 used_ann: bool,
22 diskann_io_backend: Option<IoBackend>,
23 diskann_sector_reads: u64,
24 used_scalar: bool,
25 used_fts_index: bool,
26 candidates: u64,
27}
28
29impl Collection {
30 pub fn query(&self, query: &SearchQuery) -> Result<Vec<Doc>> {
31 self.ensure_open()?;
32 let snapshot = self.snapshot_state()?;
33 let filter = parse_optional_filter(query.filter.as_deref())?;
34 let plan = snapshot.indexes.plan_candidates(
35 &snapshot.docs,
36 snapshot.revision,
37 query,
38 filter.as_ref(),
39 )?;
40 let candidates = plan.candidate_count(snapshot.docs.len());
41 enforce_query_candidates(&snapshot, candidates)?;
42 let result = execute_query_with_candidates(
43 &snapshot.schema,
44 &snapshot.docs,
45 &snapshot.indexes,
46 query,
47 plan.selection.as_ref(),
48 plan.fts_scores.as_ref(),
49 filter.as_ref(),
50 )?;
51 let has_fts = query.fts.is_some();
52 snapshot.stats.record_query(QueryObservation {
53 kind: if has_fts {
54 QueryKind::Fts
55 } else if plan.used_ann {
56 QueryKind::Ann
57 } else {
58 QueryKind::Exact
59 },
60 diskann_io_backend: plan.diskann_io_backend,
61 diskann_sector_reads: plan.diskann_sector_reads,
62 filtered: query
63 .filter
64 .as_ref()
65 .is_some_and(|value| !value.trim().is_empty()),
66 index_usage: IndexUsage::new(plan.used_scalar, plan.used_fts_index),
67 radius: query.params.get("radius").is_some(),
68 candidates,
69 });
70 Ok(result)
71 }
72
73 pub fn multi_query(&self, query: &MultiQuery) -> Result<Vec<Doc>> {
74 self.ensure_open()?;
75 if query.queries.is_empty() {
76 return Err(Error::invalid_argument(
77 "multi-query must contain at least one sub-query",
78 ));
79 }
80 let snapshot = self.snapshot_state()?;
81 let (branches, telemetry) = execute_multi_query_branches(&snapshot, query)?;
82 let output = fuse_multi_query_branches(query, branches)?;
83 let has_fts = query.queries.iter().any(|branch| branch.fts.is_some());
84 snapshot.stats.record_query(QueryObservation {
85 kind: match (telemetry.used_ann, has_fts) {
86 (true, true) => QueryKind::AnnFts,
87 (true, false) => QueryKind::Ann,
88 (false, true) => QueryKind::Fts,
89 (false, false) => QueryKind::Exact,
90 },
91 diskann_io_backend: telemetry.diskann_io_backend,
92 diskann_sector_reads: telemetry.diskann_sector_reads,
93 filtered: query.filter.is_some(),
94 index_usage: IndexUsage::new(telemetry.used_scalar, telemetry.used_fts_index),
95 radius: false,
96 candidates: telemetry.candidates,
97 });
98 Ok(output)
99 }
100
101 pub fn group_by(&self, query: &GroupBySearchQuery) -> Result<HashMap<String, Vec<Doc>>> {
102 self.ensure_open()?;
103 let route_count =
104 usize::from(!query.vector.is_empty()) + usize::from(query.binary_vector.is_some());
105 if route_count != 1 {
106 return Err(Error::invalid_argument(
107 "group-by query must select exactly one dense or binary route",
108 ));
109 }
110 let candidate_limit = query.group_count.saturating_mul(query.group_topk).max(1);
111 let candidate_limit = i32::try_from(candidate_limit)
112 .map_err(|_| Error::resource_exhausted("group-by candidate limit exceeds i32"))?;
113 let mut vector_query = if let Some(vector) = &query.binary_vector {
114 SearchQuery::binary(&query.field_name, vector, candidate_limit)?
115 } else {
116 SearchQuery::new(&query.field_name, &query.vector, candidate_limit)?
117 };
118 vector_query.include_vector = query.include_vector;
119 vector_query.output_fields.clone_from(&query.output_fields);
120 vector_query.params.clone_from(&query.params);
121 if let Some(filter) = &query.filter {
122 vector_query.set_filter(filter)?;
123 }
124 let docs = self.query(&vector_query)?;
125 let mut groups: HashMap<String, Vec<Doc>> = HashMap::new();
126 for doc in docs {
127 let key = doc.scalar_json(&query.group_by_field).map_or_else(
128 || "__null__".to_string(),
129 |value| match value {
130 Value::String(value) => value,
131 other => other.to_string(),
132 },
133 );
134 let group = groups.entry(key).or_default();
135 if group.len() < query.group_topk as usize {
136 group.push(doc);
137 }
138 }
139 if groups.len() > query.group_count as usize {
140 let mut ranked: Vec<(String, f32)> = groups
141 .iter()
142 .map(|(key, values)| {
143 (
144 key.clone(),
145 values.first().map_or(f32::NEG_INFINITY, Doc::get_score),
146 )
147 })
148 .collect();
149 ranked.sort_by(|left, right| {
150 right
151 .1
152 .partial_cmp(&left.1)
153 .unwrap_or(Ordering::Equal)
154 .then_with(|| left.0.cmp(&right.0))
155 });
156 let keep: HashSet<String> = ranked
157 .into_iter()
158 .take(query.group_count as usize)
159 .map(|(key, _)| key)
160 .collect();
161 groups.retain(|key, _| keep.contains(key));
162 }
163 Ok(groups)
164 }
165
166 pub fn group_by_query(&self, query: &GroupBySearchQuery) -> Result<HashMap<String, Vec<Doc>>> {
167 self.group_by(query)
168 }
169
170 pub fn fetch(&self, pks: &[&str]) -> Result<Vec<Doc>> {
171 self.fetch_with_options(pks, None, true)
172 }
173
174 pub fn fetch_with_options(
175 &self,
176 pks: &[&str],
177 output_fields: Option<&[&str]>,
178 include_vector: bool,
179 ) -> Result<Vec<Doc>> {
180 self.ensure_open()?;
181 let state = self
182 .inner
183 .state
184 .read()
185 .map_err(|_| Error::internal("collection state lock poisoned"))?;
186 let fields =
187 output_fields.map(|values| values.iter().map(|v| (*v).to_string()).collect::<Vec<_>>());
188 Ok(pks
189 .iter()
190 .filter_map(|pk| state.docs.get(*pk))
191 .map(|doc| doc.project(fields.as_deref(), include_vector))
192 .collect())
193 }
194
195 #[allow(clippy::iter_not_returning_iterator)]
198 pub fn iter(&self) -> Result<DocIterator> {
199 self.iter_with_options(None, true)
200 }
201
202 pub fn iter_with_options(
203 &self,
204 output_fields: Option<&[&str]>,
205 include_vector: bool,
206 ) -> Result<DocIterator> {
207 self.ensure_open()?;
208 let state = self
209 .inner
210 .state
211 .read()
212 .map_err(|_| Error::internal("collection state lock poisoned"))?;
213 let fields =
214 output_fields.map(|values| values.iter().map(|v| (*v).to_string()).collect::<Vec<_>>());
215 let docs = state
216 .docs
217 .values()
218 .map(|doc| doc.project(fields.as_deref(), include_vector))
219 .collect();
220 Ok(DocIterator::new(docs, state.revision))
221 }
222}
223
224fn execute_multi_query_branches(
225 snapshot: &CollectionSnapshot,
226 query: &MultiQuery,
227) -> Result<(Vec<Vec<Doc>>, MultiQueryTelemetry)> {
228 let mut planned_branches = Vec::with_capacity(query.queries.len());
229 let mut telemetry = MultiQueryTelemetry::default();
230 for sub in &query.queries {
231 let mut branch = sub.to_search_query()?;
232 if let Some(filter) = query.effective_filter() {
233 branch.set_filter(filter)?;
234 }
235 branch.include_vector = query.include_vector_value;
236 branch.output_fields.clone_from(&query.output_fields);
237 let filter = parse_optional_filter(branch.filter.as_deref())?;
238 let plan = snapshot.indexes.plan_candidates(
239 &snapshot.docs,
240 snapshot.revision,
241 &branch,
242 filter.as_ref(),
243 )?;
244 telemetry.used_ann |= plan.used_ann;
245 telemetry.diskann_io_backend = telemetry.diskann_io_backend.or(plan.diskann_io_backend);
246 telemetry.diskann_sector_reads = telemetry
247 .diskann_sector_reads
248 .saturating_add(plan.diskann_sector_reads);
249 telemetry.used_scalar |= plan.used_scalar;
250 telemetry.used_fts_index |= plan.used_fts_index;
251 telemetry.candidates = telemetry
252 .candidates
253 .saturating_add(plan.candidate_count(snapshot.docs.len()));
254 enforce_query_candidates(snapshot, telemetry.candidates)?;
255 planned_branches.push((branch, filter, plan));
256 }
257
258 let mut branches = Vec::with_capacity(planned_branches.len());
259 for (branch, filter, plan) in planned_branches {
260 branches.push(execute_query_with_candidates(
261 &snapshot.schema,
262 &snapshot.docs,
263 &snapshot.indexes,
264 &branch,
265 plan.selection.as_ref(),
266 plan.fts_scores.as_ref(),
267 filter.as_ref(),
268 )?);
269 }
270 Ok((branches, telemetry))
271}
272
273fn fuse_multi_query_branches(query: &MultiQuery, mut branches: Vec<Vec<Doc>>) -> Result<Vec<Doc>> {
274 let normalization = query.normalization.as_deref().unwrap_or("none");
275 for branch in &mut branches {
276 normalize_scores(branch, normalization)?;
277 }
278 let mut fused: BTreeMap<String, (f64, Doc)> = BTreeMap::new();
279 for (branch_index, branch) in branches.into_iter().enumerate() {
280 let weight = match &query.rerank {
281 RerankMethod::Weighted { weights } => weights.get(branch_index).copied().unwrap_or(1.0),
282 RerankMethod::ReciprocalRank { .. } => 1.0,
283 };
284 for (rank, doc) in branch.into_iter().enumerate() {
285 let Some(id) = doc.get_pk().map(str::to_string) else {
286 continue;
287 };
288 let score = match query.rerank {
289 RerankMethod::ReciprocalRank { rank_constant } => {
290 weight / (rank_constant + count_to_f64(rank) + 1.0)
291 }
292 RerankMethod::Weighted { .. } => weight * f64::from(doc.get_score()),
293 };
294 fused
295 .entry(id)
296 .and_modify(|entry| entry.0 += score)
297 .or_insert((score, doc));
298 }
299 }
300 let mut output: Vec<Doc> = fused
301 .into_values()
302 .map(|(score, mut doc)| {
303 doc.set_score(score_to_f32(score)?)?;
304 Ok(doc)
305 })
306 .collect::<Result<_>>()?;
307 sort_docs(&mut output);
308 let topk = usize::try_from(query.topk_value)
309 .map_err(|_| Error::invalid_argument("multi-query topk must be non-negative"))?;
310 output.truncate(topk);
311 Ok(output)
312}
313
314fn enforce_query_candidates(snapshot: &CollectionSnapshot, candidates: u64) -> Result<()> {
315 if let Err(error) = snapshot
316 .resource_limits
317 .enforce_query_candidates(candidates)
318 {
319 snapshot.stats.record_resource_limit_rejection();
320 Err(error)
321 } else {
322 Ok(())
323 }
324}