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 TryFrom<chroma_proto::Aggregate> for Aggregate {
2377    type Error = QueryConversionError;
2378
2379    fn try_from(value: chroma_proto::Aggregate) -> Result<Self, Self::Error> {
2380        match value
2381            .aggregate
2382            .ok_or(QueryConversionError::field("aggregate"))?
2383        {
2384            chroma_proto::aggregate::Aggregate::MinK(min_k) => {
2385                let keys = min_k.keys.into_iter().map(Key::from).collect();
2386                Ok(Aggregate::MinK { keys, k: min_k.k })
2387            }
2388            chroma_proto::aggregate::Aggregate::MaxK(max_k) => {
2389                let keys = max_k.keys.into_iter().map(Key::from).collect();
2390                Ok(Aggregate::MaxK { keys, k: max_k.k })
2391            }
2392        }
2393    }
2394}
2395
2396impl From<Aggregate> for chroma_proto::Aggregate {
2397    fn from(value: Aggregate) -> Self {
2398        let aggregate = match value {
2399            Aggregate::MinK { keys, k } => {
2400                chroma_proto::aggregate::Aggregate::MinK(chroma_proto::aggregate::MinK {
2401                    keys: keys.into_iter().map(|k| k.to_string()).collect(),
2402                    k,
2403                })
2404            }
2405            Aggregate::MaxK { keys, k } => {
2406                chroma_proto::aggregate::Aggregate::MaxK(chroma_proto::aggregate::MaxK {
2407                    keys: keys.into_iter().map(|k| k.to_string()).collect(),
2408                    k,
2409                })
2410            }
2411        };
2412
2413        chroma_proto::Aggregate {
2414            aggregate: Some(aggregate),
2415        }
2416    }
2417}
2418
2419impl TryFrom<chroma_proto::GroupByOperator> for GroupBy {
2420    type Error = QueryConversionError;
2421
2422    fn try_from(value: chroma_proto::GroupByOperator) -> Result<Self, Self::Error> {
2423        let keys = value.keys.into_iter().map(Key::from).collect();
2424        let aggregate = value.aggregate.map(TryInto::try_into).transpose()?;
2425
2426        Ok(Self { keys, aggregate })
2427    }
2428}
2429
2430impl TryFrom<GroupBy> for chroma_proto::GroupByOperator {
2431    type Error = QueryConversionError;
2432
2433    fn try_from(value: GroupBy) -> Result<Self, Self::Error> {
2434        let keys = value.keys.into_iter().map(|k| k.to_string()).collect();
2435        let aggregate = value.aggregate.map(Into::into);
2436
2437        Ok(Self { keys, aggregate })
2438    }
2439}
2440
2441/// A single search result record.
2442///
2443/// Contains the document ID and optionally document content, embeddings, metadata,
2444/// and search score based on what was selected in the search query.
2445///
2446/// # Fields
2447///
2448/// * `id` - Document ID (always present)
2449/// * `document` - Document text content (if selected)
2450/// * `embedding` - Vector embedding (if selected)
2451/// * `metadata` - Document metadata (if selected)
2452/// * `score` - Search score (present when ranking is used, lower = better match)
2453///
2454/// # Examples
2455///
2456/// ```
2457/// use chroma_types::operator::SearchRecord;
2458///
2459/// fn process_results(records: Vec<SearchRecord>) {
2460///     for record in records {
2461///         println!("ID: {}", record.id);
2462///
2463///         if let Some(score) = record.score {
2464///             println!("  Score: {:.3}", score);
2465///         }
2466///
2467///         if let Some(doc) = record.document {
2468///             println!("  Document: {}", doc);
2469///         }
2470///
2471///         if let Some(meta) = record.metadata {
2472///             println!("  Metadata: {:?}", meta);
2473///         }
2474///     }
2475/// }
2476/// ```
2477#[derive(Clone, Debug, Deserialize, Serialize)]
2478#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
2479pub struct SearchRecord {
2480    pub id: String,
2481    pub document: Option<String>,
2482    pub embedding: Option<Vec<f32>>,
2483    pub metadata: Option<Metadata>,
2484    pub score: Option<f32>,
2485}
2486
2487impl TryFrom<chroma_proto::SearchRecord> for SearchRecord {
2488    type Error = QueryConversionError;
2489
2490    fn try_from(value: chroma_proto::SearchRecord) -> Result<Self, Self::Error> {
2491        Ok(Self {
2492            id: value.id,
2493            document: value.document,
2494            embedding: value
2495                .embedding
2496                .map(|vec| vec.try_into().map(|(v, _)| v))
2497                .transpose()?,
2498            metadata: value.metadata.map(TryInto::try_into).transpose()?,
2499            score: value.score,
2500        })
2501    }
2502}
2503
2504impl TryFrom<SearchRecord> for chroma_proto::SearchRecord {
2505    type Error = QueryConversionError;
2506
2507    fn try_from(value: SearchRecord) -> Result<Self, Self::Error> {
2508        Ok(Self {
2509            id: value.id,
2510            document: value.document,
2511            embedding: value
2512                .embedding
2513                .map(|embedding| {
2514                    let embedding_dimension = embedding.len();
2515                    chroma_proto::Vector::try_from((
2516                        embedding,
2517                        ScalarEncoding::FLOAT32,
2518                        embedding_dimension,
2519                    ))
2520                })
2521                .transpose()?,
2522            metadata: value.metadata.map(Into::into),
2523            score: value.score,
2524        })
2525    }
2526}
2527
2528/// Results for a single search payload.
2529///
2530/// Contains all matching records for one search query.
2531///
2532/// # Fields
2533///
2534/// * `records` - Vector of search records, ordered by score (ascending)
2535///
2536/// # Examples
2537///
2538/// ```
2539/// use chroma_types::operator::{SearchPayloadResult, SearchRecord};
2540///
2541/// fn process_search_result(result: SearchPayloadResult) {
2542///     println!("Found {} results", result.records.len());
2543///
2544///     for (i, record) in result.records.iter().enumerate() {
2545///         println!("{}. {} (score: {:?})", i + 1, record.id, record.score);
2546///     }
2547/// }
2548/// ```
2549#[derive(Clone, Debug, Default)]
2550pub struct SearchPayloadResult {
2551    pub records: Vec<SearchRecord>,
2552}
2553
2554impl TryFrom<chroma_proto::SearchPayloadResult> for SearchPayloadResult {
2555    type Error = QueryConversionError;
2556
2557    fn try_from(value: chroma_proto::SearchPayloadResult) -> Result<Self, Self::Error> {
2558        Ok(Self {
2559            records: value
2560                .records
2561                .into_iter()
2562                .map(TryInto::try_into)
2563                .collect::<Result<_, _>>()?,
2564        })
2565    }
2566}
2567
2568impl TryFrom<SearchPayloadResult> for chroma_proto::SearchPayloadResult {
2569    type Error = QueryConversionError;
2570
2571    fn try_from(value: SearchPayloadResult) -> Result<Self, Self::Error> {
2572        Ok(Self {
2573            records: value
2574                .records
2575                .into_iter()
2576                .map(TryInto::try_into)
2577                .collect::<Result<Vec<_>, _>>()?,
2578        })
2579    }
2580}
2581
2582/// Results from a batch search operation.
2583///
2584/// Contains results for each search payload in the batch, maintaining the same order
2585/// as the input searches.
2586///
2587/// # Fields
2588///
2589/// * `results` - Results for each search payload (indexed by search position)
2590/// * `pulled_log_bytes` - Total bytes pulled from log (for internal metrics)
2591///
2592/// # Examples
2593///
2594/// ## Single search
2595///
2596/// ```
2597/// use chroma_types::operator::SearchResult;
2598///
2599/// fn process_single_search(result: SearchResult) {
2600///     // Single search, so results[0] contains our records
2601///     let records = &result.results[0].records;
2602///
2603///     for record in records {
2604///         println!("{}: score={:?}", record.id, record.score);
2605///     }
2606/// }
2607/// ```
2608///
2609/// ## Batch search
2610///
2611/// ```
2612/// use chroma_types::operator::SearchResult;
2613///
2614/// fn process_batch_search(result: SearchResult) {
2615///     // Multiple searches in batch
2616///     for (i, search_result) in result.results.iter().enumerate() {
2617///         println!("\nSearch {}:", i + 1);
2618///         for record in &search_result.records {
2619///             println!("  {}: score={:?}", record.id, record.score);
2620///         }
2621///     }
2622/// }
2623/// ```
2624#[derive(Clone, Debug)]
2625pub struct SearchResult {
2626    pub results: Vec<SearchPayloadResult>,
2627    pub pulled_log_bytes: u64,
2628}
2629
2630impl SearchResult {
2631    pub fn size_bytes(&self) -> u64 {
2632        self.results
2633            .iter()
2634            .flat_map(|result| {
2635                result.records.iter().map(|record| {
2636                    (record.id.len()
2637                        + record
2638                            .document
2639                            .as_ref()
2640                            .map(|doc| doc.len())
2641                            .unwrap_or_default()
2642                        + record
2643                            .embedding
2644                            .as_ref()
2645                            .map(|emb| size_of_val(&emb[..]))
2646                            .unwrap_or_default()
2647                        + record
2648                            .metadata
2649                            .as_ref()
2650                            .map(logical_size_of_metadata)
2651                            .unwrap_or_default()
2652                        + record.score.as_ref().map(size_of_val).unwrap_or_default())
2653                        as u64
2654                })
2655            })
2656            .sum()
2657    }
2658}
2659
2660impl TryFrom<chroma_proto::SearchResult> for SearchResult {
2661    type Error = QueryConversionError;
2662
2663    fn try_from(value: chroma_proto::SearchResult) -> Result<Self, Self::Error> {
2664        Ok(Self {
2665            results: value
2666                .results
2667                .into_iter()
2668                .map(TryInto::try_into)
2669                .collect::<Result<_, _>>()?,
2670            pulled_log_bytes: value.pulled_log_bytes,
2671        })
2672    }
2673}
2674
2675impl TryFrom<SearchResult> for chroma_proto::SearchResult {
2676    type Error = QueryConversionError;
2677
2678    fn try_from(value: SearchResult) -> Result<Self, Self::Error> {
2679        Ok(Self {
2680            results: value
2681                .results
2682                .into_iter()
2683                .map(TryInto::try_into)
2684                .collect::<Result<Vec<_>, _>>()?,
2685            pulled_log_bytes: value.pulled_log_bytes,
2686        })
2687    }
2688}
2689
2690/// Reciprocal Rank Fusion (RRF) - combines multiple ranking strategies.
2691///
2692/// RRF is ideal for hybrid search where you want to merge results from different
2693/// ranking methods (e.g., dense and sparse embeddings) with different score scales.
2694/// It uses rank positions instead of raw scores, making it scale-agnostic.
2695///
2696/// # Formula
2697///
2698/// ```text
2699/// score = -Σ(weight_i / (k + rank_i))
2700/// ```
2701///
2702/// Where:
2703/// - `weight_i` = weight for ranking i (default: 1.0)
2704/// - `rank_i` = rank position from ranking i (0, 1, 2...)
2705/// - `k` = smoothing parameter (default: 60)
2706///
2707/// Score is negative because Chroma uses ascending order (lower = better).
2708///
2709/// # Arguments
2710///
2711/// * `ranks` - List of ranking expressions (must have `return_rank=true`)
2712/// * `k` - Smoothing parameter (None = 60). Higher values reduce emphasis on top ranks.
2713/// * `weights` - Weight for each ranking (None = all 1.0)
2714/// * `normalize` - If true, normalize weights to sum to 1.0
2715///
2716/// # Returns
2717///
2718/// A combined RankExpr or an error if:
2719/// - `ranks` is empty
2720/// - `weights` length doesn't match `ranks` length
2721/// - `weights` sum to zero when normalizing
2722///
2723/// # Examples
2724///
2725/// ## Basic RRF with default parameters
2726///
2727/// ```
2728/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2729///
2730/// let dense = RankExpr::Knn {
2731///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2732///     key: Key::Embedding,
2733///     limit: 200,
2734///     default: None,
2735///     return_rank: true, // Required for RRF
2736/// };
2737///
2738/// let sparse = RankExpr::Knn {
2739///     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2740///     key: Key::field("sparse_embedding"),
2741///     limit: 200,
2742///     default: None,
2743///     return_rank: true, // Required for RRF
2744/// };
2745///
2746/// // Equal weights, k=60 (defaults)
2747/// let combined = rrf(vec![dense, sparse], None, None, false).unwrap();
2748/// ```
2749///
2750/// ## RRF with custom weights
2751///
2752/// ```
2753/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2754///
2755/// # let dense = RankExpr::Knn {
2756/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2757/// #     key: Key::Embedding,
2758/// #     limit: 200,
2759/// #     default: None,
2760/// #     return_rank: true,
2761/// # };
2762/// # let sparse = RankExpr::Knn {
2763/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2764/// #     key: Key::field("sparse_embedding"),
2765/// #     limit: 200,
2766/// #     default: None,
2767/// #     return_rank: true,
2768/// # };
2769/// // 70% dense, 30% sparse
2770/// let combined = rrf(
2771///     vec![dense, sparse],
2772///     Some(60),
2773///     Some(vec![0.7, 0.3]),
2774///     false,
2775/// ).unwrap();
2776/// ```
2777///
2778/// ## RRF with normalized weights
2779///
2780/// ```
2781/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2782///
2783/// # let dense = RankExpr::Knn {
2784/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2785/// #     key: Key::Embedding,
2786/// #     limit: 200,
2787/// #     default: None,
2788/// #     return_rank: true,
2789/// # };
2790/// # let sparse = RankExpr::Knn {
2791/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2792/// #     key: Key::field("sparse_embedding"),
2793/// #     limit: 200,
2794/// #     default: None,
2795/// #     return_rank: true,
2796/// # };
2797/// // Weights [75, 25] normalized to [0.75, 0.25]
2798/// let combined = rrf(
2799///     vec![dense, sparse],
2800///     Some(60),
2801///     Some(vec![75.0, 25.0]),
2802///     true, // normalize
2803/// ).unwrap();
2804/// ```
2805///
2806/// ## Adjusting the k parameter
2807///
2808/// ```
2809/// use chroma_types::operator::{RankExpr, QueryVector, Key, rrf};
2810///
2811/// # let dense = RankExpr::Knn {
2812/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2813/// #     key: Key::Embedding,
2814/// #     limit: 200,
2815/// #     default: None,
2816/// #     return_rank: true,
2817/// # };
2818/// # let sparse = RankExpr::Knn {
2819/// #     query: QueryVector::Dense(vec![0.1, 0.2, 0.3]),
2820/// #     key: Key::field("sparse_embedding"),
2821/// #     limit: 200,
2822/// #     default: None,
2823/// #     return_rank: true,
2824/// # };
2825/// // Small k (10) = heavy emphasis on top ranks
2826/// let top_heavy = rrf(vec![dense.clone(), sparse.clone()], Some(10), None, false).unwrap();
2827///
2828/// // Default k (60) = balanced
2829/// let balanced = rrf(vec![dense.clone(), sparse.clone()], Some(60), None, false).unwrap();
2830///
2831/// // Large k (200) = more uniform weighting
2832/// let uniform = rrf(vec![dense, sparse], Some(200), None, false).unwrap();
2833/// ```
2834pub fn rrf(
2835    ranks: Vec<RankExpr>,
2836    k: Option<u32>,
2837    weights: Option<Vec<f32>>,
2838    normalize: bool,
2839) -> Result<RankExpr, QueryConversionError> {
2840    let k = k.unwrap_or(60);
2841
2842    if ranks.is_empty() {
2843        return Err(QueryConversionError::validation(
2844            "RRF requires at least one rank expression",
2845        ));
2846    }
2847
2848    let weights = weights.unwrap_or_else(|| vec![1.0; ranks.len()]);
2849
2850    if weights.len() != ranks.len() {
2851        return Err(QueryConversionError::validation(format!(
2852            "RRF weights length ({}) must match ranks length ({})",
2853            weights.len(),
2854            ranks.len()
2855        )));
2856    }
2857
2858    let weights = if normalize {
2859        let sum: f32 = weights.iter().sum();
2860        if sum == 0.0 {
2861            return Err(QueryConversionError::validation(
2862                "RRF weights sum to zero, cannot normalize",
2863            ));
2864        }
2865        weights.into_iter().map(|w| w / sum).collect()
2866    } else {
2867        weights
2868    };
2869
2870    let terms: Vec<RankExpr> = weights
2871        .into_iter()
2872        .zip(ranks)
2873        .map(|(w, rank)| RankExpr::Value(w) / (RankExpr::Value(k as f32) + rank))
2874        .collect();
2875
2876    // Safe: ranks is validated as non-empty above, so terms cannot be empty.
2877    // Using unwrap_or_else as defensive programming to avoid panic.
2878    let sum = terms
2879        .into_iter()
2880        .reduce(|a, b| a + b)
2881        .unwrap_or(RankExpr::Value(0.0));
2882    Ok(-sum)
2883}
2884
2885#[cfg(test)]
2886mod tests {
2887    use super::*;
2888
2889    #[test]
2890    fn test_key_from_string() {
2891        // Test predefined keys
2892        assert_eq!(Key::from("#document"), Key::Document);
2893        assert_eq!(Key::from("#embedding"), Key::Embedding);
2894        assert_eq!(Key::from("#metadata"), Key::Metadata);
2895        assert_eq!(Key::from("#score"), Key::Score);
2896
2897        // Test metadata field keys
2898        assert_eq!(
2899            Key::from("custom_field"),
2900            Key::MetadataField("custom_field".to_string())
2901        );
2902        assert_eq!(
2903            Key::from("author"),
2904            Key::MetadataField("author".to_string())
2905        );
2906
2907        // Test String variant
2908        assert_eq!(Key::from("#embedding".to_string()), Key::Embedding);
2909        assert_eq!(
2910            Key::from("year".to_string()),
2911            Key::MetadataField("year".to_string())
2912        );
2913    }
2914
2915    #[test]
2916    fn test_query_vector_dense_proto_conversion() {
2917        let dense_vec = vec![0.1, 0.2, 0.3, 0.4, 0.5];
2918        let query_vector = QueryVector::Dense(dense_vec.clone());
2919
2920        // Convert to proto
2921        let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2922
2923        // Convert back
2924        let converted: QueryVector = proto.try_into().unwrap();
2925
2926        assert_eq!(converted, query_vector);
2927        if let QueryVector::Dense(v) = converted {
2928            assert_eq!(v, dense_vec);
2929        } else {
2930            panic!("Expected dense vector");
2931        }
2932    }
2933
2934    #[test]
2935    fn test_query_vector_sparse_proto_conversion() {
2936        let sparse = SparseVector::new(vec![0, 5, 10], vec![0.1, 0.5, 0.9]).unwrap();
2937        let query_vector = QueryVector::Sparse(sparse.clone());
2938
2939        // Convert to proto
2940        let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2941
2942        // Convert back
2943        let converted: QueryVector = proto.try_into().unwrap();
2944
2945        assert_eq!(converted, query_vector);
2946        if let QueryVector::Sparse(s) = converted {
2947            assert_eq!(s, sparse);
2948        } else {
2949            panic!("Expected sparse vector");
2950        }
2951    }
2952
2953    #[test]
2954    fn test_filter_json_deserialization() {
2955        // For the new search API, deserialization treats the entire JSON as a where clause
2956
2957        // Test 1: Simple direct metadata comparison
2958        let simple_where = r#"{"author": "John Doe"}"#;
2959        let filter: Filter = serde_json::from_str(simple_where).unwrap();
2960        assert_eq!(filter.query_ids, None);
2961        assert!(filter.where_clause.is_some());
2962
2963        // Test 2: ID filter using #id with $in operator
2964        let id_filter_json = serde_json::json!({
2965            "#id": {
2966                "$in": ["doc1", "doc2", "doc3"]
2967            }
2968        });
2969        let filter: Filter = serde_json::from_value(id_filter_json).unwrap();
2970        assert_eq!(filter.query_ids, None);
2971        assert!(filter.where_clause.is_some());
2972
2973        // Test 3: Complex nested expression with AND, OR, and various operators
2974        let complex_json = serde_json::json!({
2975            "$and": [
2976                {
2977                    "#id": {
2978                        "$in": ["doc1", "doc2", "doc3"]
2979                    }
2980                },
2981                {
2982                    "$or": [
2983                        {
2984                            "author": {
2985                                "$eq": "John Doe"
2986                            }
2987                        },
2988                        {
2989                            "author": {
2990                                "$eq": "Jane Smith"
2991                            }
2992                        }
2993                    ]
2994                },
2995                {
2996                    "year": {
2997                        "$gte": 2020
2998                    }
2999                },
3000                {
3001                    "tags": {
3002                        "$contains": "machine-learning"
3003                    }
3004                }
3005            ]
3006        });
3007
3008        let filter: Filter = serde_json::from_value(complex_json.clone()).unwrap();
3009        assert_eq!(filter.query_ids, None);
3010        assert!(filter.where_clause.is_some());
3011
3012        // Verify the structure
3013        if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3014            assert_eq!(composite.operator, crate::metadata::BooleanOperator::And);
3015            assert_eq!(composite.children.len(), 4);
3016
3017            // Check that the second child is an OR
3018            if let crate::metadata::Where::Composite(or_composite) = &composite.children[1] {
3019                assert_eq!(or_composite.operator, crate::metadata::BooleanOperator::Or);
3020                assert_eq!(or_composite.children.len(), 2);
3021            } else {
3022                panic!("Expected OR composite in second child");
3023            }
3024        } else {
3025            panic!("Expected AND composite where clause");
3026        }
3027
3028        // Test 4: Mixed operators - $ne, $lt, $gt, $lte
3029        let mixed_operators_json = serde_json::json!({
3030            "$and": [
3031                {
3032                    "status": {
3033                        "$ne": "deleted"
3034                    }
3035                },
3036                {
3037                    "score": {
3038                        "$gt": 0.5
3039                    }
3040                },
3041                {
3042                    "score": {
3043                        "$lt": 0.9
3044                    }
3045                },
3046                {
3047                    "priority": {
3048                        "$lte": 10
3049                    }
3050                }
3051            ]
3052        });
3053
3054        let filter: Filter = serde_json::from_value(mixed_operators_json).unwrap();
3055        assert_eq!(filter.query_ids, None);
3056        assert!(filter.where_clause.is_some());
3057
3058        // Test 5: Deeply nested expression
3059        let deeply_nested_json = serde_json::json!({
3060            "$or": [
3061                {
3062                    "$and": [
3063                        {
3064                            "#id": {
3065                                "$in": ["id1", "id2"]
3066                            }
3067                        },
3068                        {
3069                            "$or": [
3070                                {
3071                                    "category": "tech"
3072                                },
3073                                {
3074                                    "category": "science"
3075                                }
3076                            ]
3077                        }
3078                    ]
3079                },
3080                {
3081                    "$and": [
3082                        {
3083                            "author": "Admin"
3084                        },
3085                        {
3086                            "published": true
3087                        }
3088                    ]
3089                }
3090            ]
3091        });
3092
3093        let filter: Filter = serde_json::from_value(deeply_nested_json).unwrap();
3094        assert_eq!(filter.query_ids, None);
3095        assert!(filter.where_clause.is_some());
3096
3097        // Verify it's an OR at the top level
3098        if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3099            assert_eq!(composite.operator, crate::metadata::BooleanOperator::Or);
3100            assert_eq!(composite.children.len(), 2);
3101
3102            // Both children should be AND composites
3103            for child in &composite.children {
3104                if let crate::metadata::Where::Composite(and_composite) = child {
3105                    assert_eq!(
3106                        and_composite.operator,
3107                        crate::metadata::BooleanOperator::And
3108                    );
3109                } else {
3110                    panic!("Expected AND composite in OR children");
3111                }
3112            }
3113        } else {
3114            panic!("Expected OR composite at top level");
3115        }
3116
3117        // Test 6: Single ID filter (edge case)
3118        let single_id_json = serde_json::json!({
3119            "#id": {
3120                "$eq": "single-doc-id"
3121            }
3122        });
3123
3124        let filter: Filter = serde_json::from_value(single_id_json).unwrap();
3125        assert_eq!(filter.query_ids, None);
3126        assert!(filter.where_clause.is_some());
3127
3128        // Test 7: Empty object should create empty filter
3129        let empty_json = serde_json::json!({});
3130        let filter: Filter = serde_json::from_value(empty_json).unwrap();
3131        assert_eq!(filter.query_ids, None);
3132        // Empty object results in None where_clause
3133        assert_eq!(filter.where_clause, None);
3134
3135        // Test 8: Combining #id filter with $not_contains and numeric comparisons
3136        let advanced_json = serde_json::json!({
3137            "$and": [
3138                {
3139                    "#id": {
3140                        "$in": ["doc1", "doc2", "doc3", "doc4", "doc5"]
3141                    }
3142                },
3143                {
3144                    "tags": {
3145                        "$not_contains": "deprecated"
3146                    }
3147                },
3148                {
3149                    "$or": [
3150                        {
3151                            "$and": [
3152                                {
3153                                    "confidence": {
3154                                        "$gte": 0.8
3155                                    }
3156                                },
3157                                {
3158                                    "verified": true
3159                                }
3160                            ]
3161                        },
3162                        {
3163                            "$and": [
3164                                {
3165                                    "confidence": {
3166                                        "$gte": 0.6
3167                                    }
3168                                },
3169                                {
3170                                    "confidence": {
3171                                        "$lt": 0.8
3172                                    }
3173                                },
3174                                {
3175                                    "reviews": {
3176                                        "$gte": 5
3177                                    }
3178                                }
3179                            ]
3180                        }
3181                    ]
3182                }
3183            ]
3184        });
3185
3186        let filter: Filter = serde_json::from_value(advanced_json).unwrap();
3187        assert_eq!(filter.query_ids, None);
3188        assert!(filter.where_clause.is_some());
3189
3190        // Verify top-level structure
3191        if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3192            assert_eq!(composite.operator, crate::metadata::BooleanOperator::And);
3193            assert_eq!(composite.children.len(), 3);
3194        } else {
3195            panic!("Expected AND composite at top level");
3196        }
3197    }
3198
3199    #[test]
3200    fn test_limit_json_serialization() {
3201        let limit = Limit {
3202            offset: 10,
3203            limit: Some(20),
3204        };
3205
3206        let json = serde_json::to_string(&limit).unwrap();
3207        let deserialized: Limit = serde_json::from_str(&json).unwrap();
3208
3209        assert_eq!(deserialized.offset, limit.offset);
3210        assert_eq!(deserialized.limit, limit.limit);
3211    }
3212
3213    #[test]
3214    fn test_query_vector_json_serialization() {
3215        // Test dense vector
3216        let dense = QueryVector::Dense(vec![0.1, 0.2, 0.3]);
3217        let json = serde_json::to_string(&dense).unwrap();
3218        let deserialized: QueryVector = serde_json::from_str(&json).unwrap();
3219        assert_eq!(deserialized, dense);
3220
3221        // Test sparse vector
3222        let sparse =
3223            QueryVector::Sparse(SparseVector::new(vec![0, 5, 10], vec![0.1, 0.5, 0.9]).unwrap());
3224        let json = serde_json::to_string(&sparse).unwrap();
3225        let deserialized: QueryVector = serde_json::from_str(&json).unwrap();
3226        assert_eq!(deserialized, sparse);
3227    }
3228
3229    #[test]
3230    fn test_select_key_json_serialization() {
3231        use std::collections::HashSet;
3232
3233        // Test predefined keys
3234        let doc_key = Key::Document;
3235        assert_eq!(serde_json::to_string(&doc_key).unwrap(), "\"#document\"");
3236
3237        let embed_key = Key::Embedding;
3238        assert_eq!(serde_json::to_string(&embed_key).unwrap(), "\"#embedding\"");
3239
3240        let meta_key = Key::Metadata;
3241        assert_eq!(serde_json::to_string(&meta_key).unwrap(), "\"#metadata\"");
3242
3243        let score_key = Key::Score;
3244        assert_eq!(serde_json::to_string(&score_key).unwrap(), "\"#score\"");
3245
3246        // Test metadata key
3247        let custom_key = Key::MetadataField("custom_key".to_string());
3248        assert_eq!(
3249            serde_json::to_string(&custom_key).unwrap(),
3250            "\"custom_key\""
3251        );
3252
3253        // Test deserialization
3254        let deserialized: Key = serde_json::from_str("\"#document\"").unwrap();
3255        assert!(matches!(deserialized, Key::Document));
3256
3257        let deserialized: Key = serde_json::from_str("\"custom_field\"").unwrap();
3258        assert!(matches!(deserialized, Key::MetadataField(s) if s == "custom_field"));
3259
3260        // Test Select struct with multiple keys
3261        let mut keys = HashSet::new();
3262        keys.insert(Key::Document);
3263        keys.insert(Key::Embedding);
3264        keys.insert(Key::MetadataField("author".to_string()));
3265
3266        let select = Select { keys };
3267        let json = serde_json::to_string(&select).unwrap();
3268        let deserialized: Select = serde_json::from_str(&json).unwrap();
3269
3270        assert_eq!(deserialized.keys.len(), 3);
3271        assert!(deserialized.keys.contains(&Key::Document));
3272        assert!(deserialized.keys.contains(&Key::Embedding));
3273        assert!(deserialized
3274            .keys
3275            .contains(&Key::MetadataField("author".to_string())));
3276    }
3277
3278    #[test]
3279    fn test_merge_basic_integers() {
3280        use std::cmp::Reverse;
3281
3282        let merge = Merge { k: 5 };
3283
3284        // Input: sorted vectors of Reverse(u32) - ascending order of inner values
3285        let input = vec![
3286            vec![Reverse(1), Reverse(4), Reverse(7), Reverse(10)],
3287            vec![Reverse(2), Reverse(5), Reverse(8)],
3288            vec![Reverse(3), Reverse(6), Reverse(9), Reverse(11), Reverse(12)],
3289        ];
3290
3291        let result = merge.merge(input);
3292
3293        // Should get top-5 smallest values (largest Reverse values)
3294        assert_eq!(result.len(), 5);
3295        assert_eq!(
3296            result,
3297            vec![Reverse(1), Reverse(2), Reverse(3), Reverse(4), Reverse(5)]
3298        );
3299    }
3300
3301    #[test]
3302    fn test_merge_u32_descending() {
3303        let merge = Merge { k: 6 };
3304
3305        // Regular u32 in descending order (largest first)
3306        let input = vec![
3307            vec![100u32, 75, 50, 25],
3308            vec![90, 60, 30],
3309            vec![95, 85, 70, 40, 10],
3310        ];
3311
3312        let result = merge.merge(input);
3313
3314        // Should get top-6 largest u32 values
3315        assert_eq!(result.len(), 6);
3316        assert_eq!(result, vec![100, 95, 90, 85, 75, 70]);
3317    }
3318
3319    #[test]
3320    fn test_merge_i32_descending() {
3321        let merge = Merge { k: 5 };
3322
3323        // i32 values in descending order (including negatives)
3324        let input = vec![
3325            vec![50i32, 10, -10, -50],
3326            vec![30, 0, -30],
3327            vec![40, 20, -20, -40],
3328        ];
3329
3330        let result = merge.merge(input);
3331
3332        // Should get top-5 largest i32 values
3333        assert_eq!(result.len(), 5);
3334        assert_eq!(result, vec![50, 40, 30, 20, 10]);
3335    }
3336
3337    #[test]
3338    fn test_merge_with_duplicates() {
3339        let merge = Merge { k: 10 };
3340
3341        // Input with duplicates using regular u32 in descending order
3342        let input = vec![
3343            vec![100u32, 80, 80, 60, 40],
3344            vec![90, 80, 50, 30],
3345            vec![100, 70, 60, 20],
3346        ];
3347
3348        let result = merge.merge(input);
3349
3350        // Duplicates should be removed
3351        assert_eq!(result, vec![100, 90, 80, 70, 60, 50, 40, 30, 20]);
3352    }
3353
3354    #[test]
3355    fn test_merge_empty_vectors() {
3356        let merge = Merge { k: 5 };
3357
3358        // All empty with u32
3359        let input: Vec<Vec<u32>> = vec![vec![], vec![], vec![]];
3360        let result = merge.merge(input);
3361        assert_eq!(result.len(), 0);
3362
3363        // Some empty, some with data (u64)
3364        let input = vec![vec![], vec![1000u64, 750, 500], vec![], vec![850, 600]];
3365        let result = merge.merge(input);
3366        assert_eq!(result, vec![1000, 850, 750, 600, 500]);
3367
3368        // Single non-empty vector (i32)
3369        let input = vec![vec![], vec![100i32, 50, 25], vec![]];
3370        let result = merge.merge(input);
3371        assert_eq!(result, vec![100, 50, 25]);
3372    }
3373
3374    #[test]
3375    fn test_merge_k_boundary_conditions() {
3376        // k = 0 with u32
3377        let merge = Merge { k: 0 };
3378        let input = vec![vec![100u32, 50], vec![75, 25]];
3379        let result = merge.merge(input);
3380        assert_eq!(result.len(), 0);
3381
3382        // k = 1 with i64
3383        let merge = Merge { k: 1 };
3384        let input = vec![vec![1000i64, 500], vec![750, 250], vec![900, 100]];
3385        let result = merge.merge(input);
3386        assert_eq!(result, vec![1000]);
3387
3388        // k larger than total unique elements with u128
3389        let merge = Merge { k: 100 };
3390        let input = vec![vec![10000u128, 5000], vec![8000, 3000]];
3391        let result = merge.merge(input);
3392        assert_eq!(result, vec![10000, 8000, 5000, 3000]);
3393    }
3394
3395    #[test]
3396    fn test_merge_with_strings() {
3397        let merge = Merge { k: 4 };
3398
3399        // Strings must be sorted in descending order (largest first) for the max heap merge
3400        let input = vec![
3401            vec!["zebra".to_string(), "dog".to_string(), "apple".to_string()],
3402            vec!["elephant".to_string(), "banana".to_string()],
3403            vec!["fish".to_string(), "cat".to_string()],
3404        ];
3405
3406        let result = merge.merge(input);
3407
3408        // Should get top-4 lexicographically largest strings
3409        assert_eq!(
3410            result,
3411            vec![
3412                "zebra".to_string(),
3413                "fish".to_string(),
3414                "elephant".to_string(),
3415                "dog".to_string()
3416            ]
3417        );
3418    }
3419
3420    #[test]
3421    fn test_merge_with_custom_struct() {
3422        #[derive(Debug, Clone, Hash, PartialEq, Eq, PartialOrd, Ord)]
3423        struct Score {
3424            value: i32,
3425            id: String,
3426        }
3427
3428        let merge = Merge { k: 3 };
3429
3430        // Custom structs sorted by value (descending), then by id
3431        let input = vec![
3432            vec![
3433                Score {
3434                    value: 100,
3435                    id: "a".to_string(),
3436                },
3437                Score {
3438                    value: 80,
3439                    id: "b".to_string(),
3440                },
3441                Score {
3442                    value: 60,
3443                    id: "c".to_string(),
3444                },
3445            ],
3446            vec![
3447                Score {
3448                    value: 90,
3449                    id: "d".to_string(),
3450                },
3451                Score {
3452                    value: 70,
3453                    id: "e".to_string(),
3454                },
3455            ],
3456            vec![
3457                Score {
3458                    value: 95,
3459                    id: "f".to_string(),
3460                },
3461                Score {
3462                    value: 85,
3463                    id: "g".to_string(),
3464                },
3465            ],
3466        ];
3467
3468        let result = merge.merge(input);
3469
3470        assert_eq!(result.len(), 3);
3471        assert_eq!(
3472            result[0],
3473            Score {
3474                value: 100,
3475                id: "a".to_string()
3476            }
3477        );
3478        assert_eq!(
3479            result[1],
3480            Score {
3481                value: 95,
3482                id: "f".to_string()
3483            }
3484        );
3485        assert_eq!(
3486            result[2],
3487            Score {
3488                value: 90,
3489                id: "d".to_string()
3490            }
3491        );
3492    }
3493
3494    #[test]
3495    fn test_merge_preserves_order() {
3496        use std::cmp::Reverse;
3497
3498        let merge = Merge { k: 10 };
3499
3500        // For Reverse, smaller inner values are "larger" in ordering
3501        // So vectors should be sorted with smallest inner values first
3502        let input = vec![
3503            vec![Reverse(2), Reverse(6), Reverse(10), Reverse(14)],
3504            vec![Reverse(4), Reverse(8), Reverse(12), Reverse(16)],
3505            vec![Reverse(1), Reverse(3), Reverse(5), Reverse(7), Reverse(9)],
3506        ];
3507
3508        let result = merge.merge(input);
3509
3510        // Verify output maintains order - should be sorted by Reverse ordering
3511        // which means ascending inner values
3512        for i in 1..result.len() {
3513            assert!(
3514                result[i - 1] >= result[i],
3515                "Output should be in descending Reverse order"
3516            );
3517            assert!(
3518                result[i - 1].0 <= result[i].0,
3519                "Inner values should be in ascending order"
3520            );
3521        }
3522
3523        // Check we got the right elements
3524        assert_eq!(
3525            result,
3526            vec![
3527                Reverse(1),
3528                Reverse(2),
3529                Reverse(3),
3530                Reverse(4),
3531                Reverse(5),
3532                Reverse(6),
3533                Reverse(7),
3534                Reverse(8),
3535                Reverse(9),
3536                Reverse(10)
3537            ]
3538        );
3539    }
3540
3541    #[test]
3542    fn test_merge_single_vector() {
3543        let merge = Merge { k: 3 };
3544
3545        // Single vector input with u64
3546        let input = vec![vec![1000u64, 800, 600, 400, 200]];
3547
3548        let result = merge.merge(input);
3549
3550        assert_eq!(result, vec![1000, 800, 600]);
3551    }
3552
3553    #[test]
3554    fn test_merge_all_same_values() {
3555        let merge = Merge { k: 5 };
3556
3557        // All vectors contain the same value (using i16)
3558        let input = vec![vec![42i16, 42, 42], vec![42, 42], vec![42, 42, 42, 42]];
3559
3560        let result = merge.merge(input);
3561
3562        // Should deduplicate to single value
3563        assert_eq!(result, vec![42]);
3564    }
3565
3566    #[test]
3567    fn test_merge_mixed_types_sizes() {
3568        // Test with usize (common in real usage)
3569        let merge = Merge { k: 4 };
3570        let input = vec![
3571            vec![1000usize, 500, 100],
3572            vec![800, 300],
3573            vec![900, 600, 200],
3574        ];
3575        let result = merge.merge(input);
3576        assert_eq!(result, vec![1000, 900, 800, 600]);
3577
3578        // Test with negative integers (i32)
3579        let merge = Merge { k: 5 };
3580        let input = vec![vec![10i32, 0, -10, -20], vec![5, -5, -15], vec![15, -25]];
3581        let result = merge.merge(input);
3582        assert_eq!(result, vec![15, 10, 5, 0, -5]);
3583    }
3584
3585    #[test]
3586    fn test_merge_dedup_same_id_different_scores() {
3587        use std::cmp::Reverse;
3588
3589        // Simulates quantized distance estimation: the same offset_id appears
3590        // in multiple posting lists with different approximate distances.
3591        // Merge should keep only the first (best-scored) occurrence per id.
3592        let merge = Merge { k: 5 };
3593
3594        // Three posting lists with overlapping IDs and different estimated distances.
3595        // Using Reverse<RecordMeasure> to match the real knn_merge call site:
3596        // the heap acts as a min-heap on measure (smallest distance = best).
3597        let input: Vec<Vec<Reverse<RecordMeasure>>> = vec![
3598            vec![
3599                Reverse(RecordMeasure {
3600                    offset_id: 1,
3601                    measure: 0.10,
3602                }),
3603                Reverse(RecordMeasure {
3604                    offset_id: 4,
3605                    measure: 0.50,
3606                }),
3607                Reverse(RecordMeasure {
3608                    offset_id: 5,
3609                    measure: 0.70,
3610                }),
3611            ],
3612            vec![
3613                Reverse(RecordMeasure {
3614                    offset_id: 2,
3615                    measure: 0.20,
3616                }),
3617                Reverse(RecordMeasure {
3618                    offset_id: 1,
3619                    measure: 0.25,
3620                }), // dup id=1, worse
3621                Reverse(RecordMeasure {
3622                    offset_id: 3,
3623                    measure: 0.60,
3624                }),
3625            ],
3626            vec![
3627                Reverse(RecordMeasure {
3628                    offset_id: 3,
3629                    measure: 0.15,
3630                }), // dup id=3, better than 0.60
3631                Reverse(RecordMeasure {
3632                    offset_id: 4,
3633                    measure: 0.35,
3634                }), // dup id=4, better than 0.50
3635                Reverse(RecordMeasure {
3636                    offset_id: 2,
3637                    measure: 0.80,
3638                }), // dup id=2, worse
3639            ],
3640        ];
3641
3642        // Expected merge order (ascending distance via Reverse min-heap):
3643        //   pop 0.10 id=1 → keep
3644        //   pop 0.15 id=3 → keep
3645        //   pop 0.20 id=2 → keep
3646        //   pop 0.25 id=1 → dup, skip
3647        //   pop 0.35 id=4 → keep
3648        //   pop 0.50 id=4 → dup, skip
3649        //   pop 0.60 id=3 → dup, skip
3650        //   pop 0.70 id=5 → keep (5th unique)
3651        let result: Vec<Reverse<RecordMeasure>> = merge.merge(input);
3652        let ids: Vec<u32> = result.iter().map(|Reverse(r)| r.offset_id).collect();
3653        let measures: Vec<f32> = result.iter().map(|Reverse(r)| r.measure).collect();
3654
3655        assert_eq!(ids, vec![1, 3, 2, 4, 5]);
3656        assert_eq!(measures, vec![0.10, 0.15, 0.20, 0.35, 0.70]);
3657    }
3658
3659    #[test]
3660    fn test_aggregate_json_serialization() {
3661        // Test MinK serialization
3662        let min_k = Aggregate::MinK {
3663            keys: vec![Key::Score, Key::field("date")],
3664            k: 3,
3665        };
3666        let json = serde_json::to_value(&min_k).unwrap();
3667        assert!(json.get("$min_k").is_some());
3668        assert_eq!(json["$min_k"]["k"], 3);
3669
3670        // Test MinK deserialization
3671        let min_k_json = serde_json::json!({
3672            "$min_k": {
3673                "keys": ["#score", "date"],
3674                "k": 5
3675            }
3676        });
3677        let deserialized: Aggregate = serde_json::from_value(min_k_json).unwrap();
3678        match deserialized {
3679            Aggregate::MinK { keys, k } => {
3680                assert_eq!(k, 5);
3681                assert_eq!(keys.len(), 2);
3682                assert_eq!(keys[0], Key::Score);
3683                assert_eq!(keys[1], Key::field("date"));
3684            }
3685            _ => panic!("Expected MinK"),
3686        }
3687
3688        // Test MaxK serialization
3689        let max_k = Aggregate::MaxK {
3690            keys: vec![Key::field("timestamp")],
3691            k: 10,
3692        };
3693        let json = serde_json::to_value(&max_k).unwrap();
3694        assert!(json.get("$max_k").is_some());
3695        assert_eq!(json["$max_k"]["k"], 10);
3696
3697        // Test MaxK deserialization
3698        let max_k_json = serde_json::json!({
3699            "$max_k": {
3700                "keys": ["timestamp"],
3701                "k": 2
3702            }
3703        });
3704        let deserialized: Aggregate = serde_json::from_value(max_k_json).unwrap();
3705        match deserialized {
3706            Aggregate::MaxK { keys, k } => {
3707                assert_eq!(k, 2);
3708                assert_eq!(keys.len(), 1);
3709                assert_eq!(keys[0], Key::field("timestamp"));
3710            }
3711            _ => panic!("Expected MaxK"),
3712        }
3713    }
3714
3715    #[test]
3716    fn test_group_by_json_serialization() {
3717        // Test GroupBy with MinK
3718        let group_by = GroupBy {
3719            keys: vec![Key::field("category"), Key::field("author")],
3720            aggregate: Some(Aggregate::MinK {
3721                keys: vec![Key::Score],
3722                k: 3,
3723            }),
3724        };
3725
3726        let json = serde_json::to_value(&group_by).unwrap();
3727        assert_eq!(json["keys"].as_array().unwrap().len(), 2);
3728        assert!(json["aggregate"]["$min_k"].is_object());
3729
3730        // Test roundtrip
3731        let deserialized: GroupBy = serde_json::from_value(json).unwrap();
3732        assert_eq!(deserialized.keys.len(), 2);
3733        assert_eq!(deserialized.keys[0], Key::field("category"));
3734        assert_eq!(deserialized.keys[1], Key::field("author"));
3735        assert!(deserialized.aggregate.is_some());
3736
3737        // Test empty GroupBy
3738        let empty_group_by = GroupBy::default();
3739        let json = serde_json::to_value(&empty_group_by).unwrap();
3740        let deserialized: GroupBy = serde_json::from_value(json).unwrap();
3741        assert!(deserialized.keys.is_empty());
3742        assert!(deserialized.aggregate.is_none());
3743
3744        // Test deserialization from JSON
3745        let json = serde_json::json!({
3746            "keys": ["category"],
3747            "aggregate": {
3748                "$max_k": {
3749                    "keys": ["#score", "priority"],
3750                    "k": 5
3751                }
3752            }
3753        });
3754        let group_by: GroupBy = serde_json::from_value(json).unwrap();
3755        assert_eq!(group_by.keys.len(), 1);
3756        assert_eq!(group_by.keys[0], Key::field("category"));
3757        match group_by.aggregate {
3758            Some(Aggregate::MaxK { keys, k }) => {
3759                assert_eq!(k, 5);
3760                assert_eq!(keys.len(), 2);
3761                assert_eq!(keys[0], Key::Score);
3762            }
3763            _ => panic!("Expected MaxK aggregate"),
3764        }
3765    }
3766}