Skip to main content

a3s_vec/collection/
query_api.rs

1//! Collection query, fetch, and iterator APIs.
2
3use 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    // This name is retained for zvec API compatibility; `DocIterator` exposes
196    // fallible batch iteration rather than implementing `Iterator` directly.
197    #[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}