mirror of
https://github.com/qdrant/qdrant.git
synced 2026-08-06 18:10:58 -05:00
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:
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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>>;
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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?;
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user