use std::sync::atomic::AtomicBool;
use ahash::{AHashMap, AHashSet};
use crate::common::counter::hardware_counter::HardwareCounterCell;
use crate::common::types::DeferredBehavior;
use crate::segment::common::operation_error::OperationResult;
use crate::segment::data_types::segment_record::{SegmentRecord, SegmentRecordRaw};
use crate::segment::entry::ReadSegmentEntry;
use crate::segment::types::{PointIdType, SeqNumberType, WithPayload, WithVector};
use super::delete_points;
use super::upsert::{PointToUpsert, upsert_points_impl};
use crate::shard::operations::point_ops::{PointStructPersisted, PointStructRawPersisted};
use crate::shard::segment_holder::SegmentHolder;
pub fn sync_points(
segments: &SegmentHolder,
op_num: SeqNumberType,
from_id: Option<PointIdType>,
to_id: Option<PointIdType>,
points: &[PointStructPersisted],
hw_counter: &HardwareCounterCell,
) -> OperationResult<(usize, usize, usize)> {
sync_points_impl(segments, op_num, from_id, to_id, points, hw_counter)
}
pub fn sync_points_raw(
segments: &SegmentHolder,
op_num: SeqNumberType,
from_id: Option<PointIdType>,
to_id: Option<PointIdType>,
points: &[PointStructRawPersisted],
hw_counter: &HardwareCounterCell,
) -> OperationResult<(usize, usize, usize)> {
sync_points_impl(segments, op_num, from_id, to_id, points, hw_counter)
}
trait PointToSync: PointToUpsert {
type StoredRecord;
fn retrieve_stored(
segment: &dyn ReadSegmentEntry,
ids: &[PointIdType],
hw_counter: &HardwareCounterCell,
is_stopped: &AtomicBool,
) -> OperationResult<AHashMap<PointIdType, Self::StoredRecord>>;
fn is_equal_to(&self, stored: &Self::StoredRecord) -> bool;
}
fn sync_points_impl<P>(
segments: &SegmentHolder,
op_num: SeqNumberType,
from_id: Option<PointIdType>,
to_id: Option<PointIdType>,
points: &[P],
hw_counter: &HardwareCounterCell,
) -> OperationResult<(usize, usize, usize)>
where
P: PointToSync,
{
let id_to_point: AHashMap<PointIdType, &P> = points.iter().map(|p| (p.id(), p)).collect();
let sync_points: AHashSet<_> = points.iter().map(|p| p.id()).collect();
let stored_point_ids: AHashSet<_> = segments
.iter()
.flat_map(|(_, segment)| segment.get().read().read_range(from_id, to_id))
.collect();
let points_to_remove: Vec<_> = stored_point_ids.difference(&sync_points).copied().collect();
let deleted = delete_points(segments, op_num, points_to_remove.as_slice(), hw_counter)?;
let existing_point_ids: Vec<_> = stored_point_ids
.intersection(&sync_points)
.copied()
.collect();
let mut points_to_update: Vec<&P> = Vec::new();
let is_stopped = AtomicBool::new(false);
let _num_updated = segments.read_points(
existing_point_ids.as_slice(),
&is_stopped,
DeferredBehavior::WithDeferred,
|ids, segment| {
let stored_records = P::retrieve_stored(&**segment, ids, hw_counter, &is_stopped)?;
let mut updated = 0;
for (id, stored_record) in stored_records {
let point = id_to_point.get(&id).unwrap();
if !point.is_equal_to(&stored_record) {
points_to_update.push(*point);
updated += 1;
}
}
Ok(updated)
},
)?;
let num_updated = points_to_update.len();
let mut num_new = 0;
sync_points.difference(&stored_point_ids).for_each(|id| {
num_new += 1;
points_to_update.push(*id_to_point.get(id).unwrap());
});
let num_replaced = upsert_points_impl(segments, op_num, points_to_update, hw_counter)?;
debug_assert!(
num_replaced <= num_updated,
"number of replaced points cannot be greater than points to update ({num_replaced} <= {num_updated})",
);
Ok((deleted, num_new, num_updated))
}
impl PointToSync for PointStructPersisted {
type StoredRecord = SegmentRecord;
fn retrieve_stored(
segment: &dyn ReadSegmentEntry,
ids: &[PointIdType],
hw_counter: &HardwareCounterCell,
is_stopped: &AtomicBool,
) -> OperationResult<AHashMap<PointIdType, SegmentRecord>> {
segment.retrieve(
ids,
&WithPayload::from(true),
&WithVector::Bool(true),
hw_counter,
is_stopped,
DeferredBehavior::WithDeferred,
)
}
fn is_equal_to(&self, stored: &SegmentRecord) -> bool {
PointStructPersisted::is_equal_to(self, stored)
}
}
impl PointToSync for PointStructRawPersisted {
type StoredRecord = SegmentRecordRaw;
fn retrieve_stored(
segment: &dyn ReadSegmentEntry,
ids: &[PointIdType],
hw_counter: &HardwareCounterCell,
is_stopped: &AtomicBool,
) -> OperationResult<AHashMap<PointIdType, SegmentRecordRaw>> {
segment.retrieve_raw(
ids,
&WithPayload::from(true),
&WithVector::Bool(true),
hw_counter,
is_stopped,
DeferredBehavior::WithDeferred,
)
}
fn is_equal_to(&self, stored: &SegmentRecordRaw) -> bool {
PointStructRawPersisted::is_equal_to(self, stored)
}
}