diff --git a/lib/collection/src/collection.rs b/lib/collection/src/collection.rs index 3513009e54..c8d6ee4ec5 100644 --- a/lib/collection/src/collection.rs +++ b/lib/collection/src/collection.rs @@ -12,7 +12,7 @@ use tokio::runtime::{Handle, Runtime}; use segment::types::{ Condition, Filter, HasIdCondition, PayloadKeyType, PayloadSchemaInfo, PointIdType, ScoredPoint, - SegmentType, VectorElementType, + SegmentType, VectorElementType, WithPayload, }; use crate::collection_builder::optimizers_builder::build_optimizers; @@ -123,9 +123,10 @@ impl Collection { let limit = request .limit .unwrap_or_else(|| default_request.limit.unwrap()); - let with_payload = request + let with_payload_interface = &request .with_payload - .unwrap_or_else(|| default_request.with_payload.unwrap()); + .clone() + .unwrap_or_else(|| default_request.with_payload.clone().unwrap()); let with_vector = request .with_vector .unwrap_or_else(|| default_request.with_vector.unwrap()); @@ -154,8 +155,9 @@ impl Collection { .take(limit) .collect_vec(); + let with_payload = WithPayload::from(with_payload_interface); let mut points = segment_searcher - .retrieve(segments, &point_ids, with_payload, with_vector) + .retrieve(segments, &point_ids, &with_payload, with_vector) .await?; points.sort_by_key(|point| point.id); @@ -192,7 +194,12 @@ impl Collection { .collect_vec(); let vectors = segment_searcher - .retrieve(segments, &reference_vectors_ids, false, true) + .retrieve( + segments, + &reference_vectors_ids, + &WithPayload::from(true), + true, + ) .await?; let vectors_map: HashMap> = vectors .into_iter() @@ -244,6 +251,7 @@ impl Collection { has_id: reference_vectors_ids.iter().cloned().collect(), })]), }), + with_payload: None, params: request.params, top: request.top, }; diff --git a/lib/collection/src/collection_manager/collection_managers.rs b/lib/collection/src/collection_manager/collection_managers.rs index d5c23003fb..75bfc43181 100644 --- a/lib/collection/src/collection_manager/collection_managers.rs +++ b/lib/collection/src/collection_manager/collection_managers.rs @@ -3,7 +3,7 @@ use std::sync::Arc; use parking_lot::RwLock; use tokio::runtime::Handle; -use segment::types::{PointIdType, ScoredPoint, SeqNumberType}; +use segment::types::{PointIdType, ScoredPoint, SeqNumberType, WithPayload}; use crate::collection_manager::holders::segment_holder::SegmentHolder; use crate::operations::types::{CollectionResult, Record, SearchRequest}; @@ -23,7 +23,7 @@ pub trait CollectionSearcher { &self, segments: &RwLock, points: &[PointIdType], - with_payload: bool, + with_payload: &WithPayload, with_vector: bool, ) -> CollectionResult>; } diff --git a/lib/collection/src/collection_manager/holders/proxy_segment.rs b/lib/collection/src/collection_manager/holders/proxy_segment.rs index b3736d54bc..7be22cee28 100644 --- a/lib/collection/src/collection_manager/holders/proxy_segment.rs +++ b/lib/collection/src/collection_manager/holders/proxy_segment.rs @@ -4,7 +4,7 @@ use segment::entry::entry_point::{OperationResult, SegmentEntry}; use segment::types::{ Condition, Filter, PayloadKeyType, PayloadKeyTypeRef, PayloadType, PointIdType, ScoredPoint, SearchParams, SegmentConfig, SegmentInfo, SegmentType, SeqNumberType, TheMap, - VectorElementType, + VectorElementType, WithPayload, }; use std::cmp::max; use std::collections::HashSet; @@ -108,6 +108,7 @@ impl SegmentEntry for ProxySegment { fn search( &self, vector: &[VectorElementType], + with_payload: &WithPayload, filter: Option<&Filter>, top: usize, params: Option<&SearchParams>, @@ -124,22 +125,25 @@ impl SegmentEntry for ProxySegment { // This copy might slow process down if there will be a lot of deleted points let wrapped_filter = self.add_deleted_points_condition_to_filter(filter); - self.wrapped_segment - .get() - .read() - .search(vector, Some(&wrapped_filter), top, params)? + self.wrapped_segment.get().read().search( + vector, + with_payload, + Some(&wrapped_filter), + top, + params, + )? } else { self.wrapped_segment .get() .read() - .search(vector, filter, top, params)? + .search(vector, with_payload, filter, top, params)? }; - let mut write_result = self - .write_segment - .get() - .read() - .search(vector, filter, top, params)?; + let mut write_result = + self.write_segment + .get() + .read() + .search(vector, with_payload, filter, top, params)?; wrapped_result.append(&mut write_result); Ok(wrapped_result) @@ -456,7 +460,9 @@ mod tests { proxy_segment.delete_point(102, 1).unwrap(); let query_vector = vec![1.0, 1.0, 1.0, 1.0]; - let search_result = proxy_segment.search(&query_vector, None, 10, None).unwrap(); + let search_result = proxy_segment + .search(&query_vector, &WithPayload::default(), None, 10, None) + .unwrap(); eprintln!("search_result = {:#?}", search_result); diff --git a/lib/collection/src/collection_manager/simple_collection_searcher.rs b/lib/collection/src/collection_manager/simple_collection_searcher.rs index e58f927cd4..3dc856242f 100644 --- a/lib/collection/src/collection_manager/simple_collection_searcher.rs +++ b/lib/collection/src/collection_manager/simple_collection_searcher.rs @@ -6,7 +6,7 @@ use parking_lot::RwLock; use tokio::runtime::Handle; use segment::spaces::tools::peek_top_scores_iterable; -use segment::types::{PointIdType, ScoredPoint, SeqNumberType}; +use segment::types::{PointIdType, ScoredPoint, SeqNumberType, WithPayload, WithPayloadInterface}; use crate::collection_manager::collection_managers::CollectionSearcher; use crate::collection_manager::holders::segment_holder::{LockedSegment, SegmentHolder}; @@ -85,7 +85,7 @@ impl CollectionSearcher for SimpleCollectionSearcher { &self, segments: &RwLock, points: &[PointIdType], - with_payload: bool, + with_payload: &WithPayload, with_vector: bool, ) -> CollectionResult> { let mut point_version: HashMap = Default::default(); @@ -98,8 +98,12 @@ impl CollectionSearcher for SimpleCollectionSearcher { id, Record { id, - payload: if with_payload { - Some(segment.payload(id)?) + payload: if with_payload.enable { + if let Some(selector) = &with_payload.payload_selector { + Some(selector.process(segment.payload(id)?)) + } else { + Some(segment.payload(id)?) + } } else { None }, @@ -122,8 +126,14 @@ async fn search_in_segment( segment: LockedSegment, request: Arc, ) -> CollectionResult> { + let with_payload_interface = request + .with_payload + .as_ref() + .unwrap_or(&WithPayloadInterface::Bool(false)); + let with_payload = WithPayload::from(with_payload_interface); let res = segment.get().read().search( &request.vector, + &with_payload, request.filter.as_ref(), request.top, request.params.as_ref(), @@ -152,6 +162,7 @@ mod tests { let req = Arc::new(SearchRequest { vector: query, + with_payload: None, filter: None, params: None, top: 5, @@ -178,7 +189,7 @@ mod tests { let searcher = SimpleCollectionSearcher::new(); let records = searcher - .retrieve(&segment_holder, &[1, 2, 3], true, true) + .retrieve(&segment_holder, &[1, 2, 3], &WithPayload::from(true), true) .await .unwrap(); assert_eq!(records.len(), 3); diff --git a/lib/collection/src/collection_manager/simple_collection_updater.rs b/lib/collection/src/collection_manager/simple_collection_updater.rs index 57269df691..e7a2684519 100644 --- a/lib/collection/src/collection_manager/simple_collection_updater.rs +++ b/lib/collection/src/collection_manager/simple_collection_updater.rs @@ -4,7 +4,9 @@ use segment::types::SeqNumberType; use crate::collection_manager::collection_managers::CollectionUpdater; use crate::collection_manager::holders::segment_holder::SegmentHolder; -use crate::collection_manager::segments_updater::*; +use crate::collection_manager::segments_updater::{ + process_field_index_operation, process_payload_operation, process_point_operation, +}; use crate::operations::types::CollectionResult; use crate::operations::CollectionUpdateOperations; @@ -44,10 +46,11 @@ impl CollectionUpdater for SimpleCollectionUpdater { mod tests { use tempdir::TempDir; - use segment::types::{PayloadInterface, PayloadKeyType, PayloadVariant}; + use segment::types::{PayloadInterface, PayloadKeyType, PayloadVariant, WithPayload}; use crate::collection_manager::collection_managers::CollectionSearcher; use crate::collection_manager::fixtures::build_test_holder; + use crate::collection_manager::segments_updater::upsert_points; use crate::collection_manager::simple_collection_searcher::SimpleCollectionSearcher; use super::*; @@ -70,7 +73,7 @@ mod tests { assert!(matches!(res, Ok(1))); let records = searcher - .retrieve(&segments, &[1, 2, 500], true, true) + .retrieve(&segments, &[1, 2, 500], &WithPayload::from(true), true) .await .unwrap(); @@ -95,7 +98,7 @@ mod tests { .unwrap(); let records = searcher - .retrieve(&segments, &[1, 2, 500], true, true) + .retrieve(&segments, &[1, 2, 500], &WithPayload::from(true), true) .await .unwrap(); @@ -131,7 +134,7 @@ mod tests { .unwrap(); let res = searcher - .retrieve(&segments, &points, true, false) + .retrieve(&segments, &points, &WithPayload::from(true), false) .await .unwrap(); @@ -159,7 +162,7 @@ mod tests { .unwrap(); let res = searcher - .retrieve(&segments, &[3], true, false) + .retrieve(&segments, &[3], &WithPayload::from(true), false) .await .unwrap(); assert_eq!(res.len(), 1); @@ -168,7 +171,7 @@ mod tests { // Test clear payload let res = searcher - .retrieve(&segments, &[2], true, false) + .retrieve(&segments, &[2], &WithPayload::from(true), false) .await .unwrap(); assert_eq!(res.len(), 1); @@ -181,7 +184,7 @@ mod tests { ) .unwrap(); let res = searcher - .retrieve(&segments, &[2], true, false) + .retrieve(&segments, &[2], &WithPayload::from(true), false) .await .unwrap(); assert_eq!(res.len(), 1); diff --git a/lib/collection/src/operations/types.rs b/lib/collection/src/operations/types.rs index 32f6b1c114..6db7243298 100644 --- a/lib/collection/src/operations/types.rs +++ b/lib/collection/src/operations/types.rs @@ -11,7 +11,7 @@ use tokio::task::JoinError; use segment::entry::entry_point::OperationError; use segment::types::{ Filter, PayloadKeyType, PayloadSchemaInfo, PayloadType, PointIdType, SearchParams, - SeqNumberType, TheMap, VectorElementType, + SeqNumberType, TheMap, VectorElementType, WithPayloadInterface, }; use crate::config::CollectionConfig; @@ -81,7 +81,7 @@ pub struct UpdateResult { pub status: UpdateStatus, } -#[derive(Debug, Deserialize, Serialize, JsonSchema)] +#[derive(Debug, Deserialize, Serialize, JsonSchema, Clone)] #[serde(rename_all = "snake_case")] /// Scroll request - paginate over all points which matches given condition pub struct ScrollRequest { @@ -91,8 +91,8 @@ pub struct ScrollRequest { pub limit: Option, /// Look only for points which satisfies this conditions. If not provided - all points. pub filter: Option, - /// Return point payload with the result. Default: true - pub with_payload: Option, + /// Return point payload with the result. Default: True + pub with_payload: Option, /// Return point vector with the result. Default: false pub with_vector: Option, } @@ -103,7 +103,7 @@ impl Default for ScrollRequest { offset: Some(0), limit: Some(10), filter: None, - with_payload: Some(true), + with_payload: Some(WithPayloadInterface::Bool(true)), with_vector: Some(false), } } @@ -131,6 +131,8 @@ pub struct SearchRequest { pub params: Option, /// Max number of result to return pub top: usize, + /// Payload interface + pub with_payload: Option, } #[derive(Debug, Deserialize, Serialize, JsonSchema)] diff --git a/lib/collection/tests/collection_restore_test.rs b/lib/collection/tests/collection_restore_test.rs index e9ae21a9dd..eb455b674c 100644 --- a/lib/collection/tests/collection_restore_test.rs +++ b/lib/collection/tests/collection_restore_test.rs @@ -8,7 +8,7 @@ use collection::collection_manager::simple_collection_updater::SimpleCollectionU use collection::operations::point_ops::{PointInsertOperations, PointOperations}; use collection::operations::types::ScrollRequest; use collection::operations::CollectionUpdateOperations; -use segment::types::PayloadType; +use segment::types::{PayloadSelector, PayloadType, WithPayloadInterface}; use crate::common::simple_collection_fixture; @@ -52,7 +52,7 @@ async fn test_collection_payload_reloading() { ids: vec![0, 1], vectors: vec![vec![1.0, 0.0, 1.0, 1.0], vec![1.0, 0.0, 1.0, 0.0]], payloads: serde_json::from_str( - &r#"[{ "k": { "type": "keyword", "value": "v1" } }, { "k": "v2" }]"#, + &r#"[{ "k": { "type": "keyword", "value": "v1" } }, { "k": "v2"}]"#, ) .unwrap(), }), @@ -72,7 +72,7 @@ async fn test_collection_payload_reloading() { offset: Some(0), limit: Some(10), filter: None, - with_payload: Some(true), + with_payload: Some(WithPayloadInterface::Bool(true)), with_vector: Some(true), }, &searcher, @@ -98,3 +98,104 @@ async fn test_collection_payload_reloading() { res.points[0].payload.as_ref().unwrap().get("k") ); } + +#[tokio::test] +async fn test_collection_payload_custom_payload() { + let collection_dir = TempDir::new("collection").unwrap(); + let updater = Arc::new(SimpleCollectionUpdater::new()); + { + let collection = simple_collection_fixture(collection_dir.path()).await; + let insert_points = CollectionUpdateOperations::PointOperation( + PointOperations::UpsertPoints(PointInsertOperations::BatchPoints { + ids: vec![0, 1], + vectors: vec![vec![1.0, 0.0, 1.0, 1.0], vec![1.0, 0.0, 1.0, 0.0]], + payloads: serde_json::from_str( + &r#"[{ "k": { "type": "keyword", "value": "v1" } }, { "k": "v2" , "v": "v3", "v2": "v4"}]"#, + ) + .unwrap(), + }), + ); + collection + .update_by(insert_points, true, updater.clone()) + .await + .unwrap(); + } + + let collection = load_collection(collection_dir.path(), updater.clone()); + + let searcher = SimpleCollectionSearcher::new(); + // Test res with filter payload + let res_with_custom_payload = collection + .scroll_by( + ScrollRequest { + offset: Some(0), + limit: Some(10), + filter: None, + with_payload: Some(WithPayloadInterface::Fields(vec![String::from("v")])), + with_vector: Some(true), + }, + &searcher, + ) + .await + .unwrap(); + assert!(res_with_custom_payload.points[0] + .payload + .as_ref() + .expect("has payload") + .is_empty()); + + match res_with_custom_payload.points[1] + .payload + .as_ref() + .expect("has payload") + .get("v") + .expect("has value") + { + PayloadType::Keyword(values) => assert_eq!(&vec!["v3".to_string()], values), + _ => panic!("unexpected type"), + } + + eprintln!( + "res_with_custom_payload = {:#?}", + res_with_custom_payload.points[0].payload.as_ref().unwrap() + ); + + // Test res with filter payload dict + let res_with_custom_payload = collection + .scroll_by( + ScrollRequest { + offset: Some(0), + limit: Some(10), + filter: None, + with_payload: Some(WithPayloadInterface::Selector(PayloadSelector { + include: vec![String::from("v"), String::from("v2")], + exclude: vec![String::from("v")], + })), + with_vector: Some(false), + }, + &searcher, + ) + .await + .unwrap(); + assert!(res_with_custom_payload.points[0] + .payload + .as_ref() + .expect("has payload") + .is_empty()); + + match res_with_custom_payload.points[1] + .payload + .as_ref() + .expect("has payload") + .get("v2") + .expect("has value") + { + PayloadType::Keyword(values) => assert_eq!(&vec!["v4".to_string()], values), + _ => panic!("unexpected type"), + } + + eprintln!( + "res_with_custom_payload = {:#?}", + res_with_custom_payload.points[0].payload.as_ref().unwrap() + ); +} diff --git a/lib/collection/tests/collection_test.rs b/lib/collection/tests/collection_test.rs index 9bdf34fe5d..ac0195a97b 100644 --- a/lib/collection/tests/collection_test.rs +++ b/lib/collection/tests/collection_test.rs @@ -10,7 +10,9 @@ use collection::operations::point_ops::PointInsertOperations::{BatchPoints, Poin use collection::operations::point_ops::{PointOperations, PointStruct}; use collection::operations::types::{RecommendRequest, ScrollRequest, SearchRequest, UpdateStatus}; use collection::operations::CollectionUpdateOperations; -use segment::types::{PayloadInterface, PayloadKeyType, PayloadVariant}; +use segment::types::{ + PayloadInterface, PayloadKeyType, PayloadVariant, WithPayload, WithPayloadInterface, +}; use crate::common::simple_collection_fixture; use collection::collection_manager::collection_managers::CollectionSearcher; @@ -52,6 +54,7 @@ async fn test_collection_updater() { let search_request = SearchRequest { vector: vec![1.0, 1.0, 1.0, 1.0], + with_payload: None, filter: None, params: None, top: 3, @@ -70,6 +73,62 @@ async fn test_collection_updater() { Ok(res) => { assert_eq!(res.len(), 3); assert_eq!(res[0].id, 2); + assert_eq!(res[0].payload.len(), 0); + } + Err(err) => panic!("search failed: {:?}", err), + } +} + +#[tokio::test] +async fn test_collection_search_with_payload() { + let collection_dir = TempDir::new("collection").unwrap(); + + let collection = simple_collection_fixture(collection_dir.path()).await; + let segment_updater = Arc::new(SimpleCollectionUpdater::new()); + + let insert_points = + CollectionUpdateOperations::PointOperation(PointOperations::UpsertPoints(BatchPoints { + ids: vec![0, 1], + vectors: vec![vec![1.0, 0.0, 1.0, 1.0], vec![1.0, 0.0, 1.0, 0.0]], + payloads: serde_json::from_str( + &r#"[{ "k": { "type": "keyword", "value": "v1" } }, { "k": "v2" , "v": "v3"}]"#, + ) + .unwrap(), + })); + + let insert_result = collection + .update_by(insert_points, true, segment_updater.clone()) + .await; + + match insert_result { + Ok(res) => { + assert_eq!(res.status, UpdateStatus::Completed) + } + Err(err) => panic!("operation failed: {:?}", err), + } + + let search_request = SearchRequest { + vector: vec![1.0, 0.0, 1.0, 1.0], + with_payload: Some(WithPayloadInterface::Bool(true)), + filter: None, + params: None, + top: 3, + }; + + let segment_searcher = SimpleCollectionSearcher::new(); + let search_res = segment_searcher + .search( + collection.segments(), + Arc::new(search_request), + &Handle::current(), + ) + .await; + + match search_res { + Ok(res) => { + assert_eq!(res.len(), 2); + assert_eq!(res[0].id, 0); + assert_eq!(res[0].payload.len(), 1); } Err(err) => panic!("search failed: {:?}", err), } @@ -122,7 +181,12 @@ async fn test_collection_loading() { let loaded_collection = load_collection(collection_dir.path(), segment_updater.clone()); let segment_searcher = SimpleCollectionSearcher::new(); let retrieved = segment_searcher - .retrieve(loaded_collection.segments(), &[1, 2], true, true) + .retrieve( + loaded_collection.segments(), + &[1, 2], + &WithPayload::from(true), + true, + ) .await .unwrap(); @@ -236,7 +300,7 @@ async fn test_recommendation_api() { .await .unwrap(); assert!(result.len() > 0); - let top1 = result[0]; + let top1 = &result[0]; assert!(top1.id == 5 || top1.id == 6); } @@ -276,7 +340,7 @@ async fn test_read_api() { offset: Some(0), limit: Some(2), filter: None, - with_payload: Some(true), + with_payload: Some(WithPayloadInterface::Bool(true)), with_vector: None, }, &segment_searcher, diff --git a/lib/segment/src/entry/entry_point.rs b/lib/segment/src/entry/entry_point.rs index 3d687cc483..6adb43a18c 100644 --- a/lib/segment/src/entry/entry_point.rs +++ b/lib/segment/src/entry/entry_point.rs @@ -1,6 +1,6 @@ use crate::types::{ Filter, PayloadKeyType, PayloadKeyTypeRef, PayloadType, PointIdType, ScoredPoint, SearchParams, - SegmentConfig, SegmentInfo, SegmentType, SeqNumberType, TheMap, VectorElementType, + SegmentConfig, SegmentInfo, SegmentType, SeqNumberType, TheMap, VectorElementType, WithPayload, }; use atomicwrites::Error as AtomicIoError; use rocksdb::Error; @@ -84,6 +84,7 @@ pub trait SegmentEntry { fn search( &self, vector: &[VectorElementType], + with_payload: &WithPayload, filter: Option<&Filter>, top: usize, params: Option<&SearchParams>, diff --git a/lib/segment/src/segment.rs b/lib/segment/src/segment.rs index cffaadc845..30b6ecd0c9 100644 --- a/lib/segment/src/segment.rs +++ b/lib/segment/src/segment.rs @@ -6,7 +6,7 @@ use crate::spaces::tools::mertic_object; use crate::types::{ Filter, PayloadKeyType, PayloadKeyTypeRef, PayloadSchemaInfo, PayloadType, PointIdType, PointOffsetType, ScoredPoint, SearchParams, SegmentConfig, SegmentInfo, SegmentState, - SegmentType, SeqNumberType, TheMap, VectorElementType, + SegmentType, SeqNumberType, TheMap, VectorElementType, WithPayload, }; use crate::vector_storage::VectorStorage; use atomic_refcell::AtomicRefCell; @@ -103,6 +103,7 @@ impl SegmentEntry for Segment { fn search( &self, vector: &[VectorElementType], + with_payload: &WithPayload, filter: Option<&Filter>, top: usize, params: Option<&SearchParams>, @@ -121,21 +122,36 @@ impl SegmentEntry for Segment { .search(vector, filter, top, params); let id_mapper = self.id_mapper.borrow(); - let res = internal_result + + let res: OperationResult> = internal_result .iter() - .map(|&scored_point_offset| ScoredPoint { - id: id_mapper + .map(|&scored_point_offset| { + let point_id = id_mapper .external_id(scored_point_offset.idx) - .unwrap_or_else(|| { - panic!( + .ok_or_else(|| OperationError::ServiceError { + description: format!( "Corrupter id_mapper, no external value for {}", scored_point_offset.idx - ) - }), - score: scored_point_offset.score, + ), + })?; + let payload = if with_payload.enable { + let initial_payload = self.payload(point_id)?; + if let Some(i) = &with_payload.payload_selector { + i.process(initial_payload) + } else { + initial_payload + } + } else { + TheMap::new() + }; + Ok(ScoredPoint { + id: point_id, + score: scored_point_offset.score, + payload, + }) }) .collect(); - Ok(res) + res } fn upsert_point( @@ -700,13 +716,20 @@ mod tests { let filter_invalid: Filter = serde_json::from_str(filter_invalid_str).unwrap(); let results_with_valid_filter = segment - .search(&vec![1.0 as f32, 1.0 as f32], Some(&filter_valid), 1, None) + .search( + &vec![1.0 as f32, 1.0 as f32], + &WithPayload::default(), + Some(&filter_valid), + 1, + None, + ) .unwrap(); assert_eq!(results_with_valid_filter.len(), 1); assert_eq!(results_with_valid_filter.first().unwrap().id, 0); let results_with_invalid_filter = segment .search( &vec![1.0 as f32, 1.0 as f32], + &WithPayload::default(), Some(&filter_invalid), 1, None, diff --git a/lib/segment/src/types.rs b/lib/segment/src/types.rs index 8d7f64ca09..67c2444c2c 100644 --- a/lib/segment/src/types.rs +++ b/lib/segment/src/types.rs @@ -39,12 +39,14 @@ pub enum Order { SmallBetter, } -#[derive(Deserialize, Serialize, JsonSchema, Copy, Clone, PartialEq, Debug)] +#[derive(Deserialize, Serialize, JsonSchema, Clone, Debug)] pub struct ScoredPoint { /// Point id pub id: PointIdType, /// Points vector distance to the query vector pub score: ScoreType, + /// Payload storage + pub payload: TheMap, } impl Eq for ScoredPoint {} @@ -61,6 +63,12 @@ impl PartialOrd for ScoredPoint { } } +impl PartialEq for ScoredPoint { + fn eq(&self, other: &Self) -> bool { + (self.id, &self.score) == (other.id, &other.score) + } +} + #[derive(Debug, Deserialize, Serialize, JsonSchema, Clone, Copy, PartialEq)] #[serde(rename_all = "snake_case")] pub enum SegmentType { @@ -395,6 +403,86 @@ pub enum Condition { Filter(Filter), } +#[derive(Debug, Deserialize, Serialize, JsonSchema, Clone)] +#[serde(rename_all = "snake_case")] +#[serde(untagged)] +pub enum WithPayloadInterface { + Bool(bool), + Fields(Vec), + Selector(PayloadSelector), +} +impl From for WithPayload { + fn from(x: bool) -> Self { + WithPayload { + enable: x, + payload_selector: None, + } + } +} + +impl From<&WithPayloadInterface> for WithPayload { + fn from(interface: &WithPayloadInterface) -> Self { + match interface { + WithPayloadInterface::Bool(x) => WithPayload { + enable: *x, + payload_selector: None, + }, + WithPayloadInterface::Fields(x) => WithPayload { + enable: true, + payload_selector: Some(PayloadSelector::new_include(x.clone())), + }, + WithPayloadInterface::Selector(x) => WithPayload { + enable: true, + payload_selector: Some(x.clone()), + }, + } + } +} + +#[derive(Debug, Deserialize, Serialize, JsonSchema, Clone)] +#[serde(deny_unknown_fields)] +#[serde(rename_all = "snake_case")] +pub struct PayloadSelector { + /// Include return payload key type + pub include: Vec, + /// Post-exclude return payload key type + pub exclude: Vec, +} + +impl PayloadSelector { + pub fn new_include(vecs_payload_key_type: Vec) -> Self { + PayloadSelector { + include: vecs_payload_key_type, + exclude: Vec::new(), + } + } + + pub fn new_include_and_exclude( + include: Vec, + exclude: Vec, + ) -> Self { + PayloadSelector { include, exclude } + } + + pub fn process( + &self, + x: TheMap, + ) -> TheMap { + x.into_iter() + .filter(|(key, _)| self.include.contains(key) && !self.exclude.contains(key)) + .collect() + } +} + +#[derive(Debug, Deserialize, Serialize, JsonSchema, Clone, Default)] +#[serde(deny_unknown_fields)] +#[serde(rename_all = "snake_case")] +pub struct WithPayload { + /// Enable return payloads or not + pub enable: bool, + /// Filter include and exclude payloads + pub payload_selector: Option, +} #[derive(Debug, Deserialize, Serialize, JsonSchema, Clone)] #[serde(deny_unknown_fields)] #[serde(rename_all = "snake_case")] diff --git a/lib/segment/tests/payload_index_test.rs b/lib/segment/tests/payload_index_test.rs index 619c91fe9b..a4e657c92d 100644 --- a/lib/segment/tests/payload_index_test.rs +++ b/lib/segment/tests/payload_index_test.rs @@ -8,7 +8,7 @@ mod tests { use segment::segment_constructor::build_segment; use segment::types::{ Condition, Distance, FieldCondition, Filter, Indexes, PayloadIndexType, PayloadKeyType, - PayloadType, Range, SegmentConfig, StorageType, TheMap, + PayloadType, Range, SegmentConfig, StorageType, TheMap, WithPayload, }; use tempdir::TempDir; @@ -144,10 +144,22 @@ mod tests { let query_filter = random_filter(&mut rnd); let plain_result = plain_segment - .search(&query_vector, Some(&query_filter), 5, None) + .search( + &query_vector, + &WithPayload::default(), + Some(&query_filter), + 5, + None, + ) .unwrap(); let struct_result = struct_segment - .search(&query_vector, Some(&query_filter), 5, None) + .search( + &query_vector, + &WithPayload::default(), + Some(&query_filter), + 5, + None, + ) .unwrap(); let estimation = struct_segment diff --git a/lib/segment/tests/segment_tests.rs b/lib/segment/tests/segment_tests.rs index 71fe6d1206..454836574b 100644 --- a/lib/segment/tests/segment_tests.rs +++ b/lib/segment/tests/segment_tests.rs @@ -4,7 +4,7 @@ mod fixtures; mod tests { use crate::fixtures::segment::build_segment_1; use segment::entry::entry_point::SegmentEntry; - use segment::types::{Condition, Filter}; + use segment::types::{Condition, Filter, WithPayload}; use std::collections::HashSet; use tempdir::TempDir; @@ -18,7 +18,9 @@ mod tests { let query_vector = vec![1.0, 1.0, 1.0, 1.0]; - let res = segment.search(&query_vector, None, 1, None).unwrap(); + let res = segment + .search(&query_vector, &WithPayload::default(), None, 1, None) + .unwrap(); let best_match = res.get(0).expect("Non-empty result"); assert_eq!(best_match.id, 3); @@ -31,7 +33,9 @@ mod tests { must_not: Some(vec![Condition::HasId(ids.into())]), }; - let res = segment.search(&query_vector, Some(&frt), 1, None).unwrap(); + let res = segment + .search(&query_vector, &WithPayload::default(), Some(&frt), 1, None) + .unwrap(); let best_match = res.get(0).expect("Non-empty result"); assert_ne!(best_match.id, 3); diff --git a/lib/storage/src/content_manager/toc.rs b/lib/storage/src/content_manager/toc.rs index 7601fed584..a792b8ec34 100644 --- a/lib/storage/src/content_manager/toc.rs +++ b/lib/storage/src/content_manager/toc.rs @@ -18,7 +18,7 @@ use collection::operations::types::{ RecommendRequest, Record, ScrollRequest, ScrollResult, SearchRequest, UpdateResult, }; use collection::operations::CollectionUpdateOperations; -use segment::types::{PointIdType, ScoredPoint}; +use segment::types::{PointIdType, ScoredPoint, WithPayload}; use crate::content_manager::collections_ops::{Checker, Collections}; use crate::content_manager::errors::StorageError; @@ -317,7 +317,7 @@ impl TableOfContent { &self, collection_name: &str, points: &[PointIdType], - with_payload: bool, + with_payload: &WithPayload, with_vector: bool, ) -> Result, StorageError> { let collection = self.get_collection(collection_name).await?; diff --git a/src/actix/api/retrieve_api.rs b/src/actix/api/retrieve_api.rs index 32beeeeb1c..26bde9be1c 100644 --- a/src/actix/api/retrieve_api.rs +++ b/src/actix/api/retrieve_api.rs @@ -6,7 +6,7 @@ use schemars::JsonSchema; use serde::{Deserialize, Serialize}; use collection::operations::types::{Record, ScrollRequest, ScrollResult}; -use segment::types::PointIdType; +use segment::types::{PointIdType, WithPayload, WithPayloadInterface}; use storage::content_manager::errors::StorageError; use storage::content_manager::toc::TableOfContent; @@ -15,6 +15,7 @@ use crate::actix::helpers::process_response; #[derive(Deserialize, Serialize, JsonSchema)] pub struct PointRequest { pub ids: Vec, + pub with_payload: Option, } async fn do_get_point( @@ -22,7 +23,7 @@ async fn do_get_point( collection_name: &str, point_id: PointIdType, ) -> Result, StorageError> { - toc.retrieve(collection_name, &[point_id], true, true) + toc.retrieve(collection_name, &[point_id], &WithPayload::from(true), true) .await .map(|points| points.into_iter().next()) } @@ -32,7 +33,11 @@ async fn do_get_points( collection_name: &str, request: PointRequest, ) -> Result, StorageError> { - toc.retrieve(collection_name, &request.ids, true, true) + let with_payload_interface = &request + .with_payload + .unwrap_or(WithPayloadInterface::Bool(true)); + let with_payload = WithPayload::from(with_payload_interface); + toc.retrieve(collection_name, &request.ids, &with_payload, true) .await }