Skip to main content

spark_connect/
expression.rs

1//! Expression tree mirroring PySpark's `pyspark.sql.connect.expressions`.
2//!
3//! Defines the Expression hierarchy and conversion functions to Spark Connect protobufs.
4
5use std::sync::atomic::{AtomicU32, Ordering};
6
7use spark_connect_proto as proto;
8
9use crate::types::DataType;
10use crate::udf::CommonInlineUserDefinedFunctionExpression;
11
12/// Thread-local counter for generating unique lambda variable names.
13static LAMBDA_VAR_COUNTER: AtomicU32 = AtomicU32::new(0);
14
15/// Get the next unique lambda variable suffix.
16pub fn next_lambda_var_index() -> u32 {
17    LAMBDA_VAR_COUNTER.fetch_add(1, Ordering::SeqCst)
18}
19
20/// The base Expression type, mirroring `pyspark.sql.connect.expressions.Expression`.
21///
22/// All expressions are represented as variants of this enum. Each variant
23/// carries the data needed to construct the corresponding proto expression.
24#[derive(Debug, Clone, PartialEq)]
25pub enum Expression {
26    /// `pyspark.sql.connect.expressions.LiteralExpression` - a constant value.
27    Literal(LiteralExpression),
28    /// `pyspark.sql.connect.expressions.ColumnReference` - a reference to a column.
29    ColumnReference(ColumnReference),
30    /// `pyspark.sql.connect.expressions.UnresolvedFunction` - a function call.
31    UnresolvedFunction(UnresolvedFunction),
32    /// `pyspark.sql.connect.expressions.UnresolvedStar` - a star expression (`*`).
33    /// Carries an optional target ending in `.*` (e.g. `df.*` from `col("df.*")`).
34    UnresolvedStar(Option<String>),
35    /// `pyspark.sql.connect.expressions.ColumnAlias` / `Alias` - an aliased expression.
36    Alias(Box<Alias>),
37    /// `pyspark.sql.connect.expressions.CastExpression` - a cast expression.
38    Cast(Box<Cast>),
39    /// `spark.connect.Expression.DirectShufflePartitionID` - wraps a child expression
40    /// that evaluates to a partition id (used by `DataFrame.repartitionById`).
41    DirectShufflePartitionId(Box<Expression>),
42    /// `pyspark.sql.connect.expressions.UnresolvedRegex` - a regex column reference.
43    UnresolvedRegex(String),
44    /// `pyspark.sql.connect.expressions.SortOrder` - a sort order expression.
45    SortOrder(Box<SortOrder>),
46    /// `pyspark.sql.connect.expressions.CaseWhen` - a CASE WHEN expression.
47    CaseWhen(Box<CaseWhen>),
48    /// `pyspark.sql.connect.expressions.UnresolvedExtractValue` - `col[k]` / `getField`.
49    UnresolvedExtractValue(Box<ExtractValue>),
50    /// `pyspark.sql.connect.expressions.UpdateFields` - `withField` / `dropFields`.
51    UpdateFields(Box<UpdateFieldsExpr>),
52    /// `pyspark.sql.connect.expressions.SQLExpression` - a raw SQL expression.
53    SQLExpression(String),
54    /// `pyspark.sql.connect.expressions.CallFunction` - a direct function call.
55    CallFunction(Box<CallFunctionWrapper>),
56    /// `pyspark.sql.connect.expressions.WindowExpression` - a window function call.
57    WindowExpression(Box<WindowExpressionWrapper>),
58    /// `pyspark.sql.connect.expressions.LambdaFunction` - a lambda function.
59    LambdaFunction(Box<LambdaFunction>),
60    /// `pyspark.sql.connect.expressions.UnresolvedNamedLambdaVariable` - a lambda variable.
61    UnresolvedNamedLambdaVariable(UnresolvedNamedLambdaVariable),
62    /// `pyspark.sql.connect.expressions.CommonInlineUserDefinedFunction` - a UDF expression.
63    CommonInlineUserDefinedFunction(Box<CommonInlineUserDefinedFunctionExpression>),
64}
65
66impl Expression {
67    /// Render this expression to a human-readable string, mirroring the
68    /// `__repr__` of `pyspark.sql.connect.expressions.*`.
69    ///
70    /// This is what `Column.__repr__` wraps as `Column<'...'>`. pandas-on-Spark's
71    /// `spark_column_equals` compares these strings (after stripping backticks) to
72    /// decide column equality, so the rendering must be deterministic and mirror
73    /// PySpark's format for the common operators.
74    pub fn render(&self) -> String {
75        match self {
76            Expression::Literal(lit) => lit.render(),
77            Expression::ColumnReference(col_ref) => col_ref.name.clone(),
78            Expression::UnresolvedFunction(func) => func.render(),
79            Expression::UnresolvedStar(target) => match target {
80                Some(t) => t.clone(),
81                None => "*".to_string(),
82            },
83            Expression::Alias(alias) => {
84                let name = if alias.names.len() == 1 {
85                    alias.names[0].clone()
86                } else {
87                    format!("({})", alias.names.join(", "))
88                };
89                format!("{} AS {}", alias.child.render(), name)
90            }
91            Expression::Cast(cast) => {
92                let type_str = match &cast.target {
93                    CastTarget::Type(dt) => dt.simple_string(),
94                    CastTarget::TypeStr(s) => s.clone(),
95                };
96                format!("CAST({} AS {})", cast.child.render(), type_str)
97            }
98            Expression::DirectShufflePartitionId(child) => {
99                format!("DIRECT_SHUFFLE_PARTITION_ID({})", child.render())
100            }
101            Expression::UnresolvedRegex(col_name) => col_name.clone(),
102            Expression::SortOrder(sort) => sort.render(),
103            Expression::CaseWhen(case_when) => case_when.render(),
104            Expression::UnresolvedExtractValue(ev) => {
105                format!("{}[{}]", ev.child.render(), ev.extraction.render())
106            }
107            Expression::UpdateFields(uf) => match &uf.value_expression {
108                Some(v) => format!(
109                    "update_field({}, {}, {})",
110                    uf.struct_expression.render(),
111                    uf.field_name,
112                    v.render()
113                ),
114                None => format!(
115                    "drop_field({}, {})",
116                    uf.struct_expression.render(),
117                    uf.field_name
118                ),
119            },
120            Expression::SQLExpression(sql) => sql.clone(),
121            Expression::CallFunction(_) => format!("{self:?}"),
122            Expression::WindowExpression(_) => format!("{self:?}"),
123            Expression::LambdaFunction(lf) => format!("{lf:?}"),
124            Expression::UnresolvedNamedLambdaVariable(var) => format!("{var:?}"),
125            Expression::CommonInlineUserDefinedFunction(_) => format!("{self:?}"),
126        }
127    }
128
129    /// Converts the expression to a Spark Connect protobuf expression.
130    pub fn to_proto(&self) -> proto::Expression {
131        match self {
132            Expression::Literal(lit) => lit.to_proto(),
133            Expression::ColumnReference(col_ref) => col_ref.to_proto(),
134            Expression::UnresolvedFunction(func) => func.to_proto(),
135            Expression::UnresolvedStar(target) => {
136                let mut expr = proto::Expression::default();
137                expr.expr_type = Some(proto::expression::ExprType::UnresolvedStar(
138                    proto::expression::UnresolvedStar {
139                        unparsed_target: target.clone(),
140                        plan_id: None,
141                    },
142                ));
143                expr
144            }
145            Expression::Alias(alias) => alias.to_proto(),
146            Expression::Cast(cast) => cast.to_proto(),
147            Expression::DirectShufflePartitionId(child) => {
148                let mut expr = proto::Expression::default();
149                expr.expr_type = Some(proto::expression::ExprType::DirectShufflePartitionId(
150                    Box::new(proto::expression::DirectShufflePartitionId {
151                        child: Some(Box::new(child.to_proto())),
152                    }),
153                ));
154                expr
155            }
156            Expression::UnresolvedRegex(col_name) => {
157                let mut expr = proto::Expression::default();
158                expr.expr_type = Some(proto::expression::ExprType::UnresolvedRegex(
159                    proto::expression::UnresolvedRegex {
160                        col_name: col_name.clone(),
161                        plan_id: None,
162                    },
163                ));
164                expr
165            }
166            Expression::SortOrder(sort) => sort.to_proto(),
167            Expression::CaseWhen(case_when) => case_when.to_proto(),
168            Expression::UnresolvedExtractValue(ev) => ev.to_proto(),
169            Expression::UpdateFields(uf) => uf.to_proto(),
170            Expression::SQLExpression(sql) => {
171                let mut expr = proto::Expression::default();
172                expr.expr_type = Some(proto::expression::ExprType::ExpressionString(
173                    proto::expression::ExpressionString {
174                        expression: sql.clone(),
175                    },
176                ));
177                expr
178            }
179            Expression::CallFunction(cf) => cf.to_proto(),
180            Expression::WindowExpression(we) => we.to_proto(),
181            Expression::LambdaFunction(lf) => lf.to_proto(),
182            Expression::UnresolvedNamedLambdaVariable(var) => var.to_proto(),
183            Expression::CommonInlineUserDefinedFunction(udf) => {
184                let mut expr = proto::Expression::default();
185                expr.expr_type = Some(
186                    proto::expression::ExprType::CommonInlineUserDefinedFunction(udf.to_proto()),
187                );
188                expr
189            }
190        }
191    }
192}
193
194/// `pyspark.sql.connect.expressions.LiteralExpression`
195#[derive(Debug, Clone, PartialEq)]
196pub enum LiteralExpression {
197    Null(DataType),
198    Boolean(bool),
199    Byte(i32),
200    Short(i32),
201    Integer(i32),
202    Long(i64),
203    Float(f32),
204    Double(f64),
205    Decimal {
206        value: String,
207        precision: i32,
208        scale: i32,
209    },
210    String(String),
211    Binary(Vec<u8>),
212    Date(i32),
213    Timestamp(i64),
214    TimestampNtz(i64),
215    Time {
216        nano: i64,
217        precision: i32,
218    },
219    Array {
220        element_type: Box<DataType>,
221        elements: Vec<LiteralExpression>,
222    },
223}
224
225impl LiteralExpression {
226    pub fn to_proto(&self) -> proto::Expression {
227        let mut expr = proto::Expression::default();
228        let literal_type = match self {
229            LiteralExpression::Null(data_type) => {
230                proto::expression::literal::LiteralType::Null(data_type.to_proto())
231            }
232            LiteralExpression::Boolean(b) => proto::expression::literal::LiteralType::Boolean(*b),
233            LiteralExpression::Byte(v) => proto::expression::literal::LiteralType::Byte(*v),
234            LiteralExpression::Short(v) => proto::expression::literal::LiteralType::Short(*v),
235            LiteralExpression::Integer(v) => proto::expression::literal::LiteralType::Integer(*v),
236            LiteralExpression::Long(v) => proto::expression::literal::LiteralType::Long(*v),
237            LiteralExpression::Float(v) => proto::expression::literal::LiteralType::Float(*v),
238            LiteralExpression::Double(v) => proto::expression::literal::LiteralType::Double(*v),
239            LiteralExpression::Decimal {
240                value,
241                precision,
242                scale,
243            } => {
244                let mut decimal = proto::expression::literal::Decimal::default();
245                decimal.value = value.clone();
246                decimal.precision = Some(*precision);
247                decimal.scale = Some(*scale);
248                proto::expression::literal::LiteralType::Decimal(decimal)
249            }
250            LiteralExpression::String(v) => {
251                proto::expression::literal::LiteralType::String(v.clone())
252            }
253            LiteralExpression::Binary(v) => {
254                proto::expression::literal::LiteralType::Binary(v.clone().into())
255            }
256            LiteralExpression::Date(v) => proto::expression::literal::LiteralType::Date(*v),
257            LiteralExpression::Timestamp(v) => {
258                proto::expression::literal::LiteralType::Timestamp(*v)
259            }
260            LiteralExpression::TimestampNtz(v) => {
261                proto::expression::literal::LiteralType::TimestampNtz(*v)
262            }
263            LiteralExpression::Time { nano, precision } => {
264                let mut time = proto::expression::literal::Time::default();
265                time.nano = *nano;
266                time.precision = Some(*precision);
267                proto::expression::literal::LiteralType::Time(time)
268            }
269            LiteralExpression::Array {
270                element_type: _,
271                elements,
272            } => {
273                let mut array = proto::expression::literal::Array::default();
274                for elem in elements {
275                    let elem_proto = elem.to_proto();
276                    if let Some(proto::expression::ExprType::Literal(lit)) = elem_proto.expr_type {
277                        array.elements.push(lit);
278                    }
279                }
280                proto::expression::literal::LiteralType::Array(array)
281            }
282        };
283
284        let mut literal = proto::expression::Literal::default();
285        literal.literal_type = Some(literal_type);
286        expr.expr_type = Some(proto::expression::ExprType::Literal(literal));
287        expr
288    }
289
290    /// Render the literal value, mirroring `LiteralExpression.__repr__`.
291    pub fn render(&self) -> String {
292        match self {
293            LiteralExpression::Null(_) => "NULL".to_string(),
294            LiteralExpression::Boolean(b) => {
295                if *b {
296                    "true".to_string()
297                } else {
298                    "false".to_string()
299                }
300            }
301            LiteralExpression::Byte(v)
302            | LiteralExpression::Short(v)
303            | LiteralExpression::Integer(v) => v.to_string(),
304            LiteralExpression::Long(v) => v.to_string(),
305            LiteralExpression::Float(v) => v.to_string(),
306            LiteralExpression::Double(v) => v.to_string(),
307            LiteralExpression::Decimal { value, .. } => value.clone(),
308            LiteralExpression::String(v) => v.clone(),
309            LiteralExpression::Binary(v) => format!("{v:?}"),
310            LiteralExpression::Date(v) => v.to_string(),
311            LiteralExpression::Timestamp(v) | LiteralExpression::TimestampNtz(v) => v.to_string(),
312            LiteralExpression::Time { nano, .. } => nano.to_string(),
313            LiteralExpression::Array { elements, .. } => {
314                let inner: Vec<String> = elements.iter().map(|e| e.render()).collect();
315                format!("[{}]", inner.join(", "))
316            }
317        }
318    }
319
320    /// Construct a null literal of a specific type.
321    pub fn null(data_type: DataType) -> Self {
322        LiteralExpression::Null(data_type)
323    }
324
325    /// Construct an integer literal.
326    pub fn int(value: i32) -> Self {
327        LiteralExpression::Integer(value)
328    }
329
330    /// Construct a long literal.
331    pub fn long(value: i64) -> Self {
332        LiteralExpression::Long(value)
333    }
334
335    /// Construct a double literal.
336    pub fn double(value: f64) -> Self {
337        LiteralExpression::Double(value)
338    }
339
340    /// Construct a string literal.
341    pub fn string(value: impl Into<String>) -> Self {
342        LiteralExpression::String(value.into())
343    }
344
345    /// Construct a boolean literal.
346    pub fn boolean(value: bool) -> Self {
347        LiteralExpression::Boolean(value)
348    }
349
350    /// Construct a binary literal (byte array).
351    pub fn binary(value: Vec<u8>) -> Self {
352        LiteralExpression::Binary(value)
353    }
354}
355
356/// `pyspark.sql.connect.expressions.ColumnReference`
357/// Represents a reference to a column by name (unresolved attribute).
358#[derive(Debug, Clone, PartialEq, Eq)]
359pub struct ColumnReference {
360    /// The column name / unparsed identifier.
361    pub name: String,
362    /// Plan id this attribute is bound to. `None` for a free `col(...)`; set only
363    /// when the column is resolved against a specific DataFrame (mirrors
364    /// `ColumnReference._plan_id`).
365    pub plan_id: Option<i64>,
366    /// Whether this references a metadata column (mirrors
367    /// `DataFrame.metadataColumn`). `false` for a normal `col(...)`.
368    pub is_metadata_column: bool,
369}
370
371impl ColumnReference {
372    pub fn new(name: impl Into<String>) -> Self {
373        Self {
374            name: name.into(),
375            plan_id: None,
376            is_metadata_column: false,
377        }
378    }
379
380    /// Bind this attribute to a DataFrame plan id.
381    pub fn with_plan_id(mut self, plan_id: i64) -> Self {
382        self.plan_id = Some(plan_id);
383        self
384    }
385
386    /// Mark this reference as a metadata column (mirrors `DataFrame.metadataColumn`).
387    pub fn metadata(mut self) -> Self {
388        self.is_metadata_column = true;
389        self
390    }
391
392    pub fn to_proto(&self) -> proto::Expression {
393        let mut expr = proto::Expression::default();
394        expr.expr_type = Some(proto::expression::ExprType::UnresolvedAttribute(
395            proto::expression::UnresolvedAttribute {
396                unparsed_identifier: self.name.clone(),
397                plan_id: self.plan_id,
398                is_metadata_column: Some(self.is_metadata_column),
399            },
400        ));
401        expr
402    }
403}
404
405/// `pyspark.sql.connect.expressions.UnresolvedExtractValue` - `col[key]` / `getField`.
406#[derive(Debug, Clone, PartialEq)]
407pub struct ExtractValue {
408    pub child: Expression,
409    pub extraction: Expression,
410}
411
412impl ExtractValue {
413    pub fn new(child: Expression, extraction: Expression) -> Self {
414        Self { child, extraction }
415    }
416
417    pub fn to_proto(&self) -> proto::Expression {
418        let mut expr = proto::Expression::default();
419        expr.expr_type = Some(proto::expression::ExprType::UnresolvedExtractValue(
420            Box::new(proto::expression::UnresolvedExtractValue {
421                child: Some(Box::new(self.child.to_proto())),
422                extraction: Some(Box::new(self.extraction.to_proto())),
423            }),
424        ));
425        expr
426    }
427}
428
429/// `pyspark.sql.connect.expressions.Expression.UpdateFields` - add/replace
430/// (`withField`) or drop (`dropFields`, `value` = None) a struct field.
431#[derive(Debug, Clone, PartialEq)]
432pub struct UpdateFieldsExpr {
433    pub struct_expression: Expression,
434    pub field_name: String,
435    pub value_expression: Option<Expression>,
436}
437
438impl UpdateFieldsExpr {
439    pub fn new(
440        struct_expression: Expression,
441        field_name: impl Into<String>,
442        value_expression: Option<Expression>,
443    ) -> Self {
444        Self {
445            struct_expression,
446            field_name: field_name.into(),
447            value_expression,
448        }
449    }
450
451    pub fn to_proto(&self) -> proto::Expression {
452        let mut expr = proto::Expression::default();
453        expr.expr_type = Some(proto::expression::ExprType::UpdateFields(Box::new(
454            proto::expression::UpdateFields {
455                struct_expression: Some(Box::new(self.struct_expression.to_proto())),
456                field_name: self.field_name.clone(),
457                value_expression: self
458                    .value_expression
459                    .as_ref()
460                    .map(|e| Box::new(e.to_proto())),
461            },
462        )));
463        expr
464    }
465}
466
467/// `pyspark.sql.connect.expressions.UnresolvedFunction`
468/// Represents a function call with a name and arguments.
469#[derive(Debug, Clone, PartialEq)]
470pub struct UnresolvedFunction {
471    pub name: String,
472    pub args: Vec<Expression>,
473    pub is_distinct: bool,
474}
475
476impl UnresolvedFunction {
477    pub fn new(name: impl Into<String>, args: Vec<Expression>) -> Self {
478        Self {
479            name: name.into(),
480            args,
481            is_distinct: false,
482        }
483    }
484
485    pub fn new_distinct(name: impl Into<String>, args: Vec<Expression>) -> Self {
486        Self {
487            name: name.into(),
488            args,
489            is_distinct: true,
490        }
491    }
492
493    /// Render this function call, mirroring `UnresolvedFunction.__repr__`:
494    /// binary operators render infix as `(a op b)`, the unary `not`/negate render
495    /// prefixed, everything else as `name(arg, arg, ...)`.
496    pub fn render(&self) -> String {
497        const INFIX_OPS: &[&str] = &[
498            "+", "-", "*", "/", "%", "==", "!=", "<", "<=", ">", ">=", "and", "or", "&", "|", "^",
499            "<=>",
500        ];
501        if self.args.len() == 2 && INFIX_OPS.contains(&self.name.as_str()) {
502            return format!(
503                "({} {} {})",
504                self.args[0].render(),
505                self.name,
506                self.args[1].render()
507            );
508        }
509        if self.args.len() == 1 {
510            match self.name.as_str() {
511                "not" => return format!("(NOT {})", self.args[0].render()),
512                "negative" | "negate" => return format!("(- {})", self.args[0].render()),
513                _ => {}
514            }
515        }
516        let inner: Vec<String> = self.args.iter().map(|a| a.render()).collect();
517        format!("{}({})", self.name, inner.join(", "))
518    }
519
520    pub fn to_proto(&self) -> proto::Expression {
521        let mut expr = proto::Expression::default();
522        let mut func = proto::expression::UnresolvedFunction::default();
523        func.function_name = self.name.clone();
524        func.is_distinct = self.is_distinct;
525        for arg in &self.args {
526            func.arguments.push(arg.to_proto());
527        }
528        expr.expr_type = Some(proto::expression::ExprType::UnresolvedFunction(func));
529        expr
530    }
531}
532
533/// `pyspark.sql.connect.expressions.Alias` / `ColumnAlias`
534#[derive(Debug, Clone, PartialEq)]
535pub struct Alias {
536    pub child: Expression,
537    pub names: Vec<String>,
538    pub metadata: Option<String>,
539}
540
541impl Alias {
542    pub fn new(child: Expression, name: impl Into<String>) -> Self {
543        Self {
544            child,
545            names: vec![name.into()],
546            metadata: None,
547        }
548    }
549
550    pub fn with_metadata(mut self, metadata: String) -> Self {
551        self.metadata = Some(metadata);
552        self
553    }
554
555    pub fn to_proto(&self) -> proto::Expression {
556        let mut expr = proto::Expression::default();
557        let mut alias = proto::expression::Alias::default();
558        alias.expr = Some(Box::new(self.child.to_proto()));
559        alias.name = self.names.clone();
560        if let Some(meta) = &self.metadata {
561            alias.metadata = Some(meta.clone());
562        }
563        expr.expr_type = Some(proto::expression::ExprType::Alias(Box::new(alias)));
564        expr
565    }
566}
567
568/// `pyspark.sql.connect.expressions.CastExpression` / `Cast`
569#[derive(Debug, Clone, PartialEq)]
570pub struct Cast {
571    pub child: Expression,
572    pub target: CastTarget,
573    pub eval_mode: Option<CastEvalMode>,
574}
575
576/// The cast target: a structured `DataType` or a DDL type string. Mirrors the
577/// `cast_to_type` oneof (`type` vs `type_str`); `Column.cast("string")` uses the
578/// string form, `Column.cast(IntegerType())` the structured form.
579#[derive(Debug, Clone, PartialEq)]
580pub enum CastTarget {
581    Type(DataType),
582    TypeStr(String),
583}
584
585#[derive(Debug, Clone, Copy, PartialEq, Eq)]
586pub enum CastEvalMode {
587    Legacy,
588    Ansi,
589    Try,
590}
591
592impl Cast {
593    pub fn new(child: Expression, to_type: DataType) -> Self {
594        Self {
595            child,
596            target: CastTarget::Type(to_type),
597            eval_mode: None,
598        }
599    }
600
601    /// Cast to a DDL type string (mirrors `Column.cast("string")`).
602    pub fn new_str(child: Expression, type_str: impl Into<String>) -> Self {
603        Self {
604            child,
605            target: CastTarget::TypeStr(type_str.into()),
606            eval_mode: None,
607        }
608    }
609
610    pub fn with_eval_mode(mut self, mode: CastEvalMode) -> Self {
611        self.eval_mode = Some(mode);
612        self
613    }
614
615    pub fn to_proto(&self) -> proto::Expression {
616        let mut expr = proto::Expression::default();
617        let mut cast = proto::expression::Cast::default();
618        cast.expr = Some(Box::new(self.child.to_proto()));
619        cast.cast_to_type = Some(match &self.target {
620            CastTarget::Type(dt) => proto::expression::cast::CastToType::Type(dt.to_proto()),
621            CastTarget::TypeStr(s) => proto::expression::cast::CastToType::TypeStr(s.clone()),
622        });
623
624        if let Some(mode) = self.eval_mode {
625            cast.eval_mode = match mode {
626                CastEvalMode::Legacy => 1i32,
627                CastEvalMode::Ansi => 2i32,
628                CastEvalMode::Try => 3i32,
629            };
630        }
631
632        expr.expr_type = Some(proto::expression::ExprType::Cast(Box::new(cast)));
633        expr
634    }
635}
636
637/// `pyspark.sql.connect.expressions.SortOrder`
638#[derive(Debug, Clone, PartialEq)]
639pub struct SortOrder {
640    pub child: Expression,
641    pub ascending: bool,
642    pub null_ordering: NullOrdering,
643}
644
645#[derive(Debug, Clone, Copy, PartialEq, Eq)]
646pub enum NullOrdering {
647    First,
648    Last,
649}
650
651impl SortOrder {
652    /// Render this sort order, mirroring `SortOrder.__repr__`.
653    pub fn render(&self) -> String {
654        let dir = if self.ascending { "ASC" } else { "DESC" };
655        let nulls = match self.null_ordering {
656            NullOrdering::First => "NULLS FIRST",
657            NullOrdering::Last => "NULLS LAST",
658        };
659        format!("{} {} {}", self.child.render(), dir, nulls)
660    }
661
662    pub fn asc_nulls_first(child: Expression) -> Self {
663        Self {
664            child,
665            ascending: true,
666            null_ordering: NullOrdering::First,
667        }
668    }
669
670    pub fn asc_nulls_last(child: Expression) -> Self {
671        Self {
672            child,
673            ascending: true,
674            null_ordering: NullOrdering::Last,
675        }
676    }
677
678    pub fn desc_nulls_first(child: Expression) -> Self {
679        Self {
680            child,
681            ascending: false,
682            null_ordering: NullOrdering::First,
683        }
684    }
685
686    pub fn desc_nulls_last(child: Expression) -> Self {
687        Self {
688            child,
689            ascending: false,
690            null_ordering: NullOrdering::Last,
691        }
692    }
693
694    pub fn to_proto(&self) -> proto::Expression {
695        let mut expr = proto::Expression::default();
696        let mut sort = proto::expression::SortOrder::default();
697        sort.child = Some(Box::new(self.child.to_proto()));
698        sort.direction = if self.ascending { 1i32 } else { 2i32 };
699        sort.null_ordering = match self.null_ordering {
700            NullOrdering::First => 1i32,
701            NullOrdering::Last => 2i32,
702        };
703        expr.expr_type = Some(proto::expression::ExprType::SortOrder(Box::new(sort)));
704        expr
705    }
706}
707
708/// `pyspark.sql.connect.expressions.CaseWhen`
709#[derive(Debug, Clone, PartialEq)]
710pub struct CaseWhen {
711    pub branches: Vec<(Expression, Expression)>,
712    pub else_expr: Option<Box<Expression>>,
713}
714
715impl CaseWhen {
716    pub fn new(branches: Vec<(Expression, Expression)>) -> Self {
717        Self {
718            branches,
719            else_expr: None,
720        }
721    }
722
723    pub fn with_else(mut self, else_expr: Expression) -> Self {
724        self.else_expr = Some(Box::new(else_expr));
725        self
726    }
727
728    /// Render this CASE WHEN, mirroring `CaseWhen.__repr__`.
729    pub fn render(&self) -> String {
730        let mut parts = vec!["CASE".to_string()];
731        for (cond, value) in &self.branches {
732            parts.push(format!("WHEN {} THEN {}", cond.render(), value.render()));
733        }
734        if let Some(else_expr) = &self.else_expr {
735            parts.push(format!("ELSE {}", else_expr.render()));
736        }
737        parts.push("END".to_string());
738        parts.join(" ")
739    }
740
741    pub fn to_proto(&self) -> proto::Expression {
742        let mut args = Vec::new();
743        for (condition, value) in &self.branches {
744            args.push(condition.clone());
745            args.push(value.clone());
746        }
747        if let Some(else_expr) = &self.else_expr {
748            args.push((**else_expr).clone());
749        }
750        let func = UnresolvedFunction::new("when", args);
751        func.to_proto()
752    }
753}
754
755/// Wrapper for `spark.connect.CallFunction` protobuf message.
756///
757/// Mirrors `pyspark.sql.connect.expressions.CallFunction`: a named function call
758/// carrying its argument expressions.
759#[derive(Debug, Clone, PartialEq)]
760pub struct CallFunctionWrapper {
761    pub function_name: String,
762    pub arguments: Vec<Expression>,
763}
764
765impl CallFunctionWrapper {
766    /// Create a new CallFunctionWrapper with the given argument expressions.
767    pub fn new(function_name: impl Into<String>, arguments: Vec<Expression>) -> Self {
768        CallFunctionWrapper {
769            function_name: function_name.into(),
770            arguments,
771        }
772    }
773
774    /// Convert to protobuf.
775    pub fn to_proto(&self) -> proto::Expression {
776        let mut expr = proto::Expression::default();
777        expr.expr_type = Some(proto::expression::ExprType::CallFunction(
778            proto::CallFunction {
779                function_name: self.function_name.clone(),
780                arguments: self.arguments.iter().map(|a| a.to_proto()).collect(),
781            },
782        ));
783        expr
784    }
785}
786
787/// Wrapper for `spark.connect.Window` (window expression with OVER clause).
788#[derive(Debug, Clone, PartialEq)]
789pub struct WindowExpressionWrapper {
790    pub window_function: Expression,
791    pub partition_spec: Vec<Expression>,
792    pub order_spec: Vec<SortOrder>,
793    pub frame_spec: Option<(u32, FrameBoundary, FrameBoundary)>,
794}
795
796/// Frame boundary for window specification.
797#[derive(Debug, Clone, PartialEq)]
798pub enum FrameBoundary {
799    UnboundedPreceding,
800    Preceding(i64),
801    CurrentRow,
802    Following(i64),
803    UnboundedFollowing,
804}
805
806impl WindowExpressionWrapper {
807    /// Create a new WindowExpressionWrapper.
808    pub fn new(
809        window_function: Expression,
810        partition_spec: Vec<Expression>,
811        order_spec: Vec<SortOrder>,
812        frame_spec: Option<(u32, FrameBoundary, FrameBoundary)>,
813    ) -> Self {
814        Self {
815            window_function,
816            partition_spec,
817            order_spec,
818            frame_spec,
819        }
820    }
821
822    /// Convert to protobuf.
823    pub fn to_proto(&self) -> proto::Expression {
824        let mut expr = proto::Expression::default();
825        let mut window = proto::expression::Window::default();
826
827        // Set the window function
828        window.window_function = Some(Box::new(self.window_function.to_proto()));
829
830        // Set partition spec
831        window.partition_spec = self.partition_spec.iter().map(|e| e.to_proto()).collect();
832
833        // Set order spec
834        window.order_spec = self
835            .order_spec
836            .iter()
837            .map(|s| s.to_proto_sort_order())
838            .collect();
839
840        // Set frame spec if present
841        if let Some((frame_type, lower, upper)) = &self.frame_spec {
842            let mut frame = proto::expression::window::WindowFrame::default();
843            frame.frame_type = *frame_type as i32;
844            frame.lower = Some(Box::new(to_proto_frame_boundary(lower)));
845            frame.upper = Some(Box::new(to_proto_frame_boundary(upper)));
846            window.frame_spec = Some(Box::new(frame));
847        }
848
849        expr.expr_type = Some(proto::expression::ExprType::Window(Box::new(window)));
850        expr
851    }
852}
853
854/// Convert FrameBoundary to proto FrameBoundary.
855fn to_proto_frame_boundary(
856    boundary: &FrameBoundary,
857) -> proto::expression::window::window_frame::FrameBoundary {
858    let mut proto_boundary = proto::expression::window::window_frame::FrameBoundary::default();
859    proto_boundary.boundary = match boundary {
860        FrameBoundary::UnboundedPreceding => {
861            Some(proto::expression::window::window_frame::frame_boundary::Boundary::Unbounded(true))
862        }
863        FrameBoundary::UnboundedFollowing => {
864            Some(proto::expression::window::window_frame::frame_boundary::Boundary::Unbounded(true))
865        }
866        FrameBoundary::CurrentRow => Some(
867            proto::expression::window::window_frame::frame_boundary::Boundary::CurrentRow(true),
868        ),
869        // A ROWS-frame offset must be an IntegerType literal (Spark rejects bigint with
870        // SPECIFIED_WINDOW_FRAME_UNACCEPTED_TYPE), and preceding is a negative offset.
871        FrameBoundary::Preceding(n) => Some(
872            proto::expression::window::window_frame::frame_boundary::Boundary::Value(Box::new(
873                Expression::Literal(LiteralExpression::int(-(*n as i32))).to_proto(),
874            )),
875        ),
876        FrameBoundary::Following(n) => Some(
877            proto::expression::window::window_frame::frame_boundary::Boundary::Value(Box::new(
878                Expression::Literal(LiteralExpression::int(*n as i32)).to_proto(),
879            )),
880        ),
881    };
882    proto_boundary
883}
884
885impl SortOrder {
886    /// Convert to proto SortOrder (used by window).
887    pub fn to_proto_sort_order(&self) -> proto::expression::SortOrder {
888        let mut sort = proto::expression::SortOrder::default();
889        sort.child = Some(Box::new(self.child.to_proto()));
890        sort.direction = if self.ascending { 1i32 } else { 2i32 };
891        sort.null_ordering = match self.null_ordering {
892            NullOrdering::First => 1i32,
893            NullOrdering::Last => 2i32,
894        };
895        sort
896    }
897}
898
899/// `pyspark.sql.connect.expressions.LambdaFunction` - a lambda function with body and arguments.
900#[derive(Debug, Clone, PartialEq)]
901pub struct LambdaFunction {
902    pub function: Expression,
903    pub arguments: Vec<UnresolvedNamedLambdaVariable>,
904}
905
906impl LambdaFunction {
907    pub fn new(function: Expression, arguments: Vec<UnresolvedNamedLambdaVariable>) -> Self {
908        Self {
909            function,
910            arguments,
911        }
912    }
913
914    pub fn to_proto(&self) -> proto::Expression {
915        let mut expr = proto::Expression::default();
916        let mut lambda = proto::expression::LambdaFunction::default();
917        lambda.function = Some(Box::new(self.function.to_proto()));
918        for arg in &self.arguments {
919            lambda
920                .arguments
921                .push(proto::expression::UnresolvedNamedLambdaVariable {
922                    name_parts: vec![arg.name_parts.clone()],
923                });
924        }
925        expr.expr_type = Some(proto::expression::ExprType::LambdaFunction(Box::new(
926            lambda,
927        )));
928        expr
929    }
930}
931
932/// `pyspark.sql.connect.expressions.UnresolvedNamedLambdaVariable` - a lambda variable.
933#[derive(Debug, Clone, PartialEq, Eq)]
934pub struct UnresolvedNamedLambdaVariable {
935    pub name_parts: String,
936}
937
938impl UnresolvedNamedLambdaVariable {
939    pub fn new(name_parts: impl Into<String>) -> Self {
940        Self {
941            name_parts: name_parts.into(),
942        }
943    }
944
945    pub fn to_proto(&self) -> proto::Expression {
946        let mut expr = proto::Expression::default();
947        expr.expr_type = Some(proto::expression::ExprType::UnresolvedNamedLambdaVariable(
948            proto::expression::UnresolvedNamedLambdaVariable {
949                name_parts: vec![self.name_parts.clone()],
950            },
951        ));
952        expr
953    }
954}
955
956#[cfg(test)]
957mod tests {
958    use super::*;
959
960    fn col(name: &str) -> Expression {
961        Expression::ColumnReference(ColumnReference::new(name))
962    }
963
964    #[test]
965    fn test_render_matches_pyspark_connect_format() {
966        // Column reference and literal.
967        assert_eq!(col("x").render(), "x");
968        assert_eq!(
969            Expression::Literal(LiteralExpression::Integer(0)).render(),
970            "0"
971        );
972        assert_eq!(
973            Expression::Literal(LiteralExpression::null(DataType::Integer)).render(),
974            "NULL"
975        );
976
977        // Binary operators render infix, mirroring UnresolvedFunction.__repr__.
978        let add = Expression::UnresolvedFunction(UnresolvedFunction::new(
979            "+",
980            vec![col("x"), Expression::Literal(LiteralExpression::Integer(1))],
981        ));
982        assert_eq!(add.render(), "(x + 1)");
983
984        let eq =
985            Expression::UnresolvedFunction(UnresolvedFunction::new("==", vec![col("a"), col("b")]));
986        assert_eq!(eq.render(), "(a == b)");
987
988        // Unary not.
989        let neq = Expression::UnresolvedFunction(UnresolvedFunction::new("not", vec![eq.clone()]));
990        assert_eq!(neq.render(), "(NOT (a == b))");
991
992        // Non-operator function renders as name(args).
993        let f = Expression::UnresolvedFunction(UnresolvedFunction::new(
994            "coalesce",
995            vec![col("a"), col("b")],
996        ));
997        assert_eq!(f.render(), "coalesce(a, b)");
998
999        // Alias, cast, star.
1000        assert_eq!(
1001            Expression::Alias(Box::new(Alias::new(col("x"), "y"))).render(),
1002            "x AS y"
1003        );
1004        assert_eq!(
1005            Expression::Cast(Box::new(Cast {
1006                child: col("x"),
1007                target: CastTarget::TypeStr("int".to_string()),
1008                eval_mode: None,
1009            }))
1010            .render(),
1011            "CAST(x AS int)"
1012        );
1013        assert_eq!(Expression::UnresolvedStar(None).render(), "*");
1014
1015        // Equal expressions render identically; different ones differ
1016        // (the property spark_column_equals relies on).
1017        assert_eq!(add.render(), add.render());
1018        assert_ne!(add.render(), eq.render());
1019    }
1020
1021    #[test]
1022    fn test_literal_integer() {
1023        let lit = LiteralExpression::int(42);
1024        let proto = lit.to_proto();
1025        assert!(proto.expr_type.is_some());
1026    }
1027
1028    #[test]
1029    fn test_column_reference() {
1030        let col = ColumnReference::new("x");
1031        let proto = col.to_proto();
1032        assert!(proto.expr_type.is_some());
1033    }
1034
1035    #[test]
1036    fn test_literal_decimal() {
1037        let lit = LiteralExpression::Decimal {
1038            value: "123.45".to_string(),
1039            precision: 5,
1040            scale: 2,
1041        };
1042        let proto = lit.to_proto();
1043        assert!(proto.expr_type.is_some());
1044        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1045            if let Some(proto::expression::literal::LiteralType::Decimal(decimal)) =
1046                literal.literal_type
1047            {
1048                assert_eq!(decimal.value, "123.45");
1049                assert_eq!(decimal.precision, Some(5));
1050                assert_eq!(decimal.scale, Some(2));
1051            } else {
1052                panic!("Expected decimal literal type");
1053            }
1054        } else {
1055            panic!("Expected literal expression type");
1056        }
1057    }
1058
1059    #[test]
1060    fn test_literal_date() {
1061        let lit = LiteralExpression::Date(18993); // some days since epoch
1062        let proto = lit.to_proto();
1063        assert!(proto.expr_type.is_some());
1064        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1065            if let Some(proto::expression::literal::LiteralType::Date(days)) = literal.literal_type
1066            {
1067                assert_eq!(days, 18993);
1068            } else {
1069                panic!("Expected date literal type");
1070            }
1071        } else {
1072            panic!("Expected literal expression type");
1073        }
1074    }
1075
1076    #[test]
1077    fn test_literal_timestamp() {
1078        let lit = LiteralExpression::Timestamp(1693526400000000); // micros since epoch
1079        let proto = lit.to_proto();
1080        assert!(proto.expr_type.is_some());
1081        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1082            if let Some(proto::expression::literal::LiteralType::Timestamp(micros)) =
1083                literal.literal_type
1084            {
1085                assert_eq!(micros, 1693526400000000);
1086            } else {
1087                panic!("Expected timestamp literal type");
1088            }
1089        } else {
1090            panic!("Expected literal expression type");
1091        }
1092    }
1093
1094    // Additional tests for expression render() methods and uncovered branches
1095    #[test]
1096    fn test_literal_render_byte() {
1097        let lit = LiteralExpression::Byte(42);
1098        assert_eq!(lit.render(), "42");
1099    }
1100
1101    #[test]
1102    fn test_literal_render_short() {
1103        let lit = LiteralExpression::Short(1000);
1104        assert_eq!(lit.render(), "1000");
1105    }
1106
1107    #[test]
1108    fn test_literal_render_long() {
1109        let lit = LiteralExpression::Long(9999999999i64);
1110        assert_eq!(lit.render(), "9999999999");
1111    }
1112
1113    #[test]
1114    fn test_literal_render_float() {
1115        let lit = LiteralExpression::Float(3.14);
1116        assert_eq!(lit.render(), "3.14");
1117    }
1118
1119    #[test]
1120    fn test_literal_render_double() {
1121        let lit = LiteralExpression::Double(2.71828);
1122        assert_eq!(lit.render(), "2.71828");
1123    }
1124
1125    #[test]
1126    fn test_literal_render_string() {
1127        let lit = LiteralExpression::String("hello".to_string());
1128        assert_eq!(lit.render(), "hello");
1129    }
1130
1131    #[test]
1132    fn test_literal_render_binary() {
1133        let lit = LiteralExpression::Binary(vec![1, 2, 3]);
1134        let rendered = lit.render();
1135        assert!(rendered.contains("[1, 2, 3]"));
1136    }
1137
1138    #[test]
1139    fn test_literal_render_date() {
1140        let lit = LiteralExpression::Date(18993);
1141        assert_eq!(lit.render(), "18993");
1142    }
1143
1144    #[test]
1145    fn test_literal_render_timestamp_ntz() {
1146        let lit = LiteralExpression::TimestampNtz(1693526400000000);
1147        assert_eq!(lit.render(), "1693526400000000");
1148    }
1149
1150    #[test]
1151    fn test_literal_render_time() {
1152        let lit = LiteralExpression::Time {
1153            nano: 3600000000000i64,
1154            precision: 9,
1155        };
1156        assert_eq!(lit.render(), "3600000000000");
1157    }
1158
1159    #[test]
1160    fn test_literal_render_array() {
1161        let lit = LiteralExpression::Array {
1162            element_type: Box::new(DataType::Integer),
1163            elements: vec![
1164                LiteralExpression::int(1),
1165                LiteralExpression::int(2),
1166                LiteralExpression::int(3),
1167            ],
1168        };
1169        assert_eq!(lit.render(), "[1, 2, 3]");
1170    }
1171
1172    #[test]
1173    fn test_literal_render_array_empty() {
1174        let lit = LiteralExpression::Array {
1175            element_type: Box::new(DataType::Integer),
1176            elements: vec![],
1177        };
1178        assert_eq!(lit.render(), "[]");
1179    }
1180
1181    #[test]
1182    fn test_literal_render_decimal() {
1183        let lit = LiteralExpression::Decimal {
1184            value: "123.45".to_string(),
1185            precision: 5,
1186            scale: 2,
1187        };
1188        assert_eq!(lit.render(), "123.45");
1189    }
1190
1191    #[test]
1192    fn test_expression_render_unresolved_star_with_target() {
1193        let expr = Expression::UnresolvedStar(Some("table.*".to_string()));
1194        assert_eq!(expr.render(), "table.*");
1195    }
1196
1197    #[test]
1198    fn test_expression_render_unresolved_regex() {
1199        let expr = Expression::UnresolvedRegex("`col_.*`".to_string());
1200        assert_eq!(expr.render(), "`col_.*`");
1201    }
1202
1203    #[test]
1204    fn test_expression_render_direct_shuffle_partition_id() {
1205        let child = col("x");
1206        let expr = Expression::DirectShufflePartitionId(Box::new(child));
1207        assert_eq!(expr.render(), "DIRECT_SHUFFLE_PARTITION_ID(x)");
1208    }
1209
1210    #[test]
1211    fn test_expression_render_unresolved_extract_value() {
1212        let child = col("struct_col");
1213        let extraction = Expression::Literal(LiteralExpression::string("field"));
1214        let ev = ExtractValue::new(child, extraction);
1215        let expr = Expression::UnresolvedExtractValue(Box::new(ev));
1216        assert_eq!(expr.render(), "struct_col[field]");
1217    }
1218
1219    #[test]
1220    fn test_expression_render_update_fields_with_value() {
1221        let struct_expr = col("s");
1222        let value_expr = Expression::Literal(LiteralExpression::int(42));
1223        let uf = UpdateFieldsExpr::new(struct_expr, "f1", Some(value_expr));
1224        let expr = Expression::UpdateFields(Box::new(uf));
1225        assert_eq!(expr.render(), "update_field(s, f1, 42)");
1226    }
1227
1228    #[test]
1229    fn test_expression_render_update_fields_drop() {
1230        let struct_expr = col("s");
1231        let uf = UpdateFieldsExpr::new(struct_expr, "f1", None);
1232        let expr = Expression::UpdateFields(Box::new(uf));
1233        assert_eq!(expr.render(), "drop_field(s, f1)");
1234    }
1235
1236    #[test]
1237    fn test_sort_order_render_asc_nulls_first() {
1238        let sort = SortOrder::asc_nulls_first(col("x"));
1239        assert_eq!(sort.render(), "x ASC NULLS FIRST");
1240    }
1241
1242    #[test]
1243    fn test_sort_order_render_asc_nulls_last() {
1244        let sort = SortOrder::asc_nulls_last(col("x"));
1245        assert_eq!(sort.render(), "x ASC NULLS LAST");
1246    }
1247
1248    #[test]
1249    fn test_sort_order_render_desc_nulls_first() {
1250        let sort = SortOrder::desc_nulls_first(col("x"));
1251        assert_eq!(sort.render(), "x DESC NULLS FIRST");
1252    }
1253
1254    #[test]
1255    fn test_sort_order_render_desc_nulls_last() {
1256        let sort = SortOrder::desc_nulls_last(col("x"));
1257        assert_eq!(sort.render(), "x DESC NULLS LAST");
1258    }
1259
1260    #[test]
1261    fn test_expression_render_sort_order() {
1262        let sort = SortOrder::asc_nulls_first(col("a"));
1263        let expr = Expression::SortOrder(Box::new(sort));
1264        assert_eq!(expr.render(), "a ASC NULLS FIRST");
1265    }
1266
1267    #[test]
1268    fn test_case_when_render_single_branch() {
1269        let cw = CaseWhen::new(vec![(
1270            Expression::Literal(LiteralExpression::boolean(true)),
1271            Expression::Literal(LiteralExpression::int(1)),
1272        )]);
1273        assert_eq!(cw.render(), "CASE WHEN true THEN 1 END");
1274    }
1275
1276    #[test]
1277    fn test_case_when_render_multiple_branches() {
1278        let cw = CaseWhen::new(vec![
1279            (
1280                Expression::Literal(LiteralExpression::boolean(true)),
1281                Expression::Literal(LiteralExpression::int(1)),
1282            ),
1283            (
1284                Expression::Literal(LiteralExpression::boolean(false)),
1285                Expression::Literal(LiteralExpression::int(2)),
1286            ),
1287        ]);
1288        assert_eq!(cw.render(), "CASE WHEN true THEN 1 WHEN false THEN 2 END");
1289    }
1290
1291    #[test]
1292    fn test_case_when_render_with_else() {
1293        let cw = CaseWhen::new(vec![(
1294            Expression::Literal(LiteralExpression::boolean(true)),
1295            Expression::Literal(LiteralExpression::int(1)),
1296        )])
1297        .with_else(Expression::Literal(LiteralExpression::int(99)));
1298        assert_eq!(cw.render(), "CASE WHEN true THEN 1 ELSE 99 END");
1299    }
1300
1301    #[test]
1302    fn test_expression_render_case_when() {
1303        let cw = CaseWhen::new(vec![(
1304            col("cond"),
1305            Expression::Literal(LiteralExpression::int(1)),
1306        )]);
1307        let expr = Expression::CaseWhen(Box::new(cw));
1308        assert_eq!(expr.render(), "CASE WHEN cond THEN 1 END");
1309    }
1310
1311    #[test]
1312    fn test_unresolved_function_render_negate() {
1313        let func = UnresolvedFunction::new("negate", vec![col("x")]);
1314        assert_eq!(func.render(), "(- x)");
1315    }
1316
1317    #[test]
1318    fn test_unresolved_function_render_negative() {
1319        let func = UnresolvedFunction::new("negative", vec![col("x")]);
1320        assert_eq!(func.render(), "(- x)");
1321    }
1322
1323    #[test]
1324    fn test_unresolved_function_render_multiple_args() {
1325        let func = UnresolvedFunction::new(
1326            "concat",
1327            vec![
1328                Expression::Literal(LiteralExpression::string("a")),
1329                Expression::Literal(LiteralExpression::string("b")),
1330                Expression::Literal(LiteralExpression::string("c")),
1331            ],
1332        );
1333        assert_eq!(func.render(), "concat(a, b, c)");
1334    }
1335
1336    #[test]
1337    fn test_alias_render_multiple_names() {
1338        let alias = Alias {
1339            child: col("x"),
1340            names: vec!["a".to_string(), "b".to_string()],
1341            metadata: None,
1342        };
1343        let expr = Expression::Alias(Box::new(alias));
1344        assert_eq!(expr.render(), "x AS (a, b)");
1345    }
1346
1347    #[test]
1348    fn test_cast_render_with_datatype() {
1349        let cast = Cast::new(
1350            col("x"),
1351            DataType::String {
1352                collation: "".to_string(),
1353            },
1354        );
1355        let expr = Expression::Cast(Box::new(cast));
1356        assert_eq!(expr.render(), "CAST(x AS string)");
1357    }
1358
1359    #[test]
1360    fn test_sql_expression_render() {
1361        let expr = Expression::SQLExpression("SELECT * FROM table".to_string());
1362        assert_eq!(expr.render(), "SELECT * FROM table");
1363    }
1364
1365    #[test]
1366    fn test_literal_boolean_true_render() {
1367        let lit = LiteralExpression::boolean(true);
1368        assert_eq!(lit.render(), "true");
1369    }
1370
1371    #[test]
1372    fn test_literal_boolean_false_render() {
1373        let lit = LiteralExpression::boolean(false);
1374        assert_eq!(lit.render(), "false");
1375    }
1376
1377    #[test]
1378    fn test_unresolved_function_render_not() {
1379        let func = UnresolvedFunction::new(
1380            "not",
1381            vec![Expression::Literal(LiteralExpression::boolean(true))],
1382        );
1383        assert_eq!(func.render(), "(NOT true)");
1384    }
1385
1386    #[test]
1387    fn test_infix_operators_render() {
1388        let ops = vec![
1389            ("+", "(1 + 2)"),
1390            ("-", "(1 - 2)"),
1391            ("*", "(1 * 2)"),
1392            ("/", "(1 / 2)"),
1393            ("%", "(1 % 2)"),
1394            ("==", "(1 == 2)"),
1395            ("!=", "(1 != 2)"),
1396            ("<", "(1 < 2)"),
1397            ("<=", "(1 <= 2)"),
1398            (">", "(1 > 2)"),
1399            (">=", "(1 >= 2)"),
1400            ("&", "(1 & 2)"),
1401            ("|", "(1 | 2)"),
1402            ("^", "(1 ^ 2)"),
1403            ("<=>", "(1 <=> 2)"),
1404        ];
1405
1406        for (op_name, expected_result) in ops.iter().take(15) {
1407            let func = UnresolvedFunction::new(
1408                *op_name,
1409                vec![
1410                    Expression::Literal(LiteralExpression::int(1)),
1411                    Expression::Literal(LiteralExpression::int(2)),
1412                ],
1413            );
1414            assert_eq!(
1415                func.render(),
1416                *expected_result,
1417                "Failed for operator: {}",
1418                op_name
1419            );
1420        }
1421
1422        // Test 'and' and 'or' separately
1423        let and_func = UnresolvedFunction::new(
1424            "and",
1425            vec![
1426                Expression::Literal(LiteralExpression::boolean(true)),
1427                Expression::Literal(LiteralExpression::boolean(false)),
1428            ],
1429        );
1430        assert_eq!(and_func.render(), "(true and false)");
1431
1432        let or_func = UnresolvedFunction::new(
1433            "or",
1434            vec![
1435                Expression::Literal(LiteralExpression::boolean(true)),
1436                Expression::Literal(LiteralExpression::boolean(false)),
1437            ],
1438        );
1439        assert_eq!(or_func.render(), "(true or false)");
1440    }
1441
1442    #[test]
1443    fn test_column_reference_with_plan_id() {
1444        let mut col_ref = ColumnReference::new("x");
1445        col_ref = col_ref.with_plan_id(123);
1446        assert_eq!(col_ref.plan_id, Some(123));
1447    }
1448
1449    #[test]
1450    fn test_column_reference_metadata() {
1451        let mut col_ref = ColumnReference::new("x");
1452        col_ref = col_ref.metadata();
1453        assert!(col_ref.is_metadata_column);
1454    }
1455
1456    #[test]
1457    fn test_call_function_wrapper_render() {
1458        let cf = CallFunctionWrapper::new("my_func", vec![col("a"), col("b")]);
1459        let expr = Expression::CallFunction(Box::new(cf));
1460        // CallFunction renders as debug format
1461        let rendered = expr.render();
1462        assert!(rendered.contains("CallFunctionWrapper"));
1463    }
1464
1465    #[test]
1466    fn test_window_expression_render() {
1467        let we = WindowExpressionWrapper::new(
1468            Expression::UnresolvedFunction(UnresolvedFunction::new("sum", vec![col("x")])),
1469            vec![col("group_col")],
1470            vec![],
1471            None,
1472        );
1473        let expr = Expression::WindowExpression(Box::new(we));
1474        let rendered = expr.render();
1475        assert!(rendered.contains("WindowExpressionWrapper"));
1476    }
1477
1478    #[test]
1479    fn test_lambda_function_render() {
1480        let lf = LambdaFunction::new(col("x"), vec![UnresolvedNamedLambdaVariable::new("x")]);
1481        let expr = Expression::LambdaFunction(Box::new(lf));
1482        let rendered = expr.render();
1483        assert!(rendered.contains("LambdaFunction"));
1484    }
1485
1486    #[test]
1487    fn test_unresolved_named_lambda_variable_render() {
1488        let var = UnresolvedNamedLambdaVariable::new("x");
1489        let expr = Expression::UnresolvedNamedLambdaVariable(var);
1490        let rendered = expr.render();
1491        assert!(rendered.contains("UnresolvedNamedLambdaVariable"));
1492    }
1493
1494    #[test]
1495    fn test_to_proto_literal_null() {
1496        let lit = LiteralExpression::null(DataType::Integer);
1497        let proto = lit.to_proto();
1498        assert!(proto.expr_type.is_some());
1499        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1500            if let Some(proto::expression::literal::LiteralType::Null(_)) = literal.literal_type {
1501                // Expected
1502            } else {
1503                panic!("Expected null literal type");
1504            }
1505        } else {
1506            panic!("Expected literal expression type");
1507        }
1508    }
1509
1510    #[test]
1511    fn test_to_proto_literal_boolean() {
1512        let lit = LiteralExpression::boolean(true);
1513        let proto = lit.to_proto();
1514        assert!(proto.expr_type.is_some());
1515        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1516            if let Some(proto::expression::literal::LiteralType::Boolean(b)) = literal.literal_type
1517            {
1518                assert!(b);
1519            } else {
1520                panic!("Expected boolean literal type");
1521            }
1522        } else {
1523            panic!("Expected literal expression type");
1524        }
1525    }
1526
1527    #[test]
1528    fn test_to_proto_literal_binary() {
1529        let lit = LiteralExpression::binary(vec![1, 2, 3]);
1530        let proto = lit.to_proto();
1531        assert!(proto.expr_type.is_some());
1532        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1533            if let Some(proto::expression::literal::LiteralType::Binary(b)) = literal.literal_type {
1534                assert_eq!(b.as_ref(), [1, 2, 3]);
1535            } else {
1536                panic!("Expected binary literal type");
1537            }
1538        } else {
1539            panic!("Expected literal expression type");
1540        }
1541    }
1542
1543    #[test]
1544    fn test_to_proto_literal_string() {
1545        let lit = LiteralExpression::string("hello");
1546        let proto = lit.to_proto();
1547        assert!(proto.expr_type.is_some());
1548        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1549            if let Some(proto::expression::literal::LiteralType::String(s)) = literal.literal_type {
1550                assert_eq!(s, "hello");
1551            } else {
1552                panic!("Expected string literal type");
1553            }
1554        } else {
1555            panic!("Expected literal expression type");
1556        }
1557    }
1558
1559    #[test]
1560    fn test_to_proto_literal_array() {
1561        let lit = LiteralExpression::Array {
1562            element_type: Box::new(DataType::Integer),
1563            elements: vec![LiteralExpression::int(1), LiteralExpression::int(2)],
1564        };
1565        let proto = lit.to_proto();
1566        assert!(proto.expr_type.is_some());
1567        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1568            if let Some(proto::expression::literal::LiteralType::Array(arr)) = literal.literal_type
1569            {
1570                assert_eq!(arr.elements.len(), 2);
1571            } else {
1572                panic!("Expected array literal type");
1573            }
1574        } else {
1575            panic!("Expected literal expression type");
1576        }
1577    }
1578
1579    #[test]
1580    fn test_to_proto_unresolved_star_none() {
1581        let expr = Expression::UnresolvedStar(None);
1582        let proto = expr.to_proto();
1583        assert!(proto.expr_type.is_some());
1584        if let Some(proto::expression::ExprType::UnresolvedStar(star)) = proto.expr_type {
1585            assert!(star.unparsed_target.is_none());
1586        } else {
1587            panic!("Expected unresolved star expression type");
1588        }
1589    }
1590
1591    #[test]
1592    fn test_to_proto_unresolved_star_with_target() {
1593        let expr = Expression::UnresolvedStar(Some("table.*".to_string()));
1594        let proto = expr.to_proto();
1595        assert!(proto.expr_type.is_some());
1596        if let Some(proto::expression::ExprType::UnresolvedStar(star)) = proto.expr_type {
1597            assert_eq!(star.unparsed_target, Some("table.*".to_string()));
1598        } else {
1599            panic!("Expected unresolved star expression type");
1600        }
1601    }
1602
1603    #[test]
1604    fn test_to_proto_unresolved_regex() {
1605        let expr = Expression::UnresolvedRegex("`col_.*`".to_string());
1606        let proto = expr.to_proto();
1607        assert!(proto.expr_type.is_some());
1608        if let Some(proto::expression::ExprType::UnresolvedRegex(regex)) = proto.expr_type {
1609            assert_eq!(regex.col_name, "`col_.*`");
1610        } else {
1611            panic!("Expected unresolved regex expression type");
1612        }
1613    }
1614
1615    #[test]
1616    fn test_to_proto_direct_shuffle_partition_id() {
1617        let child = col("x");
1618        let expr = Expression::DirectShufflePartitionId(Box::new(child));
1619        let proto = expr.to_proto();
1620        assert!(proto.expr_type.is_some());
1621        if let Some(proto::expression::ExprType::DirectShufflePartitionId(dspi)) = proto.expr_type {
1622            assert!(dspi.child.is_some());
1623        } else {
1624            panic!("Expected direct shuffle partition id expression type");
1625        }
1626    }
1627
1628    #[test]
1629    fn test_to_proto_unresolved_extract_value() {
1630        let child = col("struct_col");
1631        let extraction = Expression::Literal(LiteralExpression::string("field"));
1632        let ev = ExtractValue::new(child, extraction);
1633        let expr = Expression::UnresolvedExtractValue(Box::new(ev));
1634        let proto = expr.to_proto();
1635        assert!(proto.expr_type.is_some());
1636        if let Some(proto::expression::ExprType::UnresolvedExtractValue(uev)) = proto.expr_type {
1637            assert!(uev.child.is_some());
1638            assert!(uev.extraction.is_some());
1639        } else {
1640            panic!("Expected unresolved extract value expression type");
1641        }
1642    }
1643
1644    #[test]
1645    fn test_to_proto_update_fields() {
1646        let struct_expr = col("s");
1647        let value_expr = Expression::Literal(LiteralExpression::int(42));
1648        let uf = UpdateFieldsExpr::new(struct_expr, "f1", Some(value_expr));
1649        let expr = Expression::UpdateFields(Box::new(uf));
1650        let proto = expr.to_proto();
1651        assert!(proto.expr_type.is_some());
1652        if let Some(proto::expression::ExprType::UpdateFields(uf_proto)) = proto.expr_type {
1653            assert_eq!(uf_proto.field_name, "f1");
1654            assert!(uf_proto.value_expression.is_some());
1655        } else {
1656            panic!("Expected update fields expression type");
1657        }
1658    }
1659
1660    #[test]
1661    fn test_to_proto_alias_with_metadata() {
1662        let alias = Alias::new(col("x"), "y").with_metadata("metadata_str".to_string());
1663        let expr = Expression::Alias(Box::new(alias));
1664        let proto = expr.to_proto();
1665        assert!(proto.expr_type.is_some());
1666        if let Some(proto::expression::ExprType::Alias(alias_proto)) = proto.expr_type {
1667            assert_eq!(alias_proto.name, vec!["y".to_string()]);
1668            assert_eq!(alias_proto.metadata, Some("metadata_str".to_string()));
1669        } else {
1670            panic!("Expected alias expression type");
1671        }
1672    }
1673
1674    #[test]
1675    fn test_to_proto_cast_with_datatype() {
1676        let cast = Cast::new(col("x"), DataType::Integer);
1677        let expr = Expression::Cast(Box::new(cast));
1678        let proto = expr.to_proto();
1679        assert!(proto.expr_type.is_some());
1680        if let Some(proto::expression::ExprType::Cast(cast_proto)) = proto.expr_type {
1681            assert!(cast_proto.expr.is_some());
1682            assert!(cast_proto.cast_to_type.is_some());
1683        } else {
1684            panic!("Expected cast expression type");
1685        }
1686    }
1687
1688    #[test]
1689    fn test_to_proto_cast_with_eval_mode_legacy() {
1690        let cast = Cast::new(col("x"), DataType::Integer).with_eval_mode(CastEvalMode::Legacy);
1691        let expr = Expression::Cast(Box::new(cast));
1692        let proto = expr.to_proto();
1693        assert!(proto.expr_type.is_some());
1694        if let Some(proto::expression::ExprType::Cast(cast_proto)) = proto.expr_type {
1695            assert_eq!(cast_proto.eval_mode, 1i32);
1696        } else {
1697            panic!("Expected cast expression type");
1698        }
1699    }
1700
1701    #[test]
1702    fn test_to_proto_cast_with_eval_mode_ansi() {
1703        let cast = Cast::new(col("x"), DataType::Integer).with_eval_mode(CastEvalMode::Ansi);
1704        let expr = Expression::Cast(Box::new(cast));
1705        let proto = expr.to_proto();
1706        assert!(proto.expr_type.is_some());
1707        if let Some(proto::expression::ExprType::Cast(cast_proto)) = proto.expr_type {
1708            assert_eq!(cast_proto.eval_mode, 2i32);
1709        } else {
1710            panic!("Expected cast expression type");
1711        }
1712    }
1713
1714    #[test]
1715    fn test_to_proto_cast_with_eval_mode_try() {
1716        let cast = Cast::new(col("x"), DataType::Integer).with_eval_mode(CastEvalMode::Try);
1717        let expr = Expression::Cast(Box::new(cast));
1718        let proto = expr.to_proto();
1719        assert!(proto.expr_type.is_some());
1720        if let Some(proto::expression::ExprType::Cast(cast_proto)) = proto.expr_type {
1721            assert_eq!(cast_proto.eval_mode, 3i32);
1722        } else {
1723            panic!("Expected cast expression type");
1724        }
1725    }
1726
1727    #[test]
1728    fn test_to_proto_cast_str() {
1729        let cast = Cast::new_str(col("x"), "integer");
1730        let expr = Expression::Cast(Box::new(cast));
1731        let proto = expr.to_proto();
1732        assert!(proto.expr_type.is_some());
1733        if let Some(proto::expression::ExprType::Cast(cast_proto)) = proto.expr_type {
1734            assert!(cast_proto.cast_to_type.is_some());
1735        } else {
1736            panic!("Expected cast expression type");
1737        }
1738    }
1739
1740    #[test]
1741    fn test_to_proto_sort_order() {
1742        let sort = SortOrder::asc_nulls_first(col("x"));
1743        let expr = Expression::SortOrder(Box::new(sort));
1744        let proto = expr.to_proto();
1745        assert!(proto.expr_type.is_some());
1746        if let Some(proto::expression::ExprType::SortOrder(sort_proto)) = proto.expr_type {
1747            assert_eq!(sort_proto.direction, 1i32); // ASC
1748            assert_eq!(sort_proto.null_ordering, 1i32); // FIRST
1749        } else {
1750            panic!("Expected sort order expression type");
1751        }
1752    }
1753
1754    #[test]
1755    fn test_to_proto_case_when() {
1756        let cw = CaseWhen::new(vec![(
1757            Expression::Literal(LiteralExpression::boolean(true)),
1758            Expression::Literal(LiteralExpression::int(1)),
1759        )])
1760        .with_else(Expression::Literal(LiteralExpression::int(99)));
1761        let expr = Expression::CaseWhen(Box::new(cw));
1762        let proto = expr.to_proto();
1763        assert!(proto.expr_type.is_some());
1764        if let Some(proto::expression::ExprType::UnresolvedFunction(func_proto)) = proto.expr_type {
1765            assert_eq!(func_proto.function_name, "when");
1766            assert_eq!(func_proto.arguments.len(), 3); // 2 branches + 1 else
1767        } else {
1768            panic!("Expected unresolved function expression type for case when");
1769        }
1770    }
1771
1772    #[test]
1773    fn test_to_proto_sql_expression() {
1774        let expr = Expression::SQLExpression("SELECT * FROM table".to_string());
1775        let proto = expr.to_proto();
1776        assert!(proto.expr_type.is_some());
1777        if let Some(proto::expression::ExprType::ExpressionString(es)) = proto.expr_type {
1778            assert_eq!(es.expression, "SELECT * FROM table");
1779        } else {
1780            panic!("Expected expression string type");
1781        }
1782    }
1783
1784    #[test]
1785    fn test_to_proto_call_function() {
1786        let cf = CallFunctionWrapper::new("my_func", vec![col("a"), col("b")]);
1787        let expr = Expression::CallFunction(Box::new(cf));
1788        let proto = expr.to_proto();
1789        assert!(proto.expr_type.is_some());
1790        if let Some(proto::expression::ExprType::CallFunction(cf_proto)) = proto.expr_type {
1791            assert_eq!(cf_proto.function_name, "my_func");
1792            assert_eq!(cf_proto.arguments.len(), 2);
1793        } else {
1794            panic!("Expected call function expression type");
1795        }
1796    }
1797
1798    #[test]
1799    fn test_to_proto_lambda_function() {
1800        let lf = LambdaFunction::new(col("x"), vec![UnresolvedNamedLambdaVariable::new("x")]);
1801        let expr = Expression::LambdaFunction(Box::new(lf));
1802        let proto = expr.to_proto();
1803        assert!(proto.expr_type.is_some());
1804        if let Some(proto::expression::ExprType::LambdaFunction(lf_proto)) = proto.expr_type {
1805            assert!(lf_proto.function.is_some());
1806            assert_eq!(lf_proto.arguments.len(), 1);
1807        } else {
1808            panic!("Expected lambda function expression type");
1809        }
1810    }
1811
1812    #[test]
1813    fn test_to_proto_unresolved_named_lambda_variable() {
1814        let var = UnresolvedNamedLambdaVariable::new("x");
1815        let expr = Expression::UnresolvedNamedLambdaVariable(var);
1816        let proto = expr.to_proto();
1817        assert!(proto.expr_type.is_some());
1818        if let Some(proto::expression::ExprType::UnresolvedNamedLambdaVariable(var_proto)) =
1819            proto.expr_type
1820        {
1821            assert_eq!(var_proto.name_parts.len(), 1);
1822        } else {
1823            panic!("Expected unresolved named lambda variable expression type");
1824        }
1825    }
1826
1827    #[test]
1828    fn test_window_expression_to_proto() {
1829        let we = WindowExpressionWrapper::new(
1830            Expression::UnresolvedFunction(UnresolvedFunction::new("sum", vec![col("x")])),
1831            vec![col("group_col")],
1832            vec![SortOrder::asc_nulls_last(col("sort_col"))],
1833            Some((
1834                1u32,
1835                FrameBoundary::UnboundedPreceding,
1836                FrameBoundary::CurrentRow,
1837            )),
1838        );
1839        let expr = Expression::WindowExpression(Box::new(we));
1840        let proto = expr.to_proto();
1841        assert!(proto.expr_type.is_some());
1842        if let Some(proto::expression::ExprType::Window(window_proto)) = proto.expr_type {
1843            assert!(window_proto.window_function.is_some());
1844            assert_eq!(window_proto.partition_spec.len(), 1);
1845            assert_eq!(window_proto.order_spec.len(), 1);
1846            assert!(window_proto.frame_spec.is_some());
1847        } else {
1848            panic!("Expected window expression type");
1849        }
1850    }
1851
1852    #[test]
1853    fn test_literal_byte_to_proto() {
1854        let lit = LiteralExpression::Byte(42);
1855        let proto = lit.to_proto();
1856        assert!(proto.expr_type.is_some());
1857        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1858            if let Some(proto::expression::literal::LiteralType::Byte(b)) = literal.literal_type {
1859                assert_eq!(b, 42);
1860            } else {
1861                panic!("Expected byte literal type");
1862            }
1863        } else {
1864            panic!("Expected literal expression type");
1865        }
1866    }
1867
1868    #[test]
1869    fn test_literal_short_to_proto() {
1870        let lit = LiteralExpression::Short(1000);
1871        let proto = lit.to_proto();
1872        assert!(proto.expr_type.is_some());
1873        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1874            if let Some(proto::expression::literal::LiteralType::Short(s)) = literal.literal_type {
1875                assert_eq!(s, 1000);
1876            } else {
1877                panic!("Expected short literal type");
1878            }
1879        } else {
1880            panic!("Expected literal expression type");
1881        }
1882    }
1883
1884    #[test]
1885    fn test_literal_float_to_proto() {
1886        let lit = LiteralExpression::Float(3.14);
1887        let proto = lit.to_proto();
1888        assert!(proto.expr_type.is_some());
1889        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1890            if let Some(proto::expression::literal::LiteralType::Float(f)) = literal.literal_type {
1891                assert!((f - 3.14).abs() < 0.01);
1892            } else {
1893                panic!("Expected float literal type");
1894            }
1895        } else {
1896            panic!("Expected literal expression type");
1897        }
1898    }
1899
1900    #[test]
1901    fn test_literal_double_to_proto() {
1902        let lit = LiteralExpression::Double(2.71828);
1903        let proto = lit.to_proto();
1904        assert!(proto.expr_type.is_some());
1905        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1906            if let Some(proto::expression::literal::LiteralType::Double(d)) = literal.literal_type {
1907                assert!((d - 2.71828).abs() < 0.00001);
1908            } else {
1909                panic!("Expected double literal type");
1910            }
1911        } else {
1912            panic!("Expected literal expression type");
1913        }
1914    }
1915
1916    #[test]
1917    fn test_literal_time_to_proto() {
1918        let lit = LiteralExpression::Time {
1919            nano: 3600000000000i64,
1920            precision: 9,
1921        };
1922        let proto = lit.to_proto();
1923        assert!(proto.expr_type.is_some());
1924        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1925            if let Some(proto::expression::literal::LiteralType::Time(t)) = literal.literal_type {
1926                assert_eq!(t.nano, 3600000000000i64);
1927                assert_eq!(t.precision, Some(9));
1928            } else {
1929                panic!("Expected time literal type");
1930            }
1931        } else {
1932            panic!("Expected literal expression type");
1933        }
1934    }
1935
1936    #[test]
1937    fn test_literal_timestamp_ntz_to_proto() {
1938        let lit = LiteralExpression::TimestampNtz(1693526400000000);
1939        let proto = lit.to_proto();
1940        assert!(proto.expr_type.is_some());
1941        if let Some(proto::expression::ExprType::Literal(literal)) = proto.expr_type {
1942            if let Some(proto::expression::literal::LiteralType::TimestampNtz(ts)) =
1943                literal.literal_type
1944            {
1945                assert_eq!(ts, 1693526400000000);
1946            } else {
1947                panic!("Expected timestamp ntz literal type");
1948            }
1949        } else {
1950            panic!("Expected literal expression type");
1951        }
1952    }
1953}