Skip to main content

chroma_types/execution/
operator.rs

1use serde::{de::Error, ser::SerializeMap, Deserialize, Deserializer, Serialize, Serializer};
2use serde_json::Value;
3use std::{
4    cmp::Ordering,
5    collections::{BinaryHeap, HashSet},
6    fmt,
7    hash::{Hash, Hasher},
8    ops::{Add, Div, Mul, Neg, Sub},
9};
10use thiserror::Error;
11
12use crate::{
13    chroma_proto, logical_size_of_metadata, parse_where, CollectionAndSegments, CollectionUuid,
14    ContainsOperator, DocumentExpression, DocumentOperator, Metadata, MetadataComparison,
15    MetadataExpression, MetadataSetValue, MetadataValue, PrimitiveOperator, ScalarEncoding,
16    SetOperator, SparseVector, Where,
17};
18
19use super::error::QueryConversionError;
20
21pub type InitialInput = ();
22
23/// The `Scan` opeartor pins the data used by all downstream operators
24///
25/// # Parameters
26/// - `collection_and_segments`: The consistent snapshot of collection
27#[derive(Clone, Debug)]
28pub struct Scan {
29    pub collection_and_segments: CollectionAndSegments,
30    pub shard_index: u32,
31    pub num_shards: u32,
32    /// Upper bound log offset scouted by the frontend.
33    /// 0 means the worker should scout independently.
34    pub log_upper_bound_offset: i64,
35}
36
37impl TryFrom<chroma_proto::ScanOperator> for Scan {
38    type Error = QueryConversionError;
39
40    fn try_from(value: chroma_proto::ScanOperator) -> Result<Self, Self::Error> {
41        let num_shards = value.num_shards.max(1);
42        Ok(Self {
43            collection_and_segments: CollectionAndSegments {
44                collection: value
45                    .collection
46                    .ok_or(QueryConversionError::field("collection"))?
47                    .try_into()?,
48                metadata_segment: value
49                    .metadata
50                    .ok_or(QueryConversionError::field("metadata segment"))?
51                    .try_into()?,
52                record_segment: value
53                    .record
54                    .ok_or(QueryConversionError::field("record segment"))?
55                    .try_into()?,
56                vector_segment: value
57                    .knn
58                    .ok_or(QueryConversionError::field("vector segment"))?
59                    .try_into()?,
60            },
61            shard_index: value.shard_index,
62            num_shards,
63            log_upper_bound_offset: value.log_upper_bound_offset,
64        })
65    }
66}
67
68#[derive(Debug, Error)]
69pub enum ScanToProtoError {
70    #[error("Could not convert collection to proto")]
71    CollectionToProto(#[from] crate::CollectionToProtoError),
72}
73
74impl TryFrom<Scan> for chroma_proto::ScanOperator {
75    type Error = ScanToProtoError;
76
77    fn try_from(value: Scan) -> Result<Self, Self::Error> {
78        Ok(Self {
79            collection: Some(value.collection_and_segments.collection.try_into()?),
80            knn: Some(value.collection_and_segments.vector_segment.into()),
81            metadata: Some(value.collection_and_segments.metadata_segment.into()),
82            record: Some(value.collection_and_segments.record_segment.into()),
83            shard_index: value.shard_index,
84            num_shards: value.num_shards,
85            log_upper_bound_offset: value.log_upper_bound_offset,
86        })
87    }
88}
89
90#[derive(Clone, Debug)]
91pub struct CountResult {
92    pub count: u32,
93    pub pulled_log_bytes: u64,
94}
95
96impl CountResult {
97    pub fn size_bytes(&self) -> u64 {
98        size_of_val(&self.count) as u64
99    }
100}
101
102impl From<chroma_proto::CountResult> for CountResult {
103    fn from(value: chroma_proto::CountResult) -> Self {
104        Self {
105            count: value.count,
106            pulled_log_bytes: value.pulled_log_bytes,
107        }
108    }
109}
110
111impl From<CountResult> for chroma_proto::CountResult {
112    fn from(value: CountResult) -> Self {
113        Self {
114            count: value.count,
115            pulled_log_bytes: value.pulled_log_bytes,
116        }
117    }
118}
119
120/// The `FetchLog` operator fetches logs from the log service
121///
122/// # Parameters
123/// - `start_log_offset_id`: The offset id of the first log to read
124/// - `maximum_fetch_count`: The maximum number of logs to fetch in total
125/// - `collection_uuid`: The uuid of the collection where the fetched logs should belong
126#[derive(Clone, Debug)]
127pub struct FetchLog {
128    pub collection_uuid: CollectionUuid,
129    pub maximum_fetch_count: Option<u32>,
130    pub start_log_offset_id: u32,
131}
132
133/// Filter the search results.
134///
135/// Combines document ID filtering with metadata and document content predicates.
136/// For the Search API, use `where_clause` with Key expressions.
137///
138/// # Fields
139///
140/// * `query_ids` - Optional list of document IDs to filter (legacy, prefer Where expressions)
141/// * `where_clause` - Predicate on document metadata, content, or IDs
142///
143/// # Examples
144///
145/// ## Simple metadata filter
146///
147/// ```
148/// use chroma_types::operator::{Filter, Key};
149///
150/// let filter = Filter {
151///     query_ids: None,
152///     where_clause: Some(Key::field("status").eq("published")),
153/// };
154/// ```
155///
156/// ## Combined filters
157///
158/// ```
159/// use chroma_types::operator::{Filter, Key};
160///
161/// let filter = Filter {
162///     query_ids: None,
163///     where_clause: Some(
164///         Key::field("status").eq("published")
165///             & Key::field("year").gte(2020)
166///             & Key::field("category").is_in(vec!["tech", "science"])
167///     ),
168/// };
169/// ```
170///
171/// ## Document content filter
172///
173/// ```
174/// use chroma_types::operator::{Filter, Key};
175///
176/// let filter = Filter {
177///     query_ids: None,
178///     where_clause: Some(Key::Document.contains("machine learning")),
179/// };
180/// ```
181#[derive(Clone, Debug, Default)]
182pub struct Filter {
183    pub query_ids: Option<Vec<String>>,
184    pub where_clause: Option<Where>,
185}
186
187impl Serialize for Filter {
188    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
189    where
190        S: Serializer,
191    {
192        // For the search API, serialize directly as the where clause (or empty object if None)
193        // If query_ids are present, they should be combined with the where_clause as Key::ID.is_in([...])
194
195        match (&self.query_ids, &self.where_clause) {
196            (None, None) => {
197                // No filter at all - serialize empty object
198                let map = serializer.serialize_map(Some(0))?;
199                map.end()
200            }
201            (None, Some(where_clause)) => {
202                // Only where clause - serialize it directly
203                where_clause.serialize(serializer)
204            }
205            (Some(ids), None) => {
206                // Only query_ids - create Where clause: Key::ID.is_in(ids)
207                let id_where = Where::Metadata(MetadataExpression {
208                    key: "#id".to_string(),
209                    comparison: MetadataComparison::Set(
210                        SetOperator::In,
211                        MetadataSetValue::Str(ids.clone()),
212                    ),
213                });
214                id_where.serialize(serializer)
215            }
216            (Some(ids), Some(where_clause)) => {
217                // Both present - combine with AND: Key::ID.is_in(ids) & where_clause
218                let id_where = Where::Metadata(MetadataExpression {
219                    key: "#id".to_string(),
220                    comparison: MetadataComparison::Set(
221                        SetOperator::In,
222                        MetadataSetValue::Str(ids.clone()),
223                    ),
224                });
225                let combined = Where::conjunction(vec![id_where, where_clause.clone()]);
226                combined.serialize(serializer)
227            }
228        }
229    }
230}
231
232impl<'de> Deserialize<'de> for Filter {
233    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
234    where
235        D: Deserializer<'de>,
236    {
237        // For the new search API, the entire JSON is the where clause
238        let where_json = Value::deserialize(deserializer)?;
239        let where_clause =
240            if where_json.is_null() || where_json.as_object().is_some_and(|obj| obj.is_empty()) {
241                None
242            } else {
243                Some(parse_where(&where_json).map_err(|e| D::Error::custom(e.to_string()))?)
244            };
245
246        Ok(Filter {
247            query_ids: None, // Always None for new search API
248            where_clause,
249        })
250    }
251}
252
253impl TryFrom<chroma_proto::FilterOperator> for Filter {
254    type Error = QueryConversionError;
255
256    fn try_from(value: chroma_proto::FilterOperator) -> Result<Self, Self::Error> {
257        let where_metadata = value.r#where.map(TryInto::try_into).transpose()?;
258        let where_document = value.where_document.map(TryInto::try_into).transpose()?;
259        let where_clause = match (where_metadata, where_document) {
260            (Some(w), Some(wd)) => Some(Where::conjunction(vec![w, wd])),
261            (Some(w), None) | (None, Some(w)) => Some(w),
262            _ => None,
263        };
264
265        Ok(Self {
266            query_ids: value.ids.map(|uids| uids.ids),
267            where_clause,
268        })
269    }
270}
271
272impl TryFrom<Filter> for chroma_proto::FilterOperator {
273    type Error = QueryConversionError;
274
275    fn try_from(value: Filter) -> Result<Self, Self::Error> {
276        Ok(Self {
277            ids: value.query_ids.map(|ids| chroma_proto::UserIds { ids }),
278            r#where: value.where_clause.map(TryInto::try_into).transpose()?,
279            where_document: None,
280        })
281    }
282}
283
284/// The `Knn` operator searches for the nearest neighbours of the specified embedding. This is intended to use by executor
285///
286/// # Parameters
287/// - `embedding`: The target embedding to search around
288/// - `fetch`: The number of records to fetch around the target
289#[derive(Clone, Debug)]
290pub struct Knn {
291    pub embedding: Vec<f32>,
292    pub fetch: u32,
293}
294
295impl From<KnnBatch> for Vec<Knn> {
296    fn from(value: KnnBatch) -> Self {
297        value
298            .embeddings
299            .into_iter()
300            .map(|embedding| Knn {
301                embedding,
302                fetch: value.fetch,
303            })
304            .collect()
305    }
306}
307
308/// The `KnnBatch` operator searches for the nearest neighbours of the specified embedding. This is intended to use by frontend
309///
310/// # Parameters
311/// - `embedding`: The target embedding to search around
312/// - `fetch`: The number of records to fetch around the target
313#[derive(Clone, Debug)]
314pub struct KnnBatch {
315    pub embeddings: Vec<Vec<f32>>,
316    pub fetch: u32,
317}
318
319impl TryFrom<chroma_proto::KnnOperator> for KnnBatch {
320    type Error = QueryConversionError;
321
322    fn try_from(value: chroma_proto::KnnOperator) -> Result<Self, Self::Error> {
323        Ok(Self {
324            embeddings: value
325                .embeddings
326                .into_iter()
327                .map(|vec| vec.try_into().map(|(v, _)| v))
328                .collect::<Result<_, _>>()?,
329            fetch: value.fetch,
330        })
331    }
332}
333
334impl TryFrom<KnnBatch> for chroma_proto::KnnOperator {
335    type Error = QueryConversionError;
336
337    fn try_from(value: KnnBatch) -> Result<Self, Self::Error> {
338        Ok(Self {
339            embeddings: value
340                .embeddings
341                .into_iter()
342                .map(|embedding| {
343                    let dim = embedding.len();
344                    chroma_proto::Vector::try_from((embedding, ScalarEncoding::FLOAT32, dim))
345                })
346                .collect::<Result<_, _>>()?,
347            fetch: value.fetch,
348        })
349    }
350}
351
352/// Pagination control for search results.
353///
354/// Controls how many results to return and how many to skip for pagination.
355///
356/// # Fields
357///
358/// * `offset` - Number of results to skip (default: 0)
359/// * `limit` - Maximum results to return (None = no limit)
360///
361/// # Examples
362///
363/// ```
364/// use chroma_types::operator::Limit;
365///
366/// // First page: results 0-9
367/// let limit = Limit {
368///     offset: 0,
369///     limit: Some(10),
370/// };
371///
372/// // Second page: results 10-19
373/// let limit = Limit {
374///     offset: 10,
375///     limit: Some(10),
376/// };
377///
378/// // No limit: all results
379/// let limit = Limit {
380///     offset: 0,
381///     limit: None,
382/// };
383/// ```
384#[derive(Clone, Debug, Default, Deserialize, Serialize)]
385pub struct Limit {
386    #[serde(default)]
387    pub offset: u32,
388    #[serde(default)]
389    pub limit: Option<u32>,
390}
391
392impl From<chroma_proto::LimitOperator> for Limit {
393    fn from(value: chroma_proto::LimitOperator) -> Self {
394        Self {
395            offset: value.offset,
396            limit: value.limit,
397        }
398    }
399}
400
401impl From<Limit> for chroma_proto::LimitOperator {
402    fn from(value: Limit) -> Self {
403        Self {
404            offset: value.offset,
405            limit: value.limit,
406        }
407    }
408}
409
410/// The `RecordDistance` represents a measure of embedding (identified by `offset_id`) with respect to query embedding
411#[derive(Clone, Copy, Debug)]
412pub struct RecordMeasure {
413    pub offset_id: u32,
414    pub measure: f32,
415}
416
417impl PartialEq for RecordMeasure {
418    fn eq(&self, other: &Self) -> bool {
419        self.offset_id.eq(&other.offset_id)
420    }
421}
422
423impl Eq for RecordMeasure {}
424
425impl Hash for RecordMeasure {
426    fn hash<H: Hasher>(&self, state: &mut H) {
427        self.offset_id.hash(state);
428    }
429}
430
431impl Ord for RecordMeasure {
432    fn cmp(&self, other: &Self) -> Ordering {
433        self.measure
434            .total_cmp(&other.measure)
435            .then_with(|| self.offset_id.cmp(&other.offset_id))
436    }
437}
438
439impl PartialOrd for RecordMeasure {
440    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
441        Some(self.cmp(other))
442    }
443}
444
445#[derive(Debug, Default)]
446pub struct KnnOutput {
447    pub distances: Vec<RecordMeasure>,
448}
449
450/// The `Merge` operator selects the top records from the batch vectors of records
451/// which are all sorted in descending order. If the same record occurs multiple times
452/// only one copy will remain in the final result.
453///
454/// # Parameters
455/// - `k`: The total number of records to take after merge
456///
457/// # Usage
458/// It can be used to merge the query results from different operators
459#[derive(Clone, Debug)]
460pub struct Merge {
461    pub k: u32,
462}
463
464impl Merge {
465    pub fn merge<M: Clone + Eq + Hash + Ord>(&self, input: Vec<Vec<M>>) -> Vec<M> {
466        let mut batch_iters = input.into_iter().map(Vec::into_iter).collect::<Vec<_>>();
467
468        let mut max_heap = batch_iters
469            .iter_mut()
470            .enumerate()
471            .filter_map(|(idx, itr)| itr.next().map(|rec| (rec, idx)))
472            .collect::<BinaryHeap<_>>();
473
474        let mut seen = HashSet::with_capacity(self.k as usize);
475        let mut fusion = Vec::with_capacity(self.k as usize);
476        while let Some((m, idx)) = max_heap.pop() {
477            if self.k <= fusion.len() as u32 {
478                break;
479            }
480            if let Some(next_m) = batch_iters[idx].next() {
481                max_heap.push((next_m, idx));
482            }
483            if !seen.insert(m.clone()) {
484                continue;
485            }
486            fusion.push(m);
487        }
488        fusion
489    }
490}
491
492/// The `Projection` operator retrieves record content by offset ids
493///
494/// # Parameters
495/// - `document`: Whether to retrieve document
496/// - `embedding`: Whether to retrieve embedding
497/// - `metadata`: Whether to retrieve metadata
498#[derive(Clone, Debug, Default)]
499pub struct Projection {
500    pub document: bool,
501    pub embedding: bool,
502    pub metadata: bool,
503}
504
505impl From<chroma_proto::ProjectionOperator> for Projection {
506    fn from(value: chroma_proto::ProjectionOperator) -> Self {
507        Self {
508            document: value.document,
509            embedding: value.embedding,
510            metadata: value.metadata,
511        }
512    }
513}
514
515impl From<Projection> for chroma_proto::ProjectionOperator {
516    fn from(value: Projection) -> Self {
517        Self {
518            document: value.document,
519            embedding: value.embedding,
520            metadata: value.metadata,
521        }
522    }
523}
524
525#[derive(Clone, Debug, PartialEq)]
526pub struct ProjectionRecord {
527    pub id: String,
528    pub document: Option<String>,
529    pub embedding: Option<Vec<f32>>,
530    pub metadata: Option<Metadata>,
531}
532
533impl ProjectionRecord {
534    pub fn size_bytes(&self) -> u64 {
535        (self.id.len()
536            + self
537                .document
538                .as_ref()
539                .map(|doc| doc.len())
540                .unwrap_or_default()
541            + self
542                .embedding
543                .as_ref()
544                .map(|emb| size_of_val(&emb[..]))
545                .unwrap_or_default()
546            + self
547                .metadata
548                .as_ref()
549                .map(logical_size_of_metadata)
550                .unwrap_or_default()) as u64
551    }
552}
553
554impl Eq for ProjectionRecord {}
555
556impl TryFrom<chroma_proto::ProjectionRecord> for ProjectionRecord {
557    type Error = QueryConversionError;
558
559    fn try_from(value: chroma_proto::ProjectionRecord) -> Result<Self, Self::Error> {
560        Ok(Self {
561            id: value.id,
562            document: value.document,
563            embedding: value
564                .embedding
565                .map(|vec| vec.try_into().map(|(v, _)| v))
566                .transpose()?,
567            metadata: value.metadata.map(TryInto::try_into).transpose()?,
568        })
569    }
570}
571
572impl TryFrom<ProjectionRecord> for chroma_proto::ProjectionRecord {
573    type Error = QueryConversionError;
574
575    fn try_from(value: ProjectionRecord) -> Result<Self, Self::Error> {
576        Ok(Self {
577            id: value.id,
578            document: value.document,
579            embedding: value
580                .embedding
581                .map(|embedding| {
582                    let embedding_dimension = embedding.len();
583                    chroma_proto::Vector::try_from((
584                        embedding,
585                        ScalarEncoding::FLOAT32,
586                        embedding_dimension,
587                    ))
588                })
589                .transpose()?,
590            metadata: value.metadata.map(|metadata| metadata.into()),
591        })
592    }
593}
594
595#[derive(Clone, Debug, Eq, PartialEq)]
596pub struct ProjectionOutput {
597    pub records: Vec<ProjectionRecord>,
598}
599
600#[derive(Clone, Debug, Eq, PartialEq)]
601pub struct GetResult {
602    pub pulled_log_bytes: u64,
603    pub result: ProjectionOutput,
604}
605
606impl GetResult {
607    pub fn size_bytes(&self) -> u64 {
608        self.result
609            .records
610            .iter()
611            .map(ProjectionRecord::size_bytes)
612            .sum()
613    }
614}
615
616impl TryFrom<chroma_proto::GetResult> for GetResult {
617    type Error = QueryConversionError;
618
619    fn try_from(value: chroma_proto::GetResult) -> Result<Self, Self::Error> {
620        Ok(Self {
621            pulled_log_bytes: value.pulled_log_bytes,
622            result: ProjectionOutput {
623                records: value
624                    .records
625                    .into_iter()
626                    .map(TryInto::try_into)
627                    .collect::<Result<_, _>>()?,
628            },
629        })
630    }
631}
632
633impl TryFrom<GetResult> for chroma_proto::GetResult {
634    type Error = QueryConversionError;
635
636    fn try_from(value: GetResult) -> Result<Self, Self::Error> {
637        Ok(Self {
638            pulled_log_bytes: value.pulled_log_bytes,
639            records: value
640                .result
641                .records
642                .into_iter()
643                .map(TryInto::try_into)
644                .collect::<Result<_, _>>()?,
645        })
646    }
647}
648
649/// The `KnnProjection` operator retrieves record content by offset ids
650/// It is based on `ProjectionOperator`, and it attaches the distance
651/// of the records to the target embedding to the record content
652///
653/// # Parameters
654/// - `projection`: The parameters of the `ProjectionOperator`
655/// - `distance`: Whether to attach distance information
656#[derive(Clone, Debug)]
657pub struct KnnProjection {
658    pub projection: Projection,
659    pub distance: bool,
660}
661
662impl TryFrom<chroma_proto::KnnProjectionOperator> for KnnProjection {
663    type Error = QueryConversionError;
664
665    fn try_from(value: chroma_proto::KnnProjectionOperator) -> Result<Self, Self::Error> {
666        Ok(Self {
667            projection: value
668                .projection
669                .ok_or(QueryConversionError::field("projection"))?
670                .into(),
671            distance: value.distance,
672        })
673    }
674}
675
676impl From<KnnProjection> for chroma_proto::KnnProjectionOperator {
677    fn from(value: KnnProjection) -> Self {
678        Self {
679            projection: Some(value.projection.into()),
680            distance: value.distance,
681        }
682    }
683}
684
685#[derive(Clone, Debug)]
686pub struct KnnProjectionRecord {
687    pub record: ProjectionRecord,
688    pub distance: Option<f32>,
689}
690
691impl TryFrom<chroma_proto::KnnProjectionRecord> for KnnProjectionRecord {
692    type Error = QueryConversionError;
693
694    fn try_from(value: chroma_proto::KnnProjectionRecord) -> Result<Self, Self::Error> {
695        Ok(Self {
696            record: value
697                .record
698                .ok_or(QueryConversionError::field("record"))?
699                .try_into()?,
700            distance: value.distance,
701        })
702    }
703}
704
705impl TryFrom<KnnProjectionRecord> for chroma_proto::KnnProjectionRecord {
706    type Error = QueryConversionError;
707
708    fn try_from(value: KnnProjectionRecord) -> Result<Self, Self::Error> {
709        Ok(Self {
710            record: Some(value.record.try_into()?),
711            distance: value.distance,
712        })
713    }
714}
715
716#[derive(Clone, Debug, Default)]
717pub struct KnnProjectionOutput {
718    pub records: Vec<KnnProjectionRecord>,
719}
720
721impl TryFrom<chroma_proto::KnnResult> for KnnProjectionOutput {
722    type Error = QueryConversionError;
723
724    fn try_from(value: chroma_proto::KnnResult) -> Result<Self, Self::Error> {
725        Ok(Self {
726            records: value
727                .records
728                .into_iter()
729                .map(TryInto::try_into)
730                .collect::<Result<_, _>>()?,
731        })
732    }
733}
734
735impl TryFrom<KnnProjectionOutput> for chroma_proto::KnnResult {
736    type Error = QueryConversionError;
737
738    fn try_from(value: KnnProjectionOutput) -> Result<Self, Self::Error> {
739        Ok(Self {
740            records: value
741                .records
742                .into_iter()
743                .map(TryInto::try_into)
744                .collect::<Result<_, _>>()?,
745        })
746    }
747}
748
749#[derive(Clone, Debug, Default)]
750pub struct KnnBatchResult {
751    pub pulled_log_bytes: u64,
752    pub results: Vec<KnnProjectionOutput>,
753}
754
755impl KnnBatchResult {
756    pub fn size_bytes(&self) -> u64 {
757        self.results
758            .iter()
759            .flat_map(|res| {
760                res.records
761                    .iter()
762                    .map(|rec| rec.record.size_bytes() + size_of_val(&rec.distance) as u64)
763            })
764            .sum()
765    }
766}
767
768impl TryFrom<chroma_proto::KnnBatchResult> for KnnBatchResult {
769    type Error = QueryConversionError;
770
771    fn try_from(value: chroma_proto::KnnBatchResult) -> Result<Self, Self::Error> {
772        Ok(Self {
773            pulled_log_bytes: value.pulled_log_bytes,
774            results: value
775                .results
776                .into_iter()
777                .map(TryInto::try_into)
778                .collect::<Result<_, _>>()?,
779        })
780    }
781}
782
783impl TryFrom<KnnBatchResult> for chroma_proto::KnnBatchResult {
784    type Error = QueryConversionError;
785
786    fn try_from(value: KnnBatchResult) -> Result<Self, Self::Error> {
787        Ok(Self {
788            pulled_log_bytes: value.pulled_log_bytes,
789            results: value
790                .results
791                .into_iter()
792                .map(TryInto::try_into)
793                .collect::<Result<_, _>>()?,
794        })
795    }
796}
797
798/// A query vector for KNN search.
799///
800/// Supports both dense and sparse vector formats.
801///
802/// # Variants
803///
804/// ## Dense
805///
806/// Standard dense embeddings as a vector of floats.
807///
808/// ```
809/// use chroma_types::operator::QueryVector;
810///
811/// let dense = QueryVector::Dense(vec![0.1, 0.2, 0.3, 0.4]);
812/// ```
813///
814/// ## Sparse
815///
816/// Sparse vectors with explicit indices and values.
817///
818/// ```
819/// use chroma_types::operator::QueryVector;
820/// use chroma_types::SparseVector;
821///
822/// let sparse = QueryVector::Sparse(SparseVector::new(
823///     vec![0, 5, 10, 50],      // indices
824///     vec![0.5, 0.3, 0.8, 0.2], // values
825/// ).unwrap());
826/// ```
827///
828/// # Examples
829///
830/// ## Dense vector in KNN
831///
832/// ```
833/// use chroma_types::operator::{RankExpr, QueryVector, Key};
834///
835/// let rank = RankExpr::Knn {
836///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
837///     key: Key::Embedding,
838///     limit: 100,
839///     default: None,
840///     return_rank: false,
841/// };
842/// ```
843///
844/// ## Sparse vector in KNN
845///
846/// ```
847/// use chroma_types::operator::{RankExpr, QueryVector, Key};
848/// use chroma_types::SparseVector;
849///
850/// let rank = RankExpr::Knn {
851///     query: QueryVector::Sparse(SparseVector::new(
852///         vec![1, 5, 10],
853///         vec![0.5, 0.3, 0.8],
854///     ).unwrap()),
855///     key: Key::field("sparse_embedding"),
856///     limit: 100,
857///     default: None,
858///     return_rank: false,
859/// };
860/// ```
861#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
862#[serde(untagged)]
863pub enum QueryVector {
864    Dense(Vec<f32>),
865    Sparse(SparseVector),
866}
867
868impl TryFrom<chroma_proto::QueryVector> for QueryVector {
869    type Error = QueryConversionError;
870
871    fn try_from(value: chroma_proto::QueryVector) -> Result<Self, Self::Error> {
872        let vector = value.vector.ok_or(QueryConversionError::field("vector"))?;
873        match vector {
874            chroma_proto::query_vector::Vector::Dense(dense) => {
875                Ok(QueryVector::Dense(dense.try_into().map(|(v, _)| v)?))
876            }
877            chroma_proto::query_vector::Vector::Sparse(sparse) => {
878                Ok(QueryVector::Sparse(sparse.try_into().map_err(|_| {
879                    QueryConversionError::validation("sparse vector length mismatch")
880                })?))
881            }
882        }
883    }
884}
885
886impl TryFrom<QueryVector> for chroma_proto::QueryVector {
887    type Error = QueryConversionError;
888
889    fn try_from(value: QueryVector) -> Result<Self, Self::Error> {
890        match value {
891            QueryVector::Dense(vec) => {
892                let dim = vec.len();
893                Ok(chroma_proto::QueryVector {
894                    vector: Some(chroma_proto::query_vector::Vector::Dense(
895                        chroma_proto::Vector::try_from((vec, ScalarEncoding::FLOAT32, dim))?,
896                    )),
897                })
898            }
899            QueryVector::Sparse(sparse) => Ok(chroma_proto::QueryVector {
900                vector: Some(chroma_proto::query_vector::Vector::Sparse(sparse.into())),
901            }),
902        }
903    }
904}
905
906impl From<Vec<f32>> for QueryVector {
907    fn from(vec: Vec<f32>) -> Self {
908        QueryVector::Dense(vec)
909    }
910}
911
912impl From<SparseVector> for QueryVector {
913    fn from(sparse: SparseVector) -> Self {
914        QueryVector::Sparse(sparse)
915    }
916}
917
918#[derive(Clone, Debug, PartialEq)]
919pub struct KnnQuery {
920    pub query: QueryVector,
921    pub key: Key,
922    pub limit: u32,
923}
924
925/// Wrapper for ranking expressions in search queries.
926///
927/// Contains an optional ranking expression. When None, results are returned in
928/// natural storage order without scoring.
929///
930/// # Fields
931///
932/// * `expr` - The ranking expression (None = no ranking)
933///
934/// # Examples
935///
936/// ```
937/// use chroma_types::operator::{Rank, RankExpr, QueryVector, Key};
938///
939/// // With ranking
940/// let rank = Rank {
941///     expr: Some(RankExpr::Knn {
942///         query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
943///         key: Key::Embedding,
944///         limit: 100,
945///         default: None,
946///         return_rank: false,
947///     }),
948/// };
949///
950/// // No ranking (natural order)
951/// let rank = Rank {
952///     expr: None,
953/// };
954/// ```
955#[derive(Clone, Debug, Default, Deserialize, Serialize)]
956#[serde(transparent)]
957pub struct Rank {
958    pub expr: Option<RankExpr>,
959}
960
961impl Rank {
962    pub fn knn_queries(&self) -> Vec<KnnQuery> {
963        self.expr
964            .as_ref()
965            .map(RankExpr::knn_queries)
966            .unwrap_or_default()
967    }
968}
969
970impl TryFrom<chroma_proto::RankOperator> for Rank {
971    type Error = QueryConversionError;
972
973    fn try_from(proto_rank: chroma_proto::RankOperator) -> Result<Self, Self::Error> {
974        Ok(Rank {
975            expr: proto_rank.expr.map(TryInto::try_into).transpose()?,
976        })
977    }
978}
979
980impl TryFrom<Rank> for chroma_proto::RankOperator {
981    type Error = QueryConversionError;
982
983    fn try_from(rank: Rank) -> Result<Self, Self::Error> {
984        Ok(chroma_proto::RankOperator {
985            expr: rank.expr.map(TryInto::try_into).transpose()?,
986        })
987    }
988}
989
990/// A ranking expression for scoring and ordering search results.
991///
992/// Ranking expressions determine which documents appear in results and their order.
993/// Lower scores indicate better matches (distance-based scoring).
994///
995/// # Variants
996///
997/// ## Knn - K-Nearest Neighbor Search
998///
999/// The primary ranking method for vector similarity search.
1000///
1001/// ```
1002/// use chroma_types::operator::{RankExpr, QueryVector, Key};
1003///
1004/// let rank = RankExpr::Knn {
1005///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1006///     key: Key::Embedding,
1007///     limit: 100,        // Consider top 100 candidates
1008///     default: None,     // No default score for missing documents
1009///     return_rank: false, // Return distances, not rank positions
1010/// };
1011/// ```
1012///
1013/// ## Value - Constant
1014///
1015/// Represents a constant score.
1016///
1017/// ```
1018/// use chroma_types::operator::RankExpr;
1019///
1020/// let rank = RankExpr::Value(0.5);
1021/// ```
1022///
1023/// ## Arithmetic Operations
1024///
1025/// Combine ranking expressions using standard operators (+, -, *, /).
1026///
1027/// ```
1028/// use chroma_types::operator::{RankExpr, QueryVector, Key};
1029///
1030/// let knn1 = RankExpr::Knn {
1031///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1032///     key: Key::Embedding,
1033///     limit: 100,
1034///     default: None,
1035///     return_rank: false,
1036/// };
1037///
1038/// let knn2 = RankExpr::Knn {
1039///     query: QueryVector::Dense(vec![0.2, 0.3, 0.4]),
1040///     key: Key::field("other_embedding"),
1041///     limit: 100,
1042///     default: None,
1043///     return_rank: false,
1044/// };
1045///
1046/// // Weighted combination: 70% knn1 + 30% knn2
1047/// let combined = knn1 * 0.7 + knn2 * 0.3;
1048///
1049/// // Normalized
1050/// let normalized = combined / 2.0;
1051/// ```
1052///
1053/// ## Mathematical Functions
1054///
1055/// Apply mathematical transformations to scores.
1056///
1057/// ```
1058/// use chroma_types::operator::{RankExpr, QueryVector, Key};
1059///
1060/// let knn = RankExpr::Knn {
1061///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1062///     key: Key::Embedding,
1063///     limit: 100,
1064///     default: None,
1065///     return_rank: false,
1066/// };
1067///
1068/// // Exponential - amplifies differences
1069/// let amplified = knn.clone().exp();
1070///
1071/// // Logarithm - compresses range (add constant to avoid log(0))
1072/// let compressed = (knn.clone() + 1.0).log();
1073///
1074/// // Absolute value
1075/// let absolute = knn.clone().abs();
1076///
1077/// // Min/Max - clamping
1078/// let clamped = knn.min(1.0).max(0.0);
1079/// ```
1080///
1081/// # Examples
1082///
1083/// ## Basic vector search
1084///
1085/// ```
1086/// use chroma_types::operator::{RankExpr, QueryVector, Key};
1087///
1088/// let rank = RankExpr::Knn {
1089///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1090///     key: Key::Embedding,
1091///     limit: 100,
1092///     default: None,
1093///     return_rank: false,
1094/// };
1095/// ```
1096///
1097/// ## Hybrid search with weighted combination
1098///
1099/// ```
1100/// use chroma_types::operator::{RankExpr, QueryVector, Key};
1101///
1102/// let dense = RankExpr::Knn {
1103///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1104///     key: Key::Embedding,
1105///     limit: 200,
1106///     default: None,
1107///     return_rank: false,
1108/// };
1109///
1110/// let sparse = RankExpr::Knn {
1111///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]), // Use sparse in practice
1112///     key: Key::field("sparse_embedding"),
1113///     limit: 200,
1114///     default: None,
1115///     return_rank: false,
1116/// };
1117///
1118/// // 70% semantic + 30% keyword
1119/// let hybrid = dense * 0.7 + sparse * 0.3;
1120/// ```
1121///
1122/// ## Reciprocal Rank Fusion (RRF)
1123///
1124/// Use the `rrf()` function for combining rankings with different score scales.
1125///
1126/// ```
1127/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
1128///
1129/// let dense = RankExpr::Knn {
1130///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1131///     key: Key::Embedding,
1132///     limit: 200,
1133///     default: None,
1134///     return_rank: true, // RRF requires rank positions
1135/// };
1136///
1137/// let sparse = RankExpr::Knn {
1138///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1139///     key: Key::field("sparse_embedding"),
1140///     limit: 200,
1141///     default: None,
1142///     return_rank: true, // RRF requires rank positions
1143/// };
1144///
1145/// let rrf_rank = rrf(
1146///     vec![dense, sparse],
1147///     Some(60),           // k parameter (smoothing)
1148///     Some(vec![0.7, 0.3]), // weights
1149///     false,              // normalize weights
1150/// ).unwrap();
1151/// ```
1152#[derive(Clone, Debug, Deserialize, Serialize)]
1153pub enum RankExpr {
1154    #[serde(rename = "$abs")]
1155    Absolute(Box<RankExpr>),
1156    #[serde(rename = "$div")]
1157    Division {
1158        left: Box<RankExpr>,
1159        right: Box<RankExpr>,
1160    },
1161    #[serde(rename = "$exp")]
1162    Exponentiation(Box<RankExpr>),
1163    #[serde(rename = "$knn")]
1164    Knn {
1165        query: QueryVector,
1166        #[serde(default = "RankExpr::default_knn_key")]
1167        key: Key,
1168        #[serde(default = "RankExpr::default_knn_limit")]
1169        limit: u32,
1170        #[serde(default)]
1171        default: Option<f32>,
1172        #[serde(default)]
1173        return_rank: bool,
1174    },
1175    #[serde(rename = "$log")]
1176    Logarithm(Box<RankExpr>),
1177    #[serde(rename = "$max")]
1178    Maximum(Vec<RankExpr>),
1179    #[serde(rename = "$min")]
1180    Minimum(Vec<RankExpr>),
1181    #[serde(rename = "$mul")]
1182    Multiplication(Vec<RankExpr>),
1183    #[serde(rename = "$sub")]
1184    Subtraction {
1185        left: Box<RankExpr>,
1186        right: Box<RankExpr>,
1187    },
1188    #[serde(rename = "$sum")]
1189    Summation(Vec<RankExpr>),
1190    #[serde(rename = "$val")]
1191    Value(f32),
1192}
1193
1194impl RankExpr {
1195    pub fn default_knn_key() -> Key {
1196        Key::Embedding
1197    }
1198
1199    pub fn default_knn_limit() -> u32 {
1200        16
1201    }
1202
1203    pub fn knn_queries(&self) -> Vec<KnnQuery> {
1204        match self {
1205            RankExpr::Absolute(expr)
1206            | RankExpr::Exponentiation(expr)
1207            | RankExpr::Logarithm(expr) => expr.knn_queries(),
1208            RankExpr::Division { left, right } | RankExpr::Subtraction { left, right } => left
1209                .knn_queries()
1210                .into_iter()
1211                .chain(right.knn_queries())
1212                .collect(),
1213            RankExpr::Maximum(exprs)
1214            | RankExpr::Minimum(exprs)
1215            | RankExpr::Multiplication(exprs)
1216            | RankExpr::Summation(exprs) => exprs.iter().flat_map(RankExpr::knn_queries).collect(),
1217            RankExpr::Value(_) => Vec::new(),
1218            RankExpr::Knn {
1219                query,
1220                key,
1221                limit,
1222                default: _,
1223                return_rank: _,
1224            } => vec![KnnQuery {
1225                query: query.clone(),
1226                key: key.clone(),
1227                limit: *limit,
1228            }],
1229        }
1230    }
1231
1232    /// Applies exponential transformation: e^rank.
1233    ///
1234    /// Amplifies differences between scores.
1235    ///
1236    /// # Examples
1237    ///
1238    /// ```
1239    /// use chroma_types::operator::{RankExpr, QueryVector, Key};
1240    ///
1241    /// let knn = RankExpr::Knn {
1242    ///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1243    ///     key: Key::Embedding,
1244    ///     limit: 100,
1245    ///     default: None,
1246    ///     return_rank: false,
1247    /// };
1248    ///
1249    /// let amplified = knn.exp();
1250    /// ```
1251    pub fn exp(self) -> Self {
1252        RankExpr::Exponentiation(Box::new(self))
1253    }
1254
1255    /// Applies natural logarithm transformation: ln(rank).
1256    ///
1257    /// Compresses the score range. Add a constant to avoid log(0).
1258    ///
1259    /// # Examples
1260    ///
1261    /// ```
1262    /// use chroma_types::operator::{RankExpr, QueryVector, Key};
1263    ///
1264    /// let knn = RankExpr::Knn {
1265    ///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1266    ///     key: Key::Embedding,
1267    ///     limit: 100,
1268    ///     default: None,
1269    ///     return_rank: false,
1270    /// };
1271    ///
1272    /// // Add constant to avoid log(0)
1273    /// let compressed = (knn + 1.0).log();
1274    /// ```
1275    pub fn log(self) -> Self {
1276        RankExpr::Logarithm(Box::new(self))
1277    }
1278
1279    /// Takes absolute value of the ranking expression.
1280    ///
1281    /// # Examples
1282    ///
1283    /// ```
1284    /// use chroma_types::operator::{RankExpr, QueryVector, Key};
1285    ///
1286    /// let knn1 = RankExpr::Knn {
1287    ///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1288    ///     key: Key::Embedding,
1289    ///     limit: 100,
1290    ///     default: None,
1291    ///     return_rank: false,
1292    /// };
1293    ///
1294    /// let knn2 = RankExpr::Knn {
1295    ///     query: QueryVector::Dense(vec![0.2, 0.3, 0.4]),
1296    ///     key: Key::field("other"),
1297    ///     limit: 100,
1298    ///     default: None,
1299    ///     return_rank: false,
1300    /// };
1301    ///
1302    /// // Absolute difference
1303    /// let diff = (knn1 - knn2).abs();
1304    /// ```
1305    pub fn abs(self) -> Self {
1306        RankExpr::Absolute(Box::new(self))
1307    }
1308
1309    /// Returns maximum of this expression and another.
1310    ///
1311    /// Can be chained to clamp scores to a maximum value.
1312    ///
1313    /// # Examples
1314    ///
1315    /// ```
1316    /// use chroma_types::operator::{RankExpr, QueryVector, Key};
1317    ///
1318    /// let knn = RankExpr::Knn {
1319    ///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1320    ///     key: Key::Embedding,
1321    ///     limit: 100,
1322    ///     default: None,
1323    ///     return_rank: false,
1324    /// };
1325    ///
1326    /// // Clamp to maximum of 1.0
1327    /// let clamped = knn.clone().max(1.0);
1328    ///
1329    /// // Clamp to range [0.0, 1.0]
1330    /// let range_clamped = knn.min(0.0).max(1.0);
1331    /// ```
1332    pub fn max(self, other: impl Into<RankExpr>) -> Self {
1333        let other = other.into();
1334
1335        match self {
1336            RankExpr::Maximum(mut exprs) => match other {
1337                RankExpr::Maximum(other_exprs) => {
1338                    exprs.extend(other_exprs);
1339                    RankExpr::Maximum(exprs)
1340                }
1341                _ => {
1342                    exprs.push(other);
1343                    RankExpr::Maximum(exprs)
1344                }
1345            },
1346            _ => match other {
1347                RankExpr::Maximum(mut exprs) => {
1348                    exprs.insert(0, self);
1349                    RankExpr::Maximum(exprs)
1350                }
1351                _ => RankExpr::Maximum(vec![self, other]),
1352            },
1353        }
1354    }
1355
1356    /// Returns minimum of this expression and another.
1357    ///
1358    /// Can be chained to clamp scores to a minimum value.
1359    ///
1360    /// # Examples
1361    ///
1362    /// ```
1363    /// use chroma_types::operator::{RankExpr, QueryVector, Key};
1364    ///
1365    /// let knn = RankExpr::Knn {
1366    ///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
1367    ///     key: Key::Embedding,
1368    ///     limit: 100,
1369    ///     default: None,
1370    ///     return_rank: false,
1371    /// };
1372    ///
1373    /// // Clamp to minimum of 0.0 (ensure non-negative)
1374    /// let clamped = knn.clone().min(0.0);
1375    ///
1376    /// // Clamp to range [0.0, 1.0]
1377    /// let range_clamped = knn.min(0.0).max(1.0);
1378    /// ```
1379    pub fn min(self, other: impl Into<RankExpr>) -> Self {
1380        let other = other.into();
1381
1382        match self {
1383            RankExpr::Minimum(mut exprs) => match other {
1384                RankExpr::Minimum(other_exprs) => {
1385                    exprs.extend(other_exprs);
1386                    RankExpr::Minimum(exprs)
1387                }
1388                _ => {
1389                    exprs.push(other);
1390                    RankExpr::Minimum(exprs)
1391                }
1392            },
1393            _ => match other {
1394                RankExpr::Minimum(mut exprs) => {
1395                    exprs.insert(0, self);
1396                    RankExpr::Minimum(exprs)
1397                }
1398                _ => RankExpr::Minimum(vec![self, other]),
1399            },
1400        }
1401    }
1402}
1403
1404impl Add for RankExpr {
1405    type Output = RankExpr;
1406
1407    fn add(self, rhs: Self) -> Self::Output {
1408        match self {
1409            RankExpr::Summation(mut exprs) => match rhs {
1410                RankExpr::Summation(rhs_exprs) => {
1411                    exprs.extend(rhs_exprs);
1412                    RankExpr::Summation(exprs)
1413                }
1414                _ => {
1415                    exprs.push(rhs);
1416                    RankExpr::Summation(exprs)
1417                }
1418            },
1419            _ => match rhs {
1420                RankExpr::Summation(mut exprs) => {
1421                    exprs.insert(0, self);
1422                    RankExpr::Summation(exprs)
1423                }
1424                _ => RankExpr::Summation(vec![self, rhs]),
1425            },
1426        }
1427    }
1428}
1429
1430impl Add<f32> for RankExpr {
1431    type Output = RankExpr;
1432
1433    fn add(self, rhs: f32) -> Self::Output {
1434        self + RankExpr::Value(rhs)
1435    }
1436}
1437
1438impl Add<RankExpr> for f32 {
1439    type Output = RankExpr;
1440
1441    fn add(self, rhs: RankExpr) -> Self::Output {
1442        RankExpr::Value(self) + rhs
1443    }
1444}
1445
1446impl Sub for RankExpr {
1447    type Output = RankExpr;
1448
1449    fn sub(self, rhs: Self) -> Self::Output {
1450        RankExpr::Subtraction {
1451            left: Box::new(self),
1452            right: Box::new(rhs),
1453        }
1454    }
1455}
1456
1457impl Sub<f32> for RankExpr {
1458    type Output = RankExpr;
1459
1460    fn sub(self, rhs: f32) -> Self::Output {
1461        self - RankExpr::Value(rhs)
1462    }
1463}
1464
1465impl Sub<RankExpr> for f32 {
1466    type Output = RankExpr;
1467
1468    fn sub(self, rhs: RankExpr) -> Self::Output {
1469        RankExpr::Value(self) - rhs
1470    }
1471}
1472
1473impl Mul for RankExpr {
1474    type Output = RankExpr;
1475
1476    fn mul(self, rhs: Self) -> Self::Output {
1477        match self {
1478            RankExpr::Multiplication(mut exprs) => match rhs {
1479                RankExpr::Multiplication(rhs_exprs) => {
1480                    exprs.extend(rhs_exprs);
1481                    RankExpr::Multiplication(exprs)
1482                }
1483                _ => {
1484                    exprs.push(rhs);
1485                    RankExpr::Multiplication(exprs)
1486                }
1487            },
1488            _ => match rhs {
1489                RankExpr::Multiplication(mut exprs) => {
1490                    exprs.insert(0, self);
1491                    RankExpr::Multiplication(exprs)
1492                }
1493                _ => RankExpr::Multiplication(vec![self, rhs]),
1494            },
1495        }
1496    }
1497}
1498
1499impl Mul<f32> for RankExpr {
1500    type Output = RankExpr;
1501
1502    fn mul(self, rhs: f32) -> Self::Output {
1503        self * RankExpr::Value(rhs)
1504    }
1505}
1506
1507impl Mul<RankExpr> for f32 {
1508    type Output = RankExpr;
1509
1510    fn mul(self, rhs: RankExpr) -> Self::Output {
1511        RankExpr::Value(self) * rhs
1512    }
1513}
1514
1515impl Div for RankExpr {
1516    type Output = RankExpr;
1517
1518    fn div(self, rhs: Self) -> Self::Output {
1519        RankExpr::Division {
1520            left: Box::new(self),
1521            right: Box::new(rhs),
1522        }
1523    }
1524}
1525
1526impl Div<f32> for RankExpr {
1527    type Output = RankExpr;
1528
1529    fn div(self, rhs: f32) -> Self::Output {
1530        self / RankExpr::Value(rhs)
1531    }
1532}
1533
1534impl Div<RankExpr> for f32 {
1535    type Output = RankExpr;
1536
1537    fn div(self, rhs: RankExpr) -> Self::Output {
1538        RankExpr::Value(self) / rhs
1539    }
1540}
1541
1542impl Neg for RankExpr {
1543    type Output = RankExpr;
1544
1545    fn neg(self) -> Self::Output {
1546        RankExpr::Value(-1.0) * self
1547    }
1548}
1549
1550impl From<f32> for RankExpr {
1551    fn from(v: f32) -> Self {
1552        RankExpr::Value(v)
1553    }
1554}
1555
1556impl TryFrom<chroma_proto::RankExpr> for RankExpr {
1557    type Error = QueryConversionError;
1558
1559    fn try_from(proto_expr: chroma_proto::RankExpr) -> Result<Self, Self::Error> {
1560        match proto_expr.rank {
1561            Some(chroma_proto::rank_expr::Rank::Absolute(expr)) => {
1562                Ok(RankExpr::Absolute(Box::new(RankExpr::try_from(*expr)?)))
1563            }
1564            Some(chroma_proto::rank_expr::Rank::Division(div)) => {
1565                let left = div.left.ok_or(QueryConversionError::field("left"))?;
1566                let right = div.right.ok_or(QueryConversionError::field("right"))?;
1567                Ok(RankExpr::Division {
1568                    left: Box::new(RankExpr::try_from(*left)?),
1569                    right: Box::new(RankExpr::try_from(*right)?),
1570                })
1571            }
1572            Some(chroma_proto::rank_expr::Rank::Exponentiation(expr)) => Ok(
1573                RankExpr::Exponentiation(Box::new(RankExpr::try_from(*expr)?)),
1574            ),
1575            Some(chroma_proto::rank_expr::Rank::Knn(knn)) => {
1576                let query = knn
1577                    .query
1578                    .ok_or(QueryConversionError::field("query"))?
1579                    .try_into()?;
1580                Ok(RankExpr::Knn {
1581                    query,
1582                    key: Key::from(knn.key),
1583                    limit: knn.limit,
1584                    default: knn.default,
1585                    return_rank: knn.return_rank,
1586                })
1587            }
1588            Some(chroma_proto::rank_expr::Rank::Logarithm(expr)) => {
1589                Ok(RankExpr::Logarithm(Box::new(RankExpr::try_from(*expr)?)))
1590            }
1591            Some(chroma_proto::rank_expr::Rank::Maximum(max)) => {
1592                let exprs = max
1593                    .exprs
1594                    .into_iter()
1595                    .map(RankExpr::try_from)
1596                    .collect::<Result<Vec<_>, _>>()?;
1597                Ok(RankExpr::Maximum(exprs))
1598            }
1599            Some(chroma_proto::rank_expr::Rank::Minimum(min)) => {
1600                let exprs = min
1601                    .exprs
1602                    .into_iter()
1603                    .map(RankExpr::try_from)
1604                    .collect::<Result<Vec<_>, _>>()?;
1605                Ok(RankExpr::Minimum(exprs))
1606            }
1607            Some(chroma_proto::rank_expr::Rank::Multiplication(mul)) => {
1608                let exprs = mul
1609                    .exprs
1610                    .into_iter()
1611                    .map(RankExpr::try_from)
1612                    .collect::<Result<Vec<_>, _>>()?;
1613                Ok(RankExpr::Multiplication(exprs))
1614            }
1615            Some(chroma_proto::rank_expr::Rank::Subtraction(sub)) => {
1616                let left = sub.left.ok_or(QueryConversionError::field("left"))?;
1617                let right = sub.right.ok_or(QueryConversionError::field("right"))?;
1618                Ok(RankExpr::Subtraction {
1619                    left: Box::new(RankExpr::try_from(*left)?),
1620                    right: Box::new(RankExpr::try_from(*right)?),
1621                })
1622            }
1623            Some(chroma_proto::rank_expr::Rank::Summation(sum)) => {
1624                let exprs = sum
1625                    .exprs
1626                    .into_iter()
1627                    .map(RankExpr::try_from)
1628                    .collect::<Result<Vec<_>, _>>()?;
1629                Ok(RankExpr::Summation(exprs))
1630            }
1631            Some(chroma_proto::rank_expr::Rank::Value(value)) => Ok(RankExpr::Value(value)),
1632            None => Err(QueryConversionError::field("rank")),
1633        }
1634    }
1635}
1636
1637impl TryFrom<RankExpr> for chroma_proto::RankExpr {
1638    type Error = QueryConversionError;
1639
1640    fn try_from(rank_expr: RankExpr) -> Result<Self, Self::Error> {
1641        let proto_rank = match rank_expr {
1642            RankExpr::Absolute(expr) => chroma_proto::rank_expr::Rank::Absolute(Box::new(
1643                chroma_proto::RankExpr::try_from(*expr)?,
1644            )),
1645            RankExpr::Division { left, right } => chroma_proto::rank_expr::Rank::Division(
1646                Box::new(chroma_proto::rank_expr::RankPair {
1647                    left: Some(Box::new(chroma_proto::RankExpr::try_from(*left)?)),
1648                    right: Some(Box::new(chroma_proto::RankExpr::try_from(*right)?)),
1649                }),
1650            ),
1651            RankExpr::Exponentiation(expr) => chroma_proto::rank_expr::Rank::Exponentiation(
1652                Box::new(chroma_proto::RankExpr::try_from(*expr)?),
1653            ),
1654            RankExpr::Knn {
1655                query,
1656                key,
1657                limit,
1658                default,
1659                return_rank,
1660            } => chroma_proto::rank_expr::Rank::Knn(chroma_proto::rank_expr::Knn {
1661                query: Some(query.try_into()?),
1662                key: key.to_string(),
1663                limit,
1664                default,
1665                return_rank,
1666            }),
1667            RankExpr::Logarithm(expr) => chroma_proto::rank_expr::Rank::Logarithm(Box::new(
1668                chroma_proto::RankExpr::try_from(*expr)?,
1669            )),
1670            RankExpr::Maximum(exprs) => {
1671                let proto_exprs = exprs
1672                    .into_iter()
1673                    .map(chroma_proto::RankExpr::try_from)
1674                    .collect::<Result<Vec<_>, _>>()?;
1675                chroma_proto::rank_expr::Rank::Maximum(chroma_proto::rank_expr::RankList {
1676                    exprs: proto_exprs,
1677                })
1678            }
1679            RankExpr::Minimum(exprs) => {
1680                let proto_exprs = exprs
1681                    .into_iter()
1682                    .map(chroma_proto::RankExpr::try_from)
1683                    .collect::<Result<Vec<_>, _>>()?;
1684                chroma_proto::rank_expr::Rank::Minimum(chroma_proto::rank_expr::RankList {
1685                    exprs: proto_exprs,
1686                })
1687            }
1688            RankExpr::Multiplication(exprs) => {
1689                let proto_exprs = exprs
1690                    .into_iter()
1691                    .map(chroma_proto::RankExpr::try_from)
1692                    .collect::<Result<Vec<_>, _>>()?;
1693                chroma_proto::rank_expr::Rank::Multiplication(chroma_proto::rank_expr::RankList {
1694                    exprs: proto_exprs,
1695                })
1696            }
1697            RankExpr::Subtraction { left, right } => chroma_proto::rank_expr::Rank::Subtraction(
1698                Box::new(chroma_proto::rank_expr::RankPair {
1699                    left: Some(Box::new(chroma_proto::RankExpr::try_from(*left)?)),
1700                    right: Some(Box::new(chroma_proto::RankExpr::try_from(*right)?)),
1701                }),
1702            ),
1703            RankExpr::Summation(exprs) => {
1704                let proto_exprs = exprs
1705                    .into_iter()
1706                    .map(chroma_proto::RankExpr::try_from)
1707                    .collect::<Result<Vec<_>, _>>()?;
1708                chroma_proto::rank_expr::Rank::Summation(chroma_proto::rank_expr::RankList {
1709                    exprs: proto_exprs,
1710                })
1711            }
1712            RankExpr::Value(value) => chroma_proto::rank_expr::Rank::Value(value),
1713        };
1714
1715        Ok(chroma_proto::RankExpr {
1716            rank: Some(proto_rank),
1717        })
1718    }
1719}
1720
1721/// Represents a field key in search queries.
1722///
1723/// Used for both selecting fields to return and building filter expressions.
1724/// Predefined keys access special fields, while custom keys access metadata.
1725///
1726/// # Predefined Keys
1727///
1728/// - `Key::Document` - Document text content (`#document`)
1729/// - `Key::Embedding` - Vector embeddings (`#embedding`)
1730/// - `Key::Metadata` - All metadata fields (`#metadata`)
1731/// - `Key::Score` - Search scores (`#score`)
1732///
1733/// # Custom Keys
1734///
1735/// Use `Key::field()` or `Key::from()` to reference metadata fields:
1736///
1737/// ```
1738/// use chroma_types::operator::Key;
1739///
1740/// let key = Key::field("author");
1741/// let key = Key::from("title");
1742/// ```
1743///
1744/// # Examples
1745///
1746/// ## Building filters
1747///
1748/// ```
1749/// use chroma_types::operator::Key;
1750///
1751/// // Equality
1752/// let filter = Key::field("status").eq("published");
1753///
1754/// // Comparisons
1755/// let filter = Key::field("year").gte(2020);
1756/// let filter = Key::field("score").lt(0.9);
1757///
1758/// // Set operations
1759/// let filter = Key::field("category").is_in(vec!["tech", "science"]);
1760/// let filter = Key::field("status").not_in(vec!["deleted", "archived"]);
1761///
1762/// // Document content
1763/// let filter = Key::Document.contains("machine learning");
1764/// let filter = Key::Document.regex(r"\bAPI\b");
1765///
1766/// // Combining filters
1767/// let filter = Key::field("status").eq("published")
1768///     & Key::field("year").gte(2020);
1769/// ```
1770///
1771/// ## Selecting fields
1772///
1773/// ```
1774/// use chroma_types::plan::SearchPayload;
1775/// use chroma_types::operator::Key;
1776///
1777/// let search = SearchPayload::default()
1778///     .select([
1779///         Key::Document,
1780///         Key::Score,
1781///         Key::field("title"),
1782///         Key::field("author"),
1783///     ]);
1784/// ```
1785#[derive(Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
1786#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
1787pub enum Key {
1788    // Predefined keys
1789    Document,
1790    Embedding,
1791    Metadata,
1792    Score,
1793    MetadataField(String),
1794}
1795
1796impl Serialize for Key {
1797    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
1798    where
1799        S: serde::Serializer,
1800    {
1801        match self {
1802            Key::Document => serializer.serialize_str("#document"),
1803            Key::Embedding => serializer.serialize_str("#embedding"),
1804            Key::Metadata => serializer.serialize_str("#metadata"),
1805            Key::Score => serializer.serialize_str("#score"),
1806            Key::MetadataField(field) => serializer.serialize_str(field),
1807        }
1808    }
1809}
1810
1811impl<'de> Deserialize<'de> for Key {
1812    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
1813    where
1814        D: Deserializer<'de>,
1815    {
1816        let s = String::deserialize(deserializer)?;
1817        Ok(Key::from(s))
1818    }
1819}
1820
1821impl fmt::Display for Key {
1822    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1823        match self {
1824            Key::Document => write!(f, "#document"),
1825            Key::Embedding => write!(f, "#embedding"),
1826            Key::Metadata => write!(f, "#metadata"),
1827            Key::Score => write!(f, "#score"),
1828            Key::MetadataField(field) => write!(f, "{}", field),
1829        }
1830    }
1831}
1832
1833impl From<&str> for Key {
1834    fn from(s: &str) -> Self {
1835        match s {
1836            "#document" => Key::Document,
1837            "#embedding" => Key::Embedding,
1838            "#metadata" => Key::Metadata,
1839            "#score" => Key::Score,
1840            // Any other string is treated as a metadata field key
1841            field => Key::MetadataField(field.to_string()),
1842        }
1843    }
1844}
1845
1846impl From<String> for Key {
1847    fn from(s: String) -> Self {
1848        Key::from(s.as_str())
1849    }
1850}
1851
1852impl Key {
1853    /// Creates a Key for a custom metadata field.
1854    ///
1855    /// # Examples
1856    ///
1857    /// ```
1858    /// use chroma_types::operator::Key;
1859    ///
1860    /// let status = Key::field("status");
1861    /// let year = Key::field("year");
1862    /// let author = Key::field("author");
1863    /// ```
1864    pub fn field(name: impl Into<String>) -> Self {
1865        Key::MetadataField(name.into())
1866    }
1867
1868    /// Creates an equality filter: `field == value`.
1869    ///
1870    /// # Examples
1871    ///
1872    /// ```
1873    /// use chroma_types::operator::Key;
1874    ///
1875    /// // String equality
1876    /// let filter = Key::field("status").eq("published");
1877    ///
1878    /// // Numeric equality
1879    /// let filter = Key::field("count").eq(42);
1880    ///
1881    /// // Boolean equality
1882    /// let filter = Key::field("featured").eq(true);
1883    /// ```
1884    pub fn eq<T: Into<MetadataValue>>(self, value: T) -> Where {
1885        Where::Metadata(MetadataExpression {
1886            key: self.to_string(),
1887            comparison: MetadataComparison::Primitive(PrimitiveOperator::Equal, value.into()),
1888        })
1889    }
1890
1891    /// Creates an inequality filter: `field != value`.
1892    ///
1893    /// # Examples
1894    ///
1895    /// ```
1896    /// use chroma_types::operator::Key;
1897    ///
1898    /// let filter = Key::field("status").ne("deleted");
1899    /// let filter = Key::field("count").ne(0);
1900    /// ```
1901    pub fn ne<T: Into<MetadataValue>>(self, value: T) -> Where {
1902        Where::Metadata(MetadataExpression {
1903            key: self.to_string(),
1904            comparison: MetadataComparison::Primitive(PrimitiveOperator::NotEqual, value.into()),
1905        })
1906    }
1907
1908    /// Creates a greater-than filter: `field > value` (numeric only).
1909    ///
1910    /// # Examples
1911    ///
1912    /// ```
1913    /// use chroma_types::operator::Key;
1914    ///
1915    /// let filter = Key::field("score").gt(0.5);
1916    /// let filter = Key::field("year").gt(2020);
1917    /// ```
1918    pub fn gt<T: Into<MetadataValue>>(self, value: T) -> Where {
1919        Where::Metadata(MetadataExpression {
1920            key: self.to_string(),
1921            comparison: MetadataComparison::Primitive(PrimitiveOperator::GreaterThan, value.into()),
1922        })
1923    }
1924
1925    /// Creates a greater-than-or-equal filter: `field >= value` (numeric only).
1926    ///
1927    /// # Examples
1928    ///
1929    /// ```
1930    /// use chroma_types::operator::Key;
1931    ///
1932    /// let filter = Key::field("score").gte(0.5);
1933    /// let filter = Key::field("year").gte(2020);
1934    /// ```
1935    pub fn gte<T: Into<MetadataValue>>(self, value: T) -> Where {
1936        Where::Metadata(MetadataExpression {
1937            key: self.to_string(),
1938            comparison: MetadataComparison::Primitive(
1939                PrimitiveOperator::GreaterThanOrEqual,
1940                value.into(),
1941            ),
1942        })
1943    }
1944
1945    /// Creates a less-than filter: `field < value` (numeric only).
1946    ///
1947    /// # Examples
1948    ///
1949    /// ```
1950    /// use chroma_types::operator::Key;
1951    ///
1952    /// let filter = Key::field("score").lt(0.9);
1953    /// let filter = Key::field("year").lt(2025);
1954    /// ```
1955    pub fn lt<T: Into<MetadataValue>>(self, value: T) -> Where {
1956        Where::Metadata(MetadataExpression {
1957            key: self.to_string(),
1958            comparison: MetadataComparison::Primitive(PrimitiveOperator::LessThan, value.into()),
1959        })
1960    }
1961
1962    /// Creates a less-than-or-equal filter: `field <= value` (numeric only).
1963    ///
1964    /// # Examples
1965    ///
1966    /// ```
1967    /// use chroma_types::operator::Key;
1968    ///
1969    /// let filter = Key::field("score").lte(0.9);
1970    /// let filter = Key::field("year").lte(2024);
1971    /// ```
1972    pub fn lte<T: Into<MetadataValue>>(self, value: T) -> Where {
1973        Where::Metadata(MetadataExpression {
1974            key: self.to_string(),
1975            comparison: MetadataComparison::Primitive(
1976                PrimitiveOperator::LessThanOrEqual,
1977                value.into(),
1978            ),
1979        })
1980    }
1981
1982    /// Creates a set membership filter: `field IN values`.
1983    ///
1984    /// Accepts any iterator (Vec, array, slice, etc.).
1985    ///
1986    /// # Examples
1987    ///
1988    /// ```
1989    /// use chroma_types::operator::Key;
1990    ///
1991    /// // With Vec
1992    /// let filter = Key::field("year").is_in(vec![2023, 2024, 2025]);
1993    ///
1994    /// // With array
1995    /// let filter = Key::field("category").is_in(["tech", "science", "math"]);
1996    ///
1997    /// // With owned strings
1998    /// let categories = vec!["tech".to_string(), "science".to_string()];
1999    /// let filter = Key::field("category").is_in(categories);
2000    /// ```
2001    pub fn is_in<I, T>(self, values: I) -> Where
2002    where
2003        I: IntoIterator<Item = T>,
2004        Vec<T>: Into<MetadataSetValue>,
2005    {
2006        let vec: Vec<T> = values.into_iter().collect();
2007        Where::Metadata(MetadataExpression {
2008            key: self.to_string(),
2009            comparison: MetadataComparison::Set(SetOperator::In, vec.into()),
2010        })
2011    }
2012
2013    /// Creates a set exclusion filter: `field NOT IN values`.
2014    ///
2015    /// Accepts any iterator (Vec, array, slice, etc.).
2016    ///
2017    /// # Examples
2018    ///
2019    /// ```
2020    /// use chroma_types::operator::Key;
2021    ///
2022    /// // Exclude deleted and archived
2023    /// let filter = Key::field("status").not_in(vec!["deleted", "archived"]);
2024    ///
2025    /// // Exclude specific years
2026    /// let filter = Key::field("year").not_in(vec![2019, 2020]);
2027    /// ```
2028    pub fn not_in<I, T>(self, values: I) -> Where
2029    where
2030        I: IntoIterator<Item = T>,
2031        Vec<T>: Into<MetadataSetValue>,
2032    {
2033        let vec: Vec<T> = values.into_iter().collect();
2034        Where::Metadata(MetadataExpression {
2035            key: self.to_string(),
2036            comparison: MetadataComparison::Set(SetOperator::NotIn, vec.into()),
2037        })
2038    }
2039
2040    /// Creates a document substring filter (case-sensitive).
2041    ///
2042    /// Only valid on `Key::Document`. Pattern must have at least 3 literal
2043    /// characters for accurate results.
2044    ///
2045    /// For metadata array contains, use [`contains_value`](Key::contains_value).
2046    ///
2047    /// # Examples
2048    ///
2049    /// ```
2050    /// use chroma_types::operator::Key;
2051    ///
2052    /// let filter = Key::Document.contains("machine learning");
2053    /// let filter = Key::Document.contains("API");
2054    /// ```
2055    pub fn contains<S: Into<String>>(self, text: S) -> Where {
2056        Where::Document(DocumentExpression {
2057            operator: DocumentOperator::Contains,
2058            pattern: text.into(),
2059        })
2060    }
2061
2062    /// Creates a negative document substring filter (case-sensitive).
2063    ///
2064    /// Only valid on `Key::Document`.
2065    ///
2066    /// For metadata array not-contains, use
2067    /// [`not_contains_value`](Key::not_contains_value).
2068    ///
2069    /// # Examples
2070    ///
2071    /// ```
2072    /// use chroma_types::operator::Key;
2073    ///
2074    /// let filter = Key::Document.not_contains("deprecated");
2075    /// let filter = Key::Document.not_contains("beta");
2076    /// ```
2077    pub fn not_contains<S: Into<String>>(self, text: S) -> Where {
2078        Where::Document(DocumentExpression {
2079            operator: DocumentOperator::NotContains,
2080            pattern: text.into(),
2081        })
2082    }
2083
2084    /// Checks whether a metadata array field contains the given scalar value.
2085    ///
2086    /// # Examples
2087    ///
2088    /// ```
2089    /// use chroma_types::operator::Key;
2090    ///
2091    /// let filter = Key::field("tags").contains_value("action");
2092    /// let filter = Key::field("scores").contains_value(42);
2093    /// let filter = Key::field("ratings").contains_value(4.5);
2094    /// let filter = Key::field("flags").contains_value(true);
2095    /// ```
2096    pub fn contains_value<T: Into<MetadataValue>>(self, value: T) -> Where {
2097        Where::Metadata(MetadataExpression {
2098            key: self.to_string(),
2099            comparison: MetadataComparison::ArrayContains(ContainsOperator::Contains, value.into()),
2100        })
2101    }
2102
2103    /// Checks that a metadata array field does **not** contain the given scalar
2104    /// value.
2105    ///
2106    /// # Examples
2107    ///
2108    /// ```
2109    /// use chroma_types::operator::Key;
2110    ///
2111    /// let filter = Key::field("tags").not_contains_value("draft");
2112    /// let filter = Key::field("scores").not_contains_value(0);
2113    /// ```
2114    pub fn not_contains_value<T: Into<MetadataValue>>(self, value: T) -> Where {
2115        Where::Metadata(MetadataExpression {
2116            key: self.to_string(),
2117            comparison: MetadataComparison::ArrayContains(
2118                ContainsOperator::NotContains,
2119                value.into(),
2120            ),
2121        })
2122    }
2123
2124    /// Creates a regex filter (case-sensitive, document content only).
2125    ///
2126    /// Note: Currently only works with `Key::Document`. Pattern must have at least
2127    /// 3 literal characters for accurate results.
2128    ///
2129    /// # Examples
2130    ///
2131    /// ```
2132    /// use chroma_types::operator::Key;
2133    ///
2134    /// // Match whole word "API"
2135    /// let filter = Key::Document.regex(r"\bAPI\b");
2136    ///
2137    /// // Match version pattern
2138    /// let filter = Key::Document.regex(r"v\d+\.\d+\.\d+");
2139    /// ```
2140    pub fn regex<S: Into<String>>(self, pattern: S) -> Where {
2141        Where::Document(DocumentExpression {
2142            operator: DocumentOperator::Regex,
2143            pattern: pattern.into(),
2144        })
2145    }
2146
2147    /// Creates a negative regex filter (case-sensitive, document content only).
2148    ///
2149    /// Note: Currently only works with `Key::Document`.
2150    ///
2151    /// # Examples
2152    ///
2153    /// ```
2154    /// use chroma_types::operator::Key;
2155    ///
2156    /// // Exclude beta versions
2157    /// let filter = Key::Document.not_regex(r"beta");
2158    ///
2159    /// // Exclude test documents
2160    /// let filter = Key::Document.not_regex(r"\btest\b");
2161    /// ```
2162    pub fn not_regex<S: Into<String>>(self, pattern: S) -> Where {
2163        Where::Document(DocumentExpression {
2164            operator: DocumentOperator::NotRegex,
2165            pattern: pattern.into(),
2166        })
2167    }
2168}
2169
2170/// Field selection for search results.
2171///
2172/// Specifies which fields to include in the results. IDs are always included.
2173///
2174/// # Fields
2175///
2176/// * `keys` - Set of keys to include in results
2177///
2178/// # Available Keys
2179///
2180/// * `Key::Document` - Document text content
2181/// * `Key::Embedding` - Vector embeddings
2182/// * `Key::Metadata` - All metadata fields
2183/// * `Key::Score` - Search scores
2184/// * `Key::field("name")` - Specific metadata field
2185///
2186/// # Performance
2187///
2188/// Selecting fewer fields improves performance by reducing data transfer:
2189/// - Minimal: IDs only (default, fastest)
2190/// - Moderate: Scores + specific metadata fields
2191/// - Heavy: Documents + embeddings (larger payloads)
2192///
2193/// # Examples
2194///
2195/// ```
2196/// use chroma_types::operator::{Select, Key};
2197/// use std::collections::HashSet;
2198///
2199/// // Select predefined fields
2200/// let select = Select {
2201///     keys: [Key::Document, Key::Score].into_iter().collect(),
2202/// };
2203///
2204/// // Select specific metadata fields
2205/// let select = Select {
2206///     keys: [
2207///         Key::field("title"),
2208///         Key::field("author"),
2209///         Key::Score,
2210///     ].into_iter().collect(),
2211/// };
2212///
2213/// // Select everything
2214/// let select = Select {
2215///     keys: [
2216///         Key::Document,
2217///         Key::Embedding,
2218///         Key::Metadata,
2219///         Key::Score,
2220///     ].into_iter().collect(),
2221/// };
2222/// ```
2223#[derive(Clone, Debug, Default, Deserialize, Serialize)]
2224pub struct Select {
2225    #[serde(default)]
2226    pub keys: HashSet<Key>,
2227}
2228
2229impl TryFrom<chroma_proto::SelectOperator> for Select {
2230    type Error = QueryConversionError;
2231
2232    fn try_from(value: chroma_proto::SelectOperator) -> Result<Self, Self::Error> {
2233        let keys = value
2234            .keys
2235            .into_iter()
2236            .map(|key| {
2237                // Try to deserialize each string as a Key
2238                serde_json::from_value(serde_json::Value::String(key))
2239                    .map_err(|_| QueryConversionError::field("keys"))
2240            })
2241            .collect::<Result<HashSet<_>, _>>()?;
2242
2243        Ok(Self { keys })
2244    }
2245}
2246
2247impl TryFrom<Select> for chroma_proto::SelectOperator {
2248    type Error = QueryConversionError;
2249
2250    fn try_from(value: Select) -> Result<Self, Self::Error> {
2251        let keys = value
2252            .keys
2253            .into_iter()
2254            .map(|key| {
2255                // Serialize each Key back to string
2256                serde_json::to_value(&key)
2257                    .ok()
2258                    .and_then(|v| v.as_str().map(String::from))
2259                    .ok_or(QueryConversionError::field("keys"))
2260            })
2261            .collect::<Result<Vec<_>, _>>()?;
2262
2263        Ok(Self { keys })
2264    }
2265}
2266
2267/// Aggregation function applied within each group.
2268///
2269/// Determines which records to keep from each group and their ordering.
2270///
2271/// # Variants
2272///
2273/// * `MinK` - Returns k records with minimum values (ascending order).
2274///   Use with `Key::Score` to get best matches (lower score = better in Chroma).
2275/// * `MaxK` - Returns k records with maximum values (descending order).
2276///
2277/// # Multi-level Ordering
2278///
2279/// The `keys` field supports multi-level ordering. Records are sorted by
2280/// the first key, then by the second key for ties, and so on.
2281///
2282/// # Examples
2283///
2284/// ```
2285/// use chroma_types::operator::{Aggregate, Key};
2286///
2287/// // Best 3 by score per group
2288/// let agg = Aggregate::MinK {
2289///     keys: vec![Key::Score],
2290///     k: 3,
2291/// };
2292///
2293/// // Best 3 by score, then by date for ties
2294/// let agg = Aggregate::MinK {
2295///     keys: vec![Key::Score, Key::field("date")],
2296///     k: 3,
2297/// };
2298///
2299/// // Top 5 by recency (highest date first)
2300/// let agg = Aggregate::MaxK {
2301///     keys: vec![Key::field("date")],
2302///     k: 5,
2303/// };
2304/// ```
2305#[derive(Clone, Debug, Deserialize, Serialize)]
2306pub enum Aggregate {
2307    /// Returns k records with minimum values (ascending order)
2308    #[serde(rename = "$min_k")]
2309    MinK {
2310        /// Keys for multi-level ordering
2311        keys: Vec<Key>,
2312        /// Number of records to return per group
2313        k: u32,
2314    },
2315    /// Returns k records with maximum values (descending order)
2316    #[serde(rename = "$max_k")]
2317    MaxK {
2318        /// Keys for multi-level ordering
2319        keys: Vec<Key>,
2320        /// Number of records to return per group
2321        k: u32,
2322    },
2323}
2324
2325/// Groups results by metadata keys and aggregates within each group.
2326///
2327/// Results are grouped by the specified metadata keys (like SQL GROUP BY),
2328/// then aggregated within each group using MinK or MaxK ordering.
2329/// The final output is flattened and sorted by score.
2330///
2331/// # Fields
2332///
2333/// * `keys` - Metadata keys to group by (composite grouping)
2334/// * `aggregate` - Aggregation function to apply within each group
2335///
2336/// # Behavior
2337///
2338/// * Missing metadata keys are treated as Null (forming their own group)
2339/// * Empty groups are omitted from results
2340/// * Final output is flattened (group structure not preserved)
2341/// * Results are sorted by score after aggregation
2342///
2343/// # Examples
2344///
2345/// ```
2346/// use chroma_types::operator::{GroupBy, Aggregate, Key};
2347///
2348/// // Top 3 documents per category
2349/// let group_by = GroupBy {
2350///     keys: vec![Key::field("category")],
2351///     aggregate: Some(Aggregate::MinK {
2352///         keys: vec![Key::Score],
2353///         k: 3,
2354///     }),
2355/// };
2356///
2357/// // Top 2 per (category, author) combination
2358/// let group_by = GroupBy {
2359///     keys: vec![Key::field("category"), Key::field("author")],
2360///     aggregate: Some(Aggregate::MinK {
2361///         keys: vec![Key::Score, Key::field("date")],
2362///         k: 2,
2363///     }),
2364/// };
2365/// ```
2366#[derive(Clone, Debug, Default, Deserialize, Serialize)]
2367pub struct GroupBy {
2368    /// Metadata keys to group by
2369    #[serde(default)]
2370    pub keys: Vec<Key>,
2371    /// Aggregation to apply within each group (required when keys is non-empty)
2372    #[serde(default)]
2373    pub aggregate: Option<Aggregate>,
2374}
2375
2376impl GroupBy {
2377    /// Returns true when this GroupBy has both keys and an aggregate,
2378    /// meaning it will actually perform grouping.
2379    pub fn is_active(&self) -> bool {
2380        !self.keys.is_empty() && self.aggregate.is_some()
2381    }
2382
2383    /// Returns the sort keys from the aggregate (`MinK`/`MaxK`),
2384    /// or an empty slice if no aggregate is set.
2385    pub fn aggregate_keys(&self) -> &[Key] {
2386        match &self.aggregate {
2387            Some(Aggregate::MinK { keys, .. } | Aggregate::MaxK { keys, .. }) => keys,
2388            None => &[],
2389        }
2390    }
2391
2392    /// Returns all distinct metadata `Key`s referenced by this GroupBy
2393    /// (from both grouping keys and aggregate sort keys).
2394    pub fn metadata_keys(&self) -> Vec<Key> {
2395        let mut result: Vec<Key> = self
2396            .keys
2397            .iter()
2398            .filter(|k| matches!(k, Key::MetadataField(_)))
2399            .cloned()
2400            .collect();
2401        for k in self.aggregate_keys() {
2402            if matches!(k, Key::MetadataField(_)) && !result.contains(k) {
2403                result.push(k.clone());
2404            }
2405        }
2406        result
2407    }
2408}
2409
2410impl TryFrom<chroma_proto::Aggregate> for Aggregate {
2411    type Error = QueryConversionError;
2412
2413    fn try_from(value: chroma_proto::Aggregate) -> Result<Self, Self::Error> {
2414        match value
2415            .aggregate
2416            .ok_or(QueryConversionError::field("aggregate"))?
2417        {
2418            chroma_proto::aggregate::Aggregate::MinK(min_k) => {
2419                let keys = min_k.keys.into_iter().map(Key::from).collect();
2420                Ok(Aggregate::MinK { keys, k: min_k.k })
2421            }
2422            chroma_proto::aggregate::Aggregate::MaxK(max_k) => {
2423                let keys = max_k.keys.into_iter().map(Key::from).collect();
2424                Ok(Aggregate::MaxK { keys, k: max_k.k })
2425            }
2426        }
2427    }
2428}
2429
2430impl From<Aggregate> for chroma_proto::Aggregate {
2431    fn from(value: Aggregate) -> Self {
2432        let aggregate = match value {
2433            Aggregate::MinK { keys, k } => {
2434                chroma_proto::aggregate::Aggregate::MinK(chroma_proto::aggregate::MinK {
2435                    keys: keys.into_iter().map(|k| k.to_string()).collect(),
2436                    k,
2437                })
2438            }
2439            Aggregate::MaxK { keys, k } => {
2440                chroma_proto::aggregate::Aggregate::MaxK(chroma_proto::aggregate::MaxK {
2441                    keys: keys.into_iter().map(|k| k.to_string()).collect(),
2442                    k,
2443                })
2444            }
2445        };
2446
2447        chroma_proto::Aggregate {
2448            aggregate: Some(aggregate),
2449        }
2450    }
2451}
2452
2453impl TryFrom<chroma_proto::GroupByOperator> for GroupBy {
2454    type Error = QueryConversionError;
2455
2456    fn try_from(value: chroma_proto::GroupByOperator) -> Result<Self, Self::Error> {
2457        let keys = value.keys.into_iter().map(Key::from).collect();
2458        let aggregate = value.aggregate.map(TryInto::try_into).transpose()?;
2459
2460        Ok(Self { keys, aggregate })
2461    }
2462}
2463
2464impl TryFrom<GroupBy> for chroma_proto::GroupByOperator {
2465    type Error = QueryConversionError;
2466
2467    fn try_from(value: GroupBy) -> Result<Self, Self::Error> {
2468        let keys = value.keys.into_iter().map(|k| k.to_string()).collect();
2469        let aggregate = value.aggregate.map(Into::into);
2470
2471        Ok(Self { keys, aggregate })
2472    }
2473}
2474
2475/// A single search result record.
2476///
2477/// Contains the document ID and optionally document content, embeddings, metadata,
2478/// and search score based on what was selected in the search query.
2479///
2480/// # Fields
2481///
2482/// * `id` - Document ID (always present)
2483/// * `document` - Document text content (if selected)
2484/// * `embedding` - Vector embedding (if selected)
2485/// * `metadata` - Document metadata (if selected)
2486/// * `score` - Search score (present when ranking is used, lower = better match)
2487///
2488/// # Examples
2489///
2490/// ```
2491/// use chroma_types::operator::SearchRecord;
2492///
2493/// fn process_results(records: Vec<SearchRecord>) {
2494///     for record in records {
2495///         println!("ID: {}", record.id);
2496///
2497///         if let Some(score) = record.score {
2498///             println!("  Score: {:.3}", score);
2499///         }
2500///
2501///         if let Some(doc) = record.document {
2502///             println!("  Document: {}", doc);
2503///         }
2504///
2505///         if let Some(meta) = record.metadata {
2506///             println!("  Metadata: {:?}", meta);
2507///         }
2508///     }
2509/// }
2510/// ```
2511#[derive(Clone, Debug, Deserialize, Serialize)]
2512#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
2513pub struct SearchRecord {
2514    pub id: String,
2515    pub document: Option<String>,
2516    pub embedding: Option<Vec<f32>>,
2517    pub metadata: Option<Metadata>,
2518    pub score: Option<f32>,
2519}
2520
2521impl TryFrom<chroma_proto::SearchRecord> for SearchRecord {
2522    type Error = QueryConversionError;
2523
2524    fn try_from(value: chroma_proto::SearchRecord) -> Result<Self, Self::Error> {
2525        Ok(Self {
2526            id: value.id,
2527            document: value.document,
2528            embedding: value
2529                .embedding
2530                .map(|vec| vec.try_into().map(|(v, _)| v))
2531                .transpose()?,
2532            metadata: value.metadata.map(TryInto::try_into).transpose()?,
2533            score: value.score,
2534        })
2535    }
2536}
2537
2538impl TryFrom<SearchRecord> for chroma_proto::SearchRecord {
2539    type Error = QueryConversionError;
2540
2541    fn try_from(value: SearchRecord) -> Result<Self, Self::Error> {
2542        Ok(Self {
2543            id: value.id,
2544            document: value.document,
2545            embedding: value
2546                .embedding
2547                .map(|embedding| {
2548                    let embedding_dimension = embedding.len();
2549                    chroma_proto::Vector::try_from((
2550                        embedding,
2551                        ScalarEncoding::FLOAT32,
2552                        embedding_dimension,
2553                    ))
2554                })
2555                .transpose()?,
2556            metadata: value.metadata.map(Into::into),
2557            score: value.score,
2558        })
2559    }
2560}
2561
2562/// Results for a single search payload.
2563///
2564/// Contains all matching records for one search query.
2565///
2566/// # Fields
2567///
2568/// * `records` - Vector of search records, ordered by score (ascending)
2569///
2570/// # Examples
2571///
2572/// ```
2573/// use chroma_types::operator::{SearchPayloadResult, SearchRecord};
2574///
2575/// fn process_search_result(result: SearchPayloadResult) {
2576///     println!("Found {} results", result.records.len());
2577///
2578///     for (i, record) in result.records.iter().enumerate() {
2579///         println!("{}. {} (score: {:?})", i + 1, record.id, record.score);
2580///     }
2581/// }
2582/// ```
2583#[derive(Clone, Debug, Default)]
2584pub struct SearchPayloadResult {
2585    pub records: Vec<SearchRecord>,
2586}
2587
2588impl TryFrom<chroma_proto::SearchPayloadResult> for SearchPayloadResult {
2589    type Error = QueryConversionError;
2590
2591    fn try_from(value: chroma_proto::SearchPayloadResult) -> Result<Self, Self::Error> {
2592        Ok(Self {
2593            records: value
2594                .records
2595                .into_iter()
2596                .map(TryInto::try_into)
2597                .collect::<Result<_, _>>()?,
2598        })
2599    }
2600}
2601
2602impl TryFrom<SearchPayloadResult> for chroma_proto::SearchPayloadResult {
2603    type Error = QueryConversionError;
2604
2605    fn try_from(value: SearchPayloadResult) -> Result<Self, Self::Error> {
2606        Ok(Self {
2607            records: value
2608                .records
2609                .into_iter()
2610                .map(TryInto::try_into)
2611                .collect::<Result<Vec<_>, _>>()?,
2612        })
2613    }
2614}
2615
2616/// Results from a batch search operation.
2617///
2618/// Contains results for each search payload in the batch, maintaining the same order
2619/// as the input searches.
2620///
2621/// # Fields
2622///
2623/// * `results` - Results for each search payload (indexed by search position)
2624/// * `pulled_log_bytes` - Total bytes pulled from log (for internal metrics)
2625///
2626/// # Examples
2627///
2628/// ## Single search
2629///
2630/// ```
2631/// use chroma_types::operator::SearchResult;
2632///
2633/// fn process_single_search(result: SearchResult) {
2634///     // Single search, so results[0] contains our records
2635///     let records = &result.results[0].records;
2636///
2637///     for record in records {
2638///         println!("{}: score={:?}", record.id, record.score);
2639///     }
2640/// }
2641/// ```
2642///
2643/// ## Batch search
2644///
2645/// ```
2646/// use chroma_types::operator::SearchResult;
2647///
2648/// fn process_batch_search(result: SearchResult) {
2649///     // Multiple searches in batch
2650///     for (i, search_result) in result.results.iter().enumerate() {
2651///         println!("\nSearch {}:", i + 1);
2652///         for record in &search_result.records {
2653///             println!("  {}: score={:?}", record.id, record.score);
2654///         }
2655///     }
2656/// }
2657/// ```
2658#[derive(Clone, Debug)]
2659pub struct SearchResult {
2660    pub results: Vec<SearchPayloadResult>,
2661    pub pulled_log_bytes: u64,
2662}
2663
2664impl SearchResult {
2665    pub fn size_bytes(&self) -> u64 {
2666        self.results
2667            .iter()
2668            .flat_map(|result| {
2669                result.records.iter().map(|record| {
2670                    (record.id.len()
2671                        + record
2672                            .document
2673                            .as_ref()
2674                            .map(|doc| doc.len())
2675                            .unwrap_or_default()
2676                        + record
2677                            .embedding
2678                            .as_ref()
2679                            .map(|emb| size_of_val(&emb[..]))
2680                            .unwrap_or_default()
2681                        + record
2682                            .metadata
2683                            .as_ref()
2684                            .map(logical_size_of_metadata)
2685                            .unwrap_or_default()
2686                        + record.score.as_ref().map(size_of_val).unwrap_or_default())
2687                        as u64
2688                })
2689            })
2690            .sum()
2691    }
2692}
2693
2694impl TryFrom<chroma_proto::SearchResult> for SearchResult {
2695    type Error = QueryConversionError;
2696
2697    fn try_from(value: chroma_proto::SearchResult) -> Result<Self, Self::Error> {
2698        Ok(Self {
2699            results: value
2700                .results
2701                .into_iter()
2702                .map(TryInto::try_into)
2703                .collect::<Result<_, _>>()?,
2704            pulled_log_bytes: value.pulled_log_bytes,
2705        })
2706    }
2707}
2708
2709impl TryFrom<SearchResult> for chroma_proto::SearchResult {
2710    type Error = QueryConversionError;
2711
2712    fn try_from(value: SearchResult) -> Result<Self, Self::Error> {
2713        Ok(Self {
2714            results: value
2715                .results
2716                .into_iter()
2717                .map(TryInto::try_into)
2718                .collect::<Result<Vec<_>, _>>()?,
2719            pulled_log_bytes: value.pulled_log_bytes,
2720        })
2721    }
2722}
2723
2724/// Reciprocal Rank Fusion (RRF) - combines multiple ranking strategies.
2725///
2726/// RRF is ideal for hybrid search where you want to merge results from different
2727/// ranking methods (e.g., dense and sparse embeddings) with different score scales.
2728/// It uses rank positions instead of raw scores, making it scale-agnostic.
2729///
2730/// # Formula
2731///
2732/// ```text
2733/// score = -Σ(weight_i / (k + rank_i))
2734/// ```
2735///
2736/// Where:
2737/// - `weight_i` = weight for ranking i (default: 1.0)
2738/// - `rank_i` = rank position from ranking i (0, 1, 2...)
2739/// - `k` = smoothing parameter (default: 60)
2740///
2741/// Score is negative because Chroma uses ascending order (lower = better).
2742///
2743/// # Arguments
2744///
2745/// * `ranks` - List of ranking expressions (must have `return_rank=true`)
2746/// * `k` - Smoothing parameter (None = 60). Higher values reduce emphasis on top ranks.
2747/// * `weights` - Weight for each ranking (None = all 1.0)
2748/// * `normalize` - If true, normalize weights to sum to 1.0
2749///
2750/// # Returns
2751///
2752/// A combined RankExpr or an error if:
2753/// - `ranks` is empty
2754/// - `weights` length doesn't match `ranks` length
2755/// - `weights` sum to zero when normalizing
2756///
2757/// # Examples
2758///
2759/// ## Basic RRF with default parameters
2760///
2761/// ```
2762/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2763///
2764/// let dense = RankExpr::Knn {
2765///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2766///     key: Key::Embedding,
2767///     limit: 200,
2768///     default: None,
2769///     return_rank: true, // Required for RRF
2770/// };
2771///
2772/// let sparse = RankExpr::Knn {
2773///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2774///     key: Key::field("sparse_embedding"),
2775///     limit: 200,
2776///     default: None,
2777///     return_rank: true, // Required for RRF
2778/// };
2779///
2780/// // Equal weights, k=60 (defaults)
2781/// let combined = rrf(vec![dense, sparse], None, None, false).unwrap();
2782/// ```
2783///
2784/// ## RRF with custom weights
2785///
2786/// ```
2787/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2788///
2789/// # let dense = RankExpr::Knn {
2790/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2791/// #     key: Key::Embedding,
2792/// #     limit: 200,
2793/// #     default: None,
2794/// #     return_rank: true,
2795/// # };
2796/// # let sparse = RankExpr::Knn {
2797/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2798/// #     key: Key::field("sparse_embedding"),
2799/// #     limit: 200,
2800/// #     default: None,
2801/// #     return_rank: true,
2802/// # };
2803/// // 70% dense, 30% sparse
2804/// let combined = rrf(
2805///     vec![dense, sparse],
2806///     Some(60),
2807///     Some(vec![0.7, 0.3]),
2808///     false,
2809/// ).unwrap();
2810/// ```
2811///
2812/// ## RRF with normalized weights
2813///
2814/// ```
2815/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2816///
2817/// # let dense = RankExpr::Knn {
2818/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2819/// #     key: Key::Embedding,
2820/// #     limit: 200,
2821/// #     default: None,
2822/// #     return_rank: true,
2823/// # };
2824/// # let sparse = RankExpr::Knn {
2825/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2826/// #     key: Key::field("sparse_embedding"),
2827/// #     limit: 200,
2828/// #     default: None,
2829/// #     return_rank: true,
2830/// # };
2831/// // Weights [75, 25] normalized to [0.75, 0.25]
2832/// let combined = rrf(
2833///     vec![dense, sparse],
2834///     Some(60),
2835///     Some(vec![75.0, 25.0]),
2836///     true, // normalize
2837/// ).unwrap();
2838/// ```
2839///
2840/// ## Adjusting the k parameter
2841///
2842/// ```
2843/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2844///
2845/// # let dense = RankExpr::Knn {
2846/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2847/// #     key: Key::Embedding,
2848/// #     limit: 200,
2849/// #     default: None,
2850/// #     return_rank: true,
2851/// # };
2852/// # let sparse = RankExpr::Knn {
2853/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2854/// #     key: Key::field("sparse_embedding"),
2855/// #     limit: 200,
2856/// #     default: None,
2857/// #     return_rank: true,
2858/// # };
2859/// // Small k (10) = heavy emphasis on top ranks
2860/// let top_heavy = rrf(vec![dense.clone(), sparse.clone()], Some(10), None, false).unwrap();
2861///
2862/// // Default k (60) = balanced
2863/// let balanced = rrf(vec![dense.clone(), sparse.clone()], Some(60), None, false).unwrap();
2864///
2865/// // Large k (200) = more uniform weighting
2866/// let uniform = rrf(vec![dense, sparse], Some(200), None, false).unwrap();
2867/// ```
2868pub fn rrf(
2869    ranks: Vec<RankExpr>,
2870    k: Option<u32>,
2871    weights: Option<Vec<f32>>,
2872    normalize: bool,
2873) -> Result<RankExpr, QueryConversionError> {
2874    let k = k.unwrap_or(60);
2875
2876    if ranks.is_empty() {
2877        return Err(QueryConversionError::validation(
2878            "RRF requires at least one rank expression",
2879        ));
2880    }
2881
2882    let weights = weights.unwrap_or_else(|| vec![1.0; ranks.len()]);
2883
2884    if weights.len() != ranks.len() {
2885        return Err(QueryConversionError::validation(format!(
2886            "RRF weights length ({}) must match ranks length ({})",
2887            weights.len(),
2888            ranks.len()
2889        )));
2890    }
2891
2892    let weights = if normalize {
2893        let sum: f32 = weights.iter().sum();
2894        if sum == 0.0 {
2895            return Err(QueryConversionError::validation(
2896                "RRF weights sum to zero, cannot normalize",
2897            ));
2898        }
2899        weights.into_iter().map(|w| w / sum).collect()
2900    } else {
2901        weights
2902    };
2903
2904    let terms: Vec<RankExpr> = weights
2905        .into_iter()
2906        .zip(ranks)
2907        .map(|(w, rank)| RankExpr::Value(w) / (RankExpr::Value(k as f32) + rank))
2908        .collect();
2909
2910    // Safe: ranks is validated as non-empty above, so terms cannot be empty.
2911    // Using unwrap_or_else as defensive programming to avoid panic.
2912    let sum = terms
2913        .into_iter()
2914        .reduce(|a, b| a + b)
2915        .unwrap_or(RankExpr::Value(0.0));
2916    Ok(-sum)
2917}
2918
2919#[cfg(test)]
2920mod tests {
2921    use super::*;
2922
2923    #[test]
2924    fn test_key_from_string() {
2925        // Test predefined keys
2926        assert_eq!(Key::from("#document"), Key::Document);
2927        assert_eq!(Key::from("#embedding"), Key::Embedding);
2928        assert_eq!(Key::from("#metadata"), Key::Metadata);
2929        assert_eq!(Key::from("#score"), Key::Score);
2930
2931        // Test metadata field keys
2932        assert_eq!(
2933            Key::from("custom_field"),
2934            Key::MetadataField("custom_field".to_string())
2935        );
2936        assert_eq!(
2937            Key::from("author"),
2938            Key::MetadataField("author".to_string())
2939        );
2940
2941        // Test String variant
2942        assert_eq!(Key::from("#embedding".to_string()), Key::Embedding);
2943        assert_eq!(
2944            Key::from("year".to_string()),
2945            Key::MetadataField("year".to_string())
2946        );
2947    }
2948
2949    #[test]
2950    fn test_query_vector_dense_proto_conversion() {
2951        let dense_vec = vec![0.1, 0.2, 0.3, 0.4, 0.5];
2952        let query_vector = QueryVector::Dense(dense_vec.clone());
2953
2954        // Convert to proto
2955        let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2956
2957        // Convert back
2958        let converted: QueryVector = proto.try_into().unwrap();
2959
2960        assert_eq!(converted, query_vector);
2961        if let QueryVector::Dense(v) = converted {
2962            assert_eq!(v, dense_vec);
2963        } else {
2964            panic!("Expected dense vector");
2965        }
2966    }
2967
2968    #[test]
2969    fn test_query_vector_sparse_proto_conversion() {
2970        let sparse = SparseVector::new(vec![0, 5, 10], vec![0.1, 0.5, 0.9]).unwrap();
2971        let query_vector = QueryVector::Sparse(sparse.clone());
2972
2973        // Convert to proto
2974        let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2975
2976        // Convert back
2977        let converted: QueryVector = proto.try_into().unwrap();
2978
2979        assert_eq!(converted, query_vector);
2980        if let QueryVector::Sparse(s) = converted {
2981            assert_eq!(s, sparse);
2982        } else {
2983            panic!("Expected sparse vector");
2984        }
2985    }
2986
2987    #[test]
2988    fn test_filter_json_deserialization() {
2989        // For the new search API, deserialization treats the entire JSON as a where clause
2990
2991        // Test 1: Simple direct metadata comparison
2992        let simple_where = r#"{"author": "John Doe"}"#;
2993        let filter: Filter = serde_json::from_str(simple_where).unwrap();
2994        assert_eq!(filter.query_ids, None);
2995        assert!(filter.where_clause.is_some());
2996
2997        // Test 2: ID filter using #id with $in operator
2998        let id_filter_json = serde_json::json!({
2999            "#id": {
3000                "$in": ["doc1", "doc2", "doc3"]
3001            }
3002        });
3003        let filter: Filter = serde_json::from_value(id_filter_json).unwrap();
3004        assert_eq!(filter.query_ids, None);
3005        assert!(filter.where_clause.is_some());
3006
3007        // Test 3: Complex nested expression with AND, OR, and various operators
3008        let complex_json = serde_json::json!({
3009            "$and": [
3010                {
3011                    "#id": {
3012                        "$in": ["doc1", "doc2", "doc3"]
3013                    }
3014                },
3015                {
3016                    "$or": [
3017                        {
3018                            "author": {
3019                                "$eq": "John Doe"
3020                            }
3021                        },
3022                        {
3023                            "author": {
3024                                "$eq": "Jane Smith"
3025                            }
3026                        }
3027                    ]
3028                },
3029                {
3030                    "year": {
3031                        "$gte": 2020
3032                    }
3033                },
3034                {
3035                    "tags": {
3036                        "$contains": "machine-learning"
3037                    }
3038                }
3039            ]
3040        });
3041
3042        let filter: Filter = serde_json::from_value(complex_json.clone()).unwrap();
3043        assert_eq!(filter.query_ids, None);
3044        assert!(filter.where_clause.is_some());
3045
3046        // Verify the structure
3047        if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3048            assert_eq!(composite.operator, crate::metadata::BooleanOperator::And);
3049            assert_eq!(composite.children.len(), 4);
3050
3051            // Check that the second child is an OR
3052            if let crate::metadata::Where::Composite(or_composite) = &composite.children[1] {
3053                assert_eq!(or_composite.operator, crate::metadata::BooleanOperator::Or);
3054                assert_eq!(or_composite.children.len(), 2);
3055            } else {
3056                panic!("Expected OR composite in second child");
3057            }
3058        } else {
3059            panic!("Expected AND composite where clause");
3060        }
3061
3062        // Test 4: Mixed operators - $ne, $lt, $gt, $lte
3063        let mixed_operators_json = serde_json::json!({
3064            "$and": [
3065                {
3066                    "status": {
3067                        "$ne": "deleted"
3068                    }
3069                },
3070                {
3071                    "score": {
3072                        "$gt": 0.5
3073                    }
3074                },
3075                {
3076                    "score": {
3077                        "$lt": 0.9
3078                    }
3079                },
3080                {
3081                    "priority": {
3082                        "$lte": 10
3083                    }
3084                }
3085            ]
3086        });
3087
3088        let filter: Filter = serde_json::from_value(mixed_operators_json).unwrap();
3089        assert_eq!(filter.query_ids, None);
3090        assert!(filter.where_clause.is_some());
3091
3092        // Test 5: Deeply nested expression
3093        let deeply_nested_json = serde_json::json!({
3094            "$or": [
3095                {
3096                    "$and": [
3097                        {
3098                            "#id": {
3099                                "$in": ["id1", "id2"]
3100                            }
3101                        },
3102                        {
3103                            "$or": [
3104                                {
3105                                    "category": "tech"
3106                                },
3107                                {
3108                                    "category": "science"
3109                                }
3110                            ]
3111                        }
3112                    ]
3113                },
3114                {
3115                    "$and": [
3116                        {
3117                            "author": "Admin"
3118                        },
3119                        {
3120                            "published": true
3121                        }
3122                    ]
3123                }
3124            ]
3125        });
3126
3127        let filter: Filter = serde_json::from_value(deeply_nested_json).unwrap();
3128        assert_eq!(filter.query_ids, None);
3129        assert!(filter.where_clause.is_some());
3130
3131        // Verify it's an OR at the top level
3132        if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3133            assert_eq!(composite.operator, crate::metadata::BooleanOperator::Or);
3134            assert_eq!(composite.children.len(), 2);
3135
3136            // Both children should be AND composites
3137            for child in &composite.children {
3138                if let crate::metadata::Where::Composite(and_composite) = child {
3139                    assert_eq!(
3140                        and_composite.operator,
3141                        crate::metadata::BooleanOperator::And
3142                    );
3143                } else {
3144                    panic!("Expected AND composite in OR children");
3145                }
3146            }
3147        } else {
3148            panic!("Expected OR composite at top level");
3149        }
3150
3151        // Test 6: Single ID filter (edge case)
3152        let single_id_json = serde_json::json!({
3153            "#id": {
3154                "$eq": "single-doc-id"
3155            }
3156        });
3157
3158        let filter: Filter = serde_json::from_value(single_id_json).unwrap();
3159        assert_eq!(filter.query_ids, None);
3160        assert!(filter.where_clause.is_some());
3161
3162        // Test 7: Empty object should create empty filter
3163        let empty_json = serde_json::json!({});
3164        let filter: Filter = serde_json::from_value(empty_json).unwrap();
3165        assert_eq!(filter.query_ids, None);
3166        // Empty object results in None where_clause
3167        assert_eq!(filter.where_clause, None);
3168
3169        // Test 8: Combining #id filter with $not_contains and numeric comparisons
3170        let advanced_json = serde_json::json!({
3171            "$and": [
3172                {
3173                    "#id": {
3174                        "$in": ["doc1", "doc2", "doc3", "doc4", "doc5"]
3175                    }
3176                },
3177                {
3178                    "tags": {
3179                        "$not_contains": "deprecated"
3180                    }
3181                },
3182                {
3183                    "$or": [
3184                        {
3185                            "$and": [
3186                                {
3187                                    "confidence": {
3188                                        "$gte": 0.8
3189                                    }
3190                                },
3191                                {
3192                                    "verified": true
3193                                }
3194                            ]
3195                        },
3196                        {
3197                            "$and": [
3198                                {
3199                                    "confidence": {
3200                                        "$gte": 0.6
3201                                    }
3202                                },
3203                                {
3204                                    "confidence": {
3205                                        "$lt": 0.8
3206                                    }
3207                                },
3208                                {
3209                                    "reviews": {
3210                                        "$gte": 5
3211                                    }
3212                                }
3213                            ]
3214                        }
3215                    ]
3216                }
3217            ]
3218        });
3219
3220        let filter: Filter = serde_json::from_value(advanced_json).unwrap();
3221        assert_eq!(filter.query_ids, None);
3222        assert!(filter.where_clause.is_some());
3223
3224        // Verify top-level structure
3225        if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3226            assert_eq!(composite.operator, crate::metadata::BooleanOperator::And);
3227            assert_eq!(composite.children.len(), 3);
3228        } else {
3229            panic!("Expected AND composite at top level");
3230        }
3231    }
3232
3233    #[test]
3234    fn test_limit_json_serialization() {
3235        let limit = Limit {
3236            offset: 10,
3237            limit: Some(20),
3238        };
3239
3240        let json = serde_json::to_string(&limit).unwrap();
3241        let deserialized: Limit = serde_json::from_str(&json).unwrap();
3242
3243        assert_eq!(deserialized.offset, limit.offset);
3244        assert_eq!(deserialized.limit, limit.limit);
3245    }
3246
3247    #[test]
3248    fn test_query_vector_json_serialization() {
3249        // Test dense vector
3250        let dense = QueryVector::Dense(vec![0.1, 0.2, 0.3]);
3251        let json = serde_json::to_string(&dense).unwrap();
3252        let deserialized: QueryVector = serde_json::from_str(&json).unwrap();
3253        assert_eq!(deserialized, dense);
3254
3255        // Test sparse vector
3256        let sparse =
3257            QueryVector::Sparse(SparseVector::new(vec![0, 5, 10], vec![0.1, 0.5, 0.9]).unwrap());
3258        let json = serde_json::to_string(&sparse).unwrap();
3259        let deserialized: QueryVector = serde_json::from_str(&json).unwrap();
3260        assert_eq!(deserialized, sparse);
3261    }
3262
3263    #[test]
3264    fn test_select_key_json_serialization() {
3265        use std::collections::HashSet;
3266
3267        // Test predefined keys
3268        let doc_key = Key::Document;
3269        assert_eq!(serde_json::to_string(&doc_key).unwrap(), "\"#document\"");
3270
3271        let embed_key = Key::Embedding;
3272        assert_eq!(serde_json::to_string(&embed_key).unwrap(), "\"#embedding\"");
3273
3274        let meta_key = Key::Metadata;
3275        assert_eq!(serde_json::to_string(&meta_key).unwrap(), "\"#metadata\"");
3276
3277        let score_key = Key::Score;
3278        assert_eq!(serde_json::to_string(&score_key).unwrap(), "\"#score\"");
3279
3280        // Test metadata key
3281        let custom_key = Key::MetadataField("custom_key".to_string());
3282        assert_eq!(
3283            serde_json::to_string(&custom_key).unwrap(),
3284            "\"custom_key\""
3285        );
3286
3287        // Test deserialization
3288        let deserialized: Key = serde_json::from_str("\"#document\"").unwrap();
3289        assert!(matches!(deserialized, Key::Document));
3290
3291        let deserialized: Key = serde_json::from_str("\"custom_field\"").unwrap();
3292        assert!(matches!(deserialized, Key::MetadataField(s) if s == "custom_field"));
3293
3294        // Test Select struct with multiple keys
3295        let mut keys = HashSet::new();
3296        keys.insert(Key::Document);
3297        keys.insert(Key::Embedding);
3298        keys.insert(Key::MetadataField("author".to_string()));
3299
3300        let select = Select { keys };
3301        let json = serde_json::to_string(&select).unwrap();
3302        let deserialized: Select = serde_json::from_str(&json).unwrap();
3303
3304        assert_eq!(deserialized.keys.len(), 3);
3305        assert!(deserialized.keys.contains(&Key::Document));
3306        assert!(deserialized.keys.contains(&Key::Embedding));
3307        assert!(deserialized
3308            .keys
3309            .contains(&Key::MetadataField("author".to_string())));
3310    }
3311
3312    #[test]
3313    fn test_merge_basic_integers() {
3314        use std::cmp::Reverse;
3315
3316        let merge = Merge { k: 5 };
3317
3318        // Input: sorted vectors of Reverse(u32) - ascending order of inner values
3319        let input = vec![
3320            vec![Reverse(1), Reverse(4), Reverse(7), Reverse(10)],
3321            vec![Reverse(2), Reverse(5), Reverse(8)],
3322            vec![Reverse(3), Reverse(6), Reverse(9), Reverse(11), Reverse(12)],
3323        ];
3324
3325        let result = merge.merge(input);
3326
3327        // Should get top-5 smallest values (largest Reverse values)
3328        assert_eq!(result.len(), 5);
3329        assert_eq!(
3330            result,
3331            vec![Reverse(1), Reverse(2), Reverse(3), Reverse(4), Reverse(5)]
3332        );
3333    }
3334
3335    #[test]
3336    fn test_merge_u32_descending() {
3337        let merge = Merge { k: 6 };
3338
3339        // Regular u32 in descending order (largest first)
3340        let input = vec![
3341            vec![100u32, 75, 50, 25],
3342            vec![90, 60, 30],
3343            vec![95, 85, 70, 40, 10],
3344        ];
3345
3346        let result = merge.merge(input);
3347
3348        // Should get top-6 largest u32 values
3349        assert_eq!(result.len(), 6);
3350        assert_eq!(result, vec![100, 95, 90, 85, 75, 70]);
3351    }
3352
3353    #[test]
3354    fn test_merge_i32_descending() {
3355        let merge = Merge { k: 5 };
3356
3357        // i32 values in descending order (including negatives)
3358        let input = vec![
3359            vec![50i32, 10, -10, -50],
3360            vec![30, 0, -30],
3361            vec![40, 20, -20, -40],
3362        ];
3363
3364        let result = merge.merge(input);
3365
3366        // Should get top-5 largest i32 values
3367        assert_eq!(result.len(), 5);
3368        assert_eq!(result, vec![50, 40, 30, 20, 10]);
3369    }
3370
3371    #[test]
3372    fn test_merge_with_duplicates() {
3373        let merge = Merge { k: 10 };
3374
3375        // Input with duplicates using regular u32 in descending order
3376        let input = vec![
3377            vec![100u32, 80, 80, 60, 40],
3378            vec![90, 80, 50, 30],
3379            vec![100, 70, 60, 20],
3380        ];
3381
3382        let result = merge.merge(input);
3383
3384        // Duplicates should be removed
3385        assert_eq!(result, vec![100, 90, 80, 70, 60, 50, 40, 30, 20]);
3386    }
3387
3388    #[test]
3389    fn test_merge_empty_vectors() {
3390        let merge = Merge { k: 5 };
3391
3392        // All empty with u32
3393        let input: Vec<Vec<u32>> = vec![vec![], vec![], vec![]];
3394        let result = merge.merge(input);
3395        assert_eq!(result.len(), 0);
3396
3397        // Some empty, some with data (u64)
3398        let input = vec![vec![], vec![1000u64, 750, 500], vec![], vec![850, 600]];
3399        let result = merge.merge(input);
3400        assert_eq!(result, vec![1000, 850, 750, 600, 500]);
3401
3402        // Single non-empty vector (i32)
3403        let input = vec![vec![], vec![100i32, 50, 25], vec![]];
3404        let result = merge.merge(input);
3405        assert_eq!(result, vec![100, 50, 25]);
3406    }
3407
3408    #[test]
3409    fn test_merge_k_boundary_conditions() {
3410        // k = 0 with u32
3411        let merge = Merge { k: 0 };
3412        let input = vec![vec![100u32, 50], vec![75, 25]];
3413        let result = merge.merge(input);
3414        assert_eq!(result.len(), 0);
3415
3416        // k = 1 with i64
3417        let merge = Merge { k: 1 };
3418        let input = vec![vec![1000i64, 500], vec![750, 250], vec![900, 100]];
3419        let result = merge.merge(input);
3420        assert_eq!(result, vec![1000]);
3421
3422        // k larger than total unique elements with u128
3423        let merge = Merge { k: 100 };
3424        let input = vec![vec![10000u128, 5000], vec![8000, 3000]];
3425        let result = merge.merge(input);
3426        assert_eq!(result, vec![10000, 8000, 5000, 3000]);
3427    }
3428
3429    #[test]
3430    fn test_merge_with_strings() {
3431        let merge = Merge { k: 4 };
3432
3433        // Strings must be sorted in descending order (largest first) for the max heap merge
3434        let input = vec![
3435            vec!["zebra".to_string(), "dog".to_string(), "apple".to_string()],
3436            vec!["elephant".to_string(), "banana".to_string()],
3437            vec!["fish".to_string(), "cat".to_string()],
3438        ];
3439
3440        let result = merge.merge(input);
3441
3442        // Should get top-4 lexicographically largest strings
3443        assert_eq!(
3444            result,
3445            vec![
3446                "zebra".to_string(),
3447                "fish".to_string(),
3448                "elephant".to_string(),
3449                "dog".to_string()
3450            ]
3451        );
3452    }
3453
3454    #[test]
3455    fn test_merge_with_custom_struct() {
3456        #[derive(Debug, Clone, Hash, PartialEq, Eq, PartialOrd, Ord)]
3457        struct Score {
3458            value: i32,
3459            id: String,
3460        }
3461
3462        let merge = Merge { k: 3 };
3463
3464        // Custom structs sorted by value (descending), then by id
3465        let input = vec![
3466            vec![
3467                Score {
3468                    value: 100,
3469                    id: "a".to_string(),
3470                },
3471                Score {
3472                    value: 80,
3473                    id: "b".to_string(),
3474                },
3475                Score {
3476                    value: 60,
3477                    id: "c".to_string(),
3478                },
3479            ],
3480            vec![
3481                Score {
3482                    value: 90,
3483                    id: "d".to_string(),
3484                },
3485                Score {
3486                    value: 70,
3487                    id: "e".to_string(),
3488                },
3489            ],
3490            vec![
3491                Score {
3492                    value: 95,
3493                    id: "f".to_string(),
3494                },
3495                Score {
3496                    value: 85,
3497                    id: "g".to_string(),
3498                },
3499            ],
3500        ];
3501
3502        let result = merge.merge(input);
3503
3504        assert_eq!(result.len(), 3);
3505        assert_eq!(
3506            result[0],
3507            Score {
3508                value: 100,
3509                id: "a".to_string()
3510            }
3511        );
3512        assert_eq!(
3513            result[1],
3514            Score {
3515                value: 95,
3516                id: "f".to_string()
3517            }
3518        );
3519        assert_eq!(
3520            result[2],
3521            Score {
3522                value: 90,
3523                id: "d".to_string()
3524            }
3525        );
3526    }
3527
3528    #[test]
3529    fn test_merge_preserves_order() {
3530        use std::cmp::Reverse;
3531
3532        let merge = Merge { k: 10 };
3533
3534        // For Reverse, smaller inner values are "larger" in ordering
3535        // So vectors should be sorted with smallest inner values first
3536        let input = vec![
3537            vec![Reverse(2), Reverse(6), Reverse(10), Reverse(14)],
3538            vec![Reverse(4), Reverse(8), Reverse(12), Reverse(16)],
3539            vec![Reverse(1), Reverse(3), Reverse(5), Reverse(7), Reverse(9)],
3540        ];
3541
3542        let result = merge.merge(input);
3543
3544        // Verify output maintains order - should be sorted by Reverse ordering
3545        // which means ascending inner values
3546        for i in 1..result.len() {
3547            assert!(
3548                result[i - 1] >= result[i],
3549                "Output should be in descending Reverse order"
3550            );
3551            assert!(
3552                result[i - 1].0 <= result[i].0,
3553                "Inner values should be in ascending order"
3554            );
3555        }
3556
3557        // Check we got the right elements
3558        assert_eq!(
3559            result,
3560            vec![
3561                Reverse(1),
3562                Reverse(2),
3563                Reverse(3),
3564                Reverse(4),
3565                Reverse(5),
3566                Reverse(6),
3567                Reverse(7),
3568                Reverse(8),
3569                Reverse(9),
3570                Reverse(10)
3571            ]
3572        );
3573    }
3574
3575    #[test]
3576    fn test_merge_single_vector() {
3577        let merge = Merge { k: 3 };
3578
3579        // Single vector input with u64
3580        let input = vec![vec![1000u64, 800, 600, 400, 200]];
3581
3582        let result = merge.merge(input);
3583
3584        assert_eq!(result, vec![1000, 800, 600]);
3585    }
3586
3587    #[test]
3588    fn test_merge_all_same_values() {
3589        let merge = Merge { k: 5 };
3590
3591        // All vectors contain the same value (using i16)
3592        let input = vec![vec![42i16, 42, 42], vec![42, 42], vec![42, 42, 42, 42]];
3593
3594        let result = merge.merge(input);
3595
3596        // Should deduplicate to single value
3597        assert_eq!(result, vec![42]);
3598    }
3599
3600    #[test]
3601    fn test_merge_mixed_types_sizes() {
3602        // Test with usize (common in real usage)
3603        let merge = Merge { k: 4 };
3604        let input = vec![
3605            vec![1000usize, 500, 100],
3606            vec![800, 300],
3607            vec![900, 600, 200],
3608        ];
3609        let result = merge.merge(input);
3610        assert_eq!(result, vec![1000, 900, 800, 600]);
3611
3612        // Test with negative integers (i32)
3613        let merge = Merge { k: 5 };
3614        let input = vec![vec![10i32, 0, -10, -20], vec![5, -5, -15], vec![15, -25]];
3615        let result = merge.merge(input);
3616        assert_eq!(result, vec![15, 10, 5, 0, -5]);
3617    }
3618
3619    #[test]
3620    fn test_merge_dedup_same_id_different_scores() {
3621        use std::cmp::Reverse;
3622
3623        // Simulates quantized distance estimation: the same offset_id appears
3624        // in multiple posting lists with different approximate distances.
3625        // Merge should keep only the first (best-scored) occurrence per id.
3626        let merge = Merge { k: 5 };
3627
3628        // Three posting lists with overlapping IDs and different estimated distances.
3629        // Using Reverse<RecordMeasure> to match the real knn_merge call site:
3630        // the heap acts as a min-heap on measure (smallest distance = best).
3631        let input: Vec<Vec<Reverse<RecordMeasure>>> = vec![
3632            vec![
3633                Reverse(RecordMeasure {
3634                    offset_id: 1,
3635                    measure: 0.10,
3636                }),
3637                Reverse(RecordMeasure {
3638                    offset_id: 4,
3639                    measure: 0.50,
3640                }),
3641                Reverse(RecordMeasure {
3642                    offset_id: 5,
3643                    measure: 0.70,
3644                }),
3645            ],
3646            vec![
3647                Reverse(RecordMeasure {
3648                    offset_id: 2,
3649                    measure: 0.20,
3650                }),
3651                Reverse(RecordMeasure {
3652                    offset_id: 1,
3653                    measure: 0.25,
3654                }), // dup id=1, worse
3655                Reverse(RecordMeasure {
3656                    offset_id: 3,
3657                    measure: 0.60,
3658                }),
3659            ],
3660            vec![
3661                Reverse(RecordMeasure {
3662                    offset_id: 3,
3663                    measure: 0.15,
3664                }), // dup id=3, better than 0.60
3665                Reverse(RecordMeasure {
3666                    offset_id: 4,
3667                    measure: 0.35,
3668                }), // dup id=4, better than 0.50
3669                Reverse(RecordMeasure {
3670                    offset_id: 2,
3671                    measure: 0.80,
3672                }), // dup id=2, worse
3673            ],
3674        ];
3675
3676        // Expected merge order (ascending distance via Reverse min-heap):
3677        //   pop 0.10 id=1 → keep
3678        //   pop 0.15 id=3 → keep
3679        //   pop 0.20 id=2 → keep
3680        //   pop 0.25 id=1 → dup, skip
3681        //   pop 0.35 id=4 → keep
3682        //   pop 0.50 id=4 → dup, skip
3683        //   pop 0.60 id=3 → dup, skip
3684        //   pop 0.70 id=5 → keep (5th unique)
3685        let result: Vec<Reverse<RecordMeasure>> = merge.merge(input);
3686        let ids: Vec<u32> = result.iter().map(|Reverse(r)| r.offset_id).collect();
3687        let measures: Vec<f32> = result.iter().map(|Reverse(r)| r.measure).collect();
3688
3689        assert_eq!(ids, vec![1, 3, 2, 4, 5]);
3690        assert_eq!(measures, vec![0.10, 0.15, 0.20, 0.35, 0.70]);
3691    }
3692
3693    #[test]
3694    fn test_aggregate_json_serialization() {
3695        // Test MinK serialization
3696        let min_k = Aggregate::MinK {
3697            keys: vec![Key::Score, Key::field("date")],
3698            k: 3,
3699        };
3700        let json = serde_json::to_value(&min_k).unwrap();
3701        assert!(json.get("$min_k").is_some());
3702        assert_eq!(json["$min_k"]["k"], 3);
3703
3704        // Test MinK deserialization
3705        let min_k_json = serde_json::json!({
3706            "$min_k": {
3707                "keys": ["#score", "date"],
3708                "k": 5
3709            }
3710        });
3711        let deserialized: Aggregate = serde_json::from_value(min_k_json).unwrap();
3712        match deserialized {
3713            Aggregate::MinK { keys, k } => {
3714                assert_eq!(k, 5);
3715                assert_eq!(keys.len(), 2);
3716                assert_eq!(keys[0], Key::Score);
3717                assert_eq!(keys[1], Key::field("date"));
3718            }
3719            _ => panic!("Expected MinK"),
3720        }
3721
3722        // Test MaxK serialization
3723        let max_k = Aggregate::MaxK {
3724            keys: vec![Key::field("timestamp")],
3725            k: 10,
3726        };
3727        let json = serde_json::to_value(&max_k).unwrap();
3728        assert!(json.get("$max_k").is_some());
3729        assert_eq!(json["$max_k"]["k"], 10);
3730
3731        // Test MaxK deserialization
3732        let max_k_json = serde_json::json!({
3733            "$max_k": {
3734                "keys": ["timestamp"],
3735                "k": 2
3736            }
3737        });
3738        let deserialized: Aggregate = serde_json::from_value(max_k_json).unwrap();
3739        match deserialized {
3740            Aggregate::MaxK { keys, k } => {
3741                assert_eq!(k, 2);
3742                assert_eq!(keys.len(), 1);
3743                assert_eq!(keys[0], Key::field("timestamp"));
3744            }
3745            _ => panic!("Expected MaxK"),
3746        }
3747    }
3748
3749    #[test]
3750    fn test_group_by_json_serialization() {
3751        // Test GroupBy with MinK
3752        let group_by = GroupBy {
3753            keys: vec![Key::field("category"), Key::field("author")],
3754            aggregate: Some(Aggregate::MinK {
3755                keys: vec![Key::Score],
3756                k: 3,
3757            }),
3758        };
3759
3760        let json = serde_json::to_value(&group_by).unwrap();
3761        assert_eq!(json["keys"].as_array().unwrap().len(), 2);
3762        assert!(json["aggregate"]["$min_k"].is_object());
3763
3764        // Test roundtrip
3765        let deserialized: GroupBy = serde_json::from_value(json).unwrap();
3766        assert_eq!(deserialized.keys.len(), 2);
3767        assert_eq!(deserialized.keys[0], Key::field("category"));
3768        assert_eq!(deserialized.keys[1], Key::field("author"));
3769        assert!(deserialized.aggregate.is_some());
3770
3771        // Test empty GroupBy
3772        let empty_group_by = GroupBy::default();
3773        let json = serde_json::to_value(&empty_group_by).unwrap();
3774        let deserialized: GroupBy = serde_json::from_value(json).unwrap();
3775        assert!(deserialized.keys.is_empty());
3776        assert!(deserialized.aggregate.is_none());
3777
3778        // Test deserialization from JSON
3779        let json = serde_json::json!({
3780            "keys": ["category"],
3781            "aggregate": {
3782                "$max_k": {
3783                    "keys": ["#score", "priority"],
3784                    "k": 5
3785                }
3786            }
3787        });
3788        let group_by: GroupBy = serde_json::from_value(json).unwrap();
3789        assert_eq!(group_by.keys.len(), 1);
3790        assert_eq!(group_by.keys[0], Key::field("category"));
3791        match group_by.aggregate {
3792            Some(Aggregate::MaxK { keys, k }) => {
3793                assert_eq!(k, 5);
3794                assert_eq!(keys.len(), 2);
3795                assert_eq!(keys[0], Key::Score);
3796            }
3797            _ => panic!("Expected MaxK aggregate"),
3798        }
3799    }
3800
3801    fn sparse_knn_leaf(index: u32, key: &str, return_rank: bool) -> RankExpr {
3802        RankExpr::Knn {
3803            query: QueryVector::Sparse(SparseVector::new(vec![index], vec![1.0]).unwrap()),
3804            key: Key::field(key),
3805            limit: 10,
3806            default: None,
3807            return_rank,
3808        }
3809    }
3810
3811    #[test]
3812    fn test_knn_queries_collects_all_sparse_leaves_in_dfs_order() {
3813        // Two distinct sparse keys plus a repeat of the first key with a
3814        // different query vector. All three leaves must be collected, in order,
3815        // with no deduplication by key or query type.
3816        let expr = RankExpr::Summation(vec![
3817            sparse_knn_leaf(0, "sparse_a", false),
3818            sparse_knn_leaf(1, "sparse_b", false),
3819            sparse_knn_leaf(2, "sparse_a", false),
3820        ]);
3821
3822        let leaves = expr.knn_queries();
3823        assert_eq!(leaves.len(), 3);
3824        assert_eq!(leaves[0].key, Key::field("sparse_a"));
3825        assert_eq!(leaves[1].key, Key::field("sparse_b"));
3826        assert_eq!(leaves[2].key, Key::field("sparse_a"));
3827
3828        // The two same-key leaves keep their distinct query vectors.
3829        match (&leaves[0].query, &leaves[2].query) {
3830            (QueryVector::Sparse(first), QueryVector::Sparse(third)) => {
3831                assert_ne!(first.indices, third.indices);
3832            }
3833            _ => panic!("expected sparse query vectors"),
3834        }
3835    }
3836
3837    #[test]
3838    fn test_rrf_preserves_all_sparse_leaves() {
3839        // RRF over multiple sparse leaves expands into arithmetic but must still
3840        // surface every leaf (so each gets its own per-key orchestrator).
3841        let expr = rrf(
3842            vec![
3843                sparse_knn_leaf(0, "sparse_a", true),
3844                sparse_knn_leaf(1, "sparse_b", true),
3845                sparse_knn_leaf(2, "sparse_c", true),
3846            ],
3847            None,
3848            None,
3849            false,
3850        )
3851        .expect("rrf should build");
3852
3853        let keys: Vec<Key> = expr.knn_queries().into_iter().map(|q| q.key).collect();
3854        assert_eq!(
3855            keys,
3856            vec![
3857                Key::field("sparse_a"),
3858                Key::field("sparse_b"),
3859                Key::field("sparse_c"),
3860            ]
3861        );
3862    }
3863}