use std::cmp::Ordering;
use std::collections::BinaryHeap;
use std::sync::Arc;
use futures::StreamExt;
use surrealdb_types::ToSql;
use crate::catalog::Distance;
use crate::err::EngineError;
use crate::exec::physical_expr::{EvalContext, PhysicalExpr};
use crate::exec::{
AccessMode, CardinalityHint, ContextLevel, ExecOperator, ExecutionContext, FlowResult,
OperatorMetrics, ValueBatch, ValueBatchStream, buffer_stream, monitor_stream,
};
use crate::expr::Idiom;
use crate::val::{Number, Value};
#[derive(Clone, Debug)]
pub(crate) enum KnnVectorSource {
Literal(Vec<Number>),
Deferred(Arc<dyn PhysicalExpr>),
}
struct DistanceEntry<T> {
distance: Number,
item: T,
seq: u64,
}
impl<T> PartialEq for DistanceEntry<T> {
fn eq(&self, other: &Self) -> bool {
self.cmp(other) == Ordering::Equal
}
}
impl<T> Eq for DistanceEntry<T> {}
impl<T> PartialOrd for DistanceEntry<T> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl<T> Ord for DistanceEntry<T> {
fn cmp(&self, other: &Self) -> Ordering {
other
.distance
.partial_cmp(&self.distance)
.unwrap_or(Ordering::Equal)
.then_with(|| other.seq.cmp(&self.seq))
}
}
pub(crate) struct KnnTopKHeap<T> {
k: usize,
heap: BinaryHeap<std::cmp::Reverse<DistanceEntry<T>>>,
seq: u64,
}
impl<T> KnnTopKHeap<T> {
pub(crate) fn new(k: usize) -> Self {
Self {
k,
heap: BinaryHeap::with_capacity(k + 1),
seq: 0,
}
}
pub(crate) fn offer(&mut self, distance: Number, item: T) {
let entry = DistanceEntry {
distance,
item,
seq: self.seq,
};
self.seq += 1;
if self.heap.len() >= self.k {
if let Some(worst) = self.heap.peek()
&& entry.distance < worst.0.distance
{
self.heap.push(std::cmp::Reverse(entry));
self.heap.pop();
}
} else {
self.heap.push(std::cmp::Reverse(entry));
}
}
pub(crate) fn into_sorted_nearest_first(mut self) -> Vec<(Number, T)> {
let mut entries: Vec<DistanceEntry<T>> = Vec::with_capacity(self.heap.len());
while let Some(std::cmp::Reverse(entry)) = self.heap.pop() {
entries.push(entry);
}
entries.reverse();
entries.into_iter().map(|e| (e.distance, e.item)).collect()
}
}
#[derive(Debug)]
pub struct KnnTopK {
pub(crate) input: Arc<dyn ExecOperator>,
pub(crate) field: Idiom,
pub(crate) query_vector: KnnVectorSource,
pub(crate) k: usize,
pub(crate) distance: Distance,
pub(crate) metrics: Arc<OperatorMetrics>,
pub(crate) knn_context: Option<Arc<crate::exec::function::KnnContext>>,
}
impl KnnTopK {
pub(crate) fn new(
input: Arc<dyn ExecOperator>,
field: Idiom,
query_vector: KnnVectorSource,
k: usize,
distance: Distance,
) -> Self {
Self {
input,
field,
query_vector,
k,
distance,
metrics: Arc::new(OperatorMetrics::new()),
knn_context: None,
}
}
pub(crate) fn with_knn_context(
mut self,
knn_context: Option<Arc<crate::exec::function::KnnContext>>,
) -> Self {
self.knn_context = knn_context;
self
}
}
impl ExecOperator for KnnTopK {
fn name(&self) -> &'static str {
"KnnTopK"
}
fn attrs(&self) -> Vec<(String, String)> {
let dimension = match &self.query_vector {
KnnVectorSource::Literal(v) => v.len().to_string(),
KnnVectorSource::Deferred(_) => "deferred".to_string(),
};
vec![
("field".to_string(), self.field.to_sql()),
("k".to_string(), self.k.to_string()),
("distance".to_string(), format!("{:?}", self.distance)),
("dimension".to_string(), dimension),
]
}
fn required_context(&self) -> ContextLevel {
let ctx = self.input.required_context();
match &self.query_vector {
KnnVectorSource::Deferred(expr) => ctx.max(expr.required_context()),
KnnVectorSource::Literal(_) => ctx,
}
}
fn access_mode(&self) -> AccessMode {
let mode = self.input.access_mode();
match &self.query_vector {
KnnVectorSource::Deferred(expr) => mode.combine(expr.access_mode()),
KnnVectorSource::Literal(_) => mode,
}
}
fn cardinality_hint(&self) -> CardinalityHint {
CardinalityHint::Bounded(self.k)
}
fn children(&self) -> Vec<&Arc<dyn ExecOperator>> {
vec![&self.input]
}
fn metrics(&self) -> Option<&OperatorMetrics> {
Some(&self.metrics)
}
fn execute(&self, ctx: &ExecutionContext) -> FlowResult<ValueBatchStream> {
let mut input_stream = buffer_stream(
self.input.execute(ctx)?,
self.input.access_mode(),
self.input.cardinality_hint(),
ctx.root().ctx.config.exec.operator_buffer_size,
);
let field = self.field.clone();
let query_vector = self.query_vector.clone();
let k = self.k;
let distance = self.distance.clone();
let cancellation = ctx.cancellation().clone();
let knn_context = self.knn_context.clone();
let exec_ctx = ctx.clone();
let result_stream = futures::stream::once(async move {
let query_vector: Vec<Number> = match &query_vector {
KnnVectorSource::Literal(v) => v.clone(),
KnnVectorSource::Deferred(expr) => {
let value = expr.evaluate(EvalContext::from_exec_ctx(&exec_ctx)).await?;
value
.coerce_to::<Vec<Number>>()
.map_err(|e| crate::expr::ControlFlow::Err(anyhow::Error::new(e)))?
}
};
let mut heap: KnnTopKHeap<Value> = KnnTopKHeap::new(k);
while let Some(batch_result) = input_stream.next().await {
if cancellation.is_cancelled() {
return Err(crate::expr::ControlFlow::Err(anyhow::anyhow!(
EngineError::QueryCancelled
)));
}
let batch = match batch_result {
Ok(b) => b,
Err(e) => return Err(e),
};
for value in batch.into_values() {
let record_vec = match extract_vector(&value, &field) {
Some(v) => v,
None => continue, };
let dist = match crate::idx::trees::vector::distance_compute(
&distance,
&record_vec,
&query_vector,
) {
Ok(d) => d,
Err(_) => continue, };
heap.offer(dist, value);
}
}
let entries = heap.into_sorted_nearest_first();
if let Some(ref knn_ctx) = knn_context {
for (distance, value) in &entries {
if let Value::Object(obj) = value
&& let Some(Value::RecordId(rid)) = obj.get("id")
{
knn_ctx.insert(rid.clone(), *distance).await;
}
}
}
let sorted: Vec<Value> = entries.into_iter().map(|(_, v)| v).collect();
Ok(ValueBatch::new(sorted))
});
let filtered = result_stream.filter_map(|result| async move {
match result {
Ok(batch) if batch.is_empty() => None,
other => Some(other),
}
});
Ok(monitor_stream(Box::pin(filtered), "KnnTopK", &self.metrics))
}
}
pub(crate) fn extract_vector(value: &Value, field: &Idiom) -> Option<Vec<Number>> {
match value.pick(field) {
Value::Array(arr) if !arr.is_empty() => {
let mut nums = Vec::with_capacity(arr.len());
for v in arr.iter() {
match v {
Value::Number(n) => nums.push(*n),
_ => return None,
}
}
Some(nums)
}
_ => None,
}
}