1use 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}