use std::cmp::Ordering;
use std::collections::HashMap;
use std::fmt::{self, Debug, Display, Formatter};
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use arrow::datatypes::{DataType, Field, Schema};
use datafusion_common::file_options::file_type::FileType;
use datafusion_common::{DFSchemaRef, Result, TableReference, internal_err};
use crate::{Expr, LogicalPlan, TableSource};
#[derive(Clone)]
pub struct CopyTo {
pub input: Arc<LogicalPlan>,
pub output_url: String,
pub partition_by: Vec<String>,
pub file_type: Arc<dyn FileType>,
pub options: HashMap<String, String>,
pub output_schema: DFSchemaRef,
}
impl Debug for CopyTo {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("CopyTo")
.field("input", &self.input)
.field("output_url", &self.output_url)
.field("partition_by", &self.partition_by)
.field("file_type", &"...")
.field("options", &self.options)
.field("output_schema", &self.output_schema)
.finish_non_exhaustive()
}
}
impl PartialEq for CopyTo {
fn eq(&self, other: &Self) -> bool {
self.input == other.input && self.output_url == other.output_url
}
}
impl Eq for CopyTo {}
impl PartialOrd for CopyTo {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
match self.input.partial_cmp(&other.input) {
Some(Ordering::Equal) => match self.output_url.partial_cmp(&other.output_url)
{
Some(Ordering::Equal) => {
self.partition_by.partial_cmp(&other.partition_by)
}
cmp => cmp,
},
cmp => cmp,
}
.filter(|cmp| *cmp != Ordering::Equal || self == other)
}
}
impl Hash for CopyTo {
fn hash<H: Hasher>(&self, state: &mut H) {
self.input.hash(state);
self.output_url.hash(state);
}
}
impl CopyTo {
pub fn new(
input: Arc<LogicalPlan>,
output_url: String,
partition_by: Vec<String>,
file_type: Arc<dyn FileType>,
options: HashMap<String, String>,
) -> Self {
Self {
input,
output_url,
partition_by,
file_type,
options,
output_schema: make_count_schema(),
}
}
}
#[derive(Clone)]
pub struct DmlStatement {
pub table_name: TableReference,
pub target: Arc<dyn TableSource>,
pub op: WriteOp,
pub input: Arc<LogicalPlan>,
pub output_schema: DFSchemaRef,
}
impl Eq for DmlStatement {}
impl Hash for DmlStatement {
fn hash<H: Hasher>(&self, state: &mut H) {
self.table_name.hash(state);
self.target.schema().hash(state);
self.op.hash(state);
self.input.hash(state);
self.output_schema.hash(state);
}
}
impl PartialEq for DmlStatement {
fn eq(&self, other: &Self) -> bool {
self.table_name == other.table_name
&& self.target.schema() == other.target.schema()
&& self.op == other.op
&& self.input == other.input
&& self.output_schema == other.output_schema
}
}
impl Debug for DmlStatement {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("DmlStatement")
.field("table_name", &self.table_name)
.field("target", &"...")
.field("target_schema", &self.target.schema())
.field("op", &self.op)
.field("input", &self.input)
.field("output_schema", &self.output_schema)
.finish()
}
}
impl DmlStatement {
pub fn new(
table_name: TableReference,
target: Arc<dyn TableSource>,
op: WriteOp,
input: Arc<LogicalPlan>,
) -> Self {
Self {
table_name,
target,
op,
input,
output_schema: make_count_schema(),
}
}
pub fn name(&self) -> &str {
self.op.name()
}
}
impl PartialOrd for DmlStatement {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
match self.table_name.partial_cmp(&other.table_name) {
Some(Ordering::Equal) => match self.op.partial_cmp(&other.op) {
Some(Ordering::Equal) => self.input.partial_cmp(&other.input),
cmp => cmp,
},
cmp => cmp,
}
.filter(|cmp| *cmp != Ordering::Equal || self == other)
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Hash)]
#[non_exhaustive]
pub enum WriteOp {
Insert(InsertOp),
Delete,
Update,
Ctas,
Truncate,
MergeInto(Box<MergeIntoOp>),
}
impl WriteOp {
pub fn name(&self) -> &str {
match self {
WriteOp::Insert(insert) => insert.name(),
WriteOp::Delete => "Delete",
WriteOp::Update => "Update",
WriteOp::Ctas => "Ctas",
WriteOp::Truncate => "Truncate",
WriteOp::MergeInto(_) => "MergeInto",
}
}
}
impl Display for WriteOp {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.name())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Hash)]
pub enum InsertOp {
Append,
Overwrite,
Replace,
}
impl InsertOp {
pub fn name(&self) -> &str {
match self {
InsertOp::Append => "Insert Into",
InsertOp::Overwrite => "Insert Overwrite",
InsertOp::Replace => "Replace Into",
}
}
}
impl Display for InsertOp {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.name())
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Hash)]
pub struct MergeIntoOp {
pub on: Expr,
pub clauses: Vec<MergeIntoClause>,
}
impl MergeIntoOp {
fn expr_count(&self) -> usize {
1 + self
.clauses
.iter()
.map(|c| {
c.predicate.is_some() as usize
+ match &c.action {
MergeIntoAction::Update(a) => a.len(),
MergeIntoAction::Insert { values, .. } => values.len(),
MergeIntoAction::Delete => 0,
}
})
.sum::<usize>()
}
pub fn exprs(&self) -> Vec<&Expr> {
let mut out = Vec::with_capacity(self.expr_count());
out.push(&self.on);
for clause in &self.clauses {
if let Some(predicate) = &clause.predicate {
out.push(predicate);
}
match &clause.action {
MergeIntoAction::Update(assignments) => {
out.extend(assignments.iter().map(|(_, value)| value));
}
MergeIntoAction::Insert { values, .. } => {
out.extend(values.iter());
}
MergeIntoAction::Delete => {}
}
}
out
}
pub fn with_new_exprs(&self, exprs: Vec<Expr>) -> Result<Self> {
let expected = self.expr_count();
if exprs.len() != expected {
return internal_err!(
"MergeIntoOp::with_new_exprs expected {expected} expressions, got {}",
exprs.len()
);
}
let mut iter = exprs.into_iter();
let on = iter.next().expect("non-empty by length check");
let clauses = self
.clauses
.iter()
.map(|clause| {
let predicate = clause
.predicate
.is_some()
.then(|| iter.next().expect("non-empty by length check"));
let action = match &clause.action {
MergeIntoAction::Update(assignments) => {
let assignments = assignments
.iter()
.map(|(name, _)| {
(
name.clone(),
iter.next().expect("non-empty by length check"),
)
})
.collect();
MergeIntoAction::Update(assignments)
}
MergeIntoAction::Insert { columns, values } => {
let values = values
.iter()
.map(|_| iter.next().expect("non-empty by length check"))
.collect();
MergeIntoAction::Insert {
columns: columns.clone(),
values,
}
}
MergeIntoAction::Delete => MergeIntoAction::Delete,
};
MergeIntoClause {
kind: clause.kind,
predicate,
action,
}
})
.collect();
Ok(Self { on, clauses })
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Hash)]
pub struct MergeIntoClause {
pub kind: MergeIntoClauseKind,
pub predicate: Option<Expr>,
pub action: MergeIntoAction,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Hash)]
pub enum MergeIntoClauseKind {
Matched,
NotMatched,
NotMatchedByTarget,
NotMatchedBySource,
}
impl MergeIntoClauseKind {
pub fn is_not_matched_by_target(&self) -> bool {
matches!(self, Self::NotMatched | Self::NotMatchedByTarget)
}
pub fn canonical(self) -> Self {
match self {
Self::NotMatched => Self::NotMatchedByTarget,
other => other,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Hash)]
pub enum MergeIntoAction {
Update(Vec<(String, Expr)>),
Insert {
columns: Vec<String>,
values: Vec<Expr>,
},
Delete,
}
fn make_count_schema() -> DFSchemaRef {
Arc::new(
Schema::new(vec![Field::new("count", DataType::UInt64, false)])
.try_into()
.unwrap(),
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{col, lit};
#[test]
fn write_op_merge_into_name_and_display() {
let op = WriteOp::MergeInto(Box::new(MergeIntoOp {
on: col("id").eq(col("source_id")),
clauses: vec![MergeIntoClause {
kind: MergeIntoClauseKind::Matched,
predicate: Some(col("qty").gt(lit(0_i64))),
action: MergeIntoAction::Update(vec![(
"qty".to_string(),
col("source_qty"),
)]),
}],
}));
assert_eq!(op.name(), "MergeInto");
assert_eq!(format!("{op}"), "MergeInto");
}
#[test]
fn merge_into_clause_kind_is_not_matched_by_target() {
assert!(!MergeIntoClauseKind::Matched.is_not_matched_by_target());
assert!(MergeIntoClauseKind::NotMatched.is_not_matched_by_target());
assert!(MergeIntoClauseKind::NotMatchedByTarget.is_not_matched_by_target());
assert!(!MergeIntoClauseKind::NotMatchedBySource.is_not_matched_by_target());
}
#[test]
fn merge_into_clause_kind_canonical_collapses_not_matched() {
assert_eq!(
MergeIntoClauseKind::NotMatched.canonical(),
MergeIntoClauseKind::NotMatchedByTarget
);
assert_eq!(
MergeIntoClauseKind::NotMatchedByTarget.canonical(),
MergeIntoClauseKind::NotMatchedByTarget
);
assert_eq!(
MergeIntoClauseKind::Matched.canonical(),
MergeIntoClauseKind::Matched
);
assert_eq!(
MergeIntoClauseKind::NotMatchedBySource.canonical(),
MergeIntoClauseKind::NotMatchedBySource
);
}
#[test]
fn merge_into_op_exprs_round_trip() {
let op = MergeIntoOp {
on: col("id").eq(col("source_id")),
clauses: vec![
MergeIntoClause {
kind: MergeIntoClauseKind::Matched,
predicate: Some(col("qty").gt(lit(0_i64))),
action: MergeIntoAction::Update(vec![
("qty".to_string(), col("source_qty")),
("price".to_string(), col("source_price")),
]),
},
MergeIntoClause {
kind: MergeIntoClauseKind::NotMatched,
predicate: None,
action: MergeIntoAction::Insert {
columns: vec!["id".to_string(), "qty".to_string()],
values: vec![col("source_id"), col("source_qty")],
},
},
MergeIntoClause {
kind: MergeIntoClauseKind::NotMatchedBySource,
predicate: Some(col("active").eq(lit(true))),
action: MergeIntoAction::Delete,
},
],
};
let exprs = op.exprs();
assert_eq!(exprs.len(), 7);
let owned: Vec<Expr> = exprs.into_iter().cloned().collect();
let rebuilt = op.with_new_exprs(owned).unwrap();
assert_eq!(op, rebuilt);
}
#[test]
fn merge_into_op_with_new_exprs_length_mismatch() {
let op = MergeIntoOp {
on: col("id").eq(col("source_id")),
clauses: vec![],
};
let err = op.with_new_exprs(vec![]).unwrap_err();
assert!(
err.to_string().contains("expected 1 expressions, got 0"),
"unexpected error: {err}"
);
}
}