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 | "has_table_privilege"
99 | "has_column_privilege"
100 | "has_database_privilege"
101 | "has_schema_privilege"
102 | "has_sequence_privilege"
103 | "has_function_privilege"
104 )
105 }
106 _ => false,
107 };
108 if is_builtin {
109 local.to_string()
110 } else {
111 lower
112 }
113}
114
115pub fn is_builtin_aggregate(expr: &ScalarExpr) -> bool {
116 matches!(expr, ScalarExpr::Func { name, .. } if matches!(
117 name.to_ascii_lowercase().as_str(),
118 "count"
119 | "sum"
120 | "avg"
121 | "min"
122 | "max"
123 | "string_agg"
124 | "array_agg"
125 | "bool_and"
126 | "bool_or"
127 | "stddev"
128 | "stddev_samp"
129 | "stddev_pop"
130 | "variance"
131 | "var_samp"
132 | "var_pop"
133 | "percentile_cont"
134 | "percentile_disc"
135 | "mode"
136 | "json_agg"
137 | "jsonb_agg"
138 | "json_object_agg"
139 | "jsonb_object_agg"
140 ))
141}
142
143pub enum MultiFieldMatchShape<'a> {
144 FieldsThenQuery {
145 fields: Vec<&'a ScalarExpr>,
146 query_idx: usize,
147 },
148 Pairs {
149 fields: Vec<&'a ScalarExpr>,
150 },
151}
152
153pub fn multi_field_match_shape(args: &[ScalarExpr]) -> Result<MultiFieldMatchShape<'_>, SQLError> {
154 let first_non_column = args.iter().position(|arg| {
155 !matches!(
156 arg,
157 ScalarExpr::Column(_) | ScalarExpr::QualifiedColumn { .. }
158 )
159 });
160 if let Some(query_idx) = first_non_column {
161 if query_idx >= 2 {
162 return Ok(MultiFieldMatchShape::FieldsThenQuery {
163 fields: args[..query_idx].iter().collect(),
164 query_idx,
165 });
166 }
167 }
168 if args.len() < 4 || !args.len().is_multiple_of(2) {
169 if let Some(query_idx) = first_non_column {
170 if query_idx < 2 && args.len() >= 3 {
171 return Err(SQLError::TypeMismatch(format!(
172 "multi_field_match field arguments must be column references, \
173 but argument {} is an expression; store computed text in an \
174 indexed column instead of concatenating at query time",
175 query_idx + 1
176 )));
177 }
178 }
179 return Err(SQLError::BadArity {
180 name: "multi_field_match".into(),
181 expected: ">= 3 (fields..., query[, weights...]) or even >= 4 (field, query pairs)"
182 .into(),
183 actual: args.len(),
184 });
185 }
186 Ok(MultiFieldMatchShape::Pairs {
187 fields: (0..args.len() / 2).map(|i| &args[2 * i]).collect(),
188 })
189}
190
191pub fn is_semantic_field_argument(
193 function: &str,
194 args: &[ScalarExpr],
195 argument_index: usize,
196) -> Result<bool, SQLError> {
197 let dispatch_name = crate::semantics::builtin_function_dispatch_name(function);
198 let Some(kind) = crate::registry::lookup(&dispatch_name) else {
199 return Ok(false);
200 };
201 let is_field = match kind {
202 FunctionKind::TextMatch | FunctionKind::BayesianMatch | FunctionKind::KNNMatch => {
203 argument_index == 0
204 }
205 FunctionKind::FTSMatch => argument_index == 0 && !fts_query_is_jsonpath(args.get(1)),
206 FunctionKind::BayesianMatchWithPrior => matches!(argument_index, 0 | 2),
207 FunctionKind::CalibratedVectorMatch => argument_index == 0,
208 FunctionKind::MultiFieldMatch => match multi_field_match_shape(args)? {
209 MultiFieldMatchShape::FieldsThenQuery { query_idx, .. } => argument_index < query_idx,
210 MultiFieldMatchShape::Pairs { .. } => argument_index.is_multiple_of(2),
211 },
212 FunctionKind::StagedRetrieval => {
213 !matches!(args.first(), Some(ScalarExpr::Func { .. }))
214 && argument_index.is_multiple_of(3)
215 }
216 FunctionKind::UQAFacets => true,
217 FunctionKind::ScoreBM25 | FunctionKind::ScoreBayesianBM25 => {
218 args.len() == 2 && argument_index == 0
219 }
220 FunctionKind::FuseLogOdds
221 | FunctionKind::PositiveEvidencePool
222 | FunctionKind::BayesianEvidenceFusion
223 | FunctionKind::GraphPagerank
224 | FunctionKind::GraphHits
225 | FunctionKind::GraphBetweenness
226 | FunctionKind::GraphTraverse
227 | FunctionKind::GraphNeighbors
228 | FunctionKind::DeepPredict
229 | FunctionKind::UQAHighlight
230 | FunctionKind::TraverseMatch
231 | FunctionKind::TemporalTraverse
232 | FunctionKind::RPQ
233 | FunctionKind::GraphCreate
234 | FunctionKind::GraphDrop
235 | FunctionKind::GraphExists
236 | FunctionKind::GraphLabelCreate
237 | FunctionKind::GraphLabelDrop
238 | FunctionKind::GraphAlter
239 | FunctionKind::GraphEdges
240 | FunctionKind::AttentionFusion
241 | FunctionKind::LearnedFusion
242 | FunctionKind::SparseThreshold
243 | FunctionKind::DeepLearn
244 | FunctionKind::Convolve
245 | FunctionKind::Pool
246 | FunctionKind::Flatten
247 | FunctionKind::Dense
248 | FunctionKind::Softmax
249 | FunctionKind::Layer
250 | FunctionKind::Model => false,
251 };
252 Ok(is_field)
253}
254
255pub fn fts_query_is_jsonpath(query_arg: Option<&ScalarExpr>) -> bool {
259 matches!(
260 query_arg,
261 Some(ScalarExpr::Literal(Value::Str(path))) if path.trim_start().starts_with('$')
262 )
263}
264
265pub fn contains_retrieval(expression: &ScalarExpr) -> bool {
269 match expression {
270 ScalarExpr::Func {
271 name,
272 args,
273 order_by,
274 filter,
275 ..
276 } => {
277 retrieval_function(name)
278 || args.iter().any(contains_retrieval)
279 || order_by.iter().any(|order| contains_retrieval(&order.expr))
280 || filter.as_deref().is_some_and(contains_retrieval)
281 }
282 ScalarExpr::Array(items)
283 | ScalarExpr::Row(items)
284 | ScalarExpr::And(items)
285 | ScalarExpr::Or(items) => items.iter().any(contains_retrieval),
286 ScalarExpr::Binary { lhs, rhs, .. } => contains_retrieval(lhs) || contains_retrieval(rhs),
287 ScalarExpr::UnaryMinus(inner)
288 | ScalarExpr::Not(inner)
289 | ScalarExpr::IsNull { expr: inner, .. }
290 | ScalarExpr::Cast { expr: inner, .. } => contains_retrieval(inner),
291 ScalarExpr::Between { expr, low, high } => {
292 contains_retrieval(expr) || contains_retrieval(low) || contains_retrieval(high)
293 }
294 ScalarExpr::InList { expr, list, .. } => {
295 contains_retrieval(expr) || list.iter().any(contains_retrieval)
296 }
297 ScalarExpr::WindowCall { args, spec, .. } => {
298 args.iter().any(contains_retrieval)
299 || spec.partition_by.iter().any(contains_retrieval)
300 || spec
301 .order_by
302 .iter()
303 .any(|order| contains_retrieval(&order.expr))
304 }
305 ScalarExpr::Case {
306 base,
307 when,
308 else_branch,
309 } => {
310 base.as_deref().is_some_and(contains_retrieval)
311 || when.iter().any(|(condition, result)| {
312 contains_retrieval(condition) || contains_retrieval(result)
313 })
314 || else_branch.as_deref().is_some_and(contains_retrieval)
315 }
316 ScalarExpr::InSubquery { expr, .. } => contains_retrieval(expr),
317 ScalarExpr::Default
318 | ScalarExpr::Star
319 | ScalarExpr::QualifiedStar(_)
320 | ScalarExpr::Column(_)
321 | ScalarExpr::Position(_)
322 | ScalarExpr::InternalColumn(_)
323 | ScalarExpr::QualifiedColumn { .. }
324 | ScalarExpr::Literal(_)
325 | ScalarExpr::TypedLiteral { .. }
326 | ScalarExpr::Param(_)
327 | ScalarExpr::ScalarSubquery(_)
328 | ScalarExpr::Exists { .. } => false,
329 }
330}
331
332pub fn retrieval_function(name: &str) -> bool {
333 matches!(
334 name.to_ascii_lowercase().as_str(),
335 "text_match"
336 | "bayesian_match"
337 | "fts_match"
338 | "bayesian_match_with_prior"
339 | "calibrated_vector_match"
340 | "knn_match"
341 | "fuse_log_odds"
342 | "pool_positive_evidence"
343 | "fuse_bayesian_evidence"
344 | "multi_field_match"
345 | "staged_retrieval"
346 | "attention"
347 | "fuse_attention"
348 | "fuse_multihead"
349 | "learned_fusion"
350 | "fuse_learned"
351 | "sparse_threshold"
352 | "graph_pagerank"
353 | "pagerank"
354 | "graph_hits"
355 | "hits"
356 | "graph_betweenness"
357 | "betweenness"
358 | "graph_traverse"
359 | "traverse_match"
360 | "graph_neighbors"
361 | "graph_edges"
362 | "temporal_traverse"
363 | "rpq"
364 | "deep_predict"
365 )
366}
367
368pub fn expect_column_name(expr: &ScalarExpr, label: &str) -> Result<String, SQLError> {
369 match expr {
370 ScalarExpr::Column(name) => Ok(name.clone()),
371 ScalarExpr::QualifiedColumn { column, .. } => Ok(column.clone()),
372 other => Err(SQLError::TypeMismatch(format!(
373 "{label} must be a column reference, got {other:?}"
374 ))),
375 }
376}
377
378#[cfg(test)]
379mod tests {
380 use super::builtin_function_dispatch_name;
381
382 #[test]
383 fn sequence_introspection_dispatch_preserves_schema_identity() {
384 for name in [
385 "pg_get_sequence_data",
386 "pg_sequence_last_value",
387 "pg_sequence_parameters",
388 ] {
389 assert_eq!(builtin_function_dispatch_name(name), name);
390 assert_eq!(
391 builtin_function_dispatch_name(&format!("pg_catalog.{name}")),
392 name
393 );
394 assert_eq!(
395 builtin_function_dispatch_name(&format!("PG_CATALOG.{}", name.to_uppercase())),
396 name
397 );
398 let user_function = format!("public.{name}");
399 assert_eq!(
400 builtin_function_dispatch_name(&user_function),
401 user_function
402 );
403 }
404 assert_eq!(
405 builtin_function_dispatch_name("pg_catalog.custom_function"),
406 "pg_catalog.custom_function"
407 );
408 }
409}