use std::any::Any;
use std::fmt;
use std::sync::Arc;
use arrow::datatypes::SchemaRef;
use datafusion::catalog::Session;
use datafusion::common::Statistics;
use datafusion::datasource::{TableProvider, TableType};
use datafusion::error::Result as DFResult;
use datafusion::logical_expr::{Expr, Operator, TableProviderFilterPushDown};
use datafusion::physical_plan::ExecutionPlan;
use super::ResolvedSubscription;
use super::engine::RssEngine;
use super::exec::{RssScanExec, RssTableKind};
const FEED_COLUMN: &str = "feed";
const FEED_URL_COLUMN: &str = "feed_url";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FeedKey {
Name,
Url,
}
impl FeedKey {
fn of(column: &str) -> Option<Self> {
match column {
FEED_COLUMN => Some(Self::Name),
FEED_URL_COLUMN => Some(Self::Url),
_ => None,
}
}
fn read(self, sub: &ResolvedSubscription) -> &str {
match self {
Self::Name => &sub.name,
Self::Url => &sub.url,
}
}
}
struct FeedFilter<'a> {
key: FeedKey,
values: Vec<&'a str>,
}
impl FeedFilter<'_> {
fn admits(&self, sub: &ResolvedSubscription) -> bool {
self.values.contains(&self.key.read(sub))
}
}
fn feed_filter(expr: &Expr) -> Option<FeedFilter<'_>> {
match expr {
Expr::BinaryExpr(binary) if binary.op == Operator::Eq => {
let (column, literal) = match (binary.left.as_ref(), binary.right.as_ref()) {
(Expr::Column(column), literal) | (literal, Expr::Column(column)) => {
(column, literal)
}
_ => return None,
};
Some(FeedFilter {
key: FeedKey::of(&column.name)?,
values: vec![string_literal(literal)?],
})
}
Expr::InList(in_list) if !in_list.negated => {
let Expr::Column(column) = in_list.expr.as_ref() else {
return None;
};
let key = FeedKey::of(&column.name)?;
let mut values = Vec::with_capacity(in_list.list.len());
for item in &in_list.list {
values.push(string_literal(item)?);
}
Some(FeedFilter { key, values })
}
Expr::BinaryExpr(binary) if binary.op == Operator::Or => {
let left = feed_filter(&binary.left)?;
let right = feed_filter(&binary.right)?;
if left.key != right.key {
return None;
}
let mut values = left.values;
values.extend(right.values);
Some(FeedFilter {
key: left.key,
values,
})
}
_ => None,
}
}
fn string_literal(expr: &Expr) -> Option<&str> {
match expr {
Expr::Literal(value, _) => value.try_as_str().flatten(),
_ => None,
}
}
pub(crate) fn prune_feeds(filters: &[Expr], subs: &[ResolvedSubscription]) -> Vec<String> {
let mut kept: Vec<&ResolvedSubscription> = subs.iter().collect();
for filter in filters {
if let Some(prunable) = feed_filter(filter) {
kept.retain(|sub| prunable.admits(sub));
}
}
kept.into_iter().map(|sub| sub.name.clone()).collect()
}
pub struct RssTableProvider {
engine: Arc<RssEngine>,
kind: RssTableKind,
schema: SchemaRef,
}
impl fmt::Debug for RssTableProvider {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RssTableProvider")
.field("kind", &self.kind)
.field("subscriptions", &self.engine.subscriptions().len())
.finish()
}
}
impl RssTableProvider {
pub fn feeds(engine: Arc<RssEngine>) -> Self {
Self::new(engine, RssTableKind::Feeds)
}
pub fn items(engine: Arc<RssEngine>) -> Self {
Self::new(engine, RssTableKind::Items)
}
fn new(engine: Arc<RssEngine>, kind: RssTableKind) -> Self {
Self {
engine,
kind,
schema: kind.schema(),
}
}
fn all_feeds(&self) -> Vec<String> {
self.engine
.subscriptions()
.iter()
.map(|sub| sub.name.clone())
.collect()
}
}
#[async_trait::async_trait]
impl TableProvider for RssTableProvider {
fn as_any(&self) -> &dyn Any {
self
}
fn schema(&self) -> SchemaRef {
Arc::clone(&self.schema)
}
fn table_type(&self) -> TableType {
TableType::Base
}
fn supports_filters_pushdown(
&self,
filters: &[&Expr],
) -> DFResult<Vec<TableProviderFilterPushDown>> {
Ok(filters
.iter()
.map(|filter| match self.kind {
RssTableKind::Items if feed_filter(filter).is_some() => {
TableProviderFilterPushDown::Exact
}
_ => TableProviderFilterPushDown::Unsupported,
})
.collect())
}
async fn scan(
&self,
_state: &dyn Session,
projection: Option<&Vec<usize>>,
filters: &[Expr],
limit: Option<usize>,
) -> DFResult<Arc<dyn ExecutionPlan>> {
let feeds = match self.kind {
RssTableKind::Items => prune_feeds(filters, self.engine.subscriptions()),
RssTableKind::Feeds => self.all_feeds(),
};
Ok(Arc::new(RssScanExec::new(
Arc::clone(&self.engine),
self.kind,
feeds,
projection.cloned(),
limit,
)?))
}
fn statistics(&self) -> Option<Statistics> {
None
}
}
#[cfg(test)]
mod tests {
use arrow::datatypes::DataType;
use datafusion::common::ScalarValue;
use datafusion::execution::TaskContext;
use datafusion::logical_expr::TableProviderFilterPushDown::{Exact, Unsupported};
use datafusion::logical_expr::{cast, col, lit};
use datafusion::prelude::SessionContext;
use super::*;
use crate::sources::providers::rss::schema::{feeds_schema, items_schema};
use crate::sources::providers::rss::testutil::{
MockFeedServer, MockResponse, MockResponseExt, RSS2_MINIMAL, collect_stream, feed_urls,
str_col, test_engine, total_rows,
};
fn subs(pairs: &[(&str, &str)]) -> Vec<ResolvedSubscription> {
pairs
.iter()
.map(|(name, url)| ResolvedSubscription {
name: (*name).to_string(),
url: (*url).to_string(),
})
.collect()
}
fn offline_engine(feeds: &[&str]) -> Arc<RssEngine> {
let urls: Vec<(String, String)> = feeds
.iter()
.map(|name| {
(
(*name).to_string(),
format!("http://feed.invalid/{name}.xml"),
)
})
.collect();
Arc::new(test_engine(&urls, |_| {}))
}
fn items_provider_with_feeds(feeds: &[&str]) -> RssTableProvider {
RssTableProvider::items(offline_engine(feeds))
}
fn one_line(plan: &Arc<dyn ExecutionPlan>) -> String {
datafusion::physical_plan::displayable(plan.as_ref())
.one_line()
.to_string()
.trim_end()
.to_string()
}
fn sql_context(engine: &Arc<RssEngine>) -> SessionContext {
let ctx = SessionContext::new();
ctx.register_table(
"items",
Arc::new(RssTableProvider::items(Arc::clone(engine))),
)
.expect("register items");
ctx.register_table(
"feeds",
Arc::new(RssTableProvider::feeds(Arc::clone(engine))),
)
.expect("register feeds");
ctx
}
async fn query(ctx: &SessionContext, sql: &str) -> Vec<arrow::record_batch::RecordBatch> {
ctx.sql(sql)
.await
.unwrap_or_else(|e| panic!("plan {sql:?}: {e}"))
.collect()
.await
.unwrap_or_else(|e| panic!("execute {sql:?}: {e}"))
}
fn column(batches: &[arrow::record_batch::RecordBatch], name: &str) -> Vec<String> {
batches.iter().flat_map(|b| str_col(b, name)).collect()
}
#[tokio::test]
async fn pushdown_classification_is_exact_only_for_feed_predicates() {
let p = items_provider_with_feeds(&["a", "b"]);
let feed_eq = col("feed").eq(lit("a"));
let url_in = col("feed_url").in_list(vec![lit("http://x/1")], false);
let title_eq = col("title").eq(lit("t"));
let feed_gt = col("feed").gt(lit("a"));
let neg_in = col("feed").in_list(vec![lit("a")], true);
let got = p
.supports_filters_pushdown(&[&feed_eq, &url_in, &title_eq, &feed_gt, &neg_in])
.expect("classify");
assert_eq!(
got,
vec![Exact, Exact, Unsupported, Unsupported, Unsupported]
);
}
#[tokio::test]
async fn the_feeds_table_pushes_down_nothing() {
let p = RssTableProvider::feeds(offline_engine(&["a", "b"]));
let name_eq = col("name").eq(lit("a"));
let url_in = col("url").in_list(vec![lit("http://x/1")], false);
let feed_eq = col("feed").eq(lit("a"));
let got = p
.supports_filters_pushdown(&[&name_eq, &url_in, &feed_eq])
.expect("classify");
assert_eq!(got, vec![Unsupported, Unsupported, Unsupported]);
}
#[test]
fn prune_intersects_predicates_and_maps_urls() {
let subs = subs(&[
("a", "http://x/a"),
("b", "http://x/b"),
("c", "http://x/c"),
]);
assert_eq!(prune_feeds(&[col("feed").eq(lit("b"))], &subs), vec!["b"]);
assert_eq!(
prune_feeds(
&[col("feed").in_list(vec![lit("a"), lit("c")], false)],
&subs
),
vec!["a", "c"]
);
assert_eq!(
prune_feeds(&[col("feed_url").eq(lit("http://x/b"))], &subs),
vec!["b"]
);
assert_eq!(
prune_feeds(
&[
col("feed").in_list(vec![lit("a"), lit("b")], false),
col("feed").eq(lit("b")),
],
&subs
),
vec!["b"]
);
assert!(prune_feeds(&[col("feed").eq(lit("nope"))], &subs).is_empty());
assert!(
prune_feeds(&[col("feed").eq(lit("a")), col("feed").eq(lit("b"))], &subs).is_empty()
);
assert_eq!(prune_feeds(&[lit("a").eq(col("feed"))], &subs), vec!["a"]);
assert!(
prune_feeds(
&[
col("feed_url").eq(lit("http://x/a")),
col("feed").eq(lit("b")),
],
&subs
)
.is_empty()
);
assert_eq!(prune_feeds(&[col("title").eq(lit("t"))], &subs).len(), 3);
assert_eq!(prune_feeds(&[], &subs).len(), 3);
}
#[test]
fn pruning_preserves_subscription_order() {
let subs = subs(&[
("a", "http://x/a"),
("b", "http://x/b"),
("c", "http://x/c"),
]);
assert_eq!(
prune_feeds(
&[col("feed").in_list(vec![lit("c"), lit("a")], false)],
&subs
),
vec!["a", "c"]
);
}
#[tokio::test]
async fn disjunctions_of_feed_equalities_prune_to_the_union_and_claim_exact() {
let subs = subs(&[
("a", "http://x/a"),
("b", "http://x/b"),
("c", "http://x/c"),
]);
let provider = items_provider_with_feeds(&["a", "b", "c"]);
let cases: Vec<(&str, Expr, Vec<&str>)> = vec![
(
"two equalities",
col("feed").eq(lit("a")).or(col("feed").eq(lit("b"))),
vec!["a", "b"],
),
(
"three equalities, left-deep as the simplifier builds them",
col("feed")
.eq(lit("a"))
.or(col("feed").eq(lit("b")))
.or(col("feed").eq(lit("c"))),
vec!["a", "b", "c"],
),
(
"three equalities, right-deep",
col("feed")
.eq(lit("a"))
.or(col("feed").eq(lit("b")).or(col("feed").eq(lit("c")))),
vec!["a", "b", "c"],
),
(
"the union in predicate order, reported in subscription order",
col("feed").eq(lit("c")).or(col("feed").eq(lit("a"))),
vec!["a", "c"],
),
(
"a reversed-operand leaf",
lit("a").eq(col("feed")).or(col("feed").eq(lit("c"))),
vec!["a", "c"],
),
(
"a duplicated leaf, which names one feed once",
col("feed").eq(lit("a")).or(col("feed").eq(lit("a"))),
vec!["a"],
),
(
"a leaf naming no subscription, which adds nothing",
col("feed").eq(lit("a")).or(col("feed").eq(lit("nope"))),
vec!["a"],
),
(
"every leaf naming no subscription, which is unsatisfiable",
col("feed").eq(lit("nope")).or(col("feed").eq(lit("nor"))),
vec![],
),
(
"leaves on feed_url, mapped back to subscription names",
col("feed_url")
.eq(lit("http://x/a"))
.or(col("feed_url").eq(lit("http://x/c"))),
vec!["a", "c"],
),
(
"an IN list as a leaf, whose members join the union",
col("feed")
.eq(lit("a"))
.or(col("feed").in_list(vec![lit("b"), lit("c")], false)),
vec!["a", "b", "c"],
),
];
for (why, filter, expected) in cases {
assert_eq!(
prune_feeds(std::slice::from_ref(&filter), &subs),
expected,
"{why} must prune to the union of its leaves: {filter}"
);
assert_eq!(
provider
.supports_filters_pushdown(&[&filter])
.expect("classify"),
vec![Exact],
"{why} prunes, so it must be claimed Exact: {filter}"
);
}
}
#[test]
fn a_disjunction_intersects_with_the_predicates_beside_it() {
let subs = subs(&[
("a", "http://x/a"),
("b", "http://x/b"),
("c", "http://x/c"),
]);
assert_eq!(
prune_feeds(
&[
col("feed").eq(lit("a")).or(col("feed").eq(lit("b"))),
col("feed").eq(lit("b")).or(col("feed").eq(lit("c"))),
],
&subs
),
vec!["b"]
);
}
#[tokio::test]
async fn shapes_outside_the_allowlist_prune_nothing_and_claim_nothing() {
let subs = subs(&[
("a", "http://x/a"),
("b", "http://x/b"),
("c", "http://x/c"),
]);
let provider = items_provider_with_feeds(&["a", "b", "c"]);
let cases: Vec<(&str, Expr)> = vec![
("an inequality", col("feed").gt(lit("a"))),
("a negated equality", col("feed").not_eq(lit("a"))),
("a negated IN", col("feed").in_list(vec![lit("a")], true)),
("another column", col("title").eq(lit("t"))),
("a non-string literal", col("feed").eq(lit(1_i64))),
(
"a NULL literal",
col("feed").eq(Expr::Literal(ScalarValue::Utf8(None), None)),
),
(
"a non-literal IN element",
col("feed").in_list(vec![lit("a"), col("feed_url")], false),
),
(
"a wrapped column",
cast(col("feed"), DataType::Utf8View).eq(lit("a")),
),
("a column on both sides", col("feed").eq(col("feed_url"))),
("an IS NULL", col("feed").is_null()),
(
"a disjunction with a non-feed leaf",
col("feed").eq(lit("a")).or(col("title").eq(lit("t"))),
),
(
"a disjunction with an inequality leaf",
col("feed").eq(lit("a")).or(col("feed").gt(lit("b"))),
),
(
"a disjunction with a negated leaf",
col("feed").eq(lit("a")).or(col("feed").not_eq(lit("b"))),
),
(
"a disjunction mixing the two feed columns",
col("feed")
.eq(lit("a"))
.or(col("feed_url").eq(lit("http://x/c"))),
),
(
"a conjunction nested inside a disjunction",
col("feed")
.eq(lit("a"))
.or(col("feed").eq(lit("b")).and(col("title").eq(lit("t")))),
),
(
"a disjunction whose non-feed leaf is itself a disjunction",
col("feed")
.eq(lit("a"))
.or(col("title").eq(lit("t")).or(col("feed").eq(lit("b")))),
),
];
for (why, filter) in cases {
assert_eq!(
prune_feeds(std::slice::from_ref(&filter), &subs).len(),
3,
"{why} must prune nothing, not prune to empty: {filter}"
);
assert_eq!(
provider
.supports_filters_pushdown(&[&filter])
.expect("classify"),
vec![Unsupported],
"{why} must not be claimed Exact: {filter}"
);
}
}
#[tokio::test]
async fn schema_and_table_type_match_the_kind() {
let items = items_provider_with_feeds(&["a"]);
let feeds = RssTableProvider::feeds(offline_engine(&["a"]));
assert_eq!(items.schema(), items_schema());
assert_eq!(feeds.schema(), feeds_schema());
assert_eq!(items.table_type(), TableType::Base);
assert_eq!(feeds.table_type(), TableType::Base);
assert!(items.statistics().is_none());
}
#[tokio::test]
async fn the_tables_are_read_only() {
let provider = items_provider_with_feeds(&["a"]);
let ctx = SessionContext::new();
let state = ctx.state();
let input = provider
.scan(&state, None, &[], None)
.await
.expect("build a scan to feed the insert");
let error = provider
.insert_into(
&state,
input,
datafusion::logical_expr::dml::InsertOp::Append,
)
.await
.expect_err("an rss table takes no writes");
let message = error.to_string();
assert!(
message.contains("Insert into not implemented for this table"),
"the refusal names the unimplemented operation: {message}"
);
}
#[tokio::test]
async fn scan_hands_the_exec_the_pruned_feeds_the_projection_and_the_limit() {
let provider = items_provider_with_feeds(&["a", "b", "c"]);
let ctx = SessionContext::new();
let state = ctx.state();
let plan = provider
.scan(&state, Some(&vec![0]), &[col("feed").eq(lit("b"))], Some(5))
.await
.expect("scan");
assert_eq!(
one_line(&plan),
"RssScanExec: kind=items feeds=1 limit=Some(5)"
);
assert_eq!(plan.properties().partitioning.partition_count(), 1);
assert_eq!(
plan.schema().fields().len(),
1,
"the projection reached the exec"
);
}
#[tokio::test]
async fn a_zero_match_predicate_yields_one_empty_partition() {
let provider = items_provider_with_feeds(&["a", "b"]);
let ctx = SessionContext::new();
let plan = provider
.scan(&ctx.state(), None, &[col("feed").eq(lit("nope"))], None)
.await
.expect("scan");
assert_eq!(
one_line(&plan),
"RssScanExec: kind=items feeds=0 limit=None"
);
assert_eq!(plan.properties().partitioning.partition_count(), 1);
let batches = collect_stream(plan.execute(0, Arc::new(TaskContext::default()))).await;
assert!(batches.is_empty(), "no feed to visit means no rows");
}
#[tokio::test]
async fn a_feeds_scan_visits_every_subscription_whatever_the_filters() {
let provider = RssTableProvider::feeds(offline_engine(&["a", "b", "c"]));
let ctx = SessionContext::new();
let plan = provider
.scan(&ctx.state(), None, &[col("name").eq(lit("a"))], None)
.await
.expect("scan");
assert_eq!(
one_line(&plan),
"RssScanExec: kind=feeds feeds=3 limit=None"
);
assert_eq!(plan.properties().partitioning.partition_count(), 3);
}
#[tokio::test]
async fn end_to_end_sql_prunes_to_one_fetch() {
let server = MockFeedServer::start(|_| MockResponse::xml(RSS2_MINIMAL)).await;
let engine = Arc::new(test_engine(
&feed_urls(&server, &[("a", "/a.xml"), ("b", "/b.xml")]),
|_| {},
));
let ctx = sql_context(&engine);
let batches = query(&ctx, "SELECT feed, guid FROM items WHERE feed = 'a'").await;
assert_eq!(total_rows(&batches), 1);
assert_eq!(column(&batches, "feed"), vec!["a"]);
let paths: Vec<String> = server.requests().iter().map(|r| r.path.clone()).collect();
assert_eq!(
paths,
vec!["/a.xml".to_string()],
"the pruned-away subscription must never be fetched"
);
}
#[tokio::test]
async fn end_to_end_sql_prunes_on_feed_url() {
let server = MockFeedServer::start(|_| MockResponse::xml(RSS2_MINIMAL)).await;
let feeds = feed_urls(&server, &[("a", "/a.xml"), ("b", "/b.xml")]);
let engine = Arc::new(test_engine(&feeds, |_| {}));
let ctx = sql_context(&engine);
let url = feeds[1].1.clone();
let batches = query(
&ctx,
&format!("SELECT feed FROM items WHERE feed_url = '{url}'"),
)
.await;
assert_eq!(column(&batches, "feed"), vec!["b"]);
let paths: Vec<String> = server.requests().iter().map(|r| r.path.clone()).collect();
assert_eq!(paths, vec!["/b.xml".to_string()]);
}
#[tokio::test]
async fn end_to_end_sql_with_an_unknown_feed_fetches_nothing() {
let server = MockFeedServer::start(|_| MockResponse::xml(RSS2_MINIMAL)).await;
let engine = Arc::new(test_engine(
&feed_urls(&server, &[("a", "/a.xml"), ("b", "/b.xml")]),
|_| {},
));
let ctx = sql_context(&engine);
let batches = query(&ctx, "SELECT guid FROM items WHERE feed = 'nope'").await;
assert_eq!(total_rows(&batches), 0);
assert!(
server.requests().is_empty(),
"an unsatisfiable feed predicate must issue no requests"
);
}
#[tokio::test]
async fn non_prunable_filters_are_applied_above_the_scan() {
let server = MockFeedServer::start(|_| MockResponse::xml(RSS2_MINIMAL)).await;
let engine = Arc::new(test_engine(
&feed_urls(&server, &[("a", "/a.xml"), ("b", "/b.xml")]),
|_| {},
));
let ctx = sql_context(&engine);
let none = query(&ctx, "SELECT guid FROM items WHERE title = 'nope'").await;
assert_eq!(
total_rows(&none),
0,
"a filter on a non-prunable column must still remove rows"
);
let after_a = query(&ctx, "SELECT feed FROM items WHERE feed > 'a'").await;
assert_eq!(
column(&after_a, "feed"),
vec!["b"],
"an unsupported operator on `feed` must be filtered, not ignored"
);
assert_eq!(
server.requests().len(),
2,
"neither query prunes, so both subscriptions are visited"
);
}
#[tokio::test]
async fn a_feeds_scan_is_total_over_subscriptions() {
let server = MockFeedServer::start(|req| match req.path.as_str() {
"/a.xml" => MockResponse::xml(RSS2_MINIMAL),
_ => MockResponse::status(500),
})
.await;
let engine = Arc::new(test_engine(
&feed_urls(&server, &[("a", "/a.xml"), ("b", "/b.xml")]),
|_| {},
));
let ctx = sql_context(&engine);
let all = query(&ctx, "SELECT name FROM feeds ORDER BY name").await;
assert_eq!(column(&all, "name"), vec!["a", "b"]);
let missing = query(
&ctx,
"SELECT f.name FROM feeds f LEFT JOIN items i ON i.feed = f.name \
WHERE i.feed IS NULL ORDER BY f.name",
)
.await;
assert_eq!(
column(&missing, "name"),
vec!["b"],
"the dead subscription is the one with no items"
);
}
}