use std::any::Any;
use std::sync::Arc;
use datafusion::arrow::array::*;
use datafusion::arrow::datatypes::*;
use datafusion::catalog::{Session, TableProvider};
use datafusion::datasource::TableType;
use datafusion::error::{DataFusionError, Result as DataFusionResult};
use datafusion::execution::{SendableRecordBatchStream, TaskContext};
use datafusion::logical_expr::dml::InsertOp;
use datafusion::physical_plan::execution_plan::Boundedness;
use datafusion::physical_plan::{DisplayAs, DisplayFormatType, ExecutionPlan, PlanProperties};
use datafusion::prelude::Expr;
use datafusion::sql::TableReference;
use qdrant_client::Qdrant;
use qdrant_client::qdrant::{QueryPointsBuilder, VectorsSelector};
use crate::arrow::deserialize::QdrantRecordBatchBuilder;
use crate::arrow::schema::collection_to_arrow_schema;
use crate::error::{Error, Result};
use crate::stream::QdrantQueryStream;
use crate::utils;
#[derive(Clone)]
pub struct QdrantTableProvider {
table: TableReference,
client: Arc<Qdrant>,
schema: Arc<Schema>,
}
impl std::fmt::Debug for QdrantTableProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("QdrantTableProvider")
.field("table", &self.table)
.field("client", &"Qdrant")
.field("schema", &self.schema)
.finish()
}
}
impl QdrantTableProvider {
pub async fn try_new(client: Qdrant, collection: &str) -> Result<Self> {
let info = client.collection_info(collection).await?;
let config = info
.result
.ok_or(Error::MissingCollectionInfo(collection.into()))?
.config
.ok_or(Error::MissingCollectionInfo(collection.into()))?;
let schema = collection_to_arrow_schema(collection, &config)?;
Ok(Self {
table: TableReference::bare(collection),
client: Arc::new(client),
schema: Arc::new(schema),
})
}
}
#[async_trait::async_trait]
impl TableProvider for QdrantTableProvider {
fn as_any(&self) -> &dyn Any { self }
fn schema(&self) -> SchemaRef { Arc::clone(&self.schema) }
fn table_type(&self) -> TableType { TableType::Base }
async fn scan(
&self,
_state: &dyn Session,
projection: Option<&Vec<usize>>,
filters: &[Expr],
limit: Option<usize>,
) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
let projected_schema = match projection {
Some(indices) if !indices.is_empty() => Arc::new(self.schema.project(indices)?),
_ => Arc::clone(&self.schema),
};
let vector_selector = utils::build_vector_selector(&projected_schema);
let payload_selector = utils::build_payload_selector(&projected_schema);
Ok(Arc::new(QdrantScanExec::new(
Arc::clone(&self.client),
self.table.table().to_string(),
projected_schema,
vector_selector,
payload_selector,
filters,
limit,
)))
}
async fn insert_into(
&self,
_state: &dyn Session,
_input: Arc<dyn ExecutionPlan>,
_insert_op: InsertOp,
) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
todo!()
}
}
#[derive(Clone)]
pub struct QdrantScanExec {
client: Arc<Qdrant>,
collection: String,
schema: SchemaRef, vector_selector: utils::VectorSelectorSpec,
payload_selector: bool,
filter: Arc<[Expr]>,
limit: Option<usize>,
properties: PlanProperties,
}
impl std::fmt::Debug for QdrantScanExec {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("QdrantScanExec")
.field("client", &"Qdrant")
.field("collection", &self.collection)
.field("schema", &self.schema)
.field("vector_selector", &self.vector_selector)
.field("payload_selector", &self.payload_selector)
.field("limit", &self.limit)
.finish_non_exhaustive()
}
}
impl QdrantScanExec {
pub fn new(
client: Arc<Qdrant>,
collection: String,
schema: SchemaRef,
vector_selector: utils::VectorSelectorSpec,
payload_selector: bool,
filter: &[Expr],
limit: Option<usize>,
) -> Self {
let properties = PlanProperties::new(
datafusion::physical_expr::EquivalenceProperties::new(Arc::clone(&schema)),
datafusion::physical_plan::Partitioning::UnknownPartitioning(1),
datafusion::physical_plan::execution_plan::EmissionType::Final,
Boundedness::Bounded,
);
Self {
client,
collection,
schema,
vector_selector,
payload_selector,
filter: Arc::from(filter),
limit,
properties,
}
}
}
pub(crate) async fn execute_qdrant_query(
client: Arc<Qdrant>,
collection: String,
schema: SchemaRef,
vector_selector: utils::VectorSelectorSpec,
payload_selector: bool,
_filters: &[Expr],
limit: Option<usize>,
) -> DataFusionResult<RecordBatch> {
let mut query_builder = QueryPointsBuilder::new(&collection);
match vector_selector {
utils::VectorSelectorSpec::None => {
query_builder = query_builder.with_vectors(false);
}
utils::VectorSelectorSpec::All => {
query_builder = query_builder.with_vectors(true);
}
utils::VectorSelectorSpec::Named(names) => {
query_builder = query_builder.with_vectors(VectorsSelector { names });
}
}
query_builder = query_builder.with_payload(payload_selector);
if let Some(limit_val) = limit {
query_builder = query_builder.limit(limit_val as u64);
}
let response =
client.query(query_builder).await.map_err(|e| DataFusionError::External(Box::new(e)))?;
let points = response.result;
if points.is_empty() {
return Ok(RecordBatch::new_empty(schema));
}
let mut builder = QdrantRecordBatchBuilder::new(schema, points.len());
for point in points {
builder.append_point(point); }
builder.finish()
}
impl ExecutionPlan for QdrantScanExec {
fn name(&self) -> &'static str { "QdrantScanExec" }
fn as_any(&self) -> &dyn Any { self }
fn properties(&self) -> &PlanProperties { &self.properties }
fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> { vec![] }
fn with_new_children(
self: Arc<Self>,
_children: Vec<Arc<dyn ExecutionPlan>>,
) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
Ok(self)
}
fn execute(
&self,
_partition: usize,
_context: Arc<TaskContext>,
) -> DataFusionResult<SendableRecordBatchStream> {
let client = Arc::clone(&self.client);
let collection = self.collection.clone();
let schema = Arc::clone(&self.schema);
let vector_selector = self.vector_selector.clone();
let payload_selector = self.payload_selector;
let filter = Arc::clone(&self.filter);
let limit = self.limit;
let inner = Box::pin(futures_util::stream::once(async move {
execute_qdrant_query(
client,
collection,
schema,
vector_selector,
payload_selector,
&filter,
limit,
)
.await
}));
let stream = QdrantQueryStream::new(Arc::clone(&self.schema), inner);
Ok(Box::pin(stream))
}
}
impl DisplayAs for QdrantScanExec {
fn fmt_as(&self, t: DisplayFormatType, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match t {
DisplayFormatType::Default | DisplayFormatType::Verbose => {
write!(f, "QdrantScanExec: collection={}", self.collection)?;
if let Some(limit) = self.limit {
write!(f, ", limit={limit}")?;
}
Ok(())
}
DisplayFormatType::TreeRender => {
write!(f, "QdrantScanExec: collection={}", self.collection)
}
}
}
}