use std::sync::Arc;
use datafusion::arrow::array::RecordBatch;
use datafusion::arrow::datatypes::SchemaRef;
use datafusion::dataframe::DataFrame;
use datafusion::logical_expr::JoinType;
use datafusion::physical_plan::displayable;
use datafusion::prelude::{Expr, ParquetReadOptions, lit};
use oxidelake_core::EngineError;
use oxidelake_core::params::DistanceMetric;
use oxidelake_runtime::OxideSession;
use oxidelake_runtime::udf::{cosine_distance_udf, l2_distance_udf, query_literal};
#[derive(Debug, Clone)]
pub struct OxideFrame {
inner: DataFrame,
}
impl OxideFrame {
pub fn new(inner: DataFrame) -> Self {
Self { inner }
}
pub fn into_inner(self) -> DataFrame {
self.inner
}
pub fn schema(&self) -> SchemaRef {
Arc::new(self.inner.schema().as_arrow().clone())
}
pub fn filter(self, predicate: Expr) -> Result<Self, EngineError> {
Ok(Self::new(self.inner.filter(predicate)?))
}
pub fn select(self, columns: &[&str]) -> Result<Self, EngineError> {
Ok(Self::new(self.inner.select_columns(columns)?))
}
pub fn alias(self, name: &str) -> Result<Self, EngineError> {
Ok(Self::new(self.inner.alias(name)?))
}
pub fn join(
self,
right: OxideFrame,
left_key: &str,
right_key: &str,
) -> Result<Self, EngineError> {
Ok(Self::new(self.inner.join(
right.inner,
JoinType::Inner,
&[left_key],
&[right_key],
None,
)?))
}
pub fn aggregate(
self,
group_by: Vec<Expr>,
aggregates: Vec<Expr>,
) -> Result<Self, EngineError> {
Ok(Self::new(self.inner.aggregate(group_by, aggregates)?))
}
pub fn sort(self, exprs: Vec<datafusion::logical_expr::SortExpr>) -> Result<Self, EngineError> {
Ok(Self::new(self.inner.sort(exprs)?))
}
pub fn limit(self, n: usize) -> Result<Self, EngineError> {
Ok(Self::new(self.inner.limit(0, Some(n))?))
}
pub fn vector_distance(
self,
column: &str,
query: &[f32],
metric: DistanceMetric,
output: &str,
) -> Result<Self, EngineError> {
let udf = match metric {
DistanceMetric::L2 => l2_distance_udf(),
DistanceMetric::Cosine => cosine_distance_udf(),
};
let call = udf.call(vec![
datafusion::prelude::col(column),
lit(query_literal(query)),
]);
Ok(Self::new(self.inner.with_column(output, call)?))
}
pub async fn explain(&self) -> Result<String, EngineError> {
let plan = self.inner.clone().create_physical_plan().await?;
Ok(displayable(plan.as_ref()).indent(true).to_string())
}
pub async fn collect(self) -> Result<Vec<RecordBatch>, EngineError> {
Ok(self.inner.collect().await?)
}
pub async fn show(self) -> Result<(), EngineError> {
Ok(self.inner.show().await?)
}
}
impl From<DataFrame> for OxideFrame {
fn from(inner: DataFrame) -> Self {
Self::new(inner)
}
}
pub trait OxideSessionExt {
fn read_parquet(
&self,
path: &str,
) -> impl Future<Output = Result<OxideFrame, EngineError>> + Send;
fn table(&self, name: &str) -> impl Future<Output = Result<OxideFrame, EngineError>> + Send;
fn sql_frame(
&self,
query: &str,
) -> impl Future<Output = Result<OxideFrame, EngineError>> + Send;
}
impl OxideSessionExt for OxideSession {
async fn read_parquet(&self, path: &str) -> Result<OxideFrame, EngineError> {
Ok(OxideFrame::new(
self.ctx()
.read_parquet(path, ParquetReadOptions::default())
.await?,
))
}
async fn table(&self, name: &str) -> Result<OxideFrame, EngineError> {
Ok(OxideFrame::new(self.ctx().table(name).await?))
}
async fn sql_frame(&self, query: &str) -> Result<OxideFrame, EngineError> {
Ok(OxideFrame::new(self.sql(query).await?))
}
}