Features/filter payload (#104)

* update more test

* update fmt

* reduce non usecode and update docker version

* update commend code

* update name filter

* renames and minor fixes

* fix linter

Co-authored-by: hai che <haiche@jobhop.com>
Co-authored-by: Andrey Vasnetsov <andrey@vasnetsov.com>
Co-authored-by: Andrey Vasnetsov <vasnetsov93@gmail.com>
This commit is contained in:
HaiCheViet
2021-10-12 16:07:36 +07:00
committed by GitHub
parent 370d16d878
commit f55e5aa7b7
15 changed files with 396 additions and 68 deletions

View File

@@ -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<PointIdType, Vec<VectorElementType>> = 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,
};

View File

@@ -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<SegmentHolder>,
points: &[PointIdType],
with_payload: bool,
with_payload: &WithPayload,
with_vector: bool,
) -> CollectionResult<Vec<Record>>;
}

View File

@@ -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);

View File

@@ -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<SegmentHolder>,
points: &[PointIdType],
with_payload: bool,
with_payload: &WithPayload,
with_vector: bool,
) -> CollectionResult<Vec<Record>> {
let mut point_version: HashMap<PointIdType, SeqNumberType> = 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<SearchRequest>,
) -> CollectionResult<Vec<ScoredPoint>> {
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);

View File

@@ -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);

View File

@@ -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<usize>,
/// Look only for points which satisfies this conditions. If not provided - all points.
pub filter: Option<Filter>,
/// Return point payload with the result. Default: true
pub with_payload: Option<bool>,
/// Return point payload with the result. Default: True
pub with_payload: Option<WithPayloadInterface>,
/// Return point vector with the result. Default: false
pub with_vector: Option<bool>,
}
@@ -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<SearchParams>,
/// Max number of result to return
pub top: usize,
/// Payload interface
pub with_payload: Option<WithPayloadInterface>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]

View File

@@ -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()
);
}

View File

@@ -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,

View File

@@ -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>,

View File

@@ -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<Vec<ScoredPoint>> = 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,

View File

@@ -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<PayloadKeyType, PayloadType>,
}
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<String>),
Selector(PayloadSelector),
}
impl From<bool> 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<PayloadKeyType>,
/// Post-exclude return payload key type
pub exclude: Vec<PayloadKeyType>,
}
impl PayloadSelector {
pub fn new_include(vecs_payload_key_type: Vec<PayloadKeyType>) -> Self {
PayloadSelector {
include: vecs_payload_key_type,
exclude: Vec::new(),
}
}
pub fn new_include_and_exclude(
include: Vec<PayloadKeyType>,
exclude: Vec<PayloadKeyType>,
) -> Self {
PayloadSelector { include, exclude }
}
pub fn process(
&self,
x: TheMap<PayloadKeyType, PayloadType>,
) -> TheMap<PayloadKeyType, PayloadType> {
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<PayloadSelector>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema, Clone)]
#[serde(deny_unknown_fields)]
#[serde(rename_all = "snake_case")]

View File

@@ -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

View File

@@ -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);

View File

@@ -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<Vec<Record>, StorageError> {
let collection = self.get_collection(collection_name).await?;

View File

@@ -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<PointIdType>,
pub with_payload: Option<WithPayloadInterface>,
}
async fn do_get_point(
@@ -22,7 +23,7 @@ async fn do_get_point(
collection_name: &str,
point_id: PointIdType,
) -> Result<Option<Record>, 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<Vec<Record>, 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
}