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 GroupBy {
2377 pub fn is_active(&self) -> bool {
2380 !self.keys.is_empty() && self.aggregate.is_some()
2381 }
2382
2383 pub fn aggregate_keys(&self) -> &[Key] {
2386 match &self.aggregate {
2387 Some(Aggregate::MinK { keys, .. } | Aggregate::MaxK { keys, .. }) => keys,
2388 None => &[],
2389 }
2390 }
2391
2392 pub fn metadata_keys(&self) -> Vec<Key> {
2395 let mut result: Vec<Key> = self
2396 .keys
2397 .iter()
2398 .filter(|k| matches!(k, Key::MetadataField(_)))
2399 .cloned()
2400 .collect();
2401 for k in self.aggregate_keys() {
2402 if matches!(k, Key::MetadataField(_)) && !result.contains(k) {
2403 result.push(k.clone());
2404 }
2405 }
2406 result
2407 }
2408}
2409
2410impl TryFrom<chroma_proto::Aggregate> for Aggregate {
2411 type Error = QueryConversionError;
2412
2413 fn try_from(value: chroma_proto::Aggregate) -> Result<Self, Self::Error> {
2414 match value
2415 .aggregate
2416 .ok_or(QueryConversionError::field("aggregate"))?
2417 {
2418 chroma_proto::aggregate::Aggregate::MinK(min_k) => {
2419 let keys = min_k.keys.into_iter().map(Key::from).collect();
2420 Ok(Aggregate::MinK { keys, k: min_k.k })
2421 }
2422 chroma_proto::aggregate::Aggregate::MaxK(max_k) => {
2423 let keys = max_k.keys.into_iter().map(Key::from).collect();
2424 Ok(Aggregate::MaxK { keys, k: max_k.k })
2425 }
2426 }
2427 }
2428}
2429
2430impl From<Aggregate> for chroma_proto::Aggregate {
2431 fn from(value: Aggregate) -> Self {
2432 let aggregate = match value {
2433 Aggregate::MinK { keys, k } => {
2434 chroma_proto::aggregate::Aggregate::MinK(chroma_proto::aggregate::MinK {
2435 keys: keys.into_iter().map(|k| k.to_string()).collect(),
2436 k,
2437 })
2438 }
2439 Aggregate::MaxK { keys, k } => {
2440 chroma_proto::aggregate::Aggregate::MaxK(chroma_proto::aggregate::MaxK {
2441 keys: keys.into_iter().map(|k| k.to_string()).collect(),
2442 k,
2443 })
2444 }
2445 };
2446
2447 chroma_proto::Aggregate {
2448 aggregate: Some(aggregate),
2449 }
2450 }
2451}
2452
2453impl TryFrom<chroma_proto::GroupByOperator> for GroupBy {
2454 type Error = QueryConversionError;
2455
2456 fn try_from(value: chroma_proto::GroupByOperator) -> Result<Self, Self::Error> {
2457 let keys = value.keys.into_iter().map(Key::from).collect();
2458 let aggregate = value.aggregate.map(TryInto::try_into).transpose()?;
2459
2460 Ok(Self { keys, aggregate })
2461 }
2462}
2463
2464impl TryFrom<GroupBy> for chroma_proto::GroupByOperator {
2465 type Error = QueryConversionError;
2466
2467 fn try_from(value: GroupBy) -> Result<Self, Self::Error> {
2468 let keys = value.keys.into_iter().map(|k| k.to_string()).collect();
2469 let aggregate = value.aggregate.map(Into::into);
2470
2471 Ok(Self { keys, aggregate })
2472 }
2473}
2474
2475#[derive(Clone, Debug, Deserialize, Serialize)]
2512#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
2513pub struct SearchRecord {
2514 pub id: String,
2515 pub document: Option<String>,
2516 pub embedding: Option<Vec<f32>>,
2517 pub metadata: Option<Metadata>,
2518 pub score: Option<f32>,
2519}
2520
2521impl TryFrom<chroma_proto::SearchRecord> for SearchRecord {
2522 type Error = QueryConversionError;
2523
2524 fn try_from(value: chroma_proto::SearchRecord) -> Result<Self, Self::Error> {
2525 Ok(Self {
2526 id: value.id,
2527 document: value.document,
2528 embedding: value
2529 .embedding
2530 .map(|vec| vec.try_into().map(|(v, _)| v))
2531 .transpose()?,
2532 metadata: value.metadata.map(TryInto::try_into).transpose()?,
2533 score: value.score,
2534 })
2535 }
2536}
2537
2538impl TryFrom<SearchRecord> for chroma_proto::SearchRecord {
2539 type Error = QueryConversionError;
2540
2541 fn try_from(value: SearchRecord) -> Result<Self, Self::Error> {
2542 Ok(Self {
2543 id: value.id,
2544 document: value.document,
2545 embedding: value
2546 .embedding
2547 .map(|embedding| {
2548 let embedding_dimension = embedding.len();
2549 chroma_proto::Vector::try_from((
2550 embedding,
2551 ScalarEncoding::FLOAT32,
2552 embedding_dimension,
2553 ))
2554 })
2555 .transpose()?,
2556 metadata: value.metadata.map(Into::into),
2557 score: value.score,
2558 })
2559 }
2560}
2561
2562#[derive(Clone, Debug, Default)]
2584pub struct SearchPayloadResult {
2585 pub records: Vec<SearchRecord>,
2586}
2587
2588impl TryFrom<chroma_proto::SearchPayloadResult> for SearchPayloadResult {
2589 type Error = QueryConversionError;
2590
2591 fn try_from(value: chroma_proto::SearchPayloadResult) -> Result<Self, Self::Error> {
2592 Ok(Self {
2593 records: value
2594 .records
2595 .into_iter()
2596 .map(TryInto::try_into)
2597 .collect::<Result<_, _>>()?,
2598 })
2599 }
2600}
2601
2602impl TryFrom<SearchPayloadResult> for chroma_proto::SearchPayloadResult {
2603 type Error = QueryConversionError;
2604
2605 fn try_from(value: SearchPayloadResult) -> Result<Self, Self::Error> {
2606 Ok(Self {
2607 records: value
2608 .records
2609 .into_iter()
2610 .map(TryInto::try_into)
2611 .collect::<Result<Vec<_>, _>>()?,
2612 })
2613 }
2614}
2615
2616#[derive(Clone, Debug)]
2659pub struct SearchResult {
2660 pub results: Vec<SearchPayloadResult>,
2661 pub pulled_log_bytes: u64,
2662}
2663
2664impl SearchResult {
2665 pub fn size_bytes(&self) -> u64 {
2666 self.results
2667 .iter()
2668 .flat_map(|result| {
2669 result.records.iter().map(|record| {
2670 (record.id.len()
2671 + record
2672 .document
2673 .as_ref()
2674 .map(|doc| doc.len())
2675 .unwrap_or_default()
2676 + record
2677 .embedding
2678 .as_ref()
2679 .map(|emb| size_of_val(&emb[..]))
2680 .unwrap_or_default()
2681 + record
2682 .metadata
2683 .as_ref()
2684 .map(logical_size_of_metadata)
2685 .unwrap_or_default()
2686 + record.score.as_ref().map(size_of_val).unwrap_or_default())
2687 as u64
2688 })
2689 })
2690 .sum()
2691 }
2692}
2693
2694impl TryFrom<chroma_proto::SearchResult> for SearchResult {
2695 type Error = QueryConversionError;
2696
2697 fn try_from(value: chroma_proto::SearchResult) -> Result<Self, Self::Error> {
2698 Ok(Self {
2699 results: value
2700 .results
2701 .into_iter()
2702 .map(TryInto::try_into)
2703 .collect::<Result<_, _>>()?,
2704 pulled_log_bytes: value.pulled_log_bytes,
2705 })
2706 }
2707}
2708
2709impl TryFrom<SearchResult> for chroma_proto::SearchResult {
2710 type Error = QueryConversionError;
2711
2712 fn try_from(value: SearchResult) -> Result<Self, Self::Error> {
2713 Ok(Self {
2714 results: value
2715 .results
2716 .into_iter()
2717 .map(TryInto::try_into)
2718 .collect::<Result<Vec<_>, _>>()?,
2719 pulled_log_bytes: value.pulled_log_bytes,
2720 })
2721 }
2722}
2723
2724pub fn rrf(
2869 ranks: Vec<RankExpr>,
2870 k: Option<u32>,
2871 weights: Option<Vec<f32>>,
2872 normalize: bool,
2873) -> Result<RankExpr, QueryConversionError> {
2874 let k = k.unwrap_or(60);
2875
2876 if ranks.is_empty() {
2877 return Err(QueryConversionError::validation(
2878 "RRF requires at least one rank expression",
2879 ));
2880 }
2881
2882 let weights = weights.unwrap_or_else(|| vec![1.0; ranks.len()]);
2883
2884 if weights.len() != ranks.len() {
2885 return Err(QueryConversionError::validation(format!(
2886 "RRF weights length ({}) must match ranks length ({})",
2887 weights.len(),
2888 ranks.len()
2889 )));
2890 }
2891
2892 let weights = if normalize {
2893 let sum: f32 = weights.iter().sum();
2894 if sum == 0.0 {
2895 return Err(QueryConversionError::validation(
2896 "RRF weights sum to zero, cannot normalize",
2897 ));
2898 }
2899 weights.into_iter().map(|w| w / sum).collect()
2900 } else {
2901 weights
2902 };
2903
2904 let terms: Vec<RankExpr> = weights
2905 .into_iter()
2906 .zip(ranks)
2907 .map(|(w, rank)| RankExpr::Value(w) / (RankExpr::Value(k as f32) + rank))
2908 .collect();
2909
2910 let sum = terms
2913 .into_iter()
2914 .reduce(|a, b| a + b)
2915 .unwrap_or(RankExpr::Value(0.0));
2916 Ok(-sum)
2917}
2918
2919#[cfg(test)]
2920mod tests {
2921 use super::*;
2922
2923 #[test]
2924 fn test_key_from_string() {
2925 assert_eq!(Key::from("#document"), Key::Document);
2927 assert_eq!(Key::from("#embedding"), Key::Embedding);
2928 assert_eq!(Key::from("#metadata"), Key::Metadata);
2929 assert_eq!(Key::from("#score"), Key::Score);
2930
2931 assert_eq!(
2933 Key::from("custom_field"),
2934 Key::MetadataField("custom_field".to_string())
2935 );
2936 assert_eq!(
2937 Key::from("author"),
2938 Key::MetadataField("author".to_string())
2939 );
2940
2941 assert_eq!(Key::from("#embedding".to_string()), Key::Embedding);
2943 assert_eq!(
2944 Key::from("year".to_string()),
2945 Key::MetadataField("year".to_string())
2946 );
2947 }
2948
2949 #[test]
2950 fn test_query_vector_dense_proto_conversion() {
2951 let dense_vec = vec![0.1, 0.2, 0.3, 0.4, 0.5];
2952 let query_vector = QueryVector::Dense(dense_vec.clone());
2953
2954 let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2956
2957 let converted: QueryVector = proto.try_into().unwrap();
2959
2960 assert_eq!(converted, query_vector);
2961 if let QueryVector::Dense(v) = converted {
2962 assert_eq!(v, dense_vec);
2963 } else {
2964 panic!("Expected dense vector");
2965 }
2966 }
2967
2968 #[test]
2969 fn test_query_vector_sparse_proto_conversion() {
2970 let sparse = SparseVector::new(vec![0, 5, 10], vec![0.1, 0.5, 0.9]).unwrap();
2971 let query_vector = QueryVector::Sparse(sparse.clone());
2972
2973 let proto: chroma_proto::QueryVector = query_vector.clone().try_into().unwrap();
2975
2976 let converted: QueryVector = proto.try_into().unwrap();
2978
2979 assert_eq!(converted, query_vector);
2980 if let QueryVector::Sparse(s) = converted {
2981 assert_eq!(s, sparse);
2982 } else {
2983 panic!("Expected sparse vector");
2984 }
2985 }
2986
2987 #[test]
2988 fn test_filter_json_deserialization() {
2989 let simple_where = r#"{"author": "John Doe"}"#;
2993 let filter: Filter = serde_json::from_str(simple_where).unwrap();
2994 assert_eq!(filter.query_ids, None);
2995 assert!(filter.where_clause.is_some());
2996
2997 let id_filter_json = serde_json::json!({
2999 "#id": {
3000 "$in": ["doc1", "doc2", "doc3"]
3001 }
3002 });
3003 let filter: Filter = serde_json::from_value(id_filter_json).unwrap();
3004 assert_eq!(filter.query_ids, None);
3005 assert!(filter.where_clause.is_some());
3006
3007 let complex_json = serde_json::json!({
3009 "$and": [
3010 {
3011 "#id": {
3012 "$in": ["doc1", "doc2", "doc3"]
3013 }
3014 },
3015 {
3016 "$or": [
3017 {
3018 "author": {
3019 "$eq": "John Doe"
3020 }
3021 },
3022 {
3023 "author": {
3024 "$eq": "Jane Smith"
3025 }
3026 }
3027 ]
3028 },
3029 {
3030 "year": {
3031 "$gte": 2020
3032 }
3033 },
3034 {
3035 "tags": {
3036 "$contains": "machine-learning"
3037 }
3038 }
3039 ]
3040 });
3041
3042 let filter: Filter = serde_json::from_value(complex_json.clone()).unwrap();
3043 assert_eq!(filter.query_ids, None);
3044 assert!(filter.where_clause.is_some());
3045
3046 if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3048 assert_eq!(composite.operator, crate::metadata::BooleanOperator::And);
3049 assert_eq!(composite.children.len(), 4);
3050
3051 if let crate::metadata::Where::Composite(or_composite) = &composite.children[1] {
3053 assert_eq!(or_composite.operator, crate::metadata::BooleanOperator::Or);
3054 assert_eq!(or_composite.children.len(), 2);
3055 } else {
3056 panic!("Expected OR composite in second child");
3057 }
3058 } else {
3059 panic!("Expected AND composite where clause");
3060 }
3061
3062 let mixed_operators_json = serde_json::json!({
3064 "$and": [
3065 {
3066 "status": {
3067 "$ne": "deleted"
3068 }
3069 },
3070 {
3071 "score": {
3072 "$gt": 0.5
3073 }
3074 },
3075 {
3076 "score": {
3077 "$lt": 0.9
3078 }
3079 },
3080 {
3081 "priority": {
3082 "$lte": 10
3083 }
3084 }
3085 ]
3086 });
3087
3088 let filter: Filter = serde_json::from_value(mixed_operators_json).unwrap();
3089 assert_eq!(filter.query_ids, None);
3090 assert!(filter.where_clause.is_some());
3091
3092 let deeply_nested_json = serde_json::json!({
3094 "$or": [
3095 {
3096 "$and": [
3097 {
3098 "#id": {
3099 "$in": ["id1", "id2"]
3100 }
3101 },
3102 {
3103 "$or": [
3104 {
3105 "category": "tech"
3106 },
3107 {
3108 "category": "science"
3109 }
3110 ]
3111 }
3112 ]
3113 },
3114 {
3115 "$and": [
3116 {
3117 "author": "Admin"
3118 },
3119 {
3120 "published": true
3121 }
3122 ]
3123 }
3124 ]
3125 });
3126
3127 let filter: Filter = serde_json::from_value(deeply_nested_json).unwrap();
3128 assert_eq!(filter.query_ids, None);
3129 assert!(filter.where_clause.is_some());
3130
3131 if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3133 assert_eq!(composite.operator, crate::metadata::BooleanOperator::Or);
3134 assert_eq!(composite.children.len(), 2);
3135
3136 for child in &composite.children {
3138 if let crate::metadata::Where::Composite(and_composite) = child {
3139 assert_eq!(
3140 and_composite.operator,
3141 crate::metadata::BooleanOperator::And
3142 );
3143 } else {
3144 panic!("Expected AND composite in OR children");
3145 }
3146 }
3147 } else {
3148 panic!("Expected OR composite at top level");
3149 }
3150
3151 let single_id_json = serde_json::json!({
3153 "#id": {
3154 "$eq": "single-doc-id"
3155 }
3156 });
3157
3158 let filter: Filter = serde_json::from_value(single_id_json).unwrap();
3159 assert_eq!(filter.query_ids, None);
3160 assert!(filter.where_clause.is_some());
3161
3162 let empty_json = serde_json::json!({});
3164 let filter: Filter = serde_json::from_value(empty_json).unwrap();
3165 assert_eq!(filter.query_ids, None);
3166 assert_eq!(filter.where_clause, None);
3168
3169 let advanced_json = serde_json::json!({
3171 "$and": [
3172 {
3173 "#id": {
3174 "$in": ["doc1", "doc2", "doc3", "doc4", "doc5"]
3175 }
3176 },
3177 {
3178 "tags": {
3179 "$not_contains": "deprecated"
3180 }
3181 },
3182 {
3183 "$or": [
3184 {
3185 "$and": [
3186 {
3187 "confidence": {
3188 "$gte": 0.8
3189 }
3190 },
3191 {
3192 "verified": true
3193 }
3194 ]
3195 },
3196 {
3197 "$and": [
3198 {
3199 "confidence": {
3200 "$gte": 0.6
3201 }
3202 },
3203 {
3204 "confidence": {
3205 "$lt": 0.8
3206 }
3207 },
3208 {
3209 "reviews": {
3210 "$gte": 5
3211 }
3212 }
3213 ]
3214 }
3215 ]
3216 }
3217 ]
3218 });
3219
3220 let filter: Filter = serde_json::from_value(advanced_json).unwrap();
3221 assert_eq!(filter.query_ids, None);
3222 assert!(filter.where_clause.is_some());
3223
3224 if let crate::metadata::Where::Composite(composite) = filter.where_clause.unwrap() {
3226 assert_eq!(composite.operator, crate::metadata::BooleanOperator::And);
3227 assert_eq!(composite.children.len(), 3);
3228 } else {
3229 panic!("Expected AND composite at top level");
3230 }
3231 }
3232
3233 #[test]
3234 fn test_limit_json_serialization() {
3235 let limit = Limit {
3236 offset: 10,
3237 limit: Some(20),
3238 };
3239
3240 let json = serde_json::to_string(&limit).unwrap();
3241 let deserialized: Limit = serde_json::from_str(&json).unwrap();
3242
3243 assert_eq!(deserialized.offset, limit.offset);
3244 assert_eq!(deserialized.limit, limit.limit);
3245 }
3246
3247 #[test]
3248 fn test_query_vector_json_serialization() {
3249 let dense = QueryVector::Dense(vec![0.1, 0.2, 0.3]);
3251 let json = serde_json::to_string(&dense).unwrap();
3252 let deserialized: QueryVector = serde_json::from_str(&json).unwrap();
3253 assert_eq!(deserialized, dense);
3254
3255 let sparse =
3257 QueryVector::Sparse(SparseVector::new(vec![0, 5, 10], vec![0.1, 0.5, 0.9]).unwrap());
3258 let json = serde_json::to_string(&sparse).unwrap();
3259 let deserialized: QueryVector = serde_json::from_str(&json).unwrap();
3260 assert_eq!(deserialized, sparse);
3261 }
3262
3263 #[test]
3264 fn test_select_key_json_serialization() {
3265 use std::collections::HashSet;
3266
3267 let doc_key = Key::Document;
3269 assert_eq!(serde_json::to_string(&doc_key).unwrap(), "\"#document\"");
3270
3271 let embed_key = Key::Embedding;
3272 assert_eq!(serde_json::to_string(&embed_key).unwrap(), "\"#embedding\"");
3273
3274 let meta_key = Key::Metadata;
3275 assert_eq!(serde_json::to_string(&meta_key).unwrap(), "\"#metadata\"");
3276
3277 let score_key = Key::Score;
3278 assert_eq!(serde_json::to_string(&score_key).unwrap(), "\"#score\"");
3279
3280 let custom_key = Key::MetadataField("custom_key".to_string());
3282 assert_eq!(
3283 serde_json::to_string(&custom_key).unwrap(),
3284 "\"custom_key\""
3285 );
3286
3287 let deserialized: Key = serde_json::from_str("\"#document\"").unwrap();
3289 assert!(matches!(deserialized, Key::Document));
3290
3291 let deserialized: Key = serde_json::from_str("\"custom_field\"").unwrap();
3292 assert!(matches!(deserialized, Key::MetadataField(s) if s == "custom_field"));
3293
3294 let mut keys = HashSet::new();
3296 keys.insert(Key::Document);
3297 keys.insert(Key::Embedding);
3298 keys.insert(Key::MetadataField("author".to_string()));
3299
3300 let select = Select { keys };
3301 let json = serde_json::to_string(&select).unwrap();
3302 let deserialized: Select = serde_json::from_str(&json).unwrap();
3303
3304 assert_eq!(deserialized.keys.len(), 3);
3305 assert!(deserialized.keys.contains(&Key::Document));
3306 assert!(deserialized.keys.contains(&Key::Embedding));
3307 assert!(deserialized
3308 .keys
3309 .contains(&Key::MetadataField("author".to_string())));
3310 }
3311
3312 #[test]
3313 fn test_merge_basic_integers() {
3314 use std::cmp::Reverse;
3315
3316 let merge = Merge { k: 5 };
3317
3318 let input = vec![
3320 vec![Reverse(1), Reverse(4), Reverse(7), Reverse(10)],
3321 vec![Reverse(2), Reverse(5), Reverse(8)],
3322 vec![Reverse(3), Reverse(6), Reverse(9), Reverse(11), Reverse(12)],
3323 ];
3324
3325 let result = merge.merge(input);
3326
3327 assert_eq!(result.len(), 5);
3329 assert_eq!(
3330 result,
3331 vec![Reverse(1), Reverse(2), Reverse(3), Reverse(4), Reverse(5)]
3332 );
3333 }
3334
3335 #[test]
3336 fn test_merge_u32_descending() {
3337 let merge = Merge { k: 6 };
3338
3339 let input = vec![
3341 vec![100u32, 75, 50, 25],
3342 vec![90, 60, 30],
3343 vec![95, 85, 70, 40, 10],
3344 ];
3345
3346 let result = merge.merge(input);
3347
3348 assert_eq!(result.len(), 6);
3350 assert_eq!(result, vec![100, 95, 90, 85, 75, 70]);
3351 }
3352
3353 #[test]
3354 fn test_merge_i32_descending() {
3355 let merge = Merge { k: 5 };
3356
3357 let input = vec![
3359 vec![50i32, 10, -10, -50],
3360 vec![30, 0, -30],
3361 vec![40, 20, -20, -40],
3362 ];
3363
3364 let result = merge.merge(input);
3365
3366 assert_eq!(result.len(), 5);
3368 assert_eq!(result, vec![50, 40, 30, 20, 10]);
3369 }
3370
3371 #[test]
3372 fn test_merge_with_duplicates() {
3373 let merge = Merge { k: 10 };
3374
3375 let input = vec![
3377 vec![100u32, 80, 80, 60, 40],
3378 vec![90, 80, 50, 30],
3379 vec![100, 70, 60, 20],
3380 ];
3381
3382 let result = merge.merge(input);
3383
3384 assert_eq!(result, vec![100, 90, 80, 70, 60, 50, 40, 30, 20]);
3386 }
3387
3388 #[test]
3389 fn test_merge_empty_vectors() {
3390 let merge = Merge { k: 5 };
3391
3392 let input: Vec<Vec<u32>> = vec![vec![], vec![], vec![]];
3394 let result = merge.merge(input);
3395 assert_eq!(result.len(), 0);
3396
3397 let input = vec![vec![], vec![1000u64, 750, 500], vec![], vec![850, 600]];
3399 let result = merge.merge(input);
3400 assert_eq!(result, vec![1000, 850, 750, 600, 500]);
3401
3402 let input = vec![vec![], vec![100i32, 50, 25], vec![]];
3404 let result = merge.merge(input);
3405 assert_eq!(result, vec![100, 50, 25]);
3406 }
3407
3408 #[test]
3409 fn test_merge_k_boundary_conditions() {
3410 let merge = Merge { k: 0 };
3412 let input = vec![vec![100u32, 50], vec![75, 25]];
3413 let result = merge.merge(input);
3414 assert_eq!(result.len(), 0);
3415
3416 let merge = Merge { k: 1 };
3418 let input = vec![vec![1000i64, 500], vec![750, 250], vec![900, 100]];
3419 let result = merge.merge(input);
3420 assert_eq!(result, vec![1000]);
3421
3422 let merge = Merge { k: 100 };
3424 let input = vec![vec![10000u128, 5000], vec![8000, 3000]];
3425 let result = merge.merge(input);
3426 assert_eq!(result, vec![10000, 8000, 5000, 3000]);
3427 }
3428
3429 #[test]
3430 fn test_merge_with_strings() {
3431 let merge = Merge { k: 4 };
3432
3433 let input = vec![
3435 vec!["zebra".to_string(), "dog".to_string(), "apple".to_string()],
3436 vec!["elephant".to_string(), "banana".to_string()],
3437 vec!["fish".to_string(), "cat".to_string()],
3438 ];
3439
3440 let result = merge.merge(input);
3441
3442 assert_eq!(
3444 result,
3445 vec![
3446 "zebra".to_string(),
3447 "fish".to_string(),
3448 "elephant".to_string(),
3449 "dog".to_string()
3450 ]
3451 );
3452 }
3453
3454 #[test]
3455 fn test_merge_with_custom_struct() {
3456 #[derive(Debug, Clone, Hash, PartialEq, Eq, PartialOrd, Ord)]
3457 struct Score {
3458 value: i32,
3459 id: String,
3460 }
3461
3462 let merge = Merge { k: 3 };
3463
3464 let input = vec![
3466 vec![
3467 Score {
3468 value: 100,
3469 id: "a".to_string(),
3470 },
3471 Score {
3472 value: 80,
3473 id: "b".to_string(),
3474 },
3475 Score {
3476 value: 60,
3477 id: "c".to_string(),
3478 },
3479 ],
3480 vec![
3481 Score {
3482 value: 90,
3483 id: "d".to_string(),
3484 },
3485 Score {
3486 value: 70,
3487 id: "e".to_string(),
3488 },
3489 ],
3490 vec![
3491 Score {
3492 value: 95,
3493 id: "f".to_string(),
3494 },
3495 Score {
3496 value: 85,
3497 id: "g".to_string(),
3498 },
3499 ],
3500 ];
3501
3502 let result = merge.merge(input);
3503
3504 assert_eq!(result.len(), 3);
3505 assert_eq!(
3506 result[0],
3507 Score {
3508 value: 100,
3509 id: "a".to_string()
3510 }
3511 );
3512 assert_eq!(
3513 result[1],
3514 Score {
3515 value: 95,
3516 id: "f".to_string()
3517 }
3518 );
3519 assert_eq!(
3520 result[2],
3521 Score {
3522 value: 90,
3523 id: "d".to_string()
3524 }
3525 );
3526 }
3527
3528 #[test]
3529 fn test_merge_preserves_order() {
3530 use std::cmp::Reverse;
3531
3532 let merge = Merge { k: 10 };
3533
3534 let input = vec![
3537 vec![Reverse(2), Reverse(6), Reverse(10), Reverse(14)],
3538 vec![Reverse(4), Reverse(8), Reverse(12), Reverse(16)],
3539 vec![Reverse(1), Reverse(3), Reverse(5), Reverse(7), Reverse(9)],
3540 ];
3541
3542 let result = merge.merge(input);
3543
3544 for i in 1..result.len() {
3547 assert!(
3548 result[i - 1] >= result[i],
3549 "Output should be in descending Reverse order"
3550 );
3551 assert!(
3552 result[i - 1].0 <= result[i].0,
3553 "Inner values should be in ascending order"
3554 );
3555 }
3556
3557 assert_eq!(
3559 result,
3560 vec![
3561 Reverse(1),
3562 Reverse(2),
3563 Reverse(3),
3564 Reverse(4),
3565 Reverse(5),
3566 Reverse(6),
3567 Reverse(7),
3568 Reverse(8),
3569 Reverse(9),
3570 Reverse(10)
3571 ]
3572 );
3573 }
3574
3575 #[test]
3576 fn test_merge_single_vector() {
3577 let merge = Merge { k: 3 };
3578
3579 let input = vec![vec![1000u64, 800, 600, 400, 200]];
3581
3582 let result = merge.merge(input);
3583
3584 assert_eq!(result, vec![1000, 800, 600]);
3585 }
3586
3587 #[test]
3588 fn test_merge_all_same_values() {
3589 let merge = Merge { k: 5 };
3590
3591 let input = vec![vec![42i16, 42, 42], vec![42, 42], vec![42, 42, 42, 42]];
3593
3594 let result = merge.merge(input);
3595
3596 assert_eq!(result, vec![42]);
3598 }
3599
3600 #[test]
3601 fn test_merge_mixed_types_sizes() {
3602 let merge = Merge { k: 4 };
3604 let input = vec![
3605 vec![1000usize, 500, 100],
3606 vec![800, 300],
3607 vec![900, 600, 200],
3608 ];
3609 let result = merge.merge(input);
3610 assert_eq!(result, vec![1000, 900, 800, 600]);
3611
3612 let merge = Merge { k: 5 };
3614 let input = vec![vec![10i32, 0, -10, -20], vec![5, -5, -15], vec![15, -25]];
3615 let result = merge.merge(input);
3616 assert_eq!(result, vec![15, 10, 5, 0, -5]);
3617 }
3618
3619 #[test]
3620 fn test_merge_dedup_same_id_different_scores() {
3621 use std::cmp::Reverse;
3622
3623 let merge = Merge { k: 5 };
3627
3628 let input: Vec<Vec<Reverse<RecordMeasure>>> = vec![
3632 vec![
3633 Reverse(RecordMeasure {
3634 offset_id: 1,
3635 measure: 0.10,
3636 }),
3637 Reverse(RecordMeasure {
3638 offset_id: 4,
3639 measure: 0.50,
3640 }),
3641 Reverse(RecordMeasure {
3642 offset_id: 5,
3643 measure: 0.70,
3644 }),
3645 ],
3646 vec![
3647 Reverse(RecordMeasure {
3648 offset_id: 2,
3649 measure: 0.20,
3650 }),
3651 Reverse(RecordMeasure {
3652 offset_id: 1,
3653 measure: 0.25,
3654 }), Reverse(RecordMeasure {
3656 offset_id: 3,
3657 measure: 0.60,
3658 }),
3659 ],
3660 vec![
3661 Reverse(RecordMeasure {
3662 offset_id: 3,
3663 measure: 0.15,
3664 }), Reverse(RecordMeasure {
3666 offset_id: 4,
3667 measure: 0.35,
3668 }), Reverse(RecordMeasure {
3670 offset_id: 2,
3671 measure: 0.80,
3672 }), ],
3674 ];
3675
3676 let result: Vec<Reverse<RecordMeasure>> = merge.merge(input);
3686 let ids: Vec<u32> = result.iter().map(|Reverse(r)| r.offset_id).collect();
3687 let measures: Vec<f32> = result.iter().map(|Reverse(r)| r.measure).collect();
3688
3689 assert_eq!(ids, vec![1, 3, 2, 4, 5]);
3690 assert_eq!(measures, vec![0.10, 0.15, 0.20, 0.35, 0.70]);
3691 }
3692
3693 #[test]
3694 fn test_aggregate_json_serialization() {
3695 let min_k = Aggregate::MinK {
3697 keys: vec![Key::Score, Key::field("date")],
3698 k: 3,
3699 };
3700 let json = serde_json::to_value(&min_k).unwrap();
3701 assert!(json.get("$min_k").is_some());
3702 assert_eq!(json["$min_k"]["k"], 3);
3703
3704 let min_k_json = serde_json::json!({
3706 "$min_k": {
3707 "keys": ["#score", "date"],
3708 "k": 5
3709 }
3710 });
3711 let deserialized: Aggregate = serde_json::from_value(min_k_json).unwrap();
3712 match deserialized {
3713 Aggregate::MinK { keys, k } => {
3714 assert_eq!(k, 5);
3715 assert_eq!(keys.len(), 2);
3716 assert_eq!(keys[0], Key::Score);
3717 assert_eq!(keys[1], Key::field("date"));
3718 }
3719 _ => panic!("Expected MinK"),
3720 }
3721
3722 let max_k = Aggregate::MaxK {
3724 keys: vec![Key::field("timestamp")],
3725 k: 10,
3726 };
3727 let json = serde_json::to_value(&max_k).unwrap();
3728 assert!(json.get("$max_k").is_some());
3729 assert_eq!(json["$max_k"]["k"], 10);
3730
3731 let max_k_json = serde_json::json!({
3733 "$max_k": {
3734 "keys": ["timestamp"],
3735 "k": 2
3736 }
3737 });
3738 let deserialized: Aggregate = serde_json::from_value(max_k_json).unwrap();
3739 match deserialized {
3740 Aggregate::MaxK { keys, k } => {
3741 assert_eq!(k, 2);
3742 assert_eq!(keys.len(), 1);
3743 assert_eq!(keys[0], Key::field("timestamp"));
3744 }
3745 _ => panic!("Expected MaxK"),
3746 }
3747 }
3748
3749 #[test]
3750 fn test_group_by_json_serialization() {
3751 let group_by = GroupBy {
3753 keys: vec![Key::field("category"), Key::field("author")],
3754 aggregate: Some(Aggregate::MinK {
3755 keys: vec![Key::Score],
3756 k: 3,
3757 }),
3758 };
3759
3760 let json = serde_json::to_value(&group_by).unwrap();
3761 assert_eq!(json["keys"].as_array().unwrap().len(), 2);
3762 assert!(json["aggregate"]["$min_k"].is_object());
3763
3764 let deserialized: GroupBy = serde_json::from_value(json).unwrap();
3766 assert_eq!(deserialized.keys.len(), 2);
3767 assert_eq!(deserialized.keys[0], Key::field("category"));
3768 assert_eq!(deserialized.keys[1], Key::field("author"));
3769 assert!(deserialized.aggregate.is_some());
3770
3771 let empty_group_by = GroupBy::default();
3773 let json = serde_json::to_value(&empty_group_by).unwrap();
3774 let deserialized: GroupBy = serde_json::from_value(json).unwrap();
3775 assert!(deserialized.keys.is_empty());
3776 assert!(deserialized.aggregate.is_none());
3777
3778 let json = serde_json::json!({
3780 "keys": ["category"],
3781 "aggregate": {
3782 "$max_k": {
3783 "keys": ["#score", "priority"],
3784 "k": 5
3785 }
3786 }
3787 });
3788 let group_by: GroupBy = serde_json::from_value(json).unwrap();
3789 assert_eq!(group_by.keys.len(), 1);
3790 assert_eq!(group_by.keys[0], Key::field("category"));
3791 match group_by.aggregate {
3792 Some(Aggregate::MaxK { keys, k }) => {
3793 assert_eq!(k, 5);
3794 assert_eq!(keys.len(), 2);
3795 assert_eq!(keys[0], Key::Score);
3796 }
3797 _ => panic!("Expected MaxK aggregate"),
3798 }
3799 }
3800
3801 fn sparse_knn_leaf(index: u32, key: &str, return_rank: bool) -> RankExpr {
3802 RankExpr::Knn {
3803 query: QueryVector::Sparse(SparseVector::new(vec![index], vec![1.0]).unwrap()),
3804 key: Key::field(key),
3805 limit: 10,
3806 default: None,
3807 return_rank,
3808 }
3809 }
3810
3811 #[test]
3812 fn test_knn_queries_collects_all_sparse_leaves_in_dfs_order() {
3813 let expr = RankExpr::Summation(vec![
3817 sparse_knn_leaf(0, "sparse_a", false),
3818 sparse_knn_leaf(1, "sparse_b", false),
3819 sparse_knn_leaf(2, "sparse_a", false),
3820 ]);
3821
3822 let leaves = expr.knn_queries();
3823 assert_eq!(leaves.len(), 3);
3824 assert_eq!(leaves[0].key, Key::field("sparse_a"));
3825 assert_eq!(leaves[1].key, Key::field("sparse_b"));
3826 assert_eq!(leaves[2].key, Key::field("sparse_a"));
3827
3828 match (&leaves[0].query, &leaves[2].query) {
3830 (QueryVector::Sparse(first), QueryVector::Sparse(third)) => {
3831 assert_ne!(first.indices, third.indices);
3832 }
3833 _ => panic!("expected sparse query vectors"),
3834 }
3835 }
3836
3837 #[test]
3838 fn test_rrf_preserves_all_sparse_leaves() {
3839 let expr = rrf(
3842 vec![
3843 sparse_knn_leaf(0, "sparse_a", true),
3844 sparse_knn_leaf(1, "sparse_b", true),
3845 sparse_knn_leaf(2, "sparse_c", true),
3846 ],
3847 None,
3848 None,
3849 false,
3850 )
3851 .expect("rrf should build");
3852
3853 let keys: Vec<Key> = expr.knn_queries().into_iter().map(|q| q.key).collect();
3854 assert_eq!(
3855 keys,
3856 vec![
3857 Key::field("sparse_a"),
3858 Key::field("sparse_b"),
3859 Key::field("sparse_c"),
3860 ]
3861 );
3862 }
3863}