uqa_sql/semantics/
functions.rs1use crate::registry::FunctionKind;
10use crate::{SQLError, ScalarExpr};
11use std::sync::LazyLock;
12use uqa_core::Value;
13
14pub fn merge_action_attribute() -> crate::ast::InternalColumnRef {
16 static ATTRIBUTE: LazyLock<crate::ast::InternalColumnRef> =
17 LazyLock::new(|| crate::ast::InternalRelationId::allocate().column(0));
18 *ATTRIBUTE
19}
20pub fn builtin_function_dispatch_name(name: &str) -> String {
24 let lower = name.to_ascii_lowercase();
25 let Some((schema, local)) = lower.split_once('.') else {
26 return lower;
27 };
28 let is_builtin = match schema {
29 "ag_catalog" => matches!(
30 local,
31 "cypher"
32 | "create_graph"
33 | "drop_graph"
34 | "graph_exists"
35 | "create_vlabel"
36 | "create_elabel"
37 | "drop_label"
38 | "alter_graph"
39 ),
40 "pg_catalog" => {
41 crate::registry::is_registered(local)
42 || matches!(
43 local,
44 "generate_series"
45 | "unnest"
46 | "regexp_split_to_table"
47 | "string_to_table"
48 | "json_array_elements"
49 | "jsonb_array_elements"
50 | "json_array_elements_text"
51 | "jsonb_array_elements_text"
52 | "json_each"
53 | "jsonb_each"
54 | "json_each_text"
55 | "jsonb_each_text"
56 | "json_object_keys"
57 | "jsonb_object_keys"
58 | "upper"
59 | "lower"
60 | "bit_length"
61 | "char_length"
62 | "character_length"
63 | "crc32"
64 | "crc32c"
65 | "gamma"
66 | "json_strip_nulls"
67 | "jsonb_strip_nulls"
68 | "length"
69 | "lgamma"
70 | "md5"
71 | "octet_length"
72 | "reverse"
73 | "random"
74 | "setseed"
75 | "nextval"
76 | "currval"
77 | "lastval"
78 | "setval"
79 | "current_schema"
80 | "current_schemas"
81 | "current_setting"
82 | "pg_backend_pid"
83 | "pg_listening_channels"
84 | "pg_notify"
85 | "pg_notification_queue_usage"
86 | "pg_get_expr"
87 | "pg_get_partkeydef"
88 | "pg_get_serial_sequence"
89 | "pg_get_sequence_data"
90 | "pg_sequence_last_value"
91 | "pg_sequence_parameters"
92 | "pg_get_triggerdef"
93 | "pg_get_ruledef"
94 | "pg_get_viewdef"
95 | "pg_get_indexdef"
96 | "format_type"
97 | "pg_has_role"
98 | "pg_get_userbyid"
99 | "has_table_privilege"
100 | "has_column_privilege"
101 | "has_database_privilege"
102 | "has_schema_privilege"
103 | "has_sequence_privilege"
104 | "has_function_privilege"
105 )
106 }
107 _ => false,
108 };
109 if is_builtin {
110 local.to_string()
111 } else {
112 lower
113 }
114}
115
116pub fn is_builtin_aggregate(expr: &ScalarExpr) -> bool {
117 matches!(expr, ScalarExpr::Func { name, .. } if matches!(
118 name.to_ascii_lowercase().as_str(),
119 "count"
120 | "sum"
121 | "avg"
122 | "min"
123 | "max"
124 | "string_agg"
125 | "array_agg"
126 | "bool_and"
127 | "bool_or"
128 | "stddev"
129 | "stddev_samp"
130 | "stddev_pop"
131 | "variance"
132 | "var_samp"
133 | "var_pop"
134 | "percentile_cont"
135 | "percentile_disc"
136 | "mode"
137 | "json_agg"
138 | "jsonb_agg"
139 | "json_object_agg"
140 | "jsonb_object_agg"
141 ))
142}
143
144pub enum MultiFieldMatchShape<'a> {
145 FieldsThenQuery {
146 fields: Vec<&'a ScalarExpr>,
147 query_idx: usize,
148 },
149 Pairs {
150 fields: Vec<&'a ScalarExpr>,
151 },
152}
153
154pub fn multi_field_match_shape(args: &[ScalarExpr]) -> Result<MultiFieldMatchShape<'_>, SQLError> {
155 let first_non_column = args.iter().position(|arg| {
156 !matches!(
157 arg,
158 ScalarExpr::Column(_) | ScalarExpr::QualifiedColumn { .. }
159 )
160 });
161 if let Some(query_idx) = first_non_column {
162 if query_idx >= 2 {
163 return Ok(MultiFieldMatchShape::FieldsThenQuery {
164 fields: args[..query_idx].iter().collect(),
165 query_idx,
166 });
167 }
168 }
169 if args.len() < 4 || !args.len().is_multiple_of(2) {
170 if let Some(query_idx) = first_non_column {
171 if query_idx < 2 && args.len() >= 3 {
172 return Err(SQLError::TypeMismatch(format!(
173 "multi_field_match field arguments must be column references, \
174 but argument {} is an expression; store computed text in an \
175 indexed column instead of concatenating at query time",
176 query_idx + 1
177 )));
178 }
179 }
180 return Err(SQLError::BadArity {
181 name: "multi_field_match".into(),
182 expected: ">= 3 (fields..., query[, weights...]) or even >= 4 (field, query pairs)"
183 .into(),
184 actual: args.len(),
185 });
186 }
187 Ok(MultiFieldMatchShape::Pairs {
188 fields: (0..args.len() / 2).map(|i| &args[2 * i]).collect(),
189 })
190}
191
192pub fn is_semantic_field_argument(
194 function: &str,
195 args: &[ScalarExpr],
196 argument_index: usize,
197) -> Result<bool, SQLError> {
198 let dispatch_name = crate::semantics::builtin_function_dispatch_name(function);
199 let Some(kind) = crate::registry::lookup(&dispatch_name) else {
200 return Ok(false);
201 };
202 let is_field = match kind {
203 FunctionKind::TextMatch | FunctionKind::BayesianMatch | FunctionKind::KNNMatch => {
204 argument_index == 0
205 }
206 FunctionKind::FTSMatch => argument_index == 0 && !fts_query_is_jsonpath(args.get(1)),
207 FunctionKind::BayesianMatchWithPrior => matches!(argument_index, 0 | 2),
208 FunctionKind::CalibratedVectorMatch => argument_index == 0,
209 FunctionKind::MultiFieldMatch => match multi_field_match_shape(args)? {
210 MultiFieldMatchShape::FieldsThenQuery { query_idx, .. } => argument_index < query_idx,
211 MultiFieldMatchShape::Pairs { .. } => argument_index.is_multiple_of(2),
212 },
213 FunctionKind::StagedRetrieval => {
214 !matches!(args.first(), Some(ScalarExpr::Func { .. }))
215 && argument_index.is_multiple_of(3)
216 }
217 FunctionKind::UQAFacets => true,
218 FunctionKind::ScoreBM25 | FunctionKind::ScoreBayesianBM25 => {
219 args.len() == 2 && argument_index == 0
220 }
221 FunctionKind::FuseLogOdds
222 | FunctionKind::PositiveEvidencePool
223 | FunctionKind::BayesianEvidenceFusion
224 | FunctionKind::GraphPagerank
225 | FunctionKind::GraphHits
226 | FunctionKind::GraphBetweenness
227 | FunctionKind::GraphTraverse
228 | FunctionKind::GraphNeighbors
229 | FunctionKind::DeepPredict
230 | FunctionKind::UQAHighlight
231 | FunctionKind::TraverseMatch
232 | FunctionKind::TemporalTraverse
233 | FunctionKind::RPQ
234 | FunctionKind::GraphCreate
235 | FunctionKind::GraphDrop
236 | FunctionKind::GraphExists
237 | FunctionKind::GraphLabelCreate
238 | FunctionKind::GraphLabelDrop
239 | FunctionKind::GraphAlter
240 | FunctionKind::GraphEdges
241 | FunctionKind::AttentionFusion
242 | FunctionKind::LearnedFusion
243 | FunctionKind::SparseThreshold
244 | FunctionKind::DeepLearn
245 | FunctionKind::Convolve
246 | FunctionKind::Pool
247 | FunctionKind::Flatten
248 | FunctionKind::Dense
249 | FunctionKind::Softmax
250 | FunctionKind::Layer
251 | FunctionKind::Model => false,
252 };
253 Ok(is_field)
254}
255
256pub fn fts_query_is_jsonpath(query_arg: Option<&ScalarExpr>) -> bool {
260 matches!(
261 query_arg,
262 Some(ScalarExpr::Literal(Value::Str(path))) if path.trim_start().starts_with('$')
263 )
264}
265
266pub fn contains_retrieval(expression: &ScalarExpr) -> bool {
270 match expression {
271 ScalarExpr::Func {
272 name,
273 args,
274 order_by,
275 filter,
276 ..
277 } => {
278 retrieval_function(name)
279 || args.iter().any(contains_retrieval)
280 || order_by.iter().any(|order| contains_retrieval(&order.expr))
281 || filter.as_deref().is_some_and(contains_retrieval)
282 }
283 ScalarExpr::Array(items)
284 | ScalarExpr::Row(items)
285 | ScalarExpr::And(items)
286 | ScalarExpr::Or(items) => items.iter().any(contains_retrieval),
287 ScalarExpr::Binary { lhs, rhs, .. } => contains_retrieval(lhs) || contains_retrieval(rhs),
288 ScalarExpr::UnaryMinus(inner)
289 | ScalarExpr::Not(inner)
290 | ScalarExpr::IsNull { expr: inner, .. }
291 | ScalarExpr::Cast { expr: inner, .. } => contains_retrieval(inner),
292 ScalarExpr::Between { expr, low, high } => {
293 contains_retrieval(expr) || contains_retrieval(low) || contains_retrieval(high)
294 }
295 ScalarExpr::InList { expr, list, .. } => {
296 contains_retrieval(expr) || list.iter().any(contains_retrieval)
297 }
298 ScalarExpr::WindowCall { args, spec, .. } => {
299 args.iter().any(contains_retrieval)
300 || spec.partition_by.iter().any(contains_retrieval)
301 || spec
302 .order_by
303 .iter()
304 .any(|order| contains_retrieval(&order.expr))
305 }
306 ScalarExpr::Case {
307 base,
308 when,
309 else_branch,
310 } => {
311 base.as_deref().is_some_and(contains_retrieval)
312 || when.iter().any(|(condition, result)| {
313 contains_retrieval(condition) || contains_retrieval(result)
314 })
315 || else_branch.as_deref().is_some_and(contains_retrieval)
316 }
317 ScalarExpr::InSubquery { expr, .. } => contains_retrieval(expr),
318 ScalarExpr::Default
319 | ScalarExpr::Star
320 | ScalarExpr::QualifiedStar(_)
321 | ScalarExpr::Column(_)
322 | ScalarExpr::Position(_)
323 | ScalarExpr::InternalColumn(_)
324 | ScalarExpr::QualifiedColumn { .. }
325 | ScalarExpr::Literal(_)
326 | ScalarExpr::TypedLiteral { .. }
327 | ScalarExpr::Param(_)
328 | ScalarExpr::ScalarSubquery(_)
329 | ScalarExpr::Exists { .. } => false,
330 }
331}
332
333pub fn retrieval_function(name: &str) -> bool {
334 matches!(
335 name.to_ascii_lowercase().as_str(),
336 "text_match"
337 | "bayesian_match"
338 | "fts_match"
339 | "bayesian_match_with_prior"
340 | "calibrated_vector_match"
341 | "knn_match"
342 | "fuse_log_odds"
343 | "pool_positive_evidence"
344 | "fuse_bayesian_evidence"
345 | "multi_field_match"
346 | "staged_retrieval"
347 | "attention"
348 | "fuse_attention"
349 | "fuse_multihead"
350 | "learned_fusion"
351 | "fuse_learned"
352 | "sparse_threshold"
353 | "graph_pagerank"
354 | "pagerank"
355 | "graph_hits"
356 | "hits"
357 | "graph_betweenness"
358 | "betweenness"
359 | "graph_traverse"
360 | "traverse_match"
361 | "graph_neighbors"
362 | "graph_edges"
363 | "temporal_traverse"
364 | "rpq"
365 | "deep_predict"
366 )
367}
368
369pub fn expect_column_name(expr: &ScalarExpr, label: &str) -> Result<String, SQLError> {
370 match expr {
371 ScalarExpr::Column(name) => Ok(name.clone()),
372 ScalarExpr::QualifiedColumn { column, .. } => Ok(column.clone()),
373 other => Err(SQLError::TypeMismatch(format!(
374 "{label} must be a column reference, got {other:?}"
375 ))),
376 }
377}
378
379#[cfg(test)]
380mod tests {
381 use super::builtin_function_dispatch_name;
382
383 #[test]
384 fn sequence_introspection_dispatch_preserves_schema_identity() {
385 for name in [
386 "pg_get_sequence_data",
387 "pg_sequence_last_value",
388 "pg_sequence_parameters",
389 ] {
390 assert_eq!(builtin_function_dispatch_name(name), name);
391 assert_eq!(
392 builtin_function_dispatch_name(&format!("pg_catalog.{name}")),
393 name
394 );
395 assert_eq!(
396 builtin_function_dispatch_name(&format!("PG_CATALOG.{}", name.to_uppercase())),
397 name
398 );
399 let user_function = format!("public.{name}");
400 assert_eq!(
401 builtin_function_dispatch_name(&user_function),
402 user_function
403 );
404 }
405 assert_eq!(
406 builtin_function_dispatch_name("pg_catalog.custom_function"),
407 "pg_catalog.custom_function"
408 );
409 }
410}