1use std::sync::atomic::{AtomicU32, Ordering};
6
7use spark_connect_proto as proto;
8
9use crate::types::DataType;
10use crate::udf::CommonInlineUserDefinedFunctionExpression;
11
12static LAMBDA_VAR_COUNTER: AtomicU32 = AtomicU32::new(0);
14
15pub fn next_lambda_var_index() -> u32 {
17 LAMBDA_VAR_COUNTER.fetch_add(1, Ordering::SeqCst)
18}
19
20#[derive(Debug, Clone, PartialEq)]
25pub enum Expression {
26 Literal(LiteralExpression),
28 ColumnReference(ColumnReference),
30 UnresolvedFunction(UnresolvedFunction),
32 UnresolvedStar(Option<String>),
35 Alias(Box<Alias>),
37 Cast(Box<Cast>),
39 DirectShufflePartitionId(Box<Expression>),
42 UnresolvedRegex(String),
44 SortOrder(Box<SortOrder>),
46 CaseWhen(Box<CaseWhen>),
48 UnresolvedExtractValue(Box<ExtractValue>),
50 UpdateFields(Box<UpdateFieldsExpr>),
52 SQLExpression(String),
54 CallFunction(Box<CallFunctionWrapper>),
56 WindowExpression(Box<WindowExpressionWrapper>),
58 LambdaFunction(Box<LambdaFunction>),
60 UnresolvedNamedLambdaVariable(UnresolvedNamedLambdaVariable),
62 CommonInlineUserDefinedFunction(Box<CommonInlineUserDefinedFunctionExpression>),
64}
65
66impl Expression {
67 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 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#[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 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 pub fn null(data_type: DataType) -> Self {
322 LiteralExpression::Null(data_type)
323 }
324
325 pub fn int(value: i32) -> Self {
327 LiteralExpression::Integer(value)
328 }
329
330 pub fn long(value: i64) -> Self {
332 LiteralExpression::Long(value)
333 }
334
335 pub fn double(value: f64) -> Self {
337 LiteralExpression::Double(value)
338 }
339
340 pub fn string(value: impl Into<String>) -> Self {
342 LiteralExpression::String(value.into())
343 }
344
345 pub fn boolean(value: bool) -> Self {
347 LiteralExpression::Boolean(value)
348 }
349
350 pub fn binary(value: Vec<u8>) -> Self {
352 LiteralExpression::Binary(value)
353 }
354}
355
356#[derive(Debug, Clone, PartialEq, Eq)]
359pub struct ColumnReference {
360 pub name: String,
362 pub plan_id: Option<i64>,
366 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 pub fn with_plan_id(mut self, plan_id: i64) -> Self {
382 self.plan_id = Some(plan_id);
383 self
384 }
385
386 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#[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#[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#[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 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#[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#[derive(Debug, Clone, PartialEq)]
570pub struct Cast {
571 pub child: Expression,
572 pub target: CastTarget,
573 pub eval_mode: Option<CastEvalMode>,
574}
575
576#[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 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#[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 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#[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 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#[derive(Debug, Clone, PartialEq)]
760pub struct CallFunctionWrapper {
761 pub function_name: String,
762 pub arguments: Vec<Expression>,
763}
764
765impl CallFunctionWrapper {
766 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 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#[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#[derive(Debug, Clone, PartialEq)]
798pub enum FrameBoundary {
799 UnboundedPreceding,
800 Preceding(i64),
801 CurrentRow,
802 Following(i64),
803 UnboundedFollowing,
804}
805
806impl WindowExpressionWrapper {
807 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 pub fn to_proto(&self) -> proto::Expression {
824 let mut expr = proto::Expression::default();
825 let mut window = proto::expression::Window::default();
826
827 window.window_function = Some(Box::new(self.window_function.to_proto()));
829
830 window.partition_spec = self.partition_spec.iter().map(|e| e.to_proto()).collect();
832
833 window.order_spec = self
835 .order_spec
836 .iter()
837 .map(|s| s.to_proto_sort_order())
838 .collect();
839
840 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
854fn 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 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 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#[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#[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 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 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 let neq = Expression::UnresolvedFunction(UnresolvedFunction::new("not", vec![eq.clone()]));
990 assert_eq!(neq.render(), "(NOT (a == b))");
991
992 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 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 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); 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); 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 #[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 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 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 } 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); assert_eq!(sort_proto.null_ordering, 1i32); } 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); } 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}