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#[derive(Clone, Debug)]
28pub struct Scan {
29 pub collection_and_segments: CollectionAndSegments,
30 pub shard_index: u32,
31 pub num_shards: u32,
32 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#[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#[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 match (&self.query_ids, &self.where_clause) {
196 (None, None) => {
197 let map = serializer.serialize_map(Some(0))?;
199 map.end()
200 }
201 (None, Some(where_clause)) => {
202 where_clause.serialize(serializer)
204 }
205 (Some(ids), None) => {
206 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 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 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, 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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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 pub fn exp(self) -> Self {
1252 RankExpr::Exponentiation(Box::new(self))
1253 }
1254
1255 pub fn log(self) -> Self {
1276 RankExpr::Logarithm(Box::new(self))
1277 }
1278
1279 pub fn abs(self) -> Self {
1306 RankExpr::Absolute(Box::new(self))
1307 }
1308
1309 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 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#[derive(Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
1786#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
1787pub enum Key {
1788 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 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 pub fn field(name: impl Into<String>) -> Self {
1865 Key::MetadataField(name.into())
1866 }
1867
1868 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 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 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 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 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 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 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 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 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 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 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 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 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 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#[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 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 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#[derive(Clone, Debug, Deserialize, Serialize)]
2306pub enum Aggregate {
2307 #[serde(rename = "$min_k")]
2309 MinK {
2310 keys: Vec<Key>,
2312 k: u32,
2314 },
2315 #[serde(rename = "$max_k")]
2317 MaxK {
2318 keys: Vec<Key>,
2320 k: u32,
2322 },
2323}
2324
2325#[derive(Clone, Debug, Default, Deserialize, Serialize)]
2367pub struct GroupBy {
2368 #[serde(default)]
2370 pub keys: Vec<Key>,
2371 #[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#[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#[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#[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
2690pub 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 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 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 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 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 let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2922
2923 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 let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2941
2942 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 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 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 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 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 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 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 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 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 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 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 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 assert_eq!(filter.where_clause, None);
3134
3135 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 let input: Vec<Vec<u32>> = vec![vec![], vec![], vec![]];
3360 let result = merge.merge(input);
3361 assert_eq!(result.len(), 0);
3362
3363 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 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 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 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 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 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 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 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 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 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 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 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 let input = vec![vec![42i16, 42, 42], vec![42, 42], vec![42, 42, 42, 42]];
3559
3560 let result = merge.merge(input);
3561
3562 assert_eq!(result, vec![42]);
3564 }
3565
3566 #[test]
3567 fn test_merge_mixed_types_sizes() {
3568 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 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 let merge = Merge { k: 5 };
3593
3594 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 }), 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 }), Reverse(RecordMeasure {
3632 offset_id: 4,
3633 measure: 0.35,
3634 }), Reverse(RecordMeasure {
3636 offset_id: 2,
3637 measure: 0.80,
3638 }), ],
3640 ];
3641
3642 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 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 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 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 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 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 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 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 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}