Skip to main content

reifydb_evaluate/expression/
udf_extract.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2026 ReifyDB
3
4use reifydb_core::interface::identifier::{ColumnIdentifier, ColumnObject};
5use reifydb_rql::expression::{
6	AddExpression, AliasExpression, AndExpression, BetweenExpression, CallExpression, CastExpression,
7	ColumnExpression, ContainsExpression, DivExpression, ElseIfExpression, EqExpression, Expression,
8	ExtendExpression, FieldAccessExpression, GreaterThanEqExpression, GreaterThanExpression, IfExpression,
9	InExpression, LessThanEqExpression, LessThanExpression, ListExpression, MapExpression, MulExpression,
10	NotEqExpression, OrExpression, PrefixExpression, RemExpression, SubExpression, TupleExpression, XorExpression,
11};
12use reifydb_value::fragment::Fragment;
13
14use crate::stack::{Callable, SymbolTable};
15
16pub struct ExtractedUdf {
17	pub callable: Callable,
18	pub arg_expressions: Vec<Expression>,
19	pub result_column: Fragment,
20}
21
22pub fn extract_udf_calls(
23	expr: &Expression,
24	symbols: &SymbolTable,
25	counter: &mut usize,
26) -> (Expression, Vec<ExtractedUdf>) {
27	let mut extracted = Vec::new();
28	let rewritten = rewrite_expr(expr, symbols, counter, &mut extracted);
29	(rewritten, extracted)
30}
31
32fn rw(e: &Expression, s: &SymbolTable, c: &mut usize, x: &mut Vec<ExtractedUdf>) -> Expression {
33	rewrite_expr(e, s, c, x)
34}
35
36fn rw_vec(exprs: &[Expression], s: &SymbolTable, c: &mut usize, x: &mut Vec<ExtractedUdf>) -> Vec<Expression> {
37	exprs.iter().map(|e| rewrite_expr(e, s, c, x)).collect()
38}
39
40fn rewrite_expr(
41	expr: &Expression,
42	symbols: &SymbolTable,
43	counter: &mut usize,
44	extracted: &mut Vec<ExtractedUdf>,
45) -> Expression {
46	match expr {
47		Expression::Call(call) => {
48			let rewritten_args = rw_vec(&call.args, symbols, counter, extracted);
49			let function_name = call.func.0.text();
50
51			if let Some(callable) = symbols.resolve_callable(function_name) {
52				let col_name = Fragment::internal(format!("__udf_{}", counter));
53				*counter += 1;
54
55				extracted.push(ExtractedUdf {
56					callable,
57					arg_expressions: rewritten_args,
58					result_column: col_name.clone(),
59				});
60
61				Expression::Column(ColumnExpression(ColumnIdentifier {
62					object: ColumnObject::Alias(Fragment::internal("")),
63					name: col_name,
64				}))
65			} else {
66				Expression::Call(CallExpression {
67					func: call.func.clone(),
68					args: rewritten_args,
69					fragment: call.fragment.clone(),
70				})
71			}
72		}
73
74		Expression::Add(e) => Expression::Add(AddExpression {
75			left: Box::new(rw(&e.left, symbols, counter, extracted)),
76			right: Box::new(rw(&e.right, symbols, counter, extracted)),
77			fragment: e.fragment.clone(),
78		}),
79		Expression::Sub(e) => Expression::Sub(SubExpression {
80			left: Box::new(rw(&e.left, symbols, counter, extracted)),
81			right: Box::new(rw(&e.right, symbols, counter, extracted)),
82			fragment: e.fragment.clone(),
83		}),
84		Expression::Mul(e) => Expression::Mul(MulExpression {
85			left: Box::new(rw(&e.left, symbols, counter, extracted)),
86			right: Box::new(rw(&e.right, symbols, counter, extracted)),
87			fragment: e.fragment.clone(),
88		}),
89		Expression::Div(e) => Expression::Div(DivExpression {
90			left: Box::new(rw(&e.left, symbols, counter, extracted)),
91			right: Box::new(rw(&e.right, symbols, counter, extracted)),
92			fragment: e.fragment.clone(),
93		}),
94		Expression::Rem(e) => Expression::Rem(RemExpression {
95			left: Box::new(rw(&e.left, symbols, counter, extracted)),
96			right: Box::new(rw(&e.right, symbols, counter, extracted)),
97			fragment: e.fragment.clone(),
98		}),
99		Expression::GreaterThan(e) => Expression::GreaterThan(GreaterThanExpression {
100			left: Box::new(rw(&e.left, symbols, counter, extracted)),
101			right: Box::new(rw(&e.right, symbols, counter, extracted)),
102			fragment: e.fragment.clone(),
103		}),
104		Expression::GreaterThanEqual(e) => Expression::GreaterThanEqual(GreaterThanEqExpression {
105			left: Box::new(rw(&e.left, symbols, counter, extracted)),
106			right: Box::new(rw(&e.right, symbols, counter, extracted)),
107			fragment: e.fragment.clone(),
108		}),
109		Expression::LessThan(e) => Expression::LessThan(LessThanExpression {
110			left: Box::new(rw(&e.left, symbols, counter, extracted)),
111			right: Box::new(rw(&e.right, symbols, counter, extracted)),
112			fragment: e.fragment.clone(),
113		}),
114		Expression::LessThanEqual(e) => Expression::LessThanEqual(LessThanEqExpression {
115			left: Box::new(rw(&e.left, symbols, counter, extracted)),
116			right: Box::new(rw(&e.right, symbols, counter, extracted)),
117			fragment: e.fragment.clone(),
118		}),
119		Expression::Equal(e) => Expression::Equal(EqExpression {
120			left: Box::new(rw(&e.left, symbols, counter, extracted)),
121			right: Box::new(rw(&e.right, symbols, counter, extracted)),
122			fragment: e.fragment.clone(),
123		}),
124		Expression::NotEqual(e) => Expression::NotEqual(NotEqExpression {
125			left: Box::new(rw(&e.left, symbols, counter, extracted)),
126			right: Box::new(rw(&e.right, symbols, counter, extracted)),
127			fragment: e.fragment.clone(),
128		}),
129		Expression::And(e) => Expression::And(AndExpression {
130			left: Box::new(rw(&e.left, symbols, counter, extracted)),
131			right: Box::new(rw(&e.right, symbols, counter, extracted)),
132			fragment: e.fragment.clone(),
133		}),
134		Expression::Or(e) => Expression::Or(OrExpression {
135			left: Box::new(rw(&e.left, symbols, counter, extracted)),
136			right: Box::new(rw(&e.right, symbols, counter, extracted)),
137			fragment: e.fragment.clone(),
138		}),
139		Expression::Xor(e) => Expression::Xor(XorExpression {
140			left: Box::new(rw(&e.left, symbols, counter, extracted)),
141			right: Box::new(rw(&e.right, symbols, counter, extracted)),
142			fragment: e.fragment.clone(),
143		}),
144
145		Expression::Between(e) => Expression::Between(BetweenExpression {
146			value: Box::new(rw(&e.value, symbols, counter, extracted)),
147			lower: Box::new(rw(&e.lower, symbols, counter, extracted)),
148			upper: Box::new(rw(&e.upper, symbols, counter, extracted)),
149			fragment: e.fragment.clone(),
150		}),
151		Expression::In(e) => Expression::In(InExpression {
152			value: Box::new(rw(&e.value, symbols, counter, extracted)),
153			list: Box::new(rw(&e.list, symbols, counter, extracted)),
154			negated: e.negated,
155			fragment: e.fragment.clone(),
156		}),
157		Expression::Contains(e) => Expression::Contains(ContainsExpression {
158			value: Box::new(rw(&e.value, symbols, counter, extracted)),
159			list: Box::new(rw(&e.list, symbols, counter, extracted)),
160			fragment: e.fragment.clone(),
161		}),
162		Expression::Cast(e) => Expression::Cast(CastExpression {
163			expression: Box::new(rw(&e.expression, symbols, counter, extracted)),
164			to: e.to.clone(),
165			fragment: e.fragment.clone(),
166		}),
167		Expression::Prefix(e) => Expression::Prefix(PrefixExpression {
168			operator: e.operator.clone(),
169			expression: Box::new(rw(&e.expression, symbols, counter, extracted)),
170			fragment: e.fragment.clone(),
171		}),
172		Expression::Alias(e) => Expression::Alias(AliasExpression {
173			alias: e.alias.clone(),
174			expression: Box::new(rw(&e.expression, symbols, counter, extracted)),
175			fragment: e.fragment.clone(),
176		}),
177		Expression::If(e) => Expression::If(IfExpression {
178			condition: Box::new(rw(&e.condition, symbols, counter, extracted)),
179			then_expr: Box::new(rw(&e.then_expr, symbols, counter, extracted)),
180			else_ifs: e
181				.else_ifs
182				.iter()
183				.map(|ei| ElseIfExpression {
184					condition: Box::new(rw(&ei.condition, symbols, counter, extracted)),
185					then_expr: Box::new(rw(&ei.then_expr, symbols, counter, extracted)),
186					fragment: ei.fragment.clone(),
187				})
188				.collect(),
189			else_expr: e.else_expr.as_ref().map(|b| Box::new(rw(b, symbols, counter, extracted))),
190			fragment: e.fragment.clone(),
191		}),
192		Expression::FieldAccess(e) => Expression::FieldAccess(FieldAccessExpression {
193			object: Box::new(rw(&e.object, symbols, counter, extracted)),
194			field: e.field.clone(),
195			fragment: e.fragment.clone(),
196		}),
197
198		Expression::Tuple(e) => Expression::Tuple(TupleExpression {
199			expressions: rw_vec(&e.expressions, symbols, counter, extracted),
200			fragment: e.fragment.clone(),
201		}),
202		Expression::List(e) => Expression::List(ListExpression {
203			expressions: rw_vec(&e.expressions, symbols, counter, extracted),
204			fragment: e.fragment.clone(),
205		}),
206		Expression::Map(e) => Expression::Map(MapExpression {
207			expressions: rw_vec(&e.expressions, symbols, counter, extracted),
208			fragment: e.fragment.clone(),
209		}),
210		Expression::Extend(e) => Expression::Extend(ExtendExpression {
211			expressions: rw_vec(&e.expressions, symbols, counter, extracted),
212			fragment: e.fragment.clone(),
213		}),
214
215		Expression::Constant(_)
216		| Expression::Column(_)
217		| Expression::AccessSource(_)
218		| Expression::Parameter(_)
219		| Expression::Variable(_)
220		| Expression::Type(_)
221		| Expression::SumTypeConstructor(_)
222		| Expression::IsVariant(_) => expr.clone(),
223	}
224}