use std::collections::HashSet;
use crate::segment::types::{Filter, PointIdType, VectorNameBuf};
use serde::{Deserialize, Serialize};
use strum::{EnumDiscriminants, EnumIter};
use super::point_ops::{PointIdsList, VectorStructPersisted};
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, EnumDiscriminants, Hash)]
#[strum_discriminants(derive(EnumIter))]
#[serde(rename_all = "snake_case")]
pub enum VectorOperations {
UpdateVectors(UpdateVectorsOp),
DeleteVectors(PointIdsList, Vec<VectorNameBuf>),
DeleteVectorsByFilter(Filter, Vec<VectorNameBuf>),
}
impl VectorOperations {
pub fn point_ids(&self) -> Option<Vec<PointIdType>> {
match self {
Self::UpdateVectors(op) => Some(op.points.iter().map(|point| point.id).collect()),
Self::DeleteVectors(points, _) => Some(points.points.clone()),
Self::DeleteVectorsByFilter(_, _) => None,
}
}
pub fn retain_point_ids<F>(&mut self, filter: F)
where
F: Fn(&PointIdType) -> bool,
{
match self {
Self::UpdateVectors(op) => op.points.retain(|point| filter(&point.id)),
Self::DeleteVectors(points, _) => points.points.retain(filter),
Self::DeleteVectorsByFilter(_, _) => (),
}
}
pub fn retain_vector_names(&mut self, valid: &HashSet<VectorNameBuf>) {
match self {
Self::UpdateVectors(op) => {
for point in &mut op.points {
point.vector.retain_vector_names(valid);
}
}
Self::DeleteVectors(_, names) | Self::DeleteVectorsByFilter(_, names) => {
names.retain(|name| valid.contains(name));
}
}
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Hash)]
pub struct UpdateVectorsOp {
pub points: Vec<PointVectorsPersisted>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub update_filter: Option<Filter>,
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Hash)]
pub struct PointVectorsPersisted {
pub id: PointIdType,
pub vector: VectorStructPersisted,
}