use std::any::Any;
use arrow::datatypes::DataType;
use datafusion::common::not_impl_err;
use datafusion::error::Result as DFResult;
use datafusion::logical_expr::{
ColumnarValue, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl, Signature, TypeSignature,
Volatility,
};
use datafusion::prelude::SessionContext;
#[derive(Debug, PartialEq, Eq, Hash)]
struct PgFtsUdf {
name: &'static str,
signature: Signature,
return_type: DataType,
}
impl PgFtsUdf {
fn new(name: &'static str, arities: &[usize], return_type: DataType) -> Self {
Self {
name,
signature: Signature::one_of(
arities.iter().copied().map(TypeSignature::Any).collect(),
Volatility::Immutable,
),
return_type,
}
}
}
impl ScalarUDFImpl for PgFtsUdf {
fn as_any(&self) -> &dyn Any {
self
}
fn name(&self) -> &str {
self.name
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> DFResult<DataType> {
Ok(self.return_type.clone())
}
fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> DFResult<ColumnarValue> {
not_impl_err!(
"{} is evaluated by PostgreSQL; this query was not pushed down",
self.name
)
}
}
pub fn register_pg_fts_udfs(ctx: &SessionContext) {
for udf in [
PgFtsUdf::new("to_tsvector", &[1, 2], DataType::Utf8),
PgFtsUdf::new("websearch_to_tsquery", &[1, 2], DataType::Utf8),
PgFtsUdf::new("ts_rank", &[2, 3, 4], DataType::Float32),
] {
ctx.register_udf(ScalarUDF::new_from_impl(udf));
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use arrow::array::{RecordBatch, StringArray};
use arrow::datatypes::{Field, Schema, SchemaRef};
use datafusion::common::tree_node::TreeNode;
use datafusion::datasource::MemTable;
use datafusion::execution::FunctionRegistry;
use datafusion::logical_expr::{Expr, LogicalPlan, TableProviderFilterPushDown};
use datafusion::sql::unparser::Unparser;
use datafusion::sql::unparser::dialect::PostgreSqlDialect;
use datafusion::sql::unparser::plan_to_sql;
use datafusion_table_providers::sql::sql_provider_datafusion::default_filter_pushdown;
fn docs_schema() -> SchemaRef {
Arc::new(Schema::new(vec![
Field::new("path", DataType::Utf8, true),
Field::new("body", DataType::Utf8, true),
]))
}
fn ctx_with_docs() -> SessionContext {
let ctx = SessionContext::new();
register_pg_fts_udfs(&ctx);
let schema = docs_schema();
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(StringArray::from(vec!["/a.md"])),
Arc::new(StringArray::from(vec!["cats and dogs"])),
],
)
.expect("batch");
let table = MemTable::try_new(schema, vec![vec![batch]]).expect("memtable");
ctx.register_table("docs", Arc::new(table))
.expect("register");
ctx
}
const FTS_FUNCTION_NAMES: [&str; 3] = ["to_tsvector", "websearch_to_tsquery", "ts_rank"];
fn assert_all_three_are_immutable(ctx: &SessionContext) {
for name in FTS_FUNCTION_NAMES {
let udf = ctx.udf(name).unwrap_or_else(|e| panic!("{name}: {e}"));
assert_eq!(
udf.signature().volatility,
Volatility::Immutable,
"{name} must stay Immutable"
);
}
}
#[tokio::test]
async fn the_three_functions_are_registered_immutable() {
let ctx = SessionContext::new();
register_pg_fts_udfs(&ctx);
assert_all_three_are_immutable(&ctx);
}
const SEARCH_OKF_SQL: &str = "SELECT path, ts_rank(to_tsvector('english', body), \
websearch_to_tsquery('english', 'cats')) AS rank \
FROM docs WHERE to_tsvector('english', body) @@ \
websearch_to_tsquery('english', 'cats') \
ORDER BY rank DESC LIMIT 5";
#[tokio::test]
async fn the_three_fts_functions_plan() {
let ctx = SessionContext::new();
register_pg_fts_udfs(&ctx);
for sql in [
"SELECT to_tsvector('english', 'a')",
"SELECT websearch_to_tsquery('english', 'a')",
"SELECT ts_rank(to_tsvector('english','a'), websearch_to_tsquery('english','a'))",
] {
ctx.state()
.create_logical_plan(sql)
.await
.unwrap_or_else(|e| panic!("{sql} must plan: {e}"));
}
}
#[rstest::rstest]
#[case::to_tsvector_default_regconfig("SELECT to_tsvector('a')")]
#[case::to_tsvector_explicit_regconfig("SELECT to_tsvector('english', 'a')")]
#[case::websearch_default_regconfig("SELECT websearch_to_tsquery('cats')")]
#[case::websearch_explicit_regconfig("SELECT websearch_to_tsquery('english', 'cats')")]
#[case::the_reviewers_predicate("SELECT to_tsvector('a') @@ websearch_to_tsquery('cats')")]
#[case::ts_rank_two_args(
"SELECT ts_rank(to_tsvector('english','a'), websearch_to_tsquery('english','a'))"
)]
#[case::ts_rank_with_normalization(
"SELECT ts_rank(to_tsvector('english','a'), websearch_to_tsquery('english','a'), 32)"
)]
#[case::ts_rank_with_leading_weights(
"SELECT ts_rank(ARRAY[0.1, 0.2, 0.4, 1.0], to_tsvector('english','a'), \
websearch_to_tsquery('english','a'))"
)]
#[case::ts_rank_with_weights_and_normalization(
"SELECT ts_rank(ARRAY[0.1, 0.2, 0.4, 1.0], to_tsvector('english','a'), \
websearch_to_tsquery('english','a'), 32)"
)]
#[tokio::test]
async fn every_postgres_overload_plans(#[case] sql: &str) {
let ctx = SessionContext::new();
register_pg_fts_udfs(&ctx);
ctx.state()
.create_logical_plan(sql)
.await
.unwrap_or_else(|e| panic!("{sql} must plan: {e}"));
}
#[tokio::test]
async fn the_search_okf_predicate_shape_plans_end_to_end() {
let ctx = ctx_with_docs();
ctx.state()
.create_logical_plan(SEARCH_OKF_SQL)
.await
.unwrap_or_else(|e| panic!("the search-okf predicate must plan: {e}"));
}
#[tokio::test]
async fn the_at_at_operator_type_checks_over_the_two_string_returns() {
let ctx = ctx_with_docs();
let plan = ctx
.state()
.create_logical_plan(
"SELECT path FROM docs WHERE to_tsvector('english', body) @@ \
websearch_to_tsquery('english', 'cats')",
)
.await
.expect("the @@ predicate must plan");
assert!(
format!("{}", plan.display_indent()).contains("@@"),
"the operator survives planning: {}",
plan.display_indent()
);
}
#[tokio::test]
async fn constant_folding_does_not_kill_the_plan() {
let ctx = ctx_with_docs();
assert_all_three_are_immutable(&ctx);
let optimized = ctx
.sql(SEARCH_OKF_SQL)
.await
.expect("plans")
.into_optimized_plan()
.unwrap_or_else(|e| panic!("the search-okf predicate must optimize: {e}"));
let shown = format!("{}", optimized.display_indent());
for needle in ["to_tsvector", "websearch_to_tsquery", "ts_rank"] {
assert!(
shown.contains(needle),
"{needle} survives optimization rather than being folded away: {shown}"
);
}
}
#[tokio::test]
async fn executing_one_in_datafusion_refuses_rather_than_returning_a_wrong_answer() {
let ctx = ctx_with_docs();
for (sql, needle) in [
(
"SELECT to_tsvector('english', body) FROM docs",
"to_tsvector",
),
(
"SELECT websearch_to_tsquery('english', body) FROM docs",
"websearch_to_tsquery",
),
("SELECT ts_rank(body, body) FROM docs", "ts_rank"),
] {
let err = ctx
.sql(sql)
.await
.expect("plans")
.collect()
.await
.expect_err("local evaluation must refuse, never answer");
let msg = err.to_string();
assert!(msg.contains(needle), "the error names the function: {msg}");
assert!(msg.contains("not pushed down"), "and the cause: {msg}");
}
}
async fn optimized_filter(ctx: &SessionContext, sql: &str) -> Expr {
let plan = ctx
.sql(sql)
.await
.expect("plans")
.into_optimized_plan()
.expect("optimizes");
let mut found = None;
plan.apply(|node| {
if let LogicalPlan::Filter(f) = node {
found = Some(f.predicate.clone());
}
if let LogicalPlan::TableScan(t) = node
&& let Some(e) = t.filters.first()
{
found = Some(e.clone());
}
Ok(datafusion::common::tree_node::TreeNodeRecursion::Continue)
})
.expect("walk");
found.expect("the plan carries the FTS predicate")
}
#[tokio::test]
async fn the_predicate_unparses_back_to_postgres_sql() {
let ctx = ctx_with_docs();
let predicate = optimized_filter(
&ctx,
"SELECT path FROM docs WHERE to_tsvector('english', body) @@ \
websearch_to_tsquery('english', 'cats')",
)
.await;
let rendered = Unparser::new(&PostgreSqlDialect {})
.expr_to_sql(&predicate)
.unwrap_or_else(|e| panic!("the FTS predicate must unparse: {e}"))
.to_string();
eprintln!("PUSHED PREDICATE: {rendered}");
for needle in [
"to_tsvector('english'",
"websearch_to_tsquery('english'",
" @@ ",
] {
assert!(
rendered.contains(needle),
"unparsed predicate must contain {needle}: {rendered}"
);
}
assert_eq!(
default_filter_pushdown(&[&predicate], &PostgreSqlDialect {}),
vec![TableProviderFilterPushDown::Exact],
"the predicate goes to Postgres as text, not to DataFusion"
);
}
#[tokio::test]
async fn the_whole_search_okf_statement_unparses_back_to_postgres_sql() {
let ctx = ctx_with_docs();
let plan = ctx
.state()
.create_logical_plan(SEARCH_OKF_SQL)
.await
.expect("plans");
let rendered = plan_to_sql(&plan)
.unwrap_or_else(|e| panic!("the search-okf statement must unparse: {e}"))
.to_string();
eprintln!("PUSHED STATEMENT: {rendered}");
for needle in [
"to_tsvector('english'",
"websearch_to_tsquery('english'",
"ts_rank(",
"@@",
"ORDER BY",
"LIMIT",
] {
assert!(
rendered.contains(needle),
"unparsed statement must contain {needle}: {rendered}"
);
}
}
}