diff --git a/datafusion/core/src/dataframe/mod.rs b/datafusion/core/src/dataframe/mod.rs index ed3dc5ea838b9..55d3695a68d7a 100644 --- a/datafusion/core/src/dataframe/mod.rs +++ b/datafusion/core/src/dataframe/mod.rs @@ -1380,6 +1380,107 @@ impl DataFrame { }) } + /// Join this `DataFrame` to the closest eligible row in `right`. + /// + /// Every left row is emitted exactly once, with `NULL` values for the right + /// columns when no eligible row exists. `NULL` ordered values and equality + /// keys never match. When present, `on` must contain equality comparisons + /// combined with `AND`. `match_condition` must be a single `<`, `<=`, `>`, + /// or `>=` comparison whose left and right operands reference this + /// `DataFrame` and `right`, respectively. + /// + /// # Example + /// ``` + /// # use datafusion::arrow::array::record_batch; + /// # use datafusion::error::Result; + /// # use datafusion::prelude::*; + /// # #[tokio::main] + /// # async fn main() -> Result<()> { + /// # let ctx = SessionContext::new(); + /// # let trades = ctx.read_batch(record_batch!( + /// # ("symbol", Utf8, ["A"]), + /// # ("ts", Int64, [4]) + /// # )?)?.alias("trades")?; + /// # let prices = ctx.read_batch(record_batch!( + /// # ("symbol", Utf8, ["A"]), + /// # ("ts", Int64, [2]), + /// # ("price", Int32, [20]) + /// # )?)?.alias("prices")?; + /// // For each trade, find the latest price at or before its timestamp. + /// let joined = trades.join_asof( + /// prices, + /// Some(col("trades.symbol").eq(col("prices.symbol"))), + /// col("trades.ts").gt_eq(col("prices.ts")), + /// )?; + /// # let _ = joined; + /// # Ok(()) + /// # } + /// ``` + pub fn join_asof( + self, + right: DataFrame, + on: Option, + match_condition: Expr, + ) -> Result { + let plan = LogicalPlanBuilder::from(self.plan) + .asof_join_on(right.plan, on, match_condition)? + .build()?; + Ok(DataFrame { + session_state: self.session_state, + plan, + projection_requires_validation: true, + }) + } + + /// Join this `DataFrame` to the closest eligible row in `right` using + /// same-named equality keys. + /// + /// This has the same matching behavior as [`join_asof`](Self::join_asof), + /// but accepts columns that appear under the same name on both inputs. + /// + /// # Example + /// ``` + /// # use datafusion::arrow::array::record_batch; + /// # use datafusion::error::Result; + /// # use datafusion::prelude::*; + /// # #[tokio::main] + /// # async fn main() -> Result<()> { + /// # let ctx = SessionContext::new(); + /// # let trades = ctx.read_batch(record_batch!( + /// # ("symbol", Utf8, ["A"]), + /// # ("ts", Int64, [4]) + /// # )?)?.alias("trades")?; + /// # let prices = ctx.read_batch(record_batch!( + /// # ("symbol", Utf8, ["A"]), + /// # ("ts", Int64, [2]), + /// # ("price", Int32, [20]) + /// # )?)?.alias("prices")?; + /// // Same-named equality keys can be specified once. + /// let joined = trades.join_asof_using( + /// prices, + /// vec![Column::from_name("symbol")], + /// col("trades.ts").gt_eq(col("prices.ts")), + /// )?; + /// # let _ = joined; + /// # Ok(()) + /// # } + /// ``` + pub fn join_asof_using( + self, + right: DataFrame, + using_keys: Vec, + match_condition: Expr, + ) -> Result { + let plan = LogicalPlanBuilder::from(self.plan) + .asof_join_using(right.plan, using_keys, match_condition)? + .build()?; + Ok(DataFrame { + session_state: self.session_state, + plan, + projection_requires_validation: true, + }) + } + /// Repartition a DataFrame based on a logical partitioning scheme. /// /// # Example diff --git a/datafusion/core/tests/dataframe/mod.rs b/datafusion/core/tests/dataframe/mod.rs index 4e32ea31169ab..ec2c23b531a19 100644 --- a/datafusion/core/tests/dataframe/mod.rs +++ b/datafusion/core/tests/dataframe/mod.rs @@ -1521,6 +1521,75 @@ async fn join() -> Result<()> { Ok(()) } +#[tokio::test] +async fn join_asof() -> Result<()> { + let ctx = SessionContext::new(); + let left = ctx + .read_batch(record_batch!( + ("symbol", Utf8, ["A", "A", "B"]), + ("ts", Int64, [1, 4, 2]), + ("trade_id", Int32, [1, 2, 3]) + )?)? + .alias("trades")?; + let right = ctx + .read_batch(record_batch!( + ("symbol", Utf8, ["A", "A", "B"]), + ("ts", Int64, [2, 4, 1]), + ("price", Int32, [20, 40, 101]) + )?)? + .alias("prices")?; + + let results = left + .clone() + .join_asof( + right.clone(), + Some(col("trades.symbol").eq(col("prices.symbol"))), + col("trades.ts").gt_eq(col("prices.ts")), + )? + .select(vec![col("trade_id"), col("price")])? + .sort(vec![col("trade_id").sort(true, true)])? + .collect() + .await?; + + assert_batches_eq!( + [ + "+----------+-------+", + "| trade_id | price |", + "+----------+-------+", + "| 1 | |", + "| 2 | 40 |", + "| 3 | 101 |", + "+----------+-------+", + ], + &results + ); + + let results = left + .join_asof_using( + right, + vec![datafusion_common::Column::from_name("symbol")], + col("trades.ts").gt_eq(col("prices.ts")), + )? + .select(vec![col("symbol"), col("trade_id"), col("price")])? + .sort(vec![col("trade_id").sort(true, true)])? + .collect() + .await?; + + assert_batches_eq!( + [ + "+--------+----------+-------+", + "| symbol | trade_id | price |", + "+--------+----------+-------+", + "| A | 1 | |", + "| A | 2 | 40 |", + "| B | 3 | 101 |", + "+--------+----------+-------+", + ], + &results + ); + Ok(()) +} + #[tokio::test] async fn join_coercion_unnamed() -> Result<()> { let ctx = SessionContext::new();