1use std::any::Any;
8use std::collections::HashMap;
9use std::sync::{Arc, LazyLock};
10
11use datafusion::arrow::array::{
12 Array, FixedSizeListArray, LargeListArray, ListArray, new_empty_array,
13};
14use datafusion::arrow::datatypes::{DataType, Field, FieldRef};
15use datafusion::logical_expr::expr::Placeholder;
16use datafusion::logical_expr::{
17 ColumnarValue, Expr as DfExpr, ExprSchemable, Operator, ReturnFieldArgs, ScalarFunctionArgs,
18 ScalarUDF, ScalarUDFImpl, Signature, Volatility, cast, col, lit, not, when,
19};
20use datafusion::scalar::ScalarValue;
21
22use graphforge_core::PropId;
23use graphforge_ir::expr::{BinaryOpKind, IrExpr, IrLiteral, UnaryOpKind};
24use graphforge_ir::{ExprArena, ExprId, VarId};
25use graphforge_ontology::OntologyHandle;
26
27#[derive(Debug, Clone, Default)]
38pub struct VarMap(HashMap<u32, String>);
39
40impl VarMap {
41 #[must_use]
43 pub fn new() -> Self {
44 Self::default()
45 }
46
47 pub fn insert(&mut self, var: VarId, col_name: impl Into<String>) {
49 self.0.insert(var.0, col_name.into());
50 }
51
52 pub fn clear(&mut self) {
55 self.0.clear();
56 }
57
58 #[must_use]
60 pub fn get(&self, var: VarId) -> Option<&str> {
61 self.0.get(&var.0).map(String::as_str)
62 }
63
64 pub fn var_ids(&self) -> impl Iterator<Item = VarId> + '_ {
69 self.0.keys().map(|&k| VarId(k))
70 }
71}
72
73#[derive(Debug, Clone, thiserror::Error)]
75pub enum LoweringError {
76 #[error("unknown built-in function: {0}")]
79 UnknownFunction(String),
80
81 #[error("unsupported expression: {0}")]
83 UnsupportedExpr(String),
84
85 #[error("unbound variable: VarId({0})")]
88 UnboundVar(u32),
89
90 #[error("invalid argument type: {0}")]
95 InvalidType(String),
96}
97
98#[derive(Clone, Debug)]
102pub struct NodeShape {
103 pub prop_names: Vec<String>,
107}
108
109#[derive(Clone, Copy, Debug, PartialEq, Eq)]
110enum EntityIdentityKind {
111 Node,
112 Edge,
113}
114
115pub struct ExprLowerer<'a> {
124 arena: &'a ExprArena,
125 var_map: &'a VarMap,
126 prop_names: HashMap<u32, String>,
130 node_shapes: HashMap<u32, NodeShape>,
133 type_id_to_entity_name: HashMap<u32, String>,
138 entity_name_to_type_id: HashMap<String, u32>,
143 props_authoritative: bool,
150 input_schema: Option<datafusion::common::DFSchemaRef>,
156 elem_struct_cols: Vec<String>,
165 read_target: Option<std::path::PathBuf>,
171 now: std::sync::OnceLock<chrono::NaiveDateTime>,
176}
177
178impl<'a> ExprLowerer<'a> {
179 #[must_use]
190 pub fn new(
191 arena: &'a ExprArena,
192 ontology: Option<&'a OntologyHandle>,
193 var_map: &'a VarMap,
194 ) -> Self {
195 Self {
196 arena,
197 var_map,
198 prop_names: build_prop_names(ontology),
199 node_shapes: HashMap::new(),
200 type_id_to_entity_name: HashMap::new(),
201 entity_name_to_type_id: HashMap::new(),
202 props_authoritative: false,
203 input_schema: None,
204 elem_struct_cols: Vec::new(),
205 read_target: None,
206 now: std::sync::OnceLock::new(),
207 }
208 }
209
210 #[must_use]
216 pub fn with_prop_names(
217 arena: &'a ExprArena,
218 var_map: &'a VarMap,
219 prop_names: HashMap<u32, String>,
220 ) -> Self {
221 Self {
222 arena,
223 var_map,
224 prop_names,
225 node_shapes: HashMap::new(),
226 type_id_to_entity_name: HashMap::new(),
227 entity_name_to_type_id: HashMap::new(),
228 props_authoritative: false,
229 input_schema: None,
230 elem_struct_cols: Vec::new(),
231 read_target: None,
232 now: std::sync::OnceLock::new(),
233 }
234 }
235
236 #[must_use]
243 pub fn with_prop_names_and_nodes(
244 arena: &'a ExprArena,
245 var_map: &'a VarMap,
246 prop_names: HashMap<u32, String>,
247 node_shapes: HashMap<u32, NodeShape>,
248 type_id_to_entity_name: HashMap<u32, String>,
249 props_authoritative: bool,
250 ) -> Self {
251 let entity_name_to_type_id = type_id_to_entity_name
252 .iter()
253 .map(|(id, name)| (name.clone(), *id))
254 .collect();
255 Self {
256 arena,
257 var_map,
258 prop_names,
259 node_shapes,
260 type_id_to_entity_name,
261 entity_name_to_type_id,
262 props_authoritative,
263 input_schema: None,
264 elem_struct_cols: Vec::new(),
265 read_target: None,
266 now: std::sync::OnceLock::new(),
267 }
268 }
269
270 #[must_use]
274 pub fn with_input_schema(mut self, schema: datafusion::common::DFSchemaRef) -> Self {
275 self.input_schema = Some(schema);
276 self
277 }
278
279 #[must_use]
284 pub fn with_elem_struct_col(mut self, col: String) -> Self {
285 self.elem_struct_cols.push(col);
286 self
287 }
288
289 #[must_use]
292 pub fn with_read_target(mut self, dir: std::path::PathBuf) -> Self {
293 self.read_target = Some(dir);
294 self
295 }
296
297 #[allow(
303 clippy::too_many_lines,
304 reason = "one cohesive dispatch match over every IrExpr variant plus the \
305 namespaced temporal builtins (date/time/datetime truncate); \
306 splitting the arms would scatter the lowering logic"
307 )]
308 pub fn lower(&self, id: ExprId) -> Result<DfExpr, LoweringError> {
309 match self.arena.get(id) {
310 IrExpr::Literal(lit_val) => Ok(lower_literal(lit_val)),
311
312 IrExpr::VarRef(var_id) => {
313 let col_name = self
314 .var_map
315 .get(*var_id)
316 .ok_or(LoweringError::UnboundVar(var_id.0))?;
317 if let Some(schema) = self.input_schema.as_ref()
318 && schema.field_with_unqualified_name(col_name).is_err()
319 {
320 let qual = datafusion::common::TableReference::bare(col_name);
321 if schema
322 .index_of_column_by_name(Some(&qual), "node_uuid")
323 .is_some()
324 {
325 return Ok(col(format!("{col_name}.node_uuid")));
326 }
327 if schema
328 .index_of_column_by_name(Some(&qual), "edge_uuid")
329 .is_some()
330 {
331 return Ok(col(format!("{col_name}.edge_uuid")));
332 }
333 }
334 if self
339 .input_schema
340 .as_ref()
341 .is_some_and(|schema| schema.index_of_column_by_name(None, col_name).is_some())
342 {
343 Ok(DfExpr::Column(datafusion::common::Column::new_unqualified(
344 col_name,
345 )))
346 } else {
347 Ok(col_literal(col_name))
348 }
349 }
350
351 IrExpr::PropertyAccess { base, prop } => {
352 if let Some(prop_name) = self.prop_names.get(&prop.0).cloned()
353 && let Some(out) = self.lower_static_value_access(*base, &prop_name)?
354 {
355 return Ok(out);
356 }
357 if self.props_authoritative
372 && let IrExpr::VarRef(v) = self.arena.get(*base)
373 && let Some(shape) = self.node_shapes.get(&v.0)
374 {
375 let prop_name = self
376 .prop_names
377 .get(&prop.0)
378 .cloned()
379 .unwrap_or_else(|| format!("prop_{}", prop.0));
380 let is_topology = graphforge_storage::TOPOLOGY_NODES_SCHEMA
381 .field_with_name(&prop_name)
382 .is_ok();
383 if !is_topology && !shape.prop_names.contains(&prop_name) {
384 return Ok(lit(ScalarValue::Null));
385 }
386 }
387 if let IrExpr::VarRef(v) = self.arena.get(*base)
394 && let Some(col_name) = self.var_map.get(*v)
395 && let Some(prop_name) = self.prop_names.get(&prop.0)
396 && crate::temporal::is_date_accessor(prop_name)
397 && let Some(schema) = self.input_schema.as_ref()
398 && let Ok(field) = schema.field_with_unqualified_name(col_name)
399 && is_date_struct(field.data_type())
400 {
401 return Ok(CYPHER_DATE_COMPONENT
402 .call(vec![col_literal(col_name), lit(prop_name.as_str())]));
403 }
404 if let IrExpr::VarRef(v) = self.arena.get(*base)
407 && let Some(col_name) = self.var_map.get(*v)
408 && let Some(prop_name) = self.prop_names.get(&prop.0)
409 && crate::temporal::is_duration_accessor(prop_name)
410 && let Some(schema) = self.input_schema.as_ref()
411 && let Ok(field) = schema.field_with_unqualified_name(col_name)
412 && is_duration_struct(field.data_type())
413 {
414 return Ok(CYPHER_DURATION_COMPONENT
415 .call(vec![col_literal(col_name), lit(prop_name.as_str())]));
416 }
417 if let IrExpr::VarRef(v) = self.arena.get(*base)
424 && let Some(col_name) = self.var_map.get(*v)
425 && let Some(prop_name) = self.prop_names.get(&prop.0)
426 && let Some(schema) = self.input_schema.as_ref()
427 && let Ok(field) = schema.field_with_unqualified_name(col_name)
428 && temporal_accessor_valid(field.data_type(), prop_name)
429 {
430 let args = vec![col_literal(col_name), lit(prop_name.as_str())];
431 return Ok(if crate::temporal::is_zone_str_accessor(prop_name) {
432 CYPHER_TEMPORAL_ZONE_STR.call(args)
433 } else {
434 CYPHER_TEMPORAL_COMPONENT.call(args)
435 });
436 }
437 if let IrExpr::VarRef(v) = self.arena.get(*base)
445 && let Some(col_name) = self.var_map.get(*v)
446 && let Some(schema) = self.input_schema.as_ref()
447 && let Ok(field) = schema.field_with_unqualified_name(col_name)
448 && is_plain_map_struct_type(field.data_type())
449 {
450 let prop_name = self
451 .prop_names
452 .get(&prop.0)
453 .cloned()
454 .unwrap_or_else(|| format!("prop_{}", prop.0));
455 return Ok(datafusion::functions::core::expr_fn::get_field(
456 col_literal(col_name),
457 prop_name,
458 ));
459 }
460 let base_expr = if let IrExpr::VarRef(v) = self.arena.get(*base) {
465 col_literal(self.var_map.get(*v).ok_or(LoweringError::UnboundVar(v.0))?)
466 } else {
467 self.lower(*base)?
468 };
469 if self.is_known_non_value_access_container(&base_expr) {
470 let prop_name = self
471 .prop_names
472 .get(&prop.0)
473 .cloned()
474 .unwrap_or_else(|| format!("prop_{}", prop.0));
475 return Err(LoweringError::InvalidType(format!(
476 "property access `{prop_name}` requires a map or graph element"
477 )));
478 }
479 let prop_col = self.resolve_prop_col(base_expr, *prop);
480 Ok(prop_col)
481 }
482
483 IrExpr::BinaryOp { op, left, right } => self.lower_binary(*op, *left, *right),
484
485 IrExpr::UnaryOp { op, expr } => self.lower_unary(*op, *expr),
486
487 IrExpr::FunctionCall { name, args } if name == "_node_struct" => {
488 self.lower_node_struct(args)
489 }
490
491 IrExpr::FunctionCall { name, args } if name == "_node_struct_list" => {
492 self.lower_node_struct_list(args)
493 }
494
495 IrExpr::FunctionCall { name, args } if name == "_rel_struct" => {
496 self.lower_rel_struct(args)
497 }
498
499 IrExpr::FunctionCall { name, args } if name == "_rel_struct_list" => {
500 self.lower_rel_struct_list(args)
501 }
502
503 IrExpr::FunctionCall { name, args } if name == "keys" => self.lower_keys(args),
504
505 IrExpr::FunctionCall { name, args } if name == "properties" => {
506 self.lower_properties(args)
507 }
508
509 IrExpr::FunctionCall { name, args } if name == "labels" => self.lower_labels(args),
510
511 IrExpr::FunctionCall { name, args }
512 if matches!(name.as_str(), "nodes" | "relationships") =>
513 {
514 let [arg] = args.as_slice() else {
515 return Err(LoweringError::InvalidType(format!(
516 "{name}() expects one path argument"
517 )));
518 };
519 Ok(datafusion::functions::core::expr_fn::get_field(
520 self.lower(*arg)?,
521 name,
522 ))
523 }
524
525 IrExpr::FunctionCall { name, args } if name == "_subscript" => {
526 self.lower_subscript(args)
527 }
528
529 IrExpr::FunctionCall { name, args }
530 if matches!(
531 name.as_str(),
532 "date" | "localtime" | "time" | "localdatetime" | "datetime" | "duration"
533 ) =>
534 {
535 self.lower_temporal(name, args)
536 }
537
538 IrExpr::FunctionCall { name, args }
539 if matches!(
540 name.as_str(),
541 "datetime.fromepoch" | "datetime.fromepochmillis"
542 ) =>
543 {
544 self.lower_from_epoch(name, args)
545 }
546
547 IrExpr::FunctionCall { name, args } if name == "date.truncate" => {
548 self.lower_date_truncate(args)
549 }
550 IrExpr::FunctionCall { name, args } if name == "localtime.truncate" => {
551 self.lower_localtime_truncate(args)
552 }
553 IrExpr::FunctionCall { name, args } if name == "localdatetime.truncate" => {
554 self.lower_localdatetime_truncate(args)
555 }
556 IrExpr::FunctionCall { name, args } if name == "time.truncate" => {
557 self.lower_time_truncate(args)
558 }
559 IrExpr::FunctionCall { name, args } if name == "datetime.truncate" => {
560 self.lower_datetime_truncate(args)
561 }
562 IrExpr::FunctionCall { name, args } if is_temporal_clock_fn(name) => {
568 if self.sole_arg_is_null(args) {
569 Ok(DfExpr::Literal(temporal_null_scalar(name), None))
570 } else {
571 Err(LoweringError::UnsupportedExpr(format!(
572 "{name}: temporal clock functions are not supported in a \
573 deterministic query context"
574 )))
575 }
576 }
577 IrExpr::FunctionCall { name, args }
581 if matches!(
582 name.to_ascii_lowercase().as_str(),
583 "duration.between"
584 | "duration.inmonths"
585 | "duration.indays"
586 | "duration.inseconds"
587 ) =>
588 {
589 self.lower_duration_between(&name.to_ascii_lowercase(), args)
590 }
591
592 IrExpr::FunctionCall { name, args } => {
593 let lowered = if is_path_builtin_name(name) {
594 args.iter()
595 .map(|&a| self.lower_path_builtin_arg(a))
596 .collect::<Result<Vec<_>, _>>()?
597 } else {
598 args.iter()
599 .map(|&a| self.lower(a))
600 .collect::<Result<Vec<_>, _>>()?
601 };
602 if name == "reverse"
607 && let [arg] = lowered.as_slice()
608 {
609 return Ok(if self.is_string_typed(arg) {
610 datafusion::functions::unicode::expr_fn::reverse(arg.clone())
611 } else if self.is_list_typed(arg) {
612 datafusion::functions_nested::expr_fn::array_reverse(arg.clone())
613 } else {
614 CYPHER_REVERSE.call(vec![arg.clone()])
618 });
619 }
620 resolve_builtin(name, lowered, || self.path_node_hydration())
621 .ok_or_else(|| LoweringError::UnknownFunction(name.clone()))
622 }
623
624 IrExpr::Parameter(name) => Ok(DfExpr::Placeholder(Placeholder {
625 id: format!("${name}"),
631 field: None,
632 })),
633
634 IrExpr::Case {
635 operand,
636 arms,
637 else_expr,
638 } => self.lower_case(operand.as_ref().copied(), arms, else_expr.as_ref().copied()),
639
640 IrExpr::ListLiteral(ids) => {
641 let elems: Vec<DfExpr> = ids
642 .iter()
643 .map(|&id| self.lower_value(id))
644 .collect::<Result<_, _>>()?;
645 Ok(lower_list_literal(elems, self.input_schema.as_deref()))
646 }
647
648 IrExpr::MapLiteral(entries) => self.lower_map_literal(entries),
649
650 IrExpr::Quantifier {
651 kind,
652 loop_var,
653 list,
654 predicate,
655 } => self.lower_quantifier(*kind, *loop_var, *list, *predicate),
656
657 IrExpr::ListComprehension {
658 loop_var,
659 list,
660 filter,
661 projection,
662 } => self.lower_list_comprehension(*loop_var, *list, *filter, *projection),
663 }
664 }
665
666 fn lower_value(&self, id: ExprId) -> Result<DfExpr, LoweringError> {
667 if let IrExpr::VarRef(var_id) = self.arena.get(id) {
668 let base = self
669 .var_map
670 .get(*var_id)
671 .ok_or(LoweringError::UnboundVar(var_id.0))?;
672 if self.node_shapes.contains_key(&var_id.0) || self.is_node_var(base) {
673 let prop_names = self
674 .node_shapes
675 .get(&var_id.0)
676 .map(|s| s.prop_names.clone())
677 .unwrap_or_default();
678 return Ok(node_value_struct(
679 base,
680 None,
681 &self.type_id_to_entity_name,
682 &prop_names,
683 ));
684 }
685 if self.is_edge_var(base) {
686 let props = self.edge_prop_names(base);
687 let value =
688 relationship_value_struct(base, col(format!("{base}.rel_type_name")), &props);
689 return Ok(null_unless(edge_present_qual(base), value));
690 }
691 }
692 self.lower(id)
693 }
694
695 fn lower_path_builtin_arg(&self, id: ExprId) -> Result<DfExpr, LoweringError> {
696 if let IrExpr::VarRef(v) = self.arena.get(id) {
697 return Ok(col_literal(
698 self.var_map.get(*v).ok_or(LoweringError::UnboundVar(v.0))?,
699 ));
700 }
701 self.lower(id)
702 }
703
704 fn lower_subscript(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
705 let [base_id, key_id] = args else {
706 return Err(LoweringError::UnsupportedExpr(
707 "_subscript expects two arguments".into(),
708 ));
709 };
710 let key_expr = self.lower(*key_id)?;
711 if let Some(key) = const_string_key(&key_expr) {
712 match key {
713 ConstStringKey::Null => return Ok(lit(ScalarValue::Null)),
714 ConstStringKey::Value(k) => {
715 if let Some(value) = self.lower_static_value_access(*base_id, &k)? {
716 return Ok(value);
717 }
718 }
719 }
720 }
721 if let Some(container) = self.lower_dynamic_access_container(*base_id)? {
722 return Ok(CYPHER_VALUE_ACCESS.call(vec![container, key_expr]));
723 }
724 let base = self.lower(*base_id)?;
725 if self.expr_data_type(&base).is_some_and(|dt| {
726 !matches!(
727 dt,
728 DataType::Null
729 | DataType::List(_)
730 | DataType::LargeList(_)
731 | DataType::FixedSizeList(_, _)
732 ) && !is_plain_map_struct_type(&dt)
733 && !is_het_struct_type(Some(&dt))
734 }) {
735 return Err(LoweringError::InvalidType(
736 "subscript requires a list, map, node, relationship, or null".into(),
737 ));
738 }
739 if self.is_list_typed(&base) {
740 match self.expr_data_type(&key_expr) {
741 Some(dt) if is_integer_data_type(&dt) => {
742 return Ok(datafusion::functions_nested::expr_fn::array_element(
743 base,
744 one_based_index(key_expr),
745 ));
746 }
747 Some(DataType::Null) | None => {}
748 Some(_) => {
749 return Err(LoweringError::InvalidType(
750 "list subscript index must be an integer or null".into(),
751 ));
752 }
753 }
754 }
755 Ok(CYPHER_VALUE_ACCESS.call(vec![base, key_expr]))
756 }
757
758 fn lower_static_value_access(
759 &self,
760 base: ExprId,
761 key: &str,
762 ) -> Result<Option<DfExpr>, LoweringError> {
763 if matches!(self.arena.get(base), IrExpr::Literal(IrLiteral::Null)) {
764 return Ok(Some(lit(ScalarValue::Null)));
765 }
766 if let Some(value) = self.lower_static_indexed_value_access(base, key)? {
767 return Ok(Some(value));
768 }
769 if let Some(value) = self.lower_map_literal_static_field(base, key)? {
770 return Ok(Some(value));
771 }
772 if let IrExpr::VarRef(var_id) = self.arena.get(base)
773 && let Some(base_name) = self.var_map.get(*var_id)
774 {
775 if let Some(value) = self.lower_entity_static_property(*var_id, base_name, key) {
776 return Ok(Some(value));
777 }
778 if let Some(value) = self.lower_struct_static_field(col_literal(base_name), key) {
779 return Ok(Some(value));
780 }
781 }
782
783 let base_expr = self.lower(base)?;
784 Ok(self.lower_struct_static_field(base_expr, key))
785 }
786
787 fn lower_map_literal_static_field(
788 &self,
789 base: ExprId,
790 key: &str,
791 ) -> Result<Option<DfExpr>, LoweringError> {
792 let IrExpr::MapLiteral(entries) = self.arena.get(base) else {
793 return Ok(None);
794 };
795 entries
796 .iter()
797 .find(|(k, _)| k == key)
798 .map_or(Ok(Some(lit(ScalarValue::Null))), |(_, id)| {
799 self.lower(*id).map(Some)
800 })
801 }
802
803 fn lower_entity_static_property(&self, var_id: VarId, base: &str, key: &str) -> Option<DfExpr> {
804 let qual = datafusion::common::TableReference::bare(base);
805 if let Some(schema) = self.input_schema.as_ref() {
806 if schema.index_of_column_by_name(Some(&qual), key).is_some() {
807 return Some(qualified_col(base, key));
808 }
809 if self.is_node_var(base) || self.is_edge_var(base) {
810 return Some(lit(ScalarValue::Null));
811 }
812 }
813 if let Some(shape) = self.node_shapes.get(&var_id.0) {
814 let is_topology = graphforge_storage::TOPOLOGY_NODES_SCHEMA
815 .field_with_name(key)
816 .is_ok();
817 if is_topology || shape.prop_names.iter().any(|p| p == key) {
818 return Some(qualified_col(base, key));
819 }
820 if self.props_authoritative {
821 return Some(lit(ScalarValue::Null));
822 }
823 }
824 None
825 }
826
827 fn lower_static_indexed_value_access(
828 &self,
829 base: ExprId,
830 key: &str,
831 ) -> Result<Option<DfExpr>, LoweringError> {
832 let IrExpr::FunctionCall { name, args } = self.arena.get(base) else {
833 return Ok(None);
834 };
835 if name != "_subscript" {
836 return Ok(None);
837 }
838 let [list_id, index_id] = args.as_slice() else {
839 return Ok(None);
840 };
841 let IrExpr::ListLiteral(items) = self.arena.get(*list_id) else {
842 return Ok(None);
843 };
844 let IrExpr::Literal(IrLiteral::Int(idx)) = self.arena.get(*index_id) else {
845 return Ok(None);
846 };
847 let len = i64::try_from(items.len()).map_err(|_| {
848 LoweringError::UnsupportedExpr("list literal length exceeds i64 range".into())
849 })?;
850 let pos = if *idx < 0 { len + idx } else { *idx };
851 if pos < 0 || pos >= len {
852 return Ok(Some(lit(ScalarValue::Null)));
853 }
854 let pos = usize::try_from(pos)
855 .map_err(|_| LoweringError::UnsupportedExpr("list index exceeds usize".into()))?;
856 self.lower_static_value_access(items[pos], key)
857 }
858
859 fn lower_struct_static_field(&self, base_expr: DfExpr, key: &str) -> Option<DfExpr> {
860 let dt = self.expr_data_type(&base_expr)?;
861 match dt {
862 DataType::Null => Some(lit(ScalarValue::Null)),
863 dt if is_het_struct_type(Some(&dt)) => Some(
864 ScalarUDF::new_from_impl(CypherStaticValueAccess::new(key.to_owned()))
865 .call(vec![base_expr]),
866 ),
867 DataType::Struct(fields)
868 if is_plain_map_struct_type(&DataType::Struct(fields.clone())) =>
869 {
870 if fields.iter().any(|f| f.name() == key) {
871 Some(datafusion::functions::core::expr_fn::get_field(
872 base_expr,
873 key.to_owned(),
874 ))
875 } else {
876 Some(lit(ScalarValue::Null))
877 }
878 }
879 DataType::Struct(fields) => {
880 let is_entity = fields
881 .iter()
882 .any(|field| matches!(field.name().as_str(), "node_uuid" | "edge_uuid"));
883 if is_entity && fields.iter().any(|field| field.name() == key) {
884 Some(datafusion::functions::core::expr_fn::get_field(
885 base_expr,
886 key.to_owned(),
887 ))
888 } else if is_entity {
889 Some(lit(ScalarValue::Null))
890 } else {
891 None
892 }
893 }
894 _ => None,
895 }
896 }
897
898 fn lower_dynamic_access_container(
899 &self,
900 base: ExprId,
901 ) -> Result<Option<DfExpr>, LoweringError> {
902 if matches!(self.arena.get(base), IrExpr::Literal(IrLiteral::Null)) {
903 return Ok(Some(lit(ScalarValue::Null)));
904 }
905 if matches!(self.arena.get(base), IrExpr::MapLiteral(_)) {
906 return self.lower(base).map(Some);
907 }
908 if let IrExpr::VarRef(var_id) = self.arena.get(base)
909 && let Some(base_name) = self.var_map.get(*var_id)
910 {
911 if let Some(expr) = self.entity_property_bag(*var_id, base_name) {
912 return Ok(Some(expr));
913 }
914 let base_expr = col_literal(base_name);
915 if self.expr_data_type(&base_expr).is_some_and(|dt| {
916 matches!(dt, DataType::Null)
917 || is_plain_map_struct_type(&dt)
918 || is_het_struct_type(Some(&dt))
919 }) {
920 return Ok(Some(base_expr));
921 }
922 }
923 let base_expr = self.lower(base)?;
924 Ok(self
925 .expr_data_type(&base_expr)
926 .is_some_and(|dt| {
927 matches!(dt, DataType::Null)
928 || is_plain_map_struct_type(&dt)
929 || is_het_struct_type(Some(&dt))
930 })
931 .then_some(base_expr))
932 }
933
934 fn entity_property_bag(&self, var_id: VarId, base: &str) -> Option<DfExpr> {
935 self.entity_property_bag_inner(var_id, base, false)
936 }
937
938 fn entity_property_bag_with_empty(&self, var_id: VarId, base: &str) -> Option<DfExpr> {
939 self.entity_property_bag_inner(var_id, base, true)
940 }
941
942 fn entity_property_bag_inner(
943 &self,
944 var_id: VarId,
945 base: &str,
946 empty_map_for_present_entity: bool,
947 ) -> Option<DfExpr> {
948 use datafusion::functions::core::expr_fn::named_struct;
949
950 let (prop_names, present) = if let Some(shape) = self.node_shapes.get(&var_id.0) {
951 let has_node_uuid = self.input_schema.as_ref().is_some_and(|schema| {
952 let qual = datafusion::common::TableReference::bare(base);
953 schema
954 .index_of_column_by_name(Some(&qual), "node_uuid")
955 .is_some()
956 });
957 let present = if has_node_uuid {
958 col(format!("{base}.node_uuid")).is_not_null()
959 } else {
960 lit(true)
961 };
962 (shape.prop_names.clone(), present)
963 } else if self.is_edge_var(base) {
964 (self.edge_prop_names(base), edge_present_qual(base))
965 } else {
966 return None;
967 };
968 if prop_names.is_empty() {
969 let value = if empty_map_for_present_entity {
970 ScalarUDF::new_from_impl(CypherEntityProperties::new(1)).call(vec![present.clone()])
971 } else {
972 lit(ScalarValue::Null)
973 };
974 return Some(null_unless(present, value));
975 }
976 if empty_map_for_present_entity {
977 let mut args = Vec::with_capacity(1 + prop_names.len() * 2);
978 args.push(present);
979 for prop in prop_names {
980 args.push(lit(prop.as_str()));
981 args.push(qualified_col(base, &prop));
982 }
983 return Some(
984 ScalarUDF::new_from_impl(CypherEntityProperties::new(args.len())).call(args),
985 );
986 }
987 let mut args = Vec::with_capacity(prop_names.len() * 2);
988 for prop in prop_names {
989 args.push(lit(prop.as_str()));
990 args.push(qualified_col(base, &prop));
991 }
992 Some(null_unless(present, named_struct(args)))
993 }
994
995 fn lower_quantifier(
1002 &self,
1003 kind: graphforge_ir::QuantifierKind,
1004 loop_var: VarId,
1005 list: ExprId,
1006 predicate: ExprId,
1007 ) -> Result<DfExpr, LoweringError> {
1008 let elem_name = match self.elem_struct_cols.len() {
1012 0 => "__gf_elem".to_owned(),
1013 d => format!("__gf_elem_{d}"),
1014 };
1015 let list_expr = self.lower(list)?;
1016 if let IrExpr::Literal(IrLiteral::Bool(predicate)) = self.arena.get(predicate) {
1017 let udf =
1018 ScalarUDF::new_from_impl(CypherInvariantQuantifier::new(kind, Some(*predicate)));
1019 return Ok(udf.call(vec![list_expr]));
1020 }
1021 if matches!(self.arena.get(predicate), IrExpr::Literal(IrLiteral::Null)) {
1022 let udf = ScalarUDF::new_from_impl(CypherInvariantQuantifier::new(kind, None));
1023 return Ok(udf.call(vec![list_expr]));
1024 }
1025 let mut elem_vars = VarMap::new();
1028 for v in self.var_map.var_ids() {
1029 if let Some(name) = self.var_map.get(v) {
1030 elem_vars.insert(v, name.to_owned());
1031 }
1032 }
1033 elem_vars.insert(loop_var, elem_name.as_str());
1034 let pred_lowerer = {
1035 let mut l =
1036 ExprLowerer::with_prop_names(self.arena, &elem_vars, self.prop_names.clone());
1037 if let Some(s) = self.input_schema.as_ref() {
1038 l = l.with_input_schema(s.clone());
1039 }
1040 if let Some(t) = self.read_target.as_ref() {
1041 l = l.with_read_target(t.clone());
1042 }
1043 for c in &self.elem_struct_cols {
1046 l = l.with_elem_struct_col(c.clone());
1047 }
1048 l.with_elem_struct_col(elem_name.clone())
1049 };
1050 let pred_expr = pred_lowerer.lower(predicate)?;
1051
1052 let mut outer: Vec<String> = pred_expr
1056 .column_refs()
1057 .into_iter()
1058 .map(|c| c.name.clone())
1059 .filter(|n| n != &elem_name)
1060 .collect();
1061 outer.sort();
1062 outer.dedup();
1063
1064 if outer.is_empty()
1074 && !is_empty_list_literal(&list_expr)
1075 && let Some(elem_type) = self.list_element_type(&list_expr)
1076 {
1077 use datafusion::arrow::datatypes::{Field, Schema};
1078 use datafusion::common::DFSchema;
1079 use datafusion::logical_expr::execution_props::ExecutionProps;
1080 use datafusion::physical_expr::create_physical_expr;
1081 let schema = Schema::new(vec![Field::new(&elem_name, elem_type, true)]);
1082 if let Ok(df_schema) = DFSchema::try_from(schema)
1083 && create_physical_expr(&pred_expr, &df_schema, &ExecutionProps::new()).is_err()
1084 {
1085 return Err(LoweringError::InvalidType(format!(
1086 "quantifier predicate cannot apply to the list's element type ({kind:?})"
1087 )));
1088 }
1089 }
1090
1091 let mut call_args = Vec::with_capacity(1 + outer.len());
1092 call_args.push(list_expr);
1093 for name in &outer {
1094 call_args.push(col_literal(name));
1095 }
1096 let udf =
1097 ScalarUDF::new_from_impl(CypherQuantifier::new(kind, pred_expr, elem_name, outer));
1098 Ok(udf.call(call_args))
1099 }
1100
1101 #[allow(
1107 clippy::too_many_lines,
1108 reason = "schema synthesis, correlation rebinding, and UDF construction stay aligned"
1109 )]
1110 fn lower_list_comprehension(
1111 &self,
1112 loop_var: VarId,
1113 list: ExprId,
1114 filter: Option<ExprId>,
1115 projection: Option<ExprId>,
1116 ) -> Result<DfExpr, LoweringError> {
1117 let elem_name = match self.elem_struct_cols.len() {
1121 0 => "__gf_elem".to_owned(),
1122 d => format!("__gf_elem_{d}"),
1123 };
1124 let list_expr = self.lower(list)?;
1125 let clause_schema = self.list_element_type(&list_expr).and_then(|element_type| {
1126 let mut fields = self.input_schema.as_ref().map_or_else(Vec::new, |schema| {
1127 schema
1128 .iter()
1129 .map(|(qualifier, field)| (qualifier.cloned(), Arc::clone(field)))
1130 .collect()
1131 });
1132 fields.push((
1133 None,
1134 Arc::new(datafusion::arrow::datatypes::Field::new(
1135 &elem_name,
1136 element_type,
1137 true,
1138 )),
1139 ));
1140 datafusion::common::DFSchema::new_with_metadata(fields, HashMap::new())
1141 .ok()
1142 .map(Arc::new)
1143 });
1144
1145 let mut elem_vars = VarMap::new();
1148 for v in self.var_map.var_ids() {
1149 if let Some(name) = self.var_map.get(v) {
1150 elem_vars.insert(v, name.to_owned());
1151 }
1152 }
1153 elem_vars.insert(loop_var, elem_name.as_str());
1154 let clause_lowerer = {
1155 let mut l =
1156 ExprLowerer::with_prop_names(self.arena, &elem_vars, self.prop_names.clone());
1157 if let Some(s) = clause_schema.as_ref().or(self.input_schema.as_ref()) {
1158 l = l.with_input_schema(s.clone());
1159 }
1160 if let Some(t) = self.read_target.as_ref() {
1161 l = l.with_read_target(t.clone());
1162 }
1163 for c in &self.elem_struct_cols {
1165 l = l.with_elem_struct_col(c.clone());
1166 }
1167 l.with_elem_struct_col(elem_name.clone())
1168 };
1169 let mut filter_expr = filter.map(|f| clause_lowerer.lower(f)).transpose()?;
1170 let mut projection_expr = projection.map(|p| clause_lowerer.lower(p)).transpose()?;
1171
1172 let mut outer_columns = Vec::new();
1174 for e in [filter_expr.as_ref(), projection_expr.as_ref()]
1175 .into_iter()
1176 .flatten()
1177 {
1178 for c in e.column_refs() {
1179 if c.name != elem_name && !outer_columns.contains(c) {
1180 outer_columns.push(c.clone());
1181 }
1182 }
1183 }
1184 outer_columns.sort_by_key(datafusion::common::Column::flat_name);
1185 let outer = (0..outer_columns.len())
1186 .map(|index| format!("__gf_outer_{index}"))
1187 .collect::<Vec<_>>();
1188 if !outer_columns.is_empty() {
1189 use datafusion::common::tree_node::{Transformed, TreeNode};
1190 let rewrite = |expr: DfExpr| {
1191 expr.transform_up(|expr| {
1192 let DfExpr::Column(column) = &expr else {
1193 return Ok(Transformed::no(expr));
1194 };
1195 let Some(index) = outer_columns.iter().position(|outer| outer == column) else {
1196 return Ok(Transformed::no(expr));
1197 };
1198 Ok(Transformed::yes(DfExpr::Column(
1199 datafusion::common::Column::from_name(outer[index].clone()),
1200 )))
1201 })
1202 .map(|transformed| transformed.data)
1203 };
1204 filter_expr = filter_expr
1205 .map(&rewrite)
1206 .transpose()
1207 .map_err(|error| LoweringError::UnsupportedExpr(error.to_string()))?;
1208 projection_expr = projection_expr
1209 .map(rewrite)
1210 .transpose()
1211 .map_err(|error| LoweringError::UnsupportedExpr(error.to_string()))?;
1212 }
1213
1214 if let Some(fexpr) = filter_expr.as_ref()
1221 && outer.is_empty()
1222 && let Some(elem_type) = self.list_element_type(&list_expr)
1223 {
1224 use datafusion::arrow::datatypes::{Field, Schema};
1225 use datafusion::common::DFSchema;
1226 use datafusion::logical_expr::execution_props::ExecutionProps;
1227 use datafusion::physical_expr::create_physical_expr;
1228 let schema = Schema::new(vec![Field::new(&elem_name, elem_type, true)]);
1229 if let Ok(df_schema) = DFSchema::try_from(schema)
1230 && create_physical_expr(fexpr, &df_schema, &ExecutionProps::new()).is_err()
1231 {
1232 return Err(LoweringError::InvalidType(
1233 "list comprehension filter cannot apply to the list's element type".into(),
1234 ));
1235 }
1236 }
1237
1238 let mut call_args = Vec::with_capacity(1 + outer.len());
1239 call_args.push(list_expr);
1240 for column in &outer_columns {
1241 call_args.push(DfExpr::Column(column.clone()));
1242 }
1243 let udf = ScalarUDF::new_from_impl(CypherListComp::new(
1244 filter_expr,
1245 projection_expr,
1246 elem_name,
1247 outer,
1248 ));
1249 Ok(udf.call(call_args))
1250 }
1251
1252 fn lower_map_literal(&self, entries: &[(String, ExprId)]) -> Result<DfExpr, LoweringError> {
1263 use datafusion::functions::core::expr_fn::named_struct;
1264 if entries.is_empty() {
1265 return Ok(empty_map_struct());
1267 }
1268 let mut args: Vec<DfExpr> = Vec::with_capacity(entries.len() * 2);
1269 for (key, value) in entries {
1270 args.push(lit(key.as_str()));
1271 args.push(self.lower_value(*value)?);
1272 }
1273 Ok(named_struct(args))
1274 }
1275
1276 fn sole_arg_is_null(&self, args: &[ExprId]) -> bool {
1289 matches!(args, [a] if matches!(self.arena.get(*a), IrExpr::Literal(graphforge_ir::IrLiteral::Null)))
1290 }
1291
1292 fn lower_clock_now(&self, name: &str) -> DfExpr {
1300 use chrono::Timelike;
1301 let now = *self.now.get_or_init(|| chrono::Utc::now().naive_utc());
1302 let days = crate::temporal::date_to_epoch_days(now.date()).unwrap_or(0);
1303 let nanos = i64::from(now.time().num_seconds_from_midnight()) * 1_000_000_000
1304 + i64::from(now.time().nanosecond());
1305 let scalar = match name {
1306 "date" => date_scalar(Some(days)),
1307 "localtime" => ScalarValue::Time64Nanosecond(Some(nanos)),
1308 "localdatetime" => localdatetime_scalar(Some((days, nanos))),
1309 "time" => time_scalar(Some((nanos, 0))),
1310 _ => datetime_scalar(Some((days, nanos, 0, None))),
1312 };
1313 lit(scalar)
1314 }
1315
1316 fn lower_temporal(&self, name: &str, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
1317 if self.sole_arg_is_null(args) {
1320 return Ok(DfExpr::Literal(temporal_null_scalar(name), None));
1321 }
1322 if args.is_empty() && name != "duration" {
1325 return Ok(self.lower_clock_now(name));
1326 }
1327 if let [arg] = args {
1328 if name == "date" {
1333 if let Some(days) = self.const_date(*arg) {
1334 return Ok(DfExpr::Literal(date_scalar(Some(days)), None));
1335 }
1336 if let Some(projected) = self.lower_date_runtime(*arg)? {
1340 return Ok(projected);
1341 }
1342 }
1343 if name == "localtime" {
1347 if let Some(nanos) = self.const_local_time(*arg) {
1348 return Ok(DfExpr::Literal(
1349 ScalarValue::Time64Nanosecond(Some(nanos)),
1350 None,
1351 ));
1352 }
1353 if let Some(projected) = self.lower_localtime_runtime(*arg)? {
1354 return Ok(projected);
1355 }
1356 }
1357 if name == "localdatetime" {
1363 if let Some((days, nanos)) = self.const_local_date_time(*arg) {
1364 return Ok(DfExpr::Literal(
1365 localdatetime_scalar(Some((days, nanos))),
1366 None,
1367 ));
1368 }
1369 if let Some(projected) = self.lower_localdatetime_runtime(*arg)? {
1370 return Ok(projected);
1371 }
1372 }
1373 if name == "time" {
1377 if let Some((nanos, offset)) = self.const_time(*arg) {
1378 return Ok(DfExpr::Literal(time_scalar(Some((nanos, offset))), None));
1379 }
1380 if let Some(projected) = self.lower_time_runtime(*arg)? {
1381 return Ok(projected);
1382 }
1383 }
1384 if name == "datetime" {
1390 if let Some(parts) = self.const_datetime(*arg) {
1391 return Ok(DfExpr::Literal(datetime_scalar(Some(parts)), None));
1392 }
1393 if let Some(projected) = self.lower_datetime_runtime(*arg)? {
1394 return Ok(projected);
1395 }
1396 }
1397 if name == "duration" {
1402 if let Some(dur) = self.const_duration(*arg) {
1403 return Ok(DfExpr::Literal(duration_scalar(Some(dur)), None));
1404 }
1405 let lowered = self.lower(*arg)?;
1406 if self.is_string_typed(&lowered) {
1407 return Ok(CYPHER_DURATION_PARSE.call(vec![lowered]));
1408 }
1409 }
1410 match self.arena.get(*arg) {
1411 IrExpr::Literal(IrLiteral::Str(s)) => {
1412 if let Some(rendered) = render_temporal(name, s) {
1413 return Ok(lit(rendered));
1414 }
1415 }
1416 IrExpr::MapLiteral(entries) => {
1417 if let Some(fields) = self.extract_temporal_fields(entries)
1418 && let Some(rendered) = crate::temporal::render_temporal_map(name, &fields)
1419 {
1420 return Ok(lit(rendered));
1421 }
1422 }
1423 _ => {}
1424 }
1425 }
1426 let lowered: Vec<DfExpr> = args
1427 .iter()
1428 .map(|&a| self.lower(a))
1429 .collect::<Result<_, _>>()?;
1430 resolve_builtin(name, lowered, || self.path_node_hydration())
1431 .ok_or_else(|| LoweringError::UnknownFunction(name.to_string()))
1432 }
1433
1434 fn const_date(&self, arg: ExprId) -> Option<i64> {
1437 match self.arena.get(arg) {
1438 IrExpr::Literal(IrLiteral::Str(s)) => crate::temporal::parse_date_string(s),
1439 IrExpr::MapLiteral(entries) => {
1440 let fields = self.extract_temporal_fields(entries)?;
1441 crate::temporal::date_from_map(&fields)
1442 }
1443 _ => None,
1444 }
1445 }
1446
1447 fn lower_date_runtime(&self, arg: ExprId) -> Result<Option<DfExpr>, LoweringError> {
1454 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1455 let (base, overrides) = match self.arena.get(arg) {
1456 IrExpr::MapLiteral(entries) if entries.iter().any(|(k, _)| k == "date") => {
1457 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1458 let base = self.lower(field("date").expect("checked `date` key exists"))?;
1459 let ov = |name: &str| match field(name) {
1460 Some(id) => self.lower(id),
1461 None => Ok(null_i64()),
1462 };
1463 let overrides = [
1464 ov("year")?,
1465 ov("month")?,
1466 ov("day")?,
1467 ov("week")?,
1468 ov("dayOfWeek")?,
1469 ov("ordinalDay")?,
1470 ov("quarter")?,
1471 ov("dayOfQuarter")?,
1472 ];
1473 (base, overrides)
1474 }
1475 IrExpr::MapLiteral(_) => return Ok(None),
1478 _ => (self.lower(arg)?, std::array::from_fn(|_| null_i64())),
1480 };
1481 let mut call_args = Vec::with_capacity(9);
1482 call_args.push(base);
1483 call_args.extend(overrides);
1484 Ok(Some(CYPHER_DATE_PROJECT.call(call_args)))
1485 }
1486
1487 fn const_local_time(&self, arg: ExprId) -> Option<i64> {
1490 match self.arena.get(arg) {
1491 IrExpr::Literal(IrLiteral::Str(s)) => crate::temporal::localtime_nanos_from_str(s),
1492 IrExpr::MapLiteral(entries) => {
1493 let fields = self.extract_temporal_fields(entries)?;
1494 crate::temporal::localtime_nanos_from_map(&fields)
1495 }
1496 _ => None,
1497 }
1498 }
1499
1500 fn lower_localtime_runtime(&self, arg: ExprId) -> Result<Option<DfExpr>, LoweringError> {
1507 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1508 let (base, overrides) = match self.arena.get(arg) {
1509 IrExpr::MapLiteral(entries) if entries.iter().any(|(k, _)| k == "time") => {
1510 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1511 let base = self.lower(field("time").expect("checked `time` key exists"))?;
1512 let ov = |name: &str| match field(name) {
1513 Some(id) => self.lower(id),
1514 None => Ok(null_i64()),
1515 };
1516 let overrides = [
1517 ov("hour")?,
1518 ov("minute")?,
1519 ov("second")?,
1520 ov("millisecond")?,
1521 ov("microsecond")?,
1522 ov("nanosecond")?,
1523 ];
1524 (base, overrides)
1525 }
1526 IrExpr::MapLiteral(_) => return Ok(None),
1527 _ => (self.lower(arg)?, std::array::from_fn(|_| null_i64())),
1528 };
1529 let mut call_args = Vec::with_capacity(7);
1530 call_args.push(base);
1531 call_args.extend(overrides);
1532 Ok(Some(CYPHER_LOCALTIME_PROJECT.call(call_args)))
1533 }
1534
1535 fn const_local_date_time(&self, arg: ExprId) -> Option<(i64, i64)> {
1539 match self.arena.get(arg) {
1540 IrExpr::Literal(IrLiteral::Str(s)) => crate::temporal::localdatetime_parts_from_str(s),
1541 IrExpr::MapLiteral(entries) => {
1542 let fields = self.extract_temporal_fields(entries)?;
1543 crate::temporal::localdatetime_parts_from_map(&fields)
1544 }
1545 _ => None,
1546 }
1547 }
1548
1549 fn lower_localdatetime_runtime(&self, arg: ExprId) -> Result<Option<DfExpr>, LoweringError> {
1556 let null = || DfExpr::Literal(ScalarValue::Null, None);
1557 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1558 let (date_src, time_src, overrides) =
1559 if let IrExpr::MapLiteral(entries) = self.arena.get(arg) {
1560 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1561 let lower_or = |id: Option<ExprId>, default: &dyn Fn() -> DfExpr| match id {
1562 Some(id) => self.lower(id),
1563 None => Ok(default()),
1564 };
1565 let date_anchor = field("datetime").or_else(|| field("date"));
1568 let time_anchor = field("datetime").or_else(|| field("time"));
1569 let date_src = lower_or(date_anchor, &null)?;
1570 let time_src = lower_or(time_anchor, &null)?;
1571 let ov = |name: &str| lower_or(field(name), &null_i64);
1572 let overrides = [
1573 ov("year")?,
1574 ov("month")?,
1575 ov("day")?,
1576 ov("week")?,
1577 ov("dayOfWeek")?,
1578 ov("ordinalDay")?,
1579 ov("quarter")?,
1580 ov("dayOfQuarter")?,
1581 ov("hour")?,
1582 ov("minute")?,
1583 ov("second")?,
1584 ov("millisecond")?,
1585 ov("microsecond")?,
1586 ov("nanosecond")?,
1587 ];
1588 (date_src, time_src, overrides)
1589 } else {
1590 let base = self.lower(arg)?;
1593 (base.clone(), base, std::array::from_fn(|_| null_i64()))
1594 };
1595 let mut call_args = Vec::with_capacity(16);
1596 call_args.push(date_src);
1597 call_args.push(time_src);
1598 call_args.extend(overrides);
1599 Ok(Some(CYPHER_LOCALDATETIME_PROJECT.call(call_args)))
1600 }
1601
1602 fn const_time(&self, arg: ExprId) -> Option<(i64, i32)> {
1606 match self.arena.get(arg) {
1607 IrExpr::Literal(IrLiteral::Str(s)) => crate::temporal::time_value_from_str(s),
1608 IrExpr::MapLiteral(entries) => {
1609 let fields = self.extract_temporal_fields(entries)?;
1610 crate::temporal::time_value_from_map(&fields)
1611 }
1612 _ => None,
1613 }
1614 }
1615
1616 fn lower_time_runtime(&self, arg: ExprId) -> Result<Option<DfExpr>, LoweringError> {
1623 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1624 let null_str = || DfExpr::Literal(ScalarValue::Utf8(None), None);
1625 let (base, overrides, timezone) = match self.arena.get(arg) {
1626 IrExpr::MapLiteral(entries) if entries.iter().any(|(k, _)| k == "time") => {
1627 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1628 let base = self.lower(field("time").expect("checked `time` key exists"))?;
1629 let ov = |name: &str| match field(name) {
1630 Some(id) => self.lower(id),
1631 None => Ok(null_i64()),
1632 };
1633 let overrides = [
1634 ov("hour")?,
1635 ov("minute")?,
1636 ov("second")?,
1637 ov("millisecond")?,
1638 ov("microsecond")?,
1639 ov("nanosecond")?,
1640 ];
1641 let timezone = match field("timezone") {
1642 Some(id) => self.lower(id)?,
1643 None => null_str(),
1644 };
1645 (base, overrides, timezone)
1646 }
1647 IrExpr::MapLiteral(_) | IrExpr::Literal(IrLiteral::Str(_)) => return Ok(None),
1653 _ => (
1654 self.lower(arg)?,
1655 std::array::from_fn(|_| null_i64()),
1656 null_str(),
1657 ),
1658 };
1659 let mut call_args = Vec::with_capacity(8);
1660 call_args.push(base);
1661 call_args.extend(overrides);
1662 call_args.push(timezone);
1663 Ok(Some(CYPHER_TIME_PROJECT.call(call_args)))
1664 }
1665
1666 fn const_datetime(&self, arg: ExprId) -> Option<(i64, i64, i32, Option<String>)> {
1670 match self.arena.get(arg) {
1671 IrExpr::Literal(IrLiteral::Str(s)) => crate::temporal::datetime_value_from_str(s),
1672 IrExpr::MapLiteral(entries) => {
1673 let fields = self.extract_temporal_fields(entries)?;
1674 crate::temporal::datetime_value_from_map(&fields)
1675 }
1676 _ => None,
1677 }
1678 }
1679
1680 fn const_duration(&self, arg: ExprId) -> Option<crate::temporal::DurationValue> {
1683 match self.arena.get(arg) {
1684 IrExpr::Literal(IrLiteral::Str(s)) => crate::temporal::duration_value_from_str(s),
1685 IrExpr::MapLiteral(entries) => {
1686 let fields = self.extract_temporal_fields(entries)?;
1687 crate::temporal::duration_value_from_map(&fields)
1688 }
1689 _ => None,
1690 }
1691 }
1692
1693 fn lower_datetime_runtime(&self, arg: ExprId) -> Result<Option<DfExpr>, LoweringError> {
1700 let null = || DfExpr::Literal(ScalarValue::Null, None);
1701 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1702 let null_str = || DfExpr::Literal(ScalarValue::Utf8(None), None);
1703 let (date_src, time_src, overrides, timezone) = match self.arena.get(arg) {
1704 IrExpr::MapLiteral(entries) => {
1705 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1706 let lower_or = |id: Option<ExprId>, default: &dyn Fn() -> DfExpr| match id {
1707 Some(id) => self.lower(id),
1708 None => Ok(default()),
1709 };
1710 let date_anchor = field("datetime").or_else(|| field("date"));
1711 let time_anchor = field("datetime").or_else(|| field("time"));
1712 let date_src = lower_or(date_anchor, &null)?;
1713 let time_src = lower_or(time_anchor, &null)?;
1714 let ov = |name: &str| lower_or(field(name), &null_i64);
1715 let overrides = [
1716 ov("year")?,
1717 ov("month")?,
1718 ov("day")?,
1719 ov("week")?,
1720 ov("dayOfWeek")?,
1721 ov("ordinalDay")?,
1722 ov("quarter")?,
1723 ov("dayOfQuarter")?,
1724 ov("hour")?,
1725 ov("minute")?,
1726 ov("second")?,
1727 ov("millisecond")?,
1728 ov("microsecond")?,
1729 ov("nanosecond")?,
1730 ];
1731 (
1732 date_src,
1733 time_src,
1734 overrides,
1735 lower_or(field("timezone"), &null_str)?,
1736 )
1737 }
1738 IrExpr::Literal(IrLiteral::Str(_)) => return Ok(None),
1741 _ => {
1742 let base = self.lower(arg)?;
1743 (
1744 base.clone(),
1745 base,
1746 std::array::from_fn(|_| null_i64()),
1747 null_str(),
1748 )
1749 }
1750 };
1751 let mut call_args = Vec::with_capacity(17);
1752 call_args.push(date_src);
1753 call_args.push(time_src);
1754 call_args.extend(overrides);
1755 call_args.push(timezone);
1756 Ok(Some(CYPHER_DATETIME_PROJECT.call(call_args)))
1757 }
1758
1759 fn lower_date_truncate(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
1763 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1764 let [unit_id, value_id, rest @ ..] = args else {
1765 return Err(LoweringError::UnknownFunction("date.truncate".to_string()));
1766 };
1767 let unit = self.lower(*unit_id)?;
1768 let value = self.lower(*value_id)?;
1769 let overrides: [DfExpr; 8] = match rest.first() {
1774 None => std::array::from_fn(|_| null_i64()),
1775 Some(map_id) => {
1776 let IrExpr::MapLiteral(entries) = self.arena.get(*map_id) else {
1777 return Err(LoweringError::UnsupportedExpr(
1778 "date.truncate override map must be a literal map".to_string(),
1779 ));
1780 };
1781 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1782 let ov = |name: &str| match field(name) {
1783 Some(id) => self.lower(id),
1784 None => Ok(null_i64()),
1785 };
1786 [
1787 ov("year")?,
1788 ov("month")?,
1789 ov("day")?,
1790 ov("week")?,
1791 ov("dayOfWeek")?,
1792 ov("ordinalDay")?,
1793 ov("quarter")?,
1794 ov("dayOfQuarter")?,
1795 ]
1796 }
1797 };
1798 let mut call_args = Vec::with_capacity(10);
1799 call_args.push(value);
1800 call_args.push(unit);
1801 call_args.extend(overrides);
1802 Ok(CYPHER_DATE_TRUNCATE.call(call_args))
1803 }
1804
1805 fn lower_localtime_truncate(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
1809 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1810 let [unit_id, value_id, rest @ ..] = args else {
1811 return Err(LoweringError::UnknownFunction(
1812 "localtime.truncate".to_string(),
1813 ));
1814 };
1815 let unit = self.lower(*unit_id)?;
1816 let value = self.lower(*value_id)?;
1817 let overrides: [DfExpr; 6] = match rest.first() {
1819 None => std::array::from_fn(|_| null_i64()),
1820 Some(map_id) => {
1821 let IrExpr::MapLiteral(entries) = self.arena.get(*map_id) else {
1822 return Err(LoweringError::UnsupportedExpr(
1823 "localtime.truncate override map must be a literal map".to_string(),
1824 ));
1825 };
1826 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1827 let ov = |name: &str| match field(name) {
1828 Some(id) => self.lower(id),
1829 None => Ok(null_i64()),
1830 };
1831 [
1832 ov("hour")?,
1833 ov("minute")?,
1834 ov("second")?,
1835 ov("millisecond")?,
1836 ov("microsecond")?,
1837 ov("nanosecond")?,
1838 ]
1839 }
1840 };
1841 let mut call_args = Vec::with_capacity(8);
1842 call_args.push(value);
1843 call_args.push(unit);
1844 call_args.extend(overrides);
1845 Ok(CYPHER_LOCALTIME_TRUNCATE.call(call_args))
1846 }
1847
1848 fn lower_localdatetime_truncate(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
1853 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1854 let [unit_id, value_id, rest @ ..] = args else {
1855 return Err(LoweringError::UnknownFunction(
1856 "localdatetime.truncate".to_string(),
1857 ));
1858 };
1859 let unit = self.lower(*unit_id)?;
1860 let value = self.lower(*value_id)?;
1861 let overrides: [DfExpr; 14] = match rest.first() {
1864 None => std::array::from_fn(|_| null_i64()),
1865 Some(map_id) => {
1866 let IrExpr::MapLiteral(entries) = self.arena.get(*map_id) else {
1867 return Err(LoweringError::UnsupportedExpr(
1868 "localdatetime.truncate override map must be a literal map".to_string(),
1869 ));
1870 };
1871 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1872 let ov = |name: &str| match field(name) {
1873 Some(id) => self.lower(id),
1874 None => Ok(null_i64()),
1875 };
1876 [
1877 ov("year")?,
1878 ov("month")?,
1879 ov("day")?,
1880 ov("week")?,
1881 ov("dayOfWeek")?,
1882 ov("ordinalDay")?,
1883 ov("quarter")?,
1884 ov("dayOfQuarter")?,
1885 ov("hour")?,
1886 ov("minute")?,
1887 ov("second")?,
1888 ov("millisecond")?,
1889 ov("microsecond")?,
1890 ov("nanosecond")?,
1891 ]
1892 }
1893 };
1894 let mut call_args = Vec::with_capacity(16);
1895 call_args.push(value);
1896 call_args.push(unit);
1897 call_args.extend(overrides);
1898 Ok(CYPHER_LOCALDATETIME_TRUNCATE.call(call_args))
1899 }
1900
1901 fn lower_time_truncate(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
1906 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1907 let null_str = || DfExpr::Literal(ScalarValue::Utf8(None), None);
1908 let [unit_id, value_id, rest @ ..] = args else {
1909 return Err(LoweringError::UnknownFunction("time.truncate".to_string()));
1910 };
1911 let unit = self.lower(*unit_id)?;
1912 let value = self.lower(*value_id)?;
1913 let (overrides, timezone): ([DfExpr; 6], DfExpr) = match rest.first() {
1914 None => (std::array::from_fn(|_| null_i64()), null_str()),
1915 Some(map_id) => {
1916 let IrExpr::MapLiteral(entries) = self.arena.get(*map_id) else {
1917 return Err(LoweringError::UnsupportedExpr(
1918 "time.truncate override map must be a literal map".to_string(),
1919 ));
1920 };
1921 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1922 let ov = |name: &str| match field(name) {
1923 Some(id) => self.lower(id),
1924 None => Ok(null_i64()),
1925 };
1926 let tz = match field("timezone") {
1927 Some(id) => self.lower(id)?,
1928 None => null_str(),
1929 };
1930 (
1931 [
1932 ov("hour")?,
1933 ov("minute")?,
1934 ov("second")?,
1935 ov("millisecond")?,
1936 ov("microsecond")?,
1937 ov("nanosecond")?,
1938 ],
1939 tz,
1940 )
1941 }
1942 };
1943 let mut call_args = Vec::with_capacity(9);
1944 call_args.push(value);
1945 call_args.push(unit);
1946 call_args.extend(overrides);
1947 call_args.push(timezone);
1948 Ok(CYPHER_TIME_TRUNCATE.call(call_args))
1949 }
1950
1951 fn lower_datetime_truncate(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
1956 let null_i64 = || DfExpr::Literal(ScalarValue::Int64(None), None);
1957 let null_str = || DfExpr::Literal(ScalarValue::Utf8(None), None);
1958 let [unit_id, value_id, rest @ ..] = args else {
1959 return Err(LoweringError::UnknownFunction(
1960 "datetime.truncate".to_string(),
1961 ));
1962 };
1963 let unit = self.lower(*unit_id)?;
1964 let value = self.lower(*value_id)?;
1965 let (overrides, timezone): ([DfExpr; 14], DfExpr) = match rest.first() {
1966 None => (std::array::from_fn(|_| null_i64()), null_str()),
1967 Some(map_id) => {
1968 let IrExpr::MapLiteral(entries) = self.arena.get(*map_id) else {
1969 return Err(LoweringError::UnsupportedExpr(
1970 "datetime.truncate override map must be a literal map".to_string(),
1971 ));
1972 };
1973 let field = |name: &str| entries.iter().find(|(k, _)| k == name).map(|(_, v)| *v);
1974 let ov = |name: &str| match field(name) {
1975 Some(id) => self.lower(id),
1976 None => Ok(null_i64()),
1977 };
1978 let tz = match field("timezone") {
1979 Some(id) => self.lower(id)?,
1980 None => null_str(),
1981 };
1982 (
1983 [
1984 ov("year")?,
1985 ov("month")?,
1986 ov("day")?,
1987 ov("week")?,
1988 ov("dayOfWeek")?,
1989 ov("ordinalDay")?,
1990 ov("quarter")?,
1991 ov("dayOfQuarter")?,
1992 ov("hour")?,
1993 ov("minute")?,
1994 ov("second")?,
1995 ov("millisecond")?,
1996 ov("microsecond")?,
1997 ov("nanosecond")?,
1998 ],
1999 tz,
2000 )
2001 }
2002 };
2003 let mut call_args = Vec::with_capacity(17);
2004 call_args.push(value);
2005 call_args.push(unit);
2006 call_args.extend(overrides);
2007 call_args.push(timezone);
2008 Ok(CYPHER_DATETIME_TRUNCATE.call(call_args))
2009 }
2010
2011 fn lower_duration_between(&self, name: &str, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2015 let [a_id, b_id] = args else {
2016 return Err(LoweringError::UnknownFunction(name.to_string()));
2017 };
2018 let a = self.lower(*a_id)?;
2019 let b = self.lower(*b_id)?;
2020 Ok(CYPHER_DURATION_BETWEEN.call(vec![a, b, lit(name)]))
2021 }
2022
2023 fn extract_temporal_fields(
2024 &self,
2025 entries: &[(String, ExprId)],
2026 ) -> Option<std::collections::HashMap<String, crate::temporal::TemporalField>> {
2027 let mut fields = std::collections::HashMap::with_capacity(entries.len());
2028 for (key, value) in entries {
2029 fields.insert(key.clone(), self.extract_temporal_field(*value)?);
2030 }
2031 Some(fields)
2032 }
2033
2034 fn extract_temporal_field(&self, id: ExprId) -> Option<crate::temporal::TemporalField> {
2037 use crate::temporal::TemporalField;
2038 use graphforge_ir::expr::UnaryOpKind;
2039 match self.arena.get(id) {
2040 IrExpr::Literal(IrLiteral::Int(n)) => Some(TemporalField::Int(*n)),
2041 IrExpr::Literal(IrLiteral::Float(x)) => Some(TemporalField::Float(*x)),
2042 IrExpr::Literal(IrLiteral::Str(s)) => Some(TemporalField::Str(s.clone())),
2043 IrExpr::UnaryOp {
2046 op: UnaryOpKind::Neg,
2047 expr,
2048 } => match self.arena.get(*expr) {
2049 IrExpr::Literal(IrLiteral::Int(n)) => Some(TemporalField::Int(-n)),
2050 IrExpr::Literal(IrLiteral::Float(x)) => Some(TemporalField::Float(-x)),
2051 _ => None,
2052 },
2053 IrExpr::FunctionCall { name, args } if name == "date" => {
2054 if let [a] = args.as_slice()
2055 && let IrExpr::Literal(IrLiteral::Str(s)) = self.arena.get(*a)
2056 {
2057 crate::temporal::parse_date_string(s).map(TemporalField::Date)
2058 } else {
2059 None
2060 }
2061 }
2062 _ => None,
2063 }
2064 }
2065
2066 fn lower_from_epoch(&self, name: &str, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2070 self.try_from_epoch(name, args)
2071 .map(lit)
2072 .ok_or_else(|| LoweringError::UnknownFunction(name.to_string()))
2073 }
2074
2075 fn try_from_epoch(&self, name: &str, args: &[ExprId]) -> Option<String> {
2076 match (name, args) {
2077 ("datetime.fromepoch", &[a, b]) => {
2078 crate::temporal::render_from_epoch(self.int_literal(a)?, self.int_literal(b)?)
2079 }
2080 ("datetime.fromepochmillis", &[a]) => {
2081 crate::temporal::render_from_epoch_millis(self.int_literal(a)?)
2082 }
2083 _ => None,
2084 }
2085 }
2086
2087 fn int_literal(&self, id: ExprId) -> Option<i64> {
2089 match self.arena.get(id) {
2090 IrExpr::Literal(IrLiteral::Int(n)) => Some(*n),
2091 _ => None,
2092 }
2093 }
2094
2095 fn lower_node_struct(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2105 let Some(&base_id) = args.first() else {
2108 return Err(LoweringError::UnsupportedExpr(
2109 "_node_struct expects at least one argument".into(),
2110 ));
2111 };
2112 let IrExpr::VarRef(var_id) = self.arena.get(base_id) else {
2113 return Err(LoweringError::UnsupportedExpr(
2114 "_node_struct argument must be a node variable".into(),
2115 ));
2116 };
2117 let base = self
2118 .var_map
2119 .get(*var_id)
2120 .ok_or(LoweringError::UnboundVar(var_id.0))?;
2121 let label = args.get(1).and_then(|&id| match self.arena.get(id) {
2122 IrExpr::Literal(IrLiteral::Str(s)) => Some(s.as_str()),
2123 _ => None,
2124 });
2125 let prop_names = self
2126 .node_shapes
2127 .get(&var_id.0)
2128 .map(|s| s.prop_names.clone())
2129 .unwrap_or_default();
2130 let labels = self.input_schema.as_ref().and_then(|schema| {
2131 let qualifier = datafusion::common::TableReference::bare(base);
2132 schema
2133 .index_of_column_by_name(Some(&qualifier), "labels")
2134 .is_some()
2135 .then(|| qualified_col(base, "labels"))
2136 });
2137 Ok(labels.map_or_else(
2138 || node_value_struct(base, label, &self.type_id_to_entity_name, &prop_names),
2139 |labels| node_value_struct_with_labels(base, labels, &prop_names),
2140 ))
2141 }
2142
2143 fn lower_node_struct_list(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2144 use datafusion::functions_nested::expr_fn::make_array;
2145
2146 let [first, second, edge] = args else {
2147 return Err(LoweringError::UnsupportedExpr(
2148 "_node_struct_list expects two nodes and one relationship".into(),
2149 ));
2150 };
2151 let node_vars = [first, second].map(|id| match self.arena.get(*id) {
2152 IrExpr::VarRef(var) => Ok(*var),
2153 _ => Err(LoweringError::UnsupportedExpr(
2154 "_node_struct_list node arguments must be variables".into(),
2155 )),
2156 });
2157 let [first_var, second_var] = node_vars;
2158 let node_vars = [first_var?, second_var?];
2159 let prop_names = node_vars
2160 .iter()
2161 .filter_map(|var| self.node_shapes.get(&var.0))
2162 .flat_map(|shape| shape.prop_names.iter().cloned())
2163 .collect::<std::collections::BTreeSet<_>>()
2164 .into_iter()
2165 .collect::<Vec<_>>();
2166 let nodes = node_vars
2167 .iter()
2168 .map(|var| {
2169 let base = self
2170 .var_map
2171 .get(*var)
2172 .ok_or(LoweringError::UnboundVar(var.0))?;
2173 Ok(node_value_struct(
2174 base,
2175 None,
2176 &self.type_id_to_entity_name,
2177 &prop_names,
2178 ))
2179 })
2180 .collect::<Result<Vec<_>, LoweringError>>()?;
2181 let edge = self.lower_path_builtin_arg(*edge)?;
2182 let present = edge_present(&edge).ok_or_else(|| {
2183 LoweringError::UnsupportedExpr(
2184 "_node_struct_list relationship argument must be a bound edge".into(),
2185 )
2186 })?;
2187 Ok(null_unless(present, make_array(nodes)))
2188 }
2189
2190 fn lower_rel_struct(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2191 let (base, rel_type) = self.lower_rel_struct_args(args)?;
2192 let props = self.edge_prop_names(base);
2193 let value = relationship_value_struct(base, rel_type, &props);
2194 Ok(null_unless(edge_present_qual(base), value))
2195 }
2196
2197 fn lower_rel_struct_list(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2198 use datafusion::functions_nested::expr_fn::make_array;
2199
2200 let (base, rel_type) = self.lower_rel_struct_args(args)?;
2201 let props = self.edge_prop_names(base);
2202 let value = relationship_value_struct(base, rel_type, &props);
2203 Ok(null_unless(
2204 edge_present_qual(base),
2205 make_array(vec![value]),
2206 ))
2207 }
2208
2209 fn lower_rel_struct_args(&self, args: &[ExprId]) -> Result<(&str, DfExpr), LoweringError> {
2210 let Some(&base_id) = args.first() else {
2211 return Err(LoweringError::UnsupportedExpr(
2212 "_rel_struct expects an edge variable".into(),
2213 ));
2214 };
2215 let IrExpr::VarRef(var_id) = self.arena.get(base_id) else {
2216 return Err(LoweringError::UnsupportedExpr(
2217 "_rel_struct argument must be a relationship variable".into(),
2218 ));
2219 };
2220 let base = self
2221 .var_map
2222 .get(*var_id)
2223 .ok_or(LoweringError::UnboundVar(var_id.0))?;
2224 let rel_type = match args.get(1).map(|&id| (id, self.arena.get(id))) {
2225 Some((_, IrExpr::Literal(IrLiteral::Null))) | None => {
2226 col(format!("{base}.rel_type_name"))
2227 }
2228 Some((id, _)) => self.lower(id)?,
2229 };
2230 Ok((base, rel_type))
2231 }
2232
2233 fn edge_prop_names(&self, base: &str) -> Vec<String> {
2234 let Some(schema) = self.input_schema.as_ref() else {
2235 return Vec::new();
2236 };
2237 schema
2238 .iter()
2239 .filter_map(|(qualifier, field)| {
2240 let q = qualifier?;
2241 if q.to_string() != base
2242 || is_edge_value_topology_field(field.name())
2243 || matches!(field.data_type(), DataType::Null)
2244 {
2245 return None;
2246 }
2247 Some(field.name().clone())
2248 })
2249 .collect()
2250 }
2251
2252 fn lower_labels(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2255 let Some(&base_id) = args.first() else {
2256 return Err(LoweringError::UnsupportedExpr(
2257 "labels() expects one argument".into(),
2258 ));
2259 };
2260 match self.arena.get(base_id) {
2261 IrExpr::Literal(IrLiteral::Null) => Ok(null_utf8_list()),
2262 IrExpr::VarRef(var_id) if self.node_shapes.contains_key(&var_id.0) => {
2263 let base = self
2264 .var_map
2265 .get(*var_id)
2266 .ok_or(LoweringError::UnboundVar(var_id.0))?;
2267 Ok(null_unless(
2268 col(format!("{base}.node_uuid")).is_not_null(),
2269 node_labels_list(base, None, &self.type_id_to_entity_name),
2270 ))
2271 }
2272 IrExpr::VarRef(var_id) => {
2273 let base = self
2274 .var_map
2275 .get(*var_id)
2276 .ok_or(LoweringError::UnboundVar(var_id.0))?;
2277 if self.is_node_var(base) {
2278 Ok(null_unless(
2279 col(format!("{base}.node_uuid")).is_not_null(),
2280 node_labels_list(base, None, &self.type_id_to_entity_name),
2281 ))
2282 } else {
2283 let value = self.lower(base_id)?;
2284 Ok(CYPHER_LABELS.call(vec![value]))
2285 }
2286 }
2287 _ => {
2288 let value = self.lower(base_id)?;
2289 Ok(CYPHER_LABELS.call(vec![value]))
2290 }
2291 }
2292 }
2293
2294 fn lower_keys(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2297 use datafusion::functions_nested::expr_fn::{array_concat, make_array};
2298
2299 let Some(&base_id) = args.first() else {
2300 return Err(LoweringError::UnsupportedExpr(
2301 "keys() expects one argument".into(),
2302 ));
2303 };
2304 if let Some(map_keys) = self.lower_map_keys(base_id)? {
2305 return Ok(map_keys);
2306 }
2307 let var_id = match self.arena.get(base_id) {
2308 IrExpr::VarRef(var_id) => *var_id,
2309 _ => {
2310 return Err(LoweringError::InvalidType(
2311 "keys() requires a map, node, relationship, or null".into(),
2312 ));
2313 }
2314 };
2315 let base = self
2316 .var_map
2317 .get(var_id)
2318 .ok_or(LoweringError::UnboundVar(var_id.0))?;
2319 let (prop_names, present) = if let Some(shape) = self.node_shapes.get(&var_id.0) {
2320 let has_node_uuid = self.input_schema.as_ref().is_some_and(|schema| {
2321 let qual = datafusion::common::TableReference::bare(base);
2322 schema
2323 .index_of_column_by_name(Some(&qual), "node_uuid")
2324 .is_some()
2325 });
2326 let node_present = if has_node_uuid {
2327 col(format!("{base}.node_uuid")).is_not_null()
2328 } else {
2329 lit(true)
2330 };
2331 (shape.prop_names.clone(), node_present)
2332 } else if self.is_edge_var(base) {
2333 (self.edge_prop_names(base), edge_present_qual(base))
2334 } else {
2335 return Err(LoweringError::UnsupportedExpr(
2336 "keys() requires an entity with a known shape".into(),
2337 ));
2338 };
2339 let empty = empty_utf8_list();
2342 let parts: Vec<DfExpr> = prop_names
2343 .iter()
2344 .map(|p| {
2345 when(
2346 qualified_col(base, p).is_not_null(),
2347 make_array(vec![lit(p.as_str())]),
2348 )
2349 .otherwise(empty.clone())
2350 .expect("CASE build")
2351 })
2352 .collect();
2353 let value = parts
2354 .into_iter()
2355 .reduce(|acc, part| array_concat(vec![acc, part]))
2356 .unwrap_or(empty);
2357 Ok(null_unless(present, value))
2358 }
2359
2360 fn lower_map_keys(&self, base_id: ExprId) -> Result<Option<DfExpr>, LoweringError> {
2361 match self.arena.get(base_id) {
2362 IrExpr::Literal(IrLiteral::Null) => return Ok(Some(null_utf8_list())),
2363 IrExpr::MapLiteral(_) => {
2364 return Ok(Some(CYPHER_MAP_KEYS.call(vec![self.lower(base_id)?])));
2365 }
2366 IrExpr::ListLiteral(_) => {
2367 return Err(LoweringError::InvalidType(
2368 "keys() requires a map, node, relationship, or null".into(),
2369 ));
2370 }
2371 _ => {}
2372 }
2373 if let IrExpr::VarRef(var_id) = self.arena.get(base_id)
2374 && let Some(base) = self.var_map.get(*var_id)
2375 {
2376 if self.node_shapes.contains_key(&var_id.0) || self.is_edge_var(base) {
2377 return Ok(None);
2378 }
2379 if let Some(schema) = self.input_schema.as_ref()
2380 && let Ok(field) = schema.field_with_unqualified_name(base)
2381 {
2382 if matches!(field.data_type(), DataType::Null)
2383 || is_plain_map_struct_type(field.data_type())
2384 || is_het_struct_type(Some(field.data_type()))
2385 {
2386 return Ok(Some(CYPHER_MAP_KEYS.call(vec![col_literal(base)])));
2387 }
2388 return Ok(None);
2389 }
2390 }
2391 let value = self.lower(base_id)?;
2392 if let Some(dt) = self.expr_data_type(&value) {
2393 if matches!(dt, DataType::Null)
2394 || is_plain_map_struct_type(&dt)
2395 || is_het_struct_type(Some(&dt))
2396 {
2397 return Ok(Some(CYPHER_MAP_KEYS.call(vec![value])));
2398 }
2399 return Err(LoweringError::InvalidType(
2400 "keys() requires a map, node, relationship, or null".into(),
2401 ));
2402 }
2403 Ok(Some(CYPHER_MAP_KEYS.call(vec![value])))
2404 }
2405
2406 fn lower_properties(&self, args: &[ExprId]) -> Result<DfExpr, LoweringError> {
2407 let Some(&base_id) = args.first() else {
2408 return Err(LoweringError::UnsupportedExpr(
2409 "properties() expects one argument".into(),
2410 ));
2411 };
2412 match self.arena.get(base_id) {
2413 IrExpr::Literal(IrLiteral::Null) => return Ok(lit(ScalarValue::Null)),
2414 IrExpr::MapLiteral(_) => return self.lower(base_id),
2415 IrExpr::ListLiteral(_) => {
2416 return Err(LoweringError::InvalidType(
2417 "properties() requires a map, node, relationship, or null".into(),
2418 ));
2419 }
2420 _ => {}
2421 }
2422 if let IrExpr::VarRef(var_id) = self.arena.get(base_id)
2423 && let Some(base) = self.var_map.get(*var_id)
2424 {
2425 if let Some(value) = self.entity_property_bag_with_empty(*var_id, base) {
2426 return Ok(value);
2427 }
2428 if let Some(schema) = self.input_schema.as_ref()
2429 && let Ok(field) = schema.field_with_unqualified_name(base)
2430 {
2431 return match field.data_type() {
2432 DataType::Null => Ok(lit(ScalarValue::Null)),
2433 dt if is_plain_map_struct_type(dt) => Ok(col_literal(base)),
2434 _ => Err(LoweringError::InvalidType(
2435 "properties() requires a map, node, relationship, or null".into(),
2436 )),
2437 };
2438 }
2439 }
2440 let value = self.lower(base_id)?;
2441 if let Some(dt) = self.expr_data_type(&value) {
2442 return match dt {
2443 DataType::Null => Ok(lit(ScalarValue::Null)),
2444 dt if is_plain_map_struct_type(&dt) => Ok(value),
2445 _ => Err(LoweringError::InvalidType(
2446 "properties() requires a map, node, relationship, or null".into(),
2447 )),
2448 };
2449 }
2450 Ok(value)
2451 }
2452
2453 fn is_edge_var(&self, base: &str) -> bool {
2454 let Some(schema) = self.input_schema.as_ref() else {
2455 return false;
2456 };
2457 let qual = datafusion::common::TableReference::bare(base);
2458 schema
2459 .index_of_column_by_name(Some(&qual), "edge_uuid")
2460 .is_some()
2461 }
2462
2463 fn is_node_var(&self, base: &str) -> bool {
2464 let Some(schema) = self.input_schema.as_ref() else {
2465 return false;
2466 };
2467 let qual = datafusion::common::TableReference::bare(base);
2468 schema
2469 .index_of_column_by_name(Some(&qual), "node_uuid")
2470 .is_some()
2471 }
2472
2473 fn identity_uuid_of(&self, id: ExprId) -> Option<(EntityIdentityKind, DfExpr)> {
2478 if let IrExpr::VarRef(v) = self.arena.get(id)
2479 && self.node_shapes.contains_key(&v.0)
2480 {
2481 let base = self.var_map.get(*v)?;
2482 return Some((EntityIdentityKind::Node, col(format!("{base}.node_uuid"))));
2483 }
2484 if let IrExpr::VarRef(v) = self.arena.get(id) {
2485 let base = self.var_map.get(*v)?;
2486 let qual = datafusion::common::TableReference::bare(base);
2487 if let Some(schema) = self.input_schema.as_ref()
2488 && schema
2489 .index_of_column_by_name(Some(&qual), "edge_uuid")
2490 .is_some()
2491 {
2492 return Some((EntityIdentityKind::Edge, col(format!("{base}.edge_uuid"))));
2493 }
2494 }
2495 None
2496 }
2497
2498 fn is_list_typed(&self, e: &DfExpr) -> bool {
2502 if is_list_literal(e) {
2503 return true;
2504 }
2505 if let Some(schema) = self.input_schema.as_ref()
2506 && let Ok(dt) = e.get_type(schema)
2507 {
2508 return matches!(
2509 dt,
2510 DataType::List(_) | DataType::LargeList(_) | DataType::FixedSizeList(_, _)
2511 );
2512 }
2513 false
2514 }
2515
2516 fn is_string_typed(&self, e: &DfExpr) -> bool {
2520 if matches!(
2521 e,
2522 DfExpr::Literal(ScalarValue::Utf8(_) | ScalarValue::LargeUtf8(_), _)
2523 ) {
2524 return true;
2525 }
2526 if let Some(schema) = self.input_schema.as_ref()
2527 && let Ok(dt) = e.get_type(schema)
2528 {
2529 return matches!(
2530 dt,
2531 DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View
2532 );
2533 }
2534 false
2535 }
2536
2537 fn is_known_non_string(&self, e: &DfExpr) -> bool {
2538 if is_list_literal(e)
2539 || matches!(e, DfExpr::ScalarFunction(f) if f.func.name() == "named_struct")
2540 {
2541 return true;
2542 }
2543 self.expr_data_type(e).is_some_and(|dt| {
2544 !matches!(
2545 dt,
2546 DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View | DataType::Null
2547 )
2548 })
2549 }
2550
2551 fn expr_data_type(&self, e: &DfExpr) -> Option<DataType> {
2554 if let DfExpr::Literal(sv, _) = e {
2555 return Some(sv.data_type());
2556 }
2557 self.input_schema.as_ref().and_then(|s| e.get_type(s).ok())
2558 }
2559
2560 fn is_known_non_bool(&self, e: &DfExpr) -> bool {
2569 if is_list_literal(e)
2573 || matches!(e, DfExpr::ScalarFunction(f) if f.func.name() == "named_struct")
2574 {
2575 return true;
2576 }
2577 !matches!(
2578 self.expr_data_type(e),
2579 None | Some(DataType::Boolean | DataType::Null)
2580 )
2581 }
2582
2583 fn is_known_non_numeric(&self, e: &DfExpr) -> bool {
2588 match self.expr_data_type(e) {
2589 None | Some(DataType::Null) => false,
2590 Some(dt) => !dt.is_numeric(),
2591 }
2592 }
2593
2594 fn is_known_non_list(&self, e: &DfExpr) -> bool {
2595 if matches!(e, DfExpr::ScalarFunction(f) if f.func.name() == "named_struct") {
2596 return true;
2597 }
2598 match self.expr_data_type(e) {
2599 None
2600 | Some(
2601 DataType::Null
2602 | DataType::List(_)
2603 | DataType::LargeList(_)
2604 | DataType::FixedSizeList(_, _),
2605 ) => false,
2606 Some(dt) if is_het_struct_type(Some(&dt)) => false,
2607 Some(_) => true,
2608 }
2609 }
2610
2611 fn is_temporal_typed(&self, e: &DfExpr) -> bool {
2614 match self.expr_data_type(e) {
2615 Some(DataType::Time64(_)) => true,
2616 Some(dt) => {
2617 is_date_struct(&dt)
2618 || is_localdatetime_struct(&dt)
2619 || is_time_struct(&dt)
2620 || is_datetime_struct(&dt)
2621 }
2622 None => false,
2623 }
2624 }
2625
2626 fn is_known_non_value_access_container(&self, e: &DfExpr) -> bool {
2627 match self.expr_data_type(e) {
2628 None | Some(DataType::Null) => false,
2629 Some(dt) => {
2630 !is_plain_map_struct_type(&dt)
2631 && !is_het_struct_type(Some(&dt))
2632 && !matches!(dt, DataType::Struct(_))
2633 }
2634 }
2635 }
2636
2637 fn is_duration_typed(&self, e: &DfExpr) -> bool {
2639 self.expr_data_type(e)
2640 .is_some_and(|dt| is_duration_struct(&dt))
2641 }
2642
2643 fn list_element_type(&self, list: &DfExpr) -> Option<DataType> {
2646 let schema = self
2647 .input_schema
2648 .clone()
2649 .unwrap_or_else(|| std::sync::Arc::new(datafusion::common::DFSchema::empty()));
2650 match list.get_type(&schema).ok()? {
2651 DataType::List(f) | DataType::LargeList(f) | DataType::FixedSizeList(f, _) => {
2652 Some(f.data_type().clone())
2653 }
2654 _ => None,
2655 }
2656 }
2657
2658 #[allow(
2659 clippy::too_many_lines,
2660 reason = "one arm per Cypher binary operator, several with temporal/list/string dispatch"
2661 )]
2662 fn lower_binary(
2663 &self,
2664 op: BinaryOpKind,
2665 left: ExprId,
2666 right: ExprId,
2667 ) -> Result<DfExpr, LoweringError> {
2668 if op == BinaryOpKind::In
2674 && let Some(membership) = self.lower_known_label_membership(left, right)?
2675 {
2676 return Ok(membership);
2677 }
2678
2679 if matches!(op, BinaryOpKind::Eq | BinaryOpKind::Neq)
2683 && let (Some((lk, lc)), Some((rk, rc))) =
2684 (self.identity_uuid_of(left), self.identity_uuid_of(right))
2685 {
2686 if lk != rk {
2687 let value = matches!(op, BinaryOpKind::Neq);
2688 return Ok(when(
2689 lc.clone().is_null().or(rc.clone().is_null()),
2690 lit(ScalarValue::Boolean(None)),
2691 )
2692 .otherwise(lit(value))
2693 .expect("CASE build is infallible for mismatched entity comparison"));
2694 }
2695 return Ok(if matches!(op, BinaryOpKind::Eq) {
2696 lc.eq(rc)
2697 } else {
2698 lc.not_eq(rc)
2699 });
2700 }
2701 let l = self.lower(left)?;
2702 let r = self.lower(right)?;
2703 let expr = match op {
2704 BinaryOpKind::Eq => CYPHER_EQ.call(vec![l, r]),
2710 BinaryOpKind::Neq => datafusion::logical_expr::not(CYPHER_EQ.call(vec![l, r])),
2711 BinaryOpKind::Lt | BinaryOpKind::Lte | BinaryOpKind::Gt | BinaryOpKind::Gte => {
2716 let code = match op {
2717 BinaryOpKind::Lt => 0i8,
2718 BinaryOpKind::Lte => 1i8,
2719 BinaryOpKind::Gt => 2i8,
2720 _ => 3i8,
2721 };
2722 CYPHER_CMP_PRED.call(vec![l, r, lit(code)])
2723 }
2724 BinaryOpKind::And | BinaryOpKind::Or | BinaryOpKind::Xor => {
2730 let keyword = match op {
2731 BinaryOpKind::And => "AND",
2732 BinaryOpKind::Or => "OR",
2733 BinaryOpKind::Xor => "XOR",
2734 _ => unreachable!("matched boolean operator"),
2735 };
2736 if self.is_known_non_bool(&l) || self.is_known_non_bool(&r) {
2737 return Err(LoweringError::InvalidType(format!(
2738 "{keyword} requires boolean operands"
2739 )));
2740 }
2741 match op {
2742 BinaryOpKind::And => CYPHER_AND.call(vec![l, r]),
2743 BinaryOpKind::Or => CYPHER_OR.call(vec![l, r]),
2744 BinaryOpKind::Xor => CYPHER_XOR.call(vec![l, r]),
2745 _ => unreachable!("matched boolean operator"),
2746 }
2747 }
2748 BinaryOpKind::Add => {
2749 if self.is_temporal_typed(&l) && self.is_duration_typed(&r) {
2757 CYPHER_TEMPORAL_ARITH.call(vec![l, r, lit(1i64)])
2759 } else if self.is_duration_typed(&l) && self.is_temporal_typed(&r) {
2760 CYPHER_TEMPORAL_ARITH.call(vec![r, l, lit(1i64)])
2762 } else if self.is_duration_typed(&l) && self.is_duration_typed(&r) {
2763 CYPHER_DURATION_ADD.call(vec![l, r, lit(1i64)])
2765 } else if let (Some(le), Some(re)) =
2766 (self.list_element_type(&l), self.list_element_type(&r))
2767 {
2768 if le == re && !is_het_struct_type(Some(&le)) {
2769 datafusion::functions_nested::expr_fn::array_concat(vec![l, r])
2770 } else if graph_value_types_compatible(&le, &re) {
2771 let target = self.expr_data_type(&l).ok_or_else(|| {
2772 LoweringError::UnsupportedExpr(
2773 "cannot resolve path-list type for concatenation".into(),
2774 )
2775 })?;
2776 datafusion::functions_nested::expr_fn::array_concat(vec![
2777 l,
2778 cast(r, target),
2779 ])
2780 } else {
2781 CYPHER_LIST_PLUS.call(vec![l, r])
2782 }
2783 } else if let Some(le) = self.list_element_type(&l) {
2784 if !is_het_struct_type(Some(&le))
2785 && self
2786 .expr_data_type(&r)
2787 .is_some_and(|rt| rt == le || matches!(rt, DataType::Null))
2788 {
2789 datafusion::functions_nested::expr_fn::array_append(l, r)
2790 } else {
2791 CYPHER_LIST_PLUS.call(vec![l, r])
2792 }
2793 } else if let Some(re) = self.list_element_type(&r) {
2794 if !is_het_struct_type(Some(&re))
2795 && self
2796 .expr_data_type(&l)
2797 .is_some_and(|lt| lt == re || matches!(lt, DataType::Null))
2798 {
2799 datafusion::functions_nested::expr_fn::array_prepend(l, r)
2800 } else {
2801 CYPHER_LIST_PLUS.call(vec![l, r])
2802 }
2803 } else if self.is_list_typed(&l) || self.is_list_typed(&r) {
2804 CYPHER_LIST_PLUS.call(vec![l, r])
2805 } else if (self.is_string_typed(&l) && !self.is_known_non_string(&r))
2806 || (self.is_string_typed(&r) && !self.is_known_non_string(&l))
2807 {
2808 DfExpr::BinaryExpr(datafusion::logical_expr::BinaryExpr {
2809 left: Box::new(l),
2810 op: Operator::StringConcat,
2811 right: Box::new(r),
2812 })
2813 } else {
2814 DfExpr::BinaryExpr(datafusion::logical_expr::BinaryExpr {
2815 left: Box::new(l),
2816 op: Operator::Plus,
2817 right: Box::new(r),
2818 })
2819 }
2820 }
2821 BinaryOpKind::Sub => {
2822 if self.is_temporal_typed(&l) && self.is_duration_typed(&r) {
2823 CYPHER_TEMPORAL_ARITH.call(vec![l, r, lit(-1i64)])
2825 } else if self.is_duration_typed(&l) && self.is_duration_typed(&r) {
2826 CYPHER_DURATION_ADD.call(vec![l, r, lit(-1i64)])
2828 } else {
2829 DfExpr::BinaryExpr(datafusion::logical_expr::BinaryExpr {
2830 left: Box::new(l),
2831 op: Operator::Minus,
2832 right: Box::new(r),
2833 })
2834 }
2835 }
2836 BinaryOpKind::Mul => {
2837 if self.is_duration_typed(&l) {
2839 CYPHER_DURATION_SCALE.call(vec![l, r, lit(false)])
2840 } else if self.is_duration_typed(&r) {
2841 CYPHER_DURATION_SCALE.call(vec![r, l, lit(false)])
2842 } else {
2843 DfExpr::BinaryExpr(datafusion::logical_expr::BinaryExpr {
2844 left: Box::new(l),
2845 op: Operator::Multiply,
2846 right: Box::new(r),
2847 })
2848 }
2849 }
2850 BinaryOpKind::Div => {
2851 if self.is_duration_typed(&l) {
2853 CYPHER_DURATION_SCALE.call(vec![l, r, lit(true)])
2854 } else {
2855 let (l, r) = match (self.expr_data_type(&l), self.expr_data_type(&r)) {
2856 (Some(DataType::Float64), Some(rt)) if is_integer_data_type(&rt) => {
2857 (l, cast(r, DataType::Float64))
2858 }
2859 (Some(lt), Some(DataType::Float64)) if is_integer_data_type(<) => {
2860 (cast(l, DataType::Float64), r)
2861 }
2862 _ => (l, r),
2863 };
2864 DfExpr::BinaryExpr(datafusion::logical_expr::BinaryExpr {
2865 left: Box::new(l),
2866 op: Operator::Divide,
2867 right: Box::new(r),
2868 })
2869 }
2870 }
2871 BinaryOpKind::Mod => {
2872 if self.is_known_non_numeric(&l) || self.is_known_non_numeric(&r) {
2873 return Err(LoweringError::InvalidType(
2874 "% requires numeric operands".into(),
2875 ));
2876 }
2877 DfExpr::BinaryExpr(datafusion::logical_expr::BinaryExpr {
2878 left: Box::new(l),
2879 op: Operator::Modulo,
2880 right: Box::new(r),
2881 })
2882 }
2883 BinaryOpKind::Pow => {
2884 if self.is_known_non_numeric(&l) || self.is_known_non_numeric(&r) {
2885 return Err(LoweringError::InvalidType(
2886 "^ requires numeric operands".into(),
2887 ));
2888 }
2889 datafusion::functions::math::expr_fn::power(l, r)
2890 }
2891 BinaryOpKind::In => {
2896 if self.is_known_non_list(&r) {
2897 return Err(LoweringError::InvalidType(
2898 "IN requires a list or null right-hand operand".into(),
2899 ));
2900 }
2901 CYPHER_IN.call(vec![l, r])
2902 }
2903 BinaryOpKind::StartsWith => CYPHER_STARTS_WITH.call(vec![l, r]),
2904 BinaryOpKind::EndsWith => CYPHER_ENDS_WITH.call(vec![l, r]),
2905 BinaryOpKind::Contains => CYPHER_CONTAINS.call(vec![l, r]),
2906 BinaryOpKind::RegexMatch => {
2907 datafusion::functions::regex::expr_fn::regexp_like(l, r, None)
2909 }
2910 };
2911 Ok(expr)
2912 }
2913
2914 fn lower_known_label_membership(
2920 &self,
2921 left: ExprId,
2922 right: ExprId,
2923 ) -> Result<Option<DfExpr>, LoweringError> {
2924 use datafusion::functions_nested::expr_fn::array_has;
2925
2926 let IrExpr::Literal(IrLiteral::Str(label)) = self.arena.get(left) else {
2927 return Ok(None);
2928 };
2929 let IrExpr::FunctionCall { name, args } = self.arena.get(right) else {
2930 return Ok(None);
2931 };
2932 let [arg] = args.as_slice() else {
2933 return Ok(None);
2934 };
2935 if name != "labels" {
2936 return Ok(None);
2937 }
2938 let IrExpr::VarRef(var_id) = self.arena.get(*arg) else {
2939 return Ok(None);
2940 };
2941 let Some(type_id) = self.entity_name_to_type_id.get(label) else {
2942 return Ok(None);
2943 };
2944 let base = self
2945 .var_map
2946 .get(*var_id)
2947 .ok_or(LoweringError::UnboundVar(var_id.0))?;
2948 Ok(Some(array_has(
2949 col(format!("{base}.type_ids")),
2950 lit(*type_id),
2951 )))
2952 }
2953
2954 fn lower_unary(&self, op: UnaryOpKind, expr: ExprId) -> Result<DfExpr, LoweringError> {
2955 let e = self.lower(expr)?;
2956 let result = match op {
2957 UnaryOpKind::Not => {
2958 if self.is_known_non_bool(&e) {
2959 return Err(LoweringError::InvalidType(
2960 "NOT requires a boolean operand".into(),
2961 ));
2962 }
2963 not(e)
2964 }
2965 UnaryOpKind::Neg => {
2966 if self.is_known_non_numeric(&e) {
2967 return Err(LoweringError::InvalidType(
2968 "unary minus requires a numeric operand".into(),
2969 ));
2970 }
2971 DfExpr::Negative(Box::new(e))
2972 }
2973 UnaryOpKind::IsNull => e.is_null(),
2974 UnaryOpKind::IsNotNull => e.is_not_null(),
2975 };
2976 Ok(result)
2977 }
2978
2979 fn lower_case(
2980 &self,
2981 operand: Option<ExprId>,
2982 arms: &[graphforge_ir::expr::CaseArm],
2983 else_expr: Option<ExprId>,
2984 ) -> Result<DfExpr, LoweringError> {
2985 let when_thens: Result<Vec<_>, _> = arms
2986 .iter()
2987 .map(|arm| {
2988 let when = self.lower(arm.when)?;
2989 let then = self.lower(arm.then)?;
2990 Ok((Box::new(when), Box::new(then)))
2991 })
2992 .collect();
2993 let when_thens = when_thens?;
2994 let else_expr_df = else_expr.map(|id| self.lower(id)).transpose()?;
2995
2996 Ok(DfExpr::Case(datafusion::logical_expr::expr::Case {
2997 expr: operand.map(|id| self.lower(id)).transpose()?.map(Box::new),
2998 when_then_expr: when_thens,
2999 else_expr: else_expr_df.map(Box::new),
3000 }))
3001 }
3002
3003 fn resolve_prop_col(&self, base_expr: DfExpr, prop: PropId) -> DfExpr {
3009 let prop_name = self
3010 .prop_names
3011 .get(&prop.0)
3012 .cloned()
3013 .unwrap_or_else(|| format!("prop_{}", prop.0));
3014
3015 if let DfExpr::Column(col_ref) = &base_expr
3022 && !self
3023 .elem_struct_cols
3024 .iter()
3025 .any(|c| c == col_ref.name.as_str())
3026 {
3027 return qualified_col(&col_ref.name, &prop_name);
3028 }
3029
3030 datafusion::functions::core::expr_fn::get_field(base_expr, prop_name)
3033 }
3034
3035 fn path_node_hydration(&self) -> Option<PathNodeHydration> {
3042 use datafusion::arrow::datatypes::Field;
3043 let dir = self.read_target.as_ref()?;
3044 let stems = graphforge_storage::list_property_stems(dir);
3045 let mut fields = vec![
3046 Field::new("node_uuid", DataType::FixedSizeBinary(16), false),
3047 Field::new("labels", DataType::new_list(DataType::Utf8, true), true),
3048 ];
3049 let mut seen: std::collections::HashSet<String> =
3050 fields.iter().map(|f| f.name().clone()).collect();
3051 for stem in &stems {
3052 let table = graphforge_storage::PropertyTable::open_discovered(dir, stem);
3053 for f in table.schema_ref().fields() {
3054 if f.name() == "node_uuid" || !seen.insert(f.name().clone()) {
3055 continue;
3056 }
3057 fields.push(f.as_ref().clone().with_nullable(true));
3058 }
3059 }
3060 let mut labels_by_type: Vec<(u32, String)> = self
3061 .type_id_to_entity_name
3062 .iter()
3063 .map(|(id, name)| (*id, name.clone()))
3064 .collect();
3065 labels_by_type.sort();
3066 Some(PathNodeHydration {
3067 dir: dir.clone(),
3068 labels_by_type,
3069 prop_stems: stems,
3070 fields: fields.into(),
3071 })
3072 }
3073}
3074
3075fn col_literal(name: &str) -> DfExpr {
3087 if name.contains('.') {
3088 col(name)
3089 } else {
3090 DfExpr::Column(datafusion::common::Column::new_unqualified(name))
3091 }
3092}
3093
3094pub(crate) fn qualified_col(relation: &str, name: &str) -> DfExpr {
3095 DfExpr::Column(datafusion::common::Column::new(
3096 Some(datafusion::common::TableReference::bare(relation)),
3097 name,
3098 ))
3099}
3100
3101fn build_prop_names(_ontology: Option<&OntologyHandle>) -> HashMap<u32, String> {
3111 HashMap::new()
3112}
3113
3114fn is_empty_list_literal(e: &DfExpr) -> bool {
3137 use datafusion::arrow::array::Array;
3138 matches!(e, DfExpr::Literal(ScalarValue::List(arr), _) if arr.value(0).is_empty())
3139}
3140
3141fn try_const_scalar(e: &DfExpr) -> Option<ScalarValue> {
3147 match e {
3148 DfExpr::Literal(s, _) => Some(s.clone()),
3149 DfExpr::ScalarFunction(f) if f.func.name() == "named_struct" => {
3150 let mut entries: Vec<(String, ScalarValue)> = Vec::with_capacity(f.args.len() / 2);
3151 let pairs = f.args.chunks_exact(2);
3152 if !pairs.remainder().is_empty() {
3153 return None;
3154 }
3155 for pair in pairs {
3156 let DfExpr::Literal(ScalarValue::Utf8(Some(k)), _) = &pair[0] else {
3157 return None;
3158 };
3159 entries.push((k.clone(), try_const_scalar(&pair[1])?));
3160 }
3161 const_map_scalar(&entries)
3162 }
3163 _ => None,
3164 }
3165}
3166
3167enum ConstStringKey {
3168 Null,
3169 Value(String),
3170}
3171
3172fn const_string_key(e: &DfExpr) -> Option<ConstStringKey> {
3173 match e {
3174 DfExpr::Literal(ScalarValue::Utf8(v) | ScalarValue::LargeUtf8(v), _) => Some(
3175 v.clone()
3176 .map_or(ConstStringKey::Null, ConstStringKey::Value),
3177 ),
3178 DfExpr::BinaryExpr(b) if b.op == Operator::StringConcat => {
3179 let l = const_string_key(&b.left)?;
3180 let r = const_string_key(&b.right)?;
3181 Some(match (l, r) {
3182 (ConstStringKey::Value(l), ConstStringKey::Value(r)) => {
3183 ConstStringKey::Value(format!("{l}{r}"))
3184 }
3185 _ => ConstStringKey::Null,
3186 })
3187 }
3188 _ => None,
3189 }
3190}
3191
3192fn const_map_scalar(entries: &[(String, ScalarValue)]) -> Option<ScalarValue> {
3196 use datafusion::arrow::array::{ArrayRef, StructArray};
3197 use datafusion::arrow::datatypes::{Field, Fields};
3198 use std::sync::Arc;
3199 if entries.is_empty() {
3200 return Some(ScalarValue::Struct(Arc::new(
3201 StructArray::new_empty_fields(1, None),
3202 )));
3203 }
3204 let mut fields: Vec<Field> = Vec::with_capacity(entries.len());
3205 let mut arrays: Vec<ArrayRef> = Vec::with_capacity(entries.len());
3206 for (k, sv) in entries {
3207 let arr = sv.to_array().ok()?; fields.push(Field::new(k, arr.data_type().clone(), true));
3209 arrays.push(arr);
3210 }
3211 let s = StructArray::try_new(Fields::from(fields), arrays, None).ok()?;
3212 Some(ScalarValue::Struct(Arc::new(s)))
3213}
3214
3215fn is_plain_map_struct(arr: &datafusion::arrow::array::StructArray) -> bool {
3220 use datafusion::arrow::array::Array;
3221 is_plain_map_struct_type(arr.data_type())
3222}
3223
3224fn is_plain_map_struct_type(dt: &DataType) -> bool {
3228 let DataType::Struct(fields) = dt else {
3229 return false;
3230 };
3231 let is_entity = fields.iter().any(|f| {
3233 matches!(
3234 f.name().as_str(),
3235 "node_uuid" | "src_uuid" | "dst_uuid" | "nodes" | "relationships" | "labels"
3236 )
3237 });
3238 !is_entity
3239 && !is_het_struct_type(Some(dt))
3240 && !is_date_struct(dt)
3241 && !is_localdatetime_struct(dt)
3242 && !is_duration_struct(dt)
3243 && !is_time_struct(dt)
3244 && !is_datetime_struct(dt)
3245}
3246
3247fn het_map_entry_fields(depth: usize) -> datafusion::arrow::datatypes::Fields {
3251 use datafusion::arrow::datatypes::{DataType, Field};
3252 datafusion::arrow::datatypes::Fields::from(vec![
3253 Field::new("__het_mkey", DataType::Utf8, false),
3254 Field::new("__het_mval", DataType::Struct(het_fields(depth)), true),
3255 ])
3256}
3257
3258fn het_fields(depth: usize) -> datafusion::arrow::datatypes::Fields {
3268 use datafusion::arrow::datatypes::{DataType, Field};
3269 use std::sync::Arc;
3270 let mut v = vec![
3271 Field::new("__het_key", DataType::Float64, true),
3272 Field::new("__het_tag", DataType::Int8, false),
3273 Field::new("__het_int", DataType::Int64, true),
3274 Field::new("__het_float", DataType::Float64, true),
3275 Field::new("__het_str", DataType::Utf8, true),
3276 Field::new("__het_bool", DataType::Boolean, true),
3277 ];
3278 if depth >= 1 {
3279 let inner = Field::new("item", DataType::Struct(het_fields(depth - 1)), true);
3280 v.push(Field::new(
3281 "__het_list",
3282 DataType::List(Arc::new(inner)),
3283 true,
3284 ));
3285 let entry = Field::new(
3286 "item",
3287 DataType::Struct(het_map_entry_fields(depth - 1)),
3288 true,
3289 );
3290 v.push(Field::new(
3291 "__het_map",
3292 DataType::List(Arc::new(entry)),
3293 true,
3294 ));
3295 }
3296 datafusion::arrow::datatypes::Fields::from(v)
3297}
3298
3299fn het_depth(s: &ScalarValue) -> Option<usize> {
3303 use datafusion::arrow::array::Array;
3304 let s = unwrap_het(s.clone());
3305 match &s {
3306 ScalarValue::Int64(_)
3307 | ScalarValue::Float64(_)
3308 | ScalarValue::Utf8(_)
3309 | ScalarValue::LargeUtf8(_)
3310 | ScalarValue::Utf8View(_)
3311 | ScalarValue::Boolean(_)
3312 | ScalarValue::Null => Some(0),
3313 ScalarValue::List(arr) => {
3314 let inner = arr.value(0);
3315 let mut d = 0;
3316 for i in 0..inner.len() {
3317 let e = unwrap_het(ScalarValue::try_from_array(&inner, i).ok()?);
3321 d = d.max(het_depth(&e)?);
3322 }
3323 Some(1 + d)
3324 }
3325 ScalarValue::Struct(arr) if is_plain_map_struct(arr) => {
3328 let mut d = 0;
3329 for i in 0..arr.num_columns() {
3330 let v = unwrap_het(ScalarValue::try_from_array(arr.column(i), 0).ok()?);
3331 d = d.max(het_depth(&v)?);
3332 }
3333 Some(1 + d)
3334 }
3335 _ => None,
3336 }
3337}
3338
3339fn unwrap_het(s: ScalarValue) -> ScalarValue {
3342 if let ScalarValue::Dictionary(_, value) = s {
3343 return unwrap_het(*value);
3344 }
3345 decode_het(&s).unwrap_or(s)
3346}
3347
3348#[allow(
3351 clippy::too_many_lines,
3352 reason = "one cohesive per-field array builder; splitting it would obscure the field/offset bookkeeping"
3353)]
3354fn build_het_struct(
3355 scalars: &[ScalarValue],
3356 depth: usize,
3357) -> Option<datafusion::arrow::array::StructArray> {
3358 use datafusion::arrow::array::{
3359 ArrayRef, BooleanArray, Float64Array, Int8Array, Int64Array, ListArray, StringArray,
3360 StructArray,
3361 };
3362 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer};
3363 use datafusion::arrow::datatypes::{DataType, Field};
3364 use std::sync::Arc;
3365
3366 let n = scalars.len();
3367 let mut keys: Vec<Option<f64>> = Vec::with_capacity(n);
3368 let mut tags: Vec<i8> = Vec::with_capacity(n);
3369 let mut ints: Vec<Option<i64>> = Vec::with_capacity(n);
3370 let mut floats: Vec<Option<f64>> = Vec::with_capacity(n);
3371 let mut strs: Vec<Option<String>> = Vec::with_capacity(n);
3372 let mut bools: Vec<Option<bool>> = Vec::with_capacity(n);
3373 let mut valid: Vec<bool> = Vec::with_capacity(n);
3374 let mut child_elems: Vec<ScalarValue> = Vec::new();
3376 let mut child_offsets: Vec<i32> = vec![0];
3377 let mut child_valid: Vec<bool> = Vec::new();
3378 let mut map_keys: Vec<String> = Vec::new();
3381 let mut map_vals: Vec<ScalarValue> = Vec::new();
3382 let mut map_offsets: Vec<i32> = vec![0];
3383 let mut map_valid: Vec<bool> = Vec::new();
3384 for scalar in scalars {
3385 let scalar = unwrap_het(scalar.clone());
3386 let (mut key, mut tag, mut int_v, mut float_v, mut str_v, mut bool_v, mut ok) =
3388 (None, 0i8, None, None, None, None, true);
3389 let mut child: Option<Vec<ScalarValue>> = None;
3390 let mut map_child: Option<Vec<(String, ScalarValue)>> = None;
3391 match &scalar {
3392 ScalarValue::Int64(Some(x)) => {
3393 #[allow(
3394 clippy::cast_precision_loss,
3395 reason = "key feeds only min/max ORDERING; the exact integer is preserved in __het_int"
3396 )]
3397 let k = Some(*x as f64);
3398 key = k;
3399 tag = 0;
3400 int_v = Some(*x);
3401 }
3402 ScalarValue::Float64(Some(x)) => {
3403 key = Some(*x);
3404 tag = 1;
3405 float_v = Some(*x);
3406 }
3407 ScalarValue::Utf8(Some(x))
3408 | ScalarValue::LargeUtf8(Some(x))
3409 | ScalarValue::Utf8View(Some(x)) => {
3410 tag = 2;
3411 str_v = Some(x.clone());
3412 }
3413 ScalarValue::Boolean(Some(x)) => {
3414 tag = 3;
3415 bool_v = Some(*x);
3416 }
3417 ScalarValue::List(arr) => {
3418 if depth == 0 {
3419 return None; }
3421 tag = 4;
3422 let inner = arr.value(0);
3423 let mut elems = Vec::with_capacity(inner.len());
3424 for idx in 0..inner.len() {
3425 elems.push(unwrap_het(ScalarValue::try_from_array(&inner, idx).ok()?));
3426 }
3427 child = Some(elems);
3428 }
3429 ScalarValue::Struct(arr) if is_plain_map_struct(arr) => {
3432 if depth == 0 {
3433 return None; }
3435 tag = 5;
3436 let mut kv = Vec::with_capacity(arr.num_columns());
3437 for (i, f) in arr.fields().iter().enumerate() {
3438 let v = unwrap_het(ScalarValue::try_from_array(arr.column(i), 0).ok()?);
3439 kv.push((f.name().clone(), v));
3440 }
3441 map_child = Some(kv);
3442 }
3443 ScalarValue::Int64(None)
3445 | ScalarValue::Float64(None)
3446 | ScalarValue::Utf8(None)
3447 | ScalarValue::LargeUtf8(None)
3448 | ScalarValue::Utf8View(None)
3449 | ScalarValue::Boolean(None)
3450 | ScalarValue::Null => ok = false,
3451 _ => return None, }
3453 keys.push(key);
3454 tags.push(tag);
3455 ints.push(int_v);
3456 floats.push(float_v);
3457 strs.push(str_v);
3458 bools.push(bool_v);
3459 valid.push(ok);
3460 if depth >= 1 {
3461 if let Some(elems) = child {
3462 child_elems.extend(elems);
3463 child_valid.push(true);
3464 } else {
3465 child_valid.push(false);
3466 }
3467 child_offsets.push(i32::try_from(child_elems.len()).ok()?);
3468 if let Some(kv) = map_child {
3469 for (k, v) in kv {
3470 map_keys.push(k);
3471 map_vals.push(v);
3472 }
3473 map_valid.push(true);
3474 } else {
3475 map_valid.push(false);
3476 }
3477 map_offsets.push(i32::try_from(map_keys.len()).ok()?);
3478 }
3479 }
3480
3481 let mut arrays: Vec<ArrayRef> = vec![
3482 Arc::new(Float64Array::from(keys)),
3483 Arc::new(Int8Array::from(tags)),
3484 Arc::new(Int64Array::from(ints)),
3485 Arc::new(Float64Array::from(floats)),
3486 Arc::new(StringArray::from(strs)),
3487 Arc::new(BooleanArray::from(bools)),
3488 ];
3489 if depth >= 1 {
3490 let child_struct = build_het_struct(&child_elems, depth - 1)?;
3491 let inner_field = Arc::new(Field::new(
3492 "item",
3493 DataType::Struct(het_fields(depth - 1)),
3494 true,
3495 ));
3496 let het_list = ListArray::new(
3497 inner_field,
3498 OffsetBuffer::new(child_offsets.into()),
3499 Arc::new(child_struct),
3500 Some(NullBuffer::from(child_valid)),
3501 );
3502 arrays.push(Arc::new(het_list));
3503
3504 let entry_fields = het_map_entry_fields(depth - 1);
3507 let mkey_arr = Arc::new(StringArray::from(map_keys)) as ArrayRef;
3508 let mval_struct = build_het_struct(&map_vals, depth - 1)?;
3509 let entry_struct = StructArray::new(
3510 entry_fields.clone(),
3511 vec![mkey_arr, Arc::new(mval_struct)],
3512 None,
3513 );
3514 let entry_field = Arc::new(Field::new("item", DataType::Struct(entry_fields), true));
3515 let het_map = ListArray::new(
3516 entry_field,
3517 OffsetBuffer::new(map_offsets.into()),
3518 Arc::new(entry_struct),
3519 Some(NullBuffer::from(map_valid)),
3520 );
3521 arrays.push(Arc::new(het_map));
3522 }
3523 Some(StructArray::new(
3524 het_fields(depth),
3525 arrays,
3526 Some(NullBuffer::from(valid)),
3527 ))
3528}
3529
3530fn tagged_numeric_list(scalars: &[ScalarValue]) -> Option<DfExpr> {
3538 use datafusion::arrow::array::ListArray;
3539 use datafusion::arrow::buffer::OffsetBuffer;
3540 use datafusion::arrow::datatypes::{DataType, Field};
3541 use std::sync::Arc;
3542
3543 let mut depth = 0usize;
3546 for s in scalars {
3547 depth = depth.max(het_depth(s)?);
3548 }
3549 let n = scalars.len();
3550 let elem = build_het_struct(scalars, depth)?;
3551 let list_field = Arc::new(Field::new(
3552 "item",
3553 DataType::Struct(het_fields(depth)),
3554 true,
3555 ));
3556 let list = ListArray::new(
3557 list_field,
3558 OffsetBuffer::from_lengths([n]),
3559 Arc::new(elem),
3560 None,
3561 );
3562 Some(DfExpr::Literal(ScalarValue::List(Arc::new(list)), None))
3563}
3564
3565static CYPHER_DYNAMIC_HET_LIST: LazyLock<ScalarUDF> =
3566 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDynamicHetList::new()));
3567
3568#[derive(Debug, PartialEq, Eq, Hash)]
3569struct CypherDynamicHetList {
3570 signature: Signature,
3571}
3572
3573impl CypherDynamicHetList {
3574 fn new() -> Self {
3575 Self {
3576 signature: Signature::variadic_any(Volatility::Immutable),
3577 }
3578 }
3579}
3580
3581fn dynamic_het_type(arg_types: &[DataType]) -> DataType {
3582 use datafusion::arrow::datatypes::Fields;
3583 let mut fields = Vec::with_capacity(arg_types.len() + 1);
3584 fields.push(Field::new("__het_tag", DataType::Int8, false));
3585 fields.extend(
3586 arg_types
3587 .iter()
3588 .enumerate()
3589 .map(|(i, ty)| Field::new(format!("__het_value_{i}"), ty.clone(), true)),
3590 );
3591 DataType::new_list(DataType::Struct(Fields::from(fields)), true)
3592}
3593
3594impl ScalarUDFImpl for CypherDynamicHetList {
3595 fn as_any(&self) -> &dyn Any {
3596 self
3597 }
3598
3599 fn name(&self) -> &'static str {
3600 "cypher_dynamic_het_list"
3601 }
3602
3603 fn signature(&self) -> &Signature {
3604 &self.signature
3605 }
3606
3607 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
3608 Ok(dynamic_het_type(arg_types))
3609 }
3610
3611 fn return_field_from_args(&self, args: ReturnFieldArgs) -> datafusion::error::Result<FieldRef> {
3612 let arg_types = args
3613 .arg_fields
3614 .iter()
3615 .map(|field| field.data_type().clone())
3616 .collect::<Vec<_>>();
3617 Ok(Arc::new(Field::new(
3618 self.name(),
3619 dynamic_het_type(&arg_types),
3620 false,
3621 )))
3622 }
3623
3624 fn invoke_with_args(
3625 &self,
3626 args: ScalarFunctionArgs,
3627 ) -> datafusion::error::Result<ColumnarValue> {
3628 use datafusion::arrow::array::{ArrayRef, Int8Array, Int32Array, StructArray};
3629 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer};
3630 use datafusion::arrow::compute::take;
3631 use datafusion::error::DataFusionError;
3632
3633 let rows = args.number_rows;
3634 let width = args.args.len();
3635 let width_i8 = i8::try_from(width).map_err(|_| {
3636 DataFusionError::Plan("heterogeneous list literal exceeds 127 elements".into())
3637 })?;
3638 let values = args
3639 .args
3640 .iter()
3641 .map(|value| value.to_array(rows))
3642 .collect::<datafusion::error::Result<Vec<_>>>()?;
3643 let tags = Int8Array::from_iter_values((0..rows).flat_map(|_| 0..width_i8));
3644 let mut columns: Vec<ArrayRef> = vec![Arc::new(tags)];
3645 for (value_idx, value) in values.iter().enumerate() {
3646 let indices = (0..rows)
3647 .flat_map(|row| {
3648 (0..width).map(move |element_idx| {
3649 (element_idx == value_idx)
3650 .then(|| i32::try_from(row).ok())
3651 .flatten()
3652 })
3653 })
3654 .collect::<Int32Array>();
3655 columns.push(take(value.as_ref(), &indices, None)?);
3656 }
3657 let valid = (0..rows)
3658 .flat_map(|row| values.iter().map(move |value| !value.is_null(row)))
3659 .collect::<NullBuffer>();
3660 let DataType::List(item) = args.return_field.data_type() else {
3661 return Err(DataFusionError::Internal(
3662 "dynamic heterogeneous list has a non-list return type".into(),
3663 ));
3664 };
3665 let DataType::Struct(fields) = item.data_type() else {
3666 return Err(DataFusionError::Internal(
3667 "dynamic heterogeneous list has a non-struct element type".into(),
3668 ));
3669 };
3670 let elements = StructArray::new(fields.clone(), columns, Some(valid));
3671 let list = ListArray::new(
3672 item.clone(),
3673 OffsetBuffer::from_lengths(std::iter::repeat_n(width, rows)),
3674 Arc::new(elements),
3675 None,
3676 );
3677 Ok(ColumnarValue::Array(Arc::new(list)))
3678 }
3679}
3680
3681fn lower_list_literal(
3682 elems: Vec<DfExpr>,
3683 input_schema: Option<&datafusion::common::DFSchema>,
3684) -> DfExpr {
3685 if let Some(scalars) = elems
3688 .iter()
3689 .map(try_const_scalar)
3690 .collect::<Option<Vec<ScalarValue>>>()
3691 {
3692 let elem_type = scalars
3694 .iter()
3695 .map(ScalarValue::data_type)
3696 .find(|t| *t != DataType::Null)
3697 .unwrap_or(DataType::Int64);
3698 let typed: Vec<ScalarValue> = scalars
3702 .iter()
3703 .map(|s| {
3704 if matches!(s, ScalarValue::Null) {
3705 ScalarValue::try_from(&elem_type).unwrap_or(ScalarValue::Null)
3706 } else {
3707 s.clone()
3708 }
3709 })
3710 .collect();
3711 if typed.iter().all(|s| s.data_type() == elem_type) {
3716 let list = ScalarValue::new_list(&typed, &elem_type, true);
3717 return DfExpr::Literal(ScalarValue::List(list), None);
3718 }
3719 let all_maps = scalars
3727 .iter()
3728 .any(|s| matches!(s, ScalarValue::Struct(a) if is_plain_map_struct(a)))
3729 && scalars.iter().all(|s| {
3730 s.is_null() || matches!(s, ScalarValue::Struct(a) if is_plain_map_struct(a))
3731 });
3732 if all_maps {
3733 if let Some(padded) = all_map_union_list(&scalars) {
3734 return padded;
3735 }
3736 } else if let Some(tagged) = tagged_numeric_list(&scalars) {
3737 return tagged;
3738 }
3739 }
3740 if let Some(schema) = input_schema
3741 && let Some(types) = elems
3742 .iter()
3743 .map(|elem| elem.get_type(schema).ok())
3744 .collect::<Option<Vec<_>>>()
3745 && types.windows(2).any(|pair| pair[0] != pair[1])
3746 {
3747 return CYPHER_DYNAMIC_HET_LIST.call(elems);
3748 }
3749 datafusion::functions_nested::expr_fn::make_array(elems)
3750}
3751
3752static CYPHER_LIST_PLUS: LazyLock<ScalarUDF> =
3753 LazyLock::new(|| ScalarUDF::new_from_impl(CypherListPlus::new()));
3754
3755static CYPHER_RELATIONSHIP_DISJOINT: LazyLock<ScalarUDF> =
3756 LazyLock::new(|| ScalarUDF::new_from_impl(CypherRelationshipDisjoint::new()));
3757
3758pub(crate) fn relationship_disjoint(left: DfExpr, right: DfExpr) -> DfExpr {
3759 CYPHER_RELATIONSHIP_DISJOINT.call(vec![left, right])
3760}
3761
3762#[derive(Debug, PartialEq, Eq, Hash)]
3763struct CypherRelationshipDisjoint {
3764 signature: Signature,
3765}
3766
3767impl CypherRelationshipDisjoint {
3768 fn new() -> Self {
3769 Self {
3770 signature: Signature::any(2, Volatility::Immutable),
3771 }
3772 }
3773}
3774
3775impl ScalarUDFImpl for CypherRelationshipDisjoint {
3776 fn as_any(&self) -> &dyn Any {
3777 self
3778 }
3779
3780 fn name(&self) -> &'static str {
3781 "cypher_relationship_disjoint"
3782 }
3783
3784 fn signature(&self) -> &Signature {
3785 &self.signature
3786 }
3787
3788 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
3789 Ok(DataType::Boolean)
3790 }
3791
3792 fn invoke_with_args(
3793 &self,
3794 args: ScalarFunctionArgs,
3795 ) -> datafusion::error::Result<ColumnarValue> {
3796 use datafusion::arrow::array::BooleanArray;
3797
3798 let rows = args.number_rows;
3799 let left = args.args[0].to_array(rows)?;
3800 let right = args.args[1].to_array(rows)?;
3801 let values = (0..rows)
3802 .map(|row| {
3803 let left = ScalarValue::try_from_array(&left, row)?;
3804 let right = ScalarValue::try_from_array(&right, row)?;
3805 let mut left_ids = Vec::new();
3806 let mut right_ids = Vec::new();
3807 relationship_ids(&left, &mut left_ids);
3808 relationship_ids(&right, &mut right_ids);
3809 Ok(!left_ids.iter().any(|id| right_ids.contains(id)))
3810 })
3811 .collect::<datafusion::error::Result<BooleanArray>>()?;
3812 Ok(ColumnarValue::Array(std::sync::Arc::new(values)))
3813 }
3814}
3815
3816fn relationship_ids(value: &ScalarValue, ids: &mut Vec<Vec<u8>>) {
3817 match value {
3818 ScalarValue::FixedSizeBinary(_, Some(uuid)) => ids.push(uuid.clone()),
3819 ScalarValue::List(list) if !list.is_null(0) => {
3820 let values = list.value(0);
3821 for index in 0..values.len() {
3822 if let Ok(value) = ScalarValue::try_from_array(&values, index) {
3823 relationship_ids(&value, ids);
3824 }
3825 }
3826 }
3827 ScalarValue::Struct(value) if !value.is_null(0) => {
3828 if let Some(uuid) = value.column_by_name("edge_uuid")
3829 && let Ok(uuid) = ScalarValue::try_from_array(uuid, 0)
3830 {
3831 relationship_ids(&uuid, ids);
3832 }
3833 }
3834 _ => {}
3835 }
3836}
3837
3838fn graph_value_types_compatible(left: &DataType, right: &DataType) -> bool {
3842 let (DataType::Struct(left), DataType::Struct(right)) = (left, right) else {
3843 return false;
3844 };
3845 left.len() == right.len()
3846 && left.iter().zip(right.iter()).all(|(left, right)| {
3847 left.name() == right.name()
3848 && match (left.data_type(), right.data_type()) {
3849 (DataType::Struct(_), DataType::Struct(_)) => {
3850 graph_value_types_compatible(left.data_type(), right.data_type())
3851 }
3852 (DataType::List(left), DataType::List(right)) => {
3853 left.data_type() == right.data_type()
3854 }
3855 (left, right) => left == right,
3856 }
3857 })
3858}
3859
3860#[derive(Debug, PartialEq, Eq, Hash)]
3861struct CypherListPlus {
3862 signature: Signature,
3863}
3864
3865impl CypherListPlus {
3866 fn new() -> Self {
3867 Self {
3868 signature: Signature::any(2, Volatility::Immutable),
3869 }
3870 }
3871}
3872
3873impl ScalarUDFImpl for CypherListPlus {
3874 fn as_any(&self) -> &dyn Any {
3875 self
3876 }
3877
3878 fn name(&self) -> &'static str {
3879 "cypher_list_plus"
3880 }
3881
3882 fn signature(&self) -> &Signature {
3883 &self.signature
3884 }
3885
3886 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
3887 if list_plus_has_graph_value(arg_types) {
3888 Ok(list_plus_return_type(arg_types))
3889 } else {
3890 Ok(DataType::new_list(
3891 DataType::Struct(het_fields(list_plus_depth(arg_types))),
3892 true,
3893 ))
3894 }
3895 }
3896
3897 #[allow(
3898 clippy::too_many_lines,
3899 reason = "list/list and list/element shaping share one offset and validity pass"
3900 )]
3901 fn invoke_with_args(
3902 &self,
3903 args: ScalarFunctionArgs,
3904 ) -> datafusion::error::Result<ColumnarValue> {
3905 use datafusion::arrow::array::{Array, ArrayRef, ListArray};
3906 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
3907 use datafusion::arrow::datatypes::DataType;
3908 use datafusion::error::DataFusionError;
3909 use std::sync::Arc;
3910
3911 let rows = args.number_rows;
3912 let left = args.args[0].to_array(rows)?;
3913 let right = args.args[1].to_array(rows)?;
3914 let left_is_list = list_item_type(left.data_type()).is_some();
3915 let right_is_list = list_item_type(right.data_type()).is_some();
3916 if !left_is_list && !right_is_list {
3917 return Err(DataFusionError::Execution(
3918 "list + requires at least one list operand".into(),
3919 ));
3920 }
3921
3922 if left_is_list
3923 && !right_is_list
3924 && let Some(result) =
3925 invoke_tagged_list_element_plus(&left, &right, args.return_field.data_type())?
3926 {
3927 return Ok(ColumnarValue::Array(result));
3928 }
3929
3930 let mut flat: Vec<ScalarValue> = Vec::new();
3931 let mut offsets: Vec<i32> = Vec::with_capacity(rows + 1);
3932 let mut validity: Vec<bool> = Vec::with_capacity(rows);
3933 offsets.push(0);
3934
3935 for row in 0..rows {
3936 let row_values = match (left_is_list, right_is_list) {
3937 (true, true) => match (
3938 list_elements_at(&left, row)?,
3939 list_elements_at(&right, row)?,
3940 ) {
3941 (Some(mut l), Some(r)) => {
3942 l.extend(r);
3943 Some(l)
3944 }
3945 _ => None,
3946 },
3947 (true, false) => match list_elements_at(&left, row)? {
3948 Some(mut l) => {
3949 let r = decoded_scalar_at(&right, row)?;
3950 if let Some(r) = scalar_list_elements(&r)? {
3951 l.extend(r);
3952 } else {
3953 l.push(r);
3954 }
3955 Some(l)
3956 }
3957 None => None,
3958 },
3959 (false, true) => match list_elements_at(&right, row)? {
3960 Some(mut r) => {
3961 let l = decoded_scalar_at(&left, row)?;
3962 let mut l = scalar_list_elements(&l)?.unwrap_or_else(|| vec![l]);
3963 l.append(&mut r);
3964 Some(l)
3965 }
3966 None => None,
3967 },
3968 (false, false) => unreachable!("checked above"),
3969 };
3970
3971 match row_values {
3972 Some(values) => {
3973 flat.extend(values);
3974 validity.push(true);
3975 }
3976 None => validity.push(false),
3977 }
3978 offsets.push(i32::try_from(flat.len()).map_err(|_| {
3979 DataFusionError::Execution("cypher_list_plus: list too long".into())
3980 })?);
3981 }
3982
3983 let DataType::List(field) = args.return_field.data_type() else {
3984 return Err(DataFusionError::Internal(
3985 "cypher_list_plus return type is not a list".into(),
3986 ));
3987 };
3988 if is_het_struct_type(Some(field.data_type()))
3989 && !is_dynamic_variant_struct(field.data_type())
3990 {
3991 let depth = list_plus_depth(&[left.data_type().clone(), right.data_type().clone()]);
3992 let values = build_het_struct(&flat, depth).ok_or_else(|| {
3993 DataFusionError::Execution(
3994 "cypher_list_plus: cannot encode value in heterogeneous list".into(),
3995 )
3996 })?;
3997 let out = ListArray::new(
3998 field.clone(),
3999 OffsetBuffer::new(ScalarBuffer::from(offsets)),
4000 Arc::new(values) as ArrayRef,
4001 Some(NullBuffer::from(validity)),
4002 );
4003 return Ok(ColumnarValue::Array(Arc::new(out)));
4004 }
4005 if !is_dynamic_variant_struct(field.data_type()) {
4006 let flat = flat
4007 .into_iter()
4008 .map(|value| {
4009 if value.data_type() == *field.data_type() {
4010 Ok(value)
4011 } else {
4012 value.cast_to(field.data_type())
4013 }
4014 })
4015 .collect::<datafusion::error::Result<Vec<_>>>()?;
4016 let values = if flat.is_empty() {
4017 new_empty_array(field.data_type())
4018 } else {
4019 ScalarValue::iter_to_array(flat)?
4020 };
4021 let out = ListArray::new(
4022 field.clone(),
4023 OffsetBuffer::new(ScalarBuffer::from(offsets)),
4024 values,
4025 Some(NullBuffer::from(validity)),
4026 );
4027 return Ok(ColumnarValue::Array(Arc::new(out)));
4028 }
4029 let DataType::Struct(fields) = field.data_type() else {
4030 return Err(DataFusionError::Internal(
4031 "cypher_list_plus element type is not tagged".into(),
4032 ));
4033 };
4034 let variants = fields
4035 .iter()
4036 .filter(|field| field.name().starts_with("__het_value_"))
4037 .map(|field| field.data_type().clone())
4038 .collect::<Vec<_>>();
4039 let mut tags = Vec::with_capacity(flat.len());
4040 let mut valid = Vec::with_capacity(flat.len());
4041 let mut columns = Vec::with_capacity(variants.len() + 1);
4042 for value in &flat {
4043 let tag = variants
4044 .iter()
4045 .position(|variant| {
4046 value.data_type() == *variant
4047 || graph_value_types_compatible(&value.data_type(), variant)
4048 })
4049 .unwrap_or(0);
4050 tags.push(i8::try_from(tag).map_err(|_| {
4051 DataFusionError::Execution("cypher_list_plus has too many value variants".into())
4052 })?);
4053 valid.push(!value.is_null());
4054 }
4055 columns.push(Arc::new(datafusion::arrow::array::Int8Array::from(tags.clone())) as ArrayRef);
4056 for (variant_index, variant) in variants.iter().enumerate() {
4057 let null = ScalarValue::try_new_null(variant)?;
4058 let values = flat.iter().zip(&tags).map(|(value, tag)| {
4059 if usize::try_from(*tag).ok() == Some(variant_index) {
4060 value.clone()
4061 } else {
4062 null.clone()
4063 }
4064 });
4065 columns.push(ScalarValue::iter_to_array(values)?);
4066 }
4067 let values = datafusion::arrow::array::StructArray::new(
4068 fields.clone(),
4069 columns,
4070 Some(NullBuffer::from(valid)),
4071 );
4072 let out = ListArray::new(
4073 field.clone(),
4074 OffsetBuffer::new(ScalarBuffer::from(offsets)),
4075 Arc::new(values) as ArrayRef,
4076 Some(NullBuffer::from(validity)),
4077 );
4078 Ok(ColumnarValue::Array(Arc::new(out)))
4079 }
4080}
4081
4082#[allow(
4087 clippy::too_many_lines,
4088 reason = "one range-assembly pass keeps offsets, validity, and three Arrow sources synchronized"
4089)]
4090fn invoke_tagged_list_element_plus(
4091 left: &datafusion::arrow::array::ArrayRef,
4092 right: &datafusion::arrow::array::ArrayRef,
4093 return_type: &DataType,
4094) -> datafusion::error::Result<Option<datafusion::arrow::array::ArrayRef>> {
4095 use arrow_data::transform::MutableArrayData;
4096 use datafusion::arrow::array::{Array, Int8Array, ListArray, StructArray, make_array};
4097 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
4098 use datafusion::arrow::datatypes::DataType;
4099 use datafusion::arrow::error::ArrowError;
4100 use datafusion::error::DataFusionError;
4101 use std::sync::Arc;
4102
4103 let Some(left) = left.as_any().downcast_ref::<ListArray>() else {
4104 return Ok(None);
4105 };
4106 let Some(right) = right.as_any().downcast_ref::<StructArray>() else {
4107 return Ok(None);
4108 };
4109 let DataType::List(return_field) = return_type else {
4110 return Ok(None);
4111 };
4112 if left.value_type() != right.data_type().clone()
4113 || return_field.data_type() != right.data_type()
4114 || !is_het_struct_type(Some(right.data_type()))
4115 {
4116 return Ok(None);
4117 }
4118
4119 let Some(tags) = right
4120 .column_by_name("__het_tag")
4121 .and_then(|column| column.as_any().downcast_ref::<Int8Array>())
4122 else {
4123 return Ok(None);
4124 };
4125 let Some(nested) = right
4126 .column_by_name("__het_list")
4127 .and_then(|column| column.as_any().downcast_ref::<ListArray>())
4128 else {
4129 return Ok(None);
4130 };
4131
4132 let nested_values = nested.values();
4133 let nested_offsets = nested.value_offsets();
4134 let mut promoted_ranges = vec![None; right.len()];
4135 for row in 0..right.len() {
4136 if right.is_null(row) || tags.value(row) != 4 || nested.is_null(row) {
4137 continue;
4138 }
4139 let start = usize::try_from(nested_offsets[row]).map_err(|_| {
4140 DataFusionError::ArrowError(
4141 Box::new(ArrowError::ComputeError(
4142 "negative heterogeneous-list offset".into(),
4143 )),
4144 None,
4145 )
4146 })?;
4147 let end = usize::try_from(nested_offsets[row + 1]).map_err(|_| {
4148 DataFusionError::ArrowError(
4149 Box::new(ArrowError::ComputeError(
4150 "negative heterogeneous-list offset".into(),
4151 )),
4152 None,
4153 )
4154 })?;
4155 promoted_ranges[row] = Some((start, end));
4156 }
4157 let promoted = if nested_values.is_empty() {
4158 Arc::new(right.slice(0, 0)) as datafusion::arrow::array::ArrayRef
4159 } else {
4160 promote_het_array(nested_values, right.data_type())?
4161 };
4162
4163 let left_data = left.values().to_data();
4164 let right_data = right.to_data();
4165 let promoted_data = promoted.to_data();
4166 let capacity = left.values().len() + right.len() + promoted.len();
4167 let mut values = MutableArrayData::new(
4168 vec![&left_data, &right_data, &promoted_data],
4169 true,
4170 capacity,
4171 );
4172 let left_offsets = left.value_offsets();
4173 let mut offsets = Vec::with_capacity(left.len() + 1);
4174 let mut validity = Vec::with_capacity(left.len());
4175 let mut output_len = 0usize;
4176 offsets.push(0i32);
4177 for row in 0..left.len() {
4178 if left.is_null(row) {
4179 validity.push(false);
4180 offsets.push(i32::try_from(output_len).map_err(|_| {
4181 DataFusionError::Execution("cypher_list_plus: list too long".into())
4182 })?);
4183 continue;
4184 }
4185 validity.push(true);
4186 let start = usize::try_from(left_offsets[row]).map_err(|_| {
4187 DataFusionError::Execution("cypher_list_plus: negative list offset".into())
4188 })?;
4189 let end = usize::try_from(left_offsets[row + 1]).map_err(|_| {
4190 DataFusionError::Execution("cypher_list_plus: negative list offset".into())
4191 })?;
4192 values.extend(0, start, end);
4193 output_len += end - start;
4194 if let Some((start, end)) = promoted_ranges[row] {
4195 values.extend(2, start, end);
4196 output_len += end - start;
4197 } else {
4198 values.extend(1, row, row + 1);
4199 output_len += 1;
4200 }
4201 offsets.push(
4202 i32::try_from(output_len).map_err(|_| {
4203 DataFusionError::Execution("cypher_list_plus: list too long".into())
4204 })?,
4205 );
4206 }
4207
4208 let values = make_array(values.freeze());
4209 Ok(Some(Arc::new(ListArray::new(
4210 return_field.clone(),
4211 OffsetBuffer::new(ScalarBuffer::from(offsets)),
4212 values,
4213 Some(NullBuffer::from(validity)),
4214 ))))
4215}
4216
4217fn promote_het_array(
4221 source: &datafusion::arrow::array::ArrayRef,
4222 target: &DataType,
4223) -> datafusion::error::Result<datafusion::arrow::array::ArrayRef> {
4224 use datafusion::arrow::array::{Array, ListArray, StructArray, new_null_array};
4225 use datafusion::arrow::compute::cast;
4226 use datafusion::arrow::datatypes::DataType;
4227 use datafusion::error::DataFusionError;
4228 use std::sync::Arc;
4229
4230 if source.data_type() == target {
4231 return Ok(source.clone());
4232 }
4233 match (source.data_type(), target) {
4234 (DataType::Struct(_), DataType::Struct(target_fields)) => {
4235 let source = source
4236 .as_any()
4237 .downcast_ref::<StructArray>()
4238 .ok_or_else(|| {
4239 DataFusionError::Internal(
4240 "heterogeneous value has a non-struct physical array".into(),
4241 )
4242 })?;
4243 let columns = target_fields
4244 .iter()
4245 .map(|field| {
4246 source.column_by_name(field.name()).map_or_else(
4247 || Ok(new_null_array(field.data_type(), source.len())),
4248 |column| promote_het_array(column, field.data_type()),
4249 )
4250 })
4251 .collect::<datafusion::error::Result<Vec<_>>>()?;
4252 Ok(Arc::new(StructArray::new(
4253 target_fields.clone(),
4254 columns,
4255 source.nulls().cloned(),
4256 )))
4257 }
4258 (DataType::List(_), DataType::List(target_field)) => {
4259 let source = source.as_any().downcast_ref::<ListArray>().ok_or_else(|| {
4260 DataFusionError::Internal("heterogeneous list has a non-list physical array".into())
4261 })?;
4262 let values = promote_het_array(source.values(), target_field.data_type())?;
4263 Ok(Arc::new(ListArray::new(
4264 target_field.clone(),
4265 source.offsets().clone(),
4266 values,
4267 source.nulls().cloned(),
4268 )))
4269 }
4270 _ => Ok(cast(source, target)?),
4271 }
4272}
4273
4274fn list_plus_return_type(arg_types: &[DataType]) -> DataType {
4275 use datafusion::arrow::datatypes::Fields;
4276
4277 let value_types = arg_types
4278 .iter()
4279 .flat_map(|arg_type| {
4280 let value_type = list_item_type(arg_type).unwrap_or(arg_type);
4281 dynamic_variant_types(value_type)
4282 })
4283 .filter(|value_type| !matches!(value_type, DataType::Null))
4284 .collect::<Vec<_>>();
4285 if let Some(first) = value_types.first()
4286 && is_graph_value_struct(first)
4287 && value_types
4288 .iter()
4289 .all(|value_type| graph_value_types_compatible(first, value_type))
4290 {
4291 return DataType::new_list((*first).clone(), true);
4292 }
4293
4294 let mut variants = Vec::new();
4295 for arg_type in arg_types {
4296 let value_type = list_item_type(arg_type).unwrap_or(arg_type);
4297 if let DataType::Struct(fields) = value_type
4298 && fields
4299 .iter()
4300 .any(|field| field.name().starts_with("__het_value_"))
4301 {
4302 for field in fields
4303 .iter()
4304 .filter(|field| field.name().starts_with("__het_value_"))
4305 {
4306 if !variants.contains(field.data_type()) {
4307 variants.push(field.data_type().clone());
4308 }
4309 }
4310 } else if let DataType::Struct(fields) = value_type
4311 && fields.iter().any(|field| field.name() == "__het_tag")
4312 {
4313 for (name, data_type) in [
4314 ("__het_int", DataType::Int64),
4315 ("__het_float", DataType::Float64),
4316 ("__het_str", DataType::Utf8),
4317 ("__het_bool", DataType::Boolean),
4318 ] {
4319 if fields.iter().any(|field| field.name() == name) && !variants.contains(&data_type)
4320 {
4321 variants.push(data_type);
4322 }
4323 }
4324 } else if !matches!(value_type, DataType::Null) && !variants.contains(value_type) {
4325 variants.push(value_type.clone());
4326 }
4327 }
4328 if variants.is_empty() {
4329 variants.push(DataType::Null);
4330 }
4331 let mut fields = vec![Field::new("__het_tag", DataType::Int8, false)];
4332 fields.extend(
4333 variants
4334 .into_iter()
4335 .enumerate()
4336 .map(|(index, data_type)| Field::new(format!("__het_value_{index}"), data_type, true)),
4337 );
4338 DataType::new_list(DataType::Struct(Fields::from(fields)), true)
4339}
4340
4341fn list_plus_has_graph_value(arg_types: &[DataType]) -> bool {
4342 arg_types.iter().any(|arg_type| {
4343 let value_type = list_item_type(arg_type).unwrap_or(arg_type);
4344 dynamic_variant_types(value_type)
4345 .into_iter()
4346 .any(is_graph_value_struct)
4347 })
4348}
4349
4350fn is_graph_value_struct(data_type: &DataType) -> bool {
4351 matches!(data_type, DataType::Struct(fields) if fields.iter().any(|field| {
4352 matches!(
4353 field.name().as_str(),
4354 "node_uuid" | "edge_uuid" | "nodes" | "relationships"
4355 )
4356 }))
4357}
4358
4359fn is_dynamic_variant_struct(data_type: &DataType) -> bool {
4360 matches!(data_type, DataType::Struct(fields) if fields
4361 .iter()
4362 .any(|field| field.name().starts_with("__het_value_")))
4363}
4364
4365fn dynamic_variant_types(data_type: &DataType) -> Vec<&DataType> {
4366 if let DataType::Struct(fields) = data_type {
4367 let variants = fields
4368 .iter()
4369 .filter(|field| field.name().starts_with("__het_value_"))
4370 .map(|field| field.data_type())
4371 .collect::<Vec<_>>();
4372 if !variants.is_empty() {
4373 return variants;
4374 }
4375 }
4376 vec![data_type]
4377}
4378
4379fn list_item_type(dt: &DataType) -> Option<&DataType> {
4380 match dt {
4381 DataType::List(f) | DataType::LargeList(f) | DataType::FixedSizeList(f, _) => {
4382 Some(f.data_type())
4383 }
4384 _ => None,
4385 }
4386}
4387
4388fn list_plus_depth(arg_types: &[DataType]) -> usize {
4389 arg_types
4390 .iter()
4391 .filter_map(|data_type| {
4392 list_item_type(data_type)
4393 .or(Some(data_type))
4394 .and_then(het_depth_for_data_type)
4395 })
4396 .max()
4397 .unwrap_or(0)
4398}
4399
4400fn het_depth_for_data_type(data_type: &DataType) -> Option<usize> {
4401 if is_het_struct_type(Some(data_type)) {
4402 return het_struct_type_depth(data_type);
4403 }
4404 match data_type {
4405 DataType::Null
4406 | DataType::Boolean
4407 | DataType::Int8
4408 | DataType::Int16
4409 | DataType::Int32
4410 | DataType::Int64
4411 | DataType::UInt8
4412 | DataType::UInt16
4413 | DataType::UInt32
4414 | DataType::UInt64
4415 | DataType::Float16
4416 | DataType::Float32
4417 | DataType::Float64
4418 | DataType::Utf8
4419 | DataType::LargeUtf8 => Some(0),
4420 DataType::List(field) | DataType::LargeList(field) | DataType::FixedSizeList(field, _) => {
4421 Some(1 + het_depth_for_data_type(field.data_type())?)
4422 }
4423 DataType::Struct(fields) if is_plain_map_struct_type(data_type) => fields
4424 .iter()
4425 .filter_map(|field| het_depth_for_data_type(field.data_type()))
4426 .max()
4427 .map_or(Some(1), |depth| Some(1 + depth)),
4428 _ => None,
4429 }
4430}
4431
4432fn het_struct_type_depth(data_type: &DataType) -> Option<usize> {
4433 let DataType::Struct(fields) = data_type else {
4434 return None;
4435 };
4436 let Some(list_field) = fields.iter().find(|field| field.name() == "__het_list") else {
4437 return Some(0);
4438 };
4439 match list_field.data_type() {
4440 DataType::List(inner) => het_struct_type_depth(inner.data_type()).map(|depth| depth + 1),
4441 _ => Some(0),
4442 }
4443}
4444
4445fn list_elements_at(
4446 array: &datafusion::arrow::array::ArrayRef,
4447 row: usize,
4448) -> datafusion::error::Result<Option<Vec<ScalarValue>>> {
4449 use datafusion::arrow::array::{Array, FixedSizeListArray, LargeListArray, ListArray};
4450
4451 let values = if let Some(list) = array.as_any().downcast_ref::<ListArray>() {
4452 if list.is_null(row) {
4453 return Ok(None);
4454 }
4455 list.value(row)
4456 } else if let Some(list) = array.as_any().downcast_ref::<LargeListArray>() {
4457 if list.is_null(row) {
4458 return Ok(None);
4459 }
4460 list.value(row)
4461 } else if let Some(list) = array.as_any().downcast_ref::<FixedSizeListArray>() {
4462 if list.is_null(row) {
4463 return Ok(None);
4464 }
4465 list.value(row)
4466 } else {
4467 return Ok(None);
4468 };
4469
4470 (0..values.len())
4471 .map(|i| ScalarValue::try_from_array(&values, i).map(unwrap_het))
4472 .collect::<datafusion::error::Result<Vec<_>>>()
4473 .map(Some)
4474}
4475
4476fn scalar_list_elements(
4477 value: &ScalarValue,
4478) -> datafusion::error::Result<Option<Vec<ScalarValue>>> {
4479 use datafusion::arrow::array::Array;
4480
4481 match value {
4482 ScalarValue::List(list) => {
4483 if list.is_null(0) {
4484 return Ok(None);
4485 }
4486 let values = list.value(0);
4487 (0..values.len())
4488 .map(|i| ScalarValue::try_from_array(&values, i).map(unwrap_het))
4489 .collect::<datafusion::error::Result<Vec<_>>>()
4490 .map(Some)
4491 }
4492 ScalarValue::LargeList(list) => {
4493 if list.is_null(0) {
4494 return Ok(None);
4495 }
4496 let values = list.value(0);
4497 (0..values.len())
4498 .map(|i| ScalarValue::try_from_array(&values, i).map(unwrap_het))
4499 .collect::<datafusion::error::Result<Vec<_>>>()
4500 .map(Some)
4501 }
4502 _ => Ok(None),
4503 }
4504}
4505
4506fn all_map_union_list(scalars: &[ScalarValue]) -> Option<DfExpr> {
4513 use datafusion::arrow::array::{Array, ArrayRef, StructArray, new_null_array};
4514 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer};
4515 use datafusion::arrow::compute::{cast, concat};
4516 use datafusion::arrow::datatypes::{Field, Fields};
4517 use std::collections::HashMap;
4518 use std::sync::Arc;
4519
4520 let mut order: Vec<String> = Vec::new();
4523 let mut types: HashMap<String, DataType> = HashMap::new();
4524 for s in scalars {
4525 let arr = match s {
4526 ScalarValue::Struct(a) => a,
4527 ScalarValue::Null => continue,
4528 _ => return None,
4529 };
4530 for f in arr.fields() {
4531 let t = f.data_type().clone();
4532 match types.get(f.name()) {
4533 None => {
4534 order.push(f.name().clone());
4535 types.insert(f.name().clone(), t);
4536 }
4537 Some(prev) if *prev == DataType::Null => {
4538 types.insert(f.name().clone(), t);
4539 }
4540 Some(prev) if t != DataType::Null && t != *prev => return None,
4541 _ => {}
4542 }
4543 }
4544 }
4545 let union_fields: Fields = order
4546 .iter()
4547 .map(|n| Field::new(n, types.get(n).cloned().unwrap_or(DataType::Null), true))
4548 .collect::<Vec<_>>()
4549 .into();
4550
4551 let mut columns: Vec<ArrayRef> = Vec::with_capacity(order.len());
4554 for name in &order {
4555 let ut = types.get(name).cloned().unwrap_or(DataType::Null);
4556 let mut pieces: Vec<ArrayRef> = Vec::with_capacity(scalars.len());
4557 for s in scalars {
4558 let piece = match s {
4559 ScalarValue::Struct(a) => a
4560 .column_by_name(name)
4561 .and_then(|c| cast(c, &ut).ok())
4562 .unwrap_or_else(|| new_null_array(&ut, 1)),
4563 _ => new_null_array(&ut, 1),
4564 };
4565 pieces.push(piece);
4566 }
4567 let refs: Vec<&dyn Array> = pieces.iter().map(AsRef::as_ref).collect();
4568 columns.push(concat(&refs).ok()?);
4569 }
4570 let valid: NullBuffer = scalars.iter().map(|s| !s.is_null()).collect();
4572 let elem = StructArray::try_new(union_fields, columns, Some(valid)).ok()?;
4573 let n = scalars.len();
4574 let list_field = Arc::new(Field::new("item", elem.data_type().clone(), true));
4575 let list = datafusion::arrow::array::ListArray::new(
4576 list_field,
4577 OffsetBuffer::from_lengths([n]),
4578 Arc::new(elem),
4579 None,
4580 );
4581 Some(DfExpr::Literal(ScalarValue::List(Arc::new(list)), None))
4582}
4583
4584fn lower_literal(lit_val: &IrLiteral) -> DfExpr {
4586 lit(ir_literal_to_scalar(lit_val))
4587}
4588
4589#[must_use]
4595pub fn ir_literal_to_scalar(lit_val: &IrLiteral) -> ScalarValue {
4596 match lit_val {
4597 IrLiteral::Null => ScalarValue::Null,
4598 IrLiteral::Bool(b) => ScalarValue::Boolean(Some(*b)),
4599 IrLiteral::Int(n) => ScalarValue::Int64(Some(*n)),
4600 IrLiteral::Float(f) => ScalarValue::Float64(Some(*f)),
4601 IrLiteral::Str(s) => ScalarValue::Utf8(Some(s.clone())),
4602 IrLiteral::Uuid(uuid) => ScalarValue::FixedSizeBinary(16, Some(uuid.to_vec())),
4603 IrLiteral::Duration {
4604 months,
4605 days,
4606 seconds,
4607 nanos,
4608 } => duration_scalar(Some(crate::temporal::DurationValue {
4609 months: *months,
4610 days: *days,
4611 seconds: *seconds,
4612 nanos: *nanos,
4613 })),
4614 IrLiteral::DateTime(us) => ScalarValue::TimestampMicrosecond(Some(*us), Some("UTC".into())),
4615 IrLiteral::Date(days) => date_scalar(Some(*days)),
4616 IrLiteral::LocalDateTime { days, nanos } => localdatetime_scalar(Some((*days, *nanos))),
4617 IrLiteral::Time(nanos) => ScalarValue::Time64Nanosecond(Some(*nanos)),
4618 IrLiteral::ZonedTime { nanos, offset } => time_scalar(Some((*nanos, *offset))),
4619 IrLiteral::ZonedDateTime {
4620 days,
4621 nanos,
4622 offset,
4623 zone,
4624 } => datetime_scalar(Some((*days, *nanos, *offset, zone.clone()))),
4625 IrLiteral::List(items) => {
4629 let scalars: Vec<ScalarValue> = items.iter().map(ir_literal_to_scalar).collect();
4630 let elem_type = scalars
4631 .iter()
4632 .find(|s| !s.is_null())
4633 .map_or(DataType::Null, ScalarValue::data_type);
4634 let typed: Vec<ScalarValue> = scalars
4635 .iter()
4636 .map(|s| {
4637 if matches!(s, ScalarValue::Null) {
4638 ScalarValue::try_from(&elem_type).unwrap_or(ScalarValue::Null)
4639 } else {
4640 s.clone()
4641 }
4642 })
4643 .collect();
4644 ScalarValue::List(ScalarValue::new_list(&typed, &elem_type, true))
4645 }
4646 IrLiteral::Map(entries) => {
4647 let scalars: Vec<(String, ScalarValue)> = entries
4648 .iter()
4649 .map(|(key, value)| (key.clone(), ir_literal_to_scalar(value)))
4650 .collect();
4651 const_map_scalar(&scalars).expect("IR map literal should lower to an Arrow struct")
4652 }
4653 }
4654}
4655
4656pub fn scalar_to_ir_literal(value: &ScalarValue) -> Result<IrLiteral, LoweringError> {
4673 if value.is_null() {
4675 return Ok(IrLiteral::Null);
4676 }
4677 let lit = match value {
4678 ScalarValue::Boolean(Some(b)) => IrLiteral::Bool(*b),
4679 ScalarValue::Int8(Some(n)) => IrLiteral::Int(i64::from(*n)),
4680 ScalarValue::Int16(Some(n)) => IrLiteral::Int(i64::from(*n)),
4681 ScalarValue::Int32(Some(n)) => IrLiteral::Int(i64::from(*n)),
4682 ScalarValue::Int64(Some(n)) => IrLiteral::Int(*n),
4683 ScalarValue::UInt8(Some(n)) => IrLiteral::Int(i64::from(*n)),
4684 ScalarValue::UInt16(Some(n)) => IrLiteral::Int(i64::from(*n)),
4685 ScalarValue::UInt32(Some(n)) => IrLiteral::Int(i64::from(*n)),
4686 ScalarValue::UInt64(Some(n)) => i64::try_from(*n).map(IrLiteral::Int).map_err(|_| {
4687 LoweringError::UnsupportedExpr(format!("SET value {n} exceeds the i64 range"))
4688 })?,
4689 ScalarValue::Float32(Some(f)) => IrLiteral::Float(f64::from(*f)),
4690 ScalarValue::Float64(Some(f)) => IrLiteral::Float(*f),
4691 ScalarValue::Utf8(Some(s))
4692 | ScalarValue::LargeUtf8(Some(s))
4693 | ScalarValue::Utf8View(Some(s)) => IrLiteral::Str(s.clone()),
4694 ScalarValue::FixedSizeBinary(16, Some(_)) => {
4695 return Err(LoweringError::InvalidType(
4696 "UUID values cannot be stored as graph properties".into(),
4697 ));
4698 }
4699 ScalarValue::DurationSecond(Some(s)) => duration_value_to_ir(dur_secs_nanos(*s, 0)),
4704 ScalarValue::DurationMillisecond(Some(ms)) => duration_value_to_ir(dur_secs_nanos(
4705 ms.div_euclid(1_000),
4706 ms.rem_euclid(1_000) * 1_000_000,
4707 )),
4708 ScalarValue::DurationMicrosecond(Some(us)) => duration_value_to_ir(dur_secs_nanos(
4709 us.div_euclid(1_000_000),
4710 us.rem_euclid(1_000_000) * 1_000,
4711 )),
4712 ScalarValue::DurationNanosecond(Some(ns)) => {
4713 duration_value_to_ir(crate::temporal::DurationValue::from_total_nanos(0, 0, *ns))
4714 }
4715 ScalarValue::TimestampMicrosecond(Some(us), _) => IrLiteral::DateTime(*us),
4716 ScalarValue::TimestampSecond(Some(s), _) => IrLiteral::DateTime(s * 1_000_000),
4717 ScalarValue::TimestampMillisecond(Some(ms), _) => IrLiteral::DateTime(ms * 1_000),
4718 ScalarValue::TimestampNanosecond(Some(ns), _) => IrLiteral::DateTime(ns / 1_000),
4719 ScalarValue::Struct(arr) if is_date_struct(&DataType::Struct(arr.fields().clone())) => {
4723 match date_struct_value(arr, 0) {
4724 Some(days) => IrLiteral::Date(days),
4725 None => IrLiteral::Null,
4726 }
4727 }
4728 ScalarValue::Struct(arr) if is_duration_struct(&DataType::Struct(arr.fields().clone())) => {
4730 match duration_struct_parts(arr, 0) {
4731 Some(d) => duration_value_to_ir(d),
4732 None => IrLiteral::Null,
4733 }
4734 }
4735 ScalarValue::Struct(arr)
4737 if is_localdatetime_struct(&DataType::Struct(arr.fields().clone())) =>
4738 {
4739 match localdatetime_struct_parts(arr, 0) {
4740 Some((days, nanos)) => IrLiteral::LocalDateTime { days, nanos },
4741 None => IrLiteral::Null,
4742 }
4743 }
4744 ScalarValue::Time64Nanosecond(Some(n)) => IrLiteral::Time(*n),
4746 ScalarValue::Struct(arr) if is_time_struct(&DataType::Struct(arr.fields().clone())) => {
4748 match time_struct_parts(arr, 0) {
4749 Some((nanos, offset)) => IrLiteral::ZonedTime { nanos, offset },
4750 None => IrLiteral::Null,
4751 }
4752 }
4753 ScalarValue::Struct(arr) if is_datetime_struct(&DataType::Struct(arr.fields().clone())) => {
4755 match datetime_struct_parts(arr, 0) {
4756 Some((days, nanos, offset, zone)) => IrLiteral::ZonedDateTime {
4757 days,
4758 nanos,
4759 offset,
4760 zone,
4761 },
4762 None => IrLiteral::Null,
4763 }
4764 }
4765 ScalarValue::List(arr) => list_scalar_to_ir_literal(&arr.value(0))?,
4770 ScalarValue::LargeList(arr) => list_scalar_to_ir_literal(&arr.value(0))?,
4771 other => {
4772 return Err(LoweringError::InvalidType(format!(
4775 "invalid property type for SET value: {:?}",
4776 other.data_type()
4777 )));
4778 }
4779 };
4780 Ok(lit)
4781}
4782
4783fn list_scalar_to_ir_literal(
4786 elems: &datafusion::arrow::array::ArrayRef,
4787) -> Result<IrLiteral, LoweringError> {
4788 let mut items = Vec::with_capacity(elems.len());
4789 for j in 0..elems.len() {
4790 let ev = ScalarValue::try_from_array(elems, j).map_err(|e| {
4791 LoweringError::UnsupportedExpr(format!("list element is not a scalar value: {e}"))
4792 })?;
4793 items.push(scalar_to_ir_literal(&ev)?);
4794 }
4795 Ok(IrLiteral::List(items))
4796}
4797
4798fn render_temporal(name: &str, s: &str) -> Option<String> {
4803 use crate::temporal;
4804 match name {
4805 "date" => temporal::render_date(s),
4806 "localtime" => temporal::render_local_time(s),
4807 "time" => temporal::render_time(s),
4808 "localdatetime" => temporal::render_local_date_time(s),
4809 "datetime" => temporal::render_date_time(s),
4810 "duration" => temporal::render_duration(s),
4811 _ => None,
4812 }
4813}
4814
4815#[allow(
4819 clippy::too_many_lines,
4820 reason = "a flat one-arm-per-builtin dispatch table; clearest kept inline"
4821)]
4822fn resolve_builtin(
4823 name: &str,
4824 args: Vec<DfExpr>,
4825 path_hydration: impl FnOnce() -> Option<PathNodeHydration>,
4826) -> Option<DfExpr> {
4827 use datafusion::functions::math::expr_fn as mfn;
4828 use datafusion::functions::string::expr_fn as sfn;
4829
4830 let name = name.to_ascii_lowercase();
4831 let mut a = args;
4832 match name.as_str() {
4833 "toupper" | "upper" => Some(sfn::upper(a.remove(0))),
4835 "tolower" | "lower" => Some(sfn::lower(a.remove(0))),
4836 "trim" => Some(sfn::btrim(a)),
4837 "ltrim" => Some(sfn::ltrim(a)),
4838 "rtrim" => Some(sfn::rtrim(a)),
4839 "string.concat" | "concat" => Some(sfn::concat(a)),
4840 "replace" => Some(sfn::replace(a.remove(0), a.remove(0), a.remove(0))),
4841 "substring" if a.len() == 2 || a.len() == 3 => Some(cypher_substring(a)),
4842 "char_length" | "character_length" => Some(
4844 datafusion::functions::unicode::expr_fn::char_length(a.remove(0)),
4845 ),
4846 "tostring" => Some(CYPHER_TO_STRING.call(vec![a.remove(0)])),
4852 "tointeger" => Some(CYPHER_TO_INTEGER.call(vec![a.remove(0)])),
4853 "tofloat" => Some(CYPHER_TO_FLOAT.call(vec![a.remove(0)])),
4854 "toboolean" => Some(CYPHER_TO_BOOLEAN.call(vec![a.remove(0)])),
4855
4856 "date" if a.len() == 1 => {
4862 use datafusion::functions::datetime::expr_fn::to_date;
4863 Some(to_date(vec![a.remove(0)]))
4864 }
4865
4866 "abs" => Some(mfn::abs(a.remove(0))),
4868 "ceil" => Some(mfn::ceil(a.remove(0))),
4869 "floor" => Some(mfn::floor(a.remove(0))),
4870 "round" => Some(mfn::round(a)),
4871 "sqrt" => Some(mfn::sqrt(a.remove(0))),
4872 "log" => Some(mfn::log(a.remove(0), a.remove(0))),
4873 "exp" => Some(mfn::exp(a.remove(0))),
4874 "power" => Some(mfn::power(a.remove(0), a.remove(0))),
4875 "rand" if a.is_empty() => Some(mfn::random()),
4879 "sign" if a.len() == 1 => Some(cast(mfn::signum(a.remove(0)), DataType::Int64)),
4882 "coalesce" if !a.is_empty() => Some(datafusion::functions::core::expr_fn::coalesce(a)),
4884 "tail" if a.len() == 1 => Some(datafusion::functions_nested::expr_fn::array_pop_front(
4886 a.remove(0),
4887 )),
4888 "split" if a.len() == 2 => Some(datafusion::functions_nested::expr_fn::string_to_array(
4890 a.remove(0),
4891 a.remove(0),
4892 DfExpr::Literal(ScalarValue::Utf8(None), None),
4893 )),
4894
4895 "length" => Some(datafusion::functions_nested::expr_fn::array_length(
4899 a.remove(0),
4900 )),
4901
4902 "size" => Some(CYPHER_SIZE.call(vec![a.remove(0)])),
4907
4908 "_subscript" => {
4914 let list = a.remove(0);
4915 let idx = a.remove(0);
4916 Some(datafusion::functions_nested::expr_fn::array_element(
4917 list,
4918 one_based_index(idx),
4919 ))
4920 }
4921
4922 "head" => Some(datafusion::functions_nested::expr_fn::array_element(
4924 a.remove(0),
4925 lit(1_i64),
4926 )),
4927 "last" => Some(datafusion::functions_nested::expr_fn::array_element(
4928 a.remove(0),
4929 lit(-1_i64),
4930 )),
4931
4932 "_slice" => {
4943 let list = a.remove(0);
4944 let start = a.remove(0);
4945 let end = a.remove(0);
4946 Some(cypher_slice(list, Some(start), Some(end)))
4947 }
4948 "_slice_from_start" => {
4949 let list = a.remove(0);
4950 let end = a.remove(0);
4951 Some(cypher_slice(list, None, Some(end)))
4952 }
4953 "_slice_to_end" => {
4954 let list = a.remove(0);
4955 let start = a.remove(0);
4956 Some(cypher_slice(list, Some(start), None))
4957 }
4958
4959 "range" if a.len() == 2 || a.len() == 3 => {
4965 let from = a.remove(0);
4966 let end = a.remove(0);
4967 let by = if a.is_empty() {
4968 lit(1_i64)
4969 } else {
4970 a.remove(0)
4971 };
4972 Some(CYPHER_RANGE.call(vec![from, end, by]))
4973 }
4974
4975 "type" => Some(CYPHER_REL_TYPE.call(vec![a.remove(0)])),
4978
4979 other => resolve_path_builtin(other, a, path_hydration),
4981 }
4982}
4983
4984pub(crate) fn list_index_range(list: DfExpr) -> DfExpr {
4988 let len = cast(
4989 datafusion::functions_nested::expr_fn::array_length(list),
4990 DataType::Int64,
4991 );
4992 CYPHER_RANGE.call(vec![lit(0_i64), len - lit(1_i64), lit(1_i64)])
4993}
4994
4995fn cypher_substring(mut args: Vec<DfExpr>) -> DfExpr {
4996 let original = args.remove(0);
4997 let start = args.remove(0) + lit(1_i64);
4998 let substring = if args.is_empty() {
4999 datafusion::functions::unicode::expr_fn::substr(original, start)
5000 } else {
5001 datafusion::functions::unicode::expr_fn::substring(original, start, args.remove(0))
5002 };
5003 cast(substring, DataType::Utf8)
5004}
5005
5006fn is_path_builtin_name(name: &str) -> bool {
5007 matches!(name, "_path_nodes" | "_path_fixed_length" | "_path_struct")
5008}
5009
5010fn resolve_path_builtin(
5022 name: &str,
5023 args: Vec<DfExpr>,
5024 hydration: impl FnOnce() -> Option<PathNodeHydration>,
5025) -> Option<DfExpr> {
5026 use datafusion::functions::core::expr_fn::named_struct;
5027
5028 let mut a = args;
5029 match name {
5030 "_path_nodes" => {
5036 let seed = node_uuid_col(a.remove(0));
5037 let rels = a.remove(0);
5038 Some(match hydration() {
5039 Some(h) => ScalarUDF::new_from_impl(CypherPathNodes::with_hydration(h))
5040 .call(vec![seed, rels]),
5041 None => CYPHER_PATH_NODES.call(vec![seed, rels]),
5042 })
5043 }
5044
5045 "_path_fixed_length" => {
5049 let present = edge_present(&a.remove(0))?;
5050 Some(
5051 when(present, lit(1_u64))
5052 .otherwise(lit(ScalarValue::UInt64(None)))
5053 .expect("CASE build is infallible for a single WHEN + ELSE"),
5054 )
5055 }
5056
5057 "_path_struct" => {
5062 let nodes = a.remove(0);
5063 let rels = a.remove(0);
5064 let struct_expr = named_struct(vec![
5065 lit("nodes"),
5066 nodes.clone(),
5067 lit("relationships"),
5068 rels,
5069 ]);
5070 Some(null_unless(nodes.is_not_null(), struct_expr))
5071 }
5072
5073 _ => None,
5074 }
5075}
5076
5077fn edge_present(edge: &DfExpr) -> Option<DfExpr> {
5081 let DfExpr::Column(c) = edge else {
5082 return None;
5083 };
5084 Some(col(format!("{}.edge_uuid", c.name)).is_not_null())
5085}
5086
5087fn null_unless(present: DfExpr, value: DfExpr) -> DfExpr {
5091 when(present, value)
5092 .otherwise(lit(ScalarValue::Null))
5093 .expect("CASE build is infallible for a single WHEN + ELSE")
5094}
5095
5096fn node_uuid_col(base: DfExpr) -> DfExpr {
5103 match &base {
5104 DfExpr::Column(c) => col(format!("{}.node_uuid", c.name)),
5105 _ => datafusion::functions::core::expr_fn::get_field(base, "node_uuid"),
5106 }
5107}
5108
5109fn edge_present_qual(base: &str) -> DfExpr {
5110 col(format!("{base}.edge_uuid")).is_not_null()
5111}
5112
5113fn is_edge_value_topology_field(name: &str) -> bool {
5114 matches!(
5115 name,
5116 "edge_uuid"
5117 | "src_uuid"
5118 | "dst_uuid"
5119 | "edge_id"
5120 | "src_id"
5121 | "dst_id"
5122 | "created_at"
5123 | "rel_type_name"
5124 )
5125}
5126
5127fn empty_utf8_list() -> DfExpr {
5128 DfExpr::Literal(
5129 ScalarValue::List(ScalarValue::new_list(&[], &DataType::Utf8, true)),
5130 None,
5131 )
5132}
5133
5134fn empty_map_struct() -> DfExpr {
5135 DfExpr::Literal(
5136 ScalarValue::Struct(std::sync::Arc::new(
5137 datafusion::arrow::array::StructArray::new_empty_fields(1, None),
5138 )),
5139 None,
5140 )
5141}
5142
5143fn null_utf8_list() -> DfExpr {
5144 null_unless(lit(false), empty_utf8_list())
5145}
5146
5147fn node_labels_list(base: &str, label: Option<&str>, type_id_map: &HashMap<u32, String>) -> DfExpr {
5148 use datafusion::functions_nested::expr_fn::{array_concat, array_has, make_array};
5149
5150 if type_id_map.is_empty() {
5151 return label.map_or_else(empty_utf8_list, |name| make_array(vec![lit(name)]));
5152 }
5153
5154 let mut entries: Vec<(u32, &str)> = type_id_map
5155 .iter()
5156 .map(|(id, name)| (*id, name.as_str()))
5157 .collect();
5158 entries.sort_by_key(|(id, _)| *id);
5159 let labels = col(format!("{base}.type_ids"));
5160 let parts = entries
5161 .into_iter()
5162 .map(|(id, name)| {
5163 when(
5164 array_has(labels.clone(), lit(id)),
5165 make_array(vec![lit(name)]),
5166 )
5167 .otherwise(empty_utf8_list())
5168 .expect("CASE build is infallible for a single WHEN + ELSE")
5169 })
5170 .collect();
5171 array_concat(parts)
5172}
5173
5174fn node_value_struct(
5180 base: &str,
5181 label: Option<&str>,
5182 type_id_map: &HashMap<u32, String>,
5183 prop_names: &[String],
5184) -> DfExpr {
5185 let labels = node_labels_list(base, label, type_id_map);
5186 node_value_struct_with_labels(base, labels, prop_names)
5187}
5188
5189fn node_value_struct_with_labels(base: &str, labels: DfExpr, prop_names: &[String]) -> DfExpr {
5190 use datafusion::functions::core::expr_fn::named_struct;
5191
5192 let mut fields = vec![
5193 lit("node_uuid"),
5194 col(format!("{base}.node_uuid")),
5195 lit("labels"),
5196 labels,
5197 ];
5198 for name in prop_names {
5199 fields.push(lit(name.as_str()));
5200 fields.push(qualified_col(base, name));
5201 }
5202 let value = named_struct(fields);
5203 null_unless(qualified_col(base, "node_uuid").is_not_null(), value)
5204}
5205
5206fn relationship_value_struct(base: &str, rel_type: DfExpr, prop_names: &[String]) -> DfExpr {
5212 use datafusion::functions::core::expr_fn::named_struct;
5213
5214 let mut fields = vec![
5215 lit("edge_uuid"),
5216 col(format!("{base}.edge_uuid")),
5217 lit("src_uuid"),
5218 col(format!("{base}.src_uuid")),
5219 lit("dst_uuid"),
5220 col(format!("{base}.dst_uuid")),
5221 lit("rel_type"),
5222 cast(rel_type, DataType::Utf8),
5223 ];
5224 for name in prop_names {
5225 fields.push(lit(name.as_str()));
5226 fields.push(qualified_col(base, name));
5227 }
5228 named_struct(fields)
5229}
5230
5231fn one_based_index(idx: DfExpr) -> DfExpr {
5236 when(idx.clone().gt_eq(lit(0_i64)), idx.clone() + lit(1_i64))
5237 .otherwise(idx)
5238 .expect("CASE build is infallible for a single WHEN + ELSE")
5239}
5240
5241fn is_integer_data_type(dt: &DataType) -> bool {
5242 matches!(
5243 dt,
5244 DataType::Int8
5245 | DataType::Int16
5246 | DataType::Int32
5247 | DataType::Int64
5248 | DataType::UInt8
5249 | DataType::UInt16
5250 | DataType::UInt32
5251 | DataType::UInt64
5252 )
5253}
5254
5255fn as_null_literal(e: &DfExpr) -> bool {
5257 matches!(e, DfExpr::Literal(ScalarValue::Null, _))
5258}
5259
5260fn cypher_slice(list: DfExpr, start: Option<DfExpr>, end: Option<DfExpr>) -> DfExpr {
5261 let mut present = lit(true);
5262 let begin_expr = match start {
5263 Some(start) if as_null_literal(&start) => {
5264 present = lit(false);
5265 lit(1_i64)
5266 }
5267 Some(start) => {
5268 present = present.and(start.clone().is_not_null());
5269 when(start.clone().gt_eq(lit(0_i64)), start.clone() + lit(1_i64))
5270 .otherwise(start)
5271 .expect("CASE build")
5272 }
5273 None => lit(1_i64),
5274 };
5275 let end_expr = match end {
5276 Some(end) if as_null_literal(&end) => {
5277 present = lit(false);
5278 cast(
5279 datafusion::functions_nested::expr_fn::array_length(list.clone()),
5280 DataType::Int64,
5281 )
5282 }
5283 Some(end) => {
5284 present = present.and(end.clone().is_not_null());
5285 when(end.clone().gt_eq(lit(0_i64)), end.clone())
5286 .otherwise(end - lit(1_i64))
5287 .expect("CASE build")
5288 }
5289 None => cast(
5290 datafusion::functions_nested::expr_fn::array_length(list.clone()),
5291 DataType::Int64,
5292 ),
5293 };
5294 let slice =
5295 datafusion::functions_nested::expr_fn::array_slice(list, begin_expr, end_expr, None);
5296 null_unless(present, slice)
5297}
5298
5299static CYPHER_TO_INTEGER: LazyLock<ScalarUDF> = LazyLock::new(|| {
5304 ScalarUDF::new_from_impl(CypherConversion::new(CypherConversionKind::Integer))
5305});
5306static CYPHER_TO_FLOAT: LazyLock<ScalarUDF> =
5307 LazyLock::new(|| ScalarUDF::new_from_impl(CypherConversion::new(CypherConversionKind::Float)));
5308static CYPHER_TO_BOOLEAN: LazyLock<ScalarUDF> = LazyLock::new(|| {
5309 ScalarUDF::new_from_impl(CypherConversion::new(CypherConversionKind::Boolean))
5310});
5311
5312#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
5313enum CypherConversionKind {
5314 Integer,
5315 Float,
5316 Boolean,
5317}
5318
5319#[derive(Debug, PartialEq, Eq, Hash)]
5320struct CypherConversion {
5321 kind: CypherConversionKind,
5322 signature: Signature,
5323}
5324
5325impl CypherConversion {
5326 fn new(kind: CypherConversionKind) -> Self {
5327 Self {
5328 kind,
5329 signature: Signature::any(1, Volatility::Immutable),
5330 }
5331 }
5332}
5333
5334impl ScalarUDFImpl for CypherConversion {
5335 fn as_any(&self) -> &dyn Any {
5336 self
5337 }
5338
5339 fn name(&self) -> &'static str {
5340 match self.kind {
5341 CypherConversionKind::Integer => "cypher_to_integer",
5342 CypherConversionKind::Float => "cypher_to_float",
5343 CypherConversionKind::Boolean => "cypher_to_boolean",
5344 }
5345 }
5346
5347 fn signature(&self) -> &Signature {
5348 &self.signature
5349 }
5350
5351 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
5352 Ok(match self.kind {
5353 CypherConversionKind::Integer => DataType::Int64,
5354 CypherConversionKind::Float => DataType::Float64,
5355 CypherConversionKind::Boolean => DataType::Boolean,
5356 })
5357 }
5358
5359 fn invoke_with_args(
5360 &self,
5361 args: ScalarFunctionArgs,
5362 ) -> datafusion::error::Result<ColumnarValue> {
5363 use datafusion::arrow::array::{BooleanArray, Float64Array, Int64Array};
5364
5365 let array = args.args[0].to_array(args.number_rows)?;
5366 match self.kind {
5367 CypherConversionKind::Integer => {
5368 let out: datafusion::error::Result<Int64Array> = (0..array.len())
5369 .map(|i| {
5370 let value = decoded_scalar_at(&array, i)?;
5371 to_cypher_integer(&value)
5372 })
5373 .collect();
5374 Ok(ColumnarValue::Array(std::sync::Arc::new(out?)))
5375 }
5376 CypherConversionKind::Float => {
5377 let out: datafusion::error::Result<Float64Array> = (0..array.len())
5378 .map(|i| {
5379 let value = decoded_scalar_at(&array, i)?;
5380 to_cypher_float(&value)
5381 })
5382 .collect();
5383 Ok(ColumnarValue::Array(std::sync::Arc::new(out?)))
5384 }
5385 CypherConversionKind::Boolean => {
5386 let out: datafusion::error::Result<BooleanArray> = (0..array.len())
5387 .map(|i| {
5388 let value = decoded_scalar_at(&array, i)?;
5389 to_cypher_boolean(&value)
5390 })
5391 .collect();
5392 Ok(ColumnarValue::Array(std::sync::Arc::new(out?)))
5393 }
5394 }
5395 }
5396}
5397
5398fn decoded_scalar_at(
5399 array: &datafusion::arrow::array::ArrayRef,
5400 row: usize,
5401) -> datafusion::error::Result<ScalarValue> {
5402 let value = ScalarValue::try_from_array(array, row)?;
5403 Ok(unwrap_het(value))
5404}
5405
5406fn conversion_type_error(fn_name: &str, value: &ScalarValue) -> datafusion::error::DataFusionError {
5407 datafusion::error::DataFusionError::Execution(format!(
5408 "{fn_name}() cannot convert value of type {:?}",
5409 value.data_type()
5410 ))
5411}
5412
5413fn to_cypher_integer(value: &ScalarValue) -> datafusion::error::Result<Option<i64>> {
5414 if value.is_null() {
5415 return Ok(None);
5416 }
5417 if let Some(i) = scalar_as_i128(value) {
5418 return i64::try_from(i)
5419 .map(Some)
5420 .map_err(|_| conversion_type_error("toInteger", value));
5421 }
5422 match value {
5423 ScalarValue::Float32(Some(f)) => Ok(trunc_float_to_i64(f64::from(*f))),
5424 ScalarValue::Float64(Some(f)) => Ok(trunc_float_to_i64(*f)),
5425 ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) => Ok(s
5426 .parse::<f64>()
5427 .ok()
5428 .filter(|f| f.is_finite())
5429 .and_then(trunc_float_to_i64)),
5430 _ => Err(conversion_type_error("toInteger", value)),
5431 }
5432}
5433
5434#[allow(
5435 clippy::cast_possible_truncation,
5436 clippy::cast_precision_loss,
5437 reason = "openCypher toInteger truncates finite floating values toward zero"
5438)]
5439fn trunc_float_to_i64(f: f64) -> Option<i64> {
5440 if !f.is_finite() {
5441 return None;
5442 }
5443 let truncated = f.trunc();
5444 if truncated < i64::MIN as f64 || truncated > i64::MAX as f64 {
5445 return None;
5446 }
5447 Some(truncated as i64)
5448}
5449
5450fn to_cypher_float(value: &ScalarValue) -> datafusion::error::Result<Option<f64>> {
5451 if value.is_null() {
5452 return Ok(None);
5453 }
5454 if let Some(f) = scalar_as_f64(value) {
5455 return Ok(Some(f));
5456 }
5457 match value {
5458 ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) => {
5459 Ok(s.parse::<f64>().ok().filter(|f| f.is_finite()))
5460 }
5461 _ => Err(conversion_type_error("toFloat", value)),
5462 }
5463}
5464
5465fn to_cypher_boolean(value: &ScalarValue) -> datafusion::error::Result<Option<bool>> {
5466 if value.is_null() {
5467 return Ok(None);
5468 }
5469 match value {
5470 ScalarValue::Boolean(Some(b)) => Ok(Some(*b)),
5471 ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) => match s.as_str() {
5472 "true" => Ok(Some(true)),
5473 "false" => Ok(Some(false)),
5474 _ => Ok(None),
5475 },
5476 _ => Err(conversion_type_error("toBoolean", value)),
5477 }
5478}
5479
5480fn cypher_float_string(f: f64) -> String {
5481 if f == 0.0 {
5482 "0.0".to_owned()
5483 } else if f.is_nan() {
5484 "NaN".to_owned()
5485 } else if f.is_infinite() {
5486 if f.is_sign_positive() {
5487 "Infinity".to_owned()
5488 } else {
5489 "-Infinity".to_owned()
5490 }
5491 } else {
5492 let mut s = f.to_string();
5493 if !s.contains('.') && !s.contains('e') && !s.contains('E') {
5494 s.push_str(".0");
5495 }
5496 s
5497 }
5498}
5499
5500fn to_cypher_string(value: &ScalarValue) -> datafusion::error::Result<Option<String>> {
5501 if value.is_null() {
5502 return Ok(None);
5503 }
5504 match value {
5505 ScalarValue::Int8(Some(n)) => Ok(Some(n.to_string())),
5506 ScalarValue::Int16(Some(n)) => Ok(Some(n.to_string())),
5507 ScalarValue::Int32(Some(n)) => Ok(Some(n.to_string())),
5508 ScalarValue::Int64(Some(n)) => Ok(Some(n.to_string())),
5509 ScalarValue::UInt8(Some(n)) => Ok(Some(n.to_string())),
5510 ScalarValue::UInt16(Some(n)) => Ok(Some(n.to_string())),
5511 ScalarValue::UInt32(Some(n)) => Ok(Some(n.to_string())),
5512 ScalarValue::UInt64(Some(n)) => Ok(Some(n.to_string())),
5513 ScalarValue::Float32(Some(f)) => Ok(Some(cypher_float_string(f64::from(*f)))),
5514 ScalarValue::Float64(Some(f)) => Ok(Some(cypher_float_string(*f))),
5515 ScalarValue::Boolean(Some(b)) => Ok(Some(b.to_string())),
5516 ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) => Ok(Some(s.clone())),
5517 _ => Err(conversion_type_error("toString", value)),
5518 }
5519}
5520
5521static CYPHER_SIZE: LazyLock<ScalarUDF> =
5537 LazyLock::new(|| ScalarUDF::new_from_impl(CypherSize::new()));
5538
5539pub(crate) static CYPHER_ROW_MARKER: LazyLock<ScalarUDF> =
5543 LazyLock::new(|| ScalarUDF::new_from_impl(CypherRowMarker::new()));
5544
5545#[derive(Debug, PartialEq, Eq, Hash)]
5546struct CypherRowMarker {
5547 signature: Signature,
5548}
5549
5550impl CypherRowMarker {
5551 fn new() -> Self {
5552 Self {
5553 signature: Signature::any(1, Volatility::Immutable),
5554 }
5555 }
5556}
5557
5558impl ScalarUDFImpl for CypherRowMarker {
5559 fn as_any(&self) -> &dyn Any {
5560 self
5561 }
5562
5563 fn name(&self) -> &'static str {
5564 "cypher_row_marker"
5565 }
5566
5567 fn signature(&self) -> &Signature {
5568 &self.signature
5569 }
5570
5571 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
5572 Ok(DataType::Boolean)
5573 }
5574
5575 fn invoke_with_args(
5576 &self,
5577 args: ScalarFunctionArgs,
5578 ) -> datafusion::error::Result<ColumnarValue> {
5579 let values = datafusion::arrow::array::BooleanArray::from(vec![true; args.number_rows]);
5580 Ok(ColumnarValue::Array(Arc::new(values)))
5581 }
5582}
5583
5584#[derive(Debug, PartialEq, Eq, Hash)]
5585struct CypherSize {
5586 signature: Signature,
5587}
5588
5589impl CypherSize {
5590 fn new() -> Self {
5591 Self {
5593 signature: Signature::any(1, Volatility::Immutable),
5594 }
5595 }
5596}
5597
5598impl ScalarUDFImpl for CypherSize {
5599 fn as_any(&self) -> &dyn Any {
5600 self
5601 }
5602
5603 fn name(&self) -> &'static str {
5604 "cypher_size"
5605 }
5606
5607 fn signature(&self) -> &Signature {
5608 &self.signature
5609 }
5610
5611 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
5612 Ok(DataType::Int64)
5613 }
5614
5615 fn invoke_with_args(
5616 &self,
5617 args: ScalarFunctionArgs,
5618 ) -> datafusion::error::Result<ColumnarValue> {
5619 use datafusion::arrow::array::{Array, Int64Array};
5620 use datafusion::arrow::compute::kernels::length::length;
5621 use datafusion::common::cast::{as_large_list_array, as_list_array};
5622
5623 let array = args.args[0].to_array(args.number_rows)?;
5624 let out: Int64Array = match array.data_type() {
5625 DataType::List(_) => {
5627 let list = as_list_array(&array)?;
5628 (0..list.len())
5629 .map(|i| {
5630 (!list.is_null(i))
5631 .then(|| i64::try_from(list.value(i).len()).unwrap_or(i64::MAX))
5632 })
5633 .collect()
5634 }
5635 DataType::LargeList(_) => {
5636 let list = as_large_list_array(&array)?;
5637 (0..list.len())
5638 .map(|i| {
5639 (!list.is_null(i))
5640 .then(|| i64::try_from(list.value(i).len()).unwrap_or(i64::MAX))
5641 })
5642 .collect()
5643 }
5644 DataType::Utf8 | DataType::LargeUtf8 => {
5647 let lengths = length(&array)?;
5648 let casted = datafusion::arrow::compute::cast(&lengths, &DataType::Int64)?;
5649 casted
5650 .as_any()
5651 .downcast_ref::<Int64Array>()
5652 .expect("cast to Int64 yields Int64Array")
5653 .clone()
5654 }
5655 t if is_het_struct_type(Some(t)) => (0..array.len())
5661 .map(|i| {
5662 let sv = ScalarValue::try_from_array(&array, i).ok()?;
5663 match decode_het(&sv)? {
5664 ScalarValue::List(l) => (!l.is_null(0))
5665 .then(|| i64::try_from(l.value(0).len()).unwrap_or(i64::MAX)),
5666 ScalarValue::LargeList(l) => (!l.is_null(0))
5667 .then(|| i64::try_from(l.value(0).len()).unwrap_or(i64::MAX)),
5668 ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) => {
5669 i64::try_from(s.len()).ok()
5670 }
5671 _ => None,
5672 }
5673 })
5674 .collect(),
5675 _ => (0..array.len()).map(|_| None).collect(),
5678 };
5679 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
5680 }
5681}
5682
5683static CYPHER_LABELS: LazyLock<ScalarUDF> =
5688 LazyLock::new(|| ScalarUDF::new_from_impl(CypherGraphMetadata::new(GraphMetadataKind::Labels)));
5689static CYPHER_REL_TYPE: LazyLock<ScalarUDF> = LazyLock::new(|| {
5690 ScalarUDF::new_from_impl(CypherGraphMetadata::new(
5691 GraphMetadataKind::RelationshipType,
5692 ))
5693});
5694
5695#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
5696enum GraphMetadataKind {
5697 Labels,
5698 RelationshipType,
5699}
5700
5701#[derive(Debug, PartialEq, Eq, Hash)]
5702struct CypherGraphMetadata {
5703 kind: GraphMetadataKind,
5704 signature: Signature,
5705}
5706
5707impl CypherGraphMetadata {
5708 fn new(kind: GraphMetadataKind) -> Self {
5709 Self {
5710 kind,
5711 signature: Signature::any(1, Volatility::Immutable),
5712 }
5713 }
5714}
5715
5716impl ScalarUDFImpl for CypherGraphMetadata {
5717 fn as_any(&self) -> &dyn Any {
5718 self
5719 }
5720
5721 fn name(&self) -> &'static str {
5722 match self.kind {
5723 GraphMetadataKind::Labels => "cypher_labels",
5724 GraphMetadataKind::RelationshipType => "cypher_relationship_type",
5725 }
5726 }
5727
5728 fn signature(&self) -> &Signature {
5729 &self.signature
5730 }
5731
5732 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
5733 Ok(match self.kind {
5734 GraphMetadataKind::Labels => DataType::new_list(DataType::Utf8, true),
5735 GraphMetadataKind::RelationshipType => DataType::Utf8,
5736 })
5737 }
5738
5739 fn invoke_with_args(
5740 &self,
5741 args: ScalarFunctionArgs,
5742 ) -> datafusion::error::Result<ColumnarValue> {
5743 use datafusion::arrow::array::new_empty_array;
5744 use datafusion::error::DataFusionError;
5745
5746 let rows = args.number_rows;
5747 let values = args.args[0].to_array(rows)?;
5748 let field = match self.kind {
5749 GraphMetadataKind::Labels => "labels",
5750 GraphMetadataKind::RelationshipType => "rel_type",
5751 };
5752 let identity_field = match self.kind {
5753 GraphMetadataKind::Labels => "node_uuid",
5754 GraphMetadataKind::RelationshipType => "edge_uuid",
5755 };
5756 let mut output = Vec::with_capacity(rows);
5757 for row in 0..rows {
5758 let value = ScalarValue::try_from_array(&values, row)?;
5759 let value = decode_het(&value).unwrap_or(value);
5760 if value.is_null() {
5761 output.push(ScalarValue::try_new_null(&self.return_type(&[])?)?);
5762 continue;
5763 }
5764 let ScalarValue::Struct(entity) = value else {
5765 return Err(DataFusionError::Execution(format!(
5766 "InvalidArgumentValue: {}() requires a graph element",
5767 self.name()
5768 )));
5769 };
5770 if entity.column_by_name(identity_field).is_none() {
5771 return Err(DataFusionError::Execution(format!(
5772 "InvalidArgumentValue: {}() received the wrong graph element kind",
5773 self.name()
5774 )));
5775 }
5776 let column = entity.column_by_name(field).ok_or_else(|| {
5777 DataFusionError::Execution(format!(
5778 "InvalidArgumentValue: {}() received the wrong graph element kind",
5779 self.name()
5780 ))
5781 })?;
5782 output.push(ScalarValue::try_from_array(column, 0)?);
5783 }
5784 let data_type = self.return_type(&[])?;
5785 let array = if output.is_empty() {
5786 new_empty_array(&data_type)
5787 } else {
5788 ScalarValue::iter_to_array(output)?
5789 };
5790 Ok(ColumnarValue::Array(array))
5791 }
5792}
5793
5794const ENTITY_PROPERTY_MAP_HET_DEPTH: usize = 3;
5799
5800#[derive(Debug, PartialEq, Eq, Hash)]
5801struct CypherEntityProperties {
5802 signature: Signature,
5803}
5804
5805impl CypherEntityProperties {
5806 fn new(arity: usize) -> Self {
5807 Self {
5808 signature: Signature::any(arity, Volatility::Immutable),
5809 }
5810 }
5811}
5812
5813impl ScalarUDFImpl for CypherEntityProperties {
5814 fn as_any(&self) -> &dyn Any {
5815 self
5816 }
5817
5818 fn name(&self) -> &'static str {
5819 "cypher_entity_properties"
5820 }
5821
5822 fn signature(&self) -> &Signature {
5823 &self.signature
5824 }
5825
5826 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
5827 Ok(DataType::Struct(het_fields(ENTITY_PROPERTY_MAP_HET_DEPTH)))
5828 }
5829
5830 fn invoke_with_args(
5831 &self,
5832 args: ScalarFunctionArgs,
5833 ) -> datafusion::error::Result<ColumnarValue> {
5834 use datafusion::arrow::array::ArrayRef;
5835 use datafusion::error::DataFusionError;
5836 use std::sync::Arc;
5837
5838 let rows = args.number_rows;
5839 if args.args.is_empty() || !(args.args.len() - 1).is_multiple_of(2) {
5840 return Err(DataFusionError::Plan(
5841 "properties() entity map expects present plus key/value pairs".into(),
5842 ));
5843 }
5844 let cols: Vec<ArrayRef> = args
5845 .args
5846 .iter()
5847 .map(|arg| arg.to_array(rows))
5848 .collect::<datafusion::error::Result<_>>()?;
5849 let mut maps = Vec::with_capacity(rows);
5850 for row in 0..rows {
5851 let present = match ScalarValue::try_from_array(&cols[0], row)? {
5852 ScalarValue::Boolean(Some(true)) => true,
5853 ScalarValue::Boolean(Some(false) | None) | ScalarValue::Null => false,
5854 other => {
5855 return Err(DataFusionError::Execution(format!(
5856 "properties() entity presence must be boolean, got {other:?}"
5857 )));
5858 }
5859 };
5860 if !present {
5861 maps.push(ScalarValue::Null);
5862 continue;
5863 }
5864 let mut entries = Vec::with_capacity((cols.len() - 1) / 2);
5865 for pair in cols[1..].chunks_exact(2) {
5866 let key = ScalarValue::try_from_array(&pair[0], row)?;
5867 let Some(key) = scalar_access_key(&key)? else {
5868 continue;
5869 };
5870 let value = ScalarValue::try_from_array(&pair[1], row)?;
5871 if value.is_null() {
5872 continue;
5873 }
5874 entries.push((key, unwrap_het(value)));
5875 }
5876 let map = const_map_scalar(&entries).ok_or_else(|| {
5877 DataFusionError::Execution(
5878 "properties() could not encode entity property map".into(),
5879 )
5880 })?;
5881 maps.push(map);
5882 }
5883 let out = build_het_struct(&maps, ENTITY_PROPERTY_MAP_HET_DEPTH).ok_or_else(|| {
5884 DataFusionError::Execution("properties() could not encode entity property map".into())
5885 })?;
5886 Ok(ColumnarValue::Array(Arc::new(out)))
5887 }
5888}
5889
5890static CYPHER_MAP_KEYS: LazyLock<ScalarUDF> =
5891 LazyLock::new(|| ScalarUDF::new_from_impl(CypherMapKeys::new()));
5892
5893#[derive(Debug, PartialEq, Eq, Hash)]
5894struct CypherMapKeys {
5895 signature: Signature,
5896}
5897
5898impl CypherMapKeys {
5899 fn new() -> Self {
5900 Self {
5901 signature: Signature::any(1, Volatility::Immutable),
5902 }
5903 }
5904}
5905
5906impl ScalarUDFImpl for CypherMapKeys {
5907 fn as_any(&self) -> &dyn Any {
5908 self
5909 }
5910
5911 fn name(&self) -> &'static str {
5912 "cypher_map_keys"
5913 }
5914
5915 fn signature(&self) -> &Signature {
5916 &self.signature
5917 }
5918
5919 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
5920 Ok(DataType::new_list(DataType::Utf8, true))
5921 }
5922
5923 fn invoke_with_args(
5924 &self,
5925 args: ScalarFunctionArgs,
5926 ) -> datafusion::error::Result<ColumnarValue> {
5927 use datafusion::arrow::array::{Array, ArrayRef, ListArray, StringArray, StructArray};
5928 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
5929 use datafusion::error::DataFusionError;
5930 use std::sync::Arc;
5931
5932 let rows = args.number_rows;
5933 let values = args.args[0].to_array(rows)?;
5934 if matches!(values.data_type(), DataType::Null) {
5935 let nulls = NullBuffer::from(vec![false; rows]);
5936 let list = ListArray::new(
5937 Arc::new(Field::new("item", DataType::Utf8, true)),
5938 OffsetBuffer::new(ScalarBuffer::from(vec![0_i32; rows + 1])),
5939 Arc::new(StringArray::from(Vec::<Option<String>>::new())) as ArrayRef,
5940 Some(nulls),
5941 );
5942 return Ok(ColumnarValue::Array(Arc::new(list)));
5943 }
5944 let DataType::Struct(fields) = values.data_type() else {
5945 return Err(DataFusionError::Execution(format!(
5946 "keys() requires a map, node, relationship, or null, got {:?}",
5947 values.data_type()
5948 )));
5949 };
5950 let map = values
5951 .as_any()
5952 .downcast_ref::<StructArray>()
5953 .ok_or_else(|| DataFusionError::Execution("keys() expected a struct map".into()))?;
5954 if is_het_struct_type(Some(values.data_type())) {
5955 return tagged_map_keys(map, rows);
5956 }
5957 if !is_plain_map_struct_type(values.data_type()) {
5958 return Err(DataFusionError::Execution(format!(
5959 "keys() requires a map, node, relationship, or null, got {:?}",
5960 values.data_type()
5961 )));
5962 }
5963 let names: Vec<String> = fields.iter().map(|f| f.name().clone()).collect();
5964 let mut offsets = Vec::with_capacity(rows + 1);
5965 let mut values = Vec::new();
5966 let mut valid = Vec::with_capacity(rows);
5967 offsets.push(0_i32);
5968 for row in 0..rows {
5969 if map.is_null(row) {
5970 valid.push(false);
5971 } else {
5972 valid.push(true);
5973 values.extend(names.iter().cloned().map(Some));
5974 }
5975 offsets.push(i32::try_from(values.len()).map_err(|_| {
5976 DataFusionError::Execution("keys() result exceeded i32 list offsets".into())
5977 })?);
5978 }
5979 let list = ListArray::new(
5980 Arc::new(Field::new("item", DataType::Utf8, true)),
5981 OffsetBuffer::new(ScalarBuffer::from(offsets)),
5982 Arc::new(StringArray::from(values)) as ArrayRef,
5983 Some(NullBuffer::from(valid)),
5984 );
5985 Ok(ColumnarValue::Array(Arc::new(list)))
5986 }
5987}
5988
5989fn tagged_map_keys(
5990 map: &datafusion::arrow::array::StructArray,
5991 rows: usize,
5992) -> datafusion::error::Result<ColumnarValue> {
5993 use datafusion::arrow::array::{
5994 Array, ArrayRef, Int8Array, ListArray, StringArray, StructArray,
5995 };
5996 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
5997 use datafusion::error::DataFusionError;
5998 use std::sync::Arc;
5999
6000 let tags = map
6001 .column_by_name("__het_tag")
6002 .and_then(|c| c.as_any().downcast_ref::<Int8Array>())
6003 .ok_or_else(|| DataFusionError::Plan("tagged map is missing __het_tag".into()))?;
6004 let entries = map
6005 .column_by_name("__het_map")
6006 .and_then(|c| c.as_any().downcast_ref::<ListArray>())
6007 .ok_or_else(|| DataFusionError::Plan("tagged map is missing __het_map".into()))?;
6008 let mut offsets = Vec::with_capacity(rows + 1);
6009 let mut values = Vec::new();
6010 let mut valid = Vec::with_capacity(rows);
6011 offsets.push(0_i32);
6012 for row in 0..rows {
6013 if map.is_null(row) {
6014 valid.push(false);
6015 } else {
6016 if tags.value(row) != 5 {
6017 return Err(DataFusionError::Execution(
6018 "keys() requires a map, node, relationship, or null".into(),
6019 ));
6020 }
6021 valid.push(true);
6022 if !entries.is_null(row) {
6023 let entry_values = entries.value(row);
6024 let entry_struct = entry_values
6025 .as_any()
6026 .downcast_ref::<StructArray>()
6027 .ok_or_else(|| {
6028 DataFusionError::Plan("tagged map entries must be structs".into())
6029 })?;
6030 let map_keys = entry_struct
6031 .column_by_name("__het_mkey")
6032 .and_then(|c| c.as_any().downcast_ref::<StringArray>())
6033 .ok_or_else(|| {
6034 DataFusionError::Plan("tagged map entries must carry __het_mkey".into())
6035 })?;
6036 for idx in 0..entry_struct.len() {
6037 if !map_keys.is_null(idx) {
6038 values.push(Some(map_keys.value(idx).to_owned()));
6039 }
6040 }
6041 }
6042 }
6043 offsets.push(i32::try_from(values.len()).map_err(|_| {
6044 DataFusionError::Execution("keys() result exceeded i32 list offsets".into())
6045 })?);
6046 }
6047 let list = ListArray::new(
6048 Arc::new(Field::new("item", DataType::Utf8, true)),
6049 OffsetBuffer::new(ScalarBuffer::from(offsets)),
6050 Arc::new(StringArray::from(values)) as ArrayRef,
6051 Some(NullBuffer::from(valid)),
6052 );
6053 Ok(ColumnarValue::Array(Arc::new(list)))
6054}
6055
6056static CYPHER_VALUE_ACCESS: LazyLock<ScalarUDF> =
6061 LazyLock::new(|| ScalarUDF::new_from_impl(CypherValueAccess::new()));
6062
6063#[derive(Debug, PartialEq, Eq, Hash)]
6064struct CypherStaticValueAccess {
6065 key: String,
6066 signature: Signature,
6067}
6068
6069impl CypherStaticValueAccess {
6070 fn new(key: String) -> Self {
6071 Self {
6072 key,
6073 signature: Signature::any(1, Volatility::Immutable),
6074 }
6075 }
6076}
6077
6078impl ScalarUDFImpl for CypherStaticValueAccess {
6079 fn as_any(&self) -> &dyn Any {
6080 self
6081 }
6082
6083 fn name(&self) -> &'static str {
6084 "cypher_static_value_access"
6085 }
6086
6087 fn signature(&self) -> &Signature {
6088 &self.signature
6089 }
6090
6091 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
6092 static_value_access_return_type(arg_types.first(), &self.key)
6093 }
6094
6095 fn return_field_from_args(&self, args: ReturnFieldArgs) -> datafusion::error::Result<FieldRef> {
6096 Ok(Arc::new(Field::new(
6097 self.name(),
6098 static_value_access_return_type(
6099 args.arg_fields.first().map(|field| field.data_type()),
6100 &self.key,
6101 )?,
6102 true,
6103 )))
6104 }
6105
6106 fn invoke_with_args(
6107 &self,
6108 args: ScalarFunctionArgs,
6109 ) -> datafusion::error::Result<ColumnarValue> {
6110 use datafusion::error::DataFusionError;
6111
6112 let rows = args.number_rows;
6113 let values = args.args[0].to_array(rows)?;
6114 let return_type = args.return_field.data_type();
6115 let null_value = ScalarValue::try_new_null(return_type)?;
6116 let output = (0..rows)
6117 .map(|row| {
6118 let value = ScalarValue::try_from_array(&values, row)?;
6119 let value = decode_het(&value).unwrap_or(value);
6120 if value.is_null() {
6121 return Ok(null_value.clone());
6122 }
6123 let ScalarValue::Struct(value) = value else {
6124 return Err(DataFusionError::Execution(
6125 "InvalidArgumentValue: property access requires a map or graph element"
6126 .into(),
6127 ));
6128 };
6129 let Some(column) = value.column_by_name(&self.key) else {
6130 return Ok(null_value.clone());
6131 };
6132 let result = ScalarValue::try_from_array(column, 0)?;
6133 if result.data_type() == *return_type {
6134 Ok(result)
6135 } else if result.is_null() {
6136 Ok(null_value.clone())
6137 } else {
6138 Err(DataFusionError::Execution(format!(
6139 "property `{}` has incompatible runtime type {:?}; expected {:?}",
6140 self.key,
6141 result.data_type(),
6142 return_type
6143 )))
6144 }
6145 })
6146 .collect::<datafusion::error::Result<Vec<_>>>()?;
6147 Ok(ColumnarValue::Array(ScalarValue::iter_to_array(output)?))
6148 }
6149}
6150
6151#[derive(Debug, PartialEq, Eq, Hash)]
6152struct CypherValueAccess {
6153 signature: Signature,
6154}
6155
6156impl CypherValueAccess {
6157 fn new() -> Self {
6158 Self {
6159 signature: Signature::any(2, Volatility::Immutable),
6160 }
6161 }
6162}
6163
6164impl ScalarUDFImpl for CypherValueAccess {
6165 fn as_any(&self) -> &dyn Any {
6166 self
6167 }
6168
6169 fn name(&self) -> &'static str {
6170 "cypher_value_access"
6171 }
6172
6173 fn signature(&self) -> &Signature {
6174 &self.signature
6175 }
6176
6177 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
6178 value_access_return_type(arg_types.first())
6179 }
6180
6181 fn return_field_from_args(&self, args: ReturnFieldArgs) -> datafusion::error::Result<FieldRef> {
6182 Ok(std::sync::Arc::new(Field::new(
6183 self.name(),
6184 value_access_return_type(args.arg_fields.first().map(|f| f.data_type()))?,
6185 true,
6186 )))
6187 }
6188
6189 fn invoke_with_args(
6190 &self,
6191 args: ScalarFunctionArgs,
6192 ) -> datafusion::error::Result<ColumnarValue> {
6193 use datafusion::arrow::array::StructArray;
6194 use datafusion::error::DataFusionError;
6195
6196 let rows = args.number_rows;
6197 let values = args.args[0].to_array(rows)?;
6198 let keys = args.args[1].to_array(rows)?;
6199 let return_type = args.return_field.data_type().clone();
6200 let null_value = ScalarValue::try_from(&return_type).unwrap_or(ScalarValue::Null);
6201 if matches!(values.data_type(), DataType::Null) {
6202 return Ok(ColumnarValue::Array(ScalarValue::iter_to_array(
6203 (0..rows).map(|_| null_value.clone()),
6204 )?));
6205 }
6206 if let Some(list) = ListView::from_array(&values) {
6207 let out = (0..rows)
6208 .map(|i| list_access_value(&list, &keys, i, &null_value))
6209 .collect::<datafusion::error::Result<Vec<_>>>()?;
6210 return Ok(ColumnarValue::Array(ScalarValue::iter_to_array(out)?));
6211 }
6212 let map = values
6213 .as_any()
6214 .downcast_ref::<StructArray>()
6215 .ok_or_else(|| {
6216 DataFusionError::Execution(format!(
6217 "dynamic subscript requires a list or map/entity struct, got {:?}",
6218 values.data_type()
6219 ))
6220 })?;
6221 if is_het_struct_type(Some(values.data_type())) {
6222 let out = (0..rows)
6223 .map(|i| het_map_access_value(map, &keys, i, &null_value))
6224 .collect::<datafusion::error::Result<Vec<_>>>()?;
6225 return Ok(ColumnarValue::Array(ScalarValue::iter_to_array(out)?));
6226 }
6227 let out = (0..rows)
6228 .map(|i| {
6229 if map.is_null(i) {
6230 return Ok(null_value.clone());
6231 }
6232 let key = ScalarValue::try_from_array(&keys, i)?;
6233 let Some(key) = scalar_access_key(&key)? else {
6234 return Ok(null_value.clone());
6235 };
6236 let Some(col) = map.column_by_name(&key) else {
6237 return Ok(null_value.clone());
6238 };
6239 ScalarValue::try_from_array(col, i)
6240 })
6241 .collect::<datafusion::error::Result<Vec<_>>>()?;
6242 Ok(ColumnarValue::Array(ScalarValue::iter_to_array(out)?))
6243 }
6244}
6245
6246fn value_access_return_type(dt: Option<&DataType>) -> datafusion::error::Result<DataType> {
6247 let Some(dt) = dt else {
6248 return Ok(DataType::Null);
6249 };
6250 match dt {
6251 DataType::Null => Ok(DataType::Null),
6252 DataType::List(field) | DataType::LargeList(field) | DataType::FixedSizeList(field, _) => {
6253 Ok(field.data_type().clone())
6254 }
6255 dt if is_het_struct_type(Some(dt)) => het_value_access_return_type(dt),
6256 DataType::Struct(fields) => common_struct_field_type(fields),
6257 _ => Ok(DataType::Null),
6261 }
6262}
6263
6264fn static_value_access_return_type(
6265 dt: Option<&DataType>,
6266 key: &str,
6267) -> datafusion::error::Result<DataType> {
6268 let Some(DataType::Struct(fields)) = dt else {
6269 return Ok(DataType::Null);
6270 };
6271 if !is_het_struct_type(dt) {
6272 return fields
6273 .iter()
6274 .find(|field| field.name() == key)
6275 .map_or(Ok(DataType::Null), |field| Ok(field.data_type().clone()));
6276 }
6277 if fields.iter().any(|field| field.name() == "__het_map") {
6278 return het_value_access_return_type(&DataType::Struct(fields.clone()));
6279 }
6280 let mut data_type = None;
6281 for variant in fields
6282 .iter()
6283 .filter(|field| field.name().starts_with("__het_value_"))
6284 {
6285 let DataType::Struct(value_fields) = variant.data_type() else {
6286 continue;
6287 };
6288 let Some(property) = value_fields.iter().find(|field| field.name() == key) else {
6289 continue;
6290 };
6291 if matches!(property.data_type(), DataType::Null) {
6292 continue;
6293 }
6294 match &data_type {
6295 None => data_type = Some(property.data_type().clone()),
6296 Some(existing) if existing == property.data_type() => {}
6297 Some(existing) => {
6298 return Err(datafusion::error::DataFusionError::Plan(format!(
6299 "property `{key}` has incompatible graph-value types {existing:?} and {:?}",
6300 property.data_type()
6301 )));
6302 }
6303 }
6304 }
6305 Ok(data_type.unwrap_or(DataType::Null))
6306}
6307
6308fn list_access_value(
6309 list: &ListView<'_>,
6310 keys: &datafusion::arrow::array::ArrayRef,
6311 row: usize,
6312 null_value: &ScalarValue,
6313) -> datafusion::error::Result<ScalarValue> {
6314 if list.is_null(row) {
6315 return Ok(null_value.clone());
6316 }
6317 let key = ScalarValue::try_from_array(keys, row)?;
6318 let Some(idx) = scalar_list_index(&key)? else {
6319 return Ok(null_value.clone());
6320 };
6321 let elems = list.value(row);
6322 let len = i64::try_from(elems.len()).map_err(|_| {
6323 datafusion::error::DataFusionError::Execution(
6324 "dynamic list access length exceeds i64 range".into(),
6325 )
6326 })?;
6327 let pos = if idx < 0 { len + idx } else { idx };
6328 if pos < 0 || pos >= len {
6329 return Ok(null_value.clone());
6330 }
6331 let pos = usize::try_from(pos).map_err(|_| {
6332 datafusion::error::DataFusionError::Execution(
6333 "dynamic list access index exceeds usize range".into(),
6334 )
6335 })?;
6336 ScalarValue::try_from_array(&elems, pos)
6337}
6338
6339fn scalar_list_index(s: &ScalarValue) -> datafusion::error::Result<Option<i64>> {
6340 if s.is_null() {
6341 return Ok(None);
6342 }
6343 macro_rules! signed_index {
6344 ($value:expr) => {
6345 $value.map(i64::from)
6346 };
6347 }
6348 let idx = match s {
6349 ScalarValue::Int8(v) => signed_index!(*v),
6350 ScalarValue::Int16(v) => signed_index!(*v),
6351 ScalarValue::Int32(v) => signed_index!(*v),
6352 ScalarValue::Int64(v) => *v,
6353 ScalarValue::UInt8(v) => v.map(i64::from),
6354 ScalarValue::UInt16(v) => v.map(i64::from),
6355 ScalarValue::UInt32(v) => v.map(i64::from),
6356 ScalarValue::UInt64(v) => v.map(i64::try_from).transpose().map_err(|_| {
6357 datafusion::error::DataFusionError::Execution(
6358 "dynamic list access index exceeds i64 range".into(),
6359 )
6360 })?,
6361 other => {
6362 return Err(datafusion::error::DataFusionError::Execution(format!(
6363 "dynamic list access index must be an integer, got {other:?}"
6364 )));
6365 }
6366 };
6367 Ok(idx)
6368}
6369
6370fn het_value_access_return_type(dt: &DataType) -> datafusion::error::Result<DataType> {
6371 use datafusion::error::DataFusionError;
6372
6373 let DataType::Struct(fields) = dt else {
6374 unreachable!("caller checked het struct type")
6375 };
6376 let Some(map_field) = fields.iter().find(|f| f.name() == "__het_map") else {
6377 return Err(DataFusionError::Plan(
6378 "dynamic value access requires a tagged map element".into(),
6379 ));
6380 };
6381 let DataType::List(entry_field) = map_field.data_type() else {
6382 return Err(DataFusionError::Plan(
6383 "tagged map field must be a list".into(),
6384 ));
6385 };
6386 let DataType::Struct(entry_fields) = entry_field.data_type() else {
6387 return Err(DataFusionError::Plan(
6388 "tagged map entries must be structs".into(),
6389 ));
6390 };
6391 entry_fields
6392 .iter()
6393 .find(|f| f.name() == "__het_mval")
6394 .map(|f| f.data_type().clone())
6395 .ok_or_else(|| DataFusionError::Plan("tagged map entries must carry __het_mval".into()))
6396}
6397
6398fn het_map_access_value(
6399 map: &datafusion::arrow::array::StructArray,
6400 keys: &datafusion::arrow::array::ArrayRef,
6401 row: usize,
6402 null_value: &ScalarValue,
6403) -> datafusion::error::Result<ScalarValue> {
6404 use datafusion::arrow::array::{Array, Int8Array, ListArray, StringArray, StructArray};
6405 use datafusion::error::DataFusionError;
6406
6407 if map.is_null(row) {
6408 return Ok(null_value.clone());
6409 }
6410 let key = ScalarValue::try_from_array(keys, row)?;
6411 let Some(key) = scalar_access_key(&key)? else {
6412 return Ok(null_value.clone());
6413 };
6414 let tag = map
6415 .column_by_name("__het_tag")
6416 .and_then(|c| c.as_any().downcast_ref::<Int8Array>())
6417 .ok_or_else(|| DataFusionError::Plan("tagged value is missing __het_tag".into()))?;
6418 if tag.value(row) != 5 {
6419 return Err(DataFusionError::Execution(
6420 "invalid argument type: dynamic value access requires a map".into(),
6421 ));
6422 }
6423 let entries = map
6424 .column_by_name("__het_map")
6425 .and_then(|c| c.as_any().downcast_ref::<ListArray>())
6426 .ok_or_else(|| DataFusionError::Plan("tagged map value is missing __het_map".into()))?;
6427 if entries.is_null(row) {
6428 return Ok(null_value.clone());
6429 }
6430 let entry_values = entries.value(row);
6431 let entry_struct = entry_values
6432 .as_any()
6433 .downcast_ref::<StructArray>()
6434 .ok_or_else(|| DataFusionError::Plan("tagged map entries must be structs".into()))?;
6435 let map_keys = entry_struct
6436 .column_by_name("__het_mkey")
6437 .and_then(|c| c.as_any().downcast_ref::<StringArray>())
6438 .ok_or_else(|| DataFusionError::Plan("tagged map entries must carry __het_mkey".into()))?;
6439 let map_values = entry_struct
6440 .column_by_name("__het_mval")
6441 .ok_or_else(|| DataFusionError::Plan("tagged map entries must carry __het_mval".into()))?;
6442 for idx in 0..entry_struct.len() {
6443 if !map_keys.is_null(idx) && map_keys.value(idx) == key {
6444 return ScalarValue::try_from_array(map_values, idx);
6445 }
6446 }
6447 Ok(null_value.clone())
6448}
6449
6450fn common_struct_field_type(
6451 fields: &datafusion::arrow::datatypes::Fields,
6452) -> datafusion::error::Result<DataType> {
6453 let mut dtype: Option<DataType> = None;
6454 for field in fields {
6455 let field_type = field.data_type();
6456 if matches!(field_type, DataType::Null) {
6457 continue;
6458 }
6459 match &dtype {
6460 None => dtype = Some(field_type.clone()),
6461 Some(prev) if prev == field_type => {}
6462 Some(prev) => {
6463 return Err(datafusion::error::DataFusionError::Plan(format!(
6464 "dynamic value access over mixed field types is not supported: {prev:?} and {field_type:?}"
6465 )));
6466 }
6467 }
6468 }
6469 Ok(dtype.unwrap_or(DataType::Null))
6470}
6471
6472fn scalar_access_key(s: &ScalarValue) -> datafusion::error::Result<Option<String>> {
6473 if s.is_null() {
6474 return Ok(None);
6475 }
6476 match s {
6477 ScalarValue::Utf8(v) | ScalarValue::LargeUtf8(v) | ScalarValue::Utf8View(v) => {
6478 Ok(v.clone())
6479 }
6480 other => Err(datafusion::error::DataFusionError::Execution(format!(
6481 "dynamic map/property access key must be a string, got {other:?}"
6482 ))),
6483 }
6484}
6485
6486static CYPHER_REVERSE: LazyLock<ScalarUDF> =
6496 LazyLock::new(|| ScalarUDF::new_from_impl(CypherReverse::new()));
6497
6498#[derive(Debug, PartialEq, Eq, Hash)]
6499struct CypherReverse {
6500 signature: Signature,
6501}
6502
6503impl CypherReverse {
6504 fn new() -> Self {
6505 Self {
6506 signature: Signature::any(1, Volatility::Immutable),
6507 }
6508 }
6509}
6510
6511impl ScalarUDFImpl for CypherReverse {
6512 fn as_any(&self) -> &dyn Any {
6513 self
6514 }
6515 fn name(&self) -> &'static str {
6516 "cypher_reverse"
6517 }
6518 fn signature(&self) -> &Signature {
6519 &self.signature
6520 }
6521 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
6522 Ok(arg_types.first().cloned().unwrap_or(DataType::Null))
6524 }
6525
6526 fn invoke_with_args(
6527 &self,
6528 args: ScalarFunctionArgs,
6529 ) -> datafusion::error::Result<ColumnarValue> {
6530 use datafusion::arrow::array::{
6531 Array, LargeStringArray, ListArray, StringArray, UInt32Array,
6532 };
6533 use datafusion::common::cast::as_list_array;
6534 use datafusion::error::DataFusionError;
6535 use std::sync::Arc;
6536
6537 let array = args.args[0].to_array(args.number_rows)?;
6538 match array.data_type() {
6539 DataType::Utf8 => {
6540 let s = array
6541 .as_any()
6542 .downcast_ref::<StringArray>()
6543 .ok_or_else(|| {
6544 DataFusionError::Internal("cypher_reverse: not a string array".into())
6545 })?;
6546 let out: StringArray = (0..s.len())
6547 .map(|i| (!s.is_null(i)).then(|| s.value(i).chars().rev().collect::<String>()))
6548 .collect();
6549 Ok(ColumnarValue::Array(Arc::new(out)))
6550 }
6551 DataType::LargeUtf8 => {
6552 let s = array
6553 .as_any()
6554 .downcast_ref::<LargeStringArray>()
6555 .ok_or_else(|| {
6556 DataFusionError::Internal("cypher_reverse: not a large string array".into())
6557 })?;
6558 let out: LargeStringArray = (0..s.len())
6559 .map(|i| (!s.is_null(i)).then(|| s.value(i).chars().rev().collect::<String>()))
6560 .collect();
6561 Ok(ColumnarValue::Array(Arc::new(out)))
6562 }
6563 DataType::List(_) => {
6564 let list = as_list_array(&array)?;
6565 let values = list.values();
6566 let offsets = list.offsets();
6567 let mut idx: Vec<u32> = Vec::with_capacity(values.len());
6570 for w in offsets.windows(2) {
6571 let (start, end) = (w[0], w[1]);
6572 for j in (start..end).rev() {
6573 idx.push(u32::try_from(j).unwrap_or(0));
6574 }
6575 }
6576 let taken =
6577 datafusion::arrow::compute::take(values, &UInt32Array::from(idx), None)?;
6578 let field = match list.data_type() {
6579 DataType::List(f) => Arc::clone(f),
6580 _ => unreachable!("matched List above"),
6581 };
6582 let reversed = ListArray::new(field, offsets.clone(), taken, list.nulls().cloned());
6583 Ok(ColumnarValue::Array(Arc::new(reversed)))
6584 }
6585 other => Err(DataFusionError::Plan(format!(
6586 "reverse() expects a string or list, got {other:?}"
6587 ))),
6588 }
6589 }
6590}
6591
6592static CYPHER_AND: LazyLock<ScalarUDF> =
6597 LazyLock::new(|| ScalarUDF::new_from_impl(CypherBoolOp::new(CypherBoolOpKind::And)));
6598static CYPHER_OR: LazyLock<ScalarUDF> =
6599 LazyLock::new(|| ScalarUDF::new_from_impl(CypherBoolOp::new(CypherBoolOpKind::Or)));
6600static CYPHER_XOR: LazyLock<ScalarUDF> =
6601 LazyLock::new(|| ScalarUDF::new_from_impl(CypherBoolOp::new(CypherBoolOpKind::Xor)));
6602
6603#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
6604enum CypherBoolOpKind {
6605 And,
6606 Or,
6607 Xor,
6608}
6609
6610#[derive(Debug, PartialEq, Eq, Hash)]
6611struct CypherBoolOp {
6612 signature: Signature,
6613 kind: CypherBoolOpKind,
6614}
6615
6616impl CypherBoolOp {
6617 fn new(kind: CypherBoolOpKind) -> Self {
6618 Self {
6619 signature: Signature::any(2, Volatility::Immutable),
6620 kind,
6621 }
6622 }
6623}
6624
6625impl ScalarUDFImpl for CypherBoolOp {
6626 fn as_any(&self) -> &dyn Any {
6627 self
6628 }
6629 fn name(&self) -> &'static str {
6630 match self.kind {
6631 CypherBoolOpKind::And => "cypher_and",
6632 CypherBoolOpKind::Or => "cypher_or",
6633 CypherBoolOpKind::Xor => "cypher_xor",
6634 }
6635 }
6636 fn signature(&self) -> &Signature {
6637 &self.signature
6638 }
6639 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
6640 Ok(DataType::Boolean)
6641 }
6642 fn invoke_with_args(
6643 &self,
6644 args: ScalarFunctionArgs,
6645 ) -> datafusion::error::Result<ColumnarValue> {
6646 use datafusion::arrow::array::BooleanArray;
6647 let rows = args.number_rows;
6648 let lhs = args.args[0].to_array(rows)?;
6649 let rhs = args.args[1].to_array(rows)?;
6650 let out: BooleanArray = (0..rows)
6651 .map(|i| {
6652 let l = ScalarValue::try_from_array(&lhs, i)?;
6653 let r = ScalarValue::try_from_array(&rhs, i)?;
6654 let l = scalar_as_bool(&l)?;
6655 let r = scalar_as_bool(&r)?;
6656 Ok::<Option<bool>, datafusion::error::DataFusionError>(match self.kind {
6657 CypherBoolOpKind::And => match (l, r) {
6658 (Some(false), _) | (_, Some(false)) => Some(false),
6659 (Some(true), Some(true)) => Some(true),
6660 _ => None,
6661 },
6662 CypherBoolOpKind::Or => match (l, r) {
6663 (Some(true), _) | (_, Some(true)) => Some(true),
6664 (Some(false), Some(false)) => Some(false),
6665 _ => None,
6666 },
6667 CypherBoolOpKind::Xor => match (l, r) {
6668 (Some(left), Some(right)) => Some(left ^ right),
6669 _ => None,
6670 },
6671 })
6672 })
6673 .collect::<datafusion::error::Result<Vec<_>>>()?
6674 .into_iter()
6675 .collect();
6676 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
6677 }
6678}
6679
6680fn scalar_as_bool(s: &ScalarValue) -> datafusion::error::Result<Option<bool>> {
6681 let s = unwrap_het(s.clone());
6682 if s.is_null() {
6683 return Ok(None);
6684 }
6685 match s {
6686 ScalarValue::Boolean(v) => Ok(v),
6687 other => Err(datafusion::error::DataFusionError::Plan(format!(
6688 "expected boolean operand, got {other:?}"
6689 ))),
6690 }
6691}
6692
6693static CYPHER_CMP_PRED: LazyLock<ScalarUDF> =
6694 LazyLock::new(|| ScalarUDF::new_from_impl(CypherCmpPred::new()));
6695
6696#[derive(Debug, PartialEq, Eq, Hash)]
6697struct CypherCmpPred {
6698 signature: Signature,
6699}
6700
6701impl CypherCmpPred {
6702 fn new() -> Self {
6703 Self {
6704 signature: Signature::any(3, Volatility::Immutable),
6705 }
6706 }
6707}
6708
6709impl ScalarUDFImpl for CypherCmpPred {
6710 fn as_any(&self) -> &dyn Any {
6711 self
6712 }
6713 fn name(&self) -> &'static str {
6714 "cypher_cmp_pred"
6715 }
6716 fn signature(&self) -> &Signature {
6717 &self.signature
6718 }
6719 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
6720 Ok(DataType::Boolean)
6721 }
6722 fn invoke_with_args(
6723 &self,
6724 args: ScalarFunctionArgs,
6725 ) -> datafusion::error::Result<ColumnarValue> {
6726 use datafusion::arrow::array::BooleanArray;
6727 let rows = args.number_rows;
6728 let lhs = args.args[0].to_array(rows)?;
6729 let rhs = args.args[1].to_array(rows)?;
6730 let op = args.args[2].to_array(rows)?;
6731 let out: BooleanArray = (0..rows)
6732 .map(|i| {
6733 let l = ScalarValue::try_from_array(&lhs, i)?;
6734 let r = ScalarValue::try_from_array(&rhs, i)?;
6735 let op = ScalarValue::try_from_array(&op, i)?;
6736 let op = scalar_as_i8(&op)?;
6737 Ok::<Option<bool>, datafusion::error::DataFusionError>(cypher_compare_pred(
6738 &l, &r, op,
6739 ))
6740 })
6741 .collect::<datafusion::error::Result<Vec<_>>>()?
6742 .into_iter()
6743 .collect();
6744 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
6745 }
6746}
6747
6748fn scalar_as_i8(s: &ScalarValue) -> datafusion::error::Result<i8> {
6749 match s {
6750 ScalarValue::Int8(Some(v)) => Ok(*v),
6751 ScalarValue::Int64(Some(v)) => i8::try_from(*v).map_err(|_| {
6752 datafusion::error::DataFusionError::Plan(format!(
6753 "comparison opcode {v} is outside i8 range"
6754 ))
6755 }),
6756 other => Err(datafusion::error::DataFusionError::Plan(format!(
6757 "comparison opcode must be an integer, got {other:?}"
6758 ))),
6759 }
6760}
6761
6762fn cypher_compare_pred(l: &ScalarValue, r: &ScalarValue, op: i8) -> Option<bool> {
6763 let l = unwrap_het(l.clone());
6764 let r = unwrap_het(r.clone());
6765 if l.is_null() || r.is_null() {
6766 return None;
6767 }
6768 if is_numeric_scalar(&l) && is_numeric_scalar(&r) {
6769 let lf = scalar_as_f64(&l)?;
6770 let rf = scalar_as_f64(&r)?;
6771 if lf.is_nan() || rf.is_nan() {
6772 return Some(false);
6773 }
6774 }
6775 let cmp = cypher_compare(&l, &r)?;
6776 Some(match op {
6777 0 => cmp < 0,
6778 1 => cmp <= 0,
6779 2 => cmp > 0,
6780 3 => cmp >= 0,
6781 _ => false,
6782 })
6783}
6784
6785fn is_numeric_scalar(s: &ScalarValue) -> bool {
6786 scalar_as_f64(s).is_some()
6787}
6788
6789static CYPHER_STARTS_WITH: LazyLock<ScalarUDF> =
6790 LazyLock::new(|| ScalarUDF::new_from_impl(CypherStringPredicate::new(StringPredicate::Starts)));
6791static CYPHER_ENDS_WITH: LazyLock<ScalarUDF> =
6792 LazyLock::new(|| ScalarUDF::new_from_impl(CypherStringPredicate::new(StringPredicate::Ends)));
6793static CYPHER_CONTAINS: LazyLock<ScalarUDF> = LazyLock::new(|| {
6794 ScalarUDF::new_from_impl(CypherStringPredicate::new(StringPredicate::Contains))
6795});
6796
6797#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
6798enum StringPredicate {
6799 Starts,
6800 Ends,
6801 Contains,
6802}
6803
6804#[derive(Debug, PartialEq, Eq, Hash)]
6805struct CypherStringPredicate {
6806 signature: Signature,
6807 kind: StringPredicate,
6808}
6809
6810impl CypherStringPredicate {
6811 fn new(kind: StringPredicate) -> Self {
6812 Self {
6813 signature: Signature::any(2, Volatility::Immutable),
6814 kind,
6815 }
6816 }
6817}
6818
6819impl ScalarUDFImpl for CypherStringPredicate {
6820 fn as_any(&self) -> &dyn Any {
6821 self
6822 }
6823 fn name(&self) -> &'static str {
6824 match self.kind {
6825 StringPredicate::Starts => "cypher_starts_with",
6826 StringPredicate::Ends => "cypher_ends_with",
6827 StringPredicate::Contains => "cypher_contains",
6828 }
6829 }
6830 fn signature(&self) -> &Signature {
6831 &self.signature
6832 }
6833 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
6834 Ok(DataType::Boolean)
6835 }
6836 fn invoke_with_args(
6837 &self,
6838 args: ScalarFunctionArgs,
6839 ) -> datafusion::error::Result<ColumnarValue> {
6840 use datafusion::arrow::array::BooleanArray;
6841 let rows = args.number_rows;
6842 let lhs = args.args[0].to_array(rows)?;
6843 let rhs = args.args[1].to_array(rows)?;
6844 let out: BooleanArray = (0..rows)
6845 .map(|i| {
6846 let l = ScalarValue::try_from_array(&lhs, i).ok()?;
6847 let r = ScalarValue::try_from_array(&rhs, i).ok()?;
6848 let l = scalar_as_string(&l)?;
6849 let r = scalar_as_string(&r)?;
6850 Some(match self.kind {
6851 StringPredicate::Starts => l.starts_with(&r),
6852 StringPredicate::Ends => l.ends_with(&r),
6853 StringPredicate::Contains => l.contains(&r),
6854 })
6855 })
6856 .collect();
6857 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
6858 }
6859}
6860
6861fn scalar_as_string(s: &ScalarValue) -> Option<String> {
6862 let s = unwrap_het(s.clone());
6863 if s.is_null() {
6864 return None;
6865 }
6866 match s {
6867 ScalarValue::Utf8(v) | ScalarValue::LargeUtf8(v) => v,
6868 _ => None,
6869 }
6870}
6871
6872static CYPHER_RANGE: LazyLock<ScalarUDF> =
6873 LazyLock::new(|| ScalarUDF::new_from_impl(CypherRange::new()));
6874pub(crate) static CYPHER_ORDER_KEY: LazyLock<ScalarUDF> =
6875 LazyLock::new(|| ScalarUDF::new_from_impl(CypherOrderKey::new()));
6876
6877#[derive(Debug, PartialEq, Eq, Hash)]
6878struct CypherRange {
6879 signature: Signature,
6880}
6881
6882impl CypherRange {
6883 fn new() -> Self {
6884 Self {
6885 signature: Signature::any(3, Volatility::Volatile),
6889 }
6890 }
6891}
6892
6893impl ScalarUDFImpl for CypherRange {
6894 fn as_any(&self) -> &dyn Any {
6895 self
6896 }
6897 fn name(&self) -> &'static str {
6898 "cypher_range"
6899 }
6900 fn signature(&self) -> &Signature {
6901 &self.signature
6902 }
6903 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
6904 Ok(DataType::new_list(DataType::Int64, true))
6905 }
6906 fn invoke_with_args(
6907 &self,
6908 args: ScalarFunctionArgs,
6909 ) -> datafusion::error::Result<ColumnarValue> {
6910 use datafusion::arrow::array::{Int64Builder, ListBuilder};
6911 let rows = args.number_rows;
6912 let starts = args.args[0].to_array(rows)?;
6913 let ends = args.args[1].to_array(rows)?;
6914 let steps = args.args[2].to_array(rows)?;
6915 let mut out = ListBuilder::new(Int64Builder::new());
6916 for i in 0..rows {
6917 let start = ScalarValue::try_from_array(&starts, i)?;
6918 let end = ScalarValue::try_from_array(&ends, i)?;
6919 let step = ScalarValue::try_from_array(&steps, i)?;
6920 if start.is_null() || end.is_null() || step.is_null() {
6921 out.append_null();
6922 continue;
6923 }
6924 let start = scalar_as_i64_arg(&start, "range start")?;
6925 let end = scalar_as_i64_arg(&end, "range end")?;
6926 let step = scalar_as_i64_arg(&step, "range step")?;
6927 if step == 0 {
6928 return Err(datafusion::error::DataFusionError::Plan(
6929 "range step must not be zero".into(),
6930 ));
6931 }
6932 if (step > 0 && start > end) || (step < 0 && start < end) {
6933 out.append(true);
6934 continue;
6935 }
6936 let mut cur = start;
6937 loop {
6938 out.values().append_value(cur);
6939 if cur == end {
6940 break;
6941 }
6942 let Some(next) = cur.checked_add(step) else {
6943 return Err(datafusion::error::DataFusionError::Plan(
6944 "range overflowed i64".into(),
6945 ));
6946 };
6947 if (step > 0 && next > end) || (step < 0 && next < end) {
6948 break;
6949 }
6950 cur = next;
6951 }
6952 out.append(true);
6953 }
6954 Ok(ColumnarValue::Array(std::sync::Arc::new(out.finish())))
6955 }
6956}
6957
6958fn scalar_as_i64_arg(s: &ScalarValue, name: &str) -> datafusion::error::Result<i64> {
6959 match s {
6960 ScalarValue::Int8(Some(v)) => Ok(i64::from(*v)),
6961 ScalarValue::Int16(Some(v)) => Ok(i64::from(*v)),
6962 ScalarValue::Int32(Some(v)) => Ok(i64::from(*v)),
6963 ScalarValue::Int64(Some(v)) => Ok(*v),
6964 ScalarValue::UInt8(Some(v)) => Ok(i64::from(*v)),
6965 ScalarValue::UInt16(Some(v)) => Ok(i64::from(*v)),
6966 ScalarValue::UInt32(Some(v)) => Ok(i64::from(*v)),
6967 ScalarValue::UInt64(Some(v)) => i64::try_from(*v).map_err(|_| {
6968 datafusion::error::DataFusionError::Plan(format!("{name} exceeds i64::MAX"))
6969 }),
6970 other => Err(datafusion::error::DataFusionError::Plan(format!(
6971 "{name} must be an integer, got {other:?}"
6972 ))),
6973 }
6974}
6975
6976#[derive(Debug, PartialEq, Eq, Hash)]
6977struct CypherOrderKey {
6978 signature: Signature,
6979}
6980
6981impl CypherOrderKey {
6982 fn new() -> Self {
6983 Self {
6984 signature: Signature::any(1, Volatility::Immutable),
6985 }
6986 }
6987}
6988
6989impl ScalarUDFImpl for CypherOrderKey {
6990 fn as_any(&self) -> &dyn Any {
6991 self
6992 }
6993 fn name(&self) -> &'static str {
6994 "cypher_order_key"
6995 }
6996 fn signature(&self) -> &Signature {
6997 &self.signature
6998 }
6999 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
7000 Ok(DataType::Utf8)
7001 }
7002 fn invoke_with_args(
7003 &self,
7004 args: ScalarFunctionArgs,
7005 ) -> datafusion::error::Result<ColumnarValue> {
7006 use datafusion::arrow::array::StringArray;
7007 let rows = args.number_rows;
7008 let values = args.args[0].to_array(rows)?;
7009 let out: StringArray = (0..rows)
7010 .map(|i| {
7011 let v = ScalarValue::try_from_array(&values, i).ok()?;
7012 Some(cypher_order_key(&v))
7013 })
7014 .collect();
7015 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
7016 }
7017}
7018
7019pub(crate) fn needs_cypher_order_key_type(t: &DataType) -> bool {
7020 matches!(
7021 t,
7022 DataType::List(_) | DataType::LargeList(_) | DataType::FixedSizeList(_, _)
7023 ) || (matches!(t, DataType::Struct(_))
7024 && !is_date_struct(t)
7025 && !is_localdatetime_struct(t)
7026 && !is_duration_struct(t))
7027}
7028
7029fn cypher_order_key(v: &ScalarValue) -> String {
7030 let v = unwrap_het(v.clone());
7031 if v.is_null() {
7032 return "99:null".to_string();
7033 }
7034 match &v {
7035 ScalarValue::Struct(s) if is_time_struct(&v.data_type()) => {
7036 let Some((nanos, offset)) = time_struct_parts(s, 0) else {
7037 return "99:null".to_string();
7038 };
7039 let instant = i128::from(nanos) - i128::from(offset) * 1_000_000_000;
7040 format!("55:time:{}", ordered_i128_key(instant))
7041 }
7042 ScalarValue::Struct(s) if is_datetime_struct(&v.data_type()) => {
7043 let Some((days, nanos, offset, _)) = datetime_struct_parts(s, 0) else {
7044 return "99:null".to_string();
7045 };
7046 let instant = i128::from(days) * 86_400_000_000_000 + i128::from(nanos)
7047 - i128::from(offset) * 1_000_000_000;
7048 format!("55:datetime:{}", ordered_i128_key(instant))
7049 }
7050 ScalarValue::Struct(s) if is_path_struct(s) => "50:path".to_string(),
7051 ScalarValue::Struct(s) if is_rel_struct(s) => "30:rel".to_string(),
7052 ScalarValue::Struct(s) if is_node_struct(s) => "20:node".to_string(),
7053 ScalarValue::Struct(_) => "10:map".to_string(),
7054 ScalarValue::List(a) => format!("40:list:{}", cypher_list_order_key(&a.value(0))),
7055 ScalarValue::LargeList(a) => format!("40:list:{}", cypher_list_order_key(&a.value(0))),
7056 ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) => format!("60:str:{s}"),
7057 ScalarValue::Boolean(Some(b)) => format!("70:bool:{}", u8::from(*b)),
7058 ScalarValue::Int64(Some(n)) => {
7059 #[allow(
7060 clippy::cast_precision_loss,
7061 reason = "Cypher numeric order shares a number bucket across ints and floats"
7062 )]
7063 let n = *n as f64;
7064 format!("80:num:{}", ordered_f64_key(n))
7065 }
7066 ScalarValue::Float64(Some(f)) if f.is_nan() => "90:nan".to_string(),
7067 ScalarValue::Float64(Some(f)) => format!("80:num:{}", ordered_f64_key(*f)),
7068 _ => "98:other".to_string(),
7069 }
7070}
7071
7072fn cypher_list_order_key(values: &datafusion::arrow::array::ArrayRef) -> String {
7073 let mut out = String::new();
7074 for i in 0..values.len() {
7075 let v = ScalarValue::try_from_array(values, i).unwrap_or(ScalarValue::Null);
7076 out.push_str(&cypher_order_key(&v));
7077 out.push('|');
7078 }
7079 out.push_str("00:end");
7080 out
7081}
7082
7083fn ordered_f64_key(f: f64) -> String {
7084 let bits = f.to_bits();
7085 let key = if (bits >> 63) == 0 {
7086 bits | (1 << 63)
7087 } else {
7088 !bits
7089 };
7090 format!("{key:016x}")
7091}
7092
7093fn ordered_i128_key(value: i128) -> String {
7094 let key = value.cast_unsigned() ^ (1_u128 << 127);
7095 format!("{key:032x}")
7096}
7097
7098fn is_node_struct(s: &datafusion::arrow::array::StructArray) -> bool {
7099 s.column_by_name("node_uuid").is_some() || s.column_by_name("labels").is_some()
7100}
7101
7102fn is_rel_struct(s: &datafusion::arrow::array::StructArray) -> bool {
7103 s.column_by_name("edge_uuid").is_some()
7104 || s.column_by_name("src_uuid").is_some()
7105 || s.column_by_name("dst_uuid").is_some()
7106}
7107
7108fn is_path_struct(s: &datafusion::arrow::array::StructArray) -> bool {
7109 s.column_by_name("nodes").is_some() || s.column_by_name("relationships").is_some()
7110}
7111
7112static CYPHER_EQ: LazyLock<ScalarUDF> = LazyLock::new(|| ScalarUDF::new_from_impl(CypherEq::new()));
7122
7123#[derive(Debug, PartialEq, Eq, Hash)]
7124struct CypherEq {
7125 signature: Signature,
7126}
7127
7128impl CypherEq {
7129 fn new() -> Self {
7130 Self {
7131 signature: Signature::any(2, Volatility::Immutable),
7132 }
7133 }
7134}
7135
7136impl ScalarUDFImpl for CypherEq {
7137 fn as_any(&self) -> &dyn Any {
7138 self
7139 }
7140
7141 fn name(&self) -> &'static str {
7142 "cypher_eq"
7143 }
7144
7145 fn signature(&self) -> &Signature {
7146 &self.signature
7147 }
7148
7149 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
7150 Ok(DataType::Boolean)
7151 }
7152
7153 fn simplify(
7163 &self,
7164 args: Vec<DfExpr>,
7165 info: &datafusion::logical_expr::simplify::SimplifyContext,
7166 ) -> datafusion::error::Result<datafusion::logical_expr::simplify::ExprSimplifyResult> {
7167 use datafusion::logical_expr::simplify::ExprSimplifyResult;
7168 let [l, r] = args.as_slice() else {
7169 return Ok(ExprSimplifyResult::Original(args));
7170 };
7171 let nested = |t: &DataType| {
7172 matches!(
7173 t,
7174 DataType::List(_) | DataType::LargeList(_) | DataType::Struct(_)
7175 )
7176 };
7177 let floaty =
7178 |t: &DataType| matches!(t, DataType::Float16 | DataType::Float32 | DataType::Float64);
7179 let placeholder = |e: &DfExpr| matches!(e, DfExpr::Placeholder(_));
7183 let keep_udf = if placeholder(l) || placeholder(r) {
7184 false
7185 } else if let (Ok(lt), Ok(rt)) = (info.get_data_type(l), info.get_data_type(r)) {
7186 nested(<) || nested(&rt) || floaty(<) || floaty(&rt) || lt != rt
7187 } else {
7188 false };
7190 if keep_udf {
7191 return Ok(ExprSimplifyResult::Original(args));
7192 }
7193 let [l, r]: [DfExpr; 2] = args.try_into().expect("checked length 2 above");
7194 Ok(ExprSimplifyResult::Simplified(l.eq(r)))
7195 }
7196
7197 fn invoke_with_args(
7198 &self,
7199 args: ScalarFunctionArgs,
7200 ) -> datafusion::error::Result<ColumnarValue> {
7201 use datafusion::arrow::array::BooleanArray;
7202 use datafusion::arrow::compute::kernels::cmp::eq;
7203
7204 let rows = args.number_rows;
7205 let lhs = args.args[0].to_array(rows)?;
7206 let rhs = args.args[1].to_array(rows)?;
7207 let (lt, rt) = (lhs.data_type(), rhs.data_type());
7208
7209 let nested = |t: &DataType| {
7210 matches!(
7211 t,
7212 DataType::List(_) | DataType::LargeList(_) | DataType::Struct(_)
7213 )
7214 };
7215
7216 let floaty =
7217 |t: &DataType| matches!(t, DataType::Float16 | DataType::Float32 | DataType::Float64);
7218
7219 if lt == rt && !nested(lt) && !floaty(lt) {
7223 let res = eq(&lhs, &rhs)?;
7224 return Ok(ColumnarValue::Array(std::sync::Arc::new(res)));
7225 }
7226
7227 let out: BooleanArray = (0..rows)
7231 .map(|i| {
7232 let l = ScalarValue::try_from_array(&lhs, i).ok()?;
7233 let r = ScalarValue::try_from_array(&rhs, i).ok()?;
7234 cypher_value_eq(&l, &r)
7235 })
7236 .collect();
7237 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
7238 }
7239}
7240
7241static CYPHER_IN: LazyLock<ScalarUDF> = LazyLock::new(|| ScalarUDF::new_from_impl(CypherIn::new()));
7247
7248#[derive(Debug, PartialEq, Eq, Hash)]
7249struct CypherIn {
7250 signature: Signature,
7251}
7252
7253impl CypherIn {
7254 fn new() -> Self {
7255 Self {
7256 signature: Signature::any(2, Volatility::Immutable),
7257 }
7258 }
7259}
7260
7261enum ListView<'a> {
7262 Fixed(&'a FixedSizeListArray),
7263 List(&'a ListArray),
7264 Large(&'a LargeListArray),
7265}
7266
7267impl ListView<'_> {
7268 fn from_array(array: &datafusion::arrow::array::ArrayRef) -> Option<ListView<'_>> {
7269 if let Some(list) = array.as_any().downcast_ref::<FixedSizeListArray>() {
7270 Some(ListView::Fixed(list))
7271 } else if let Some(list) = array.as_any().downcast_ref::<ListArray>() {
7272 Some(ListView::List(list))
7273 } else {
7274 array
7275 .as_any()
7276 .downcast_ref::<LargeListArray>()
7277 .map(ListView::Large)
7278 }
7279 }
7280
7281 fn is_null(&self, row: usize) -> bool {
7282 match self {
7283 Self::Fixed(a) => a.is_null(row),
7284 Self::List(a) => a.is_null(row),
7285 Self::Large(a) => a.is_null(row),
7286 }
7287 }
7288
7289 fn value(&self, row: usize) -> datafusion::arrow::array::ArrayRef {
7290 match self {
7291 Self::Fixed(a) => a.value(row),
7292 Self::List(a) => a.value(row),
7293 Self::Large(a) => a.value(row),
7294 }
7295 }
7296}
7297
7298fn cypher_in_elems(lhs: &ScalarValue, elems: &datafusion::arrow::array::ArrayRef) -> Option<bool> {
7299 let mut saw_null = false;
7300 for j in 0..elems.len() {
7301 let ev = ScalarValue::try_from_array(elems, j).ok()?;
7302 match cypher_value_eq(lhs, &ev) {
7303 Some(true) => return Some(true),
7304 None => saw_null = true,
7305 Some(false) => {}
7306 }
7307 }
7308 if saw_null { None } else { Some(false) }
7309}
7310
7311fn cypher_in_tagged_list(
7312 lhs: &ScalarValue,
7313 rhs: &datafusion::arrow::array::ArrayRef,
7314 row: usize,
7315) -> Option<bool> {
7316 let rv = ScalarValue::try_from_array(rhs, row).ok()?;
7317 match unwrap_het(rv) {
7318 ScalarValue::List(list) => {
7319 if list.is_null(0) {
7320 None
7321 } else {
7322 cypher_in_elems(lhs, &list.value(0))
7323 }
7324 }
7325 ScalarValue::LargeList(list) => {
7326 if list.is_null(0) {
7327 None
7328 } else {
7329 cypher_in_elems(lhs, &list.value(0))
7330 }
7331 }
7332 _ => None,
7333 }
7334}
7335
7336impl ScalarUDFImpl for CypherIn {
7337 fn as_any(&self) -> &dyn Any {
7338 self
7339 }
7340 fn name(&self) -> &'static str {
7341 "cypher_in"
7342 }
7343 fn signature(&self) -> &Signature {
7344 &self.signature
7345 }
7346 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
7347 Ok(DataType::Boolean)
7348 }
7349
7350 fn invoke_with_args(
7351 &self,
7352 args: ScalarFunctionArgs,
7353 ) -> datafusion::error::Result<ColumnarValue> {
7354 use datafusion::arrow::array::BooleanArray;
7355 let rows = args.number_rows;
7356 let lhs = args.args[0].to_array(rows)?;
7357 let rhs = args.args[1].to_array(rows)?;
7358 if is_het_struct_type(Some(rhs.data_type())) {
7359 let out: BooleanArray = (0..rows)
7360 .map(|i| {
7361 let lv = ScalarValue::try_from_array(&lhs, i).ok()?;
7362 cypher_in_tagged_list(&lv, &rhs, i)
7363 })
7364 .collect();
7365 return Ok(ColumnarValue::Array(std::sync::Arc::new(out)));
7366 }
7367 let list = if let Some(list) = rhs.as_any().downcast_ref::<FixedSizeListArray>() {
7368 ListView::Fixed(list)
7369 } else if let Some(list) = rhs.as_any().downcast_ref::<ListArray>() {
7370 ListView::List(list)
7371 } else if let Some(list) = rhs.as_any().downcast_ref::<LargeListArray>() {
7372 ListView::Large(list)
7373 } else {
7374 return Ok(ColumnarValue::Array(std::sync::Arc::new(
7375 BooleanArray::new_null(rows),
7376 )));
7377 };
7378 let out: BooleanArray = (0..rows)
7379 .map(|i| {
7380 if list.is_null(i) {
7381 return None; }
7383 let lv = ScalarValue::try_from_array(&lhs, i).ok()?;
7384 let elems = list.value(i);
7385 cypher_in_elems(&lv, &elems)
7386 })
7387 .collect();
7388 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
7389 }
7390}
7391
7392fn cypher_compare(a: &ScalarValue, b: &ScalarValue) -> Option<i8> {
7400 let to_i8 = |o: std::cmp::Ordering| o as i8;
7401 let a = unwrap_het(a.clone());
7402 let b = unwrap_het(b.clone());
7403 if a.is_null() || b.is_null() {
7404 return None;
7405 }
7406 match (&a, &b) {
7407 _ if is_numeric_scalar(&a) && is_numeric_scalar(&b) => {
7408 if let (Some(x), Some(y)) = (scalar_as_i128(&a), scalar_as_i128(&b)) {
7409 Some(to_i8(x.cmp(&y)))
7410 } else {
7411 scalar_as_f64(&a)?
7412 .partial_cmp(&scalar_as_f64(&b)?)
7413 .map(to_i8)
7414 }
7415 }
7416 (ScalarValue::Utf8(Some(x)), ScalarValue::Utf8(Some(y))) => Some(to_i8(x.cmp(y))),
7417 (ScalarValue::Boolean(Some(x)), ScalarValue::Boolean(Some(y))) => Some(to_i8(x.cmp(y))),
7418 (ScalarValue::Time64Nanosecond(Some(x)), ScalarValue::Time64Nanosecond(Some(y))) => {
7419 Some(to_i8(x.cmp(y)))
7420 }
7421 (ScalarValue::List(x), ScalarValue::List(y)) => {
7422 cypher_seq_compare(&x.value(0), &y.value(0))
7423 }
7424 (ScalarValue::Struct(x), ScalarValue::Struct(y))
7425 if is_date_struct(&a.data_type()) && is_date_struct(&b.data_type()) =>
7426 {
7427 Some(to_i8(
7428 date_struct_value(x, 0)?.cmp(&date_struct_value(y, 0)?),
7429 ))
7430 }
7431 (ScalarValue::Struct(x), ScalarValue::Struct(y))
7432 if is_localdatetime_struct(&a.data_type())
7433 && is_localdatetime_struct(&b.data_type()) =>
7434 {
7435 Some(to_i8(
7436 localdatetime_struct_parts(x, 0)?.cmp(&localdatetime_struct_parts(y, 0)?),
7437 ))
7438 }
7439 (ScalarValue::Struct(x), ScalarValue::Struct(y))
7442 if is_time_struct(&a.data_type()) && is_time_struct(&b.data_type()) =>
7443 {
7444 let (xn, xo) = time_struct_parts(x, 0)?;
7445 let (yn, yo) = time_struct_parts(y, 0)?;
7446 let xi = i128::from(xn) - i128::from(xo) * 1_000_000_000;
7447 let yi = i128::from(yn) - i128::from(yo) * 1_000_000_000;
7448 Some(to_i8(xi.cmp(&yi)))
7449 }
7450 (ScalarValue::Struct(x), ScalarValue::Struct(y))
7451 if is_datetime_struct(&a.data_type()) && is_datetime_struct(&b.data_type()) =>
7452 {
7453 let (xd, xn, xo, _) = datetime_struct_parts(x, 0)?;
7454 let (yd, yn, yo, _) = datetime_struct_parts(y, 0)?;
7455 let xi = i128::from(xd) * 86_400_000_000_000 + i128::from(xn)
7456 - i128::from(xo) * 1_000_000_000;
7457 let yi = i128::from(yd) * 86_400_000_000_000 + i128::from(yn)
7458 - i128::from(yo) * 1_000_000_000;
7459 Some(to_i8(xi.cmp(&yi)))
7460 }
7461 _ => None, }
7463}
7464
7465fn cypher_seq_compare(
7467 a: &datafusion::arrow::array::ArrayRef,
7468 b: &datafusion::arrow::array::ArrayRef,
7469) -> Option<i8> {
7470 let common = a.len().min(b.len());
7471 for i in 0..common {
7472 let av = ScalarValue::try_from_array(a, i).ok()?;
7473 let bv = ScalarValue::try_from_array(b, i).ok()?;
7474 match cypher_compare(&av, &bv)? {
7475 0 => {}
7476 c => return Some(c),
7477 }
7478 }
7479 Some(a.len().cmp(&b.len()) as i8) }
7481
7482fn is_list_literal(e: &DfExpr) -> bool {
7484 matches!(e, DfExpr::Literal(ScalarValue::List(_), _))
7485}
7486
7487pub(crate) fn is_het_struct_type(t: Option<&DataType>) -> bool {
7490 matches!(t, Some(DataType::Struct(fields)) if fields.iter().any(|f| f.name() == "__het_tag"))
7491}
7492
7493fn cypher_order(a: &ScalarValue, b: &ScalarValue) -> std::cmp::Ordering {
7500 use std::cmp::Ordering;
7501 let a = unwrap_het(a.clone());
7502 let b = unwrap_het(b.clone());
7503 let rank = |v: &ScalarValue| -> u8 {
7504 match v {
7505 ScalarValue::List(_) | ScalarValue::LargeList(_) => 1,
7506 ScalarValue::Utf8(_) | ScalarValue::LargeUtf8(_) => 2,
7507 ScalarValue::Boolean(_) => 3,
7508 _ if is_numeric_scalar(v) => 4,
7509 ScalarValue::Struct(_) => 5,
7510 _ => 0,
7511 }
7512 };
7513 let (ra, rb) = (rank(&a), rank(&b));
7514 if ra != rb {
7515 return ra.cmp(&rb);
7516 }
7517 match (&a, &b) {
7518 (ScalarValue::List(x), ScalarValue::List(y)) => cypher_seq_order(&x.value(0), &y.value(0)),
7519 (ScalarValue::Utf8(Some(x)), ScalarValue::Utf8(Some(y))) => x.cmp(y),
7520 (ScalarValue::Boolean(Some(x)), ScalarValue::Boolean(Some(y))) => x.cmp(y),
7521 (ScalarValue::Struct(x), ScalarValue::Struct(y)) => cypher_map_order(x, y),
7522 _ => match (scalar_as_f64(&a), scalar_as_f64(&b)) {
7523 (Some(x), Some(y)) => x.partial_cmp(&y).unwrap_or(Ordering::Equal),
7524 _ => Ordering::Equal,
7525 },
7526 }
7527}
7528
7529fn cypher_map_order(
7533 a: &datafusion::arrow::array::StructArray,
7534 b: &datafusion::arrow::array::StructArray,
7535) -> std::cmp::Ordering {
7536 let sorted = |s: &datafusion::arrow::array::StructArray| -> Vec<(String, ScalarValue)> {
7537 let mut e: Vec<(String, ScalarValue)> = s
7538 .fields()
7539 .iter()
7540 .enumerate()
7541 .map(|(i, f)| {
7542 (
7543 f.name().clone(),
7544 ScalarValue::try_from_array(s.column(i), 0).unwrap_or(ScalarValue::Null),
7545 )
7546 })
7547 .collect();
7548 e.sort_by(|x, y| x.0.cmp(&y.0));
7549 e
7550 };
7551 let (ea, eb) = (sorted(a), sorted(b));
7552 for ((ka, va), (kb, vb)) in ea.iter().zip(eb.iter()) {
7553 match ka.cmp(kb) {
7554 std::cmp::Ordering::Equal => {}
7555 c => return c,
7556 }
7557 match cypher_order(va, vb) {
7558 std::cmp::Ordering::Equal => {}
7559 c => return c,
7560 }
7561 }
7562 ea.len().cmp(&eb.len())
7563}
7564
7565fn cypher_seq_order(
7568 a: &datafusion::arrow::array::ArrayRef,
7569 b: &datafusion::arrow::array::ArrayRef,
7570) -> std::cmp::Ordering {
7571 let common = a.len().min(b.len());
7572 for i in 0..common {
7573 let av = ScalarValue::try_from_array(a, i).unwrap_or(ScalarValue::Null);
7574 let bv = ScalarValue::try_from_array(b, i).unwrap_or(ScalarValue::Null);
7575 match cypher_order(&av, &bv) {
7576 std::cmp::Ordering::Equal => {}
7577 c => return c,
7578 }
7579 }
7580 a.len().cmp(&b.len())
7581}
7582
7583pub(crate) static CYPHER_MAX: LazyLock<datafusion::logical_expr::AggregateUDF> =
7587 LazyLock::new(|| {
7588 datafusion::logical_expr::AggregateUDF::new_from_impl(CypherExtreme::new(true))
7589 });
7590pub(crate) static CYPHER_MIN: LazyLock<datafusion::logical_expr::AggregateUDF> =
7591 LazyLock::new(|| {
7592 datafusion::logical_expr::AggregateUDF::new_from_impl(CypherExtreme::new(false))
7593 });
7594pub(crate) static CYPHER_COLLECT: LazyLock<datafusion::logical_expr::AggregateUDF> =
7595 LazyLock::new(|| {
7596 datafusion::logical_expr::AggregateUDF::new_from_impl(CypherCollect::new(false))
7597 });
7598pub(crate) static CYPHER_COLLECT_DISTINCT: LazyLock<datafusion::logical_expr::AggregateUDF> =
7599 LazyLock::new(|| {
7600 datafusion::logical_expr::AggregateUDF::new_from_impl(CypherCollect::new(true))
7601 });
7602pub(crate) static CYPHER_PERCENTILE_DISC: LazyLock<datafusion::logical_expr::AggregateUDF> =
7603 LazyLock::new(|| {
7604 datafusion::logical_expr::AggregateUDF::new_from_impl(CypherPercentile::new(false))
7605 });
7606pub(crate) static CYPHER_PERCENTILE_CONT: LazyLock<datafusion::logical_expr::AggregateUDF> =
7607 LazyLock::new(|| {
7608 datafusion::logical_expr::AggregateUDF::new_from_impl(CypherPercentile::new(true))
7609 });
7610
7611#[derive(Debug)]
7612struct CypherExtreme {
7613 signature: Signature,
7614 is_max: bool,
7615}
7616impl CypherExtreme {
7617 fn new(is_max: bool) -> Self {
7618 Self {
7619 signature: Signature::any(1, Volatility::Immutable),
7620 is_max,
7621 }
7622 }
7623}
7624impl PartialEq for CypherExtreme {
7625 fn eq(&self, o: &Self) -> bool {
7626 self.is_max == o.is_max
7627 }
7628}
7629impl Eq for CypherExtreme {}
7630impl std::hash::Hash for CypherExtreme {
7631 fn hash<H: std::hash::Hasher>(&self, st: &mut H) {
7632 self.is_max.hash(st);
7633 }
7634}
7635impl datafusion::logical_expr::AggregateUDFImpl for CypherExtreme {
7636 fn as_any(&self) -> &dyn Any {
7637 self
7638 }
7639 fn name(&self) -> &str {
7640 if self.is_max {
7641 "cypher_max"
7642 } else {
7643 "cypher_min"
7644 }
7645 }
7646 fn signature(&self) -> &Signature {
7647 &self.signature
7648 }
7649 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
7650 Ok(arg_types[0].clone())
7651 }
7652 fn accumulator(
7653 &self,
7654 args: datafusion::logical_expr::function::AccumulatorArgs,
7655 ) -> datafusion::error::Result<Box<dyn datafusion::logical_expr::Accumulator>> {
7656 Ok(Box::new(ExtremeAcc {
7657 is_max: self.is_max,
7658 dtype: args.return_field.data_type().clone(),
7659 best: None,
7660 }))
7661 }
7662 fn state_fields(
7663 &self,
7664 args: datafusion::logical_expr::function::StateFieldsArgs,
7665 ) -> datafusion::error::Result<Vec<datafusion::arrow::datatypes::FieldRef>> {
7666 use datafusion::arrow::datatypes::Field;
7667 Ok(vec![std::sync::Arc::new(Field::new(
7668 "best",
7669 args.return_field.data_type().clone(),
7670 true,
7671 ))])
7672 }
7673}
7674
7675#[derive(Debug)]
7676struct ExtremeAcc {
7677 is_max: bool,
7678 dtype: DataType,
7679 best: Option<ScalarValue>,
7680}
7681impl datafusion::logical_expr::Accumulator for ExtremeAcc {
7682 fn update_batch(
7683 &mut self,
7684 values: &[datafusion::arrow::array::ArrayRef],
7685 ) -> datafusion::error::Result<()> {
7686 use datafusion::arrow::array::Array;
7687 let arr = &values[0];
7688 for i in 0..arr.len() {
7689 if arr.is_null(i) {
7690 continue;
7691 }
7692 let v = ScalarValue::try_from_array(arr, i)?;
7693 if v.is_null() {
7694 continue;
7695 }
7696 let take = match &self.best {
7697 None => true,
7698 Some(b) => {
7699 let ord = cypher_order(&v, b);
7700 (self.is_max && ord == std::cmp::Ordering::Greater)
7701 || (!self.is_max && ord == std::cmp::Ordering::Less)
7702 }
7703 };
7704 if take {
7705 self.best = Some(v);
7706 }
7707 }
7708 Ok(())
7709 }
7710 fn evaluate(&mut self) -> datafusion::error::Result<ScalarValue> {
7711 match &self.best {
7712 Some(b) => Ok(b.clone()),
7713 None => ScalarValue::try_from(&self.dtype),
7714 }
7715 }
7716 fn size(&self) -> usize {
7717 std::mem::size_of_val(self) + self.best.as_ref().map_or(0, ScalarValue::size)
7718 }
7719 fn state(&mut self) -> datafusion::error::Result<Vec<ScalarValue>> {
7720 Ok(vec![self.evaluate()?])
7721 }
7722 fn merge_batch(
7723 &mut self,
7724 states: &[datafusion::arrow::array::ArrayRef],
7725 ) -> datafusion::error::Result<()> {
7726 self.update_batch(states)
7727 }
7728}
7729
7730#[derive(Debug)]
7731struct CypherCollect {
7732 signature: Signature,
7733 distinct: bool,
7734}
7735
7736impl CypherCollect {
7737 fn new(distinct: bool) -> Self {
7738 Self {
7739 signature: Signature::any(1, Volatility::Immutable),
7740 distinct,
7741 }
7742 }
7743}
7744
7745impl PartialEq for CypherCollect {
7746 fn eq(&self, o: &Self) -> bool {
7747 self.distinct == o.distinct
7748 }
7749}
7750
7751impl Eq for CypherCollect {}
7752
7753impl std::hash::Hash for CypherCollect {
7754 fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
7755 self.distinct.hash(state);
7756 }
7757}
7758
7759impl datafusion::logical_expr::AggregateUDFImpl for CypherCollect {
7760 fn as_any(&self) -> &dyn Any {
7761 self
7762 }
7763 fn name(&self) -> &str {
7764 if self.distinct {
7765 "cypher_collect_distinct"
7766 } else {
7767 "cypher_collect"
7768 }
7769 }
7770 fn signature(&self) -> &Signature {
7771 &self.signature
7772 }
7773 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
7774 Ok(DataType::new_list(arg_types[0].clone(), true))
7775 }
7776 fn accumulator(
7777 &self,
7778 args: datafusion::logical_expr::function::AccumulatorArgs,
7779 ) -> datafusion::error::Result<Box<dyn datafusion::logical_expr::Accumulator>> {
7780 let DataType::List(field) = args.return_field.data_type() else {
7781 return Err(datafusion::error::DataFusionError::Plan(
7782 "cypher_collect return type must be a list".into(),
7783 ));
7784 };
7785 Ok(Box::new(CollectAcc {
7786 distinct: self.distinct,
7787 elem_type: field.data_type().clone(),
7788 values: Vec::new(),
7789 }))
7790 }
7791 fn state_fields(
7792 &self,
7793 args: datafusion::logical_expr::function::StateFieldsArgs,
7794 ) -> datafusion::error::Result<Vec<datafusion::arrow::datatypes::FieldRef>> {
7795 use datafusion::arrow::datatypes::Field;
7796 Ok(vec![std::sync::Arc::new(Field::new(
7797 "values",
7798 args.return_field.data_type().clone(),
7799 true,
7800 ))])
7801 }
7802}
7803
7804#[derive(Debug)]
7805struct CollectAcc {
7806 distinct: bool,
7807 elem_type: DataType,
7808 values: Vec<ScalarValue>,
7809}
7810
7811impl CollectAcc {
7812 fn push_value(&mut self, v: ScalarValue) {
7813 if v.is_null() {
7814 return;
7815 }
7816 if self.distinct
7817 && self
7818 .values
7819 .iter()
7820 .any(|seen| cypher_value_eq(seen, &v) == Some(true))
7821 {
7822 return;
7823 }
7824 self.values.push(v);
7825 }
7826
7827 fn as_list(&self) -> ScalarValue {
7828 ScalarValue::List(ScalarValue::new_list(&self.values, &self.elem_type, true))
7829 }
7830}
7831
7832impl datafusion::logical_expr::Accumulator for CollectAcc {
7833 fn update_batch(
7834 &mut self,
7835 values: &[datafusion::arrow::array::ArrayRef],
7836 ) -> datafusion::error::Result<()> {
7837 use datafusion::arrow::array::Array;
7838 let arr = &values[0];
7839 for i in 0..arr.len() {
7840 if arr.is_null(i) {
7841 continue;
7842 }
7843 self.push_value(ScalarValue::try_from_array(arr, i)?);
7844 }
7845 Ok(())
7846 }
7847 fn evaluate(&mut self) -> datafusion::error::Result<ScalarValue> {
7848 Ok(self.as_list())
7849 }
7850 fn size(&self) -> usize {
7851 std::mem::size_of_val(self) + self.values.iter().map(ScalarValue::size).sum::<usize>()
7852 }
7853 fn state(&mut self) -> datafusion::error::Result<Vec<ScalarValue>> {
7854 Ok(vec![self.as_list()])
7855 }
7856 fn merge_batch(
7857 &mut self,
7858 states: &[datafusion::arrow::array::ArrayRef],
7859 ) -> datafusion::error::Result<()> {
7860 use datafusion::arrow::array::{Array, ListArray};
7861 let arr = &states[0];
7862 let Some(list) = arr.as_any().downcast_ref::<ListArray>() else {
7863 return Err(datafusion::error::DataFusionError::Plan(
7864 "cypher_collect state must be a list".into(),
7865 ));
7866 };
7867 for row in 0..list.len() {
7868 if list.is_null(row) {
7869 continue;
7870 }
7871 let values = list.value(row);
7872 for i in 0..values.len() {
7873 if values.is_null(i) {
7874 continue;
7875 }
7876 self.push_value(ScalarValue::try_from_array(&values, i)?);
7877 }
7878 }
7879 Ok(())
7880 }
7881}
7882
7883#[derive(Debug)]
7884struct CypherPercentile {
7885 signature: Signature,
7886 continuous: bool,
7887}
7888
7889impl CypherPercentile {
7890 fn new(continuous: bool) -> Self {
7891 Self {
7892 signature: Signature::any(2, Volatility::Immutable),
7893 continuous,
7894 }
7895 }
7896}
7897
7898impl PartialEq for CypherPercentile {
7899 fn eq(&self, o: &Self) -> bool {
7900 self.continuous == o.continuous
7901 }
7902}
7903
7904impl Eq for CypherPercentile {}
7905
7906impl std::hash::Hash for CypherPercentile {
7907 fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
7908 self.continuous.hash(state);
7909 }
7910}
7911
7912impl datafusion::logical_expr::AggregateUDFImpl for CypherPercentile {
7913 fn as_any(&self) -> &dyn Any {
7914 self
7915 }
7916
7917 fn name(&self) -> &str {
7918 if self.continuous {
7919 "cypher_percentile_cont"
7920 } else {
7921 "cypher_percentile_disc"
7922 }
7923 }
7924
7925 fn signature(&self) -> &Signature {
7926 &self.signature
7927 }
7928
7929 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
7930 let Some(value_type) = arg_types.first() else {
7931 return Err(datafusion::error::DataFusionError::Plan(
7932 "percentile aggregate requires value and percentile arguments".into(),
7933 ));
7934 };
7935 if !is_percentile_numeric_type(value_type) {
7936 return Err(datafusion::error::DataFusionError::Plan(format!(
7937 "percentile value expression must be numeric, got {value_type}"
7938 )));
7939 }
7940 if self.continuous {
7941 Ok(DataType::Float64)
7942 } else {
7943 Ok(value_type.clone())
7944 }
7945 }
7946
7947 fn accumulator(
7948 &self,
7949 args: datafusion::logical_expr::function::AccumulatorArgs,
7950 ) -> datafusion::error::Result<Box<dyn datafusion::logical_expr::Accumulator>> {
7951 let value_type = args.expr_fields.first().map_or_else(
7952 || args.return_field.data_type().clone(),
7953 |f| f.data_type().clone(),
7954 );
7955 Ok(Box::new(PercentileAcc {
7956 continuous: self.continuous,
7957 value_type,
7958 result_type: args.return_field.data_type().clone(),
7959 values: Vec::new(),
7960 percentile: None,
7961 }))
7962 }
7963
7964 fn state_fields(
7965 &self,
7966 args: datafusion::logical_expr::function::StateFieldsArgs,
7967 ) -> datafusion::error::Result<Vec<datafusion::arrow::datatypes::FieldRef>> {
7968 use datafusion::arrow::datatypes::Field;
7969 let value_type = args.input_fields.first().map_or_else(
7970 || args.return_field.data_type().clone(),
7971 |f| f.data_type().clone(),
7972 );
7973 Ok(vec![
7974 std::sync::Arc::new(Field::new(
7975 "values",
7976 DataType::new_list(value_type, true),
7977 true,
7978 )),
7979 std::sync::Arc::new(Field::new("percentile", DataType::Float64, true)),
7980 ])
7981 }
7982}
7983
7984#[derive(Debug)]
7985struct PercentileAcc {
7986 continuous: bool,
7987 value_type: DataType,
7988 result_type: DataType,
7989 values: Vec<ScalarValue>,
7990 percentile: Option<f64>,
7991}
7992
7993impl PercentileAcc {
7994 fn push_value(&mut self, v: ScalarValue) -> datafusion::error::Result<()> {
7995 if v.is_null() {
7996 return Ok(());
7997 }
7998 if scalar_as_f64(&v).is_none() {
7999 return Err(datafusion::error::DataFusionError::Execution(format!(
8000 "percentile value expression must be numeric, got {}",
8001 v.data_type()
8002 )));
8003 }
8004 self.values.push(v);
8005 Ok(())
8006 }
8007
8008 fn observe_percentile(&mut self, p: Option<f64>) -> datafusion::error::Result<()> {
8009 let Some(p) = p else {
8010 return Ok(());
8011 };
8012 if !p.is_finite() || !(0.0..=1.0).contains(&p) {
8013 return Err(datafusion::error::DataFusionError::Execution(format!(
8014 "percentile argument must be a finite number between 0.0 and 1.0 inclusive, got {p}"
8015 )));
8016 }
8017 match self.percentile {
8018 Some(existing) if (existing - p).abs() > f64::EPSILON => {
8019 Err(datafusion::error::DataFusionError::Execution(
8020 "percentile argument must be constant within an aggregate group".into(),
8021 ))
8022 }
8023 Some(_) => Ok(()),
8024 None => {
8025 self.percentile = Some(p);
8026 Ok(())
8027 }
8028 }
8029 }
8030
8031 fn null_result(&self) -> datafusion::error::Result<ScalarValue> {
8032 ScalarValue::try_from(&self.result_type)
8033 }
8034
8035 fn percentile_scalar(
8036 values: &[datafusion::arrow::array::ArrayRef],
8037 row: usize,
8038 ) -> datafusion::error::Result<Option<f64>> {
8039 use datafusion::arrow::array::Array;
8040 let arr = &values[1];
8041 if arr.is_null(row) {
8042 return Ok(None);
8043 }
8044 let scalar = ScalarValue::try_from_array(arr, row)?;
8045 if scalar.is_null() {
8046 Ok(None)
8047 } else {
8048 scalar_as_f64(&scalar).map(Some).ok_or_else(|| {
8049 datafusion::error::DataFusionError::Execution(format!(
8050 "percentile argument must be numeric, got {}",
8051 scalar.data_type()
8052 ))
8053 })
8054 }
8055 }
8056}
8057
8058impl datafusion::logical_expr::Accumulator for PercentileAcc {
8059 fn update_batch(
8060 &mut self,
8061 values: &[datafusion::arrow::array::ArrayRef],
8062 ) -> datafusion::error::Result<()> {
8063 use datafusion::arrow::array::Array;
8064 let value_arr = &values[0];
8065 for row in 0..value_arr.len() {
8066 self.observe_percentile(Self::percentile_scalar(values, row)?)?;
8067 if value_arr.is_null(row) {
8068 continue;
8069 }
8070 self.push_value(ScalarValue::try_from_array(value_arr, row)?)?;
8071 }
8072 Ok(())
8073 }
8074
8075 fn evaluate(&mut self) -> datafusion::error::Result<ScalarValue> {
8076 let Some(percentile) = self.percentile else {
8077 return self.null_result();
8078 };
8079 if self.values.is_empty() {
8080 return self.null_result();
8081 }
8082 let mut values: Vec<(f64, ScalarValue)> = self
8083 .values
8084 .iter()
8085 .filter_map(|v| scalar_as_f64(v).map(|f| (f, v.clone())))
8086 .collect();
8087 if values.is_empty() {
8088 return self.null_result();
8089 }
8090 values.sort_by(|(l, _), (r, _)| l.total_cmp(r));
8091
8092 if self.continuous {
8093 let len = values.len();
8094 if len == 1 {
8095 return Ok(ScalarValue::Float64(Some(values[0].0)));
8096 }
8097 let (lower_index, upper_index, fraction) = percentile_cont_indices(percentile, len);
8098 let result = if lower_index == upper_index {
8099 values[lower_index].0
8100 } else {
8101 let lower = values[lower_index].0;
8102 let upper = values[upper_index].0;
8103 lower + (upper - lower) * fraction
8104 };
8105 Ok(ScalarValue::Float64(Some(result)))
8106 } else {
8107 let index = percentile_disc_index(percentile, values.len());
8108 Ok(values[index].1.clone())
8109 }
8110 }
8111
8112 fn size(&self) -> usize {
8113 std::mem::size_of_val(self) + self.values.iter().map(ScalarValue::size).sum::<usize>()
8114 }
8115
8116 fn state(&mut self) -> datafusion::error::Result<Vec<ScalarValue>> {
8117 Ok(vec![
8118 ScalarValue::List(ScalarValue::new_list(&self.values, &self.value_type, true)),
8119 ScalarValue::Float64(self.percentile),
8120 ])
8121 }
8122
8123 fn merge_batch(
8124 &mut self,
8125 states: &[datafusion::arrow::array::ArrayRef],
8126 ) -> datafusion::error::Result<()> {
8127 use datafusion::arrow::array::{Array, Float64Array, ListArray};
8128 let values = &states[0];
8129 let Some(lists) = values.as_any().downcast_ref::<ListArray>() else {
8130 return Err(datafusion::error::DataFusionError::Plan(
8131 "percentile state values must be a list".into(),
8132 ));
8133 };
8134 let Some(percentiles) = states[1].as_any().downcast_ref::<Float64Array>() else {
8135 return Err(datafusion::error::DataFusionError::Plan(
8136 "percentile state percentile must be Float64".into(),
8137 ));
8138 };
8139 for row in 0..lists.len() {
8140 self.observe_percentile(if percentiles.is_null(row) {
8141 None
8142 } else {
8143 Some(percentiles.value(row))
8144 })?;
8145 if lists.is_null(row) {
8146 continue;
8147 }
8148 let values = lists.value(row);
8149 for i in 0..values.len() {
8150 if values.is_null(i) {
8151 continue;
8152 }
8153 self.push_value(ScalarValue::try_from_array(&values, i)?)?;
8154 }
8155 }
8156 Ok(())
8157 }
8158}
8159
8160#[allow(
8161 clippy::cast_possible_truncation,
8162 clippy::cast_precision_loss,
8163 clippy::cast_sign_loss,
8164 reason = "percentile ranks are defined by converting bounded [0, 1] floats into sorted indexes"
8165)]
8166fn percentile_cont_indices(percentile: f64, len: usize) -> (usize, usize, f64) {
8167 let index = percentile * ((len - 1) as f64);
8168 let lower = index.floor() as usize;
8169 let upper = index.ceil() as usize;
8170 (lower, upper, index.fract())
8171}
8172
8173#[allow(
8174 clippy::cast_possible_truncation,
8175 clippy::cast_precision_loss,
8176 clippy::cast_sign_loss,
8177 reason = "percentile ranks are defined by converting bounded [0, 1] floats into sorted indexes"
8178)]
8179fn percentile_disc_index(percentile: f64, len: usize) -> usize {
8180 if percentile <= f64::EPSILON {
8181 0
8182 } else {
8183 ((percentile * (len as f64)).ceil() as usize)
8184 .saturating_sub(1)
8185 .min(len - 1)
8186 }
8187}
8188
8189fn is_percentile_numeric_type(dt: &DataType) -> bool {
8190 matches!(
8191 dt,
8192 DataType::Int8
8193 | DataType::Null
8194 | DataType::Int16
8195 | DataType::Int32
8196 | DataType::Int64
8197 | DataType::UInt8
8198 | DataType::UInt16
8199 | DataType::UInt32
8200 | DataType::UInt64
8201 | DataType::Float32
8202 | DataType::Float64
8203 )
8204}
8205
8206fn cypher_value_eq(l: &ScalarValue, r: &ScalarValue) -> Option<bool> {
8212 if l.is_null() || r.is_null() {
8213 return None;
8214 }
8215 if let Some(dl) = decode_het(l) {
8219 return cypher_value_eq(&dl, r);
8220 }
8221 if let Some(dr) = decode_het(r) {
8222 return cypher_value_eq(l, &dr);
8223 }
8224 match (l, r) {
8225 (ScalarValue::List(a), ScalarValue::List(b)) => cypher_seq_eq(&a.value(0), &b.value(0)),
8226 (ScalarValue::LargeList(a), ScalarValue::LargeList(b)) => {
8227 cypher_seq_eq(&a.value(0), &b.value(0))
8228 }
8229 (ScalarValue::Struct(a), ScalarValue::Struct(b)) => {
8234 if is_entity_struct(a) || is_entity_struct(b) {
8235 Some(l == r)
8236 } else {
8237 cypher_struct_eq(a, b)
8238 }
8239 }
8240 _ if is_numeric_scalar(l) && is_numeric_scalar(r) => {
8241 if let (Some(li), Some(ri)) = (scalar_as_i128(l), scalar_as_i128(r)) {
8242 return Some(li == ri);
8243 }
8244 let (Some(lf), Some(rf)) = (scalar_as_f64(l), scalar_as_f64(r)) else {
8245 return None;
8246 };
8247 Some(!lf.is_nan() && !rf.is_nan() && lf == rf)
8248 }
8249 _ if std::mem::discriminant(l) == std::mem::discriminant(r) => Some(l == r),
8251 _ => Some(false),
8252 }
8253}
8254
8255fn decode_het(s: &ScalarValue) -> Option<ScalarValue> {
8262 use datafusion::arrow::array::{
8263 Array, ArrayRef, BooleanArray, Float64Array, Int8Array, Int64Array, ListArray, StringArray,
8264 StructArray,
8265 };
8266 use datafusion::arrow::datatypes::{Field, Fields};
8267 use std::sync::Arc;
8268 let ScalarValue::Struct(arr) = s else {
8269 return None;
8270 };
8271 arr.column_by_name("__het_tag")?; if arr.is_null(0) {
8273 return Some(ScalarValue::Null);
8274 }
8275 let tag = arr
8276 .column_by_name("__het_tag")?
8277 .as_any()
8278 .downcast_ref::<Int8Array>()?
8279 .value(0);
8280 let col = |name: &str| arr.column_by_name(name);
8281 if let Some(value) = col(&format!("__het_value_{tag}")) {
8282 return ScalarValue::try_from_array(value, 0).ok();
8283 }
8284 let v = match tag {
8285 0 => ScalarValue::Int64(Some(
8286 col("__het_int")?
8287 .as_any()
8288 .downcast_ref::<Int64Array>()?
8289 .value(0),
8290 )),
8291 1 => ScalarValue::Float64(Some(
8292 col("__het_float")?
8293 .as_any()
8294 .downcast_ref::<Float64Array>()?
8295 .value(0),
8296 )),
8297 2 => ScalarValue::Utf8(Some(
8298 col("__het_str")?
8299 .as_any()
8300 .downcast_ref::<StringArray>()?
8301 .value(0)
8302 .to_string(),
8303 )),
8304 3 => ScalarValue::Boolean(Some(
8305 col("__het_bool")?
8306 .as_any()
8307 .downcast_ref::<BooleanArray>()?
8308 .value(0),
8309 )),
8310 4 => ScalarValue::try_from_array(col("__het_list")?, 0).ok()?,
8311 5 => {
8312 let entries = col("__het_map")?
8313 .as_any()
8314 .downcast_ref::<ListArray>()?
8315 .value(0);
8316 let es = entries.as_any().downcast_ref::<StructArray>()?;
8317 if es.is_empty() {
8318 return Some(ScalarValue::Struct(Arc::new(
8319 StructArray::new_empty_fields(1, None),
8320 )));
8321 }
8322 let mkeys = es
8323 .column_by_name("__het_mkey")?
8324 .as_any()
8325 .downcast_ref::<StringArray>()?;
8326 let mvals = es.column_by_name("__het_mval")?;
8327 let mut fields: Vec<Field> = Vec::with_capacity(es.len());
8328 let mut cols: Vec<ArrayRef> = Vec::with_capacity(es.len());
8329 for i in 0..es.len() {
8330 let varr = mvals.slice(i, 1);
8333 fields.push(Field::new(mkeys.value(i), varr.data_type().clone(), true));
8334 cols.push(varr);
8335 }
8336 ScalarValue::Struct(Arc::new(
8337 StructArray::try_new(Fields::from(fields), cols, None).ok()?,
8338 ))
8339 }
8340 _ => return None,
8341 };
8342 Some(v)
8343}
8344
8345#[must_use]
8347pub fn decode_het_scalar(value: &ScalarValue) -> Option<ScalarValue> {
8348 decode_het(value)
8349}
8350
8351fn is_entity_struct(s: &datafusion::arrow::array::StructArray) -> bool {
8355 s.fields().iter().any(|f| {
8356 matches!(
8357 f.name().as_str(),
8358 "node_uuid" | "src_uuid" | "dst_uuid" | "nodes" | "relationships" | "labels"
8359 )
8360 })
8361}
8362
8363static CYPHER_DATE_COMPONENT: LazyLock<ScalarUDF> =
8372 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDateComponent::new()));
8373
8374#[derive(Debug, PartialEq, Eq, Hash)]
8375struct CypherDateComponent {
8376 signature: Signature,
8377}
8378
8379impl CypherDateComponent {
8380 fn new() -> Self {
8381 Self {
8382 signature: Signature::any(2, Volatility::Immutable),
8383 }
8384 }
8385}
8386
8387impl ScalarUDFImpl for CypherDateComponent {
8388 fn as_any(&self) -> &dyn Any {
8389 self
8390 }
8391
8392 fn name(&self) -> &'static str {
8393 "cypher_date_component"
8394 }
8395
8396 fn signature(&self) -> &Signature {
8397 &self.signature
8398 }
8399
8400 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
8401 Ok(DataType::Int64)
8402 }
8403
8404 fn invoke_with_args(
8405 &self,
8406 args: ScalarFunctionArgs,
8407 ) -> datafusion::error::Result<ColumnarValue> {
8408 use crate::temporal::date_component;
8409 use datafusion::arrow::array::{Array, Int64Array, StringArray, StructArray};
8410
8411 let rows = args.number_rows;
8412 let dates = args.args[0].to_array(rows)?;
8413 let names = args.args[1].to_array(rows)?;
8414 let d = dates.as_any().downcast_ref::<StructArray>();
8415 let n = names.as_any().downcast_ref::<StringArray>();
8416 let out: Int64Array = (0..rows)
8417 .map(|i| {
8418 let (d, n) = (d?, n?);
8419 if n.is_null(i) {
8420 return None;
8421 }
8422 date_component(date_struct_value(d, i)?, n.value(i))
8423 })
8424 .collect();
8425 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
8426 }
8427}
8428
8429static CYPHER_DURATION_COMPONENT: LazyLock<ScalarUDF> =
8436 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDurationComponent::new()));
8437
8438#[derive(Debug, PartialEq, Eq, Hash)]
8439struct CypherDurationComponent {
8440 signature: Signature,
8441}
8442
8443impl CypherDurationComponent {
8444 fn new() -> Self {
8445 Self {
8446 signature: Signature::any(2, Volatility::Immutable),
8447 }
8448 }
8449}
8450
8451impl ScalarUDFImpl for CypherDurationComponent {
8452 fn as_any(&self) -> &dyn Any {
8453 self
8454 }
8455
8456 fn name(&self) -> &'static str {
8457 "cypher_duration_component"
8458 }
8459
8460 fn signature(&self) -> &Signature {
8461 &self.signature
8462 }
8463
8464 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
8465 Ok(DataType::Int64)
8466 }
8467
8468 fn invoke_with_args(
8469 &self,
8470 args: ScalarFunctionArgs,
8471 ) -> datafusion::error::Result<ColumnarValue> {
8472 use crate::temporal::duration_component;
8473 use datafusion::arrow::array::{Array, Int64Array, StringArray, StructArray};
8474
8475 let rows = args.number_rows;
8476 let durs = args.args[0].to_array(rows)?;
8477 let names = args.args[1].to_array(rows)?;
8478 let d = durs.as_any().downcast_ref::<StructArray>();
8479 let n = names.as_any().downcast_ref::<StringArray>();
8480 let out: Int64Array = (0..rows)
8481 .map(|i| {
8482 let (d, n) = (d?, n?);
8483 if d.is_null(i) || n.is_null(i) {
8484 return None;
8485 }
8486 duration_component(&duration_struct_parts(d, i)?, n.value(i))
8487 })
8488 .collect();
8489 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
8490 }
8491}
8492
8493fn temporal_accessor_valid(dt: &DataType, name: &str) -> bool {
8503 use crate::temporal::{
8504 is_date_accessor, is_epoch_accessor, is_time_accessor, is_zone_int_accessor,
8505 is_zone_str_accessor,
8506 };
8507 match dt {
8508 DataType::Time64(_) => is_time_accessor(name),
8509 DataType::Struct(_) if is_time_struct(dt) => {
8510 is_time_accessor(name) || is_zone_int_accessor(name) || is_zone_str_accessor(name)
8511 }
8512 DataType::Struct(_) if is_localdatetime_struct(dt) => {
8513 is_date_accessor(name) || is_time_accessor(name)
8514 }
8515 DataType::Struct(_) if is_datetime_struct(dt) => {
8516 is_date_accessor(name)
8517 || is_time_accessor(name)
8518 || is_zone_int_accessor(name)
8519 || is_zone_str_accessor(name)
8520 || is_epoch_accessor(name)
8521 }
8522 _ => false,
8523 }
8524}
8525
8526static CYPHER_TEMPORAL_COMPONENT: LazyLock<ScalarUDF> =
8531 LazyLock::new(|| ScalarUDF::new_from_impl(CypherTemporalComponent::new()));
8532
8533#[derive(Debug, PartialEq, Eq, Hash)]
8534struct CypherTemporalComponent {
8535 signature: Signature,
8536}
8537
8538impl CypherTemporalComponent {
8539 fn new() -> Self {
8540 Self {
8541 signature: Signature::any(2, Volatility::Immutable),
8542 }
8543 }
8544}
8545
8546impl ScalarUDFImpl for CypherTemporalComponent {
8547 fn as_any(&self) -> &dyn Any {
8548 self
8549 }
8550
8551 fn name(&self) -> &'static str {
8552 "cypher_temporal_component"
8553 }
8554
8555 fn signature(&self) -> &Signature {
8556 &self.signature
8557 }
8558
8559 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
8560 Ok(DataType::Int64)
8561 }
8562
8563 fn invoke_with_args(
8564 &self,
8565 args: ScalarFunctionArgs,
8566 ) -> datafusion::error::Result<ColumnarValue> {
8567 use crate::temporal::{
8568 date_component, epoch_component, is_date_accessor, is_time_accessor,
8569 is_zone_int_accessor, time_component, zone_int_component,
8570 };
8571 use datafusion::arrow::array::{
8572 Array, Int64Array, StringArray, StructArray, Time64NanosecondArray,
8573 };
8574 use datafusion::arrow::datatypes::TimeUnit;
8575
8576 let rows = args.number_rows;
8577 let vals = args.args[0].to_array(rows)?;
8578 let names = args.args[1].to_array(rows)?;
8579 let n = names.as_any().downcast_ref::<StringArray>();
8580 let out: Int64Array = (0..rows)
8581 .map(|i| {
8582 let n = n?;
8583 if vals.is_null(i) || n.is_null(i) {
8584 return None;
8585 }
8586 let name = n.value(i);
8587 match vals.data_type() {
8588 DataType::Time64(TimeUnit::Nanosecond) => {
8589 let v = vals.as_any().downcast_ref::<Time64NanosecondArray>()?;
8590 time_component(v.value(i), name)
8591 }
8592 DataType::Struct(_) => {
8593 let s = vals.as_any().downcast_ref::<StructArray>()?;
8594 if is_time_struct(vals.data_type()) {
8595 let (nanos, offset) = time_struct_parts(s, i)?;
8596 if is_zone_int_accessor(name) {
8597 zone_int_component(offset, name)
8598 } else {
8599 time_component(nanos, name)
8600 }
8601 } else if is_localdatetime_struct(vals.data_type()) {
8602 let (days, nanos) = localdatetime_struct_parts(s, i)?;
8603 if is_date_accessor(name) {
8604 date_component(days, name)
8605 } else {
8606 time_component(nanos, name)
8607 }
8608 } else if is_datetime_struct(vals.data_type()) {
8609 let (days, nanos, offset, _) = datetime_struct_parts(s, i)?;
8610 if is_date_accessor(name) {
8611 date_component(days, name)
8612 } else if is_time_accessor(name) {
8613 time_component(nanos, name)
8614 } else if is_zone_int_accessor(name) {
8615 zone_int_component(offset, name)
8616 } else {
8617 epoch_component(days, nanos, offset, name)
8618 }
8619 } else {
8620 None
8621 }
8622 }
8623 _ => None,
8624 }
8625 })
8626 .collect();
8627 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
8628 }
8629}
8630
8631static CYPHER_TEMPORAL_ZONE_STR: LazyLock<ScalarUDF> =
8634 LazyLock::new(|| ScalarUDF::new_from_impl(CypherTemporalZoneStr::new()));
8635
8636#[derive(Debug, PartialEq, Eq, Hash)]
8637struct CypherTemporalZoneStr {
8638 signature: Signature,
8639}
8640
8641impl CypherTemporalZoneStr {
8642 fn new() -> Self {
8643 Self {
8644 signature: Signature::any(2, Volatility::Immutable),
8645 }
8646 }
8647}
8648
8649impl ScalarUDFImpl for CypherTemporalZoneStr {
8650 fn as_any(&self) -> &dyn Any {
8651 self
8652 }
8653
8654 fn name(&self) -> &'static str {
8655 "cypher_temporal_zone_str"
8656 }
8657
8658 fn signature(&self) -> &Signature {
8659 &self.signature
8660 }
8661
8662 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
8663 Ok(DataType::Utf8)
8664 }
8665
8666 fn invoke_with_args(
8667 &self,
8668 args: ScalarFunctionArgs,
8669 ) -> datafusion::error::Result<ColumnarValue> {
8670 use crate::temporal::zone_str_component;
8671 use datafusion::arrow::array::{Array, StringArray, StructArray};
8672
8673 let rows = args.number_rows;
8674 let vals = args.args[0].to_array(rows)?;
8675 let names = args.args[1].to_array(rows)?;
8676 let n = names.as_any().downcast_ref::<StringArray>();
8677 let out: StringArray = (0..rows)
8678 .map(|i| {
8679 let n = n?;
8680 if vals.is_null(i) || n.is_null(i) {
8681 return None;
8682 }
8683 let name = n.value(i);
8684 let s = vals.as_any().downcast_ref::<StructArray>()?;
8685 if is_time_struct(vals.data_type()) {
8686 let (_, offset) = time_struct_parts(s, i)?;
8687 zone_str_component(offset, None, name)
8688 } else if is_datetime_struct(vals.data_type()) {
8689 let (_, _, offset, zone) = datetime_struct_parts(s, i)?;
8690 zone_str_component(offset, zone.as_deref(), name)
8691 } else {
8692 None
8693 }
8694 })
8695 .collect();
8696 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
8697 }
8698}
8699
8700static CYPHER_DURATION_BETWEEN: LazyLock<ScalarUDF> =
8709 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDurationBetween::new()));
8710
8711#[derive(Debug, PartialEq, Eq, Hash)]
8712struct CypherDurationBetween {
8713 signature: Signature,
8714}
8715
8716impl CypherDurationBetween {
8717 fn new() -> Self {
8718 Self {
8719 signature: Signature::any(3, Volatility::Immutable),
8720 }
8721 }
8722}
8723
8724fn between_operand(
8727 arr: &datafusion::arrow::array::ArrayRef,
8728 i: usize,
8729) -> Option<crate::temporal::BetweenOperand> {
8730 use datafusion::arrow::array::{Array, StructArray, Time64NanosecondArray};
8731 use datafusion::arrow::datatypes::TimeUnit;
8732 if arr.is_null(i) {
8733 return None;
8734 }
8735 match arr.data_type() {
8736 DataType::Time64(TimeUnit::Nanosecond) => Some((
8737 None,
8738 arr.as_any()
8739 .downcast_ref::<Time64NanosecondArray>()?
8740 .value(i),
8741 None,
8742 None,
8743 )),
8744 DataType::Struct(_) => {
8745 let s = arr.as_any().downcast_ref::<StructArray>()?;
8746 if is_date_struct(arr.data_type()) {
8747 Some((Some(date_struct_value(s, i)?), 0, None, None))
8748 } else if is_datetime_struct(arr.data_type()) {
8749 let (days, nanos, offset, zone) = datetime_struct_parts(s, i)?;
8752 Some((Some(days), nanos, Some(offset), zone))
8753 } else if is_time_struct(arr.data_type()) {
8754 let (nanos, offset) = time_struct_parts(s, i)?;
8755 Some((None, nanos, Some(offset), None))
8756 } else {
8757 let (days, nanos) = localdatetime_struct_parts(s, i)?;
8758 Some((Some(days), nanos, None, None))
8759 }
8760 }
8761 _ => None,
8762 }
8763}
8764
8765impl ScalarUDFImpl for CypherDurationBetween {
8766 fn as_any(&self) -> &dyn Any {
8767 self
8768 }
8769
8770 fn name(&self) -> &'static str {
8771 "cypher_duration_between"
8772 }
8773
8774 fn signature(&self) -> &Signature {
8775 &self.signature
8776 }
8777
8778 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
8779 Ok(DataType::Struct(
8780 graphforge_storage::schemas::duration_struct_fields(),
8781 ))
8782 }
8783
8784 fn invoke_with_args(
8785 &self,
8786 args: ScalarFunctionArgs,
8787 ) -> datafusion::error::Result<ColumnarValue> {
8788 use crate::temporal::{BetweenMode, duration_between};
8789 use datafusion::arrow::array::{Array, StringArray};
8790 use datafusion::arrow::compute::cast;
8791 use datafusion::error::DataFusionError;
8792
8793 let rows = args.number_rows;
8794 let cols = udf_argument_arrays(&args)?;
8795 let mode_arr = cast(&cols[2], &DataType::Utf8).map_err(DataFusionError::from)?;
8796 let modes = mode_arr.as_any().downcast_ref::<StringArray>();
8797
8798 let parts: Vec<Option<crate::temporal::DurationValue>> = (0..rows)
8799 .map(|i| {
8800 let m = modes?;
8801 if m.is_null(i) {
8802 return None;
8803 }
8804 let mode = match m.value(i) {
8805 "duration.between" => BetweenMode::Between,
8806 "duration.inmonths" => BetweenMode::Months,
8807 "duration.indays" => BetweenMode::Days,
8808 "duration.inseconds" => BetweenMode::Seconds,
8809 _ => return None,
8810 };
8811 let a = between_operand(&cols[0], i)?;
8812 let b = between_operand(&cols[1], i)?;
8813 duration_between(&a, &b, mode)
8814 })
8815 .collect();
8816 Ok(ColumnarValue::Array(std::sync::Arc::new(
8817 build_duration_struct(&parts),
8818 )))
8819 }
8820}
8821
8822static CYPHER_TEMPORAL_ARITH: LazyLock<ScalarUDF> =
8833 LazyLock::new(|| ScalarUDF::new_from_impl(CypherTemporalArith::new()));
8834
8835#[derive(Debug, PartialEq, Eq, Hash)]
8836struct CypherTemporalArith {
8837 signature: Signature,
8838}
8839
8840impl CypherTemporalArith {
8841 fn new() -> Self {
8842 Self {
8843 signature: Signature::any(3, Volatility::Immutable),
8844 }
8845 }
8846}
8847
8848impl ScalarUDFImpl for CypherTemporalArith {
8849 fn as_any(&self) -> &dyn Any {
8850 self
8851 }
8852
8853 fn name(&self) -> &'static str {
8854 "cypher_temporal_arith"
8855 }
8856
8857 fn signature(&self) -> &Signature {
8858 &self.signature
8859 }
8860
8861 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
8862 Ok(arg_types.first().cloned().unwrap_or(DataType::Null))
8864 }
8865
8866 #[allow(
8867 clippy::too_many_lines,
8868 reason = "one cohesive per-temporal-type dispatch (date/localtime/time/\
8869 localdatetime/datetime) applying a signed duration"
8870 )]
8871 fn invoke_with_args(
8872 &self,
8873 args: ScalarFunctionArgs,
8874 ) -> datafusion::error::Result<ColumnarValue> {
8875 use crate::temporal::{
8876 date_plus_duration, datetime_plus_duration, localtime_plus_duration,
8877 };
8878 use datafusion::arrow::array::{
8879 Array, ArrayRef, Int64Array, StructArray, Time64NanosecondArray,
8880 };
8881 use datafusion::arrow::datatypes::TimeUnit;
8882
8883 let rows = args.number_rows;
8884 let cols = udf_argument_arrays(&args)?;
8885 let temporal = &cols[0];
8886 let dur = cols[1].as_any().downcast_ref::<StructArray>();
8887 let signs = cols[2].as_any().downcast_ref::<Int64Array>();
8888
8889 let signed = |i: usize| -> Option<crate::temporal::DurationValue> {
8891 let (d, sg) = (dur?, signs?);
8892 if d.is_null(i) || sg.is_null(i) {
8893 return None;
8894 }
8895 let dv = duration_struct_parts(d, i)?;
8896 Some(if sg.value(i) < 0 {
8897 crate::temporal::DurationValue {
8898 months: -dv.months,
8899 days: -dv.days,
8900 seconds: -dv.seconds,
8901 nanos: -dv.nanos,
8902 }
8903 } else {
8904 dv
8905 })
8906 };
8907 let sub_day_nanos = |d: &crate::temporal::DurationValue| {
8912 d.seconds.rem_euclid(86_400) * 1_000_000_000 + d.nanos
8913 };
8914
8915 let out: ArrayRef = match temporal.data_type() {
8916 DataType::Struct(_) if is_date_struct(temporal.data_type()) => {
8917 let t = temporal.as_any().downcast_ref::<StructArray>();
8918 let days: Vec<Option<i64>> = (0..rows)
8919 .map(|i| {
8920 let t = t?;
8921 let d = date_struct_value(t, i)?;
8922 let dv = signed(i)?;
8923 Some(date_plus_duration(d, &dv))
8924 })
8925 .collect();
8926 std::sync::Arc::new(build_date_struct(&days))
8927 }
8928 DataType::Time64(TimeUnit::Nanosecond) => {
8929 let t = temporal.as_any().downcast_ref::<Time64NanosecondArray>();
8930 let a: Time64NanosecondArray = (0..rows)
8931 .map(|i| {
8932 let t = t?;
8933 if t.is_null(i) {
8934 return None;
8935 }
8936 let dv = signed(i)?;
8937 Some(localtime_plus_duration(t.value(i), sub_day_nanos(&dv)))
8938 })
8939 .collect();
8940 std::sync::Arc::new(a)
8941 }
8942 DataType::Struct(_) if is_time_struct(temporal.data_type()) => {
8943 let s = temporal.as_any().downcast_ref::<StructArray>();
8944 let parts: Vec<Option<(i64, i32)>> = (0..rows)
8945 .map(|i| {
8946 let s = s?;
8947 let (nanos, offset) = time_struct_parts(s, i)?;
8948 let dv = signed(i)?;
8949 Some((localtime_plus_duration(nanos, sub_day_nanos(&dv)), offset))
8950 })
8951 .collect();
8952 std::sync::Arc::new(build_time_struct(&parts))
8953 }
8954 DataType::Struct(_) if is_datetime_struct(temporal.data_type()) => {
8955 let s = temporal.as_any().downcast_ref::<StructArray>();
8956 let parts: Vec<DateTimeRow> = (0..rows)
8957 .map(|i| {
8958 let s = s?;
8959 let (days, nanos, offset, zone) = datetime_struct_parts(s, i)?;
8960 let dv = signed(i)?;
8961 let (date, no) = datetime_plus_duration(days, nanos, &dv);
8962 Some((date, no, offset, zone))
8963 })
8964 .collect();
8965 std::sync::Arc::new(build_datetime_struct(&parts))
8966 }
8967 DataType::Struct(_) if is_localdatetime_struct(temporal.data_type()) => {
8969 let s = temporal.as_any().downcast_ref::<StructArray>();
8970 let parts: Vec<Option<(i64, i64)>> = (0..rows)
8971 .map(|i| {
8972 let s = s?;
8973 let (days, nanos) = localdatetime_struct_parts(s, i)?;
8974 let dv = signed(i)?;
8975 let (date, no) = datetime_plus_duration(days, nanos, &dv);
8976 Some((date, no))
8977 })
8978 .collect();
8979 std::sync::Arc::new(build_localdatetime_struct(&parts))
8980 }
8981 other => {
8982 return Err(datafusion::error::DataFusionError::Internal(format!(
8983 "cypher_temporal_arith: left operand is not a temporal value ({other:?})"
8984 )));
8985 }
8986 };
8987 Ok(ColumnarValue::Array(out))
8988 }
8989}
8990
8991static CYPHER_DURATION_PARSE: LazyLock<ScalarUDF> =
9000 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDurationParse::new()));
9001
9002#[derive(Debug, PartialEq, Eq, Hash)]
9003struct CypherDurationParse {
9004 signature: Signature,
9005}
9006
9007impl CypherDurationParse {
9008 fn new() -> Self {
9009 Self {
9010 signature: Signature::any(1, Volatility::Immutable),
9011 }
9012 }
9013}
9014
9015impl ScalarUDFImpl for CypherDurationParse {
9016 fn as_any(&self) -> &dyn Any {
9017 self
9018 }
9019 fn name(&self) -> &'static str {
9020 "cypher_duration_parse"
9021 }
9022 fn signature(&self) -> &Signature {
9023 &self.signature
9024 }
9025 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
9026 Ok(DataType::Struct(
9027 graphforge_storage::schemas::duration_struct_fields(),
9028 ))
9029 }
9030
9031 fn invoke_with_args(
9032 &self,
9033 args: ScalarFunctionArgs,
9034 ) -> datafusion::error::Result<ColumnarValue> {
9035 use datafusion::arrow::array::{Array, StringArray};
9036 use datafusion::arrow::compute::cast;
9037
9038 let rows = args.number_rows;
9039 let arr = cast(&args.args[0].to_array(rows)?, &DataType::Utf8)?;
9040 let s = arr.as_any().downcast_ref::<StringArray>();
9041 let parts: Vec<Option<crate::temporal::DurationValue>> = (0..rows)
9042 .map(|i| {
9043 let s = s?;
9044 if s.is_null(i) {
9045 return Option::None;
9046 }
9047 crate::temporal::duration_value_from_str(s.value(i))
9048 })
9049 .collect();
9050 Ok(ColumnarValue::Array(std::sync::Arc::new(
9051 build_duration_struct(&parts),
9052 )))
9053 }
9054}
9055
9056static CYPHER_DURATION_ADD: LazyLock<ScalarUDF> =
9059 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDurationAdd::new()));
9060
9061#[derive(Debug, PartialEq, Eq, Hash)]
9062struct CypherDurationAdd {
9063 signature: Signature,
9064}
9065
9066impl CypherDurationAdd {
9067 fn new() -> Self {
9068 Self {
9069 signature: Signature::any(3, Volatility::Immutable),
9070 }
9071 }
9072}
9073
9074impl ScalarUDFImpl for CypherDurationAdd {
9075 fn as_any(&self) -> &dyn Any {
9076 self
9077 }
9078
9079 fn name(&self) -> &'static str {
9080 "cypher_duration_add"
9081 }
9082
9083 fn signature(&self) -> &Signature {
9084 &self.signature
9085 }
9086
9087 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
9088 Ok(DataType::Struct(
9089 graphforge_storage::schemas::duration_struct_fields(),
9090 ))
9091 }
9092
9093 fn invoke_with_args(
9094 &self,
9095 args: ScalarFunctionArgs,
9096 ) -> datafusion::error::Result<ColumnarValue> {
9097 use datafusion::arrow::array::{Array, Int64Array, StructArray};
9098
9099 let rows = args.number_rows;
9100 let cols = udf_argument_arrays(&args)?;
9101 let a = cols[0].as_any().downcast_ref::<StructArray>();
9102 let b = cols[1].as_any().downcast_ref::<StructArray>();
9103 let signs = cols[2].as_any().downcast_ref::<Int64Array>();
9104
9105 let parts: Vec<Option<crate::temporal::DurationValue>> = (0..rows)
9106 .map(|i| {
9107 let (a, b, sg) = (a?, b?, signs?);
9108 let av = duration_struct_parts(a, i)?;
9109 let bv = duration_struct_parts(b, i)?;
9110 let s: i64 = if sg.is_null(i) || sg.value(i) >= 0 {
9111 1
9112 } else {
9113 -1
9114 };
9115 let nanos_sum = av.nanos + s * bv.nanos;
9121 let seconds = av.seconds + s * bv.seconds + nanos_sum.div_euclid(1_000_000_000);
9122 Some(crate::temporal::DurationValue {
9123 months: av.months + s * bv.months,
9124 days: av.days + s * bv.days,
9125 seconds,
9126 nanos: nanos_sum.rem_euclid(1_000_000_000),
9127 })
9128 })
9129 .collect();
9130 Ok(ColumnarValue::Array(std::sync::Arc::new(
9131 build_duration_struct(&parts),
9132 )))
9133 }
9134}
9135
9136static CYPHER_DURATION_SCALE: LazyLock<ScalarUDF> =
9140 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDurationScale::new()));
9141
9142#[derive(Debug, PartialEq, Eq, Hash)]
9143struct CypherDurationScale {
9144 signature: Signature,
9145}
9146
9147impl CypherDurationScale {
9148 fn new() -> Self {
9149 Self {
9150 signature: Signature::any(3, Volatility::Immutable),
9151 }
9152 }
9153}
9154
9155impl ScalarUDFImpl for CypherDurationScale {
9156 fn as_any(&self) -> &dyn Any {
9157 self
9158 }
9159 fn name(&self) -> &'static str {
9160 "cypher_duration_scale"
9161 }
9162 fn signature(&self) -> &Signature {
9163 &self.signature
9164 }
9165 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
9166 Ok(DataType::Struct(
9167 graphforge_storage::schemas::duration_struct_fields(),
9168 ))
9169 }
9170
9171 fn invoke_with_args(
9172 &self,
9173 args: ScalarFunctionArgs,
9174 ) -> datafusion::error::Result<ColumnarValue> {
9175 use datafusion::arrow::array::{Array, BooleanArray, Float64Array, StructArray};
9176 use datafusion::arrow::compute::cast;
9177
9178 let rows = args.number_rows;
9179 let cols = udf_argument_arrays(&args)?;
9180 let dur = cols[0].as_any().downcast_ref::<StructArray>();
9181 let num = cast(&cols[1], &DataType::Float64)?;
9183 let num = num.as_any().downcast_ref::<Float64Array>();
9184 let is_div = cols[2].as_any().downcast_ref::<BooleanArray>();
9185
9186 let parts: Vec<Option<crate::temporal::DurationValue>> = (0..rows)
9187 .map(|i| {
9188 let (dur, num, is_div) = (dur?, num?, is_div?);
9189 if num.is_null(i) {
9190 return Option::None; }
9192 let dv = duration_struct_parts(dur, i)?;
9193 let divide = !is_div.is_null(i) && is_div.value(i);
9194 Some(crate::temporal::scale_duration(&dv, num.value(i), divide))
9195 })
9196 .collect();
9197 Ok(ColumnarValue::Array(std::sync::Arc::new(
9198 build_duration_struct(&parts),
9199 )))
9200 }
9201}
9202
9203#[allow(
9211 clippy::match_same_arms,
9212 reason = "per-quantifier arms read clearest grouped by kind, even where two \
9213 fallback bodies coincide (all/none both default true)"
9214)]
9215fn reduce_quantifier(
9216 kind: graphforge_ir::QuantifierKind,
9217 bools: &datafusion::arrow::array::BooleanArray,
9218 n: usize,
9219) -> Option<bool> {
9220 use datafusion::arrow::array::Array;
9221 use graphforge_ir::QuantifierKind as Q;
9222 let (mut any_true, mut any_false, mut any_null, mut count_true) = (false, false, false, 0u32);
9223 for j in 0..n {
9224 if bools.is_null(j) {
9225 any_null = true;
9226 } else if bools.value(j) {
9227 any_true = true;
9228 count_true += 1;
9229 } else {
9230 any_false = true;
9231 }
9232 }
9233 match kind {
9236 Q::All if any_false => Some(false),
9237 Q::All => (!any_null).then_some(true),
9238 Q::Any if any_true => Some(true),
9239 Q::Any => (!any_null).then_some(false),
9240 Q::None if any_true => Some(false),
9241 Q::None => (!any_null).then_some(true),
9242 Q::Single if count_true > 1 => Some(false),
9243 Q::Single => (!any_null).then_some(count_true == 1),
9244 }
9245}
9246
9247#[derive(Debug, PartialEq, Eq, Hash)]
9252struct CypherQuantifier {
9253 kind: graphforge_ir::QuantifierKind,
9254 predicate: DfExpr,
9255 elem_name: String,
9256 outer_names: Vec<String>,
9257 signature: Signature,
9258}
9259
9260impl CypherQuantifier {
9261 fn new(
9262 kind: graphforge_ir::QuantifierKind,
9263 predicate: DfExpr,
9264 elem_name: String,
9265 outer_names: Vec<String>,
9266 ) -> Self {
9267 let arity = 1 + outer_names.len();
9268 Self {
9269 kind,
9270 predicate,
9271 elem_name,
9272 outer_names,
9273 signature: Signature::any(arity, Volatility::Volatile),
9276 }
9277 }
9278}
9279
9280#[derive(Debug, PartialEq, Eq, Hash)]
9284struct CypherInvariantQuantifier {
9285 kind: graphforge_ir::QuantifierKind,
9286 predicate: Option<bool>,
9287 signature: Signature,
9288}
9289
9290#[cfg(test)]
9291static INVARIANT_QUANTIFIER_ROWS: std::sync::atomic::AtomicUsize =
9292 std::sync::atomic::AtomicUsize::new(0);
9293
9294impl CypherInvariantQuantifier {
9295 fn new(kind: graphforge_ir::QuantifierKind, predicate: Option<bool>) -> Self {
9296 Self {
9297 kind,
9298 predicate,
9299 signature: Signature::any(1, Volatility::Immutable),
9300 }
9301 }
9302}
9303
9304fn reduce_invariant_quantifier(
9305 kind: graphforge_ir::QuantifierKind,
9306 predicate: Option<bool>,
9307 len: usize,
9308) -> Option<bool> {
9309 use graphforge_ir::QuantifierKind as Q;
9310 if len == 0 {
9311 return Some(matches!(kind, Q::All | Q::None));
9312 }
9313 match (kind, predicate) {
9314 (_, None) => None,
9315 (Q::All | Q::Any, Some(value)) => Some(value),
9316 (Q::None, Some(value)) => Some(!value),
9317 (Q::Single, Some(true)) => Some(len == 1),
9318 (Q::Single, Some(false)) => Some(false),
9319 }
9320}
9321
9322impl ScalarUDFImpl for CypherInvariantQuantifier {
9323 fn as_any(&self) -> &dyn Any {
9324 self
9325 }
9326 fn name(&self) -> &'static str {
9327 "cypher_invariant_quantifier"
9328 }
9329 fn signature(&self) -> &Signature {
9330 &self.signature
9331 }
9332 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
9333 Ok(DataType::Boolean)
9334 }
9335 fn invoke_with_args(
9336 &self,
9337 args: ScalarFunctionArgs,
9338 ) -> datafusion::error::Result<ColumnarValue> {
9339 use datafusion::arrow::array::{Array, BooleanArray, ListArray};
9340 use datafusion::error::DataFusionError;
9341
9342 let rows = args.number_rows;
9343 #[cfg(test)]
9344 INVARIANT_QUANTIFIER_ROWS.fetch_add(rows, std::sync::atomic::Ordering::SeqCst);
9345 let list = args.args[0].to_array(rows)?;
9346 let list = list.as_any().downcast_ref::<ListArray>().ok_or_else(|| {
9347 DataFusionError::Internal("cypher_invariant_quantifier: argument is not a list".into())
9348 })?;
9349 let values = (0..rows).map(|row| {
9350 if list.is_null(row) {
9351 None
9352 } else {
9353 reduce_invariant_quantifier(self.kind, self.predicate, list.value(row).len())
9354 }
9355 });
9356 Ok(ColumnarValue::Array(Arc::new(
9357 values.collect::<BooleanArray>(),
9358 )))
9359 }
9360}
9361
9362impl ScalarUDFImpl for CypherQuantifier {
9363 fn as_any(&self) -> &dyn Any {
9364 self
9365 }
9366 fn name(&self) -> &'static str {
9367 "cypher_quantifier"
9368 }
9369 fn signature(&self) -> &Signature {
9370 &self.signature
9371 }
9372 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
9373 Ok(DataType::Boolean)
9374 }
9375
9376 fn invoke_with_args(
9377 &self,
9378 args: ScalarFunctionArgs,
9379 ) -> datafusion::error::Result<ColumnarValue> {
9380 use datafusion::arrow::array::{Array, ArrayRef, BooleanArray, ListArray, RecordBatch};
9381 use datafusion::arrow::datatypes::{Field, Schema};
9382 use datafusion::common::DFSchema;
9383 use datafusion::error::DataFusionError;
9384 use datafusion::logical_expr::execution_props::ExecutionProps;
9385 use datafusion::physical_expr::create_physical_expr;
9386 use std::sync::Arc;
9387
9388 let rows = args.number_rows;
9389 let cols: Vec<ArrayRef> = args
9390 .args
9391 .iter()
9392 .map(|a| a.to_array(rows))
9393 .collect::<datafusion::error::Result<_>>()?;
9394 let list = cols[0]
9395 .as_any()
9396 .downcast_ref::<ListArray>()
9397 .ok_or_else(|| {
9398 DataFusionError::Internal("cypher_quantifier: first argument is not a list".into())
9399 })?;
9400 let elem_type = match list.data_type() {
9401 DataType::List(f) | DataType::LargeList(f) => f.data_type().clone(),
9402 _ => DataType::Null,
9403 };
9404
9405 let mut fields = vec![Field::new(&self.elem_name, elem_type, true)];
9407 for (i, name) in self.outer_names.iter().enumerate() {
9408 fields.push(Field::new(name, cols[i + 1].data_type().clone(), true));
9409 }
9410 let schema = Arc::new(Schema::new(fields));
9411 let df_schema = DFSchema::try_from(schema.as_ref().clone())?;
9412 let phys = create_physical_expr(&self.predicate, &df_schema, &ExecutionProps::new());
9417
9418 let mut out = BooleanArray::builder(rows);
9419 for row in 0..rows {
9420 if list.is_null(row) {
9421 out.append_null(); continue;
9423 }
9424 let elems = list.value(row);
9425 let n = elems.len();
9426 if n == 0 {
9427 out.append_option(reduce_quantifier(self.kind, &BooleanArray::new_null(0), 0));
9431 continue;
9432 }
9433 let phys = phys
9434 .as_ref()
9435 .map_err(|e| DataFusionError::Execution(e.to_string()))?;
9436 let verdict = (|| {
9437 let mut batch_cols: Vec<ArrayRef> = Vec::with_capacity(1 + self.outer_names.len());
9438 batch_cols.push(elems);
9439 for i in 0..self.outer_names.len() {
9440 let sv = ScalarValue::try_from_array(&cols[i + 1], row).ok()?;
9441 batch_cols.push(sv.to_array_of_size(n).ok()?);
9442 }
9443 let batch = RecordBatch::try_new(Arc::clone(&schema), batch_cols).ok()?;
9444 let evaluated = phys.evaluate(&batch).ok()?.into_array(n).ok()?;
9445 if evaluated.data_type() == &DataType::Null {
9448 return reduce_quantifier(self.kind, &BooleanArray::new_null(n), n);
9449 }
9450 let bools = evaluated.as_any().downcast_ref::<BooleanArray>()?;
9451 reduce_quantifier(self.kind, bools, n)
9452 })();
9453 out.append_option(verdict);
9454 }
9455 Ok(ColumnarValue::Array(std::sync::Arc::new(out.finish())))
9456 }
9457}
9458
9459#[derive(Debug, PartialEq, Eq, Hash)]
9467struct CypherListComp {
9468 filter: Option<DfExpr>,
9469 projection: Option<DfExpr>,
9470 elem_name: String,
9471 outer_names: Vec<String>,
9472 signature: Signature,
9473}
9474
9475impl CypherListComp {
9476 fn new(
9477 filter: Option<DfExpr>,
9478 projection: Option<DfExpr>,
9479 elem_name: String,
9480 outer_names: Vec<String>,
9481 ) -> Self {
9482 let arity = 1 + outer_names.len();
9483 Self {
9484 filter,
9485 projection,
9486 elem_name,
9487 outer_names,
9488 signature: Signature::any(arity, Volatility::Volatile),
9491 }
9492 }
9493
9494 fn item_type(
9499 &self,
9500 elem_type: &DataType,
9501 outer_types: &[DataType],
9502 ) -> datafusion::error::Result<DataType> {
9503 use datafusion::arrow::datatypes::{Field, Schema};
9504 use datafusion::common::DFSchema;
9505 use datafusion::logical_expr::execution_props::ExecutionProps;
9506 use datafusion::physical_expr::create_physical_expr;
9507 let Some(proj) = &self.projection else {
9508 return Ok(elem_type.clone());
9509 };
9510 let mut fields = vec![Field::new(&self.elem_name, elem_type.clone(), true)];
9511 for (name, dt) in self.outer_names.iter().zip(outer_types) {
9512 fields.push(Field::new(name, dt.clone(), true));
9513 }
9514 let schema = Schema::new(fields);
9515 let df_schema = DFSchema::try_from(schema.clone())?;
9516 let phys = create_physical_expr(proj, &df_schema, &ExecutionProps::new())?;
9517 phys.data_type(&schema)
9518 }
9519
9520 #[allow(
9521 clippy::too_many_lines,
9522 reason = "one flatten/filter/project/reassemble pass keeps volatile evaluation and row offsets aligned"
9523 )]
9524 fn invoke_uncorrelated(
9525 list: &datafusion::arrow::array::ListArray,
9526 schema: datafusion::arrow::datatypes::SchemaRef,
9527 filter_phys: Option<&Arc<dyn datafusion::physical_expr::PhysicalExpr>>,
9528 proj_phys: Option<&Arc<dyn datafusion::physical_expr::PhysicalExpr>>,
9529 item_type: &DataType,
9530 ) -> datafusion::error::Result<ColumnarValue> {
9531 use datafusion::arrow::array::{
9532 Array, ArrayRef, BooleanArray, ListArray, UInt32Array, new_empty_array,
9533 };
9534 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
9535 use datafusion::arrow::compute::{cast, filter_record_batch, take};
9536 use datafusion::arrow::datatypes::Field;
9537 use datafusion::arrow::record_batch::RecordBatch;
9538 use datafusion::error::DataFusionError;
9539
9540 let rows = list.len();
9541 let mut validity = Vec::with_capacity(rows);
9542 let mut lengths = Vec::with_capacity(rows);
9543 let mut indices = Vec::new();
9544 let offsets = list.value_offsets();
9545 for row in 0..rows {
9546 let valid = list.is_valid(row);
9547 validity.push(valid);
9548 let length = if valid {
9549 usize::try_from(offsets[row + 1] - offsets[row]).map_err(|_| {
9550 DataFusionError::Internal("negative list-comprehension length".into())
9551 })?
9552 } else {
9553 0
9554 };
9555 lengths.push(length);
9556 if valid {
9557 for index in offsets[row]..offsets[row + 1] {
9558 indices.push(u32::try_from(index).map_err(|_| {
9559 DataFusionError::Internal(
9560 "list-comprehension element index exceeds u32::MAX".into(),
9561 )
9562 })?);
9563 }
9564 }
9565 }
9566
9567 let flat: ArrayRef = if indices.len() == list.values().len()
9568 && indices.first().is_none_or(|first| *first == 0)
9569 {
9570 Arc::clone(list.values())
9571 } else {
9572 take(list.values(), &UInt32Array::from(indices), None)?
9573 };
9574 let total = flat.len();
9575 let batch = RecordBatch::try_new(schema, vec![flat])?;
9576 let mask = if let Some(filter) = filter_phys
9577 && total > 0
9578 {
9579 let evaluated = filter.evaluate(&batch)?.into_array(total)?;
9580 let evaluated = evaluated
9581 .as_any()
9582 .downcast_ref::<BooleanArray>()
9583 .ok_or_else(|| {
9584 DataFusionError::Internal(
9585 "cypher_list_comprehension: filter did not evaluate to boolean".into(),
9586 )
9587 })?;
9588 Some(
9589 (0..total)
9590 .map(|index| evaluated.is_valid(index) && evaluated.value(index))
9591 .collect::<BooleanArray>(),
9592 )
9593 } else {
9594 None
9595 };
9596 let kept = if let Some(mask) = &mask {
9597 filter_record_batch(&batch, mask)?
9598 } else {
9599 batch
9600 };
9601 let kept_rows = kept.num_rows();
9602 let projected = if let Some(projection) = proj_phys {
9603 if kept_rows == 0 {
9604 new_empty_array(item_type)
9605 } else {
9606 projection.evaluate(&kept)?.into_array(kept_rows)?
9607 }
9608 } else {
9609 Arc::clone(kept.column(0))
9610 };
9611 let projected = if projected.data_type() == item_type {
9612 projected
9613 } else {
9614 cast(&projected, item_type)?
9615 };
9616
9617 let mut output_offsets = Vec::with_capacity(rows + 1);
9618 output_offsets.push(0i32);
9619 let mut input_offset = 0usize;
9620 let mut output_offset = 0i32;
9621 for length in lengths {
9622 let kept = mask.as_ref().map_or(length, |mask| {
9623 (input_offset..input_offset + length)
9624 .filter(|index| mask.value(*index))
9625 .count()
9626 });
9627 input_offset += length;
9628 let kept = i32::try_from(kept).map_err(|_| {
9629 DataFusionError::Internal(
9630 "cypher_list_comprehension: list length exceeds i32::MAX".into(),
9631 )
9632 })?;
9633 output_offset = output_offset.checked_add(kept).ok_or_else(|| {
9634 DataFusionError::Internal(
9635 "cypher_list_comprehension: total list length exceeds i32::MAX".into(),
9636 )
9637 })?;
9638 output_offsets.push(output_offset);
9639 }
9640 let list = ListArray::try_new(
9641 Arc::new(Field::new("item", item_type.clone(), true)),
9642 OffsetBuffer::new(ScalarBuffer::from(output_offsets)),
9643 projected,
9644 Some(NullBuffer::from(validity)),
9645 )?;
9646 Ok(ColumnarValue::Array(Arc::new(list)))
9647 }
9648}
9649
9650impl ScalarUDFImpl for CypherListComp {
9651 fn as_any(&self) -> &dyn Any {
9652 self
9653 }
9654 fn name(&self) -> &'static str {
9655 "cypher_list_comprehension"
9656 }
9657 fn signature(&self) -> &Signature {
9658 &self.signature
9659 }
9660 fn return_type(&self, arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
9661 use datafusion::arrow::datatypes::Field;
9662 let elem_type = match arg_types.first() {
9667 Some(DataType::List(f)) => f.data_type().clone(),
9668 _ => DataType::Null,
9669 };
9670 let item = self.item_type(&elem_type, arg_types.get(1..).unwrap_or(&[]))?;
9671 Ok(DataType::List(std::sync::Arc::new(Field::new(
9672 "item", item, true,
9673 ))))
9674 }
9675
9676 #[allow(
9677 clippy::too_many_lines,
9678 reason = "the per-row filter/project/reassemble loop reads clearest inline"
9679 )]
9680 fn invoke_with_args(
9681 &self,
9682 args: ScalarFunctionArgs,
9683 ) -> datafusion::error::Result<ColumnarValue> {
9684 use datafusion::arrow::array::{
9685 Array, ArrayRef, BooleanArray, ListArray, RecordBatch, new_empty_array,
9686 };
9687 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
9688 use datafusion::arrow::compute::{cast, concat, filter_record_batch};
9689 use datafusion::arrow::datatypes::{Field, Schema};
9690 use datafusion::common::DFSchema;
9691 use datafusion::error::DataFusionError;
9692 use datafusion::logical_expr::execution_props::ExecutionProps;
9693 use datafusion::physical_expr::create_physical_expr;
9694 use std::sync::Arc;
9695
9696 let rows = args.number_rows;
9697 let cols: Vec<ArrayRef> = args
9698 .args
9699 .iter()
9700 .map(|a| a.to_array(rows))
9701 .collect::<datafusion::error::Result<_>>()?;
9702 let list = cols[0]
9703 .as_any()
9704 .downcast_ref::<ListArray>()
9705 .ok_or_else(|| {
9706 DataFusionError::Internal(
9707 "cypher_list_comprehension: first argument is not a list".into(),
9708 )
9709 })?;
9710 let elem_type = match list.data_type() {
9711 DataType::List(f) => f.data_type().clone(),
9712 _ => DataType::Null,
9713 };
9714 let outer_types: Vec<DataType> = (0..self.outer_names.len())
9715 .map(|i| cols[i + 1].data_type().clone())
9716 .collect();
9717 let item_type = self.item_type(&elem_type, &outer_types)?;
9718
9719 let mut fields = vec![Field::new(&self.elem_name, elem_type, true)];
9721 for (i, name) in self.outer_names.iter().enumerate() {
9722 fields.push(Field::new(name, cols[i + 1].data_type().clone(), true));
9723 }
9724 let schema = Arc::new(Schema::new(fields));
9725 let df_schema = DFSchema::try_from(schema.as_ref().clone())?;
9726 let props = ExecutionProps::new();
9727 let filter_phys = self
9728 .filter
9729 .as_ref()
9730 .map(|f| create_physical_expr(f, &df_schema, &props))
9731 .transpose()?;
9732 let proj_phys = self
9733 .projection
9734 .as_ref()
9735 .map(|p| create_physical_expr(p, &df_schema, &props))
9736 .transpose()?;
9737
9738 if self.outer_names.is_empty() {
9739 return Self::invoke_uncorrelated(
9740 list,
9741 schema,
9742 filter_phys.as_ref(),
9743 proj_phys.as_ref(),
9744 &item_type,
9745 );
9746 }
9747
9748 let mut pieces: Vec<ArrayRef> = Vec::new();
9749 let mut offsets: Vec<i32> = Vec::with_capacity(rows + 1);
9750 offsets.push(0);
9751 let mut validity: Vec<bool> = Vec::with_capacity(rows);
9752 let mut cur: i32 = 0;
9753
9754 for row in 0..rows {
9755 if list.is_null(row) {
9756 validity.push(false);
9757 offsets.push(cur);
9758 continue;
9759 }
9760 validity.push(true);
9761 let elems = list.value(row);
9762 let n = elems.len();
9763 let mut batch_cols: Vec<ArrayRef> = Vec::with_capacity(1 + self.outer_names.len());
9764 batch_cols.push(elems);
9765 for i in 0..self.outer_names.len() {
9766 let sv = ScalarValue::try_from_array(&cols[i + 1], row)?;
9767 batch_cols.push(sv.to_array_of_size(n)?);
9768 }
9769 let batch = RecordBatch::try_new(Arc::clone(&schema), batch_cols)?;
9770
9771 let kept = if let Some(fp) = &filter_phys {
9773 let mask = fp.evaluate(&batch)?.into_array(n)?;
9774 let mask = mask
9775 .as_any()
9776 .downcast_ref::<BooleanArray>()
9777 .ok_or_else(|| {
9778 DataFusionError::Internal(
9779 "cypher_list_comprehension: filter did not evaluate to boolean".into(),
9780 )
9781 })?;
9782 let clean: BooleanArray =
9783 (0..n).map(|j| mask.is_valid(j) && mask.value(j)).collect();
9784 filter_record_batch(&batch, &clean)?
9785 } else {
9786 batch
9787 };
9788
9789 let projected: ArrayRef = if let Some(pp) = &proj_phys {
9791 let m = kept.num_rows();
9792 if m == 0 {
9793 new_empty_array(&item_type)
9794 } else {
9795 pp.evaluate(&kept)?.into_array(m)?
9796 }
9797 } else {
9798 Arc::clone(kept.column(0))
9799 };
9800 let projected = if projected.data_type() == &item_type {
9801 projected
9802 } else {
9803 cast(&projected, &item_type)?
9804 };
9805
9806 let len = i32::try_from(projected.len()).map_err(|_| {
9807 DataFusionError::Internal("cypher_list_comprehension: list too long".into())
9808 })?;
9809 cur = cur.checked_add(len).ok_or_else(|| {
9810 DataFusionError::Internal(
9811 "cypher_list_comprehension: total list length exceeds i32::MAX".into(),
9812 )
9813 })?;
9814 offsets.push(cur);
9815 pieces.push(projected);
9816 }
9817
9818 let child: ArrayRef = if pieces.is_empty() {
9819 new_empty_array(&item_type)
9820 } else {
9821 let refs: Vec<&dyn Array> = pieces.iter().map(AsRef::as_ref).collect();
9822 concat(&refs)?
9823 };
9824 let field = Arc::new(Field::new("item", item_type, true));
9825 let list_arr = ListArray::try_new(
9826 field,
9827 OffsetBuffer::new(ScalarBuffer::from(offsets)),
9828 child,
9829 Some(NullBuffer::from(validity)),
9830 )?;
9831 Ok(ColumnarValue::Array(Arc::new(list_arr)))
9832 }
9833}
9834
9835fn udf_argument_arrays(
9840 args: &ScalarFunctionArgs,
9841) -> datafusion::error::Result<Vec<datafusion::arrow::array::ArrayRef>> {
9842 args.args
9843 .iter()
9844 .map(|value| value.to_array(args.number_rows))
9845 .collect()
9846}
9847
9848fn cast_argument_arrays(
9849 arrays: &[datafusion::arrow::array::ArrayRef],
9850 data_type: &DataType,
9851) -> datafusion::error::Result<Vec<datafusion::arrow::array::ArrayRef>> {
9852 arrays
9853 .iter()
9854 .map(|array| datafusion::arrow::compute::cast(array, data_type).map_err(Into::into))
9855 .collect()
9856}
9857
9858static CYPHER_DATE_PROJECT: LazyLock<ScalarUDF> =
9864 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDateProject::new()));
9865
9866#[derive(Debug, PartialEq, Eq, Hash)]
9867struct CypherDateProject {
9868 signature: Signature,
9869}
9870
9871impl CypherDateProject {
9872 fn new() -> Self {
9873 Self {
9874 signature: Signature::any(9, Volatility::Immutable),
9875 }
9876 }
9877}
9878
9879impl ScalarUDFImpl for CypherDateProject {
9880 fn as_any(&self) -> &dyn Any {
9881 self
9882 }
9883
9884 fn name(&self) -> &'static str {
9885 "cypher_date_project"
9886 }
9887
9888 fn signature(&self) -> &Signature {
9889 &self.signature
9890 }
9891
9892 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
9893 Ok(DataType::Struct(
9894 graphforge_storage::schemas::date_struct_fields(),
9895 ))
9896 }
9897
9898 fn invoke_with_args(
9899 &self,
9900 args: ScalarFunctionArgs,
9901 ) -> datafusion::error::Result<ColumnarValue> {
9902 use crate::temporal::{DateOverrides, parse_date_or_datetime_prefix};
9903 use datafusion::arrow::array::{Array, StringArray, StructArray};
9904 use datafusion::arrow::compute::cast;
9905
9906 let rows = args.number_rows;
9907 let cols = udf_argument_arrays(&args)?;
9908 let base = if is_date_struct(cols[0].data_type())
9911 || is_localdatetime_struct(cols[0].data_type())
9912 || is_time_struct(cols[0].data_type())
9913 || is_datetime_struct(cols[0].data_type())
9914 {
9915 std::sync::Arc::clone(&cols[0])
9916 } else {
9917 cast(&cols[0], &DataType::Utf8).map_err(datafusion::error::DataFusionError::from)?
9918 };
9919 let ov: Vec<_> = cols[1..9]
9921 .iter()
9922 .map(|c| cast(c, &DataType::Int64))
9923 .collect::<std::result::Result<Vec<_>, _>>()
9924 .map_err(datafusion::error::DataFusionError::from)?;
9925
9926 let base_date = |i: usize| -> Option<i64> {
9927 if base.is_null(i) {
9928 return None;
9929 }
9930 match base.data_type() {
9931 DataType::Struct(_) => {
9932 let s = base.as_any().downcast_ref::<StructArray>()?;
9933 if is_date_struct(base.data_type()) {
9934 date_struct_value(s, i)
9935 } else {
9936 Some(localdatetime_struct_parts(s, i)?.0)
9938 }
9939 }
9940 DataType::Utf8 => parse_date_or_datetime_prefix(
9941 base.as_any().downcast_ref::<StringArray>()?.value(i),
9942 ),
9943 _ => None,
9944 }
9945 };
9946
9947 let out: Vec<Option<i64>> = (0..rows)
9948 .map(|i| {
9949 let overrides = DateOverrides {
9950 year: optional_i64_at(&ov[0], i),
9951 month: optional_i64_at(&ov[1], i),
9952 day: optional_i64_at(&ov[2], i),
9953 week: optional_i64_at(&ov[3], i),
9954 day_of_week: optional_i64_at(&ov[4], i),
9955 ordinal_day: optional_i64_at(&ov[5], i),
9956 quarter: optional_i64_at(&ov[6], i),
9957 day_of_quarter: optional_i64_at(&ov[7], i),
9958 };
9959 crate::temporal::project_date(base_date(i)?, &overrides)
9960 })
9961 .collect();
9962 Ok(ColumnarValue::Array(std::sync::Arc::new(
9963 build_date_struct(&out),
9964 )))
9965 }
9966}
9967
9968static CYPHER_LOCALTIME_PROJECT: LazyLock<ScalarUDF> =
9979 LazyLock::new(|| ScalarUDF::new_from_impl(CypherLocalTimeProject::new()));
9980
9981#[derive(Debug, PartialEq, Eq, Hash)]
9982struct CypherLocalTimeProject {
9983 signature: Signature,
9984}
9985
9986impl CypherLocalTimeProject {
9987 fn new() -> Self {
9988 Self {
9989 signature: Signature::any(7, Volatility::Immutable),
9990 }
9991 }
9992}
9993
9994impl ScalarUDFImpl for CypherLocalTimeProject {
9995 fn as_any(&self) -> &dyn Any {
9996 self
9997 }
9998
9999 fn name(&self) -> &'static str {
10000 "cypher_localtime_project"
10001 }
10002
10003 fn signature(&self) -> &Signature {
10004 &self.signature
10005 }
10006
10007 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
10008 Ok(DataType::Time64(
10009 datafusion::arrow::datatypes::TimeUnit::Nanosecond,
10010 ))
10011 }
10012
10013 fn invoke_with_args(
10014 &self,
10015 args: ScalarFunctionArgs,
10016 ) -> datafusion::error::Result<ColumnarValue> {
10017 use crate::temporal::{LocalTimeOverrides, project_localtime, time_of_day_nanos_any};
10018 use datafusion::arrow::array::{
10019 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
10020 };
10021 use datafusion::arrow::compute::cast;
10022 use datafusion::arrow::datatypes::TimeUnit;
10023 use datafusion::error::DataFusionError;
10024
10025 let rows = args.number_rows;
10026 let cols = udf_argument_arrays(&args)?;
10027 let base: ArrayRef =
10030 if matches!(cols[0].data_type(), DataType::Time64(TimeUnit::Nanosecond))
10031 || is_localdatetime_struct(cols[0].data_type())
10032 || is_time_struct(cols[0].data_type())
10033 || is_datetime_struct(cols[0].data_type())
10034 {
10035 std::sync::Arc::clone(&cols[0])
10036 } else {
10037 cast(&cols[0], &DataType::Utf8).map_err(DataFusionError::from)?
10038 };
10039 let ov = cast_argument_arrays(&cols[1..7], &DataType::Int64)?;
10040
10041 let base_nanos = |i: usize| -> Option<i64> {
10042 if base.is_null(i) {
10043 return None;
10044 }
10045 match base.data_type() {
10046 DataType::Time64(TimeUnit::Nanosecond) => Some(
10047 base.as_any()
10048 .downcast_ref::<Time64NanosecondArray>()?
10049 .value(i),
10050 ),
10051 DataType::Struct(_) => {
10053 let s = base.as_any().downcast_ref::<StructArray>()?;
10054 if is_time_struct(base.data_type()) {
10055 Some(time_struct_parts(s, i)?.0)
10056 } else {
10057 Some(localdatetime_struct_parts(s, i)?.1)
10058 }
10059 }
10060 DataType::Utf8 => {
10061 time_of_day_nanos_any(base.as_any().downcast_ref::<StringArray>()?.value(i))
10062 }
10063 _ => None,
10064 }
10065 };
10066
10067 let out: Time64NanosecondArray = (0..rows)
10068 .map(|i| {
10069 let overrides = LocalTimeOverrides {
10070 hour: optional_i64_at(&ov[0], i),
10071 minute: optional_i64_at(&ov[1], i),
10072 second: optional_i64_at(&ov[2], i),
10073 millisecond: optional_i64_at(&ov[3], i),
10074 microsecond: optional_i64_at(&ov[4], i),
10075 nanosecond: optional_i64_at(&ov[5], i),
10076 };
10077 project_localtime(base_nanos(i)?, &overrides)
10078 })
10079 .collect();
10080 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
10081 }
10082}
10083
10084static CYPHER_LOCALTIME_TRUNCATE: LazyLock<ScalarUDF> =
10094 LazyLock::new(|| ScalarUDF::new_from_impl(CypherLocalTimeTruncate::new()));
10095
10096#[derive(Debug, PartialEq, Eq, Hash)]
10097struct CypherLocalTimeTruncate {
10098 signature: Signature,
10099}
10100
10101impl CypherLocalTimeTruncate {
10102 fn new() -> Self {
10103 Self {
10104 signature: Signature::any(8, Volatility::Immutable),
10105 }
10106 }
10107}
10108
10109impl ScalarUDFImpl for CypherLocalTimeTruncate {
10110 fn as_any(&self) -> &dyn Any {
10111 self
10112 }
10113
10114 fn name(&self) -> &'static str {
10115 "cypher_localtime_truncate"
10116 }
10117
10118 fn signature(&self) -> &Signature {
10119 &self.signature
10120 }
10121
10122 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
10123 Ok(DataType::Time64(
10124 datafusion::arrow::datatypes::TimeUnit::Nanosecond,
10125 ))
10126 }
10127
10128 fn invoke_with_args(
10129 &self,
10130 args: ScalarFunctionArgs,
10131 ) -> datafusion::error::Result<ColumnarValue> {
10132 use crate::temporal::{
10133 LocalTimeOverrides, project_localtime, time_of_day_nanos_any, truncate_time_nanos,
10134 };
10135 use datafusion::arrow::array::{
10136 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
10137 };
10138 use datafusion::arrow::compute::cast;
10139 use datafusion::arrow::datatypes::TimeUnit;
10140 use datafusion::error::DataFusionError;
10141
10142 let rows = args.number_rows;
10143 let cols = udf_argument_arrays(&args)?;
10144 let base: ArrayRef =
10145 if matches!(cols[0].data_type(), DataType::Time64(TimeUnit::Nanosecond))
10146 || is_localdatetime_struct(cols[0].data_type())
10147 || is_time_struct(cols[0].data_type())
10148 || is_datetime_struct(cols[0].data_type())
10149 {
10150 std::sync::Arc::clone(&cols[0])
10151 } else {
10152 cast(&cols[0], &DataType::Utf8).map_err(DataFusionError::from)?
10153 };
10154 let units_arr = cast(&cols[1], &DataType::Utf8).map_err(DataFusionError::from)?;
10155 let units = units_arr.as_any().downcast_ref::<StringArray>();
10156 let ov = cast_argument_arrays(&cols[2..8], &DataType::Int64)?;
10157
10158 let base_nanos = |i: usize| -> Option<i64> {
10159 if base.is_null(i) {
10160 return None;
10161 }
10162 match base.data_type() {
10163 DataType::Time64(TimeUnit::Nanosecond) => Some(
10164 base.as_any()
10165 .downcast_ref::<Time64NanosecondArray>()?
10166 .value(i),
10167 ),
10168 DataType::Struct(_) => {
10169 let s = base.as_any().downcast_ref::<StructArray>()?;
10170 if is_time_struct(base.data_type()) {
10171 Some(time_struct_parts(s, i)?.0)
10172 } else {
10173 Some(localdatetime_struct_parts(s, i)?.1)
10174 }
10175 }
10176 DataType::Utf8 => {
10177 time_of_day_nanos_any(base.as_any().downcast_ref::<StringArray>()?.value(i))
10178 }
10179 _ => None,
10180 }
10181 };
10182
10183 let out: Time64NanosecondArray = (0..rows)
10184 .map(|i| {
10185 let u = units?;
10186 if u.is_null(i) {
10187 return None;
10188 }
10189 let truncated = truncate_time_nanos(base_nanos(i)?, u.value(i))?;
10190 let overrides = LocalTimeOverrides {
10191 hour: optional_i64_at(&ov[0], i),
10192 minute: optional_i64_at(&ov[1], i),
10193 second: optional_i64_at(&ov[2], i),
10194 millisecond: optional_i64_at(&ov[3], i),
10195 microsecond: optional_i64_at(&ov[4], i),
10196 nanosecond: optional_i64_at(&ov[5], i),
10197 };
10198 project_localtime(truncated, &overrides)
10199 })
10200 .collect();
10201 Ok(ColumnarValue::Array(std::sync::Arc::new(out)))
10202 }
10203}
10204
10205fn date_fields() -> datafusion::arrow::datatypes::Fields {
10214 graphforge_storage::schemas::date_struct_fields()
10215}
10216
10217fn is_date_struct(dt: &DataType) -> bool {
10220 matches!(dt, DataType::Struct(fields)
10221 if fields.len() == 1
10222 && fields[0].name() == "epoch_day"
10223 && *fields[0].data_type() == DataType::Int64)
10224}
10225
10226fn build_date_struct(rows: &[Option<i64>]) -> datafusion::arrow::array::StructArray {
10229 use datafusion::arrow::array::Int64Array;
10230 use datafusion::arrow::buffer::NullBuffer;
10231 let days: Int64Array = rows.iter().copied().collect();
10232 let nulls = rows.iter().map(Option::is_some).collect::<NullBuffer>();
10233 datafusion::arrow::array::StructArray::new(
10234 date_fields(),
10235 vec![std::sync::Arc::new(days)],
10236 Some(nulls),
10237 )
10238}
10239
10240fn date_scalar(days: Option<i64>) -> ScalarValue {
10242 ScalarValue::Struct(std::sync::Arc::new(build_date_struct(&[days])))
10243}
10244
10245fn date_struct_value(arr: &datafusion::arrow::array::StructArray, i: usize) -> Option<i64> {
10247 use datafusion::arrow::array::{Array, Int64Array};
10248 if arr.is_null(i) {
10249 return None;
10250 }
10251 let days = arr.column(0).as_any().downcast_ref::<Int64Array>()?;
10252 days.is_valid(i).then(|| days.value(i))
10253}
10254
10255fn optional_i64_at(array: &datafusion::arrow::array::ArrayRef, row: usize) -> Option<i64> {
10259 use datafusion::arrow::array::{Array, Int64Array};
10260
10261 let values = array.as_any().downcast_ref::<Int64Array>()?;
10262 (!values.is_null(row)).then(|| values.value(row))
10263}
10264
10265fn localdatetime_fields() -> datafusion::arrow::datatypes::Fields {
10269 graphforge_storage::schemas::localdatetime_struct_fields()
10270}
10271
10272fn is_localdatetime_struct(dt: &DataType) -> bool {
10276 use datafusion::arrow::datatypes::TimeUnit;
10277 matches!(dt, DataType::Struct(fields)
10278 if fields.len() == 2
10279 && fields[0].name() == "date"
10280 && *fields[0].data_type() == DataType::Int64
10281 && fields[1].name() == "time"
10282 && *fields[1].data_type() == DataType::Time64(TimeUnit::Nanosecond))
10283}
10284
10285fn build_localdatetime_struct(
10288 rows: &[Option<(i64, i64)>],
10289) -> datafusion::arrow::array::StructArray {
10290 use datafusion::arrow::array::{Int64Array, Time64NanosecondArray};
10291 use datafusion::arrow::buffer::NullBuffer;
10292 let days: Int64Array = rows.iter().map(|r| r.map(|(d, _)| d)).collect();
10293 let nanos: Time64NanosecondArray = rows.iter().map(|r| r.map(|(_, n)| n)).collect();
10294 let nulls = rows.iter().map(Option::is_some).collect::<NullBuffer>();
10295 datafusion::arrow::array::StructArray::new(
10296 localdatetime_fields(),
10297 vec![std::sync::Arc::new(days), std::sync::Arc::new(nanos)],
10298 Some(nulls),
10299 )
10300}
10301
10302fn localdatetime_scalar(parts: Option<(i64, i64)>) -> ScalarValue {
10304 ScalarValue::Struct(std::sync::Arc::new(build_localdatetime_struct(&[parts])))
10305}
10306
10307fn is_duration_struct(dt: &DataType) -> bool {
10310 matches!(dt, DataType::Struct(fields)
10311 if fields.len() == 4
10312 && fields[0].name() == "months"
10313 && fields[1].name() == "days"
10314 && fields[2].name() == "seconds"
10315 && fields[3].name() == "nanos")
10316}
10317
10318fn is_temporal_clock_fn(name: &str) -> bool {
10322 matches!(
10324 name.to_ascii_lowercase().split_once('.'),
10325 Some((
10326 "date" | "localtime" | "time" | "localdatetime" | "datetime",
10327 "transaction" | "statement" | "realtime",
10328 ))
10329 )
10330}
10331
10332fn temporal_null_scalar(name: &str) -> ScalarValue {
10338 let lower = name.to_ascii_lowercase();
10339 let base = lower.split('.').next().unwrap_or(&lower);
10340 match base {
10341 "date" => date_scalar(None),
10342 "localtime" => ScalarValue::Time64Nanosecond(None),
10343 "time" => time_scalar(None),
10344 "localdatetime" => localdatetime_scalar(None),
10345 "datetime" => datetime_scalar(None),
10346 "duration" => duration_scalar(None),
10347 _ => ScalarValue::Null,
10348 }
10349}
10350
10351fn build_duration_struct(
10356 rows: &[Option<crate::temporal::DurationValue>],
10357) -> datafusion::arrow::array::StructArray {
10358 use datafusion::arrow::array::Int64Array;
10359 use datafusion::arrow::buffer::NullBuffer;
10360 let months: Int64Array = rows.iter().map(|r| r.map(|d| d.months)).collect();
10361 let days: Int64Array = rows.iter().map(|r| r.map(|d| d.days)).collect();
10362 let seconds: Int64Array = rows.iter().map(|r| r.map(|d| d.seconds)).collect();
10363 let nanos: Int64Array = rows.iter().map(|r| r.map(|d| d.nanos)).collect();
10364 let nulls = rows.iter().map(Option::is_some).collect::<NullBuffer>();
10365 datafusion::arrow::array::StructArray::new(
10366 graphforge_storage::schemas::duration_struct_fields(),
10367 vec![
10368 std::sync::Arc::new(months),
10369 std::sync::Arc::new(days),
10370 std::sync::Arc::new(seconds),
10371 std::sync::Arc::new(nanos),
10372 ],
10373 Some(nulls),
10374 )
10375}
10376
10377fn duration_struct_parts(
10380 arr: &datafusion::arrow::array::StructArray,
10381 i: usize,
10382) -> Option<crate::temporal::DurationValue> {
10383 use datafusion::arrow::array::{Array, Int64Array};
10384 if arr.is_null(i) {
10385 return None;
10386 }
10387 let col = |idx: usize| arr.column(idx).as_any().downcast_ref::<Int64Array>();
10388 Some(crate::temporal::DurationValue {
10389 months: col(0)?.value(i),
10390 days: col(1)?.value(i),
10391 seconds: col(2)?.value(i),
10392 nanos: col(3)?.value(i),
10393 })
10394}
10395
10396fn duration_scalar(parts: Option<crate::temporal::DurationValue>) -> ScalarValue {
10398 ScalarValue::Struct(std::sync::Arc::new(build_duration_struct(&[parts])))
10399}
10400
10401fn dur_secs_nanos(seconds: i64, nanos: i64) -> crate::temporal::DurationValue {
10404 crate::temporal::DurationValue {
10405 months: 0,
10406 days: 0,
10407 seconds,
10408 nanos,
10409 }
10410}
10411
10412fn duration_value_to_ir(d: crate::temporal::DurationValue) -> IrLiteral {
10414 IrLiteral::Duration {
10415 months: d.months,
10416 days: d.days,
10417 seconds: d.seconds,
10418 nanos: d.nanos,
10419 }
10420}
10421
10422fn localdatetime_struct_parts(
10428 arr: &datafusion::arrow::array::StructArray,
10429 i: usize,
10430) -> Option<(i64, i64)> {
10431 use datafusion::arrow::array::{Array, Int64Array, Time64NanosecondArray};
10432 if arr.is_null(i) {
10433 return None;
10434 }
10435 let d = arr.column(0).as_any().downcast_ref::<Int64Array>()?;
10436 let t = arr
10437 .column(1)
10438 .as_any()
10439 .downcast_ref::<Time64NanosecondArray>()?;
10440 (!d.is_null(i) && !t.is_null(i)).then(|| (d.value(i), t.value(i)))
10441}
10442
10443static CYPHER_LOCALDATETIME_PROJECT: LazyLock<ScalarUDF> =
10452 LazyLock::new(|| ScalarUDF::new_from_impl(CypherLocalDateTimeProject::new()));
10453
10454#[derive(Debug, PartialEq, Eq, Hash)]
10455struct CypherLocalDateTimeProject {
10456 signature: Signature,
10457}
10458
10459impl CypherLocalDateTimeProject {
10460 fn new() -> Self {
10461 Self {
10462 signature: Signature::any(16, Volatility::Immutable),
10463 }
10464 }
10465}
10466
10467impl ScalarUDFImpl for CypherLocalDateTimeProject {
10468 fn as_any(&self) -> &dyn Any {
10469 self
10470 }
10471
10472 fn name(&self) -> &'static str {
10473 "cypher_localdatetime_project"
10474 }
10475
10476 fn signature(&self) -> &Signature {
10477 &self.signature
10478 }
10479
10480 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
10481 Ok(DataType::Struct(localdatetime_fields()))
10482 }
10483
10484 #[allow(
10485 clippy::too_many_lines,
10486 reason = "one cohesive per-row projection: two typed base extractions \
10487 (date + time) plus 14 component overrides — splitting it would \
10488 scatter the row logic across helpers without aiding clarity"
10489 )]
10490 fn invoke_with_args(
10491 &self,
10492 args: ScalarFunctionArgs,
10493 ) -> datafusion::error::Result<ColumnarValue> {
10494 use crate::temporal::{
10495 DateOverrides, LocalTimeOverrides, parse_date_or_datetime_prefix, project_date,
10496 project_localtime, time_of_day_nanos_any,
10497 };
10498 use datafusion::arrow::array::{
10499 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
10500 };
10501 use datafusion::arrow::compute::cast;
10502 use datafusion::arrow::datatypes::TimeUnit;
10503 use datafusion::error::DataFusionError;
10504
10505 let rows = args.number_rows;
10506 let cols = udf_argument_arrays(&args)?;
10507 let typed_or_utf8 = |a: &ArrayRef| -> datafusion::error::Result<ArrayRef> {
10511 if matches!(a.data_type(), DataType::Time64(TimeUnit::Nanosecond))
10512 || is_date_struct(a.data_type())
10513 || is_localdatetime_struct(a.data_type())
10514 || is_time_struct(a.data_type())
10515 || is_datetime_struct(a.data_type())
10516 {
10517 Ok(std::sync::Arc::clone(a))
10518 } else {
10519 cast(a, &DataType::Utf8).map_err(DataFusionError::from)
10520 }
10521 };
10522 let date_src = typed_or_utf8(&cols[0])?;
10523 let time_src = typed_or_utf8(&cols[1])?;
10524 let ov = cast_argument_arrays(&cols[2..16], &DataType::Int64)?;
10525
10526 let base_date = |i: usize| -> Option<i64> {
10527 if date_src.is_null(i) {
10528 return Some(0); }
10530 match date_src.data_type() {
10531 DataType::Struct(_) => {
10532 let s = date_src.as_any().downcast_ref::<StructArray>()?;
10533 if is_date_struct(date_src.data_type()) {
10534 date_struct_value(s, i)
10535 } else {
10536 Some(localdatetime_struct_parts(s, i)?.0)
10538 }
10539 }
10540 DataType::Utf8 => parse_date_or_datetime_prefix(
10541 date_src.as_any().downcast_ref::<StringArray>()?.value(i),
10542 ),
10543 _ => None,
10544 }
10545 };
10546 let base_time = |i: usize| -> Option<i64> {
10547 if time_src.is_null(i) {
10548 return Some(0); }
10550 match time_src.data_type() {
10551 DataType::Time64(TimeUnit::Nanosecond) => Some(
10552 time_src
10553 .as_any()
10554 .downcast_ref::<Time64NanosecondArray>()?
10555 .value(i),
10556 ),
10557 DataType::Struct(_) => {
10558 let s = time_src.as_any().downcast_ref::<StructArray>()?;
10559 if is_date_struct(time_src.data_type()) {
10564 Some(0)
10565 } else if is_time_struct(time_src.data_type()) {
10566 Some(time_struct_parts(s, i)?.0)
10567 } else {
10568 Some(localdatetime_struct_parts(s, i)?.1)
10569 }
10570 }
10571 DataType::Utf8 => {
10572 time_of_day_nanos_any(time_src.as_any().downcast_ref::<StringArray>()?.value(i))
10573 }
10574 _ => None,
10575 }
10576 };
10577
10578 let parts: Vec<Option<(i64, i64)>> = (0..rows)
10579 .map(|i| {
10580 let date_overrides = DateOverrides {
10581 year: optional_i64_at(&ov[0], i),
10582 month: optional_i64_at(&ov[1], i),
10583 day: optional_i64_at(&ov[2], i),
10584 week: optional_i64_at(&ov[3], i),
10585 day_of_week: optional_i64_at(&ov[4], i),
10586 ordinal_day: optional_i64_at(&ov[5], i),
10587 quarter: optional_i64_at(&ov[6], i),
10588 day_of_quarter: optional_i64_at(&ov[7], i),
10589 };
10590 let time_overrides = LocalTimeOverrides {
10591 hour: optional_i64_at(&ov[8], i),
10592 minute: optional_i64_at(&ov[9], i),
10593 second: optional_i64_at(&ov[10], i),
10594 millisecond: optional_i64_at(&ov[11], i),
10595 microsecond: optional_i64_at(&ov[12], i),
10596 nanosecond: optional_i64_at(&ov[13], i),
10597 };
10598 let date = project_date(base_date(i)?, &date_overrides)?;
10599 let time = project_localtime(base_time(i)?, &time_overrides)?;
10600 Some((date, time))
10601 })
10602 .collect();
10603 Ok(ColumnarValue::Array(std::sync::Arc::new(
10604 build_localdatetime_struct(&parts),
10605 )))
10606 }
10607}
10608
10609static CYPHER_LOCALDATETIME_TRUNCATE: LazyLock<ScalarUDF> =
10621 LazyLock::new(|| ScalarUDF::new_from_impl(CypherLocalDateTimeTruncate::new()));
10622
10623#[derive(Debug, PartialEq, Eq, Hash)]
10624struct CypherLocalDateTimeTruncate {
10625 signature: Signature,
10626}
10627
10628impl CypherLocalDateTimeTruncate {
10629 fn new() -> Self {
10630 Self {
10631 signature: Signature::any(16, Volatility::Immutable),
10632 }
10633 }
10634}
10635
10636impl ScalarUDFImpl for CypherLocalDateTimeTruncate {
10637 fn as_any(&self) -> &dyn Any {
10638 self
10639 }
10640
10641 fn name(&self) -> &'static str {
10642 "cypher_localdatetime_truncate"
10643 }
10644
10645 fn signature(&self) -> &Signature {
10646 &self.signature
10647 }
10648
10649 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
10650 Ok(DataType::Struct(localdatetime_fields()))
10651 }
10652
10653 #[allow(
10654 clippy::too_many_lines,
10655 reason = "one cohesive per-row truncation: typed source extraction, the \
10656 date/time granularity split, and 14 component overrides"
10657 )]
10658 fn invoke_with_args(
10659 &self,
10660 args: ScalarFunctionArgs,
10661 ) -> datafusion::error::Result<ColumnarValue> {
10662 use crate::temporal::{
10663 DateOverrides, LocalTimeOverrides, parse_date_or_datetime_prefix, project_date,
10664 project_localtime, time_of_day_nanos_any, truncate_date, truncate_time_nanos,
10665 };
10666 use datafusion::arrow::array::{
10667 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
10668 };
10669 use datafusion::arrow::compute::cast;
10670 use datafusion::arrow::datatypes::TimeUnit;
10671 use datafusion::error::DataFusionError;
10672
10673 let rows = args.number_rows;
10674 let cols = udf_argument_arrays(&args)?;
10675 let value: ArrayRef =
10677 if matches!(cols[0].data_type(), DataType::Time64(TimeUnit::Nanosecond))
10678 || is_date_struct(cols[0].data_type())
10679 || is_localdatetime_struct(cols[0].data_type())
10680 || is_time_struct(cols[0].data_type())
10681 || is_datetime_struct(cols[0].data_type())
10682 {
10683 std::sync::Arc::clone(&cols[0])
10684 } else {
10685 cast(&cols[0], &DataType::Utf8).map_err(DataFusionError::from)?
10686 };
10687 let units_arr = cast(&cols[1], &DataType::Utf8).map_err(DataFusionError::from)?;
10688 let units = units_arr.as_any().downcast_ref::<StringArray>();
10689 let ov = cast_argument_arrays(&cols[2..16], &DataType::Int64)?;
10690
10691 let base_date = |i: usize| -> Option<i64> {
10692 if value.is_null(i) {
10693 return None;
10694 }
10695 match value.data_type() {
10696 DataType::Struct(_) => {
10697 let s = value.as_any().downcast_ref::<StructArray>()?;
10698 if is_date_struct(value.data_type()) {
10699 date_struct_value(s, i)
10700 } else {
10701 Some(localdatetime_struct_parts(s, i)?.0)
10702 }
10703 }
10704 DataType::Utf8 => parse_date_or_datetime_prefix(
10705 value.as_any().downcast_ref::<StringArray>()?.value(i),
10706 ),
10707 _ => None,
10708 }
10709 };
10710 let base_time = |i: usize| -> Option<i64> {
10711 if value.is_null(i) {
10712 return None;
10713 }
10714 match value.data_type() {
10715 DataType::Time64(TimeUnit::Nanosecond) => Some(
10716 value
10717 .as_any()
10718 .downcast_ref::<Time64NanosecondArray>()?
10719 .value(i),
10720 ),
10721 DataType::Struct(_) => {
10722 let s = value.as_any().downcast_ref::<StructArray>()?;
10723 if is_date_struct(value.data_type()) {
10724 Some(0) } else if is_time_struct(value.data_type()) {
10726 Some(time_struct_parts(s, i)?.0)
10727 } else {
10728 Some(localdatetime_struct_parts(s, i)?.1)
10729 }
10730 }
10731 DataType::Utf8 => {
10732 time_of_day_nanos_any(value.as_any().downcast_ref::<StringArray>()?.value(i))
10733 }
10734 _ => None,
10735 }
10736 };
10737
10738 let parts: Vec<Option<(i64, i64)>> = (0..rows)
10739 .map(|i| {
10740 let u = units?;
10741 if u.is_null(i) {
10742 return None;
10743 }
10744 let (date, time) = match truncate_date(base_date(i)?, u.value(i)) {
10747 Some(d) => (d, 0i64),
10748 None => (
10749 base_date(i)?,
10750 truncate_time_nanos(base_time(i)?, u.value(i))?,
10751 ),
10752 };
10753 let date_overrides = DateOverrides {
10754 year: optional_i64_at(&ov[0], i),
10755 month: optional_i64_at(&ov[1], i),
10756 day: optional_i64_at(&ov[2], i),
10757 week: optional_i64_at(&ov[3], i),
10758 day_of_week: optional_i64_at(&ov[4], i),
10759 ordinal_day: optional_i64_at(&ov[5], i),
10760 quarter: optional_i64_at(&ov[6], i),
10761 day_of_quarter: optional_i64_at(&ov[7], i),
10762 };
10763 let time_overrides = LocalTimeOverrides {
10764 hour: optional_i64_at(&ov[8], i),
10765 minute: optional_i64_at(&ov[9], i),
10766 second: optional_i64_at(&ov[10], i),
10767 millisecond: optional_i64_at(&ov[11], i),
10768 microsecond: optional_i64_at(&ov[12], i),
10769 nanosecond: optional_i64_at(&ov[13], i),
10770 };
10771 let date = project_date(date, &date_overrides)?;
10772 let time = project_localtime(time, &time_overrides)?;
10773 Some((date, time))
10774 })
10775 .collect();
10776 Ok(ColumnarValue::Array(std::sync::Arc::new(
10777 build_localdatetime_struct(&parts),
10778 )))
10779 }
10780}
10781
10782fn time_fields() -> datafusion::arrow::datatypes::Fields {
10789 graphforge_storage::schemas::time_struct_fields()
10790}
10791
10792fn is_time_struct(dt: &DataType) -> bool {
10795 use datafusion::arrow::datatypes::TimeUnit;
10796 matches!(dt, DataType::Struct(fields)
10797 if fields.len() == 2
10798 && fields[0].name() == "time"
10799 && *fields[0].data_type() == DataType::Time64(TimeUnit::Nanosecond)
10800 && fields[1].name() == "offset"
10801 && *fields[1].data_type() == DataType::Int32)
10802}
10803
10804fn build_time_struct(rows: &[Option<(i64, i32)>]) -> datafusion::arrow::array::StructArray {
10807 use datafusion::arrow::array::{Int32Array, Time64NanosecondArray};
10808 use datafusion::arrow::buffer::NullBuffer;
10809 let nanos: Time64NanosecondArray = rows.iter().map(|r| r.map(|(n, _)| n)).collect();
10810 let offset: Int32Array = rows.iter().map(|r| r.map(|(_, o)| o)).collect();
10811 let nulls = rows.iter().map(Option::is_some).collect::<NullBuffer>();
10812 datafusion::arrow::array::StructArray::new(
10813 time_fields(),
10814 vec![std::sync::Arc::new(nanos), std::sync::Arc::new(offset)],
10815 Some(nulls),
10816 )
10817}
10818
10819fn time_scalar(parts: Option<(i64, i32)>) -> ScalarValue {
10821 ScalarValue::Struct(std::sync::Arc::new(build_time_struct(&[parts])))
10822}
10823
10824fn time_struct_parts(arr: &datafusion::arrow::array::StructArray, i: usize) -> Option<(i64, i32)> {
10827 use datafusion::arrow::array::{Array, Int32Array, Time64NanosecondArray};
10828 if arr.is_null(i) {
10829 return None;
10830 }
10831 let t = arr
10832 .column(0)
10833 .as_any()
10834 .downcast_ref::<Time64NanosecondArray>()?;
10835 let o = arr.column(1).as_any().downcast_ref::<Int32Array>()?;
10836 (!t.is_null(i) && !o.is_null(i)).then(|| (t.value(i), o.value(i)))
10837}
10838
10839static CYPHER_TIME_PROJECT: LazyLock<ScalarUDF> =
10847 LazyLock::new(|| ScalarUDF::new_from_impl(CypherTimeProject::new()));
10848
10849#[derive(Debug, PartialEq, Eq, Hash)]
10850struct CypherTimeProject {
10851 signature: Signature,
10852}
10853
10854impl CypherTimeProject {
10855 fn new() -> Self {
10856 Self {
10857 signature: Signature::any(8, Volatility::Immutable),
10858 }
10859 }
10860}
10861
10862impl ScalarUDFImpl for CypherTimeProject {
10863 fn as_any(&self) -> &dyn Any {
10864 self
10865 }
10866
10867 fn name(&self) -> &'static str {
10868 "cypher_time_project"
10869 }
10870
10871 fn signature(&self) -> &Signature {
10872 &self.signature
10873 }
10874
10875 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
10876 Ok(DataType::Struct(time_fields()))
10877 }
10878
10879 fn invoke_with_args(
10880 &self,
10881 args: ScalarFunctionArgs,
10882 ) -> datafusion::error::Result<ColumnarValue> {
10883 use crate::temporal::{
10884 LocalTimeOverrides, parse_offset_seconds, project_localtime, project_time,
10885 time_of_day_with_offset,
10886 };
10887 use datafusion::arrow::array::{
10888 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
10889 };
10890 use datafusion::arrow::compute::cast;
10891 use datafusion::arrow::datatypes::TimeUnit;
10892 use datafusion::error::DataFusionError;
10893
10894 let rows = args.number_rows;
10895 let cols = udf_argument_arrays(&args)?;
10896 let base: ArrayRef =
10897 if matches!(cols[0].data_type(), DataType::Time64(TimeUnit::Nanosecond))
10898 || is_time_struct(cols[0].data_type())
10899 || is_localdatetime_struct(cols[0].data_type())
10900 || is_datetime_struct(cols[0].data_type())
10901 {
10902 std::sync::Arc::clone(&cols[0])
10903 } else {
10904 cast(&cols[0], &DataType::Utf8).map_err(DataFusionError::from)?
10905 };
10906 let ov = cast_argument_arrays(&cols[1..7], &DataType::Int64)?;
10907 let tz_arr = cast(&cols[7], &DataType::Utf8).map_err(DataFusionError::from)?;
10908 let tz = tz_arr.as_any().downcast_ref::<StringArray>();
10909
10910 let base_parts = |i: usize| -> Option<(i64, Option<i32>)> {
10913 if base.is_null(i) {
10914 return None;
10915 }
10916 match base.data_type() {
10917 DataType::Time64(TimeUnit::Nanosecond) => Some((
10918 base.as_any()
10919 .downcast_ref::<Time64NanosecondArray>()?
10920 .value(i),
10921 None,
10922 )),
10923 DataType::Struct(_) => {
10924 let s = base.as_any().downcast_ref::<StructArray>()?;
10925 if is_time_struct(base.data_type()) {
10926 let (n, o) = time_struct_parts(s, i)?;
10927 Some((n, Some(o)))
10928 } else if is_datetime_struct(base.data_type()) {
10929 let (_, n, o, _) = datetime_struct_parts(s, i)?;
10932 Some((n, Some(o)))
10933 } else {
10934 Some((localdatetime_struct_parts(s, i)?.1, None))
10936 }
10937 }
10938 DataType::Utf8 => {
10939 time_of_day_with_offset(base.as_any().downcast_ref::<StringArray>()?.value(i))
10940 }
10941 _ => None,
10942 }
10943 };
10944
10945 let parts: Vec<Option<(i64, i32)>> = (0..rows)
10946 .map(|i| {
10947 let (base_nanos, base_offset) = base_parts(i)?;
10948 let overrides = LocalTimeOverrides {
10949 hour: optional_i64_at(&ov[0], i),
10950 minute: optional_i64_at(&ov[1], i),
10951 second: optional_i64_at(&ov[2], i),
10952 millisecond: optional_i64_at(&ov[3], i),
10953 microsecond: optional_i64_at(&ov[4], i),
10954 nanosecond: optional_i64_at(&ov[5], i),
10955 };
10956 let nanos = project_localtime(base_nanos, &overrides)?;
10957 let new_offset = match tz {
10959 Some(a) if !a.is_null(i) => Some(parse_offset_seconds(a.value(i))?),
10960 _ => None,
10961 };
10962 Some(project_time(nanos, base_offset, new_offset))
10963 })
10964 .collect();
10965 Ok(ColumnarValue::Array(std::sync::Arc::new(
10966 build_time_struct(&parts),
10967 )))
10968 }
10969}
10970
10971static CYPHER_TIME_TRUNCATE: LazyLock<ScalarUDF> =
10980 LazyLock::new(|| ScalarUDF::new_from_impl(CypherTimeTruncate::new()));
10981
10982#[derive(Debug, PartialEq, Eq, Hash)]
10983struct CypherTimeTruncate {
10984 signature: Signature,
10985}
10986
10987impl CypherTimeTruncate {
10988 fn new() -> Self {
10989 Self {
10990 signature: Signature::any(9, Volatility::Immutable),
10991 }
10992 }
10993}
10994
10995impl ScalarUDFImpl for CypherTimeTruncate {
10996 fn as_any(&self) -> &dyn Any {
10997 self
10998 }
10999
11000 fn name(&self) -> &'static str {
11001 "cypher_time_truncate"
11002 }
11003
11004 fn signature(&self) -> &Signature {
11005 &self.signature
11006 }
11007
11008 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
11009 Ok(DataType::Struct(time_fields()))
11010 }
11011
11012 fn invoke_with_args(
11013 &self,
11014 args: ScalarFunctionArgs,
11015 ) -> datafusion::error::Result<ColumnarValue> {
11016 use crate::temporal::{
11017 LocalTimeOverrides, parse_offset_seconds, project_localtime, project_time,
11018 time_of_day_with_offset, truncate_time_nanos,
11019 };
11020 use datafusion::arrow::array::{
11021 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
11022 };
11023 use datafusion::arrow::compute::cast;
11024 use datafusion::arrow::datatypes::TimeUnit;
11025 use datafusion::error::DataFusionError;
11026
11027 let rows = args.number_rows;
11028 let cols = udf_argument_arrays(&args)?;
11029 let base: ArrayRef =
11030 if matches!(cols[0].data_type(), DataType::Time64(TimeUnit::Nanosecond))
11031 || is_time_struct(cols[0].data_type())
11032 || is_localdatetime_struct(cols[0].data_type())
11033 || is_datetime_struct(cols[0].data_type())
11034 {
11035 std::sync::Arc::clone(&cols[0])
11036 } else {
11037 cast(&cols[0], &DataType::Utf8).map_err(DataFusionError::from)?
11038 };
11039 let units_arr = cast(&cols[1], &DataType::Utf8).map_err(DataFusionError::from)?;
11040 let units = units_arr.as_any().downcast_ref::<StringArray>();
11041 let ov = cast_argument_arrays(&cols[2..8], &DataType::Int64)?;
11042 let tz_arr = cast(&cols[8], &DataType::Utf8).map_err(DataFusionError::from)?;
11043 let tz = tz_arr.as_any().downcast_ref::<StringArray>();
11044
11045 let base_parts = |i: usize| -> Option<(i64, Option<i32>)> {
11046 if base.is_null(i) {
11047 return None;
11048 }
11049 match base.data_type() {
11050 DataType::Time64(TimeUnit::Nanosecond) => Some((
11051 base.as_any()
11052 .downcast_ref::<Time64NanosecondArray>()?
11053 .value(i),
11054 None,
11055 )),
11056 DataType::Struct(_) => {
11057 let s = base.as_any().downcast_ref::<StructArray>()?;
11058 if is_time_struct(base.data_type()) {
11059 let (n, o) = time_struct_parts(s, i)?;
11060 Some((n, Some(o)))
11061 } else if is_datetime_struct(base.data_type()) {
11062 let (_, n, o, _) = datetime_struct_parts(s, i)?;
11063 Some((n, Some(o)))
11064 } else {
11065 Some((localdatetime_struct_parts(s, i)?.1, None))
11066 }
11067 }
11068 DataType::Utf8 => {
11069 time_of_day_with_offset(base.as_any().downcast_ref::<StringArray>()?.value(i))
11070 }
11071 _ => None,
11072 }
11073 };
11074
11075 let parts: Vec<Option<(i64, i32)>> = (0..rows)
11076 .map(|i| {
11077 let u = units?;
11078 if u.is_null(i) {
11079 return None;
11080 }
11081 let (base_nanos, base_offset) = base_parts(i)?;
11082 let truncated = truncate_time_nanos(base_nanos, u.value(i))?;
11083 let overrides = LocalTimeOverrides {
11084 hour: optional_i64_at(&ov[0], i),
11085 minute: optional_i64_at(&ov[1], i),
11086 second: optional_i64_at(&ov[2], i),
11087 millisecond: optional_i64_at(&ov[3], i),
11088 microsecond: optional_i64_at(&ov[4], i),
11089 nanosecond: optional_i64_at(&ov[5], i),
11090 };
11091 let nanos = project_localtime(truncated, &overrides)?;
11092 let new_offset = match tz {
11093 Some(a) if !a.is_null(i) => Some(parse_offset_seconds(a.value(i))?),
11094 _ => None,
11095 };
11096 let eff_offset = if new_offset.is_some() {
11101 None
11102 } else {
11103 base_offset
11104 };
11105 Some(project_time(nanos, eff_offset, new_offset))
11106 })
11107 .collect();
11108 Ok(ColumnarValue::Array(std::sync::Arc::new(
11109 build_time_struct(&parts),
11110 )))
11111 }
11112}
11113
11114type DateTimeRow = Option<(i64, i64, i32, Option<String>)>;
11121
11122fn datetime_fields() -> datafusion::arrow::datatypes::Fields {
11127 graphforge_storage::schemas::datetime_struct_fields()
11128}
11129
11130fn is_datetime_struct(dt: &DataType) -> bool {
11132 use datafusion::arrow::datatypes::TimeUnit;
11133 matches!(dt, DataType::Struct(fields)
11134 if fields.len() == 4
11135 && fields[0].name() == "date" && *fields[0].data_type() == DataType::Int64
11136 && fields[1].name() == "time"
11137 && *fields[1].data_type() == DataType::Time64(TimeUnit::Nanosecond)
11138 && fields[2].name() == "offset" && *fields[2].data_type() == DataType::Int32
11139 && fields[3].name() == "zone" && *fields[3].data_type() == DataType::Utf8)
11140}
11141
11142fn build_datetime_struct(rows: &[DateTimeRow]) -> datafusion::arrow::array::StructArray {
11145 use datafusion::arrow::array::{Int32Array, Int64Array, StringArray, Time64NanosecondArray};
11146 use datafusion::arrow::buffer::NullBuffer;
11147 let days: Int64Array = rows.iter().map(|r| r.as_ref().map(|t| t.0)).collect();
11148 let nanos: Time64NanosecondArray = rows.iter().map(|r| r.as_ref().map(|t| t.1)).collect();
11149 let offset: Int32Array = rows.iter().map(|r| r.as_ref().map(|t| t.2)).collect();
11150 let zone: StringArray = rows
11154 .iter()
11155 .map(|r| r.as_ref().map(|t| t.3.clone().unwrap_or_default()))
11156 .collect();
11157 let nulls = rows.iter().map(Option::is_some).collect::<NullBuffer>();
11158 datafusion::arrow::array::StructArray::new(
11159 datetime_fields(),
11160 vec![
11161 std::sync::Arc::new(days),
11162 std::sync::Arc::new(nanos),
11163 std::sync::Arc::new(offset),
11164 std::sync::Arc::new(zone),
11165 ],
11166 Some(nulls),
11167 )
11168}
11169
11170fn datetime_scalar(parts: DateTimeRow) -> ScalarValue {
11172 ScalarValue::Struct(std::sync::Arc::new(build_datetime_struct(&[parts])))
11173}
11174
11175fn datetime_struct_parts(arr: &datafusion::arrow::array::StructArray, i: usize) -> DateTimeRow {
11178 use datafusion::arrow::array::{
11179 Array, Int32Array, Int64Array, StringArray, Time64NanosecondArray,
11180 };
11181 if arr.is_null(i) {
11182 return None;
11183 }
11184 let days = arr.column(0).as_any().downcast_ref::<Int64Array>()?;
11185 let nanos = arr
11186 .column(1)
11187 .as_any()
11188 .downcast_ref::<Time64NanosecondArray>()?;
11189 let offset = arr.column(2).as_any().downcast_ref::<Int32Array>()?;
11190 let zone = arr.column(3).as_any().downcast_ref::<StringArray>()?;
11191 if days.is_null(i) || nanos.is_null(i) || offset.is_null(i) {
11192 return None;
11193 }
11194 let zone_label =
11196 (!zone.is_null(i) && !zone.value(i).is_empty()).then(|| zone.value(i).to_string());
11197 Some((days.value(i), nanos.value(i), offset.value(i), zone_label))
11198}
11199
11200static CYPHER_DATETIME_PROJECT: LazyLock<ScalarUDF> =
11208 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDateTimeProject::new()));
11209
11210#[derive(Debug, PartialEq, Eq, Hash)]
11211struct CypherDateTimeProject {
11212 signature: Signature,
11213}
11214
11215impl CypherDateTimeProject {
11216 fn new() -> Self {
11217 Self {
11218 signature: Signature::any(17, Volatility::Immutable),
11219 }
11220 }
11221}
11222
11223impl ScalarUDFImpl for CypherDateTimeProject {
11224 fn as_any(&self) -> &dyn Any {
11225 self
11226 }
11227
11228 fn name(&self) -> &'static str {
11229 "cypher_datetime_project"
11230 }
11231
11232 fn signature(&self) -> &Signature {
11233 &self.signature
11234 }
11235
11236 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
11237 Ok(DataType::Struct(datetime_fields()))
11238 }
11239
11240 #[allow(
11241 clippy::too_many_lines,
11242 reason = "one cohesive per-row projection: typed date/time/zone source \
11243 extraction, 14 component overrides, and zone re-resolution"
11244 )]
11245 fn invoke_with_args(
11246 &self,
11247 args: ScalarFunctionArgs,
11248 ) -> datafusion::error::Result<ColumnarValue> {
11249 use crate::temporal::{
11250 DateOverrides, LocalTimeOverrides, parse_date_or_datetime_prefix, project_date,
11251 project_datetime, project_localtime, time_offset_zone,
11252 };
11253 use datafusion::arrow::array::{
11254 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
11255 };
11256 use datafusion::arrow::compute::cast;
11257 use datafusion::arrow::datatypes::TimeUnit;
11258 use datafusion::error::DataFusionError;
11259
11260 let rows = args.number_rows;
11261 let cols = udf_argument_arrays(&args)?;
11262 let typed_or_utf8 = |a: &ArrayRef| -> datafusion::error::Result<ArrayRef> {
11263 if matches!(a.data_type(), DataType::Time64(TimeUnit::Nanosecond))
11264 || is_date_struct(a.data_type())
11265 || is_localdatetime_struct(a.data_type())
11266 || is_time_struct(a.data_type())
11267 || is_datetime_struct(a.data_type())
11268 {
11269 Ok(std::sync::Arc::clone(a))
11270 } else {
11271 cast(a, &DataType::Utf8).map_err(DataFusionError::from)
11272 }
11273 };
11274 let date_src = typed_or_utf8(&cols[0])?;
11275 let time_src = typed_or_utf8(&cols[1])?;
11276 let ov = cast_argument_arrays(&cols[2..16], &DataType::Int64)?;
11277 let tz_arr = cast(&cols[16], &DataType::Utf8).map_err(DataFusionError::from)?;
11278 let tz = tz_arr.as_any().downcast_ref::<StringArray>();
11279
11280 let base_date = |i: usize| -> Option<i64> {
11281 if date_src.is_null(i) {
11282 return Some(0); }
11284 match date_src.data_type() {
11285 DataType::Struct(_) => {
11286 let s = date_src.as_any().downcast_ref::<StructArray>()?;
11287 if is_date_struct(date_src.data_type()) {
11288 date_struct_value(s, i)
11289 } else if is_datetime_struct(date_src.data_type()) {
11290 Some(datetime_struct_parts(s, i)?.0)
11291 } else {
11292 Some(localdatetime_struct_parts(s, i)?.0)
11293 }
11294 }
11295 DataType::Utf8 => parse_date_or_datetime_prefix(
11296 date_src.as_any().downcast_ref::<StringArray>()?.value(i),
11297 ),
11298 _ => None,
11299 }
11300 };
11301 let base_time = |i: usize| -> Option<(i64, Option<i32>, Option<String>)> {
11303 if time_src.is_null(i) {
11304 return Some((0, None, None));
11305 }
11306 match time_src.data_type() {
11307 DataType::Time64(TimeUnit::Nanosecond) => Some((
11308 time_src
11309 .as_any()
11310 .downcast_ref::<Time64NanosecondArray>()?
11311 .value(i),
11312 None,
11313 None,
11314 )),
11315 DataType::Struct(_) => {
11316 let s = time_src.as_any().downcast_ref::<StructArray>()?;
11317 if is_date_struct(time_src.data_type()) {
11318 None
11321 } else if is_datetime_struct(time_src.data_type()) {
11322 let (_, n, o, z) = datetime_struct_parts(s, i)?;
11323 Some((n, Some(o), z))
11324 } else if is_time_struct(time_src.data_type()) {
11325 let (n, o) = time_struct_parts(s, i)?;
11326 Some((n, Some(o), None))
11327 } else {
11328 Some((localdatetime_struct_parts(s, i)?.1, None, None))
11329 }
11330 }
11331 DataType::Utf8 => {
11332 time_offset_zone(time_src.as_any().downcast_ref::<StringArray>()?.value(i))
11333 }
11334 _ => None,
11335 }
11336 };
11337
11338 let parts: Vec<DateTimeRow> = (0..rows)
11339 .map(|i| {
11340 let date_overrides = DateOverrides {
11341 year: optional_i64_at(&ov[0], i),
11342 month: optional_i64_at(&ov[1], i),
11343 day: optional_i64_at(&ov[2], i),
11344 week: optional_i64_at(&ov[3], i),
11345 day_of_week: optional_i64_at(&ov[4], i),
11346 ordinal_day: optional_i64_at(&ov[5], i),
11347 quarter: optional_i64_at(&ov[6], i),
11348 day_of_quarter: optional_i64_at(&ov[7], i),
11349 };
11350 let time_overrides = LocalTimeOverrides {
11351 hour: optional_i64_at(&ov[8], i),
11352 minute: optional_i64_at(&ov[9], i),
11353 second: optional_i64_at(&ov[10], i),
11354 millisecond: optional_i64_at(&ov[11], i),
11355 microsecond: optional_i64_at(&ov[12], i),
11356 nanosecond: optional_i64_at(&ov[13], i),
11357 };
11358 let (base_nanos, src_offset, src_zone) = base_time(i)?;
11359 let date = project_date(base_date(i)?, &date_overrides)?;
11360 let nanos = project_localtime(base_nanos, &time_overrides)?;
11361 let new_tz = tz.and_then(|a| (!a.is_null(i)).then(|| a.value(i)));
11362 let (date, nanos, offset, zone) =
11363 project_datetime(date, nanos, src_offset, src_zone.as_deref(), new_tz)?;
11364 Some((date, nanos, offset, zone))
11365 })
11366 .collect();
11367 Ok(ColumnarValue::Array(std::sync::Arc::new(
11368 build_datetime_struct(&parts),
11369 )))
11370 }
11371}
11372
11373static CYPHER_DATETIME_TRUNCATE: LazyLock<ScalarUDF> =
11384 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDateTimeTruncate::new()));
11385
11386#[derive(Debug, PartialEq, Eq, Hash)]
11387struct CypherDateTimeTruncate {
11388 signature: Signature,
11389}
11390
11391impl CypherDateTimeTruncate {
11392 fn new() -> Self {
11393 Self {
11394 signature: Signature::any(17, Volatility::Immutable),
11395 }
11396 }
11397}
11398
11399impl ScalarUDFImpl for CypherDateTimeTruncate {
11400 fn as_any(&self) -> &dyn Any {
11401 self
11402 }
11403
11404 fn name(&self) -> &'static str {
11405 "cypher_datetime_truncate"
11406 }
11407
11408 fn signature(&self) -> &Signature {
11409 &self.signature
11410 }
11411
11412 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
11413 Ok(DataType::Struct(datetime_fields()))
11414 }
11415
11416 #[allow(
11417 clippy::too_many_lines,
11418 reason = "one cohesive per-row truncation: typed source extraction, the \
11419 date/time granularity split, 14 overrides, and zone re-resolution"
11420 )]
11421 fn invoke_with_args(
11422 &self,
11423 args: ScalarFunctionArgs,
11424 ) -> datafusion::error::Result<ColumnarValue> {
11425 use crate::temporal::{
11426 DateOverrides, LocalTimeOverrides, parse_date_or_datetime_prefix, project_date,
11427 project_datetime, project_localtime, time_offset_zone, truncate_date,
11428 truncate_time_nanos,
11429 };
11430 use datafusion::arrow::array::{
11431 Array, ArrayRef, StringArray, StructArray, Time64NanosecondArray,
11432 };
11433 use datafusion::arrow::compute::cast;
11434 use datafusion::arrow::datatypes::TimeUnit;
11435 use datafusion::error::DataFusionError;
11436
11437 let rows = args.number_rows;
11438 let cols = udf_argument_arrays(&args)?;
11439 let value: ArrayRef =
11441 if matches!(cols[0].data_type(), DataType::Time64(TimeUnit::Nanosecond))
11442 || is_date_struct(cols[0].data_type())
11443 || is_localdatetime_struct(cols[0].data_type())
11444 || is_time_struct(cols[0].data_type())
11445 || is_datetime_struct(cols[0].data_type())
11446 {
11447 std::sync::Arc::clone(&cols[0])
11448 } else {
11449 cast(&cols[0], &DataType::Utf8).map_err(DataFusionError::from)?
11450 };
11451 let units_arr = cast(&cols[1], &DataType::Utf8).map_err(DataFusionError::from)?;
11452 let units = units_arr.as_any().downcast_ref::<StringArray>();
11453 let ov = cast_argument_arrays(&cols[2..16], &DataType::Int64)?;
11454 let tz_arr = cast(&cols[16], &DataType::Utf8).map_err(DataFusionError::from)?;
11455 let tz = tz_arr.as_any().downcast_ref::<StringArray>();
11456
11457 let base_date = |i: usize| -> Option<i64> {
11458 if value.is_null(i) {
11459 return None;
11460 }
11461 match value.data_type() {
11462 DataType::Struct(_) => {
11463 let s = value.as_any().downcast_ref::<StructArray>()?;
11464 if is_date_struct(value.data_type()) {
11465 date_struct_value(s, i)
11466 } else if is_datetime_struct(value.data_type()) {
11467 Some(datetime_struct_parts(s, i)?.0)
11468 } else {
11469 Some(localdatetime_struct_parts(s, i)?.0)
11470 }
11471 }
11472 DataType::Utf8 => parse_date_or_datetime_prefix(
11473 value.as_any().downcast_ref::<StringArray>()?.value(i),
11474 ),
11475 _ => None,
11476 }
11477 };
11478 let base_time = |i: usize| -> Option<(i64, Option<i32>, Option<String>)> {
11480 if value.is_null(i) {
11481 return None;
11482 }
11483 match value.data_type() {
11484 DataType::Time64(TimeUnit::Nanosecond) => Some((
11485 value
11486 .as_any()
11487 .downcast_ref::<Time64NanosecondArray>()?
11488 .value(i),
11489 None,
11490 None,
11491 )),
11492 DataType::Struct(_) => {
11493 let s = value.as_any().downcast_ref::<StructArray>()?;
11494 if is_date_struct(value.data_type()) {
11495 Some((0, None, None)) } else if is_datetime_struct(value.data_type()) {
11497 let (_, n, o, z) = datetime_struct_parts(s, i)?;
11498 Some((n, Some(o), z))
11499 } else if is_time_struct(value.data_type()) {
11500 let (n, o) = time_struct_parts(s, i)?;
11501 Some((n, Some(o), None))
11502 } else {
11503 Some((localdatetime_struct_parts(s, i)?.1, None, None))
11504 }
11505 }
11506 DataType::Utf8 => {
11507 time_offset_zone(value.as_any().downcast_ref::<StringArray>()?.value(i))
11508 }
11509 _ => None,
11510 }
11511 };
11512
11513 let parts: Vec<DateTimeRow> = (0..rows)
11514 .map(|i| {
11515 let u = units?;
11516 if u.is_null(i) {
11517 return None;
11518 }
11519 let (bt_nanos, src_offset, src_zone) = base_time(i)?;
11520 let (date0, nanos0) = match truncate_date(base_date(i)?, u.value(i)) {
11523 Some(d) => (d, 0i64),
11524 None => (base_date(i)?, truncate_time_nanos(bt_nanos, u.value(i))?),
11525 };
11526 let date_overrides = DateOverrides {
11527 year: optional_i64_at(&ov[0], i),
11528 month: optional_i64_at(&ov[1], i),
11529 day: optional_i64_at(&ov[2], i),
11530 week: optional_i64_at(&ov[3], i),
11531 day_of_week: optional_i64_at(&ov[4], i),
11532 ordinal_day: optional_i64_at(&ov[5], i),
11533 quarter: optional_i64_at(&ov[6], i),
11534 day_of_quarter: optional_i64_at(&ov[7], i),
11535 };
11536 let time_overrides = LocalTimeOverrides {
11537 hour: optional_i64_at(&ov[8], i),
11538 minute: optional_i64_at(&ov[9], i),
11539 second: optional_i64_at(&ov[10], i),
11540 millisecond: optional_i64_at(&ov[11], i),
11541 microsecond: optional_i64_at(&ov[12], i),
11542 nanosecond: optional_i64_at(&ov[13], i),
11543 };
11544 let date = project_date(date0, &date_overrides)?;
11545 let nanos = project_localtime(nanos0, &time_overrides)?;
11546 let new_tz = tz.and_then(|a| (!a.is_null(i)).then(|| a.value(i)));
11547 let src_offset = if new_tz.is_some() { None } else { src_offset };
11555 let (date, nanos, offset, zone) =
11556 project_datetime(date, nanos, src_offset, src_zone.as_deref(), new_tz)?;
11557 Some((date, nanos, offset, zone))
11558 })
11559 .collect();
11560 Ok(ColumnarValue::Array(std::sync::Arc::new(
11561 build_datetime_struct(&parts),
11562 )))
11563 }
11564}
11565
11566static CYPHER_TO_STRING: LazyLock<ScalarUDF> =
11578 LazyLock::new(|| ScalarUDF::new_from_impl(CypherToString::new()));
11579
11580#[derive(Debug, PartialEq, Eq, Hash)]
11581struct CypherToString {
11582 signature: Signature,
11583}
11584
11585impl CypherToString {
11586 fn new() -> Self {
11587 Self {
11588 signature: Signature::any(1, Volatility::Immutable),
11589 }
11590 }
11591}
11592
11593impl ScalarUDFImpl for CypherToString {
11594 fn as_any(&self) -> &dyn Any {
11595 self
11596 }
11597
11598 fn name(&self) -> &'static str {
11599 "cypher_to_string"
11600 }
11601
11602 fn signature(&self) -> &Signature {
11603 &self.signature
11604 }
11605
11606 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
11607 Ok(DataType::Utf8)
11608 }
11609
11610 fn invoke_with_args(
11611 &self,
11612 args: ScalarFunctionArgs,
11613 ) -> datafusion::error::Result<ColumnarValue> {
11614 use crate::temporal::{
11615 format_date, render_localdatetime, render_localtime_nanos, render_time_value,
11616 };
11617 use datafusion::arrow::array::{Array, StringArray, StructArray, Time64NanosecondArray};
11618 use datafusion::arrow::compute::cast;
11619 use datafusion::arrow::datatypes::TimeUnit;
11620
11621 let arr = args.args[0].to_array(args.number_rows)?;
11622 let rows = arr.len();
11623 let render = |f: &dyn Fn(usize) -> Option<String>| -> ColumnarValue {
11625 let out: StringArray = (0..rows)
11626 .map(|i| if arr.is_null(i) { None } else { f(i) })
11627 .collect();
11628 ColumnarValue::Array(std::sync::Arc::new(out))
11629 };
11630
11631 let result = match arr.data_type() {
11632 DataType::Struct(_) if is_date_struct(arr.data_type()) => {
11633 let s = arr.as_any().downcast_ref::<StructArray>().unwrap();
11634 render(&|i| date_struct_value(s, i).map(format_date))
11635 }
11636 DataType::Time64(TimeUnit::Nanosecond) => {
11637 let a = arr
11638 .as_any()
11639 .downcast_ref::<Time64NanosecondArray>()
11640 .unwrap();
11641 render(&|i| Some(render_localtime_nanos(a.value(i))))
11642 }
11643 DataType::Struct(_) if is_localdatetime_struct(arr.data_type()) => {
11644 let s = arr.as_any().downcast_ref::<StructArray>().unwrap();
11645 render(&|i| {
11646 localdatetime_struct_parts(s, i).map(|(d, n)| render_localdatetime(d, n))
11647 })
11648 }
11649 DataType::Struct(_) if is_time_struct(arr.data_type()) => {
11650 let s = arr.as_any().downcast_ref::<StructArray>().unwrap();
11651 render(&|i| time_struct_parts(s, i).map(|(n, o)| render_time_value(n, o)))
11652 }
11653 DataType::Struct(_) if is_datetime_struct(arr.data_type()) => {
11654 let s = arr.as_any().downcast_ref::<StructArray>().unwrap();
11655 render(&|i| {
11656 datetime_struct_parts(s, i).map(|(d, n, o, z)| {
11657 crate::temporal::render_datetime_value(d, n, o, z.as_deref())
11658 })
11659 })
11660 }
11661 DataType::Struct(_) if is_duration_struct(arr.data_type()) => {
11662 let s = arr.as_any().downcast_ref::<StructArray>().unwrap();
11663 render(&|i| {
11664 duration_struct_parts(s, i).map(|d| crate::temporal::render_duration_value(&d))
11665 })
11666 }
11667 DataType::Struct(_) if is_het_struct_type(Some(arr.data_type())) => {
11668 let out: datafusion::error::Result<StringArray> = (0..rows)
11669 .map(|i| {
11670 let value = decoded_scalar_at(&arr, i)?;
11671 to_cypher_string(&value)
11672 })
11673 .collect();
11674 ColumnarValue::Array(std::sync::Arc::new(out?))
11675 }
11676 _ => ColumnarValue::Array(
11678 cast(&arr, &DataType::Utf8).map_err(datafusion::error::DataFusionError::from)?,
11679 ),
11680 };
11681 Ok(result)
11682 }
11683}
11684
11685fn scalar_as_i128(v: &ScalarValue) -> Option<i128> {
11687 match v {
11688 ScalarValue::Int8(Some(n)) => Some(i128::from(*n)),
11689 ScalarValue::Int16(Some(n)) => Some(i128::from(*n)),
11690 ScalarValue::Int32(Some(n)) => Some(i128::from(*n)),
11691 ScalarValue::Int64(Some(n)) => Some(i128::from(*n)),
11692 ScalarValue::UInt8(Some(n)) => Some(i128::from(*n)),
11693 ScalarValue::UInt16(Some(n)) => Some(i128::from(*n)),
11694 ScalarValue::UInt32(Some(n)) => Some(i128::from(*n)),
11695 ScalarValue::UInt64(Some(n)) => Some(i128::from(*n)),
11696 _ => None,
11697 }
11698}
11699
11700fn scalar_as_f64(v: &ScalarValue) -> Option<f64> {
11702 #[allow(
11703 clippy::cast_precision_loss,
11704 reason = "only used for mixed integer/float numeric semantics; pure integers use i128"
11705 )]
11706 match v {
11707 ScalarValue::Int8(Some(n)) => Some(f64::from(*n)),
11708 ScalarValue::Int16(Some(n)) => Some(f64::from(*n)),
11709 ScalarValue::Int32(Some(n)) => Some(f64::from(*n)),
11710 ScalarValue::Int64(Some(n)) => Some(*n as f64),
11711 ScalarValue::UInt8(Some(n)) => Some(f64::from(*n)),
11712 ScalarValue::UInt16(Some(n)) => Some(f64::from(*n)),
11713 ScalarValue::UInt32(Some(n)) => Some(f64::from(*n)),
11714 ScalarValue::UInt64(Some(n)) => Some(*n as f64),
11715 ScalarValue::Float32(Some(f)) => Some(f64::from(*f)),
11716 ScalarValue::Float64(Some(f)) => Some(*f),
11717 _ => None,
11718 }
11719}
11720
11721fn cypher_seq_eq(
11723 a: &datafusion::arrow::array::ArrayRef,
11724 b: &datafusion::arrow::array::ArrayRef,
11725) -> Option<bool> {
11726 if a.len() != b.len() {
11727 return Some(false); }
11729 let mut saw_null = false;
11730 for i in 0..a.len() {
11731 let av = ScalarValue::try_from_array(a, i).ok()?;
11732 let bv = ScalarValue::try_from_array(b, i).ok()?;
11733 match cypher_value_eq(&av, &bv) {
11734 Some(false) => return Some(false),
11735 None => saw_null = true,
11736 Some(true) => {}
11737 }
11738 }
11739 if saw_null { None } else { Some(true) }
11740}
11741
11742fn cypher_struct_eq(
11744 a: &datafusion::arrow::array::StructArray,
11745 b: &datafusion::arrow::array::StructArray,
11746) -> Option<bool> {
11747 let mut ka: Vec<&str> = a.fields().iter().map(|f| f.name().as_str()).collect();
11749 let mut kb: Vec<&str> = b.fields().iter().map(|f| f.name().as_str()).collect();
11750 ka.sort_unstable();
11751 kb.sort_unstable();
11752 if ka != kb {
11753 return Some(false);
11754 }
11755 let mut saw_null = false;
11756 for key in ka {
11757 let av = ScalarValue::try_from_array(a.column_by_name(key)?, 0).ok()?;
11758 let bv = ScalarValue::try_from_array(b.column_by_name(key)?, 0).ok()?;
11759 match cypher_value_eq(&av, &bv) {
11760 Some(false) => return Some(false),
11761 None => saw_null = true,
11762 Some(true) => {}
11763 }
11764 }
11765 if saw_null { None } else { Some(true) }
11766}
11767
11768static CYPHER_DATE_TRUNCATE: LazyLock<ScalarUDF> =
11778 LazyLock::new(|| ScalarUDF::new_from_impl(CypherDateTruncate::new()));
11779
11780#[derive(Debug, PartialEq, Eq, Hash)]
11781struct CypherDateTruncate {
11782 signature: Signature,
11783}
11784
11785impl CypherDateTruncate {
11786 fn new() -> Self {
11787 Self {
11788 signature: Signature::any(10, Volatility::Immutable),
11789 }
11790 }
11791}
11792
11793impl ScalarUDFImpl for CypherDateTruncate {
11794 fn as_any(&self) -> &dyn Any {
11795 self
11796 }
11797
11798 fn name(&self) -> &'static str {
11799 "cypher_date_truncate"
11800 }
11801
11802 fn signature(&self) -> &Signature {
11803 &self.signature
11804 }
11805
11806 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
11807 Ok(DataType::Struct(
11808 graphforge_storage::schemas::date_struct_fields(),
11809 ))
11810 }
11811
11812 fn invoke_with_args(
11813 &self,
11814 args: ScalarFunctionArgs,
11815 ) -> datafusion::error::Result<ColumnarValue> {
11816 use crate::temporal::{
11817 DateOverrides, parse_date_or_datetime_prefix, project_date, truncate_date,
11818 };
11819 use datafusion::arrow::array::{Array, ArrayRef, StringArray, StructArray};
11820 use datafusion::arrow::compute::cast;
11821 use datafusion::error::DataFusionError;
11822
11823 let rows = args.number_rows;
11824 let cols = udf_argument_arrays(&args)?;
11825 let value: ArrayRef = if is_date_struct(cols[0].data_type())
11830 || is_localdatetime_struct(cols[0].data_type())
11831 || is_time_struct(cols[0].data_type())
11832 || is_datetime_struct(cols[0].data_type())
11833 {
11834 std::sync::Arc::clone(&cols[0])
11835 } else {
11836 cast(&cols[0], &DataType::Utf8).map_err(DataFusionError::from)?
11837 };
11838 let units_arr = cast(&cols[1], &DataType::Utf8).map_err(DataFusionError::from)?;
11839 let units = units_arr.as_any().downcast_ref::<StringArray>();
11840 let ov = cast_argument_arrays(&cols[2..10], &DataType::Int64)?;
11841
11842 let base_date = |i: usize| -> Option<i64> {
11843 if value.is_null(i) {
11844 return None;
11845 }
11846 match value.data_type() {
11847 DataType::Struct(_) => {
11848 let s = value.as_any().downcast_ref::<StructArray>()?;
11849 if is_date_struct(value.data_type()) {
11850 date_struct_value(s, i)
11851 } else {
11852 Some(localdatetime_struct_parts(s, i)?.0)
11854 }
11855 }
11856 DataType::Utf8 => parse_date_or_datetime_prefix(
11857 value.as_any().downcast_ref::<StringArray>()?.value(i),
11858 ),
11859 _ => None,
11860 }
11861 };
11862
11863 let out: Vec<Option<i64>> = (0..rows)
11864 .map(|i| {
11865 let u = units?;
11866 if u.is_null(i) {
11867 return None;
11868 }
11869 let truncated = truncate_date(base_date(i)?, u.value(i))?;
11870 let overrides = DateOverrides {
11871 year: optional_i64_at(&ov[0], i),
11872 month: optional_i64_at(&ov[1], i),
11873 day: optional_i64_at(&ov[2], i),
11874 week: optional_i64_at(&ov[3], i),
11875 day_of_week: optional_i64_at(&ov[4], i),
11876 ordinal_day: optional_i64_at(&ov[5], i),
11877 quarter: optional_i64_at(&ov[6], i),
11878 day_of_quarter: optional_i64_at(&ov[7], i),
11879 };
11880 project_date(truncated, &overrides)
11881 })
11882 .collect();
11883 Ok(ColumnarValue::Array(std::sync::Arc::new(
11884 build_date_struct(&out),
11885 )))
11886 }
11887}
11888
11889static CYPHER_PATH_NODES: LazyLock<ScalarUDF> =
11907 LazyLock::new(|| ScalarUDF::new_from_impl(CypherPathNodes::new()));
11908
11909fn path_node_struct_fields() -> datafusion::arrow::datatypes::Fields {
11915 use datafusion::arrow::datatypes::Field;
11916 vec![Field::new(
11917 "node_uuid",
11918 DataType::FixedSizeBinary(16),
11919 false,
11920 )]
11921 .into()
11922}
11923
11924#[derive(Debug, PartialEq, Eq, Hash)]
11928struct PathNodeHydration {
11929 dir: std::path::PathBuf,
11931 labels_by_type: Vec<(u32, String)>,
11933 prop_stems: Vec<String>,
11936 fields: datafusion::arrow::datatypes::Fields,
11938}
11939
11940#[derive(Debug, PartialEq, Eq, Hash)]
11941struct CypherPathNodes {
11942 signature: Signature,
11943 hydrate: Option<PathNodeHydration>,
11944}
11945
11946impl CypherPathNodes {
11947 fn new() -> Self {
11948 Self {
11950 signature: Signature::any(2, Volatility::Immutable),
11951 hydrate: None,
11952 }
11953 }
11954
11955 fn with_hydration(hydrate: PathNodeHydration) -> Self {
11956 Self {
11957 signature: Signature::any(2, Volatility::Immutable),
11958 hydrate: Some(hydrate),
11959 }
11960 }
11961
11962 fn element_fields(&self) -> datafusion::arrow::datatypes::Fields {
11963 self.hydrate
11964 .as_ref()
11965 .map_or_else(path_node_struct_fields, |h| h.fields.clone())
11966 }
11967}
11968
11969impl ScalarUDFImpl for CypherPathNodes {
11970 fn as_any(&self) -> &dyn Any {
11971 self
11972 }
11973
11974 fn name(&self) -> &'static str {
11975 "cypher_path_nodes"
11976 }
11977
11978 fn signature(&self) -> &Signature {
11979 &self.signature
11980 }
11981
11982 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
11983 Ok(DataType::new_list(
11984 DataType::Struct(self.element_fields()),
11985 true,
11986 ))
11987 }
11988
11989 fn invoke_with_args(
11990 &self,
11991 args: ScalarFunctionArgs,
11992 ) -> datafusion::error::Result<ColumnarValue> {
11993 use datafusion::arrow::array::{
11994 Array, ArrayRef, FixedSizeBinaryArray, ListArray, StructArray, new_empty_array,
11995 };
11996 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer};
11997 use datafusion::arrow::datatypes::Field;
11998 use datafusion::common::cast::as_list_array;
11999 use datafusion::error::DataFusionError;
12000 use std::sync::Arc;
12001
12002 let exec_err = |m: String| DataFusionError::Execution(m);
12003 let as_fsb16 =
12004 |array: &dyn Array, what: &str| -> datafusion::error::Result<FixedSizeBinaryArray> {
12005 array
12006 .as_any()
12007 .downcast_ref::<FixedSizeBinaryArray>()
12008 .filter(|a| a.value_length() == 16)
12009 .cloned()
12010 .ok_or_else(|| {
12011 exec_err(format!(
12012 "cypher_path_nodes: expected FixedSizeBinary(16) {what}, got {:?}",
12013 array.data_type()
12014 ))
12015 })
12016 };
12017
12018 let seeds = args.args[0].to_array(args.number_rows)?;
12019 let seeds = as_fsb16(seeds.as_ref(), "start-node uuid")?;
12020 let rels = args.args[1].to_array(args.number_rows)?;
12021 let rels = as_list_array(&rels)?.clone();
12022
12023 let mut flat: Vec<[u8; 16]> = Vec::new();
12026 let mut lengths: Vec<usize> = Vec::with_capacity(rels.len());
12027 let mut valid: Vec<bool> = Vec::with_capacity(rels.len());
12028 for row in 0..rels.len() {
12029 if seeds.is_null(row) || rels.is_null(row) {
12030 lengths.push(0);
12031 valid.push(false);
12032 continue;
12033 }
12034 let start = flat.len();
12035 let mut cur = [0u8; 16];
12036 cur.copy_from_slice(seeds.value(row));
12037 flat.push(cur);
12038
12039 let edges = rels.value(row);
12040 let edges = edges
12041 .as_any()
12042 .downcast_ref::<StructArray>()
12043 .ok_or_else(|| {
12044 exec_err("cypher_path_nodes: relationship-list items must be structs".into())
12045 })?;
12046 let src = edges.column_by_name("src_uuid").ok_or_else(|| {
12049 exec_err("cypher_path_nodes: relationship struct has no src_uuid".into())
12050 })?;
12051 let src = as_fsb16(src.as_ref(), "src_uuid")?;
12052 let dst = edges.column_by_name("dst_uuid").ok_or_else(|| {
12053 exec_err("cypher_path_nodes: relationship struct has no dst_uuid".into())
12054 })?;
12055 let dst = as_fsb16(dst.as_ref(), "dst_uuid")?;
12056
12057 for i in 0..edges.len() {
12058 let s = src.value(i);
12059 let d = dst.value(i);
12060 let next = if cur == s {
12061 d
12062 } else if cur == d {
12063 s
12064 } else {
12065 return Err(exec_err(format!(
12066 "cypher_path_nodes: edge {i} is disconnected from the \
12067 path (corrupt traversal emission)"
12068 )));
12069 };
12070 cur.copy_from_slice(next);
12071 flat.push(cur);
12072 }
12073 lengths.push(flat.len() - start);
12074 valid.push(true);
12075 }
12076
12077 let uuid_child: ArrayRef = if flat.is_empty() {
12080 new_empty_array(&DataType::FixedSizeBinary(16))
12081 } else {
12082 Arc::new(
12083 FixedSizeBinaryArray::try_from_iter(flat.iter())
12084 .map_err(|e| exec_err(e.to_string()))?,
12085 )
12086 };
12087 let fields = self.element_fields();
12088 let mut children: Vec<ArrayRef> = vec![uuid_child];
12089 if let Some(h) = self.hydrate.as_ref() {
12090 children.extend(hydrate_path_node_children(h, &flat)?);
12091 }
12092
12093 let struct_arr = StructArray::try_new(fields.clone(), children, None)
12094 .map_err(|e| exec_err(e.to_string()))?;
12095 let offsets = OffsetBuffer::<i32>::from_lengths(lengths);
12096 let item = Arc::new(Field::new("item", DataType::Struct(fields), true));
12097 let list = ListArray::try_new(
12098 item,
12099 offsets,
12100 Arc::new(struct_arr),
12101 Some(NullBuffer::from(valid)),
12102 )
12103 .map_err(|e| exec_err(e.to_string()))?;
12104 Ok(ColumnarValue::Array(Arc::new(list)))
12105 }
12106}
12107
12108fn hydrate_path_node_children(
12111 h: &PathNodeHydration,
12112 flat: &[[u8; 16]],
12113) -> datafusion::error::Result<Vec<datafusion::arrow::array::ArrayRef>> {
12114 let mut children = vec![path_node_labels_child(h, flat)?];
12115 children.extend(path_node_prop_children(h, flat)?);
12116 Ok(children)
12117}
12118
12119fn hydration_fsb16(
12121 b: &datafusion::arrow::array::RecordBatch,
12122 name: &str,
12123) -> datafusion::error::Result<datafusion::arrow::array::FixedSizeBinaryArray> {
12124 use datafusion::arrow::array::FixedSizeBinaryArray;
12125 b.column_by_name(name)
12126 .and_then(|c| c.as_any().downcast_ref::<FixedSizeBinaryArray>().cloned())
12127 .filter(|a| a.value_length() == 16)
12128 .ok_or_else(|| {
12129 datafusion::error::DataFusionError::Execution(format!(
12130 "cypher_path_nodes: no FixedSizeBinary(16) {name} column"
12131 ))
12132 })
12133}
12134
12135fn path_node_labels_child(
12140 h: &PathNodeHydration,
12141 flat: &[[u8; 16]],
12142) -> datafusion::error::Result<datafusion::arrow::array::ArrayRef> {
12143 use datafusion::arrow::array::{Array, ListBuilder, StringBuilder};
12144 use datafusion::error::DataFusionError;
12145 use std::collections::HashMap;
12146
12147 let exec_err = |m: String| DataFusionError::Execution(m);
12148 let node_batches =
12149 graphforge_storage::read_nodes(&h.dir).map_err(|e| exec_err(e.to_string()))?;
12150 let mut label_of: HashMap<[u8; 16], usize> = HashMap::new();
12151 for b in &node_batches {
12152 let uuids = hydration_fsb16(b, "node_uuid")?;
12153 let type_ids = b
12154 .column_by_name("type_id")
12155 .and_then(|c| {
12156 c.as_any()
12157 .downcast_ref::<datafusion::arrow::array::UInt32Array>()
12158 .cloned()
12159 })
12160 .ok_or_else(|| exec_err("cypher_path_nodes: no UInt32 type_id column".into()))?;
12161 for r in 0..b.num_rows() {
12162 if uuids.is_null(r) || type_ids.is_null(r) {
12163 continue;
12164 }
12165 let mut u = [0u8; 16];
12166 u.copy_from_slice(uuids.value(r));
12167 if let Ok(i) = h
12168 .labels_by_type
12169 .binary_search_by_key(&type_ids.value(r), |(id, _)| *id)
12170 {
12171 label_of.insert(u, i);
12172 }
12173 }
12174 }
12175 let mut labels_b = ListBuilder::new(StringBuilder::new());
12176 for u in flat {
12177 match label_of.get(u) {
12178 Some(&i) => labels_b.values().append_value(&h.labels_by_type[i].1),
12179 None => labels_b.values().append_null(),
12180 }
12181 labels_b.append(true);
12182 }
12183 Ok(std::sync::Arc::new(labels_b.finish()))
12184}
12185
12186fn path_node_prop_children(
12191 h: &PathNodeHydration,
12192 flat: &[[u8; 16]],
12193) -> datafusion::error::Result<Vec<datafusion::arrow::array::ArrayRef>> {
12194 use datafusion::arrow::array::{Array, ArrayRef, UInt32Array, new_null_array};
12195 use datafusion::arrow::compute::kernels::zip::zip;
12196 use datafusion::arrow::compute::{concat_batches, is_not_null, take};
12197 use datafusion::error::DataFusionError;
12198 use std::collections::HashMap;
12199
12200 let exec_err = |m: String| DataFusionError::Execution(m);
12201
12202 let mut batches = Vec::with_capacity(h.prop_stems.len());
12206 for stem in &h.prop_stems {
12207 let bs = graphforge_storage::read_properties(&h.dir, stem)
12208 .map_err(|e| exec_err(e.to_string()))?;
12209 if let Some(first) = bs.first() {
12210 batches
12211 .push(concat_batches(&first.schema(), &bs).map_err(|e| exec_err(e.to_string()))?);
12212 }
12213 }
12214 let mut uuid_to_loc: HashMap<[u8; 16], (usize, u32)> = HashMap::new();
12215 for (bi, b) in batches.iter().enumerate() {
12216 let key = hydration_fsb16(b, "node_uuid")?;
12217 for r in 0..key.len() {
12218 if key.is_null(r) {
12219 continue;
12220 }
12221 let mut u = [0u8; 16];
12222 u.copy_from_slice(key.value(r));
12223 uuid_to_loc.entry(u).or_insert((
12224 bi,
12225 u32::try_from(r).map_err(|_| exec_err(format!("property row {r} exceeds u32")))?,
12226 ));
12227 }
12228 }
12229 let take_by_batch: Vec<UInt32Array> = (0..batches.len())
12230 .map(|bi| {
12231 flat.iter()
12232 .map(|u| match uuid_to_loc.get(u) {
12233 Some(&(owner, row)) if owner == bi => Some(row),
12234 _ => None,
12235 })
12236 .collect()
12237 })
12238 .collect();
12239
12240 let mut children = Vec::with_capacity(h.fields.len().saturating_sub(2));
12241 for field in h.fields.iter().skip(2) {
12242 let mut child: ArrayRef = new_null_array(field.data_type(), flat.len());
12243 for (bi, b) in batches.iter().enumerate() {
12244 let Some(col) = b.column_by_name(field.name()) else {
12247 continue;
12248 };
12249 let taken = take(col, &take_by_batch[bi], None).map_err(|e| exec_err(e.to_string()))?;
12250 child = if batches.len() == 1 {
12251 taken
12252 } else {
12253 let mask = is_not_null(&taken).map_err(|e| exec_err(e.to_string()))?;
12254 zip(&mask, &taken, &child).map_err(|e| exec_err(e.to_string()))?
12255 };
12256 }
12257 children.push(child);
12258 }
12259 Ok(children)
12260}
12261
12262#[cfg(test)]
12267mod tests {
12268 use super::*;
12269 use graphforge_core::PropId;
12270 use graphforge_ir::expr::{BinaryOpKind, IrExpr, IrLiteral, UnaryOpKind};
12271 use graphforge_ir::{ExprArena, VarId};
12272 use std::sync::atomic::{AtomicUsize, Ordering};
12273
12274 static VOLATILE_CALLS: AtomicUsize = AtomicUsize::new(0);
12275 static VOLATILE_ROWS: AtomicUsize = AtomicUsize::new(0);
12276
12277 fn invoke_test_udf<U: ScalarUDFImpl>(
12278 udf: &U,
12279 values: Vec<ScalarValue>,
12280 ) -> datafusion::error::Result<datafusion::arrow::array::ArrayRef> {
12281 use datafusion::arrow::datatypes::Field;
12282 use datafusion::config::ConfigOptions;
12283
12284 let types = values
12285 .iter()
12286 .map(ScalarValue::data_type)
12287 .collect::<Vec<_>>();
12288 let return_type = udf.return_type(&types)?;
12289 let result = udf.invoke_with_args(ScalarFunctionArgs {
12290 args: values.into_iter().map(ColumnarValue::Scalar).collect(),
12291 arg_fields: types
12292 .iter()
12293 .enumerate()
12294 .map(|(index, data_type)| {
12295 Arc::new(Field::new(format!("arg_{index}"), data_type.clone(), true))
12296 })
12297 .collect(),
12298 number_rows: 1,
12299 return_field: Arc::new(Field::new("result", return_type.clone(), true)),
12300 config_options: Arc::new(ConfigOptions::default()),
12301 })?;
12302 let array = match result {
12303 ColumnarValue::Array(array) => array,
12304 ColumnarValue::Scalar(value) => value.to_array_of_size(1)?,
12305 };
12306 assert_eq!(array.data_type(), &return_type);
12307 Ok(array)
12308 }
12309
12310 fn invoke_test_udf_with_return_type<U: ScalarUDFImpl>(
12311 udf: &U,
12312 values: Vec<ScalarValue>,
12313 return_type: DataType,
12314 ) -> datafusion::error::Result<ColumnarValue> {
12315 use datafusion::arrow::datatypes::Field;
12316 use datafusion::config::ConfigOptions;
12317
12318 let types = values
12319 .iter()
12320 .map(ScalarValue::data_type)
12321 .collect::<Vec<_>>();
12322 udf.invoke_with_args(ScalarFunctionArgs {
12323 args: values.into_iter().map(ColumnarValue::Scalar).collect(),
12324 arg_fields: types
12325 .iter()
12326 .enumerate()
12327 .map(|(index, data_type)| {
12328 Arc::new(Field::new(format!("arg_{index}"), data_type.clone(), true))
12329 })
12330 .collect(),
12331 number_rows: 1,
12332 return_field: Arc::new(Field::new("result", return_type, true)),
12333 config_options: Arc::new(ConfigOptions::default()),
12334 })
12335 }
12336
12337 #[derive(Debug, PartialEq, Eq, Hash)]
12338 struct CountingVolatilePredicate {
12339 signature: Signature,
12340 }
12341
12342 impl CountingVolatilePredicate {
12343 fn new() -> Self {
12344 Self {
12345 signature: Signature::nullary(Volatility::Volatile),
12346 }
12347 }
12348 }
12349
12350 impl ScalarUDFImpl for CountingVolatilePredicate {
12351 fn as_any(&self) -> &dyn Any {
12352 self
12353 }
12354
12355 fn name(&self) -> &'static str {
12356 "counting_volatile_predicate"
12357 }
12358
12359 fn signature(&self) -> &Signature {
12360 &self.signature
12361 }
12362
12363 fn return_type(&self, _arg_types: &[DataType]) -> datafusion::error::Result<DataType> {
12364 Ok(DataType::Boolean)
12365 }
12366
12367 fn invoke_with_args(
12368 &self,
12369 args: ScalarFunctionArgs,
12370 ) -> datafusion::error::Result<ColumnarValue> {
12371 use datafusion::arrow::array::BooleanArray;
12372 VOLATILE_CALLS.fetch_add(1, Ordering::SeqCst);
12373 VOLATILE_ROWS.fetch_add(args.number_rows, Ordering::SeqCst);
12374 Ok(ColumnarValue::Array(std::sync::Arc::new(
12375 BooleanArray::from(vec![true; args.number_rows]),
12376 )))
12377 }
12378 }
12379
12380 #[test]
12381 fn temporal_clock_fn_recognition() {
12382 for base in ["date", "localtime", "time", "localdatetime", "datetime"] {
12384 for clock in ["transaction", "statement", "realtime"] {
12385 assert!(is_temporal_clock_fn(&format!("{base}.{clock}")));
12386 }
12387 }
12388 assert!(is_temporal_clock_fn("Date.Realtime"));
12390 assert!(is_temporal_clock_fn("DATETIME.TRANSACTION"));
12391 assert!(!is_temporal_clock_fn("datetime.truncate"));
12393 assert!(!is_temporal_clock_fn("duration.realtime")); assert!(!is_temporal_clock_fn("date"));
12395 assert!(!is_temporal_clock_fn("foo.realtime"));
12396
12397 assert_eq!(
12399 temporal_null_scalar("date").data_type(),
12400 date_scalar(None).data_type()
12401 );
12402 assert_eq!(
12403 temporal_null_scalar("localtime"),
12404 ScalarValue::Time64Nanosecond(None)
12405 );
12406 assert_eq!(
12407 temporal_null_scalar("datetime.realtime").data_type(),
12408 datetime_scalar(None).data_type()
12409 );
12410 assert_eq!(temporal_null_scalar("unknown"), ScalarValue::Null);
12411 }
12412
12413 #[test]
12414 fn temporal_struct_builders_and_extractors_preserve_values_and_nulls() {
12415 use crate::temporal::DurationValue;
12416 use datafusion::arrow::array::{ArrayRef, Int32Array, StructArray};
12417 use datafusion::arrow::datatypes::{Field, Fields};
12418
12419 let dates = build_date_struct(&[Some(19_723), None]);
12420 assert_eq!(date_struct_value(&dates, 0), Some(19_723));
12421 assert_eq!(date_struct_value(&dates, 1), None);
12422
12423 let local = build_localdatetime_struct(&[(Some((19_723, 45_000))), None]);
12424 assert_eq!(
12425 localdatetime_struct_parts(&local, 0),
12426 Some((19_723, 45_000))
12427 );
12428 assert_eq!(localdatetime_struct_parts(&local, 1), None);
12429
12430 let duration = DurationValue {
12431 months: 14,
12432 days: -2,
12433 seconds: 90,
12434 nanos: 123,
12435 };
12436 let durations = build_duration_struct(&[Some(duration), None]);
12437 assert_eq!(duration_struct_parts(&durations, 0), Some(duration));
12438 assert_eq!(duration_struct_parts(&durations, 1), None);
12439 assert_eq!(
12440 dur_secs_nanos(-3, 7),
12441 DurationValue {
12442 months: 0,
12443 days: 0,
12444 seconds: -3,
12445 nanos: 7,
12446 }
12447 );
12448 assert_eq!(
12449 duration_value_to_ir(duration),
12450 IrLiteral::Duration {
12451 months: 14,
12452 days: -2,
12453 seconds: 90,
12454 nanos: 123,
12455 }
12456 );
12457
12458 let times = build_time_struct(&[Some((86_399_000_000_000, -25_200)), None]);
12459 assert_eq!(
12460 time_struct_parts(×, 0),
12461 Some((86_399_000_000_000, -25_200))
12462 );
12463 assert_eq!(time_struct_parts(×, 1), None);
12464
12465 let datetimes = build_datetime_struct(&[
12466 Some((19_723, 45_000, 3_600, Some("Europe/Paris".into()))),
12467 Some((19_724, 46_000, 0, None)),
12468 None,
12469 ]);
12470 assert_eq!(
12471 datetime_struct_parts(&datetimes, 0),
12472 Some((19_723, 45_000, 3_600, Some("Europe/Paris".into())))
12473 );
12474 assert_eq!(
12475 datetime_struct_parts(&datetimes, 1),
12476 Some((19_724, 46_000, 0, None))
12477 );
12478 assert_eq!(datetime_struct_parts(&datetimes, 2), None);
12479
12480 let wrong_fields = Fields::from(vec![Field::new("epoch_day", DataType::Int32, true)]);
12483 let wrong_children: Vec<ArrayRef> = vec![Arc::new(Int32Array::from(vec![Some(1)]))];
12484 let wrong_date = StructArray::new(wrong_fields, wrong_children, None);
12485 assert_eq!(date_struct_value(&wrong_date, 0), None);
12486
12487 let overrides: ArrayRef = Arc::new(datafusion::arrow::array::Int64Array::from(vec![
12488 Some(8),
12489 None,
12490 ]));
12491 assert_eq!(optional_i64_at(&overrides, 0), Some(8));
12492 assert_eq!(optional_i64_at(&overrides, 1), None);
12493 let wrong_override: ArrayRef = Arc::new(Int32Array::from(vec![Some(8)]));
12494 assert_eq!(optional_i64_at(&wrong_override, 0), None);
12495 }
12496
12497 #[test]
12498 fn temporal_project_and_truncate_udfs_execute_all_typed_families() {
12499 use datafusion::arrow::array::Array;
12500
12501 let null_ints = |count| vec![ScalarValue::Int64(None); count];
12502 let assert_value = |array: datafusion::arrow::array::ArrayRef| {
12503 assert_eq!(array.len(), 1);
12504 assert!(!array.is_null(0));
12505 };
12506
12507 let mut args = vec![ScalarValue::Utf8(Some("2024-02-29".into()))];
12508 args.extend(null_ints(8));
12509 assert_value(invoke_test_udf(&CypherDateProject::new(), args).unwrap());
12510 let mut null_args = vec![ScalarValue::Utf8(None)];
12511 null_args.extend(null_ints(8));
12512 let null_date = invoke_test_udf(&CypherDateProject::new(), null_args).unwrap();
12513 assert!(null_date.is_null(0));
12514 let mut typed_args = vec![date_scalar(Some(19_782))];
12515 typed_args.extend(null_ints(8));
12516 assert_value(invoke_test_udf(&CypherDateProject::new(), typed_args).unwrap());
12517
12518 let mut args = vec![ScalarValue::Utf8(Some("12:34:56.123".into()))];
12519 args.extend(null_ints(6));
12520 assert_value(invoke_test_udf(&CypherLocalTimeProject::new(), args).unwrap());
12521 let mut typed_args = vec![ScalarValue::Time64Nanosecond(Some(45_296_123_000_000))];
12522 typed_args.extend(null_ints(6));
12523 assert_value(invoke_test_udf(&CypherLocalTimeProject::new(), typed_args).unwrap());
12524
12525 let mut args = vec![
12526 ScalarValue::Utf8(Some("12:34:56.123".into())),
12527 ScalarValue::Utf8(Some("second".into())),
12528 ];
12529 args.extend(null_ints(6));
12530 assert_value(invoke_test_udf(&CypherLocalTimeTruncate::new(), args).unwrap());
12531
12532 let mut args = vec![
12533 ScalarValue::Utf8(Some("2024-02-29".into())),
12534 ScalarValue::Utf8(Some("12:34:56.123".into())),
12535 ];
12536 args.extend(null_ints(14));
12537 assert_value(invoke_test_udf(&CypherLocalDateTimeProject::new(), args).unwrap());
12538 let mut typed_args = vec![
12539 date_scalar(Some(19_782)),
12540 ScalarValue::Time64Nanosecond(Some(45_296_123_000_000)),
12541 ];
12542 typed_args.extend(null_ints(14));
12543 assert_value(invoke_test_udf(&CypherLocalDateTimeProject::new(), typed_args).unwrap());
12544
12545 let mut args = vec![
12546 ScalarValue::Utf8(Some("2024-02-29T12:34:56.123".into())),
12547 ScalarValue::Utf8(Some("day".into())),
12548 ];
12549 args.extend(null_ints(14));
12550 assert_value(invoke_test_udf(&CypherLocalDateTimeTruncate::new(), args).unwrap());
12551
12552 let mut args = vec![ScalarValue::Utf8(Some("12:34:56+01:00".into()))];
12553 args.extend(null_ints(6));
12554 args.push(ScalarValue::Utf8(None));
12555 assert_value(invoke_test_udf(&CypherTimeProject::new(), args).unwrap());
12556 let mut typed_args = vec![time_scalar(Some((45_296_000_000_000, 3_600)))];
12557 typed_args.extend(null_ints(6));
12558 typed_args.push(ScalarValue::Utf8(None));
12559 assert_value(invoke_test_udf(&CypherTimeProject::new(), typed_args).unwrap());
12560
12561 let mut args = vec![
12562 ScalarValue::Utf8(Some("12:34:56+01:00".into())),
12563 ScalarValue::Utf8(Some("minute".into())),
12564 ];
12565 args.extend(null_ints(6));
12566 args.push(ScalarValue::Utf8(None));
12567 assert_value(invoke_test_udf(&CypherTimeTruncate::new(), args).unwrap());
12568
12569 let mut args = vec![
12570 ScalarValue::Utf8(Some("2024-02-29".into())),
12571 ScalarValue::Utf8(Some("12:34:56+01:00".into())),
12572 ];
12573 args.extend(null_ints(14));
12574 args.push(ScalarValue::Utf8(None));
12575 assert_value(invoke_test_udf(&CypherDateTimeProject::new(), args).unwrap());
12576 let mut typed_args = vec![
12577 date_scalar(Some(19_782)),
12578 time_scalar(Some((45_296_000_000_000, 3_600))),
12579 ];
12580 typed_args.extend(null_ints(14));
12581 typed_args.push(ScalarValue::Utf8(None));
12582 assert_value(invoke_test_udf(&CypherDateTimeProject::new(), typed_args).unwrap());
12583
12584 let mut args = vec![
12585 ScalarValue::Utf8(Some("2024-02-29T12:34:56+01:00".into())),
12586 ScalarValue::Utf8(Some("hour".into())),
12587 ];
12588 args.extend(null_ints(14));
12589 args.push(ScalarValue::Utf8(None));
12590 assert_value(invoke_test_udf(&CypherDateTimeTruncate::new(), args).unwrap());
12591
12592 let mut args = vec![
12593 ScalarValue::Utf8(Some("2024-02-29".into())),
12594 ScalarValue::Utf8(Some("month".into())),
12595 ];
12596 args.extend(null_ints(8));
12597 assert_value(invoke_test_udf(&CypherDateTruncate::new(), args).unwrap());
12598 }
12599
12600 #[test]
12601 fn quantifier_three_valued_reduce() {
12602 use datafusion::arrow::array::BooleanArray;
12603 use graphforge_ir::QuantifierKind::{All, Any, None, Single};
12604 let b = |v: Vec<Option<bool>>| BooleanArray::from(v);
12605 let r = |k, v: Vec<Option<bool>>| {
12606 let arr = b(v.clone());
12607 reduce_quantifier(k, &arr, v.len())
12608 };
12609 assert_eq!(r(All, vec![]), Some(true));
12611 assert_eq!(r(None, vec![]), Some(true));
12612 assert_eq!(r(Any, vec![]), Some(false));
12613 assert_eq!(r(Single, vec![]), Some(false));
12614 assert_eq!(r(All, vec![Some(true), Some(true)]), Some(true));
12616 assert_eq!(r(All, vec![Some(true), Some(false)]), Some(false));
12617 assert_eq!(r(Any, vec![Some(false), Some(true)]), Some(true));
12618 assert_eq!(r(None, vec![Some(false), Some(false)]), Some(true));
12619 assert_eq!(r(Single, vec![Some(true), Some(false)]), Some(true));
12620 assert_eq!(r(Single, vec![Some(true), Some(true)]), Some(false));
12621 assert_eq!(r(All, vec![Some(true), Option::None]), Option::None); assert_eq!(r(All, vec![Some(false), Option::None]), Some(false)); assert_eq!(r(Any, vec![Some(false), Option::None]), Option::None);
12625 assert_eq!(r(Any, vec![Some(true), Option::None]), Some(true)); assert_eq!(
12627 r(Single, vec![Some(true), Some(true), Option::None]),
12628 Some(false)
12629 ); assert_eq!(r(Single, vec![Some(true), Option::None]), Option::None);
12631 }
12632
12633 #[test]
12634 fn invariant_quantifier_truth_matrix_preserves_cardinality_and_nulls() {
12635 use graphforge_ir::QuantifierKind::{All, Any, None as NoneQ, Single};
12636
12637 for kind in [All, Any, NoneQ, Single] {
12638 for predicate in [Some(true), Some(false), Option::None] {
12639 for length in [0, 1, 4] {
12640 assert_eq!(
12641 reduce_invariant_quantifier(kind, predicate, length),
12642 match predicate {
12643 Some(value) => {
12644 let values = datafusion::arrow::array::BooleanArray::from(vec![
12645 value;
12646 length
12647 ]);
12648 reduce_quantifier(kind, &values, length)
12649 }
12650 Option::None => {
12651 let values =
12652 datafusion::arrow::array::BooleanArray::new_null(length);
12653 reduce_quantifier(kind, &values, length)
12654 }
12655 },
12656 "{kind:?}, predicate={predicate:?}, length={length}"
12657 );
12658 }
12659 }
12660 }
12661 }
12662
12663 #[test]
12664 fn invariant_quantifier_scaling_counts_rows_not_heterogeneous_elements() {
12665 use datafusion::arrow::array::{Array, ArrayRef, BooleanArray, ListArray};
12666 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
12667 use datafusion::arrow::datatypes::Field;
12668 use datafusion::config::ConfigOptions;
12669 use datafusion::scalar::ScalarValue as S;
12670 use graphforge_ir::QuantifierKind::{All, Any, None as NoneQ, Single};
12671 use std::sync::Arc;
12672
12673 let pattern = [
12674 S::Int64(Some(1)),
12675 S::Null,
12676 S::Boolean(Some(true)),
12677 S::Utf8(Some("x".to_owned())),
12678 ];
12679 let flat = pattern
12680 .iter()
12681 .cloned()
12682 .chain(pattern.iter().cloned().cycle().take(40))
12683 .collect::<Vec<_>>();
12684 let values: ArrayRef = Arc::new(build_het_struct(&flat, 0).unwrap());
12685 let item = Arc::new(Field::new("item", values.data_type().clone(), true));
12686 let list: ArrayRef = Arc::new(ListArray::new(
12687 item,
12688 OffsetBuffer::new(ScalarBuffer::from(vec![0i32, 4, 44, 44])),
12689 values,
12690 Some(NullBuffer::from(vec![true, true, false])),
12691 ));
12692 INVARIANT_QUANTIFIER_ROWS.store(0, Ordering::SeqCst);
12693
12694 for kind in [All, Any, NoneQ, Single] {
12695 for predicate in [Some(true), Some(false)] {
12696 let udf = CypherInvariantQuantifier::new(kind, predicate);
12697 let output = udf
12698 .invoke_with_args(ScalarFunctionArgs {
12699 args: vec![ColumnarValue::Array(Arc::clone(&list))],
12700 arg_fields: vec![Arc::new(Field::new(
12701 "list",
12702 list.data_type().clone(),
12703 true,
12704 ))],
12705 number_rows: 3,
12706 return_field: Arc::new(Field::new("out", DataType::Boolean, true)),
12707 config_options: Arc::new(ConfigOptions::default()),
12708 })
12709 .unwrap()
12710 .into_array(3)
12711 .unwrap();
12712 let output = output.as_any().downcast_ref::<BooleanArray>().unwrap();
12713 assert_eq!(
12714 output.value(0),
12715 reduce_invariant_quantifier(kind, predicate, 4).unwrap()
12716 );
12717 assert_eq!(
12718 output.value(1),
12719 reduce_invariant_quantifier(kind, predicate, 40).unwrap()
12720 );
12721 assert!(output.is_null(2));
12722 }
12723 }
12724 assert_eq!(
12725 INVARIANT_QUANTIFIER_ROWS.load(Ordering::SeqCst),
12726 4 * 2 * 3,
12727 "1x and 10x element counts must keep invariant work at one fold per list row"
12728 );
12729 }
12730
12731 #[test]
12732 fn invariant_quantifier_lowering_uses_cardinality_only_udf() {
12733 use graphforge_ir::QuantifierKind::None as NoneQ;
12734
12735 for (predicate, expected) in [
12736 (IrLiteral::Bool(true), Some(true)),
12737 (IrLiteral::Bool(false), Some(false)),
12738 (IrLiteral::Null, Option::None),
12739 ] {
12740 let mut arena = ExprArena::new();
12741 let list = arena.push(IrExpr::VarRef(VarId(0)));
12742 let predicate = arena.push(IrExpr::Literal(predicate));
12743 let quantifier = arena.push(IrExpr::Quantifier {
12744 kind: NoneQ,
12745 loop_var: VarId(1),
12746 list,
12747 predicate,
12748 });
12749 let mut vars = VarMap::new();
12750 vars.insert(VarId(0), "list");
12751 let lowered = make_lowerer(&arena, &vars).lower(quantifier).unwrap();
12752 let DfExpr::ScalarFunction(function) = lowered else {
12753 panic!("invariant quantifier must lower to a scalar UDF")
12754 };
12755 let invariant = function
12756 .func
12757 .inner()
12758 .as_any()
12759 .downcast_ref::<CypherInvariantQuantifier>()
12760 .expect("cardinality-only quantifier UDF");
12761 assert_eq!(invariant.predicate, expected);
12762 }
12763 }
12764
12765 #[test]
12766 fn uncorrelated_list_comprehension_batches_volatile_predicate_once() {
12767 use datafusion::arrow::array::{Array, Int64Builder, ListArray, ListBuilder};
12768 use datafusion::arrow::datatypes::Field;
12769 use datafusion::config::ConfigOptions;
12770 use std::sync::Arc;
12771
12772 VOLATILE_CALLS.store(0, Ordering::SeqCst);
12773 VOLATILE_ROWS.store(0, Ordering::SeqCst);
12774 let mut builder = ListBuilder::new(Int64Builder::new());
12775 builder.append_null();
12776 builder.append(true);
12777 builder.values().append_value(1);
12778 builder.values().append_value(2);
12779 builder.append(true);
12780 builder.values().append_value(3);
12781 builder.append(true);
12782 let input = Arc::new(builder.finish()) as datafusion::arrow::array::ArrayRef;
12783 let predicate = ScalarUDF::new_from_impl(CountingVolatilePredicate::new()).call(vec![]);
12784 let udf = CypherListComp::new(Some(predicate), None, "__gf_elem".to_owned(), vec![]);
12785 let return_type = udf.return_type(&[input.data_type().clone()]).unwrap();
12786 let output = udf
12787 .invoke_with_args(ScalarFunctionArgs {
12788 args: vec![ColumnarValue::Array(input)],
12789 arg_fields: vec![Arc::new(Field::new(
12790 "list",
12791 DataType::new_list(DataType::Int64, true),
12792 true,
12793 ))],
12794 number_rows: 4,
12795 return_field: Arc::new(Field::new("out", return_type, true)),
12796 config_options: Arc::new(ConfigOptions::default()),
12797 })
12798 .unwrap()
12799 .into_array(4)
12800 .unwrap();
12801 let output = output.as_any().downcast_ref::<ListArray>().unwrap();
12802
12803 assert!(output.is_null(0));
12804 assert_eq!(output.value(1).len(), 0);
12805 assert_eq!(output.value(2).len(), 2);
12806 assert_eq!(output.value(3).len(), 1);
12807 assert_eq!(VOLATILE_CALLS.load(Ordering::SeqCst), 1);
12808 assert_eq!(VOLATILE_ROWS.load(Ordering::SeqCst), 3);
12809 }
12810
12811 #[test]
12812 fn correlated_list_comprehension_filters_and_projects_with_outer_value() {
12813 use datafusion::arrow::array::{Array, ListArray};
12814 use datafusion::logical_expr::{col, lit};
12815
12816 let input = ScalarValue::List(ScalarValue::new_list(
12817 &[
12818 ScalarValue::Int64(Some(1)),
12819 ScalarValue::Int64(Some(3)),
12820 ScalarValue::Int64(Some(5)),
12821 ],
12822 &DataType::Int64,
12823 true,
12824 ));
12825 let udf = CypherListComp::new(
12826 Some(col("__gf_elem").gt(col("threshold"))),
12827 Some(col("__gf_elem") + col("threshold") + lit(0_i64)),
12828 "__gf_elem".into(),
12829 vec!["threshold".into()],
12830 );
12831 let output = invoke_test_udf(&udf, vec![input, ScalarValue::Int64(Some(2))]).unwrap();
12832 let output = output.as_any().downcast_ref::<ListArray>().expect("List");
12833 assert!(!output.is_null(0));
12834 let values = output.value(0);
12835 assert_eq!(
12836 (0..values.len())
12837 .map(|row| ScalarValue::try_from_array(&values, row).unwrap())
12838 .collect::<Vec<_>>(),
12839 vec![ScalarValue::Int64(Some(5)), ScalarValue::Int64(Some(7))]
12840 );
12841 }
12842
12843 #[test]
12844 fn percentile_and_comparison_error_contract_matrix_is_exact() {
12845 use datafusion::logical_expr::AggregateUDFImpl;
12846
12847 let continuous = CypherPercentile::new(true);
12848 assert_eq!(
12849 continuous.return_type(&[]).unwrap_err().to_string(),
12850 "Error during planning: percentile aggregate requires value and percentile arguments"
12851 );
12852 assert_eq!(
12853 continuous
12854 .return_type(&[DataType::Utf8, DataType::Float64])
12855 .unwrap_err()
12856 .to_string(),
12857 "Error during planning: percentile value expression must be numeric, got Utf8"
12858 );
12859 assert_eq!(
12860 continuous
12861 .return_type(&[DataType::Int32, DataType::Float64])
12862 .unwrap(),
12863 DataType::Float64
12864 );
12865 assert_eq!(
12866 CypherPercentile::new(false)
12867 .return_type(&[DataType::Int32, DataType::Float64])
12868 .unwrap(),
12869 DataType::Int32
12870 );
12871
12872 let mut accumulator = PercentileAcc {
12873 continuous: true,
12874 value_type: DataType::Int64,
12875 result_type: DataType::Float64,
12876 values: vec![],
12877 percentile: None,
12878 };
12879 accumulator.push_value(ScalarValue::Null).unwrap();
12880 assert_eq!(
12881 accumulator
12882 .push_value(ScalarValue::Utf8(Some("bad".into())))
12883 .unwrap_err()
12884 .to_string(),
12885 "Execution error: percentile value expression must be numeric, got Utf8"
12886 );
12887 for invalid in [f64::NAN, -0.1, 1.1] {
12888 assert!(
12889 accumulator
12890 .observe_percentile(Some(invalid))
12891 .unwrap_err()
12892 .to_string()
12893 .contains("finite number between 0.0 and 1.0")
12894 );
12895 }
12896 accumulator.observe_percentile(Some(0.25)).unwrap();
12897 assert_eq!(
12898 accumulator
12899 .observe_percentile(Some(0.75))
12900 .unwrap_err()
12901 .to_string(),
12902 "Execution error: percentile argument must be constant within an aggregate group"
12903 );
12904
12905 assert_eq!(scalar_as_i8(&ScalarValue::Int8(Some(-7))).unwrap(), -7);
12906 assert_eq!(scalar_as_i8(&ScalarValue::Int64(Some(7))).unwrap(), 7);
12907 assert!(
12908 scalar_as_i8(&ScalarValue::Int64(Some(128)))
12909 .unwrap_err()
12910 .to_string()
12911 .contains("outside i8 range")
12912 );
12913 assert_eq!(
12914 scalar_as_i8(&ScalarValue::Utf8(Some("1".into())))
12915 .unwrap_err()
12916 .to_string(),
12917 "Error during planning: comparison opcode must be an integer, got Utf8(\"1\")"
12918 );
12919 }
12920
12921 #[test]
12922 fn shared_temporal_cast_helper_preserves_arrow_values_nulls_and_error() {
12923 use datafusion::arrow::array::{Array, ArrayRef, Int32Array, Int64Array};
12924
12925 let integers: ArrayRef = Arc::new(Int32Array::from(vec![Some(7), None]));
12926 let casted = cast_argument_arrays(&[integers], &DataType::Int64).unwrap();
12927 let casted = casted[0]
12928 .as_any()
12929 .downcast_ref::<Int64Array>()
12930 .expect("Int64");
12931 assert_eq!(casted.value(0), 7);
12932 assert!(casted.is_null(1));
12933
12934 let invalid = ScalarValue::List(ScalarValue::new_list(
12935 &[ScalarValue::Int64(Some(1))],
12936 &DataType::Int64,
12937 true,
12938 ))
12939 .to_array_of_size(1)
12940 .unwrap();
12941 let direct_error = datafusion::arrow::compute::cast(&invalid, &DataType::Int64)
12942 .map_err(datafusion::error::DataFusionError::from)
12943 .unwrap_err()
12944 .to_string();
12945 let helper_error = cast_argument_arrays(&[invalid], &DataType::Int64)
12946 .unwrap_err()
12947 .to_string();
12948 assert_eq!(helper_error, direct_error);
12949 }
12950
12951 #[test]
12952 fn list_comprehension_rejects_non_list_input_through_public_udf_contract() {
12953 let udf = CypherListComp::new(None, None, "__gf_elem".into(), vec![]);
12954 let error = invoke_test_udf(&udf, vec![ScalarValue::Int64(Some(1))]).unwrap_err();
12955 let datafusion::error::DataFusionError::Internal(message) = error else {
12956 panic!("expected DataFusion internal contract error")
12957 };
12958 assert_eq!(
12959 message,
12960 "cypher_list_comprehension: first argument is not a list"
12961 );
12962 }
12963
12964 #[test]
12965 fn public_map_and_value_access_error_null_and_success_matrix() {
12966 use datafusion::arrow::array::{Array, ListArray};
12967
12968 let null_keys = invoke_test_udf(&CypherMapKeys::new(), vec![ScalarValue::Null]).unwrap();
12969 let null_keys = null_keys
12970 .as_any()
12971 .downcast_ref::<ListArray>()
12972 .expect("List");
12973 assert!(null_keys.is_null(0));
12974
12975 let keys_error =
12976 invoke_test_udf(&CypherMapKeys::new(), vec![ScalarValue::Int64(Some(1))]).unwrap_err();
12977 assert_eq!(
12978 keys_error.to_string(),
12979 "Execution error: keys() requires a map, node, relationship, or null, got Int64"
12980 );
12981
12982 let map = const_map_scalar(&[
12983 ("answer".into(), ScalarValue::Int64(Some(42))),
12984 ("empty".into(), ScalarValue::Null),
12985 ])
12986 .expect("map scalar");
12987 let keys = invoke_test_udf(&CypherMapKeys::new(), vec![map.clone()]).unwrap();
12988 let keys = keys.as_any().downcast_ref::<ListArray>().expect("List");
12989 assert_eq!(keys.value(0).len(), 2);
12990
12991 let answer = invoke_test_udf(
12992 &CypherStaticValueAccess::new("answer".into()),
12993 vec![map.clone()],
12994 )
12995 .unwrap();
12996 assert_eq!(
12997 ScalarValue::try_from_array(&answer, 0).unwrap(),
12998 ScalarValue::Int64(Some(42))
12999 );
13000 let missing =
13001 invoke_test_udf(&CypherStaticValueAccess::new("missing".into()), vec![map]).unwrap();
13002 assert!(ScalarValue::try_from_array(&missing, 0).unwrap().is_null());
13003
13004 let static_error = invoke_test_udf(
13005 &CypherStaticValueAccess::new("answer".into()),
13006 vec![ScalarValue::Int64(Some(1))],
13007 )
13008 .unwrap_err();
13009 assert_eq!(
13010 static_error.to_string(),
13011 "Execution error: InvalidArgumentValue: property access requires a map or graph element"
13012 );
13013
13014 let dynamic_error = invoke_test_udf(
13015 &CypherValueAccess::new(),
13016 vec![
13017 ScalarValue::Int64(Some(1)),
13018 ScalarValue::Utf8(Some("answer".into())),
13019 ],
13020 )
13021 .unwrap_err();
13022 assert_eq!(
13023 dynamic_error.to_string(),
13024 "Execution error: dynamic subscript requires a list or map/entity struct, got Int64"
13025 );
13026 }
13027
13028 #[test]
13029 fn entity_properties_and_percentile_state_validation_errors_are_exact() {
13030 use datafusion::arrow::array::{ArrayRef, Float64Array, Int64Array};
13031 use datafusion::logical_expr::Accumulator;
13032
13033 let arity_error = invoke_test_udf(&CypherEntityProperties::new(0), vec![]).unwrap_err();
13034 assert_eq!(
13035 arity_error.to_string(),
13036 "Error during planning: properties() entity map expects present plus key/value pairs"
13037 );
13038 let presence_error = invoke_test_udf(
13039 &CypherEntityProperties::new(3),
13040 vec![
13041 ScalarValue::Int64(Some(1)),
13042 ScalarValue::Utf8(Some("name".into())),
13043 ScalarValue::Utf8(Some("Ada".into())),
13044 ],
13045 )
13046 .unwrap_err();
13047 assert_eq!(
13048 presence_error.to_string(),
13049 "Execution error: properties() entity presence must be boolean, got Int64(1)"
13050 );
13051
13052 let mut accumulator = PercentileAcc {
13053 continuous: true,
13054 value_type: DataType::Int64,
13055 result_type: DataType::Float64,
13056 values: vec![],
13057 percentile: None,
13058 };
13059 let wrong_values: ArrayRef = Arc::new(Int64Array::from(vec![1]));
13060 let percentile: ArrayRef = Arc::new(Float64Array::from(vec![0.5]));
13061 assert_eq!(
13062 accumulator
13063 .merge_batch(&[wrong_values, percentile])
13064 .unwrap_err()
13065 .to_string(),
13066 "Error during planning: percentile state values must be a list"
13067 );
13068 }
13069
13070 fn invoke_cypher_quantifier(
13073 kind: graphforge_ir::QuantifierKind,
13074 predicate: DfExpr,
13075 list: datafusion::arrow::array::ArrayRef,
13076 ) -> datafusion::error::Result<datafusion::arrow::array::BooleanArray> {
13077 use datafusion::arrow::array::BooleanArray;
13078 use datafusion::arrow::datatypes::Field;
13079 use datafusion::config::ConfigOptions;
13080 use std::sync::Arc;
13081
13082 let n = list.len();
13083 let field = Arc::new(Field::new("l", list.data_type().clone(), true));
13084 let ret = Arc::new(Field::new("q", DataType::Boolean, true));
13085 let args = ScalarFunctionArgs {
13086 args: vec![ColumnarValue::Array(list)],
13087 arg_fields: vec![field],
13088 number_rows: n,
13089 return_field: ret,
13090 config_options: Arc::new(ConfigOptions::default()),
13091 };
13092 let udf = CypherQuantifier::new(kind, predicate, "__gf_elem".to_owned(), vec![]);
13093 let arr = match udf.invoke_with_args(args)? {
13094 ColumnarValue::Array(a) => a,
13095 ColumnarValue::Scalar(s) => s.to_array_of_size(n)?,
13096 };
13097 Ok(arr
13098 .as_any()
13099 .downcast_ref::<BooleanArray>()
13100 .expect("boolean verdicts")
13101 .clone())
13102 }
13103
13104 #[test]
13105 fn cypher_quantifier_empty_list_yields_identity_without_predicate() {
13106 use datafusion::arrow::array::{Array, ListArray};
13107 use datafusion::arrow::datatypes::Int64Type;
13108 use graphforge_ir::QuantifierKind::{All, Any, None as NoneQ, Single};
13109
13110 let empty = || {
13112 std::sync::Arc::new(ListArray::from_iter_primitive::<Int64Type, _, _>(vec![
13113 Some(Vec::<Option<i64>>::new()),
13114 ])) as datafusion::arrow::array::ArrayRef
13115 };
13116
13117 let field_pred =
13121 datafusion::functions::core::expr_fn::get_field(col_literal("__gf_elem"), "a")
13122 .eq(DfExpr::Literal(ScalarValue::Int64(Some(2)), Option::None));
13123 let elem_pred = col_literal("__gf_elem");
13124
13125 for (kind, expected) in [(All, true), (NoneQ, true), (Any, false), (Single, false)] {
13126 for pred in [field_pred.clone(), elem_pred.clone()] {
13127 let out = invoke_cypher_quantifier(kind, pred, empty()).expect("no error");
13128 assert!(!out.is_null(0), "{kind:?}: empty list must be definitive");
13129 assert_eq!(out.value(0), expected, "{kind:?} identity");
13130 }
13131 }
13132 }
13133
13134 #[test]
13135 fn cypher_quantifier_mixed_batch_and_unplannable_predicate() {
13136 use datafusion::arrow::array::{Array, BooleanBuilder, ListBuilder};
13137 use graphforge_ir::QuantifierKind::Any;
13138
13139 let mut b = ListBuilder::new(BooleanBuilder::new());
13141 b.append_null();
13142 b.append(true);
13143 b.values().append_value(true);
13144 b.values().append_value(false);
13145 b.append(true);
13146 let list = std::sync::Arc::new(b.finish()) as datafusion::arrow::array::ArrayRef;
13147
13148 let out = invoke_cypher_quantifier(Any, col_literal("__gf_elem"), list.clone())
13150 .expect("no error");
13151 assert!(out.is_null(0), "null list → null");
13152 assert!(!out.is_null(1) && !out.value(1), "any over [] is false");
13153 assert!(out.value(2), "any over [true, false] is true");
13154
13155 let bad = datafusion::functions::core::expr_fn::get_field(col_literal("__gf_elem"), "a");
13157 assert!(
13158 invoke_cypher_quantifier(Any, bad, list).is_err(),
13159 "non-empty row against an unbuildable predicate must error"
13160 );
13161 }
13162
13163 fn make_lowerer<'a>(arena: &'a ExprArena, var_map: &'a VarMap) -> ExprLowerer<'a> {
13165 ExprLowerer::new(arena, None, var_map)
13166 }
13167
13168 #[test]
13173 fn boolean_and_numeric_operators_reject_known_bad_operands() {
13174 use graphforge_ir::expr::{BinaryOpKind as B, IrExpr, IrLiteral, UnaryOpKind as U};
13175 let vm = VarMap::new();
13176
13177 let reject = |build: &dyn Fn(&mut ExprArena) -> ExprId| {
13179 let mut a = ExprArena::new();
13180 let id = build(&mut a);
13181 let err = make_lowerer(&a, &vm).lower(id).expect_err("should reject");
13182 assert!(
13183 matches!(err, LoweringError::InvalidType(_)),
13184 "expected InvalidType, got {err:?}"
13185 );
13186 };
13187 reject(&|a| {
13189 let l = a.push(IrExpr::Literal(IrLiteral::Int(1)));
13190 let r = a.push(IrExpr::Literal(IrLiteral::Bool(true)));
13191 a.push(IrExpr::BinaryOp {
13192 op: B::And,
13193 left: l,
13194 right: r,
13195 })
13196 });
13197 reject(&|a| {
13199 let l = a.push(IrExpr::Literal(IrLiteral::Int(1)));
13200 let r = a.push(IrExpr::Literal(IrLiteral::Bool(true)));
13201 a.push(IrExpr::BinaryOp {
13202 op: B::Xor,
13203 left: l,
13204 right: r,
13205 })
13206 });
13207 reject(&|a| {
13209 let e = a.push(IrExpr::Literal(IrLiteral::Str("x".into())));
13210 a.push(IrExpr::UnaryOp {
13211 op: U::Not,
13212 expr: e,
13213 })
13214 });
13215 reject(&|a| {
13217 let e = a.push(IrExpr::Literal(IrLiteral::Bool(true)));
13218 a.push(IrExpr::UnaryOp {
13219 op: U::Neg,
13220 expr: e,
13221 })
13222 });
13223 reject(&|a| {
13225 let l = a.push(IrExpr::Literal(IrLiteral::Str("a".into())));
13226 let r = a.push(IrExpr::Literal(IrLiteral::Int(2)));
13227 a.push(IrExpr::BinaryOp {
13228 op: B::Mod,
13229 left: l,
13230 right: r,
13231 })
13232 });
13233 reject(&|a| {
13235 let v = a.push(IrExpr::Literal(IrLiteral::Int(1)));
13236 let m = a.push(IrExpr::MapLiteral(vec![("k".into(), v)]));
13237 a.push(IrExpr::UnaryOp {
13238 op: U::Not,
13239 expr: m,
13240 })
13241 });
13242 }
13243
13244 #[test]
13245 fn operator_type_check_accepts_valid_and_unknown_operands() {
13246 use graphforge_ir::VarId;
13247 use graphforge_ir::expr::{BinaryOpKind as B, IrExpr, IrLiteral, UnaryOpKind as U};
13248 let accept = |vm: &VarMap, build: &dyn Fn(&mut ExprArena) -> ExprId| {
13249 let mut a = ExprArena::new();
13250 let id = build(&mut a);
13251 make_lowerer(&a, vm)
13252 .lower(id)
13253 .expect("valid/unknown operands must lower cleanly");
13254 };
13255 accept(&VarMap::new(), &|a| {
13257 let l = a.push(IrExpr::Literal(IrLiteral::Bool(true)));
13258 let r = a.push(IrExpr::Literal(IrLiteral::Bool(false)));
13259 a.push(IrExpr::BinaryOp {
13260 op: B::And,
13261 left: l,
13262 right: r,
13263 })
13264 });
13265 accept(&VarMap::new(), &|a| {
13267 let l = a.push(IrExpr::Literal(IrLiteral::Null));
13268 let r = a.push(IrExpr::Literal(IrLiteral::Bool(true)));
13269 a.push(IrExpr::BinaryOp {
13270 op: B::And,
13271 left: l,
13272 right: r,
13273 })
13274 });
13275 let mut vm = VarMap::new();
13277 vm.insert(VarId(0), "x");
13278 accept(&vm, &|a| {
13279 let e = a.push(IrExpr::VarRef(VarId(0)));
13280 a.push(IrExpr::UnaryOp {
13281 op: U::Not,
13282 expr: e,
13283 })
13284 });
13285 }
13286
13287 #[test]
13288 fn nested_quantifiers_get_per_depth_elem_columns() {
13289 use graphforge_ir::QuantifierKind::{None as NoneQ, Single};
13290
13291 let mut arena = ExprArena::new();
13295 let list_outer = arena.push(IrExpr::VarRef(VarId(0)));
13296 let list_inner = arena.push(IrExpr::VarRef(VarId(0)));
13297 let x = arena.push(IrExpr::VarRef(VarId(1)));
13298 let y = arena.push(IrExpr::VarRef(VarId(2)));
13299 let sum = arena.push(IrExpr::BinaryOp {
13300 op: BinaryOpKind::Add,
13301 left: x,
13302 right: y,
13303 });
13304 let fifteen = arena.push(IrExpr::Literal(IrLiteral::Int(15)));
13305 let eq = arena.push(IrExpr::BinaryOp {
13306 op: BinaryOpKind::Eq,
13307 left: sum,
13308 right: fifteen,
13309 });
13310 let inner = arena.push(IrExpr::Quantifier {
13311 kind: Single,
13312 loop_var: VarId(2),
13313 list: list_inner,
13314 predicate: eq,
13315 });
13316 let outer = arena.push(IrExpr::Quantifier {
13317 kind: NoneQ,
13318 loop_var: VarId(1),
13319 list: list_outer,
13320 predicate: inner,
13321 });
13322
13323 let mut vm = VarMap::new();
13324 vm.insert(VarId(0), "list");
13325 let lowered = make_lowerer(&arena, &vm).lower(outer).expect("lower");
13326
13327 let as_quant = |e: &DfExpr| -> Option<(String, Vec<String>, Vec<DfExpr>)> {
13328 let DfExpr::ScalarFunction(f) = e else {
13329 return Option::None;
13330 };
13331 let q = f.func.inner().as_any().downcast_ref::<CypherQuantifier>()?;
13332 Some((q.elem_name.clone(), q.outer_names.clone(), f.args.clone()))
13333 };
13334
13335 let (outer_elem, outer_outers, _) = as_quant(&lowered).expect("outer quantifier UDF");
13338 assert_eq!(outer_elem, "__gf_elem");
13339 assert_eq!(outer_outers, vec!["list".to_owned()]);
13340
13341 let DfExpr::ScalarFunction(outer_fn) = &lowered else {
13344 panic!("outer is a scalar function")
13345 };
13346 let outer_q = outer_fn
13347 .func
13348 .inner()
13349 .as_any()
13350 .downcast_ref::<CypherQuantifier>()
13351 .unwrap();
13352 let (inner_elem, inner_outers, inner_args) =
13353 as_quant(&outer_q.predicate).expect("inner quantifier UDF");
13354 assert_eq!(inner_elem, "__gf_elem_1");
13355 assert_eq!(
13356 inner_outers,
13357 vec!["__gf_elem".to_owned()],
13358 "the outer element is an outer column of the inner UDF (the list \
13359 argument resolves in the enclosing batch, not via outer_names)"
13360 );
13361 let arg_names: Vec<String> = inner_args.iter().map(|a| a.to_string()).collect();
13363 assert_eq!(arg_names, vec!["list", "__gf_elem"]);
13364 }
13365
13366 #[test]
13367 fn hydrated_path_nodes_return_type_carries_labels_and_props() {
13368 use datafusion::arrow::datatypes::Field;
13369
13370 let bare = CypherPathNodes::new().return_type(&[]).unwrap();
13372 let DataType::List(item) = &bare else {
13373 panic!("list return, got {bare:?}")
13374 };
13375 let DataType::Struct(fields) = item.data_type() else {
13376 panic!("struct element")
13377 };
13378 assert_eq!(fields.len(), 1);
13379 assert_eq!(fields[0].name(), "node_uuid");
13380
13381 let hydrated = CypherPathNodes::with_hydration(PathNodeHydration {
13384 dir: std::path::PathBuf::from("/nonexistent"),
13385 labels_by_type: vec![(0, "A".to_owned())],
13386 prop_stems: vec!["_untyped".to_owned()],
13387 fields: vec![
13388 Field::new("node_uuid", DataType::FixedSizeBinary(16), false),
13389 Field::new("labels", DataType::new_list(DataType::Utf8, true), true),
13390 Field::new("name", DataType::Utf8, true),
13391 ]
13392 .into(),
13393 })
13394 .return_type(&[])
13395 .unwrap();
13396 let DataType::List(item) = &hydrated else {
13397 panic!("list return, got {hydrated:?}")
13398 };
13399 let DataType::Struct(fields) = item.data_type() else {
13400 panic!("struct element")
13401 };
13402 let names: Vec<&str> = fields.iter().map(|f| f.name().as_str()).collect();
13403 assert_eq!(names, vec!["node_uuid", "labels", "name"]);
13404 }
13405
13406 #[test]
13407 fn elem_struct_col_routes_property_to_get_field() {
13408 let mut arena = ExprArena::new();
13415 let base = arena.push(IrExpr::VarRef(VarId(0)));
13416 let access = arena.push(IrExpr::PropertyAccess {
13417 base,
13418 prop: PropId(0),
13419 });
13420 let mut vm = VarMap::new();
13421 vm.insert(VarId(0), "__gf_elem");
13422 let mut prop_names = HashMap::new();
13423 prop_names.insert(0u32, "a".to_owned());
13424
13425 let dotted = ExprLowerer::with_prop_names(&arena, &vm, prop_names.clone())
13429 .lower(access)
13430 .unwrap();
13431 assert!(
13432 matches!(&dotted, DfExpr::Column(_)) && dotted.to_string() == "__gf_elem.a",
13433 "expected dotted column, got {dotted:?}"
13434 );
13435
13436 let via_get_field = ExprLowerer::with_prop_names(&arena, &vm, prop_names.clone())
13439 .with_elem_struct_col("__gf_elem".to_owned())
13440 .lower(access)
13441 .unwrap();
13442 assert!(
13443 !matches!(&via_get_field, DfExpr::Column(_)),
13444 "element field access must not be a dotted column: {via_get_field:?}"
13445 );
13446 assert!(
13447 via_get_field.to_string().contains("get_field"),
13448 "expected a get_field call, got {via_get_field}"
13449 );
13450 }
13451
13452 #[test]
13453 fn map_column_field_access_uses_get_field_via_schema() {
13454 use datafusion::arrow::datatypes::{DataType, Field, Fields, Schema};
13459 use datafusion::common::DFSchema;
13460
13461 let map_ty = DataType::Struct(Fields::from(vec![
13462 Field::new("list", DataType::new_list(DataType::Int64, true), true),
13463 Field::new("fixed", DataType::Boolean, true),
13464 ]));
13465 let schema = Schema::new(vec![Field::new("input", map_ty, true)]);
13466 let df_schema = std::sync::Arc::new(DFSchema::try_from(schema).unwrap());
13467
13468 let mut arena = ExprArena::new();
13469 let base = arena.push(IrExpr::VarRef(VarId(0)));
13470 let access = arena.push(IrExpr::PropertyAccess {
13471 base,
13472 prop: PropId(0),
13473 });
13474 let mut vm = VarMap::new();
13475 vm.insert(VarId(0), "input");
13476 let mut prop_names = HashMap::new();
13477 prop_names.insert(0u32, "list".to_owned());
13478
13479 let out = ExprLowerer::with_prop_names(&arena, &vm, prop_names)
13480 .with_input_schema(df_schema)
13481 .lower(access)
13482 .unwrap();
13483 assert!(
13484 !matches!(&out, DfExpr::Column(_)),
13485 "map field must not be a dotted column: {out:?}"
13486 );
13487 assert!(
13488 out.to_string().contains("get_field"),
13489 "expected a get_field call, got {out}"
13490 );
13491 }
13492
13493 #[test]
13494 fn list_plus_uses_native_ops_for_homogeneous_schema_types() {
13495 use datafusion::arrow::datatypes::{DataType, Field, Schema};
13500 use datafusion::common::DFSchema;
13501 use graphforge_ir::expr::BinaryOpKind;
13502
13503 let schema = Schema::new(vec![
13504 Field::new("xs", DataType::new_list(DataType::Int64, true), true),
13505 Field::new("y", DataType::Int64, true),
13506 Field::new("ys", DataType::new_list(DataType::Int64, true), true),
13507 Field::new("s", DataType::Utf8, true),
13508 Field::new("ss", DataType::new_list(DataType::Utf8, true), true),
13509 ]);
13510 let df_schema = std::sync::Arc::new(DFSchema::try_from(schema).unwrap());
13511
13512 let mut arena = ExprArena::new();
13513 let xs = arena.push(IrExpr::VarRef(VarId(0)));
13514 let y = arena.push(IrExpr::VarRef(VarId(1)));
13515 let ys = arena.push(IrExpr::VarRef(VarId(2)));
13516 let s = arena.push(IrExpr::VarRef(VarId(3)));
13517 let ss = arena.push(IrExpr::VarRef(VarId(4)));
13518 let append = arena.push(IrExpr::BinaryOp {
13519 op: BinaryOpKind::Add,
13520 left: xs,
13521 right: y,
13522 });
13523 let concat = arena.push(IrExpr::BinaryOp {
13524 op: BinaryOpKind::Add,
13525 left: xs,
13526 right: ys,
13527 });
13528 let hetero_append = arena.push(IrExpr::BinaryOp {
13529 op: BinaryOpKind::Add,
13530 left: xs,
13531 right: s,
13532 });
13533 let hetero_concat = arena.push(IrExpr::BinaryOp {
13534 op: BinaryOpKind::Add,
13535 left: xs,
13536 right: ss,
13537 });
13538 let mut vm = VarMap::new();
13539 vm.insert(VarId(0), "xs");
13540 vm.insert(VarId(1), "y");
13541 vm.insert(VarId(2), "ys");
13542 vm.insert(VarId(3), "s");
13543 vm.insert(VarId(4), "ss");
13544
13545 let lowerer =
13546 ExprLowerer::with_prop_names(&arena, &vm, HashMap::new()).with_input_schema(df_schema);
13547 assert!(
13548 lowerer
13549 .lower(append)
13550 .unwrap()
13551 .to_string()
13552 .contains("array_append"),
13553 "list + element should append"
13554 );
13555 assert!(
13556 lowerer
13557 .lower(concat)
13558 .unwrap()
13559 .to_string()
13560 .contains("array_concat"),
13561 "list + list should concat"
13562 );
13563 assert!(
13564 lowerer
13565 .lower(hetero_append)
13566 .unwrap()
13567 .to_string()
13568 .contains("cypher_list_plus"),
13569 "list + heterogeneous element should use tagged list-plus"
13570 );
13571 assert!(
13572 lowerer
13573 .lower(hetero_concat)
13574 .unwrap()
13575 .to_string()
13576 .contains("cypher_list_plus"),
13577 "list + heterogeneous list should use tagged list-plus"
13578 );
13579 }
13580
13581 #[test]
13582 fn cypher_list_plus_concats_decoded_list_element() {
13583 use datafusion::arrow::array::{Array, ArrayRef, ListArray};
13584 use datafusion::arrow::datatypes::Field;
13585 use datafusion::config::ConfigOptions;
13586 use datafusion::scalar::ScalarValue as S;
13587 use std::sync::Arc;
13588
13589 let list = |items: Vec<S>| S::List(S::new_list(&items, &DataType::Int64, true));
13590 let left_items = vec![
13591 list(vec![S::Int64(Some(1))]),
13592 list(vec![S::Int64(Some(2)), S::Int64(Some(3))]),
13593 list(vec![S::Int64(Some(4)), S::Int64(Some(5))]),
13594 ];
13595 let left = S::List(S::new_list(&left_items, &left_items[0].data_type(), true))
13596 .to_array()
13597 .unwrap();
13598
13599 let right_list = list(vec![S::Int64(Some(8)), S::Int64(Some(9))]);
13600 let right: ArrayRef = Arc::new(build_het_struct(&[right_list], 1).unwrap());
13601 let udf = CypherListPlus::new();
13602 let arg_types = vec![left.data_type().clone(), right.data_type().clone()];
13603 let return_type = udf.return_type(&arg_types).unwrap();
13604 let args = ScalarFunctionArgs {
13605 args: vec![ColumnarValue::Array(left), ColumnarValue::Array(right)],
13606 arg_fields: vec![
13607 Arc::new(Field::new(
13608 "l",
13609 DataType::new_list(DataType::Int64, true),
13610 true,
13611 )),
13612 Arc::new(Field::new("r", DataType::Struct(het_fields(1)), true)),
13613 ],
13614 number_rows: 1,
13615 return_field: Arc::new(Field::new("out", return_type, true)),
13616 config_options: Arc::new(ConfigOptions::default()),
13617 };
13618 let out = match udf.invoke_with_args(args).unwrap() {
13619 ColumnarValue::Array(a) => a,
13620 ColumnarValue::Scalar(s) => s.to_array().unwrap(),
13621 };
13622 let out = out.as_any().downcast_ref::<ListArray>().unwrap();
13623 assert!(!out.is_null(0));
13624 let values = out.value(0);
13625 assert_eq!(values.len(), 5);
13626 assert_eq!(
13627 decode_het(&ScalarValue::try_from_array(&values, 3).unwrap()),
13628 Some(S::Int64(Some(8)))
13629 );
13630 assert_eq!(
13631 decode_het(&ScalarValue::try_from_array(&values, 4).unwrap()),
13632 Some(S::Int64(Some(9)))
13633 );
13634 }
13635
13636 #[test]
13637 fn tagged_list_element_plus_preserves_dynamic_concat_and_null_rows() {
13638 use datafusion::arrow::array::{Array, ArrayRef, ListArray};
13639 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
13640 use datafusion::arrow::datatypes::Field;
13641 use datafusion::scalar::ScalarValue as S;
13642 use std::sync::Arc;
13643
13644 let values: ArrayRef = Arc::new(
13645 build_het_struct(
13646 &[S::Int64(Some(1)), S::Boolean(Some(true)), S::Int64(Some(2))],
13647 1,
13648 )
13649 .unwrap(),
13650 );
13651 let item = Arc::new(Field::new("item", values.data_type().clone(), true));
13652 let left: ArrayRef = Arc::new(ListArray::new(
13653 item.clone(),
13654 OffsetBuffer::new(ScalarBuffer::from(vec![0i32, 2, 3, 3, 3])),
13655 values,
13656 Some(NullBuffer::from(vec![true, true, true, false])),
13657 ));
13658 let nested = S::List(S::new_list(
13659 &[S::Int64(Some(8)), S::Int64(Some(7))],
13660 &DataType::Int64,
13661 true,
13662 ));
13663 let right: ArrayRef = Arc::new(
13664 build_het_struct(&[S::Int64(Some(9)), nested, S::Null, S::Int64(Some(1))], 1).unwrap(),
13665 );
13666 let return_type = DataType::List(item);
13667 let output = invoke_tagged_list_element_plus(&left, &right, &return_type)
13668 .unwrap()
13669 .expect("tagged fast path");
13670 let output = output.as_any().downcast_ref::<ListArray>().unwrap();
13671 let decode_row = |row: usize| {
13672 let values = output.value(row);
13673 (0..values.len())
13674 .map(|index| {
13675 decode_het(&ScalarValue::try_from_array(&values, index).unwrap()).unwrap()
13676 })
13677 .collect::<Vec<_>>()
13678 };
13679
13680 assert_eq!(
13681 decode_row(0),
13682 vec![S::Int64(Some(1)), S::Boolean(Some(true)), S::Int64(Some(9))]
13683 );
13684 assert_eq!(
13685 decode_row(1),
13686 vec![S::Int64(Some(2)), S::Int64(Some(8)), S::Int64(Some(7))]
13687 );
13688 assert_eq!(decode_row(2), vec![S::Null]);
13689 assert!(output.is_null(3));
13690 }
13691
13692 #[test]
13693 fn cypher_conversions_decode_tagged_values() {
13694 use datafusion::arrow::array::{Array, ArrayRef, Float64Array, Int64Array, StringArray};
13695 use datafusion::arrow::datatypes::Field;
13696 use datafusion::config::ConfigOptions;
13697 use datafusion::scalar::ScalarValue as S;
13698 use std::sync::Arc;
13699
13700 let values: ArrayRef = Arc::new(
13701 build_het_struct(
13702 &[
13703 S::Int64(Some(2)),
13704 S::Float64(Some(2.9)),
13705 S::Utf8(Some("foo".to_owned())),
13706 ],
13707 0,
13708 )
13709 .unwrap(),
13710 );
13711 let invoke = |kind: CypherConversionKind| {
13712 let udf = CypherConversion::new(kind);
13713 let return_type = udf.return_type(&[values.data_type().clone()]).unwrap();
13714 let args = ScalarFunctionArgs {
13715 args: vec![ColumnarValue::Array(Arc::clone(&values))],
13716 arg_fields: vec![Arc::new(Field::new("v", values.data_type().clone(), true))],
13717 number_rows: values.len(),
13718 return_field: Arc::new(Field::new("out", return_type, true)),
13719 config_options: Arc::new(ConfigOptions::default()),
13720 };
13721 match udf.invoke_with_args(args).unwrap() {
13722 ColumnarValue::Array(a) => a,
13723 ColumnarValue::Scalar(s) => s.to_array_of_size(values.len()).unwrap(),
13724 }
13725 };
13726
13727 let ints = invoke(CypherConversionKind::Integer);
13728 let ints = ints.as_any().downcast_ref::<Int64Array>().unwrap();
13729 assert_eq!(ints.value(0), 2);
13730 assert_eq!(ints.value(1), 2);
13731 assert!(ints.is_null(2));
13732
13733 let floats = invoke(CypherConversionKind::Float);
13734 let floats = floats.as_any().downcast_ref::<Float64Array>().unwrap();
13735 assert_eq!(floats.value(0), 2.0);
13736 assert_eq!(floats.value(1), 2.9);
13737 assert!(floats.is_null(2));
13738
13739 let strings = match CYPHER_TO_STRING
13740 .invoke_with_args(ScalarFunctionArgs {
13741 args: vec![ColumnarValue::Array(values)],
13742 arg_fields: vec![Arc::new(Field::new(
13743 "v",
13744 DataType::Struct(het_fields(0)),
13745 true,
13746 ))],
13747 number_rows: 3,
13748 return_field: Arc::new(Field::new("out", DataType::Utf8, true)),
13749 config_options: Arc::new(ConfigOptions::default()),
13750 })
13751 .unwrap()
13752 {
13753 ColumnarValue::Array(a) => a,
13754 ColumnarValue::Scalar(s) => s.to_array_of_size(3).unwrap(),
13755 };
13756 let strings = strings.as_any().downcast_ref::<StringArray>().unwrap();
13757 assert_eq!(strings.value(0), "2");
13758 assert_eq!(strings.value(1), "2.9");
13759 assert_eq!(strings.value(2), "foo");
13760 }
13761
13762 #[test]
13763 fn het_list_map_roundtrip_and_order() {
13764 use datafusion::scalar::ScalarValue as S;
13765 let map = const_map_scalar(&[
13767 ("a".to_owned(), S::Int64(Some(2))),
13768 ("b".to_owned(), S::Boolean(Some(true))),
13769 ])
13770 .expect("map scalar");
13771 let S::Struct(m) = &map else {
13772 panic!("map is a struct")
13773 };
13774 assert!(is_plain_map_struct(m));
13775 let S::Struct(d) = date_scalar(Some(0)) else {
13776 panic!("date is a struct")
13777 };
13778 assert!(!is_plain_map_struct(&d));
13779
13780 let scalars = vec![S::Int64(Some(1)), map.clone()];
13783 assert_eq!(het_depth(&scalars[0]), Some(0));
13784 assert_eq!(het_depth(&scalars[1]), Some(1));
13785 let depth = scalars.iter().filter_map(het_depth).max().unwrap();
13786 let elem = build_het_struct(&scalars, depth).expect("build tagged struct");
13787 let e0 = ScalarValue::try_from_array(&elem, 0).unwrap();
13788 assert_eq!(decode_het(&e0), Some(S::Int64(Some(1))));
13789 let e1 = ScalarValue::try_from_array(&elem, 1).unwrap();
13790 let decoded = decode_het(&e1).expect("decode map element");
13791 assert_eq!(cypher_value_eq(&decoded, &map), Some(true));
13792
13793 let m1 = const_map_scalar(&[("a".to_owned(), S::Int64(Some(1)))]).unwrap();
13796 let m2 = const_map_scalar(&[("a".to_owned(), S::Int64(Some(2)))]).unwrap();
13797 assert_eq!(cypher_order(&m1, &m2), std::cmp::Ordering::Less);
13798 assert_eq!(
13799 cypher_order(&S::Int64(Some(99)), &m1),
13800 std::cmp::Ordering::Less
13801 );
13802 }
13803
13804 #[test]
13805 fn empty_map_scalar_has_one_row_and_no_fields() {
13806 use datafusion::arrow::array::Array;
13807 use datafusion::scalar::ScalarValue as S;
13808
13809 let map = const_map_scalar(&[]).expect("empty map scalar");
13810 let S::Struct(values) = map else {
13811 panic!("empty map is a struct")
13812 };
13813 assert_eq!(values.len(), 1);
13814 assert_eq!(values.num_columns(), 0);
13815 assert!(is_plain_map_struct(&values));
13816
13817 let scalars = vec![S::Int64(Some(1)), S::Struct(values.clone())];
13818 let encoded = build_het_struct(&scalars, 1).expect("encode empty map");
13819 let tagged = S::try_from_array(&encoded, 1).expect("tagged empty map");
13820 let S::Struct(decoded) = decode_het(&tagged).expect("decode empty map") else {
13821 panic!("decoded empty map is a struct")
13822 };
13823 assert_eq!(decoded.len(), 1);
13824 assert_eq!(decoded.num_columns(), 0);
13825 }
13826
13827 #[test]
13828 fn dictionary_scalar_is_normalized_for_heterogeneous_encoding() {
13829 use datafusion::arrow::datatypes::DataType;
13830 use datafusion::scalar::ScalarValue as S;
13831
13832 let dictionary = S::Dictionary(
13833 Box::new(DataType::Int32),
13834 Box::new(S::Utf8(Some("value".to_owned()))),
13835 );
13836 let scalars = vec![S::Int64(Some(1)), dictionary];
13837 assert_eq!(het_depth(&scalars[1]), Some(0));
13838
13839 let encoded = build_het_struct(&scalars, 0).expect("encode dictionary scalar");
13840 let tagged = S::try_from_array(&encoded, 1).expect("tagged dictionary scalar");
13841 assert_eq!(decode_het(&tagged), Some(S::Utf8(Some("value".to_owned()))));
13842 }
13843
13844 #[test]
13845 fn cypher_value_eq_scalars() {
13846 use datafusion::scalar::ScalarValue as S;
13847 let i = |n| S::Int64(Some(n));
13848
13849 assert_eq!(cypher_value_eq(&i(1), &i(1)), Some(true));
13850 assert_eq!(cypher_value_eq(&i(1), &i(2)), Some(false));
13851 assert_eq!(
13853 cypher_value_eq(&i(1), &S::Utf8(Some("1".into()))),
13854 Some(false)
13855 );
13856 assert_eq!(cypher_value_eq(&i(1), &S::Float64(Some(1.0))), Some(true));
13858 assert_eq!(cypher_value_eq(&i(1), &S::UInt64(Some(1))), Some(true));
13859 assert_eq!(
13860 cypher_value_eq(&S::Float64(Some(f64::NAN)), &S::Float64(Some(f64::NAN))),
13861 Some(false)
13862 );
13863 assert_eq!(cypher_value_eq(&i(1), &S::Int64(None)), None);
13865 assert_eq!(cypher_value_eq(&S::Null, &i(1)), None);
13866 }
13867
13868 #[test]
13869 fn cypher_comparison_predicate_handles_nan_and_cross_type() {
13870 use datafusion::scalar::ScalarValue as S;
13871
13872 assert_eq!(
13873 cypher_compare_pred(&S::Float64(Some(f64::NAN)), &S::Int64(Some(1)), 2),
13874 Some(false)
13875 );
13876 assert_eq!(
13877 cypher_compare_pred(&S::Utf8(Some("1".to_owned())), &S::Int64(Some(1)), 0,),
13878 None
13879 );
13880 assert_eq!(
13881 cypher_compare_pred(&S::Int64(Some(1)), &S::Float64(Some(2.0)), 0),
13882 Some(true)
13883 );
13884 assert_eq!(
13885 cypher_compare_pred(&S::UInt64(Some(1)), &S::Int64(Some(2)), 0),
13886 Some(true)
13887 );
13888 }
13889
13890 #[test]
13895 fn literal_null() {
13896 let mut arena = ExprArena::new();
13897 let id = arena.push(IrExpr::Literal(IrLiteral::Null));
13898 let vm = VarMap::new();
13899 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13900 assert!(matches!(result, DfExpr::Literal(ScalarValue::Null, _)));
13901 }
13902
13903 #[test]
13904 fn literal_bool() {
13905 let mut arena = ExprArena::new();
13906 let id = arena.push(IrExpr::Literal(IrLiteral::Bool(true)));
13907 let vm = VarMap::new();
13908 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13909 assert!(matches!(
13910 result,
13911 DfExpr::Literal(ScalarValue::Boolean(Some(true)), _)
13912 ));
13913 }
13914
13915 #[test]
13916 fn literal_int() {
13917 let mut arena = ExprArena::new();
13918 let id = arena.push(IrExpr::Literal(IrLiteral::Int(42)));
13919 let vm = VarMap::new();
13920 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13921 assert!(matches!(
13922 result,
13923 DfExpr::Literal(ScalarValue::Int64(Some(42)), _)
13924 ));
13925 }
13926
13927 #[test]
13928 fn literal_float() {
13929 let mut arena = ExprArena::new();
13930 let id = arena.push(IrExpr::Literal(IrLiteral::Float(2.71)));
13931 let vm = VarMap::new();
13932 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13933 assert!(matches!(
13934 result,
13935 DfExpr::Literal(ScalarValue::Float64(Some(_)), _)
13936 ));
13937 }
13938
13939 #[test]
13940 fn literal_str() {
13941 let mut arena = ExprArena::new();
13942 let id = arena.push(IrExpr::Literal(IrLiteral::Str("hello".into())));
13943 let vm = VarMap::new();
13944 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13945 assert!(matches!(
13946 result,
13947 DfExpr::Literal(ScalarValue::Utf8(Some(_)), _)
13948 ));
13949 }
13950
13951 #[test]
13952 fn literal_duration() {
13953 let mut arena = ExprArena::new();
13954 let id = arena.push(IrExpr::Literal(IrLiteral::Duration {
13955 months: 0,
13956 days: 0,
13957 seconds: 1,
13958 nanos: 0,
13959 }));
13960 let vm = VarMap::new();
13961 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13962 assert!(matches!(result, DfExpr::Literal(ScalarValue::Struct(_), _)));
13963 }
13964
13965 #[test]
13966 fn literal_datetime() {
13967 let mut arena = ExprArena::new();
13968 let id = arena.push(IrExpr::Literal(IrLiteral::DateTime(0)));
13969 let vm = VarMap::new();
13970 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13971 assert!(matches!(
13972 result,
13973 DfExpr::Literal(ScalarValue::TimestampMicrosecond(Some(0), Some(_)), _)
13974 ));
13975 }
13976
13977 #[test]
13982 fn var_ref_bound() {
13983 let mut arena = ExprArena::new();
13984 let id = arena.push(IrExpr::VarRef(VarId(0)));
13985 let mut vm = VarMap::new();
13986 vm.insert(VarId(0), "node_id");
13987 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
13988 assert!(matches!(result, DfExpr::Column(_)));
13989 if let DfExpr::Column(col) = result {
13990 assert_eq!(col.name, "node_id");
13991 }
13992 }
13993
13994 #[test]
13995 fn var_ref_unbound_returns_error() {
13996 let mut arena = ExprArena::new();
13997 let id = arena.push(IrExpr::VarRef(VarId(99)));
13998 let vm = VarMap::new();
13999 let result = make_lowerer(&arena, &vm).lower(id);
14000 assert!(matches!(result, Err(LoweringError::UnboundVar(99))));
14001 }
14002
14003 #[test]
14008 fn unary_not() {
14009 let mut arena = ExprArena::new();
14010 let inner = arena.push(IrExpr::Literal(IrLiteral::Bool(true)));
14011 let id = arena.push(IrExpr::UnaryOp {
14012 op: UnaryOpKind::Not,
14013 expr: inner,
14014 });
14015 let vm = VarMap::new();
14016 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14017 assert!(matches!(result, DfExpr::Not(_)));
14018 }
14019
14020 #[test]
14021 fn unary_is_null() {
14022 let mut arena = ExprArena::new();
14023 let mut vm = VarMap::new();
14024 vm.insert(VarId(0), "x");
14025 let inner = arena.push(IrExpr::VarRef(VarId(0)));
14026 let id = arena.push(IrExpr::UnaryOp {
14027 op: UnaryOpKind::IsNull,
14028 expr: inner,
14029 });
14030 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14031 assert!(matches!(result, DfExpr::IsNull(_)));
14032 }
14033
14034 #[test]
14035 fn unary_is_not_null() {
14036 let mut arena = ExprArena::new();
14037 let mut vm = VarMap::new();
14038 vm.insert(VarId(0), "x");
14039 let inner = arena.push(IrExpr::VarRef(VarId(0)));
14040 let id = arena.push(IrExpr::UnaryOp {
14041 op: UnaryOpKind::IsNotNull,
14042 expr: inner,
14043 });
14044 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14045 assert!(matches!(result, DfExpr::IsNotNull(_)));
14046 }
14047
14048 #[test]
14053 fn binary_eq() {
14054 let mut arena = ExprArena::new();
14055 let l = arena.push(IrExpr::Literal(IrLiteral::Int(1)));
14056 let r = arena.push(IrExpr::Literal(IrLiteral::Int(2)));
14057 let id = arena.push(IrExpr::BinaryOp {
14058 op: BinaryOpKind::Eq,
14059 left: l,
14060 right: r,
14061 });
14062 let vm = VarMap::new();
14063 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14064 let DfExpr::ScalarFunction(sf) = result else {
14067 panic!("expected a cypher_eq scalar-function call, got {result:?}");
14068 };
14069 assert_eq!(sf.func.name(), "cypher_eq");
14070 assert_eq!(sf.args.len(), 2);
14071 }
14072
14073 #[test]
14074 fn binary_in_list() {
14075 let mut arena = ExprArena::new();
14076 let mut vm = VarMap::new();
14077 vm.insert(VarId(0), "x");
14078 let l = arena.push(IrExpr::VarRef(VarId(0)));
14079 let one = arena.push(IrExpr::Literal(IrLiteral::Int(1)));
14080 let two = arena.push(IrExpr::Literal(IrLiteral::Int(2)));
14081 let list = arena.push(IrExpr::ListLiteral(vec![one, two]));
14082 let id = arena.push(IrExpr::BinaryOp {
14083 op: BinaryOpKind::In,
14084 left: l,
14085 right: list,
14086 });
14087 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14088 let DfExpr::ScalarFunction(sf) = result else {
14092 panic!("expected a cypher_in scalar-function call, got {result:?}");
14093 };
14094 assert_eq!(sf.func.name(), "cypher_in");
14095 assert_eq!(sf.args.len(), 2);
14096 }
14097
14098 fn label_membership_expr(arena: &mut ExprArena, left: IrExpr, node_var: VarId) -> ExprId {
14099 let left = arena.push(left);
14100 let node = arena.push(IrExpr::VarRef(node_var));
14101 let labels = arena.push(IrExpr::FunctionCall {
14102 name: "labels".into(),
14103 args: vec![node],
14104 });
14105 arena.push(IrExpr::BinaryOp {
14106 op: BinaryOpKind::In,
14107 left,
14108 right: labels,
14109 })
14110 }
14111
14112 fn label_lowerer<'a>(arena: &'a ExprArena, var_map: &'a VarMap) -> ExprLowerer<'a> {
14113 ExprLowerer::with_prop_names_and_nodes(
14114 arena,
14115 var_map,
14116 HashMap::new(),
14117 HashMap::from([(0, NodeShape { prop_names: vec![] })]),
14118 HashMap::from([(7, "Known".to_owned())]),
14119 false,
14120 )
14121 }
14122
14123 #[test]
14124 fn known_literal_in_labels_lowers_to_direct_type_id_membership() {
14125 let mut arena = ExprArena::new();
14126 let id = label_membership_expr(
14127 &mut arena,
14128 IrExpr::Literal(IrLiteral::Str("Known".into())),
14129 VarId(0),
14130 );
14131 let mut vm = VarMap::new();
14132 vm.insert(VarId(0), "var_0");
14133
14134 let result = label_lowerer(&arena, &vm).lower(id).unwrap();
14135 let DfExpr::ScalarFunction(sf) = &result else {
14136 panic!("expected array_has scalar function, got {result:?}");
14137 };
14138 assert_eq!(sf.func.name(), "array_has");
14139 let rendered = result.to_string();
14140 assert!(rendered.contains("var_0.type_ids"));
14141 assert!(rendered.contains("UInt32(7)"));
14142 assert!(!rendered.contains("cypher_in"));
14143 assert!(!rendered.contains("array_concat"));
14144 }
14145
14146 #[test]
14147 fn unknown_literal_in_labels_retains_generic_membership() {
14148 let mut arena = ExprArena::new();
14149 let id = label_membership_expr(
14150 &mut arena,
14151 IrExpr::Literal(IrLiteral::Str("Unknown".into())),
14152 VarId(0),
14153 );
14154 let mut vm = VarMap::new();
14155 vm.insert(VarId(0), "var_0");
14156
14157 let result = label_lowerer(&arena, &vm).lower(id).unwrap();
14158 let DfExpr::ScalarFunction(sf) = result else {
14159 panic!("expected cypher_in scalar function");
14160 };
14161 assert_eq!(sf.func.name(), "cypher_in");
14162 }
14163
14164 #[test]
14165 fn dynamic_in_labels_retains_generic_membership() {
14166 let mut arena = ExprArena::new();
14167 let id = label_membership_expr(&mut arena, IrExpr::VarRef(VarId(1)), VarId(0));
14168 let mut vm = VarMap::new();
14169 vm.insert(VarId(0), "var_0");
14170 vm.insert(VarId(1), "label_name");
14171
14172 let result = label_lowerer(&arena, &vm).lower(id).unwrap();
14173 let DfExpr::ScalarFunction(sf) = result else {
14174 panic!("expected cypher_in scalar function");
14175 };
14176 assert_eq!(sf.func.name(), "cypher_in");
14177 }
14178
14179 #[test]
14184 fn compound_predicate() {
14185 let mut arena = ExprArena::new();
14186 let mut vm = VarMap::new();
14187 vm.insert(VarId(0), "a.age"); vm.insert(VarId(1), "b.name"); let a_age = arena.push(IrExpr::VarRef(VarId(0)));
14191 let thirty = arena.push(IrExpr::Literal(IrLiteral::Int(30)));
14192 let gt = arena.push(IrExpr::BinaryOp {
14193 op: BinaryOpKind::Gt,
14194 left: a_age,
14195 right: thirty,
14196 });
14197
14198 let b_name = arena.push(IrExpr::VarRef(VarId(1)));
14199 let param = arena.push(IrExpr::Parameter("name".into()));
14200 let eq = arena.push(IrExpr::BinaryOp {
14201 op: BinaryOpKind::Eq,
14202 left: b_name,
14203 right: param,
14204 });
14205
14206 let and = arena.push(IrExpr::BinaryOp {
14207 op: BinaryOpKind::And,
14208 left: gt,
14209 right: eq,
14210 });
14211
14212 let result = make_lowerer(&arena, &vm).lower(and).unwrap();
14213 assert!(matches!(result, DfExpr::ScalarFunction(_)));
14214 if let DfExpr::ScalarFunction(sf) = result {
14215 assert_eq!(sf.func.name(), "cypher_and");
14216 }
14217 }
14218
14219 #[test]
14220 fn xor_chain_lowering_has_one_udf_per_source_operator() {
14221 use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion};
14222
14223 let lower_and_count = |operands: usize| {
14224 let mut arena = ExprArena::new();
14225 let mut root = arena.push(IrExpr::Literal(IrLiteral::Bool(true)));
14226 for index in 1..operands {
14227 let right = match index % 3 {
14228 0 => IrLiteral::Bool(true),
14229 1 => IrLiteral::Bool(false),
14230 _ => IrLiteral::Null,
14231 };
14232 let right = arena.push(IrExpr::Literal(right));
14233 root = arena.push(IrExpr::BinaryOp {
14234 op: BinaryOpKind::Xor,
14235 left: root,
14236 right,
14237 });
14238 }
14239
14240 let lowered = make_lowerer(&arena, &VarMap::new())
14241 .lower(root)
14242 .expect("XOR chain lowers");
14243 let mut udf_count = 0;
14244 lowered
14245 .apply(|expr| {
14246 if matches!(
14247 expr,
14248 DfExpr::ScalarFunction(function)
14249 if function.func.name() == "cypher_xor"
14250 ) {
14251 udf_count += 1;
14252 }
14253 Ok(TreeNodeRecursion::Continue)
14254 })
14255 .expect("expression traversal succeeds");
14256 udf_count
14257 };
14258
14259 let eleven = lower_and_count(11);
14260 let twenty_two = lower_and_count(22);
14261 assert_eq!(eleven, 10);
14262 assert_eq!(twenty_two, 21);
14263 assert!(
14264 twenty_two <= eleven * 3,
14265 "doubling operands must keep deterministic lowering work within 3x"
14266 );
14267 }
14268
14269 #[test]
14270 fn cypher_xor_implements_three_valued_truth_table() {
14271 use datafusion::arrow::array::{Array, BooleanArray};
14272 use datafusion::arrow::datatypes::Field;
14273 use datafusion::config::ConfigOptions;
14274
14275 let left = BooleanArray::from(vec![
14276 Some(false),
14277 Some(false),
14278 Some(false),
14279 Some(true),
14280 Some(true),
14281 Some(true),
14282 None,
14283 None,
14284 None,
14285 ]);
14286 let right = BooleanArray::from(vec![
14287 Some(false),
14288 Some(true),
14289 None,
14290 Some(false),
14291 Some(true),
14292 None,
14293 Some(false),
14294 Some(true),
14295 None,
14296 ]);
14297 let expected = [
14298 Some(false),
14299 Some(true),
14300 None,
14301 Some(true),
14302 Some(false),
14303 None,
14304 None,
14305 None,
14306 None,
14307 ];
14308 let field = Arc::new(Field::new("value", DataType::Boolean, true));
14309 let arguments = ScalarFunctionArgs {
14310 args: vec![
14311 ColumnarValue::Array(Arc::new(left)),
14312 ColumnarValue::Array(Arc::new(right)),
14313 ],
14314 arg_fields: vec![Arc::clone(&field), field],
14315 number_rows: expected.len(),
14316 return_field: Arc::new(Field::new("xor", DataType::Boolean, true)),
14317 config_options: Arc::new(ConfigOptions::default()),
14318 };
14319 let result = CypherBoolOp::new(CypherBoolOpKind::Xor)
14320 .invoke_with_args(arguments)
14321 .expect("XOR evaluates");
14322 let ColumnarValue::Array(result) = result else {
14323 panic!("array inputs must produce an array")
14324 };
14325 let result = result
14326 .as_any()
14327 .downcast_ref::<BooleanArray>()
14328 .expect("XOR result is boolean");
14329
14330 for (index, expected) in expected.into_iter().enumerate() {
14331 let actual = (!result.is_null(index)).then(|| result.value(index));
14332 assert_eq!(actual, expected, "truth-table row {index}");
14333 }
14334 }
14335
14336 #[test]
14341 fn function_call_to_upper() {
14342 let mut arena = ExprArena::new();
14343 let mut vm = VarMap::new();
14344 vm.insert(VarId(0), "n.name");
14345 let arg = arena.push(IrExpr::VarRef(VarId(0)));
14346 let id = arena.push(IrExpr::FunctionCall {
14347 name: "toUpper".into(),
14348 args: vec![arg],
14349 });
14350 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14351 assert!(matches!(result, DfExpr::ScalarFunction(_)));
14353 }
14354
14355 #[test]
14356 fn function_call_unknown_returns_error() {
14357 let mut arena = ExprArena::new();
14358 let id = arena.push(IrExpr::FunctionCall {
14359 name: "unknownFn".into(),
14360 args: vec![],
14361 });
14362 let vm = VarMap::new();
14363 let result = make_lowerer(&arena, &vm).lower(id);
14364 assert!(matches!(result, Err(LoweringError::UnknownFunction(_))));
14365 }
14366
14367 fn lower_rel_fn(name: &str, int_args: &[i64]) -> String {
14374 let mut arena = ExprArena::new();
14375 let mut vm = VarMap::new();
14376 vm.insert(VarId(0), "var_1.rels");
14377 let mut args = vec![arena.push(IrExpr::VarRef(VarId(0)))];
14378 for &n in int_args {
14379 args.push(arena.push(IrExpr::Literal(IrLiteral::Int(n))));
14380 }
14381 let id = arena.push(IrExpr::FunctionCall {
14382 name: name.into(),
14383 args,
14384 });
14385 let expr = make_lowerer(&arena, &vm).lower(id).expect("lower");
14386 format!("{expr}")
14387 }
14388
14389 #[test]
14390 fn subscript_lowers_to_cypher_value_access() {
14391 let s = lower_rel_fn("_subscript", &[0]);
14394 assert!(s.contains("cypher_value_access"), "got {s}");
14395 assert!(s.contains("var_1.rels"), "got {s}");
14396 }
14397
14398 #[test]
14399 fn head_and_last_lower_to_array_element() {
14400 assert!(lower_rel_fn("head", &[]).contains("array_element"));
14401 assert!(lower_rel_fn("last", &[]).contains("array_element"));
14402 }
14403
14404 #[test]
14405 fn slice_lowers_to_array_slice() {
14406 let s = lower_rel_fn("_slice", &[0, 2]);
14408 assert!(s.contains("array_slice"), "got {s}");
14409 assert!(s.contains("var_1.rels"), "got {s}");
14410 }
14411
14412 #[test]
14413 fn slice_with_null_bounds_uses_array_length() {
14414 let mut arena = ExprArena::new();
14417 let mut vm = VarMap::new();
14418 vm.insert(VarId(0), "var_1.rels");
14419 let list = arena.push(IrExpr::VarRef(VarId(0)));
14420 let start = arena.push(IrExpr::Literal(IrLiteral::Int(1)));
14421 let end_null = arena.push(IrExpr::Literal(IrLiteral::Null));
14422 let id = arena.push(IrExpr::FunctionCall {
14423 name: "_slice".into(),
14424 args: vec![list, start, end_null],
14425 });
14426 let expr = make_lowerer(&arena, &vm).lower(id).expect("lower");
14427 let s = format!("{expr}");
14428 assert!(s.contains("array_slice"), "got {s}");
14429 assert!(
14430 s.contains("array_length"),
14431 "an unbounded end must default to array_length: {s}"
14432 );
14433 }
14434
14435 #[test]
14436 fn type_of_element_lowers_to_runtime_graph_metadata_dispatch() {
14437 let mut arena = ExprArena::new();
14438 let mut vm = VarMap::new();
14439 vm.insert(VarId(0), "var_1.rels");
14440 let list = arena.push(IrExpr::VarRef(VarId(0)));
14441 let idx = arena.push(IrExpr::Literal(IrLiteral::Int(0)));
14442 let elem = arena.push(IrExpr::FunctionCall {
14443 name: "_subscript".into(),
14444 args: vec![list, idx],
14445 });
14446 let id = arena.push(IrExpr::FunctionCall {
14447 name: "type".into(),
14448 args: vec![elem],
14449 });
14450 let expr = make_lowerer(&arena, &vm).lower(id).expect("lower");
14451 let s = format!("{expr}");
14452 assert!(
14453 s.contains("cypher_relationship_type"),
14454 "must dispatch graph metadata by runtime value: {s}"
14455 );
14456 assert!(
14457 s.contains("cypher_value_access"),
14458 "over the indexed element: {s}"
14459 );
14460 }
14461
14462 fn invoke_graph_metadata(
14463 kind: GraphMetadataKind,
14464 value: ScalarValue,
14465 ) -> datafusion::error::Result<ScalarValue> {
14466 use datafusion::config::ConfigOptions;
14467
14468 let udf = CypherGraphMetadata::new(kind);
14469 let return_type = udf.return_type(&[])?;
14470 let result = udf.invoke_with_args(ScalarFunctionArgs {
14471 args: vec![ColumnarValue::Scalar(value.clone())],
14472 arg_fields: vec![Arc::new(Field::new("value", value.data_type(), true))],
14473 number_rows: 1,
14474 return_field: Arc::new(Field::new("metadata", return_type, true)),
14475 config_options: Arc::new(ConfigOptions::default()),
14476 })?;
14477 match result {
14478 ColumnarValue::Array(array) => ScalarValue::try_from_array(&array, 0),
14479 ColumnarValue::Scalar(value) => Ok(value),
14480 }
14481 }
14482
14483 #[test]
14484 fn graph_metadata_runtime_dispatch_validates_entity_kind() {
14485 use datafusion::arrow::array::{Int64Array, StringArray, StructArray};
14486 use datafusion::arrow::datatypes::Fields;
14487
14488 let labels = ScalarValue::List(ScalarValue::new_list(
14489 &[ScalarValue::Utf8(Some("Person".into()))],
14490 &DataType::Utf8,
14491 true,
14492 ));
14493 let node = ScalarValue::Struct(Arc::new(StructArray::new(
14494 Fields::from(vec![
14495 Field::new("node_uuid", DataType::Int64, false),
14496 Field::new("labels", labels.data_type(), true),
14497 Field::new("rel_type", DataType::Utf8, true),
14498 ]),
14499 vec![
14500 Arc::new(Int64Array::from(vec![1])),
14501 match &labels {
14502 ScalarValue::List(array) => Arc::clone(array) as _,
14503 _ => unreachable!(),
14504 },
14505 Arc::new(StringArray::from(vec![Some("property, not metadata")])),
14506 ],
14507 None,
14508 )));
14509 let relationship = ScalarValue::Struct(Arc::new(StructArray::new(
14510 Fields::from(vec![
14511 Field::new("edge_uuid", DataType::Int64, false),
14512 Field::new("rel_type", DataType::Utf8, false),
14513 ]),
14514 vec![
14515 Arc::new(Int64Array::from(vec![2])),
14516 Arc::new(StringArray::from(vec!["KNOWS"])),
14517 ],
14518 None,
14519 )));
14520
14521 let actual_labels =
14522 invoke_graph_metadata(GraphMetadataKind::Labels, node.clone()).expect("labels(node)");
14523 assert_eq!(actual_labels, labels);
14524 assert_eq!(
14525 invoke_graph_metadata(GraphMetadataKind::RelationshipType, relationship.clone())
14526 .expect("type(relationship)"),
14527 ScalarValue::Utf8(Some("KNOWS".into()))
14528 );
14529 assert!(invoke_graph_metadata(GraphMetadataKind::RelationshipType, node).is_err());
14530 assert!(invoke_graph_metadata(GraphMetadataKind::Labels, relationship).is_err());
14531 assert!(
14532 invoke_graph_metadata(GraphMetadataKind::Labels, ScalarValue::Null)
14533 .expect("labels(null)")
14534 .is_null()
14535 );
14536
14537 let colliding_map = ScalarValue::Struct(Arc::new(StructArray::new(
14538 Fields::from(vec![Field::new("labels", labels.data_type(), true)]),
14539 vec![match labels {
14540 ScalarValue::List(array) => array as _,
14541 _ => unreachable!(),
14542 }],
14543 None,
14544 )));
14545 assert!(invoke_graph_metadata(GraphMetadataKind::Labels, colliding_map).is_err());
14546 }
14547
14548 #[test]
14549 fn size_lowers_to_cypher_size_udf() {
14550 let mut arena = ExprArena::new();
14551 let mut vm = VarMap::new();
14552 vm.insert(VarId(0), "var_1.rels");
14553 let list = arena.push(IrExpr::VarRef(VarId(0)));
14554 let id = arena.push(IrExpr::FunctionCall {
14555 name: "size".into(),
14556 args: vec![list],
14557 });
14558 let expr = make_lowerer(&arena, &vm).lower(id).expect("lower");
14559 assert!(format!("{expr}").contains("cypher_size"), "got {expr}");
14560 }
14561
14562 #[test]
14563 fn cypher_in_decodes_tagged_list_rhs() {
14564 use datafusion::arrow::array::{Array, BooleanArray};
14565 use datafusion::scalar::ScalarValue as S;
14566 use std::sync::Arc;
14567
14568 let lhs = Arc::new(BooleanArray::from(vec![
14569 Some(true),
14570 Some(true),
14571 Some(true),
14572 Some(true),
14573 ]));
14574 let rhs_values = vec![
14575 S::List(S::new_list(
14576 &[S::Boolean(Some(true))],
14577 &DataType::Boolean,
14578 true,
14579 )),
14580 S::List(S::new_list(
14581 &[S::Boolean(Some(false))],
14582 &DataType::Boolean,
14583 true,
14584 )),
14585 S::List(S::new_list(&[S::Boolean(None)], &DataType::Boolean, true)),
14586 S::List(S::new_list(&[], &DataType::Boolean, true)),
14587 ];
14588 let depth = rhs_values.iter().filter_map(het_depth).max().unwrap();
14589 let rhs = Arc::new(build_het_struct(&rhs_values, depth).expect("tagged RHS"));
14590 let out = invoke_cypher_in(lhs, rhs);
14591 let bools = out.as_any().downcast_ref::<BooleanArray>().unwrap();
14592
14593 assert!(bools.value(0), "true IN [true]");
14594 assert!(!bools.value(1), "true IN [false]");
14595 assert!(bools.is_null(2), "true IN [null] -> null");
14596 assert!(!bools.value(3), "true IN []");
14597 }
14598
14599 fn invoke_cypher_in(
14600 lhs: datafusion::arrow::array::ArrayRef,
14601 rhs: datafusion::arrow::array::ArrayRef,
14602 ) -> datafusion::arrow::array::ArrayRef {
14603 use std::sync::Arc;
14604
14605 use datafusion::arrow::datatypes::Field;
14606 use datafusion::config::ConfigOptions;
14607
14608 let n = lhs.len();
14609 let lhs_field = Arc::new(Field::new("lhs", lhs.data_type().clone(), true));
14610 let rhs_field = Arc::new(Field::new("rhs", rhs.data_type().clone(), true));
14611 let ret = Arc::new(Field::new("in", DataType::Boolean, true));
14612 let args = ScalarFunctionArgs {
14613 args: vec![ColumnarValue::Array(lhs), ColumnarValue::Array(rhs)],
14614 arg_fields: vec![lhs_field, rhs_field],
14615 number_rows: n,
14616 return_field: ret,
14617 config_options: Arc::new(ConfigOptions::default()),
14618 };
14619 match CypherIn::new().invoke_with_args(args).unwrap() {
14620 ColumnarValue::Array(a) => a,
14621 ColumnarValue::Scalar(s) => s.to_array_of_size(n).unwrap(),
14622 }
14623 }
14624
14625 #[test]
14630 fn cypher_size_counts_list_elements() {
14631 use datafusion::arrow::array::{Int64Array, ListArray};
14632 use datafusion::arrow::datatypes::Int32Type;
14633
14634 let arr = ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
14636 Some(vec![Some(10), Some(20), Some(30)]),
14637 Some(vec![]),
14638 ]);
14639 let out = invoke_cypher_size(std::sync::Arc::new(arr));
14640 let counts = out.as_any().downcast_ref::<Int64Array>().unwrap();
14641 assert_eq!(counts.value(0), 3);
14642 assert_eq!(counts.value(1), 0);
14643 }
14644
14645 #[test]
14646 fn cypher_size_counts_string_chars() {
14647 use datafusion::arrow::array::{Array, Int64Array, StringArray};
14648
14649 let arr = StringArray::from(vec![Some("abc"), Some(""), None]);
14650 let out = invoke_cypher_size(std::sync::Arc::new(arr));
14651 let counts = out.as_any().downcast_ref::<Int64Array>().unwrap();
14652 assert_eq!(counts.value(0), 3);
14653 assert_eq!(counts.value(1), 0);
14654 assert!(counts.is_null(2), "null string → null size");
14655 }
14656
14657 #[test]
14658 fn cypher_size_counts_het_tagged_elements() {
14659 use datafusion::arrow::array::{Array, Int64Array};
14660 use datafusion::scalar::ScalarValue as S;
14661
14662 let scalars = vec![
14666 S::List(S::new_list(
14667 &[S::Int64(Some(1)), S::Int64(Some(2)), S::Int64(Some(3))],
14668 &DataType::Int64,
14669 true,
14670 )),
14671 S::Utf8(Some("ab".to_owned())),
14672 S::Boolean(Some(true)),
14673 S::Null,
14674 ];
14675 let depth = scalars.iter().filter_map(het_depth).max().unwrap();
14676 let elems = build_het_struct(&scalars, depth).expect("build tagged struct");
14677 let out = invoke_cypher_size(std::sync::Arc::new(elems));
14678 let counts = out.as_any().downcast_ref::<Int64Array>().unwrap();
14679 assert_eq!(counts.value(0), 3, "tag-4 list element → element count");
14680 assert_eq!(counts.value(1), 2, "tag-2 string element → char count");
14681 assert!(counts.is_null(2), "non-list/string element → null");
14682 assert!(counts.is_null(3), "null element → null");
14683 }
14684
14685 fn invoke_cypher_size(
14688 array: datafusion::arrow::array::ArrayRef,
14689 ) -> datafusion::arrow::array::ArrayRef {
14690 use std::sync::Arc;
14691
14692 use datafusion::arrow::datatypes::Field;
14693 use datafusion::config::ConfigOptions;
14694
14695 let n = array.len();
14696 let field = Arc::new(Field::new("x", array.data_type().clone(), true));
14697 let ret = Arc::new(Field::new("size", DataType::Int64, true));
14698 let args = ScalarFunctionArgs {
14699 args: vec![ColumnarValue::Array(array)],
14700 arg_fields: vec![field],
14701 number_rows: n,
14702 return_field: ret,
14703 config_options: Arc::new(ConfigOptions::default()),
14704 };
14705 match CypherSize::new().invoke_with_args(args).unwrap() {
14706 ColumnarValue::Array(a) => a,
14707 ColumnarValue::Scalar(s) => s.to_array_of_size(n).unwrap(),
14708 }
14709 }
14710
14711 #[test]
14716 fn path_nodes_lowers_to_udf_over_node_uuid() {
14717 let mut arena = ExprArena::new();
14718 let mut vm = VarMap::new();
14719 vm.insert(VarId(0), "var_0"); vm.insert(VarId(1), "var_1.rels"); let start = arena.push(IrExpr::VarRef(VarId(0)));
14722 let rels = arena.push(IrExpr::VarRef(VarId(1)));
14723 let id = arena.push(IrExpr::FunctionCall {
14724 name: "_path_nodes".into(),
14725 args: vec![start, rels],
14726 });
14727 let expr = make_lowerer(&arena, &vm).lower(id).expect("lower");
14728 let s = format!("{expr}");
14729 assert!(s.contains("cypher_path_nodes"), "got {s}");
14730 assert!(
14731 s.contains("var_0.node_uuid"),
14732 "seed is the uuid column: {s}"
14733 );
14734 }
14735
14736 fn uuid16(b: u8) -> Vec<u8> {
14738 vec![b; 16]
14739 }
14740
14741 fn edge_list(rows: &[Option<&[(u8, u8)]>]) -> datafusion::arrow::array::ArrayRef {
14745 use datafusion::arrow::array::{FixedSizeBinaryBuilder, ListBuilder, StructBuilder};
14746 use datafusion::arrow::datatypes::Field;
14747
14748 let fields: datafusion::arrow::datatypes::Fields = vec![
14749 Field::new("src_uuid", DataType::FixedSizeBinary(16), false),
14750 Field::new("dst_uuid", DataType::FixedSizeBinary(16), false),
14751 ]
14752 .into();
14753 let mut b = ListBuilder::new(StructBuilder::new(
14754 fields,
14755 vec![
14756 Box::new(FixedSizeBinaryBuilder::new(16)),
14757 Box::new(FixedSizeBinaryBuilder::new(16)),
14758 ],
14759 ));
14760 for row in rows {
14761 let Some(edges) = row else {
14762 b.append_null();
14763 continue;
14764 };
14765 for (src, dst) in *edges {
14766 b.values()
14767 .field_builder::<FixedSizeBinaryBuilder>(0)
14768 .unwrap()
14769 .append_value(uuid16(*src))
14770 .unwrap();
14771 b.values()
14772 .field_builder::<FixedSizeBinaryBuilder>(1)
14773 .unwrap()
14774 .append_value(uuid16(*dst))
14775 .unwrap();
14776 b.values().append(true);
14777 }
14778 b.append(true);
14779 }
14780 std::sync::Arc::new(b.finish())
14781 }
14782
14783 fn seed_uuids(vals: &[Option<u8>]) -> datafusion::arrow::array::ArrayRef {
14785 use datafusion::arrow::array::FixedSizeBinaryBuilder;
14786 let mut b = FixedSizeBinaryBuilder::new(16);
14787 for v in vals {
14788 match v {
14789 Some(x) => b.append_value(uuid16(*x)).unwrap(),
14790 None => b.append_null(),
14791 }
14792 }
14793 std::sync::Arc::new(b.finish())
14794 }
14795
14796 fn invoke_path_nodes(
14797 seed: datafusion::arrow::array::ArrayRef,
14798 rels: datafusion::arrow::array::ArrayRef,
14799 ) -> datafusion::error::Result<datafusion::arrow::array::ArrayRef> {
14800 use std::sync::Arc;
14801
14802 use datafusion::arrow::datatypes::Field;
14803 use datafusion::config::ConfigOptions;
14804
14805 let udf = CypherPathNodes::new();
14806 let n = seed.len();
14807 let args = ScalarFunctionArgs {
14808 args: vec![
14809 ColumnarValue::Array(Arc::clone(&seed)),
14810 ColumnarValue::Array(Arc::clone(&rels)),
14811 ],
14812 arg_fields: vec![
14813 Arc::new(Field::new("seed", seed.data_type().clone(), true)),
14814 Arc::new(Field::new("rels", rels.data_type().clone(), true)),
14815 ],
14816 number_rows: n,
14817 return_field: Arc::new(Field::new("nodes", udf.return_type(&[])?, true)),
14818 config_options: Arc::new(ConfigOptions::default()),
14819 };
14820 udf.invoke_with_args(args).map(|v| match v {
14821 ColumnarValue::Array(a) => a,
14822 ColumnarValue::Scalar(s) => s.to_array_of_size(n).unwrap(),
14823 })
14824 }
14825
14826 fn path_node_bytes(out: &datafusion::arrow::array::ArrayRef, i: usize) -> Option<Vec<u8>> {
14829 use datafusion::arrow::array::{Array, FixedSizeBinaryArray, ListArray, StructArray};
14830 let list = out.as_any().downcast_ref::<ListArray>().unwrap();
14831 if list.is_null(i) {
14832 return None;
14833 }
14834 let items = list.value(i);
14835 let items = items.as_any().downcast_ref::<StructArray>().unwrap();
14836 let uuids = items.column_by_name("node_uuid").unwrap();
14837 let uuids = uuids
14838 .as_any()
14839 .downcast_ref::<FixedSizeBinaryArray>()
14840 .unwrap();
14841 Some((0..uuids.len()).map(|j| uuids.value(j)[0]).collect())
14842 }
14843
14844 #[test]
14845 fn path_nodes_walks_forward_chain() {
14846 let out = invoke_path_nodes(
14847 seed_uuids(&[Some(1)]),
14848 edge_list(&[Some(&[(1, 2), (2, 3)])]),
14849 )
14850 .unwrap();
14851 assert_eq!(path_node_bytes(&out, 0), Some(vec![1, 2, 3]));
14852 }
14853
14854 #[test]
14855 fn path_nodes_flips_reversed_storage_orientation() {
14856 let out = invoke_path_nodes(seed_uuids(&[Some(1)]), edge_list(&[Some(&[(2, 1)])])).unwrap();
14859 assert_eq!(path_node_bytes(&out, 0), Some(vec![1, 2]));
14860 }
14861
14862 #[test]
14863 fn path_nodes_mixed_orientation_walk() {
14864 let out = invoke_path_nodes(
14866 seed_uuids(&[Some(1)]),
14867 edge_list(&[Some(&[(1, 2), (3, 2)])]),
14868 )
14869 .unwrap();
14870 assert_eq!(path_node_bytes(&out, 0), Some(vec![1, 2, 3]));
14871 }
14872
14873 #[test]
14874 fn path_nodes_self_loop_stays_put() {
14875 let out = invoke_path_nodes(seed_uuids(&[Some(1)]), edge_list(&[Some(&[(1, 1)])])).unwrap();
14876 assert_eq!(path_node_bytes(&out, 0), Some(vec![1, 1]));
14877 }
14878
14879 #[test]
14880 fn path_nodes_zero_hop_is_seed_only() {
14881 let out = invoke_path_nodes(seed_uuids(&[Some(7)]), edge_list(&[Some(&[])])).unwrap();
14882 assert_eq!(path_node_bytes(&out, 0), Some(vec![7]));
14883 }
14884
14885 #[test]
14886 fn path_nodes_null_seed_or_list_is_null() {
14887 let out = invoke_path_nodes(
14889 seed_uuids(&[None, Some(1)]),
14890 edge_list(&[Some(&[(1, 2)]), None]),
14891 )
14892 .unwrap();
14893 assert_eq!(path_node_bytes(&out, 0), None);
14894 assert_eq!(path_node_bytes(&out, 1), None);
14895 }
14896
14897 #[test]
14898 fn path_nodes_disconnected_edge_errors() {
14899 let err = invoke_path_nodes(seed_uuids(&[Some(1)]), edge_list(&[Some(&[(5, 6)])]))
14900 .expect_err("an edge touching neither endpoint is a corrupt emission");
14901 assert!(err.to_string().contains("disconnected"), "got {err}");
14902 }
14903
14904 #[test]
14905 fn path_nodes_output_matches_declared_return_type() {
14906 let out = invoke_path_nodes(seed_uuids(&[Some(1)]), edge_list(&[Some(&[(1, 2)])])).unwrap();
14909 assert_eq!(
14910 out.data_type(),
14911 &CypherPathNodes::new().return_type(&[]).unwrap()
14912 );
14913 }
14914
14915 #[test]
14920 fn parameter_produces_placeholder() {
14921 let mut arena = ExprArena::new();
14922 let id = arena.push(IrExpr::Parameter("eid".into()));
14923 let vm = VarMap::new();
14924 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14925 if let DfExpr::Placeholder(p) = result {
14926 assert_eq!(p.id, "$eid");
14928 } else {
14929 panic!("expected Placeholder, got {result:?}");
14930 }
14931 }
14932
14933 #[test]
14938 fn list_literal_of_ints_folds_to_scalar_list() {
14939 use datafusion::arrow::array::Array;
14940 let mut arena = ExprArena::new();
14942 let e1 = arena.push(IrExpr::Literal(IrLiteral::Int(1)));
14943 let e2 = arena.push(IrExpr::Literal(IrLiteral::Int(2)));
14944 let e3 = arena.push(IrExpr::Literal(IrLiteral::Int(3)));
14945 let id = arena.push(IrExpr::ListLiteral(vec![e1, e2, e3]));
14946 let vm = VarMap::new();
14947 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14948 let DfExpr::Literal(ScalarValue::List(arr), _) = result else {
14949 panic!("expected a ScalarValue::List literal, got {result:?}");
14950 };
14951 assert_eq!(arr.len(), 1);
14953 assert_eq!(arr.value(0).len(), 3);
14954 assert_eq!(arr.value(0).data_type(), &DataType::Int64);
14955 }
14956
14957 #[test]
14958 fn empty_list_literal_folds_to_empty_int64_list() {
14959 use datafusion::arrow::array::Array;
14960 let mut arena = ExprArena::new();
14961 let id = arena.push(IrExpr::ListLiteral(vec![]));
14962 let vm = VarMap::new();
14963 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14964 let DfExpr::Literal(ScalarValue::List(arr), _) = result else {
14965 panic!("expected an empty ScalarValue::List literal, got {result:?}");
14966 };
14967 assert_eq!(arr.len(), 1);
14968 assert_eq!(arr.value(0).len(), 0, "no elements");
14969 }
14970
14971 #[test]
14972 fn list_literal_with_expression_element_uses_make_array() {
14973 let mut arena = ExprArena::new();
14975 let var = arena.push(IrExpr::VarRef(VarId(0)));
14976 let age = arena.push(IrExpr::PropertyAccess {
14977 base: var,
14978 prop: PropId(0),
14979 });
14980 let one = arena.push(IrExpr::Literal(IrLiteral::Int(1)));
14981 let id = arena.push(IrExpr::ListLiteral(vec![age, one]));
14982 let mut vm = VarMap::new();
14983 vm.insert(VarId(0), "var_0");
14984 let result = make_lowerer(&arena, &vm).lower(id).unwrap();
14985 let DfExpr::ScalarFunction(f) = result else {
14986 panic!("expected a make_array ScalarFunction, got {result:?}");
14987 };
14988 assert_eq!(f.name(), "make_array");
14989 assert_eq!(f.args.len(), 2);
14990 }
14991
14992 #[test]
14997 fn scalar_to_ir_literal_round_trips_each_kind() {
14998 for lit in [
15001 IrLiteral::Bool(true),
15002 IrLiteral::Int(42),
15003 IrLiteral::Float(1.5),
15004 IrLiteral::Str("hi".into()),
15005 IrLiteral::Duration {
15006 months: 14,
15007 days: 3,
15008 seconds: 5,
15009 nanos: 1000,
15010 },
15011 IrLiteral::DateTime(1_700_000_000_000_000),
15012 ] {
15013 let scalar = ir_literal_to_scalar(&lit);
15014 assert_eq!(scalar_to_ir_literal(&scalar).unwrap(), lit);
15015 }
15016 }
15017
15018 #[test]
15019 fn scalar_to_ir_literal_rejects_graph_identity_as_a_property_value() {
15020 let scalar = ir_literal_to_scalar(&IrLiteral::Uuid([0x42; 16]));
15021 assert!(matches!(
15022 scalar_to_ir_literal(&scalar),
15023 Err(LoweringError::InvalidType(message))
15024 if message == "UUID values cannot be stored as graph properties"
15025 ));
15026 }
15027
15028 #[test]
15029 fn scalar_to_ir_literal_null_variants_map_to_null() {
15030 assert_eq!(
15031 scalar_to_ir_literal(&ScalarValue::Null).unwrap(),
15032 IrLiteral::Null
15033 );
15034 assert_eq!(
15035 scalar_to_ir_literal(&ScalarValue::Int64(None)).unwrap(),
15036 IrLiteral::Null
15037 );
15038 }
15039
15040 #[test]
15041 fn scalar_to_ir_literal_widens_smaller_ints() {
15042 assert_eq!(
15043 scalar_to_ir_literal(&ScalarValue::Int32(Some(7))).unwrap(),
15044 IrLiteral::Int(7)
15045 );
15046 assert_eq!(
15047 scalar_to_ir_literal(&ScalarValue::UInt8(Some(255))).unwrap(),
15048 IrLiteral::Int(255)
15049 );
15050 }
15051
15052 #[test]
15053 fn scalar_to_ir_literal_normalizes_all_native_widths_and_rejects_overflow() {
15054 let cases = [
15055 (ScalarValue::Int8(Some(-8)), IrLiteral::Int(-8)),
15056 (ScalarValue::Int16(Some(-16)), IrLiteral::Int(-16)),
15057 (ScalarValue::UInt16(Some(16)), IrLiteral::Int(16)),
15058 (ScalarValue::UInt32(Some(32)), IrLiteral::Int(32)),
15059 (ScalarValue::UInt64(Some(64)), IrLiteral::Int(64)),
15060 (ScalarValue::Float32(Some(1.25)), IrLiteral::Float(1.25)),
15061 (
15062 ScalarValue::LargeUtf8(Some("large".into())),
15063 IrLiteral::Str("large".into()),
15064 ),
15065 (
15066 ScalarValue::Utf8View(Some("view".into())),
15067 IrLiteral::Str("view".into()),
15068 ),
15069 (
15070 ScalarValue::TimestampSecond(Some(2), None),
15071 IrLiteral::DateTime(2_000_000),
15072 ),
15073 (
15074 ScalarValue::TimestampMillisecond(Some(3), None),
15075 IrLiteral::DateTime(3_000),
15076 ),
15077 (
15078 ScalarValue::TimestampNanosecond(Some(4_000), None),
15079 IrLiteral::DateTime(4),
15080 ),
15081 (ScalarValue::Time64Nanosecond(Some(5)), IrLiteral::Time(5)),
15082 ];
15083 for (scalar, expected) in cases {
15084 assert_eq!(scalar_to_ir_literal(&scalar).unwrap(), expected);
15085 }
15086 assert!(matches!(
15087 scalar_to_ir_literal(&ScalarValue::UInt64(Some(u64::MAX))),
15088 Err(LoweringError::UnsupportedExpr(message)) if message.contains("exceeds the i64 range")
15089 ));
15090 assert!(matches!(
15091 scalar_to_ir_literal(&ScalarValue::Binary(Some(vec![1, 2]))),
15092 Err(LoweringError::InvalidType(message)) if message.contains("invalid property type")
15093 ));
15094 }
15095
15096 #[test]
15097 fn dynamic_access_helpers_cover_null_bounds_types_and_schema_errors() {
15098 use datafusion::arrow::datatypes::{Field, Fields};
15099
15100 for (scalar, expected) in [
15101 (ScalarValue::Int8(Some(-1)), Some(-1)),
15102 (ScalarValue::Int16(Some(2)), Some(2)),
15103 (ScalarValue::Int32(Some(3)), Some(3)),
15104 (ScalarValue::Int64(Some(4)), Some(4)),
15105 (ScalarValue::UInt8(Some(5)), Some(5)),
15106 (ScalarValue::UInt16(Some(6)), Some(6)),
15107 (ScalarValue::UInt32(Some(7)), Some(7)),
15108 (ScalarValue::UInt64(Some(8)), Some(8)),
15109 (ScalarValue::Null, None),
15110 ] {
15111 assert_eq!(scalar_list_index(&scalar).unwrap(), expected);
15112 }
15113 assert!(scalar_list_index(&ScalarValue::UInt64(Some(u64::MAX))).is_err());
15114 assert!(scalar_list_index(&ScalarValue::Utf8(Some("one".into()))).is_err());
15115 assert_eq!(scalar_access_key(&ScalarValue::Null).unwrap(), None);
15116 assert_eq!(
15117 scalar_access_key(&ScalarValue::LargeUtf8(Some("key".into()))).unwrap(),
15118 Some("key".into())
15119 );
15120 assert!(scalar_access_key(&ScalarValue::Int64(Some(1))).is_err());
15121
15122 let homogeneous = Fields::from(vec![
15123 Field::new("a", DataType::Null, true),
15124 Field::new("b", DataType::Int64, true),
15125 Field::new("c", DataType::Int64, false),
15126 ]);
15127 assert_eq!(
15128 common_struct_field_type(&homogeneous).unwrap(),
15129 DataType::Int64
15130 );
15131 let mixed = Fields::from(vec![
15132 Field::new("a", DataType::Int64, true),
15133 Field::new("b", DataType::Utf8, true),
15134 ]);
15135 assert!(common_struct_field_type(&mixed).is_err());
15136
15137 for dtype in [
15138 DataType::Struct(Fields::empty()),
15139 DataType::Struct(Fields::from(vec![Field::new(
15140 "__het_map",
15141 DataType::Utf8,
15142 true,
15143 )])),
15144 DataType::Struct(Fields::from(vec![Field::new(
15145 "__het_map",
15146 DataType::List(Arc::new(Field::new("item", DataType::Utf8, true))),
15147 true,
15148 )])),
15149 ] {
15150 assert!(het_value_access_return_type(&dtype).is_err());
15151 }
15152 }
15153
15154 #[test]
15155 fn heterogeneous_map_access_returns_exact_values_and_rejects_non_maps() {
15156 use datafusion::arrow::array::{ArrayRef, StringArray};
15157 use datafusion::scalar::ScalarValue as S;
15158
15159 let map = const_map_scalar(&[
15160 ("answer".to_owned(), S::Int64(Some(42))),
15161 ("empty".to_owned(), S::Int64(None)),
15162 ])
15163 .expect("map scalar");
15164 let encoded = build_het_struct(&[map], 1).expect("tagged map");
15165 let return_type =
15166 het_value_access_return_type(encoded.data_type()).expect("map value type");
15167 let null_value = S::try_from(&return_type).expect("typed null");
15168 let keys = |key: Option<&str>| -> ArrayRef { Arc::new(StringArray::from(vec![key])) };
15169
15170 let found = het_map_access_value(&encoded, &keys(Some("answer")), 0, &null_value)
15171 .expect("existing map key");
15172 assert_eq!(decode_het(&found), Some(S::Int64(Some(42))));
15173
15174 let stored_null = het_map_access_value(&encoded, &keys(Some("empty")), 0, &null_value)
15175 .expect("stored null");
15176 assert_eq!(decode_het(&stored_null), Some(S::Null));
15177 assert_eq!(
15178 het_map_access_value(&encoded, &keys(Some("missing")), 0, &null_value)
15179 .expect("missing key"),
15180 null_value
15181 );
15182 assert_eq!(
15183 het_map_access_value(&encoded, &keys(None), 0, &null_value).expect("null key"),
15184 null_value
15185 );
15186
15187 let non_map = build_het_struct(&[S::Int64(Some(7))], 0).expect("tagged integer");
15188 let error = het_map_access_value(&non_map, &keys(Some("answer")), 0, &null_value)
15189 .expect_err("a tagged integer is not dynamically property-readable");
15190 assert_eq!(
15191 error.to_string(),
15192 "Execution error: invalid argument type: dynamic value access requires a map"
15193 );
15194 }
15195
15196 #[test]
15197 fn dynamic_struct_access_observes_missing_null_type_and_row_null_semantics() {
15198 use datafusion::arrow::array::{Array, ArrayRef, Int64Array, StringArray, StructArray};
15199 use datafusion::arrow::buffer::NullBuffer;
15200 use datafusion::arrow::datatypes::{Field, Fields};
15201 use datafusion::config::ConfigOptions;
15202
15203 let values: ArrayRef = Arc::new(StructArray::new(
15204 Fields::from(vec![Field::new("score", DataType::Int64, true)]),
15205 vec![Arc::new(Int64Array::from(vec![Some(9), None, Some(11)]))],
15206 Some(NullBuffer::from(vec![true, true, false])),
15207 ));
15208 let invoke = |keys: ArrayRef| -> datafusion::error::Result<ArrayRef> {
15209 let udf = CypherValueAccess::new();
15210 let result = udf.invoke_with_args(ScalarFunctionArgs {
15211 args: vec![
15212 ColumnarValue::Array(Arc::clone(&values)),
15213 ColumnarValue::Array(keys),
15214 ],
15215 arg_fields: vec![
15216 Arc::new(Field::new("value", values.data_type().clone(), true)),
15217 Arc::new(Field::new("key", DataType::Utf8, true)),
15218 ],
15219 number_rows: 3,
15220 return_field: Arc::new(Field::new("out", DataType::Int64, true)),
15221 config_options: Arc::new(ConfigOptions::default()),
15222 })?;
15223 match result {
15224 ColumnarValue::Array(array) => Ok(array),
15225 ColumnarValue::Scalar(value) => value.to_array_of_size(3),
15226 }
15227 };
15228
15229 let result = invoke(Arc::new(StringArray::from(vec![
15230 Some("score"),
15231 Some("missing"),
15232 Some("score"),
15233 ])))
15234 .expect("dynamic struct access");
15235 let result = result.as_any().downcast_ref::<Int64Array>().expect("Int64");
15236 assert_eq!(result.value(0), 9);
15237 assert!(result.is_null(1), "an absent property is null");
15238 assert!(result.is_null(2), "a null graph-element row is null");
15239
15240 let bad_keys: ArrayRef = Arc::new(Int64Array::from(vec![1, 2, 3]));
15241 let error = invoke(bad_keys).expect_err("numeric property key");
15242 assert!(
15243 error
15244 .to_string()
15245 .contains("dynamic map/property access key must be a string"),
15246 "{error}"
15247 );
15248 }
15249
15250 #[test]
15251 fn temporal_accessor_type_matrix_distinguishes_values_from_properties() {
15252 assert!(!temporal_accessor_valid(&DataType::Date32, "year"));
15254 assert!(!temporal_accessor_valid(&DataType::Date32, "timezone"));
15255 assert!(temporal_accessor_valid(
15256 &ScalarValue::Time64Nanosecond(None).data_type(),
15257 "nanosecond"
15258 ));
15259 assert!(!temporal_accessor_valid(
15260 &duration_scalar(None).data_type(),
15261 "monthsOfYear"
15262 ));
15263 assert!(temporal_accessor_valid(
15264 &localdatetime_scalar(None).data_type(),
15265 "year"
15266 ));
15267 assert!(temporal_accessor_valid(
15268 &datetime_scalar(None).data_type(),
15269 "offsetSeconds"
15270 ));
15271 assert!(!temporal_accessor_valid(&DataType::Utf8, "year"));
15272 assert!(!temporal_accessor_valid(&DataType::Int64, "day"));
15273 }
15274
15275 #[test]
15276 fn scalar_to_ir_literal_round_trips_a_list() {
15277 for lit in [
15280 IrLiteral::List(vec![IrLiteral::Int(1), IrLiteral::Int(2)]),
15281 IrLiteral::List(vec![IrLiteral::Date(5428), IrLiteral::Date(5429)]),
15282 ] {
15283 let scalar = ir_literal_to_scalar(&lit);
15284 assert_eq!(scalar_to_ir_literal(&scalar).unwrap(), lit);
15285 }
15286 }
15287
15288 #[test]
15289 fn zoned_temporal_order_keys_compare_absolute_instants() {
15290 let hour = 3_600_000_000_000_i64;
15291 let early_time = time_scalar(Some((12 * hour + 35 * 60_000_000_000, 5 * 3_600)));
15292 let late_time = time_scalar(Some((10 * hour + 35 * 60_000_000_000, -8 * 3_600)));
15293 assert!(cypher_order_key(&early_time) < cypher_order_key(&late_time));
15294
15295 let earlier_datetime = datetime_scalar(Some((5_000, 12 * hour, 3_600, None)));
15296 let later_datetime = datetime_scalar(Some((5_000, 12 * hour, 0, None)));
15297 assert!(cypher_order_key(&earlier_datetime) < cypher_order_key(&later_datetime));
15298 }
15299
15300 #[test]
15301 fn zoned_temporal_structs_require_cypher_order_keys() {
15302 assert!(needs_cypher_order_key_type(&time_scalar(None).data_type()));
15303 assert!(needs_cypher_order_key_type(
15304 &datetime_scalar(None).data_type()
15305 ));
15306 }
15307
15308 #[test]
15309 fn dynamic_heterogeneous_list_preserves_graph_value_payloads() {
15310 use datafusion::arrow::array::{Int64Array, StructArray};
15311 use datafusion::arrow::datatypes::{Field, Fields};
15312 use datafusion::config::ConfigOptions;
15313
15314 let node_array = StructArray::new(
15315 Fields::from(vec![Field::new("node_uuid", DataType::Int64, false)]),
15316 vec![Arc::new(Int64Array::from(vec![7]))],
15317 None,
15318 );
15319 let node = ScalarValue::Struct(Arc::new(node_array));
15320 let number = ScalarValue::Int64(Some(42));
15321 let arg_types = vec![node.data_type(), number.data_type()];
15322 let args = ScalarFunctionArgs {
15323 args: vec![
15324 ColumnarValue::Scalar(node.clone()),
15325 ColumnarValue::Scalar(number.clone()),
15326 ],
15327 arg_fields: arg_types
15328 .iter()
15329 .enumerate()
15330 .map(|(i, ty)| Arc::new(Field::new(format!("arg_{i}"), ty.clone(), true)))
15331 .collect(),
15332 number_rows: 1,
15333 return_field: Arc::new(Field::new("out", dynamic_het_type(&arg_types), false)),
15334 config_options: Arc::new(ConfigOptions::default()),
15335 };
15336 let out = match CypherDynamicHetList::new()
15337 .invoke_with_args(args)
15338 .expect("dynamic heterogeneous list")
15339 {
15340 ColumnarValue::Array(array) => array,
15341 ColumnarValue::Scalar(value) => value.to_array().expect("scalar list"),
15342 };
15343 let list = out.as_any().downcast_ref::<ListArray>().expect("List");
15344 let values = list.value(0);
15345 let first = ScalarValue::try_from_array(&values, 0).expect("node element");
15346 let second = ScalarValue::try_from_array(&values, 1).expect("number element");
15347 assert_eq!(decode_het(&first), Some(node));
15348 assert_eq!(decode_het(&second), Some(number));
15349 assert!(cypher_order_key(&first).starts_with("20:node"));
15350 assert!(cypher_order_key(&second).starts_with("80:num"));
15351 }
15352
15353 #[test]
15354 fn cypher_reverse_runtime_dispatches_strings_lists_and_type_errors() {
15355 use datafusion::arrow::array::{Array, LargeStringArray, ListArray};
15356 use datafusion::arrow::datatypes::Field;
15357 use datafusion::config::ConfigOptions;
15358
15359 let invoke = |value: ScalarValue| {
15360 let data_type = value.data_type();
15361 CypherReverse::new().invoke_with_args(ScalarFunctionArgs {
15362 args: vec![ColumnarValue::Scalar(value)],
15363 arg_fields: vec![Arc::new(Field::new("value", data_type.clone(), true))],
15364 number_rows: 1,
15365 return_field: Arc::new(Field::new("out", data_type, true)),
15366 config_options: Arc::new(ConfigOptions::default()),
15367 })
15368 };
15369
15370 let text = invoke(ScalarValue::LargeUtf8(Some("Áda".into()))).unwrap();
15371 let text = match text {
15372 ColumnarValue::Array(array) => array,
15373 ColumnarValue::Scalar(value) => value.to_array_of_size(1).unwrap(),
15374 };
15375 let text = text.as_any().downcast_ref::<LargeStringArray>().unwrap();
15376 assert_eq!(text.value(0), "ad́A");
15377
15378 use datafusion::arrow::array::StringArray;
15379 let utf8 = invoke(ScalarValue::Utf8(Some("Graph".into()))).unwrap();
15380 let utf8 = match utf8 {
15381 ColumnarValue::Array(array) => array,
15382 ColumnarValue::Scalar(value) => value.to_array_of_size(1).unwrap(),
15383 };
15384 let utf8 = utf8.as_any().downcast_ref::<StringArray>().unwrap();
15385 assert_eq!(utf8.value(0), "hparG");
15386
15387 let list = ScalarValue::List(ScalarValue::new_list(
15388 &[
15389 ScalarValue::Int64(Some(1)),
15390 ScalarValue::Int64(Some(2)),
15391 ScalarValue::Int64(Some(3)),
15392 ],
15393 &DataType::Int64,
15394 true,
15395 ));
15396 let reversed = invoke(list).unwrap();
15397 let reversed = match reversed {
15398 ColumnarValue::Array(array) => array,
15399 ColumnarValue::Scalar(value) => value.to_array_of_size(1).unwrap(),
15400 };
15401 let reversed = reversed.as_any().downcast_ref::<ListArray>().unwrap();
15402 let values = reversed.value(0);
15403 assert_eq!(
15404 (0..values.len())
15405 .map(|row| ScalarValue::try_from_array(&values, row).unwrap())
15406 .collect::<Vec<_>>(),
15407 [
15408 ScalarValue::Int64(Some(3)),
15409 ScalarValue::Int64(Some(2)),
15410 ScalarValue::Int64(Some(1)),
15411 ]
15412 );
15413
15414 let error = invoke(ScalarValue::Int64(Some(7))).unwrap_err();
15415 assert_eq!(
15416 error.to_string(),
15417 "Error during planning: reverse() expects a string or list, got Int64"
15418 );
15419 }
15420
15421 #[test]
15422 fn cypher_list_plus_runtime_covers_each_operand_shape() {
15423 use datafusion::arrow::array::{Array, ListArray};
15424 use datafusion::arrow::datatypes::Field;
15425 use datafusion::config::ConfigOptions;
15426
15427 let list = |values: &[i64]| {
15428 ScalarValue::List(ScalarValue::new_list(
15429 &values
15430 .iter()
15431 .copied()
15432 .map(|value| ScalarValue::Int64(Some(value)))
15433 .collect::<Vec<_>>(),
15434 &DataType::Int64,
15435 true,
15436 ))
15437 };
15438 let invoke = |left: ScalarValue, right: ScalarValue| {
15439 let udf = CypherListPlus::new();
15440 let types = [left.data_type(), right.data_type()];
15441 let return_type = udf.return_type(&types).unwrap();
15442 udf.invoke_with_args(ScalarFunctionArgs {
15443 args: vec![ColumnarValue::Scalar(left), ColumnarValue::Scalar(right)],
15444 arg_fields: types
15445 .iter()
15446 .enumerate()
15447 .map(|(index, data_type)| {
15448 Arc::new(Field::new(format!("arg_{index}"), data_type.clone(), true))
15449 })
15450 .collect(),
15451 number_rows: 1,
15452 return_field: Arc::new(Field::new("out", return_type, true)),
15453 config_options: Arc::new(ConfigOptions::default()),
15454 })
15455 };
15456 let values = |result: ColumnarValue| {
15457 let array = match result {
15458 ColumnarValue::Array(array) => array,
15459 ColumnarValue::Scalar(value) => value.to_array_of_size(1).unwrap(),
15460 };
15461 let list = array.as_any().downcast_ref::<ListArray>().unwrap();
15462 let values = list.value(0);
15463 (0..values.len())
15464 .map(|row| {
15465 let value = ScalarValue::try_from_array(&values, row).unwrap();
15466 decode_het(&value).unwrap_or(value)
15467 })
15468 .collect::<Vec<_>>()
15469 };
15470
15471 assert_eq!(
15472 values(invoke(list(&[1, 2]), list(&[3, 4])).unwrap()),
15473 [1, 2, 3, 4]
15474 .map(|value| ScalarValue::Int64(Some(value)))
15475 .to_vec()
15476 );
15477 assert_eq!(
15478 values(invoke(list(&[1, 2]), ScalarValue::Int64(Some(3))).unwrap()),
15479 [1, 2, 3]
15480 .map(|value| ScalarValue::Int64(Some(value)))
15481 .to_vec()
15482 );
15483 assert_eq!(
15484 values(invoke(ScalarValue::Int64(Some(1)), list(&[2, 3])).unwrap()),
15485 [1, 2, 3]
15486 .map(|value| ScalarValue::Int64(Some(value)))
15487 .to_vec()
15488 );
15489 let error = invoke(ScalarValue::Int64(Some(1)), ScalarValue::Int64(Some(2))).unwrap_err();
15490 assert_eq!(
15491 error.to_string(),
15492 "Execution error: list + requires at least one list operand"
15493 );
15494 }
15495
15496 #[test]
15497 fn cypher_list_plus_executes_large_list_operands_and_null_rows() {
15498 use datafusion::arrow::array::{Array, ArrayRef, Int64Array, LargeListArray, ListArray};
15499 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
15500 use datafusion::arrow::datatypes::Field;
15501
15502 let large = |values: &[i64], valid: bool| {
15503 let values: ArrayRef = Arc::new(Int64Array::from(values.to_vec()));
15504 ScalarValue::LargeList(Arc::new(LargeListArray::new(
15505 Arc::new(Field::new("item", DataType::Int64, true)),
15506 OffsetBuffer::new(ScalarBuffer::from(vec![
15507 0_i64,
15508 i64::try_from(values.len()).unwrap(),
15509 ])),
15510 values,
15511 Some(NullBuffer::from(vec![valid])),
15512 )))
15513 };
15514 let values = |array: ArrayRef| {
15515 let list = array.as_any().downcast_ref::<ListArray>().expect("List");
15516 if list.is_null(0) {
15517 return None;
15518 }
15519 let values = list.value(0);
15520 Some(
15521 (0..values.len())
15522 .map(|row| {
15523 let value = ScalarValue::try_from_array(&values, row).unwrap();
15524 decode_het(&value).unwrap_or(value)
15525 })
15526 .collect::<Vec<_>>(),
15527 )
15528 };
15529
15530 assert_eq!(
15531 values(
15532 invoke_test_udf(
15533 &CypherListPlus::new(),
15534 vec![large(&[1, 2], true), ScalarValue::Int64(Some(3))],
15535 )
15536 .unwrap()
15537 ),
15538 Some(vec![
15539 ScalarValue::Int64(Some(1)),
15540 ScalarValue::Int64(Some(2)),
15541 ScalarValue::Int64(Some(3)),
15542 ])
15543 );
15544 assert_eq!(
15545 values(
15546 invoke_test_udf(
15547 &CypherListPlus::new(),
15548 vec![ScalarValue::Int64(Some(0)), large(&[1, 2], true)],
15549 )
15550 .unwrap()
15551 ),
15552 Some(vec![
15553 ScalarValue::Int64(Some(0)),
15554 ScalarValue::Int64(Some(1)),
15555 ScalarValue::Int64(Some(2)),
15556 ])
15557 );
15558 assert_eq!(
15559 values(
15560 invoke_test_udf(
15561 &CypherListPlus::new(),
15562 vec![large(&[1], false), large(&[2], true)],
15563 )
15564 .unwrap()
15565 ),
15566 None
15567 );
15568 }
15569
15570 #[test]
15571 fn scalar_udf_runtime_truth_tables_strings_and_ranges() {
15572 use datafusion::arrow::array::{Array, BooleanArray, ListArray};
15573 use datafusion::arrow::datatypes::Field;
15574 use datafusion::config::ConfigOptions;
15575
15576 fn invoke<U: ScalarUDFImpl>(
15577 udf: &U,
15578 values: Vec<ScalarValue>,
15579 return_type: DataType,
15580 ) -> datafusion::error::Result<ColumnarValue> {
15581 let fields = values
15582 .iter()
15583 .enumerate()
15584 .map(|(index, value)| {
15585 Arc::new(Field::new(format!("arg_{index}"), value.data_type(), true))
15586 })
15587 .collect();
15588 udf.invoke_with_args(ScalarFunctionArgs {
15589 args: values.into_iter().map(ColumnarValue::Scalar).collect(),
15590 arg_fields: fields,
15591 number_rows: 1,
15592 return_field: Arc::new(Field::new("out", return_type, true)),
15593 config_options: Arc::new(ConfigOptions::default()),
15594 })
15595 }
15596
15597 let booleans = [None, Some(false), Some(true)];
15598 for kind in [
15599 CypherBoolOpKind::And,
15600 CypherBoolOpKind::Or,
15601 CypherBoolOpKind::Xor,
15602 ] {
15603 for left in booleans {
15604 for right in booleans {
15605 let output = invoke(
15606 &CypherBoolOp::new(kind),
15607 vec![ScalarValue::Boolean(left), ScalarValue::Boolean(right)],
15608 DataType::Boolean,
15609 )
15610 .unwrap();
15611 let output = match output {
15612 ColumnarValue::Array(array) => array,
15613 ColumnarValue::Scalar(value) => value.to_array_of_size(1).unwrap(),
15614 };
15615 let output = output.as_any().downcast_ref::<BooleanArray>().unwrap();
15616 let actual = (!output.is_null(0)).then(|| output.value(0));
15617 let expected = match kind {
15618 CypherBoolOpKind::And => match (left, right) {
15619 (Some(false), _) | (_, Some(false)) => Some(false),
15620 (Some(true), Some(true)) => Some(true),
15621 _ => None,
15622 },
15623 CypherBoolOpKind::Or => match (left, right) {
15624 (Some(true), _) | (_, Some(true)) => Some(true),
15625 (Some(false), Some(false)) => Some(false),
15626 _ => None,
15627 },
15628 CypherBoolOpKind::Xor => left.zip(right).map(|(l, r)| l ^ r),
15629 };
15630 assert_eq!(actual, expected);
15631 }
15632 }
15633 }
15634 assert!(
15635 invoke(
15636 &CypherBoolOp::new(CypherBoolOpKind::And),
15637 vec![
15638 ScalarValue::Int64(Some(1)),
15639 ScalarValue::Boolean(Some(true))
15640 ],
15641 DataType::Boolean,
15642 )
15643 .unwrap_err()
15644 .to_string()
15645 .contains("expected boolean operand")
15646 );
15647
15648 for (kind, expected) in [
15649 (StringPredicate::Starts, true),
15650 (StringPredicate::Ends, false),
15651 (StringPredicate::Contains, true),
15652 ] {
15653 let output = invoke(
15654 &CypherStringPredicate::new(kind),
15655 vec![
15656 ScalarValue::LargeUtf8(Some("GraphForge".into())),
15657 ScalarValue::Utf8(Some("Graph".into())),
15658 ],
15659 DataType::Boolean,
15660 )
15661 .unwrap();
15662 let output = match output {
15663 ColumnarValue::Array(array) => array,
15664 ColumnarValue::Scalar(value) => value.to_array_of_size(1).unwrap(),
15665 };
15666 assert_eq!(
15667 output
15668 .as_any()
15669 .downcast_ref::<BooleanArray>()
15670 .unwrap()
15671 .value(0),
15672 expected
15673 );
15674 }
15675
15676 let range_type = DataType::new_list(DataType::Int64, true);
15677 for (start, end, step, expected) in [(1, 5, 2, vec![1, 3, 5]), (5, 1, -2, vec![5, 3, 1])] {
15678 let output = invoke(
15679 &CypherRange::new(),
15680 vec![
15681 ScalarValue::Int64(Some(start)),
15682 ScalarValue::Int64(Some(end)),
15683 ScalarValue::Int64(Some(step)),
15684 ],
15685 range_type.clone(),
15686 )
15687 .unwrap();
15688 let output = match output {
15689 ColumnarValue::Array(array) => array,
15690 ColumnarValue::Scalar(value) => value.to_array_of_size(1).unwrap(),
15691 };
15692 let list = output
15693 .as_any()
15694 .downcast_ref::<ListArray>()
15695 .unwrap()
15696 .value(0);
15697 assert_eq!(
15698 (0..list.len())
15699 .map(|row| ScalarValue::try_from_array(&list, row).unwrap())
15700 .collect::<Vec<_>>(),
15701 expected
15702 .into_iter()
15703 .map(|value| ScalarValue::Int64(Some(value)))
15704 .collect::<Vec<_>>()
15705 );
15706 }
15707 for (step, fragment) in [(0, "must not be zero"), (2, "overflowed i64")] {
15708 let start = if step == 0 { 1 } else { i64::MAX - 1 };
15709 let end = if step == 0 { 2 } else { i64::MAX };
15710 let error = invoke(
15711 &CypherRange::new(),
15712 vec![
15713 ScalarValue::Int64(Some(start)),
15714 ScalarValue::Int64(Some(end)),
15715 ScalarValue::Int64(Some(step)),
15716 ],
15717 range_type.clone(),
15718 )
15719 .unwrap_err();
15720 assert!(error.to_string().contains(fragment));
15721 }
15722 }
15723
15724 #[test]
15725 fn scalar_conversion_helpers_exhaust_every_numeric_width_null_and_error_contract() {
15726 let integers = [
15727 (ScalarValue::Int8(Some(-8)), -8_i64),
15728 (ScalarValue::Int16(Some(-16)), -16),
15729 (ScalarValue::Int32(Some(-32)), -32),
15730 (ScalarValue::Int64(Some(-64)), -64),
15731 (ScalarValue::UInt8(Some(8)), 8),
15732 (ScalarValue::UInt16(Some(16)), 16),
15733 (ScalarValue::UInt32(Some(32)), 32),
15734 (ScalarValue::UInt64(Some(64)), 64),
15735 ];
15736 for (value, expected) in &integers {
15737 assert_eq!(scalar_as_i128(value), Some(i128::from(*expected)));
15738 assert_eq!(scalar_as_f64(value), Some(*expected as f64));
15739 assert_eq!(to_cypher_integer(value).unwrap(), Some(*expected));
15740 assert_eq!(to_cypher_float(value).unwrap(), Some(*expected as f64));
15741 assert_eq!(to_cypher_string(value).unwrap(), Some(expected.to_string()));
15742 }
15743
15744 for (value, integer, float, text) in [
15745 (
15746 ScalarValue::Float32(Some(12.75)),
15747 Some(12),
15748 Some(12.75),
15749 Some("12.75".to_owned()),
15750 ),
15751 (
15752 ScalarValue::Float64(Some(-12.75)),
15753 Some(-12),
15754 Some(-12.75),
15755 Some("-12.75".to_owned()),
15756 ),
15757 (
15758 ScalarValue::Utf8(Some("42.9".into())),
15759 Some(42),
15760 Some(42.9),
15761 Some("42.9".to_owned()),
15762 ),
15763 (
15764 ScalarValue::LargeUtf8(Some("-3".into())),
15765 Some(-3),
15766 Some(-3.0),
15767 Some("-3".to_owned()),
15768 ),
15769 ] {
15770 assert_eq!(to_cypher_integer(&value).unwrap(), integer);
15771 assert_eq!(to_cypher_float(&value).unwrap(), float);
15772 assert_eq!(to_cypher_string(&value).unwrap(), text);
15773 }
15774
15775 for null in [
15776 ScalarValue::Null,
15777 ScalarValue::Int64(None),
15778 ScalarValue::Float64(None),
15779 ScalarValue::Utf8(None),
15780 ScalarValue::Boolean(None),
15781 ] {
15782 assert_eq!(to_cypher_integer(&null).unwrap(), None);
15783 assert_eq!(to_cypher_float(&null).unwrap(), None);
15784 assert_eq!(to_cypher_boolean(&null).unwrap(), None);
15785 assert_eq!(to_cypher_string(&null).unwrap(), None);
15786 }
15787
15788 for invalid_float in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY, f64::MAX] {
15789 assert_eq!(trunc_float_to_i64(invalid_float), None);
15790 }
15791 assert_eq!(trunc_float_to_i64(-9.99), Some(-9));
15792 for invalid_text in ["", "not-a-number", "NaN", "inf"] {
15793 let value = ScalarValue::Utf8(Some(invalid_text.into()));
15794 assert_eq!(to_cypher_integer(&value).unwrap(), None);
15795 assert_eq!(to_cypher_float(&value).unwrap(), None);
15796 }
15797 assert_eq!(
15798 to_cypher_boolean(&ScalarValue::Boolean(Some(true))).unwrap(),
15799 Some(true)
15800 );
15801 assert_eq!(
15802 to_cypher_boolean(&ScalarValue::Utf8(Some("true".into()))).unwrap(),
15803 Some(true)
15804 );
15805 assert_eq!(
15806 to_cypher_boolean(&ScalarValue::LargeUtf8(Some("false".into()))).unwrap(),
15807 Some(false)
15808 );
15809 assert_eq!(
15810 to_cypher_boolean(&ScalarValue::Utf8(Some("TRUE".into()))).unwrap(),
15811 None
15812 );
15813
15814 for invalid in [
15815 ScalarValue::Boolean(Some(true)),
15816 ScalarValue::Binary(Some(vec![1])),
15817 ] {
15818 assert!(to_cypher_integer(&invalid).is_err());
15819 assert!(to_cypher_float(&invalid).is_err());
15820 }
15821 assert!(to_cypher_boolean(&ScalarValue::Int64(Some(1))).is_err());
15822 assert!(to_cypher_string(&ScalarValue::Binary(Some(vec![1]))).is_err());
15823 assert!(to_cypher_integer(&ScalarValue::UInt64(Some(u64::MAX))).is_err());
15824 }
15825
15826 #[test]
15827 fn canonical_float_strings_and_scalar_range_arguments_cover_boundaries() {
15828 for (value, expected) in [
15829 (0.0, "0.0"),
15830 (-0.0, "0.0"),
15831 (f64::NAN, "NaN"),
15832 (f64::INFINITY, "Infinity"),
15833 (f64::NEG_INFINITY, "-Infinity"),
15834 (1.0, "1.0"),
15835 (1.5, "1.5"),
15836 (1e20, "100000000000000000000.0"),
15837 ] {
15838 assert_eq!(cypher_float_string(value), expected);
15839 }
15840
15841 for (value, expected) in [
15842 (ScalarValue::Int8(Some(-1)), -1),
15843 (ScalarValue::Int16(Some(-2)), -2),
15844 (ScalarValue::Int32(Some(-3)), -3),
15845 (ScalarValue::Int64(Some(-4)), -4),
15846 (ScalarValue::UInt8(Some(1)), 1),
15847 (ScalarValue::UInt16(Some(2)), 2),
15848 (ScalarValue::UInt32(Some(3)), 3),
15849 (ScalarValue::UInt64(Some(4)), 4),
15850 ] {
15851 assert_eq!(scalar_as_i64_arg(&value, "bound").unwrap(), expected);
15852 }
15853 assert!(
15854 scalar_as_i64_arg(&ScalarValue::UInt64(Some(u64::MAX)), "bound")
15855 .unwrap_err()
15856 .to_string()
15857 .contains("exceeds i64::MAX")
15858 );
15859 assert!(
15860 scalar_as_i64_arg(&ScalarValue::Utf8(Some("1".into())), "bound")
15861 .unwrap_err()
15862 .to_string()
15863 .contains("must be an integer")
15864 );
15865 }
15866
15867 #[test]
15868 fn temporal_literal_render_dispatch_and_ir_scalar_round_trip_matrix() {
15869 for (name, input) in [
15870 ("date", "2024-02-29"),
15871 ("localtime", "12:34:56"),
15872 ("time", "12:34:56+01:00"),
15873 ("localdatetime", "2024-02-29T12:34:56"),
15874 ("datetime", "2024-02-29T12:34:56Z"),
15875 ("duration", "P1M2DT3S"),
15876 ] {
15877 assert!(render_temporal(name, input).is_some(), "{name}({input})");
15878 }
15879 assert_eq!(render_temporal("unknown", "2024-01-01"), None);
15880 assert_eq!(render_temporal("date", "not-a-date"), None);
15881
15882 let literals = [
15883 IrLiteral::Null,
15884 IrLiteral::Bool(true),
15885 IrLiteral::Int(-7),
15886 IrLiteral::Float(1.25),
15887 IrLiteral::Str("value".into()),
15888 IrLiteral::Duration {
15889 months: 1,
15890 days: 2,
15891 seconds: 3,
15892 nanos: 4,
15893 },
15894 IrLiteral::DateTime(123),
15895 IrLiteral::Date(20_000),
15896 IrLiteral::LocalDateTime {
15897 days: 20_000,
15898 nanos: 123,
15899 },
15900 IrLiteral::Time(456),
15901 IrLiteral::ZonedTime {
15902 nanos: 789,
15903 offset: 3_600,
15904 },
15905 IrLiteral::ZonedDateTime {
15906 days: 20_000,
15907 nanos: 999,
15908 offset: -3_600,
15909 zone: Some("America/Denver".into()),
15910 },
15911 IrLiteral::List(vec![IrLiteral::Int(1), IrLiteral::Null]),
15912 IrLiteral::Map(vec![("answer".into(), IrLiteral::Int(42))]),
15913 ];
15914 for literal in literals {
15915 let scalar = ir_literal_to_scalar(&literal);
15916 if !matches!(literal, IrLiteral::Map(_)) {
15917 assert_eq!(scalar_to_ir_literal(&scalar).unwrap(), literal);
15918 }
15919 }
15920
15921 for (scalar, expected) in [
15922 (
15923 ScalarValue::DurationSecond(Some(-2)),
15924 IrLiteral::Duration {
15925 months: 0,
15926 days: 0,
15927 seconds: -2,
15928 nanos: 0,
15929 },
15930 ),
15931 (
15932 ScalarValue::DurationMillisecond(Some(-1)),
15933 IrLiteral::Duration {
15934 months: 0,
15935 days: 0,
15936 seconds: -1,
15937 nanos: 999_000_000,
15938 },
15939 ),
15940 (
15941 ScalarValue::DurationMicrosecond(Some(-1)),
15942 IrLiteral::Duration {
15943 months: 0,
15944 days: 0,
15945 seconds: -1,
15946 nanos: 999_999_000,
15947 },
15948 ),
15949 (
15950 ScalarValue::DurationNanosecond(Some(-1)),
15951 IrLiteral::Duration {
15952 months: 0,
15953 days: 0,
15954 seconds: -1,
15955 nanos: 999_999_999,
15956 },
15957 ),
15958 ] {
15959 assert_eq!(scalar_to_ir_literal(&scalar).unwrap(), expected);
15960 }
15961 }
15962
15963 #[test]
15964 fn temporal_truncate_lowering_covers_arity_default_literal_override_and_rejection_paths() {
15965 for name in [
15966 "date.truncate",
15967 "localtime.truncate",
15968 "localdatetime.truncate",
15969 "time.truncate",
15970 "datetime.truncate",
15971 ] {
15972 let mut missing = ExprArena::new();
15973 let call = missing.push(IrExpr::FunctionCall {
15974 name: name.into(),
15975 args: vec![],
15976 });
15977 assert!(matches!(
15978 make_lowerer(&missing, &VarMap::new()).lower(call),
15979 Err(LoweringError::UnknownFunction(function)) if function == name
15980 ));
15981
15982 let mut defaults = ExprArena::new();
15983 let unit = defaults.push(IrExpr::Literal(IrLiteral::Str("day".into())));
15984 let value = defaults.push(IrExpr::Literal(IrLiteral::Null));
15985 let call = defaults.push(IrExpr::FunctionCall {
15986 name: name.into(),
15987 args: vec![unit, value],
15988 });
15989 let lowered = make_lowerer(&defaults, &VarMap::new()).lower(call).unwrap();
15990 assert!(format!("{lowered}").contains("truncate"));
15991
15992 let mut overrides = ExprArena::new();
15993 let unit = overrides.push(IrExpr::Literal(IrLiteral::Str("day".into())));
15994 let value = overrides.push(IrExpr::Literal(IrLiteral::Null));
15995 let one = overrides.push(IrExpr::Literal(IrLiteral::Int(1)));
15996 let zone = overrides.push(IrExpr::Literal(IrLiteral::Str("UTC".into())));
15997 let map = overrides.push(IrExpr::MapLiteral(vec![
15998 ("year".into(), one),
15999 ("month".into(), one),
16000 ("day".into(), one),
16001 ("week".into(), one),
16002 ("dayOfWeek".into(), one),
16003 ("ordinalDay".into(), one),
16004 ("quarter".into(), one),
16005 ("dayOfQuarter".into(), one),
16006 ("hour".into(), one),
16007 ("minute".into(), one),
16008 ("second".into(), one),
16009 ("millisecond".into(), one),
16010 ("microsecond".into(), one),
16011 ("nanosecond".into(), one),
16012 ("timezone".into(), zone),
16013 ]));
16014 let call = overrides.push(IrExpr::FunctionCall {
16015 name: name.into(),
16016 args: vec![unit, value, map],
16017 });
16018 assert!(make_lowerer(&overrides, &VarMap::new()).lower(call).is_ok());
16019
16020 let mut dynamic = ExprArena::new();
16021 let unit = dynamic.push(IrExpr::Literal(IrLiteral::Str("day".into())));
16022 let value = dynamic.push(IrExpr::Literal(IrLiteral::Null));
16023 let parameter = dynamic.push(IrExpr::Parameter("overrides".into()));
16024 let call = dynamic.push(IrExpr::FunctionCall {
16025 name: name.into(),
16026 args: vec![unit, value, parameter],
16027 });
16028 assert!(
16029 make_lowerer(&dynamic, &VarMap::new())
16030 .lower(call)
16031 .unwrap_err()
16032 .to_string()
16033 .contains("override map must be a literal map")
16034 );
16035 }
16036
16037 for name in [
16038 "duration.between",
16039 "duration.inmonths",
16040 "duration.indays",
16041 "duration.inseconds",
16042 ] {
16043 let mut arena = ExprArena::new();
16044 let call = arena.push(IrExpr::FunctionCall {
16045 name: name.into(),
16046 args: vec![],
16047 });
16048 assert!(matches!(
16049 make_lowerer(&arena, &VarMap::new()).lower(call),
16050 Err(LoweringError::UnknownFunction(function)) if function == name
16051 ));
16052
16053 let left = arena.push(IrExpr::Literal(IrLiteral::Null));
16054 let right = arena.push(IrExpr::Literal(IrLiteral::Null));
16055 let call = arena.push(IrExpr::FunctionCall {
16056 name: name.into(),
16057 args: vec![left, right],
16058 });
16059 assert!(make_lowerer(&arena, &VarMap::new()).lower(call).is_ok());
16060 }
16061 }
16062
16063 #[test]
16064 fn cypher_value_comparison_and_order_helpers_cover_cross_type_edges() {
16065 use datafusion::scalar::ScalarValue as S;
16066
16067 for value in [
16068 S::Int8(Some(1)),
16069 S::Int16(Some(1)),
16070 S::Int32(Some(1)),
16071 S::Int64(Some(1)),
16072 S::UInt8(Some(1)),
16073 S::UInt16(Some(1)),
16074 S::UInt32(Some(1)),
16075 S::UInt64(Some(1)),
16076 S::Float32(Some(1.0)),
16077 S::Float64(Some(1.0)),
16078 ] {
16079 assert_eq!(scalar_as_f64(&value), Some(1.0));
16080 }
16081 assert_eq!(scalar_as_i128(&S::Float64(Some(1.0))), None);
16082 assert_eq!(scalar_as_f64(&S::Boolean(Some(true))), None);
16083
16084 let one = S::List(S::new_list(&[S::Int64(Some(1))], &DataType::Int64, true));
16085 let one_null = S::List(S::new_list(
16086 &[S::Int64(Some(1)), S::Int64(None)],
16087 &DataType::Int64,
16088 true,
16089 ));
16090 let two = S::List(S::new_list(
16091 &[S::Int64(Some(1)), S::Int64(Some(2))],
16092 &DataType::Int64,
16093 true,
16094 ));
16095 assert_eq!(cypher_value_eq(&one, &one), Some(true));
16096 assert_eq!(cypher_value_eq(&one, &two), Some(false));
16097 assert_eq!(cypher_value_eq(&one_null, &one_null), None);
16098 assert_eq!(cypher_value_eq(&S::Null, &S::Int64(Some(1))), None);
16099 assert_eq!(
16100 cypher_value_eq(&S::Int64(Some(1)), &S::Float64(Some(1.0))),
16101 Some(true)
16102 );
16103 assert_eq!(
16104 cypher_value_eq(&S::Utf8(Some("a".into())), &S::Utf8(Some("b".into()))),
16105 Some(false)
16106 );
16107
16108 assert!(cypher_order_key(&S::Null).starts_with("99:null"));
16109 assert!(cypher_order_key(&S::Utf8(Some("a".into()))).starts_with("60:str"));
16110 assert!(cypher_order_key(&S::Boolean(Some(true))).starts_with("70:bool"));
16111 assert!(cypher_order_key(&S::Float64(Some(f64::NAN))).starts_with("90:nan"));
16112 assert!(cypher_order_key(&S::Binary(Some(vec![1]))).starts_with("98:other"));
16113 assert!(cypher_order_key(&one).starts_with("40:list"));
16114 }
16115
16116 #[test]
16117 fn expression_lowering_error_and_static_access_matrix_reaches_contract_branches() {
16118 let lower_call = |name: &str, args: Vec<IrExpr>| {
16119 let mut arena = ExprArena::new();
16120 let args = args
16121 .into_iter()
16122 .map(|expr| arena.push(expr))
16123 .collect::<Vec<_>>();
16124 let call = arena.push(IrExpr::FunctionCall {
16125 name: name.into(),
16126 args,
16127 });
16128 make_lowerer(&arena, &VarMap::new()).lower(call)
16129 };
16130
16131 for (name, expected) in [
16132 ("_subscript", "expects two arguments"),
16133 ("_node_struct", "expects at least one argument"),
16134 (
16135 "_node_struct_list",
16136 "expects two nodes and one relationship",
16137 ),
16138 ("_rel_struct", "expects an edge variable"),
16139 ("_rel_struct_list", "expects an edge variable"),
16140 ("keys", "expects one argument"),
16141 ("properties", "expects one argument"),
16142 ("labels", "expects one argument"),
16143 ] {
16144 assert!(
16145 lower_call(name, vec![])
16146 .unwrap_err()
16147 .to_string()
16148 .contains(expected),
16149 "{name}"
16150 );
16151 }
16152 for name in ["nodes", "relationships"] {
16153 assert!(
16154 lower_call(name, vec![])
16155 .unwrap_err()
16156 .to_string()
16157 .contains("expects one path argument")
16158 );
16159 }
16160
16161 assert!(
16162 lower_call("_node_struct", vec![IrExpr::Literal(IrLiteral::Int(1))])
16163 .unwrap_err()
16164 .to_string()
16165 .contains("must be a node variable")
16166 );
16167 assert!(
16168 lower_call("_rel_struct", vec![IrExpr::Literal(IrLiteral::Int(1))])
16169 .unwrap_err()
16170 .to_string()
16171 .contains("must be a relationship variable")
16172 );
16173 assert!(
16174 lower_call(
16175 "_node_struct_list",
16176 vec![
16177 IrExpr::Literal(IrLiteral::Int(1)),
16178 IrExpr::Literal(IrLiteral::Int(2)),
16179 IrExpr::Literal(IrLiteral::Null),
16180 ],
16181 )
16182 .unwrap_err()
16183 .to_string()
16184 .contains("node arguments must be variables")
16185 );
16186
16187 for name in ["keys", "properties"] {
16188 assert!(
16189 lower_call(name, vec![IrExpr::ListLiteral(vec![])],)
16190 .unwrap_err()
16191 .to_string()
16192 .contains("requires a map, node, relationship, or null")
16193 );
16194 assert!(lower_call(name, vec![IrExpr::Literal(IrLiteral::Null)]).is_ok());
16195 assert!(lower_call(name, vec![IrExpr::MapLiteral(vec![])],).is_ok());
16196 }
16197 assert!(lower_call("labels", vec![IrExpr::Literal(IrLiteral::Null)]).is_ok());
16198
16199 let mut arena = ExprArena::new();
16200 let null = arena.push(IrExpr::Literal(IrLiteral::Null));
16201 let null_key = arena.push(IrExpr::Literal(IrLiteral::Null));
16202 let access = arena.push(IrExpr::FunctionCall {
16203 name: "_subscript".into(),
16204 args: vec![null, null_key],
16205 });
16206 let null_access = make_lowerer(&arena, &VarMap::new()).lower(access).unwrap();
16207 assert!(format!("{null_access}").contains("cypher_value_access"));
16208
16209 let mut arena = ExprArena::new();
16210 let answer = arena.push(IrExpr::Literal(IrLiteral::Int(42)));
16211 let map = arena.push(IrExpr::MapLiteral(vec![("answer".into(), answer)]));
16212 let key = arena.push(IrExpr::Literal(IrLiteral::Str("answer".into())));
16213 let missing = arena.push(IrExpr::Literal(IrLiteral::Str("missing".into())));
16214 let found = arena.push(IrExpr::FunctionCall {
16215 name: "_subscript".into(),
16216 args: vec![map, key],
16217 });
16218 let absent = arena.push(IrExpr::FunctionCall {
16219 name: "_subscript".into(),
16220 args: vec![map, missing],
16221 });
16222 assert_eq!(
16223 format!(
16224 "{}",
16225 make_lowerer(&arena, &VarMap::new()).lower(found).unwrap()
16226 ),
16227 "Int64(42)"
16228 );
16229 assert!(matches!(
16230 make_lowerer(&arena, &VarMap::new()).lower(absent).unwrap(),
16231 DfExpr::Literal(ScalarValue::Null, _)
16232 ));
16233
16234 let mut arena = ExprArena::new();
16235 let scalar = arena.push(IrExpr::Literal(IrLiteral::Int(1)));
16236 let index = arena.push(IrExpr::Literal(IrLiteral::Int(0)));
16237 let invalid = arena.push(IrExpr::FunctionCall {
16238 name: "_subscript".into(),
16239 args: vec![scalar, index],
16240 });
16241 assert!(
16242 make_lowerer(&arena, &VarMap::new())
16243 .lower(invalid)
16244 .unwrap_err()
16245 .to_string()
16246 .contains("subscript requires a list")
16247 );
16248
16249 for argument in [
16250 IrExpr::Literal(IrLiteral::Str("abc".into())),
16251 IrExpr::ListLiteral(vec![]),
16252 IrExpr::Parameter("value".into()),
16253 ] {
16254 assert!(lower_call("reverse", vec![argument]).is_ok());
16255 }
16256 }
16257
16258 #[test]
16259 fn static_nested_list_map_access_handles_negative_oob_and_nonliteral_indices() {
16260 let mut arena = ExprArena::new();
16261 let one = arena.push(IrExpr::Literal(IrLiteral::Int(1)));
16262 let two = arena.push(IrExpr::Literal(IrLiteral::Int(2)));
16263 let first_map = arena.push(IrExpr::MapLiteral(vec![("value".into(), one)]));
16264 let second_map = arena.push(IrExpr::MapLiteral(vec![("value".into(), two)]));
16265 let list = arena.push(IrExpr::ListLiteral(vec![first_map, second_map]));
16266 let negative = arena.push(IrExpr::Literal(IrLiteral::Int(-1)));
16267 let oob = arena.push(IrExpr::Literal(IrLiteral::Int(9)));
16268 let dynamic = arena.push(IrExpr::Parameter("index".into()));
16269 let key = arena.push(IrExpr::Literal(IrLiteral::Str("value".into())));
16270
16271 for (index, expected) in [(negative, Some("Int64(2)")), (oob, None)] {
16272 let indexed = arena.push(IrExpr::FunctionCall {
16273 name: "_subscript".into(),
16274 args: vec![list, index],
16275 });
16276 let field = arena.push(IrExpr::FunctionCall {
16277 name: "_subscript".into(),
16278 args: vec![indexed, key],
16279 });
16280 let lowered = make_lowerer(&arena, &VarMap::new()).lower(field).unwrap();
16281 match expected {
16282 Some(expected) => assert_eq!(format!("{lowered}"), expected),
16283 None => assert!(matches!(lowered, DfExpr::Literal(ScalarValue::Null, _))),
16284 }
16285 }
16286
16287 let indexed = arena.push(IrExpr::FunctionCall {
16288 name: "_subscript".into(),
16289 args: vec![list, dynamic],
16290 });
16291 let field = arena.push(IrExpr::FunctionCall {
16292 name: "_subscript".into(),
16293 args: vec![indexed, key],
16294 });
16295 assert!(
16296 format!(
16297 "{}",
16298 make_lowerer(&arena, &VarMap::new()).lower(field).unwrap()
16299 )
16300 .contains("cypher_value_access")
16301 );
16302 }
16303
16304 #[test]
16305 fn aggregate_accumulators_cover_update_merge_state_and_empty_contracts() {
16306 use datafusion::arrow::array::{
16307 ArrayRef, Float64Array, Int64Array, ListArray, StringArray,
16308 };
16309 use datafusion::logical_expr::Accumulator;
16310
16311 let ints: ArrayRef = Arc::new(Int64Array::from(vec![Some(4), None, Some(-2), Some(9)]));
16312 for (is_max, expected) in [(true, 9), (false, -2)] {
16313 let mut acc = ExtremeAcc {
16314 is_max,
16315 dtype: DataType::Int64,
16316 best: None,
16317 };
16318 assert_eq!(acc.evaluate().unwrap(), ScalarValue::Int64(None));
16319 acc.update_batch(std::slice::from_ref(&ints)).unwrap();
16320 assert_eq!(acc.evaluate().unwrap(), ScalarValue::Int64(Some(expected)));
16321 assert_eq!(
16322 acc.state().unwrap(),
16323 vec![ScalarValue::Int64(Some(expected))]
16324 );
16325 assert!(acc.size() >= std::mem::size_of::<ExtremeAcc>());
16326 let merged: ArrayRef =
16327 Arc::new(Int64Array::from(vec![Some(if is_max { 12 } else { -7 })]));
16328 acc.merge_batch(&[merged]).unwrap();
16329 assert_eq!(
16330 acc.evaluate().unwrap(),
16331 ScalarValue::Int64(Some(if is_max { 12 } else { -7 }))
16332 );
16333 }
16334
16335 for distinct in [false, true] {
16336 let mut acc = CollectAcc {
16337 distinct,
16338 elem_type: DataType::Int64,
16339 values: Vec::new(),
16340 };
16341 acc.update_batch(std::slice::from_ref(&ints)).unwrap();
16342 let merge: ArrayRef = Arc::new(ListArray::from_iter_primitive::<
16343 datafusion::arrow::datatypes::Int64Type,
16344 _,
16345 _,
16346 >([Some(vec![Some(4), Some(11)])]));
16347 acc.merge_batch(&[merge]).unwrap();
16348 let ScalarValue::List(values) = acc.evaluate().unwrap() else {
16349 panic!("collect must return a list")
16350 };
16351 let expected_len = if distinct { 4 } else { 5 };
16352 assert_eq!(values.value(0).len(), expected_len);
16353 assert_eq!(acc.state().unwrap().len(), 1);
16354 assert!(acc.size() >= std::mem::size_of::<CollectAcc>());
16355 let bad: ArrayRef = Arc::new(Int64Array::from(vec![1]));
16356 assert!(
16357 acc.merge_batch(&[bad])
16358 .unwrap_err()
16359 .to_string()
16360 .contains("must be a list")
16361 );
16362 }
16363
16364 for continuous in [false, true] {
16365 let mut acc = PercentileAcc {
16366 continuous,
16367 value_type: DataType::Int64,
16368 result_type: if continuous {
16369 DataType::Float64
16370 } else {
16371 DataType::Int64
16372 },
16373 values: Vec::new(),
16374 percentile: None,
16375 };
16376 assert!(acc.evaluate().unwrap().is_null());
16377 let p: ArrayRef = Arc::new(Float64Array::from(vec![Some(0.5); 4]));
16378 acc.update_batch(&[Arc::clone(&ints), p]).unwrap();
16379 assert_eq!(
16380 acc.evaluate().unwrap(),
16381 if continuous {
16382 ScalarValue::Float64(Some(4.0))
16383 } else {
16384 ScalarValue::Int64(Some(4))
16385 }
16386 );
16387 assert_eq!(acc.state().unwrap().len(), 2);
16388 assert!(acc.size() >= std::mem::size_of::<PercentileAcc>());
16389 assert!(acc.observe_percentile(Some(f64::NAN)).is_err());
16390 assert!(acc.observe_percentile(Some(0.75)).is_err());
16391 let bad_values: ArrayRef = Arc::new(StringArray::from(vec!["not-list"]));
16392 let good_p: ArrayRef = Arc::new(Float64Array::from(vec![0.5]));
16393 assert!(
16394 acc.merge_batch(&[bad_values, good_p])
16395 .unwrap_err()
16396 .to_string()
16397 .contains("must be a list")
16398 );
16399 }
16400 }
16401
16402 #[test]
16403 fn duration_and_temporal_runtime_udfs_cover_each_value_family_and_nulls() {
16404 let d1 = crate::temporal::DurationValue {
16405 months: 1,
16406 days: 2,
16407 seconds: 3,
16408 nanos: 750_000_000,
16409 };
16410 let d2 = crate::temporal::DurationValue {
16411 months: 2,
16412 days: 3,
16413 seconds: 4,
16414 nanos: 500_000_000,
16415 };
16416 let d1 = duration_scalar(Some(d1));
16417 let d2 = duration_scalar(Some(d2));
16418
16419 let parsed = invoke_test_udf(
16420 &CypherDurationParse::new(),
16421 vec![ScalarValue::Utf8(Some("P1M2DT3.5S".into()))],
16422 )
16423 .unwrap();
16424 assert!(duration_struct_parts(parsed.as_any().downcast_ref().unwrap(), 0).is_some());
16425 let invalid = invoke_test_udf(
16426 &CypherDurationParse::new(),
16427 vec![ScalarValue::Utf8(Some("invalid".into()))],
16428 )
16429 .unwrap();
16430 assert!(duration_struct_parts(invalid.as_any().downcast_ref().unwrap(), 0).is_none());
16431
16432 for sign in [1, -1] {
16433 let out = invoke_test_udf(
16434 &CypherDurationAdd::new(),
16435 vec![d1.clone(), d2.clone(), ScalarValue::Int64(Some(sign))],
16436 )
16437 .unwrap();
16438 assert!(duration_struct_parts(out.as_any().downcast_ref().unwrap(), 0).is_some());
16439 }
16440 for (factor, divide) in [(2.0, false), (2.0, true)] {
16441 let out = invoke_test_udf(
16442 &CypherDurationScale::new(),
16443 vec![
16444 d1.clone(),
16445 ScalarValue::Float64(Some(factor)),
16446 ScalarValue::Boolean(Some(divide)),
16447 ],
16448 )
16449 .unwrap();
16450 assert!(duration_struct_parts(out.as_any().downcast_ref().unwrap(), 0).is_some());
16451 }
16452
16453 let temporal_values = [
16454 date_scalar(Some(20_000)),
16455 ScalarValue::Time64Nanosecond(Some(10)),
16456 time_scalar(Some((10, 3_600))),
16457 localdatetime_scalar(Some((20_000, 10))),
16458 datetime_scalar(Some((20_000, 10, 0, Some("UTC".into())))),
16459 ];
16460 for temporal in temporal_values {
16461 for sign in [1, -1] {
16462 let out = invoke_test_udf(
16463 &CypherTemporalArith::new(),
16464 vec![temporal.clone(), d1.clone(), ScalarValue::Int64(Some(sign))],
16465 )
16466 .unwrap();
16467 assert_eq!(out.data_type(), &temporal.data_type());
16468 assert!(!out.is_null(0));
16469 }
16470 }
16471 assert!(
16472 invoke_test_udf(
16473 &CypherTemporalArith::new(),
16474 vec![ScalarValue::Int64(Some(1)), d1, ScalarValue::Int64(Some(1))],
16475 )
16476 .unwrap_err()
16477 .to_string()
16478 .contains("not a temporal value")
16479 );
16480 }
16481
16482 #[test]
16483 fn exact_zero_large_list_legacy_variant_order_and_percentile_branches() {
16484 use datafusion::arrow::array::{Array, ArrayRef, Int64Array, LargeListArray};
16485 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
16486 use datafusion::arrow::datatypes::{Field, Fields};
16487 use datafusion::logical_expr::Accumulator;
16488
16489 let large = |values: Vec<i64>, valid: bool| {
16490 let len = i64::try_from(values.len()).unwrap();
16491 ScalarValue::LargeList(Arc::new(LargeListArray::new(
16492 Arc::new(Field::new("item", DataType::Int64, true)),
16493 OffsetBuffer::new(ScalarBuffer::from(vec![0, len])),
16494 Arc::new(Int64Array::from(values)) as ArrayRef,
16495 Some(NullBuffer::from(vec![valid])),
16496 )))
16497 };
16498 assert_eq!(
16499 scalar_list_elements(&large(vec![1, 2], true)).unwrap(),
16500 Some(vec![
16501 ScalarValue::Int64(Some(1)),
16502 ScalarValue::Int64(Some(2))
16503 ])
16504 );
16505 assert_eq!(scalar_list_elements(&large(vec![], false)).unwrap(), None);
16506
16507 let size = invoke_test_udf(&CypherSize::new(), vec![large(vec![1, 2, 3], true)]).unwrap();
16508 assert_eq!(
16509 ScalarValue::try_from_array(&size, 0).unwrap(),
16510 ScalarValue::Int64(Some(3))
16511 );
16512 let null_size = invoke_test_udf(&CypherSize::new(), vec![large(vec![], false)]).unwrap();
16513 assert!(null_size.is_null(0));
16514
16515 let reversed = invoke_test_udf(
16516 &CypherReverse::new(),
16517 vec![ScalarValue::LargeUtf8(Some("a😀b".into()))],
16518 )
16519 .unwrap();
16520 assert_eq!(
16521 ScalarValue::try_from_array(&reversed, 0).unwrap(),
16522 ScalarValue::LargeUtf8(Some("b😀a".into()))
16523 );
16524
16525 let shorter: ArrayRef = Arc::new(Int64Array::from(vec![1, 2]));
16526 let longer: ArrayRef = Arc::new(Int64Array::from(vec![1, 2, 3]));
16527 let different: ArrayRef = Arc::new(Int64Array::from(vec![1, 9]));
16528 assert_eq!(
16529 cypher_seq_order(&shorter, &longer),
16530 std::cmp::Ordering::Less
16531 );
16532 assert_eq!(
16533 cypher_seq_order(&different, &shorter),
16534 std::cmp::Ordering::Greater
16535 );
16536 assert_eq!(
16537 cypher_seq_order(&shorter, &shorter),
16538 std::cmp::Ordering::Equal
16539 );
16540
16541 let legacy = DataType::Struct(Fields::from(vec![
16542 Field::new("__het_tag", DataType::Int8, false),
16543 Field::new("__het_int", DataType::Int64, true),
16544 Field::new("__het_float", DataType::Float64, true),
16545 Field::new("__het_str", DataType::Utf8, true),
16546 Field::new("__het_bool", DataType::Boolean, true),
16547 ]));
16548 let return_type = list_plus_return_type(&[
16549 DataType::new_list(legacy, true),
16550 DataType::new_list(DataType::Int64, true),
16551 ]);
16552 let DataType::List(item) = return_type else {
16553 panic!("list return")
16554 };
16555 let DataType::Struct(variants) = item.data_type() else {
16556 panic!("variant struct")
16557 };
16558 assert!(
16559 variants
16560 .iter()
16561 .any(|field| field.data_type() == &DataType::Int64)
16562 );
16563 assert!(
16564 variants
16565 .iter()
16566 .any(|field| field.data_type() == &DataType::Utf8)
16567 );
16568
16569 let mut percentile = PercentileAcc {
16570 continuous: true,
16571 value_type: DataType::Int64,
16572 result_type: DataType::Float64,
16573 values: Vec::new(),
16574 percentile: None,
16575 };
16576 let values: ArrayRef = Arc::new(Int64Array::from(vec![1]));
16577 let bad_percentile: ArrayRef =
16578 Arc::new(datafusion::arrow::array::StringArray::from(vec!["half"]));
16579 assert!(
16580 percentile
16581 .update_batch(&[values, bad_percentile])
16582 .unwrap_err()
16583 .to_string()
16584 .contains("must be numeric")
16585 );
16586 }
16587
16588 #[test]
16589 fn exact_zero_temporal_udf_metadata_contracts_are_total() {
16590 fn check<U: ScalarUDFImpl + 'static>(udf: U) {
16591 assert!(udf.as_any().is::<U>());
16592 assert!(!udf.name().is_empty());
16593 let _ = udf.signature();
16594 assert!(udf.return_type(&[DataType::Null]).is_ok());
16595 }
16596
16597 check(CypherDurationBetween::new());
16598 check(CypherTemporalArith::new());
16599 check(CypherDurationParse::new());
16600 check(CypherDurationAdd::new());
16601 check(CypherDurationScale::new());
16602 check(CypherDateProject::new());
16603 check(CypherLocalTimeProject::new());
16604 check(CypherLocalTimeTruncate::new());
16605 check(CypherLocalDateTimeProject::new());
16606 check(CypherLocalDateTimeTruncate::new());
16607 check(CypherTimeProject::new());
16608 check(CypherTimeTruncate::new());
16609 check(CypherDateTimeProject::new());
16610 check(CypherDateTimeTruncate::new());
16611 check(CypherToString::new());
16612 check(CypherDateTruncate::new());
16613 }
16614
16615 #[test]
16616 fn exact_zero_access_map_metadata_and_type_helpers() {
16617 use datafusion::arrow::datatypes::{Field, Fields};
16618
16619 let map = const_map_scalar(&[
16620 ("k".into(), ScalarValue::Int64(Some(7))),
16621 ("other".into(), ScalarValue::Int64(None)),
16622 ])
16623 .unwrap();
16624 let accessed =
16625 invoke_test_udf(&CypherStaticValueAccess::new("k".into()), vec![map.clone()]).unwrap();
16626 assert_eq!(
16627 ScalarValue::try_from_array(&accessed, 0).unwrap(),
16628 ScalarValue::Int64(Some(7))
16629 );
16630 let missing = invoke_test_udf(
16631 &CypherStaticValueAccess::new("missing".into()),
16632 vec![map.clone()],
16633 )
16634 .unwrap();
16635 assert!(ScalarValue::try_from_array(&missing, 0).unwrap().is_null());
16636
16637 let keys = invoke_test_udf(&CypherMapKeys::new(), vec![map]).unwrap();
16638 let ScalarValue::List(keys) = ScalarValue::try_from_array(&keys, 0).unwrap() else {
16639 panic!("keys must return a list")
16640 };
16641 assert_eq!(keys.value(0).len(), 2);
16642
16643 let props = invoke_test_udf(
16644 &CypherEntityProperties::new(3),
16645 vec![
16646 ScalarValue::Boolean(Some(true)),
16647 ScalarValue::Utf8(Some("k".into())),
16648 ScalarValue::Int64(Some(7)),
16649 ],
16650 )
16651 .unwrap();
16652 assert!(!props.is_null(0));
16653 let absent = invoke_test_udf(
16654 &CypherEntityProperties::new(3),
16655 vec![
16656 ScalarValue::Boolean(Some(false)),
16657 ScalarValue::Utf8(Some("k".into())),
16658 ScalarValue::Int64(Some(7)),
16659 ],
16660 )
16661 .unwrap();
16662 assert!(absent.is_null(0));
16663
16664 let null_access = invoke_test_udf(
16665 &CypherValueAccess::new(),
16666 vec![ScalarValue::Null, ScalarValue::Utf8(Some("k".into()))],
16667 )
16668 .unwrap();
16669 assert!(
16670 ScalarValue::try_from_array(&null_access, 0)
16671 .unwrap()
16672 .is_null()
16673 );
16674
16675 for (value, expected) in [
16676 (ScalarValue::Int8(Some(-1)), Some(-1)),
16677 (ScalarValue::Int16(Some(-2)), Some(-2)),
16678 (ScalarValue::Int32(Some(-3)), Some(-3)),
16679 (ScalarValue::Int64(Some(-4)), Some(-4)),
16680 (ScalarValue::UInt8(Some(1)), Some(1)),
16681 (ScalarValue::UInt16(Some(2)), Some(2)),
16682 (ScalarValue::UInt32(Some(3)), Some(3)),
16683 (ScalarValue::UInt64(Some(4)), Some(4)),
16684 (ScalarValue::Int64(None), None),
16685 ] {
16686 assert_eq!(scalar_list_index(&value).unwrap(), expected);
16687 }
16688 assert!(scalar_list_index(&ScalarValue::UInt64(Some(u64::MAX))).is_err());
16689 assert!(scalar_list_index(&ScalarValue::Utf8(Some("0".into()))).is_err());
16690
16691 let nested_a = DataType::Struct(Fields::from(vec![
16692 Field::new("value", DataType::Int64, false),
16693 Field::new("items", DataType::new_list(DataType::Utf8, true), true),
16694 ]));
16695 let nested_b = DataType::Struct(Fields::from(vec![
16696 Field::new("value", DataType::Int64, true),
16697 Field::new("items", DataType::new_list(DataType::Utf8, true), false),
16698 ]));
16699 assert!(graph_value_types_compatible(&nested_a, &nested_b));
16700 assert!(!graph_value_types_compatible(&nested_a, &DataType::Int64));
16701 assert!(!graph_value_types_compatible(
16702 &nested_a,
16703 &DataType::Struct(Fields::from(vec![Field::new(
16704 "other",
16705 DataType::Int64,
16706 true
16707 )]))
16708 ));
16709
16710 for (name, value) in [
16711 ("date", "2020-01-02"),
16712 ("localtime", "12:34:56"),
16713 ("time", "12:34:56Z"),
16714 ("localdatetime", "2020-01-02T12:34:56"),
16715 ("datetime", "2020-01-02T12:34:56Z"),
16716 ("duration", "P1D"),
16717 ] {
16718 assert!(render_temporal(name, value).is_some());
16719 }
16720 assert_eq!(render_temporal("unknown", "P1D"), None);
16721 }
16722
16723 #[test]
16724 fn exact_zero_dynamic_heterogeneous_list_builds_row_aligned_variants() {
16725 use datafusion::arrow::array::{Array, ListArray, StructArray};
16726
16727 let output = invoke_test_udf(
16728 &CypherDynamicHetList::new(),
16729 vec![
16730 ScalarValue::Int64(Some(7)),
16731 ScalarValue::Utf8(Some("seven".into())),
16732 ScalarValue::Boolean(None),
16733 ],
16734 )
16735 .unwrap();
16736 let lists = output.as_any().downcast_ref::<ListArray>().unwrap();
16737 assert_eq!(lists.len(), 1);
16738 assert_eq!(lists.value_length(0), 3);
16739 let values = lists.value(0);
16740 let variants = values.as_any().downcast_ref::<StructArray>().unwrap();
16741 assert_eq!(variants.len(), 3);
16742 assert!(!variants.is_null(0));
16743 assert!(!variants.is_null(1));
16744 assert!(variants.is_null(2));
16745 }
16746
16747 #[test]
16748 fn exact_zero_list_plus_handles_each_operand_shape_and_null_propagation() {
16749 use datafusion::arrow::array::{Array, ArrayRef, Int64Array, ListArray};
16750 use datafusion::arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
16751
16752 let list = |values: &[i64]| {
16753 ScalarValue::List(ScalarValue::new_list(
16754 &values
16755 .iter()
16756 .copied()
16757 .map(|value| ScalarValue::Int64(Some(value)))
16758 .collect::<Vec<_>>(),
16759 &DataType::Int64,
16760 true,
16761 ))
16762 };
16763 for (left, right, expected) in [
16764 (list(&[1, 2]), list(&[3, 4]), vec![1, 2, 3, 4]),
16765 (list(&[1, 2]), ScalarValue::Int64(Some(3)), vec![1, 2, 3]),
16766 (ScalarValue::Int64(Some(1)), list(&[2, 3]), vec![1, 2, 3]),
16767 ] {
16768 let output = invoke_test_udf(&CypherListPlus::new(), vec![left, right]).unwrap();
16769 let lists = output.as_any().downcast_ref::<ListArray>().unwrap();
16770 let values = lists.value(0);
16771 assert_eq!(
16772 (0..values.len())
16773 .map(|row| { unwrap_het(ScalarValue::try_from_array(&values, row).unwrap()) })
16774 .collect::<Vec<_>>(),
16775 expected
16776 .into_iter()
16777 .map(|value| ScalarValue::Int64(Some(value)))
16778 .collect::<Vec<_>>()
16779 );
16780 }
16781
16782 let null_list = ScalarValue::List(Arc::new(ListArray::new(
16783 Arc::new(Field::new("item", DataType::Int64, true)),
16784 OffsetBuffer::new(ScalarBuffer::from(vec![0, 0])),
16785 Arc::new(Int64Array::from(Vec::<i64>::new())) as ArrayRef,
16786 Some(NullBuffer::from(vec![false])),
16787 )));
16788 let output = invoke_test_udf(
16789 &CypherListPlus::new(),
16790 vec![null_list, ScalarValue::Int64(Some(1))],
16791 )
16792 .unwrap();
16793 assert!(output.is_null(0));
16794
16795 assert!(
16796 invoke_test_udf(
16797 &CypherListPlus::new(),
16798 vec![ScalarValue::Int64(Some(1)), ScalarValue::Int64(Some(2))],
16799 )
16800 .unwrap_err()
16801 .to_string()
16802 .contains("at least one list operand")
16803 );
16804 }
16805
16806 #[test]
16807 fn exact_zero_total_order_keys_distinguish_core_cypher_value_domains() {
16808 let list = ScalarValue::List(ScalarValue::new_list(
16809 &[ScalarValue::Int64(Some(1)), ScalarValue::Int64(Some(2))],
16810 &DataType::Int64,
16811 true,
16812 ));
16813 let values = [
16814 ScalarValue::Null,
16815 list,
16816 ScalarValue::Utf8(Some("text".into())),
16817 ScalarValue::Boolean(Some(true)),
16818 ScalarValue::Int64(Some(-2)),
16819 ScalarValue::Float64(Some(2.5)),
16820 ScalarValue::Float64(Some(f64::NAN)),
16821 ];
16822 let keys = values.iter().map(cypher_order_key).collect::<Vec<_>>();
16823 assert!(keys[0].starts_with("99:null"));
16824 assert!(keys[1].starts_with("40:list"));
16825 assert!(keys[2].starts_with("60:str"));
16826 assert!(keys[3].starts_with("70:bool"));
16827 assert!(keys[4].starts_with("80:num"));
16828 assert!(keys[5].starts_with("80:num"));
16829 assert!(keys[6].starts_with("90:nan"));
16830 assert_eq!(
16831 cypher_order(
16832 &ScalarValue::Utf8(Some("a".into())),
16833 &ScalarValue::Boolean(Some(false))
16834 ),
16835 std::cmp::Ordering::Less
16836 );
16837 }
16838
16839 #[test]
16840 fn exact_zero_dynamic_list_access_supports_negative_null_and_missing_indexes() {
16841 let values = ScalarValue::List(ScalarValue::new_list(
16842 &[
16843 ScalarValue::Utf8(Some("first".into())),
16844 ScalarValue::Utf8(Some("second".into())),
16845 ],
16846 &DataType::Utf8,
16847 true,
16848 ));
16849 for (index, expected) in [
16850 (ScalarValue::Int64(Some(0)), Some("first")),
16851 (ScalarValue::Int64(Some(-1)), Some("second")),
16852 (ScalarValue::Int64(Some(9)), None),
16853 (ScalarValue::Int64(Some(-9)), None),
16854 (ScalarValue::Int64(None), None),
16855 ] {
16856 let output =
16857 invoke_test_udf(&CypherValueAccess::new(), vec![values.clone(), index]).unwrap();
16858 assert_eq!(
16859 ScalarValue::try_from_array(&output, 0).unwrap(),
16860 ScalarValue::Utf8(expected.map(str::to_owned))
16861 );
16862 }
16863 assert!(
16864 invoke_test_udf(
16865 &CypherValueAccess::new(),
16866 vec![values, ScalarValue::Utf8(Some("not-an-index".into())),],
16867 )
16868 .unwrap_err()
16869 .to_string()
16870 .contains("index must be an integer")
16871 );
16872 }
16873
16874 #[test]
16875 fn exact_zero_udf_return_shape_guards_report_contract_errors() {
16876 let too_wide = (0..128)
16877 .map(|value| ScalarValue::Int64(Some(value)))
16878 .collect::<Vec<_>>();
16879 assert!(
16880 invoke_test_udf(&CypherDynamicHetList::new(), too_wide)
16881 .unwrap_err()
16882 .to_string()
16883 .contains("exceeds 127 elements")
16884 );
16885 assert!(
16886 invoke_test_udf_with_return_type(
16887 &CypherDynamicHetList::new(),
16888 vec![ScalarValue::Int64(Some(1))],
16889 DataType::Int64,
16890 )
16891 .unwrap_err()
16892 .to_string()
16893 .contains("non-list return type")
16894 );
16895 assert!(
16896 invoke_test_udf_with_return_type(
16897 &CypherDynamicHetList::new(),
16898 vec![ScalarValue::Int64(Some(1))],
16899 DataType::new_list(DataType::Int64, true),
16900 )
16901 .unwrap_err()
16902 .to_string()
16903 .contains("non-struct element type")
16904 );
16905 assert!(
16906 invoke_test_udf_with_return_type(
16907 &CypherListPlus::new(),
16908 vec![
16909 ScalarValue::List(ScalarValue::new_list(
16910 &[ScalarValue::Int64(Some(1))],
16911 &DataType::Int64,
16912 true,
16913 )),
16914 ScalarValue::Int64(Some(2)),
16915 ],
16916 DataType::Int64,
16917 )
16918 .unwrap_err()
16919 .to_string()
16920 .contains("return type is not a list")
16921 );
16922 }
16923
16924 #[test]
16925 fn exact_zero_tagged_append_rejects_incompatible_arrow_shapes() {
16926 use datafusion::arrow::array::{ArrayRef, Int64Array};
16927
16928 let scalar: ArrayRef = Arc::new(Int64Array::from(vec![1]));
16929 assert!(
16930 invoke_tagged_list_element_plus(&scalar, &scalar, &DataType::Int64)
16931 .unwrap()
16932 .is_none()
16933 );
16934 let list = ScalarValue::List(ScalarValue::new_list(
16935 &[ScalarValue::Int64(Some(1))],
16936 &DataType::Int64,
16937 true,
16938 ))
16939 .to_array_of_size(1)
16940 .unwrap();
16941 assert!(
16942 invoke_tagged_list_element_plus(&list, &scalar, &DataType::Int64)
16943 .unwrap()
16944 .is_none()
16945 );
16946
16947 let map = const_map_scalar(&[("value".into(), ScalarValue::Int64(Some(1)))])
16948 .unwrap()
16949 .to_array_of_size(1)
16950 .unwrap();
16951 assert!(
16952 invoke_tagged_list_element_plus(&list, &map, &DataType::Int64)
16953 .unwrap()
16954 .is_none()
16955 );
16956 assert!(
16957 invoke_tagged_list_element_plus(
16958 &list,
16959 &map,
16960 &DataType::new_list(map.data_type().clone(), true),
16961 )
16962 .unwrap()
16963 .is_none()
16964 );
16965 }
16966
16967 #[test]
16968 fn exact_zero_heterogeneous_depth_and_builder_type_matrix_is_total() {
16969 use datafusion::arrow::datatypes::{Field, Fields};
16970
16971 let primitives = [
16972 DataType::Null,
16973 DataType::Boolean,
16974 DataType::Int8,
16975 DataType::Int16,
16976 DataType::Int32,
16977 DataType::Int64,
16978 DataType::UInt8,
16979 DataType::UInt16,
16980 DataType::UInt32,
16981 DataType::UInt64,
16982 DataType::Float16,
16983 DataType::Float32,
16984 DataType::Float64,
16985 DataType::Utf8,
16986 DataType::LargeUtf8,
16987 ];
16988 for data_type in primitives {
16989 assert_eq!(het_depth_for_data_type(&data_type), Some(0));
16990 }
16991 assert_eq!(
16992 het_depth_for_data_type(&DataType::new_list(
16993 DataType::new_list(DataType::Int64, true),
16994 true,
16995 )),
16996 Some(2)
16997 );
16998 assert_eq!(het_depth_for_data_type(&DataType::Binary), None);
16999
17000 let map_type = DataType::Struct(Fields::from(vec![Field::new(
17001 "value",
17002 DataType::new_list(DataType::Int64, true),
17003 true,
17004 )]));
17005 assert_eq!(het_depth_for_data_type(&map_type), Some(2));
17006 assert!(build_het_struct(&[ScalarValue::Binary(Some(vec![1]))], 0).is_none());
17007
17008 let nested_list = ScalarValue::List(ScalarValue::new_list(
17009 &[ScalarValue::Int64(Some(1))],
17010 &DataType::Int64,
17011 true,
17012 ));
17013 assert!(build_het_struct(std::slice::from_ref(&nested_list), 0).is_none());
17014 assert!(build_het_struct(&[const_map_scalar(&[]).unwrap()], 0).is_none());
17015 let built = build_het_struct(
17016 &[
17017 ScalarValue::Int64(Some(1)),
17018 ScalarValue::Float64(Some(2.0)),
17019 ScalarValue::LargeUtf8(Some("three".into())),
17020 ScalarValue::Boolean(Some(true)),
17021 nested_list,
17022 const_map_scalar(&[("k".into(), ScalarValue::Int64(Some(4)))]).unwrap(),
17023 ScalarValue::Null,
17024 ],
17025 1,
17026 )
17027 .unwrap();
17028 assert_eq!(built.len(), 7);
17029 assert!(built.is_null(6));
17030 }
17031
17032 #[test]
17033 fn exact_zero_map_union_rejects_non_maps_and_conflicting_key_types() {
17034 assert!(all_map_union_list(&[ScalarValue::Int64(Some(1))]).is_none());
17035 let int_map = const_map_scalar(&[("key".into(), ScalarValue::Int64(Some(1)))]).unwrap();
17036 let text_map =
17037 const_map_scalar(&[("key".into(), ScalarValue::Utf8(Some("one".into())))]).unwrap();
17038 assert!(all_map_union_list(&[int_map, text_map]).is_none());
17039
17040 let left = const_map_scalar(&[("left".into(), ScalarValue::Int64(Some(1)))]).unwrap();
17041 let right =
17042 const_map_scalar(&[("right".into(), ScalarValue::Utf8(Some("r".into())))]).unwrap();
17043 let union = all_map_union_list(&[left, ScalarValue::Null, right]).unwrap();
17044 let DfExpr::Literal(ScalarValue::List(values), None) = union else {
17045 panic!("map union must const-fold to a list")
17046 };
17047 assert_eq!(values.value(0).len(), 3);
17048 assert!(values.value(0).is_null(1));
17049 }
17050
17051 #[test]
17052 fn exact_zero_map_and_subscript_error_guards_are_precise() {
17053 let null_keys = invoke_test_udf(&CypherMapKeys::new(), vec![ScalarValue::Null]).unwrap();
17054 assert!(null_keys.is_null(0));
17055 assert!(
17056 invoke_test_udf(&CypherMapKeys::new(), vec![ScalarValue::Int64(Some(1))])
17057 .unwrap_err()
17058 .to_string()
17059 .contains("keys() requires a map")
17060 );
17061 assert!(
17062 invoke_test_udf(
17063 &CypherStaticValueAccess::new("key".into()),
17064 vec![ScalarValue::Int64(Some(1))],
17065 )
17066 .unwrap_err()
17067 .to_string()
17068 .contains("property access requires a map")
17069 );
17070 assert!(
17071 invoke_test_udf(
17072 &CypherValueAccess::new(),
17073 vec![ScalarValue::Int64(Some(1)), ScalarValue::Int64(Some(0)),],
17074 )
17075 .unwrap_err()
17076 .to_string()
17077 .contains("requires a list or map")
17078 );
17079
17080 let map = const_map_scalar(&[("key".into(), ScalarValue::Int64(Some(7)))]).unwrap();
17081 let mismatch = invoke_test_udf_with_return_type(
17082 &CypherStaticValueAccess::new("key".into()),
17083 vec![map],
17084 DataType::Utf8,
17085 )
17086 .unwrap_err();
17087 assert!(mismatch.to_string().contains("incompatible runtime type"));
17088
17089 assert_eq!(value_access_return_type(None).unwrap(), DataType::Null);
17090 assert_eq!(
17091 value_access_return_type(Some(&DataType::Null)).unwrap(),
17092 DataType::Null
17093 );
17094 assert_eq!(
17095 value_access_return_type(Some(&DataType::new_list(DataType::Int64, true))).unwrap(),
17096 DataType::Int64
17097 );
17098 assert_eq!(
17099 static_value_access_return_type(None, "key").unwrap(),
17100 DataType::Null
17101 );
17102 }
17103
17104 #[test]
17105 fn exact_zero_heterogeneous_promotion_preserves_struct_and_list_validity() {
17106 use datafusion::arrow::array::{Array, Float64Array, ListArray, StructArray};
17107 use datafusion::arrow::datatypes::{Field, Fields};
17108
17109 let source_map = const_map_scalar(&[("present".into(), ScalarValue::Int64(Some(7)))])
17110 .unwrap()
17111 .to_array_of_size(1)
17112 .unwrap();
17113 assert!(Arc::ptr_eq(
17114 &source_map,
17115 &promote_het_array(&source_map, source_map.data_type()).unwrap()
17116 ));
17117 let target = DataType::Struct(Fields::from(vec![
17118 Field::new("present", DataType::Float64, true),
17119 Field::new("missing", DataType::Utf8, true),
17120 ]));
17121 let promoted = promote_het_array(&source_map, &target).unwrap();
17122 let promoted = promoted.as_any().downcast_ref::<StructArray>().unwrap();
17123 assert_eq!(promoted.num_columns(), 2);
17124 assert_eq!(
17125 promoted
17126 .column_by_name("present")
17127 .unwrap()
17128 .as_any()
17129 .downcast_ref::<Float64Array>()
17130 .unwrap()
17131 .value(0),
17132 7.0
17133 );
17134 assert!(promoted.column_by_name("missing").unwrap().is_null(0));
17135
17136 let source_list = ScalarValue::List(ScalarValue::new_list(
17137 &[ScalarValue::Int64(Some(1)), ScalarValue::Int64(None)],
17138 &DataType::Int64,
17139 true,
17140 ))
17141 .to_array_of_size(1)
17142 .unwrap();
17143 let target_list = DataType::new_list(DataType::Float64, true);
17144 let promoted = promote_het_array(&source_list, &target_list).unwrap();
17145 let promoted = promoted.as_any().downcast_ref::<ListArray>().unwrap();
17146 assert_eq!(promoted.value_length(0), 2);
17147 assert!(promoted.value(0).is_null(1));
17148 }
17149
17150 #[test]
17151 fn exact_zero_uncorrelated_list_comprehension_preserves_null_and_empty_rows() {
17152 use datafusion::arrow::array::{Array, Int64Builder, ListArray, ListBuilder};
17153 use datafusion::arrow::datatypes::Field;
17154 use datafusion::config::ConfigOptions;
17155
17156 let mut builder = ListBuilder::new(Int64Builder::new());
17157 builder.append_null();
17158 builder.append(true);
17159 builder.values().append_value(1);
17160 builder.values().append_value(2);
17161 builder.append(true);
17162 let input = Arc::new(builder.finish()) as datafusion::arrow::array::ArrayRef;
17163 let udf = CypherListComp::new(None, None, "__gf_elem".into(), vec![]);
17164 let return_type = udf.return_type(&[input.data_type().clone()]).unwrap();
17165 let output = udf
17166 .invoke_with_args(ScalarFunctionArgs {
17167 args: vec![ColumnarValue::Array(input)],
17168 arg_fields: vec![Arc::new(Field::new(
17169 "list",
17170 DataType::new_list(DataType::Int64, true),
17171 true,
17172 ))],
17173 number_rows: 3,
17174 return_field: Arc::new(Field::new("out", return_type, true)),
17175 config_options: Arc::new(ConfigOptions::default()),
17176 })
17177 .unwrap()
17178 .into_array(3)
17179 .unwrap();
17180 let output = output.as_any().downcast_ref::<ListArray>().unwrap();
17181 assert!(output.is_null(0));
17182 assert_eq!(output.value_length(1), 0);
17183 assert_eq!(output.value_length(2), 2);
17184 }
17185}