use std::collections::{HashMap, HashSet};
use std::fmt::{Debug, Formatter};
use std::hash::{Hash, Hasher};
use std::mem;
#[cfg(feature = "api")]
use api::grpc::RawPayload;
use crate::common::validation::validate_multi_vector;
use itertools::Itertools as _;
use ordered_float::OrderedFloat;
use schemars::JsonSchema;
use crate::segment::common::operation_error::OperationError;
use crate::segment::common::utils::unordered_hash_unique;
use crate::segment::data_types::named_vectors::NamedVectors;
use crate::segment::data_types::segment_record::{SegmentRecord, SegmentRecordRaw};
use crate::segment::data_types::vectors::{
BatchVectorStructInternal, DEFAULT_VECTOR_NAME, DenseVector, MultiDenseVector,
MultiDenseVectorInternal, VectorInternal, VectorRef, VectorStructInternal,
};
use crate::segment::types::{Filter, Payload, PointIdType, VectorNameBuf};
use serde::{Deserialize, Serialize};
use smallvec::SmallVec;
use crate::sparse::common::types::{DimId, DimWeight};
use strum::{EnumDiscriminants, EnumIter};
use validator::{Validate, ValidationErrors};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Deserialize, Serialize, Hash)]
#[serde(rename_all = "snake_case")]
pub enum UpdateMode {
#[default]
Upsert,
InsertOnly,
UpdateOnly,
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, JsonSchema, Validate, Hash)]
#[serde(rename_all = "snake_case")]
pub struct PointIdsList {
pub points: Vec<PointIdType>,
#[cfg(feature = "api")]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub shard_key: Option<api::rest::ShardKeySelector>,
}
impl From<Vec<PointIdType>> for PointIdsList {
fn from(points: Vec<PointIdType>) -> Self {
Self {
points,
#[cfg(feature = "api")]
shard_key: None,
}
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, EnumDiscriminants, Hash)]
#[strum_discriminants(derive(EnumIter))]
#[serde(rename_all = "snake_case")]
pub enum PointOperations {
UpsertPoints(PointInsertOperationsInternal),
UpsertPointsConditional(ConditionalInsertOperationInternal),
DeletePoints { ids: Vec<PointIdType> },
DeletePointsByFilter(Filter),
SyncPoints(PointSyncOperation),
UpsertPointsRaw(Vec<PointStructRawPersisted>),
SyncPointsRaw(PointSyncRawOperation),
}
impl PointOperations {
pub fn point_ids(&self) -> Option<Vec<PointIdType>> {
match self {
Self::UpsertPoints(op) => Some(op.point_ids()),
Self::UpsertPointsConditional(op) => Some(op.points_op.point_ids()),
Self::DeletePoints { ids } => Some(ids.clone()),
Self::DeletePointsByFilter(_) => None,
Self::SyncPoints(op) => Some(op.points.iter().map(|point| point.id).collect()),
Self::UpsertPointsRaw(points) => Some(points.iter().map(|point| point.id).collect()),
Self::SyncPointsRaw(op) => Some(op.points.iter().map(|point| point.id).collect()),
}
}
pub fn retain_point_ids<F>(&mut self, filter: F)
where
F: Fn(&PointIdType) -> bool,
{
match self {
Self::UpsertPoints(op) => op.retain_point_ids(filter),
Self::UpsertPointsConditional(op) => {
op.points_op.retain_point_ids(filter);
}
Self::DeletePoints { ids } => ids.retain(filter),
Self::DeletePointsByFilter(_) => (),
Self::SyncPoints(op) => op.points.retain(|point| filter(&point.id)),
Self::UpsertPointsRaw(points) => points.retain(|point| filter(&point.id)),
Self::SyncPointsRaw(op) => op.points.retain(|point| filter(&point.id)),
}
}
pub fn retain_vector_names(&mut self, valid: &HashSet<VectorNameBuf>) {
match self {
Self::UpsertPoints(op) => op.retain_vector_names(valid),
Self::UpsertPointsConditional(op) => op.points_op.retain_vector_names(valid),
Self::SyncPoints(op) => {
for point in &mut op.points {
point.vector.retain_vector_names(valid);
}
}
Self::UpsertPointsRaw(points) => {
for point in points {
point.vectors.retain(|(name, _)| valid.contains(name));
}
}
Self::SyncPointsRaw(op) => {
for point in &mut op.points {
point.vectors.retain(|(name, _)| valid.contains(name));
}
}
Self::DeletePoints { .. } | Self::DeletePointsByFilter(_) => (),
}
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, EnumDiscriminants, Hash)]
#[strum_discriminants(derive(EnumIter))]
#[serde(rename_all = "snake_case")]
pub enum PointInsertOperationsInternal {
#[serde(rename = "batch")]
PointsBatch(BatchPersisted),
#[serde(rename = "points")]
PointsList(Vec<PointStructPersisted>),
}
impl PointInsertOperationsInternal {
pub fn point_ids(&self) -> Vec<PointIdType> {
match self {
Self::PointsBatch(batch) => batch.ids.clone(),
Self::PointsList(points) => points.iter().map(|point| point.id).collect(),
}
}
pub fn into_point_vec(self) -> Vec<PointStructPersisted> {
match self {
PointInsertOperationsInternal::PointsBatch(batch) => {
let batch_vectors = BatchVectorStructInternal::from(batch.vectors);
let all_vectors = batch_vectors.into_all_vectors(batch.ids.len());
let vectors_iter = batch.ids.into_iter().zip(all_vectors);
match batch.payloads {
None => vectors_iter
.map(|(id, vectors)| PointStructPersisted {
id,
vector: VectorStructInternal::from(vectors).into(),
payload: None,
})
.collect(),
Some(payloads) => vectors_iter
.zip(payloads)
.map(|((id, vectors), payload)| PointStructPersisted {
id,
vector: VectorStructInternal::from(vectors).into(),
payload,
})
.collect(),
}
}
PointInsertOperationsInternal::PointsList(points) => points,
}
}
pub fn retain_vector_names(&mut self, valid: &HashSet<VectorNameBuf>) {
match self {
Self::PointsBatch(batch) => batch.vectors.retain_vector_names(valid),
Self::PointsList(points) => {
for point in points {
point.vector.retain_vector_names(valid);
}
}
}
}
pub fn retain_point_ids<F>(&mut self, filter: F)
where
F: Fn(&PointIdType) -> bool,
{
match self {
Self::PointsBatch(batch) => {
let mut retain_indices = HashSet::new();
retain_with_index(&mut batch.ids, |index, id| {
if filter(id) {
retain_indices.insert(index);
true
} else {
false
}
});
match &mut batch.vectors {
BatchVectorStructPersisted::Single(vectors) => {
retain_with_index(vectors, |index, _| retain_indices.contains(&index));
}
BatchVectorStructPersisted::MultiDense(vectors) => {
retain_with_index(vectors, |index, _| retain_indices.contains(&index));
}
BatchVectorStructPersisted::Named(vectors) => {
for vectors in vectors.values_mut() {
retain_with_index(vectors, |index, _| retain_indices.contains(&index));
}
}
}
if let Some(payload) = &mut batch.payloads {
retain_with_index(payload, |index, _| retain_indices.contains(&index));
}
}
Self::PointsList(points) => points.retain(|point| filter(&point.id)),
}
}
}
impl From<BatchPersisted> for PointInsertOperationsInternal {
fn from(batch: BatchPersisted) -> Self {
PointInsertOperationsInternal::PointsBatch(batch)
}
}
impl From<Vec<PointStructPersisted>> for PointInsertOperationsInternal {
fn from(points: Vec<PointStructPersisted>) -> Self {
PointInsertOperationsInternal::PointsList(points)
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Hash)]
pub struct ConditionalInsertOperationInternal {
pub points_op: PointInsertOperationsInternal,
pub condition: Filter,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub update_mode: Option<UpdateMode>,
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Hash)]
pub struct PointSyncOperation {
pub from_id: Option<PointIdType>,
pub to_id: Option<PointIdType>,
pub points: Vec<PointStructPersisted>,
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Hash)]
pub struct PointSyncRawOperation {
pub from_id: Option<PointIdType>,
pub to_id: Option<PointIdType>,
pub points: Vec<PointStructRawPersisted>,
}
pub type RawVectorsPersisted = SmallVec<[(VectorNameBuf, Vec<u8>); 1]>;
#[derive(Clone, PartialEq, Deserialize, Serialize, Hash)]
#[serde(rename_all = "snake_case")]
pub struct PointStructRawPersisted {
pub id: PointIdType,
#[serde(with = "raw_vectors_serde")]
pub vectors: RawVectorsPersisted,
pub payload: Option<Payload>,
}
mod raw_vectors_serde {
use std::fmt;
use crate::segment::types::VectorNameBuf;
use serde::de::{self, Deserializer, SeqAccess, Visitor};
use serde::ser::{SerializeSeq, Serializer};
use super::RawVectorsPersisted;
const MAX_RAW_VECTOR_PREALLOC: usize = 128 * 1024 * 1024;
struct BytesRef<'a>(&'a [u8]);
impl serde::Serialize for BytesRef<'_> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_bytes(self.0)
}
}
pub fn serialize<S: Serializer>(
vectors: &[(VectorNameBuf, Vec<u8>)],
serializer: S,
) -> Result<S::Ok, S::Error> {
let mut seq = serializer.serialize_seq(Some(vectors.len()))?;
for (name, bytes) in vectors {
seq.serialize_element(&(name, BytesRef(bytes)))?;
}
seq.end()
}
struct ByteVec(Vec<u8>);
impl<'de> serde::Deserialize<'de> for ByteVec {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct ByteVecVisitor;
impl<'de> Visitor<'de> for ByteVecVisitor {
type Value = Vec<u8>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a byte string")
}
fn visit_bytes<E: de::Error>(self, value: &[u8]) -> Result<Self::Value, E> {
Ok(value.to_vec())
}
fn visit_byte_buf<E: de::Error>(self, value: Vec<u8>) -> Result<Self::Value, E> {
Ok(value)
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Self::Value, A::Error> {
let capacity = seq.size_hint().unwrap_or(0).min(MAX_RAW_VECTOR_PREALLOC);
let mut bytes = Vec::with_capacity(capacity);
while let Some(byte) = seq.next_element()? {
bytes.push(byte);
}
Ok(bytes)
}
}
deserializer
.deserialize_byte_buf(ByteVecVisitor)
.map(ByteVec)
}
}
pub fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<RawVectorsPersisted, D::Error> {
let raw: Vec<(VectorNameBuf, ByteVec)> = serde::Deserialize::deserialize(deserializer)?;
Ok(raw
.into_iter()
.map(|(name, bytes)| (name, bytes.0))
.collect())
}
}
impl From<SegmentRecordRaw> for PointStructRawPersisted {
fn from(record: SegmentRecordRaw) -> Self {
let SegmentRecordRaw {
id,
vectors,
payload,
} = record;
Self {
id,
vectors: vectors.unwrap_or_default(),
payload,
}
}
}
impl PointStructRawPersisted {
pub fn is_equal_to(&self, segment_record: &SegmentRecordRaw) -> bool {
let SegmentRecordRaw {
id,
vectors,
payload,
} = segment_record;
if &self.id != id {
return false;
}
let segment_vectors = vectors.as_deref().unwrap_or(&[]);
if self.vectors.len() != segment_vectors.len() {
return false;
}
for (name, bytes) in segment_vectors {
let own_bytes = self
.vectors
.iter()
.find(|(own_name, _)| own_name == name)
.map(|(_, bytes)| bytes);
if own_bytes != Some(bytes) {
return false;
}
}
let self_payload = self.payload.as_ref().filter(|p| !p.is_empty());
let segment_payload = payload.as_ref().filter(|p| !p.is_empty());
self_payload == segment_payload
}
}
#[cfg(feature = "api")]
impl From<PointStructRawPersisted> for api::grpc::qdrant::PointStructRaw {
fn from(value: PointStructRawPersisted) -> Self {
let PointStructRawPersisted {
id,
vectors,
payload,
} = value;
Self {
id: Some(id.into()),
vectors: vectors.into_iter().collect(),
payload: payload
.map(api::conversions::json::payload_to_proto)
.unwrap_or_default(),
raw_payload: None,
}
}
}
#[cfg(feature = "api")]
impl TryFrom<api::grpc::qdrant::PointStructRaw> for PointStructRawPersisted {
type Error = tonic::Status;
fn try_from(value: api::grpc::qdrant::PointStructRaw) -> Result<Self, Self::Error> {
let api::grpc::qdrant::PointStructRaw {
id,
vectors,
payload,
raw_payload,
} = value;
let id = id
.ok_or_else(|| tonic::Status::invalid_argument("Empty id is not allowed"))?
.try_into()?;
let payload = match raw_payload {
Some(raw_payload) => decode_payload(raw_payload)?,
None => api::conversions::json::proto_to_payloads(payload)?,
};
let payload = (!payload.is_empty()).then_some(payload);
Ok(Self {
id,
vectors: vectors.into_iter().collect(),
payload,
})
}
}
#[cfg(feature = "api")]
fn decode_payload(raw_payload: RawPayload) -> Result<Payload, tonic::Status> {
match raw_payload.encoding() {
api::grpc::RawPayloadEncoding::JsonBytes => {
serde_json::from_slice(&raw_payload.payload_bytes).map_err(|err| {
tonic::Status::invalid_argument(format!("Malformed raw payload blob: {err}"))
})
}
}
}
impl Debug for PointStructRawPersisted {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
let vectors = self
.vectors
.iter()
.map(|(name, bytes)| format!("{name}: {} bytes", bytes.len()))
.join(", ");
write!(
f,
"PointStructRawPersisted {{ id: {}, vectors: [{vectors}], payload: {:?} }}",
self.id, self.payload,
)
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Hash)]
#[serde(rename_all = "snake_case")]
pub struct BatchPersisted {
pub ids: Vec<PointIdType>,
pub vectors: BatchVectorStructPersisted,
pub payloads: Option<Vec<Option<Payload>>>,
}
#[cfg(feature = "api")]
impl TryFrom<BatchPersisted> for Vec<api::grpc::qdrant::PointStruct> {
type Error = tonic::Status;
fn try_from(batch: BatchPersisted) -> Result<Self, Self::Error> {
let BatchPersisted {
ids,
vectors,
payloads,
} = batch;
let mut points = Vec::with_capacity(ids.len());
let batch_vectors = BatchVectorStructInternal::from(vectors);
let all_vectors = batch_vectors.into_all_vectors(ids.len());
for (i, p_id) in ids.into_iter().enumerate() {
let id = Some(p_id.into());
let vector = all_vectors.get(i).cloned();
let payload = payloads.as_ref().and_then(|payloads| {
payloads.get(i).map(|payload| match payload {
None => HashMap::new(),
Some(payload) => api::conversions::json::payload_to_proto(payload.clone()),
})
});
let vectors: Option<VectorStructInternal> = vector.map(NamedVectors::into);
let point = api::grpc::qdrant::PointStruct {
id,
vectors: vectors.map(api::grpc::qdrant::Vectors::from),
payload: payload.unwrap_or_default(),
};
points.push(point);
}
Ok(points)
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize)]
#[serde(untagged, rename_all = "snake_case")]
pub enum BatchVectorStructPersisted {
Single(Vec<DenseVector>),
MultiDense(Vec<MultiDenseVector>),
Named(HashMap<VectorNameBuf, Vec<VectorPersisted>>),
}
impl BatchVectorStructPersisted {
pub fn retain_vector_names(&mut self, valid: &HashSet<VectorNameBuf>) {
if let BatchVectorStructPersisted::Named(named) = self {
named.retain(|name, _| valid.contains(name));
}
}
}
impl Hash for BatchVectorStructPersisted {
fn hash<H: Hasher>(&self, state: &mut H) {
mem::discriminant(self).hash(state);
match self {
BatchVectorStructPersisted::Single(dense) => {
for vector in dense {
for v in vector {
OrderedFloat(*v).hash(state);
}
}
}
BatchVectorStructPersisted::MultiDense(multidense) => {
for vector in multidense {
for v in vector {
for element in v {
OrderedFloat(*element).hash(state);
}
}
}
}
BatchVectorStructPersisted::Named(named) => unordered_hash_unique(state, named.iter()),
}
}
}
impl From<BatchVectorStructPersisted> for BatchVectorStructInternal {
fn from(value: BatchVectorStructPersisted) -> Self {
match value {
BatchVectorStructPersisted::Single(vector) => BatchVectorStructInternal::Single(vector),
BatchVectorStructPersisted::MultiDense(vectors) => {
BatchVectorStructInternal::MultiDense(
vectors
.into_iter()
.map(MultiDenseVectorInternal::new_unchecked)
.collect(),
)
}
BatchVectorStructPersisted::Named(vectors) => BatchVectorStructInternal::Named(
vectors
.into_iter()
.map(|(k, v)| (k, v.into_iter().map(VectorInternal::from).collect()))
.collect(),
),
}
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Validate, Hash)]
#[serde(rename_all = "snake_case")]
pub struct PointStructPersisted {
pub id: PointIdType,
pub vector: VectorStructPersisted,
pub payload: Option<Payload>,
}
impl PointStructPersisted {
pub fn get_vectors(&self) -> NamedVectors<'_> {
let mut named_vectors = NamedVectors::default();
match &self.vector {
VectorStructPersisted::Single(vector) => named_vectors.insert(
DEFAULT_VECTOR_NAME.to_owned(),
VectorInternal::from(vector.clone()),
),
VectorStructPersisted::MultiDense(vector) => named_vectors.insert(
DEFAULT_VECTOR_NAME.to_owned(),
VectorInternal::from(MultiDenseVectorInternal::new_unchecked(vector.clone())),
),
VectorStructPersisted::Named(vectors) => {
for (name, vector) in vectors {
named_vectors.insert(name.clone(), VectorInternal::from(vector.clone()));
}
}
}
named_vectors
}
pub fn is_equal_to(&self, segment_record: &SegmentRecord) -> bool {
let SegmentRecord {
id,
vectors,
payload,
} = segment_record;
if &self.id != id {
return false;
}
let self_vectors = self.get_vectors();
if let Some(segment_vectors) = vectors {
if self_vectors.len() != segment_vectors.len() {
return false;
}
for (name, vec) in segment_vectors {
if self_vectors.get(name) != Some(VectorRef::from(vec)) {
return false;
}
}
} else if !self_vectors.is_empty() {
return false;
}
let self_payload = self.payload.as_ref().filter(|p| !p.is_empty());
let segment_payload = payload.as_ref().filter(|p| !p.is_empty());
self_payload == segment_payload
}
}
#[cfg(feature = "api")]
impl TryFrom<api::rest::schema::Record> for PointStructPersisted {
type Error = String;
fn try_from(record: api::rest::schema::Record) -> Result<Self, Self::Error> {
let api::rest::schema::Record {
id,
payload,
vector,
shard_key: _,
order_value: _,
} = record;
if vector.is_none() {
return Err("Vector is empty".to_string());
}
Ok(Self {
id,
payload,
vector: VectorStructPersisted::from(vector.unwrap()),
})
}
}
#[cfg(feature = "api")]
impl TryFrom<PointStructPersisted> for api::grpc::qdrant::PointStruct {
type Error = tonic::Status;
fn try_from(value: PointStructPersisted) -> Result<Self, Self::Error> {
let PointStructPersisted {
id,
vector,
payload,
} = value;
let vectors_internal = VectorStructInternal::try_from(vector).map_err(|e| {
tonic::Status::invalid_argument(format!("Failed to convert vectors: {e}"))
})?;
let vectors = api::grpc::qdrant::Vectors::from(vectors_internal);
let converted_payload = match payload {
None => HashMap::new(),
Some(payload) => api::conversions::json::payload_to_proto(payload),
};
Ok(Self {
id: Some(id.into()),
vectors: Some(vectors),
payload: converted_payload,
})
}
}
#[derive(Clone, PartialEq, Deserialize, Serialize)]
#[serde(untagged, rename_all = "snake_case")]
pub enum VectorStructPersisted {
Single(DenseVector),
MultiDense(MultiDenseVector),
Named(HashMap<VectorNameBuf, VectorPersisted>),
}
impl VectorStructPersisted {
pub fn retain_vector_names(&mut self, valid: &HashSet<VectorNameBuf>) {
if let VectorStructPersisted::Named(named) = self {
named.retain(|name, _| valid.contains(name));
}
}
}
impl std::hash::Hash for VectorStructPersisted {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
mem::discriminant(self).hash(state);
match self {
VectorStructPersisted::Single(vec) => {
for v in vec {
OrderedFloat(*v).hash(state);
}
}
VectorStructPersisted::MultiDense(multi_vec) => {
for vec in multi_vec {
for v in vec {
OrderedFloat(*v).hash(state);
}
}
}
VectorStructPersisted::Named(map) => {
unordered_hash_unique(state, map.iter());
}
}
}
}
impl Debug for VectorStructPersisted {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
VectorStructPersisted::Single(vector) => {
let first_elements = vector.iter().take(4).join(", ");
write!(f, "Single([{}, ... x {}])", first_elements, vector.len())
}
VectorStructPersisted::MultiDense(vector) => {
let first_vectors = vector
.iter()
.take(4)
.map(|v| {
let first_elements = v.iter().take(4).join(", ");
format!("[{}, ... x {}]", first_elements, v.len())
})
.join(", ");
write!(f, "MultiDense([{}, ... x {})", first_vectors, vector.len())
}
VectorStructPersisted::Named(vectors) => write!(f, "Named(( ")
.and_then(|_| {
for (name, vector) in vectors {
write!(f, "{name}: {vector:?}, ")?;
}
Ok(())
})
.and_then(|_| write!(f, "))")),
}
}
}
impl VectorStructPersisted {
pub fn is_empty(&self) -> bool {
match self {
VectorStructPersisted::Single(vector) => vector.is_empty(),
VectorStructPersisted::MultiDense(vector) => vector.is_empty(),
VectorStructPersisted::Named(vectors) => vectors.values().all(|v| match v {
VectorPersisted::Dense(vector) => vector.is_empty(),
VectorPersisted::Sparse(vector) => vector.indices.is_empty(),
VectorPersisted::MultiDense(vector) => vector.is_empty(),
}),
}
}
}
impl Validate for VectorStructPersisted {
fn validate(&self) -> Result<(), ValidationErrors> {
match self {
VectorStructPersisted::Single(_) => Ok(()),
VectorStructPersisted::MultiDense(v) => validate_multi_vector(v),
VectorStructPersisted::Named(v) => crate::common::validation::validate_iter(v.values()),
}
}
}
impl From<DenseVector> for VectorStructPersisted {
fn from(value: DenseVector) -> Self {
VectorStructPersisted::Single(value)
}
}
impl From<VectorStructInternal> for VectorStructPersisted {
fn from(value: VectorStructInternal) -> Self {
match value {
VectorStructInternal::Single(vector) => VectorStructPersisted::Single(vector),
VectorStructInternal::MultiDense(vector) => {
VectorStructPersisted::MultiDense(vector.into_multi_vectors())
}
VectorStructInternal::Named(vectors) => VectorStructPersisted::Named(
vectors
.into_iter()
.map(|(k, v)| (k, VectorPersisted::from(v)))
.collect(),
),
}
}
}
#[cfg(feature = "api")]
impl From<api::rest::VectorStructOutput> for VectorStructPersisted {
fn from(value: api::rest::VectorStructOutput) -> Self {
match value {
api::rest::VectorStructOutput::Single(vector) => VectorStructPersisted::Single(vector),
api::rest::VectorStructOutput::MultiDense(vector) => {
VectorStructPersisted::MultiDense(vector)
}
api::rest::VectorStructOutput::Named(vectors) => VectorStructPersisted::Named(
vectors
.into_iter()
.map(|(k, v)| (k, VectorPersisted::from(v)))
.collect(),
),
}
}
}
impl TryFrom<VectorStructPersisted> for VectorStructInternal {
type Error = OperationError;
fn try_from(value: VectorStructPersisted) -> Result<Self, Self::Error> {
let vector_struct = match value {
VectorStructPersisted::Single(vector) => VectorStructInternal::Single(vector),
VectorStructPersisted::MultiDense(vector) => {
VectorStructInternal::MultiDense(MultiDenseVectorInternal::try_from(vector)?)
}
VectorStructPersisted::Named(vectors) => VectorStructInternal::Named(
vectors
.into_iter()
.map(|(k, v)| (k, VectorInternal::from(v)))
.collect(),
),
};
Ok(vector_struct)
}
}
impl From<VectorStructPersisted> for NamedVectors<'_> {
fn from(value: VectorStructPersisted) -> Self {
match value {
VectorStructPersisted::Single(vector) => {
NamedVectors::from_pairs([(DEFAULT_VECTOR_NAME.to_owned(), vector)])
}
VectorStructPersisted::MultiDense(vector) => {
let mut named_vector = NamedVectors::default();
let multivec = MultiDenseVectorInternal::new_unchecked(vector);
named_vector.insert(
DEFAULT_VECTOR_NAME.to_owned(),
crate::segment::data_types::vectors::VectorInternal::from(multivec),
);
named_vector
}
VectorStructPersisted::Named(vectors) => {
let mut named_vector = NamedVectors::default();
for (name, vector) in vectors {
named_vector.insert(
name,
crate::segment::data_types::vectors::VectorInternal::from(vector),
);
}
named_vector
}
}
}
}
#[derive(Clone, PartialEq, Deserialize, Serialize)]
#[serde(untagged, rename_all = "snake_case")]
pub enum VectorPersisted {
Dense(DenseVector),
Sparse(crate::sparse::common::sparse_vector::SparseVector),
MultiDense(MultiDenseVector),
}
impl Hash for VectorPersisted {
fn hash<H: Hasher>(&self, state: &mut H) {
mem::discriminant(self).hash(state);
match self {
VectorPersisted::Dense(vec) => {
for v in vec {
OrderedFloat(*v).hash(state);
}
}
VectorPersisted::Sparse(sparse) => {
sparse.hash(state);
}
VectorPersisted::MultiDense(multi_vec) => {
for vec in multi_vec {
for v in vec {
OrderedFloat(*v).hash(state);
}
}
}
}
}
}
impl VectorPersisted {
pub fn new_sparse(indices: Vec<DimId>, values: Vec<DimWeight>) -> Self {
Self::Sparse(crate::sparse::common::sparse_vector::SparseVector { indices, values })
}
pub fn empty_sparse() -> Self {
Self::new_sparse(vec![], vec![])
}
}
impl Debug for VectorPersisted {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
match self {
VectorPersisted::Dense(vector) => {
let first_elements = vector.iter().take(4).join(", ");
write!(f, "Dense([{}, ... x {}])", first_elements, vector.len())
}
VectorPersisted::Sparse(vector) => {
let first_elements = vector
.indices
.iter()
.zip(vector.values.iter())
.take(4)
.map(|(k, v)| format!("{k}->{v}"))
.join(", ");
write!(
f,
"Sparse([{}, ... x {})",
first_elements,
vector.indices.len()
)
}
VectorPersisted::MultiDense(vector) => {
let first_vectors = vector
.iter()
.take(4)
.map(|v| {
let first_elements = v.iter().take(4).join(", ");
format!("[{}, ... x {}]", first_elements, v.len())
})
.join(", ");
write!(f, "MultiDense([{}, ... x {})", first_vectors, vector.len())
}
}
}
}
impl Validate for VectorPersisted {
fn validate(&self) -> Result<(), ValidationErrors> {
match self {
VectorPersisted::Dense(_) => Ok(()),
VectorPersisted::Sparse(v) => v.validate(),
VectorPersisted::MultiDense(m) => validate_multi_vector(m),
}
}
}
impl From<VectorInternal> for VectorPersisted {
fn from(value: VectorInternal) -> Self {
match value {
VectorInternal::Dense(vector) => VectorPersisted::Dense(vector),
VectorInternal::Sparse(vector) => VectorPersisted::Sparse(vector),
VectorInternal::MultiDense(vector) => {
VectorPersisted::MultiDense(vector.into_multi_vectors())
}
}
}
}
#[cfg(feature = "api")]
impl From<api::rest::VectorOutput> for VectorPersisted {
fn from(value: api::rest::VectorOutput) -> Self {
match value {
api::rest::VectorOutput::Dense(vector) => VectorPersisted::Dense(vector),
api::rest::VectorOutput::Sparse(vector) => VectorPersisted::Sparse(vector),
api::rest::VectorOutput::MultiDense(vector) => VectorPersisted::MultiDense(vector),
}
}
}
impl From<VectorPersisted> for VectorInternal {
fn from(value: VectorPersisted) -> Self {
match value {
VectorPersisted::Dense(vector) => VectorInternal::Dense(vector),
VectorPersisted::Sparse(vector) => VectorInternal::Sparse(vector),
VectorPersisted::MultiDense(vector) => {
VectorInternal::MultiDense(MultiDenseVectorInternal::new_unchecked(vector))
}
}
}
}
fn retain_with_index<T, F>(vec: &mut Vec<T>, mut filter: F)
where
F: FnMut(usize, &T) -> bool,
{
let mut index = 0;
vec.retain(|item| {
let retain = filter(index, item);
index += 1;
retain
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn raw_persisted_vectors_use_compact_byte_string() {
let blob: Vec<u8> = (0..4096u32).map(|i| i as u8).collect();
let point = PointStructRawPersisted {
id: 1.into(),
vectors: vec![("dense".to_string(), blob.clone())].into(),
payload: None,
};
let encoded = serde_cbor::to_vec(&point).unwrap();
assert!(
encoded.len() < blob.len() + 128,
"expected compact byte-string encoding, got {} bytes for a {}-byte blob",
encoded.len(),
blob.len(),
);
let decoded: PointStructRawPersisted = serde_cbor::from_slice(&encoded).unwrap();
assert!(decoded == point, "round-trip mismatch");
}
fn dense(v: f32) -> VectorPersisted {
VectorPersisted::Dense(vec![v])
}
#[test]
fn retain_vector_names_strips_list_batch_and_delete() {
let valid: HashSet<VectorNameBuf> = ["a".to_string()].into_iter().collect();
let mut list =
PointOperations::UpsertPoints(PointInsertOperationsInternal::PointsList(vec![
PointStructPersisted {
id: 1.into(),
vector: VectorStructPersisted::Named(
[("a".to_string(), dense(0.1)), ("b".to_string(), dense(0.2))]
.into_iter()
.collect(),
),
payload: None,
},
]));
list.retain_vector_names(&valid);
let PointOperations::UpsertPoints(PointInsertOperationsInternal::PointsList(points)) =
&list
else {
unreachable!()
};
let VectorStructPersisted::Named(named) = &points[0].vector else {
unreachable!()
};
assert_eq!(named.keys().cloned().collect::<Vec<_>>(), vec!["a"]);
let mut batch = PointOperations::UpsertPoints(PointInsertOperationsInternal::PointsBatch(
BatchPersisted {
ids: vec![1.into(), 2.into()],
vectors: BatchVectorStructPersisted::Named(
[
("a".to_string(), vec![dense(0.1), dense(0.2)]),
("b".to_string(), vec![dense(0.3), dense(0.4)]),
]
.into_iter()
.collect(),
),
payloads: None,
},
));
batch.retain_vector_names(&valid);
let PointOperations::UpsertPoints(PointInsertOperationsInternal::PointsBatch(batch)) =
&batch
else {
unreachable!()
};
let BatchVectorStructPersisted::Named(named) = &batch.vectors else {
unreachable!()
};
assert_eq!(named.keys().cloned().collect::<Vec<_>>(), vec!["a"]);
assert_eq!(batch.ids.len(), 2);
}
}