use std::sync::Arc;
use surrealdb_types::{SqlFormat, ToSql, write_sql};
use crate::exec::function::KnnContext;
use crate::exec::physical_expr::{EvalContext, PhysicalExpr};
use crate::exec::{AccessMode, BoxFut, ContextLevel};
use crate::expr::FlowResult;
use crate::expr::operator::BinaryOperator;
use crate::val::Value;
pub struct KnnMembershipOp {
pub(crate) left: Arc<dyn PhysicalExpr>,
pub(crate) right: Arc<dyn PhysicalExpr>,
pub(crate) operator: BinaryOperator,
knn_ctx: Option<Arc<KnnContext>>,
}
impl KnnMembershipOp {
pub(crate) fn new(
left: Arc<dyn PhysicalExpr>,
right: Arc<dyn PhysicalExpr>,
operator: BinaryOperator,
knn_ctx: Option<Arc<KnnContext>>,
) -> Self {
Self {
left,
right,
operator,
knn_ctx,
}
}
}
impl PhysicalExpr for KnnMembershipOp {
fn name(&self) -> &'static str {
"KnnMembershipOp"
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn required_context(&self) -> ContextLevel {
self.left.required_context().max(self.right.required_context())
}
fn evaluate<'a>(&'a self, ctx: EvalContext<'a>) -> BoxFut<'a, FlowResult<Value>> {
Box::pin(async move {
let Some(knn_ctx) = &self.knn_ctx else {
return Ok(Value::Bool(false));
};
let rid = match ctx.current_value {
Some(Value::Object(obj)) => match obj.get("id") {
Some(Value::RecordId(rid)) => rid.clone(),
_ => return Ok(Value::Bool(false)),
},
Some(Value::RecordId(rid)) => rid.clone(),
_ => return Ok(Value::Bool(false)),
};
Ok(Value::Bool(knn_ctx.get(&rid).await.is_some()))
})
}
fn access_mode(&self) -> AccessMode {
self.left.access_mode().combine(self.right.access_mode())
}
}
impl ToSql for KnnMembershipOp {
fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
write_sql!(f, fmt, "{} {} {}", self.left, self.operator, self.right);
}
}
impl std::fmt::Debug for KnnMembershipOp {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KnnMembershipOp")
.field("operator", &self.operator)
.field("bound", &self.knn_ctx.is_some())
.finish()
}
}